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

428 lines
18 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.

import logging
import json
import os
import random
import datetime
import networkx as nx
from .core.abstract_input_event import BaseKeyEvent, BaseKillAppEvent,BaseSetTextEvent
class UTG(object):
"""
UI transition graph
"""
def __init__(self, device, random_input):
self.logger = logging.getLogger(self.__class__.__name__)
self.device = device
self.random_input = random_input
self.G = nx.DiGraph()
self.G2 = nx.DiGraph() # graph with same-structure states clustered
self.transitions = []
self.effective_event_strs = set()
self.ineffective_event_strs = set()
self.reached_activities = set()
# Domain tracking data for visualization
self.transition_domains = {} # edge_key -> list of (event_id, new_domains_count)
self.node_domains = {} # state_str -> list of new_domains
self.first_state = None
self.last_state = None
self._pending_states = [] # 延迟保存的 state 列表
self.start_time = datetime.datetime.now()
@property
def first_state_str(self):
return self.first_state.state_str if self.first_state else None
@property
def last_state_str(self):
return self.last_state.state_str if self.last_state else None
@property
def effective_event_count(self):
return len(self.effective_event_strs)
@property
def num_transitions(self):
return len(self.transitions)
def add_transition(self, event, old_state, new_state, new_domains=None, is_guiagent_event=False):
self.add_node(old_state)
# 记录新关联的域名列表到新状态
self.add_node(new_state, new_domains=new_domains)
# make sure the states and event are not None
if not old_state or not new_state or not event:
return
if isinstance(event, BaseKillAppEvent) or isinstance(event, BaseKeyEvent) or isinstance(event, BaseSetTextEvent):
return
event_str = event.get_event_str(old_state)
self.transitions.append((old_state, event, new_state))
if old_state.state_str == new_state.state_str:
self.ineffective_event_strs.add(event_str)
# delete the transitions including the event from utg
for new_state_str in self.G[old_state.state_str]:
if event_str in self.G[old_state.state_str][new_state_str]["events"]:
self.G[old_state.state_str][new_state_str]["events"].pop(event_str)
if event_str in self.effective_event_strs:
self.effective_event_strs.remove(event_str)
return
self.effective_event_strs.add(event_str)
new_domains_count = len(new_domains) if new_domains else 0
if (old_state.state_str, new_state.state_str) not in self.G.edges():
self.G.add_edge(old_state.state_str, new_state.state_str, events={})
self.G[old_state.state_str][new_state.state_str]["events"][event_str] = {
"event": event,
"id": self.effective_event_count,
"new_domains_count": new_domains_count, # 新增:记录该步新增域名数量
"is_guiagent_event": is_guiagent_event # 新增标记是否为GuiAgent事件
}
# 记录 edge 的域名统计
edge_key = f"{old_state.state_str}-->{new_state.state_str}"
if edge_key not in self.transition_domains:
self.transition_domains[edge_key] = []
self.transition_domains[edge_key].append({
"event_id": self.effective_event_count,
"new_domains_count": new_domains_count
})
if (old_state.structure_str, new_state.structure_str) not in self.G2.edges():
self.G2.add_edge(old_state.structure_str, new_state.structure_str, events={})
self.G2[old_state.structure_str][new_state.structure_str]["events"][event_str] = {
"event": event,
"id": self.effective_event_count
}
self.last_state = new_state
def remove_transition(self, event, old_state, new_state):
event_str = event.get_event_str(old_state)
if (old_state.state_str, new_state.state_str) in self.G.edges():
events = self.G[old_state.state_str][new_state.state_str]["events"]
if event_str in events.keys():
events.pop(event_str)
if len(events) == 0:
self.G.remove_edge(old_state.state_str, new_state.state_str)
if (old_state.structure_str, new_state.structure_str) in self.G2.edges():
events = self.G2[old_state.structure_str][new_state.structure_str]["events"]
if event_str in events.keys():
events.pop(event_str)
if len(events) == 0:
self.G2.remove_edge(old_state.structure_str, new_state.structure_str)
def add_node(self, state, new_domains=None):
if not state:
return
if state.state_str not in self.G.nodes():
self._pending_states.append(state)
self.G.add_node(state.state_str, state=state)
if self.first_state is None:
self.first_state = state
# 记录到达该状态时发现的新增域名(支持多次到达时累加)
if new_domains:
if state.state_str not in self.node_domains:
self.node_domains[state.state_str] = []
for d in new_domains:
if d not in self.node_domains[state.state_str]:
self.node_domains[state.state_str].append(d)
self.node_domains[state.state_str].sort()
if state.structure_str not in self.G2.nodes():
self.G2.add_node(state.structure_str, states=[])
self.G2.nodes[state.structure_str]['states'].append(state)
if self.device._app and state.foreground_page and state.foreground_page.startswith(self.device.app_identifier):
self.reached_activities.add(state.foreground_page)
def finalize(self):
"""
探索结束后调用,一次性完成所有文件写入:
1. 批量保存所有待写入的 stateJSON + 截图)
2. 输出 utg.js 及更新 state/event JSON 的域名字段
所有文件操作使用 flush + fsync 确保数据落盘
"""
self.logger.info(f"Finalizing UTG: saving {len(self._pending_states)} states and generating utg.js")
for state in self._pending_states:
state.save2dir()
self._pending_states.clear()
self.__output_utg()
self.logger.info("UTG finalization completed successfully")
def __output_utg(self):
"""
Output current UTG to a js file
"""
if not self.device.output_dir:
return
def list_to_html_table(dict_data):
table = "<table class=\"table\">\n"
for (key, value) in dict_data:
table += "<tr><th>%s</th><td>%s</td></tr>\n" % (key, value)
table += "</table>"
return table
utg_file_path = os.path.join(self.device.output_dir, "utg.js")
utg_nodes = []
utg_edges = []
for state_str in self.G.nodes():
state = self.G.nodes[state_str]["state"]
# Handle None foreground_page gracefully
if state.foreground_page:
package_name = state.foreground_page.split("/")[0]
activity_name = state.foreground_page.split("/")[1] if "/" in state.foreground_page else state.foreground_page
else:
package_name = "unknown"
activity_name = "unknown"
short_activity_name = activity_name.split(".")[-1]
state_desc = list_to_html_table([
("package", package_name),
("activity", activity_name),
("state_str", state.state_str),
("structure_str", state.structure_str)
])
# 防御性处理WDA 超时时 screenshot_path 可能为 None
if state.screenshot_path:
image_path = os.path.relpath(state.screenshot_path, self.device.output_dir).replace('\\', '/')
else:
image_path = ""
utg_node = {
"id": state_str,
"shape": "image",
"image": image_path,
"label": short_activity_name,
# "group": state.foreground_activity,
"package": package_name,
"activity": activity_name,
"state_str": state_str,
"structure_str": state.structure_str,
"title": state_desc,
"content": "\n".join([package_name, activity_name, state.state_str, state.search_content]),
"new_domains": self.node_domains.get(state_str, []) # 新增:该状态的新增域名列表
}
# 对于有新增域名的节点,在标签中显示域名数量
domains_count = len(self.node_domains.get(state_str, []))
if domains_count > 0:
utg_node["label"] += f"\n[+{domains_count} domains]"
utg_node["font"] = "14px Arial green"
if state.state_str == self.first_state_str:
utg_node["label"] += "\n<FIRST>"
utg_node["font"] = "14px Arial red"
if state.state_str == self.last_state_str:
utg_node["label"] += "\n<LAST>"
utg_node["font"] = "14px Arial red"
utg_nodes.append(utg_node)
for state_transition in self.G.edges():
from_state = state_transition[0]
to_state = state_transition[1]
events = self.G[from_state][to_state]["events"]
event_short_descs = []
event_list = []
for event_str, event_info in sorted(iter(events.items()), key=lambda x: x[1]["id"]):
event_short_descs.append((event_info["id"], event_str))
view_images = ["views/view_" + view["view_str"] + ".png"
for view in event_info["event"].get_views()]
event_list.append({
"event_str": event_str,
"event_id": event_info["id"],
"event_type": event_info["event"].event_type.value if hasattr(event_info["event"].event_type, 'value') else str(event_info["event"].event_type),
"view_images": view_images,
"new_domains_count": event_info.get("new_domains_count", 0), # 新增:该事件的新增域名数量
"is_guiagent_event": event_info.get("is_guiagent_event", False) # 新增是否为GuiAgent事件
})
# 检查这条边是否包含GuiAgent事件
has_guiagent_event = any(e.get("is_guiagent_event", False) for e in event_list)
utg_edge = {
"from": from_state,
"to": to_state,
"id": from_state + "-->" + to_state,
"title": list_to_html_table(event_short_descs),
"label": ", ".join([str(x["event_id"]) for x in event_list]),
"events": event_list
}
# GuiAgent事件使用黄色强调
if has_guiagent_event:
utg_edge["color"] = "#FFD700" # 金色/黄色
utg_edge["width"] = 3 # 加粗边
# # Highlight last transition
# if state_transition == self.last_transition:
# utg_edge["color"] = "red"
utg_edges.append(utg_edge)
utg = {
"nodes": utg_nodes,
"edges": utg_edges,
"num_nodes": len(utg_nodes),
"num_edges": len(utg_edges),
"num_effective_events": len(self.effective_event_strs),
"num_reached_activities": len(self.reached_activities),
"test_date": self.start_time.strftime("%Y-%m-%d %H:%M:%S"),
"time_spent": (datetime.datetime.now() - self.start_time).total_seconds(),
"num_transitions": self.num_transitions,
"app_package": self.device.app_identifier,
"app_main_activity": self.device._app.main_activity if self.device._app else "",
"app_num_total_activities": len(self.device._app.activities) if self.device._app else 0,
}
utg_json = json.dumps(utg, indent=2)
with open(utg_file_path, "w") as utg_file:
utg_file.write("var utg = \n")
utg_file.write(utg_json)
utg_file.flush()
os.fsync(utg_file.fileno())
# 更新 states/*.json 文件,添加 new_domains 字段
self.__update_state_json_files()
# 更新 events/*.json 文件,添加 new_domains_count 字段
self.__update_event_json_files()
def __update_state_json_files(self):
"""
更新 states/*.json 文件,为每个状态添加 new_domains 字段
使用 flush + fsync 确保数据落盘
"""
if not self.device.output_dir:
return
states_dir = os.path.join(self.device.output_dir, "states")
if not os.path.exists(states_dir):
return
for state_str in self.G.nodes():
state = self.G.nodes[state_str]["state"]
new_domains = self.node_domains.get(state_str, [])
# 找到对应的 state json 文件
state_json_path = os.path.join(states_dir, f"state_{state.tag}.json")
if os.path.exists(state_json_path):
try:
with open(state_json_path, "r") as f:
state_data = json.load(f)
# 添加 new_domains 字段
state_data["new_domains"] = new_domains
with open(state_json_path, "w") as f:
json.dump(state_data, f, indent=2)
f.flush()
os.fsync(f.fileno())
except Exception as e:
self.logger.error(f"Failed to update state json {state_json_path}: {e}")
def __update_event_json_files(self):
"""
更新 events/*.json 文件,为每个事件添加 new_domains_count 字段
使用 flush + fsync 确保数据落盘
"""
if not self.device.output_dir:
return
events_dir = os.path.join(self.device.output_dir, "events")
if not os.path.exists(events_dir):
return
# 遍历所有边和事件
for state_transition in self.G.edges():
from_state_str = state_transition[0]
to_state_str = state_transition[1]
events = self.G[from_state_str][to_state_str]["events"]
for event_str, event_info in events.items():
event = event_info["event"]
new_domains_count = event_info.get("new_domains_count", 0)
# 获取事件的 tag如果有
if hasattr(event, 'tag'):
event_tag = event.tag
else:
# 尝试从 to_state 获取 tag
to_state = self.G.nodes[to_state_str]["state"]
event_tag = to_state.tag
event_json_path = os.path.join(events_dir, f"event_{event_tag}.json")
if os.path.exists(event_json_path):
try:
with open(event_json_path, "r") as f:
event_data = json.load(f)
# 添加 new_domains_count 字段
event_data["new_domains_count"] = new_domains_count
with open(event_json_path, "w") as f:
json.dump(event_data, f, indent=2)
f.flush()
os.fsync(f.fileno())
except Exception as e:
self.logger.error(f"Failed to update event json {event_json_path}: {e}")
def is_event_explored(self, event, state):
event_str = event.get_event_str(state)
return event_str in self.effective_event_strs or event_str in self.ineffective_event_strs
def get_G2_nav_steps(self, from_state, to_state):
if from_state is None or to_state is None:
return None
from_state_str = from_state.structure_str
to_state_str = to_state.structure_str
try:
nav_steps = []
state_strs = nx.shortest_path(G=self.G2, source=from_state_str, target=to_state_str)
if not isinstance(state_strs, list) or len(state_strs) < 2:
return None
start_state_str = state_strs[0]
for state_str in state_strs[1:]:
edge = self.G2[start_state_str][state_str]
edge_event_strs = list(edge["events"].keys())
start_state = random.choice(self.G2.nodes[start_state_str]['states'])
event_str = random.choice(edge_event_strs)
event = edge["events"][event_str]["event"]
nav_steps.append((start_state, event))
start_state_str = state_str
if nav_steps is None:
return None
# return nav_steps
# simplify the path
simple_nav_steps = []
last_state, last_action = nav_steps[-1]
for state, action in nav_steps:
if state.structure_str == last_state.structure_str:
simple_nav_steps.append((state, last_action))
break
simple_nav_steps.append((state, action))
return simple_nav_steps
except nx.NetworkXNoPath:
self.logger.debug(f"No path between {from_state_str[:32]} and {to_state_str[:32]}")
return None
except Exception as e:
self.logger.error(f"Failed to get G2 nav steps: {e}")
return None