ADK-agents/watch_session.py

271 lines
8.9 KiB
Python
Raw Permalink 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.

"""
Session 实时监控脚本
输入 session_id实时打印该会话中的所有新事件用户输入、模型回复、工具调用等
用法:
python watch_session.py --session <session_id>
python watch_session.py -s <session_id> --agent my_agent
python watch_session.py -s test_001 --poll 2.0
支持的参数:
--session / -s : 会话 ID必填
--agent / -a : Agent 名称,默认 my_agent可选my_agent / luna_agent / qwen_agent
--user / -u : 用户 ID默认 codebuddy
--poll / -p : 轮询间隔(秒),默认 1.5
--url : API Server 地址,默认根据 agent 自动选择
"""
import argparse
import json
import sys
import time
from datetime import datetime
import httpx
# Agent 对应的默认 API 地址
AGENT_URLS = {
"my_agent": "http://127.0.0.1:8001",
"luna_agent": "http://127.0.0.1:8002",
"qwen_agent": "http://127.0.0.1:8003",
}
# Agent 名称别名
AGENT_ALIASES = {
"my": "my_agent",
"default": "my_agent",
"aq": "my_agent",
"luna": "luna_agent",
"gpt": "luna_agent",
"qwen": "qwen_agent",
"astron": "qwen_agent",
}
def resolve_agent(name: str) -> str:
"""解析 agent 名称"""
name = name.strip().lower()
if name in AGENT_URLS:
return name
if name in AGENT_ALIASES:
return AGENT_ALIASES[name]
for full_name in AGENT_URLS:
if name in full_name:
return full_name
raise ValueError(
f"未知的 agent: {name}\n"
f"可用: {list(AGENT_URLS.keys())}\n"
f"别名: {list(AGENT_ALIASES.keys())}"
)
def get_api_url(agent_name: str, custom_url: str | None) -> str:
"""获取 API 地址"""
if custom_url:
return custom_url.rstrip("/")
return AGENT_URLS[agent_name]
def fetch_session(api_url: str, app_name: str, user_id: str, session_id: str) -> dict | None:
"""获取会话数据"""
try:
resp = httpx.get(
f"{api_url}/apps/{app_name}/users/{user_id}/sessions/{session_id}",
timeout=10.0,
)
if resp.status_code == 200:
return resp.json()
if resp.status_code == 404:
return None
print(f"[警告] 获取会话失败 (HTTP {resp.status_code}): {resp.text[:200]}")
return None
except Exception as e:
print(f"[警告] 连接 API Server 失败: {e}")
return None
def format_event(event: dict, index: int) -> str:
"""格式化单个事件为可读字符串"""
content = event.get("content", {})
role = content.get("role", "?")
author = event.get("author", "")
parts = content.get("parts", [])
timestamp = event.get("timestamp", 0)
time_str = ""
if timestamp:
try:
time_str = datetime.fromtimestamp(timestamp).strftime("%H:%M:%S")
except Exception:
time_str = str(timestamp)
role_label = {
"user": "👤 用户",
"model": "🤖 模型",
"function": "🔧 工具",
}.get(role, f"{role}")
author_str = f" [{author}]" if author else ""
header = f"\n{'' * 60}\n[{time_str}] {role_label}{author_str} #{index}\n{'' * 60}"
lines = [header]
for part in parts:
if "text" in part:
text = part["text"]
# thoughts 单独标注
if part.get("thought"):
lines.append(f"💭 [思考中]\n{text}\n")
else:
lines.append(f"{text}\n")
elif "functionCall" in part:
call = part["functionCall"]
args_str = json.dumps(call.get("args", {}), ensure_ascii=False, indent=2)
# 太长就截断
if len(args_str) > 500:
args_str = args_str[:500] + f"\n... (共 {len(args_str)} 字符,已截断)"
lines.append(f"📞 调用工具: {call.get('name', '?')}\n{args_str}\n")
elif "functionResponse" in part:
resp = part["functionResponse"]
resp_name = resp.get("name", "?")
resp_content = resp.get("content", [])
# 提取文本内容
text_parts = []
for c in resp_content:
if isinstance(c, dict) and c.get("type") == "text":
text_parts.append(c.get("text", ""))
elif isinstance(c, str):
text_parts.append(c)
result_text = "\n".join(text_parts) if text_parts else str(resp_content)
# 太长就截断
if len(result_text) > 800:
result_text = result_text[:800] + f"\n... (共 {len(result_text)} 字符,已截断)"
lines.append(f"✅ 工具返回: {resp_name}\n{result_text}\n")
elif "code" in part:
code = part["code"]
lines.append(f"📝 代码片段:\n```\n{code}\n```\n")
elif "executableCode" in part:
ec = part["executableCode"]
lines.append(f"💻 可执行代码 ({ec.get('language', '?')}):\n```\n{ec.get('code', '')[:500]}\n```\n")
else:
part_types = list(part.keys())
lines.append(f"[其他内容] 类型: {part_types}\n")
return "\n".join(lines)
def extract_events(session_data: dict) -> list[dict]:
"""从会话数据中提取事件列表"""
return session_data.get("events", []) or []
def watch_session(
api_url: str,
app_name: str,
user_id: str,
session_id: str,
poll_interval: float,
):
"""实时监控会话"""
print(f"🔍 开始监控会话")
print(f" Agent: {app_name}")
print(f" API: {api_url}")
print(f" 用户: {user_id}")
print(f" 会话ID: {session_id}")
print(f" 轮询间隔: {poll_interval}s")
print(f" 按 Ctrl+C 退出\n")
last_event_count = 0
# 首次获取,如果有历史事件,问要不要回放
session = fetch_session(api_url, app_name, user_id, session_id)
if session is None:
print(f"会话 [{session_id}] 不存在,请检查 session_id 和 agent 是否正确。")
print(f"提示: 确认 {app_name} 的 API Server 是否已启动({api_url}")
return
events = extract_events(session)
existing_count = len(events)
if existing_count > 0:
print(f"📜 该会话已有 {existing_count} 条历史事件。")
try:
choice = input("是否打印历史事件?(y/n默认 n): ").strip().lower()
except (EOFError, KeyboardInterrupt):
print("\n已退出。")
return
if choice in ("y", "yes"):
for i, event in enumerate(events, 1):
print(format_event(event, i))
last_event_count = existing_count
print(f"\n✅ 历史事件回放完毕,共 {existing_count} 条。")
print(f" 现在开始监控新事件...\n")
else:
last_event_count = existing_count
print(f" 跳过历史,从第 {existing_count + 1} 条开始监控新事件...\n")
else:
print("📭 该会话目前没有事件,等待新事件...\n")
# 开始轮询
try:
while True:
time.sleep(poll_interval)
session = fetch_session(api_url, app_name, user_id, session_id)
if session is None:
continue
events = extract_events(session)
current_count = len(events)
if current_count > last_event_count:
# 有新事件
for i in range(last_event_count, current_count):
print(format_event(events[i], i + 1))
last_event_count = current_count
# 检测是否结束(最后一条是 model role 的 final 事件)
# 这里不自动退出,继续轮询,因为可能有多轮对话
except KeyboardInterrupt:
print(f"\n\n👋 已停止监控。共检测到 {last_event_count} 条事件。")
def main():
parser = argparse.ArgumentParser(description="Session 实时监控工具")
parser.add_argument("--session", "-s", required=True, help="会话 ID")
parser.add_argument("--agent", "-a", default="my_agent",
help="Agent 名称(默认 my_agent")
parser.add_argument("--user", "-u", default="codebuddy",
help="用户 ID默认 codebuddy")
parser.add_argument("--poll", "-p", type=float, default=1.5,
help="轮询间隔秒数(默认 1.5")
parser.add_argument("--url", default=None,
help="自定义 API Server 地址(覆盖默认)")
args = parser.parse_args()
try:
agent_name = resolve_agent(args.agent)
except ValueError as e:
print(str(e))
sys.exit(1)
api_url = get_api_url(agent_name, args.url)
watch_session(
api_url=api_url,
app_name=agent_name,
user_id=args.user,
session_id=args.session,
poll_interval=args.poll,
)
if __name__ == "__main__":
main()