ADK-agents/task_receiver.py
2026-08-05 17:19:18 +08:00

152 lines
6.2 KiB
Python
Raw 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 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 非裸 LlmAgentrun_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