Update Heicode sub-mode runtime changes
This commit is contained in:
Vendored
BIN
Binary file not shown.
Vendored
BIN
Binary file not shown.
@@ -48,6 +48,9 @@ docker build -t your-agent:latest .
|
||||
```
|
||||
your_agent/
|
||||
├── Dockerfile
|
||||
├── common/
|
||||
│ ├── __init__.py
|
||||
│ └── agent_callback_utils.py # callback 工具
|
||||
├── requirements.txt
|
||||
├── run_api_server.py # 启动脚本
|
||||
└── src/
|
||||
@@ -65,3 +68,12 @@ your_agent/
|
||||
| LITELLM_GATEWAY_URL | 是 | LiteLLM Gateway URL |
|
||||
| LITELLM_MODEL | 否 | 模型名称,默认 taiji/gpt-4o-mini |
|
||||
| API_PORT | 否 | 端口,默认 8000 |
|
||||
| POD_NAME | 否 | Agent 名称,用于 callback 中的 `agentName` |
|
||||
| USER_ID | 否 | 用户 ID,用于 callback 中的 `userId` |
|
||||
| AGENT_CALLBACK_URL | 否 | 回调地址,默认指向 Agent Manager 计费回调接口 |
|
||||
|
||||
## Callback 模板说明
|
||||
|
||||
- 模板已内置 `common/agent_callback_utils.py`
|
||||
- `src/server/api_server.py` 已示范在 `tools/call` 和业务 API 中使用 `CallbackContextManager`
|
||||
- 以后新增业务接口时,优先复用 `run_with_callback(...)` 来包裹真实工具调用
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
|
||||
@@ -0,0 +1,151 @@
|
||||
"""
|
||||
Agent回调工具 - 用于向Agent Manager回调运行时长记录
|
||||
"""
|
||||
import os
|
||||
import time
|
||||
import logging
|
||||
import requests
|
||||
from typing import Optional, List
|
||||
from datetime import datetime, timezone
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class AgentCallbackHandler:
|
||||
"""Agent回调处理器"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
agent_name: Optional[str] = None,
|
||||
user_id: Optional[str] = None,
|
||||
callback_url: Optional[str] = None
|
||||
):
|
||||
self.agent_name = agent_name or os.getenv("POD_NAME", "unknown-agent")
|
||||
self.user_id = user_id or os.getenv("USER_ID", "")
|
||||
self.callback_url = callback_url or os.getenv(
|
||||
"AGENT_CALLBACK_URL",
|
||||
"http://mcp-server.taiji-ai.svc.cluster.local:8000/api/v1/billing/agent-callback"
|
||||
)
|
||||
|
||||
self.start_time: Optional[datetime] = None
|
||||
self.tools_used: List[str] = []
|
||||
self.request_id: Optional[str] = None
|
||||
|
||||
logger.info(
|
||||
"AgentCallbackHandler initialized: agent=%s, callback_url=%s",
|
||||
self.agent_name,
|
||||
self.callback_url,
|
||||
)
|
||||
|
||||
def start_request(self, request_id: Optional[str] = None, user_id: Optional[str] = None):
|
||||
self.start_time = datetime.now(timezone.utc)
|
||||
self.tools_used = []
|
||||
self.request_id = request_id or f"req-{int(time.time())}"
|
||||
|
||||
if user_id:
|
||||
self.user_id = user_id
|
||||
|
||||
logger.info("Request started: request_id=%s, user_id=%s", self.request_id, self.user_id)
|
||||
|
||||
def add_tool_used(self, tool_name: str):
|
||||
if tool_name not in self.tools_used:
|
||||
self.tools_used.append(tool_name)
|
||||
logger.debug("Tool used: %s", tool_name)
|
||||
|
||||
def end_request(self, tools_used: Optional[List[str]] = None) -> bool:
|
||||
if not self.start_time:
|
||||
logger.warning("Cannot end request: no start time recorded")
|
||||
return False
|
||||
|
||||
if not self.user_id:
|
||||
logger.warning("Cannot send callback: user_id not set")
|
||||
return False
|
||||
|
||||
end_time = datetime.now(timezone.utc)
|
||||
running_time = (end_time - self.start_time).total_seconds()
|
||||
final_tools_used = tools_used if tools_used is not None else self.tools_used
|
||||
|
||||
success = self._send_callback(
|
||||
running_time_seconds=int(running_time),
|
||||
start_time=self.start_time,
|
||||
end_time=end_time,
|
||||
tools_used=final_tools_used
|
||||
)
|
||||
|
||||
self.start_time = None
|
||||
self.tools_used = []
|
||||
self.request_id = None
|
||||
|
||||
return success
|
||||
|
||||
def _send_callback(
|
||||
self,
|
||||
running_time_seconds: int,
|
||||
start_time: datetime,
|
||||
end_time: datetime,
|
||||
tools_used: List[str]
|
||||
) -> bool:
|
||||
try:
|
||||
payload = {
|
||||
"agentName": self.agent_name,
|
||||
"userId": self.user_id,
|
||||
"podRunningTimeSeconds": running_time_seconds,
|
||||
"toolsUsed": tools_used,
|
||||
"startTime": start_time.isoformat(),
|
||||
"endTime": end_time.isoformat(),
|
||||
"requestId": self.request_id
|
||||
}
|
||||
|
||||
logger.info("Sending callback: %s", payload)
|
||||
|
||||
response = requests.post(
|
||||
self.callback_url,
|
||||
json=payload,
|
||||
timeout=5
|
||||
)
|
||||
|
||||
if response.status_code == 200:
|
||||
logger.info("Callback sent successfully: %s", response.json())
|
||||
return True
|
||||
|
||||
logger.error("Callback failed with status %s: %s", response.status_code, response.text)
|
||||
return False
|
||||
|
||||
except requests.exceptions.RequestException as e:
|
||||
logger.error("Failed to send callback: %s", str(e))
|
||||
return False
|
||||
except Exception as e:
|
||||
logger.error("Unexpected error sending callback: %s", str(e))
|
||||
return False
|
||||
|
||||
|
||||
class CallbackContextManager:
|
||||
"""回调上下文管理器 - 使用with语句自动处理开始和结束"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
handler: AgentCallbackHandler,
|
||||
request_id: Optional[str] = None,
|
||||
user_id: Optional[str] = None,
|
||||
tools_used: Optional[List[str]] = None
|
||||
):
|
||||
self.handler = handler
|
||||
self.request_id = request_id
|
||||
self.user_id = user_id
|
||||
self.tools_used = tools_used or []
|
||||
|
||||
def __enter__(self):
|
||||
self.handler.start_request(
|
||||
request_id=self.request_id,
|
||||
user_id=self.user_id
|
||||
)
|
||||
return self
|
||||
|
||||
def __exit__(self, exc_type, exc_val, exc_tb):
|
||||
self.handler.end_request(tools_used=self.tools_used)
|
||||
return False
|
||||
|
||||
def add_tool(self, tool_name: str):
|
||||
self.handler.add_tool_used(tool_name)
|
||||
if tool_name not in self.tools_used:
|
||||
self.tools_used.append(tool_name)
|
||||
@@ -11,3 +11,4 @@ uvicorn[standard]>=0.27.0
|
||||
|
||||
# HTTP Client
|
||||
aiohttp>=3.9.0
|
||||
requests>=2.31.0
|
||||
|
||||
@@ -14,18 +14,24 @@ from fastapi.middleware.cors import CORSMiddleware
|
||||
from fastapi.responses import StreamingResponse, JSONResponse
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from common.agent_callback_utils import AgentCallbackHandler, CallbackContextManager
|
||||
from .mcp_server import TOOL_MAP, TOOL_LIST
|
||||
|
||||
# ==================== 配置 ====================
|
||||
|
||||
SERVER_NAME = "Your Agent API" # 修改为你的 Agent 名称
|
||||
POD_NAME = os.getenv("POD_NAME", "your-agent")
|
||||
USER_ID = os.getenv("USER_ID", "")
|
||||
callback_handler: Optional[AgentCallbackHandler] = None
|
||||
|
||||
|
||||
# ==================== FastAPI 应用 ====================
|
||||
|
||||
@asynccontextmanager
|
||||
async def lifespan(app: FastAPI):
|
||||
global callback_handler
|
||||
print(f"🚀 {SERVER_NAME} 启动")
|
||||
callback_handler = AgentCallbackHandler(agent_name=POD_NAME, user_id=USER_ID)
|
||||
yield
|
||||
print(f"🛑 {SERVER_NAME} 关闭")
|
||||
|
||||
@@ -79,13 +85,18 @@ async def root():
|
||||
return {
|
||||
"service": SERVER_NAME,
|
||||
"status": "running",
|
||||
"tools": list(TOOL_MAP.keys())
|
||||
"tools": list(TOOL_MAP.keys()),
|
||||
"callback_enabled": callback_handler is not None
|
||||
}
|
||||
|
||||
|
||||
@app.get("/health")
|
||||
async def health():
|
||||
return {"status": "healthy", "service": SERVER_NAME}
|
||||
return {
|
||||
"status": "healthy",
|
||||
"service": SERVER_NAME,
|
||||
"callback_enabled": callback_handler is not None
|
||||
}
|
||||
|
||||
|
||||
# ==================== MCP 端点 ====================
|
||||
@@ -93,6 +104,27 @@ async def health():
|
||||
sessions: Dict[str, Dict] = {}
|
||||
|
||||
|
||||
async def run_with_callback(
|
||||
tool_name: str,
|
||||
func,
|
||||
*args,
|
||||
user_id: Optional[str] = None,
|
||||
request_id: Optional[str] = None,
|
||||
**kwargs
|
||||
):
|
||||
"""统一包装 callback 逻辑,便于后续新 Agent 直接复用。"""
|
||||
if not callback_handler:
|
||||
return await func(*args, **kwargs)
|
||||
|
||||
with CallbackContextManager(
|
||||
handler=callback_handler,
|
||||
user_id=user_id or USER_ID,
|
||||
request_id=request_id or f"{tool_name}-{uuid.uuid4().hex}"
|
||||
) as ctx:
|
||||
ctx.add_tool(tool_name)
|
||||
return await func(*args, **kwargs)
|
||||
|
||||
|
||||
async def handle_mcp_request(data: Dict, session_id: str = None, api_key: str = None) -> Dict:
|
||||
"""处理 MCP JSON-RPC 请求"""
|
||||
method = data.get("method")
|
||||
@@ -132,7 +164,13 @@ async def handle_mcp_request(data: Dict, session_id: str = None, api_key: str =
|
||||
os.environ['OPENAI_API_KEY'] = api_key
|
||||
|
||||
try:
|
||||
result = await TOOL_MAP[tool_name](**args)
|
||||
result = await run_with_callback(
|
||||
tool_name,
|
||||
TOOL_MAP[tool_name],
|
||||
user_id=args.get("user_id"),
|
||||
request_id=req_id or f"mcp-{tool_name}-{uuid.uuid4().hex}",
|
||||
**args
|
||||
)
|
||||
finally:
|
||||
if old_key:
|
||||
os.environ['OPENAI_API_KEY'] = old_key
|
||||
@@ -223,7 +261,13 @@ async def api_query(request: QueryRequest, api_key: str = Depends(verify_api_key
|
||||
os.environ['OPENAI_API_KEY'] = api_key
|
||||
|
||||
try:
|
||||
result = await TOOL_MAP['your_tool'](query=request.query, option=request.option)
|
||||
result = await run_with_callback(
|
||||
"your_tool",
|
||||
TOOL_MAP['your_tool'],
|
||||
query=request.query,
|
||||
option=request.option,
|
||||
request_id=f"api-your-tool-{uuid.uuid4().hex}"
|
||||
)
|
||||
return QueryResponse(success=True, result=result)
|
||||
finally:
|
||||
if old_key:
|
||||
|
||||
@@ -21,6 +21,14 @@ import structlog
|
||||
from agent import LiteLLMAgent
|
||||
from config import get_config, AgentConfig, A2AConfig
|
||||
|
||||
try:
|
||||
from common.agent_callback_utils import AgentCallbackHandler, CallbackContextManager
|
||||
CALLBACK_ENABLED = True
|
||||
except ImportError:
|
||||
CALLBACK_ENABLED = False
|
||||
AgentCallbackHandler = None
|
||||
CallbackContextManager = None
|
||||
|
||||
# 配置日志
|
||||
logger = structlog.get_logger()
|
||||
|
||||
@@ -29,6 +37,7 @@ 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")
|
||||
USER_ID = os.getenv("USER_ID", "")
|
||||
|
||||
# ============== A2A 协议数据模型 ==============
|
||||
|
||||
@@ -161,6 +170,10 @@ class A2AAgentServer:
|
||||
litellm_config=self.llm_config,
|
||||
agent_config=self.agent_config
|
||||
)
|
||||
|
||||
self.callback_handler = None
|
||||
if CALLBACK_ENABLED and AgentCallbackHandler:
|
||||
self.callback_handler = AgentCallbackHandler(agent_name=POD_NAME, user_id=USER_ID)
|
||||
|
||||
# 任务存储
|
||||
self.tasks: Dict[str, A2ATask] = {}
|
||||
@@ -370,11 +383,24 @@ class A2AAgentServer:
|
||||
|
||||
# 调用Agent获取响应
|
||||
logger.info("处理消息", task_id=task_id, message_preview=user_text[:50])
|
||||
|
||||
response_text = await agent.chat(
|
||||
message=user_text,
|
||||
conversation_id=context_id
|
||||
)
|
||||
callback_user_id = params.get("user_id") or USER_ID
|
||||
|
||||
if self.callback_handler:
|
||||
with CallbackContextManager(
|
||||
handler=self.callback_handler,
|
||||
user_id=callback_user_id,
|
||||
request_id=task_id
|
||||
) as ctx:
|
||||
ctx.add_tool("a2a_chat")
|
||||
response_text = await agent.chat(
|
||||
message=user_text,
|
||||
conversation_id=context_id
|
||||
)
|
||||
else:
|
||||
response_text = await agent.chat(
|
||||
message=user_text,
|
||||
conversation_id=context_id
|
||||
)
|
||||
|
||||
# 如果创建了新Agent,关闭它
|
||||
if api_key or model:
|
||||
@@ -435,6 +461,7 @@ class A2AAgentServer:
|
||||
try:
|
||||
# 获取Agent实例
|
||||
agent = self._get_agent(api_key, model)
|
||||
callback_user_id = params.get("user_id") or USER_ID
|
||||
|
||||
# 发送任务开始事件
|
||||
start_event = {
|
||||
@@ -444,27 +471,52 @@ class A2AAgentServer:
|
||||
}
|
||||
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
|
||||
if self.callback_handler:
|
||||
with CallbackContextManager(
|
||||
handler=self.callback_handler,
|
||||
user_id=callback_user_id,
|
||||
request_id=task_id
|
||||
) as ctx:
|
||||
ctx.add_tool("a2a_chat_stream")
|
||||
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"
|
||||
else:
|
||||
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"
|
||||
yield f"data: {json.dumps(delta_event)}\n\n"
|
||||
|
||||
# 发送完成事件
|
||||
complete_event = {
|
||||
|
||||
@@ -40,13 +40,19 @@ class LiteLLMConfig:
|
||||
max_tokens: int = 4096
|
||||
|
||||
def __post_init__(self):
|
||||
self.chat_endpoint = f"{self.base_url}/chat/completions"
|
||||
|
||||
self.base_url = (
|
||||
os.getenv("LITELLM_BASE_URL")
|
||||
or os.getenv("LLM_BASE_URL")
|
||||
or os.getenv("OPENAI_BASE_URL")
|
||||
or self.base_url
|
||||
).rstrip("/")
|
||||
|
||||
# 从环境变量读取(如果未直接提供)
|
||||
if self.api_key is None:
|
||||
self.api_key = os.getenv("LITELLM_API_KEY")
|
||||
if self.model is None:
|
||||
self.model = os.getenv("MODEL_NAME") or os.getenv("LITELLM_MODEL", "gpt-4")
|
||||
self.chat_endpoint = f"{self.base_url}/chat/completions"
|
||||
|
||||
def validate(self) -> bool:
|
||||
"""验证配置是否完整"""
|
||||
|
||||
@@ -1,6 +1,11 @@
|
||||
# Ad Creator Agent - API 文档
|
||||
# 广告创意生成智能体
|
||||
|
||||
多模态广告创意生成 Agent,通过素材(文字描述/参考图片)生成广告图片或视频。
|
||||
Ad Creator Agent 提供多模态广告创意生成能力,通过素材(文字描述/参考图片)生成广告图片或视频。
|
||||
生成的文件自动上传至 Azure Blob Storage,返回带 SAS token 的公开可访问 URL。
|
||||
|
||||
本项目包含 **一个 Agent 服务**,同时通过 HTTP API 与 MCP(Model Context Protocol)对外提供能力。
|
||||
|
||||
**Ad Creator Agent**:广告文案生成、广告图片生成、广告视频生成、智能对话
|
||||
|
||||
## 基本信息
|
||||
|
||||
@@ -9,7 +14,8 @@
|
||||
| 镜像 | `agnettaiji.azurecr.io/ai-agents/ad-creator-agent:latest` |
|
||||
| 端口 | `8000` |
|
||||
| 模板名 | `ad_creator_agent` |
|
||||
| 框架 | API (FastAPI) |
|
||||
| 框架 | API (FastAPI) + MCP |
|
||||
| 存储 | Azure Blob Storage (`multimodal` 容器) |
|
||||
|
||||
## 支持的模型
|
||||
|
||||
@@ -26,13 +32,8 @@
|
||||
|
||||
所有写操作端点均需传入 API Key,支持以下两种方式:
|
||||
|
||||
```
|
||||
api-key: sk-xxx
|
||||
```
|
||||
|
||||
```
|
||||
Authorization: Bearer sk-xxx
|
||||
```
|
||||
- `api-key: sk-xxx`
|
||||
- `Authorization: Bearer sk-xxx`
|
||||
|
||||
如果部署时配置了 `LLM_API_KEY` 环境变量,可省略请求头中的 Key。
|
||||
|
||||
@@ -41,125 +42,114 @@ Authorization: Bearer sk-xxx
|
||||
| 变量名 | 说明 | 默认值 |
|
||||
|--------|------|--------|
|
||||
| `LLM_API_KEY` | LiteLLM API Key | (必填或请求头传入) |
|
||||
| `LLM_BASE_URL` | LiteLLM Base URL | `https://litellm.graystone-fb459c5d.southeastasia.azurecontainerapps.io/v1` |
|
||||
| `LLM_BASE_URL` | LiteLLM Base URL | 已内置 |
|
||||
| `DEFAULT_IMAGE_MODEL` | 默认图片模型 | `taiji/gemini-3-pro-image-preview` |
|
||||
| `DEFAULT_TEXT_MODEL` | 默认文案模型 | `taiji/gpt-4o-mini` |
|
||||
| `DEFAULT_VIDEO_MODEL` | 默认视频模型 | `taiji/sora-2` |
|
||||
| `SERVICE_PORT` | 服务端口 | `8000` |
|
||||
| `OUTPUT_DIR` | 文件输出目录 | `/app/outputs` |
|
||||
| `AZURE_STORAGE_CONNECTION_STRING` | Azure Blob 连接字符串 | 已内置 |
|
||||
| `AZURE_BLOB_CONTAINER` | Blob 容器名称 | `multimodal` |
|
||||
| `AZURE_BLOB_SAS_TOKEN` | Blob 读取 SAS Token | 已内置(有效期至 2028) |
|
||||
|
||||
---
|
||||
|
||||
## API 端点
|
||||
## 功能概览
|
||||
|
||||
### 1. 健康检查
|
||||
提供广告素材的 **文案生成、图片生成、视频生成与智能对话** 能力,返回可直接访问的 Blob URL。
|
||||
|
||||
**GET** `/health`
|
||||
支持能力:
|
||||
|
||||
- 广告文案生成(结构化 JSON:标题/正文/CTA/hashtags/配图 prompt)
|
||||
- 广告图片生成(Gemini / GPT Image / DALL-E,支持参考图片)
|
||||
- 一键完整广告(文案 + 配图联动)
|
||||
- 广告视频生成(Sora)
|
||||
- 智能对话(自动理解需求并生成图片)
|
||||
- 文件管理(列出 / 下载 / 清理)
|
||||
|
||||
---
|
||||
|
||||
## 1⃣ generate-image — 生成广告图片
|
||||
|
||||
### 功能说明
|
||||
|
||||
根据文字描述生成广告图片,自动上传至 Blob Storage,返回可直接访问的公开 URL。
|
||||
|
||||
### REST API 调用
|
||||
|
||||
```bash
|
||||
curl http://<AGENT_URL>/health
|
||||
```
|
||||
|
||||
**响应示例:**
|
||||
POST /api/v1/generate-image
|
||||
Content-Type: application/json
|
||||
```
|
||||
|
||||
```json
|
||||
{
|
||||
"status": "healthy",
|
||||
"service": "Ad Creator Agent",
|
||||
"pod_name": "test-ad-creator",
|
||||
"models": {
|
||||
"image": "taiji/gemini-3-pro-image-preview",
|
||||
"text": "taiji/gpt-4o-mini",
|
||||
"video": "taiji/sora-2"
|
||||
},
|
||||
"callback_enabled": false,
|
||||
"timestamp": "2026-03-02T14:52:15.109589"
|
||||
"prompt": "A premium headphone floating against dark gradient background with golden light accents",
|
||||
"aspect_ratio": "1:1",
|
||||
"quality": "high",
|
||||
"style": "luxury",
|
||||
"brand_name": "SoundElite"
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
### MCP 调用
|
||||
|
||||
### 2. 生成广告图片
|
||||
|
||||
**POST** `/api/v1/generate-image`
|
||||
|
||||
通过文字描述生成广告图片,可指定模型、风格、宽高比等。
|
||||
|
||||
**请求体:**
|
||||
|
||||
| 字段 | 类型 | 必填 | 说明 |
|
||||
|------|------|------|------|
|
||||
| `prompt` | string | 是 | 广告图片描述/创意需求 |
|
||||
| `model` | string | 否 | 模型名称,默认 `taiji/gemini-3-pro-image-preview` |
|
||||
| `aspect_ratio` | string | 否 | 宽高比: `1:1`, `16:9`, `9:16`, `4:3`, `3:4`(Gemini) |
|
||||
| `size` | string | 否 | 图片尺寸(仅 GPT/DALL-E): `1024x1024`, `1024x1792`, `1792x1024` |
|
||||
| `quality` | string | 否 | 质量: `low`, `medium`, `high`(默认 `high`) |
|
||||
| `style` | string | 否 | 广告风格: `modern`, `minimalist`, `luxury`, `playful`, `tech`, `vintage` |
|
||||
| `brand_name` | string | 否 | 品牌名称 |
|
||||
| `reference_image_b64` | string | 否 | 参考图片 base64(仅 Gemini 支持) |
|
||||
|
||||
**示例 - Gemini 生成:**
|
||||
|
||||
```bash
|
||||
curl -X POST http://<AGENT_URL>/api/v1/generate-image \
|
||||
-H "Content-Type: application/json" \
|
||||
-H "api-key: sk-xxx" \
|
||||
-d '{
|
||||
"prompt": "A premium headphone floating against dark gradient background with golden light accents",
|
||||
"aspect_ratio": "1:1",
|
||||
"quality": "high",
|
||||
"style": "luxury",
|
||||
"brand_name": "SoundElite"
|
||||
}'
|
||||
```json
|
||||
{
|
||||
"jsonrpc": "2.0",
|
||||
"id": 1,
|
||||
"method": "tools/call",
|
||||
"params": {
|
||||
"name": "generate_ad_image",
|
||||
"arguments": {
|
||||
"prompt": "A premium headphone floating against dark gradient background",
|
||||
"model": "taiji/gemini-3-pro-image-preview",
|
||||
"aspect_ratio": "1:1",
|
||||
"style": "luxury",
|
||||
"brand_name": "SoundElite"
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
**示例 - GPT Image 生成:**
|
||||
### 参数说明
|
||||
|
||||
```bash
|
||||
curl -X POST http://<AGENT_URL>/api/v1/generate-image \
|
||||
-H "Content-Type: application/json" \
|
||||
-H "api-key: sk-xxx" \
|
||||
-d '{
|
||||
"prompt": "A vibrant Instagram ad for a coffee brand with warm morning light",
|
||||
"model": "taiji/gpt-image-1",
|
||||
"size": "1024x1024",
|
||||
"quality": "high"
|
||||
}'
|
||||
```
|
||||
| 参数 | 类型 | 必需 | 默认值 | 说明 |
|
||||
|------|------|------|--------|------|
|
||||
| prompt | string | ✅ | - | 广告图片描述(英文效果更好) |
|
||||
| model | string | ❌ | gemini-3-pro-image-preview | 图片生成模型 |
|
||||
| aspect_ratio | string | ❌ | 1:1 | 宽高比: 1:1, 16:9, 9:16, 4:3, 3:4(Gemini) |
|
||||
| size | string | ❌ | 1024x1024 | 图片尺寸(仅 GPT/DALL-E) |
|
||||
| quality | string | ❌ | high | 质量: low, medium, high |
|
||||
| style | string | ❌ | null | 风格: modern, minimalist, luxury, playful, tech, vintage |
|
||||
| brand_name | string | ❌ | null | 品牌名称 |
|
||||
| reference_image_b64 | string | ❌ | null | 参考图片 base64(仅 Gemini 支持) |
|
||||
|
||||
**响应示例:**
|
||||
### 返回结果
|
||||
|
||||
```json
|
||||
{
|
||||
"success": true,
|
||||
"file_path": "/app/outputs/images/ad_gemini_20260302_145310_209307.jpg",
|
||||
"filename": "ad_gemini_20260302_145310_209307.jpg",
|
||||
"url": "/api/v1/files/ad_gemini_20260302_145310_209307.jpg",
|
||||
"filename": "ad_gemini_20260302_171758_512832.jpg",
|
||||
"url": "https://agnettool.blob.core.windows.net/multimodal/ad_gemini_20260302_171758_512832.jpg?sp=r&st=...",
|
||||
"model": "taiji/gemini-3-pro-image-preview"
|
||||
}
|
||||
```
|
||||
|
||||
> 返回的 `url` 可直接在浏览器中打开查看图片。
|
||||
|
||||
---
|
||||
|
||||
### 3. 上传参考图片并生成广告图
|
||||
## 2⃣ generate-image-upload — 上传参考图片并生成
|
||||
|
||||
**POST** `/api/v1/generate-image-upload`
|
||||
### 功能说明
|
||||
|
||||
支持 `multipart/form-data` 上传参考图片,结合文字描述生成广告图。
|
||||
通过 `multipart/form-data` 上传参考图片,结合文字描述生成广告图。
|
||||
|
||||
**表单字段:**
|
||||
### REST API 调用
|
||||
|
||||
| 字段 | 类型 | 必填 | 说明 |
|
||||
|------|------|------|------|
|
||||
| `prompt` | string | 是 | 广告图片描述 |
|
||||
| `reference_image` | file | 否 | 参考图片文件 |
|
||||
| `model` | string | 否 | 模型名称 |
|
||||
| `aspect_ratio` | string | 否 | 宽高比 |
|
||||
| `quality` | string | 否 | 质量 |
|
||||
| `style` | string | 否 | 广告风格 |
|
||||
| `brand_name` | string | 否 | 品牌名称 |
|
||||
|
||||
**示例:**
|
||||
```
|
||||
POST /api/v1/generate-image-upload
|
||||
Content-Type: multipart/form-data
|
||||
```
|
||||
|
||||
```bash
|
||||
curl -X POST http://<AGENT_URL>/api/v1/generate-image-upload \
|
||||
@@ -170,93 +160,148 @@ curl -X POST http://<AGENT_URL>/api/v1/generate-image-upload \
|
||||
-F "aspect_ratio=16:9"
|
||||
```
|
||||
|
||||
### 参数说明
|
||||
|
||||
| 参数 | 类型 | 必需 | 默认值 | 说明 |
|
||||
|------|------|------|--------|------|
|
||||
| prompt | string | ✅ | - | 广告图片描述 |
|
||||
| reference_image | file | ❌ | null | 参考图片文件 |
|
||||
| model | string | ❌ | gemini | 模型名称 |
|
||||
| aspect_ratio | string | ❌ | 1:1 | 宽高比 |
|
||||
| quality | string | ❌ | high | 质量 |
|
||||
| style | string | ❌ | null | 广告风格 |
|
||||
| brand_name | string | ❌ | null | 品牌名称 |
|
||||
|
||||
---
|
||||
|
||||
### 4. 生成广告文案
|
||||
## 3⃣ generate-copy — 生成广告文案
|
||||
|
||||
**POST** `/api/v1/generate-copy`
|
||||
### 功能说明
|
||||
|
||||
根据产品信息,由 LLM 生成结构化广告文案(标题、正文、CTA、hashtags)以及用于图片生成的英文 prompt。
|
||||
|
||||
**请求体:**
|
||||
### REST API 调用
|
||||
|
||||
| 字段 | 类型 | 必填 | 说明 |
|
||||
|------|------|------|------|
|
||||
| `product` | string | 是 | 产品/服务描述 |
|
||||
| `target_audience` | string | 否 | 目标受众 |
|
||||
| `tone` | string | 否 | 语气: `professional`, `casual`, `humorous`, `urgent`, `luxury` |
|
||||
| `platform` | string | 否 | 投放平台: `instagram`, `facebook`, `tiktok`, `billboard`, `general` |
|
||||
| `language` | string | 否 | 语言: `zh`, `en`, `ja`(默认 `zh`) |
|
||||
|
||||
**示例:**
|
||||
|
||||
```bash
|
||||
curl -X POST http://<AGENT_URL>/api/v1/generate-copy \
|
||||
-H "Content-Type: application/json" \
|
||||
-H "api-key: sk-xxx" \
|
||||
-d '{
|
||||
"product": "高端无线降噪耳机,主打沉浸式音乐体验",
|
||||
"target_audience": "音乐爱好者和商务人士",
|
||||
"tone": "luxury",
|
||||
"platform": "instagram",
|
||||
"language": "zh"
|
||||
}'
|
||||
```
|
||||
POST /api/v1/generate-copy
|
||||
Content-Type: application/json
|
||||
```
|
||||
|
||||
**响应示例:**
|
||||
```json
|
||||
{
|
||||
"product": "高端无线降噪耳机,主打沉浸式音乐体验",
|
||||
"target_audience": "音乐爱好者和商务人士",
|
||||
"tone": "luxury",
|
||||
"platform": "instagram",
|
||||
"language": "zh"
|
||||
}
|
||||
```
|
||||
|
||||
### MCP 调用
|
||||
|
||||
```json
|
||||
{
|
||||
"jsonrpc": "2.0",
|
||||
"id": 2,
|
||||
"method": "tools/call",
|
||||
"params": {
|
||||
"name": "generate_ad_copy",
|
||||
"arguments": {
|
||||
"product": "高端无线降噪耳机",
|
||||
"target_audience": "音乐爱好者",
|
||||
"tone": "luxury",
|
||||
"platform": "instagram",
|
||||
"language": "zh"
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
### 参数说明
|
||||
|
||||
| 参数 | 类型 | 必需 | 默认值 | 说明 |
|
||||
|------|------|------|--------|------|
|
||||
| product | string | ✅ | - | 产品/服务描述 |
|
||||
| target_audience | string | ❌ | null | 目标受众 |
|
||||
| tone | string | ❌ | professional | 语气: professional, casual, humorous, urgent, luxury |
|
||||
| platform | string | ❌ | general | 投放平台: instagram, facebook, tiktok, billboard, general |
|
||||
| language | string | ❌ | zh | 语言: zh, en, ja |
|
||||
|
||||
### 返回结果
|
||||
|
||||
```json
|
||||
{
|
||||
"success": true,
|
||||
"headline": "沉浸高端音质",
|
||||
"body_copy": "体验非凡音质,尽享音乐带来的宁静与专注...",
|
||||
"body_copy": "体验非凡音质,尽享音乐带来的宁静与专注。我们的高端无线降噪耳机,专为追求极致的您设计。",
|
||||
"cta": "立即体验",
|
||||
"image_prompt": "A luxurious setting featuring a sleek wireless headphone...",
|
||||
"image_prompt": "A luxurious setting featuring a sleek wireless headphone on polished wood...",
|
||||
"hashtags": ["#高端耳机", "#沉浸音乐", "#商务生活"]
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### 5. 一键生成完整广告(文案 + 图片)
|
||||
## 4⃣ generate-ad — 一键生成完整广告
|
||||
|
||||
**POST** `/api/v1/generate-ad`
|
||||
### 功能说明
|
||||
|
||||
自动生成广告文案,并基于文案中的图片 prompt 自动生成配图。
|
||||
一次调用完成 **文案生成 → 图片 prompt 提取 → 图片生成 → 上传**,返回完整广告方案。
|
||||
|
||||
**请求体:**
|
||||
### REST API 调用
|
||||
|
||||
| 字段 | 类型 | 必填 | 说明 |
|
||||
|------|------|------|------|
|
||||
| `product` | string | 是 | 产品/服务描述 |
|
||||
| `image_model` | string | 否 | 图片生成模型 |
|
||||
| `aspect_ratio` | string | 否 | 宽高比 |
|
||||
| `style` | string | 否 | 广告风格 |
|
||||
| `brand_name` | string | 否 | 品牌名称 |
|
||||
| `target_audience` | string | 否 | 目标受众 |
|
||||
| `tone` | string | 否 | 语气 |
|
||||
| `platform` | string | 否 | 投放平台 |
|
||||
| `language` | string | 否 | 语言 |
|
||||
| `reference_image_b64` | string | 否 | 参考图片 base64 |
|
||||
|
||||
**示例:**
|
||||
|
||||
```bash
|
||||
curl -X POST http://<AGENT_URL>/api/v1/generate-ad \
|
||||
-H "Content-Type: application/json" \
|
||||
-H "api-key: sk-xxx" \
|
||||
-d '{
|
||||
"product": "新能源电动汽车,零排放、高续航、智能驾驶",
|
||||
"target_audience": "环保意识强的中产家庭",
|
||||
"tone": "professional",
|
||||
"platform": "facebook",
|
||||
"language": "zh",
|
||||
"style": "tech",
|
||||
"brand_name": "GreenDrive"
|
||||
}'
|
||||
```
|
||||
POST /api/v1/generate-ad
|
||||
Content-Type: application/json
|
||||
```
|
||||
|
||||
**响应示例:**
|
||||
```json
|
||||
{
|
||||
"product": "新能源电动汽车,零排放、高续航、智能驾驶",
|
||||
"target_audience": "环保意识强的中产家庭",
|
||||
"tone": "professional",
|
||||
"platform": "facebook",
|
||||
"language": "zh",
|
||||
"style": "tech",
|
||||
"brand_name": "GreenDrive"
|
||||
}
|
||||
```
|
||||
|
||||
### MCP 调用
|
||||
|
||||
```json
|
||||
{
|
||||
"jsonrpc": "2.0",
|
||||
"id": 3,
|
||||
"method": "tools/call",
|
||||
"params": {
|
||||
"name": "generate_full_ad",
|
||||
"arguments": {
|
||||
"product": "新能源电动汽车",
|
||||
"style": "tech",
|
||||
"brand_name": "GreenDrive",
|
||||
"language": "zh"
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
### 参数说明
|
||||
|
||||
| 参数 | 类型 | 必需 | 默认值 | 说明 |
|
||||
|------|------|------|--------|------|
|
||||
| product | string | ✅ | - | 产品/服务描述 |
|
||||
| image_model | string | ❌ | gemini | 图片生成模型 |
|
||||
| aspect_ratio | string | ❌ | 1:1 | 宽高比 |
|
||||
| style | string | ❌ | null | 广告风格 |
|
||||
| brand_name | string | ❌ | null | 品牌名称 |
|
||||
| target_audience | string | ❌ | null | 目标受众 |
|
||||
| tone | string | ❌ | professional | 语气 |
|
||||
| platform | string | ❌ | general | 投放平台 |
|
||||
| language | string | ❌ | zh | 语言 |
|
||||
| reference_image_b64 | string | ❌ | null | 参考图片 base64 |
|
||||
|
||||
### 返回结果
|
||||
|
||||
```json
|
||||
{
|
||||
@@ -264,15 +309,15 @@ curl -X POST http://<AGENT_URL>/api/v1/generate-ad \
|
||||
"copy": {
|
||||
"success": true,
|
||||
"headline": "开启绿色出行新生活",
|
||||
"body_copy": "选择我们的新能源电动汽车...",
|
||||
"body_copy": "选择我们的新能源电动汽车,为您的家庭带来零排放和高续航的驾驶体验。",
|
||||
"cta": "立即了解更多",
|
||||
"image_prompt": "A futuristic electric vehicle...",
|
||||
"image_prompt": "A futuristic electric vehicle on a modern highway...",
|
||||
"hashtags": ["#新能源车", "#绿色出行", "#智能驾驶"]
|
||||
},
|
||||
"image": {
|
||||
"success": true,
|
||||
"filename": "ad_gemini_20260302_145504_262223.jpg",
|
||||
"url": "/api/v1/files/ad_gemini_20260302_145504_262223.jpg",
|
||||
"url": "https://agnettool.blob.core.windows.net/multimodal/ad_gemini_20260302_145504_262223.jpg?sp=r&st=...",
|
||||
"model": "taiji/gemini-3-pro-image-preview"
|
||||
},
|
||||
"timestamp": "2026-03-02T14:55:04.262223"
|
||||
@@ -281,60 +326,64 @@ curl -X POST http://<AGENT_URL>/api/v1/generate-ad \
|
||||
|
||||
---
|
||||
|
||||
### 6. 生成广告视频
|
||||
## 5⃣ generate-video — 生成广告视频
|
||||
|
||||
**POST** `/api/v1/generate-video`
|
||||
### 功能说明
|
||||
|
||||
使用 Sora 模型生成广告短视频。
|
||||
使用 Sora 模型生成广告短视频,上传至 Blob 并返回 URL。
|
||||
|
||||
**请求体:**
|
||||
### REST API 调用
|
||||
|
||||
| 字段 | 类型 | 必填 | 说明 |
|
||||
|------|------|------|------|
|
||||
| `prompt` | string | 是 | 视频描述/创意需求 |
|
||||
| `model` | string | 否 | 视频模型(默认 `taiji/sora-2`) |
|
||||
| `aspect_ratio` | string | 否 | 宽高比: `16:9`, `9:16`, `1:1` |
|
||||
| `duration` | string | 否 | 视频时长秒数(默认 `5`) |
|
||||
|
||||
**示例:**
|
||||
|
||||
```bash
|
||||
curl -X POST http://<AGENT_URL>/api/v1/generate-video \
|
||||
-H "Content-Type: application/json" \
|
||||
-H "api-key: sk-xxx" \
|
||||
-d '{
|
||||
"prompt": "A sleek electric car driving through a futuristic city at sunset, cinematic style",
|
||||
"aspect_ratio": "16:9",
|
||||
"duration": "5"
|
||||
}'
|
||||
```
|
||||
POST /api/v1/generate-video
|
||||
Content-Type: application/json
|
||||
```
|
||||
|
||||
```json
|
||||
{
|
||||
"prompt": "A sleek electric car driving through a futuristic city at sunset, cinematic style",
|
||||
"aspect_ratio": "16:9",
|
||||
"duration": "5"
|
||||
}
|
||||
```
|
||||
|
||||
### 参数说明
|
||||
|
||||
| 参数 | 类型 | 必需 | 默认值 | 说明 |
|
||||
|------|------|------|--------|------|
|
||||
| prompt | string | ✅ | - | 视频描述/创意需求 |
|
||||
| model | string | ❌ | taiji/sora-2 | 视频模型 |
|
||||
| aspect_ratio | string | ❌ | 16:9 | 宽高比: 16:9, 9:16, 1:1 |
|
||||
| duration | string | ❌ | 5 | 视频时长秒数 |
|
||||
|
||||
---
|
||||
|
||||
### 7. 智能对话
|
||||
## 6⃣ chat — 智能对话
|
||||
|
||||
**POST** `/chat`
|
||||
### 功能说明
|
||||
|
||||
与 AI 广告创意总监对话。系统会理解需求,自动决定是否生成图片。
|
||||
与 AI 广告创意总监对话。系统会理解用户需求,自动决定是否生成图片。
|
||||
|
||||
**请求体:**
|
||||
### REST API 调用
|
||||
|
||||
| 字段 | 类型 | 必填 | 说明 |
|
||||
|------|------|------|------|
|
||||
| `message` | string | 是 | 用户消息 |
|
||||
|
||||
**示例:**
|
||||
|
||||
```bash
|
||||
curl -X POST http://<AGENT_URL>/chat \
|
||||
-H "Content-Type: application/json" \
|
||||
-H "api-key: sk-xxx" \
|
||||
-d '{
|
||||
"message": "帮我为一款蓝牙音箱做一个抖音封面图,要有科技感"
|
||||
}'
|
||||
```
|
||||
POST /chat
|
||||
Content-Type: application/json
|
||||
```
|
||||
|
||||
**响应示例:**
|
||||
```json
|
||||
{
|
||||
"message": "帮我为一款蓝牙音箱做一个抖音封面图,要有科技感和年轻活力"
|
||||
}
|
||||
```
|
||||
|
||||
### 参数说明
|
||||
|
||||
| 参数 | 类型 | 必需 | 说明 |
|
||||
|------|------|------|------|
|
||||
| message | string | ✅ | 用户消息 |
|
||||
|
||||
### 返回结果
|
||||
|
||||
```json
|
||||
{
|
||||
@@ -342,7 +391,7 @@ curl -X POST http://<AGENT_URL>/chat \
|
||||
"image": {
|
||||
"success": true,
|
||||
"filename": "ad_gemini_20260302_145539_866923.jpg",
|
||||
"url": "/api/v1/files/ad_gemini_20260302_145539_866923.jpg",
|
||||
"url": "https://agnettool.blob.core.windows.net/multimodal/ad_gemini_20260302_145539_866923.jpg?sp=r&st=...",
|
||||
"model": "taiji/gemini-3-pro-image-preview"
|
||||
},
|
||||
"timestamp": "2026-03-02T14:55:39.866923"
|
||||
@@ -351,36 +400,40 @@ curl -X POST http://<AGENT_URL>/chat \
|
||||
|
||||
---
|
||||
|
||||
### 8. 下载生成的文件
|
||||
## 7⃣ list-files — 列出已生成的文件
|
||||
|
||||
**GET** `/api/v1/files/{filename}`
|
||||
### REST API 调用
|
||||
|
||||
```bash
|
||||
curl -O http://<AGENT_URL>/api/v1/files/ad_gemini_20260302_145310_209307.jpg
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### 9. 列出已生成的文件
|
||||
|
||||
**GET** `/api/v1/list-files?file_type=all`
|
||||
GET /api/v1/list-files?file_type=all
|
||||
```
|
||||
|
||||
参数 `file_type` 可选值: `all`, `image`, `video`
|
||||
|
||||
```bash
|
||||
curl http://<AGENT_URL>/api/v1/list-files
|
||||
### MCP 调用
|
||||
|
||||
```json
|
||||
{
|
||||
"jsonrpc": "2.0",
|
||||
"id": 5,
|
||||
"method": "tools/call",
|
||||
"params": {
|
||||
"name": "list_generated_files",
|
||||
"arguments": { "file_type": "all" }
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
**响应示例:**
|
||||
### 返回结果
|
||||
|
||||
```json
|
||||
{
|
||||
"images": [
|
||||
{
|
||||
"filename": "ad_gemini_20260302_145539_866923.jpg",
|
||||
"url": "/api/v1/files/ad_gemini_20260302_145539_866923.jpg",
|
||||
"size_bytes": 589722,
|
||||
"created_at": "2026-03-02T14:55:39.865520"
|
||||
"filename": "ad_gemini_20260302_171758_512832.jpg",
|
||||
"url": "https://agnettool.blob.core.windows.net/multimodal/ad_gemini_20260302_171758_512832.jpg?sp=r&st=...",
|
||||
"size_bytes": 543592,
|
||||
"created_at": "2026-03-02T17:17:58+00:00"
|
||||
}
|
||||
],
|
||||
"videos": []
|
||||
@@ -389,35 +442,66 @@ curl http://<AGENT_URL>/api/v1/list-files
|
||||
|
||||
---
|
||||
|
||||
### 10. 清理旧文件
|
||||
## 8⃣ 其他端点
|
||||
|
||||
**POST** `/api/v1/cleanup?max_age_hours=24`
|
||||
### 下载/访问文件
|
||||
|
||||
删除超过指定时间的旧文件。
|
||||
|
||||
```bash
|
||||
curl -X POST "http://<AGENT_URL>/api/v1/cleanup?max_age_hours=24"
|
||||
```
|
||||
GET /api/v1/files/{filename}
|
||||
```
|
||||
|
||||
---
|
||||
Blob 模式下返回 302 跳转到 Blob 公开 URL。也可以直接使用生成时返回的 Blob URL。
|
||||
|
||||
### 11. 状态查看
|
||||
### 清理旧文件
|
||||
|
||||
**GET** `/status`
|
||||
|
||||
```bash
|
||||
curl http://<AGENT_URL>/status
|
||||
```
|
||||
POST /api/v1/cleanup?max_age_hours=24
|
||||
```
|
||||
|
||||
**响应示例:**
|
||||
从 Blob Storage 删除超过指定时间的旧文件。
|
||||
|
||||
### 健康检查
|
||||
|
||||
```
|
||||
GET /health
|
||||
```
|
||||
|
||||
### 状态查看
|
||||
|
||||
```
|
||||
GET /status
|
||||
```
|
||||
|
||||
```json
|
||||
{
|
||||
"status": "running",
|
||||
"pod_name": "test-ad-creator",
|
||||
"generated_images": 4,
|
||||
"pod_name": "ad-creator-v2",
|
||||
"storage": "azure_blob",
|
||||
"generated_images": 6,
|
||||
"generated_videos": 0,
|
||||
"timestamp": "2026-03-02T15:01:43.636444"
|
||||
"timestamp": "2026-03-02T17:20:00.000000"
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 统一错误格式
|
||||
|
||||
成功:
|
||||
|
||||
```json
|
||||
{
|
||||
"success": true,
|
||||
"data": {}
|
||||
}
|
||||
```
|
||||
|
||||
失败:
|
||||
|
||||
```json
|
||||
{
|
||||
"success": false,
|
||||
"error": "错误描述"
|
||||
}
|
||||
```
|
||||
|
||||
@@ -425,7 +509,7 @@ curl http://<AGENT_URL>/status
|
||||
|
||||
## 通过 Agent Manager 部署
|
||||
|
||||
### 1. 注册模板
|
||||
### 注册模板
|
||||
|
||||
```bash
|
||||
curl -X POST http://20.212.121.126/templates/create \
|
||||
@@ -444,7 +528,9 @@ curl -X POST http://20.212.121.126/templates/create \
|
||||
}'
|
||||
```
|
||||
|
||||
### 2. 创建实例
|
||||
### 创建实例
|
||||
|
||||
Blob Storage 凭证已内置,只需传 LLM API Key:
|
||||
|
||||
```bash
|
||||
curl -X POST http://20.212.121.126/agents \
|
||||
@@ -459,7 +545,7 @@ curl -X POST http://20.212.121.126/agents \
|
||||
}'
|
||||
```
|
||||
|
||||
### 3. 删除实例
|
||||
### 删除实例
|
||||
|
||||
```bash
|
||||
curl -X DELETE http://20.212.121.126/agents/my-ad-creator
|
||||
|
||||
@@ -12,11 +12,13 @@ RUN apt-get update && apt-get install -y \
|
||||
COPY common/requirements_a2a.txt /app/
|
||||
|
||||
# 安装Python依赖
|
||||
RUN pip install --no-cache-dir -r requirements_a2a.txt
|
||||
RUN pip install --no-cache-dir -r requirements_a2a.txt requests
|
||||
|
||||
# 复制应用代码和共享工具
|
||||
COPY agents/azure_blob_agent_a2a/azure_blob_agent_a2a.py /app/
|
||||
COPY common/api_key_utils.py /app/common/
|
||||
COPY common/agent_callback_utils.py /app/common/
|
||||
RUN touch /app/common/__init__.py
|
||||
|
||||
# 暴露端口
|
||||
EXPOSE 8000
|
||||
|
||||
@@ -12,7 +12,19 @@ from fastapi import FastAPI, HTTPException, Header
|
||||
from pydantic import BaseModel, Field
|
||||
from azure.storage.blob import BlobServiceClient, ContainerClient
|
||||
import uvicorn
|
||||
from api_key_utils import get_api_key
|
||||
|
||||
try:
|
||||
from common.api_key_utils import get_api_key
|
||||
except ImportError:
|
||||
from api_key_utils import get_api_key
|
||||
|
||||
try:
|
||||
from common.agent_callback_utils import AgentCallbackHandler, CallbackContextManager
|
||||
CALLBACK_ENABLED = True
|
||||
except ImportError:
|
||||
CALLBACK_ENABLED = False
|
||||
AgentCallbackHandler = None
|
||||
CallbackContextManager = None
|
||||
|
||||
# 配置日志
|
||||
logging.basicConfig(
|
||||
@@ -56,6 +68,7 @@ AGENT_CAPABILITIES = json.loads(os.getenv("AGENT_CAPABILITIES", '["blob_storage"
|
||||
# 全局存储客户端
|
||||
blob_service_client: Optional[BlobServiceClient] = None
|
||||
connection_string: Optional[str] = None
|
||||
callback_handler: Optional[AgentCallbackHandler] = None
|
||||
|
||||
# A2A Agent 注册表 (其他可协作的 Agent)
|
||||
registered_agents: Dict[str, Dict] = {}
|
||||
@@ -459,7 +472,20 @@ async def handle_a2a_message(message: A2AMessage):
|
||||
|
||||
try:
|
||||
handler = ACTION_HANDLERS[action]
|
||||
result = await handler(message.parameters)
|
||||
callback_user_id = (
|
||||
(message.context or {}).get("user_id")
|
||||
or USER_ID
|
||||
)
|
||||
if CALLBACK_ENABLED and callback_handler:
|
||||
with CallbackContextManager(
|
||||
handler=callback_handler,
|
||||
user_id=callback_user_id,
|
||||
request_id=message.message_id
|
||||
) as ctx:
|
||||
ctx.add_tool(action)
|
||||
result = await handler(message.parameters)
|
||||
else:
|
||||
result = await handler(message.parameters)
|
||||
|
||||
return {
|
||||
"message_id": message.message_id,
|
||||
@@ -505,16 +531,34 @@ async def query_storage(request: A2AQueryRequest):
|
||||
|
||||
# 简单的规则匹配
|
||||
if "容器" in query and ("列出" in query or "显示" in query or "有哪些" in query):
|
||||
result = await A2AActionHandler.handle_list_containers({})
|
||||
action_used = "list_containers"
|
||||
elif "统计" in query or "有多少" in query or "占用" in query:
|
||||
result = await A2AActionHandler.handle_get_stats({})
|
||||
action_used = "get_stats"
|
||||
elif request.container_name:
|
||||
if "文件" in query or "blob" in query.lower():
|
||||
result = await A2AActionHandler.handle_list_blobs({"container_name": request.container_name})
|
||||
action_used = "list_blobs"
|
||||
elif request.container_name and ("文件" in query or "blob" in query.lower()):
|
||||
action_used = "list_blobs"
|
||||
|
||||
callback_user_id = (request.context or {}).get("user_id") or USER_ID
|
||||
if action_used and CALLBACK_ENABLED and callback_handler:
|
||||
with CallbackContextManager(
|
||||
handler=callback_handler,
|
||||
user_id=callback_user_id,
|
||||
request_id=f"a2a-query-{action_used}-{int(datetime.now().timestamp())}"
|
||||
) as ctx:
|
||||
ctx.add_tool(action_used)
|
||||
if action_used == "list_containers":
|
||||
result = await A2AActionHandler.handle_list_containers({})
|
||||
elif action_used == "get_stats":
|
||||
result = await A2AActionHandler.handle_get_stats({})
|
||||
elif action_used == "list_blobs":
|
||||
result = await A2AActionHandler.handle_list_blobs({"container_name": request.container_name})
|
||||
|
||||
elif action_used == "list_containers":
|
||||
result = await A2AActionHandler.handle_list_containers({})
|
||||
elif action_used == "get_stats":
|
||||
result = await A2AActionHandler.handle_get_stats({})
|
||||
elif action_used == "list_blobs":
|
||||
result = await A2AActionHandler.handle_list_blobs({"container_name": request.container_name})
|
||||
|
||||
return {
|
||||
"status": "success" if result else "info",
|
||||
"query": request.query,
|
||||
@@ -636,6 +680,7 @@ def init_storage_connection():
|
||||
|
||||
def main():
|
||||
"""启动服务"""
|
||||
global callback_handler
|
||||
logger.info(f"🚀 启动 Azure Blob Storage AI Agent (A2A)")
|
||||
logger.info(f" - Framework: {AGENT_FRAMEWORK}")
|
||||
logger.info(f" - Agent ID: {AGENT_ID}")
|
||||
@@ -651,6 +696,12 @@ def main():
|
||||
# 初始化存储连接
|
||||
init_storage_connection()
|
||||
|
||||
if CALLBACK_ENABLED and AgentCallbackHandler:
|
||||
callback_handler = AgentCallbackHandler(agent_name=POD_NAME, user_id=USER_ID)
|
||||
logger.info(f"回调功能: 已启用 ({callback_handler.callback_url})")
|
||||
else:
|
||||
logger.info("回调功能: 未启用")
|
||||
|
||||
uvicorn.run(
|
||||
app,
|
||||
host=SERVICE_HOST,
|
||||
|
||||
@@ -12,10 +12,12 @@ RUN apt-get update && apt-get install -y \
|
||||
COPY common/requirements_mcp.txt /app/
|
||||
|
||||
# 安装Python依赖
|
||||
RUN pip install --no-cache-dir -r requirements_mcp.txt
|
||||
RUN pip install --no-cache-dir -r requirements_mcp.txt requests
|
||||
|
||||
# 复制应用代码
|
||||
COPY agents/azure_blob_agent_mcp/azure_blob_agent_mcp.py /app/
|
||||
COPY common/agent_callback_utils.py /app/common/
|
||||
RUN touch /app/common/__init__.py
|
||||
|
||||
# 暴露端口
|
||||
EXPOSE 8000
|
||||
|
||||
@@ -13,6 +13,14 @@ from azure.storage.blob import BlobServiceClient, ContainerClient
|
||||
import uvicorn
|
||||
import asyncio
|
||||
|
||||
try:
|
||||
from common.agent_callback_utils import AgentCallbackHandler, CallbackContextManager
|
||||
CALLBACK_ENABLED = True
|
||||
except ImportError:
|
||||
CALLBACK_ENABLED = False
|
||||
AgentCallbackHandler = None
|
||||
CallbackContextManager = None
|
||||
|
||||
# 配置日志
|
||||
logging.basicConfig(
|
||||
level=logging.INFO,
|
||||
@@ -50,6 +58,7 @@ NAMESPACE = os.getenv("NAMESPACE", "ai-agents")
|
||||
# 全局存储客户端
|
||||
blob_service_client: Optional[BlobServiceClient] = None
|
||||
connection_string: Optional[str] = None
|
||||
callback_handler: Optional[AgentCallbackHandler] = None
|
||||
|
||||
# MCP 工具注册表
|
||||
mcp_tools: Dict[str, Any] = {}
|
||||
@@ -486,7 +495,16 @@ async def call_mcp_tool(request: MCPToolRequest):
|
||||
|
||||
try:
|
||||
tool = mcp_tools[tool_name]
|
||||
result = await tool.execute(request.parameters)
|
||||
if CALLBACK_ENABLED and callback_handler:
|
||||
with CallbackContextManager(
|
||||
handler=callback_handler,
|
||||
user_id=USER_ID,
|
||||
request_id=f"blob-mcp-{tool_name}-{int(datetime.now().timestamp())}"
|
||||
) as ctx:
|
||||
ctx.add_tool(tool_name)
|
||||
result = await tool.execute(request.parameters)
|
||||
else:
|
||||
result = await tool.execute(request.parameters)
|
||||
|
||||
return {
|
||||
"tool": tool_name,
|
||||
@@ -510,18 +528,28 @@ async def query_storage(request: MCPQueryRequest):
|
||||
try:
|
||||
query = request.query.lower()
|
||||
result = None
|
||||
tool_name = None
|
||||
|
||||
# 简单的规则匹配 (实际应使用 LLM 进行意图识别)
|
||||
if "容器" in query and ("列出" in query or "显示" in query or "有哪些" in query):
|
||||
tool = mcp_tools["list_containers"]
|
||||
result = await tool.execute({})
|
||||
tool_name = "list_containers"
|
||||
elif "统计" in query or "有多少" in query or "占用" in query:
|
||||
tool = mcp_tools["get_storage_stats"]
|
||||
result = await tool.execute({})
|
||||
elif request.container_name:
|
||||
if "文件" in query or "blob" in query.lower():
|
||||
tool = mcp_tools["list_blobs"]
|
||||
result = await tool.execute({"container_name": request.container_name})
|
||||
tool_name = "get_storage_stats"
|
||||
elif request.container_name and ("文件" in query or "blob" in query.lower()):
|
||||
tool_name = "list_blobs"
|
||||
|
||||
if tool_name:
|
||||
params = {"container_name": request.container_name} if tool_name == "list_blobs" else {}
|
||||
if CALLBACK_ENABLED and callback_handler:
|
||||
with CallbackContextManager(
|
||||
handler=callback_handler,
|
||||
user_id=USER_ID,
|
||||
request_id=f"blob-query-{tool_name}-{int(datetime.now().timestamp())}"
|
||||
) as ctx:
|
||||
ctx.add_tool(tool_name)
|
||||
result = await mcp_tools[tool_name].execute(params)
|
||||
else:
|
||||
result = await mcp_tools[tool_name].execute(params)
|
||||
|
||||
if result:
|
||||
return {
|
||||
@@ -595,6 +623,7 @@ def init_storage_connection():
|
||||
|
||||
def main():
|
||||
"""启动服务"""
|
||||
global callback_handler
|
||||
logger.info(f"🚀 启动 Azure Blob Storage AI Agent (MCP)")
|
||||
logger.info(f" - Framework: {AGENT_FRAMEWORK}")
|
||||
logger.info(f" - Pod名称: {POD_NAME}")
|
||||
@@ -609,6 +638,12 @@ def main():
|
||||
|
||||
# 初始化存储连接
|
||||
init_storage_connection()
|
||||
|
||||
if CALLBACK_ENABLED and AgentCallbackHandler:
|
||||
callback_handler = AgentCallbackHandler(agent_name=POD_NAME, user_id=USER_ID)
|
||||
logger.info(f"回调功能: 已启用 ({callback_handler.callback_url})")
|
||||
else:
|
||||
logger.info("回调功能: 未启用")
|
||||
|
||||
uvicorn.run(
|
||||
app,
|
||||
|
||||
@@ -0,0 +1,298 @@
|
||||
# Code Manager Agent API 文档
|
||||
|
||||
Base URL: `http://<HOST>:8000`
|
||||
|
||||
所有业务接口需要在请求头中传递 API Key:
|
||||
```
|
||||
api-key: <YOUR_API_KEY>
|
||||
# 或
|
||||
Authorization: Bearer <YOUR_API_KEY>
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 健康检查
|
||||
|
||||
### GET /
|
||||
|
||||
返回服务基本信息。
|
||||
|
||||
**响应示例**
|
||||
```json
|
||||
{
|
||||
"service": "Code Manager Agent API",
|
||||
"status": "running",
|
||||
"tools": ["git_pull", "git_push", "update_code", "ssh_exec", "ssh_git_clone_and_test"]
|
||||
}
|
||||
```
|
||||
|
||||
### GET /health
|
||||
|
||||
```json
|
||||
{"status": "healthy", "service": "Code Manager Agent API"}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## MCP 接口
|
||||
|
||||
### POST /mcp
|
||||
|
||||
MCP JSON-RPC HTTP 端点,兼容 MCP 协议客户端。
|
||||
|
||||
**请求体(tools/list)**
|
||||
```json
|
||||
{"jsonrpc": "2.0", "id": 1, "method": "tools/list", "params": {}}
|
||||
```
|
||||
|
||||
**请求体(tools/call)**
|
||||
```json
|
||||
{
|
||||
"jsonrpc": "2.0",
|
||||
"id": 2,
|
||||
"method": "tools/call",
|
||||
"params": {
|
||||
"name": "git_pull",
|
||||
"arguments": {
|
||||
"username": "your_gitee_user",
|
||||
"password": "your_gitee_password"
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
### GET /sse
|
||||
|
||||
SSE 连接端点,返回 session ID。
|
||||
|
||||
### POST /sse/{session_id}
|
||||
|
||||
通过 SSE session 发送 MCP 请求(格式同 POST /mcp)。
|
||||
|
||||
---
|
||||
|
||||
## 业务 REST 接口
|
||||
|
||||
### POST /api/v1/git/pull
|
||||
|
||||
从 Gitee 仓库拉取最新代码。
|
||||
|
||||
**请求体**
|
||||
```json
|
||||
{
|
||||
"username": "your_gitee_user",
|
||||
"password": "your_gitee_password",
|
||||
"local_path": "/workspace",
|
||||
"branch": "main"
|
||||
}
|
||||
```
|
||||
|
||||
| 字段 | 类型 | 必需 | 说明 |
|
||||
|------|------|------|------|
|
||||
| username | string | 是 | Gitee 用户名 |
|
||||
| password | string | 是 | Gitee 密码 |
|
||||
| local_path | string | 否 | 本地仓库路径,默认 WORK_DIR |
|
||||
| branch | string | 否 | 分支名,默认当前分支 |
|
||||
|
||||
**响应示例**
|
||||
```json
|
||||
{
|
||||
"success": true,
|
||||
"stdout": "Already up to date.",
|
||||
"stderr": ""
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### POST /api/v1/git/push
|
||||
|
||||
提交并推送代码到 Gitee 仓库。
|
||||
|
||||
**请求体**
|
||||
```json
|
||||
{
|
||||
"username": "your_gitee_user",
|
||||
"password": "your_gitee_password",
|
||||
"local_path": "/workspace",
|
||||
"branch": "main",
|
||||
"commit_message": "feat: update agent code"
|
||||
}
|
||||
```
|
||||
|
||||
| 字段 | 类型 | 必需 | 说明 |
|
||||
|------|------|------|------|
|
||||
| username | string | 是 | Gitee 用户名 |
|
||||
| password | string | 是 | Gitee 密码 |
|
||||
| local_path | string | 否 | 本地仓库路径 |
|
||||
| branch | string | 否 | 目标分支 |
|
||||
| commit_message | string | 否 | 提交信息,为空则只 push 不 commit |
|
||||
|
||||
**响应示例**
|
||||
```json
|
||||
{
|
||||
"success": true,
|
||||
"logs": [
|
||||
{"step": "git add", "returncode": 0, "stdout": "", "stderr": ""},
|
||||
{"step": "git commit", "returncode": 0, "stdout": "[main abc1234] feat: update", "stderr": ""},
|
||||
{"step": "git push", "returncode": 0, "stdout": "", "stderr": ""}
|
||||
]
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### POST /api/v1/code/update
|
||||
|
||||
**Vibe Coding Subagent** — 接收自然语言任务,自主探索代码库、读文件、用 `edit_file`/`write_file`/`run_bash` 多轮迭代完成变更并写回磁盘。设计参考 [pi-mono coding agent](https://github.com/badlogic/pi-mono)。
|
||||
|
||||
Agent 内部工具循环:
|
||||
1. `read_file` — 按需读取任意文件
|
||||
2. `list_files` — glob 搜索文件
|
||||
3. `write_file` — 新建或全量覆写文件
|
||||
4. `edit_file` — 精确替换文件中的某段代码(surgical edit)
|
||||
5. `run_bash` — 运行 shell 命令验证(如 pytest、lint)
|
||||
6. `finish` — 宣布完成并输出摘要
|
||||
|
||||
**请求体**
|
||||
```json
|
||||
{
|
||||
"task": "给 login 函数增加 JWT 验证,失败时返回 401",
|
||||
"file_path": "src/auth/login.py",
|
||||
"local_path": "/workspace",
|
||||
"context_files": ["src/auth/models.py", "requirements.txt"]
|
||||
}
|
||||
```
|
||||
|
||||
| 字段 | 类型 | 必需 | 说明 |
|
||||
|------|------|------|------|
|
||||
| task | string | 是 | 自然语言任务描述 |
|
||||
| file_path | string | 是 | 任务入口文件(相对仓库根目录),agent 会自行探索 |
|
||||
| local_path | string | 否 | 本地仓库根路径,默认 WORK_DIR |
|
||||
| context_files | array | 否 | 初始上下文文件列表(只读提示),帮助 agent 更快定位 |
|
||||
|
||||
**响应示例**
|
||||
```json
|
||||
{
|
||||
"success": true,
|
||||
"task": "给 login 函数增加 JWT 验证",
|
||||
"files_changed": ["src/auth/login.py", "requirements.txt"],
|
||||
"tool_log": [
|
||||
{"tool": "read_file", "path": "src/auth/login.py", "bytes": 1240},
|
||||
{"tool": "edit_file", "path": "src/auth/login.py"},
|
||||
{"tool": "edit_file", "path": "requirements.txt"},
|
||||
{"tool": "run_bash", "command": "python -m pytest tests/test_auth.py", "returncode": 0},
|
||||
{"tool": "finish", "summary": "Added JWT validation to login(); updated requirements.txt with PyJWT>=2.8"}
|
||||
],
|
||||
"summary": "Added JWT validation to login(); updated requirements.txt with PyJWT>=2.8"
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### POST /api/v1/ssh/exec
|
||||
|
||||
通过 SSH 连接远程机器并执行命令。
|
||||
|
||||
**请求体**
|
||||
```json
|
||||
{
|
||||
"host": "192.168.1.100",
|
||||
"username": "ubuntu",
|
||||
"command": "ls -la /workspace",
|
||||
"password": "ssh_password",
|
||||
"port": 22
|
||||
}
|
||||
```
|
||||
|
||||
| 字段 | 类型 | 必需 | 说明 |
|
||||
|------|------|------|------|
|
||||
| host | string | 是 | 远程主机 IP 或域名 |
|
||||
| username | string | 是 | SSH 用户名 |
|
||||
| command | string | 是 | 要执行的命令 |
|
||||
| password | string | 否 | SSH 密码(与 ssh_key_path 二选一)|
|
||||
| ssh_key_path | string | 否 | SSH 私钥文件路径 |
|
||||
| port | integer | 否 | SSH 端口,默认 22 |
|
||||
|
||||
**响应示例**
|
||||
```json
|
||||
{
|
||||
"success": true,
|
||||
"exit_code": 0,
|
||||
"stdout": "total 48\ndrwxr-xr-x ...",
|
||||
"stderr": ""
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### POST /api/v1/ssh/clone-and-test
|
||||
|
||||
SSH 连接到测试机器,git clone 代码仓库,然后执行测试命令。
|
||||
|
||||
**请求体**
|
||||
```json
|
||||
{
|
||||
"host": "192.168.1.100",
|
||||
"ssh_username": "ubuntu",
|
||||
"remote_work_dir": "/home/ubuntu/test",
|
||||
"test_command": "pip install -r requirements.txt && python -m pytest",
|
||||
"gitee_username": "your_gitee_user",
|
||||
"gitee_password": "your_gitee_password",
|
||||
"ssh_password": "ssh_password",
|
||||
"ssh_port": 22,
|
||||
"branch": "main"
|
||||
}
|
||||
```
|
||||
|
||||
| 字段 | 类型 | 必需 | 说明 |
|
||||
|------|------|------|------|
|
||||
| host | string | 是 | 测试机器 IP 或域名 |
|
||||
| ssh_username | string | 是 | SSH 用户名 |
|
||||
| remote_work_dir | string | 是 | 远程机器工作目录 |
|
||||
| test_command | string | 是 | 测试命令(在仓库目录内执行)|
|
||||
| gitee_username | string | 是 | Gitee 用户名 |
|
||||
| gitee_password | string | 是 | Gitee 密码 |
|
||||
| ssh_password | string | 否 | SSH 密码(与 ssh_key_path 二选一)|
|
||||
| ssh_key_path | string | 否 | SSH 私钥文件路径 |
|
||||
| ssh_port | integer | 否 | SSH 端口,默认 22 |
|
||||
| branch | string | 否 | 要 clone 的分支 |
|
||||
|
||||
**响应示例**
|
||||
```json
|
||||
{
|
||||
"success": true,
|
||||
"logs": [
|
||||
{"step": "mkdir", "exit_code": 0, "stdout": "", "stderr": ""},
|
||||
{"step": "git clone", "exit_code": 0, "stdout": "Cloning into 'agent_management'...", "stderr": ""},
|
||||
{"step": "test", "exit_code": 0, "stdout": "All tests passed.", "stderr": ""}
|
||||
]
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## OpenClaw 接入
|
||||
|
||||
在 OpenClaw 工具配置中添加:
|
||||
|
||||
```json
|
||||
{
|
||||
"mcpServers": {
|
||||
"code_manager": {
|
||||
"url": "http://<HOST>:8000/mcp",
|
||||
"transport": "http",
|
||||
"headers": {
|
||||
"api-key": "<YOUR_API_KEY>"
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
可用工具将自动暴露给 OpenClaw,工具名称为:
|
||||
- `git_pull`
|
||||
- `git_push`
|
||||
- `update_code`
|
||||
- `ssh_exec`
|
||||
- `ssh_git_clone_and_test`
|
||||
@@ -0,0 +1,20 @@
|
||||
FROM python:3.12-slim
|
||||
|
||||
WORKDIR /app
|
||||
|
||||
ENV PYTHONUNBUFFERED=1
|
||||
ENV PYTHONDONTWRITEBYTECODE=1
|
||||
|
||||
RUN apt-get update && apt-get install -y gcc curl git openssh-client && rm -rf /var/lib/apt/lists/*
|
||||
|
||||
COPY requirements.txt .
|
||||
RUN pip install --no-cache-dir -r requirements.txt
|
||||
|
||||
COPY . .
|
||||
|
||||
EXPOSE 8000
|
||||
|
||||
HEALTHCHECK --interval=30s --timeout=10s --start-period=5s --retries=3 \
|
||||
CMD curl -f http://localhost:8000/health || exit 1
|
||||
|
||||
CMD ["python", "run_api_server.py"]
|
||||
@@ -0,0 +1,130 @@
|
||||
# Code Manager Agent
|
||||
|
||||
基于 **Pydantic AI + FastMCP** 的代码仓库管理 Agent。
|
||||
|
||||
支持:
|
||||
- Gitee 仓库 pull / push(HTTP 用户名+密码认证)
|
||||
- 本地文件更新
|
||||
- SSH 远程执行命令
|
||||
- SSH 到测试机器 git clone 并运行测试
|
||||
|
||||
代码仓库地址:`http://gitee.ath.cx:3000/zhanggangyong/agent_management`
|
||||
|
||||
---
|
||||
|
||||
## 快速开始
|
||||
|
||||
### 本地运行
|
||||
|
||||
```bash
|
||||
cd agent_templates/agents/code_manager_agent
|
||||
pip install -r requirements.txt
|
||||
export GITEE_REPO_URL=http://gitee.ath.cx:3000/zhanggangyong/agent_management
|
||||
export WORK_DIR=/path/to/local/repo
|
||||
python run_api_server.py
|
||||
```
|
||||
|
||||
### Docker
|
||||
|
||||
```bash
|
||||
docker build -t code-manager-agent:latest .
|
||||
docker run -p 8000:8000 \
|
||||
-e GITEE_REPO_URL=http://gitee.ath.cx:3000/zhanggangyong/agent_management \
|
||||
-e WORK_DIR=/workspace \
|
||||
code-manager-agent:latest
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 环境变量
|
||||
|
||||
| 变量 | 必需 | 默认值 | 说明 |
|
||||
|------|------|--------|------|
|
||||
| GITEE_REPO_URL | 否 | `http://gitee.ath.cx:3000/zhanggangyong/agent_management` | Gitee 仓库地址 |
|
||||
| WORK_DIR | 否 | `/workspace` | 本地仓库根路径 |
|
||||
| API_PORT | 否 | `8000` | 服务端口 |
|
||||
| POD_NAME | 否 | `code-manager-agent` | Agent 名称(用于 callback)|
|
||||
| USER_ID | 否 | `` | 用户 ID(用于 callback)|
|
||||
| AGENT_CALLBACK_URL | 否 | Agent Manager 默认回调 | 计费回调地址 |
|
||||
|
||||
---
|
||||
|
||||
## 项目结构
|
||||
|
||||
```
|
||||
code_manager_agent/
|
||||
├── Dockerfile
|
||||
├── README.md
|
||||
├── API_DOC.md
|
||||
├── requirements.txt
|
||||
├── run_api_server.py
|
||||
├── common/
|
||||
│ ├── __init__.py
|
||||
│ └── agent_callback_utils.py
|
||||
└── src/
|
||||
├── __init__.py
|
||||
└── server/
|
||||
├── __init__.py
|
||||
├── api_server.py # FastAPI + MCP HTTP
|
||||
└── mcp_server.py # MCP 工具定义
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 工具列表
|
||||
|
||||
| 工具 | 说明 |
|
||||
|------|------|
|
||||
| `git_pull` | 从 Gitee 拉取最新代码 |
|
||||
| `git_push` | 提交并推送代码到 Gitee |
|
||||
| `update_code` | 更新本地仓库中的指定文件 |
|
||||
| `ssh_exec` | SSH 连接远程机器执行命令 |
|
||||
| `ssh_git_clone_and_test` | SSH 到测试机 clone 代码并运行测试 |
|
||||
|
||||
---
|
||||
|
||||
## 注册到 Agent Manager
|
||||
|
||||
在 `k8s_manager.py` 中添加:
|
||||
|
||||
```python
|
||||
# TEMPLATE_PORTS
|
||||
"code_manager_agent": 8000,
|
||||
|
||||
# image_map
|
||||
"code_manager_agent": "agnettaiji.azurecr.io/ai-agents/code-manager-agent:latest",
|
||||
```
|
||||
|
||||
在 `app.py` 的 `valid_templates` 中添加 `"code_manager_agent"`。
|
||||
|
||||
---
|
||||
|
||||
## OpenClaw 接入说明
|
||||
|
||||
将以下配置添加到 OpenClaw 的 MCP 服务列表:
|
||||
|
||||
```json
|
||||
{
|
||||
"name": "code_manager_agent",
|
||||
"url": "http://<HOST>:8000/mcp",
|
||||
"transport": "http",
|
||||
"headers": {
|
||||
"api-key": "<YOUR_API_KEY>"
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
或使用 SSE 传输:
|
||||
|
||||
```json
|
||||
{
|
||||
"name": "code_manager_agent",
|
||||
"url": "http://<HOST>:8000/sse",
|
||||
"transport": "sse",
|
||||
"headers": {
|
||||
"api-key": "<YOUR_API_KEY>"
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
详细 API 说明请参考 [API_DOC.md](./API_DOC.md)。
|
||||
@@ -0,0 +1 @@
|
||||
|
||||
@@ -0,0 +1,151 @@
|
||||
"""
|
||||
Agent回调工具 - 用于向Agent Manager回调运行时长记录
|
||||
"""
|
||||
import os
|
||||
import time
|
||||
import logging
|
||||
import requests
|
||||
from typing import Optional, List
|
||||
from datetime import datetime, timezone
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class AgentCallbackHandler:
|
||||
"""Agent回调处理器"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
agent_name: Optional[str] = None,
|
||||
user_id: Optional[str] = None,
|
||||
callback_url: Optional[str] = None
|
||||
):
|
||||
self.agent_name = agent_name or os.getenv("POD_NAME", "unknown-agent")
|
||||
self.user_id = user_id or os.getenv("USER_ID", "")
|
||||
self.callback_url = callback_url or os.getenv(
|
||||
"AGENT_CALLBACK_URL",
|
||||
"http://mcp-server.taiji-ai.svc.cluster.local:8000/api/v1/billing/agent-callback"
|
||||
)
|
||||
|
||||
self.start_time: Optional[datetime] = None
|
||||
self.tools_used: List[str] = []
|
||||
self.request_id: Optional[str] = None
|
||||
|
||||
logger.info(
|
||||
"AgentCallbackHandler initialized: agent=%s, callback_url=%s",
|
||||
self.agent_name,
|
||||
self.callback_url,
|
||||
)
|
||||
|
||||
def start_request(self, request_id: Optional[str] = None, user_id: Optional[str] = None):
|
||||
self.start_time = datetime.now(timezone.utc)
|
||||
self.tools_used = []
|
||||
self.request_id = request_id or f"req-{int(time.time())}"
|
||||
|
||||
if user_id:
|
||||
self.user_id = user_id
|
||||
|
||||
logger.info("Request started: request_id=%s, user_id=%s", self.request_id, self.user_id)
|
||||
|
||||
def add_tool_used(self, tool_name: str):
|
||||
if tool_name not in self.tools_used:
|
||||
self.tools_used.append(tool_name)
|
||||
logger.debug("Tool used: %s", tool_name)
|
||||
|
||||
def end_request(self, tools_used: Optional[List[str]] = None) -> bool:
|
||||
if not self.start_time:
|
||||
logger.warning("Cannot end request: no start time recorded")
|
||||
return False
|
||||
|
||||
if not self.user_id:
|
||||
logger.warning("Cannot send callback: user_id not set")
|
||||
return False
|
||||
|
||||
end_time = datetime.now(timezone.utc)
|
||||
running_time = (end_time - self.start_time).total_seconds()
|
||||
final_tools_used = tools_used if tools_used is not None else self.tools_used
|
||||
|
||||
success = self._send_callback(
|
||||
running_time_seconds=int(running_time),
|
||||
start_time=self.start_time,
|
||||
end_time=end_time,
|
||||
tools_used=final_tools_used
|
||||
)
|
||||
|
||||
self.start_time = None
|
||||
self.tools_used = []
|
||||
self.request_id = None
|
||||
|
||||
return success
|
||||
|
||||
def _send_callback(
|
||||
self,
|
||||
running_time_seconds: int,
|
||||
start_time: datetime,
|
||||
end_time: datetime,
|
||||
tools_used: List[str]
|
||||
) -> bool:
|
||||
try:
|
||||
payload = {
|
||||
"agentName": self.agent_name,
|
||||
"userId": self.user_id,
|
||||
"podRunningTimeSeconds": running_time_seconds,
|
||||
"toolsUsed": tools_used,
|
||||
"startTime": start_time.isoformat(),
|
||||
"endTime": end_time.isoformat(),
|
||||
"requestId": self.request_id
|
||||
}
|
||||
|
||||
logger.info("Sending callback: %s", payload)
|
||||
|
||||
response = requests.post(
|
||||
self.callback_url,
|
||||
json=payload,
|
||||
timeout=5
|
||||
)
|
||||
|
||||
if response.status_code == 200:
|
||||
logger.info("Callback sent successfully: %s", response.json())
|
||||
return True
|
||||
|
||||
logger.error("Callback failed with status %s: %s", response.status_code, response.text)
|
||||
return False
|
||||
|
||||
except requests.exceptions.RequestException as e:
|
||||
logger.error("Failed to send callback: %s", str(e))
|
||||
return False
|
||||
except Exception as e:
|
||||
logger.error("Unexpected error sending callback: %s", str(e))
|
||||
return False
|
||||
|
||||
|
||||
class CallbackContextManager:
|
||||
"""回调上下文管理器 - 使用with语句自动处理开始和结束"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
handler: AgentCallbackHandler,
|
||||
request_id: Optional[str] = None,
|
||||
user_id: Optional[str] = None,
|
||||
tools_used: Optional[List[str]] = None
|
||||
):
|
||||
self.handler = handler
|
||||
self.request_id = request_id
|
||||
self.user_id = user_id
|
||||
self.tools_used = tools_used or []
|
||||
|
||||
def __enter__(self):
|
||||
self.handler.start_request(
|
||||
request_id=self.request_id,
|
||||
user_id=self.user_id
|
||||
)
|
||||
return self
|
||||
|
||||
def __exit__(self, exc_type, exc_val, exc_tb):
|
||||
self.handler.end_request(tools_used=self.tools_used)
|
||||
return False
|
||||
|
||||
def add_tool(self, tool_name: str):
|
||||
self.handler.add_tool_used(tool_name)
|
||||
if tool_name not in self.tools_used:
|
||||
self.tools_used.append(tool_name)
|
||||
@@ -0,0 +1,20 @@
|
||||
# Pydantic AI
|
||||
pydantic-ai>=0.0.14
|
||||
|
||||
# MCP
|
||||
mcp>=0.9.0
|
||||
fastmcp>=0.1.0
|
||||
|
||||
# FastAPI
|
||||
fastapi>=0.109.0
|
||||
uvicorn[standard]>=0.27.0
|
||||
|
||||
# HTTP Client
|
||||
aiohttp>=3.9.0
|
||||
requests>=2.31.0
|
||||
|
||||
# Git operations
|
||||
gitpython>=3.1.40
|
||||
|
||||
# SSH
|
||||
paramiko>=3.4.0
|
||||
@@ -0,0 +1,17 @@
|
||||
#!/usr/bin/env python
|
||||
"""启动 Code Manager Agent API 服务器"""
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).parent))
|
||||
|
||||
if __name__ == '__main__':
|
||||
from src.server.api_server import app
|
||||
import uvicorn
|
||||
import os
|
||||
|
||||
host = os.getenv('API_HOST', '0.0.0.0')
|
||||
port = int(os.getenv('API_PORT', '8000'))
|
||||
|
||||
print(f"Code Manager Agent API: http://{host}:{port}")
|
||||
uvicorn.run(app, host=host, port=port, log_level="info")
|
||||
@@ -0,0 +1 @@
|
||||
"""Agent 源代码包"""
|
||||
@@ -0,0 +1 @@
|
||||
"""服务器模块"""
|
||||
@@ -0,0 +1,284 @@
|
||||
"""
|
||||
Code Manager Agent - HTTP API 服务器
|
||||
|
||||
提供 REST API 和 MCP HTTP/SSE 端点。
|
||||
"""
|
||||
import json
|
||||
import uuid
|
||||
import os
|
||||
from typing import Optional, Dict, Any
|
||||
from contextlib import asynccontextmanager
|
||||
|
||||
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 common.agent_callback_utils import AgentCallbackHandler, CallbackContextManager
|
||||
from .mcp_server import TOOL_MAP, TOOL_LIST
|
||||
|
||||
# ==================== 配置 ====================
|
||||
|
||||
SERVER_NAME = "Code Manager Agent API"
|
||||
POD_NAME = os.getenv("POD_NAME", "code-manager-agent")
|
||||
USER_ID = os.getenv("USER_ID", "")
|
||||
callback_handler: Optional[AgentCallbackHandler] = None
|
||||
|
||||
|
||||
# ==================== FastAPI 应用 ====================
|
||||
|
||||
@asynccontextmanager
|
||||
async def lifespan(app: FastAPI):
|
||||
global callback_handler
|
||||
print(f"{SERVER_NAME} 启动")
|
||||
callback_handler = AgentCallbackHandler(agent_name=POD_NAME, user_id=USER_ID)
|
||||
yield
|
||||
print(f"{SERVER_NAME} 关闭")
|
||||
|
||||
app = FastAPI(
|
||||
title=SERVER_NAME,
|
||||
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
|
||||
raise HTTPException(status_code=401, detail="缺少 API Key")
|
||||
|
||||
|
||||
def get_api_key_from_request(request: Request) -> Optional[str]:
|
||||
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
|
||||
|
||||
|
||||
# ==================== 健康检查 ====================
|
||||
|
||||
@app.get("/")
|
||||
async def root():
|
||||
return {
|
||||
"service": SERVER_NAME,
|
||||
"status": "running",
|
||||
"tools": list(TOOL_MAP.keys()),
|
||||
}
|
||||
|
||||
|
||||
@app.get("/health")
|
||||
async def health():
|
||||
return {"status": "healthy", "service": SERVER_NAME}
|
||||
|
||||
|
||||
# ==================== MCP 端点 ====================
|
||||
|
||||
sessions: Dict[str, Dict] = {}
|
||||
|
||||
|
||||
async def run_with_callback(tool_name: str, func, *args, user_id: Optional[str] = None, request_id: Optional[str] = None, **kwargs):
|
||||
if not callback_handler:
|
||||
return await func(*args, **kwargs)
|
||||
with CallbackContextManager(
|
||||
handler=callback_handler,
|
||||
user_id=user_id or USER_ID,
|
||||
request_id=request_id or f"{tool_name}-{uuid.uuid4().hex}"
|
||||
) as ctx:
|
||||
ctx.add_tool(tool_name)
|
||||
return await func(*args, **kwargs)
|
||||
|
||||
|
||||
async def handle_mcp_request(data: Dict, session_id: str = None, api_key: str = None) -> Dict:
|
||||
method = data.get("method")
|
||||
params = data.get("params", {})
|
||||
req_id = data.get("id")
|
||||
|
||||
if method == "tools/call" and (not api_key or api_key == "sk"):
|
||||
return {"jsonrpc": "2.0", "id": req_id, "error": {"code": -32001, "message": "缺少 API Key"}}
|
||||
|
||||
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")
|
||||
arguments = params.get("arguments", {})
|
||||
if tool_name not in TOOL_MAP:
|
||||
return {"jsonrpc": "2.0", "id": req_id, "error": {"code": -32602, "message": f"未知工具: {tool_name}"}}
|
||||
old_key = os.environ.get("OPENAI_API_KEY")
|
||||
os.environ["OPENAI_API_KEY"] = api_key
|
||||
try:
|
||||
result = await run_with_callback(
|
||||
tool_name, TOOL_MAP[tool_name],
|
||||
request_id=f"mcp-{tool_name}-{uuid.uuid4().hex}",
|
||||
**arguments
|
||||
)
|
||||
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)}]}}
|
||||
else:
|
||||
return {"jsonrpc": "2.0", "id": req_id, "error": {"code": -32601, "message": f"未知方法: {method}"}}
|
||||
except Exception as e:
|
||||
return {"jsonrpc": "2.0", "id": req_id, "error": {"code": -32603, "message": str(e)}}
|
||||
|
||||
|
||||
@app.post("/mcp")
|
||||
async def mcp_http(request: Request):
|
||||
api_key = get_api_key_from_request(request)
|
||||
data = await request.json()
|
||||
result = await handle_mcp_request(data, api_key=api_key)
|
||||
return JSONResponse(content=result)
|
||||
|
||||
|
||||
@app.get("/sse")
|
||||
async def mcp_sse(request: Request):
|
||||
session_id = str(uuid.uuid4())
|
||||
sessions[session_id] = {}
|
||||
api_key = get_api_key_from_request(request)
|
||||
|
||||
async def event_stream():
|
||||
yield f"data: {json.dumps({'type': 'session', 'sessionId': session_id})}\n\n"
|
||||
|
||||
return StreamingResponse(event_stream(), media_type="text/event-stream")
|
||||
|
||||
|
||||
@app.post("/sse/{session_id}")
|
||||
async def mcp_sse_message(session_id: str, request: Request):
|
||||
api_key = get_api_key_from_request(request)
|
||||
data = await request.json()
|
||||
result = await handle_mcp_request(data, session_id=session_id, api_key=api_key)
|
||||
return JSONResponse(content=result)
|
||||
|
||||
|
||||
# ==================== 业务 API 端点 ====================
|
||||
|
||||
class GitRequest(BaseModel):
|
||||
username: str
|
||||
password: str
|
||||
local_path: Optional[str] = None
|
||||
branch: Optional[str] = None
|
||||
commit_message: Optional[str] = None
|
||||
|
||||
|
||||
class UpdateCodeRequest(BaseModel):
|
||||
task: str
|
||||
file_path: str
|
||||
local_path: Optional[str] = None
|
||||
context_files: Optional[list] = None
|
||||
|
||||
|
||||
class SSHExecRequest(BaseModel):
|
||||
host: str
|
||||
username: str
|
||||
command: str
|
||||
password: Optional[str] = None
|
||||
ssh_key_path: Optional[str] = None
|
||||
port: int = 22
|
||||
|
||||
|
||||
class SSHCloneTestRequest(BaseModel):
|
||||
host: str
|
||||
ssh_username: str
|
||||
remote_work_dir: str
|
||||
test_command: str
|
||||
gitee_username: str
|
||||
gitee_password: str
|
||||
ssh_password: Optional[str] = None
|
||||
ssh_key_path: Optional[str] = None
|
||||
ssh_port: int = 22
|
||||
branch: Optional[str] = None
|
||||
|
||||
|
||||
@app.post("/api/v1/git/pull")
|
||||
async def api_git_pull(req: GitRequest, api_key: str = Depends(verify_api_key)):
|
||||
result = await run_with_callback(
|
||||
"git_pull", TOOL_MAP["git_pull"],
|
||||
username=req.username, password=req.password,
|
||||
local_path=req.local_path, branch=req.branch,
|
||||
request_id=f"api-git-pull-{uuid.uuid4().hex}"
|
||||
)
|
||||
return JSONResponse(content=json.loads(result))
|
||||
|
||||
|
||||
@app.post("/api/v1/git/push")
|
||||
async def api_git_push(req: GitRequest, api_key: str = Depends(verify_api_key)):
|
||||
result = await run_with_callback(
|
||||
"git_push", TOOL_MAP["git_push"],
|
||||
username=req.username, password=req.password,
|
||||
local_path=req.local_path, branch=req.branch,
|
||||
commit_message=req.commit_message,
|
||||
request_id=f"api-git-push-{uuid.uuid4().hex}"
|
||||
)
|
||||
return JSONResponse(content=json.loads(result))
|
||||
|
||||
|
||||
@app.post("/api/v1/code/update")
|
||||
async def api_update_code(req: UpdateCodeRequest, api_key: str = Depends(verify_api_key)):
|
||||
result = await run_with_callback(
|
||||
"update_code", TOOL_MAP["update_code"],
|
||||
task=req.task, file_path=req.file_path,
|
||||
local_path=req.local_path, context_files=req.context_files,
|
||||
api_key=api_key,
|
||||
request_id=f"api-update-code-{uuid.uuid4().hex}"
|
||||
)
|
||||
return JSONResponse(content=json.loads(result))
|
||||
|
||||
|
||||
@app.post("/api/v1/ssh/exec")
|
||||
async def api_ssh_exec(req: SSHExecRequest, api_key: str = Depends(verify_api_key)):
|
||||
result = await run_with_callback(
|
||||
"ssh_exec", TOOL_MAP["ssh_exec"],
|
||||
host=req.host, username=req.username, command=req.command,
|
||||
password=req.password, ssh_key_path=req.ssh_key_path, port=req.port,
|
||||
request_id=f"api-ssh-exec-{uuid.uuid4().hex}"
|
||||
)
|
||||
return JSONResponse(content=json.loads(result))
|
||||
|
||||
|
||||
@app.post("/api/v1/ssh/clone-and-test")
|
||||
async def api_ssh_clone_test(req: SSHCloneTestRequest, api_key: str = Depends(verify_api_key)):
|
||||
result = await run_with_callback(
|
||||
"ssh_git_clone_and_test", TOOL_MAP["ssh_git_clone_and_test"],
|
||||
host=req.host, ssh_username=req.ssh_username,
|
||||
remote_work_dir=req.remote_work_dir, test_command=req.test_command,
|
||||
gitee_username=req.gitee_username, gitee_password=req.gitee_password,
|
||||
ssh_password=req.ssh_password, ssh_key_path=req.ssh_key_path,
|
||||
ssh_port=req.ssh_port, branch=req.branch,
|
||||
request_id=f"api-ssh-clone-test-{uuid.uuid4().hex}"
|
||||
)
|
||||
return JSONResponse(content=json.loads(result))
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
import uvicorn
|
||||
uvicorn.run(app, host="0.0.0.0", port=8000)
|
||||
@@ -0,0 +1,591 @@
|
||||
"""
|
||||
Code Manager Agent - MCP 服务器
|
||||
|
||||
提供代码仓库管理工具:
|
||||
- Gitee 仓库 pull/push(用户名+密码认证)
|
||||
- 本地代码更新
|
||||
- SSH 远程连接、git clone 与调试
|
||||
"""
|
||||
import json
|
||||
import os
|
||||
import subprocess
|
||||
import glob as glob_module
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Optional, List
|
||||
|
||||
import paramiko
|
||||
from mcp.server.fastmcp import FastMCP
|
||||
from pydantic_ai import Agent, RunContext
|
||||
|
||||
# ==================== 配置 ====================
|
||||
|
||||
GITEE_REPO_URL = os.getenv(
|
||||
"GITEE_REPO_URL",
|
||||
"http://gitee.ath.cx:3000/zhanggangyong/agent_management"
|
||||
)
|
||||
WORK_DIR = os.getenv("WORK_DIR", "/workspace")
|
||||
|
||||
# LLM 配置(供 vibe coding subagent 使用)
|
||||
_BASE_URL = os.getenv("OPENAI_BASE_URL",
|
||||
os.getenv("LLM_BASE_URL", "https://litellm.graystone-fb459c5d.southeastasia.azurecontainerapps.io/v1"))
|
||||
_API_KEY = os.getenv("OPENAI_API_KEY", "sk")
|
||||
os.environ.setdefault("OPENAI_API_KEY", _API_KEY)
|
||||
os.environ.setdefault("OPENAI_BASE_URL", _BASE_URL)
|
||||
|
||||
|
||||
def _get_model_name() -> str:
|
||||
model = os.getenv("MODEL_NAME", os.getenv("LITELLM_MODEL", "taiji/gpt-4o-mini"))
|
||||
return model if ":" in model else f"openai:{model}"
|
||||
|
||||
|
||||
CODE_AGENT_SYSTEM_PROMPT = """\
|
||||
You are an expert software engineer acting as a vibe coding agent.
|
||||
You have access to these tools:
|
||||
|
||||
- read_file(path): Read the full content of a file (relative to repo root).
|
||||
- list_files(pattern): Glob-list files matching a pattern (e.g. "src/**/*.py").
|
||||
- write_file(path, content): Write (or overwrite) a file with the given complete content.
|
||||
- edit_file(path, old_str, new_str): Replace the FIRST occurrence of old_str with new_str in a file.
|
||||
Use this for surgical edits; always verify the old_str is unique enough.
|
||||
- run_bash(command): Run a shell command in the repo root and get stdout/stderr.
|
||||
|
||||
Workflow:
|
||||
1. Start by reading the relevant files to understand the codebase.
|
||||
2. Use list_files to explore when you don't know which files to touch.
|
||||
3. Make changes using write_file (new files or full rewrites) or edit_file (surgical changes).
|
||||
4. Use run_bash to verify (e.g. run tests, linters) if applicable.
|
||||
5. When done, call finish(summary) with a human-readable summary of all changes made.
|
||||
|
||||
Rules:
|
||||
- Never output code blocks as plain text — always use the write_file or edit_file tools.
|
||||
- Preserve existing code style and conventions.
|
||||
- Prefer edit_file for small targeted changes; use write_file for new files or large rewrites.
|
||||
- Always finish with the finish() tool call.
|
||||
"""
|
||||
|
||||
|
||||
@dataclass
|
||||
class CodingAgentContext:
|
||||
repo_root: str
|
||||
changed_files: List[str] = field(default_factory=list)
|
||||
log: List[dict] = field(default_factory=list)
|
||||
|
||||
# ==================== MCP 服务器 ====================
|
||||
|
||||
server = FastMCP("Code Manager Agent")
|
||||
|
||||
|
||||
# ==================== 工具函数 ====================
|
||||
|
||||
def _run_cmd(cmd: list[str], cwd: Optional[str] = None, env: Optional[dict] = None) -> dict:
|
||||
"""执行本地命令,返回 stdout/stderr/returncode"""
|
||||
merged_env = {**os.environ, **(env or {})}
|
||||
result = subprocess.run(
|
||||
cmd, cwd=cwd, env=merged_env,
|
||||
capture_output=True, text=True, timeout=120
|
||||
)
|
||||
return {
|
||||
"returncode": result.returncode,
|
||||
"stdout": result.stdout.strip(),
|
||||
"stderr": result.stderr.strip(),
|
||||
}
|
||||
|
||||
|
||||
def _inject_credentials(repo_url: str, username: str, password: str) -> str:
|
||||
"""将用户名/密码注入 HTTP(S) URL"""
|
||||
if repo_url.startswith("http://"):
|
||||
return repo_url.replace("http://", f"http://{username}:{password}@", 1)
|
||||
if repo_url.startswith("https://"):
|
||||
return repo_url.replace("https://", f"https://{username}:{password}@", 1)
|
||||
return repo_url
|
||||
|
||||
|
||||
# ==================== MCP 工具定义 ====================
|
||||
|
||||
@server.tool()
|
||||
async def git_pull(
|
||||
username: str,
|
||||
password: str,
|
||||
local_path: Optional[str] = None,
|
||||
branch: Optional[str] = None,
|
||||
) -> str:
|
||||
"""
|
||||
从 Gitee 仓库拉取最新代码(HTTP 用户名+密码认证)。
|
||||
|
||||
Args:
|
||||
username: Gitee 用户名
|
||||
password: Gitee 密码
|
||||
local_path: 本地仓库路径,默认使用 WORK_DIR 环境变量
|
||||
branch: 分支名,默认拉取当前分支
|
||||
|
||||
Returns:
|
||||
操作结果(JSON 格式)
|
||||
"""
|
||||
try:
|
||||
cwd = local_path or WORK_DIR
|
||||
if not os.path.isdir(os.path.join(cwd, ".git")):
|
||||
return json.dumps({"success": False, "error": f"{cwd} 不是一个 git 仓库"}, ensure_ascii=False)
|
||||
|
||||
auth_url = _inject_credentials(GITEE_REPO_URL, username, password)
|
||||
# 设置 remote url(含凭据),pull 后恢复原始 url
|
||||
_run_cmd(["git", "remote", "set-url", "origin", auth_url], cwd=cwd)
|
||||
|
||||
cmd = ["git", "pull", "origin"]
|
||||
if branch:
|
||||
cmd.append(branch)
|
||||
result = _run_cmd(cmd, cwd=cwd)
|
||||
|
||||
# 恢复不含密码的 url
|
||||
_run_cmd(["git", "remote", "set-url", "origin", GITEE_REPO_URL], cwd=cwd)
|
||||
|
||||
return json.dumps({
|
||||
"success": result["returncode"] == 0,
|
||||
"stdout": result["stdout"],
|
||||
"stderr": result["stderr"],
|
||||
}, ensure_ascii=False, indent=2)
|
||||
except Exception as e:
|
||||
return json.dumps({"success": False, "error": str(e)}, ensure_ascii=False)
|
||||
|
||||
|
||||
@server.tool()
|
||||
async def git_push(
|
||||
username: str,
|
||||
password: str,
|
||||
local_path: Optional[str] = None,
|
||||
branch: Optional[str] = None,
|
||||
commit_message: Optional[str] = None,
|
||||
) -> str:
|
||||
"""
|
||||
将本地代码提交并推送到 Gitee 仓库(HTTP 用户名+密码认证)。
|
||||
|
||||
Args:
|
||||
username: Gitee 用户名
|
||||
password: Gitee 密码
|
||||
local_path: 本地仓库路径,默认使用 WORK_DIR
|
||||
branch: 目标分支,默认推送当前分支
|
||||
commit_message: 提交信息,若为空则只 push 不 commit
|
||||
|
||||
Returns:
|
||||
操作结果(JSON 格式)
|
||||
"""
|
||||
try:
|
||||
cwd = local_path or WORK_DIR
|
||||
if not os.path.isdir(os.path.join(cwd, ".git")):
|
||||
return json.dumps({"success": False, "error": f"{cwd} 不是一个 git 仓库"}, ensure_ascii=False)
|
||||
|
||||
logs = []
|
||||
|
||||
if commit_message:
|
||||
r = _run_cmd(["git", "add", "-A"], cwd=cwd)
|
||||
logs.append({"step": "git add", **r})
|
||||
r = _run_cmd(["git", "commit", "-m", commit_message], cwd=cwd)
|
||||
logs.append({"step": "git commit", **r})
|
||||
if r["returncode"] != 0 and "nothing to commit" not in r["stdout"]:
|
||||
return json.dumps({"success": False, "logs": logs}, ensure_ascii=False, indent=2)
|
||||
|
||||
auth_url = _inject_credentials(GITEE_REPO_URL, username, password)
|
||||
_run_cmd(["git", "remote", "set-url", "origin", auth_url], cwd=cwd)
|
||||
|
||||
push_cmd = ["git", "push", "origin"]
|
||||
if branch:
|
||||
push_cmd.append(branch)
|
||||
r = _run_cmd(push_cmd, cwd=cwd)
|
||||
logs.append({"step": "git push", **r})
|
||||
|
||||
_run_cmd(["git", "remote", "set-url", "origin", GITEE_REPO_URL], cwd=cwd)
|
||||
|
||||
return json.dumps({
|
||||
"success": r["returncode"] == 0,
|
||||
"logs": logs,
|
||||
}, ensure_ascii=False, indent=2)
|
||||
except Exception as e:
|
||||
return json.dumps({"success": False, "error": str(e)}, ensure_ascii=False)
|
||||
|
||||
|
||||
# ==================== Vibe Coding Subagent ====================
|
||||
|
||||
def _make_coding_agent(repo_root: str) -> Agent:
|
||||
"""构建带 read/write/edit/bash/finish 工具的 coding agent"""
|
||||
coding_agent: Agent[CodingAgentContext] = Agent(
|
||||
_get_model_name(),
|
||||
system_prompt=CODE_AGENT_SYSTEM_PROMPT,
|
||||
deps_type=CodingAgentContext,
|
||||
)
|
||||
|
||||
@coding_agent.tool
|
||||
async def read_file(ctx: RunContext[CodingAgentContext], path: str) -> str:
|
||||
"""Read a file. path is relative to repo root."""
|
||||
abs_path = os.path.join(ctx.deps.repo_root, path)
|
||||
try:
|
||||
with open(abs_path, "r", encoding="utf-8", errors="replace") as f:
|
||||
content = f.read()
|
||||
ctx.deps.log.append({"tool": "read_file", "path": path, "bytes": len(content)})
|
||||
return content
|
||||
except FileNotFoundError:
|
||||
return f"(file not found: {path})"
|
||||
|
||||
@coding_agent.tool
|
||||
async def list_files(ctx: RunContext[CodingAgentContext], pattern: str) -> str:
|
||||
"""Glob-list files matching pattern relative to repo root. Returns newline-separated paths."""
|
||||
base = ctx.deps.repo_root
|
||||
matches = glob_module.glob(os.path.join(base, pattern), recursive=True)
|
||||
rel = [os.path.relpath(m, base) for m in sorted(matches)]
|
||||
ctx.deps.log.append({"tool": "list_files", "pattern": pattern, "count": len(rel)})
|
||||
return "\n".join(rel) if rel else "(no matches)"
|
||||
|
||||
@coding_agent.tool
|
||||
async def write_file(ctx: RunContext[CodingAgentContext], path: str, content: str) -> str:
|
||||
"""Write complete content to a file (creates or overwrites). path is relative to repo root."""
|
||||
abs_path = os.path.join(ctx.deps.repo_root, path)
|
||||
os.makedirs(os.path.dirname(abs_path), exist_ok=True)
|
||||
with open(abs_path, "w", encoding="utf-8") as f:
|
||||
f.write(content)
|
||||
if path not in ctx.deps.changed_files:
|
||||
ctx.deps.changed_files.append(path)
|
||||
ctx.deps.log.append({"tool": "write_file", "path": path, "bytes": len(content.encode())})
|
||||
return f"Written {len(content.encode())} bytes to {path}"
|
||||
|
||||
@coding_agent.tool
|
||||
async def edit_file(ctx: RunContext[CodingAgentContext], path: str, old_str: str, new_str: str) -> str:
|
||||
"""Replace the FIRST occurrence of old_str with new_str in a file. path is relative to repo root."""
|
||||
abs_path = os.path.join(ctx.deps.repo_root, path)
|
||||
try:
|
||||
with open(abs_path, "r", encoding="utf-8", errors="replace") as f:
|
||||
original = f.read()
|
||||
except FileNotFoundError:
|
||||
return f"Error: file not found: {path}"
|
||||
if old_str not in original:
|
||||
return f"Error: old_str not found in {path}. No changes made."
|
||||
updated = original.replace(old_str, new_str, 1)
|
||||
with open(abs_path, "w", encoding="utf-8") as f:
|
||||
f.write(updated)
|
||||
if path not in ctx.deps.changed_files:
|
||||
ctx.deps.changed_files.append(path)
|
||||
ctx.deps.log.append({"tool": "edit_file", "path": path})
|
||||
return f"Edited {path} successfully."
|
||||
|
||||
@coding_agent.tool
|
||||
async def run_bash(ctx: RunContext[CodingAgentContext], command: str) -> str:
|
||||
"""Run a shell command in the repo root. Returns stdout + stderr."""
|
||||
r = _run_cmd(["bash", "-c", command], cwd=ctx.deps.repo_root)
|
||||
ctx.deps.log.append({"tool": "run_bash", "command": command, "returncode": r["returncode"]})
|
||||
output = ""
|
||||
if r["stdout"]:
|
||||
output += r["stdout"]
|
||||
if r["stderr"]:
|
||||
output += ("\n" if output else "") + r["stderr"]
|
||||
return output or f"(exit code {r['returncode']})"
|
||||
|
||||
@coding_agent.tool
|
||||
async def finish(ctx: RunContext[CodingAgentContext], summary: str) -> str:
|
||||
"""Call this when all changes are done. Provide a human-readable summary of what was changed."""
|
||||
ctx.deps.log.append({"tool": "finish", "summary": summary})
|
||||
return f"DONE: {summary}"
|
||||
|
||||
return coding_agent
|
||||
|
||||
|
||||
@server.tool()
|
||||
async def update_code(
|
||||
task: str,
|
||||
file_path: str,
|
||||
local_path: Optional[str] = None,
|
||||
context_files: Optional[List[str]] = None,
|
||||
api_key: Optional[str] = None,
|
||||
) -> str:
|
||||
"""
|
||||
Vibe coding subagent:接收自然语言任务,自主读取文件、理解代码、
|
||||
通过 read/write/edit/bash 工具多轮迭代完成代码变更并写回磁盘。
|
||||
参考 pi-mono 的 coding agent 设计。
|
||||
|
||||
Args:
|
||||
task: 自然语言任务描述,例如 "给 login 函数增加 JWT 验证"
|
||||
file_path: 任务入口文件(相对于仓库根目录),agent 会自行探索相关文件
|
||||
local_path: 本地仓库根路径,默认使用 WORK_DIR
|
||||
context_files: 可选的初始上下文文件列表,agent 启动时预先加载
|
||||
api_key: LLM API Key,不传则使用环境变量
|
||||
|
||||
Returns:
|
||||
JSON,包含修改的文件列表、工具调用日志、整体摘要
|
||||
"""
|
||||
try:
|
||||
repo_root = local_path or WORK_DIR
|
||||
|
||||
if api_key:
|
||||
os.environ["OPENAI_API_KEY"] = api_key
|
||||
|
||||
deps = CodingAgentContext(repo_root=repo_root)
|
||||
coding_agent = _make_coding_agent(repo_root)
|
||||
|
||||
# 构建初始 prompt
|
||||
initial_prompt_parts = [
|
||||
f"TASK: {task}",
|
||||
f"REPO ROOT: {repo_root}",
|
||||
f"START BY READING: {file_path}",
|
||||
]
|
||||
if context_files:
|
||||
initial_prompt_parts.append("ALSO CONSIDER: " + ", ".join(context_files))
|
||||
initial_prompt_parts.append(
|
||||
"\nExplore the codebase as needed, make all required changes, then call finish()."
|
||||
)
|
||||
user_prompt = "\n".join(initial_prompt_parts)
|
||||
|
||||
result = await coding_agent.run(user_prompt, deps=deps)
|
||||
|
||||
# 从日志中提取 finish summary
|
||||
summary = ""
|
||||
for entry in reversed(deps.log):
|
||||
if entry.get("tool") == "finish":
|
||||
summary = entry.get("summary", "")
|
||||
break
|
||||
|
||||
return json.dumps({
|
||||
"success": True,
|
||||
"task": task,
|
||||
"files_changed": deps.changed_files,
|
||||
"tool_log": deps.log,
|
||||
"summary": summary,
|
||||
}, ensure_ascii=False, indent=2)
|
||||
|
||||
except Exception as e:
|
||||
return json.dumps({"success": False, "error": str(e)}, ensure_ascii=False)
|
||||
|
||||
|
||||
@server.tool()
|
||||
async def ssh_exec(
|
||||
host: str,
|
||||
username: str,
|
||||
command: str,
|
||||
password: Optional[str] = None,
|
||||
ssh_key_path: Optional[str] = None,
|
||||
port: int = 22,
|
||||
) -> str:
|
||||
"""
|
||||
通过 SSH 连接远程机器并执行命令。
|
||||
|
||||
Args:
|
||||
host: 远程主机 IP 或域名
|
||||
username: SSH 用户名
|
||||
command: 要执行的 shell 命令
|
||||
password: SSH 密码(与 ssh_key_path 二选一)
|
||||
ssh_key_path: SSH 私钥文件路径(与 password 二选一)
|
||||
port: SSH 端口,默认 22
|
||||
|
||||
Returns:
|
||||
命令执行结果(JSON 格式)
|
||||
"""
|
||||
try:
|
||||
client = paramiko.SSHClient()
|
||||
client.set_missing_host_key_policy(paramiko.AutoAddPolicy())
|
||||
|
||||
connect_kwargs: dict = {"hostname": host, "port": port, "username": username, "timeout": 30}
|
||||
if ssh_key_path:
|
||||
connect_kwargs["key_filename"] = ssh_key_path
|
||||
elif password:
|
||||
connect_kwargs["password"] = password
|
||||
else:
|
||||
return json.dumps({"success": False, "error": "需要提供 password 或 ssh_key_path"}, ensure_ascii=False)
|
||||
|
||||
client.connect(**connect_kwargs)
|
||||
_, stdout, stderr = client.exec_command(command, timeout=120)
|
||||
out = stdout.read().decode(errors="replace").strip()
|
||||
err = stderr.read().decode(errors="replace").strip()
|
||||
exit_code = stdout.channel.recv_exit_status()
|
||||
client.close()
|
||||
|
||||
return json.dumps({
|
||||
"success": exit_code == 0,
|
||||
"exit_code": exit_code,
|
||||
"stdout": out,
|
||||
"stderr": err,
|
||||
}, ensure_ascii=False, indent=2)
|
||||
except Exception as e:
|
||||
return json.dumps({"success": False, "error": str(e)}, ensure_ascii=False)
|
||||
|
||||
|
||||
@server.tool()
|
||||
async def ssh_git_clone_and_test(
|
||||
host: str,
|
||||
ssh_username: str,
|
||||
remote_work_dir: str,
|
||||
test_command: str,
|
||||
gitee_username: str,
|
||||
gitee_password: str,
|
||||
ssh_password: Optional[str] = None,
|
||||
ssh_key_path: Optional[str] = None,
|
||||
ssh_port: int = 22,
|
||||
branch: Optional[str] = None,
|
||||
) -> str:
|
||||
"""
|
||||
SSH 连接到测试机器,git clone 代码仓库,然后执行测试命令。
|
||||
|
||||
Args:
|
||||
host: 测试机器 IP 或域名
|
||||
ssh_username: SSH 用户名
|
||||
remote_work_dir: 远程机器上的工作目录(clone 目标目录的父目录)
|
||||
test_command: clone 完成后要执行的测试命令(在仓库目录内执行)
|
||||
gitee_username: Gitee 用户名(用于 clone 认证)
|
||||
gitee_password: Gitee 密码(用于 clone 认证)
|
||||
ssh_password: SSH 密码(与 ssh_key_path 二选一)
|
||||
ssh_key_path: SSH 私钥文件路径
|
||||
ssh_port: SSH 端口,默认 22
|
||||
branch: 要 clone 的分支,默认主分支
|
||||
|
||||
Returns:
|
||||
各步骤执行结果(JSON 格式)
|
||||
"""
|
||||
try:
|
||||
client = paramiko.SSHClient()
|
||||
client.set_missing_host_key_policy(paramiko.AutoAddPolicy())
|
||||
|
||||
connect_kwargs: dict = {"hostname": host, "port": ssh_port, "username": ssh_username, "timeout": 30}
|
||||
if ssh_key_path:
|
||||
connect_kwargs["key_filename"] = ssh_key_path
|
||||
elif ssh_password:
|
||||
connect_kwargs["password"] = ssh_password
|
||||
else:
|
||||
return json.dumps({"success": False, "error": "需要提供 ssh_password 或 ssh_key_path"}, ensure_ascii=False)
|
||||
|
||||
client.connect(**connect_kwargs)
|
||||
|
||||
def run_remote(cmd: str) -> dict:
|
||||
_, stdout, stderr = client.exec_command(cmd, timeout=180)
|
||||
out = stdout.read().decode(errors="replace").strip()
|
||||
err = stderr.read().decode(errors="replace").strip()
|
||||
code = stdout.channel.recv_exit_status()
|
||||
return {"exit_code": code, "stdout": out, "stderr": err}
|
||||
|
||||
logs = []
|
||||
|
||||
# 1. 确保工作目录存在
|
||||
r = run_remote(f"mkdir -p {remote_work_dir}")
|
||||
logs.append({"step": "mkdir", **r})
|
||||
|
||||
# 2. 确定 repo 名称,拼接 clone url
|
||||
repo_name = GITEE_REPO_URL.rstrip("/").split("/")[-1]
|
||||
auth_url = _inject_credentials(GITEE_REPO_URL, gitee_username, gitee_password)
|
||||
clone_cmd = f"cd {remote_work_dir} && rm -rf {repo_name} && git clone"
|
||||
if branch:
|
||||
clone_cmd += f" -b {branch}"
|
||||
clone_cmd += f" {auth_url}"
|
||||
|
||||
r = run_remote(clone_cmd)
|
||||
logs.append({"step": "git clone", "exit_code": r["exit_code"],
|
||||
"stdout": r["stdout"], "stderr": r["stderr"]})
|
||||
if r["exit_code"] != 0:
|
||||
client.close()
|
||||
return json.dumps({"success": False, "logs": logs}, ensure_ascii=False, indent=2)
|
||||
|
||||
# 3. 执行测试命令
|
||||
r = run_remote(f"cd {remote_work_dir}/{repo_name} && {test_command}")
|
||||
logs.append({"step": "test", **r})
|
||||
|
||||
client.close()
|
||||
return json.dumps({
|
||||
"success": r["exit_code"] == 0,
|
||||
"logs": logs,
|
||||
}, ensure_ascii=False, indent=2)
|
||||
except Exception as e:
|
||||
return json.dumps({"success": False, "error": str(e)}, ensure_ascii=False)
|
||||
|
||||
|
||||
# ==================== 工具映射(供 API 使用)====================
|
||||
|
||||
TOOL_MAP = {
|
||||
"git_pull": git_pull,
|
||||
"git_push": git_push,
|
||||
"update_code": update_code,
|
||||
"ssh_exec": ssh_exec,
|
||||
"ssh_git_clone_and_test": ssh_git_clone_and_test,
|
||||
}
|
||||
|
||||
TOOL_LIST = [
|
||||
{
|
||||
"name": "git_pull",
|
||||
"description": "从 Gitee 仓库拉取最新代码(HTTP 用户名+密码认证)",
|
||||
"inputSchema": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"username": {"type": "string", "description": "Gitee 用户名"},
|
||||
"password": {"type": "string", "description": "Gitee 密码"},
|
||||
"local_path": {"type": "string", "description": "本地仓库路径"},
|
||||
"branch": {"type": "string", "description": "分支名"},
|
||||
},
|
||||
"required": ["username", "password"],
|
||||
},
|
||||
},
|
||||
{
|
||||
"name": "git_push",
|
||||
"description": "提交并推送本地代码到 Gitee 仓库(HTTP 用户名+密码认证)",
|
||||
"inputSchema": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"username": {"type": "string", "description": "Gitee 用户名"},
|
||||
"password": {"type": "string", "description": "Gitee 密码"},
|
||||
"local_path": {"type": "string", "description": "本地仓库路径"},
|
||||
"branch": {"type": "string", "description": "目标分支"},
|
||||
"commit_message": {"type": "string", "description": "提交信息"},
|
||||
},
|
||||
"required": ["username", "password"],
|
||||
},
|
||||
},
|
||||
{
|
||||
"name": "update_code",
|
||||
"description": "Vibe coding subagent:根据自然语言任务描述,自动读取文件、调用 LLM 生成代码并写回磁盘",
|
||||
"inputSchema": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"task": {"type": "string", "description": "自然语言任务描述,如 '给 login 函数增加 JWT 验证'"},
|
||||
"file_path": {"type": "string", "description": "主要修改目标文件路径(相对于仓库根目录)"},
|
||||
"local_path": {"type": "string", "description": "本地仓库根路径,默认 WORK_DIR"},
|
||||
"context_files": {
|
||||
"type": "array",
|
||||
"items": {"type": "string"},
|
||||
"description": "额外上下文文件列表(只读),帮助 agent 理解依赖关系"
|
||||
},
|
||||
"api_key": {"type": "string", "description": "LLM API Key,不传则使用环境变量"},
|
||||
},
|
||||
"required": ["task", "file_path"],
|
||||
},
|
||||
},
|
||||
{
|
||||
"name": "ssh_exec",
|
||||
"description": "通过 SSH 连接远程机器并执行命令",
|
||||
"inputSchema": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"host": {"type": "string", "description": "远程主机 IP 或域名"},
|
||||
"username": {"type": "string", "description": "SSH 用户名"},
|
||||
"command": {"type": "string", "description": "要执行的命令"},
|
||||
"password": {"type": "string", "description": "SSH 密码"},
|
||||
"ssh_key_path": {"type": "string", "description": "SSH 私钥文件路径"},
|
||||
"port": {"type": "integer", "description": "SSH 端口,默认 22"},
|
||||
},
|
||||
"required": ["host", "username", "command"],
|
||||
},
|
||||
},
|
||||
{
|
||||
"name": "ssh_git_clone_and_test",
|
||||
"description": "SSH 到测试机器,git clone 代码仓库后执行测试命令",
|
||||
"inputSchema": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"host": {"type": "string", "description": "测试机器 IP 或域名"},
|
||||
"ssh_username": {"type": "string", "description": "SSH 用户名"},
|
||||
"remote_work_dir": {"type": "string", "description": "远程工作目录"},
|
||||
"test_command": {"type": "string", "description": "测试命令"},
|
||||
"gitee_username": {"type": "string", "description": "Gitee 用户名"},
|
||||
"gitee_password": {"type": "string", "description": "Gitee 密码"},
|
||||
"ssh_password": {"type": "string", "description": "SSH 密码"},
|
||||
"ssh_key_path": {"type": "string", "description": "SSH 私钥文件路径"},
|
||||
"ssh_port": {"type": "integer", "description": "SSH 端口,默认 22"},
|
||||
"branch": {"type": "string", "description": "要 clone 的分支"},
|
||||
},
|
||||
"required": ["host", "ssh_username", "remote_work_dir", "test_command", "gitee_username", "gitee_password"],
|
||||
},
|
||||
},
|
||||
]
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
server.run()
|
||||
Binary file not shown.
@@ -53,12 +53,14 @@
|
||||
{
|
||||
"prompt": "做一份产品发布会的 5 页 PPT,主题是智能手表",
|
||||
"output_type": "ppt",
|
||||
"title": "可选标题,不填则由模型推断"
|
||||
"title": "可选标题,不填则由模型推断",
|
||||
"model": "taiji/gpt-5.2"
|
||||
}
|
||||
```
|
||||
|
||||
- `output_type`:`ppt` | `word` | `table`
|
||||
- `title`:可选
|
||||
- `model`:可选,不传时使用 `DEFAULT_LLM_MODEL`
|
||||
|
||||
### 响应示例
|
||||
|
||||
|
||||
@@ -9,7 +9,7 @@ import json
|
||||
import uuid
|
||||
import logging
|
||||
import aiohttp
|
||||
from typing import Optional, Dict
|
||||
from typing import Optional, Dict, Any, List
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
from io import BytesIO, StringIO
|
||||
@@ -33,6 +33,9 @@ except ImportError:
|
||||
|
||||
# python-pptx, python-docx, openpyxl
|
||||
from pptx import Presentation
|
||||
from pptx.dml.color import RGBColor
|
||||
from pptx.enum.shapes import MSO_AUTO_SHAPE_TYPE
|
||||
from pptx.enum.text import MSO_VERTICAL_ANCHOR, PP_ALIGN
|
||||
from pptx.util import Inches, Pt
|
||||
from docx import Document
|
||||
from openpyxl import Workbook
|
||||
@@ -120,6 +123,7 @@ class GenerateRequest(BaseModel):
|
||||
prompt: str = Field(..., description="描述要生成的内容,例如:做一个产品发布会的5页PPT / 写一份项目周报 / 做一个销售数据表")
|
||||
output_type: str = Field("ppt", description="输出类型: ppt, word, table")
|
||||
title: Optional[str] = Field(None, description="文档标题(可选,不填则由 LLM 根据 prompt 推断)")
|
||||
model: Optional[str] = Field(None, description="LLM 模型名称(可选,默认使用环境变量 DEFAULT_LLM_MODEL)")
|
||||
user_id: Optional[str] = None
|
||||
|
||||
|
||||
@@ -127,6 +131,7 @@ class GeneratePptRequest(BaseModel):
|
||||
prompt: str = Field(..., description="PPT 内容描述,例如:产品介绍、季度总结、培训大纲")
|
||||
title: Optional[str] = None
|
||||
num_slides: Optional[int] = Field(5, description="页数建议")
|
||||
model: Optional[str] = Field(None, description="LLM 模型名称(可选,默认使用环境变量 DEFAULT_LLM_MODEL)")
|
||||
user_id: Optional[str] = None
|
||||
|
||||
|
||||
@@ -156,10 +161,11 @@ async def call_llm_json(
|
||||
user_content: str,
|
||||
api_key: str,
|
||||
max_tokens: int = 2000,
|
||||
model: Optional[str] = None,
|
||||
) -> dict:
|
||||
"""调用 LLM 并解析为 JSON"""
|
||||
payload = {
|
||||
"model": DEFAULT_LLM_MODEL,
|
||||
"model": model or DEFAULT_LLM_MODEL,
|
||||
"messages": [
|
||||
{"role": "system", "content": system_prompt},
|
||||
{"role": "user", "content": user_content},
|
||||
@@ -192,46 +198,259 @@ async def call_llm_json(
|
||||
# ==================== 生成逻辑 ====================
|
||||
|
||||
|
||||
def _as_text(value: Any, default: str = "") -> str:
|
||||
if value is None:
|
||||
return default
|
||||
if isinstance(value, str):
|
||||
return value.strip() or default
|
||||
return str(value).strip() or default
|
||||
|
||||
|
||||
def _as_list(value: Any) -> List[Any]:
|
||||
if value is None:
|
||||
return []
|
||||
if isinstance(value, list):
|
||||
return value
|
||||
return [value]
|
||||
|
||||
|
||||
def _rgb_from_hex(value: Optional[str], fallback: tuple[int, int, int]) -> RGBColor:
|
||||
raw = _as_text(value)
|
||||
if raw.startswith("#"):
|
||||
raw = raw[1:]
|
||||
if len(raw) == 6:
|
||||
try:
|
||||
return RGBColor(int(raw[0:2], 16), int(raw[2:4], 16), int(raw[4:6], 16))
|
||||
except ValueError:
|
||||
pass
|
||||
return RGBColor(*fallback)
|
||||
|
||||
|
||||
def _get_ppt_palette(data: dict) -> Dict[str, RGBColor]:
|
||||
theme = data.get("theme") if isinstance(data.get("theme"), dict) else {}
|
||||
return {
|
||||
"primary": _rgb_from_hex(theme.get("primary"), (29, 78, 137)),
|
||||
"accent": _rgb_from_hex(theme.get("accent"), (56, 189, 248)),
|
||||
"background": _rgb_from_hex(theme.get("background"), (245, 247, 250)),
|
||||
"surface": _rgb_from_hex(theme.get("surface"), (255, 255, 255)),
|
||||
"text": _rgb_from_hex(theme.get("text"), (24, 24, 27)),
|
||||
"muted": _rgb_from_hex(theme.get("muted"), (82, 82, 91)),
|
||||
}
|
||||
|
||||
|
||||
def _normalize_ppt_data(data: dict) -> dict:
|
||||
slides = []
|
||||
for raw_slide in _as_list(data.get("slides")):
|
||||
if isinstance(raw_slide, str):
|
||||
raw_slide = {"title": raw_slide, "bullets": []}
|
||||
if not isinstance(raw_slide, dict):
|
||||
continue
|
||||
|
||||
bullets = [_as_text(item) for item in _as_list(raw_slide.get("bullets") or raw_slide.get("content")) if _as_text(item)]
|
||||
left_bullets = [_as_text(item) for item in _as_list(raw_slide.get("left_bullets")) if _as_text(item)]
|
||||
right_bullets = [_as_text(item) for item in _as_list(raw_slide.get("right_bullets")) if _as_text(item)]
|
||||
|
||||
stats = []
|
||||
for stat in _as_list(raw_slide.get("stats"))[:3]:
|
||||
if isinstance(stat, dict):
|
||||
stats.append(
|
||||
{
|
||||
"label": _as_text(stat.get("label")),
|
||||
"value": _as_text(stat.get("value")),
|
||||
"note": _as_text(stat.get("note")),
|
||||
}
|
||||
)
|
||||
else:
|
||||
stats.append({"label": "", "value": _as_text(stat), "note": ""})
|
||||
|
||||
layout = _as_text(raw_slide.get("layout")).lower()
|
||||
if not layout:
|
||||
if left_bullets or right_bullets:
|
||||
layout = "two_column"
|
||||
elif stats:
|
||||
layout = "highlight"
|
||||
else:
|
||||
layout = "content"
|
||||
|
||||
slides.append(
|
||||
{
|
||||
"layout": layout,
|
||||
"title": _as_text(raw_slide.get("title"), "未命名页面"),
|
||||
"subtitle": _as_text(raw_slide.get("subtitle")),
|
||||
"key_message": _as_text(raw_slide.get("key_message")),
|
||||
"bullets": bullets[:6],
|
||||
"left_title": _as_text(raw_slide.get("left_title"), "要点"),
|
||||
"left_bullets": left_bullets[:4],
|
||||
"right_title": _as_text(raw_slide.get("right_title"), "说明"),
|
||||
"right_bullets": right_bullets[:4],
|
||||
"stats": stats,
|
||||
"takeaway": _as_text(raw_slide.get("takeaway")),
|
||||
}
|
||||
)
|
||||
|
||||
return {
|
||||
"title": _as_text(data.get("title"), "未命名演示"),
|
||||
"subtitle": _as_text(data.get("subtitle"), "由 Doc Creator Agent 自动生成"),
|
||||
"slides": slides or [{"layout": "content", "title": "核心内容", "bullets": ["请提供更具体的业务目标和受众,以便生成更好的演示稿。"]}],
|
||||
}
|
||||
|
||||
|
||||
def _add_box(slide, shape_type, left, top, width, height, fill_color: RGBColor, line_color: Optional[RGBColor] = None):
|
||||
shape = slide.shapes.add_shape(shape_type, left, top, width, height)
|
||||
shape.fill.solid()
|
||||
shape.fill.fore_color.rgb = fill_color
|
||||
shape.line.color.rgb = line_color or fill_color
|
||||
shape.line.width = Pt(1)
|
||||
return shape
|
||||
|
||||
|
||||
def _add_text_block(
|
||||
slide,
|
||||
left,
|
||||
top,
|
||||
width,
|
||||
height,
|
||||
lines: List[str],
|
||||
font_size: int,
|
||||
color: RGBColor,
|
||||
bold: bool = False,
|
||||
align=PP_ALIGN.LEFT,
|
||||
space_after: int = 8,
|
||||
):
|
||||
textbox = slide.shapes.add_textbox(left, top, width, height)
|
||||
text_frame = textbox.text_frame
|
||||
text_frame.clear()
|
||||
text_frame.word_wrap = True
|
||||
text_frame.vertical_anchor = MSO_VERTICAL_ANCHOR.TOP
|
||||
text_frame.margin_left = Pt(0)
|
||||
text_frame.margin_right = Pt(0)
|
||||
text_frame.margin_top = Pt(0)
|
||||
text_frame.margin_bottom = Pt(0)
|
||||
|
||||
for idx, line in enumerate([line for line in lines if _as_text(line)]):
|
||||
paragraph = text_frame.paragraphs[0] if idx == 0 else text_frame.add_paragraph()
|
||||
paragraph.alignment = align
|
||||
paragraph.space_after = Pt(space_after)
|
||||
run = paragraph.add_run()
|
||||
run.text = line
|
||||
run.font.size = Pt(font_size)
|
||||
run.font.bold = bold
|
||||
run.font.color.rgb = color
|
||||
return textbox
|
||||
|
||||
|
||||
def _add_footer(slide, title: str, index: int, palette: Dict[str, RGBColor]):
|
||||
_add_box(slide, MSO_AUTO_SHAPE_TYPE.RECTANGLE, Inches(0), Inches(7.18), Inches(13.333), Inches(0.18), palette["accent"])
|
||||
_add_text_block(slide, Inches(0.65), Inches(6.88), Inches(10.8), Inches(0.25), [title], 10, palette["muted"])
|
||||
_add_text_block(slide, Inches(12.2), Inches(6.84), Inches(0.5), Inches(0.28), [str(index)], 11, palette["primary"], bold=True, align=PP_ALIGN.RIGHT)
|
||||
|
||||
|
||||
def _add_cover_slide(prs: Presentation, data: dict, palette: Dict[str, RGBColor]):
|
||||
slide = prs.slides.add_slide(prs.slide_layouts[6])
|
||||
_add_box(slide, MSO_AUTO_SHAPE_TYPE.RECTANGLE, Inches(0), Inches(0), Inches(13.333), Inches(7.5), palette["primary"])
|
||||
_add_box(slide, MSO_AUTO_SHAPE_TYPE.RECTANGLE, Inches(0), Inches(0), Inches(0.35), Inches(7.5), palette["accent"])
|
||||
_add_box(slide, MSO_AUTO_SHAPE_TYPE.ROUNDED_RECTANGLE, Inches(9.85), Inches(0.65), Inches(2.35), Inches(0.5), palette["accent"])
|
||||
_add_text_block(slide, Inches(10.15), Inches(0.8), Inches(1.8), Inches(0.2), ["DOC CREATOR"], 11, palette["surface"], bold=True, align=PP_ALIGN.CENTER)
|
||||
_add_text_block(slide, Inches(0.9), Inches(1.5), Inches(10.5), Inches(1.7), [_as_text(data.get("title"), "未命名演示")], 26, palette["surface"], bold=True)
|
||||
subtitle_lines = [_as_text(data.get("subtitle"), "结构化内容自动生成演示稿")]
|
||||
subtitle_lines.append(datetime.now().strftime("%Y-%m-%d"))
|
||||
_add_text_block(slide, Inches(0.95), Inches(3.35), Inches(8.8), Inches(1.1), subtitle_lines, 15, palette["surface"])
|
||||
|
||||
|
||||
def _render_bullet_card(slide, left, top, width, height, title: str, bullets: List[str], palette: Dict[str, RGBColor]):
|
||||
_add_box(slide, MSO_AUTO_SHAPE_TYPE.ROUNDED_RECTANGLE, left, top, width, height, palette["surface"], palette["accent"])
|
||||
_add_text_block(slide, left + Inches(0.28), top + Inches(0.22), width - Inches(0.56), Inches(0.38), [title], 14, palette["primary"], bold=True)
|
||||
bullet_lines = [f"• {item}" for item in bullets if _as_text(item)]
|
||||
_add_text_block(slide, left + Inches(0.28), top + Inches(0.7), width - Inches(0.56), height - Inches(0.95), bullet_lines, 15, palette["text"], space_after=10)
|
||||
|
||||
|
||||
def _render_agenda_slide(slide, slide_data: dict, palette: Dict[str, RGBColor]):
|
||||
bullets = slide_data.get("bullets") or ["核心背景", "关键分析", "行动建议"]
|
||||
start_top = Inches(2.0)
|
||||
for idx, bullet in enumerate(bullets[:5]):
|
||||
card_top = start_top + Inches(idx * 0.8)
|
||||
_add_box(slide, MSO_AUTO_SHAPE_TYPE.ROUNDED_RECTANGLE, Inches(0.9), card_top, Inches(0.75), Inches(0.48), palette["primary"])
|
||||
_add_text_block(slide, Inches(1.13), card_top + Inches(0.12), Inches(0.25), Inches(0.18), [str(idx + 1)], 13, palette["surface"], bold=True, align=PP_ALIGN.CENTER)
|
||||
_add_box(slide, MSO_AUTO_SHAPE_TYPE.ROUNDED_RECTANGLE, Inches(1.85), card_top, Inches(10.1), Inches(0.48), palette["surface"], palette["accent"])
|
||||
_add_text_block(slide, Inches(2.15), card_top + Inches(0.1), Inches(9.5), Inches(0.2), [bullet], 16, palette["text"])
|
||||
|
||||
|
||||
def _render_content_slide(slide, slide_data: dict, palette: Dict[str, RGBColor]):
|
||||
bullets = slide_data.get("bullets") or ["补充项目目标、对象和场景,让内容更贴近实际汇报。"]
|
||||
_render_bullet_card(slide, Inches(0.8), Inches(2.0), Inches(11.75), Inches(3.75), slide_data.get("subtitle") or "核心内容", bullets[:5], palette)
|
||||
|
||||
|
||||
def _render_two_column_slide(slide, slide_data: dict, palette: Dict[str, RGBColor]):
|
||||
left_bullets = slide_data.get("left_bullets") or slide_data.get("bullets", [])[:4]
|
||||
right_bullets = slide_data.get("right_bullets") or slide_data.get("bullets", [])[4:8]
|
||||
_render_bullet_card(slide, Inches(0.8), Inches(2.0), Inches(5.6), Inches(3.8), slide_data.get("left_title") or "左侧观点", left_bullets or ["请补充左侧分析要点"], palette)
|
||||
_render_bullet_card(slide, Inches(6.9), Inches(2.0), Inches(5.6), Inches(3.8), slide_data.get("right_title") or "右侧观点", right_bullets or ["请补充右侧分析要点"], palette)
|
||||
|
||||
|
||||
def _render_highlight_slide(slide, slide_data: dict, palette: Dict[str, RGBColor]):
|
||||
key_message = slide_data.get("key_message") or (slide_data.get("bullets") or ["突出一个最重要的结论"])[0]
|
||||
_add_box(slide, MSO_AUTO_SHAPE_TYPE.ROUNDED_RECTANGLE, Inches(0.8), Inches(2.0), Inches(7.0), Inches(2.15), palette["primary"], palette["primary"])
|
||||
_add_text_block(slide, Inches(1.1), Inches(2.35), Inches(6.4), Inches(1.3), [key_message], 24, palette["surface"], bold=True)
|
||||
|
||||
stats = slide_data.get("stats") or []
|
||||
stat_left = 8.15
|
||||
for idx, stat in enumerate(stats[:3]):
|
||||
top = Inches(2.0 + idx * 1.18)
|
||||
_add_box(slide, MSO_AUTO_SHAPE_TYPE.ROUNDED_RECTANGLE, Inches(stat_left), top, Inches(4.05), Inches(0.95), palette["surface"], palette["accent"])
|
||||
value = stat.get("value") or stat.get("label") or f"亮点 {idx + 1}"
|
||||
label = stat.get("label") or "指标"
|
||||
note = stat.get("note")
|
||||
_add_text_block(slide, Inches(stat_left + 0.25), top + Inches(0.15), Inches(2.0), Inches(0.3), [label], 11, palette["muted"])
|
||||
_add_text_block(slide, Inches(stat_left + 0.25), top + Inches(0.38), Inches(3.4), Inches(0.3), [value], 20, palette["primary"], bold=True)
|
||||
if note:
|
||||
_add_text_block(slide, Inches(stat_left + 0.25), top + Inches(0.7), Inches(3.4), Inches(0.18), [note], 10, palette["muted"])
|
||||
|
||||
extra_bullets = slide_data.get("bullets", [])[1:4]
|
||||
if extra_bullets:
|
||||
_render_bullet_card(slide, Inches(0.8), Inches(4.5), Inches(11.4), Inches(1.35), "支撑要点", extra_bullets, palette)
|
||||
|
||||
|
||||
def _render_summary_bar(slide, takeaway: str, palette: Dict[str, RGBColor]):
|
||||
if not takeaway:
|
||||
return
|
||||
_add_box(slide, MSO_AUTO_SHAPE_TYPE.ROUNDED_RECTANGLE, Inches(0.8), Inches(6.05), Inches(11.7), Inches(0.72), palette["accent"], palette["accent"])
|
||||
_add_text_block(slide, Inches(1.1), Inches(6.24), Inches(11.1), Inches(0.24), [f"结论: {takeaway}"], 14, palette["surface"], bold=True)
|
||||
|
||||
|
||||
def _build_ppt(data: dict) -> bytes:
|
||||
"""从结构化数据生成 PPTX 字节"""
|
||||
"""从结构化数据生成更适合汇报场景的 PPTX 字节"""
|
||||
normalized = _normalize_ppt_data(data)
|
||||
palette = _get_ppt_palette(data)
|
||||
|
||||
prs = Presentation()
|
||||
prs.slide_width = Inches(10)
|
||||
prs.slide_width = Inches(13.333)
|
||||
prs.slide_height = Inches(7.5)
|
||||
title_slide_layout = prs.slide_layouts[0]
|
||||
content_layout = prs.slide_layouts[6] # blank
|
||||
|
||||
# 标题页
|
||||
slide = prs.slides.add_slide(title_slide_layout)
|
||||
title = data.get("title", "未命名演示")
|
||||
slide.shapes.title.text = title
|
||||
if slide.placeholders[1]:
|
||||
slide.placeholders[1].text = data.get("subtitle", "")
|
||||
_add_cover_slide(prs, normalized, palette)
|
||||
|
||||
for index, slide_data in enumerate(normalized.get("slides", []), start=1):
|
||||
slide = prs.slides.add_slide(prs.slide_layouts[6])
|
||||
_add_box(slide, MSO_AUTO_SHAPE_TYPE.RECTANGLE, Inches(0), Inches(0), Inches(13.333), Inches(7.5), palette["background"])
|
||||
_add_box(slide, MSO_AUTO_SHAPE_TYPE.RECTANGLE, Inches(0), Inches(0), Inches(13.333), Inches(0.2), palette["primary"])
|
||||
|
||||
_add_text_block(slide, Inches(0.8), Inches(0.65), Inches(11.0), Inches(0.5), [slide_data.get("title") or f"第 {index} 页"], 24, palette["primary"], bold=True)
|
||||
if slide_data.get("key_message") and slide_data.get("layout") not in {"highlight", "summary"}:
|
||||
_add_box(slide, MSO_AUTO_SHAPE_TYPE.ROUNDED_RECTANGLE, Inches(0.8), Inches(1.25), Inches(11.2), Inches(0.52), palette["surface"], palette["accent"])
|
||||
_add_text_block(slide, Inches(1.08), Inches(1.4), Inches(10.5), Inches(0.2), [slide_data["key_message"]], 13, palette["muted"], bold=True)
|
||||
|
||||
layout = slide_data.get("layout")
|
||||
if layout == "agenda":
|
||||
_render_agenda_slide(slide, slide_data, palette)
|
||||
elif layout == "two_column":
|
||||
_render_two_column_slide(slide, slide_data, palette)
|
||||
elif layout in {"highlight", "summary"}:
|
||||
_render_highlight_slide(slide, slide_data, palette)
|
||||
else:
|
||||
_render_content_slide(slide, slide_data, palette)
|
||||
|
||||
_render_summary_bar(slide, slide_data.get("takeaway"), palette)
|
||||
_add_footer(slide, normalized["title"], index + 1, palette)
|
||||
|
||||
# 内容页
|
||||
slides_data = data.get("slides", [])
|
||||
for s in slides_data:
|
||||
slide = prs.slides.add_slide(content_layout)
|
||||
slide_title = s.get("title", "")
|
||||
bullets = s.get("bullets", s.get("content", []))
|
||||
if isinstance(bullets, str):
|
||||
bullets = [bullets]
|
||||
left = Inches(0.5)
|
||||
top = Inches(0.8)
|
||||
w, h = Inches(9), Inches(1.2)
|
||||
tx = slide.shapes.add_textbox(left, top, w, h)
|
||||
tf = tx.text_frame
|
||||
p = tf.paragraphs[0]
|
||||
p.text = slide_title
|
||||
p.font.size = Pt(28)
|
||||
p.font.bold = True
|
||||
for b in bullets:
|
||||
top += Inches(0.9)
|
||||
tx = slide.shapes.add_textbox(left, top, w, Inches(1.5))
|
||||
tf = tx.text_frame
|
||||
tf.word_wrap = True
|
||||
p = tf.paragraphs[0]
|
||||
p.text = b if isinstance(b, str) else str(b)
|
||||
p.font.size = Pt(18)
|
||||
buf = BytesIO()
|
||||
prs.save(buf)
|
||||
buf.seek(0)
|
||||
@@ -325,9 +544,27 @@ async def get_api_key(
|
||||
|
||||
PPT_JSON_SCHEMA = """{
|
||||
"title": "演示文稿主标题",
|
||||
"subtitle": "可选副标题",
|
||||
"subtitle": "一句话副标题,点明背景或目标",
|
||||
"theme": {
|
||||
"primary": "#1D4E89",
|
||||
"accent": "#38BDF8",
|
||||
"background": "#F5F7FA"
|
||||
},
|
||||
"slides": [
|
||||
{ "title": "每页标题", "bullets": ["要点1", "要点2", "要点3"] }
|
||||
{
|
||||
"layout": "agenda | content | two_column | highlight | summary",
|
||||
"title": "结论式页面标题",
|
||||
"key_message": "这一页最重要的一句话",
|
||||
"bullets": ["要点1", "要点2", "要点3"],
|
||||
"left_title": "左栏标题",
|
||||
"left_bullets": ["左栏要点1", "左栏要点2"],
|
||||
"right_title": "右栏标题",
|
||||
"right_bullets": ["右栏要点1", "右栏要点2"],
|
||||
"stats": [
|
||||
{"label": "指标名", "value": "数值", "note": "补充说明"}
|
||||
],
|
||||
"takeaway": "本页结论"
|
||||
}
|
||||
]
|
||||
}"""
|
||||
|
||||
@@ -348,6 +585,23 @@ TABLE_JSON_SCHEMA = """{
|
||||
}"""
|
||||
|
||||
|
||||
def _build_ppt_system_prompt(num_slides: Optional[int] = None) -> str:
|
||||
target_slides = num_slides or 5
|
||||
return f"""你是一个资深咨询顾问兼演示设计师,要把用户需求整理成一份可直接汇报的 PPT 结构。
|
||||
必须只返回一个 JSON 对象,不要返回 markdown,不要解释。
|
||||
格式严格如下:
|
||||
{PPT_JSON_SCHEMA}
|
||||
|
||||
生成要求:
|
||||
1. slides 不包含封面页,系统会自动生成封面;你只需要生成内容页。
|
||||
2. 总页数建议为 {target_slides} 页左右,至少包含 1 页 agenda 或 summary。
|
||||
3. 标题必须结论导向,避免“背景介绍”这类空泛标题。
|
||||
4. 每页 bullets 控制在 3-5 条,每条一句短句,适合展示,不要写成长段落。
|
||||
5. 需要对比时用 two_column;有关键数字或亮点时优先用 highlight。
|
||||
6. takeaway 必须是本页一句明确结论,不能重复 title。
|
||||
7. 如果用户没有指定风格,默认输出专业、简洁、适合业务汇报的内容。"""
|
||||
|
||||
|
||||
@app.post("/api/v1/generate")
|
||||
async def api_generate(request: GenerateRequest, api_key: str = Depends(get_api_key)):
|
||||
"""根据 prompt 和 output_type 生成文件(ppt / word / table)"""
|
||||
@@ -357,14 +611,11 @@ async def api_generate(request: GenerateRequest, api_key: str = Depends(get_api_
|
||||
raise HTTPException(status_code=400, detail="output_type 只能是 ppt, word, table")
|
||||
|
||||
if output_type == "ppt":
|
||||
system_prompt = f"""你是一个专业的演示文稿策划。根据用户的描述,生成一份 PPT 大纲。
|
||||
必须只返回一个 JSON 对象,不要其他文字。格式严格如下(可增加 slides 数量):
|
||||
{PPT_JSON_SCHEMA}
|
||||
bullets 为每页的要点列表。"""
|
||||
system_prompt = _build_ppt_system_prompt()
|
||||
user_content = f"用户需求:{prompt}"
|
||||
if request.title:
|
||||
user_content += f"\n主标题请使用:{request.title}"
|
||||
data = await call_llm_json(system_prompt, user_content, api_key)
|
||||
data = await call_llm_json(system_prompt, user_content, api_key, model=request.model)
|
||||
raw = _build_ppt(data)
|
||||
ext = "pptx"
|
||||
elif output_type == "word":
|
||||
@@ -375,7 +626,7 @@ sections 可多条,paragraphs 为每段的文字。"""
|
||||
user_content = f"用户需求:{prompt}"
|
||||
if request.title:
|
||||
user_content += f"\n文档标题请使用:{request.title}"
|
||||
data = await call_llm_json(system_prompt, user_content, api_key)
|
||||
data = await call_llm_json(system_prompt, user_content, api_key, model=request.model)
|
||||
raw = _build_word(data)
|
||||
ext = "docx"
|
||||
else:
|
||||
@@ -386,7 +637,7 @@ headers 和 rows 的列数要一致。"""
|
||||
user_content = f"用户需求:{prompt}"
|
||||
if request.title:
|
||||
user_content += f"\n表头或第一行标题可体现:{request.title}"
|
||||
data = await call_llm_json(system_prompt, user_content, api_key)
|
||||
data = await call_llm_json(system_prompt, user_content, api_key, model=request.model)
|
||||
raw = _build_table(data, "xlsx")
|
||||
ext = "xlsx"
|
||||
|
||||
@@ -408,14 +659,11 @@ headers 和 rows 的列数要一致。"""
|
||||
@app.post("/api/v1/generate-ppt")
|
||||
async def api_generate_ppt(request: GeneratePptRequest, api_key: str = Depends(get_api_key)):
|
||||
"""根据 prompt 生成 PPT"""
|
||||
system_prompt = f"""你是一个专业的演示文稿策划。根据用户的描述,生成 PPT 大纲。
|
||||
必须只返回一个 JSON 对象,不要其他文字。格式严格如下:
|
||||
{PPT_JSON_SCHEMA}
|
||||
slides 数量建议 {request.num_slides or 5} 页左右。"""
|
||||
system_prompt = _build_ppt_system_prompt(request.num_slides)
|
||||
user_content = f"用户需求:{request.prompt}"
|
||||
if request.title:
|
||||
user_content += f"\n主标题请使用:{request.title}"
|
||||
data = await call_llm_json(system_prompt, user_content, api_key)
|
||||
data = await call_llm_json(system_prompt, user_content, api_key, model=request.model)
|
||||
raw = _build_ppt(data)
|
||||
ts = datetime.now().strftime("%Y%m%d_%H%M%S")
|
||||
filename = f"doc_ppt_{ts}.pptx"
|
||||
@@ -554,6 +802,7 @@ MCP_TOOL_LIST = [
|
||||
"prompt": {"type": "string", "description": "描述要生成的内容,如:做一份产品发布会的5页PPT、写一份项目周报、做销售数据表"},
|
||||
"output_type": {"type": "string", "description": "输出类型: ppt, word, table"},
|
||||
"title": {"type": "string", "description": "可选文档标题"},
|
||||
"model": {"type": "string", "description": "可选 LLM 模型名称,默认使用部署时配置的 DEFAULT_LLM_MODEL"},
|
||||
},
|
||||
"required": ["prompt"],
|
||||
},
|
||||
@@ -567,6 +816,7 @@ MCP_TOOL_LIST = [
|
||||
"prompt": {"type": "string", "description": "PPT 内容描述"},
|
||||
"title": {"type": "string", "description": "可选标题"},
|
||||
"num_slides": {"type": "integer", "description": "建议页数"},
|
||||
"model": {"type": "string", "description": "可选 LLM 模型名称,默认使用部署时配置的 DEFAULT_LLM_MODEL"},
|
||||
},
|
||||
"required": ["prompt"],
|
||||
},
|
||||
@@ -614,6 +864,7 @@ async def _mcp_generate_document(api_key: str, **kwargs) -> str:
|
||||
prompt=kwargs["prompt"],
|
||||
output_type=kwargs.get("output_type", "ppt"),
|
||||
title=kwargs.get("title"),
|
||||
model=kwargs.get("model"),
|
||||
)
|
||||
result = await api_generate(req, api_key)
|
||||
return json.dumps(result, ensure_ascii=False, indent=2)
|
||||
@@ -621,7 +872,12 @@ async def _mcp_generate_document(api_key: str, **kwargs) -> str:
|
||||
|
||||
@_register_mcp("generate_ppt")
|
||||
async def _mcp_generate_ppt(api_key: str, **kwargs) -> str:
|
||||
req = GeneratePptRequest(prompt=kwargs["prompt"], title=kwargs.get("title"), num_slides=kwargs.get("num_slides"))
|
||||
req = GeneratePptRequest(
|
||||
prompt=kwargs["prompt"],
|
||||
title=kwargs.get("title"),
|
||||
num_slides=kwargs.get("num_slides"),
|
||||
model=kwargs.get("model"),
|
||||
)
|
||||
result = await api_generate_ppt(req, api_key)
|
||||
return json.dumps(result, ensure_ascii=False, indent=2)
|
||||
|
||||
|
||||
@@ -16,13 +16,15 @@ RUN apt-get update && apt-get install -y \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
|
||||
# 复制依赖文件
|
||||
COPY requirements.txt ./requirements.txt
|
||||
COPY agent_templates/agents/facebook_agent/requirements.txt ./requirements.txt
|
||||
|
||||
# 安装 Python 依赖
|
||||
RUN pip install --no-cache-dir -r requirements.txt
|
||||
RUN pip install --no-cache-dir -r requirements.txt requests
|
||||
|
||||
# 复制应用代码
|
||||
COPY . .
|
||||
COPY agent_templates/agents/facebook_agent/ /app/
|
||||
COPY agent_templates/common/agent_callback_utils.py /app/common/
|
||||
RUN touch /app/common/__init__.py
|
||||
|
||||
# 暴露端口
|
||||
# 8000: API 服务端口
|
||||
|
||||
@@ -4,6 +4,7 @@ FastAPI服务 - Facebook搜索智能Agent
|
||||
"""
|
||||
|
||||
import json
|
||||
import os
|
||||
import uuid
|
||||
from typing import Dict, Any, Optional, AsyncGenerator
|
||||
from fastapi import FastAPI, HTTPException, Request, Header, Depends
|
||||
@@ -28,6 +29,14 @@ except ImportError:
|
||||
from models.schemas import SearchRequest, SearchResponse
|
||||
from mcp_server import search_facebook, initialize_agent
|
||||
|
||||
try:
|
||||
from common.agent_callback_utils import AgentCallbackHandler, CallbackContextManager
|
||||
CALLBACK_ENABLED = True
|
||||
except ImportError:
|
||||
CALLBACK_ENABLED = False
|
||||
AgentCallbackHandler = None
|
||||
CallbackContextManager = None
|
||||
|
||||
|
||||
# ==================== FastAPI应用 ====================
|
||||
|
||||
@@ -50,6 +59,9 @@ app.add_middleware(
|
||||
# 全局变量
|
||||
config: Optional[Config] = None
|
||||
agent: Optional[FacebookAgent] = None
|
||||
callback_handler: Optional[AgentCallbackHandler] = None
|
||||
POD_NAME = os.getenv("POD_NAME", "facebook-agent")
|
||||
USER_ID = os.getenv("USER_ID", "")
|
||||
|
||||
# MCP 工具映射
|
||||
TOOL_MAP = {
|
||||
@@ -178,7 +190,16 @@ async def handle_mcp_request(request_data: Dict[str, Any], session_id: Optional[
|
||||
tool_func = TOOL_MAP[tool_name]
|
||||
|
||||
# 调用工具(异步)
|
||||
result = await tool_func(**arguments)
|
||||
if CALLBACK_ENABLED and callback_handler:
|
||||
with CallbackContextManager(
|
||||
handler=callback_handler,
|
||||
user_id=USER_ID,
|
||||
request_id=f"facebook-mcp-{tool_name}-{request_id or uuid.uuid4().hex}"
|
||||
) as ctx:
|
||||
ctx.add_tool(tool_name)
|
||||
result = await tool_func(**arguments)
|
||||
else:
|
||||
result = await tool_func(**arguments)
|
||||
finally:
|
||||
# 恢复原来的 API key
|
||||
if api_key and agent is not None and 'old_api_key' in locals():
|
||||
@@ -236,7 +257,7 @@ def setup_logger():
|
||||
@app.on_event("startup")
|
||||
async def startup_event():
|
||||
"""应用启动时初始化"""
|
||||
global config, agent
|
||||
global config, agent, callback_handler
|
||||
|
||||
try:
|
||||
# 加载配置
|
||||
@@ -248,6 +269,8 @@ async def startup_event():
|
||||
|
||||
# 创建Agent
|
||||
agent = FacebookAgent(config)
|
||||
if CALLBACK_ENABLED and AgentCallbackHandler:
|
||||
callback_handler = AgentCallbackHandler(agent_name=POD_NAME, user_id=USER_ID)
|
||||
|
||||
logger.info("=" * 60)
|
||||
logger.info("Facebook搜索智能Agent API 启动成功")
|
||||
@@ -377,8 +400,16 @@ async def search(request: SearchRequest, api_key: str = Depends(verify_api_key))
|
||||
agent.deps.llm_client = LiteLLMClient(agent.config)
|
||||
|
||||
try:
|
||||
# 执行搜索
|
||||
response = await agent.search(request)
|
||||
if CALLBACK_ENABLED and callback_handler:
|
||||
with CallbackContextManager(
|
||||
handler=callback_handler,
|
||||
user_id=USER_ID,
|
||||
request_id=f"facebook-search-{uuid.uuid4().hex}"
|
||||
) as ctx:
|
||||
ctx.add_tool("search_facebook")
|
||||
response = await agent.search(request)
|
||||
else:
|
||||
response = await agent.search(request)
|
||||
finally:
|
||||
# 恢复原来的 API key
|
||||
if api_key and 'old_api_key' in locals():
|
||||
|
||||
BIN
Binary file not shown.
@@ -14,6 +14,14 @@ from a2a.utils import new_agent_text_message
|
||||
from agent import SearchAgentWrapper
|
||||
from config import get_config
|
||||
|
||||
try:
|
||||
from common.agent_callback_utils import AgentCallbackHandler, CallbackContextManager
|
||||
CALLBACK_ENABLED = True
|
||||
except ImportError:
|
||||
CALLBACK_ENABLED = False
|
||||
AgentCallbackHandler = None
|
||||
CallbackContextManager = None
|
||||
|
||||
|
||||
class SearchAgentExecutor(AgentExecutor):
|
||||
"""
|
||||
@@ -47,6 +55,12 @@ class SearchAgentExecutor(AgentExecutor):
|
||||
|
||||
self.default_api_key = default_api_key
|
||||
self.default_model = default_model
|
||||
self.callback_handler = None
|
||||
if CALLBACK_ENABLED and AgentCallbackHandler:
|
||||
self.callback_handler = AgentCallbackHandler(
|
||||
agent_name=os.getenv("POD_NAME", "search-agent-a2a"),
|
||||
user_id=os.getenv("USER_ID", "")
|
||||
)
|
||||
|
||||
logger.info(
|
||||
"SearchAgentExecutor 初始化完成",
|
||||
@@ -123,8 +137,18 @@ class SearchAgentExecutor(AgentExecutor):
|
||||
agent = SearchAgentWrapper(api_key=api_key, model=model)
|
||||
|
||||
try:
|
||||
# 执行搜索
|
||||
response = await agent.search(query=user_text)
|
||||
callback_user_id = metadata.get("user_id") or os.getenv("USER_ID", "")
|
||||
|
||||
if self.callback_handler:
|
||||
with CallbackContextManager(
|
||||
handler=self.callback_handler,
|
||||
user_id=callback_user_id,
|
||||
request_id=getattr(context, "task_id", None)
|
||||
) as ctx:
|
||||
ctx.add_tool("search")
|
||||
response = await agent.search(query=user_text)
|
||||
else:
|
||||
response = await agent.search(query=user_text)
|
||||
|
||||
# 构建答案文本(包含来源信息)
|
||||
answer_parts = [response.answer.content]
|
||||
|
||||
@@ -37,6 +37,10 @@ RUN if [ -f /app/search_agent_requirements.txt ]; then \
|
||||
# 复制search_agent_A2A目录
|
||||
COPY agents/search_agent/search_agent_A2A/ /app/
|
||||
|
||||
# 复制回调工具
|
||||
COPY common/agent_callback_utils.py /app/common/
|
||||
RUN touch /app/common/__init__.py
|
||||
|
||||
# 复制search_agent核心代码
|
||||
COPY agents/search_agent/search_agent/ /app/search_agent/
|
||||
|
||||
|
||||
@@ -22,11 +22,20 @@ from loguru import logger
|
||||
from agent import SearchAgentWrapper
|
||||
from mcp_config import get_config, AgentConfig, MCPConfig
|
||||
|
||||
try:
|
||||
from common.agent_callback_utils import AgentCallbackHandler, CallbackContextManager
|
||||
CALLBACK_ENABLED = True
|
||||
except ImportError:
|
||||
CALLBACK_ENABLED = False
|
||||
AgentCallbackHandler = None
|
||||
CallbackContextManager = None
|
||||
|
||||
# 环境变量配置
|
||||
SERVICE_HOST = os.getenv("SERVICE_HOST", "0.0.0.0")
|
||||
SERVICE_PORT = int(os.getenv("SERVICE_PORT", "8080"))
|
||||
POD_NAME = os.getenv("POD_NAME", "search-agent-mcp")
|
||||
TEMPLATE_TYPE = os.getenv("TEMPLATE_TYPE", "search_agent_MCP")
|
||||
USER_ID = os.getenv("USER_ID", "")
|
||||
|
||||
# ============== MCP 协议数据模型 ==============
|
||||
|
||||
@@ -100,6 +109,10 @@ class MCPSearchAgentServer:
|
||||
|
||||
# 任务存储
|
||||
self.tasks: Dict[str, Dict[str, Any]] = {}
|
||||
|
||||
self.callback_handler = None
|
||||
if CALLBACK_ENABLED and AgentCallbackHandler:
|
||||
self.callback_handler = AgentCallbackHandler(agent_name=POD_NAME, user_id=USER_ID)
|
||||
|
||||
# 创建FastAPI应用
|
||||
self.app = self._create_app()
|
||||
@@ -360,8 +373,18 @@ class MCPSearchAgentServer:
|
||||
|
||||
# 调用Agent获取响应
|
||||
logger.info("处理搜索请求", task_id=task_id, query_preview=query[:50])
|
||||
|
||||
response = await agent.search(query=query)
|
||||
callback_user_id = params.get("user_id") or USER_ID
|
||||
|
||||
if self.callback_handler:
|
||||
with CallbackContextManager(
|
||||
handler=self.callback_handler,
|
||||
user_id=callback_user_id,
|
||||
request_id=task_id
|
||||
) as ctx:
|
||||
ctx.add_tool("search")
|
||||
response = await agent.search(query=query)
|
||||
else:
|
||||
response = await agent.search(query=query)
|
||||
|
||||
# 关闭 Agent(每个请求都创建新的 Agent)
|
||||
await agent.close()
|
||||
@@ -462,6 +485,7 @@ class MCPSearchAgentServer:
|
||||
try:
|
||||
# 获取Agent实例
|
||||
agent = self._get_agent(api_key, model)
|
||||
callback_user_id = params.get("user_id") or USER_ID
|
||||
|
||||
# 发送任务开始事件
|
||||
start_event = {
|
||||
@@ -474,8 +498,16 @@ class MCPSearchAgentServer:
|
||||
}
|
||||
yield f"data: {json.dumps(start_event)}\n\n"
|
||||
|
||||
# 执行搜索
|
||||
response = await agent.search(query=query)
|
||||
if self.callback_handler:
|
||||
with CallbackContextManager(
|
||||
handler=self.callback_handler,
|
||||
user_id=callback_user_id,
|
||||
request_id=task_id
|
||||
) as ctx:
|
||||
ctx.add_tool("search_stream")
|
||||
response = await agent.search(query=query)
|
||||
else:
|
||||
response = await agent.search(query=query)
|
||||
|
||||
# 构建答案文本
|
||||
answer_parts = [response.answer.content]
|
||||
|
||||
@@ -1 +0,0 @@
|
||||
../../search_agent/search_agent
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,253 @@
|
||||
# 🔍 智能AI搜索Agent
|
||||
|
||||
一个基于大语言模型的智能搜索代理,能够理解用户查询意图、自动规划搜索策略、从多个来源获取信息,并生成高质量、有来源引用的答案。
|
||||
|
||||
## ✨ 功能特点
|
||||
|
||||
| 能力 | 描述 |
|
||||
|------|------|
|
||||
| 🧠 查询理解 | 分析用户意图,提取关键实体,生成扩展查询 |
|
||||
| 📋 搜索规划 | 智能分解问题,制定搜索策略 |
|
||||
| 🔎 多源搜索 | 支持Web搜索和新闻搜索 |
|
||||
| 📄 内容提取 | 智能提取网页核心内容 |
|
||||
| 🎯 结果排序 | 基于相关性重排搜索结果 |
|
||||
| ✍️ 答案生成 | 综合信息生成结构化回答 |
|
||||
| 🔄 自我反思 | 评估答案质量,决定是否迭代 |
|
||||
|
||||
## 🛠️ 技术栈
|
||||
|
||||
| 组件 | 选型 | 说明 |
|
||||
|------|------|------|
|
||||
| LLM | xchat52 (GPT-5.2) | 主推理引擎 |
|
||||
| Web搜索 | Serper API | Google搜索代理 |
|
||||
| 内容提取 | Jina Reader | 网页转Markdown |
|
||||
| 重排序 | Jina Reranker | 结果相关性排序 |
|
||||
| 框架 | Python原生 + asyncio | 异步高效执行 |
|
||||
|
||||
## 📁 项目结构
|
||||
|
||||
```
|
||||
search_agent/
|
||||
├── main.py # 程序入口
|
||||
├── config.py # 配置管理
|
||||
├── requirements.txt # Python依赖
|
||||
├── .env # 环境变量配置
|
||||
│
|
||||
├── agent/
|
||||
│ ├── __init__.py
|
||||
│ ├── search_agent.py # 主Agent类
|
||||
│ └── prompts.py # Prompt模板
|
||||
│
|
||||
├── modules/
|
||||
│ ├── __init__.py
|
||||
│ ├── query_analyzer.py # 查询理解模块
|
||||
│ ├── search_planner.py # 搜索规划模块
|
||||
│ ├── search_executor.py # 搜索执行模块
|
||||
│ ├── content_extractor.py # 内容提取模块
|
||||
│ ├── result_processor.py # 结果处理模块
|
||||
│ ├── answer_generator.py # 答案生成模块
|
||||
│ └── reflector.py # 反思迭代模块
|
||||
│
|
||||
├── tools/
|
||||
│ ├── __init__.py
|
||||
│ ├── serper.py # Serper API封装
|
||||
│ ├── jina_reader.py # Jina Reader封装
|
||||
│ └── jina_reranker.py # Jina Reranker封装
|
||||
│
|
||||
├── models/
|
||||
│ ├── __init__.py
|
||||
│ └── schemas.py # 数据模型定义
|
||||
│
|
||||
└── utils/
|
||||
├── __init__.py
|
||||
├── llm_client.py # LLM客户端
|
||||
└── helpers.py # 工具函数
|
||||
```
|
||||
|
||||
## 🚀 快速开始
|
||||
|
||||
### 1. 安装依赖
|
||||
|
||||
```bash
|
||||
cd search_agent
|
||||
pip install -r requirements.txt
|
||||
```
|
||||
|
||||
### 2. 配置环境变量
|
||||
|
||||
创建 `.env` 文件:
|
||||
|
||||
```bash
|
||||
# LLM配置 (xchat52)
|
||||
LLM_BASE_URL=https://apis.openroutex.com/openai/deployments/xchat52
|
||||
LLM_API_KEY=你的API密钥
|
||||
LLM_MODEL=xchat52
|
||||
|
||||
# Serper配置 (Google搜索)
|
||||
SERPER_API_KEY=你的Serper_API_KEY
|
||||
|
||||
# Jina配置 (内容提取和重排序)
|
||||
JINA_API_KEY=你的Jina_API_KEY
|
||||
|
||||
# Agent配置
|
||||
MAX_ITERATIONS=3 # 最大迭代次数
|
||||
MAX_RESULTS_PER_QUERY=10 # 每次搜索返回结果数
|
||||
CONTENT_MAX_LENGTH=5000 # 提取内容最大长度
|
||||
|
||||
# 日志配置
|
||||
LOG_LEVEL=INFO
|
||||
TIMEOUT=30
|
||||
```
|
||||
|
||||
### 3. 运行程序
|
||||
|
||||
**交互模式**(推荐):
|
||||
```bash
|
||||
python main.py
|
||||
```
|
||||
|
||||
**单次查询**:
|
||||
```bash
|
||||
python main.py "你的问题"
|
||||
```
|
||||
|
||||
## 📖 使用示例
|
||||
|
||||
```
|
||||
🔍 智能AI搜索Agent
|
||||
======================================================================
|
||||
输入您的问题进行搜索,输入 'quit' 或 'exit' 退出
|
||||
======================================================================
|
||||
|
||||
🔎 请输入问题: 什么是大语言模型?
|
||||
|
||||
======================================================================
|
||||
📝 答案:
|
||||
======================================================================
|
||||
## 大语言模型(LLM)是什么?
|
||||
|
||||
**大语言模型(Large Language Model, LLM)**是一类用**海量文本数据**进行
|
||||
**预训练**的**超大规模深度学习模型**...
|
||||
|
||||
----------------------------------------------------------------------
|
||||
📚 来源:
|
||||
----------------------------------------------------------------------
|
||||
[1] 大语言模型 (LLM)
|
||||
🔗 https://www.ibm.com/cn-zh/think/topics/large-language-models
|
||||
[2] 什么是 LLM(大型语言模型)?
|
||||
🔗 https://aws.amazon.com/cn/what-is/large-language-model/
|
||||
...
|
||||
|
||||
----------------------------------------------------------------------
|
||||
📊 统计:
|
||||
----------------------------------------------------------------------
|
||||
• 置信度: high
|
||||
• 迭代次数: 1
|
||||
• 参考来源数: 10
|
||||
• 搜索查询数: 3
|
||||
======================================================================
|
||||
```
|
||||
|
||||
## 🔄 工作流程
|
||||
|
||||
```
|
||||
用户查询
|
||||
│
|
||||
▼
|
||||
┌───────────────────┐
|
||||
│ 查询理解 │ ──▶ 分析意图、提取实体、生成扩展查询
|
||||
└───────────────────┘
|
||||
│
|
||||
▼
|
||||
┌───────────────────┐
|
||||
│ 搜索规划 │ ──▶ 制定搜索策略(Web/新闻、并行/串行)
|
||||
└───────────────────┘
|
||||
│
|
||||
▼
|
||||
┌───────────────────┐
|
||||
│ 搜索执行 │ ──▶ 调用Serper API执行搜索
|
||||
└───────────────────┘
|
||||
│
|
||||
▼
|
||||
┌───────────────────┐
|
||||
│ 内容提取 │ ──▶ 使用Jina Reader提取网页内容
|
||||
└───────────────────┘
|
||||
│
|
||||
▼
|
||||
┌───────────────────┐
|
||||
│ 结果处理 │ ──▶ 去重 + Jina Reranker重排序
|
||||
└───────────────────┘
|
||||
│
|
||||
▼
|
||||
┌───────────────────┐
|
||||
│ 答案生成 │ ──▶ LLM综合生成结构化答案
|
||||
└───────────────────┘
|
||||
│
|
||||
▼
|
||||
┌───────────────────┐
|
||||
│ 反思评估 │ ──▶ 评估完整性,决定是否继续迭代
|
||||
└───────────────────┘
|
||||
│
|
||||
├──(完整)──▶ 返回最终答案
|
||||
│
|
||||
└──(不完整)──▶ 补充搜索(回到搜索规划)
|
||||
```
|
||||
|
||||
## ⚙️ 配置说明
|
||||
|
||||
| 配置项 | 默认值 | 说明 |
|
||||
|--------|--------|------|
|
||||
| `MAX_ITERATIONS` | 3 | 最大迭代次数,防止无限循环 |
|
||||
| `MAX_RESULTS_PER_QUERY` | 10 | 每次搜索返回的结果数量 |
|
||||
| `CONTENT_MAX_LENGTH` | 5000 | 提取内容的最大字符数 |
|
||||
| `LOG_LEVEL` | INFO | 日志级别 (DEBUG/INFO/WARNING/ERROR) |
|
||||
| `TIMEOUT` | 30 | API请求超时时间(秒)|
|
||||
|
||||
## 🔧 API说明
|
||||
|
||||
### Serper API
|
||||
- **Web搜索**: `POST https://google.serper.dev/search`
|
||||
- **新闻搜索**: `POST https://google.serper.dev/news`
|
||||
- [获取API Key](https://serper.dev/)
|
||||
|
||||
### Jina API
|
||||
- **内容提取**: `GET https://r.jina.ai/{URL}`
|
||||
- **重排序**: `POST https://api.jina.ai/v1/rerank`
|
||||
- [获取API Key](https://jina.ai/)
|
||||
|
||||
### LLM API (Azure OpenAI风格)
|
||||
- **Chat**: `POST {BASE_URL}/chat/completions?api-version=2024-10-21`
|
||||
|
||||
## 📝 编程接口
|
||||
|
||||
```python
|
||||
import asyncio
|
||||
from config import Config
|
||||
from agent.search_agent import SearchAgent
|
||||
|
||||
async def main():
|
||||
# 加载配置
|
||||
config = Config.from_env()
|
||||
|
||||
# 创建Agent
|
||||
agent = SearchAgent(config)
|
||||
|
||||
# 执行搜索
|
||||
response = await agent.search("你的问题")
|
||||
|
||||
# 获取答案
|
||||
print(response.answer.content)
|
||||
print(response.answer.sources)
|
||||
print(response.answer.confidence)
|
||||
|
||||
asyncio.run(main())
|
||||
```
|
||||
|
||||
## 📄 License
|
||||
|
||||
MIT License
|
||||
|
||||
## 🤝 贡献
|
||||
|
||||
欢迎提交Issue和Pull Request!
|
||||
|
||||
@@ -0,0 +1,6 @@
|
||||
"""
|
||||
Search Agent 核心模块
|
||||
"""
|
||||
|
||||
__version__ = "1.0.0"
|
||||
|
||||
@@ -0,0 +1,18 @@
|
||||
"""
|
||||
Agent模块
|
||||
"""
|
||||
|
||||
from .search_agent import SearchAgent
|
||||
from .prompts import (
|
||||
QUERY_ANALYSIS_PROMPT,
|
||||
ANSWER_GENERATION_PROMPT,
|
||||
REFLECTION_PROMPT,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"SearchAgent",
|
||||
"QUERY_ANALYSIS_PROMPT",
|
||||
"ANSWER_GENERATION_PROMPT",
|
||||
"REFLECTION_PROMPT",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,126 @@
|
||||
"""
|
||||
Prompt模板汇总
|
||||
集中管理所有LLM Prompt模板
|
||||
"""
|
||||
|
||||
# ==================== 查询分析 Prompt ====================
|
||||
QUERY_ANALYSIS_PROMPT = """你是一个查询分析专家。分析用户的搜索查询,提取以下信息。
|
||||
|
||||
请输出JSON格式:
|
||||
{
|
||||
"intent": "查询意图,必须是以下之一: fact_check(事实核查), comparison(对比分析), how_to(操作指南), news(新闻资讯), research(深度研究)",
|
||||
"entities": ["关键实体列表,提取查询中的核心概念、人名、产品名等"],
|
||||
"expanded_queries": ["扩展查询1", "扩展查询2", "扩展查询3"],
|
||||
"need_news": true或false,
|
||||
"time_filter": "时间过滤器,null表示不限时间,qdr:d(过去24小时), qdr:w(过去一周), qdr:m(过去一月), qdr:y(过去一年)"
|
||||
}
|
||||
|
||||
扩展查询要求:
|
||||
1. 生成2-4个扩展查询,包含不同角度或同义表达
|
||||
2. 至少包含一个英文查询(如果原查询是中文)
|
||||
3. 保持查询的核心意图
|
||||
|
||||
时间过滤器选择规则:
|
||||
- 查询涉及"最新"、"近期"、"今年"等时效性词语 → 设置相应的时间过滤器
|
||||
- 查询涉及具体年份(如"2024年") → qdr:y
|
||||
- 一般性查询 → null"""
|
||||
|
||||
|
||||
# ==================== 搜索规划 Prompt ====================
|
||||
SEARCH_PLANNING_PROMPT = """你是一个搜索规划专家。根据查询分析结果,制定搜索计划。
|
||||
|
||||
输入信息:
|
||||
- 原始查询
|
||||
- 查询意图
|
||||
- 关键实体
|
||||
- 是否需要新闻
|
||||
|
||||
输出搜索任务列表,每个任务包含:
|
||||
- query: 搜索词
|
||||
- source: web 或 news
|
||||
- time_filter: 时间过滤器(可选)
|
||||
|
||||
搜索策略规则:
|
||||
1. 简单事实查询 → 单次Web搜索
|
||||
2. 时效性查询 → Web搜索 + 新闻搜索
|
||||
3. 复杂分析查询 → 多个扩展查询
|
||||
4. 对比类查询 → 分别搜索各对比对象"""
|
||||
|
||||
|
||||
# ==================== 答案生成 Prompt ====================
|
||||
ANSWER_GENERATION_PROMPT = """你是一个专业的信息整合专家。根据以下搜索结果,回答用户的问题。
|
||||
|
||||
## 要求
|
||||
1. 综合多个来源的信息,给出全面准确的回答
|
||||
2. 使用清晰的结构组织答案(标题、列表、重点标注等)
|
||||
3. 在答案中标注信息来源,格式:[来源1]、[来源2]
|
||||
4. 如果信息有冲突,说明不同观点
|
||||
5. 如果信息不足以完整回答问题,明确指出缺失的部分
|
||||
6. 回答使用中文
|
||||
|
||||
## 输出JSON格式
|
||||
{
|
||||
"answer": "结构化的答案(Markdown格式,包含来源引用)",
|
||||
"sources": [
|
||||
{"index": 1, "title": "来源标题", "url": "来源URL"},
|
||||
{"index": 2, "title": "来源标题", "url": "来源URL"}
|
||||
],
|
||||
"confidence": "high/medium/low,基于信息质量和一致性判断"
|
||||
}"""
|
||||
|
||||
|
||||
# ==================== 反思评估 Prompt ====================
|
||||
REFLECTION_PROMPT = """你是一个质量评估专家。评估以下答案是否充分回答了用户的问题。
|
||||
|
||||
## 评估维度
|
||||
1. **完整性**: 答案是否覆盖了问题的所有方面?
|
||||
2. **准确性**: 答案内容是否有明确的来源支持?
|
||||
3. **深度**: 答案是否提供了足够的细节和解释?
|
||||
|
||||
## 输出JSON格式
|
||||
{
|
||||
"completeness": 0.0-1.0,
|
||||
"missing_aspects": ["如果有缺失,列出缺失的方面"],
|
||||
"needs_more_search": true或false,
|
||||
"suggested_queries": ["如果需要补充搜索,建议的搜索词"]
|
||||
}
|
||||
|
||||
## 判断标准
|
||||
- completeness >= 0.8 且没有重要信息缺失 → needs_more_search = false
|
||||
- completeness < 0.8 或有重要信息缺失 → needs_more_search = true
|
||||
- 建议的搜索词应该针对缺失的方面"""
|
||||
|
||||
|
||||
# ==================== 工具函数 ====================
|
||||
def format_query_analysis_prompt(query: str) -> str:
|
||||
"""格式化查询分析Prompt"""
|
||||
return f"{QUERY_ANALYSIS_PROMPT}\n\n用户查询: {query}"
|
||||
|
||||
|
||||
def format_answer_generation_prompt(query: str, documents: str) -> str:
|
||||
"""格式化答案生成Prompt"""
|
||||
return f"""{ANSWER_GENERATION_PROMPT}
|
||||
|
||||
## 用户问题
|
||||
{query}
|
||||
|
||||
## 搜索结果
|
||||
{documents}"""
|
||||
|
||||
|
||||
def format_reflection_prompt(query: str, answer: str, sources_count: int, confidence: str) -> str:
|
||||
"""格式化反思评估Prompt"""
|
||||
return f"""{REFLECTION_PROMPT}
|
||||
|
||||
## 用户问题
|
||||
{query}
|
||||
|
||||
## 生成的答案
|
||||
{answer}
|
||||
|
||||
## 答案的来源数量
|
||||
{sources_count} 个来源
|
||||
|
||||
## 答案的置信度
|
||||
{confidence}"""
|
||||
|
||||
+209
@@ -0,0 +1,209 @@
|
||||
"""
|
||||
搜索Agent主类
|
||||
协调各模块执行智能搜索
|
||||
"""
|
||||
|
||||
from typing import List, Optional
|
||||
from loguru import logger
|
||||
|
||||
from search_agent.config import Config
|
||||
from search_agent.models.schemas import (
|
||||
QueryAnalysis,
|
||||
SearchPlan,
|
||||
SearchResult,
|
||||
Document,
|
||||
RankedDocument,
|
||||
Answer,
|
||||
AgentResponse,
|
||||
)
|
||||
from search_agent.modules.query_analyzer import QueryAnalyzer
|
||||
from search_agent.modules.search_planner import SearchPlanner
|
||||
from search_agent.modules.search_executor import SearchExecutor
|
||||
from search_agent.modules.content_extractor import ContentExtractor
|
||||
from search_agent.modules.result_processor import ResultProcessor
|
||||
from search_agent.modules.answer_generator import AnswerGenerator
|
||||
from search_agent.modules.reflector import Reflector
|
||||
|
||||
|
||||
class SearchAgent:
|
||||
"""智能搜索Agent"""
|
||||
|
||||
def __init__(self, config: Config):
|
||||
"""
|
||||
初始化搜索Agent
|
||||
|
||||
Args:
|
||||
config: 配置对象
|
||||
"""
|
||||
self.config = config
|
||||
|
||||
# 初始化各模块
|
||||
self.query_analyzer = QueryAnalyzer(config)
|
||||
self.search_planner = SearchPlanner(config)
|
||||
self.search_executor = SearchExecutor(config)
|
||||
self.content_extractor = ContentExtractor(config)
|
||||
self.result_processor = ResultProcessor(config)
|
||||
self.answer_generator = AnswerGenerator(config)
|
||||
self.reflector = Reflector(config)
|
||||
|
||||
logger.info("SearchAgent 初始化完成")
|
||||
|
||||
async def search(self, query: str) -> AgentResponse:
|
||||
"""
|
||||
执行智能搜索
|
||||
|
||||
Args:
|
||||
query: 用户查询
|
||||
|
||||
Returns:
|
||||
AgentResponse对象
|
||||
"""
|
||||
logger.info(f"="*60)
|
||||
logger.info(f"开始搜索: {query}")
|
||||
logger.info(f"="*60)
|
||||
|
||||
iteration = 0
|
||||
all_documents: List[Document] = []
|
||||
all_queries: List[str] = []
|
||||
|
||||
# 1. 查询理解
|
||||
analysis = await self.query_analyzer.analyze(query)
|
||||
logger.info(f"查询分析完成: intent={analysis.intent.value}")
|
||||
|
||||
answer: Optional[Answer] = None
|
||||
|
||||
while iteration < self.config.max_iterations:
|
||||
iteration += 1
|
||||
logger.info(f"\n--- 迭代 {iteration}/{self.config.max_iterations} ---")
|
||||
|
||||
# 2. 搜索规划
|
||||
if iteration == 1:
|
||||
plan = await self.search_planner.plan(analysis)
|
||||
else:
|
||||
# 后续迭代使用建议的补充查询
|
||||
plan = self.search_planner.plan_supplementary(
|
||||
query,
|
||||
analysis.expanded_queries
|
||||
)
|
||||
|
||||
all_queries.extend([t.query for t in plan.tasks])
|
||||
logger.info(f"搜索计划: {len(plan.tasks)} 个任务")
|
||||
|
||||
# 3. 执行搜索
|
||||
search_results = await self.search_executor.execute(plan)
|
||||
logger.info(f"搜索结果: {len(search_results)} 条")
|
||||
|
||||
if not search_results:
|
||||
logger.warning("没有搜索结果")
|
||||
if answer is None:
|
||||
answer = self.answer_generator._empty_answer()
|
||||
break
|
||||
|
||||
# 4. 内容提取
|
||||
documents = await self.content_extractor.extract_batch(
|
||||
search_results,
|
||||
max_urls=10
|
||||
)
|
||||
all_documents.extend(documents)
|
||||
logger.info(f"提取文档: {len(documents)} 个")
|
||||
|
||||
if not documents:
|
||||
logger.warning("没有成功提取到文档内容")
|
||||
continue
|
||||
|
||||
# 5. 结果处理(去重+重排序)
|
||||
ranked_docs = await self.result_processor.process(
|
||||
query=query,
|
||||
documents=all_documents,
|
||||
top_k=5
|
||||
)
|
||||
logger.info(f"排序结果: {len(ranked_docs)} 个")
|
||||
|
||||
if not ranked_docs:
|
||||
logger.warning("没有有效的排序结果")
|
||||
continue
|
||||
|
||||
# 6. 生成答案
|
||||
answer = await self.answer_generator.generate(
|
||||
query=query,
|
||||
documents=ranked_docs
|
||||
)
|
||||
logger.info(f"答案生成完成: confidence={answer.confidence}")
|
||||
|
||||
# 7. 反思评估
|
||||
assessment = await self.reflector.assess(query, answer)
|
||||
|
||||
# 8. 判断是否继续迭代
|
||||
if not self.reflector.should_continue(assessment, iteration):
|
||||
break
|
||||
|
||||
# 更新分析,准备下一轮搜索
|
||||
if assessment.suggested_queries:
|
||||
analysis.expanded_queries = assessment.suggested_queries
|
||||
logger.info(f"补充搜索: {assessment.suggested_queries}")
|
||||
|
||||
# 确保有答案返回
|
||||
if answer is None:
|
||||
answer = self.answer_generator._empty_answer()
|
||||
|
||||
# 去重统计
|
||||
unique_urls = set(d.url for d in all_documents)
|
||||
|
||||
response = AgentResponse(
|
||||
answer=answer,
|
||||
iterations=iteration,
|
||||
total_sources_consulted=len(unique_urls),
|
||||
search_queries_used=list(set(all_queries))
|
||||
)
|
||||
|
||||
logger.info(f"\n{'='*60}")
|
||||
logger.info(f"搜索完成!")
|
||||
logger.info(f"迭代次数: {iteration}")
|
||||
logger.info(f"参考来源: {len(unique_urls)}")
|
||||
logger.info(f"搜索查询: {len(response.search_queries_used)}")
|
||||
logger.info(f"{'='*60}\n")
|
||||
|
||||
return response
|
||||
|
||||
async def quick_search(self, query: str) -> Answer:
|
||||
"""
|
||||
快速搜索(单次迭代)
|
||||
|
||||
Args:
|
||||
query: 用户查询
|
||||
|
||||
Returns:
|
||||
Answer对象
|
||||
"""
|
||||
# 简化分析
|
||||
analysis = await self.query_analyzer.analyze(query)
|
||||
|
||||
# 只执行一次搜索
|
||||
plan = await self.search_planner.plan(analysis)
|
||||
plan.tasks = plan.tasks[:2] # 限制搜索任务数量
|
||||
|
||||
# 执行搜索
|
||||
search_results = await self.search_executor.execute(plan)
|
||||
|
||||
if not search_results:
|
||||
return self.answer_generator._empty_answer()
|
||||
|
||||
# 提取内容
|
||||
documents = await self.content_extractor.extract_batch(
|
||||
search_results,
|
||||
max_urls=5
|
||||
)
|
||||
|
||||
if not documents:
|
||||
return self.answer_generator._empty_answer()
|
||||
|
||||
# 处理结果
|
||||
ranked_docs = await self.result_processor.process(
|
||||
query=query,
|
||||
documents=documents,
|
||||
top_k=3
|
||||
)
|
||||
|
||||
# 生成答案
|
||||
return await self.answer_generator.generate(query, ranked_docs)
|
||||
|
||||
@@ -0,0 +1,81 @@
|
||||
"""
|
||||
配置管理模块
|
||||
负责加载和管理所有配置项
|
||||
"""
|
||||
|
||||
import os
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional
|
||||
from dotenv import load_dotenv
|
||||
|
||||
|
||||
@dataclass
|
||||
class Config:
|
||||
"""Agent配置类"""
|
||||
|
||||
# LLM配置
|
||||
llm_base_url: str
|
||||
llm_api_key: str
|
||||
llm_model: str
|
||||
|
||||
# Serper配置
|
||||
serper_api_key: str
|
||||
|
||||
# Jina配置
|
||||
jina_api_key: str
|
||||
|
||||
# Agent配置
|
||||
max_iterations: int
|
||||
max_results_per_query: int
|
||||
content_max_length: int
|
||||
|
||||
# 可选配置
|
||||
log_level: str = "INFO"
|
||||
timeout: int = 30
|
||||
|
||||
@classmethod
|
||||
def from_env(cls, env_path: Optional[str] = None) -> "Config":
|
||||
"""从环境变量加载配置"""
|
||||
if env_path:
|
||||
load_dotenv(env_path)
|
||||
else:
|
||||
load_dotenv()
|
||||
|
||||
return cls(
|
||||
# LLM配置
|
||||
llm_base_url=os.getenv("LLM_BASE_URL", ""),
|
||||
llm_api_key=os.getenv("LLM_API_KEY", ""),
|
||||
llm_model=os.getenv("MODEL_NAME", "xchat52"),
|
||||
|
||||
# Serper配置
|
||||
serper_api_key=os.getenv("SERPER_API_KEY", ""),
|
||||
|
||||
# Jina配置
|
||||
jina_api_key=os.getenv("JINA_API_KEY", ""),
|
||||
|
||||
# Agent配置
|
||||
max_iterations=int(os.getenv("MAX_ITERATIONS", "3")),
|
||||
max_results_per_query=int(os.getenv("MAX_RESULTS_PER_QUERY", "10")),
|
||||
content_max_length=int(os.getenv("CONTENT_MAX_LENGTH", "5000")),
|
||||
|
||||
# 可选配置
|
||||
log_level=os.getenv("LOG_LEVEL", "INFO"),
|
||||
timeout=int(os.getenv("TIMEOUT", "30"))
|
||||
)
|
||||
|
||||
def validate(self) -> bool:
|
||||
"""验证配置是否完整"""
|
||||
required_fields = [
|
||||
("llm_base_url", self.llm_base_url),
|
||||
("llm_api_key", self.llm_api_key),
|
||||
("serper_api_key", self.serper_api_key),
|
||||
("jina_api_key", self.jina_api_key),
|
||||
]
|
||||
|
||||
missing = [name for name, value in required_fields if not value]
|
||||
|
||||
if missing:
|
||||
raise ValueError(f"缺少必要的配置项: {', '.join(missing)}")
|
||||
|
||||
return True
|
||||
|
||||
@@ -0,0 +1,125 @@
|
||||
"""
|
||||
智能AI搜索Agent - 程序入口
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import sys
|
||||
from loguru import logger
|
||||
|
||||
from search_agent.config import Config
|
||||
from search_agent.agent.search_agent import SearchAgent
|
||||
|
||||
|
||||
def setup_logging(level: str = "INFO"):
|
||||
"""配置日志"""
|
||||
logger.remove()
|
||||
logger.add(
|
||||
sys.stderr,
|
||||
level=level,
|
||||
format="<green>{time:HH:mm:ss}</green> | <level>{level: <8}</level> | <cyan>{message}</cyan>"
|
||||
)
|
||||
|
||||
|
||||
def print_response(response):
|
||||
"""格式化输出响应"""
|
||||
print("\n" + "=" * 70)
|
||||
print("📝 答案:")
|
||||
print("=" * 70)
|
||||
print(response.answer.content)
|
||||
|
||||
print("\n" + "-" * 70)
|
||||
print("📚 来源:")
|
||||
print("-" * 70)
|
||||
for source in response.answer.sources:
|
||||
print(f" [{source.index}] {source.title}")
|
||||
print(f" 🔗 {source.url}")
|
||||
|
||||
print("\n" + "-" * 70)
|
||||
print("📊 统计:")
|
||||
print("-" * 70)
|
||||
print(f" • 置信度: {response.answer.confidence}")
|
||||
print(f" • 迭代次数: {response.iterations}")
|
||||
print(f" • 参考来源数: {response.total_sources_consulted}")
|
||||
print(f" • 搜索查询数: {len(response.search_queries_used)}")
|
||||
print("=" * 70 + "\n")
|
||||
|
||||
|
||||
async def main():
|
||||
"""主函数"""
|
||||
# 加载配置
|
||||
config = Config.from_env()
|
||||
|
||||
# 配置日志
|
||||
setup_logging(config.log_level)
|
||||
|
||||
# 验证配置
|
||||
try:
|
||||
config.validate()
|
||||
except ValueError as e:
|
||||
logger.error(f"配置错误: {e}")
|
||||
logger.info("请检查 .env 文件中的配置项")
|
||||
return
|
||||
|
||||
# 创建Agent
|
||||
agent = SearchAgent(config)
|
||||
|
||||
# 交互式搜索
|
||||
print("\n" + "=" * 70)
|
||||
print("🔍 智能AI搜索Agent")
|
||||
print("=" * 70)
|
||||
print("输入您的问题进行搜索,输入 'quit' 或 'exit' 退出")
|
||||
print("=" * 70 + "\n")
|
||||
|
||||
while True:
|
||||
try:
|
||||
query = input("🔎 请输入问题: ").strip()
|
||||
|
||||
if not query:
|
||||
continue
|
||||
|
||||
if query.lower() in ['quit', 'exit', 'q']:
|
||||
print("\n👋 再见!")
|
||||
break
|
||||
|
||||
# 执行搜索
|
||||
response = await agent.search(query)
|
||||
|
||||
# 输出结果
|
||||
print_response(response)
|
||||
|
||||
except KeyboardInterrupt:
|
||||
print("\n\n👋 再见!")
|
||||
break
|
||||
except Exception as e:
|
||||
logger.error(f"搜索出错: {e}")
|
||||
continue
|
||||
|
||||
|
||||
async def search_once(query: str):
|
||||
"""
|
||||
单次搜索(用于脚本调用)
|
||||
|
||||
Args:
|
||||
query: 搜索查询
|
||||
"""
|
||||
config = Config.from_env()
|
||||
setup_logging(config.log_level)
|
||||
config.validate()
|
||||
|
||||
agent = SearchAgent(config)
|
||||
response = await agent.search(query)
|
||||
print_response(response)
|
||||
|
||||
return response
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
# 检查命令行参数
|
||||
if len(sys.argv) > 1:
|
||||
# 命令行传入查询
|
||||
query = " ".join(sys.argv[1:])
|
||||
asyncio.run(search_once(query))
|
||||
else:
|
||||
# 交互模式
|
||||
asyncio.run(main())
|
||||
|
||||
@@ -0,0 +1,34 @@
|
||||
"""
|
||||
数据模型模块
|
||||
"""
|
||||
|
||||
from .schemas import (
|
||||
SearchSource,
|
||||
Intent,
|
||||
QueryAnalysis,
|
||||
SearchTask,
|
||||
SearchPlan,
|
||||
SearchResult,
|
||||
Document,
|
||||
RankedDocument,
|
||||
Source,
|
||||
Answer,
|
||||
QualityAssessment,
|
||||
AgentResponse,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"SearchSource",
|
||||
"Intent",
|
||||
"QueryAnalysis",
|
||||
"SearchTask",
|
||||
"SearchPlan",
|
||||
"SearchResult",
|
||||
"Document",
|
||||
"RankedDocument",
|
||||
"Source",
|
||||
"Answer",
|
||||
"QualityAssessment",
|
||||
"AgentResponse",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,202 @@
|
||||
"""
|
||||
数据模型定义
|
||||
定义Agent使用的所有数据结构
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import List, Optional
|
||||
from enum import Enum
|
||||
|
||||
|
||||
class SearchSource(Enum):
|
||||
"""搜索来源枚举"""
|
||||
WEB = "web"
|
||||
NEWS = "news"
|
||||
|
||||
|
||||
class Intent(Enum):
|
||||
"""查询意图枚举"""
|
||||
FACT_CHECK = "fact_check" # 事实核查
|
||||
COMPARISON = "comparison" # 对比分析
|
||||
HOW_TO = "how_to" # 操作指南
|
||||
NEWS = "news" # 新闻资讯
|
||||
RESEARCH = "research" # 深度研究
|
||||
|
||||
|
||||
@dataclass
|
||||
class QueryAnalysis:
|
||||
"""查询分析结果"""
|
||||
original_query: str # 原始查询
|
||||
intent: Intent # 查询意图
|
||||
entities: List[str] # 关键实体
|
||||
expanded_queries: List[str] # 扩展查询列表
|
||||
need_news: bool # 是否需要新闻搜索
|
||||
time_filter: Optional[str] = None # 时间过滤器
|
||||
|
||||
def to_dict(self) -> dict:
|
||||
"""转换为字典"""
|
||||
return {
|
||||
"original_query": self.original_query,
|
||||
"intent": self.intent.value,
|
||||
"entities": self.entities,
|
||||
"expanded_queries": self.expanded_queries,
|
||||
"need_news": self.need_news,
|
||||
"time_filter": self.time_filter
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class SearchTask:
|
||||
"""搜索任务"""
|
||||
query: str # 搜索查询
|
||||
source: SearchSource # 搜索来源
|
||||
time_filter: Optional[str] = None # 时间过滤器
|
||||
num_results: int = 10 # 结果数量
|
||||
|
||||
def to_dict(self) -> dict:
|
||||
"""转换为字典"""
|
||||
return {
|
||||
"query": self.query,
|
||||
"source": self.source.value,
|
||||
"time_filter": self.time_filter,
|
||||
"num_results": self.num_results
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class SearchPlan:
|
||||
"""搜索计划"""
|
||||
tasks: List[SearchTask] # 搜索任务列表
|
||||
strategy: str = "parallel" # 执行策略: parallel/sequential
|
||||
|
||||
def to_dict(self) -> dict:
|
||||
"""转换为字典"""
|
||||
return {
|
||||
"tasks": [t.to_dict() for t in self.tasks],
|
||||
"strategy": self.strategy
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class SearchResult:
|
||||
"""搜索结果"""
|
||||
title: str # 标题
|
||||
url: str # URL
|
||||
snippet: str # 摘要
|
||||
source: SearchSource # 来源类型
|
||||
position: int # 排名位置
|
||||
date: Optional[str] = None # 日期(新闻)
|
||||
|
||||
def to_dict(self) -> dict:
|
||||
"""转换为字典"""
|
||||
return {
|
||||
"title": self.title,
|
||||
"url": self.url,
|
||||
"snippet": self.snippet,
|
||||
"source": self.source.value,
|
||||
"position": self.position,
|
||||
"date": self.date
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class Document:
|
||||
"""提取的文档内容"""
|
||||
url: str # URL
|
||||
title: str # 标题
|
||||
content: str # 内容
|
||||
source: SearchSource # 来源类型
|
||||
|
||||
def to_dict(self) -> dict:
|
||||
"""转换为字典"""
|
||||
return {
|
||||
"url": self.url,
|
||||
"title": self.title,
|
||||
"content": self.content,
|
||||
"source": self.source.value
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class RankedDocument:
|
||||
"""排序后的文档"""
|
||||
document: Document # 文档
|
||||
relevance_score: float # 相关性分数
|
||||
rank: int # 排名
|
||||
|
||||
def to_dict(self) -> dict:
|
||||
"""转换为字典"""
|
||||
return {
|
||||
"document": self.document.to_dict(),
|
||||
"relevance_score": self.relevance_score,
|
||||
"rank": self.rank
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class Source:
|
||||
"""来源引用"""
|
||||
index: int # 索引
|
||||
title: str # 标题
|
||||
url: str # URL
|
||||
|
||||
def to_dict(self) -> dict:
|
||||
"""转换为字典"""
|
||||
return {
|
||||
"index": self.index,
|
||||
"title": self.title,
|
||||
"url": self.url
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class Answer:
|
||||
"""生成的答案"""
|
||||
content: str # Markdown格式的答案内容
|
||||
sources: List[Source] # 来源列表
|
||||
confidence: str # 置信度: high/medium/low
|
||||
|
||||
def to_dict(self) -> dict:
|
||||
"""转换为字典"""
|
||||
return {
|
||||
"content": self.content,
|
||||
"sources": [s.to_dict() for s in self.sources],
|
||||
"confidence": self.confidence
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class QualityAssessment:
|
||||
"""质量评估"""
|
||||
completeness: float # 完整性 0-1
|
||||
missing_aspects: List[str] # 缺失的方面
|
||||
needs_more_search: bool # 是否需要更多搜索
|
||||
suggested_queries: List[str] # 建议的补充搜索
|
||||
|
||||
def to_dict(self) -> dict:
|
||||
"""转换为字典"""
|
||||
return {
|
||||
"completeness": self.completeness,
|
||||
"missing_aspects": self.missing_aspects,
|
||||
"needs_more_search": self.needs_more_search,
|
||||
"suggested_queries": self.suggested_queries
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class AgentResponse:
|
||||
"""Agent最终响应"""
|
||||
answer: Answer # 答案
|
||||
iterations: int # 迭代次数
|
||||
total_sources_consulted: int # 参考来源总数
|
||||
search_queries_used: List[str] # 使用的搜索查询
|
||||
|
||||
def to_dict(self) -> dict:
|
||||
"""转换为字典"""
|
||||
return {
|
||||
"answer": self.answer.to_dict(),
|
||||
"iterations": self.iterations,
|
||||
"total_sources_consulted": self.total_sources_consulted,
|
||||
"search_queries_used": self.search_queries_used
|
||||
}
|
||||
|
||||
@@ -0,0 +1,22 @@
|
||||
"""
|
||||
核心模块
|
||||
"""
|
||||
|
||||
from .query_analyzer import QueryAnalyzer
|
||||
from .search_planner import SearchPlanner
|
||||
from .search_executor import SearchExecutor
|
||||
from .content_extractor import ContentExtractor
|
||||
from .result_processor import ResultProcessor
|
||||
from .answer_generator import AnswerGenerator
|
||||
from .reflector import Reflector
|
||||
|
||||
__all__ = [
|
||||
"QueryAnalyzer",
|
||||
"SearchPlanner",
|
||||
"SearchExecutor",
|
||||
"ContentExtractor",
|
||||
"ResultProcessor",
|
||||
"AnswerGenerator",
|
||||
"Reflector",
|
||||
]
|
||||
|
||||
+151
@@ -0,0 +1,151 @@
|
||||
"""
|
||||
答案生成模块
|
||||
综合多个来源的信息生成结构化答案
|
||||
"""
|
||||
|
||||
from typing import List
|
||||
from loguru import logger
|
||||
|
||||
from search_agent.config import Config
|
||||
from search_agent.models.schemas import RankedDocument, Answer, Source
|
||||
from search_agent.utils.llm_client import LLMClient
|
||||
from search_agent.utils.helpers import format_documents_for_prompt
|
||||
|
||||
|
||||
# 答案生成Prompt
|
||||
ANSWER_GENERATION_PROMPT = """你是一个专业的信息整合专家。根据以下搜索结果,回答用户的问题。
|
||||
|
||||
## 要求
|
||||
1. 综合多个来源的信息,给出全面准确的回答
|
||||
2. 使用清晰的结构组织答案(标题、列表、重点标注等)
|
||||
3. 在答案中标注信息来源,格式:[来源1]、[来源2]
|
||||
4. 如果信息有冲突,说明不同观点
|
||||
5. 如果信息不足以完整回答问题,明确指出缺失的部分
|
||||
6. 回答使用中文
|
||||
|
||||
## 输出JSON格式
|
||||
{
|
||||
"answer": "结构化的答案(Markdown格式,包含来源引用)",
|
||||
"sources": [
|
||||
{"index": 1, "title": "来源标题", "url": "来源URL"},
|
||||
{"index": 2, "title": "来源标题", "url": "来源URL"}
|
||||
],
|
||||
"confidence": "high/medium/low,基于信息质量和一致性判断"
|
||||
}"""
|
||||
|
||||
|
||||
class AnswerGenerator:
|
||||
"""答案生成模块"""
|
||||
|
||||
def __init__(self, config: Config):
|
||||
"""
|
||||
初始化答案生成器
|
||||
|
||||
Args:
|
||||
config: 配置对象
|
||||
"""
|
||||
self.config = config
|
||||
self.llm = LLMClient(
|
||||
base_url=config.llm_base_url,
|
||||
api_key=config.llm_api_key,
|
||||
model=config.llm_model,
|
||||
timeout=120 # 答案生成可能需要更长时间
|
||||
)
|
||||
|
||||
async def generate(
|
||||
self,
|
||||
query: str,
|
||||
documents: List[RankedDocument]
|
||||
) -> Answer:
|
||||
"""
|
||||
根据文档生成答案
|
||||
|
||||
Args:
|
||||
query: 用户查询
|
||||
documents: 排序后的文档列表
|
||||
|
||||
Returns:
|
||||
Answer对象
|
||||
"""
|
||||
if not documents:
|
||||
return self._empty_answer()
|
||||
|
||||
logger.info(f"开始生成答案,使用 {len(documents)} 个文档")
|
||||
|
||||
# 格式化文档
|
||||
formatted_docs = format_documents_for_prompt(
|
||||
documents,
|
||||
max_length=self.config.content_max_length // len(documents)
|
||||
)
|
||||
|
||||
user_message = f"""## 用户问题
|
||||
{query}
|
||||
|
||||
## 搜索结果
|
||||
{formatted_docs}"""
|
||||
|
||||
try:
|
||||
result = await self.llm.chat_json(
|
||||
system_prompt=ANSWER_GENERATION_PROMPT,
|
||||
user_message=user_message,
|
||||
temperature=0.5
|
||||
)
|
||||
|
||||
# 解析来源
|
||||
sources = [
|
||||
Source(
|
||||
index=s.get("index", i + 1),
|
||||
title=s.get("title", ""),
|
||||
url=s.get("url", "")
|
||||
)
|
||||
for i, s in enumerate(result.get("sources", []))
|
||||
]
|
||||
|
||||
answer = Answer(
|
||||
content=result.get("answer", ""),
|
||||
sources=sources,
|
||||
confidence=result.get("confidence", "medium")
|
||||
)
|
||||
|
||||
logger.info(f"答案生成完成,置信度: {answer.confidence}")
|
||||
return answer
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"答案生成失败: {e}")
|
||||
return self._fallback_answer(query, documents)
|
||||
|
||||
def _empty_answer(self) -> Answer:
|
||||
"""生成空答案(无文档时)"""
|
||||
return Answer(
|
||||
content="抱歉,未能找到相关信息来回答您的问题。",
|
||||
sources=[],
|
||||
confidence="low"
|
||||
)
|
||||
|
||||
def _fallback_answer(
|
||||
self,
|
||||
query: str,
|
||||
documents: List[RankedDocument]
|
||||
) -> Answer:
|
||||
"""后备答案生成(LLM失败时)"""
|
||||
# 简单汇总文档内容
|
||||
content_parts = [f"关于「{query}」,以下是搜索到的相关信息:\n"]
|
||||
|
||||
sources = []
|
||||
for i, doc in enumerate(documents[:5], 1):
|
||||
actual_doc = doc.document
|
||||
content_parts.append(f"### 来源 [{i}]: {actual_doc.title}\n")
|
||||
content_parts.append(f"{actual_doc.content[:500]}...\n\n")
|
||||
|
||||
sources.append(Source(
|
||||
index=i,
|
||||
title=actual_doc.title,
|
||||
url=actual_doc.url
|
||||
))
|
||||
|
||||
return Answer(
|
||||
content="".join(content_parts),
|
||||
sources=sources,
|
||||
confidence="low"
|
||||
)
|
||||
|
||||
+102
@@ -0,0 +1,102 @@
|
||||
"""
|
||||
内容提取模块
|
||||
使用Jina Reader提取网页内容
|
||||
"""
|
||||
|
||||
from typing import List
|
||||
from loguru import logger
|
||||
|
||||
from search_agent.config import Config
|
||||
from search_agent.models.schemas import Document, SearchResult, SearchSource
|
||||
from search_agent.tools.jina_reader import JinaReaderClient
|
||||
|
||||
|
||||
class ContentExtractor:
|
||||
"""内容提取模块"""
|
||||
|
||||
def __init__(self, config: Config):
|
||||
"""
|
||||
初始化内容提取器
|
||||
|
||||
Args:
|
||||
config: 配置对象
|
||||
"""
|
||||
self.config = config
|
||||
self.jina_reader = JinaReaderClient(
|
||||
api_key=config.jina_api_key,
|
||||
timeout=config.timeout,
|
||||
max_content_length=config.content_max_length
|
||||
)
|
||||
|
||||
async def extract(self, search_result: SearchResult) -> Document | None:
|
||||
"""
|
||||
从搜索结果提取内容
|
||||
|
||||
Args:
|
||||
search_result: 搜索结果
|
||||
|
||||
Returns:
|
||||
Document对象,如果提取失败则返回None
|
||||
"""
|
||||
return await self.jina_reader.extract_content(
|
||||
url=search_result.url,
|
||||
source=search_result.source
|
||||
)
|
||||
|
||||
async def extract_batch(
|
||||
self,
|
||||
search_results: List[SearchResult],
|
||||
max_urls: int = 10
|
||||
) -> List[Document]:
|
||||
"""
|
||||
批量提取内容
|
||||
|
||||
Args:
|
||||
search_results: 搜索结果列表
|
||||
max_urls: 最大提取URL数量
|
||||
|
||||
Returns:
|
||||
Document列表
|
||||
"""
|
||||
# 去重并限制数量
|
||||
seen_urls = set()
|
||||
unique_results = []
|
||||
|
||||
for result in search_results:
|
||||
if result.url not in seen_urls and len(unique_results) < max_urls:
|
||||
seen_urls.add(result.url)
|
||||
unique_results.append(result)
|
||||
|
||||
logger.info(f"开始提取 {len(unique_results)} 个URL的内容")
|
||||
|
||||
# 提取内容
|
||||
urls = [r.url for r in unique_results]
|
||||
# 保存source信息以便后续使用
|
||||
url_to_source = {r.url: r.source for r in unique_results}
|
||||
|
||||
documents = await self.jina_reader.extract_batch(urls)
|
||||
|
||||
# 更新document的source信息
|
||||
for doc in documents:
|
||||
if doc.url in url_to_source:
|
||||
doc.source = url_to_source[doc.url]
|
||||
|
||||
return documents
|
||||
|
||||
async def extract_urls(
|
||||
self,
|
||||
urls: List[str],
|
||||
source: SearchSource = SearchSource.WEB
|
||||
) -> List[Document]:
|
||||
"""
|
||||
直接从URL列表提取内容
|
||||
|
||||
Args:
|
||||
urls: URL列表
|
||||
source: 来源类型
|
||||
|
||||
Returns:
|
||||
Document列表
|
||||
"""
|
||||
return await self.jina_reader.extract_batch(urls, source)
|
||||
|
||||
+117
@@ -0,0 +1,117 @@
|
||||
"""
|
||||
查询理解模块
|
||||
负责分析用户查询意图、提取关键实体、生成扩展查询
|
||||
"""
|
||||
|
||||
from typing import Optional
|
||||
from loguru import logger
|
||||
|
||||
from search_agent.config import Config
|
||||
from search_agent.models.schemas import QueryAnalysis, Intent
|
||||
from search_agent.utils.llm_client import LLMClient
|
||||
|
||||
|
||||
# 查询分析Prompt
|
||||
QUERY_ANALYSIS_PROMPT = """你是一个查询分析专家。分析用户的搜索查询,提取以下信息。
|
||||
|
||||
请输出JSON格式:
|
||||
{
|
||||
"intent": "查询意图,必须是以下之一: fact_check(事实核查), comparison(对比分析), how_to(操作指南), news(新闻资讯), research(深度研究)",
|
||||
"entities": ["关键实体列表,提取查询中的核心概念、人名、产品名等"],
|
||||
"expanded_queries": ["扩展查询1", "扩展查询2", "扩展查询3"],
|
||||
"need_news": true或false,
|
||||
"time_filter": "时间过滤器,null表示不限时间,qdr:d(过去24小时), qdr:w(过去一周), qdr:m(过去一月), qdr:y(过去一年)"
|
||||
}
|
||||
|
||||
扩展查询要求:
|
||||
1. 生成2-4个扩展查询,包含不同角度或同义表达
|
||||
2. 至少包含一个英文查询(如果原查询是中文)
|
||||
3. 保持查询的核心意图
|
||||
|
||||
时间过滤器选择规则:
|
||||
- 查询涉及"最新"、"近期"、"今年"等时效性词语 → 设置相应的时间过滤器
|
||||
- 查询涉及具体年份(如"2024年") → qdr:y
|
||||
- 一般性查询 → null"""
|
||||
|
||||
|
||||
class QueryAnalyzer:
|
||||
"""查询理解模块"""
|
||||
|
||||
def __init__(self, config: Config):
|
||||
"""
|
||||
初始化查询分析器
|
||||
|
||||
Args:
|
||||
config: 配置对象
|
||||
"""
|
||||
self.config = config
|
||||
self.llm = LLMClient(
|
||||
base_url=config.llm_base_url,
|
||||
api_key=config.llm_api_key,
|
||||
model=config.llm_model
|
||||
)
|
||||
|
||||
async def analyze(self, query: str) -> QueryAnalysis:
|
||||
"""
|
||||
分析用户查询
|
||||
|
||||
Args:
|
||||
query: 用户查询字符串
|
||||
|
||||
Returns:
|
||||
QueryAnalysis对象
|
||||
"""
|
||||
logger.info(f"开始分析查询: {query}")
|
||||
|
||||
try:
|
||||
result = await self.llm.chat_json(
|
||||
system_prompt=QUERY_ANALYSIS_PROMPT,
|
||||
user_message=f"用户查询: {query}",
|
||||
temperature=0.3
|
||||
)
|
||||
|
||||
# 解析意图
|
||||
intent_str = result.get("intent", "research")
|
||||
intent = self._parse_intent(intent_str)
|
||||
|
||||
# 构建分析结果
|
||||
analysis = QueryAnalysis(
|
||||
original_query=query,
|
||||
intent=intent,
|
||||
entities=result.get("entities", []),
|
||||
expanded_queries=result.get("expanded_queries", [query]),
|
||||
need_news=result.get("need_news", False),
|
||||
time_filter=result.get("time_filter")
|
||||
)
|
||||
|
||||
logger.info(f"查询分析完成: intent={intent.value}, entities={analysis.entities}")
|
||||
return analysis
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"查询分析失败: {e}")
|
||||
# 返回默认分析结果
|
||||
return self._default_analysis(query)
|
||||
|
||||
def _parse_intent(self, intent_str: str) -> Intent:
|
||||
"""解析意图字符串为枚举"""
|
||||
intent_mapping = {
|
||||
"fact_check": Intent.FACT_CHECK,
|
||||
"comparison": Intent.COMPARISON,
|
||||
"how_to": Intent.HOW_TO,
|
||||
"news": Intent.NEWS,
|
||||
"research": Intent.RESEARCH
|
||||
}
|
||||
|
||||
return intent_mapping.get(intent_str.lower(), Intent.RESEARCH)
|
||||
|
||||
def _default_analysis(self, query: str) -> QueryAnalysis:
|
||||
"""生成默认的查询分析结果"""
|
||||
return QueryAnalysis(
|
||||
original_query=query,
|
||||
intent=Intent.RESEARCH,
|
||||
entities=[],
|
||||
expanded_queries=[query],
|
||||
need_news=False,
|
||||
time_filter=None
|
||||
)
|
||||
|
||||
+166
@@ -0,0 +1,166 @@
|
||||
"""
|
||||
反思迭代模块
|
||||
评估答案质量,决定是否需要补充搜索
|
||||
"""
|
||||
|
||||
from typing import List
|
||||
from loguru import logger
|
||||
|
||||
from search_agent.config import Config
|
||||
from search_agent.models.schemas import Answer, QualityAssessment
|
||||
from search_agent.utils.llm_client import LLMClient
|
||||
|
||||
|
||||
# 反思评估Prompt
|
||||
REFLECTION_PROMPT = """你是一个质量评估专家。评估以下答案是否充分回答了用户的问题。
|
||||
|
||||
## 评估维度
|
||||
1. **完整性**: 答案是否覆盖了问题的所有方面?
|
||||
2. **准确性**: 答案内容是否有明确的来源支持?
|
||||
3. **深度**: 答案是否提供了足够的细节和解释?
|
||||
|
||||
## 输出JSON格式
|
||||
{
|
||||
"completeness": 0.0-1.0,
|
||||
"missing_aspects": ["如果有缺失,列出缺失的方面"],
|
||||
"needs_more_search": true或false,
|
||||
"suggested_queries": ["如果需要补充搜索,建议的搜索词"]
|
||||
}
|
||||
|
||||
## 判断标准
|
||||
- completeness >= 0.8 且没有重要信息缺失 → needs_more_search = false
|
||||
- completeness < 0.8 或有重要信息缺失 → needs_more_search = true
|
||||
- 建议的搜索词应该针对缺失的方面"""
|
||||
|
||||
|
||||
class Reflector:
|
||||
"""反思迭代模块"""
|
||||
|
||||
# 质量阈值
|
||||
COMPLETENESS_THRESHOLD = 0.8
|
||||
|
||||
def __init__(self, config: Config):
|
||||
"""
|
||||
初始化反思器
|
||||
|
||||
Args:
|
||||
config: 配置对象
|
||||
"""
|
||||
self.config = config
|
||||
self.llm = LLMClient(
|
||||
base_url=config.llm_base_url,
|
||||
api_key=config.llm_api_key,
|
||||
model=config.llm_model
|
||||
)
|
||||
|
||||
async def assess(
|
||||
self,
|
||||
query: str,
|
||||
answer: Answer
|
||||
) -> QualityAssessment:
|
||||
"""
|
||||
评估答案质量
|
||||
|
||||
Args:
|
||||
query: 原始查询
|
||||
answer: 生成的答案
|
||||
|
||||
Returns:
|
||||
QualityAssessment对象
|
||||
"""
|
||||
logger.info("开始评估答案质量")
|
||||
|
||||
# 如果答案置信度已经很低,直接建议补充搜索
|
||||
if answer.confidence == "low" and not answer.content:
|
||||
return QualityAssessment(
|
||||
completeness=0.0,
|
||||
missing_aspects=["缺少相关信息"],
|
||||
needs_more_search=True,
|
||||
suggested_queries=[query]
|
||||
)
|
||||
|
||||
user_message = f"""## 用户问题
|
||||
{query}
|
||||
|
||||
## 生成的答案
|
||||
{answer.content}
|
||||
|
||||
## 答案的来源数量
|
||||
{len(answer.sources)} 个来源
|
||||
|
||||
## 答案的置信度
|
||||
{answer.confidence}"""
|
||||
|
||||
try:
|
||||
result = await self.llm.chat_json(
|
||||
system_prompt=REFLECTION_PROMPT,
|
||||
user_message=user_message,
|
||||
temperature=0.3
|
||||
)
|
||||
|
||||
assessment = QualityAssessment(
|
||||
completeness=float(result.get("completeness", 0.5)),
|
||||
missing_aspects=result.get("missing_aspects", []),
|
||||
needs_more_search=result.get("needs_more_search", False),
|
||||
suggested_queries=result.get("suggested_queries", [])
|
||||
)
|
||||
|
||||
logger.info(
|
||||
f"质量评估: completeness={assessment.completeness:.2f}, "
|
||||
f"needs_more_search={assessment.needs_more_search}"
|
||||
)
|
||||
|
||||
return assessment
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"质量评估失败: {e}")
|
||||
return self._default_assessment(answer)
|
||||
|
||||
def _default_assessment(self, answer: Answer) -> QualityAssessment:
|
||||
"""默认评估结果"""
|
||||
# 根据答案置信度估计完整性
|
||||
confidence_score = {
|
||||
"high": 0.9,
|
||||
"medium": 0.7,
|
||||
"low": 0.4
|
||||
}.get(answer.confidence, 0.5)
|
||||
|
||||
return QualityAssessment(
|
||||
completeness=confidence_score,
|
||||
missing_aspects=[],
|
||||
needs_more_search=confidence_score < self.COMPLETENESS_THRESHOLD,
|
||||
suggested_queries=[]
|
||||
)
|
||||
|
||||
def should_continue(
|
||||
self,
|
||||
assessment: QualityAssessment,
|
||||
current_iteration: int
|
||||
) -> bool:
|
||||
"""
|
||||
判断是否应该继续迭代
|
||||
|
||||
Args:
|
||||
assessment: 质量评估结果
|
||||
current_iteration: 当前迭代次数
|
||||
|
||||
Returns:
|
||||
是否继续迭代
|
||||
"""
|
||||
# 达到最大迭代次数
|
||||
if current_iteration >= self.config.max_iterations:
|
||||
logger.info(f"达到最大迭代次数 ({self.config.max_iterations}),停止迭代")
|
||||
return False
|
||||
|
||||
# 完整性达标
|
||||
if assessment.completeness >= self.COMPLETENESS_THRESHOLD:
|
||||
logger.info(f"完整性达标 ({assessment.completeness:.2f}),停止迭代")
|
||||
return False
|
||||
|
||||
# 没有建议的补充搜索
|
||||
if not assessment.suggested_queries:
|
||||
logger.info("没有建议的补充搜索,停止迭代")
|
||||
return False
|
||||
|
||||
return assessment.needs_more_search
|
||||
|
||||
+107
@@ -0,0 +1,107 @@
|
||||
"""
|
||||
结果处理模块
|
||||
负责结果去重、相关性排序、筛选
|
||||
"""
|
||||
|
||||
from typing import List
|
||||
from loguru import logger
|
||||
|
||||
from search_agent.config import Config
|
||||
from search_agent.models.schemas import Document, RankedDocument
|
||||
from search_agent.tools.jina_reranker import JinaRerankerClient
|
||||
from search_agent.utils.helpers import deduplicate_by_url
|
||||
|
||||
|
||||
class ResultProcessor:
|
||||
"""结果处理模块"""
|
||||
|
||||
def __init__(self, config: Config):
|
||||
"""
|
||||
初始化结果处理器
|
||||
|
||||
Args:
|
||||
config: 配置对象
|
||||
"""
|
||||
self.config = config
|
||||
self.reranker = JinaRerankerClient(
|
||||
api_key=config.jina_api_key,
|
||||
timeout=config.timeout
|
||||
)
|
||||
|
||||
async def process(
|
||||
self,
|
||||
query: str,
|
||||
documents: List[Document],
|
||||
top_k: int = 5
|
||||
) -> List[RankedDocument]:
|
||||
"""
|
||||
处理文档:去重 + 重排序 + 筛选
|
||||
|
||||
Args:
|
||||
query: 原始查询
|
||||
documents: 文档列表
|
||||
top_k: 返回前k个结果
|
||||
|
||||
Returns:
|
||||
排序后的RankedDocument列表
|
||||
"""
|
||||
if not documents:
|
||||
logger.warning("没有文档需要处理")
|
||||
return []
|
||||
|
||||
logger.info(f"开始处理 {len(documents)} 个文档")
|
||||
|
||||
# 1. 去重
|
||||
unique_docs = self._deduplicate(documents)
|
||||
logger.debug(f"去重后: {len(unique_docs)} 个文档")
|
||||
|
||||
# 2. 过滤空内容
|
||||
valid_docs = [d for d in unique_docs if d.content and len(d.content.strip()) > 50]
|
||||
logger.debug(f"有效文档: {len(valid_docs)} 个")
|
||||
|
||||
if not valid_docs:
|
||||
logger.warning("没有有效文档")
|
||||
return []
|
||||
|
||||
# 3. 重排序
|
||||
ranked_docs = await self.reranker.rerank(
|
||||
query=query,
|
||||
documents=valid_docs,
|
||||
top_k=top_k,
|
||||
content_max_length=self.config.content_max_length // 5 # 使用较短内容进行排序
|
||||
)
|
||||
|
||||
logger.info(f"处理完成,返回 {len(ranked_docs)} 个排序结果")
|
||||
return ranked_docs
|
||||
|
||||
def _deduplicate(self, documents: List[Document]) -> List[Document]:
|
||||
"""去重文档"""
|
||||
return deduplicate_by_url(documents, "url")
|
||||
|
||||
async def process_without_rerank(
|
||||
self,
|
||||
documents: List[Document],
|
||||
top_k: int = 5
|
||||
) -> List[RankedDocument]:
|
||||
"""
|
||||
处理文档(不进行重排序)
|
||||
|
||||
Args:
|
||||
documents: 文档列表
|
||||
top_k: 返回前k个结果
|
||||
|
||||
Returns:
|
||||
RankedDocument列表(按原始顺序)
|
||||
"""
|
||||
unique_docs = self._deduplicate(documents)
|
||||
valid_docs = [d for d in unique_docs if d.content and len(d.content.strip()) > 50]
|
||||
|
||||
return [
|
||||
RankedDocument(
|
||||
document=doc,
|
||||
relevance_score=1.0 - (i * 0.1),
|
||||
rank=i + 1
|
||||
)
|
||||
for i, doc in enumerate(valid_docs[:top_k])
|
||||
]
|
||||
|
||||
+89
@@ -0,0 +1,89 @@
|
||||
"""
|
||||
搜索执行模块
|
||||
执行搜索计划,调用Serper API
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
from typing import List
|
||||
from loguru import logger
|
||||
|
||||
from search_agent.config import Config
|
||||
from search_agent.models.schemas import SearchPlan, SearchTask, SearchResult
|
||||
from search_agent.tools.serper import SerperClient
|
||||
|
||||
|
||||
class SearchExecutor:
|
||||
"""搜索执行模块"""
|
||||
|
||||
def __init__(self, config: Config):
|
||||
"""
|
||||
初始化搜索执行器
|
||||
|
||||
Args:
|
||||
config: 配置对象
|
||||
"""
|
||||
self.config = config
|
||||
self.serper = SerperClient(
|
||||
api_key=config.serper_api_key,
|
||||
timeout=config.timeout
|
||||
)
|
||||
|
||||
async def execute(self, plan: SearchPlan) -> List[SearchResult]:
|
||||
"""
|
||||
执行搜索计划
|
||||
|
||||
Args:
|
||||
plan: 搜索计划
|
||||
|
||||
Returns:
|
||||
搜索结果列表
|
||||
"""
|
||||
logger.info(f"开始执行搜索计划: {len(plan.tasks)} 个任务")
|
||||
|
||||
if plan.strategy == "parallel":
|
||||
results = await self._execute_parallel(plan.tasks)
|
||||
else:
|
||||
results = await self._execute_sequential(plan.tasks)
|
||||
|
||||
logger.info(f"搜索完成,共获取 {len(results)} 条结果")
|
||||
return results
|
||||
|
||||
async def _execute_parallel(self, tasks: List[SearchTask]) -> List[SearchResult]:
|
||||
"""并行执行搜索任务"""
|
||||
coroutines = [self._execute_task(task) for task in tasks]
|
||||
results_list = await asyncio.gather(*coroutines, return_exceptions=True)
|
||||
|
||||
# 合并结果
|
||||
all_results = []
|
||||
for results in results_list:
|
||||
if isinstance(results, list):
|
||||
all_results.extend(results)
|
||||
elif isinstance(results, Exception):
|
||||
logger.warning(f"搜索任务失败: {results}")
|
||||
|
||||
return all_results
|
||||
|
||||
async def _execute_sequential(self, tasks: List[SearchTask]) -> List[SearchResult]:
|
||||
"""串行执行搜索任务"""
|
||||
all_results = []
|
||||
|
||||
for task in tasks:
|
||||
try:
|
||||
results = await self._execute_task(task)
|
||||
all_results.extend(results)
|
||||
except Exception as e:
|
||||
logger.warning(f"搜索任务失败: {e}")
|
||||
|
||||
return all_results
|
||||
|
||||
async def _execute_task(self, task: SearchTask) -> List[SearchResult]:
|
||||
"""执行单个搜索任务"""
|
||||
logger.debug(f"执行搜索: {task.query} [{task.source.value}]")
|
||||
|
||||
return await self.serper.search(
|
||||
query=task.query,
|
||||
source=task.source,
|
||||
num_results=task.num_results,
|
||||
time_filter=task.time_filter
|
||||
)
|
||||
|
||||
+137
@@ -0,0 +1,137 @@
|
||||
"""
|
||||
搜索规划模块
|
||||
根据查询分析结果制定搜索计划
|
||||
"""
|
||||
|
||||
from typing import List
|
||||
from loguru import logger
|
||||
|
||||
from search_agent.config import Config
|
||||
from search_agent.models.schemas import (
|
||||
QueryAnalysis,
|
||||
SearchPlan,
|
||||
SearchTask,
|
||||
SearchSource,
|
||||
Intent
|
||||
)
|
||||
|
||||
|
||||
class SearchPlanner:
|
||||
"""搜索规划模块"""
|
||||
|
||||
def __init__(self, config: Config):
|
||||
"""
|
||||
初始化搜索规划器
|
||||
|
||||
Args:
|
||||
config: 配置对象
|
||||
"""
|
||||
self.config = config
|
||||
self.max_results = config.max_results_per_query
|
||||
|
||||
async def plan(self, analysis: QueryAnalysis) -> SearchPlan:
|
||||
"""
|
||||
根据查询分析制定搜索计划
|
||||
|
||||
Args:
|
||||
analysis: 查询分析结果
|
||||
|
||||
Returns:
|
||||
SearchPlan对象
|
||||
"""
|
||||
logger.info(f"开始制定搜索计划: intent={analysis.intent.value}")
|
||||
|
||||
tasks = []
|
||||
|
||||
# 根据意图确定搜索策略
|
||||
strategy = self._determine_strategy(analysis)
|
||||
|
||||
# 构建搜索任务
|
||||
tasks.extend(self._create_web_tasks(analysis))
|
||||
|
||||
if analysis.need_news:
|
||||
tasks.extend(self._create_news_tasks(analysis))
|
||||
|
||||
plan = SearchPlan(
|
||||
tasks=tasks,
|
||||
strategy=strategy
|
||||
)
|
||||
|
||||
logger.info(f"搜索计划: {len(tasks)} 个任务, 策略={strategy}")
|
||||
return plan
|
||||
|
||||
def _determine_strategy(self, analysis: QueryAnalysis) -> str:
|
||||
"""确定执行策略"""
|
||||
# 大多数情况使用并行策略
|
||||
if analysis.intent == Intent.COMPARISON:
|
||||
# 对比类查询可能需要串行以获取更相关的结果
|
||||
return "parallel"
|
||||
return "parallel"
|
||||
|
||||
def _create_web_tasks(self, analysis: QueryAnalysis) -> List[SearchTask]:
|
||||
"""创建Web搜索任务"""
|
||||
tasks = []
|
||||
|
||||
# 原始查询
|
||||
tasks.append(SearchTask(
|
||||
query=analysis.original_query,
|
||||
source=SearchSource.WEB,
|
||||
time_filter=analysis.time_filter,
|
||||
num_results=self.max_results
|
||||
))
|
||||
|
||||
# 扩展查询(限制数量避免过多请求)
|
||||
for query in analysis.expanded_queries[:2]:
|
||||
if query != analysis.original_query:
|
||||
tasks.append(SearchTask(
|
||||
query=query,
|
||||
source=SearchSource.WEB,
|
||||
time_filter=analysis.time_filter,
|
||||
num_results=self.max_results
|
||||
))
|
||||
|
||||
return tasks
|
||||
|
||||
def _create_news_tasks(self, analysis: QueryAnalysis) -> List[SearchTask]:
|
||||
"""创建新闻搜索任务"""
|
||||
tasks = []
|
||||
|
||||
# 新闻搜索使用原始查询
|
||||
tasks.append(SearchTask(
|
||||
query=analysis.original_query,
|
||||
source=SearchSource.NEWS,
|
||||
time_filter=analysis.time_filter or "qdr:m", # 默认过去一个月
|
||||
num_results=self.max_results
|
||||
))
|
||||
|
||||
return tasks
|
||||
|
||||
def plan_supplementary(
|
||||
self,
|
||||
original_query: str,
|
||||
suggested_queries: List[str]
|
||||
) -> SearchPlan:
|
||||
"""
|
||||
创建补充搜索计划
|
||||
|
||||
Args:
|
||||
original_query: 原始查询
|
||||
suggested_queries: 建议的补充查询
|
||||
|
||||
Returns:
|
||||
SearchPlan对象
|
||||
"""
|
||||
tasks = []
|
||||
|
||||
for query in suggested_queries[:3]: # 限制补充搜索数量
|
||||
tasks.append(SearchTask(
|
||||
query=query,
|
||||
source=SearchSource.WEB,
|
||||
num_results=self.max_results
|
||||
))
|
||||
|
||||
return SearchPlan(
|
||||
tasks=tasks,
|
||||
strategy="parallel"
|
||||
)
|
||||
|
||||
@@ -0,0 +1,19 @@
|
||||
# HTTP客户端
|
||||
aiohttp>=3.9.0
|
||||
requests>=2.31.0
|
||||
|
||||
# 环境变量
|
||||
python-dotenv>=1.0.0
|
||||
|
||||
# JSON处理
|
||||
orjson>=3.9.0
|
||||
|
||||
# 类型提示
|
||||
typing-extensions>=4.9.0
|
||||
|
||||
# 日志
|
||||
loguru>=0.7.0
|
||||
|
||||
# 异步工具
|
||||
asyncio-throttle>=1.0.2
|
||||
|
||||
@@ -0,0 +1,16 @@
|
||||
import requests
|
||||
import json
|
||||
|
||||
url = "https://google.serper.dev/search"
|
||||
|
||||
payload = json.dumps({
|
||||
"q": "apple inc"
|
||||
})
|
||||
headers = {
|
||||
'X-API-KEY': '8253b4f240b520194065312f90e85f9be0fa205f',
|
||||
'Content-Type': 'application/json'
|
||||
}
|
||||
|
||||
response = requests.request("POST", url, headers=headers, data=payload)
|
||||
|
||||
print(response.text)
|
||||
@@ -0,0 +1,14 @@
|
||||
"""
|
||||
外部API工具封装模块
|
||||
"""
|
||||
|
||||
from .serper import SerperClient
|
||||
from .jina_reader import JinaReaderClient
|
||||
from .jina_reranker import JinaRerankerClient
|
||||
|
||||
__all__ = [
|
||||
"SerperClient",
|
||||
"JinaReaderClient",
|
||||
"JinaRerankerClient",
|
||||
]
|
||||
|
||||
+180
@@ -0,0 +1,180 @@
|
||||
"""
|
||||
Jina Reader API封装
|
||||
提供网页内容提取功能
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
from typing import List, Optional
|
||||
import aiohttp
|
||||
from loguru import logger
|
||||
|
||||
from search_agent.models.schemas import Document, SearchSource
|
||||
|
||||
|
||||
class JinaReaderClient:
|
||||
"""Jina Reader API客户端"""
|
||||
|
||||
BASE_URL = "https://r.jina.ai"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
api_key: str,
|
||||
timeout: int = 30,
|
||||
max_concurrent: int = 5,
|
||||
max_content_length: int = 5000
|
||||
):
|
||||
"""
|
||||
初始化Jina Reader客户端
|
||||
|
||||
Args:
|
||||
api_key: Jina API密钥
|
||||
timeout: 请求超时时间(秒)
|
||||
max_concurrent: 最大并发请求数
|
||||
max_content_length: 最大内容长度
|
||||
"""
|
||||
self.api_key = api_key
|
||||
self.timeout = timeout
|
||||
self.max_concurrent = max_concurrent
|
||||
self.max_content_length = max_content_length
|
||||
self._semaphore = asyncio.Semaphore(max_concurrent)
|
||||
|
||||
async def extract_content(
|
||||
self,
|
||||
url: str,
|
||||
source: SearchSource = SearchSource.WEB
|
||||
) -> Optional[Document]:
|
||||
"""
|
||||
提取单个URL的内容
|
||||
|
||||
Args:
|
||||
url: 要提取的网页URL
|
||||
source: 来源类型
|
||||
|
||||
Returns:
|
||||
Document对象,如果提取失败则返回None
|
||||
"""
|
||||
reader_url = f"{self.BASE_URL}/{url}"
|
||||
|
||||
headers = {
|
||||
"Authorization": f"Bearer {self.api_key}",
|
||||
"Accept": "application/json"
|
||||
}
|
||||
|
||||
try:
|
||||
async with self._semaphore:
|
||||
async with aiohttp.ClientSession() as session:
|
||||
async with session.get(
|
||||
reader_url,
|
||||
headers=headers,
|
||||
timeout=aiohttp.ClientTimeout(total=self.timeout)
|
||||
) as response:
|
||||
if response.status != 200:
|
||||
logger.warning(f"Jina Reader提取失败 [{response.status}]: {url}")
|
||||
return None
|
||||
|
||||
# Jina Reader可能返回JSON或纯文本
|
||||
content_type = response.headers.get("Content-Type", "")
|
||||
|
||||
if "application/json" in content_type:
|
||||
result = await response.json()
|
||||
# 处理嵌套的data字段
|
||||
if "data" in result:
|
||||
result = result["data"]
|
||||
content = result.get("content", "")
|
||||
title = result.get("title", "")
|
||||
else:
|
||||
# 纯文本响应(Markdown格式)
|
||||
content = await response.text()
|
||||
# 从内容中提取标题(第一行通常是标题)
|
||||
lines = content.strip().split("\n")
|
||||
title = lines[0].lstrip("#").strip() if lines else ""
|
||||
|
||||
# 限制内容长度
|
||||
if len(content) > self.max_content_length:
|
||||
content = content[:self.max_content_length]
|
||||
|
||||
logger.debug(f"提取成功: {url[:50]}... 内容长度: {len(content)}")
|
||||
|
||||
return Document(
|
||||
url=url,
|
||||
title=title,
|
||||
content=content,
|
||||
source=source
|
||||
)
|
||||
|
||||
except aiohttp.ClientError as e:
|
||||
logger.warning(f"Jina Reader网络错误 [{url}]: {e}")
|
||||
return None
|
||||
except asyncio.TimeoutError:
|
||||
logger.warning(f"Jina Reader超时: {url}")
|
||||
return None
|
||||
except Exception as e:
|
||||
logger.warning(f"Jina Reader异常 [{url}]: {e}")
|
||||
return None
|
||||
|
||||
async def extract_batch(
|
||||
self,
|
||||
urls: List[str],
|
||||
source: SearchSource = SearchSource.WEB
|
||||
) -> List[Document]:
|
||||
"""
|
||||
批量提取多个URL的内容
|
||||
|
||||
Args:
|
||||
urls: URL列表
|
||||
source: 来源类型
|
||||
|
||||
Returns:
|
||||
成功提取的Document列表
|
||||
"""
|
||||
logger.info(f"批量提取 {len(urls)} 个URL的内容")
|
||||
|
||||
tasks = [
|
||||
self.extract_content(url, source)
|
||||
for url in urls
|
||||
]
|
||||
|
||||
results = await asyncio.gather(*tasks, return_exceptions=True)
|
||||
|
||||
# 过滤掉失败的结果
|
||||
documents = []
|
||||
for result in results:
|
||||
if isinstance(result, Document):
|
||||
documents.append(result)
|
||||
elif isinstance(result, Exception):
|
||||
logger.warning(f"提取异常: {result}")
|
||||
|
||||
logger.info(f"成功提取 {len(documents)}/{len(urls)} 个文档")
|
||||
return documents
|
||||
|
||||
async def extract_with_retry(
|
||||
self,
|
||||
url: str,
|
||||
source: SearchSource = SearchSource.WEB,
|
||||
max_retries: int = 2,
|
||||
retry_delay: float = 1.0
|
||||
) -> Optional[Document]:
|
||||
"""
|
||||
带重试的内容提取
|
||||
|
||||
Args:
|
||||
url: 要提取的网页URL
|
||||
source: 来源类型
|
||||
max_retries: 最大重试次数
|
||||
retry_delay: 重试延迟(秒)
|
||||
|
||||
Returns:
|
||||
Document对象,如果最终失败则返回None
|
||||
"""
|
||||
for attempt in range(max_retries + 1):
|
||||
result = await self.extract_content(url, source)
|
||||
|
||||
if result is not None:
|
||||
return result
|
||||
|
||||
if attempt < max_retries:
|
||||
logger.debug(f"重试提取 [{attempt + 1}/{max_retries}]: {url}")
|
||||
await asyncio.sleep(retry_delay)
|
||||
|
||||
return None
|
||||
|
||||
+191
@@ -0,0 +1,191 @@
|
||||
"""
|
||||
Jina Reranker API封装
|
||||
提供搜索结果重排序功能
|
||||
"""
|
||||
|
||||
from typing import List, Tuple
|
||||
import aiohttp
|
||||
from loguru import logger
|
||||
|
||||
from search_agent.models.schemas import Document, RankedDocument
|
||||
|
||||
|
||||
class JinaRerankerClient:
|
||||
"""Jina Reranker API客户端"""
|
||||
|
||||
BASE_URL = "https://api.jina.ai/v1/rerank"
|
||||
MODEL = "jina-reranker-v2-base-multilingual"
|
||||
|
||||
def __init__(self, api_key: str, timeout: int = 30):
|
||||
"""
|
||||
初始化Jina Reranker客户端
|
||||
|
||||
Args:
|
||||
api_key: Jina API密钥
|
||||
timeout: 请求超时时间(秒)
|
||||
"""
|
||||
self.api_key = api_key
|
||||
self.timeout = timeout
|
||||
|
||||
async def rerank(
|
||||
self,
|
||||
query: str,
|
||||
documents: List[Document],
|
||||
top_k: int = 5,
|
||||
content_max_length: int = 1000
|
||||
) -> List[RankedDocument]:
|
||||
"""
|
||||
对文档进行相关性重排序
|
||||
|
||||
Args:
|
||||
query: 查询字符串
|
||||
documents: 文档列表
|
||||
top_k: 返回前k个结果
|
||||
content_max_length: 用于排序的内容最大长度
|
||||
|
||||
Returns:
|
||||
排序后的RankedDocument列表
|
||||
"""
|
||||
if not documents:
|
||||
return []
|
||||
|
||||
# 准备文档内容(截断到合适长度)
|
||||
doc_contents = [
|
||||
doc.content[:content_max_length] if doc.content else doc.title
|
||||
for doc in documents
|
||||
]
|
||||
|
||||
headers = {
|
||||
"Authorization": f"Bearer {self.api_key}",
|
||||
"Content-Type": "application/json"
|
||||
}
|
||||
|
||||
payload = {
|
||||
"model": self.MODEL,
|
||||
"query": query,
|
||||
"documents": doc_contents,
|
||||
"top_n": min(top_k, len(documents))
|
||||
}
|
||||
|
||||
try:
|
||||
async with aiohttp.ClientSession() as session:
|
||||
async with session.post(
|
||||
self.BASE_URL,
|
||||
headers=headers,
|
||||
json=payload,
|
||||
timeout=aiohttp.ClientTimeout(total=self.timeout)
|
||||
) as response:
|
||||
if response.status != 200:
|
||||
error_text = await response.text()
|
||||
logger.error(f"Jina Reranker API错误: {response.status} - {error_text}")
|
||||
# 如果重排序失败,返回原始顺序
|
||||
return self._fallback_ranking(documents, top_k)
|
||||
|
||||
result = await response.json()
|
||||
return self._parse_rerank_results(documents, result, top_k)
|
||||
|
||||
except aiohttp.ClientError as e:
|
||||
logger.error(f"Jina Reranker网络错误: {e}")
|
||||
return self._fallback_ranking(documents, top_k)
|
||||
except Exception as e:
|
||||
logger.error(f"Jina Reranker异常: {e}")
|
||||
return self._fallback_ranking(documents, top_k)
|
||||
|
||||
def _parse_rerank_results(
|
||||
self,
|
||||
documents: List[Document],
|
||||
response: dict,
|
||||
top_k: int
|
||||
) -> List[RankedDocument]:
|
||||
"""解析重排序结果"""
|
||||
results = []
|
||||
|
||||
reranked = response.get("results", [])
|
||||
|
||||
for rank, item in enumerate(reranked[:top_k], 1):
|
||||
index = item.get("index", 0)
|
||||
score = item.get("relevance_score", 0.0)
|
||||
|
||||
if 0 <= index < len(documents):
|
||||
ranked_doc = RankedDocument(
|
||||
document=documents[index],
|
||||
relevance_score=score,
|
||||
rank=rank
|
||||
)
|
||||
results.append(ranked_doc)
|
||||
|
||||
logger.debug(f"重排序返回 {len(results)} 个结果")
|
||||
return results
|
||||
|
||||
def _fallback_ranking(
|
||||
self,
|
||||
documents: List[Document],
|
||||
top_k: int
|
||||
) -> List[RankedDocument]:
|
||||
"""后备排序:保持原始顺序"""
|
||||
logger.warning("使用后备排序(原始顺序)")
|
||||
|
||||
return [
|
||||
RankedDocument(
|
||||
document=doc,
|
||||
relevance_score=1.0 - (i * 0.1), # 模拟递减分数
|
||||
rank=i + 1
|
||||
)
|
||||
for i, doc in enumerate(documents[:top_k])
|
||||
]
|
||||
|
||||
async def rerank_texts(
|
||||
self,
|
||||
query: str,
|
||||
texts: List[str],
|
||||
top_k: int = 5
|
||||
) -> List[Tuple[int, float]]:
|
||||
"""
|
||||
对纯文本列表进行重排序
|
||||
|
||||
Args:
|
||||
query: 查询字符串
|
||||
texts: 文本列表
|
||||
top_k: 返回前k个结果
|
||||
|
||||
Returns:
|
||||
(原始索引, 相关性分数) 的列表
|
||||
"""
|
||||
if not texts:
|
||||
return []
|
||||
|
||||
headers = {
|
||||
"Authorization": f"Bearer {self.api_key}",
|
||||
"Content-Type": "application/json"
|
||||
}
|
||||
|
||||
payload = {
|
||||
"model": self.MODEL,
|
||||
"query": query,
|
||||
"documents": texts,
|
||||
"top_n": min(top_k, len(texts))
|
||||
}
|
||||
|
||||
try:
|
||||
async with aiohttp.ClientSession() as session:
|
||||
async with session.post(
|
||||
self.BASE_URL,
|
||||
headers=headers,
|
||||
json=payload,
|
||||
timeout=aiohttp.ClientTimeout(total=self.timeout)
|
||||
) as response:
|
||||
if response.status != 200:
|
||||
logger.error(f"Reranker API错误: {response.status}")
|
||||
return [(i, 1.0 - i * 0.1) for i in range(min(top_k, len(texts)))]
|
||||
|
||||
result = await response.json()
|
||||
|
||||
return [
|
||||
(item["index"], item["relevance_score"])
|
||||
for item in result.get("results", [])[:top_k]
|
||||
]
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Reranker异常: {e}")
|
||||
return [(i, 1.0 - i * 0.1) for i in range(min(top_k, len(texts)))]
|
||||
|
||||
@@ -0,0 +1,212 @@
|
||||
"""
|
||||
Serper API封装
|
||||
提供Google搜索和新闻搜索功能
|
||||
"""
|
||||
|
||||
from typing import List, Optional, Dict, Any
|
||||
import aiohttp
|
||||
from loguru import logger
|
||||
|
||||
from search_agent.models.schemas import SearchResult, SearchSource
|
||||
|
||||
|
||||
class SerperClient:
|
||||
"""Serper API客户端"""
|
||||
|
||||
BASE_URL = "https://google.serper.dev"
|
||||
|
||||
ENDPOINTS = {
|
||||
"web": "/search",
|
||||
"news": "/news"
|
||||
}
|
||||
|
||||
def __init__(self, api_key: str, timeout: int = 30):
|
||||
"""
|
||||
初始化Serper客户端
|
||||
|
||||
Args:
|
||||
api_key: Serper API密钥
|
||||
timeout: 请求超时时间(秒)
|
||||
"""
|
||||
self.api_key = api_key
|
||||
self.timeout = timeout
|
||||
|
||||
async def _request(
|
||||
self,
|
||||
endpoint: str,
|
||||
payload: Dict[str, Any]
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
发送请求到Serper API
|
||||
|
||||
Args:
|
||||
endpoint: API端点
|
||||
payload: 请求体
|
||||
|
||||
Returns:
|
||||
API响应
|
||||
"""
|
||||
url = f"{self.BASE_URL}{endpoint}"
|
||||
|
||||
headers = {
|
||||
"X-API-KEY": self.api_key,
|
||||
"Content-Type": "application/json"
|
||||
}
|
||||
|
||||
try:
|
||||
async with aiohttp.ClientSession() as session:
|
||||
async with session.post(
|
||||
url,
|
||||
headers=headers,
|
||||
json=payload,
|
||||
timeout=aiohttp.ClientTimeout(total=self.timeout)
|
||||
) as response:
|
||||
if response.status != 200:
|
||||
error_text = await response.text()
|
||||
logger.error(f"Serper API错误: {response.status} - {error_text}")
|
||||
raise Exception(f"Serper API请求失败: {response.status}")
|
||||
|
||||
return await response.json()
|
||||
|
||||
except aiohttp.ClientError as e:
|
||||
logger.error(f"Serper请求网络错误: {e}")
|
||||
raise
|
||||
|
||||
async def search_web(
|
||||
self,
|
||||
query: str,
|
||||
num_results: int = 10,
|
||||
gl: str = "cn",
|
||||
hl: str = "zh-cn",
|
||||
time_filter: Optional[str] = None
|
||||
) -> List[SearchResult]:
|
||||
"""
|
||||
执行Web搜索
|
||||
|
||||
Args:
|
||||
query: 搜索查询
|
||||
num_results: 返回结果数量
|
||||
gl: 地区代码
|
||||
hl: 语言代码
|
||||
time_filter: 时间过滤器 (qdr:d/qdr:w/qdr:m/qdr:y)
|
||||
|
||||
Returns:
|
||||
搜索结果列表
|
||||
"""
|
||||
payload = {
|
||||
"q": query,
|
||||
"num": num_results,
|
||||
"gl": gl,
|
||||
"hl": hl
|
||||
}
|
||||
|
||||
if time_filter:
|
||||
payload["tbs"] = time_filter
|
||||
|
||||
logger.info(f"执行Web搜索: {query}")
|
||||
|
||||
result = await self._request(self.ENDPOINTS["web"], payload)
|
||||
|
||||
return self._parse_web_results(result)
|
||||
|
||||
async def search_news(
|
||||
self,
|
||||
query: str,
|
||||
num_results: int = 10,
|
||||
gl: str = "cn",
|
||||
hl: str = "zh-cn",
|
||||
time_filter: Optional[str] = None
|
||||
) -> List[SearchResult]:
|
||||
"""
|
||||
执行新闻搜索
|
||||
|
||||
Args:
|
||||
query: 搜索查询
|
||||
num_results: 返回结果数量
|
||||
gl: 地区代码
|
||||
hl: 语言代码
|
||||
time_filter: 时间过滤器
|
||||
|
||||
Returns:
|
||||
搜索结果列表
|
||||
"""
|
||||
payload = {
|
||||
"q": query,
|
||||
"num": num_results,
|
||||
"gl": gl,
|
||||
"hl": hl
|
||||
}
|
||||
|
||||
if time_filter:
|
||||
payload["tbs"] = time_filter
|
||||
|
||||
logger.info(f"执行新闻搜索: {query}")
|
||||
|
||||
result = await self._request(self.ENDPOINTS["news"], payload)
|
||||
|
||||
return self._parse_news_results(result)
|
||||
|
||||
def _parse_web_results(self, response: Dict[str, Any]) -> List[SearchResult]:
|
||||
"""解析Web搜索结果"""
|
||||
results = []
|
||||
|
||||
organic = response.get("organic", [])
|
||||
|
||||
for item in organic:
|
||||
result = SearchResult(
|
||||
title=item.get("title", ""),
|
||||
url=item.get("link", ""),
|
||||
snippet=item.get("snippet", ""),
|
||||
source=SearchSource.WEB,
|
||||
position=item.get("position", 0),
|
||||
date=None
|
||||
)
|
||||
results.append(result)
|
||||
|
||||
logger.debug(f"Web搜索返回 {len(results)} 条结果")
|
||||
return results
|
||||
|
||||
def _parse_news_results(self, response: Dict[str, Any]) -> List[SearchResult]:
|
||||
"""解析新闻搜索结果"""
|
||||
results = []
|
||||
|
||||
news = response.get("news", [])
|
||||
|
||||
for i, item in enumerate(news, 1):
|
||||
result = SearchResult(
|
||||
title=item.get("title", ""),
|
||||
url=item.get("link", ""),
|
||||
snippet=item.get("snippet", ""),
|
||||
source=SearchSource.NEWS,
|
||||
position=i,
|
||||
date=item.get("date")
|
||||
)
|
||||
results.append(result)
|
||||
|
||||
logger.debug(f"新闻搜索返回 {len(results)} 条结果")
|
||||
return results
|
||||
|
||||
async def search(
|
||||
self,
|
||||
query: str,
|
||||
source: SearchSource,
|
||||
num_results: int = 10,
|
||||
time_filter: Optional[str] = None
|
||||
) -> List[SearchResult]:
|
||||
"""
|
||||
统一搜索接口
|
||||
|
||||
Args:
|
||||
query: 搜索查询
|
||||
source: 搜索来源类型
|
||||
num_results: 返回结果数量
|
||||
time_filter: 时间过滤器
|
||||
|
||||
Returns:
|
||||
搜索结果列表
|
||||
"""
|
||||
if source == SearchSource.NEWS:
|
||||
return await self.search_news(query, num_results, time_filter=time_filter)
|
||||
else:
|
||||
return await self.search_web(query, num_results, time_filter=time_filter)
|
||||
|
||||
@@ -0,0 +1,22 @@
|
||||
"""
|
||||
工具函数模块
|
||||
"""
|
||||
|
||||
from .llm_client import LLMClient
|
||||
from .helpers import (
|
||||
flatten,
|
||||
deduplicate_by_url,
|
||||
truncate_text,
|
||||
extract_json_from_text,
|
||||
format_documents_for_prompt,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"LLMClient",
|
||||
"flatten",
|
||||
"deduplicate_by_url",
|
||||
"truncate_text",
|
||||
"extract_json_from_text",
|
||||
"format_documents_for_prompt",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,197 @@
|
||||
"""
|
||||
通用工具函数
|
||||
"""
|
||||
|
||||
import re
|
||||
import json
|
||||
from typing import List, TypeVar, Optional, Dict, Any
|
||||
|
||||
T = TypeVar('T')
|
||||
|
||||
|
||||
def flatten(nested_list: List[List[T]]) -> List[T]:
|
||||
"""
|
||||
将嵌套列表展平为一维列表
|
||||
|
||||
Args:
|
||||
nested_list: 嵌套列表
|
||||
|
||||
Returns:
|
||||
展平后的一维列表
|
||||
"""
|
||||
return [item for sublist in nested_list for item in sublist]
|
||||
|
||||
|
||||
def deduplicate_by_url(items: List[Any], url_attr: str = "url") -> List[Any]:
|
||||
"""
|
||||
根据URL去重
|
||||
|
||||
Args:
|
||||
items: 包含URL属性的对象列表
|
||||
url_attr: URL属性名
|
||||
|
||||
Returns:
|
||||
去重后的列表
|
||||
"""
|
||||
seen_urls = set()
|
||||
unique_items = []
|
||||
|
||||
for item in items:
|
||||
url = getattr(item, url_attr, None) or item.get(url_attr)
|
||||
if url and url not in seen_urls:
|
||||
seen_urls.add(url)
|
||||
unique_items.append(item)
|
||||
|
||||
return unique_items
|
||||
|
||||
|
||||
def truncate_text(text: str, max_length: int, suffix: str = "...") -> str:
|
||||
"""
|
||||
截断文本到指定长度
|
||||
|
||||
Args:
|
||||
text: 原始文本
|
||||
max_length: 最大长度
|
||||
suffix: 截断后缀
|
||||
|
||||
Returns:
|
||||
截断后的文本
|
||||
"""
|
||||
if len(text) <= max_length:
|
||||
return text
|
||||
|
||||
return text[:max_length - len(suffix)] + suffix
|
||||
|
||||
|
||||
def extract_json_from_text(text: str) -> Optional[Dict[str, Any]]:
|
||||
"""
|
||||
从文本中提取JSON对象
|
||||
|
||||
Args:
|
||||
text: 可能包含JSON的文本
|
||||
|
||||
Returns:
|
||||
提取的JSON字典,如果提取失败则返回None
|
||||
"""
|
||||
# 尝试直接解析
|
||||
try:
|
||||
return json.loads(text)
|
||||
except json.JSONDecodeError:
|
||||
pass
|
||||
|
||||
# 尝试提取```json ... ```块
|
||||
json_block_pattern = r'```(?:json)?\s*([\s\S]*?)```'
|
||||
matches = re.findall(json_block_pattern, text)
|
||||
|
||||
for match in matches:
|
||||
try:
|
||||
return json.loads(match.strip())
|
||||
except json.JSONDecodeError:
|
||||
continue
|
||||
|
||||
# 尝试提取{ ... }块
|
||||
brace_pattern = r'\{[\s\S]*\}'
|
||||
matches = re.findall(brace_pattern, text)
|
||||
|
||||
for match in matches:
|
||||
try:
|
||||
return json.loads(match)
|
||||
except json.JSONDecodeError:
|
||||
continue
|
||||
|
||||
return None
|
||||
|
||||
|
||||
def format_documents_for_prompt(documents: List[Any], max_length: int = 2000) -> str:
|
||||
"""
|
||||
格式化文档列表为Prompt中使用的文本
|
||||
|
||||
Args:
|
||||
documents: 文档列表(RankedDocument或Document对象)
|
||||
max_length: 每个文档的最大内容长度
|
||||
|
||||
Returns:
|
||||
格式化后的文本
|
||||
"""
|
||||
formatted_parts = []
|
||||
|
||||
for i, doc in enumerate(documents, 1):
|
||||
# 支持RankedDocument和Document两种类型
|
||||
if hasattr(doc, 'document'):
|
||||
# RankedDocument
|
||||
actual_doc = doc.document
|
||||
score = f" (相关性: {doc.relevance_score:.2f})"
|
||||
else:
|
||||
# Document
|
||||
actual_doc = doc
|
||||
score = ""
|
||||
|
||||
content = truncate_text(actual_doc.content, max_length)
|
||||
|
||||
part = f"""### 来源 [{i}]{score}
|
||||
**标题**: {actual_doc.title}
|
||||
**URL**: {actual_doc.url}
|
||||
**内容**:
|
||||
{content}
|
||||
"""
|
||||
formatted_parts.append(part)
|
||||
|
||||
return "\n---\n".join(formatted_parts)
|
||||
|
||||
|
||||
def clean_url(url: str) -> str:
|
||||
"""
|
||||
清理和标准化URL
|
||||
|
||||
Args:
|
||||
url: 原始URL
|
||||
|
||||
Returns:
|
||||
清理后的URL
|
||||
"""
|
||||
# 移除末尾的斜杠
|
||||
url = url.rstrip("/")
|
||||
|
||||
# 移除锚点
|
||||
if "#" in url:
|
||||
url = url.split("#")[0]
|
||||
|
||||
return url
|
||||
|
||||
|
||||
def is_valid_url(url: str) -> bool:
|
||||
"""
|
||||
验证URL是否有效
|
||||
|
||||
Args:
|
||||
url: URL字符串
|
||||
|
||||
Returns:
|
||||
是否有效
|
||||
"""
|
||||
url_pattern = re.compile(
|
||||
r'^https?://' # http:// or https://
|
||||
r'(?:(?:[A-Z0-9](?:[A-Z0-9-]{0,61}[A-Z0-9])?\.)+[A-Z]{2,6}\.?|' # domain
|
||||
r'localhost|' # localhost
|
||||
r'\d{1,3}\.\d{1,3}\.\d{1,3}\.\d{1,3})' # IP
|
||||
r'(?::\d+)?' # optional port
|
||||
r'(?:/?|[/?]\S+)$', re.IGNORECASE)
|
||||
|
||||
return bool(url_pattern.match(url))
|
||||
|
||||
|
||||
def merge_dicts(base: Dict, override: Dict) -> Dict:
|
||||
"""
|
||||
合并两个字典,override中的值会覆盖base中的值
|
||||
|
||||
Args:
|
||||
base: 基础字典
|
||||
override: 覆盖字典
|
||||
|
||||
Returns:
|
||||
合并后的字典
|
||||
"""
|
||||
result = base.copy()
|
||||
result.update(override)
|
||||
return result
|
||||
|
||||
+159
@@ -0,0 +1,159 @@
|
||||
"""
|
||||
LLM客户端模块
|
||||
封装与xchat52 LLM的交互(支持Azure OpenAI风格API)
|
||||
"""
|
||||
|
||||
import json
|
||||
from typing import Optional, List, Dict, Any
|
||||
import aiohttp
|
||||
from loguru import logger
|
||||
|
||||
|
||||
class LLMClient:
|
||||
"""LLM客户端,用于与xchat52 API交互"""
|
||||
|
||||
# API版本
|
||||
API_VERSION = "2024-10-21"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
base_url: str,
|
||||
api_key: str,
|
||||
model: str = "xchat52",
|
||||
timeout: int = 60
|
||||
):
|
||||
self.base_url = base_url.rstrip("/")
|
||||
self.api_key = api_key
|
||||
self.model = model
|
||||
self.timeout = timeout
|
||||
|
||||
async def chat(
|
||||
self,
|
||||
messages: List[Dict[str, str]],
|
||||
temperature: float = 0.7,
|
||||
max_tokens: int = 4096,
|
||||
response_format: Optional[Dict[str, str]] = None
|
||||
) -> str:
|
||||
"""
|
||||
发送聊天请求到LLM
|
||||
|
||||
Args:
|
||||
messages: 消息列表,格式 [{"role": "user", "content": "..."}]
|
||||
temperature: 温度参数
|
||||
max_tokens: 最大token数
|
||||
response_format: 响应格式(如 {"type": "json_object"})
|
||||
|
||||
Returns:
|
||||
LLM的响应文本
|
||||
"""
|
||||
# Azure OpenAI 风格的URL
|
||||
url = f"{self.base_url}/chat/completions?api-version={self.API_VERSION}"
|
||||
|
||||
# Azure OpenAI 使用 api-key 头
|
||||
headers = {
|
||||
"api-key": self.api_key,
|
||||
"Content-Type": "application/json"
|
||||
}
|
||||
|
||||
payload = {
|
||||
"model": self.model,
|
||||
"messages": messages,
|
||||
"temperature": temperature,
|
||||
"max_completion_tokens": max_tokens # 新版API使用 max_completion_tokens
|
||||
}
|
||||
|
||||
if response_format:
|
||||
payload["response_format"] = response_format
|
||||
|
||||
try:
|
||||
async with aiohttp.ClientSession() as session:
|
||||
async with session.post(
|
||||
url,
|
||||
headers=headers,
|
||||
json=payload,
|
||||
timeout=aiohttp.ClientTimeout(total=self.timeout)
|
||||
) as response:
|
||||
if response.status != 200:
|
||||
error_text = await response.text()
|
||||
logger.error(f"LLM API错误: {response.status} - {error_text}")
|
||||
raise Exception(f"LLM API请求失败: {response.status}")
|
||||
|
||||
result = await response.json()
|
||||
return result["choices"][0]["message"]["content"]
|
||||
|
||||
except aiohttp.ClientError as e:
|
||||
logger.error(f"LLM请求网络错误: {e}")
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.error(f"LLM请求异常: {e}")
|
||||
raise
|
||||
|
||||
async def chat_with_system(
|
||||
self,
|
||||
system_prompt: str,
|
||||
user_message: str,
|
||||
temperature: float = 0.7,
|
||||
max_tokens: int = 4096,
|
||||
response_format: Optional[Dict[str, str]] = None
|
||||
) -> str:
|
||||
"""
|
||||
使用系统提示和用户消息进行对话
|
||||
|
||||
Args:
|
||||
system_prompt: 系统提示
|
||||
user_message: 用户消息
|
||||
temperature: 温度参数
|
||||
max_tokens: 最大token数
|
||||
response_format: 响应格式
|
||||
|
||||
Returns:
|
||||
LLM的响应文本
|
||||
"""
|
||||
messages = [
|
||||
{"role": "system", "content": system_prompt},
|
||||
{"role": "user", "content": user_message}
|
||||
]
|
||||
|
||||
return await self.chat(
|
||||
messages=messages,
|
||||
temperature=temperature,
|
||||
max_tokens=max_tokens,
|
||||
response_format=response_format
|
||||
)
|
||||
|
||||
async def chat_json(
|
||||
self,
|
||||
system_prompt: str,
|
||||
user_message: str,
|
||||
temperature: float = 0.3
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
请求JSON格式的响应
|
||||
|
||||
Args:
|
||||
system_prompt: 系统提示
|
||||
user_message: 用户消息
|
||||
temperature: 温度参数(JSON响应建议使用较低温度)
|
||||
|
||||
Returns:
|
||||
解析后的JSON字典
|
||||
"""
|
||||
from .helpers import extract_json_from_text
|
||||
|
||||
response = await self.chat_with_system(
|
||||
system_prompt=system_prompt,
|
||||
user_message=user_message,
|
||||
temperature=temperature,
|
||||
response_format={"type": "json_object"}
|
||||
)
|
||||
|
||||
try:
|
||||
return json.loads(response)
|
||||
except json.JSONDecodeError:
|
||||
# 尝试从文本中提取JSON
|
||||
extracted = extract_json_from_text(response)
|
||||
if extracted:
|
||||
return extracted
|
||||
logger.error(f"无法解析LLM响应为JSON: {response[:200]}")
|
||||
raise ValueError("LLM响应不是有效的JSON格式")
|
||||
|
||||
@@ -22,6 +22,10 @@ RUN pip install --no-cache-dir \
|
||||
# 复制search_agent_MCP目录
|
||||
COPY agents/search_agent/search_agent_MCP/ /app/
|
||||
|
||||
# 复制回调工具
|
||||
COPY common/agent_callback_utils.py /app/common/
|
||||
RUN touch /app/common/__init__.py
|
||||
|
||||
# 复制search_agent核心代码
|
||||
COPY agents/search_agent/search_agent/ /app/search_agent/
|
||||
|
||||
|
||||
@@ -16,11 +16,13 @@ RUN apt-get update && apt-get install -y \
|
||||
RUN ffmpeg -version
|
||||
|
||||
# 安装 Python 依赖
|
||||
COPY requirements.txt .
|
||||
RUN pip install --no-cache-dir -r requirements.txt
|
||||
COPY agent_templates/agents/video_generator_agent/requirements.txt .
|
||||
RUN pip install --no-cache-dir -r requirements.txt requests
|
||||
|
||||
# 复制应用代码
|
||||
COPY . .
|
||||
COPY agent_templates/agents/video_generator_agent/ /app/
|
||||
COPY agent_templates/common/agent_callback_utils.py /app/common/
|
||||
RUN touch /app/common/__init__.py
|
||||
|
||||
# 创建输出目录
|
||||
RUN mkdir -p /app/outputs/images /app/outputs/videos
|
||||
|
||||
@@ -22,19 +22,33 @@ from pydantic import BaseModel, Field
|
||||
from src.server.mcp_server import TOOL_MAP, TOOL_LIST
|
||||
from src.utils.file_manager import FileManager
|
||||
|
||||
try:
|
||||
from common.agent_callback_utils import AgentCallbackHandler, CallbackContextManager
|
||||
CALLBACK_ENABLED = True
|
||||
except ImportError:
|
||||
CALLBACK_ENABLED = False
|
||||
AgentCallbackHandler = None
|
||||
CallbackContextManager = None
|
||||
|
||||
# ==================== 配置 ====================
|
||||
|
||||
SERVER_NAME = "Video Generator Agent"
|
||||
OUTPUT_DIR = os.getenv('OUTPUT_DIR', '/app/outputs')
|
||||
POD_NAME = os.getenv("POD_NAME", "video-generator-agent")
|
||||
USER_ID = os.getenv("USER_ID", "")
|
||||
|
||||
file_manager = FileManager(base_dir=OUTPUT_DIR)
|
||||
callback_handler: Optional[AgentCallbackHandler] = None
|
||||
|
||||
# ==================== FastAPI 应用 ====================
|
||||
|
||||
@asynccontextmanager
|
||||
async def lifespan(app: FastAPI):
|
||||
global callback_handler
|
||||
print(f"🚀 {SERVER_NAME} 启动")
|
||||
print(f"📁 输出目录: {OUTPUT_DIR}")
|
||||
if CALLBACK_ENABLED and AgentCallbackHandler:
|
||||
callback_handler = AgentCallbackHandler(agent_name=POD_NAME, user_id=USER_ID)
|
||||
yield
|
||||
print(f"🛑 {SERVER_NAME} 关闭")
|
||||
|
||||
@@ -145,7 +159,16 @@ async def handle_mcp_request(data: Dict, session_id: str = None, api_key: str =
|
||||
os.environ['OPENAI_API_KEY'] = api_key
|
||||
|
||||
try:
|
||||
result = await TOOL_MAP[tool_name](**args)
|
||||
if CALLBACK_ENABLED and callback_handler:
|
||||
with CallbackContextManager(
|
||||
handler=callback_handler,
|
||||
user_id=USER_ID,
|
||||
request_id=f"video-mcp-{tool_name}-{req_id or uuid.uuid4().hex}"
|
||||
) 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
|
||||
@@ -238,11 +261,24 @@ async def api_generate_image(request: ImageGenerationRequest, api_key: str = Dep
|
||||
os.environ['OPENAI_API_KEY'] = api_key
|
||||
|
||||
try:
|
||||
result_str = await TOOL_MAP['generate_image'](
|
||||
description=request.description,
|
||||
size=request.size,
|
||||
quality=request.quality
|
||||
)
|
||||
if CALLBACK_ENABLED and callback_handler:
|
||||
with CallbackContextManager(
|
||||
handler=callback_handler,
|
||||
user_id=USER_ID,
|
||||
request_id=f"video-image-{uuid.uuid4().hex}"
|
||||
) as ctx:
|
||||
ctx.add_tool("generate_image")
|
||||
result_str = await TOOL_MAP['generate_image'](
|
||||
description=request.description,
|
||||
size=request.size,
|
||||
quality=request.quality
|
||||
)
|
||||
else:
|
||||
result_str = await TOOL_MAP['generate_image'](
|
||||
description=request.description,
|
||||
size=request.size,
|
||||
quality=request.quality
|
||||
)
|
||||
result = json.loads(result_str)
|
||||
|
||||
if not result.get("success"):
|
||||
@@ -267,12 +303,26 @@ async def api_generate_video(request: VideoGenerationRequest, api_key: str = Dep
|
||||
os.environ['OPENAI_API_KEY'] = api_key
|
||||
|
||||
try:
|
||||
result_str = await TOOL_MAP['generate_video'](
|
||||
descriptions=request.descriptions,
|
||||
duration_per_image=request.duration_per_image,
|
||||
fps=request.fps,
|
||||
transition=request.transition
|
||||
)
|
||||
if CALLBACK_ENABLED and callback_handler:
|
||||
with CallbackContextManager(
|
||||
handler=callback_handler,
|
||||
user_id=USER_ID,
|
||||
request_id=f"video-generate-{uuid.uuid4().hex}"
|
||||
) as ctx:
|
||||
ctx.add_tool("generate_video")
|
||||
result_str = await TOOL_MAP['generate_video'](
|
||||
descriptions=request.descriptions,
|
||||
duration_per_image=request.duration_per_image,
|
||||
fps=request.fps,
|
||||
transition=request.transition
|
||||
)
|
||||
else:
|
||||
result_str = await TOOL_MAP['generate_video'](
|
||||
descriptions=request.descriptions,
|
||||
duration_per_image=request.duration_per_image,
|
||||
fps=request.fps,
|
||||
transition=request.transition
|
||||
)
|
||||
result = json.loads(result_str)
|
||||
|
||||
if not result.get("success"):
|
||||
@@ -333,7 +383,16 @@ async def api_download_file(filename: str):
|
||||
async def api_cleanup(max_age_hours: int = 24):
|
||||
"""清理旧文件"""
|
||||
try:
|
||||
result_str = await TOOL_MAP['cleanup_old_files'](max_age_hours=max_age_hours)
|
||||
if CALLBACK_ENABLED and callback_handler:
|
||||
with CallbackContextManager(
|
||||
handler=callback_handler,
|
||||
user_id=USER_ID,
|
||||
request_id=f"video-cleanup-{uuid.uuid4().hex}"
|
||||
) as ctx:
|
||||
ctx.add_tool("cleanup_old_files")
|
||||
result_str = await TOOL_MAP['cleanup_old_files'](max_age_hours=max_age_hours)
|
||||
else:
|
||||
result_str = await TOOL_MAP['cleanup_old_files'](max_age_hours=max_age_hours)
|
||||
result = json.loads(result_str)
|
||||
return result
|
||||
except Exception as e:
|
||||
|
||||
Reference in New Issue
Block a user