121 lines
4.5 KiB
Python
121 lines
4.5 KiB
Python
"""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 非裸 LlmAgent),run_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"})
|