""" A2A协议兼容的Agent服务 实现Google Agent2Agent协议规范 支持从请求传入 API key,也支持从环境变量获取 """ import asyncio import json import uuid import os from typing import Optional, Dict, Any, AsyncGenerator from datetime import datetime from contextlib import asynccontextmanager from fastapi import FastAPI, HTTPException, Request, Response from fastapi.responses import StreamingResponse, JSONResponse from fastapi.middleware.cors import CORSMiddleware from pydantic import BaseModel, Field import structlog from agent import LiteLLMAgent from config import get_config, AgentConfig, A2AConfig # 配置日志 logger = structlog.get_logger() # 环境变量配置 SERVICE_HOST = os.getenv("SERVICE_HOST", "0.0.0.0") SERVICE_PORT = int(os.getenv("SERVICE_PORT", "8000")) POD_NAME = os.getenv("POD_NAME", "a2a-litellm-agent") TEMPLATE_TYPE = os.getenv("TEMPLATE_TYPE", "a2a_litellm_agent") # ============== A2A 协议数据模型 ============== class A2APart(BaseModel): """A2A消息部分""" kind: str = "text" text: Optional[str] = None data: Optional[Dict[str, Any]] = None mime_type: Optional[str] = None class A2AMessage(BaseModel): """A2A消息""" role: str parts: list[A2APart] messageId: str = Field(default_factory=lambda: uuid.uuid4().hex) class A2AMessageSendParams(BaseModel): """A2A发送消息参数""" message: A2AMessage configuration: Optional[Dict[str, Any]] = None api_key: Optional[str] = Field(None, description="LiteLLM API密钥(可选,优先使用,否则从环境变量获取)") model: Optional[str] = Field(None, description="模型名称(可选,优先使用,否则从环境变量获取)") class A2ARequest(BaseModel): """A2A JSON-RPC请求""" jsonrpc: str = "2.0" id: str method: str params: Optional[Dict[str, Any]] = None class A2AArtifact(BaseModel): """A2A响应工件""" artifactId: str = Field(default_factory=lambda: uuid.uuid4().hex) name: str = "response" parts: list[A2APart] class A2ATaskStatus(BaseModel): """A2A任务状态""" state: str # submitted, working, input-required, completed, failed, canceled timestamp: str = Field(default_factory=lambda: datetime.utcnow().isoformat() + "Z") message: Optional[str] = None class A2ATask(BaseModel): """A2A任务""" kind: str = "task" id: str = Field(default_factory=lambda: uuid.uuid4().hex) contextId: str = Field(default_factory=lambda: uuid.uuid4().hex) status: A2ATaskStatus artifacts: Optional[list[A2AArtifact]] = None class A2AResponse(BaseModel): """A2A JSON-RPC响应""" jsonrpc: str = "2.0" id: str result: Optional[A2ATask] = None error: Optional[Dict[str, Any]] = None class A2AStreamEvent(BaseModel): """A2A流式事件""" kind: str taskId: str contextId: str data: Optional[Dict[str, Any]] = None # ============== Agent Card ============== class AgentSkill(BaseModel): """Agent技能""" id: str name: str description: str inputSchema: Optional[Dict[str, Any]] = None outputSchema: Optional[Dict[str, Any]] = None class AgentCapabilities(BaseModel): """Agent能力""" text: bool = True streaming: bool = True push_notifications: bool = False forms: bool = False files: bool = False class AgentCard(BaseModel): """A2A Agent Card - 描述Agent能力""" name: str description: str version: str url: str capabilities: AgentCapabilities skills: list[AgentSkill] authentication: Optional[Dict[str, Any]] = None # ============== A2A Server ============== class A2AAgentServer: """A2A协议Agent服务器""" def __init__( self, api_key: Optional[str] = None, model: Optional[str] = None ): """ 初始化A2A Agent服务器 Args: api_key: LiteLLM API密钥(可选,优先使用,否则从环境变量获取) model: 模型名称(可选,优先使用,否则从环境变量获取) """ # 获取配置 self.llm_config, self.agent_config, self.a2a_config = get_config(api_key, model) # 创建Agent(使用默认配置) self.default_agent = LiteLLMAgent( litellm_config=self.llm_config, agent_config=self.agent_config ) # 任务存储 self.tasks: Dict[str, A2ATask] = {} # 创建FastAPI应用 self.app = self._create_app() def _create_app(self) -> FastAPI: """创建FastAPI应用""" @asynccontextmanager async def lifespan(app: FastAPI): logger.info("A2A Agent服务启动", agent_name=self.agent_config.name) yield await self.default_agent.close() logger.info("A2A Agent服务关闭") app = FastAPI( title=f"{self.agent_config.name} - A2A Agent", description=self.agent_config.description, version=self.agent_config.version, lifespan=lifespan ) # CORS中间件 app.add_middleware( CORSMiddleware, allow_origins=["*"], allow_credentials=True, allow_methods=["*"], allow_headers=["*"], ) # 注册路由 self._register_routes(app) return app def _get_agent(self, api_key: Optional[str] = None, model: Optional[str] = None) -> LiteLLMAgent: """ 获取Agent实例 如果提供了api_key或model,创建新的Agent实例 否则使用默认Agent """ if api_key or model: # 创建新的配置和Agent llm_config, agent_config, _ = get_config(api_key, model) return LiteLLMAgent(litellm_config=llm_config, agent_config=agent_config) return self.default_agent def _register_routes(self, app: FastAPI): """注册A2A协议路由""" @app.get("/") async def root(): """服务根路径""" return { "name": self.agent_config.name, "version": self.agent_config.version, "protocol": "A2A", "status": "running", "pod_name": POD_NAME, "template_type": TEMPLATE_TYPE } @app.get("/health") async def health_check(): """健康检查""" return { "status": "healthy", "pod_name": POD_NAME, "template_type": TEMPLATE_TYPE, "configured": self.llm_config.api_key is not None, "timestamp": datetime.utcnow().isoformat() } @app.get("/.well-known/agent.json") async def get_agent_card(request: Request): """获取Agent Card (A2A发现协议)""" base_url = str(request.base_url).rstrip("/") card = AgentCard( name=self.agent_config.name, description=self.agent_config.description, version=self.agent_config.version, url=base_url, capabilities=AgentCapabilities( text=True, streaming=self.agent_config.enable_streaming, push_notifications=False ), skills=[ AgentSkill( id="general-assistant", name="通用助手", description="回答问题、提供建议、协助完成各种任务" ), AgentSkill( id="code-helper", name="代码助手", description="编写、解释和调试代码" ) ] ) return card.model_dump() @app.post("/message/send") async def send_message(request: Request): """A2A message/send 端点""" body = await request.json() # 解析JSON-RPC请求 try: rpc_request = A2ARequest(**body) except Exception as e: return JSONResponse({ "jsonrpc": "2.0", "id": body.get("id", "unknown"), "error": { "code": -32600, "message": f"Invalid Request: {str(e)}" } }) # 处理 message/send 方法 if rpc_request.method == "message/send": return await self._handle_message_send(rpc_request) elif rpc_request.method == "message/stream": return await self._handle_message_stream(rpc_request) else: return JSONResponse({ "jsonrpc": "2.0", "id": rpc_request.id, "error": { "code": -32601, "message": f"Method not found: {rpc_request.method}" } }) @app.post("/message/stream") async def stream_message(request: Request): """A2A message/stream 端点 (SSE流式响应)""" body = await request.json() try: rpc_request = A2ARequest(**body) except Exception as e: return JSONResponse({ "jsonrpc": "2.0", "id": body.get("id", "unknown"), "error": { "code": -32600, "message": f"Invalid Request: {str(e)}" } }) return await self._handle_message_stream(rpc_request) @app.get("/tasks/{task_id}") async def get_task(task_id: str): """获取任务状态""" if task_id not in self.tasks: raise HTTPException(status_code=404, detail="Task not found") return self.tasks[task_id].model_dump() async def _handle_message_send(self, request: A2ARequest) -> JSONResponse: """处理 message/send 请求""" params = request.params or {} message_data = params.get("message", {}) # 提取API key和model(如果提供) api_key = params.get("api_key") or os.getenv("LITELLM_API_KEY") model = params.get("model") or os.getenv("MODEL_NAME") or os.getenv("LITELLM_MODEL") # 提取用户消息文本 user_text = "" parts = message_data.get("parts", []) for part in parts: if part.get("kind") == "text": user_text += part.get("text", "") if not user_text: return JSONResponse({ "jsonrpc": "2.0", "id": request.id, "error": { "code": -32602, "message": "Invalid params: no text content found" } }) # 创建任务 task_id = uuid.uuid4().hex context_id = params.get("contextId", uuid.uuid4().hex) task = A2ATask( id=task_id, contextId=context_id, status=A2ATaskStatus(state="working") ) self.tasks[task_id] = task try: # 获取Agent实例(如果提供了api_key或model,使用新的实例) agent = self._get_agent(api_key, model) # 调用Agent获取响应 logger.info("处理消息", task_id=task_id, message_preview=user_text[:50]) response_text = await agent.chat( message=user_text, conversation_id=context_id ) # 如果创建了新Agent,关闭它 if api_key or model: await agent.close() # 更新任务状态 task.status = A2ATaskStatus(state="completed") task.artifacts = [ A2AArtifact( name="response", parts=[A2APart(kind="text", text=response_text)] ) ] self.tasks[task_id] = task return JSONResponse({ "jsonrpc": "2.0", "id": request.id, "result": task.model_dump() }) except Exception as e: logger.error("处理消息失败", error=str(e)) task.status = A2ATaskStatus(state="failed", message=str(e)) self.tasks[task_id] = task return JSONResponse({ "jsonrpc": "2.0", "id": request.id, "error": { "code": -32000, "message": f"Agent error: {str(e)}" } }) async def _handle_message_stream(self, request: A2ARequest) -> StreamingResponse: """处理 message/stream 请求 (SSE)""" params = request.params or {} message_data = params.get("message", {}) # 提取API key和model(如果提供) api_key = params.get("api_key") or os.getenv("LITELLM_API_KEY") model = params.get("model") or os.getenv("MODEL_NAME") or os.getenv("LITELLM_MODEL") # 提取用户消息 user_text = "" parts = message_data.get("parts", []) for part in parts: if part.get("kind") == "text": user_text += part.get("text", "") task_id = uuid.uuid4().hex context_id = params.get("contextId", uuid.uuid4().hex) async def event_generator() -> AsyncGenerator[str, None]: """生成SSE事件流""" agent = None try: # 获取Agent实例 agent = self._get_agent(api_key, model) # 发送任务开始事件 start_event = { "kind": "task-start", "taskId": task_id, "contextId": context_id } yield f"data: {json.dumps(start_event)}\n\n" # 获取流式响应 stream = await agent.chat( message=user_text, conversation_id=context_id, stream=True ) full_response = "" async for chunk in stream: full_response += chunk # 发送文本增量事件 delta_event = { "kind": "artifact-delta", "taskId": task_id, "contextId": context_id, "data": { "kind": "text", "text": chunk } } yield f"data: {json.dumps(delta_event)}\n\n" # 发送完成事件 complete_event = { "kind": "task-complete", "taskId": task_id, "contextId": context_id, "data": { "status": "completed", "artifacts": [{ "name": "response", "parts": [{"kind": "text", "text": full_response}] }] } } yield f"data: {json.dumps(complete_event)}\n\n" except Exception as e: # 发送错误事件 error_event = { "kind": "task-error", "taskId": task_id, "contextId": context_id, "data": { "error": str(e) } } yield f"data: {json.dumps(error_event)}\n\n" finally: # 如果创建了新Agent,关闭它 if agent and (api_key or model): await agent.close() return StreamingResponse( event_generator(), media_type="text/event-stream", headers={ "Cache-Control": "no-cache", "Connection": "keep-alive", "X-Accel-Buffering": "no" } ) def run(self, host: Optional[str] = None, port: Optional[int] = None): """运行服务器""" import uvicorn host = host or self.agent_config.host port = port or self.agent_config.port logger.info(f"启动A2A Agent服务", host=host, port=port) uvicorn.run(self.app, host=host, port=port) def create_app(api_key: Optional[str] = None, model: Optional[str] = None) -> FastAPI: """ 创建FastAPI应用(用于uvicorn启动) 使用方式: uvicorn a2a_server:app --host 0.0.0.0 --port 8000 或设置环境变量后: export LITELLM_API_KEY="your-key" export MODEL_NAME="your-model" uvicorn a2a_server:app --host 0.0.0.0 --port 8000 """ server = A2AAgentServer(api_key=api_key, model=model) return server.app # uvicorn 启动入口 # 环境变量: LITELLM_API_KEY, MODEL_NAME (或 LITELLM_MODEL) app = create_app()