540 lines
17 KiB
Python
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", "8080"))
|
|
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 8080
|
|
|
|
或设置环境变量后:
|
|
export LITELLM_API_KEY="your-key"
|
|
export MODEL_NAME="your-model"
|
|
uvicorn a2a_server:app --host 0.0.0.0 --port 8080
|
|
"""
|
|
server = A2AAgentServer(api_key=api_key, model=model)
|
|
return server.app
|
|
|
|
|
|
# uvicorn 启动入口
|
|
# 环境变量: LITELLM_API_KEY, MODEL_NAME (或 LITELLM_MODEL)
|
|
app = create_app()
|