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 = "
\n"
for (key, value) in dict_data:
table += "| %s | %s |
\n" % (key, value)
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"
utg_node["font"] = "14px Arial red"
if state.state_str == self.last_state_str:
utg_node["label"] += "\n"
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