Initial commit: test-resource-agent with 0 tools - src/server/api_server.py
This commit is contained in:
@@ -0,0 +1,403 @@
|
||||
"""
|
||||
HTTP API 服务器 - test-resource-agent
|
||||
|
||||
Agent with 0 external tools
|
||||
提供 REST API 和 MCP HTTP/SSE 端点。
|
||||
集成回调功能用于计费。
|
||||
自动生成时间: 2026-01-30T11:49:56.097652
|
||||
"""
|
||||
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 = "test-resource-agent"
|
||||
POD_NAME = os.getenv("POD_NAME", "test-resource-agent")
|
||||
USER_ID = os.getenv("USER_ID", "")
|
||||
|
||||
# 配置日志
|
||||
logging.basicConfig(level=logging.INFO)
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# 本 Agent 包含的工具列表
|
||||
AGENT_TOOLS = []
|
||||
|
||||
# ==================== 回调处理器 ====================
|
||||
|
||||
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 0 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)
|
||||
Reference in New Issue
Block a user