Update Heicode sub-mode runtime changes

This commit is contained in:
elipitc
2026-05-31 18:00:17 +08:00
parent f8464fe606
commit ffed09647a
154 changed files with 16388 additions and 444 deletions
BIN
View File
Binary file not shown.
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:
"""验证配置是否完整"""
+320 -234
View File
@@ -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 服务端口
+35 -4
View File
@@ -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():
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}"""
@@ -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",
]
@@ -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"
)
@@ -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)
@@ -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
)
@@ -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
@@ -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])
]
@@ -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
)
@@ -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",
]
@@ -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
@@ -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
@@ -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: