271 lines
8.9 KiB
Python
271 lines
8.9 KiB
Python
"""
|
||
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()
|