ADK-agents/agents/my_agent/task_receiver.py

121 lines
4.5 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完成后回传结果。
契约(网关 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 agents.my_agent.app import dev_app
from google.adk.runners import InMemoryRunner
from google.genai import types as genai_types
logger = logging.getLogger(__name__)
router = APIRouter(tags=["tasks"])
# 与 api_server.py 一致的导入路径(独立运行时兜底)
import os
import sys
PROJECT_ROOT = os.path.dirname(os.path.abspath(__file__))
if PROJECT_ROOT not in sys.path:
sys.path.insert(0, PROJECT_ROOT)
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)
async def _run_agent_once(prompt: str, request_id: str) -> tuple[str, str]:
"""用 ADK InMemoryRunner 运行一次 agent返回 (接受任务后的首条回复, 最终总结)。
dev_app 是 App 容器(根 agent 非裸 LlmAgentrun_async 不会自动创建
session需先用 runner.session_service 显式创建。
"""
runner = InMemoryRunner(app=dev_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 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。"""
try:
prompt = _payload_to_prompt(payload)
reply, summary = asyncio.run(_run_agent_once(prompt, request_id))
gateway_client.report_result(
request_id,
agent_id=dev_app.name,
status="success",
progress=100,
result={"reply": reply, "output": summary},
)
except Exception as e:
logger.exception("agent task failed request=%s", request_id)
gateway_client.report_result(
request_id,
agent_id=dev_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 = await request.json()
if body.get("auth") != gateway_client.GATEWAY_AUTH:
raise HTTPException(status_code=401, detail="invalid auth")
payload = body.get("payload") or {}
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"})