diff --git a/task_receiver.py b/task_receiver.py new file mode 100644 index 0000000..cca0ffb --- /dev/null +++ b/task_receiver.py @@ -0,0 +1,152 @@ +"""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 threading + +from fastapi import APIRouter, HTTPException, Request +from fastapi.responses import JSONResponse + +import gateway_client +from google.adk.runners import InMemoryRunner +from google.genai import types as genai_types + +logger = logging.getLogger(__name__) + + +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]: + """用 ADK InMemoryRunner 运行一次 agent,返回 (接受任务后的首条回复, 最终总结)。 + + app 是 App 容器(根 agent 非裸 LlmAgent),run_async 不会自动创建 + session,需先用 runner.session_service 显式创建。 + """ + runner = InMemoryRunner(app=app) + session_id = f"task-{request_id}" + await runner.session_service.create_session( + app_name=runner.app_name, + user_id="gateway", + session_id=session_id, + ) + texts: list[str] = [] + final_text = "" + async for event in runner.run_async( + user_id="gateway", + session_id=session_id, + new_message=genai_types.Content(role="user", parts=[genai_types.Part(text=prompt)]), + ): + # 可中断:每收到一个事件检查一次停止标志 + if gateway_client.is_stop_requested(request_id): + raise _TaskCancelled() + if event.is_final_response(): + if event.content and event.content.parts: + for part in event.content.parts: + text = getattr(part, "text", None) + if text: + final_text += text + break + if event.content and event.content.parts: + for part in event.content.parts: + text = getattr(part, "text", None) + if text: + texts.append(text) + # 首条回复 = 接受任务后的第一条回应;最终总结 = final response + reply = (texts[0] if texts else final_text).strip() + summary = final_text.strip() or "\n".join(t for t in texts if t).strip() or "(无输出)" + return reply, summary + + def _execute_and_report(request_id: str, payload: dict) -> None: + """后台线程:执行 agent,成功后回传 success,异常回传 failed,被取消时回传 cancelled。 + + 执行过程中每步都检查停止标志(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 = asyncio.run(_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 _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。 + """ + 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) + threading.Thread( + target=_execute_and_report, + args=(request_id, payload), + name=f"task-{request_id[:8]}", + daemon=True, + ).start() + 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 \ No newline at end of file