428 lines
18 KiB
Python
428 lines
18 KiB
Python
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. 批量保存所有待写入的 state(JSON + 截图)
|
||
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
|
||
|