Initial commit: ai-bbc-agent-test with 1 tools - src/server/api_server.py

This commit is contained in:
2026-01-30 11:46:49 +00:00
parent 409ff76929
commit d90a82e69d
+403
View File
@@ -0,0 +1,403 @@
"""
HTTP API 服务器 - ai-bbc-agent-test
Agent with 1 external tools
提供 REST API 和 MCP HTTP/SSE 端点。
集成回调功能用于计费。
自动生成时间: 2026-01-30T11:46:45.454768
"""
import json
import uuid
import os
import logging
from typing import Optional, Dict, Any, AsyncGenerator, List
from contextlib import asynccontextmanager
from datetime import datetime
from fastapi import FastAPI, HTTPException, Request, Header, Depends
from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import StreamingResponse, JSONResponse
from pydantic import BaseModel, Field
from .mcp_server import TOOL_MAP, TOOL_LIST
from .agent_callback_utils import AgentCallbackHandler, CallbackContextManager
# ==================== 配置 ====================
SERVER_NAME = "ai-bbc-agent-test"
POD_NAME = os.getenv("POD_NAME", "ai-bbc-agent-test")
USER_ID = os.getenv("USER_ID", "")
# 配置日志
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)
# 本 Agent 包含的工具列表
AGENT_TOOLS = ["ai_bbc_jina"]
# ==================== 回调处理器 ====================
callback_handler: Optional[AgentCallbackHandler] = None
def get_callback_handler() -> AgentCallbackHandler:
"""获取或创建回调处理器(单例)"""
global callback_handler
if callback_handler is None:
callback_handler = AgentCallbackHandler(
agent_name=POD_NAME,
user_id=USER_ID
)
return callback_handler
# ==================== FastAPI 应用 ====================
@asynccontextmanager
async def lifespan(app: FastAPI):
logger.info(f"🚀 {SERVER_NAME} 启动")
logger.info(f" 包含工具: {', '.join(AGENT_TOOLS)}")
logger.info(f" Pod 名称: {POD_NAME}")
logger.info(f" 用户 ID: {USER_ID or '未设置'}")
yield
logger.info(f"🛑 {SERVER_NAME} 关闭")
app = FastAPI(
title=SERVER_NAME,
description="Agent with 1 external tools",
version="1.0.0",
lifespan=lifespan
)
app.add_middleware(
CORSMiddleware,
allow_origins=["*"],
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
)
# ==================== API Key 验证 ====================
async def verify_api_key(
api_key: Optional[str] = Header(None, alias="api-key"),
authorization: Optional[str] = Header(None)
) -> str:
"""验证 API Key"""
if api_key and api_key.strip() and api_key.strip() != "sk":
return api_key.strip()
if authorization:
key = authorization[7:].strip() if authorization.startswith("Bearer ") else authorization.strip()
if key and key != "sk":
return key
# 允许无 API Key 访问(使用默认配置)
return os.getenv("OPENAI_API_KEY", "sk")
def get_api_key_from_request(request: Request) -> Optional[str]:
"""从请求头提取 API Key(不验证)"""
api_key = request.headers.get("api-key") or request.headers.get("api_key")
if not api_key:
auth = request.headers.get("Authorization")
if auth:
api_key = auth[7:] if auth.startswith("Bearer ") else auth
return api_key or os.getenv("OPENAI_API_KEY", "sk")
# ==================== 健康检查 ====================
@app.get("/")
async def root():
return {
"service": SERVER_NAME,
"status": "running",
"tools": list(TOOL_MAP.keys()),
"tools_count": len(TOOL_MAP),
"pod_name": POD_NAME
}
@app.get("/health")
async def health():
return {"status": "healthy", "service": SERVER_NAME, "tools_count": len(TOOL_MAP)}
# ==================== MCP 端点 ====================
sessions: Dict[str, Dict] = {}
async def handle_mcp_request(data: Dict, session_id: str = None, api_key: str = None, user_id: str = None) -> Dict:
"""处理 MCP JSON-RPC 请求(带回调)"""
method = data.get("method")
params = data.get("params", {})
req_id = data.get("id")
try:
if method == "initialize":
session_id = session_id or str(uuid.uuid4())
sessions[session_id] = {"initialized": True}
return {
"jsonrpc": "2.0", "id": req_id,
"result": {
"protocolVersion": "2024-11-05",
"capabilities": {"tools": {}},
"serverInfo": {"name": SERVER_NAME, "version": "1.0.0"}
}
}
elif method == "tools/list":
return {"jsonrpc": "2.0", "id": req_id, "result": {"tools": TOOL_LIST}}
elif method == "tools/call":
tool_name = params.get("name")
args = params.get("arguments", {})
if tool_name not in TOOL_MAP:
raise ValueError(f"Unknown tool: {tool_name}")
# 设置 API Key 到环境变量
old_key = os.environ.get('OPENAI_API_KEY')
if api_key:
os.environ['OPENAI_API_KEY'] = api_key
# 使用回调上下文管理器(如果有 user_id)
effective_user_id = user_id or USER_ID
try:
if effective_user_id:
handler = get_callback_handler()
with CallbackContextManager(
handler=handler,
user_id=effective_user_id,
request_id=f"mcp-{req_id}-{int(datetime.utcnow().timestamp())}"
) as ctx:
ctx.add_tool(tool_name)
result = await TOOL_MAP[tool_name](**args)
else:
result = await TOOL_MAP[tool_name](**args)
finally:
if old_key:
os.environ['OPENAI_API_KEY'] = old_key
return {
"jsonrpc": "2.0", "id": req_id,
"result": {"content": [{"type": "text", "text": str(result)}]}
}
elif method == "ping":
return {"jsonrpc": "2.0", "id": req_id, "result": {}}
else:
raise ValueError(f"Unknown method: {method}")
except Exception as e:
return {"jsonrpc": "2.0", "id": req_id, "error": {"code": -32603, "message": str(e)}}
@app.post("/mcp")
async def mcp_endpoint(request: Request):
"""MCP HTTP 端点"""
try:
body = await request.json()
session_id = request.headers.get("x-mcp-session-id")
api_key = get_api_key_from_request(request)
user_id = request.headers.get("x-user-id") or USER_ID
response = await handle_mcp_request(body, session_id, api_key, user_id)
return JSONResponse(content=response, headers={"x-mcp-session-id": session_id or ""})
except Exception as e:
return JSONResponse(status_code=400, content={"jsonrpc": "2.0", "error": {"code": -32700, "message": str(e)}})
@app.get("/mcp/sse")
async def mcp_sse(request: Request):
"""MCP SSE 端点"""
session_id = request.headers.get("x-mcp-session-id") or str(uuid.uuid4())
async def stream() -> AsyncGenerator[str, None]:
yield f"data: {json.dumps({'type': 'connection', 'sessionId': session_id})}\n\n"
import asyncio
while True:
await asyncio.sleep(30)
yield f"data: {json.dumps({'type': 'ping'})}\n\n"
return StreamingResponse(stream(), media_type="text/event-stream",
headers={"Cache-Control": "no-cache", "x-mcp-session-id": session_id})
@app.post("/mcp/sse")
async def mcp_sse_post(request: Request):
"""MCP SSE POST 端点"""
try:
body = await request.json()
session_id = request.headers.get("x-mcp-session-id") or str(uuid.uuid4())
api_key = get_api_key_from_request(request)
user_id = request.headers.get("x-user-id") or USER_ID
async def stream() -> AsyncGenerator[str, None]:
response = await handle_mcp_request(body, session_id, api_key, user_id)
yield f"data: {json.dumps(response)}\n\n"
return StreamingResponse(stream(), media_type="text/event-stream",
headers={"Cache-Control": "no-cache", "x-mcp-session-id": session_id})
except Exception as e:
return JSONResponse(status_code=400, content={"jsonrpc": "2.0", "error": {"code": -32700, "message": str(e)}})
# ==================== 业务 API ====================
class ToolCallRequest(BaseModel):
"""工具调用请求"""
tool_name: str = Field(..., description="工具名称")
parameters: Dict[str, Any] = Field(default={}, description="工具参数")
user_id: Optional[str] = Field(None, description="用户ID(用于计费回调)")
class MultiToolCallRequest(BaseModel):
"""批量工具调用请求"""
calls: List[ToolCallRequest] = Field(..., description="工具调用列表")
user_id: Optional[str] = Field(None, description="用户ID(用于计费回调)")
class ToolCallResponse(BaseModel):
"""工具调用响应"""
success: bool
result: Optional[Any] = None
error: Optional[str] = None
tools_used: Optional[List[str]] = None
class MultiToolCallResponse(BaseModel):
"""批量工具调用响应"""
success: bool
results: List[ToolCallResponse]
tools_used: List[str]
total_calls: int
@app.get("/tools")
async def list_tools():
"""列出可用工具"""
return {
"tools": [
{"name": t["name"], "description": t["description"]}
for t in TOOL_LIST
],
"count": len(TOOL_LIST)
}
@app.post("/tools/call", response_model=ToolCallResponse)
async def call_tool(request: ToolCallRequest, api_key: str = Depends(verify_api_key)):
"""调用单个工具(带计费回调)"""
if request.tool_name not in TOOL_MAP:
raise HTTPException(status_code=404, detail=f"工具 {request.tool_name} 不存在")
effective_user_id = request.user_id or USER_ID
tools_used = [request.tool_name]
try:
old_key = os.environ.get('OPENAI_API_KEY')
os.environ['OPENAI_API_KEY'] = api_key
try:
# 使用回调上下文管理器
if effective_user_id:
handler = get_callback_handler()
with CallbackContextManager(
handler=handler,
user_id=effective_user_id,
request_id=f"api-{int(datetime.utcnow().timestamp())}"
) as ctx:
ctx.add_tool(request.tool_name)
result = await TOOL_MAP[request.tool_name](**request.parameters)
else:
result = await TOOL_MAP[request.tool_name](**request.parameters)
return ToolCallResponse(
success=True,
result=json.loads(result) if isinstance(result, str) else result,
tools_used=tools_used
)
finally:
if old_key:
os.environ['OPENAI_API_KEY'] = old_key
except Exception as e:
logger.error(f"工具调用失败: {e}")
return ToolCallResponse(success=False, error=str(e), tools_used=tools_used)
@app.post("/tools/batch-call", response_model=MultiToolCallResponse)
async def batch_call_tools(request: MultiToolCallRequest, api_key: str = Depends(verify_api_key)):
"""批量调用多个工具(带计费回调)"""
effective_user_id = request.user_id or USER_ID
results = []
tools_used = []
# 设置 API Key
old_key = os.environ.get('OPENAI_API_KEY')
os.environ['OPENAI_API_KEY'] = api_key
try:
# 使用回调上下文管理器
if effective_user_id:
handler = get_callback_handler()
with CallbackContextManager(
handler=handler,
user_id=effective_user_id,
request_id=f"batch-{int(datetime.utcnow().timestamp())}"
) as ctx:
for call in request.calls:
if call.tool_name not in TOOL_MAP:
results.append(ToolCallResponse(
success=False,
error=f"工具 {call.tool_name} 不存在"
))
continue
try:
ctx.add_tool(call.tool_name)
tools_used.append(call.tool_name)
result = await TOOL_MAP[call.tool_name](**call.parameters)
results.append(ToolCallResponse(
success=True,
result=json.loads(result) if isinstance(result, str) else result
))
except Exception as e:
results.append(ToolCallResponse(success=False, error=str(e)))
else:
for call in request.calls:
if call.tool_name not in TOOL_MAP:
results.append(ToolCallResponse(
success=False,
error=f"工具 {call.tool_name} 不存在"
))
continue
try:
tools_used.append(call.tool_name)
result = await TOOL_MAP[call.tool_name](**call.parameters)
results.append(ToolCallResponse(
success=True,
result=json.loads(result) if isinstance(result, str) else result
))
except Exception as e:
results.append(ToolCallResponse(success=False, error=str(e)))
finally:
if old_key:
os.environ['OPENAI_API_KEY'] = old_key
return MultiToolCallResponse(
success=all(r.success for r in results),
results=results,
tools_used=list(set(tools_used)),
total_calls=len(request.calls)
)
if __name__ == '__main__':
import uvicorn
uvicorn.run(app, host="0.0.0.0", port=8000)