140 lines
5.0 KiB
Python
140 lines
5.0 KiB
Python
"""
|
||
Dev Agent API Server
|
||
使用 ADK 官方 ApiServer 构建 REST API 服务,支持:
|
||
- REST API 调用 agent(/run、/run_sse)
|
||
- 会话管理(创建/获取/删除,SQLite 持久化)
|
||
- Swagger UI 交互式文档(/docs)
|
||
- 上下文自动压缩
|
||
- 长期记忆(InMemory,后续可换向量库)
|
||
|
||
启动方式:
|
||
python api_server.py
|
||
|
||
主要端点:
|
||
GET /list-apps 列出所有 agent
|
||
POST /run 同步运行 agent
|
||
POST /run_sse 流式运行 agent(SSE)
|
||
GET /apps/{app}/users/{user}/sessions/{session} 获取会话
|
||
POST /apps/{app}/users/{user}/sessions/{session} 创建会话
|
||
GET /docs Swagger UI
|
||
"""
|
||
import os
|
||
import sys
|
||
|
||
# 脚本所在目录(作为 .env / data 等相对路径的基准)
|
||
PROJECT_ROOT = os.path.dirname(os.path.abspath(__file__))
|
||
if PROJECT_ROOT not in sys.path:
|
||
sys.path.insert(0, PROJECT_ROOT)
|
||
|
||
# 项目根目录(往上两级),确保 from agents.xxx.xxx import 可用
|
||
_REPO_ROOT = os.path.abspath(os.path.join(PROJECT_ROOT, "../.."))
|
||
if _REPO_ROOT not in sys.path:
|
||
sys.path.insert(0, _REPO_ROOT)
|
||
|
||
from dotenv import load_dotenv
|
||
load_dotenv(os.path.join(PROJECT_ROOT, "", ".env"))
|
||
|
||
# 强制 UTF-8
|
||
os.environ["PYTHONUTF8"] = "1"
|
||
|
||
import uvicorn
|
||
from google.adk.cli.api_server import ApiServer
|
||
from google.adk.cli.utils.base_agent_loader import BaseAgentLoader
|
||
from google.adk.sessions.sqlite_session_service import SqliteSessionService
|
||
from google.adk.artifacts.in_memory_artifact_service import InMemoryArtifactService
|
||
from google.adk.auth.credential_service.in_memory_credential_service import InMemoryCredentialService
|
||
from google.adk.evaluation.in_memory_eval_sets_manager import InMemoryEvalSetsManager
|
||
from google.adk.evaluation.local_eval_set_results_manager import LocalEvalSetResultsManager
|
||
from google.adk.memory.in_memory_memory_service import InMemoryMemoryService
|
||
from agents.my_agent.app import dev_app
|
||
|
||
# A2A 网关接入(注册 + 心跳),放最底部 import 以免循环依赖
|
||
import gateway_client
|
||
|
||
|
||
# 配置
|
||
HOST = os.getenv("API_SERVER_HOST", "0.0.0.0")
|
||
PORT = int(os.getenv("API_SERVER_PORT", "8001"))
|
||
|
||
# 数据目录
|
||
DATA_DIR = os.path.join(PROJECT_ROOT, "../../data")
|
||
os.makedirs(DATA_DIR, exist_ok=True)
|
||
|
||
|
||
class DevAgentLoader(BaseAgentLoader):
|
||
"""自定义 agent 加载器,直接返回我们的 App 对象(带 compaction 配置)"""
|
||
|
||
def load_agent(self, agent_name: str):
|
||
if agent_name == dev_app.name:
|
||
return dev_app
|
||
raise ValueError(f"Agent not found: {agent_name}")
|
||
|
||
def list_agents(self) -> list[str]:
|
||
return [dev_app.name]
|
||
|
||
|
||
def create_api_server() -> ApiServer:
|
||
"""构造 ApiServer 实例"""
|
||
# 会话服务:SQLite 持久化
|
||
session_service = SqliteSessionService(
|
||
db_path=os.path.join(DATA_DIR, "sessions.db")
|
||
)
|
||
|
||
# 工件服务
|
||
artifact_service = InMemoryArtifactService()
|
||
|
||
# 认证服务(暂不需要,内存版占位)
|
||
credential_service = InMemoryCredentialService()
|
||
|
||
# 记忆服务(内存版占位:满足 ApiServer 必填参数,不含记忆工具/回调,不会注入记忆)
|
||
memory_service = InMemoryMemoryService()
|
||
|
||
# 评测集管理(暂不需要,占位)
|
||
eval_sets_manager = InMemoryEvalSetsManager()
|
||
eval_set_results_manager = LocalEvalSetResultsManager(agents_dir=DATA_DIR)
|
||
|
||
return ApiServer(
|
||
agent_loader=DevAgentLoader(),
|
||
session_service=session_service,
|
||
memory_service=memory_service,
|
||
artifact_service=artifact_service,
|
||
credential_service=credential_service,
|
||
eval_sets_manager=eval_sets_manager,
|
||
eval_set_results_manager=eval_set_results_manager,
|
||
agents_dir=os.path.join(PROJECT_ROOT, ""),
|
||
auto_create_session=True,
|
||
)
|
||
|
||
|
||
def main():
|
||
api_server = create_api_server()
|
||
fastapi_app = api_server.get_fast_api_app()
|
||
|
||
# 挂载 A2A 网关任务接收端点(POST /tasks/{request_id})
|
||
from task_receiver import create_task_router
|
||
fastapi_app.include_router(create_task_router(dev_app))
|
||
|
||
# 向 A2A 网关注册并启动心跳(注册失败不阻塞服务启动)
|
||
gateway_ok = gateway_client.register_agent(dev_app.name, f"http://127.0.0.1:{PORT}")
|
||
if gateway_ok:
|
||
gateway_client.start_heartbeat(dev_app.name)
|
||
else:
|
||
print("[gateway] 注册失败,跳过心跳(网关可能未启动或 auth 不对)")
|
||
|
||
print("=" * 60)
|
||
print("Dev Agent API Server 启动中...")
|
||
print(f" 监听地址: http://{HOST}:{PORT}")
|
||
print(f" Swagger UI: http://{HOST}:{PORT}/docs")
|
||
print(f" 同步运行: POST http://{HOST}:{PORT}/run")
|
||
print(f" 流式运行: POST http://{HOST}:{PORT}/run_sse")
|
||
print(f" 列出agent: GET http://{HOST}:{PORT}/list-apps")
|
||
print(f" 会话持久化: SQLite ({DATA_DIR}/sessions.db)")
|
||
print(f" 上下文压缩: 已启用")
|
||
print("=" * 60)
|
||
|
||
uvicorn.run(fastapi_app, host=HOST, port=PORT, log_level="info")
|
||
|
||
|
||
if __name__ == "__main__":
|
||
main()
|