""" 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)]