Files
agent_management/agent_templates/agents/a2a_litellm_agent/a2a_server.py
T
zhanggangyong 33955b68dc fix: 统一 agent 默认端口为 8000(search_agent 系列保持 8080)
- 修改 template_manager.py、k8s_manager.py 中的端口映射
- 更新 jina_search_agent、azure_blob_agent 系列、a2a_litellm_agent 的代码和 Dockerfile 为 8000
- 添加端口修改脚本和测试脚本

Made-with: Cursor
2026-03-02 15:13:07 +00:00

540 lines
17 KiB
Python

"""
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()