ADK-agents/task_receiver.py
2026-08-05 23:26:00 +08:00

195 lines
8.2 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.

"""A2A 网关任务接收端点(通用版):接收网关主动推送的任务,后台执行指定 agent完成后回传结果。
本文件为工厂模块,供任意 agent 复用。每个 agent 传入自己的 ADK App 对象即可:
from task_receiver import create_task_router
fastapi_app.include_router(create_task_router(dev_app))
契约(网关 relay.py dispatch_command 推送):
POST {endpoint}/tasks/{request_id}
body: {"auth": GATEWAY_AUTH, "request_id": str, "payload": {...}}
成功响应 202立即确认执行完成后由后台线程回传网关 /api/agent/result。
"""
import asyncio
import logging
import os
from fastapi import APIRouter, HTTPException, Request
from fastapi.responses import JSONResponse
import gateway_client
from google.adk.runners import Runner
from google.adk.sessions.sqlite_session_service import SqliteSessionService
from google.adk.artifacts.in_memory_artifact_service import InMemoryArtifactService
from google.adk.agents.run_config import RunConfig, StreamingMode
from google.genai import types as genai_types
logger = logging.getLogger(__name__)
def _sessions_db_path() -> str:
"""返回会话数据库路径(与 chat.py 一致,位于项目 data 目录)。"""
here = os.path.dirname(os.path.abspath(__file__))
data_dir = os.path.join(here, "data")
os.makedirs(data_dir, exist_ok=True)
return os.path.join(data_dir, "sessions.db")
def _first_sentence(text: str) -> str:
"""从文本中提取第一个完整句子(用于"Agent 接受任务回复"展示)。
ADK 流式事件会把回复拆成多个 text 片段,第一个片段常只有一两个字。这里把
累积文本按句子结束符切分,返回第一句完整内容;若没有句子结束符则回退为
完整文本。
"""
if not text:
return ""
for sep in ("", "", "", "!", "?", "\n", ";"):
idx = text.find(sep)
if idx != -1:
return text[: idx + 1].strip()
return text.strip()
def _payload_to_prompt(payload: dict) -> str:
"""将网关任务 payload 转换为 agent 的用户指令。"""
if not payload:
return "请执行任务并汇报结果。"
if "prompt" in payload and payload["prompt"]:
return str(payload["prompt"])
if "cmd" in payload and payload["cmd"]:
return f"请执行以下命令并汇报执行结果:\n{payload['cmd']}"
# 兜底:序列化整个 payload
return "请根据以下任务载荷执行并汇报结果:\n" + str(payload)
def create_task_router(app):
"""根据指定的 ADK App 创建网关任务接收 router。
Args:
app: ADK App 容器(如 agents.my_agent.app.dev_app需具备 .name 属性。
"""
router = APIRouter(tags=["tasks"])
class _TaskCancelled(Exception):
pass
async def _run_agent_once(prompt: str, request_id: str) -> tuple[str, str]:
"""运行一次 agent返回 (接受任务后的首条回复, 最终总结)。
必须使用 Runner + SqliteSessionService + streaming_mode=SSE与 chat.py
一致InMemoryRunner 无法驱动带 compaction 配置的 App 容器,会导致
LLM 不调用、回复为空("Root node was cancelled")。
"""
runner = Runner(
app=app,
session_service=SqliteSessionService(db_path=_sessions_db_path()),
artifact_service=InMemoryArtifactService(),
auto_create_session=True,
)
session_id = f"task-{request_id}"
message = genai_types.Content(parts=[genai_types.Part(text=prompt)])
texts: list[str] = []
final_text = ""
async for event in runner.run_async(
user_id="gateway",
session_id=session_id,
new_message=message,
run_config=RunConfig(streaming_mode=StreamingMode.SSE),
):
# 可中断:每收到一个事件检查一次停止标志
if gateway_client.is_stop_requested(request_id):
raise _TaskCancelled()
# 只收集非思考thought的用户可见文本过滤掉 thought 片段
if event.content and event.content.parts:
for part in event.content.parts:
text = getattr(part, "text", None)
is_thought = getattr(part, "thought", False)
if not text or is_thought:
continue
if event.is_final_response():
final_text += text
else:
texts.append(text)
# 最终总结 = final response 文本;首条回复 = 累积中间文本直到完整句子
summary = final_text.strip() or "".join(texts).strip() or "(无输出)"
reply = _first_sentence("".join(texts)) or summary
return reply, summary
async def _execute_and_report(request_id: str, payload: dict) -> None:
"""后台执行:执行 agent成功后回传 success异常回传 failed被取消时回传 cancelled。
此协程通过 asyncio.create_task 在主事件循环中调度,与 agent 的 MCP
session / opentelemetry 上下文保持同一事件循环,避免跨线程/跨 loop 导致的
"Root node was cancelled" / "Failed to detach context" 崩溃。
执行过程中每步都检查停止标志gateway_client.is_stop_requested一旦收到
取消指令task_stop即中断并回传失败cancelled网关 on_result 终态保护
会将其置回就绪。
"""
gateway_client.clear_stop_requested(request_id)
try:
prompt = _payload_to_prompt(payload)
reply, summary = await _run_agent_once(prompt, request_id)
if gateway_client.is_stop_requested(request_id):
raise _TaskCancelled()
gateway_client.report_result(
request_id,
agent_id=app.name,
status="success",
progress=100,
result={"reply": reply, "output": summary},
)
except asyncio.CancelledError:
logger.info("agent task cancelled (loop) request=%s", request_id)
gateway_client.report_result(
request_id,
agent_id=app.name,
status="failed",
progress=100,
error_info="cancelled by user",
)
except _TaskCancelled:
logger.info("agent task cancelled request=%s", request_id)
gateway_client.report_result(
request_id,
agent_id=app.name,
status="failed",
progress=100,
error_info="cancelled by user",
)
except Exception as e:
logger.exception("agent task failed request=%s", request_id)
gateway_client.report_result(
request_id,
agent_id=app.name,
status="failed",
progress=100,
error_info=str(e),
)
@router.post("/tasks/{request_id}")
async def receive_task(request_id: str, request: Request):
"""接收网关推送的任务,立即 202 确认,后台执行。
body 中若携带 cli_session_id则同时启动该会话的 SSE 停止指令订阅线程,
用于接收网关取消任务时下发的 task_stop。
执行在 asyncio.create_task 中调度(与 MCP session 同事件循环),
不再使用新线程 + asyncio.run避免跨事件循环导致 agent 崩溃。
"""
body = await request.json()
if body.get("auth") != gateway_client.GATEWAY_AUTH:
raise HTTPException(status_code=401, detail="invalid auth")
payload = body.get("payload") or {}
cli_session_id = body.get("cli_session_id")
if cli_session_id:
gateway_client.start_stop_listener(cli_session_id)
asyncio.create_task(
_execute_and_report(request_id, payload),
name=f"task-{request_id[:8]}",
)
logger.info("task received request=%s payload=%s", request_id, payload)
return JSONResponse(status_code=202, content={"ok": True, "request_id": request_id, "status": "accepted"})
return router