898 lines
41 KiB
Python
898 lines
41 KiB
Python
"""
|
||
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)]
|