autool/DroidBot/input_policy.py
2026-06-17 19:44:18 +08:00

898 lines
41 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""
Input Policy Module for DroidBot
Contains MemoryGuidedPolicy for intelligent UI exploration.
"""
import sys
import json
import logging
import random
import collections
import copy
import time
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.nn.utils.rnn import pad_sequence
from .core import PlatformFactory, Platform
from .core.abstract_input_event import BaseTouchEvent, BaseKeyEvent, EventType, BaseKillAppEvent,BaseSetTextEvent
from .utg import UTG
from .guiagent_bridge import GuiAgentBridge
from .exceptions import FATAL_EXCEPTIONS, InputInterruptedException, ExplorationStuckException
# 注意:日志配置现在由统一的 logging_config 模块管理
# logging.basicConfig(level=logging.INFO, format="%(asctime)s %(name)-12s %(levelname)-8s %(message)s")
# Helper function to get event classes from device
def get_event_class(device, event_type):
"""Get platform-specific event class from device"""
platform_name = device.get_platform_name()
platform = Platform(platform_name)
return PlatformFactory.get_event_class(platform, event_type)
# Policy constants
POLICY_NONE = "none"
POLICY_MEMORY_GUIDED = "memory_guided"
POLICY_MANUAL = "manual"
# Memory-guided policy constants
DEBUG = True
ACTION_INEFFECTIVE = 'ineffective'
CLOSER_ACTION_ENCOURAGEMENT = 0.01
RANDOM_EXPLORE_PROB = 0.4
N_ACTIONS_TRAINING = 32
MAX_NAV_STEPS = 10
# ==================== Neural Network Models ====================
class TextEncoder:
"""Text encoder using BERT or spacy for view text embedding."""
def __init__(self, method='spacy'):
self.method = method
self.embed_size = -1
self._initialized = False
self._nlp = None
self._tokenizer = None
self._text_encoder = None
if method == 'spacy':
self.embed_size = 300
if method == 'bert':
self.embed_size = 768
def _initialize(self):
"""Lazy initialization of the encoder models."""
if self._initialized:
return
if self.method == 'spacy':
import spacy
self._nlp = spacy.load("en_core_web_md")
if self.method == 'bert':
import os
from transformers import BertTokenizer, BertModel
local_model_path = os.path.join(os.path.dirname(__file__), 'cv', 'huggingface_models')
self._tokenizer = BertTokenizer.from_pretrained(local_model_path)
self._text_encoder = BertModel.from_pretrained(local_model_path)
self._initialized = True
def encode(self, text):
self._initialize()
if not text:
return np.zeros(self.embed_size)
if self.method == 'spacy':
doc = self._nlp(text)
return doc.vector
if self.method == 'bert':
encoding = self._tokenizer([text], return_tensors='pt', padding=True, truncation=True)
input_ids = encoding['input_ids']
attention_mask = encoding['attention_mask']
text_encoder_out = self._text_encoder(input_ids, attention_mask=attention_mask)
text_emb = text_encoder_out['pooler_output'][0]
return text_emb.detach().cpu().numpy()
class BertLayerNorm(nn.Module):
"""TF-style LayerNorm."""
def __init__(self, hidden_size, eps=1e-5):
super(BertLayerNorm, self).__init__()
self.weight = nn.Parameter(torch.ones(hidden_size))
self.bias = nn.Parameter(torch.zeros(hidden_size))
self.variance_epsilon = eps
def forward(self, x):
u = x.mean(-1, keepdim=True)
s = (x - u).pow(2).mean(-1, keepdim=True)
x = (x - u) / torch.sqrt(s + self.variance_epsilon)
return self.weight * x + self.bias
class AbsolutePositionalEncoding(nn.Module):
"""Absolute positional encoding for UI elements."""
def __init__(self, d_model, pos_max=128):
super().__init__()
self.pos_max = pos_max
self.d_model = d_model
nhid = d_model
self.x_position_embeddings = nn.Embedding(self.pos_max, nhid)
self.y_position_embeddings = nn.Embedding(self.pos_max, nhid)
self.h_position_embeddings = nn.Embedding(self.pos_max, nhid)
self.w_position_embeddings = nn.Embedding(self.pos_max, nhid)
def forward(self, pos_enc):
l_emb = self.x_position_embeddings(pos_enc[:, 0])
r_emb = self.x_position_embeddings(pos_enc[:, 1])
t_emb = self.y_position_embeddings(pos_enc[:, 2])
b_emb = self.y_position_embeddings(pos_enc[:, 3])
w_emb = self.w_position_embeddings(pos_enc[:, 4])
h_emb = self.h_position_embeddings(pos_enc[:, 5])
pos_emb = l_emb + r_emb + t_emb + b_emb + w_emb + h_emb
return pos_emb
class UIEmbedTransformer(nn.Module):
"""Transformer-based UI element embedder."""
def __init__(self, nhid=64, nhead=2, nlayers=2, dropout=0.8):
super().__init__()
from torch.nn import TransformerEncoder, TransformerEncoderLayer
self.pos_max = 128
self.text_encoder = TextEncoder(method='bert')
dim_feedforward = 256
encoder_layers = TransformerEncoderLayer(nhid, nhead, dim_feedforward, dropout)
self.transformer_encoder = TransformerEncoder(encoder_layers, nlayers)
self.meta2hid = nn.Linear(12, nhid)
self.text2hid = nn.Linear(self.text_encoder.embed_size, nhid)
self.pos2hid = AbsolutePositionalEncoding(d_model=nhid, pos_max=self.pos_max)
self.layer_norm = BertLayerNorm(nhid)
self.dropout = nn.Dropout(dropout)
def forward(self, state_encs):
state_encs, attn_mask = self.encode_state_batch(state_encs)
output = self.transformer_encoder(state_encs, src_key_padding_mask=attn_mask)
output = output.permute(1, 0, 2)
return output
def encode_state(self, state, views):
meta_enc = torch.stack([self._encode_view_meta(state, view) for view in views])
pos_enc = torch.stack([self._encode_view_pos(state, view) for view in views])
text_enc = torch.stack([self._encode_view_text(state, view) for view in views])
return meta_enc, pos_enc, text_enc
def encode_state_batch(self, state_encs):
embs = []
for state_enc in state_encs:
meta_enc, pos_enc, text_enc = state_enc
meta_emb = self.meta2hid(meta_enc)
pos_emb = self.pos2hid(pos_enc)
text_emb = self.text2hid(text_enc)
emb = meta_emb + pos_emb + text_emb
emb = self.layer_norm(emb)
emb = self.dropout(emb)
embs.append(emb)
embs_pad = pad_sequence(embs, batch_first=False)
attn_mask = embs_pad.sum(axis=2).t() == 0
return embs_pad, attn_mask
def _encode_view_meta(self, state, view):
view_children = view['children'] if 'children' in view else []
is_parent = 1 if len(view_children) > 0 else -1
view_text = view['text'] if 'text' in view else None
is_text = 1 if view_text and len(view_text) > 0 else -1
enabled = 1 if 'enabled' in view and view['enabled'] else -1
visible = 1 if 'visible' in view and view['visible'] else -1
clickable = 1 if 'clickable' in view and view['clickable'] else -1
long_clickable = 1 if 'long_clickable' in view and view['long_clickable'] else -1
checkable = 1 if 'checkable' in view and view['checkable'] else -1
checked = 1 if 'checked' in view and view['checked'] else -1
selected = 1 if 'selected' in view and view['selected'] else -1
editable = 1 if 'editable' in view and view['editable'] else -1
is_password = 1 if 'is_password' in view and view['is_password'] else -1
scrollable = 1 if 'scrollable' in view and view['scrollable'] else -1
meta_enc = np.array([
is_parent, is_text, is_password, visible,
enabled, checked, selected,
clickable, long_clickable, checkable, editable, scrollable
])
return torch.Tensor(meta_enc)
def _encode_view_pos(self, state, view):
screen_w = state.width
screen_h = state.height
[[l,t], [r,b]] = view['bounds'] if 'bounds' in view else [[0,0], [0,0]]
l, r, t, b = l/screen_w, r/screen_w, t/screen_h, b/screen_h
if l > r:
l, r = r, l
if t > b:
t, b = b, t
l = max(0, min(1, l))
r = max(0, min(1, r))
t = max(0, min(1, t))
b = max(0, min(1, b))
pos_max = self.pos_max - 1
l, r, t, b = int(pos_max*l), int(pos_max*r), int(pos_max*t), int(pos_max*b)
w = abs(l - r)
h = abs(t - b)
return torch.LongTensor(np.array([l, r, t, b, w, h]))
def _encode_view_text(self, state, view):
view_text = view['text'] if 'text' in view else None
emb = self.text_encoder.encode(view_text)
return torch.Tensor(emb)
# ==================== Memory Class ====================
class Memory:
"""Memory for storing and learning from UI transitions."""
def __init__(self, utg, device):
self.logger = logging.getLogger(self.__class__.__name__)
self.utg = utg
self.device = device
self.known_states = collections.OrderedDict()
self.known_transitions = collections.OrderedDict()
self.known_structures = collections.OrderedDict()
self.model = UIEmbedTransformer()
def _memorize_state(self, state):
if not self.device.is_foreground():
return None
if state.state_str not in self.known_states:
views = state.views
views_str = [view['view_str'] for view in views]
state_enc = self.model.encode_state(state, views)
embedder = self.model
embedder.eval()
with torch.no_grad():
views_emb = self.model.forward([state_enc])
views_emb = views_emb.detach().cpu()[0]
self.known_states[state.state_str] = {
'state': state,
'views': views,
'views_str': views_str,
'state_enc': state_enc,
'views_emb': views_emb
}
return self.known_states[state.state_str]
def save_transition(self, action, from_state, to_state):
if not from_state or not to_state:
return
from_state_info = self._memorize_state(from_state)
if isinstance(action, BaseKillAppEvent) or isinstance(action, BaseKeyEvent) or isinstance(action, BaseSetTextEvent):
return
if not hasattr(action, 'view') or action.view is None:
return
action_str = action.get_event_str(state=from_state)
if action_str in self.known_transitions and self.known_transitions[action_str]['to_state'] == to_state:
return
if from_state_info is None:
return
view = action.view
if view['view_str'] not in from_state_info['views_str']:
self.logger.warning(f"View {view['view_str']} not found in from_state's views")
return
view_idx = from_state_info['views_str'].index(view['view_str'])
action_target = ACTION_INEFFECTIVE \
if from_state.structure_str == to_state.structure_str \
else to_state.structure_str
action_effect = action_target
self.known_transitions[action_str] = {
'from_state': from_state,
'to_state': to_state,
'action': action,
'view_idx': view_idx,
'action_effect': action_effect
}
def save_structure(self, state):
structure_str = state.structure_str
is_new_structure = False
if structure_str not in self.known_structures:
self.known_structures[structure_str] = []
is_new_structure = True
self.known_structures[structure_str].append(state)
return is_new_structure
def _select_transitions_for_training(self, size):
if len(self.known_transitions) <= size:
return list(self.known_transitions.keys())
effect2actions = {}
for k,v in self.known_transitions.items():
action_effect = v['action_effect']
if action_effect not in effect2actions:
effect2actions[action_effect] = []
effect2actions[action_effect].append(k)
action_strs = []
action_probs = []
prob_per_effect = 1.0 / len(effect2actions)
for effect in effect2actions:
prob_per_action = prob_per_effect / len(effect2actions[effect])
for action_str in effect2actions[effect]:
action_strs.append(action_str)
action_probs.append(prob_per_action)
action_probs = np.array(action_probs) / sum(action_probs)
selected_actions = np.random.choice(action_strs, size=size, replace=False, p=action_probs)
return selected_actions
def encode_action_pairs(self, action_strs=None):
if action_strs is None:
action_strs = list(self.known_transitions.keys())
state_strs = [self.known_transitions[action_str]['from_state'].state_str for action_str in action_strs]
state_encs = [self.known_states[state_str]['state_enc'] for state_str in state_strs]
action_pairs = []
for i, action_str1 in enumerate(action_strs):
state_str1 = self.known_transitions[action_str1]['from_state'].state_str
state_idx1 = state_strs.index(state_str1)
view_idx1 = self.known_transitions[action_str1]['view_idx']
for j, action_str2 in enumerate(action_strs[i+1:]):
state_str2 = self.known_transitions[action_str2]['from_state'].state_str
state_idx2 = state_strs.index(state_str2)
view_idx2 = self.known_transitions[action_str2]['view_idx']
action_effect1 = self.known_transitions[action_str1]['action_effect']
action_effect2 = self.known_transitions[action_str2]['action_effect']
effect_same = 1 if action_effect1 == action_effect2 else 0
action_pairs.append((state_idx1, view_idx1, state_idx2, view_idx2, effect_same))
return state_encs, action_pairs
def get_known_actions_emb(self):
actions_emb = []
for action_str in self.known_transitions:
action_info = self.known_transitions[action_str]
state_str = action_info['from_state'].state_str
view_idx = action_info['view_idx']
if state_str not in self.known_states:
continue
action_emb = self.known_states[state_str]['views_emb'][view_idx]
actions_emb.append(action_emb)
return torch.stack(actions_emb) if len(actions_emb) > 0 else None
@staticmethod
def action_info_str(action_info):
state_activity = action_info['from_state'].foreground_page
view_sig = action_info['action'].view.get('signature', '')
action_effect = action_info['action_effect']
return f'{state_activity}-{view_sig}-{action_effect}'
def train_model(self):
if len(self.known_transitions.keys()) < 2:
return
embedder = self.model
optimizer = torch.optim.Adam(embedder.parameters(), lr=1e-3)
n_iterations = 10
def compute_loss(ele_embed, action_pairs):
pos_emb_u, pos_emb_v = [], []
neg_emb_u, neg_emb_v = [], []
for state_idx1, view_idx1, state_idx2, view_idx2, effect_same in action_pairs:
emb_u = ele_embed[state_idx1, view_idx1]
emb_v = ele_embed[state_idx2, view_idx2]
if effect_same:
pos_emb_u.append(emb_u)
pos_emb_v.append(emb_v)
else:
neg_emb_u.append(emb_u)
neg_emb_v.append(emb_v)
pos_score, neg_score = 0, 0
if len(pos_emb_u) > 0 and len(pos_emb_v) > 0:
pos_emb_u = torch.stack(pos_emb_u)
pos_emb_v = torch.stack(pos_emb_v)
pos_score = torch.cosine_similarity(pos_emb_u, pos_emb_v)
pos_score = F.logsigmoid(pos_score).mean()
if len(neg_emb_u) > 0 and len(neg_emb_v) > 0:
neg_emb_u = torch.stack(neg_emb_u)
neg_emb_v = torch.stack(neg_emb_v)
neg_score = torch.cosine_similarity(neg_emb_u, neg_emb_v)
neg_score = F.logsigmoid(-neg_score).mean()
return -pos_score - neg_score
def train():
embedder.train()
action_strs = self._select_transitions_for_training(size=N_ACTIONS_TRAINING)
state_encs, action_pairs = self.encode_action_pairs(action_strs)
ele_embed = embedder.forward(state_encs)
loss = compute_loss(ele_embed, action_pairs)
optimizer.zero_grad()
loss.backward()
optimizer.step()
return loss.item(), len(action_pairs)
for i in range(n_iterations):
epoch_start_time = time.time()
loss, n_pairs = train()
elapsed = time.time() - epoch_start_time
print(f'| iter: {i:3d} | time: {elapsed:6.2f}s | #pairs: {n_pairs:6d} | loss: {loss:8.4f}')
with torch.no_grad():
embedder.eval()
state_encs = [v['state_enc'] for k,v in self.known_states.items()]
ele_embed = embedder(state_encs)
ele_embed = ele_embed.detach().cpu()
for i, (k,v) in enumerate(self.known_states.items()):
self.known_states[k]['views_emb'] = ele_embed[i]
def get_unexplored_actions(self, current_state):
action_strs = set()
structure_strs = set()
self._memorize_state(current_state)
for state_str, state_info in reversed(self.known_states.items()):
state = state_info['state']
if state.structure_str in structure_strs:
continue
structure_strs.add(state.structure_str)
for action in state.get_possible_input():
if not isinstance(action, BaseTouchEvent):
continue
action_str = action.get_event_str(state=state)
if action_str in action_strs:
continue
if self.utg.is_event_explored(action, state):
continue
action_strs.add(action_str)
yield state, action
def get_action_emb(self, state, action):
state_str = state.state_str
if state_str not in self.known_states:
return None
view_str = action.view['view_str']
if view_str not in self.known_states[state_str]['views_str']:
return None
view_idx = self.known_states[state_str]['views_str'].index(view_str)
return self.known_states[state_str]['views_emb'][view_idx]
# ==================== Input Policy Classes ====================
class InputPolicy(object):
"""Base class for input policies."""
def __init__(self, device,enable_guiagent=False):
self.logger = logging.getLogger(self.__class__.__name__)
self.device = device
self.action_count = 0
self.enable_guiagent = enable_guiagent
def start(self, input_manager):
"""Start producing events."""
self.input_manager = input_manager # Save reference for subclasses
# 传递input_manager引用到guiagent_bridge如果存在
if hasattr(self, 'guiagent_bridge') and self.guiagent_bridge:
self.guiagent_bridge.input_manager = input_manager
self.action_count = 0
# === 启动序列 (循环外,只执行一次) ===
# 1. 杀掉旧进程
KillAppEvent = get_event_class(self.device, 'kill_app')
kill_event = KillAppEvent(app=self.device.app_identifier)
input_manager.add_event(kill_event)
self.action_count += 1
# 2. 启动应用
self.device.start_app()
# 3. cv模式 执行初始化任务 (登录/进入游戏)
if getattr(self.device, 'cv_mode', False) and self.enable_guiagent:
if hasattr(self.device, 'run_initial_setup'):
self.device.run_initial_setup()
# === 正常探索循环 ===
while input_manager.enabled and self.action_count < input_manager.event_count:
try:
event = self.generate_event()
if event is not None:
input_manager.add_event(event)
self.action_count += 1
except InputInterruptedException:
self.logger.info("InputInterruptedException caught, stopping event generation.")
break
except FATAL_EXCEPTIONS:
raise
except Exception as e:
self.logger.error(f'Non-fatal exception in event loop: {e}')
import traceback
traceback.print_exc()
continue
def generate_event(self):
"""Generate an event."""
raise NotImplementedError
class NoneInputPolicy(InputPolicy):
"""Do not send any event - for manual testing."""
def __init__(self, device, enable_guiagent=False):
super(NoneInputPolicy, self).__init__(device)
self.enable_guiagent = enable_guiagent
def generate_event(self):
"""Generate an event."""
import time
time.sleep(2)
return None
class ManualPolicy(InputPolicy):
"""
Manually control the device.
"""
def __init__(self, device, enable_guiagent=False):
super(ManualPolicy, self).__init__(device)
self.enable_guiagent = enable_guiagent
def generate_event(self):
"""
No event is generated in manual policy
"""
return None
class UtgBasedInputPolicy(InputPolicy):
"""State-based input policy using UTG."""
def __init__(self, device, random_input, enable_guiagent=False, app_name=None, traffic_monitor=None):
super(UtgBasedInputPolicy, self).__init__(device,enable_guiagent=enable_guiagent)
self.random_input = random_input
self.enable_guiagent = enable_guiagent
self.app_name = app_name
self.traffic_monitor = traffic_monitor # 新增:流量监控引用
self.last_event = None
self.last_state = None
self.current_state = None
self.utg = UTG(device=device, random_input=random_input)
# Initialize GuiAgent bridge
self.guiagent_bridge = None
if self.enable_guiagent:
self.guiagent_bridge = GuiAgentBridge(device, app_name=self.app_name, utg=self.utg)
self.logger.info("GuiAgent bridge initialized")
def generate_event(self):
"""Generate an event."""
# 复用上一步 EventLog.stop() 缓存的状态,避免重复调用 get_current_state
if hasattr(self.device, '_last_state') and self.device._last_state is not None:
self.current_state = self.device._last_state
else:
print("Warning: No last state available, using current state as start state.")
self.current_state = self.device.get_current_state()
# 首次调用时设置 _last_state供 EventLog.start() 复用
if hasattr(self.device, '_last_state'):
self.device._last_state = self.current_state
if self.current_state is None:
import time
time.sleep(5)
self.logger.warning("Current state is None, waiting for 5 seconds")
KeyEvent = get_event_class(self.device, 'key')
return KeyEvent(key_name="BACK")
self.__update_utg()
event = self.generate_event_based_on_utg()
self.last_event = event
self.last_state = self.current_state
return event
def __update_utg(self):
# 在调用 UTG 方法之前,一次性获取新增域名
# 前10步不检测域名等待流量文件生成避免获取上一次测试的流量
new_domains_count = 0
new_domains_list = []
EARLY_PHASE_STEPS = 10
if self.traffic_monitor and self.action_count > EARLY_PHASE_STEPS:
new_domains_count, new_domains_list = self.traffic_monitor.get_new_domains_since_last_step()
if new_domains_count > 0:
self.logger.info(f"Step found {new_domains_count} new domains")
# 将域名信息传递给 UTG
self.utg.add_transition(self.last_event, self.last_state, self.current_state, new_domains=new_domains_list)
def generate_event_based_on_utg(self):
"""Generate an event based on UTG - to be overridden."""
raise NotImplementedError
class MemoryGuidedPolicy(UtgBasedInputPolicy):
"""Memory-guided exploration policy with neural network learning."""
def __init__(self, device, random_input, enable_guiagent=False, app_name=None, traffic_monitor=None):
super(MemoryGuidedPolicy, self).__init__(device, random_input, enable_guiagent=enable_guiagent, app_name=app_name, traffic_monitor=traffic_monitor)
self.logger = logging.getLogger(self.__class__.__name__)
self.memory = Memory(utg=self.utg, device=device)
self.num_actions_train = 10
self._nav_steps = []
self.last_login_register_step = -100 # Initialize to allow immediate trigger
self.guiagent_message = None # 记录登录/注册消息成功或失败用于worker.report汇报
self.login_count = 0 # 记录进入登录场景的次数
self.register_count = 0 # 记录进入注册场景的次数
self.stuck_reason_code = 0 # 记录卡住原因代码用于worker.report汇报
# 新状态停滞检测变量
self._last_state_count = 0
self._no_new_state_steps = 0
self.NO_NEW_STATE_THRESHOLD = 50
self.NO_NEW_STATE_STEPS_THRESHOLD = 300
def generate_event_based_on_utg(self):
"""Generate an event based on current UTG."""
current_state = self.current_state
try:
self.memory.save_transition(self.last_event, self.last_state, current_state)
except FATAL_EXCEPTIONS:
raise
except Exception as e:
self.logger.error(f'failed to save transition: {e}')
import traceback
traceback.print_exc()
# 非致命异常,不抛出
if self.action_count % self.num_actions_train == 0:
self.memory.train_model()
self.logger.debug("Current state: %s" % current_state.state_str)
# 检测连续无新增状态
total_steps = self.input_manager.total_exploring_steps
if self._check_no_new_state_stuck():
self.logger.warning(f"连续 {self._no_new_state_steps} 步无新增状态(总步数: {total_steps}")
# 在阈值以内,尝试智能体介入
if self.enable_guiagent and total_steps < self.NO_NEW_STATE_STEPS_THRESHOLD:
self.logger.info(f"触发 GuiAgent 处理探索停滞")
success, message, reason_code = self.guiagent_bridge.handle_with_guiagent(
category="explore_stuck",
context={"state": current_state, "steps": self._no_new_state_steps, "total_steps": total_steps}
)
if message:
if self.guiagent_message:
self.guiagent_message += f"; {message}"
else:
self.guiagent_message = message
if reason_code:
self.stuck_reason_code = reason_code
if message and ("失败" in message):
self.logger.error(f"GuiAgent 判断探索无法继续: {message}")
raise ExplorationStuckException(f"{message}")
elif message and ("成功" in message):
self._no_new_state_steps = 0
self._last_state_count = sum(
1 for state_str in self.utg.G.nodes()
if self.utg.G.nodes[state_str]["state"].foreground_page
and self.utg.G.nodes[state_str]["state"].foreground_page.startswith(self.device.app_identifier)
)
self.logger.info(f"重置卡住检测计数器,继续探索")
return None
else:
# 超过阈值,直接抛出异常
self.logger.error(f"总步数 {total_steps} 超过阈值 {self.NO_NEW_STATE_STEPS_THRESHOLD},探索停滞")
raise InputInterruptedException(f"总步数 {total_steps} 超过阈值 {self.NO_NEW_STATE_STEPS_THRESHOLD},探索停滞")
nav_action = self.navigate(current_state)
if nav_action:
return nav_action
self._nav_steps = []
self.memory.save_structure(current_state)
if self.action_count >= self.num_actions_train \
and len(self._nav_steps) == 0 \
and np.random.uniform() > RANDOM_EXPLORE_PROB:
(target_state, target_action), candidates = self.pick_target(current_state)
if target_state:
if target_state.state_str == current_state.state_str:
self.logger.info(f"executing action selected from {len(candidates)} candidates")
# Check GuiAgent keywords before returning
if self.enable_guiagent :
if hasattr(target_action, 'view') and target_action.view is not None:
view = target_action.view
view_text = self.guiagent_bridge.get_text_within_bounds(view, current_state.views)
if view_text:
is_editable = view.get('editable', False)
if is_editable:
self.logger.info(f"即将操作的控件是可输入的: {view_text}. 执行自定义输入序列。")
SetTextEvent = get_event_class(self.device, 'set_text')
return SetTextEvent(text='tp', view=view)
category = self.guiagent_bridge.check_keywords(view_text)
if category:
# Special handling for login/register cooldown
if category in ["login", "register"]:
current_step = self.input_manager.total_exploring_steps if hasattr(self, 'input_manager') else 0
if current_step - self.last_login_register_step < 50:
self.logger.info(f"Skipping {category} due to cooldown (last: {self.last_login_register_step}, current: {current_step})")
return target_action
else:
self.last_login_register_step = current_step
self.logger.info(f"即将操作的控件包含关键词: {view_text}, 类别: {category}")
if category == 'login':
self.login_count += 1
self.logger.info(f"进入登录场景,累计次数: {self.login_count}")
elif category == 'register':
self.register_count += 1
self.logger.info(f"进入注册场景,累计次数: {self.register_count}")
success, guiagent_message, _ = self.guiagent_bridge.handle_with_guiagent(category, {"view": view, "state": current_state})
if guiagent_message:
if self.guiagent_message:
self.guiagent_message += f"; {guiagent_message}"
else:
self.guiagent_message = guiagent_message
if success:
new_state = self.device.get_current_state()
if new_state and new_state.state_str != current_state.state_str:
self.logger.info("GuiAgent处理成功状态已改变")
return None
else:
self.logger.warning("GuiAgent处理失败继续使用DroidBot策略执行原操作")
return target_action
self._nav_steps = self.get_shortest_nav_steps(current_state, target_state, target_action)
nav_action = self.navigate(current_state)
if nav_action:
return nav_action
self._nav_steps = []
self.logger.info("trying random action")
possible_events = current_state.get_possible_input()
# self.logger.info(possible_events)
random.shuffle(possible_events)
# 当系统界面等无可交互元素时possible_events 可能为空
if not possible_events:
self.logger.warning("possible_events 为空,执行 BACK 返回")
KeyEvent = get_event_class(self.device, 'key')
return KeyEvent(key_name="BACK")
selected_event = possible_events[0]
if self.enable_guiagent:
if hasattr(selected_event, 'view') and selected_event.view is not None:
view = selected_event.view
view_text = self.guiagent_bridge.get_text_within_bounds(view, current_state.views)
if view_text:
is_editable = view.get('editable', False)
if is_editable:
self.logger.info(f"即将操作的控件是可输入的: {view_text}. 执行自定义输入序列。")
SetTextEvent = get_event_class(self.device, 'set_text')
return SetTextEvent(text='test', view=view)
category = self.guiagent_bridge.check_keywords(view_text)
if category:
# Special handling for login/register cooldown
if category in ["login", "register"]:
current_step = self.input_manager.total_exploring_steps if hasattr(self, 'input_manager') else 0
if current_step - self.last_login_register_step < 100:
self.logger.info(f"Skipping {category} due to cooldown (last: {self.last_login_register_step}, current: {current_step})")
return selected_event
else:
self.last_login_register_step = current_step
self.logger.info(f"即将操作的控件包含关键词: {view_text}, 类别: {category}")
if category == 'login':
self.login_count += 1
self.logger.info(f"进入登录场景,累计次数: {self.login_count}")
elif category == 'register':
self.register_count += 1
self.logger.info(f"进入注册场景,累计次数: {self.register_count}")
success, guiagent_message, _ = self.guiagent_bridge.handle_with_guiagent(category, {"view": view, "state": current_state})
if guiagent_message:
if self.guiagent_message:
self.guiagent_message += f"; {guiagent_message}"
else:
self.guiagent_message = guiagent_message
if success:
new_state = self.device.get_current_state()
if new_state and new_state.state_str != current_state.state_str:
self.logger.info("GuiAgent处理成功状态已改变")
return None
else:
self.logger.warning("GuiAgent处理失败继续使用DroidBot策略执行原操作")
return selected_event
def pick_target(self, current_state):
state_action_pairs = list(self.memory.get_unexplored_actions(current_state))
best_target = None, None
best_score = -np.inf
known_actions_emb = self.memory.get_known_actions_emb()
if known_actions_emb is None:
return best_target, state_action_pairs
scores = []
for i, (state, action) in enumerate(state_action_pairs):
action_emb = self.memory.get_action_emb(state, action)
similarities = torch.cosine_similarity(action_emb.repeat((known_actions_emb.size(0), 1)), known_actions_emb)
max_sim, max_sim_idx = similarities.max(0)
score = -max_sim
if state.state_str == current_state.state_str:
score += CLOSER_ACTION_ENCOURAGEMENT
if DEBUG:
action_info_str = f'{state.foreground_page}-{action.view.get("signature", "")}'
scores.append((i, score, action_info_str, similarities, action_emb))
if score > best_score:
best_score = score
best_target = state, action
return best_target, state_action_pairs
def _check_no_new_state_stuck(self):
"""
检测是否连续多步无新增状态
只统计属于当前应用的状态foreground_page 以 app_identifier 开头)
:return: True 如果连续 NO_NEW_STATE_THRESHOLD 步无新增状态
"""
app_identifier = self.device.app_identifier
current_state_count = sum(
1 for state_str in self.utg.G.nodes()
if self.utg.G.nodes[state_str]["state"].foreground_page
and self.utg.G.nodes[state_str]["state"].foreground_page.startswith(app_identifier)
)
if current_state_count > self._last_state_count:
self._last_state_count = current_state_count
self._no_new_state_steps = 0
self.logger.debug(f"发现新状态,应用内状态总数: {current_state_count},重置计数器")
else:
self._no_new_state_steps += 1
self.logger.debug(f"无新增状态,连续步数: {self._no_new_state_steps}/{self.NO_NEW_STATE_THRESHOLD}")
return self._no_new_state_steps >= self.NO_NEW_STATE_THRESHOLD
def navigate(self, current_state):
if self._nav_steps and len(self._nav_steps) > 0:
nav_state, nav_action = self._nav_steps[0]
self._nav_steps = self._nav_steps[1:]
nav_action_ = self._get_nav_action(current_state, nav_state, nav_action)
if nav_action_:
self.logger.info(f"navigating, {len(self._nav_steps)} steps left")
return nav_action_
else:
self.logger.warning("navigation failed")
self.utg.remove_transition(self.last_event, self.last_state, nav_state)
def _get_nav_action(self, current_state, nav_state, nav_action):
try:
if current_state.structure_str != nav_state.structure_str:
return None
if not isinstance(nav_action, BaseTouchEvent):
return nav_action
if nav_action.__class__.__name__ == 'ScrollEvent':
return copy.deepcopy(nav_action)
nav_view = nav_action.view
nav_view_idx = nav_state.views.index(nav_view)
new_view = current_state.views[nav_view_idx]
new_action = copy.deepcopy(nav_action)
new_action.view = new_view
return new_action
except Exception as e:
self.logger.error(f'exception during _get_nav_action: {e}')
return nav_action
def get_shortest_nav_steps(self, current_state, target_state, target_action):
normal_nav_steps = self.utg.get_G2_nav_steps(current_state, target_state)
restart_nav_steps = self.utg.get_G2_nav_steps(self.utg.first_state, target_state)
normal_nav_steps_len = len(normal_nav_steps) if normal_nav_steps else MAX_NAV_STEPS
restart_nav_steps_len = len(restart_nav_steps) + 1 if restart_nav_steps else MAX_NAV_STEPS
if normal_nav_steps_len >= MAX_NAV_STEPS and restart_nav_steps_len >= MAX_NAV_STEPS:
self.logger.warning(f'cannot find a path to {target_state.structure_str} {target_state.foreground_page}')
target_state_str = target_state.state_str
self.memory.known_states.pop(target_state_str, None)
action_strs_to_remove = []
for action_str in self.memory.known_transitions:
action_from_state = self.memory.known_transitions[action_str]['from_state']
action_to_state = self.memory.known_transitions[action_str]['to_state']
if action_from_state.state_str == target_state_str or action_to_state.state_str == target_state_str:
action_strs_to_remove.append(action_str)
for action_str in action_strs_to_remove:
self.memory.known_transitions.pop(action_str, None)
return None
elif normal_nav_steps_len >= MAX_NAV_STEPS:
return None
else:
nav_steps = normal_nav_steps
return nav_steps + [(target_state, target_action)]