Compare commits
14
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
bf4c53b0b9 | ||
|
|
2e07f40cd4 | ||
|
|
a34d5a081e | ||
|
|
33955b68dc | ||
|
|
ca6a30bdab | ||
|
|
abe3f690ec | ||
|
|
70d9328afc | ||
|
|
02279f9344 | ||
|
|
a115ee68b0 | ||
|
|
6bc7873bc3 | ||
|
|
e67aec11ef | ||
|
|
e3849bd538 | ||
|
|
336f4c2e82 | ||
|
|
749e97cbe2 |
@@ -12,6 +12,7 @@ import re
|
||||
import requests
|
||||
from typing import Dict, List, Optional, Any
|
||||
from datetime import datetime
|
||||
from k8s_manager import sanitize_k8s_name
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -37,10 +38,10 @@ class AgentCodeGenerator:
|
||||
"ACR_LOGIN_SERVER": "agnettaiji.azurecr.io",
|
||||
"ACR_USERNAME": "agnettaiji",
|
||||
"ACR_PASSWORD": "hDpX5t34N5ZmnKdtqyjYL5co/SnXJrmD20CRpGpWaG+ACRCw2wGM",
|
||||
"AZ_CLIENT_ID": "fb306798-2cfe-4ac9-ba48-eab7bc71bcfe",
|
||||
"AZ_CLIENT_SECRET": "cVK8Q~xlfBwm2_t2TC24yrTukWV4F3G~eIjBBa0D",
|
||||
"AZ_CLIENT_ID": "f2dd1cb2-02f6-4efb-bc72-d148f6e01545",
|
||||
"AZ_CLIENT_SECRET": "UVU8Q~Hcrf5KeLi2RvUXB2rcuKFEjRCCrf_JrbwA",
|
||||
"AZ_TENANT_ID": "263c3ff6-1be5-4141-8308-b188464fb297",
|
||||
"AZ_SUBSCRIPTION_ID": "c6c47e4c-f5f4-49f8-b26f-7728862c17d6",
|
||||
"AZ_SUBSCRIPTION_ID": "45d7a360-af09-40fc-9afc-56dc475245ec",
|
||||
"AZ_RG": "taiji-ai-pda",
|
||||
"AZ_AKS": "taiji-ai-pda",
|
||||
"AZURE_DNS_ZONE": "taijiagnet.com"
|
||||
@@ -1646,7 +1647,7 @@ class CallbackContextManager:
|
||||
replicas: 副本数量
|
||||
tool_api_keys: 工具 API 密钥列表(将注入到容器环境变量)
|
||||
"""
|
||||
k8s_name = agent_name.lower().replace("_", "-").replace(" ", "-")
|
||||
k8s_name = sanitize_k8s_name(agent_name)
|
||||
image_repo = f"{self.acr_namespace}/{k8s_name}"
|
||||
|
||||
# 生成工具 API Key 环境变量配置
|
||||
@@ -1914,7 +1915,7 @@ jobs:
|
||||
) -> str:
|
||||
"""生成 README.md"""
|
||||
tools_doc = "\n".join([f"- **{t.get('name')}**: {t.get('description', '')}" for t in tools])
|
||||
k8s_name = agent_name.lower().replace("_", "-").replace(" ", "-")
|
||||
k8s_name = sanitize_k8s_name(agent_name)
|
||||
|
||||
deploy_doc = ""
|
||||
if auto_deploy:
|
||||
@@ -2053,7 +2054,7 @@ MIT License
|
||||
"""
|
||||
files = {}
|
||||
|
||||
k8s_name = agent_name.lower().replace("_", "-").replace(" ", "-")
|
||||
k8s_name = sanitize_k8s_name(agent_name)
|
||||
|
||||
# 生成 src/server/mcp_server.py
|
||||
files["src/server/mcp_server.py"] = self.generate_mcp_server(
|
||||
|
||||
@@ -24,7 +24,7 @@ COPY agents/a2a_litellm_agent/*.py /app/
|
||||
|
||||
# 设置环境变量
|
||||
ENV SERVICE_HOST=0.0.0.0
|
||||
ENV SERVICE_PORT=8080
|
||||
ENV SERVICE_PORT=8000
|
||||
ENV POD_NAME=a2a-litellm-agent
|
||||
ENV TEMPLATE_TYPE=a2a_litellm_agent
|
||||
ENV PYTHONUNBUFFERED=1
|
||||
@@ -34,10 +34,10 @@ ENV AGENT_CALLBACK_URL=http://mcp-server.taiji-ai.svc.cluster.local:8002/api/v1/
|
||||
|
||||
# 健康检查
|
||||
HEALTHCHECK --interval=30s --timeout=10s --start-period=30s --retries=3 \
|
||||
CMD python3 -c "import urllib.request; urllib.request.urlopen('http://localhost:8080/health').read()" || exit 1
|
||||
CMD python3 -c "import urllib.request; urllib.request.urlopen('http://localhost:8000/health').read()" || exit 1
|
||||
|
||||
# 暴露端口
|
||||
EXPOSE 8080
|
||||
EXPOSE 8000
|
||||
|
||||
# 启动命令
|
||||
CMD ["python", "main.py"]
|
||||
|
||||
@@ -26,7 +26,7 @@ logger = structlog.get_logger()
|
||||
|
||||
# 环境变量配置
|
||||
SERVICE_HOST = os.getenv("SERVICE_HOST", "0.0.0.0")
|
||||
SERVICE_PORT = int(os.getenv("SERVICE_PORT", "8080"))
|
||||
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")
|
||||
|
||||
@@ -523,12 +523,12 @@ def create_app(api_key: Optional[str] = None, model: Optional[str] = None) -> Fa
|
||||
创建FastAPI应用(用于uvicorn启动)
|
||||
|
||||
使用方式:
|
||||
uvicorn a2a_server:app --host 0.0.0.0 --port 8080
|
||||
uvicorn a2a_server:app --host 0.0.0.0 --port 8000
|
||||
|
||||
或设置环境变量后:
|
||||
export LITELLM_API_KEY="your-key"
|
||||
export MODEL_NAME="your-model"
|
||||
uvicorn a2a_server:app --host 0.0.0.0 --port 8080
|
||||
uvicorn a2a_server:app --host 0.0.0.0 --port 8000
|
||||
"""
|
||||
server = A2AAgentServer(api_key=api_key, model=model)
|
||||
return server.app
|
||||
|
||||
@@ -70,7 +70,7 @@ class AgentConfig:
|
||||
version: str = "1.0.0"
|
||||
|
||||
# 服务端口
|
||||
port: int = 8080
|
||||
port: int = 8000
|
||||
|
||||
# 服务主机
|
||||
host: str = "0.0.0.0"
|
||||
|
||||
@@ -8,7 +8,7 @@ from a2a_server import create_app
|
||||
|
||||
# 环境变量配置
|
||||
SERVICE_HOST = os.getenv("SERVICE_HOST", "0.0.0.0")
|
||||
SERVICE_PORT = int(os.getenv("SERVICE_PORT", "8080"))
|
||||
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")
|
||||
|
||||
|
||||
@@ -0,0 +1,466 @@
|
||||
# Ad Creator Agent - API 文档
|
||||
|
||||
多模态广告创意生成 Agent,通过素材(文字描述/参考图片)生成广告图片或视频。
|
||||
|
||||
## 基本信息
|
||||
|
||||
| 项目 | 值 |
|
||||
|------|------|
|
||||
| 镜像 | `agnettaiji.azurecr.io/ai-agents/ad-creator-agent:latest` |
|
||||
| 端口 | `8000` |
|
||||
| 模板名 | `ad_creator_agent` |
|
||||
| 框架 | API (FastAPI) |
|
||||
|
||||
## 支持的模型
|
||||
|
||||
| 用途 | 模型 | 备注 |
|
||||
|------|------|------|
|
||||
| 图片生成(默认) | `taiji/gemini-3-pro-image-preview` | 支持参考图片输入 |
|
||||
| 图片生成 | `taiji/gpt-image-1` | OpenAI GPT Image |
|
||||
| 图片生成 | `taiji/gpt-image-1-mini` | 轻量版,速度更快 |
|
||||
| 图片生成 | `taiji/dall-e-3` | DALL-E 3 |
|
||||
| 文案生成 | `taiji/gpt-4o-mini` | 广告文案 + 图片 prompt |
|
||||
| 视频生成 | `taiji/sora-2` | Sora 视频生成 |
|
||||
|
||||
## 认证方式
|
||||
|
||||
所有写操作端点均需传入 API Key,支持以下两种方式:
|
||||
|
||||
```
|
||||
api-key: sk-xxx
|
||||
```
|
||||
|
||||
```
|
||||
Authorization: Bearer sk-xxx
|
||||
```
|
||||
|
||||
如果部署时配置了 `LLM_API_KEY` 环境变量,可省略请求头中的 Key。
|
||||
|
||||
## 环境变量
|
||||
|
||||
| 变量名 | 说明 | 默认值 |
|
||||
|--------|------|--------|
|
||||
| `LLM_API_KEY` | LiteLLM API Key | (必填或请求头传入) |
|
||||
| `LLM_BASE_URL` | LiteLLM Base URL | `https://litellm.graystone-fb459c5d.southeastasia.azurecontainerapps.io/v1` |
|
||||
| `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` |
|
||||
|
||||
---
|
||||
|
||||
## API 端点
|
||||
|
||||
### 1. 健康检查
|
||||
|
||||
**GET** `/health`
|
||||
|
||||
```bash
|
||||
curl http://<AGENT_URL>/health
|
||||
```
|
||||
|
||||
**响应示例:**
|
||||
|
||||
```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"
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### 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"
|
||||
}'
|
||||
```
|
||||
|
||||
**示例 - 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"
|
||||
}'
|
||||
```
|
||||
|
||||
**响应示例:**
|
||||
|
||||
```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",
|
||||
"model": "taiji/gemini-3-pro-image-preview"
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### 3. 上传参考图片并生成广告图
|
||||
|
||||
**POST** `/api/v1/generate-image-upload`
|
||||
|
||||
支持 `multipart/form-data` 上传参考图片,结合文字描述生成广告图。
|
||||
|
||||
**表单字段:**
|
||||
|
||||
| 字段 | 类型 | 必填 | 说明 |
|
||||
|------|------|------|------|
|
||||
| `prompt` | string | 是 | 广告图片描述 |
|
||||
| `reference_image` | file | 否 | 参考图片文件 |
|
||||
| `model` | string | 否 | 模型名称 |
|
||||
| `aspect_ratio` | string | 否 | 宽高比 |
|
||||
| `quality` | string | 否 | 质量 |
|
||||
| `style` | string | 否 | 广告风格 |
|
||||
| `brand_name` | string | 否 | 品牌名称 |
|
||||
|
||||
**示例:**
|
||||
|
||||
```bash
|
||||
curl -X POST http://<AGENT_URL>/api/v1/generate-image-upload \
|
||||
-H "api-key: sk-xxx" \
|
||||
-F "prompt=基于这张产品图,生成一张高端产品广告海报" \
|
||||
-F "reference_image=@product_photo.jpg" \
|
||||
-F "style=luxury" \
|
||||
-F "aspect_ratio=16:9"
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### 4. 生成广告文案
|
||||
|
||||
**POST** `/api/v1/generate-copy`
|
||||
|
||||
根据产品信息,由 LLM 生成结构化广告文案(标题、正文、CTA、hashtags)以及用于图片生成的英文 prompt。
|
||||
|
||||
**请求体:**
|
||||
|
||||
| 字段 | 类型 | 必填 | 说明 |
|
||||
|------|------|------|------|
|
||||
| `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"
|
||||
}'
|
||||
```
|
||||
|
||||
**响应示例:**
|
||||
|
||||
```json
|
||||
{
|
||||
"success": true,
|
||||
"headline": "沉浸高端音质",
|
||||
"body_copy": "体验非凡音质,尽享音乐带来的宁静与专注...",
|
||||
"cta": "立即体验",
|
||||
"image_prompt": "A luxurious setting featuring a sleek wireless headphone...",
|
||||
"hashtags": ["#高端耳机", "#沉浸音乐", "#商务生活"]
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### 5. 一键生成完整广告(文案 + 图片)
|
||||
|
||||
**POST** `/api/v1/generate-ad`
|
||||
|
||||
自动生成广告文案,并基于文案中的图片 prompt 自动生成配图。
|
||||
|
||||
**请求体:**
|
||||
|
||||
| 字段 | 类型 | 必填 | 说明 |
|
||||
|------|------|------|------|
|
||||
| `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"
|
||||
}'
|
||||
```
|
||||
|
||||
**响应示例:**
|
||||
|
||||
```json
|
||||
{
|
||||
"success": true,
|
||||
"copy": {
|
||||
"success": true,
|
||||
"headline": "开启绿色出行新生活",
|
||||
"body_copy": "选择我们的新能源电动汽车...",
|
||||
"cta": "立即了解更多",
|
||||
"image_prompt": "A futuristic electric vehicle...",
|
||||
"hashtags": ["#新能源车", "#绿色出行", "#智能驾驶"]
|
||||
},
|
||||
"image": {
|
||||
"success": true,
|
||||
"filename": "ad_gemini_20260302_145504_262223.jpg",
|
||||
"url": "/api/v1/files/ad_gemini_20260302_145504_262223.jpg",
|
||||
"model": "taiji/gemini-3-pro-image-preview"
|
||||
},
|
||||
"timestamp": "2026-03-02T14:55:04.262223"
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### 6. 生成广告视频
|
||||
|
||||
**POST** `/api/v1/generate-video`
|
||||
|
||||
使用 Sora 模型生成广告短视频。
|
||||
|
||||
**请求体:**
|
||||
|
||||
| 字段 | 类型 | 必填 | 说明 |
|
||||
|------|------|------|------|
|
||||
| `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"
|
||||
}'
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### 7. 智能对话
|
||||
|
||||
**POST** `/chat`
|
||||
|
||||
与 AI 广告创意总监对话。系统会理解需求,自动决定是否生成图片。
|
||||
|
||||
**请求体:**
|
||||
|
||||
| 字段 | 类型 | 必填 | 说明 |
|
||||
|------|------|------|------|
|
||||
| `message` | string | 是 | 用户消息 |
|
||||
|
||||
**示例:**
|
||||
|
||||
```bash
|
||||
curl -X POST http://<AGENT_URL>/chat \
|
||||
-H "Content-Type: application/json" \
|
||||
-H "api-key: sk-xxx" \
|
||||
-d '{
|
||||
"message": "帮我为一款蓝牙音箱做一个抖音封面图,要有科技感"
|
||||
}'
|
||||
```
|
||||
|
||||
**响应示例:**
|
||||
|
||||
```json
|
||||
{
|
||||
"response": "为这款蓝牙音箱设计封面图的建议...",
|
||||
"image": {
|
||||
"success": true,
|
||||
"filename": "ad_gemini_20260302_145539_866923.jpg",
|
||||
"url": "/api/v1/files/ad_gemini_20260302_145539_866923.jpg",
|
||||
"model": "taiji/gemini-3-pro-image-preview"
|
||||
},
|
||||
"timestamp": "2026-03-02T14:55:39.866923"
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### 8. 下载生成的文件
|
||||
|
||||
**GET** `/api/v1/files/{filename}`
|
||||
|
||||
```bash
|
||||
curl -O http://<AGENT_URL>/api/v1/files/ad_gemini_20260302_145310_209307.jpg
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### 9. 列出已生成的文件
|
||||
|
||||
**GET** `/api/v1/list-files?file_type=all`
|
||||
|
||||
参数 `file_type` 可选值: `all`, `image`, `video`
|
||||
|
||||
```bash
|
||||
curl http://<AGENT_URL>/api/v1/list-files
|
||||
```
|
||||
|
||||
**响应示例:**
|
||||
|
||||
```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"
|
||||
}
|
||||
],
|
||||
"videos": []
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### 10. 清理旧文件
|
||||
|
||||
**POST** `/api/v1/cleanup?max_age_hours=24`
|
||||
|
||||
删除超过指定时间的旧文件。
|
||||
|
||||
```bash
|
||||
curl -X POST "http://<AGENT_URL>/api/v1/cleanup?max_age_hours=24"
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### 11. 状态查看
|
||||
|
||||
**GET** `/status`
|
||||
|
||||
```bash
|
||||
curl http://<AGENT_URL>/status
|
||||
```
|
||||
|
||||
**响应示例:**
|
||||
|
||||
```json
|
||||
{
|
||||
"status": "running",
|
||||
"pod_name": "test-ad-creator",
|
||||
"generated_images": 4,
|
||||
"generated_videos": 0,
|
||||
"timestamp": "2026-03-02T15:01:43.636444"
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 通过 Agent Manager 部署
|
||||
|
||||
### 1. 注册模板
|
||||
|
||||
```bash
|
||||
curl -X POST http://20.212.121.126/templates/create \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"name": "ad_creator_agent",
|
||||
"display_name": "Ad Creator Agent",
|
||||
"description": "多模态广告创意生成 Agent",
|
||||
"image": "agnettaiji.azurecr.io/ai-agents/ad-creator-agent:latest",
|
||||
"port": 8000,
|
||||
"agent_type": "platform",
|
||||
"agent_framework": "api",
|
||||
"env_requirements": {
|
||||
"LLM_API_KEY": "LiteLLM API Key"
|
||||
}
|
||||
}'
|
||||
```
|
||||
|
||||
### 2. 创建实例
|
||||
|
||||
```bash
|
||||
curl -X POST http://20.212.121.126/agents \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"name": "my-ad-creator",
|
||||
"template": "ad_creator_agent",
|
||||
"config": { "user_id": "your-user-id" },
|
||||
"env": {
|
||||
"LLM_API_KEY": "sk-your-api-key"
|
||||
}
|
||||
}'
|
||||
```
|
||||
|
||||
### 3. 删除实例
|
||||
|
||||
```bash
|
||||
curl -X DELETE http://20.212.121.126/agents/my-ad-creator
|
||||
```
|
||||
@@ -0,0 +1,35 @@
|
||||
FROM python:3.11-slim
|
||||
|
||||
WORKDIR /app
|
||||
|
||||
RUN apt-get update && apt-get install -y \
|
||||
curl \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
|
||||
RUN pip install --no-cache-dir \
|
||||
fastapi==0.109.0 \
|
||||
uvicorn[standard]==0.27.0 \
|
||||
pydantic==2.5.3 \
|
||||
aiohttp>=3.9.0 \
|
||||
python-multipart>=0.0.6
|
||||
|
||||
COPY common/agent_callback_utils.py /app/common/
|
||||
RUN touch /app/common/__init__.py
|
||||
|
||||
COPY agents/ad_creator_agent/ad_creator_agent.py /app/
|
||||
|
||||
ENV PYTHONUNBUFFERED=1
|
||||
ENV SERVICE_HOST=0.0.0.0
|
||||
ENV SERVICE_PORT=8000
|
||||
ENV OUTPUT_DIR=/app/outputs
|
||||
|
||||
ENV AGENT_CALLBACK_URL=http://mcp-server:8002/api/v1/billing/agent-callback
|
||||
|
||||
RUN mkdir -p /app/outputs/images /app/outputs/videos
|
||||
|
||||
HEALTHCHECK --interval=30s --timeout=10s --start-period=20s --retries=3 \
|
||||
CMD python3 -c "import urllib.request; urllib.request.urlopen('http://localhost:8000/health').read()" || exit 1
|
||||
|
||||
EXPOSE 8000
|
||||
|
||||
CMD ["python3", "-u", "ad_creator_agent.py"]
|
||||
@@ -0,0 +1,793 @@
|
||||
"""
|
||||
Ad Creator Agent - 多模态广告创意生成 Agent
|
||||
通过素材(文字描述/参考图片)生成广告图片或视频
|
||||
支持模型:Gemini 3 Pro Image / GPT Image 1 / DALL-E 3 / Sora 2
|
||||
"""
|
||||
import os
|
||||
import sys
|
||||
import re
|
||||
import uuid
|
||||
import json
|
||||
import base64
|
||||
import logging
|
||||
import aiohttp
|
||||
from typing import Optional, List, Dict, Any
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
from enum import Enum
|
||||
|
||||
from fastapi import FastAPI, HTTPException, Header, Depends, UploadFile, File, Form, Request
|
||||
from fastapi.middleware.cors import CORSMiddleware
|
||||
from fastapi.responses import FileResponse, JSONResponse
|
||||
from pydantic import BaseModel, Field
|
||||
import uvicorn
|
||||
|
||||
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
|
||||
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, format='%(asctime)s - %(name)s - %(levelname)s - %(message)s')
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# ==================== 环境变量 ====================
|
||||
|
||||
SERVICE_HOST = os.getenv("SERVICE_HOST", "0.0.0.0")
|
||||
SERVICE_PORT = int(os.getenv("SERVICE_PORT", "8000"))
|
||||
POD_NAME = os.getenv("POD_NAME", "ad-creator-agent")
|
||||
USER_ID = os.getenv("USER_ID", "")
|
||||
OUTPUT_DIR = os.getenv("OUTPUT_DIR", "/app/outputs")
|
||||
|
||||
LLM_BASE_URL = os.getenv(
|
||||
"LLM_BASE_URL",
|
||||
"https://litellm.graystone-fb459c5d.southeastasia.azurecontainerapps.io/v1"
|
||||
)
|
||||
LLM_API_KEY = os.getenv("LLM_API_KEY", "")
|
||||
|
||||
DEFAULT_IMAGE_MODEL = os.getenv("DEFAULT_IMAGE_MODEL", "taiji/gemini-3-pro-image-preview")
|
||||
DEFAULT_TEXT_MODEL = os.getenv("DEFAULT_TEXT_MODEL", "taiji/gpt-4o-mini")
|
||||
DEFAULT_VIDEO_MODEL = os.getenv("DEFAULT_VIDEO_MODEL", "taiji/sora-2")
|
||||
|
||||
Path(OUTPUT_DIR).mkdir(parents=True, exist_ok=True)
|
||||
Path(f"{OUTPUT_DIR}/images").mkdir(parents=True, exist_ok=True)
|
||||
Path(f"{OUTPUT_DIR}/videos").mkdir(parents=True, exist_ok=True)
|
||||
|
||||
|
||||
# ==================== 模型枚举 ====================
|
||||
|
||||
class ImageModel(str, Enum):
|
||||
GEMINI = "taiji/gemini-3-pro-image-preview"
|
||||
GPT_IMAGE = "taiji/gpt-image-1"
|
||||
GPT_IMAGE_MINI = "taiji/gpt-image-1-mini"
|
||||
DALLE3 = "taiji/dall-e-3"
|
||||
|
||||
|
||||
class AspectRatio(str, Enum):
|
||||
SQUARE = "1:1"
|
||||
LANDSCAPE = "16:9"
|
||||
PORTRAIT = "9:16"
|
||||
WIDE = "4:3"
|
||||
TALL = "3:4"
|
||||
|
||||
|
||||
# ==================== 请求/响应模型 ====================
|
||||
|
||||
class GenerateImageRequest(BaseModel):
|
||||
prompt: str = Field(..., description="广告图片描述/创意需求")
|
||||
model: Optional[str] = Field(None, description="图片生成模型,默认 gemini-3-pro-image-preview")
|
||||
aspect_ratio: Optional[str] = Field("1:1", description="宽高比: 1:1, 16:9, 9:16, 4:3, 3:4")
|
||||
size: Optional[str] = Field("1024x1024", description="图片尺寸(仅 GPT/DALL-E)")
|
||||
quality: Optional[str] = Field("high", description="质量: low, medium, high")
|
||||
style: Optional[str] = Field(None, description="广告风格: modern, minimalist, luxury, playful, tech, vintage")
|
||||
brand_name: Optional[str] = Field(None, description="品牌名称")
|
||||
reference_image_b64: Optional[str] = Field(None, description="参考图片 base64 (用于风格参考或产品素材)")
|
||||
user_id: Optional[str] = Field(None, description="用户ID")
|
||||
|
||||
|
||||
class GenerateAdCopyRequest(BaseModel):
|
||||
product: str = Field(..., description="产品/服务描述")
|
||||
target_audience: Optional[str] = Field(None, description="目标受众")
|
||||
tone: Optional[str] = Field("professional", description="语气: professional, casual, humorous, urgent, luxury")
|
||||
platform: Optional[str] = Field("general", description="投放平台: instagram, facebook, tiktok, billboard, general")
|
||||
language: Optional[str] = Field("zh", description="语言: zh, en, ja")
|
||||
user_id: Optional[str] = Field(None, description="用户ID")
|
||||
|
||||
|
||||
class GenerateAdRequest(BaseModel):
|
||||
"""完整广告生成(文案 + 图片)"""
|
||||
product: str = Field(..., description="产品/服务描述")
|
||||
image_model: Optional[str] = Field(None, description="图片生成模型")
|
||||
aspect_ratio: Optional[str] = Field("1:1", description="宽高比")
|
||||
style: Optional[str] = Field(None, description="广告风格")
|
||||
brand_name: Optional[str] = Field(None, description="品牌名称")
|
||||
target_audience: Optional[str] = Field(None, description="目标受众")
|
||||
tone: Optional[str] = Field("professional", description="语气")
|
||||
platform: Optional[str] = Field("general", description="投放平台")
|
||||
language: Optional[str] = Field("zh", description="语言")
|
||||
reference_image_b64: Optional[str] = Field(None, description="参考图片 base64")
|
||||
user_id: Optional[str] = Field(None, description="用户ID")
|
||||
|
||||
|
||||
class GenerateVideoRequest(BaseModel):
|
||||
prompt: str = Field(..., description="视频描述/创意需求")
|
||||
model: Optional[str] = Field(None, description="视频模型, 默认 sora-2")
|
||||
aspect_ratio: Optional[str] = Field("16:9", description="宽高比")
|
||||
duration: Optional[str] = Field("5", description="视频时长秒数")
|
||||
user_id: Optional[str] = Field(None, description="用户ID")
|
||||
|
||||
|
||||
class ChatRequest(BaseModel):
|
||||
message: str = Field(..., description="用户消息")
|
||||
user_id: Optional[str] = Field(None, description="用户ID")
|
||||
|
||||
|
||||
# ==================== FastAPI ====================
|
||||
|
||||
app = FastAPI(
|
||||
title="Ad Creator Agent",
|
||||
description="多模态广告创意生成 - 通过素材生成广告图片或视频(支持 Gemini / GPT Image / DALL-E / Sora)",
|
||||
version="1.0.0"
|
||||
)
|
||||
|
||||
app.add_middleware(
|
||||
CORSMiddleware,
|
||||
allow_origins=["*"],
|
||||
allow_credentials=True,
|
||||
allow_methods=["*"],
|
||||
allow_headers=["*"],
|
||||
)
|
||||
|
||||
callback_handler: Optional[Any] = None
|
||||
|
||||
|
||||
# ==================== 依赖 ====================
|
||||
|
||||
async def get_api_key(
|
||||
api_key: Optional[str] = Header(None, alias="api-key"),
|
||||
authorization: Optional[str] = Header(None),
|
||||
) -> str:
|
||||
if api_key and api_key.strip():
|
||||
return api_key.strip()
|
||||
if authorization:
|
||||
key = authorization[7:].strip() if authorization.startswith("Bearer ") else authorization.strip()
|
||||
if key:
|
||||
return key
|
||||
if LLM_API_KEY:
|
||||
return LLM_API_KEY
|
||||
raise HTTPException(status_code=401, detail="请在请求头提供 api-key 或 Authorization Bearer token")
|
||||
|
||||
|
||||
# ==================== 生命周期 ====================
|
||||
|
||||
@app.on_event("startup")
|
||||
async def startup_event():
|
||||
global callback_handler
|
||||
if CALLBACK_ENABLED:
|
||||
callback_handler = AgentCallbackHandler(agent_name=POD_NAME, user_id=USER_ID)
|
||||
logger.info(f"回调处理器已初始化: {callback_handler.callback_url}")
|
||||
else:
|
||||
logger.warning("回调模块未加载")
|
||||
logger.info(f"Ad Creator Agent 启动: port={SERVICE_PORT}, output={OUTPUT_DIR}")
|
||||
logger.info(f"默认模型: image={DEFAULT_IMAGE_MODEL}, text={DEFAULT_TEXT_MODEL}, video={DEFAULT_VIDEO_MODEL}")
|
||||
|
||||
|
||||
# ==================== 核心:图片生成 ====================
|
||||
|
||||
async def generate_image_gemini(
|
||||
prompt: str,
|
||||
api_key: str,
|
||||
aspect_ratio: str = "1:1",
|
||||
quality: str = "high",
|
||||
reference_image_b64: str = None,
|
||||
) -> dict:
|
||||
"""通过 Gemini chat/completions 生成图片,返回 {success, file_path, filename, url}"""
|
||||
|
||||
quality_map = {"low": "1K", "medium": "1K", "high": "2K"}
|
||||
image_size = quality_map.get(quality, "2K")
|
||||
|
||||
messages_content: Any
|
||||
if reference_image_b64:
|
||||
messages_content = [
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {"url": f"data:image/png;base64,{reference_image_b64}"},
|
||||
},
|
||||
{"type": "text", "text": prompt},
|
||||
]
|
||||
else:
|
||||
messages_content = prompt
|
||||
|
||||
payload = {
|
||||
"model": DEFAULT_IMAGE_MODEL,
|
||||
"stream": False,
|
||||
"messages": [{"role": "user", "content": messages_content}],
|
||||
"extra_body": {
|
||||
"google": {
|
||||
"image_config": {
|
||||
"aspect_ratio": aspect_ratio,
|
||||
"image_size": image_size,
|
||||
}
|
||||
}
|
||||
},
|
||||
}
|
||||
|
||||
headers = {"Authorization": f"Bearer {api_key}", "Content-Type": "application/json"}
|
||||
base = LLM_BASE_URL.rstrip("/")
|
||||
url = f"{base}/chat/completions"
|
||||
|
||||
async with aiohttp.ClientSession() as session:
|
||||
async with session.post(url, json=payload, headers=headers, timeout=aiohttp.ClientTimeout(total=120)) as resp:
|
||||
if resp.status != 200:
|
||||
text = await resp.text()
|
||||
return {"success": False, "error": f"Gemini API error {resp.status}: {text[:500]}"}
|
||||
|
||||
data = await resp.json()
|
||||
content = data.get("choices", [{}])[0].get("message", {}).get("content", "")
|
||||
|
||||
match = re.search(r"data:image/([\w+]+);base64,([A-Za-z0-9+/=\s]+)", content)
|
||||
if not match:
|
||||
return {"success": False, "error": "Gemini 未返回图片数据", "text_response": content[:500]}
|
||||
|
||||
img_format = match.group(1).replace("+xml", "")
|
||||
ext = "jpg" if img_format == "jpeg" else img_format
|
||||
b64_data = match.group(2).replace("\n", "").replace(" ", "")
|
||||
|
||||
ts = datetime.now().strftime("%Y%m%d_%H%M%S_%f")
|
||||
filename = f"ad_gemini_{ts}.{ext}"
|
||||
file_path = os.path.join(OUTPUT_DIR, "images", filename)
|
||||
|
||||
with open(file_path, "wb") as f:
|
||||
f.write(base64.b64decode(b64_data))
|
||||
|
||||
logger.info(f"Gemini 图片已生成: {file_path} ({os.path.getsize(file_path)} bytes)")
|
||||
return {
|
||||
"success": True,
|
||||
"file_path": file_path,
|
||||
"filename": filename,
|
||||
"url": f"/api/v1/files/{filename}",
|
||||
"model": DEFAULT_IMAGE_MODEL,
|
||||
}
|
||||
|
||||
|
||||
async def generate_image_openai(
|
||||
prompt: str,
|
||||
api_key: str,
|
||||
model: str = "taiji/gpt-image-1",
|
||||
size: str = "1024x1024",
|
||||
quality: str = "high",
|
||||
) -> dict:
|
||||
"""通过 OpenAI images/generations 接口生成图片(GPT Image / DALL-E)"""
|
||||
|
||||
quality_map = {"low": "low", "medium": "medium", "high": "high"}
|
||||
payload = {
|
||||
"model": model,
|
||||
"prompt": prompt,
|
||||
"n": 1,
|
||||
"size": size,
|
||||
"quality": quality_map.get(quality, "high"),
|
||||
}
|
||||
|
||||
headers = {"Authorization": f"Bearer {api_key}", "Content-Type": "application/json"}
|
||||
base = LLM_BASE_URL.rstrip("/")
|
||||
url = f"{base}/images/generations"
|
||||
|
||||
async with aiohttp.ClientSession() as session:
|
||||
async with session.post(url, json=payload, headers=headers, timeout=aiohttp.ClientTimeout(total=120)) as resp:
|
||||
if resp.status != 200:
|
||||
text = await resp.text()
|
||||
return {"success": False, "error": f"OpenAI API error {resp.status}: {text[:500]}"}
|
||||
|
||||
data = await resp.json()
|
||||
items = data.get("data", [])
|
||||
if not items:
|
||||
return {"success": False, "error": "API 未返回图片数据"}
|
||||
|
||||
item = items[0]
|
||||
b64_data = item.get("b64_json")
|
||||
image_url = item.get("url")
|
||||
|
||||
ts = datetime.now().strftime("%Y%m%d_%H%M%S_%f")
|
||||
model_tag = model.split("/")[-1].replace("-", "")
|
||||
filename = f"ad_{model_tag}_{ts}.png"
|
||||
file_path = os.path.join(OUTPUT_DIR, "images", filename)
|
||||
|
||||
if b64_data:
|
||||
with open(file_path, "wb") as f:
|
||||
f.write(base64.b64decode(b64_data))
|
||||
elif image_url:
|
||||
async with session.get(image_url, timeout=aiohttp.ClientTimeout(total=30)) as dl_resp:
|
||||
if dl_resp.status != 200:
|
||||
return {"success": False, "error": f"下载图片失败: HTTP {dl_resp.status}"}
|
||||
with open(file_path, "wb") as f:
|
||||
f.write(await dl_resp.read())
|
||||
else:
|
||||
return {"success": False, "error": "API 响应中无图片数据"}
|
||||
|
||||
logger.info(f"OpenAI 图片已生成: {file_path} ({os.path.getsize(file_path)} bytes)")
|
||||
return {
|
||||
"success": True,
|
||||
"file_path": file_path,
|
||||
"filename": filename,
|
||||
"url": f"/api/v1/files/{filename}",
|
||||
"model": model,
|
||||
}
|
||||
|
||||
|
||||
async def generate_image_dispatch(
|
||||
prompt: str,
|
||||
api_key: str,
|
||||
model: str = None,
|
||||
aspect_ratio: str = "1:1",
|
||||
size: str = "1024x1024",
|
||||
quality: str = "high",
|
||||
reference_image_b64: str = None,
|
||||
) -> dict:
|
||||
"""根据模型分发到对应生成函数"""
|
||||
model = model or DEFAULT_IMAGE_MODEL
|
||||
|
||||
if "gemini" in model.lower():
|
||||
return await generate_image_gemini(
|
||||
prompt=prompt,
|
||||
api_key=api_key,
|
||||
aspect_ratio=aspect_ratio,
|
||||
quality=quality,
|
||||
reference_image_b64=reference_image_b64,
|
||||
)
|
||||
else:
|
||||
return await generate_image_openai(
|
||||
prompt=prompt,
|
||||
api_key=api_key,
|
||||
model=model,
|
||||
size=size,
|
||||
quality=quality,
|
||||
)
|
||||
|
||||
|
||||
# ==================== 核心:文案生成 ====================
|
||||
|
||||
async def generate_ad_copy(
|
||||
product: str,
|
||||
api_key: str,
|
||||
target_audience: str = None,
|
||||
tone: str = "professional",
|
||||
platform: str = "general",
|
||||
language: str = "zh",
|
||||
) -> dict:
|
||||
"""用 LLM 生成广告文案和图片 prompt"""
|
||||
|
||||
lang_map = {"zh": "中文", "en": "English", "ja": "日本語"}
|
||||
lang_label = lang_map.get(language, language)
|
||||
|
||||
system_prompt = f"""你是一位顶尖的广告创意总监。请根据产品信息生成广告创意方案。
|
||||
回复必须使用 {lang_label}。
|
||||
回复格式必须严格为以下 JSON(不包含 markdown 代码块标记):
|
||||
{{
|
||||
"headline": "广告标题(10字以内)",
|
||||
"body_copy": "广告正文(30-60字)",
|
||||
"cta": "行动号召按钮文字",
|
||||
"image_prompt": "用于 AI 生成广告配图的英文详细描述(100-200 words, 描述画面构图、色调、元素)",
|
||||
"hashtags": ["标签1", "标签2", "标签3"]
|
||||
}}"""
|
||||
|
||||
user_msg = f"产品/服务: {product}"
|
||||
if target_audience:
|
||||
user_msg += f"\n目标受众: {target_audience}"
|
||||
user_msg += f"\n语气风格: {tone}\n投放平台: {platform}"
|
||||
|
||||
payload = {
|
||||
"model": DEFAULT_TEXT_MODEL,
|
||||
"messages": [
|
||||
{"role": "system", "content": system_prompt},
|
||||
{"role": "user", "content": user_msg},
|
||||
],
|
||||
"max_tokens": 800,
|
||||
"temperature": 0.8,
|
||||
}
|
||||
|
||||
headers = {"Authorization": f"Bearer {api_key}", "Content-Type": "application/json"}
|
||||
base = LLM_BASE_URL.rstrip("/")
|
||||
url = f"{base}/chat/completions"
|
||||
|
||||
async with aiohttp.ClientSession() as session:
|
||||
async with session.post(url, json=payload, headers=headers, timeout=aiohttp.ClientTimeout(total=30)) as resp:
|
||||
if resp.status != 200:
|
||||
text = await resp.text()
|
||||
return {"success": False, "error": f"LLM API error {resp.status}: {text[:300]}"}
|
||||
data = await resp.json()
|
||||
content = data.get("choices", [{}])[0].get("message", {}).get("content", "")
|
||||
|
||||
try:
|
||||
cleaned = content.strip()
|
||||
if cleaned.startswith("```"):
|
||||
cleaned = re.sub(r"^```\w*\n?", "", cleaned)
|
||||
cleaned = re.sub(r"\n?```$", "", cleaned)
|
||||
copy_data = json.loads(cleaned)
|
||||
except json.JSONDecodeError:
|
||||
return {
|
||||
"success": True,
|
||||
"raw_text": content,
|
||||
"headline": "",
|
||||
"body_copy": content[:200],
|
||||
"cta": "",
|
||||
"image_prompt": "",
|
||||
"hashtags": [],
|
||||
}
|
||||
|
||||
return {
|
||||
"success": True,
|
||||
"headline": copy_data.get("headline", ""),
|
||||
"body_copy": copy_data.get("body_copy", ""),
|
||||
"cta": copy_data.get("cta", ""),
|
||||
"image_prompt": copy_data.get("image_prompt", ""),
|
||||
"hashtags": copy_data.get("hashtags", []),
|
||||
}
|
||||
|
||||
|
||||
# ==================== 核心:视频生成 ====================
|
||||
|
||||
async def generate_video_sora(
|
||||
prompt: str,
|
||||
api_key: str,
|
||||
model: str = None,
|
||||
aspect_ratio: str = "16:9",
|
||||
duration: str = "5",
|
||||
) -> dict:
|
||||
"""通过 Sora 生成视频(OpenAI 兼容接口)"""
|
||||
model = model or DEFAULT_VIDEO_MODEL
|
||||
|
||||
size_map = {
|
||||
"16:9": "1920x1080",
|
||||
"9:16": "1080x1920",
|
||||
"1:1": "1080x1080",
|
||||
}
|
||||
size = size_map.get(aspect_ratio, "1920x1080")
|
||||
|
||||
payload = {
|
||||
"model": model,
|
||||
"input": prompt,
|
||||
"size": size,
|
||||
"duration": int(duration),
|
||||
"n": 1,
|
||||
}
|
||||
|
||||
headers = {"Authorization": f"Bearer {api_key}", "Content-Type": "application/json"}
|
||||
base = LLM_BASE_URL.rstrip("/")
|
||||
url = f"{base}/videos/generations"
|
||||
|
||||
async with aiohttp.ClientSession() as session:
|
||||
async with session.post(url, json=payload, headers=headers, timeout=aiohttp.ClientTimeout(total=300)) as resp:
|
||||
if resp.status != 200:
|
||||
text = await resp.text()
|
||||
return {"success": False, "error": f"Sora API error {resp.status}: {text[:500]}"}
|
||||
|
||||
data = await resp.json()
|
||||
items = data.get("data", [])
|
||||
if not items:
|
||||
return {"success": False, "error": "Sora 未返回视频数据"}
|
||||
|
||||
video_url = items[0].get("url")
|
||||
b64_data = items[0].get("b64_json")
|
||||
|
||||
ts = datetime.now().strftime("%Y%m%d_%H%M%S_%f")
|
||||
filename = f"ad_video_{ts}.mp4"
|
||||
file_path = os.path.join(OUTPUT_DIR, "videos", filename)
|
||||
|
||||
if b64_data:
|
||||
with open(file_path, "wb") as f:
|
||||
f.write(base64.b64decode(b64_data))
|
||||
elif video_url:
|
||||
async with session.get(video_url, timeout=aiohttp.ClientTimeout(total=120)) as dl_resp:
|
||||
if dl_resp.status != 200:
|
||||
return {"success": False, "error": f"下载视频失败: HTTP {dl_resp.status}"}
|
||||
with open(file_path, "wb") as f:
|
||||
f.write(await dl_resp.read())
|
||||
else:
|
||||
return {"success": False, "error": "Sora 响应中无视频数据"}
|
||||
|
||||
logger.info(f"视频已生成: {file_path} ({os.path.getsize(file_path)} bytes)")
|
||||
return {
|
||||
"success": True,
|
||||
"file_path": file_path,
|
||||
"filename": filename,
|
||||
"url": f"/api/v1/files/{filename}",
|
||||
"model": model,
|
||||
}
|
||||
|
||||
|
||||
# ==================== API 端点 ====================
|
||||
|
||||
@app.get("/")
|
||||
@app.get("/health")
|
||||
async def health_check():
|
||||
return {
|
||||
"status": "healthy",
|
||||
"service": "Ad Creator Agent",
|
||||
"pod_name": POD_NAME,
|
||||
"models": {
|
||||
"image": DEFAULT_IMAGE_MODEL,
|
||||
"text": DEFAULT_TEXT_MODEL,
|
||||
"video": DEFAULT_VIDEO_MODEL,
|
||||
},
|
||||
"callback_enabled": CALLBACK_ENABLED,
|
||||
"timestamp": datetime.utcnow().isoformat(),
|
||||
}
|
||||
|
||||
|
||||
@app.get("/status")
|
||||
async def status():
|
||||
images = list(Path(f"{OUTPUT_DIR}/images").glob("*"))
|
||||
videos = list(Path(f"{OUTPUT_DIR}/videos").glob("*"))
|
||||
return {
|
||||
"status": "running",
|
||||
"pod_name": POD_NAME,
|
||||
"generated_images": len(images),
|
||||
"generated_videos": len(videos),
|
||||
"timestamp": datetime.utcnow().isoformat(),
|
||||
}
|
||||
|
||||
|
||||
# ---------- 图片生成 ----------
|
||||
|
||||
@app.post("/api/v1/generate-image")
|
||||
async def api_generate_image(request: GenerateImageRequest, api_key: str = Depends(get_api_key)):
|
||||
"""生成广告图片"""
|
||||
prompt = request.prompt
|
||||
if request.style:
|
||||
prompt = f"[{request.style} style] {prompt}"
|
||||
if request.brand_name:
|
||||
prompt = f"{prompt}. Brand: {request.brand_name}"
|
||||
|
||||
result = await generate_image_dispatch(
|
||||
prompt=prompt,
|
||||
api_key=api_key,
|
||||
model=request.model,
|
||||
aspect_ratio=request.aspect_ratio or "1:1",
|
||||
size=request.size or "1024x1024",
|
||||
quality=request.quality or "high",
|
||||
reference_image_b64=request.reference_image_b64,
|
||||
)
|
||||
|
||||
if not result.get("success"):
|
||||
raise HTTPException(status_code=500, detail=result.get("error", "图片生成失败"))
|
||||
return result
|
||||
|
||||
|
||||
@app.post("/api/v1/generate-image-upload")
|
||||
async def api_generate_image_with_upload(
|
||||
prompt: str = Form(..., description="广告图片描述"),
|
||||
model: Optional[str] = Form(None),
|
||||
aspect_ratio: Optional[str] = Form("1:1"),
|
||||
quality: Optional[str] = Form("high"),
|
||||
style: Optional[str] = Form(None),
|
||||
brand_name: Optional[str] = Form(None),
|
||||
reference_image: Optional[UploadFile] = File(None, description="参考图片文件"),
|
||||
api_key: str = Depends(get_api_key),
|
||||
):
|
||||
"""上传参考图片并生成广告图(multipart/form-data)"""
|
||||
ref_b64 = None
|
||||
if reference_image:
|
||||
content = await reference_image.read()
|
||||
ref_b64 = base64.b64encode(content).decode("utf-8")
|
||||
|
||||
full_prompt = prompt
|
||||
if style:
|
||||
full_prompt = f"[{style} style] {full_prompt}"
|
||||
if brand_name:
|
||||
full_prompt = f"{full_prompt}. Brand: {brand_name}"
|
||||
|
||||
result = await generate_image_dispatch(
|
||||
prompt=full_prompt,
|
||||
api_key=api_key,
|
||||
model=model,
|
||||
aspect_ratio=aspect_ratio or "1:1",
|
||||
size="1024x1024",
|
||||
quality=quality or "high",
|
||||
reference_image_b64=ref_b64,
|
||||
)
|
||||
|
||||
if not result.get("success"):
|
||||
raise HTTPException(status_code=500, detail=result.get("error", "图片生成失败"))
|
||||
return result
|
||||
|
||||
|
||||
# ---------- 文案生成 ----------
|
||||
|
||||
@app.post("/api/v1/generate-copy")
|
||||
async def api_generate_copy(request: GenerateAdCopyRequest, api_key: str = Depends(get_api_key)):
|
||||
"""生成广告文案"""
|
||||
result = await generate_ad_copy(
|
||||
product=request.product,
|
||||
api_key=api_key,
|
||||
target_audience=request.target_audience,
|
||||
tone=request.tone or "professional",
|
||||
platform=request.platform or "general",
|
||||
language=request.language or "zh",
|
||||
)
|
||||
if not result.get("success"):
|
||||
raise HTTPException(status_code=500, detail=result.get("error", "文案生成失败"))
|
||||
return result
|
||||
|
||||
|
||||
# ---------- 完整广告生成(文案 + 图片) ----------
|
||||
|
||||
@app.post("/api/v1/generate-ad")
|
||||
async def api_generate_ad(request: GenerateAdRequest, api_key: str = Depends(get_api_key)):
|
||||
"""一键生成完整广告(文案 + 配图)"""
|
||||
copy_result = await generate_ad_copy(
|
||||
product=request.product,
|
||||
api_key=api_key,
|
||||
target_audience=request.target_audience,
|
||||
tone=request.tone or "professional",
|
||||
platform=request.platform or "general",
|
||||
language=request.language or "zh",
|
||||
)
|
||||
|
||||
image_prompt = copy_result.get("image_prompt", "")
|
||||
if not image_prompt:
|
||||
image_prompt = f"Advertisement for: {request.product}"
|
||||
|
||||
if request.style:
|
||||
image_prompt = f"[{request.style} style] {image_prompt}"
|
||||
if request.brand_name:
|
||||
image_prompt = f"{image_prompt}. Brand: {request.brand_name}"
|
||||
|
||||
image_result = await generate_image_dispatch(
|
||||
prompt=image_prompt,
|
||||
api_key=api_key,
|
||||
model=request.image_model,
|
||||
aspect_ratio=request.aspect_ratio or "1:1",
|
||||
quality="high",
|
||||
reference_image_b64=request.reference_image_b64,
|
||||
)
|
||||
|
||||
return {
|
||||
"success": True,
|
||||
"copy": copy_result,
|
||||
"image": image_result,
|
||||
"timestamp": datetime.utcnow().isoformat(),
|
||||
}
|
||||
|
||||
|
||||
# ---------- 视频生成 ----------
|
||||
|
||||
@app.post("/api/v1/generate-video")
|
||||
async def api_generate_video(request: GenerateVideoRequest, api_key: str = Depends(get_api_key)):
|
||||
"""生成广告视频"""
|
||||
result = await generate_video_sora(
|
||||
prompt=request.prompt,
|
||||
api_key=api_key,
|
||||
model=request.model,
|
||||
aspect_ratio=request.aspect_ratio or "16:9",
|
||||
duration=request.duration or "5",
|
||||
)
|
||||
if not result.get("success"):
|
||||
raise HTTPException(status_code=500, detail=result.get("error", "视频生成失败"))
|
||||
return result
|
||||
|
||||
|
||||
# ---------- 文件管理 ----------
|
||||
|
||||
@app.get("/api/v1/files/{filename}")
|
||||
async def download_file(filename: str):
|
||||
"""下载生成的文件"""
|
||||
for subdir in ["images", "videos"]:
|
||||
path = os.path.join(OUTPUT_DIR, subdir, filename)
|
||||
if os.path.exists(path):
|
||||
ext = filename.rsplit(".", 1)[-1].lower()
|
||||
media_types = {
|
||||
"png": "image/png", "jpg": "image/jpeg", "jpeg": "image/jpeg",
|
||||
"webp": "image/webp", "mp4": "video/mp4", "gif": "image/gif",
|
||||
}
|
||||
return FileResponse(path, media_type=media_types.get(ext, "application/octet-stream"), filename=filename)
|
||||
raise HTTPException(status_code=404, detail="文件不存在")
|
||||
|
||||
|
||||
@app.get("/api/v1/list-files")
|
||||
async def list_files(file_type: str = "all"):
|
||||
"""列出已生成的文件"""
|
||||
result = {"images": [], "videos": []}
|
||||
|
||||
if file_type in ("all", "image"):
|
||||
img_dir = Path(f"{OUTPUT_DIR}/images")
|
||||
for f in sorted(img_dir.glob("*"), key=lambda p: p.stat().st_mtime, reverse=True):
|
||||
result["images"].append({
|
||||
"filename": f.name,
|
||||
"url": f"/api/v1/files/{f.name}",
|
||||
"size_bytes": f.stat().st_size,
|
||||
"created_at": datetime.fromtimestamp(f.stat().st_mtime).isoformat(),
|
||||
})
|
||||
|
||||
if file_type in ("all", "video"):
|
||||
vid_dir = Path(f"{OUTPUT_DIR}/videos")
|
||||
for f in sorted(vid_dir.glob("*"), key=lambda p: p.stat().st_mtime, reverse=True):
|
||||
result["videos"].append({
|
||||
"filename": f.name,
|
||||
"url": f"/api/v1/files/{f.name}",
|
||||
"size_bytes": f.stat().st_size,
|
||||
"created_at": datetime.fromtimestamp(f.stat().st_mtime).isoformat(),
|
||||
})
|
||||
|
||||
return result
|
||||
|
||||
|
||||
@app.post("/api/v1/cleanup")
|
||||
async def cleanup_files(max_age_hours: int = 24):
|
||||
"""清理超过指定时间的旧文件"""
|
||||
import time
|
||||
|
||||
cutoff = time.time() - max_age_hours * 3600
|
||||
deleted = 0
|
||||
for subdir in ["images", "videos"]:
|
||||
d = Path(f"{OUTPUT_DIR}/{subdir}")
|
||||
for f in d.glob("*"):
|
||||
if f.stat().st_mtime < cutoff:
|
||||
f.unlink()
|
||||
deleted += 1
|
||||
return {"deleted": deleted, "max_age_hours": max_age_hours}
|
||||
|
||||
|
||||
# ---------- 智能对话 ----------
|
||||
|
||||
@app.post("/chat")
|
||||
async def chat(request: ChatRequest, api_key: str = Depends(get_api_key)):
|
||||
"""智能对话 - 理解需求后自动生成广告创意"""
|
||||
|
||||
system_prompt = """你是一位资深广告创意总监 AI 助手。你可以帮用户:
|
||||
1. 分析产品卖点,构思广告创意方案
|
||||
2. 生成广告文案和配图描述
|
||||
3. 推荐合适的广告风格和投放策略
|
||||
|
||||
如果用户想生成图片,在回复末尾添加一行:[GENERATE_IMAGE: 你的英文图片生成 prompt]
|
||||
如果用户只是聊天咨询,正常回复即可。"""
|
||||
|
||||
payload = {
|
||||
"model": DEFAULT_TEXT_MODEL,
|
||||
"messages": [
|
||||
{"role": "system", "content": system_prompt},
|
||||
{"role": "user", "content": request.message},
|
||||
],
|
||||
"max_tokens": 1000,
|
||||
"temperature": 0.8,
|
||||
}
|
||||
|
||||
headers = {"Authorization": f"Bearer {api_key}", "Content-Type": "application/json"}
|
||||
base = LLM_BASE_URL.rstrip("/")
|
||||
url = f"{base}/chat/completions"
|
||||
|
||||
async with aiohttp.ClientSession() as session:
|
||||
async with session.post(url, json=payload, headers=headers, timeout=aiohttp.ClientTimeout(total=30)) as resp:
|
||||
if resp.status != 200:
|
||||
text = await resp.text()
|
||||
raise HTTPException(status_code=500, detail=f"LLM error {resp.status}: {text[:300]}")
|
||||
data = await resp.json()
|
||||
llm_reply = data.get("choices", [{}])[0].get("message", {}).get("content", "")
|
||||
|
||||
image_result = None
|
||||
img_match = re.search(r"\[GENERATE_IMAGE:\s*(.+?)\]", llm_reply)
|
||||
if img_match:
|
||||
img_prompt = img_match.group(1)
|
||||
llm_reply = llm_reply.replace(img_match.group(0), "").strip()
|
||||
image_result = await generate_image_dispatch(prompt=img_prompt, api_key=api_key)
|
||||
|
||||
return {
|
||||
"response": llm_reply,
|
||||
"image": image_result,
|
||||
"timestamp": datetime.utcnow().isoformat(),
|
||||
}
|
||||
|
||||
|
||||
# ==================== 主入口 ====================
|
||||
|
||||
def main():
|
||||
logger.info(f"启动 Ad Creator Agent - {POD_NAME}")
|
||||
logger.info(f"回调功能: {'已启用' if CALLBACK_ENABLED else '未启用'}")
|
||||
uvicorn.run(app, host=SERVICE_HOST, port=SERVICE_PORT, log_level="info")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -26,11 +26,11 @@ COPY common/api_key_utils.py /app/common/
|
||||
# 设置环境变量
|
||||
ENV PYTHONUNBUFFERED=1
|
||||
ENV SERVICE_HOST=0.0.0.0
|
||||
ENV SERVICE_PORT=8080
|
||||
ENV SERVICE_PORT=8000
|
||||
|
||||
# 健康检查 - 使用Python避免僵尸进程
|
||||
HEALTHCHECK --interval=30s --timeout=10s --start-period=40s --retries=3 \
|
||||
CMD python3 -c "import urllib.request; urllib.request.urlopen('http://localhost:8080/health').read()" || exit 1
|
||||
CMD python3 -c "import urllib.request; urllib.request.urlopen('http://localhost:8000/health').read()" || exit 1
|
||||
|
||||
# 运行agent (直接使用Python,避免shell)
|
||||
CMD ["python3", "-u", "azure_blob_agent.py"]
|
||||
|
||||
@@ -25,7 +25,7 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
# 环境变量配置
|
||||
SERVICE_HOST = os.getenv("SERVICE_HOST", "0.0.0.0")
|
||||
SERVICE_PORT = int(os.getenv("SERVICE_PORT", "8080"))
|
||||
SERVICE_PORT = int(os.getenv("SERVICE_PORT", "8000"))
|
||||
POD_NAME = os.getenv("POD_NAME", "azure-blob-agent")
|
||||
TEMPLATE_TYPE = os.getenv("TEMPLATE_TYPE", "azure_blob_agent")
|
||||
|
||||
|
||||
@@ -19,11 +19,11 @@ COPY agents/azure_blob_agent_a2a/azure_blob_agent_a2a.py /app/
|
||||
COPY common/api_key_utils.py /app/common/
|
||||
|
||||
# 暴露端口
|
||||
EXPOSE 8080
|
||||
EXPOSE 8000
|
||||
|
||||
# 健康检查
|
||||
HEALTHCHECK --interval=30s --timeout=10s --start-period=5s --retries=3 \
|
||||
CMD python -c "import requests; requests.get('http://localhost:8080/health', timeout=5)"
|
||||
CMD python -c "import requests; requests.get('http://localhost:8000/health', timeout=5)"
|
||||
|
||||
# 启动应用
|
||||
CMD ["python", "azure_blob_agent_a2a.py"]
|
||||
|
||||
@@ -23,7 +23,7 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
# 环境变量配置
|
||||
SERVICE_HOST = os.getenv("SERVICE_HOST", "0.0.0.0")
|
||||
SERVICE_PORT = int(os.getenv("SERVICE_PORT", "8080"))
|
||||
SERVICE_PORT = int(os.getenv("SERVICE_PORT", "8000"))
|
||||
POD_NAME = os.getenv("POD_NAME", "azure-blob-agent-a2a")
|
||||
TEMPLATE_TYPE = os.getenv("TEMPLATE_TYPE", "azure_blob_agent_a2a")
|
||||
AGENT_FRAMEWORK = os.getenv("AGENT_FRAMEWORK", "a2a")
|
||||
|
||||
@@ -18,11 +18,11 @@ RUN pip install --no-cache-dir -r requirements_mcp.txt
|
||||
COPY agents/azure_blob_agent_mcp/azure_blob_agent_mcp.py /app/
|
||||
|
||||
# 暴露端口
|
||||
EXPOSE 8080
|
||||
EXPOSE 8000
|
||||
|
||||
# 健康检查
|
||||
HEALTHCHECK --interval=30s --timeout=10s --start-period=5s --retries=3 \
|
||||
CMD python -c "import requests; requests.get('http://localhost:8080/health', timeout=5)"
|
||||
CMD python -c "import requests; requests.get('http://localhost:8000/health', timeout=5)"
|
||||
|
||||
# 启动应用
|
||||
CMD ["python", "azure_blob_agent_mcp.py"]
|
||||
|
||||
@@ -22,7 +22,7 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
# 环境变量配置
|
||||
SERVICE_HOST = os.getenv("SERVICE_HOST", "0.0.0.0")
|
||||
SERVICE_PORT = int(os.getenv("SERVICE_PORT", "8080"))
|
||||
SERVICE_PORT = int(os.getenv("SERVICE_PORT", "8000"))
|
||||
POD_NAME = os.getenv("POD_NAME", "azure-blob-agent-mcp")
|
||||
TEMPLATE_TYPE = os.getenv("TEMPLATE_TYPE", "azure_blob_agent_mcp")
|
||||
AGENT_FRAMEWORK = os.getenv("AGENT_FRAMEWORK", "mcp")
|
||||
|
||||
@@ -0,0 +1,39 @@
|
||||
# Chain Analysis Agent Dockerfile
|
||||
# 链上数据分析 Agent - 分析地址活动、交易模式、资金流向
|
||||
|
||||
FROM python:3.11-slim
|
||||
|
||||
WORKDIR /app
|
||||
|
||||
# 安装系统依赖
|
||||
RUN apt-get update && apt-get install -y \
|
||||
curl \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
|
||||
# 复制 common 模块
|
||||
COPY common/ ./common/
|
||||
|
||||
# 复制 Agent 代码
|
||||
COPY chain_analysis_agent.py .
|
||||
COPY requirements.txt .
|
||||
|
||||
# 安装 Python 依赖
|
||||
RUN pip install --no-cache-dir -r requirements.txt
|
||||
|
||||
# 环境变量
|
||||
ENV PYTHONUNBUFFERED=1
|
||||
ENV SERVICE_HOST=0.0.0.0
|
||||
ENV SERVICE_PORT=8000
|
||||
ENV POD_NAME=chain-analysis-agent
|
||||
ENV LLM_BASE_URL=https://litellm.graystone-fb459c5d.southeastasia.azurecontainerapps.io/v1
|
||||
ENV LLM_MODEL=taiji/gpt-4o-mini
|
||||
|
||||
# 健康检查
|
||||
HEALTHCHECK --interval=30s --timeout=10s --start-period=5s --retries=3 \
|
||||
CMD curl -f http://localhost:8000/health || exit 1
|
||||
|
||||
# 暴露端口
|
||||
EXPOSE 8000
|
||||
|
||||
# 运行
|
||||
CMD ["python", "chain_analysis_agent.py"]
|
||||
@@ -0,0 +1,874 @@
|
||||
"""
|
||||
Chain Analysis Agent - 链上数据分析 Agent
|
||||
分析区块链地址活动、交易模式、资金流向、合约交互等
|
||||
支持 Ethereum, BSC, Polygon 等 EVM 兼容链
|
||||
"""
|
||||
import os
|
||||
import sys
|
||||
import logging
|
||||
import aiohttp
|
||||
from typing import Optional, List, Dict, Any
|
||||
from datetime import datetime, timedelta
|
||||
from collections import defaultdict
|
||||
|
||||
from fastapi import FastAPI, HTTPException, Query, Header, Request
|
||||
from fastapi.middleware.cors import CORSMiddleware
|
||||
from pydantic import BaseModel, Field
|
||||
import uvicorn
|
||||
|
||||
# 添加 common 模块路径
|
||||
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
|
||||
# 导入回调工具
|
||||
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, format='%(asctime)s - %(name)s - %(levelname)s - %(message)s')
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# 环境变量
|
||||
SERVICE_HOST = os.getenv("SERVICE_HOST", "0.0.0.0")
|
||||
SERVICE_PORT = int(os.getenv("SERVICE_PORT", "8080"))
|
||||
POD_NAME = os.getenv("POD_NAME", "chain-analysis-agent")
|
||||
USER_ID = os.getenv("USER_ID", "")
|
||||
|
||||
# LLM 配置
|
||||
LLM_BASE_URL = os.getenv("LLM_BASE_URL", "https://litellm.graystone-fb459c5d.southeastasia.azurecontainerapps.io/v1")
|
||||
LLM_MODEL = os.getenv("LLM_MODEL", "taiji/gpt-4o-mini")
|
||||
|
||||
# 支持的区块链网络配置 (Etherscan V2 API)
|
||||
CHAIN_CONFIGS = {
|
||||
"ethereum": {
|
||||
"name": "Ethereum",
|
||||
"symbol": "ETH",
|
||||
"decimals": 18,
|
||||
"chainid": 1,
|
||||
"api_url": "https://api.etherscan.io/v2/api",
|
||||
"explorer_url": "https://etherscan.io"
|
||||
},
|
||||
"bsc": {
|
||||
"name": "BNB Smart Chain",
|
||||
"symbol": "BNB",
|
||||
"decimals": 18,
|
||||
"chainid": 56,
|
||||
"api_url": "https://api.etherscan.io/v2/api",
|
||||
"explorer_url": "https://bscscan.com"
|
||||
},
|
||||
"polygon": {
|
||||
"name": "Polygon",
|
||||
"symbol": "POL",
|
||||
"decimals": 18,
|
||||
"chainid": 137,
|
||||
"api_url": "https://api.etherscan.io/v2/api",
|
||||
"explorer_url": "https://polygonscan.com"
|
||||
},
|
||||
"arbitrum": {
|
||||
"name": "Arbitrum",
|
||||
"symbol": "ETH",
|
||||
"decimals": 18,
|
||||
"chainid": 42161,
|
||||
"api_url": "https://api.etherscan.io/v2/api",
|
||||
"explorer_url": "https://arbiscan.io"
|
||||
},
|
||||
"optimism": {
|
||||
"name": "Optimism",
|
||||
"symbol": "ETH",
|
||||
"decimals": 18,
|
||||
"chainid": 10,
|
||||
"api_url": "https://api.etherscan.io/v2/api",
|
||||
"explorer_url": "https://optimistic.etherscan.io"
|
||||
},
|
||||
"base": {
|
||||
"name": "Base",
|
||||
"symbol": "ETH",
|
||||
"decimals": 18,
|
||||
"chainid": 8453,
|
||||
"api_url": "https://api.etherscan.io/v2/api",
|
||||
"explorer_url": "https://basescan.org"
|
||||
}
|
||||
}
|
||||
|
||||
# FastAPI 应用
|
||||
app = FastAPI(
|
||||
title="Chain Analysis Agent",
|
||||
description="链上数据分析 - 分析地址活动、交易模式、资金流向",
|
||||
version="1.0.0"
|
||||
)
|
||||
|
||||
app.add_middleware(
|
||||
CORSMiddleware,
|
||||
allow_origins=["*"],
|
||||
allow_credentials=True,
|
||||
allow_methods=["*"],
|
||||
allow_headers=["*"],
|
||||
)
|
||||
|
||||
# 回调处理器
|
||||
callback_handler: Optional[AgentCallbackHandler] = None
|
||||
|
||||
|
||||
# ==================== 请求/响应模型 ====================
|
||||
|
||||
class AddressAnalysisRequest(BaseModel):
|
||||
"""地址分析请求"""
|
||||
address: str = Field(..., description="钱包地址")
|
||||
chain: str = Field("ethereum", description="区块链网络")
|
||||
days: int = Field(30, ge=1, le=365, description="分析天数")
|
||||
user_id: Optional[str] = Field(None, description="用户ID")
|
||||
|
||||
|
||||
class TransactionPatternRequest(BaseModel):
|
||||
"""交易模式分析请求"""
|
||||
address: str = Field(..., description="钱包地址")
|
||||
chain: str = Field("ethereum", description="区块链网络")
|
||||
user_id: Optional[str] = Field(None, description="用户ID")
|
||||
|
||||
|
||||
class FundFlowRequest(BaseModel):
|
||||
"""资金流向分析请求"""
|
||||
address: str = Field(..., description="钱包地址")
|
||||
chain: str = Field("ethereum", description="区块链网络")
|
||||
limit: int = Field(100, ge=10, le=500, description="交易数量")
|
||||
user_id: Optional[str] = Field(None, description="用户ID")
|
||||
|
||||
|
||||
class ContractInteractionRequest(BaseModel):
|
||||
"""合约交互分析请求"""
|
||||
address: str = Field(..., description="钱包地址")
|
||||
chain: str = Field("ethereum", description="区块链网络")
|
||||
user_id: Optional[str] = Field(None, description="用户ID")
|
||||
|
||||
|
||||
class ChatRequest(BaseModel):
|
||||
"""Chat 请求"""
|
||||
message: str = Field(..., description="用户消息")
|
||||
chain: str = Field("ethereum", description="默认区块链网络")
|
||||
user_id: Optional[str] = Field(None, description="用户ID")
|
||||
|
||||
|
||||
class ChatResponse(BaseModel):
|
||||
"""Chat 响应"""
|
||||
response: str
|
||||
analysis: Optional[Dict[str, Any]] = None
|
||||
timestamp: str
|
||||
|
||||
|
||||
class HealthResponse(BaseModel):
|
||||
"""健康检查响应"""
|
||||
status: str
|
||||
pod_name: str
|
||||
supported_chains: List[str]
|
||||
callback_enabled: bool
|
||||
timestamp: str
|
||||
|
||||
|
||||
# ==================== 生命周期 ====================
|
||||
|
||||
@app.on_event("startup")
|
||||
async def startup_event():
|
||||
"""应用启动时初始化回调处理器"""
|
||||
global callback_handler
|
||||
|
||||
if CALLBACK_ENABLED and AgentCallbackHandler:
|
||||
try:
|
||||
callback_handler = AgentCallbackHandler(
|
||||
agent_name=POD_NAME,
|
||||
user_id=USER_ID
|
||||
)
|
||||
logger.info(f"回调处理器已初始化: agent={POD_NAME}, user={USER_ID}")
|
||||
except Exception as e:
|
||||
logger.warning(f"回调处理器初始化失败: {e}")
|
||||
|
||||
logger.info(f"Chain Analysis Agent 启动完成 - {POD_NAME}")
|
||||
logger.info(f"支持的区块链: {list(CHAIN_CONFIGS.keys())}")
|
||||
|
||||
|
||||
# ==================== 核心分析功能 ====================
|
||||
|
||||
async def fetch_all_transactions(address: str, chain: str, api_key: str, limit: int = 200) -> List[Dict]:
|
||||
"""获取所有交易用于分析"""
|
||||
if chain not in CHAIN_CONFIGS:
|
||||
return []
|
||||
|
||||
config = CHAIN_CONFIGS[chain]
|
||||
url = config["api_url"]
|
||||
|
||||
params = {
|
||||
"chainid": config["chainid"],
|
||||
"module": "account",
|
||||
"action": "txlist",
|
||||
"address": address,
|
||||
"startblock": 0,
|
||||
"endblock": 99999999,
|
||||
"page": 1,
|
||||
"offset": limit,
|
||||
"sort": "desc",
|
||||
"apikey": api_key
|
||||
}
|
||||
|
||||
try:
|
||||
async with aiohttp.ClientSession() as session:
|
||||
async with session.get(url, params=params, timeout=aiohttp.ClientTimeout(total=20)) as response:
|
||||
if response.status == 200:
|
||||
data = await response.json()
|
||||
if data.get("status") == "1":
|
||||
return data.get("result", [])
|
||||
except Exception as e:
|
||||
logger.error(f"获取交易失败: {e}")
|
||||
return []
|
||||
|
||||
|
||||
async def fetch_internal_transactions(address: str, chain: str, api_key: str) -> List[Dict]:
|
||||
"""获取内部交易"""
|
||||
if chain not in CHAIN_CONFIGS:
|
||||
return []
|
||||
|
||||
config = CHAIN_CONFIGS[chain]
|
||||
url = config["api_url"]
|
||||
|
||||
params = {
|
||||
"chainid": config["chainid"],
|
||||
"module": "account",
|
||||
"action": "txlistinternal",
|
||||
"address": address,
|
||||
"startblock": 0,
|
||||
"endblock": 99999999,
|
||||
"page": 1,
|
||||
"offset": 100,
|
||||
"sort": "desc",
|
||||
"apikey": api_key
|
||||
}
|
||||
|
||||
try:
|
||||
async with aiohttp.ClientSession() as session:
|
||||
async with session.get(url, params=params, timeout=aiohttp.ClientTimeout(total=15)) as response:
|
||||
if response.status == 200:
|
||||
data = await response.json()
|
||||
if data.get("status") == "1":
|
||||
return data.get("result", [])
|
||||
except Exception as e:
|
||||
logger.error(f"获取内部交易失败: {e}")
|
||||
return []
|
||||
|
||||
|
||||
async def fetch_balance(address: str, chain: str, api_key: str) -> float:
|
||||
"""获取余额"""
|
||||
if chain not in CHAIN_CONFIGS:
|
||||
return 0.0
|
||||
|
||||
config = CHAIN_CONFIGS[chain]
|
||||
url = config["api_url"]
|
||||
|
||||
params = {
|
||||
"chainid": config["chainid"],
|
||||
"module": "account",
|
||||
"action": "balance",
|
||||
"address": address,
|
||||
"tag": "latest",
|
||||
"apikey": api_key
|
||||
}
|
||||
|
||||
try:
|
||||
async with aiohttp.ClientSession() as session:
|
||||
async with session.get(url, params=params, timeout=aiohttp.ClientTimeout(total=10)) as response:
|
||||
if response.status == 200:
|
||||
data = await response.json()
|
||||
if data.get("status") == "1":
|
||||
balance_wei = int(data.get("result", 0))
|
||||
return balance_wei / (10 ** config["decimals"])
|
||||
except Exception as e:
|
||||
logger.error(f"获取余额失败: {e}")
|
||||
return 0.0
|
||||
|
||||
|
||||
def analyze_address_activity(transactions: List[Dict], address: str, chain: str, days: int = 30) -> Dict[str, Any]:
|
||||
"""分析地址活动"""
|
||||
config = CHAIN_CONFIGS.get(chain, CHAIN_CONFIGS["ethereum"])
|
||||
address_lower = address.lower()
|
||||
|
||||
now = datetime.utcnow()
|
||||
cutoff = now - timedelta(days=days)
|
||||
|
||||
# 统计数据
|
||||
total_sent = 0.0
|
||||
total_received = 0.0
|
||||
tx_count_in = 0
|
||||
tx_count_out = 0
|
||||
unique_addresses = set()
|
||||
failed_tx = 0
|
||||
daily_activity = defaultdict(lambda: {"in": 0, "out": 0, "count": 0})
|
||||
|
||||
for tx in transactions:
|
||||
try:
|
||||
timestamp = datetime.fromtimestamp(int(tx.get("timeStamp", 0)))
|
||||
if timestamp < cutoff:
|
||||
continue
|
||||
|
||||
value_wei = int(tx.get("value", 0))
|
||||
value = value_wei / (10 ** config["decimals"])
|
||||
|
||||
day_key = timestamp.strftime("%Y-%m-%d")
|
||||
daily_activity[day_key]["count"] += 1
|
||||
|
||||
if tx.get("isError") == "1":
|
||||
failed_tx += 1
|
||||
continue
|
||||
|
||||
from_addr = tx.get("from", "").lower()
|
||||
to_addr = tx.get("to", "").lower()
|
||||
|
||||
if from_addr == address_lower:
|
||||
# 发出
|
||||
total_sent += value
|
||||
tx_count_out += 1
|
||||
daily_activity[day_key]["out"] += value
|
||||
if to_addr:
|
||||
unique_addresses.add(to_addr)
|
||||
elif to_addr == address_lower:
|
||||
# 收到
|
||||
total_received += value
|
||||
tx_count_in += 1
|
||||
daily_activity[day_key]["in"] += value
|
||||
unique_addresses.add(from_addr)
|
||||
except Exception as e:
|
||||
logger.error(f"解析交易失败: {e}")
|
||||
|
||||
# 计算活跃天数
|
||||
active_days = len(daily_activity)
|
||||
|
||||
return {
|
||||
"address": address,
|
||||
"chain": chain,
|
||||
"period_days": days,
|
||||
"summary": {
|
||||
"total_sent": round(total_sent, 6),
|
||||
"total_received": round(total_received, 6),
|
||||
"net_flow": round(total_received - total_sent, 6),
|
||||
"tx_count_in": tx_count_in,
|
||||
"tx_count_out": tx_count_out,
|
||||
"total_tx": tx_count_in + tx_count_out,
|
||||
"failed_tx": failed_tx,
|
||||
"unique_addresses": len(unique_addresses),
|
||||
"active_days": active_days
|
||||
},
|
||||
"symbol": config["symbol"],
|
||||
"daily_activity": dict(sorted(daily_activity.items(), reverse=True)[:7]) # 最近7天
|
||||
}
|
||||
|
||||
|
||||
def analyze_transaction_patterns(transactions: List[Dict], address: str, chain: str) -> Dict[str, Any]:
|
||||
"""分析交易模式"""
|
||||
config = CHAIN_CONFIGS.get(chain, CHAIN_CONFIGS["ethereum"])
|
||||
address_lower = address.lower()
|
||||
|
||||
# 时间分布
|
||||
hourly_distribution = defaultdict(int)
|
||||
daily_distribution = defaultdict(int)
|
||||
|
||||
# 金额分布
|
||||
value_ranges = {
|
||||
"micro": 0, # < 0.01
|
||||
"small": 0, # 0.01 - 0.1
|
||||
"medium": 0, # 0.1 - 1
|
||||
"large": 0, # 1 - 10
|
||||
"whale": 0 # > 10
|
||||
}
|
||||
|
||||
# 交互地址频率
|
||||
address_frequency = defaultdict(int)
|
||||
|
||||
# 交易间隔
|
||||
timestamps = []
|
||||
|
||||
for tx in transactions:
|
||||
try:
|
||||
timestamp = datetime.fromtimestamp(int(tx.get("timeStamp", 0)))
|
||||
timestamps.append(timestamp)
|
||||
|
||||
hourly_distribution[timestamp.hour] += 1
|
||||
daily_distribution[timestamp.strftime("%A")] += 1
|
||||
|
||||
value_wei = int(tx.get("value", 0))
|
||||
value = value_wei / (10 ** config["decimals"])
|
||||
|
||||
if value < 0.01:
|
||||
value_ranges["micro"] += 1
|
||||
elif value < 0.1:
|
||||
value_ranges["small"] += 1
|
||||
elif value < 1:
|
||||
value_ranges["medium"] += 1
|
||||
elif value < 10:
|
||||
value_ranges["large"] += 1
|
||||
else:
|
||||
value_ranges["whale"] += 1
|
||||
|
||||
from_addr = tx.get("from", "").lower()
|
||||
to_addr = tx.get("to", "").lower()
|
||||
|
||||
counterparty = to_addr if from_addr == address_lower else from_addr
|
||||
if counterparty:
|
||||
address_frequency[counterparty] += 1
|
||||
except Exception as e:
|
||||
logger.error(f"解析交易失败: {e}")
|
||||
|
||||
# 计算交易间隔
|
||||
avg_interval = None
|
||||
if len(timestamps) > 1:
|
||||
timestamps.sort(reverse=True)
|
||||
intervals = []
|
||||
for i in range(len(timestamps) - 1):
|
||||
interval = (timestamps[i] - timestamps[i+1]).total_seconds() / 3600 # 小时
|
||||
intervals.append(interval)
|
||||
avg_interval = round(sum(intervals) / len(intervals), 2)
|
||||
|
||||
# 前5个交互地址
|
||||
top_addresses = sorted(address_frequency.items(), key=lambda x: x[1], reverse=True)[:5]
|
||||
|
||||
return {
|
||||
"address": address,
|
||||
"chain": chain,
|
||||
"patterns": {
|
||||
"hourly_distribution": dict(hourly_distribution),
|
||||
"daily_distribution": dict(daily_distribution),
|
||||
"value_distribution": value_ranges,
|
||||
"avg_interval_hours": avg_interval,
|
||||
"top_counterparties": [{"address": addr, "tx_count": count} for addr, count in top_addresses]
|
||||
},
|
||||
"behavior_summary": generate_behavior_summary(hourly_distribution, value_ranges, avg_interval)
|
||||
}
|
||||
|
||||
|
||||
def generate_behavior_summary(hourly: Dict, values: Dict, interval: Optional[float]) -> str:
|
||||
"""生成行为摘要"""
|
||||
summary_parts = []
|
||||
|
||||
# 活跃时段
|
||||
if hourly:
|
||||
peak_hour = max(hourly, key=hourly.get)
|
||||
summary_parts.append(f"活跃高峰时段: {peak_hour}:00 UTC")
|
||||
|
||||
# 交易规模
|
||||
total_tx = sum(values.values())
|
||||
if total_tx > 0:
|
||||
whale_ratio = values["whale"] / total_tx * 100
|
||||
if whale_ratio > 20:
|
||||
summary_parts.append("大额交易频繁(可能是机构或巨鲸)")
|
||||
elif values["micro"] / total_tx > 0.5:
|
||||
summary_parts.append("以小额交易为主(可能是频繁交易者或机器人)")
|
||||
|
||||
# 交易频率
|
||||
if interval:
|
||||
if interval < 1:
|
||||
summary_parts.append("高频交易(可能是自动化程序)")
|
||||
elif interval > 168: # 一周
|
||||
summary_parts.append("低频交易(普通持有者)")
|
||||
|
||||
return "; ".join(summary_parts) if summary_parts else "交易模式正常"
|
||||
|
||||
|
||||
def analyze_fund_flow(transactions: List[Dict], address: str, chain: str) -> Dict[str, Any]:
|
||||
"""分析资金流向"""
|
||||
config = CHAIN_CONFIGS.get(chain, CHAIN_CONFIGS["ethereum"])
|
||||
address_lower = address.lower()
|
||||
|
||||
inflow = defaultdict(float) # 资金来源
|
||||
outflow = defaultdict(float) # 资金去向
|
||||
|
||||
for tx in transactions:
|
||||
try:
|
||||
if tx.get("isError") == "1":
|
||||
continue
|
||||
|
||||
value_wei = int(tx.get("value", 0))
|
||||
value = value_wei / (10 ** config["decimals"])
|
||||
|
||||
if value == 0:
|
||||
continue
|
||||
|
||||
from_addr = tx.get("from", "").lower()
|
||||
to_addr = tx.get("to", "").lower()
|
||||
|
||||
if from_addr == address_lower and to_addr:
|
||||
outflow[to_addr] += value
|
||||
elif to_addr == address_lower:
|
||||
inflow[from_addr] += value
|
||||
except Exception as e:
|
||||
logger.error(f"解析交易失败: {e}")
|
||||
|
||||
# 排序获取 Top 10
|
||||
top_inflow = sorted(inflow.items(), key=lambda x: x[1], reverse=True)[:10]
|
||||
top_outflow = sorted(outflow.items(), key=lambda x: x[1], reverse=True)[:10]
|
||||
|
||||
total_in = sum(inflow.values())
|
||||
total_out = sum(outflow.values())
|
||||
|
||||
return {
|
||||
"address": address,
|
||||
"chain": chain,
|
||||
"fund_flow": {
|
||||
"total_inflow": round(total_in, 6),
|
||||
"total_outflow": round(total_out, 6),
|
||||
"net_flow": round(total_in - total_out, 6),
|
||||
"inflow_sources": len(inflow),
|
||||
"outflow_destinations": len(outflow),
|
||||
"top_inflow": [
|
||||
{"address": addr, "amount": round(amt, 6), "symbol": config["symbol"]}
|
||||
for addr, amt in top_inflow
|
||||
],
|
||||
"top_outflow": [
|
||||
{"address": addr, "amount": round(amt, 6), "symbol": config["symbol"]}
|
||||
for addr, amt in top_outflow
|
||||
]
|
||||
},
|
||||
"symbol": config["symbol"]
|
||||
}
|
||||
|
||||
|
||||
def analyze_contract_interactions(transactions: List[Dict], address: str, chain: str) -> Dict[str, Any]:
|
||||
"""分析合约交互"""
|
||||
config = CHAIN_CONFIGS.get(chain, CHAIN_CONFIGS["ethereum"])
|
||||
address_lower = address.lower()
|
||||
|
||||
contract_interactions = defaultdict(lambda: {"count": 0, "methods": set(), "value": 0.0})
|
||||
|
||||
for tx in transactions:
|
||||
try:
|
||||
from_addr = tx.get("from", "").lower()
|
||||
to_addr = tx.get("to", "").lower()
|
||||
|
||||
# 只分析发出的交易且有 input data 的(合约调用)
|
||||
if from_addr != address_lower:
|
||||
continue
|
||||
|
||||
input_data = tx.get("input", "")
|
||||
if input_data and input_data != "0x" and len(input_data) >= 10:
|
||||
method_id = input_data[:10]
|
||||
value_wei = int(tx.get("value", 0))
|
||||
value = value_wei / (10 ** config["decimals"])
|
||||
|
||||
contract_interactions[to_addr]["count"] += 1
|
||||
contract_interactions[to_addr]["methods"].add(method_id)
|
||||
contract_interactions[to_addr]["value"] += value
|
||||
except Exception as e:
|
||||
logger.error(f"解析交易失败: {e}")
|
||||
|
||||
# 排序
|
||||
sorted_contracts = sorted(
|
||||
contract_interactions.items(),
|
||||
key=lambda x: x[1]["count"],
|
||||
reverse=True
|
||||
)[:10]
|
||||
|
||||
return {
|
||||
"address": address,
|
||||
"chain": chain,
|
||||
"contract_interactions": {
|
||||
"total_contracts": len(contract_interactions),
|
||||
"top_contracts": [
|
||||
{
|
||||
"contract": addr,
|
||||
"interaction_count": data["count"],
|
||||
"unique_methods": len(data["methods"]),
|
||||
"total_value": round(data["value"], 6),
|
||||
"symbol": config["symbol"],
|
||||
"explorer_url": f"{config['explorer_url']}/address/{addr}"
|
||||
}
|
||||
for addr, data in sorted_contracts
|
||||
]
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
async def chat_with_llm(message: str, context: str, api_key: str) -> str:
|
||||
"""调用 LLM 生成分析报告"""
|
||||
try:
|
||||
async with aiohttp.ClientSession() as session:
|
||||
payload = {
|
||||
"model": LLM_MODEL,
|
||||
"messages": [
|
||||
{
|
||||
"role": "system",
|
||||
"content": """你是一个资深的区块链数据分析师,擅长:
|
||||
1. 分析钱包地址的链上行为模式
|
||||
2. 识别交易特征(高频交易、巨鲸、机器人等)
|
||||
3. 追踪资金流向和来源
|
||||
4. 分析合约交互行为
|
||||
5. 提供风险评估和投资建议
|
||||
|
||||
请根据链上数据提供专业、深入的分析报告,用简洁易懂的语言表达。"""
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": f"链上分析数据:\n{context}\n\n分析请求: {message}"
|
||||
}
|
||||
],
|
||||
"max_tokens": 800,
|
||||
"temperature": 0.7
|
||||
}
|
||||
headers = {
|
||||
"Authorization": f"Bearer {api_key}",
|
||||
"Content-Type": "application/json"
|
||||
}
|
||||
|
||||
async with session.post(
|
||||
f"{LLM_BASE_URL}/chat/completions",
|
||||
json=payload,
|
||||
headers=headers,
|
||||
timeout=aiohttp.ClientTimeout(total=30)
|
||||
) as response:
|
||||
if response.status == 200:
|
||||
data = await response.json()
|
||||
return data.get("choices", [{}])[0].get("message", {}).get("content", "抱歉,无法生成分析报告")
|
||||
else:
|
||||
error = await response.text()
|
||||
logger.error(f"LLM 请求失败: {response.status} - {error}")
|
||||
return f"LLM 服务错误: {response.status}"
|
||||
except Exception as e:
|
||||
logger.error(f"LLM 调用失败: {e}")
|
||||
return f"分析失败: {str(e)}"
|
||||
|
||||
|
||||
def extract_address_from_message(message: str) -> Optional[str]:
|
||||
"""从消息中提取以太坊地址"""
|
||||
import re
|
||||
pattern = r'0x[a-fA-F0-9]{40}'
|
||||
match = re.search(pattern, message)
|
||||
return match.group(0) if match else None
|
||||
|
||||
|
||||
# ==================== API 端点 ====================
|
||||
|
||||
@app.get("/", response_model=dict)
|
||||
async def root():
|
||||
"""服务状态"""
|
||||
return {
|
||||
"service": "Chain Analysis Agent",
|
||||
"description": "链上数据分析 - 分析地址活动、交易模式、资金流向",
|
||||
"status": "running",
|
||||
"supported_chains": list(CHAIN_CONFIGS.keys()),
|
||||
"tools": ["address_analysis", "transaction_patterns", "fund_flow", "contract_interactions", "chat"]
|
||||
}
|
||||
|
||||
|
||||
@app.get("/health", response_model=HealthResponse)
|
||||
async def health_check():
|
||||
"""健康检查"""
|
||||
return HealthResponse(
|
||||
status="healthy",
|
||||
pod_name=POD_NAME,
|
||||
supported_chains=list(CHAIN_CONFIGS.keys()),
|
||||
callback_enabled=CALLBACK_ENABLED,
|
||||
timestamp=datetime.utcnow().isoformat()
|
||||
)
|
||||
|
||||
|
||||
@app.post("/address-analysis")
|
||||
async def address_analysis(
|
||||
request: AddressAnalysisRequest,
|
||||
api_key: Optional[str] = Header(None, alias="api-key"),
|
||||
etherscan_key: Optional[str] = Header(None, alias="etherscan-key")
|
||||
):
|
||||
"""地址活动分析"""
|
||||
scan_key = etherscan_key or api_key
|
||||
if not scan_key:
|
||||
raise HTTPException(status_code=401, detail="请在请求头中提供 etherscan-key 或 api-key")
|
||||
|
||||
transactions = await fetch_all_transactions(request.address, request.chain, scan_key)
|
||||
|
||||
if not transactions:
|
||||
raise HTTPException(status_code=404, detail="未找到交易记录")
|
||||
|
||||
result = analyze_address_activity(transactions, request.address, request.chain, request.days)
|
||||
balance = await fetch_balance(request.address, request.chain, scan_key)
|
||||
result["current_balance"] = round(balance, 8)
|
||||
|
||||
return {
|
||||
**result,
|
||||
"timestamp": datetime.utcnow().isoformat()
|
||||
}
|
||||
|
||||
|
||||
@app.post("/transaction-patterns")
|
||||
async def transaction_patterns(
|
||||
request: TransactionPatternRequest,
|
||||
api_key: Optional[str] = Header(None, alias="api-key"),
|
||||
etherscan_key: Optional[str] = Header(None, alias="etherscan-key")
|
||||
):
|
||||
"""交易模式分析"""
|
||||
scan_key = etherscan_key or api_key
|
||||
if not scan_key:
|
||||
raise HTTPException(status_code=401, detail="请在请求头中提供 etherscan-key 或 api-key")
|
||||
|
||||
transactions = await fetch_all_transactions(request.address, request.chain, scan_key)
|
||||
|
||||
if not transactions:
|
||||
raise HTTPException(status_code=404, detail="未找到交易记录")
|
||||
|
||||
result = analyze_transaction_patterns(transactions, request.address, request.chain)
|
||||
|
||||
return {
|
||||
**result,
|
||||
"timestamp": datetime.utcnow().isoformat()
|
||||
}
|
||||
|
||||
|
||||
@app.post("/fund-flow")
|
||||
async def fund_flow(
|
||||
request: FundFlowRequest,
|
||||
api_key: Optional[str] = Header(None, alias="api-key"),
|
||||
etherscan_key: Optional[str] = Header(None, alias="etherscan-key")
|
||||
):
|
||||
"""资金流向分析"""
|
||||
scan_key = etherscan_key or api_key
|
||||
if not scan_key:
|
||||
raise HTTPException(status_code=401, detail="请在请求头中提供 etherscan-key 或 api-key")
|
||||
|
||||
transactions = await fetch_all_transactions(request.address, request.chain, scan_key, request.limit)
|
||||
|
||||
if not transactions:
|
||||
raise HTTPException(status_code=404, detail="未找到交易记录")
|
||||
|
||||
result = analyze_fund_flow(transactions, request.address, request.chain)
|
||||
|
||||
return {
|
||||
**result,
|
||||
"timestamp": datetime.utcnow().isoformat()
|
||||
}
|
||||
|
||||
|
||||
@app.post("/contract-interactions")
|
||||
async def contract_interactions(
|
||||
request: ContractInteractionRequest,
|
||||
api_key: Optional[str] = Header(None, alias="api-key"),
|
||||
etherscan_key: Optional[str] = Header(None, alias="etherscan-key")
|
||||
):
|
||||
"""合约交互分析"""
|
||||
scan_key = etherscan_key or api_key
|
||||
if not scan_key:
|
||||
raise HTTPException(status_code=401, detail="请在请求头中提供 etherscan-key 或 api-key")
|
||||
|
||||
transactions = await fetch_all_transactions(request.address, request.chain, scan_key)
|
||||
|
||||
if not transactions:
|
||||
raise HTTPException(status_code=404, detail="未找到交易记录")
|
||||
|
||||
result = analyze_contract_interactions(transactions, request.address, request.chain)
|
||||
|
||||
return {
|
||||
**result,
|
||||
"timestamp": datetime.utcnow().isoformat()
|
||||
}
|
||||
|
||||
|
||||
@app.post("/chat", response_model=ChatResponse)
|
||||
async def chat(
|
||||
request: ChatRequest,
|
||||
api_key: Optional[str] = Header(None, alias="api-key"),
|
||||
etherscan_key: Optional[str] = Header(None, alias="etherscan-key"),
|
||||
llm_key: Optional[str] = Header(None, alias="llm-key"),
|
||||
authorization: Optional[str] = Header(None)
|
||||
):
|
||||
"""智能对话 - 支持自然语言分析链上数据
|
||||
|
||||
api_key 通过请求头传递:
|
||||
- api-key 或 etherscan-key: 区块链浏览器 API Key
|
||||
- llm-key 或 Authorization: LLM API Key
|
||||
"""
|
||||
# 获取区块链 API Key
|
||||
scan_key = etherscan_key or api_key
|
||||
if not scan_key:
|
||||
raise HTTPException(status_code=401, detail="请在请求头中提供 etherscan-key 或 api-key")
|
||||
|
||||
# 获取 LLM API Key
|
||||
llm_api_key = llm_key
|
||||
if not llm_api_key and authorization:
|
||||
if authorization.startswith("Bearer "):
|
||||
llm_api_key = authorization[7:]
|
||||
else:
|
||||
llm_api_key = authorization
|
||||
|
||||
if not llm_api_key:
|
||||
raise HTTPException(status_code=401, detail="请在请求头中提供 llm-key 或 Authorization")
|
||||
|
||||
# 从消息中提取地址
|
||||
address = extract_address_from_message(request.message)
|
||||
|
||||
analysis_data = {}
|
||||
if address:
|
||||
transactions = await fetch_all_transactions(address, request.chain, scan_key)
|
||||
|
||||
if transactions:
|
||||
# 执行全面分析
|
||||
analysis_data["activity"] = analyze_address_activity(transactions, address, request.chain)
|
||||
analysis_data["patterns"] = analyze_transaction_patterns(transactions, address, request.chain)
|
||||
analysis_data["fund_flow"] = analyze_fund_flow(transactions, address, request.chain)
|
||||
analysis_data["contracts"] = analyze_contract_interactions(transactions, address, request.chain)
|
||||
analysis_data["balance"] = await fetch_balance(address, request.chain, scan_key)
|
||||
|
||||
# 构建上下文
|
||||
if analysis_data:
|
||||
context_parts = []
|
||||
if "activity" in analysis_data:
|
||||
s = analysis_data["activity"]["summary"]
|
||||
context_parts.append(f"地址: {address}")
|
||||
context_parts.append(f"当前余额: {analysis_data['balance']:.6f} ETH")
|
||||
context_parts.append(f"30天活动: 收入 {s['total_received']:.4f} ETH, 支出 {s['total_sent']:.4f} ETH")
|
||||
context_parts.append(f"交易统计: 入账 {s['tx_count_in']} 笔, 出账 {s['tx_count_out']} 笔")
|
||||
if "patterns" in analysis_data:
|
||||
p = analysis_data["patterns"]
|
||||
context_parts.append(f"行为特征: {p['behavior_summary']}")
|
||||
if "fund_flow" in analysis_data:
|
||||
f = analysis_data["fund_flow"]["fund_flow"]
|
||||
context_parts.append(f"资金来源数: {f['inflow_sources']}, 去向数: {f['outflow_destinations']}")
|
||||
if "contracts" in analysis_data:
|
||||
c = analysis_data["contracts"]["contract_interactions"]
|
||||
context_parts.append(f"交互合约数: {c['total_contracts']}")
|
||||
context = "\n".join(context_parts)
|
||||
else:
|
||||
context = "未检测到有效的钱包地址,请提供 0x 开头的以太坊地址"
|
||||
|
||||
# 调用 LLM 生成分析报告
|
||||
llm_response = await chat_with_llm(request.message, context, llm_api_key)
|
||||
|
||||
return ChatResponse(
|
||||
response=llm_response,
|
||||
analysis=analysis_data if analysis_data else {"detected_address": address},
|
||||
timestamp=datetime.utcnow().isoformat()
|
||||
)
|
||||
|
||||
|
||||
@app.get("/chains")
|
||||
async def list_chains():
|
||||
"""列出支持的区块链"""
|
||||
return {
|
||||
"chains": [
|
||||
{
|
||||
"id": chain_id,
|
||||
"name": config["name"],
|
||||
"symbol": config["symbol"],
|
||||
"explorer": config["explorer_url"]
|
||||
}
|
||||
for chain_id, config in CHAIN_CONFIGS.items()
|
||||
]
|
||||
}
|
||||
|
||||
|
||||
# ==================== 主入口 ====================
|
||||
|
||||
def main():
|
||||
"""主函数"""
|
||||
logger.info(f"启动 Chain Analysis Agent - {POD_NAME}")
|
||||
logger.info(f"支持的区块链: {list(CHAIN_CONFIGS.keys())}")
|
||||
logger.info(f"回调功能: {'已启用' if CALLBACK_ENABLED else '未启用'}")
|
||||
|
||||
uvicorn.run(app, host=SERVICE_HOST, port=SERVICE_PORT, log_level="info")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,6 @@
|
||||
fastapi>=0.104.0
|
||||
uvicorn>=0.24.0
|
||||
aiohttp>=3.9.0
|
||||
pydantic>=2.0.0
|
||||
python-multipart>=0.0.6
|
||||
httpx>=0.25.0
|
||||
@@ -0,0 +1,39 @@
|
||||
# Chain Explorer Agent Dockerfile
|
||||
# 链上数据查询 Agent - 查询地址余额、交易记录、代币信息
|
||||
|
||||
FROM python:3.11-slim
|
||||
|
||||
WORKDIR /app
|
||||
|
||||
# 安装系统依赖
|
||||
RUN apt-get update && apt-get install -y \
|
||||
curl \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
|
||||
# 复制 common 模块
|
||||
COPY common/ ./common/
|
||||
|
||||
# 复制 Agent 代码
|
||||
COPY chain_explorer_agent.py .
|
||||
COPY requirements.txt .
|
||||
|
||||
# 安装 Python 依赖
|
||||
RUN pip install --no-cache-dir -r requirements.txt
|
||||
|
||||
# 环境变量
|
||||
ENV PYTHONUNBUFFERED=1
|
||||
ENV SERVICE_HOST=0.0.0.0
|
||||
ENV SERVICE_PORT=8000
|
||||
ENV POD_NAME=chain-explorer-agent
|
||||
ENV LLM_BASE_URL=https://litellm.graystone-fb459c5d.southeastasia.azurecontainerapps.io/v1
|
||||
ENV LLM_MODEL=taiji/gpt-4o-mini
|
||||
|
||||
# 健康检查
|
||||
HEALTHCHECK --interval=30s --timeout=10s --start-period=5s --retries=3 \
|
||||
CMD curl -f http://localhost:8000/health || exit 1
|
||||
|
||||
# 暴露端口
|
||||
EXPOSE 8000
|
||||
|
||||
# 运行
|
||||
CMD ["python", "chain_explorer_agent.py"]
|
||||
@@ -0,0 +1,613 @@
|
||||
"""
|
||||
Chain Explorer Agent - 链上数据查询 Agent
|
||||
查询区块链地址余额、交易记录、代币信息等
|
||||
支持 Ethereum, BSC, Polygon 等 EVM 兼容链
|
||||
"""
|
||||
import os
|
||||
import sys
|
||||
import logging
|
||||
import aiohttp
|
||||
from typing import Optional, List, Dict, Any
|
||||
from datetime import datetime
|
||||
from decimal import Decimal
|
||||
|
||||
from fastapi import FastAPI, HTTPException, Query, Header, Request
|
||||
from fastapi.middleware.cors import CORSMiddleware
|
||||
from pydantic import BaseModel, Field
|
||||
import uvicorn
|
||||
|
||||
# 添加 common 模块路径
|
||||
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
|
||||
# 导入回调工具
|
||||
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, format='%(asctime)s - %(name)s - %(levelname)s - %(message)s')
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# 环境变量
|
||||
SERVICE_HOST = os.getenv("SERVICE_HOST", "0.0.0.0")
|
||||
SERVICE_PORT = int(os.getenv("SERVICE_PORT", "8080"))
|
||||
POD_NAME = os.getenv("POD_NAME", "chain-explorer-agent")
|
||||
USER_ID = os.getenv("USER_ID", "")
|
||||
|
||||
# LLM 配置
|
||||
LLM_BASE_URL = os.getenv("LLM_BASE_URL", "https://litellm.graystone-fb459c5d.southeastasia.azurecontainerapps.io/v1")
|
||||
LLM_MODEL = os.getenv("LLM_MODEL", "taiji/gpt-4o-mini")
|
||||
|
||||
# 支持的区块链网络配置 (Etherscan V2 API)
|
||||
CHAIN_CONFIGS = {
|
||||
"ethereum": {
|
||||
"name": "Ethereum",
|
||||
"symbol": "ETH",
|
||||
"decimals": 18,
|
||||
"chainid": 1,
|
||||
"api_url": "https://api.etherscan.io/v2/api",
|
||||
"explorer_url": "https://etherscan.io"
|
||||
},
|
||||
"bsc": {
|
||||
"name": "BNB Smart Chain",
|
||||
"symbol": "BNB",
|
||||
"decimals": 18,
|
||||
"chainid": 56,
|
||||
"api_url": "https://api.etherscan.io/v2/api",
|
||||
"explorer_url": "https://bscscan.com"
|
||||
},
|
||||
"polygon": {
|
||||
"name": "Polygon",
|
||||
"symbol": "POL",
|
||||
"decimals": 18,
|
||||
"chainid": 137,
|
||||
"api_url": "https://api.etherscan.io/v2/api",
|
||||
"explorer_url": "https://polygonscan.com"
|
||||
},
|
||||
"arbitrum": {
|
||||
"name": "Arbitrum",
|
||||
"symbol": "ETH",
|
||||
"decimals": 18,
|
||||
"chainid": 42161,
|
||||
"api_url": "https://api.etherscan.io/v2/api",
|
||||
"explorer_url": "https://arbiscan.io"
|
||||
},
|
||||
"optimism": {
|
||||
"name": "Optimism",
|
||||
"symbol": "ETH",
|
||||
"decimals": 18,
|
||||
"chainid": 10,
|
||||
"api_url": "https://api.etherscan.io/v2/api",
|
||||
"explorer_url": "https://optimistic.etherscan.io"
|
||||
},
|
||||
"base": {
|
||||
"name": "Base",
|
||||
"symbol": "ETH",
|
||||
"decimals": 18,
|
||||
"chainid": 8453,
|
||||
"api_url": "https://api.etherscan.io/v2/api",
|
||||
"explorer_url": "https://basescan.org"
|
||||
}
|
||||
}
|
||||
|
||||
# FastAPI 应用
|
||||
app = FastAPI(
|
||||
title="Chain Explorer Agent",
|
||||
description="链上数据查询 - 查询地址余额、交易记录、代币信息",
|
||||
version="1.0.0"
|
||||
)
|
||||
|
||||
app.add_middleware(
|
||||
CORSMiddleware,
|
||||
allow_origins=["*"],
|
||||
allow_credentials=True,
|
||||
allow_methods=["*"],
|
||||
allow_headers=["*"],
|
||||
)
|
||||
|
||||
# 回调处理器
|
||||
callback_handler: Optional[AgentCallbackHandler] = None
|
||||
|
||||
|
||||
# ==================== 请求/响应模型 ====================
|
||||
|
||||
class BalanceRequest(BaseModel):
|
||||
"""余额查询请求"""
|
||||
address: str = Field(..., description="钱包地址")
|
||||
chain: str = Field("ethereum", description="区块链网络: ethereum, bsc, polygon, arbitrum, optimism")
|
||||
user_id: Optional[str] = Field(None, description="用户ID")
|
||||
|
||||
|
||||
class BalanceResponse(BaseModel):
|
||||
"""余额响应"""
|
||||
address: str
|
||||
chain: str
|
||||
balance: str
|
||||
balance_formatted: str
|
||||
symbol: str
|
||||
usd_value: Optional[float] = None
|
||||
timestamp: str
|
||||
|
||||
|
||||
class TransactionRequest(BaseModel):
|
||||
"""交易查询请求"""
|
||||
address: str = Field(..., description="钱包地址")
|
||||
chain: str = Field("ethereum", description="区块链网络")
|
||||
page: int = Field(1, ge=1, description="页码")
|
||||
limit: int = Field(10, ge=1, le=100, description="每页数量")
|
||||
user_id: Optional[str] = Field(None, description="用户ID")
|
||||
|
||||
|
||||
class TokenBalanceRequest(BaseModel):
|
||||
"""代币余额查询请求"""
|
||||
address: str = Field(..., description="钱包地址")
|
||||
chain: str = Field("ethereum", description="区块链网络")
|
||||
user_id: Optional[str] = Field(None, description="用户ID")
|
||||
|
||||
|
||||
class ChatRequest(BaseModel):
|
||||
"""Chat 请求"""
|
||||
message: str = Field(..., description="用户消息")
|
||||
chain: str = Field("ethereum", description="默认区块链网络")
|
||||
user_id: Optional[str] = Field(None, description="用户ID")
|
||||
|
||||
|
||||
class ChatResponse(BaseModel):
|
||||
"""Chat 响应"""
|
||||
response: str
|
||||
data: Optional[Dict[str, Any]] = None
|
||||
timestamp: str
|
||||
|
||||
|
||||
class HealthResponse(BaseModel):
|
||||
"""健康检查响应"""
|
||||
status: str
|
||||
pod_name: str
|
||||
supported_chains: List[str]
|
||||
callback_enabled: bool
|
||||
timestamp: str
|
||||
|
||||
|
||||
# ==================== 生命周期 ====================
|
||||
|
||||
@app.on_event("startup")
|
||||
async def startup_event():
|
||||
"""应用启动时初始化回调处理器"""
|
||||
global callback_handler
|
||||
|
||||
if CALLBACK_ENABLED and AgentCallbackHandler:
|
||||
try:
|
||||
callback_handler = AgentCallbackHandler(
|
||||
agent_name=POD_NAME,
|
||||
user_id=USER_ID
|
||||
)
|
||||
logger.info(f"回调处理器已初始化: agent={POD_NAME}, user={USER_ID}")
|
||||
except Exception as e:
|
||||
logger.warning(f"回调处理器初始化失败: {e}")
|
||||
|
||||
logger.info(f"Chain Explorer Agent 启动完成 - {POD_NAME}")
|
||||
logger.info(f"支持的区块链: {list(CHAIN_CONFIGS.keys())}")
|
||||
|
||||
|
||||
# ==================== 核心功能 ====================
|
||||
|
||||
async def fetch_balance(address: str, chain: str, api_key: str) -> Dict[str, Any]:
|
||||
"""获取地址余额"""
|
||||
if chain not in CHAIN_CONFIGS:
|
||||
return {"success": False, "error": f"不支持的区块链: {chain}"}
|
||||
|
||||
config = CHAIN_CONFIGS[chain]
|
||||
url = config["api_url"]
|
||||
|
||||
params = {
|
||||
"chainid": config["chainid"],
|
||||
"module": "account",
|
||||
"action": "balance",
|
||||
"address": address,
|
||||
"tag": "latest",
|
||||
"apikey": api_key
|
||||
}
|
||||
|
||||
try:
|
||||
async with aiohttp.ClientSession() as session:
|
||||
async with session.get(url, params=params, timeout=aiohttp.ClientTimeout(total=15)) as response:
|
||||
if response.status == 200:
|
||||
data = await response.json()
|
||||
if data.get("status") == "1":
|
||||
balance_wei = int(data.get("result", 0))
|
||||
balance_eth = balance_wei / (10 ** config["decimals"])
|
||||
return {
|
||||
"success": True,
|
||||
"address": address,
|
||||
"chain": chain,
|
||||
"chain_name": config["name"],
|
||||
"balance_wei": str(balance_wei),
|
||||
"balance": round(balance_eth, 8),
|
||||
"symbol": config["symbol"],
|
||||
"explorer_url": f"{config['explorer_url']}/address/{address}"
|
||||
}
|
||||
else:
|
||||
return {"success": False, "error": data.get("message", "API 错误")}
|
||||
else:
|
||||
return {"success": False, "error": f"HTTP {response.status}"}
|
||||
except Exception as e:
|
||||
logger.error(f"获取余额失败: {e}")
|
||||
return {"success": False, "error": str(e)}
|
||||
|
||||
|
||||
async def fetch_transactions(address: str, chain: str, api_key: str, page: int = 1, limit: int = 10) -> Dict[str, Any]:
|
||||
"""获取交易记录"""
|
||||
if chain not in CHAIN_CONFIGS:
|
||||
return {"success": False, "error": f"不支持的区块链: {chain}"}
|
||||
|
||||
config = CHAIN_CONFIGS[chain]
|
||||
url = config["api_url"]
|
||||
|
||||
params = {
|
||||
"chainid": config["chainid"],
|
||||
"module": "account",
|
||||
"action": "txlist",
|
||||
"address": address,
|
||||
"startblock": 0,
|
||||
"endblock": 99999999,
|
||||
"page": page,
|
||||
"offset": limit,
|
||||
"sort": "desc",
|
||||
"apikey": api_key
|
||||
}
|
||||
|
||||
try:
|
||||
async with aiohttp.ClientSession() as session:
|
||||
async with session.get(url, params=params, timeout=aiohttp.ClientTimeout(total=15)) as response:
|
||||
if response.status == 200:
|
||||
data = await response.json()
|
||||
if data.get("status") == "1":
|
||||
transactions = []
|
||||
for tx in data.get("result", []):
|
||||
value_wei = int(tx.get("value", 0))
|
||||
value_eth = value_wei / (10 ** config["decimals"])
|
||||
transactions.append({
|
||||
"hash": tx.get("hash"),
|
||||
"block": tx.get("blockNumber"),
|
||||
"timestamp": datetime.fromtimestamp(int(tx.get("timeStamp", 0))).isoformat(),
|
||||
"from": tx.get("from"),
|
||||
"to": tx.get("to"),
|
||||
"value": round(value_eth, 8),
|
||||
"symbol": config["symbol"],
|
||||
"gas_used": tx.get("gasUsed"),
|
||||
"gas_price": tx.get("gasPrice"),
|
||||
"is_error": tx.get("isError") == "1",
|
||||
"tx_url": f"{config['explorer_url']}/tx/{tx.get('hash')}"
|
||||
})
|
||||
return {
|
||||
"success": True,
|
||||
"address": address,
|
||||
"chain": chain,
|
||||
"transactions": transactions,
|
||||
"count": len(transactions),
|
||||
"page": page
|
||||
}
|
||||
else:
|
||||
return {"success": False, "error": data.get("message", "API 错误")}
|
||||
else:
|
||||
return {"success": False, "error": f"HTTP {response.status}"}
|
||||
except Exception as e:
|
||||
logger.error(f"获取交易失败: {e}")
|
||||
return {"success": False, "error": str(e)}
|
||||
|
||||
|
||||
async def fetch_token_balances(address: str, chain: str, api_key: str) -> Dict[str, Any]:
|
||||
"""获取 ERC20 代币余额"""
|
||||
if chain not in CHAIN_CONFIGS:
|
||||
return {"success": False, "error": f"不支持的区块链: {chain}"}
|
||||
|
||||
config = CHAIN_CONFIGS[chain]
|
||||
url = config["api_url"]
|
||||
|
||||
params = {
|
||||
"chainid": config["chainid"],
|
||||
"module": "account",
|
||||
"action": "tokentx",
|
||||
"address": address,
|
||||
"page": 1,
|
||||
"offset": 100,
|
||||
"sort": "desc",
|
||||
"apikey": api_key
|
||||
}
|
||||
|
||||
try:
|
||||
async with aiohttp.ClientSession() as session:
|
||||
async with session.get(url, params=params, timeout=aiohttp.ClientTimeout(total=15)) as response:
|
||||
if response.status == 200:
|
||||
data = await response.json()
|
||||
if data.get("status") == "1":
|
||||
# 统计代币
|
||||
token_map = {}
|
||||
for tx in data.get("result", []):
|
||||
contract = tx.get("contractAddress")
|
||||
if contract not in token_map:
|
||||
token_map[contract] = {
|
||||
"contract": contract,
|
||||
"name": tx.get("tokenName"),
|
||||
"symbol": tx.get("tokenSymbol"),
|
||||
"decimals": int(tx.get("tokenDecimal", 18)),
|
||||
"tx_count": 0
|
||||
}
|
||||
token_map[contract]["tx_count"] += 1
|
||||
|
||||
tokens = list(token_map.values())
|
||||
return {
|
||||
"success": True,
|
||||
"address": address,
|
||||
"chain": chain,
|
||||
"tokens": tokens,
|
||||
"token_count": len(tokens)
|
||||
}
|
||||
else:
|
||||
return {"success": True, "address": address, "chain": chain, "tokens": [], "token_count": 0}
|
||||
else:
|
||||
return {"success": False, "error": f"HTTP {response.status}"}
|
||||
except Exception as e:
|
||||
logger.error(f"获取代币失败: {e}")
|
||||
return {"success": False, "error": str(e)}
|
||||
|
||||
|
||||
async def chat_with_llm(message: str, context: str, api_key: str) -> str:
|
||||
"""调用 LLM 生成响应"""
|
||||
try:
|
||||
async with aiohttp.ClientSession() as session:
|
||||
payload = {
|
||||
"model": LLM_MODEL,
|
||||
"messages": [
|
||||
{
|
||||
"role": "system",
|
||||
"content": """你是一个专业的区块链数据分析师。你可以:
|
||||
1. 查询钱包地址的余额和交易记录
|
||||
2. 分析地址的链上活动
|
||||
3. 解答关于以太坊、BSC、Polygon等EVM链的问题
|
||||
|
||||
请根据提供的链上数据,用简洁专业的语言回答用户问题。"""
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": f"链上数据:\n{context}\n\n用户问题: {message}"
|
||||
}
|
||||
],
|
||||
"max_tokens": 500,
|
||||
"temperature": 0.7
|
||||
}
|
||||
headers = {
|
||||
"Authorization": f"Bearer {api_key}",
|
||||
"Content-Type": "application/json"
|
||||
}
|
||||
|
||||
async with session.post(
|
||||
f"{LLM_BASE_URL}/chat/completions",
|
||||
json=payload,
|
||||
headers=headers,
|
||||
timeout=aiohttp.ClientTimeout(total=30)
|
||||
) as response:
|
||||
if response.status == 200:
|
||||
data = await response.json()
|
||||
return data.get("choices", [{}])[0].get("message", {}).get("content", "抱歉,无法生成回复")
|
||||
else:
|
||||
error = await response.text()
|
||||
logger.error(f"LLM 请求失败: {response.status} - {error}")
|
||||
return f"LLM 服务错误: {response.status}"
|
||||
except Exception as e:
|
||||
logger.error(f"LLM 调用失败: {e}")
|
||||
return f"调用失败: {str(e)}"
|
||||
|
||||
|
||||
def extract_address_from_message(message: str) -> Optional[str]:
|
||||
"""从消息中提取以太坊地址"""
|
||||
import re
|
||||
# 匹配以太坊地址格式 (0x开头,40个十六进制字符)
|
||||
pattern = r'0x[a-fA-F0-9]{40}'
|
||||
match = re.search(pattern, message)
|
||||
return match.group(0) if match else None
|
||||
|
||||
|
||||
# ==================== API 端点 ====================
|
||||
|
||||
@app.get("/", response_model=dict)
|
||||
async def root():
|
||||
"""服务状态"""
|
||||
return {
|
||||
"service": "Chain Explorer Agent",
|
||||
"description": "链上数据查询 - 查询地址余额、交易记录、代币信息",
|
||||
"status": "running",
|
||||
"supported_chains": list(CHAIN_CONFIGS.keys()),
|
||||
"tools": ["balance", "transactions", "tokens", "chat"]
|
||||
}
|
||||
|
||||
|
||||
@app.get("/health", response_model=HealthResponse)
|
||||
async def health_check():
|
||||
"""健康检查"""
|
||||
return HealthResponse(
|
||||
status="healthy",
|
||||
pod_name=POD_NAME,
|
||||
supported_chains=list(CHAIN_CONFIGS.keys()),
|
||||
callback_enabled=CALLBACK_ENABLED,
|
||||
timestamp=datetime.utcnow().isoformat()
|
||||
)
|
||||
|
||||
|
||||
@app.post("/balance")
|
||||
async def get_balance(
|
||||
request: BalanceRequest,
|
||||
api_key: Optional[str] = Header(None, alias="api-key"),
|
||||
etherscan_key: Optional[str] = Header(None, alias="etherscan-key")
|
||||
):
|
||||
"""查询地址余额"""
|
||||
scan_key = etherscan_key or api_key
|
||||
if not scan_key:
|
||||
raise HTTPException(status_code=401, detail="请在请求头中提供 etherscan-key 或 api-key")
|
||||
|
||||
result = await fetch_balance(request.address, request.chain, scan_key)
|
||||
|
||||
if not result["success"]:
|
||||
raise HTTPException(status_code=400, detail=result["error"])
|
||||
|
||||
return {
|
||||
**result,
|
||||
"timestamp": datetime.utcnow().isoformat()
|
||||
}
|
||||
|
||||
|
||||
@app.post("/transactions")
|
||||
async def get_transactions(
|
||||
request: TransactionRequest,
|
||||
api_key: Optional[str] = Header(None, alias="api-key"),
|
||||
etherscan_key: Optional[str] = Header(None, alias="etherscan-key")
|
||||
):
|
||||
"""查询交易记录"""
|
||||
scan_key = etherscan_key or api_key
|
||||
if not scan_key:
|
||||
raise HTTPException(status_code=401, detail="请在请求头中提供 etherscan-key 或 api-key")
|
||||
|
||||
result = await fetch_transactions(request.address, request.chain, scan_key, request.page, request.limit)
|
||||
|
||||
if not result["success"]:
|
||||
raise HTTPException(status_code=400, detail=result["error"])
|
||||
|
||||
return {
|
||||
**result,
|
||||
"timestamp": datetime.utcnow().isoformat()
|
||||
}
|
||||
|
||||
|
||||
@app.post("/tokens")
|
||||
async def get_token_balances(
|
||||
request: TokenBalanceRequest,
|
||||
api_key: Optional[str] = Header(None, alias="api-key"),
|
||||
etherscan_key: Optional[str] = Header(None, alias="etherscan-key")
|
||||
):
|
||||
"""查询代币余额"""
|
||||
scan_key = etherscan_key or api_key
|
||||
if not scan_key:
|
||||
raise HTTPException(status_code=401, detail="请在请求头中提供 etherscan-key 或 api-key")
|
||||
|
||||
result = await fetch_token_balances(request.address, request.chain, scan_key)
|
||||
|
||||
if not result["success"]:
|
||||
raise HTTPException(status_code=400, detail=result["error"])
|
||||
|
||||
return {
|
||||
**result,
|
||||
"timestamp": datetime.utcnow().isoformat()
|
||||
}
|
||||
|
||||
|
||||
@app.post("/chat", response_model=ChatResponse)
|
||||
async def chat(
|
||||
request: ChatRequest,
|
||||
api_key: Optional[str] = Header(None, alias="api-key"),
|
||||
etherscan_key: Optional[str] = Header(None, alias="etherscan-key"),
|
||||
llm_key: Optional[str] = Header(None, alias="llm-key"),
|
||||
authorization: Optional[str] = Header(None)
|
||||
):
|
||||
"""智能对话 - 支持自然语言查询链上数据
|
||||
|
||||
api_key 通过请求头传递:
|
||||
- api-key: 区块链浏览器 API Key (Etherscan 等)
|
||||
- etherscan-key: Etherscan API Key (优先)
|
||||
- llm-key: LLM API Key (用于 AI 分析)
|
||||
- Authorization: Bearer LLM-API-Key
|
||||
"""
|
||||
# 获取区块链 API Key
|
||||
scan_key = etherscan_key or api_key
|
||||
if not scan_key:
|
||||
raise HTTPException(status_code=401, detail="请在请求头中提供 etherscan-key 或 api-key")
|
||||
|
||||
# 获取 LLM API Key
|
||||
llm_api_key = llm_key
|
||||
if not llm_api_key and authorization:
|
||||
if authorization.startswith("Bearer "):
|
||||
llm_api_key = authorization[7:]
|
||||
else:
|
||||
llm_api_key = authorization
|
||||
|
||||
if not llm_api_key:
|
||||
raise HTTPException(status_code=401, detail="请在请求头中提供 llm-key 或 Authorization 用于 AI 分析")
|
||||
|
||||
# 从消息中提取地址
|
||||
address = extract_address_from_message(request.message)
|
||||
|
||||
chain_data = {}
|
||||
if address:
|
||||
# 获取余额
|
||||
balance_result = await fetch_balance(address, request.chain, scan_key)
|
||||
if balance_result["success"]:
|
||||
chain_data["balance"] = balance_result
|
||||
|
||||
# 获取最近交易
|
||||
tx_result = await fetch_transactions(address, request.chain, scan_key, 1, 5)
|
||||
if tx_result["success"]:
|
||||
chain_data["recent_transactions"] = tx_result["transactions"][:5]
|
||||
|
||||
# 获取代币
|
||||
token_result = await fetch_token_balances(address, request.chain, scan_key)
|
||||
if token_result["success"]:
|
||||
chain_data["tokens"] = token_result["tokens"][:10]
|
||||
|
||||
# 构建上下文
|
||||
if chain_data:
|
||||
context_parts = []
|
||||
if "balance" in chain_data:
|
||||
b = chain_data["balance"]
|
||||
context_parts.append(f"地址: {b['address']}\n余额: {b['balance']} {b['symbol']} ({b['chain_name']})")
|
||||
if "recent_transactions" in chain_data:
|
||||
context_parts.append(f"最近交易数: {len(chain_data['recent_transactions'])}")
|
||||
for tx in chain_data["recent_transactions"][:3]:
|
||||
context_parts.append(f" - {tx['value']} {tx['symbol']} @ {tx['timestamp'][:10]}")
|
||||
if "tokens" in chain_data:
|
||||
context_parts.append(f"持有代币种类: {len(chain_data['tokens'])}")
|
||||
context = "\n".join(context_parts)
|
||||
else:
|
||||
context = "未检测到有效的钱包地址,请提供 0x 开头的以太坊地址"
|
||||
|
||||
# 调用 LLM 生成回复
|
||||
llm_response = await chat_with_llm(request.message, context, llm_api_key)
|
||||
|
||||
return ChatResponse(
|
||||
response=llm_response,
|
||||
data=chain_data if chain_data else {"detected_address": address},
|
||||
timestamp=datetime.utcnow().isoformat()
|
||||
)
|
||||
|
||||
|
||||
@app.get("/chains")
|
||||
async def list_chains():
|
||||
"""列出支持的区块链"""
|
||||
return {
|
||||
"chains": [
|
||||
{
|
||||
"id": chain_id,
|
||||
"name": config["name"],
|
||||
"symbol": config["symbol"],
|
||||
"explorer": config["explorer_url"]
|
||||
}
|
||||
for chain_id, config in CHAIN_CONFIGS.items()
|
||||
]
|
||||
}
|
||||
|
||||
|
||||
# ==================== 主入口 ====================
|
||||
|
||||
def main():
|
||||
"""主函数"""
|
||||
logger.info(f"启动 Chain Explorer Agent - {POD_NAME}")
|
||||
logger.info(f"支持的区块链: {list(CHAIN_CONFIGS.keys())}")
|
||||
logger.info(f"回调功能: {'已启用' if CALLBACK_ENABLED else '未启用'}")
|
||||
|
||||
uvicorn.run(app, host=SERVICE_HOST, port=SERVICE_PORT, log_level="info")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,6 @@
|
||||
fastapi>=0.104.0
|
||||
uvicorn>=0.24.0
|
||||
aiohttp>=3.9.0
|
||||
pydantic>=2.0.0
|
||||
python-multipart>=0.0.6
|
||||
httpx>=0.25.0
|
||||
@@ -25,7 +25,7 @@ COPY agents/jina_search_agent/jina_search_agent.py /app/
|
||||
# 环境变量
|
||||
ENV PYTHONUNBUFFERED=1
|
||||
ENV SERVICE_HOST=0.0.0.0
|
||||
ENV SERVICE_PORT=8080
|
||||
ENV SERVICE_PORT=8000
|
||||
|
||||
# 默认 Jina API Key (硬编码)
|
||||
ENV JINA_API_KEY=jina_e26dc30420a44a1e859216528065b203TkMRmsoz-FgMDQC5FZX9jr5oF2CI
|
||||
@@ -35,8 +35,8 @@ ENV AGENT_CALLBACK_URL=http://mcp-server:8002/api/v1/billing/agent-callback
|
||||
|
||||
# 健康检查
|
||||
HEALTHCHECK --interval=30s --timeout=10s --start-period=20s --retries=3 \
|
||||
CMD python3 -c "import urllib.request; urllib.request.urlopen('http://localhost:8080/health').read()" || exit 1
|
||||
CMD python3 -c "import urllib.request; urllib.request.urlopen('http://localhost:8000/health').read()" || exit 1
|
||||
|
||||
EXPOSE 8080
|
||||
EXPOSE 8000
|
||||
|
||||
CMD ["python3", "-u", "jina_search_agent.py"]
|
||||
|
||||
@@ -31,7 +31,7 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
# 环境变量
|
||||
SERVICE_HOST = os.getenv("SERVICE_HOST", "0.0.0.0")
|
||||
SERVICE_PORT = int(os.getenv("SERVICE_PORT", "8080"))
|
||||
SERVICE_PORT = int(os.getenv("SERVICE_PORT", "8000"))
|
||||
POD_NAME = os.getenv("POD_NAME", "jina-search-agent")
|
||||
USER_ID = os.getenv("USER_ID", "")
|
||||
JINA_API_KEY = os.getenv("JINA_API_KEY", "")
|
||||
|
||||
@@ -0,0 +1,39 @@
|
||||
FROM python:3.11-slim
|
||||
|
||||
WORKDIR /app
|
||||
|
||||
# 安装系统依赖
|
||||
RUN apt-get update && apt-get install -y \
|
||||
curl \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
|
||||
# 安装 Python 依赖
|
||||
RUN pip install --no-cache-dir \
|
||||
fastapi==0.109.0 \
|
||||
uvicorn[standard]==0.27.0 \
|
||||
pydantic==2.5.3 \
|
||||
requests>=2.31.0 \
|
||||
aiohttp>=3.9.0
|
||||
|
||||
# 复制 common 模块(回调工具)
|
||||
COPY common/agent_callback_utils.py /app/common/
|
||||
RUN touch /app/common/__init__.py
|
||||
|
||||
# 复制应用代码
|
||||
COPY agents/stock_analysis_agent/stock_analysis_agent.py /app/
|
||||
|
||||
# 环境变量
|
||||
ENV PYTHONUNBUFFERED=1
|
||||
ENV SERVICE_HOST=0.0.0.0
|
||||
ENV SERVICE_PORT=8080
|
||||
|
||||
# 回调配置
|
||||
ENV AGENT_CALLBACK_URL=http://mcp-server:8002/api/v1/billing/agent-callback
|
||||
|
||||
# 健康检查
|
||||
HEALTHCHECK --interval=30s --timeout=10s --start-period=20s --retries=3 \
|
||||
CMD python3 -c "import urllib.request; urllib.request.urlopen('http://localhost:8080/health').read()" || exit 1
|
||||
|
||||
EXPOSE 8080
|
||||
|
||||
CMD ["python3", "-u", "stock_analysis_agent.py"]
|
||||
@@ -0,0 +1,703 @@
|
||||
"""
|
||||
Stock Analysis Agent - 美股技术分析 Agent
|
||||
提供股票技术指标分析、趋势判断和投资建议
|
||||
"""
|
||||
import os
|
||||
import sys
|
||||
import logging
|
||||
import aiohttp
|
||||
from typing import Optional, List, Dict, Any
|
||||
from datetime import datetime, timedelta
|
||||
|
||||
from fastapi import FastAPI, HTTPException, Query, Header, Request
|
||||
from fastapi.middleware.cors import CORSMiddleware
|
||||
from pydantic import BaseModel, Field
|
||||
import uvicorn
|
||||
|
||||
# 添加 common 模块路径
|
||||
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
|
||||
# 导入回调工具
|
||||
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, format='%(asctime)s - %(name)s - %(levelname)s - %(message)s')
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# 环境变量
|
||||
SERVICE_HOST = os.getenv("SERVICE_HOST", "0.0.0.0")
|
||||
SERVICE_PORT = int(os.getenv("SERVICE_PORT", "8080"))
|
||||
POD_NAME = os.getenv("POD_NAME", "stock-analysis-agent")
|
||||
USER_ID = os.getenv("USER_ID", "")
|
||||
|
||||
# LLM 配置
|
||||
LLM_BASE_URL = os.getenv("LLM_BASE_URL", "https://litellm.graystone-fb459c5d.southeastasia.azurecontainerapps.io/v1")
|
||||
LLM_MODEL = os.getenv("LLM_MODEL", "taiji/gpt-4o-mini")
|
||||
|
||||
# FastAPI 应用
|
||||
app = FastAPI(
|
||||
title="Stock Analysis Agent",
|
||||
description="美股技术分析 - 提供技术指标、趋势分析和投资建议",
|
||||
version="1.0.0"
|
||||
)
|
||||
|
||||
app.add_middleware(
|
||||
CORSMiddleware,
|
||||
allow_origins=["*"],
|
||||
allow_credentials=True,
|
||||
allow_methods=["*"],
|
||||
allow_headers=["*"],
|
||||
)
|
||||
|
||||
# 回调处理器
|
||||
callback_handler: Optional[AgentCallbackHandler] = None
|
||||
|
||||
|
||||
# ==================== 请求/响应模型 ====================
|
||||
|
||||
class TechnicalIndicators(BaseModel):
|
||||
"""技术指标"""
|
||||
sma_20: Optional[float] = Field(None, description="20日均线")
|
||||
sma_50: Optional[float] = Field(None, description="50日均线")
|
||||
sma_200: Optional[float] = Field(None, description="200日均线")
|
||||
rsi_14: Optional[float] = Field(None, description="14日RSI")
|
||||
macd: Optional[float] = Field(None, description="MACD")
|
||||
macd_signal: Optional[float] = Field(None, description="MACD信号线")
|
||||
bollinger_upper: Optional[float] = Field(None, description="布林带上轨")
|
||||
bollinger_lower: Optional[float] = Field(None, description="布林带下轨")
|
||||
volume_avg_20: Optional[float] = Field(None, description="20日平均成交量")
|
||||
|
||||
|
||||
class AnalysisRequest(BaseModel):
|
||||
"""分析请求"""
|
||||
symbol: str = Field(..., description="股票代码")
|
||||
user_id: Optional[str] = Field(None, description="用户ID(用于计费回调)")
|
||||
|
||||
|
||||
class AnalysisResponse(BaseModel):
|
||||
"""分析响应"""
|
||||
symbol: str
|
||||
current_price: float
|
||||
indicators: TechnicalIndicators
|
||||
trend: str # bullish, bearish, neutral
|
||||
signal: str # buy, sell, hold
|
||||
support_level: float
|
||||
resistance_level: float
|
||||
analysis_summary: str
|
||||
risk_level: str # low, medium, high
|
||||
timestamp: str
|
||||
|
||||
|
||||
class CompareRequest(BaseModel):
|
||||
"""对比分析请求"""
|
||||
symbols: List[str] = Field(..., description="股票代码列表(最多5个)")
|
||||
user_id: Optional[str] = Field(None, description="用户ID")
|
||||
|
||||
|
||||
class StockComparison(BaseModel):
|
||||
"""股票对比"""
|
||||
symbol: str
|
||||
price: float
|
||||
change_percent: float
|
||||
pe_ratio: Optional[float] = None
|
||||
market_cap: Optional[float] = None
|
||||
trend: str
|
||||
recommendation: str
|
||||
|
||||
|
||||
class CompareResponse(BaseModel):
|
||||
"""对比分析响应"""
|
||||
comparisons: List[StockComparison]
|
||||
best_pick: str
|
||||
analysis: str
|
||||
timestamp: str
|
||||
|
||||
|
||||
class HealthResponse(BaseModel):
|
||||
"""健康检查响应"""
|
||||
status: str
|
||||
pod_name: str
|
||||
callback_enabled: bool
|
||||
timestamp: str
|
||||
|
||||
|
||||
class ChatRequest(BaseModel):
|
||||
"""Chat 请求"""
|
||||
message: str = Field(..., description="用户消息")
|
||||
user_id: Optional[str] = Field(None, description="用户ID")
|
||||
|
||||
|
||||
class ChatResponse(BaseModel):
|
||||
"""Chat 响应"""
|
||||
response: str
|
||||
data: Optional[Dict[str, Any]] = None
|
||||
timestamp: str
|
||||
|
||||
|
||||
# ==================== 生命周期 ====================
|
||||
|
||||
@app.on_event("startup")
|
||||
async def startup_event():
|
||||
"""应用启动时初始化回调处理器"""
|
||||
global callback_handler
|
||||
|
||||
if CALLBACK_ENABLED:
|
||||
callback_handler = AgentCallbackHandler(
|
||||
agent_name=POD_NAME,
|
||||
user_id=USER_ID
|
||||
)
|
||||
logger.info(f"回调处理器已初始化: callback_url={callback_handler.callback_url}")
|
||||
else:
|
||||
logger.warning("回调模块未加载,计费回调功能不可用")
|
||||
|
||||
|
||||
# ==================== 辅助函数 ====================
|
||||
|
||||
async def fetch_historical_data(symbol: str, period: str = "3mo") -> List[Dict[str, Any]]:
|
||||
"""获取历史数据"""
|
||||
url = f"https://query1.finance.yahoo.com/v8/finance/chart/{symbol}"
|
||||
params = {
|
||||
"interval": "1d",
|
||||
"range": period
|
||||
}
|
||||
|
||||
headers = {
|
||||
"User-Agent": "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36"
|
||||
}
|
||||
|
||||
try:
|
||||
async with aiohttp.ClientSession() as session:
|
||||
async with session.get(url, params=params, headers=headers, timeout=15) as response:
|
||||
if response.status == 200:
|
||||
data = await response.json()
|
||||
result = data.get("chart", {}).get("result", [])
|
||||
|
||||
if not result:
|
||||
return []
|
||||
|
||||
quote_data = result[0]
|
||||
timestamps = quote_data.get("timestamp", [])
|
||||
indicators = quote_data.get("indicators", {}).get("quote", [{}])[0]
|
||||
|
||||
prices = []
|
||||
closes = indicators.get("close", [])
|
||||
highs = indicators.get("high", [])
|
||||
lows = indicators.get("low", [])
|
||||
volumes = indicators.get("volume", [])
|
||||
|
||||
for i, ts in enumerate(timestamps):
|
||||
if closes[i] is not None:
|
||||
prices.append({
|
||||
"date": datetime.fromtimestamp(ts).isoformat(),
|
||||
"close": closes[i],
|
||||
"high": highs[i] if i < len(highs) else None,
|
||||
"low": lows[i] if i < len(lows) else None,
|
||||
"volume": volumes[i] if i < len(volumes) else None
|
||||
})
|
||||
|
||||
return prices
|
||||
except Exception as e:
|
||||
logger.error(f"获取历史数据失败: {symbol} - {e}")
|
||||
|
||||
return []
|
||||
|
||||
|
||||
def calculate_sma(prices: List[float], period: int) -> Optional[float]:
|
||||
"""计算简单移动平均线"""
|
||||
if len(prices) < period:
|
||||
return None
|
||||
return sum(prices[-period:]) / period
|
||||
|
||||
|
||||
def calculate_rsi(prices: List[float], period: int = 14) -> Optional[float]:
|
||||
"""计算相对强弱指标 RSI"""
|
||||
if len(prices) < period + 1:
|
||||
return None
|
||||
|
||||
gains = []
|
||||
losses = []
|
||||
|
||||
for i in range(1, len(prices)):
|
||||
change = prices[i] - prices[i-1]
|
||||
if change > 0:
|
||||
gains.append(change)
|
||||
losses.append(0)
|
||||
else:
|
||||
gains.append(0)
|
||||
losses.append(abs(change))
|
||||
|
||||
if len(gains) < period:
|
||||
return None
|
||||
|
||||
avg_gain = sum(gains[-period:]) / period
|
||||
avg_loss = sum(losses[-period:]) / period
|
||||
|
||||
if avg_loss == 0:
|
||||
return 100
|
||||
|
||||
rs = avg_gain / avg_loss
|
||||
rsi = 100 - (100 / (1 + rs))
|
||||
|
||||
return round(rsi, 2)
|
||||
|
||||
|
||||
def calculate_macd(prices: List[float]) -> Dict[str, Optional[float]]:
|
||||
"""计算 MACD"""
|
||||
if len(prices) < 26:
|
||||
return {"macd": None, "signal": None}
|
||||
|
||||
# EMA 12
|
||||
ema_12 = prices[-12:]
|
||||
ema_12_val = sum(ema_12) / 12
|
||||
|
||||
# EMA 26
|
||||
ema_26 = prices[-26:]
|
||||
ema_26_val = sum(ema_26) / 26
|
||||
|
||||
macd = ema_12_val - ema_26_val
|
||||
signal = macd * 0.9 # 简化计算
|
||||
|
||||
return {"macd": round(macd, 4), "signal": round(signal, 4)}
|
||||
|
||||
|
||||
def calculate_bollinger_bands(prices: List[float], period: int = 20) -> Dict[str, Optional[float]]:
|
||||
"""计算布林带"""
|
||||
if len(prices) < period:
|
||||
return {"upper": None, "lower": None}
|
||||
|
||||
sma = sum(prices[-period:]) / period
|
||||
|
||||
# 计算标准差
|
||||
squared_diff = sum((p - sma) ** 2 for p in prices[-period:])
|
||||
std_dev = (squared_diff / period) ** 0.5
|
||||
|
||||
return {
|
||||
"upper": round(sma + 2 * std_dev, 2),
|
||||
"lower": round(sma - 2 * std_dev, 2)
|
||||
}
|
||||
|
||||
|
||||
async def analyze_stock(symbol: str) -> Dict[str, Any]:
|
||||
"""分析股票"""
|
||||
historical = await fetch_historical_data(symbol, "3mo")
|
||||
|
||||
if not historical:
|
||||
# 返回模拟数据
|
||||
return generate_mock_analysis(symbol)
|
||||
|
||||
closes = [p["close"] for p in historical if p["close"]]
|
||||
volumes = [p["volume"] for p in historical if p["volume"]]
|
||||
|
||||
current_price = closes[-1] if closes else 100
|
||||
|
||||
# 计算技术指标
|
||||
sma_20 = calculate_sma(closes, 20)
|
||||
sma_50 = calculate_sma(closes, 50)
|
||||
sma_200 = calculate_sma(closes, 200) if len(closes) >= 200 else None
|
||||
rsi = calculate_rsi(closes, 14)
|
||||
macd_data = calculate_macd(closes)
|
||||
bollinger = calculate_bollinger_bands(closes, 20)
|
||||
volume_avg = sum(volumes[-20:]) / 20 if len(volumes) >= 20 else None
|
||||
|
||||
# 确定趋势
|
||||
trend = "neutral"
|
||||
if sma_20 and sma_50:
|
||||
if current_price > sma_20 > sma_50:
|
||||
trend = "bullish"
|
||||
elif current_price < sma_20 < sma_50:
|
||||
trend = "bearish"
|
||||
|
||||
# 确定信号
|
||||
signal = "hold"
|
||||
if rsi:
|
||||
if rsi < 30 and trend != "bearish":
|
||||
signal = "buy"
|
||||
elif rsi > 70 and trend != "bullish":
|
||||
signal = "sell"
|
||||
elif trend == "bullish" and current_price > sma_20:
|
||||
signal = "buy"
|
||||
elif trend == "bearish" and current_price < sma_20:
|
||||
signal = "sell"
|
||||
|
||||
# 支撑位和阻力位
|
||||
recent_lows = [p["low"] for p in historical[-20:] if p["low"]]
|
||||
recent_highs = [p["high"] for p in historical[-20:] if p["high"]]
|
||||
|
||||
support = min(recent_lows) if recent_lows else current_price * 0.95
|
||||
resistance = max(recent_highs) if recent_highs else current_price * 1.05
|
||||
|
||||
# 风险评估
|
||||
if rsi and (rsi < 20 or rsi > 80):
|
||||
risk_level = "high"
|
||||
elif trend == "neutral":
|
||||
risk_level = "medium"
|
||||
else:
|
||||
risk_level = "low"
|
||||
|
||||
# 生成分析摘要
|
||||
summary = generate_analysis_summary(symbol, current_price, trend, signal, rsi, sma_20, sma_50)
|
||||
|
||||
return {
|
||||
"symbol": symbol.upper(),
|
||||
"current_price": round(current_price, 2),
|
||||
"indicators": {
|
||||
"sma_20": round(sma_20, 2) if sma_20 else None,
|
||||
"sma_50": round(sma_50, 2) if sma_50 else None,
|
||||
"sma_200": round(sma_200, 2) if sma_200 else None,
|
||||
"rsi_14": rsi,
|
||||
"macd": macd_data["macd"],
|
||||
"macd_signal": macd_data["signal"],
|
||||
"bollinger_upper": bollinger["upper"],
|
||||
"bollinger_lower": bollinger["lower"],
|
||||
"volume_avg_20": int(volume_avg) if volume_avg else None
|
||||
},
|
||||
"trend": trend,
|
||||
"signal": signal,
|
||||
"support_level": round(support, 2),
|
||||
"resistance_level": round(resistance, 2),
|
||||
"analysis_summary": summary,
|
||||
"risk_level": risk_level
|
||||
}
|
||||
|
||||
|
||||
def generate_mock_analysis(symbol: str) -> Dict[str, Any]:
|
||||
"""生成模拟分析数据"""
|
||||
import random
|
||||
price = random.uniform(50, 500)
|
||||
|
||||
return {
|
||||
"symbol": symbol.upper(),
|
||||
"current_price": round(price, 2),
|
||||
"indicators": {
|
||||
"sma_20": round(price * 0.98, 2),
|
||||
"sma_50": round(price * 0.95, 2),
|
||||
"sma_200": round(price * 0.90, 2),
|
||||
"rsi_14": random.uniform(30, 70),
|
||||
"macd": random.uniform(-2, 2),
|
||||
"macd_signal": random.uniform(-1.5, 1.5),
|
||||
"bollinger_upper": round(price * 1.05, 2),
|
||||
"bollinger_lower": round(price * 0.95, 2),
|
||||
"volume_avg_20": random.randint(10000000, 100000000)
|
||||
},
|
||||
"trend": random.choice(["bullish", "bearish", "neutral"]),
|
||||
"signal": random.choice(["buy", "sell", "hold"]),
|
||||
"support_level": round(price * 0.93, 2),
|
||||
"resistance_level": round(price * 1.07, 2),
|
||||
"analysis_summary": f"{symbol.upper()} is showing mixed signals. Monitor closely for breakout opportunities.",
|
||||
"risk_level": random.choice(["low", "medium", "high"])
|
||||
}
|
||||
|
||||
|
||||
def generate_analysis_summary(symbol: str, price: float, trend: str, signal: str,
|
||||
rsi: Optional[float], sma_20: Optional[float], sma_50: Optional[float]) -> str:
|
||||
"""生成分析摘要"""
|
||||
summary_parts = [f"{symbol.upper()} is currently trading at ${price:.2f}."]
|
||||
|
||||
if trend == "bullish":
|
||||
summary_parts.append("The stock shows a bullish trend with price above key moving averages.")
|
||||
elif trend == "bearish":
|
||||
summary_parts.append("The stock is in a bearish trend, trading below key moving averages.")
|
||||
else:
|
||||
summary_parts.append("The stock is consolidating with no clear directional bias.")
|
||||
|
||||
if rsi:
|
||||
if rsi < 30:
|
||||
summary_parts.append(f"RSI at {rsi:.1f} indicates oversold conditions - potential buying opportunity.")
|
||||
elif rsi > 70:
|
||||
summary_parts.append(f"RSI at {rsi:.1f} indicates overbought conditions - caution advised.")
|
||||
else:
|
||||
summary_parts.append(f"RSI at {rsi:.1f} is in neutral territory.")
|
||||
|
||||
if signal == "buy":
|
||||
summary_parts.append("Technical signals suggest a buying opportunity.")
|
||||
elif signal == "sell":
|
||||
summary_parts.append("Technical signals suggest considering profit-taking.")
|
||||
else:
|
||||
summary_parts.append("Recommend holding current positions and monitoring for clearer signals.")
|
||||
|
||||
return " ".join(summary_parts)
|
||||
|
||||
|
||||
# ==================== API 端点 ====================
|
||||
|
||||
@app.get("/health", response_model=HealthResponse)
|
||||
@app.get("/", response_model=HealthResponse)
|
||||
async def health_check():
|
||||
"""健康检查"""
|
||||
return HealthResponse(
|
||||
status="healthy",
|
||||
pod_name=POD_NAME,
|
||||
callback_enabled=CALLBACK_ENABLED,
|
||||
timestamp=datetime.utcnow().isoformat()
|
||||
)
|
||||
|
||||
|
||||
@app.post("/analyze", response_model=AnalysisResponse)
|
||||
async def analyze(request: AnalysisRequest):
|
||||
"""分析单个股票"""
|
||||
if CALLBACK_ENABLED and callback_handler and request.user_id:
|
||||
with CallbackContextManager(
|
||||
handler=callback_handler,
|
||||
user_id=request.user_id,
|
||||
request_id=f"stock-analysis-{int(datetime.utcnow().timestamp())}"
|
||||
) as ctx:
|
||||
ctx.add_tool("stock_analysis")
|
||||
ctx.add_tool("technical_indicators")
|
||||
|
||||
result = await analyze_stock(request.symbol)
|
||||
|
||||
return AnalysisResponse(
|
||||
symbol=result["symbol"],
|
||||
current_price=result["current_price"],
|
||||
indicators=TechnicalIndicators(**result["indicators"]),
|
||||
trend=result["trend"],
|
||||
signal=result["signal"],
|
||||
support_level=result["support_level"],
|
||||
resistance_level=result["resistance_level"],
|
||||
analysis_summary=result["analysis_summary"],
|
||||
risk_level=result["risk_level"],
|
||||
timestamp=datetime.utcnow().isoformat()
|
||||
)
|
||||
else:
|
||||
result = await analyze_stock(request.symbol)
|
||||
|
||||
return AnalysisResponse(
|
||||
symbol=result["symbol"],
|
||||
current_price=result["current_price"],
|
||||
indicators=TechnicalIndicators(**result["indicators"]),
|
||||
trend=result["trend"],
|
||||
signal=result["signal"],
|
||||
support_level=result["support_level"],
|
||||
resistance_level=result["resistance_level"],
|
||||
analysis_summary=result["analysis_summary"],
|
||||
risk_level=result["risk_level"],
|
||||
timestamp=datetime.utcnow().isoformat()
|
||||
)
|
||||
|
||||
|
||||
@app.get("/analyze")
|
||||
async def analyze_get(
|
||||
symbol: str = Query(..., description="股票代码"),
|
||||
user_id: Optional[str] = Query(None, description="用户ID")
|
||||
):
|
||||
"""GET 方式分析股票"""
|
||||
request = AnalysisRequest(symbol=symbol, user_id=user_id)
|
||||
return await analyze(request)
|
||||
|
||||
|
||||
@app.post("/compare", response_model=CompareResponse)
|
||||
async def compare_stocks(request: CompareRequest):
|
||||
"""对比多个股票"""
|
||||
if len(request.symbols) > 5:
|
||||
raise HTTPException(status_code=400, detail="最多支持5个股票对比")
|
||||
|
||||
comparisons = []
|
||||
best_score = -1
|
||||
best_pick = ""
|
||||
|
||||
for symbol in request.symbols:
|
||||
result = await analyze_stock(symbol)
|
||||
|
||||
# 计算简单评分
|
||||
score = 0
|
||||
if result["trend"] == "bullish":
|
||||
score += 2
|
||||
elif result["trend"] == "neutral":
|
||||
score += 1
|
||||
|
||||
if result["signal"] == "buy":
|
||||
score += 2
|
||||
elif result["signal"] == "hold":
|
||||
score += 1
|
||||
|
||||
if result["risk_level"] == "low":
|
||||
score += 2
|
||||
elif result["risk_level"] == "medium":
|
||||
score += 1
|
||||
|
||||
if score > best_score:
|
||||
best_score = score
|
||||
best_pick = symbol
|
||||
|
||||
comparisons.append(StockComparison(
|
||||
symbol=result["symbol"],
|
||||
price=result["current_price"],
|
||||
change_percent=0, # 需要额外计算
|
||||
pe_ratio=None,
|
||||
market_cap=None,
|
||||
trend=result["trend"],
|
||||
recommendation=result["signal"]
|
||||
))
|
||||
|
||||
analysis = f"Based on technical analysis, {best_pick.upper()} shows the strongest signals among the compared stocks."
|
||||
|
||||
return CompareResponse(
|
||||
comparisons=comparisons,
|
||||
best_pick=best_pick.upper(),
|
||||
analysis=analysis,
|
||||
timestamp=datetime.utcnow().isoformat()
|
||||
)
|
||||
|
||||
|
||||
@app.get("/screener")
|
||||
async def stock_screener(
|
||||
trend: Optional[str] = Query(None, description="筛选趋势: bullish, bearish, neutral"),
|
||||
signal: Optional[str] = Query(None, description="筛选信号: buy, sell, hold")
|
||||
):
|
||||
"""股票筛选器"""
|
||||
# 分析一组热门股票
|
||||
popular = ["AAPL", "MSFT", "GOOGL", "AMZN", "TSLA", "NVDA", "META", "AMD", "NFLX", "DIS"]
|
||||
|
||||
results = []
|
||||
for symbol in popular:
|
||||
analysis = await analyze_stock(symbol)
|
||||
|
||||
# 应用筛选条件
|
||||
if trend and analysis["trend"] != trend:
|
||||
continue
|
||||
if signal and analysis["signal"] != signal:
|
||||
continue
|
||||
|
||||
results.append({
|
||||
"symbol": analysis["symbol"],
|
||||
"price": analysis["current_price"],
|
||||
"trend": analysis["trend"],
|
||||
"signal": analysis["signal"],
|
||||
"risk_level": analysis["risk_level"]
|
||||
})
|
||||
|
||||
return {
|
||||
"filters": {"trend": trend, "signal": signal},
|
||||
"results": results,
|
||||
"count": len(results),
|
||||
"timestamp": datetime.utcnow().isoformat()
|
||||
}
|
||||
|
||||
|
||||
# ==================== Chat 功能 ====================
|
||||
|
||||
async def chat_with_llm(message: str, context: str, api_key: str) -> str:
|
||||
"""调用 LLM 生成响应"""
|
||||
try:
|
||||
async with aiohttp.ClientSession() as session:
|
||||
payload = {
|
||||
"model": LLM_MODEL,
|
||||
"messages": [
|
||||
{
|
||||
"role": "system",
|
||||
"content": """你是一个专业的美股技术分析师。你可以:
|
||||
1. 分析股票技术指标(SMA, RSI, MACD, 布林带等)
|
||||
2. 判断股票趋势(看涨/看跌/中性)
|
||||
3. 提供买卖信号和投资建议
|
||||
4. 评估风险等级
|
||||
|
||||
请根据提供的技术分析数据,用简洁专业的语言回答用户问题。注意:投资有风险,建议仅供参考。"""
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": f"技术分析数据:\n{context}\n\n用户问题: {message}"
|
||||
}
|
||||
],
|
||||
"max_tokens": 600,
|
||||
"temperature": 0.7
|
||||
}
|
||||
headers = {
|
||||
"Authorization": f"Bearer {api_key}",
|
||||
"Content-Type": "application/json"
|
||||
}
|
||||
|
||||
async with session.post(
|
||||
f"{LLM_BASE_URL}/chat/completions",
|
||||
json=payload,
|
||||
headers=headers,
|
||||
timeout=aiohttp.ClientTimeout(total=30)
|
||||
) as response:
|
||||
if response.status == 200:
|
||||
data = await response.json()
|
||||
return data.get("choices", [{}])[0].get("message", {}).get("content", "抱歉,无法生成回复")
|
||||
else:
|
||||
error = await response.text()
|
||||
logger.error(f"LLM 请求失败: {response.status} - {error}")
|
||||
return f"LLM 服务错误: {response.status}"
|
||||
except Exception as e:
|
||||
logger.error(f"LLM 调用失败: {e}")
|
||||
return f"调用失败: {str(e)}"
|
||||
|
||||
|
||||
@app.post("/chat", response_model=ChatResponse)
|
||||
async def chat(
|
||||
request: ChatRequest,
|
||||
api_key: Optional[str] = Header(None, alias="api-key"),
|
||||
authorization: Optional[str] = Header(None)
|
||||
):
|
||||
"""智能对话 - 获取技术分析并提供投资建议
|
||||
|
||||
api_key 通过请求头传递:
|
||||
- api-key: your-api-key
|
||||
- 或 Authorization: Bearer your-api-key
|
||||
"""
|
||||
# 从 Header 获取 api_key
|
||||
if not api_key and authorization:
|
||||
if authorization.startswith("Bearer "):
|
||||
api_key = authorization[7:]
|
||||
else:
|
||||
api_key = authorization
|
||||
|
||||
if not api_key:
|
||||
raise HTTPException(status_code=401, detail="请在请求头中提供 api-key 或 Authorization")
|
||||
|
||||
# 从消息中提取股票代码
|
||||
import re
|
||||
symbols = re.findall(r'\b([A-Z]{1,5})\b', request.message.upper())
|
||||
common_words = {"I", "A", "THE", "IS", "IT", "TO", "OF", "AND", "FOR", "IN", "ON", "AT", "BY", "BUY", "SELL"}
|
||||
symbols = [s for s in symbols if s not in common_words][:3]
|
||||
|
||||
if not symbols:
|
||||
symbols = ["AAPL"] # 默认分析苹果
|
||||
|
||||
# 获取技术分析数据
|
||||
analysis_data = []
|
||||
for symbol in symbols:
|
||||
analysis = await analyze_stock(symbol)
|
||||
if analysis:
|
||||
analysis_data.append(analysis)
|
||||
|
||||
# 构建上下文
|
||||
if analysis_data:
|
||||
context = "\n".join([
|
||||
f"{a['symbol']}: 价格${a['current_price']:.2f}, 趋势:{a['trend']}, "
|
||||
f"信号:{a['signal']}, RSI:{a['indicators'].get('rsi_14', 'N/A')}, "
|
||||
f"风险:{a['risk_level']}"
|
||||
for a in analysis_data
|
||||
])
|
||||
else:
|
||||
context = "暂无技术分析数据"
|
||||
|
||||
# 调用 LLM 生成回复
|
||||
llm_response = await chat_with_llm(request.message, context, api_key)
|
||||
|
||||
return ChatResponse(
|
||||
response=llm_response,
|
||||
data={"analysis": analysis_data, "symbols": symbols},
|
||||
timestamp=datetime.utcnow().isoformat()
|
||||
)
|
||||
|
||||
|
||||
# ==================== 主入口 ====================
|
||||
|
||||
def main():
|
||||
"""主函数"""
|
||||
logger.info(f"启动 Stock Analysis Agent - {POD_NAME}")
|
||||
logger.info(f"回调功能: {'已启用' if CALLBACK_ENABLED else '未启用'}")
|
||||
|
||||
uvicorn.run(app, host=SERVICE_HOST, port=SERVICE_PORT, log_level="info")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,39 @@
|
||||
FROM python:3.11-slim
|
||||
|
||||
WORKDIR /app
|
||||
|
||||
# 安装系统依赖
|
||||
RUN apt-get update && apt-get install -y \
|
||||
curl \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
|
||||
# 安装 Python 依赖
|
||||
RUN pip install --no-cache-dir \
|
||||
fastapi==0.109.0 \
|
||||
uvicorn[standard]==0.27.0 \
|
||||
pydantic==2.5.3 \
|
||||
requests>=2.31.0 \
|
||||
aiohttp>=3.9.0
|
||||
|
||||
# 复制 common 模块(回调工具)
|
||||
COPY common/agent_callback_utils.py /app/common/
|
||||
RUN touch /app/common/__init__.py
|
||||
|
||||
# 复制应用代码
|
||||
COPY agents/stock_news_agent/stock_news_agent.py /app/
|
||||
|
||||
# 环境变量
|
||||
ENV PYTHONUNBUFFERED=1
|
||||
ENV SERVICE_HOST=0.0.0.0
|
||||
ENV SERVICE_PORT=8080
|
||||
|
||||
# 回调配置
|
||||
ENV AGENT_CALLBACK_URL=http://mcp-server:8002/api/v1/billing/agent-callback
|
||||
|
||||
# 健康检查
|
||||
HEALTHCHECK --interval=30s --timeout=10s --start-period=20s --retries=3 \
|
||||
CMD python3 -c "import urllib.request; urllib.request.urlopen('http://localhost:8080/health').read()" || exit 1
|
||||
|
||||
EXPOSE 8080
|
||||
|
||||
CMD ["python3", "-u", "stock_news_agent.py"]
|
||||
@@ -0,0 +1,527 @@
|
||||
"""
|
||||
Stock News Agent - 美股新闻资讯 Agent
|
||||
获取美股相关新闻、市场动态和公司公告
|
||||
"""
|
||||
import os
|
||||
import sys
|
||||
import logging
|
||||
import aiohttp
|
||||
from typing import Optional, List, Dict, Any
|
||||
from datetime import datetime, timedelta
|
||||
|
||||
from fastapi import FastAPI, HTTPException, Query, Header, Request
|
||||
from fastapi.middleware.cors import CORSMiddleware
|
||||
from pydantic import BaseModel, Field
|
||||
import uvicorn
|
||||
|
||||
# 添加 common 模块路径
|
||||
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
|
||||
# 导入回调工具
|
||||
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, format='%(asctime)s - %(name)s - %(levelname)s - %(message)s')
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# 环境变量
|
||||
SERVICE_HOST = os.getenv("SERVICE_HOST", "0.0.0.0")
|
||||
SERVICE_PORT = int(os.getenv("SERVICE_PORT", "8080"))
|
||||
POD_NAME = os.getenv("POD_NAME", "stock-news-agent")
|
||||
USER_ID = os.getenv("USER_ID", "")
|
||||
|
||||
# News API (可选)
|
||||
NEWS_API_KEY = os.getenv("NEWS_API_KEY", "")
|
||||
|
||||
# LLM 配置
|
||||
LLM_BASE_URL = os.getenv("LLM_BASE_URL", "https://litellm.graystone-fb459c5d.southeastasia.azurecontainerapps.io/v1")
|
||||
LLM_MODEL = os.getenv("LLM_MODEL", "taiji/gpt-4o-mini")
|
||||
|
||||
# FastAPI 应用
|
||||
app = FastAPI(
|
||||
title="Stock News Agent",
|
||||
description="美股新闻资讯 - 获取股票相关新闻、市场动态和分析报告",
|
||||
version="1.0.0"
|
||||
)
|
||||
|
||||
app.add_middleware(
|
||||
CORSMiddleware,
|
||||
allow_origins=["*"],
|
||||
allow_credentials=True,
|
||||
allow_methods=["*"],
|
||||
allow_headers=["*"],
|
||||
)
|
||||
|
||||
# 回调处理器
|
||||
callback_handler: Optional[AgentCallbackHandler] = None
|
||||
|
||||
|
||||
# ==================== 请求/响应模型 ====================
|
||||
|
||||
class NewsItem(BaseModel):
|
||||
"""新闻条目"""
|
||||
title: str
|
||||
description: Optional[str] = None
|
||||
url: str
|
||||
source: str
|
||||
published_at: str
|
||||
sentiment: Optional[str] = None # positive, negative, neutral
|
||||
|
||||
|
||||
class NewsRequest(BaseModel):
|
||||
"""新闻查询请求"""
|
||||
symbol: Optional[str] = Field(None, description="股票代码,如 AAPL")
|
||||
query: Optional[str] = Field(None, description="搜索关键词")
|
||||
limit: int = Field(10, ge=1, le=50, description="返回新闻数量")
|
||||
user_id: Optional[str] = Field(None, description="用户ID(用于计费回调)")
|
||||
|
||||
|
||||
class NewsResponse(BaseModel):
|
||||
"""新闻查询响应"""
|
||||
symbol: Optional[str] = None
|
||||
query: Optional[str] = None
|
||||
news: List[NewsItem]
|
||||
total_count: int
|
||||
timestamp: str
|
||||
|
||||
|
||||
class MarketSummaryResponse(BaseModel):
|
||||
"""市场概要响应"""
|
||||
market_status: str
|
||||
top_gainers: List[Dict[str, Any]]
|
||||
top_losers: List[Dict[str, Any]]
|
||||
most_active: List[Dict[str, Any]]
|
||||
timestamp: str
|
||||
|
||||
|
||||
class HealthResponse(BaseModel):
|
||||
"""健康检查响应"""
|
||||
status: str
|
||||
pod_name: str
|
||||
news_api_configured: bool
|
||||
callback_enabled: bool
|
||||
timestamp: str
|
||||
|
||||
|
||||
class ChatRequest(BaseModel):
|
||||
"""Chat 请求"""
|
||||
message: str = Field(..., description="用户消息")
|
||||
user_id: Optional[str] = Field(None, description="用户ID")
|
||||
|
||||
|
||||
class ChatResponse(BaseModel):
|
||||
"""Chat 响应"""
|
||||
response: str
|
||||
data: Optional[Dict[str, Any]] = None
|
||||
timestamp: str
|
||||
|
||||
|
||||
# ==================== 生命周期 ====================
|
||||
|
||||
@app.on_event("startup")
|
||||
async def startup_event():
|
||||
"""应用启动时初始化回调处理器"""
|
||||
global callback_handler
|
||||
|
||||
if CALLBACK_ENABLED:
|
||||
callback_handler = AgentCallbackHandler(
|
||||
agent_name=POD_NAME,
|
||||
user_id=USER_ID
|
||||
)
|
||||
logger.info(f"回调处理器已初始化: callback_url={callback_handler.callback_url}")
|
||||
else:
|
||||
logger.warning("回调模块未加载,计费回调功能不可用")
|
||||
|
||||
|
||||
# ==================== 辅助函数 ====================
|
||||
|
||||
async def fetch_yahoo_news(symbol: str = None, query: str = None, limit: int = 10) -> List[Dict[str, Any]]:
|
||||
"""从 Yahoo Finance 获取新闻"""
|
||||
news_list = []
|
||||
|
||||
# 构建搜索词
|
||||
search_term = symbol if symbol else (query if query else "stock market")
|
||||
|
||||
# Yahoo Finance RSS 新闻源
|
||||
url = f"https://query1.finance.yahoo.com/v1/finance/search"
|
||||
params = {
|
||||
"q": search_term,
|
||||
"newsCount": limit,
|
||||
"enableFuzzyQuery": False,
|
||||
"quotesQueryId": "tss_match_phrase_query"
|
||||
}
|
||||
|
||||
headers = {
|
||||
"User-Agent": "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36"
|
||||
}
|
||||
|
||||
try:
|
||||
async with aiohttp.ClientSession() as session:
|
||||
async with session.get(url, params=params, headers=headers, timeout=15) as response:
|
||||
if response.status == 200:
|
||||
data = await response.json()
|
||||
news_data = data.get("news", [])
|
||||
|
||||
for item in news_data[:limit]:
|
||||
news_list.append({
|
||||
"title": item.get("title", ""),
|
||||
"description": item.get("summary", ""),
|
||||
"url": item.get("link", ""),
|
||||
"source": item.get("publisher", "Yahoo Finance"),
|
||||
"published_at": datetime.fromtimestamp(
|
||||
item.get("providerPublishTime", datetime.now().timestamp())
|
||||
).isoformat(),
|
||||
"sentiment": analyze_sentiment(item.get("title", "") + " " + item.get("summary", ""))
|
||||
})
|
||||
except Exception as e:
|
||||
logger.error(f"获取新闻失败: {e}")
|
||||
|
||||
# 如果没有获取到新闻,返回模拟数据
|
||||
if not news_list:
|
||||
news_list = generate_sample_news(symbol or query or "market", limit)
|
||||
|
||||
return news_list
|
||||
|
||||
|
||||
def analyze_sentiment(text: str) -> str:
|
||||
"""简单的情感分析"""
|
||||
positive_words = ["surge", "gain", "rise", "up", "bullish", "growth", "profit", "beat", "record", "high"]
|
||||
negative_words = ["fall", "drop", "decline", "down", "bearish", "loss", "miss", "low", "crash", "sell"]
|
||||
|
||||
text_lower = text.lower()
|
||||
positive_count = sum(1 for word in positive_words if word in text_lower)
|
||||
negative_count = sum(1 for word in negative_words if word in text_lower)
|
||||
|
||||
if positive_count > negative_count:
|
||||
return "positive"
|
||||
elif negative_count > positive_count:
|
||||
return "negative"
|
||||
else:
|
||||
return "neutral"
|
||||
|
||||
|
||||
def generate_sample_news(topic: str, limit: int) -> List[Dict[str, Any]]:
|
||||
"""生成示例新闻(当 API 不可用时)"""
|
||||
sample_news = [
|
||||
{
|
||||
"title": f"{topic.upper()} Stock Shows Strong Momentum in Pre-Market Trading",
|
||||
"description": f"Analysts remain bullish on {topic.upper()} as the stock shows continued strength.",
|
||||
"url": "https://finance.yahoo.com/",
|
||||
"source": "Yahoo Finance",
|
||||
"published_at": datetime.utcnow().isoformat(),
|
||||
"sentiment": "positive"
|
||||
},
|
||||
{
|
||||
"title": f"Market Analysis: {topic.upper()} Technical Indicators Point to Potential Breakout",
|
||||
"description": "Technical analysts identify key support and resistance levels for upcoming trading sessions.",
|
||||
"url": "https://finance.yahoo.com/",
|
||||
"source": "Market Watch",
|
||||
"published_at": (datetime.utcnow() - timedelta(hours=2)).isoformat(),
|
||||
"sentiment": "positive"
|
||||
},
|
||||
{
|
||||
"title": f"Institutional Investors Increase Holdings in {topic.upper()}",
|
||||
"description": "Latest 13F filings reveal increased institutional interest in the stock.",
|
||||
"url": "https://finance.yahoo.com/",
|
||||
"source": "Bloomberg",
|
||||
"published_at": (datetime.utcnow() - timedelta(hours=4)).isoformat(),
|
||||
"sentiment": "positive"
|
||||
},
|
||||
{
|
||||
"title": f"Wall Street Analysts Update Price Targets for {topic.upper()}",
|
||||
"description": "Multiple analysts revise their price targets following recent earnings report.",
|
||||
"url": "https://finance.yahoo.com/",
|
||||
"source": "CNBC",
|
||||
"published_at": (datetime.utcnow() - timedelta(hours=6)).isoformat(),
|
||||
"sentiment": "neutral"
|
||||
},
|
||||
{
|
||||
"title": f"Options Activity Surges for {topic.upper()} Ahead of Key Events",
|
||||
"description": "Unusual options activity detected as traders position for upcoming catalysts.",
|
||||
"url": "https://finance.yahoo.com/",
|
||||
"source": "Seeking Alpha",
|
||||
"published_at": (datetime.utcnow() - timedelta(hours=8)).isoformat(),
|
||||
"sentiment": "neutral"
|
||||
}
|
||||
]
|
||||
return sample_news[:limit]
|
||||
|
||||
|
||||
async def get_market_movers() -> Dict[str, Any]:
|
||||
"""获取市场涨跌排行"""
|
||||
# 使用 Yahoo Finance 获取市场数据
|
||||
url = "https://query1.finance.yahoo.com/v1/finance/trending/US"
|
||||
|
||||
headers = {
|
||||
"User-Agent": "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36"
|
||||
}
|
||||
|
||||
try:
|
||||
async with aiohttp.ClientSession() as session:
|
||||
async with session.get(url, headers=headers, timeout=15) as response:
|
||||
if response.status == 200:
|
||||
data = await response.json()
|
||||
quotes = data.get("finance", {}).get("result", [{}])[0].get("quotes", [])
|
||||
|
||||
return {
|
||||
"top_gainers": [{"symbol": q.get("symbol")} for q in quotes[:5]],
|
||||
"top_losers": [],
|
||||
"most_active": [{"symbol": q.get("symbol")} for q in quotes[:5]]
|
||||
}
|
||||
except Exception as e:
|
||||
logger.error(f"获取市场数据失败: {e}")
|
||||
|
||||
# 返回默认数据
|
||||
return {
|
||||
"top_gainers": [
|
||||
{"symbol": "NVDA", "change_percent": 5.2},
|
||||
{"symbol": "TSLA", "change_percent": 3.8},
|
||||
{"symbol": "AMD", "change_percent": 2.9}
|
||||
],
|
||||
"top_losers": [
|
||||
{"symbol": "INTC", "change_percent": -2.1},
|
||||
{"symbol": "BA", "change_percent": -1.8}
|
||||
],
|
||||
"most_active": [
|
||||
{"symbol": "AAPL", "volume": "85M"},
|
||||
{"symbol": "TSLA", "volume": "72M"},
|
||||
{"symbol": "NVDA", "volume": "65M"}
|
||||
]
|
||||
}
|
||||
|
||||
|
||||
# ==================== API 端点 ====================
|
||||
|
||||
@app.get("/health", response_model=HealthResponse)
|
||||
@app.get("/", response_model=HealthResponse)
|
||||
async def health_check():
|
||||
"""健康检查"""
|
||||
return HealthResponse(
|
||||
status="healthy",
|
||||
pod_name=POD_NAME,
|
||||
news_api_configured=bool(NEWS_API_KEY),
|
||||
callback_enabled=CALLBACK_ENABLED,
|
||||
timestamp=datetime.utcnow().isoformat()
|
||||
)
|
||||
|
||||
|
||||
@app.post("/news", response_model=NewsResponse)
|
||||
async def get_news(request: NewsRequest):
|
||||
"""获取股票新闻"""
|
||||
if not request.symbol and not request.query:
|
||||
raise HTTPException(status_code=400, detail="请提供股票代码(symbol)或搜索关键词(query)")
|
||||
|
||||
# 使用回调上下文管理器
|
||||
if CALLBACK_ENABLED and callback_handler and request.user_id:
|
||||
with CallbackContextManager(
|
||||
handler=callback_handler,
|
||||
user_id=request.user_id,
|
||||
request_id=f"stock-news-{int(datetime.utcnow().timestamp())}"
|
||||
) as ctx:
|
||||
ctx.add_tool("stock_news")
|
||||
ctx.add_tool("news_aggregation")
|
||||
|
||||
news_data = await fetch_yahoo_news(request.symbol, request.query, request.limit)
|
||||
|
||||
news_items = [NewsItem(**item) for item in news_data]
|
||||
|
||||
return NewsResponse(
|
||||
symbol=request.symbol,
|
||||
query=request.query,
|
||||
news=news_items,
|
||||
total_count=len(news_items),
|
||||
timestamp=datetime.utcnow().isoformat()
|
||||
)
|
||||
else:
|
||||
news_data = await fetch_yahoo_news(request.symbol, request.query, request.limit)
|
||||
news_items = [NewsItem(**item) for item in news_data]
|
||||
|
||||
return NewsResponse(
|
||||
symbol=request.symbol,
|
||||
query=request.query,
|
||||
news=news_items,
|
||||
total_count=len(news_items),
|
||||
timestamp=datetime.utcnow().isoformat()
|
||||
)
|
||||
|
||||
|
||||
@app.get("/news")
|
||||
async def get_news_get(
|
||||
symbol: Optional[str] = Query(None, description="股票代码"),
|
||||
query: Optional[str] = Query(None, description="搜索关键词"),
|
||||
limit: int = Query(10, ge=1, le=50, description="返回数量"),
|
||||
user_id: Optional[str] = Query(None, description="用户ID")
|
||||
):
|
||||
"""GET 方式获取新闻"""
|
||||
request = NewsRequest(symbol=symbol, query=query, limit=limit, user_id=user_id)
|
||||
return await get_news(request)
|
||||
|
||||
|
||||
@app.get("/market-summary", response_model=MarketSummaryResponse)
|
||||
async def get_market_summary():
|
||||
"""获取市场概要"""
|
||||
movers = await get_market_movers()
|
||||
|
||||
# 判断市场状态(简单逻辑)
|
||||
now = datetime.utcnow()
|
||||
hour = now.hour
|
||||
weekday = now.weekday()
|
||||
|
||||
if weekday >= 5: # 周末
|
||||
market_status = "closed"
|
||||
elif 13 <= hour < 21: # UTC 时间对应美东 9:30-16:00
|
||||
market_status = "open"
|
||||
elif 9 <= hour < 13: # 盘前
|
||||
market_status = "pre-market"
|
||||
elif 21 <= hour < 25: # 盘后
|
||||
market_status = "after-hours"
|
||||
else:
|
||||
market_status = "closed"
|
||||
|
||||
return MarketSummaryResponse(
|
||||
market_status=market_status,
|
||||
top_gainers=movers["top_gainers"],
|
||||
top_losers=movers["top_losers"],
|
||||
most_active=movers["most_active"],
|
||||
timestamp=datetime.utcnow().isoformat()
|
||||
)
|
||||
|
||||
|
||||
@app.get("/trending")
|
||||
async def get_trending_news():
|
||||
"""获取热门财经新闻"""
|
||||
news_data = await fetch_yahoo_news(query="stock market US", limit=20)
|
||||
news_items = [NewsItem(**item) for item in news_data]
|
||||
|
||||
return NewsResponse(
|
||||
query="trending",
|
||||
news=news_items,
|
||||
total_count=len(news_items),
|
||||
timestamp=datetime.utcnow().isoformat()
|
||||
)
|
||||
|
||||
|
||||
# ==================== Chat 功能 ====================
|
||||
|
||||
async def chat_with_llm(message: str, context: str, api_key: str) -> str:
|
||||
"""调用 LLM 生成响应"""
|
||||
try:
|
||||
async with aiohttp.ClientSession() as session:
|
||||
payload = {
|
||||
"model": LLM_MODEL,
|
||||
"messages": [
|
||||
{
|
||||
"role": "system",
|
||||
"content": """你是一个专业的美股新闻分析师。你可以:
|
||||
1. 获取和分析股票相关新闻
|
||||
2. 解读市场动态和公司公告
|
||||
3. 提供新闻情感分析和市场影响评估
|
||||
|
||||
请根据提供的新闻数据,用简洁专业的语言回答用户问题。"""
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": f"最新新闻:\n{context}\n\n用户问题: {message}"
|
||||
}
|
||||
],
|
||||
"max_tokens": 500,
|
||||
"temperature": 0.7
|
||||
}
|
||||
headers = {
|
||||
"Authorization": f"Bearer {api_key}",
|
||||
"Content-Type": "application/json"
|
||||
}
|
||||
|
||||
async with session.post(
|
||||
f"{LLM_BASE_URL}/chat/completions",
|
||||
json=payload,
|
||||
headers=headers,
|
||||
timeout=aiohttp.ClientTimeout(total=30)
|
||||
) as response:
|
||||
if response.status == 200:
|
||||
data = await response.json()
|
||||
return data.get("choices", [{}])[0].get("message", {}).get("content", "抱歉,无法生成回复")
|
||||
else:
|
||||
error = await response.text()
|
||||
logger.error(f"LLM 请求失败: {response.status} - {error}")
|
||||
return f"LLM 服务错误: {response.status}"
|
||||
except Exception as e:
|
||||
logger.error(f"LLM 调用失败: {e}")
|
||||
return f"调用失败: {str(e)}"
|
||||
|
||||
|
||||
@app.post("/chat", response_model=ChatResponse)
|
||||
async def chat(
|
||||
request: ChatRequest,
|
||||
api_key: Optional[str] = Header(None, alias="api-key"),
|
||||
authorization: Optional[str] = Header(None)
|
||||
):
|
||||
"""智能对话 - 获取新闻并提供分析
|
||||
|
||||
api_key 通过请求头传递:
|
||||
- api-key: your-api-key
|
||||
- 或 Authorization: Bearer your-api-key
|
||||
"""
|
||||
# 从 Header 获取 api_key
|
||||
if not api_key and authorization:
|
||||
if authorization.startswith("Bearer "):
|
||||
api_key = authorization[7:]
|
||||
else:
|
||||
api_key = authorization
|
||||
|
||||
if not api_key:
|
||||
raise HTTPException(status_code=401, detail="请在请求头中提供 api-key 或 Authorization")
|
||||
|
||||
# 从消息中提取关键词
|
||||
import re
|
||||
symbols = re.findall(r'\b([A-Z]{1,5})\b', request.message.upper())
|
||||
common_words = {"I", "A", "THE", "IS", "IT", "TO", "OF", "AND", "FOR", "IN", "ON", "AT", "BY", "NEWS", "WHAT"}
|
||||
symbols = [s for s in symbols if s not in common_words][:3]
|
||||
|
||||
# 获取相关新闻
|
||||
news_data = []
|
||||
if symbols:
|
||||
for symbol in symbols:
|
||||
data = await fetch_yahoo_news(symbol=symbol, limit=3)
|
||||
news_data.extend(data)
|
||||
else:
|
||||
news_data = await fetch_yahoo_news(query="stock market", limit=5)
|
||||
|
||||
# 构建上下文
|
||||
if news_data:
|
||||
context = "\n".join([
|
||||
f"- {item['title']} ({item['source']}, {item['published_at'][:10]})"
|
||||
for item in news_data[:5]
|
||||
])
|
||||
else:
|
||||
context = "暂无相关新闻"
|
||||
|
||||
# 调用 LLM 生成回复
|
||||
llm_response = await chat_with_llm(request.message, context, api_key)
|
||||
|
||||
return ChatResponse(
|
||||
response=llm_response,
|
||||
data={"news_count": len(news_data), "symbols": symbols},
|
||||
timestamp=datetime.utcnow().isoformat()
|
||||
)
|
||||
|
||||
|
||||
# ==================== 主入口 ====================
|
||||
|
||||
def main():
|
||||
"""主函数"""
|
||||
logger.info(f"启动 Stock News Agent - {POD_NAME}")
|
||||
logger.info(f"News API: {'已配置' if NEWS_API_KEY else '使用免费源'}")
|
||||
logger.info(f"回调功能: {'已启用' if CALLBACK_ENABLED else '未启用'}")
|
||||
|
||||
uvicorn.run(app, host=SERVICE_HOST, port=SERVICE_PORT, log_level="info")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,39 @@
|
||||
FROM python:3.11-slim
|
||||
|
||||
WORKDIR /app
|
||||
|
||||
# 安装系统依赖
|
||||
RUN apt-get update && apt-get install -y \
|
||||
curl \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
|
||||
# 安装 Python 依赖
|
||||
RUN pip install --no-cache-dir \
|
||||
fastapi==0.109.0 \
|
||||
uvicorn[standard]==0.27.0 \
|
||||
pydantic==2.5.3 \
|
||||
requests>=2.31.0 \
|
||||
aiohttp>=3.9.0
|
||||
|
||||
# 复制 common 模块(回调工具)
|
||||
COPY common/agent_callback_utils.py /app/common/
|
||||
RUN touch /app/common/__init__.py
|
||||
|
||||
# 复制应用代码
|
||||
COPY agents/stock_quote_agent/stock_quote_agent.py /app/
|
||||
|
||||
# 环境变量
|
||||
ENV PYTHONUNBUFFERED=1
|
||||
ENV SERVICE_HOST=0.0.0.0
|
||||
ENV SERVICE_PORT=8080
|
||||
|
||||
# 回调配置
|
||||
ENV AGENT_CALLBACK_URL=http://mcp-server:8002/api/v1/billing/agent-callback
|
||||
|
||||
# 健康检查
|
||||
HEALTHCHECK --interval=30s --timeout=10s --start-period=20s --retries=3 \
|
||||
CMD python3 -c "import urllib.request; urllib.request.urlopen('http://localhost:8080/health').read()" || exit 1
|
||||
|
||||
EXPOSE 8080
|
||||
|
||||
CMD ["python3", "-u", "stock_quote_agent.py"]
|
||||
@@ -0,0 +1,480 @@
|
||||
"""
|
||||
Stock Quote Agent - 美股实时行情查询 Agent
|
||||
使用 Yahoo Finance API 获取美股实时行情数据
|
||||
"""
|
||||
import os
|
||||
import sys
|
||||
import logging
|
||||
import aiohttp
|
||||
from typing import Optional, List, Dict, Any
|
||||
from datetime import datetime
|
||||
|
||||
from fastapi import FastAPI, HTTPException, Query, Header, Request
|
||||
from fastapi.middleware.cors import CORSMiddleware
|
||||
from pydantic import BaseModel, Field
|
||||
import uvicorn
|
||||
|
||||
# 添加 common 模块路径
|
||||
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
|
||||
# 导入回调工具
|
||||
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, format='%(asctime)s - %(name)s - %(levelname)s - %(message)s')
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# 环境变量
|
||||
SERVICE_HOST = os.getenv("SERVICE_HOST", "0.0.0.0")
|
||||
SERVICE_PORT = int(os.getenv("SERVICE_PORT", "8080"))
|
||||
POD_NAME = os.getenv("POD_NAME", "stock-quote-agent")
|
||||
USER_ID = os.getenv("USER_ID", "")
|
||||
|
||||
# Yahoo Finance API (通过 RapidAPI)
|
||||
RAPIDAPI_KEY = os.getenv("RAPIDAPI_KEY", "")
|
||||
YAHOO_FINANCE_HOST = "yahoo-finance15.p.rapidapi.com"
|
||||
|
||||
# LLM 配置
|
||||
LLM_BASE_URL = os.getenv("LLM_BASE_URL", "https://litellm.graystone-fb459c5d.southeastasia.azurecontainerapps.io/v1")
|
||||
LLM_MODEL = os.getenv("LLM_MODEL", "taiji/gpt-4o-mini")
|
||||
|
||||
# FastAPI 应用
|
||||
app = FastAPI(
|
||||
title="Stock Quote Agent",
|
||||
description="美股实时行情查询 - 获取股票价格、涨跌幅、成交量等数据",
|
||||
version="1.0.0"
|
||||
)
|
||||
|
||||
app.add_middleware(
|
||||
CORSMiddleware,
|
||||
allow_origins=["*"],
|
||||
allow_credentials=True,
|
||||
allow_methods=["*"],
|
||||
allow_headers=["*"],
|
||||
)
|
||||
|
||||
# 回调处理器
|
||||
callback_handler: Optional[AgentCallbackHandler] = None
|
||||
|
||||
|
||||
# ==================== 请求/响应模型 ====================
|
||||
|
||||
class QuoteRequest(BaseModel):
|
||||
"""行情查询请求"""
|
||||
symbol: str = Field(..., description="股票代码,如 AAPL, TSLA, MSFT")
|
||||
user_id: Optional[str] = Field(None, description="用户ID(用于计费回调)")
|
||||
|
||||
|
||||
class QuoteResponse(BaseModel):
|
||||
"""行情查询响应"""
|
||||
symbol: str
|
||||
name: str
|
||||
price: float
|
||||
change: float
|
||||
change_percent: float
|
||||
volume: int
|
||||
market_cap: Optional[float] = None
|
||||
pe_ratio: Optional[float] = None
|
||||
high_52week: Optional[float] = None
|
||||
low_52week: Optional[float] = None
|
||||
timestamp: str
|
||||
|
||||
|
||||
class BatchQuoteRequest(BaseModel):
|
||||
"""批量行情查询请求"""
|
||||
symbols: List[str] = Field(..., description="股票代码列表")
|
||||
user_id: Optional[str] = Field(None, description="用户ID(用于计费回调)")
|
||||
|
||||
|
||||
class BatchQuoteResponse(BaseModel):
|
||||
"""批量行情查询响应"""
|
||||
quotes: List[QuoteResponse]
|
||||
success_count: int
|
||||
failed_count: int
|
||||
timestamp: str
|
||||
|
||||
|
||||
class HealthResponse(BaseModel):
|
||||
"""健康检查响应"""
|
||||
status: str
|
||||
pod_name: str
|
||||
api_configured: bool
|
||||
callback_enabled: bool
|
||||
timestamp: str
|
||||
|
||||
|
||||
class ChatRequest(BaseModel):
|
||||
"""Chat 请求"""
|
||||
message: str = Field(..., description="用户消息")
|
||||
user_id: Optional[str] = Field(None, description="用户ID")
|
||||
|
||||
|
||||
class ChatResponse(BaseModel):
|
||||
"""Chat 响应"""
|
||||
response: str
|
||||
data: Optional[Dict[str, Any]] = None
|
||||
timestamp: str
|
||||
|
||||
|
||||
# ==================== 生命周期 ====================
|
||||
|
||||
@app.on_event("startup")
|
||||
async def startup_event():
|
||||
"""应用启动时初始化回调处理器"""
|
||||
global callback_handler
|
||||
|
||||
if CALLBACK_ENABLED:
|
||||
callback_handler = AgentCallbackHandler(
|
||||
agent_name=POD_NAME,
|
||||
user_id=USER_ID
|
||||
)
|
||||
logger.info(f"回调处理器已初始化: callback_url={callback_handler.callback_url}")
|
||||
else:
|
||||
logger.warning("回调模块未加载,计费回调功能不可用")
|
||||
|
||||
|
||||
# ==================== 辅助函数 ====================
|
||||
|
||||
async def fetch_stock_quote(symbol: str) -> Dict[str, Any]:
|
||||
"""获取股票行情数据"""
|
||||
# 使用免费的 Yahoo Finance API 替代方案
|
||||
url = f"https://query1.finance.yahoo.com/v8/finance/chart/{symbol}"
|
||||
params = {
|
||||
"interval": "1d",
|
||||
"range": "1d"
|
||||
}
|
||||
|
||||
headers = {
|
||||
"User-Agent": "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36"
|
||||
}
|
||||
|
||||
try:
|
||||
async with aiohttp.ClientSession() as session:
|
||||
async with session.get(url, params=params, headers=headers, timeout=15) as response:
|
||||
if response.status == 200:
|
||||
data = await response.json()
|
||||
result = data.get("chart", {}).get("result", [])
|
||||
|
||||
if not result:
|
||||
return {"success": False, "error": f"未找到股票: {symbol}"}
|
||||
|
||||
quote_data = result[0]
|
||||
meta = quote_data.get("meta", {})
|
||||
indicators = quote_data.get("indicators", {}).get("quote", [{}])[0]
|
||||
|
||||
current_price = meta.get("regularMarketPrice", 0)
|
||||
previous_close = meta.get("previousClose", 0)
|
||||
change = current_price - previous_close if previous_close else 0
|
||||
change_percent = (change / previous_close * 100) if previous_close else 0
|
||||
|
||||
return {
|
||||
"success": True,
|
||||
"symbol": symbol.upper(),
|
||||
"name": meta.get("shortName", symbol),
|
||||
"price": current_price,
|
||||
"change": round(change, 2),
|
||||
"change_percent": round(change_percent, 2),
|
||||
"volume": indicators.get("volume", [0])[-1] if indicators.get("volume") else 0,
|
||||
"market_cap": meta.get("marketCap"),
|
||||
"pe_ratio": None,
|
||||
"high_52week": meta.get("fiftyTwoWeekHigh"),
|
||||
"low_52week": meta.get("fiftyTwoWeekLow")
|
||||
}
|
||||
else:
|
||||
return {"success": False, "error": f"API 请求失败: HTTP {response.status}"}
|
||||
except Exception as e:
|
||||
logger.error(f"获取行情失败: {symbol} - {e}")
|
||||
return {"success": False, "error": str(e)}
|
||||
|
||||
|
||||
# ==================== API 端点 ====================
|
||||
|
||||
@app.get("/health", response_model=HealthResponse)
|
||||
@app.get("/", response_model=HealthResponse)
|
||||
async def health_check():
|
||||
"""健康检查"""
|
||||
return HealthResponse(
|
||||
status="healthy",
|
||||
pod_name=POD_NAME,
|
||||
api_configured=True, # 使用免费 API
|
||||
callback_enabled=CALLBACK_ENABLED,
|
||||
timestamp=datetime.utcnow().isoformat()
|
||||
)
|
||||
|
||||
|
||||
@app.post("/quote", response_model=QuoteResponse)
|
||||
async def get_quote(request: QuoteRequest):
|
||||
"""获取单个股票行情"""
|
||||
# 使用回调上下文管理器
|
||||
if CALLBACK_ENABLED and callback_handler and request.user_id:
|
||||
with CallbackContextManager(
|
||||
handler=callback_handler,
|
||||
user_id=request.user_id,
|
||||
request_id=f"stock-quote-{int(datetime.utcnow().timestamp())}"
|
||||
) as ctx:
|
||||
ctx.add_tool("stock_quote")
|
||||
ctx.add_tool("yahoo_finance")
|
||||
|
||||
result = await fetch_stock_quote(request.symbol)
|
||||
|
||||
if not result["success"]:
|
||||
raise HTTPException(status_code=500, detail=result.get("error"))
|
||||
|
||||
return QuoteResponse(
|
||||
symbol=result["symbol"],
|
||||
name=result["name"],
|
||||
price=result["price"],
|
||||
change=result["change"],
|
||||
change_percent=result["change_percent"],
|
||||
volume=result["volume"],
|
||||
market_cap=result.get("market_cap"),
|
||||
pe_ratio=result.get("pe_ratio"),
|
||||
high_52week=result.get("high_52week"),
|
||||
low_52week=result.get("low_52week"),
|
||||
timestamp=datetime.utcnow().isoformat()
|
||||
)
|
||||
else:
|
||||
result = await fetch_stock_quote(request.symbol)
|
||||
|
||||
if not result["success"]:
|
||||
raise HTTPException(status_code=500, detail=result.get("error"))
|
||||
|
||||
return QuoteResponse(
|
||||
symbol=result["symbol"],
|
||||
name=result["name"],
|
||||
price=result["price"],
|
||||
change=result["change"],
|
||||
change_percent=result["change_percent"],
|
||||
volume=result["volume"],
|
||||
market_cap=result.get("market_cap"),
|
||||
pe_ratio=result.get("pe_ratio"),
|
||||
high_52week=result.get("high_52week"),
|
||||
low_52week=result.get("low_52week"),
|
||||
timestamp=datetime.utcnow().isoformat()
|
||||
)
|
||||
|
||||
|
||||
@app.get("/quote")
|
||||
async def get_quote_get(
|
||||
symbol: str = Query(..., description="股票代码"),
|
||||
user_id: Optional[str] = Query(None, description="用户ID")
|
||||
):
|
||||
"""GET 方式获取行情"""
|
||||
request = QuoteRequest(symbol=symbol, user_id=user_id)
|
||||
return await get_quote(request)
|
||||
|
||||
|
||||
@app.post("/batch", response_model=BatchQuoteResponse)
|
||||
async def batch_get_quotes(request: BatchQuoteRequest):
|
||||
"""批量获取股票行情"""
|
||||
if CALLBACK_ENABLED and callback_handler and request.user_id:
|
||||
with CallbackContextManager(
|
||||
handler=callback_handler,
|
||||
user_id=request.user_id,
|
||||
request_id=f"stock-batch-{int(datetime.utcnow().timestamp())}"
|
||||
) as ctx:
|
||||
ctx.add_tool("stock_quote")
|
||||
ctx.add_tool("batch_quote")
|
||||
|
||||
quotes = []
|
||||
success_count = 0
|
||||
failed_count = 0
|
||||
|
||||
for symbol in request.symbols:
|
||||
result = await fetch_stock_quote(symbol)
|
||||
if result["success"]:
|
||||
quotes.append(QuoteResponse(
|
||||
symbol=result["symbol"],
|
||||
name=result["name"],
|
||||
price=result["price"],
|
||||
change=result["change"],
|
||||
change_percent=result["change_percent"],
|
||||
volume=result["volume"],
|
||||
market_cap=result.get("market_cap"),
|
||||
pe_ratio=result.get("pe_ratio"),
|
||||
high_52week=result.get("high_52week"),
|
||||
low_52week=result.get("low_52week"),
|
||||
timestamp=datetime.utcnow().isoformat()
|
||||
))
|
||||
success_count += 1
|
||||
else:
|
||||
failed_count += 1
|
||||
|
||||
return BatchQuoteResponse(
|
||||
quotes=quotes,
|
||||
success_count=success_count,
|
||||
failed_count=failed_count,
|
||||
timestamp=datetime.utcnow().isoformat()
|
||||
)
|
||||
else:
|
||||
quotes = []
|
||||
success_count = 0
|
||||
failed_count = 0
|
||||
|
||||
for symbol in request.symbols:
|
||||
result = await fetch_stock_quote(symbol)
|
||||
if result["success"]:
|
||||
quotes.append(QuoteResponse(
|
||||
symbol=result["symbol"],
|
||||
name=result["name"],
|
||||
price=result["price"],
|
||||
change=result["change"],
|
||||
change_percent=result["change_percent"],
|
||||
volume=result["volume"],
|
||||
market_cap=result.get("market_cap"),
|
||||
pe_ratio=result.get("pe_ratio"),
|
||||
high_52week=result.get("high_52week"),
|
||||
low_52week=result.get("low_52week"),
|
||||
timestamp=datetime.utcnow().isoformat()
|
||||
))
|
||||
success_count += 1
|
||||
else:
|
||||
failed_count += 1
|
||||
|
||||
return BatchQuoteResponse(
|
||||
quotes=quotes,
|
||||
success_count=success_count,
|
||||
failed_count=failed_count,
|
||||
timestamp=datetime.utcnow().isoformat()
|
||||
)
|
||||
|
||||
|
||||
@app.get("/popular")
|
||||
async def get_popular_stocks():
|
||||
"""获取热门股票行情"""
|
||||
popular_symbols = ["AAPL", "MSFT", "GOOGL", "AMZN", "TSLA", "NVDA", "META"]
|
||||
request = BatchQuoteRequest(symbols=popular_symbols)
|
||||
return await batch_get_quotes(request)
|
||||
|
||||
|
||||
async def chat_with_llm(message: str, context: str, api_key: str) -> str:
|
||||
"""调用 LLM 生成响应"""
|
||||
try:
|
||||
async with aiohttp.ClientSession() as session:
|
||||
payload = {
|
||||
"model": LLM_MODEL,
|
||||
"messages": [
|
||||
{
|
||||
"role": "system",
|
||||
"content": """你是一个专业的美股分析助手。你可以:
|
||||
1. 查询股票实时行情(价格、涨跌幅、成交量)
|
||||
2. 分析股票数据并给出建议
|
||||
3. 解答关于美股市场的问题
|
||||
|
||||
请根据提供的股票数据,用简洁专业的语言回答用户问题。"""
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": f"当前股票数据:\n{context}\n\n用户问题: {message}"
|
||||
}
|
||||
],
|
||||
"max_tokens": 500,
|
||||
"temperature": 0.7
|
||||
}
|
||||
headers = {
|
||||
"Authorization": f"Bearer {api_key}",
|
||||
"Content-Type": "application/json"
|
||||
}
|
||||
|
||||
async with session.post(
|
||||
f"{LLM_BASE_URL}/chat/completions",
|
||||
json=payload,
|
||||
headers=headers,
|
||||
timeout=aiohttp.ClientTimeout(total=30)
|
||||
) as response:
|
||||
if response.status == 200:
|
||||
data = await response.json()
|
||||
return data.get("choices", [{}])[0].get("message", {}).get("content", "抱歉,无法生成回复")
|
||||
else:
|
||||
error = await response.text()
|
||||
logger.error(f"LLM 请求失败: {response.status} - {error}")
|
||||
return f"LLM 服务错误: {response.status}"
|
||||
except Exception as e:
|
||||
logger.error(f"LLM 调用失败: {e}")
|
||||
return f"调用失败: {str(e)}"
|
||||
|
||||
|
||||
def extract_symbols_from_message(message: str) -> List[str]:
|
||||
"""从消息中提取股票代码"""
|
||||
import re
|
||||
# 匹配常见美股代码格式(1-5个大写字母)
|
||||
symbols = re.findall(r'\b([A-Z]{1,5})\b', message.upper())
|
||||
# 过滤常见词汇
|
||||
common_words = {"I", "A", "THE", "IS", "IT", "TO", "OF", "AND", "FOR", "IN", "ON", "AT", "BY"}
|
||||
return [s for s in symbols if s not in common_words][:5] # 最多5个
|
||||
|
||||
|
||||
@app.post("/chat", response_model=ChatResponse)
|
||||
async def chat(
|
||||
request: ChatRequest,
|
||||
api_key: Optional[str] = Header(None, alias="api-key"),
|
||||
authorization: Optional[str] = Header(None)
|
||||
):
|
||||
"""智能对话 - 支持自然语言查询股票信息
|
||||
|
||||
api_key 通过请求头传递:
|
||||
- api-key: your-api-key
|
||||
- 或 Authorization: Bearer your-api-key
|
||||
"""
|
||||
# 从 Header 获取 api_key
|
||||
if not api_key and authorization:
|
||||
if authorization.startswith("Bearer "):
|
||||
api_key = authorization[7:]
|
||||
else:
|
||||
api_key = authorization
|
||||
|
||||
if not api_key:
|
||||
raise HTTPException(status_code=401, detail="请在请求头中提供 api-key 或 Authorization")
|
||||
|
||||
# 从消息中提取股票代码
|
||||
symbols = extract_symbols_from_message(request.message)
|
||||
|
||||
# 如果没有提取到,默认查询热门股票
|
||||
if not symbols:
|
||||
symbols = ["AAPL", "TSLA", "NVDA"]
|
||||
|
||||
# 获取股票数据
|
||||
stock_data = []
|
||||
for symbol in symbols:
|
||||
result = await fetch_stock_quote(symbol)
|
||||
if result.get("success"):
|
||||
stock_data.append(result)
|
||||
|
||||
# 构建上下文
|
||||
if stock_data:
|
||||
context = "\n".join([
|
||||
f"{d['symbol']} ({d['name']}): ${d['price']:.2f}, 涨跌: {d['change_percent']:+.2f}%, "
|
||||
f"成交量: {d['volume']:,}, 52周范围: ${d.get('low_52week', 0):.2f}-${d.get('high_52week', 0):.2f}"
|
||||
for d in stock_data
|
||||
])
|
||||
else:
|
||||
context = "暂无股票数据"
|
||||
|
||||
# 调用 LLM 生成回复
|
||||
llm_response = await chat_with_llm(request.message, context, api_key)
|
||||
|
||||
return ChatResponse(
|
||||
response=llm_response,
|
||||
data={"stocks": stock_data, "symbols_detected": symbols},
|
||||
timestamp=datetime.utcnow().isoformat()
|
||||
)
|
||||
|
||||
|
||||
# ==================== 主入口 ====================
|
||||
|
||||
def main():
|
||||
"""主函数"""
|
||||
logger.info(f"启动 Stock Quote Agent - {POD_NAME}")
|
||||
logger.info(f"回调功能: {'已启用' if CALLBACK_ENABLED else '未启用'}")
|
||||
|
||||
uvicorn.run(app, host=SERVICE_HOST, port=SERVICE_PORT, log_level="info")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,24 @@
|
||||
__pycache__
|
||||
*.pyc
|
||||
*.pyo
|
||||
*.pyd
|
||||
.Python
|
||||
*.so
|
||||
*.egg
|
||||
*.egg-info
|
||||
dist
|
||||
build
|
||||
.pytest_cache
|
||||
.coverage
|
||||
htmlcov
|
||||
.venv
|
||||
venv
|
||||
outputs/
|
||||
*.mp4
|
||||
*.png
|
||||
*.jpg
|
||||
*.jpeg
|
||||
.DS_Store
|
||||
.env
|
||||
.vscode
|
||||
.idea
|
||||
@@ -0,0 +1,294 @@
|
||||
# Video Generator Agent - 部署和测试指南
|
||||
|
||||
## ✅ 已完成的工作
|
||||
|
||||
### 1. 核心功能实现
|
||||
- ✅ 图片生成工具 (`image_generator.py`)
|
||||
- ✅ 视频处理工具 (`video_processor.py`) - 使用 FFmpeg
|
||||
- ✅ 文件管理工具 (`file_manager.py`)
|
||||
- ✅ MCP 服务器 (`mcp_server.py`) - 4个工具
|
||||
- ✅ API 服务器 (`api_server.py`) - REST API + MCP HTTP/SSE
|
||||
|
||||
### 2. Docker 配置
|
||||
- ✅ Dockerfile(包含 FFmpeg)
|
||||
- ✅ requirements.txt
|
||||
- ✅ .dockerignore
|
||||
- ✅ 健康检查配置
|
||||
|
||||
### 3. 测试和文档
|
||||
- ✅ 测试脚本 (`test_video_agent.py`)
|
||||
- ✅ 核心功能测试 (`test_core_functions.py`)
|
||||
- ✅ README.md 完整文档
|
||||
- ✅ Docker 镜像构建成功
|
||||
- ✅ FFmpeg 验证通过
|
||||
- ✅ 模块导入验证通过
|
||||
|
||||
## 🚀 快速部署
|
||||
|
||||
### 方法 1: Docker 部署(推荐)
|
||||
|
||||
```bash
|
||||
# 1. 构建镜像
|
||||
cd agent_templates/agents/video_generator_agent
|
||||
docker build -t video-generator-agent:latest .
|
||||
|
||||
# 2. 运行容器
|
||||
docker run -d \
|
||||
--name video-generator \
|
||||
-p 8000:8000 \
|
||||
-e OPENAI_API_KEY="sk-i9AwAgXDqqxsA9Ym4AjSPg" \
|
||||
-e MODEL_NAME="dall-e-3" \
|
||||
-v $(pwd)/outputs:/app/outputs \
|
||||
video-generator-agent:latest
|
||||
|
||||
# 3. 查看日志
|
||||
docker logs -f video-generator
|
||||
|
||||
# 4. 测试健康检查
|
||||
curl http://localhost:8000/health
|
||||
```
|
||||
|
||||
### 方法 2: 本地运行
|
||||
|
||||
```bash
|
||||
# 1. 创建虚拟环境
|
||||
python3 -m venv venv
|
||||
source venv/bin/activate
|
||||
|
||||
# 2. 安装依赖
|
||||
pip install -r requirements.txt
|
||||
|
||||
# 3. 安装 FFmpeg(如果未安装)
|
||||
# Ubuntu/Debian:
|
||||
sudo apt-get install ffmpeg
|
||||
|
||||
# macOS:
|
||||
brew install ffmpeg
|
||||
|
||||
# 4. 设置环境变量
|
||||
export OPENAI_API_KEY="sk-i9AwAgXDqqxsA9Ym4AjSPg"
|
||||
export MODEL_NAME="dall-e-3"
|
||||
|
||||
# 5. 启动服务
|
||||
python run_api_server.py
|
||||
```
|
||||
|
||||
## 🧪 测试
|
||||
|
||||
### 1. 健康检查
|
||||
|
||||
```bash
|
||||
curl http://localhost:8000/health
|
||||
```
|
||||
|
||||
预期输出:
|
||||
```json
|
||||
{
|
||||
"status": "healthy",
|
||||
"service": "Video Generator Agent"
|
||||
}
|
||||
```
|
||||
|
||||
### 2. 查看服务信息
|
||||
|
||||
```bash
|
||||
curl http://localhost:8000/
|
||||
```
|
||||
|
||||
### 3. 生成单张图片
|
||||
|
||||
```bash
|
||||
curl -X POST http://localhost:8000/api/v1/generate-image \
|
||||
-H "api-key: sk-i9AwAgXDqqxsA9Ym4AjSPg" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"description": "a beautiful sunset over mountains",
|
||||
"size": "1024x1024"
|
||||
}'
|
||||
```
|
||||
|
||||
### 4. 生成视频
|
||||
|
||||
```bash
|
||||
curl -X POST http://localhost:8000/api/v1/generate-video \
|
||||
-H "api-key: sk-i9AwAgXDqqxsA9Ym4AjSPg" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"descriptions": [
|
||||
"sunrise over mountains",
|
||||
"a peaceful lake at noon",
|
||||
"starry night sky"
|
||||
],
|
||||
"duration_per_image": 3,
|
||||
"transition": "fade"
|
||||
}'
|
||||
```
|
||||
|
||||
### 5. 运行完整测试套件
|
||||
|
||||
```bash
|
||||
# 使用 Docker
|
||||
docker run --rm \
|
||||
-e API_BASE_URL=http://host.docker.internal:8000 \
|
||||
-e API_KEY=sk-i9AwAgXDqqxsA9Ym4AjSPg \
|
||||
--network host \
|
||||
video-generator-agent:latest \
|
||||
python test_video_agent.py
|
||||
|
||||
# 或本地运行
|
||||
python test_video_agent.py
|
||||
```
|
||||
|
||||
## ⚠️ 重要注意事项
|
||||
|
||||
### 1. 模型配置
|
||||
|
||||
**当前问题**: 模型 `taiji/gemini-3-pro-image-preview` 不支持图片生成。
|
||||
|
||||
**解决方案**: 使用以下支持的模型之一:
|
||||
|
||||
```bash
|
||||
# OpenAI DALL-E 模型
|
||||
export MODEL_NAME="dall-e-3"
|
||||
export MODEL_NAME="dall-e-2"
|
||||
|
||||
# 或其他 LiteLLM 支持的图片生成模型
|
||||
```
|
||||
|
||||
### 2. API Key
|
||||
|
||||
确保使用有效的 API Key:
|
||||
```bash
|
||||
export OPENAI_API_KEY="sk-i9AwAgXDqqxsA9Ym4AjSPg"
|
||||
```
|
||||
|
||||
### 3. FFmpeg 依赖
|
||||
|
||||
视频拼接功能需要 FFmpeg。Docker 镜像已包含,本地运行需要手动安装。
|
||||
|
||||
验证 FFmpeg:
|
||||
```bash
|
||||
ffmpeg -version
|
||||
```
|
||||
|
||||
### 4. 存储空间
|
||||
|
||||
生成的图片和视频会占用存储空间。建议:
|
||||
- 定期清理旧文件
|
||||
- 使用 Docker volume 持久化数据
|
||||
- 监控磁盘使用情况
|
||||
|
||||
## 📊 测试结果
|
||||
|
||||
### Docker 构建
|
||||
- ✅ 镜像构建成功
|
||||
- ✅ FFmpeg 7.1.3 已安装
|
||||
- ✅ Python 依赖已安装
|
||||
- ✅ 健康检查配置正确
|
||||
|
||||
### 模块验证
|
||||
- ✅ 所有 Python 模块导入成功
|
||||
- ✅ FFmpeg 可执行
|
||||
- ✅ 服务启动正常
|
||||
|
||||
### 功能状态
|
||||
- ✅ REST API 端点正常
|
||||
- ✅ MCP 端点正常
|
||||
- ✅ 健康检查通过
|
||||
- ⚠️ 图片生成需要正确的模型配置
|
||||
- ✅ 视频拼接功能就绪(FFmpeg)
|
||||
- ✅ 文件管理功能正常
|
||||
|
||||
## 🔧 故障排除
|
||||
|
||||
### 问题 1: 图片生成失败
|
||||
|
||||
**错误**: `not supported model for image generation`
|
||||
|
||||
**解决**:
|
||||
```bash
|
||||
# 更改模型为支持的模型
|
||||
docker run -d \
|
||||
-e MODEL_NAME="dall-e-3" \
|
||||
...
|
||||
```
|
||||
|
||||
### 问题 2: 端口被占用
|
||||
|
||||
**错误**: `Bind for 0.0.0.0:8000 failed: port is already allocated`
|
||||
|
||||
**解决**:
|
||||
```bash
|
||||
# 使用不同端口
|
||||
docker run -d -p 8765:8000 ...
|
||||
```
|
||||
|
||||
### 问题 3: FFmpeg 未找到
|
||||
|
||||
**错误**: `FFmpeg not available`
|
||||
|
||||
**解决**:
|
||||
```bash
|
||||
# Ubuntu/Debian
|
||||
sudo apt-get install ffmpeg
|
||||
|
||||
# macOS
|
||||
brew install ffmpeg
|
||||
```
|
||||
|
||||
## 📝 下一步
|
||||
|
||||
1. **配置正确的图片生成模型**
|
||||
- 联系 LiteLLM Gateway 管理员
|
||||
- 确认可用的图片生成模型
|
||||
- 更新 MODEL_NAME 环境变量
|
||||
|
||||
2. **运行完整测试**
|
||||
```bash
|
||||
python test_video_agent.py
|
||||
```
|
||||
|
||||
3. **集成到 Agent Manager**
|
||||
- 在 `k8s_manager.py` 中添加配置
|
||||
- 在 `app.py` 中注册模板
|
||||
- 部署到 Kubernetes
|
||||
|
||||
4. **生产环境优化**
|
||||
- 配置持久化存储
|
||||
- 设置资源限制
|
||||
- 配置日志收集
|
||||
- 添加监控告警
|
||||
|
||||
## 📚 相关文档
|
||||
|
||||
- [README.md](./README.md) - 完整使用文档
|
||||
- [test_video_agent.py](./test_video_agent.py) - 测试脚本
|
||||
- [Dockerfile](./Dockerfile) - Docker 配置
|
||||
|
||||
## ✨ 功能亮点
|
||||
|
||||
1. **完整的 MCP 支持** - 4个工具,支持 HTTP 和 SSE
|
||||
2. **FFmpeg 视频处理** - 支持多种转场效果
|
||||
3. **文件管理** - 自动清理、存储统计
|
||||
4. **Docker 化** - 开箱即用,包含所有依赖
|
||||
5. **完善的文档** - README、测试脚本、部署指南
|
||||
|
||||
## 🎯 总结
|
||||
|
||||
Video Generator Agent 已经完成开发和基础测试:
|
||||
|
||||
✅ **已完成**:
|
||||
- 核心功能实现(图片生成、视频拼接、文件管理)
|
||||
- MCP 和 REST API 服务器
|
||||
- Docker 配置和构建
|
||||
- 测试脚本和文档
|
||||
- FFmpeg 集成和验证
|
||||
|
||||
⚠️ **待配置**:
|
||||
- 正确的图片生成模型(当前模型不支持)
|
||||
|
||||
🚀 **可以部署**:
|
||||
- Docker 镜像已就绪
|
||||
- 服务可以启动
|
||||
- 视频拼接功能完整
|
||||
- 只需配置正确的模型即可使用图片生成功能
|
||||
@@ -0,0 +1,33 @@
|
||||
FROM python:3.12-slim
|
||||
|
||||
WORKDIR /app
|
||||
|
||||
ENV PYTHONUNBUFFERED=1
|
||||
ENV PYTHONDONTWRITEBYTECODE=1
|
||||
|
||||
# 安装系统依赖(包括 FFmpeg)
|
||||
RUN apt-get update && apt-get install -y \
|
||||
ffmpeg \
|
||||
gcc \
|
||||
curl \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
|
||||
# 验证 FFmpeg 安装
|
||||
RUN ffmpeg -version
|
||||
|
||||
# 安装 Python 依赖
|
||||
COPY requirements.txt .
|
||||
RUN pip install --no-cache-dir -r requirements.txt
|
||||
|
||||
# 复制应用代码
|
||||
COPY . .
|
||||
|
||||
# 创建输出目录
|
||||
RUN mkdir -p /app/outputs/images /app/outputs/videos
|
||||
|
||||
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,372 @@
|
||||
# Video Generator Agent
|
||||
|
||||
根据文本描述生成图片并拼接为视频的 AI Agent。
|
||||
|
||||
## 功能特性
|
||||
|
||||
- 🎨 **图片生成**: 使用 Gemini 3 Pro Image Preview 根据描述生成高质量图片
|
||||
- 🎬 **视频拼接**: 使用 FFmpeg 将多张图片拼接成视频
|
||||
- ✨ **转场效果**: 支持多种转场效果(fade, wipeleft, wiperight, slideup, slidedown)
|
||||
- 📁 **文件管理**: 自动管理生成的图片和视频文件
|
||||
- 🔌 **MCP 协议**: 支持 MCP (Model Context Protocol) 工具调用
|
||||
- 🌐 **REST API**: 提供完整的 HTTP API 接口
|
||||
|
||||
## 快速开始
|
||||
|
||||
### 1. 本地运行
|
||||
|
||||
```bash
|
||||
# 安装依赖
|
||||
pip install -r requirements.txt
|
||||
|
||||
# 设置环境变量
|
||||
export OPENAI_API_KEY="sk-i9AwAgXDqqxsA9Ym4AjSPg"
|
||||
export MODEL_NAME="taiji/gemini-3-pro-image-preview"
|
||||
|
||||
# 启动服务
|
||||
python run_api_server.py
|
||||
```
|
||||
|
||||
服务将在 `http://localhost:8000` 启动。
|
||||
|
||||
### 2. Docker 运行
|
||||
|
||||
```bash
|
||||
# 构建镜像
|
||||
docker build -t video-generator-agent:latest .
|
||||
|
||||
# 运行容器
|
||||
docker run -d \
|
||||
-p 8000:8000 \
|
||||
-e OPENAI_API_KEY="sk-i9AwAgXDqqxsA9Ym4AjSPg" \
|
||||
-e MODEL_NAME="taiji/gemini-3-pro-image-preview" \
|
||||
-v $(pwd)/outputs:/app/outputs \
|
||||
video-generator-agent:latest
|
||||
```
|
||||
|
||||
### 3. 测试
|
||||
|
||||
```bash
|
||||
# 运行测试脚本
|
||||
python test_video_agent.py
|
||||
|
||||
# 或指定自定义 API URL
|
||||
API_BASE_URL=http://localhost:8000 python test_video_agent.py
|
||||
```
|
||||
|
||||
## API 使用示例
|
||||
|
||||
### 生成单张图片
|
||||
|
||||
```bash
|
||||
curl -X POST http://localhost:8000/api/v1/generate-image \
|
||||
-H "api-key: sk-i9AwAgXDqqxsA9Ym4AjSPg" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"description": "a beautiful sunset over mountains",
|
||||
"size": "1024x1024",
|
||||
"quality": "standard"
|
||||
}'
|
||||
```
|
||||
|
||||
响应示例:
|
||||
```json
|
||||
{
|
||||
"success": true,
|
||||
"file_path": "/app/outputs/images/image_20260226_123456.png",
|
||||
"filename": "image_20260226_123456.png",
|
||||
"url": "/api/v1/files/image_20260226_123456.png",
|
||||
"description": "a beautiful sunset over mountains"
|
||||
}
|
||||
```
|
||||
|
||||
### 生成视频
|
||||
|
||||
```bash
|
||||
curl -X POST http://localhost:8000/api/v1/generate-video \
|
||||
-H "api-key: sk-i9AwAgXDqqxsA9Ym4AjSPg" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"descriptions": [
|
||||
"sunrise over mountains",
|
||||
"a peaceful lake at noon",
|
||||
"starry night sky"
|
||||
],
|
||||
"duration_per_image": 3,
|
||||
"fps": 30,
|
||||
"transition": "fade"
|
||||
}'
|
||||
```
|
||||
|
||||
响应示例:
|
||||
```json
|
||||
{
|
||||
"success": true,
|
||||
"video": {
|
||||
"file_path": "/app/outputs/videos/video_20260226_123456.mp4",
|
||||
"url": "/api/v1/files/video_20260226_123456.mp4",
|
||||
"filename": "video_20260226_123456.mp4",
|
||||
"duration": 9.0,
|
||||
"image_count": 3
|
||||
},
|
||||
"images": [
|
||||
{
|
||||
"index": 1,
|
||||
"description": "sunrise over mountains",
|
||||
"file_path": "/app/outputs/images/image_20260226_123456_1.png",
|
||||
"url": "/api/v1/files/image_20260226_123456_1.png"
|
||||
}
|
||||
],
|
||||
"settings": {
|
||||
"duration_per_image": 3,
|
||||
"fps": 30,
|
||||
"transition": "fade"
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
### 下载文件
|
||||
|
||||
```bash
|
||||
# 下载图片
|
||||
curl -O http://localhost:8000/api/v1/files/image_20260226_123456.png
|
||||
|
||||
# 下载视频
|
||||
curl -O http://localhost:8000/api/v1/files/video_20260226_123456.mp4
|
||||
```
|
||||
|
||||
### 列出文件
|
||||
|
||||
```bash
|
||||
# 列出所有文件
|
||||
curl http://localhost:8000/api/v1/list-files?file_type=all
|
||||
|
||||
# 只列出图片
|
||||
curl http://localhost:8000/api/v1/list-files?file_type=image
|
||||
|
||||
# 只列出视频
|
||||
curl http://localhost:8000/api/v1/list-files?file_type=video
|
||||
```
|
||||
|
||||
## MCP 工具
|
||||
|
||||
Agent 提供以下 MCP 工具:
|
||||
|
||||
### 1. generate_image
|
||||
|
||||
根据描述生成单张图片。
|
||||
|
||||
```json
|
||||
{
|
||||
"name": "generate_image",
|
||||
"arguments": {
|
||||
"description": "a futuristic city at night",
|
||||
"size": "1024x1024",
|
||||
"quality": "standard"
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
### 2. generate_video
|
||||
|
||||
根据多个描述生成视频。
|
||||
|
||||
```json
|
||||
{
|
||||
"name": "generate_video",
|
||||
"arguments": {
|
||||
"descriptions": [
|
||||
"scene 1 description",
|
||||
"scene 2 description",
|
||||
"scene 3 description"
|
||||
],
|
||||
"duration_per_image": 3,
|
||||
"fps": 30,
|
||||
"transition": "fade"
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
### 3. list_generated_files
|
||||
|
||||
列出已生成的文件。
|
||||
|
||||
```json
|
||||
{
|
||||
"name": "list_generated_files",
|
||||
"arguments": {
|
||||
"file_type": "all"
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
### 4. cleanup_old_files
|
||||
|
||||
清理旧文件。
|
||||
|
||||
```json
|
||||
{
|
||||
"name": "cleanup_old_files",
|
||||
"arguments": {
|
||||
"max_age_hours": 24
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
## MCP 端点
|
||||
|
||||
### HTTP 端点
|
||||
|
||||
```bash
|
||||
# 初始化
|
||||
curl -X POST http://localhost:8000/mcp \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"jsonrpc": "2.0",
|
||||
"method": "initialize",
|
||||
"id": 1
|
||||
}'
|
||||
|
||||
# 列出工具
|
||||
curl -X POST http://localhost:8000/mcp \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"jsonrpc": "2.0",
|
||||
"method": "tools/list",
|
||||
"id": 2
|
||||
}'
|
||||
|
||||
# 调用工具
|
||||
curl -X POST http://localhost:8000/mcp \
|
||||
-H "api-key: sk-i9AwAgXDqqxsA9Ym4AjSPg" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"jsonrpc": "2.0",
|
||||
"method": "tools/call",
|
||||
"params": {
|
||||
"name": "generate_image",
|
||||
"arguments": {
|
||||
"description": "a beautiful landscape"
|
||||
}
|
||||
},
|
||||
"id": 3
|
||||
}'
|
||||
```
|
||||
|
||||
### SSE 端点
|
||||
|
||||
```bash
|
||||
# 连接 SSE
|
||||
curl -N http://localhost:8000/mcp/sse
|
||||
```
|
||||
|
||||
## 环境变量
|
||||
|
||||
| 变量 | 必需 | 默认值 | 说明 |
|
||||
|------|------|--------|------|
|
||||
| `OPENAI_API_KEY` | 是 | `sk-i9AwAgXDqqxsA9Ym4AjSPg` | API Key |
|
||||
| `OPENAI_BASE_URL` | 否 | LiteLLM Gateway URL | API Base URL |
|
||||
| `MODEL_NAME` | 否 | `taiji/gemini-3-pro-image-preview` | 模型名称 |
|
||||
| `API_PORT` | 否 | `8000` | 服务端口 |
|
||||
| `OUTPUT_DIR` | 否 | `/app/outputs` | 输出目录 |
|
||||
|
||||
## 转场效果
|
||||
|
||||
支持的转场效果:
|
||||
|
||||
- `fade` - 淡入淡出
|
||||
- `wipeleft` - 左擦除
|
||||
- `wiperight` - 右擦除
|
||||
- `slideup` - 上滑
|
||||
- `slidedown` - 下滑
|
||||
|
||||
## 图片尺寸
|
||||
|
||||
支持的图片尺寸:
|
||||
|
||||
- `256x256`
|
||||
- `512x512`
|
||||
- `1024x1024` (默认)
|
||||
- `1792x1024`
|
||||
- `1024x1792`
|
||||
|
||||
## 项目结构
|
||||
|
||||
```
|
||||
video_generator_agent/
|
||||
├── Dockerfile # Docker 配置
|
||||
├── requirements.txt # Python 依赖
|
||||
├── run_api_server.py # 启动脚本
|
||||
├── test_video_agent.py # 测试脚本
|
||||
├── README.md # 文档
|
||||
└── src/
|
||||
├── __init__.py
|
||||
├── utils/ # 工具模块
|
||||
│ ├── __init__.py
|
||||
│ ├── image_generator.py # 图片生成
|
||||
│ ├── video_processor.py # 视频处理
|
||||
│ └── file_manager.py # 文件管理
|
||||
└── server/ # 服务器模块
|
||||
├── __init__.py
|
||||
├── mcp_server.py # MCP 工具定义
|
||||
└── api_server.py # FastAPI 服务器
|
||||
```
|
||||
|
||||
## 技术栈
|
||||
|
||||
- **Pydantic AI**: Agent 框架
|
||||
- **FastMCP**: MCP 服务器
|
||||
- **FastAPI**: REST API 框架
|
||||
- **FFmpeg**: 视频处理
|
||||
- **Gemini 3 Pro Image Preview**: 图片生成模型
|
||||
- **aiohttp**: 异步 HTTP 客户端
|
||||
- **Pillow**: 图片处理
|
||||
|
||||
## 故障排除
|
||||
|
||||
### FFmpeg 未安装
|
||||
|
||||
如果遇到 FFmpeg 相关错误,请确保已安装 FFmpeg:
|
||||
|
||||
```bash
|
||||
# Ubuntu/Debian
|
||||
sudo apt-get install ffmpeg
|
||||
|
||||
# macOS
|
||||
brew install ffmpeg
|
||||
|
||||
# 验证安装
|
||||
ffmpeg -version
|
||||
```
|
||||
|
||||
### API Key 错误
|
||||
|
||||
确保设置了正确的 API Key:
|
||||
|
||||
```bash
|
||||
export OPENAI_API_KEY="sk-i9AwAgXDqqxsA9Ym4AjSPg"
|
||||
```
|
||||
|
||||
### 端口被占用
|
||||
|
||||
如果端口 8000 被占用,可以更改端口:
|
||||
|
||||
```bash
|
||||
export API_PORT=8080
|
||||
python run_api_server.py
|
||||
```
|
||||
|
||||
## 性能建议
|
||||
|
||||
- 图片生成通常需要 10-30 秒
|
||||
- 视频拼接时间取决于图片数量(每张图片约 2-5 秒)
|
||||
- 建议每个视频不超过 10 个场景
|
||||
- 使用转场效果会增加处理时间
|
||||
|
||||
## 许可证
|
||||
|
||||
MIT License
|
||||
|
||||
## 作者
|
||||
|
||||
Video Generator Agent Team
|
||||
@@ -0,0 +1,19 @@
|
||||
# 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
|
||||
|
||||
# Image Processing
|
||||
Pillow>=10.0.0
|
||||
|
||||
# OpenAI (for image generation)
|
||||
openai>=1.0.0
|
||||
@@ -0,0 +1,20 @@
|
||||
#!/usr/bin/env python
|
||||
"""启动 Video Generator 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"🚀 启动 Video Generator Agent API: http://{host}:{port}")
|
||||
print(f"📖 API 文档: http://{host}:{port}/docs")
|
||||
print(f"🔧 MCP 端点: http://{host}:{port}/mcp")
|
||||
|
||||
uvicorn.run(app, host=host, port=port, log_level="info")
|
||||
@@ -0,0 +1,2 @@
|
||||
"""Video Generator Agent - 视频生成 Agent"""
|
||||
__version__ = "1.0.0"
|
||||
@@ -0,0 +1 @@
|
||||
"""Server modules"""
|
||||
@@ -0,0 +1,345 @@
|
||||
"""
|
||||
HTTP API 服务器 - 视频生成 Agent
|
||||
|
||||
提供 REST API 和 MCP HTTP/SSE 端点。
|
||||
"""
|
||||
import json
|
||||
import uuid
|
||||
import os
|
||||
import sys
|
||||
from typing import Optional, Dict, Any, AsyncGenerator, List
|
||||
from contextlib import asynccontextmanager
|
||||
from pathlib import Path
|
||||
|
||||
# 添加父目录到路径
|
||||
sys.path.insert(0, str(Path(__file__).parent.parent.parent))
|
||||
|
||||
from fastapi import FastAPI, HTTPException, Request, Header, Depends
|
||||
from fastapi.middleware.cors import CORSMiddleware
|
||||
from fastapi.responses import StreamingResponse, JSONResponse, FileResponse
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from src.server.mcp_server import TOOL_MAP, TOOL_LIST
|
||||
from src.utils.file_manager import FileManager
|
||||
|
||||
# ==================== 配置 ====================
|
||||
|
||||
SERVER_NAME = "Video Generator Agent"
|
||||
OUTPUT_DIR = os.getenv('OUTPUT_DIR', '/app/outputs')
|
||||
|
||||
file_manager = FileManager(base_dir=OUTPUT_DIR)
|
||||
|
||||
# ==================== FastAPI 应用 ====================
|
||||
|
||||
@asynccontextmanager
|
||||
async def lifespan(app: FastAPI):
|
||||
print(f"🚀 {SERVER_NAME} 启动")
|
||||
print(f"📁 输出目录: {OUTPUT_DIR}")
|
||||
yield
|
||||
print(f"🛑 {SERVER_NAME} 关闭")
|
||||
|
||||
app = FastAPI(
|
||||
title=SERVER_NAME,
|
||||
description="根据描述生成图片并拼接为视频",
|
||||
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(不验证)"""
|
||||
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():
|
||||
storage_info = file_manager.get_storage_info()
|
||||
return {
|
||||
"service": SERVER_NAME,
|
||||
"status": "running",
|
||||
"version": "1.0.0",
|
||||
"tools": list(TOOL_MAP.keys()),
|
||||
"storage": storage_info
|
||||
}
|
||||
|
||||
|
||||
@app.get("/health")
|
||||
async def health():
|
||||
return {"status": "healthy", "service": SERVER_NAME}
|
||||
|
||||
|
||||
# ==================== MCP 端点 ====================
|
||||
|
||||
sessions: Dict[str, Dict] = {}
|
||||
|
||||
|
||||
async def handle_mcp_request(data: Dict, session_id: str = None, api_key: str = None) -> Dict:
|
||||
"""处理 MCP JSON-RPC 请求"""
|
||||
method = data.get("method")
|
||||
params = data.get("params", {})
|
||||
req_id = data.get("id")
|
||||
|
||||
# tools/call 需要验证 API Key
|
||||
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")
|
||||
args = params.get("arguments", {})
|
||||
|
||||
if tool_name not in TOOL_MAP:
|
||||
raise ValueError(f"Unknown tool: {tool_name}")
|
||||
|
||||
# 设置 API Key 到环境变量
|
||||
old_key = os.environ.get('OPENAI_API_KEY')
|
||||
if api_key:
|
||||
os.environ['OPENAI_API_KEY'] = api_key
|
||||
|
||||
try:
|
||||
result = await TOOL_MAP[tool_name](**args)
|
||||
finally:
|
||||
if old_key:
|
||||
os.environ['OPENAI_API_KEY'] = old_key
|
||||
|
||||
return {
|
||||
"jsonrpc": "2.0", "id": req_id,
|
||||
"result": {"content": [{"type": "text", "text": str(result)}]}
|
||||
}
|
||||
|
||||
elif method == "ping":
|
||||
return {"jsonrpc": "2.0", "id": req_id, "result": {}}
|
||||
|
||||
else:
|
||||
raise ValueError(f"Unknown method: {method}")
|
||||
|
||||
except Exception as e:
|
||||
return {"jsonrpc": "2.0", "id": req_id, "error": {"code": -32603, "message": str(e)}}
|
||||
|
||||
|
||||
@app.post("/mcp")
|
||||
async def mcp_endpoint(request: Request):
|
||||
"""MCP HTTP 端点"""
|
||||
try:
|
||||
body = await request.json()
|
||||
session_id = request.headers.get("x-mcp-session-id")
|
||||
api_key = get_api_key_from_request(request)
|
||||
response = await handle_mcp_request(body, session_id, api_key)
|
||||
return JSONResponse(content=response, headers={"x-mcp-session-id": session_id or ""})
|
||||
except Exception as e:
|
||||
return JSONResponse(status_code=400, content={"jsonrpc": "2.0", "error": {"code": -32700, "message": str(e)}})
|
||||
|
||||
|
||||
@app.get("/mcp/sse")
|
||||
async def mcp_sse(request: Request):
|
||||
"""MCP SSE 端点"""
|
||||
session_id = request.headers.get("x-mcp-session-id") or str(uuid.uuid4())
|
||||
|
||||
async def stream() -> AsyncGenerator[str, None]:
|
||||
yield f"data: {json.dumps({'type': 'connection', 'sessionId': session_id})}\n\n"
|
||||
import asyncio
|
||||
while True:
|
||||
await asyncio.sleep(30)
|
||||
yield f"data: {json.dumps({'type': 'ping'})}\n\n"
|
||||
|
||||
return StreamingResponse(stream(), media_type="text/event-stream",
|
||||
headers={"Cache-Control": "no-cache", "x-mcp-session-id": session_id})
|
||||
|
||||
|
||||
@app.post("/mcp/sse")
|
||||
async def mcp_sse_post(request: Request):
|
||||
"""MCP SSE POST 端点"""
|
||||
try:
|
||||
body = await request.json()
|
||||
session_id = request.headers.get("x-mcp-session-id") or str(uuid.uuid4())
|
||||
api_key = get_api_key_from_request(request)
|
||||
|
||||
async def stream() -> AsyncGenerator[str, None]:
|
||||
response = await handle_mcp_request(body, session_id, api_key)
|
||||
yield f"data: {json.dumps(response)}\n\n"
|
||||
|
||||
return StreamingResponse(stream(), media_type="text/event-stream",
|
||||
headers={"Cache-Control": "no-cache", "x-mcp-session-id": session_id})
|
||||
except Exception as e:
|
||||
return JSONResponse(status_code=400, content={"jsonrpc": "2.0", "error": {"code": -32700, "message": str(e)}})
|
||||
|
||||
|
||||
# ==================== 业务 API ====================
|
||||
|
||||
class ImageGenerationRequest(BaseModel):
|
||||
"""图片生成请求"""
|
||||
description: str = Field(..., description="图片描述")
|
||||
size: str = Field("1024x1024", description="图片尺寸")
|
||||
quality: str = Field("standard", description="图片质量")
|
||||
|
||||
|
||||
class VideoGenerationRequest(BaseModel):
|
||||
"""视频生成请求"""
|
||||
descriptions: List[str] = Field(..., description="场景描述列表")
|
||||
duration_per_image: int = Field(3, description="每张图片显示秒数")
|
||||
fps: int = Field(30, description="视频帧率")
|
||||
transition: Optional[str] = Field(None, description="转场效果")
|
||||
|
||||
|
||||
@app.post("/api/v1/generate-image")
|
||||
async def api_generate_image(request: ImageGenerationRequest, api_key: str = Depends(verify_api_key)):
|
||||
"""生成单张图片"""
|
||||
try:
|
||||
# 设置 API Key
|
||||
old_key = os.environ.get('OPENAI_API_KEY')
|
||||
os.environ['OPENAI_API_KEY'] = api_key
|
||||
|
||||
try:
|
||||
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"):
|
||||
raise HTTPException(status_code=500, detail=result.get("error", "Unknown error"))
|
||||
|
||||
return result
|
||||
finally:
|
||||
if old_key:
|
||||
os.environ['OPENAI_API_KEY'] = old_key
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
@app.post("/api/v1/generate-video")
|
||||
async def api_generate_video(request: VideoGenerationRequest, api_key: str = Depends(verify_api_key)):
|
||||
"""生成视频"""
|
||||
try:
|
||||
# 设置 API Key
|
||||
old_key = os.environ.get('OPENAI_API_KEY')
|
||||
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
|
||||
)
|
||||
result = json.loads(result_str)
|
||||
|
||||
if not result.get("success"):
|
||||
raise HTTPException(status_code=500, detail=result.get("error", "Unknown error"))
|
||||
|
||||
return result
|
||||
finally:
|
||||
if old_key:
|
||||
os.environ['OPENAI_API_KEY'] = old_key
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
@app.get("/api/v1/list-files")
|
||||
async def api_list_files(file_type: str = "all"):
|
||||
"""列出文件"""
|
||||
try:
|
||||
result_str = await TOOL_MAP['list_generated_files'](file_type=file_type)
|
||||
result = json.loads(result_str)
|
||||
return result
|
||||
except Exception as e:
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
@app.get("/api/v1/files/{filename}")
|
||||
async def api_download_file(filename: str):
|
||||
"""下载文件"""
|
||||
try:
|
||||
file_path = file_manager.get_file_path(filename)
|
||||
|
||||
if not file_path or not os.path.exists(file_path):
|
||||
raise HTTPException(status_code=404, detail="File not found")
|
||||
|
||||
# 确定 media type
|
||||
if filename.endswith('.mp4'):
|
||||
media_type = "video/mp4"
|
||||
elif filename.endswith('.png'):
|
||||
media_type = "image/png"
|
||||
elif filename.endswith('.jpg') or filename.endswith('.jpeg'):
|
||||
media_type = "image/jpeg"
|
||||
else:
|
||||
media_type = "application/octet-stream"
|
||||
|
||||
return FileResponse(
|
||||
file_path,
|
||||
media_type=media_type,
|
||||
filename=filename
|
||||
)
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
@app.post("/api/v1/cleanup")
|
||||
async def api_cleanup(max_age_hours: int = 24):
|
||||
"""清理旧文件"""
|
||||
try:
|
||||
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:
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
import uvicorn
|
||||
uvicorn.run(app, host="0.0.0.0", port=8000)
|
||||
@@ -0,0 +1,308 @@
|
||||
"""
|
||||
MCP 服务器 - 视频生成 Agent 工具
|
||||
"""
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
from typing import Optional, List
|
||||
from pathlib import Path
|
||||
|
||||
# 添加父目录到路径
|
||||
sys.path.insert(0, str(Path(__file__).parent.parent.parent))
|
||||
|
||||
from mcp.server.fastmcp import FastMCP
|
||||
from pydantic_ai import Agent
|
||||
|
||||
from src.utils.image_generator import ImageGenerator
|
||||
from src.utils.video_processor import VideoProcessor
|
||||
from src.utils.file_manager import FileManager
|
||||
|
||||
# ==================== 配置 ====================
|
||||
|
||||
# LiteLLM Gateway 配置
|
||||
_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-i9AwAgXDqqxsA9Ym4AjSPg')
|
||||
|
||||
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/gemini-3-pro-image-preview'))
|
||||
return model if ':' in model else f'openai:{model}'
|
||||
|
||||
MODEL_NAME = _get_model_name()
|
||||
|
||||
# 输出目录
|
||||
OUTPUT_DIR = os.getenv('OUTPUT_DIR', '/app/outputs')
|
||||
|
||||
# ==================== 初始化工具 ====================
|
||||
|
||||
image_generator = ImageGenerator(api_key=_API_KEY, base_url=_BASE_URL, model='taiji/gemini-3-pro-image-preview')
|
||||
video_processor = VideoProcessor()
|
||||
file_manager = FileManager(base_dir=OUTPUT_DIR)
|
||||
|
||||
# ==================== MCP 服务器 ====================
|
||||
|
||||
server = FastMCP('Video Generator Agent')
|
||||
|
||||
SYSTEM_PROMPT = '''你是一个专业的视频生成 AI 助手。
|
||||
你可以根据用户的描述生成图片,并将多张图片拼接成视频。
|
||||
请根据用户的需求提供帮助。'''
|
||||
|
||||
|
||||
def get_agent() -> Agent:
|
||||
"""创建 Agent 实例"""
|
||||
return Agent(MODEL_NAME, system_prompt=SYSTEM_PROMPT)
|
||||
|
||||
|
||||
# ==================== MCP 工具定义 ====================
|
||||
|
||||
@server.tool()
|
||||
async def generate_image(
|
||||
description: str,
|
||||
size: str = "1024x1024",
|
||||
quality: str = "standard"
|
||||
) -> str:
|
||||
"""
|
||||
根据描述生成单张图片
|
||||
|
||||
Args:
|
||||
description: 图片描述(英文效果更好)
|
||||
size: 图片尺寸,可选 256x256, 512x512, 1024x1024, 1792x1024, 1024x1792
|
||||
quality: 图片质量,可选 standard, hd
|
||||
|
||||
Returns:
|
||||
JSON 格式的结果,包含图片路径和下载 URL
|
||||
"""
|
||||
try:
|
||||
result = await image_generator.generate_image(
|
||||
description=description,
|
||||
output_dir=os.path.join(OUTPUT_DIR, "images"),
|
||||
size=size,
|
||||
quality=quality
|
||||
)
|
||||
|
||||
return json.dumps(result, ensure_ascii=False, indent=2)
|
||||
|
||||
except Exception as e:
|
||||
return json.dumps({
|
||||
"success": False,
|
||||
"error": str(e)
|
||||
}, ensure_ascii=False)
|
||||
|
||||
|
||||
@server.tool()
|
||||
async def generate_video(
|
||||
descriptions: List[str],
|
||||
duration_per_image: int = 3,
|
||||
fps: int = 30,
|
||||
transition: Optional[str] = None
|
||||
) -> str:
|
||||
"""
|
||||
根据多个描述生成视频
|
||||
|
||||
Args:
|
||||
descriptions: 场景描述列表(每个描述对应一张图片)
|
||||
duration_per_image: 每张图片显示秒数(默认 3 秒)
|
||||
fps: 视频帧率(默认 30)
|
||||
transition: 转场效果,可选 fade, wipeleft, wiperight, slideup, slidedown
|
||||
|
||||
Returns:
|
||||
JSON 格式的结果,包含视频路径和下载 URL
|
||||
"""
|
||||
try:
|
||||
if not descriptions or len(descriptions) == 0:
|
||||
return json.dumps({
|
||||
"success": False,
|
||||
"error": "至少需要提供一个场景描述"
|
||||
}, ensure_ascii=False)
|
||||
|
||||
# 1. 生成所有图片
|
||||
image_paths = []
|
||||
generated_images = []
|
||||
|
||||
for i, desc in enumerate(descriptions):
|
||||
print(f"生成图片 {i+1}/{len(descriptions)}: {desc[:50]}...")
|
||||
|
||||
result = await image_generator.generate_image(
|
||||
description=desc,
|
||||
output_dir=os.path.join(OUTPUT_DIR, "images")
|
||||
)
|
||||
|
||||
if not result.get("success"):
|
||||
return json.dumps({
|
||||
"success": False,
|
||||
"error": f"生成第 {i+1} 张图片失败: {result.get('error')}",
|
||||
"generated_images": generated_images
|
||||
}, ensure_ascii=False)
|
||||
|
||||
image_paths.append(result["file_path"])
|
||||
generated_images.append({
|
||||
"index": i + 1,
|
||||
"description": desc,
|
||||
"file_path": result["file_path"],
|
||||
"url": result["url"]
|
||||
})
|
||||
|
||||
# 2. 拼接视频
|
||||
print(f"拼接视频: {len(image_paths)} 张图片...")
|
||||
|
||||
video_result = await video_processor.create_video_from_images(
|
||||
image_paths=image_paths,
|
||||
output_dir=os.path.join(OUTPUT_DIR, "videos"),
|
||||
duration_per_image=duration_per_image,
|
||||
fps=fps,
|
||||
transition=transition
|
||||
)
|
||||
|
||||
if not video_result.get("success"):
|
||||
return json.dumps({
|
||||
"success": False,
|
||||
"error": f"视频拼接失败: {video_result.get('error')}",
|
||||
"generated_images": generated_images
|
||||
}, ensure_ascii=False)
|
||||
|
||||
# 3. 返回结果
|
||||
return json.dumps({
|
||||
"success": True,
|
||||
"video": {
|
||||
"file_path": video_result["file_path"],
|
||||
"url": video_result["url"],
|
||||
"filename": video_result["filename"],
|
||||
"duration": video_result["duration"],
|
||||
"image_count": video_result["image_count"]
|
||||
},
|
||||
"images": generated_images,
|
||||
"settings": {
|
||||
"duration_per_image": duration_per_image,
|
||||
"fps": fps,
|
||||
"transition": transition
|
||||
}
|
||||
}, ensure_ascii=False, indent=2)
|
||||
|
||||
except Exception as e:
|
||||
return json.dumps({
|
||||
"success": False,
|
||||
"error": str(e)
|
||||
}, ensure_ascii=False)
|
||||
|
||||
|
||||
@server.tool()
|
||||
async def list_generated_files(file_type: str = "all") -> str:
|
||||
"""
|
||||
列出已生成的文件
|
||||
|
||||
Args:
|
||||
file_type: 文件类型,可选 all, image, video
|
||||
|
||||
Returns:
|
||||
JSON 格式的文件列表
|
||||
"""
|
||||
try:
|
||||
files = file_manager.list_files(file_type=file_type)
|
||||
storage_info = file_manager.get_storage_info()
|
||||
|
||||
return json.dumps({
|
||||
"success": True,
|
||||
"files": files,
|
||||
"storage": storage_info,
|
||||
"total_count": len(files)
|
||||
}, ensure_ascii=False, indent=2)
|
||||
|
||||
except Exception as e:
|
||||
return json.dumps({
|
||||
"success": False,
|
||||
"error": str(e)
|
||||
}, ensure_ascii=False)
|
||||
|
||||
|
||||
@server.tool()
|
||||
async def cleanup_old_files(max_age_hours: int = 24) -> str:
|
||||
"""
|
||||
清理旧文件
|
||||
|
||||
Args:
|
||||
max_age_hours: 保留最近多少小时的文件(默认 24 小时)
|
||||
|
||||
Returns:
|
||||
JSON 格式的清理结果
|
||||
"""
|
||||
try:
|
||||
result = file_manager.cleanup_old_files(max_age_hours=max_age_hours)
|
||||
return json.dumps(result, ensure_ascii=False, indent=2)
|
||||
|
||||
except Exception as e:
|
||||
return json.dumps({
|
||||
"success": False,
|
||||
"error": str(e)
|
||||
}, ensure_ascii=False)
|
||||
|
||||
|
||||
# ==================== 工具映射(供 API 使用)====================
|
||||
|
||||
TOOL_MAP = {
|
||||
'generate_image': generate_image,
|
||||
'generate_video': generate_video,
|
||||
'list_generated_files': list_generated_files,
|
||||
'cleanup_old_files': cleanup_old_files,
|
||||
}
|
||||
|
||||
TOOL_LIST = [
|
||||
{
|
||||
"name": "generate_image",
|
||||
"description": "根据描述生成单张图片",
|
||||
"inputSchema": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"description": {"type": "string", "description": "图片描述"},
|
||||
"size": {"type": "string", "description": "图片尺寸", "default": "1024x1024"},
|
||||
"quality": {"type": "string", "description": "图片质量", "default": "standard"}
|
||||
},
|
||||
"required": ["description"]
|
||||
}
|
||||
},
|
||||
{
|
||||
"name": "generate_video",
|
||||
"description": "根据多个描述生成视频",
|
||||
"inputSchema": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"descriptions": {
|
||||
"type": "array",
|
||||
"items": {"type": "string"},
|
||||
"description": "场景描述列表"
|
||||
},
|
||||
"duration_per_image": {"type": "integer", "description": "每张图片显示秒数", "default": 3},
|
||||
"fps": {"type": "integer", "description": "视频帧率", "default": 30},
|
||||
"transition": {"type": "string", "description": "转场效果"}
|
||||
},
|
||||
"required": ["descriptions"]
|
||||
}
|
||||
},
|
||||
{
|
||||
"name": "list_generated_files",
|
||||
"description": "列出已生成的文件",
|
||||
"inputSchema": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"file_type": {"type": "string", "description": "文件类型 (all/image/video)", "default": "all"}
|
||||
}
|
||||
}
|
||||
},
|
||||
{
|
||||
"name": "cleanup_old_files",
|
||||
"description": "清理旧文件",
|
||||
"inputSchema": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"max_age_hours": {"type": "integer", "description": "保留最近多少小时的文件", "default": 24}
|
||||
}
|
||||
}
|
||||
}
|
||||
]
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
server.run()
|
||||
@@ -0,0 +1,6 @@
|
||||
"""Utility modules for video generation"""
|
||||
from .image_generator import ImageGenerator
|
||||
from .video_processor import VideoProcessor
|
||||
from .file_manager import FileManager
|
||||
|
||||
__all__ = ['ImageGenerator', 'VideoProcessor', 'FileManager']
|
||||
@@ -0,0 +1,218 @@
|
||||
"""
|
||||
文件管理工具
|
||||
管理生成的图片和视频文件
|
||||
"""
|
||||
import os
|
||||
import logging
|
||||
from pathlib import Path
|
||||
from typing import List, Dict, Optional
|
||||
from datetime import datetime, timedelta
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class FileManager:
|
||||
"""文件管理器"""
|
||||
|
||||
def __init__(self, base_dir: str = "/app/outputs"):
|
||||
"""
|
||||
初始化文件管理器
|
||||
|
||||
Args:
|
||||
base_dir: 基础目录
|
||||
"""
|
||||
self.base_dir = base_dir
|
||||
self.images_dir = os.path.join(base_dir, "images")
|
||||
self.videos_dir = os.path.join(base_dir, "videos")
|
||||
|
||||
# 创建目录
|
||||
Path(self.images_dir).mkdir(parents=True, exist_ok=True)
|
||||
Path(self.videos_dir).mkdir(parents=True, exist_ok=True)
|
||||
|
||||
logger.info(f"FileManager initialized: base_dir={base_dir}")
|
||||
|
||||
def get_file_path(self, filename: str) -> Optional[str]:
|
||||
"""
|
||||
获取文件完整路径
|
||||
|
||||
Args:
|
||||
filename: 文件名
|
||||
|
||||
Returns:
|
||||
str: 文件路径,如果不存在返回 None
|
||||
"""
|
||||
# 先在 images 目录查找
|
||||
image_path = os.path.join(self.images_dir, filename)
|
||||
if os.path.exists(image_path):
|
||||
return image_path
|
||||
|
||||
# 再在 videos 目录查找
|
||||
video_path = os.path.join(self.videos_dir, filename)
|
||||
if os.path.exists(video_path):
|
||||
return video_path
|
||||
|
||||
return None
|
||||
|
||||
def list_files(self, file_type: str = "all") -> List[Dict]:
|
||||
"""
|
||||
列出文件
|
||||
|
||||
Args:
|
||||
file_type: 文件类型 (all, image, video)
|
||||
|
||||
Returns:
|
||||
List[Dict]: 文件列表
|
||||
"""
|
||||
files = []
|
||||
|
||||
try:
|
||||
if file_type in ["all", "image"]:
|
||||
files.extend(self._list_directory(self.images_dir, "image"))
|
||||
|
||||
if file_type in ["all", "video"]:
|
||||
files.extend(self._list_directory(self.videos_dir, "video"))
|
||||
|
||||
# 按修改时间排序(最新的在前)
|
||||
files.sort(key=lambda x: x["modified_time"], reverse=True)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error listing files: {e}")
|
||||
|
||||
return files
|
||||
|
||||
def _list_directory(self, directory: str, file_type: str) -> List[Dict]:
|
||||
"""
|
||||
列出目录中的文件
|
||||
|
||||
Args:
|
||||
directory: 目录路径
|
||||
file_type: 文件类型标签
|
||||
|
||||
Returns:
|
||||
List[Dict]: 文件信息列表
|
||||
"""
|
||||
files = []
|
||||
|
||||
try:
|
||||
if not os.path.exists(directory):
|
||||
return files
|
||||
|
||||
for filename in os.listdir(directory):
|
||||
file_path = os.path.join(directory, filename)
|
||||
|
||||
if not os.path.isfile(file_path):
|
||||
continue
|
||||
|
||||
stat = os.stat(file_path)
|
||||
|
||||
files.append({
|
||||
"filename": filename,
|
||||
"type": file_type,
|
||||
"size": stat.st_size,
|
||||
"size_mb": round(stat.st_size / 1024 / 1024, 2),
|
||||
"modified_time": stat.st_mtime,
|
||||
"modified_date": datetime.fromtimestamp(stat.st_mtime).isoformat(),
|
||||
"url": f"/api/v1/files/{filename}"
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error listing directory {directory}: {e}")
|
||||
|
||||
return files
|
||||
|
||||
def cleanup_old_files(self, max_age_hours: int = 24) -> Dict:
|
||||
"""
|
||||
清理旧文件
|
||||
|
||||
Args:
|
||||
max_age_hours: 最大保留时间(小时)
|
||||
|
||||
Returns:
|
||||
Dict: 清理结果
|
||||
"""
|
||||
try:
|
||||
cutoff_time = datetime.now() - timedelta(hours=max_age_hours)
|
||||
cutoff_timestamp = cutoff_time.timestamp()
|
||||
|
||||
deleted_count = 0
|
||||
freed_space = 0
|
||||
|
||||
# 清理图片
|
||||
for filename in os.listdir(self.images_dir):
|
||||
file_path = os.path.join(self.images_dir, filename)
|
||||
if os.path.isfile(file_path):
|
||||
stat = os.stat(file_path)
|
||||
if stat.st_mtime < cutoff_timestamp:
|
||||
freed_space += stat.st_size
|
||||
os.remove(file_path)
|
||||
deleted_count += 1
|
||||
logger.info(f"Deleted old image: {filename}")
|
||||
|
||||
# 清理视频
|
||||
for filename in os.listdir(self.videos_dir):
|
||||
file_path = os.path.join(self.videos_dir, filename)
|
||||
if os.path.isfile(file_path):
|
||||
stat = os.stat(file_path)
|
||||
if stat.st_mtime < cutoff_timestamp:
|
||||
freed_space += stat.st_size
|
||||
os.remove(file_path)
|
||||
deleted_count += 1
|
||||
logger.info(f"Deleted old video: {filename}")
|
||||
|
||||
return {
|
||||
"success": True,
|
||||
"deleted_count": deleted_count,
|
||||
"freed_space_mb": round(freed_space / 1024 / 1024, 2)
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error cleaning up files: {e}")
|
||||
return {
|
||||
"success": False,
|
||||
"error": str(e)
|
||||
}
|
||||
|
||||
def get_storage_info(self) -> Dict:
|
||||
"""
|
||||
获取存储信息
|
||||
|
||||
Returns:
|
||||
Dict: 存储信息
|
||||
"""
|
||||
try:
|
||||
total_size = 0
|
||||
image_count = 0
|
||||
video_count = 0
|
||||
|
||||
# 统计图片
|
||||
if os.path.exists(self.images_dir):
|
||||
for filename in os.listdir(self.images_dir):
|
||||
file_path = os.path.join(self.images_dir, filename)
|
||||
if os.path.isfile(file_path):
|
||||
total_size += os.path.getsize(file_path)
|
||||
image_count += 1
|
||||
|
||||
# 统计视频
|
||||
if os.path.exists(self.videos_dir):
|
||||
for filename in os.listdir(self.videos_dir):
|
||||
file_path = os.path.join(self.videos_dir, filename)
|
||||
if os.path.isfile(file_path):
|
||||
total_size += os.path.getsize(file_path)
|
||||
video_count += 1
|
||||
|
||||
return {
|
||||
"total_size_mb": round(total_size / 1024 / 1024, 2),
|
||||
"image_count": image_count,
|
||||
"video_count": video_count,
|
||||
"total_files": image_count + video_count
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error getting storage info: {e}")
|
||||
return {
|
||||
"total_size_mb": 0,
|
||||
"image_count": 0,
|
||||
"video_count": 0,
|
||||
"total_files": 0,
|
||||
"error": str(e)
|
||||
}
|
||||
@@ -0,0 +1,197 @@
|
||||
"""
|
||||
图片生成工具
|
||||
使用 OpenAI 兼容接口调用 Gemini 3 Pro Image Preview
|
||||
"""
|
||||
import os
|
||||
import aiohttp
|
||||
import logging
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
from datetime import datetime
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class ImageGenerator:
|
||||
"""图片生成器"""
|
||||
|
||||
def __init__(self, api_key: str = None, base_url: str = None, model: str = None):
|
||||
"""
|
||||
初始化图片生成器
|
||||
|
||||
Args:
|
||||
api_key: API Key
|
||||
base_url: API Base URL
|
||||
model: 模型名称
|
||||
"""
|
||||
self.api_key = api_key or os.getenv('OPENAI_API_KEY', 'sk')
|
||||
self.base_url = base_url or os.getenv('OPENAI_BASE_URL',
|
||||
'https://litellm.graystone-fb459c5d.southeastasia.azurecontainerapps.io/v1')
|
||||
self.model = model or os.getenv('MODEL_NAME', 'taiji/gemini-3-pro-image-preview')
|
||||
|
||||
# 确保 base_url 不以 /v1 结尾(我们会手动添加)
|
||||
if self.base_url.endswith('/v1'):
|
||||
self.base_url = self.base_url[:-3]
|
||||
|
||||
logger.info(f"ImageGenerator initialized: model={self.model}, base_url={self.base_url}")
|
||||
|
||||
async def generate_image(
|
||||
self,
|
||||
description: str,
|
||||
output_dir: str = "/app/outputs/images",
|
||||
size: str = "1024x1024",
|
||||
quality: str = "standard"
|
||||
) -> dict:
|
||||
"""
|
||||
生成图片
|
||||
|
||||
Args:
|
||||
description: 图片描述
|
||||
output_dir: 输出目录
|
||||
size: 图片尺寸
|
||||
quality: 图片质量
|
||||
|
||||
Returns:
|
||||
dict: {
|
||||
"success": bool,
|
||||
"file_path": str,
|
||||
"url": str,
|
||||
"description": str
|
||||
}
|
||||
"""
|
||||
try:
|
||||
# 创建输出目录
|
||||
Path(output_dir).mkdir(parents=True, exist_ok=True)
|
||||
|
||||
# 调用 API 生成图片
|
||||
logger.info(f"Generating image: {description[:50]}...")
|
||||
|
||||
async with aiohttp.ClientSession() as session:
|
||||
# 使用 OpenAI 兼容的图片生成接口
|
||||
url = f"{self.base_url}/v1/images/generations"
|
||||
headers = {
|
||||
"Authorization": f"Bearer {self.api_key}",
|
||||
"Content-Type": "application/json"
|
||||
}
|
||||
payload = {
|
||||
"model": self.model,
|
||||
"prompt": description,
|
||||
"n": 1,
|
||||
"size": size,
|
||||
"quality": quality
|
||||
}
|
||||
|
||||
async with session.post(url, json=payload, headers=headers, timeout=aiohttp.ClientTimeout(total=60)) as response:
|
||||
if response.status != 200:
|
||||
error_text = await response.text()
|
||||
logger.error(f"API error: {response.status} - {error_text}")
|
||||
return {
|
||||
"success": False,
|
||||
"error": f"API error: {response.status} - {error_text}"
|
||||
}
|
||||
|
||||
data = await response.json()
|
||||
|
||||
# 获取图片 URL
|
||||
if not data.get("data") or len(data["data"]) == 0:
|
||||
return {
|
||||
"success": False,
|
||||
"error": "No image data returned from API"
|
||||
}
|
||||
|
||||
image_url = data["data"][0].get("url")
|
||||
if not image_url:
|
||||
return {
|
||||
"success": False,
|
||||
"error": "No image URL in response"
|
||||
}
|
||||
|
||||
# 下载图片
|
||||
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S_%f")
|
||||
filename = f"image_{timestamp}.png"
|
||||
file_path = os.path.join(output_dir, filename)
|
||||
|
||||
download_result = await self._download_image(image_url, file_path)
|
||||
if not download_result["success"]:
|
||||
return download_result
|
||||
|
||||
logger.info(f"Image generated successfully: {file_path}")
|
||||
|
||||
return {
|
||||
"success": True,
|
||||
"file_path": file_path,
|
||||
"filename": filename,
|
||||
"url": f"/api/v1/files/{filename}",
|
||||
"description": description
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error generating image: {e}", exc_info=True)
|
||||
return {
|
||||
"success": False,
|
||||
"error": str(e)
|
||||
}
|
||||
|
||||
async def _download_image(self, url: str, save_path: str) -> dict:
|
||||
"""
|
||||
下载图片到本地
|
||||
|
||||
Args:
|
||||
url: 图片 URL
|
||||
save_path: 保存路径
|
||||
|
||||
Returns:
|
||||
dict: {"success": bool, "file_path": str, "error": str}
|
||||
"""
|
||||
try:
|
||||
async with aiohttp.ClientSession() as session:
|
||||
async with session.get(url, timeout=aiohttp.ClientTimeout(total=30)) as response:
|
||||
if response.status != 200:
|
||||
return {
|
||||
"success": False,
|
||||
"error": f"Failed to download image: HTTP {response.status}"
|
||||
}
|
||||
|
||||
content = await response.read()
|
||||
|
||||
# 保存文件
|
||||
with open(save_path, 'wb') as f:
|
||||
f.write(content)
|
||||
|
||||
# 验证文件
|
||||
if not os.path.exists(save_path) or os.path.getsize(save_path) == 0:
|
||||
return {
|
||||
"success": False,
|
||||
"error": "Downloaded file is empty or not saved"
|
||||
}
|
||||
|
||||
return {
|
||||
"success": True,
|
||||
"file_path": save_path
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error downloading image: {e}")
|
||||
return {
|
||||
"success": False,
|
||||
"error": str(e)
|
||||
}
|
||||
|
||||
def validate_image(self, file_path: str) -> bool:
|
||||
"""
|
||||
验证图片文件是否有效
|
||||
|
||||
Args:
|
||||
file_path: 文件路径
|
||||
|
||||
Returns:
|
||||
bool: 是否有效
|
||||
"""
|
||||
try:
|
||||
from PIL import Image
|
||||
with Image.open(file_path) as img:
|
||||
img.verify()
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.error(f"Image validation failed: {e}")
|
||||
return False
|
||||
@@ -0,0 +1,327 @@
|
||||
"""
|
||||
视频处理工具
|
||||
使用 FFmpeg 将图片拼接成视频
|
||||
"""
|
||||
import os
|
||||
import subprocess
|
||||
import logging
|
||||
from pathlib import Path
|
||||
from typing import List, Optional
|
||||
from datetime import datetime
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class VideoProcessor:
|
||||
"""视频处理器"""
|
||||
|
||||
def __init__(self):
|
||||
"""初始化视频处理器"""
|
||||
# 检查 FFmpeg 是否可用
|
||||
try:
|
||||
result = subprocess.run(['ffmpeg', '-version'],
|
||||
capture_output=True, text=True, timeout=5)
|
||||
if result.returncode == 0:
|
||||
logger.info("FFmpeg is available")
|
||||
else:
|
||||
logger.warning("FFmpeg check returned non-zero exit code")
|
||||
except Exception as e:
|
||||
logger.error(f"FFmpeg not available: {e}")
|
||||
|
||||
async def create_video_from_images(
|
||||
self,
|
||||
image_paths: List[str],
|
||||
output_dir: str = "/app/outputs/videos",
|
||||
duration_per_image: int = 3,
|
||||
fps: int = 30,
|
||||
transition: Optional[str] = None
|
||||
) -> dict:
|
||||
"""
|
||||
从图片创建视频
|
||||
|
||||
Args:
|
||||
image_paths: 图片路径列表
|
||||
output_dir: 输出目录
|
||||
duration_per_image: 每张图片显示秒数
|
||||
fps: 帧率
|
||||
transition: 转场效果 (fade, wipeleft, wiperight, slideup, slidedown)
|
||||
|
||||
Returns:
|
||||
dict: {
|
||||
"success": bool,
|
||||
"file_path": str,
|
||||
"url": str,
|
||||
"duration": float,
|
||||
"image_count": int
|
||||
}
|
||||
"""
|
||||
try:
|
||||
# 验证输入
|
||||
if not image_paths or len(image_paths) == 0:
|
||||
return {
|
||||
"success": False,
|
||||
"error": "No images provided"
|
||||
}
|
||||
|
||||
# 验证所有图片文件存在
|
||||
for img_path in image_paths:
|
||||
if not os.path.exists(img_path):
|
||||
return {
|
||||
"success": False,
|
||||
"error": f"Image not found: {img_path}"
|
||||
}
|
||||
|
||||
# 创建输出目录
|
||||
Path(output_dir).mkdir(parents=True, exist_ok=True)
|
||||
|
||||
# 生成输出文件名
|
||||
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
|
||||
output_filename = f"video_{timestamp}.mp4"
|
||||
output_path = os.path.join(output_dir, output_filename)
|
||||
|
||||
logger.info(f"Creating video from {len(image_paths)} images...")
|
||||
|
||||
# 根据是否需要转场效果选择不同的方法
|
||||
if transition and len(image_paths) > 1:
|
||||
result = await self._create_video_with_transitions(
|
||||
image_paths, output_path, duration_per_image, fps, transition
|
||||
)
|
||||
else:
|
||||
result = await self._create_simple_video(
|
||||
image_paths, output_path, duration_per_image, fps
|
||||
)
|
||||
|
||||
if not result["success"]:
|
||||
return result
|
||||
|
||||
# 获取视频信息
|
||||
video_info = self.get_video_info(output_path)
|
||||
|
||||
return {
|
||||
"success": True,
|
||||
"file_path": output_path,
|
||||
"filename": output_filename,
|
||||
"url": f"/api/v1/files/{output_filename}",
|
||||
"duration": video_info.get("duration", 0),
|
||||
"image_count": len(image_paths)
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error creating video: {e}", exc_info=True)
|
||||
return {
|
||||
"success": False,
|
||||
"error": str(e)
|
||||
}
|
||||
|
||||
async def _create_simple_video(
|
||||
self,
|
||||
image_paths: List[str],
|
||||
output_path: str,
|
||||
duration_per_image: int,
|
||||
fps: int
|
||||
) -> dict:
|
||||
"""
|
||||
创建简单视频(无转场效果)
|
||||
|
||||
使用 concat demuxer 方法
|
||||
"""
|
||||
try:
|
||||
# 创建临时文件列表
|
||||
temp_dir = os.path.dirname(output_path)
|
||||
concat_file = os.path.join(temp_dir, f"concat_{datetime.now().strftime('%Y%m%d_%H%M%S')}.txt")
|
||||
|
||||
# 写入文件列表
|
||||
with open(concat_file, 'w') as f:
|
||||
for img_path in image_paths:
|
||||
# FFmpeg concat 格式
|
||||
f.write(f"file '{img_path}'\n")
|
||||
f.write(f"duration {duration_per_image}\n")
|
||||
# 最后一张图片需要再写一次(FFmpeg 要求)
|
||||
f.write(f"file '{image_paths[-1]}'\n")
|
||||
|
||||
# 构建 FFmpeg 命令
|
||||
cmd = [
|
||||
'ffmpeg',
|
||||
'-f', 'concat',
|
||||
'-safe', '0',
|
||||
'-i', concat_file,
|
||||
'-vsync', 'vfr',
|
||||
'-pix_fmt', 'yuv420p',
|
||||
'-c:v', 'libx264',
|
||||
'-r', str(fps),
|
||||
'-y', # 覆盖输出文件
|
||||
output_path
|
||||
]
|
||||
|
||||
logger.info(f"Running FFmpeg command: {' '.join(cmd)}")
|
||||
|
||||
# 执行命令
|
||||
result = subprocess.run(
|
||||
cmd,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=300 # 5分钟超时
|
||||
)
|
||||
|
||||
# 清理临时文件
|
||||
try:
|
||||
os.remove(concat_file)
|
||||
except:
|
||||
pass
|
||||
|
||||
if result.returncode != 0:
|
||||
logger.error(f"FFmpeg error: {result.stderr}")
|
||||
return {
|
||||
"success": False,
|
||||
"error": f"FFmpeg failed: {result.stderr[:500]}"
|
||||
}
|
||||
|
||||
# 验证输出文件
|
||||
if not os.path.exists(output_path) or os.path.getsize(output_path) == 0:
|
||||
return {
|
||||
"success": False,
|
||||
"error": "Output video file is empty or not created"
|
||||
}
|
||||
|
||||
logger.info(f"Video created successfully: {output_path}")
|
||||
return {"success": True}
|
||||
|
||||
except subprocess.TimeoutExpired:
|
||||
return {
|
||||
"success": False,
|
||||
"error": "FFmpeg command timed out"
|
||||
}
|
||||
except Exception as e:
|
||||
logger.error(f"Error in _create_simple_video: {e}")
|
||||
return {
|
||||
"success": False,
|
||||
"error": str(e)
|
||||
}
|
||||
|
||||
async def _create_video_with_transitions(
|
||||
self,
|
||||
image_paths: List[str],
|
||||
output_path: str,
|
||||
duration_per_image: int,
|
||||
fps: int,
|
||||
transition: str
|
||||
) -> dict:
|
||||
"""
|
||||
创建带转场效果的视频
|
||||
|
||||
使用 xfade filter
|
||||
"""
|
||||
try:
|
||||
# 转场持续时间(秒)
|
||||
transition_duration = 1
|
||||
|
||||
# 构建 filter_complex
|
||||
# 每张图片显示时间 = duration_per_image
|
||||
# 转场开始时间 = duration_per_image - transition_duration
|
||||
|
||||
inputs = []
|
||||
filter_parts = []
|
||||
|
||||
# 添加所有输入
|
||||
for i, img_path in enumerate(image_paths):
|
||||
inputs.extend(['-loop', '1', '-t', str(duration_per_image), '-i', img_path])
|
||||
|
||||
# 构建 xfade filter chain
|
||||
if len(image_paths) == 2:
|
||||
# 两张图片的简单情况
|
||||
offset = duration_per_image - transition_duration
|
||||
filter_complex = f"[0][1]xfade=transition={transition}:duration={transition_duration}:offset={offset}[v]"
|
||||
else:
|
||||
# 多张图片需要链式 xfade
|
||||
current_label = "0"
|
||||
for i in range(1, len(image_paths)):
|
||||
offset = i * duration_per_image - i * transition_duration
|
||||
next_label = f"v{i}" if i < len(image_paths) - 1 else "v"
|
||||
filter_parts.append(f"[{current_label}][{i}]xfade=transition={transition}:duration={transition_duration}:offset={offset}[{next_label}]")
|
||||
current_label = next_label
|
||||
filter_complex = ";".join(filter_parts)
|
||||
|
||||
# 构建完整命令
|
||||
cmd = ['ffmpeg'] + inputs + [
|
||||
'-filter_complex', filter_complex,
|
||||
'-map', '[v]',
|
||||
'-c:v', 'libx264',
|
||||
'-pix_fmt', 'yuv420p',
|
||||
'-r', str(fps),
|
||||
'-y',
|
||||
output_path
|
||||
]
|
||||
|
||||
logger.info(f"Running FFmpeg with transitions: {' '.join(cmd[:20])}...")
|
||||
|
||||
# 执行命令
|
||||
result = subprocess.run(
|
||||
cmd,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=300
|
||||
)
|
||||
|
||||
if result.returncode != 0:
|
||||
logger.error(f"FFmpeg error: {result.stderr}")
|
||||
# 如果转场失败,回退到简单模式
|
||||
logger.warning("Falling back to simple video creation")
|
||||
return await self._create_simple_video(image_paths, output_path, duration_per_image, fps)
|
||||
|
||||
# 验证输出文件
|
||||
if not os.path.exists(output_path) or os.path.getsize(output_path) == 0:
|
||||
return {
|
||||
"success": False,
|
||||
"error": "Output video file is empty or not created"
|
||||
}
|
||||
|
||||
logger.info(f"Video with transitions created successfully: {output_path}")
|
||||
return {"success": True}
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error in _create_video_with_transitions: {e}")
|
||||
# 回退到简单模式
|
||||
return await self._create_simple_video(image_paths, output_path, duration_per_image, fps)
|
||||
|
||||
def get_video_info(self, video_path: str) -> dict:
|
||||
"""
|
||||
获取视频信息
|
||||
|
||||
Args:
|
||||
video_path: 视频路径
|
||||
|
||||
Returns:
|
||||
dict: 视频信息
|
||||
"""
|
||||
try:
|
||||
cmd = [
|
||||
'ffprobe',
|
||||
'-v', 'quiet',
|
||||
'-print_format', 'json',
|
||||
'-show_format',
|
||||
'-show_streams',
|
||||
video_path
|
||||
]
|
||||
|
||||
result = subprocess.run(cmd, capture_output=True, text=True, timeout=10)
|
||||
|
||||
if result.returncode == 0:
|
||||
import json
|
||||
data = json.loads(result.stdout)
|
||||
|
||||
duration = float(data.get('format', {}).get('duration', 0))
|
||||
size = int(data.get('format', {}).get('size', 0))
|
||||
|
||||
return {
|
||||
"duration": duration,
|
||||
"size": size,
|
||||
"format": data.get('format', {}).get('format_name', 'unknown')
|
||||
}
|
||||
except Exception as e:
|
||||
logger.error(f"Error getting video info: {e}")
|
||||
|
||||
return {
|
||||
"duration": 0,
|
||||
"size": 0,
|
||||
"format": "unknown"
|
||||
}
|
||||
@@ -0,0 +1,134 @@
|
||||
#!/usr/bin/env python
|
||||
"""
|
||||
快速测试脚本 - 验证核心功能
|
||||
"""
|
||||
import asyncio
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
# 添加路径
|
||||
sys.path.insert(0, str(Path(__file__).parent))
|
||||
|
||||
# 设置环境变量
|
||||
os.environ['OPENAI_API_KEY'] = 'sk-i9AwAgXDqqxsA9Ym4AjSPg'
|
||||
os.environ['OPENAI_BASE_URL'] = 'https://litellm.graystone-fb459c5d.southeastasia.azurecontainerapps.io/v1'
|
||||
os.environ['OUTPUT_DIR'] = './outputs'
|
||||
|
||||
from src.utils.video_processor import VideoProcessor
|
||||
from src.utils.file_manager import FileManager
|
||||
|
||||
async def test_video_processor():
|
||||
"""测试视频处理器(使用模拟图片)"""
|
||||
print("\n" + "="*60)
|
||||
print("测试: 视频处理器(FFmpeg)")
|
||||
print("="*60)
|
||||
|
||||
# 创建测试图片
|
||||
from PIL import Image
|
||||
import os
|
||||
|
||||
os.makedirs("./outputs/images", exist_ok=True)
|
||||
|
||||
image_paths = []
|
||||
colors = [(255, 0, 0), (0, 255, 0), (0, 0, 255)] # 红、绿、蓝
|
||||
|
||||
for i, color in enumerate(colors):
|
||||
img = Image.new('RGB', (1024, 1024), color)
|
||||
path = f"./outputs/images/test_image_{i}.png"
|
||||
img.save(path)
|
||||
image_paths.append(path)
|
||||
print(f"✅ 创建测试图片: {path}")
|
||||
|
||||
# 测试视频拼接
|
||||
processor = VideoProcessor()
|
||||
result = await processor.create_video_from_images(
|
||||
image_paths=image_paths,
|
||||
output_dir="./outputs/videos",
|
||||
duration_per_image=2,
|
||||
fps=30,
|
||||
transition="fade"
|
||||
)
|
||||
|
||||
if result.get("success"):
|
||||
print(f"\n✅ 视频生成成功!")
|
||||
print(f" 文件: {result['file_path']}")
|
||||
print(f" 时长: {result['duration']} 秒")
|
||||
print(f" 图片数: {result['image_count']}")
|
||||
return True
|
||||
else:
|
||||
print(f"\n❌ 视频生成失败: {result.get('error')}")
|
||||
return False
|
||||
|
||||
|
||||
async def test_file_manager():
|
||||
"""测试文件管理器"""
|
||||
print("\n" + "="*60)
|
||||
print("测试: 文件管理器")
|
||||
print("="*60)
|
||||
|
||||
manager = FileManager(base_dir="./outputs")
|
||||
|
||||
# 列出文件
|
||||
files = manager.list_files("all")
|
||||
print(f"\n📁 文件列表 (共 {len(files)} 个):")
|
||||
for file in files:
|
||||
print(f" - {file['filename']} ({file['type']}, {file['size_mb']} MB)")
|
||||
|
||||
# 存储信息
|
||||
storage = manager.get_storage_info()
|
||||
print(f"\n💾 存储信息:")
|
||||
print(f" - 总大小: {storage['total_size_mb']} MB")
|
||||
print(f" - 图片数: {storage['image_count']}")
|
||||
print(f" - 视频数: {storage['video_count']}")
|
||||
|
||||
return True
|
||||
|
||||
|
||||
async def main():
|
||||
"""运行测试"""
|
||||
print("\n" + "="*60)
|
||||
print("🧪 Video Generator Agent - 核心功能测试")
|
||||
print("="*60)
|
||||
|
||||
results = []
|
||||
|
||||
# 测试视频处理
|
||||
results.append(("视频处理器", await test_video_processor()))
|
||||
|
||||
# 测试文件管理
|
||||
results.append(("文件管理器", await test_file_manager()))
|
||||
|
||||
# 汇总
|
||||
print("\n" + "="*60)
|
||||
print("📊 测试结果")
|
||||
print("="*60)
|
||||
|
||||
passed = sum(1 for _, result in results if result)
|
||||
total = len(results)
|
||||
|
||||
for name, result in results:
|
||||
status = "✅ 通过" if result else "❌ 失败"
|
||||
print(f"{status} - {name}")
|
||||
|
||||
print(f"\n总计: {passed}/{total} 通过")
|
||||
|
||||
if passed == total:
|
||||
print("\n🎉 核心功能测试通过!")
|
||||
print("\n📝 注意: 图片生成功能需要正确的模型配置")
|
||||
print(" 当前模型 'taiji/gemini-3-pro-image-preview' 可能不支持")
|
||||
print(" 建议使用: dall-e-3, dall-e-2 或其他支持的模型")
|
||||
return 0
|
||||
else:
|
||||
return 1
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
try:
|
||||
exit_code = asyncio.run(main())
|
||||
sys.exit(exit_code)
|
||||
except Exception as e:
|
||||
print(f"\n❌ 测试失败: {e}")
|
||||
import traceback
|
||||
traceback.print_exc()
|
||||
sys.exit(1)
|
||||
@@ -0,0 +1,252 @@
|
||||
#!/usr/bin/env python
|
||||
"""
|
||||
Video Generator Agent 测试脚本
|
||||
测试图片生成和视频拼接功能
|
||||
"""
|
||||
import asyncio
|
||||
import aiohttp
|
||||
import json
|
||||
import sys
|
||||
import os
|
||||
|
||||
# 配置
|
||||
API_BASE_URL = os.getenv("API_BASE_URL", "http://localhost:8000")
|
||||
API_KEY = os.getenv("API_KEY", "sk-i9AwAgXDqqxsA9Ym4AjSPg")
|
||||
|
||||
|
||||
async def test_health_check():
|
||||
"""测试健康检查"""
|
||||
print("\n" + "="*60)
|
||||
print("测试 1: 健康检查")
|
||||
print("="*60)
|
||||
|
||||
try:
|
||||
async with aiohttp.ClientSession() as session:
|
||||
async with session.get(f"{API_BASE_URL}/health") as response:
|
||||
data = await response.json()
|
||||
print(f"✅ 状态: {response.status}")
|
||||
print(f"📊 响应: {json.dumps(data, indent=2, ensure_ascii=False)}")
|
||||
return response.status == 200
|
||||
except Exception as e:
|
||||
print(f"❌ 错误: {e}")
|
||||
return False
|
||||
|
||||
|
||||
async def test_generate_single_image():
|
||||
"""测试生成单张图片"""
|
||||
print("\n" + "="*60)
|
||||
print("测试 2: 生成单张图片")
|
||||
print("="*60)
|
||||
|
||||
try:
|
||||
payload = {
|
||||
"description": "a beautiful sunset over mountains with orange and purple sky",
|
||||
"size": "1024x1024",
|
||||
"quality": "standard"
|
||||
}
|
||||
|
||||
headers = {
|
||||
"api-key": API_KEY,
|
||||
"Content-Type": "application/json"
|
||||
}
|
||||
|
||||
print(f"📝 请求: {json.dumps(payload, indent=2, ensure_ascii=False)}")
|
||||
|
||||
async with aiohttp.ClientSession() as session:
|
||||
async with session.post(
|
||||
f"{API_BASE_URL}/api/v1/generate-image",
|
||||
json=payload,
|
||||
headers=headers,
|
||||
timeout=aiohttp.ClientTimeout(total=120)
|
||||
) as response:
|
||||
data = await response.json()
|
||||
print(f"✅ 状态: {response.status}")
|
||||
print(f"📊 响应: {json.dumps(data, indent=2, ensure_ascii=False)}")
|
||||
|
||||
if data.get("success"):
|
||||
print(f"🖼️ 图片 URL: {API_BASE_URL}{data.get('url')}")
|
||||
return True
|
||||
else:
|
||||
print(f"❌ 失败: {data.get('error')}")
|
||||
return False
|
||||
except Exception as e:
|
||||
print(f"❌ 错误: {e}")
|
||||
return False
|
||||
|
||||
|
||||
async def test_generate_video():
|
||||
"""测试生成视频"""
|
||||
print("\n" + "="*60)
|
||||
print("测试 3: 生成视频(3个场景)")
|
||||
print("="*60)
|
||||
|
||||
try:
|
||||
payload = {
|
||||
"descriptions": [
|
||||
"a peaceful morning sunrise over a calm lake with mist",
|
||||
"a busy city street at noon with people and cars",
|
||||
"a starry night sky with the milky way visible"
|
||||
],
|
||||
"duration_per_image": 3,
|
||||
"fps": 30,
|
||||
"transition": "fade"
|
||||
}
|
||||
|
||||
headers = {
|
||||
"api-key": API_KEY,
|
||||
"Content-Type": "application/json"
|
||||
}
|
||||
|
||||
print(f"📝 请求: {json.dumps(payload, indent=2, ensure_ascii=False)}")
|
||||
print("⏳ 生成视频中,这可能需要 1-3 分钟...")
|
||||
|
||||
async with aiohttp.ClientSession() as session:
|
||||
async with session.post(
|
||||
f"{API_BASE_URL}/api/v1/generate-video",
|
||||
json=payload,
|
||||
headers=headers,
|
||||
timeout=aiohttp.ClientTimeout(total=300)
|
||||
) as response:
|
||||
data = await response.json()
|
||||
print(f"✅ 状态: {response.status}")
|
||||
print(f"📊 响应: {json.dumps(data, indent=2, ensure_ascii=False)}")
|
||||
|
||||
if data.get("success"):
|
||||
video_info = data.get("video", {})
|
||||
print(f"\n🎬 视频信息:")
|
||||
print(f" - URL: {API_BASE_URL}{video_info.get('url')}")
|
||||
print(f" - 时长: {video_info.get('duration')} 秒")
|
||||
print(f" - 图片数: {video_info.get('image_count')}")
|
||||
return True
|
||||
else:
|
||||
print(f"❌ 失败: {data.get('error')}")
|
||||
return False
|
||||
except Exception as e:
|
||||
print(f"❌ 错误: {e}")
|
||||
return False
|
||||
|
||||
|
||||
async def test_list_files():
|
||||
"""测试列出文件"""
|
||||
print("\n" + "="*60)
|
||||
print("测试 4: 列出生成的文件")
|
||||
print("="*60)
|
||||
|
||||
try:
|
||||
async with aiohttp.ClientSession() as session:
|
||||
async with session.get(f"{API_BASE_URL}/api/v1/list-files?file_type=all") as response:
|
||||
data = await response.json()
|
||||
print(f"✅ 状态: {response.status}")
|
||||
|
||||
if data.get("success"):
|
||||
files = data.get("files", [])
|
||||
storage = data.get("storage", {})
|
||||
|
||||
print(f"\n📁 存储信息:")
|
||||
print(f" - 总大小: {storage.get('total_size_mb')} MB")
|
||||
print(f" - 图片数: {storage.get('image_count')}")
|
||||
print(f" - 视频数: {storage.get('video_count')}")
|
||||
|
||||
print(f"\n📄 文件列表 (共 {len(files)} 个):")
|
||||
for i, file in enumerate(files[:5], 1): # 只显示前5个
|
||||
print(f" {i}. {file['filename']} ({file['type']}, {file['size_mb']} MB)")
|
||||
|
||||
if len(files) > 5:
|
||||
print(f" ... 还有 {len(files) - 5} 个文件")
|
||||
|
||||
return True
|
||||
else:
|
||||
print(f"❌ 失败: {data.get('error')}")
|
||||
return False
|
||||
except Exception as e:
|
||||
print(f"❌ 错误: {e}")
|
||||
return False
|
||||
|
||||
|
||||
async def test_mcp_tools():
|
||||
"""测试 MCP 工具调用"""
|
||||
print("\n" + "="*60)
|
||||
print("测试 5: MCP 工具调用")
|
||||
print("="*60)
|
||||
|
||||
try:
|
||||
# 测试 tools/list
|
||||
print("\n📋 测试 tools/list:")
|
||||
payload = {
|
||||
"jsonrpc": "2.0",
|
||||
"method": "tools/list",
|
||||
"id": 1
|
||||
}
|
||||
|
||||
async with aiohttp.ClientSession() as session:
|
||||
async with session.post(
|
||||
f"{API_BASE_URL}/mcp",
|
||||
json=payload,
|
||||
timeout=aiohttp.ClientTimeout(total=30)
|
||||
) as response:
|
||||
data = await response.json()
|
||||
print(f"✅ 状态: {response.status}")
|
||||
|
||||
if "result" in data:
|
||||
tools = data["result"].get("tools", [])
|
||||
print(f"🔧 可用工具 ({len(tools)} 个):")
|
||||
for tool in tools:
|
||||
print(f" - {tool['name']}: {tool['description']}")
|
||||
return True
|
||||
else:
|
||||
print(f"❌ 失败: {data.get('error')}")
|
||||
return False
|
||||
except Exception as e:
|
||||
print(f"❌ 错误: {e}")
|
||||
return False
|
||||
|
||||
|
||||
async def main():
|
||||
"""运行所有测试"""
|
||||
print("\n" + "="*60)
|
||||
print("🧪 Video Generator Agent 测试套件")
|
||||
print("="*60)
|
||||
print(f"🌐 API URL: {API_BASE_URL}")
|
||||
print(f"🔑 API Key: {API_KEY[:20]}...")
|
||||
|
||||
results = []
|
||||
|
||||
# 运行测试
|
||||
results.append(("健康检查", await test_health_check()))
|
||||
results.append(("生成单张图片", await test_generate_single_image()))
|
||||
results.append(("生成视频", await test_generate_video()))
|
||||
results.append(("列出文件", await test_list_files()))
|
||||
results.append(("MCP 工具", await test_mcp_tools()))
|
||||
|
||||
# 汇总结果
|
||||
print("\n" + "="*60)
|
||||
print("📊 测试结果汇总")
|
||||
print("="*60)
|
||||
|
||||
passed = sum(1 for _, result in results if result)
|
||||
total = len(results)
|
||||
|
||||
for name, result in results:
|
||||
status = "✅ 通过" if result else "❌ 失败"
|
||||
print(f"{status} - {name}")
|
||||
|
||||
print(f"\n总计: {passed}/{total} 通过")
|
||||
|
||||
if passed == total:
|
||||
print("\n🎉 所有测试通过!")
|
||||
return 0
|
||||
else:
|
||||
print(f"\n⚠️ {total - passed} 个测试失败")
|
||||
return 1
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
try:
|
||||
exit_code = asyncio.run(main())
|
||||
sys.exit(exit_code)
|
||||
except KeyboardInterrupt:
|
||||
print("\n\n⚠️ 测试被中断")
|
||||
sys.exit(1)
|
||||
except Exception as e:
|
||||
print(f"\n\n❌ 测试失败: {e}")
|
||||
sys.exit(1)
|
||||
@@ -9,7 +9,7 @@ from sqlalchemy.orm import Session
|
||||
from datetime import datetime
|
||||
import logging
|
||||
|
||||
from k8s_manager import K8sManager
|
||||
from k8s_manager import K8sManager, sanitize_k8s_name
|
||||
from database import (
|
||||
get_db, Template, Agent, Quota, AgentMetric,
|
||||
AgentType, AgentStatus, parse_resource_string
|
||||
@@ -328,6 +328,12 @@ async def create_agent(request: CreateAgentRequest, db: Session = Depends(get_db
|
||||
detail=f"无效的框架类型。支持的框架: {', '.join(valid_frameworks)}"
|
||||
)
|
||||
|
||||
# DNS-1035 名称合规化:确保名称可以用作 K8s 资源名称
|
||||
original_name = request.name
|
||||
request.name = sanitize_k8s_name(request.name)
|
||||
if original_name != request.name:
|
||||
logger.info(f"🔄 Agent 名称已合规化: '{original_name}' -> '{request.name}' (DNS-1035)")
|
||||
|
||||
# 合并环境变量到config
|
||||
config_data = request.config.copy()
|
||||
config_data["agent_framework"] = framework # 添加框架类型到配置
|
||||
@@ -389,8 +395,8 @@ async def create_agent(request: CreateAgentRequest, db: Session = Depends(get_db
|
||||
config_data=config_data
|
||||
)
|
||||
|
||||
# 步骤3: 获取服务端口
|
||||
service_port = k8s_manager.TEMPLATE_PORTS.get(request.template)
|
||||
# 步骤3: 获取服务端口(从数据库动态获取)
|
||||
service_port = template_manager.get_port(request.template)
|
||||
|
||||
# 步骤4: 创建 LoadBalancer Service(AKS 会自动分配外网 IP)
|
||||
service_info = None
|
||||
@@ -587,8 +593,11 @@ async def delete_agent(agent_name: str, db: Session = Depends(get_db)):
|
||||
try:
|
||||
logger.info(f"收到删除Agent请求: {agent_name}")
|
||||
|
||||
# DNS-1035 名称合规化
|
||||
agent_name = sanitize_k8s_name(agent_name)
|
||||
|
||||
# 保护机制:防止删除 agent-manager 命名空间
|
||||
computed_namespace = f"agent-{agent_name}".lower().strip('-')[:63]
|
||||
computed_namespace = f"agent-{agent_name}"[:63].rstrip('-')
|
||||
if computed_namespace == "agent-manager":
|
||||
logger.error(f"❌ 禁止删除 agent-manager 命名空间!agent_name={agent_name}, computed_namespace={computed_namespace}")
|
||||
raise HTTPException(
|
||||
|
||||
@@ -0,0 +1,56 @@
|
||||
#!/bin/bash
|
||||
# 构建三个美股 Agent 镜像并推送到 ACR
|
||||
# 使用方法: ./build_stock_agents.sh
|
||||
|
||||
set -e
|
||||
|
||||
echo "=========================================="
|
||||
echo "构建美股 Agent 镜像并推送到 ACR"
|
||||
echo "=========================================="
|
||||
|
||||
# 登录 ACR
|
||||
echo ""
|
||||
echo "=== 步骤 1: 登录 ACR ==="
|
||||
az acr login --name agnettaiji
|
||||
|
||||
cd /home/taiji/tools/agent-manager
|
||||
|
||||
# 构建 stock-quote-agent
|
||||
echo ""
|
||||
echo "=== 步骤 2: 构建 stock-quote-agent (ARM64) ==="
|
||||
docker buildx build --platform linux/arm64 \
|
||||
-t agnettaiji.azurecr.io/ai-agents/stock-quote-agent:latest \
|
||||
-f agent_templates/agents/stock_quote_agent/stock_quote_agent.Dockerfile \
|
||||
agent_templates/ --push
|
||||
echo "✅ stock-quote-agent 推送完成"
|
||||
|
||||
# 构建 stock-news-agent
|
||||
echo ""
|
||||
echo "=== 步骤 3: 构建 stock-news-agent (ARM64) ==="
|
||||
docker buildx build --platform linux/arm64 \
|
||||
-t agnettaiji.azurecr.io/ai-agents/stock-news-agent:latest \
|
||||
-f agent_templates/agents/stock_news_agent/stock_news_agent.Dockerfile \
|
||||
agent_templates/ --push
|
||||
echo "✅ stock-news-agent 推送完成"
|
||||
|
||||
# 构建 stock-analysis-agent
|
||||
echo ""
|
||||
echo "=== 步骤 4: 构建 stock-analysis-agent (ARM64) ==="
|
||||
docker buildx build --platform linux/arm64 \
|
||||
-t agnettaiji.azurecr.io/ai-agents/stock-analysis-agent:latest \
|
||||
-f agent_templates/agents/stock_analysis_agent/stock_analysis_agent.Dockerfile \
|
||||
agent_templates/ --push
|
||||
echo "✅ stock-analysis-agent 推送完成"
|
||||
|
||||
echo ""
|
||||
echo "=========================================="
|
||||
echo "✅ 所有镜像构建并推送完成!"
|
||||
echo "=========================================="
|
||||
echo ""
|
||||
echo "镜像列表:"
|
||||
echo " - agnettaiji.azurecr.io/ai-agents/stock-quote-agent:latest"
|
||||
echo " - agnettaiji.azurecr.io/ai-agents/stock-news-agent:latest"
|
||||
echo " - agnettaiji.azurecr.io/ai-agents/stock-analysis-agent:latest"
|
||||
echo ""
|
||||
echo "下一步: 重启 agent-manager 以加载新模板"
|
||||
echo " kubectl rollout restart deployment/agent-manager -n agent-manager"
|
||||
Executable
+148
@@ -0,0 +1,148 @@
|
||||
#!/bin/bash
|
||||
# 部署三个美股 Agent:构建镜像、添加模板、测试创建
|
||||
# 使用方法: ./deploy_stock_agents.sh
|
||||
|
||||
set -e
|
||||
|
||||
AGENT_MANAGER_URL="http://20.212.121.126"
|
||||
|
||||
echo "=========================================="
|
||||
echo "部署美股 Agent 到 Agent Manager"
|
||||
echo "=========================================="
|
||||
|
||||
# ==================== 步骤 1: 构建并推送镜像 ====================
|
||||
echo ""
|
||||
echo "=== 步骤 1: 登录 ACR ==="
|
||||
az acr login --name agnettaiji
|
||||
|
||||
cd /home/taiji/tools/agent-manager
|
||||
|
||||
echo ""
|
||||
echo "=== 步骤 2: 构建 stock-quote-agent (ARM64) ==="
|
||||
docker buildx build --platform linux/arm64 \
|
||||
-t agnettaiji.azurecr.io/ai-agents/stock-quote-agent:latest \
|
||||
-f agent_templates/agents/stock_quote_agent/stock_quote_agent.Dockerfile \
|
||||
agent_templates/ --push
|
||||
echo "✅ stock-quote-agent 推送完成"
|
||||
|
||||
echo ""
|
||||
echo "=== 步骤 3: 构建 stock-news-agent (ARM64) ==="
|
||||
docker buildx build --platform linux/arm64 \
|
||||
-t agnettaiji.azurecr.io/ai-agents/stock-news-agent:latest \
|
||||
-f agent_templates/agents/stock_news_agent/stock_news_agent.Dockerfile \
|
||||
agent_templates/ --push
|
||||
echo "✅ stock-news-agent 推送完成"
|
||||
|
||||
echo ""
|
||||
echo "=== 步骤 4: 构建 stock-analysis-agent (ARM64) ==="
|
||||
docker buildx build --platform linux/arm64 \
|
||||
-t agnettaiji.azurecr.io/ai-agents/stock-analysis-agent:latest \
|
||||
-f agent_templates/agents/stock_analysis_agent/stock_analysis_agent.Dockerfile \
|
||||
agent_templates/ --push
|
||||
echo "✅ stock-analysis-agent 推送完成"
|
||||
|
||||
# ==================== 步骤 2: 通过 API 添加模板 ====================
|
||||
echo ""
|
||||
echo "=== 步骤 5: 通过 API 添加模板 ==="
|
||||
|
||||
# 添加 stock_quote_agent 模板
|
||||
echo "添加 stock_quote_agent 模板..."
|
||||
curl -s -X POST "${AGENT_MANAGER_URL}/templates/create" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"name": "stock_quote_agent",
|
||||
"display_name": "Stock Quote Agent",
|
||||
"description": "美股实时行情查询 Agent - 获取股票价格、涨跌幅、成交量等数据",
|
||||
"image": "agnettaiji.azurecr.io/ai-agents/stock-quote-agent:latest",
|
||||
"port": 8080,
|
||||
"agent_framework": "api",
|
||||
"env_requirements": {}
|
||||
}' | jq .
|
||||
echo ""
|
||||
|
||||
# 添加 stock_news_agent 模板
|
||||
echo "添加 stock_news_agent 模板..."
|
||||
curl -s -X POST "${AGENT_MANAGER_URL}/templates/create" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"name": "stock_news_agent",
|
||||
"display_name": "Stock News Agent",
|
||||
"description": "美股新闻资讯 Agent - 获取股票相关新闻和市场情绪分析",
|
||||
"image": "agnettaiji.azurecr.io/ai-agents/stock-news-agent:latest",
|
||||
"port": 8080,
|
||||
"agent_framework": "api",
|
||||
"env_requirements": {}
|
||||
}' | jq .
|
||||
echo ""
|
||||
|
||||
# 添加 stock_analysis_agent 模板
|
||||
echo "添加 stock_analysis_agent 模板..."
|
||||
curl -s -X POST "${AGENT_MANAGER_URL}/templates/create" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"name": "stock_analysis_agent",
|
||||
"display_name": "Stock Analysis Agent",
|
||||
"description": "美股技术分析 Agent - 提供技术指标、趋势分析和投资建议",
|
||||
"image": "agnettaiji.azurecr.io/ai-agents/stock-analysis-agent:latest",
|
||||
"port": 8080,
|
||||
"agent_framework": "api",
|
||||
"env_requirements": {}
|
||||
}' | jq .
|
||||
|
||||
echo ""
|
||||
echo "=== 步骤 6: 验证模板添加成功 ==="
|
||||
echo "查询所有模板..."
|
||||
curl -s "${AGENT_MANAGER_URL}/templates" | jq '.[] | select(.name | startswith("stock_"))'
|
||||
|
||||
# ==================== 步骤 3: 测试创建 Agent ====================
|
||||
echo ""
|
||||
echo "=== 步骤 7: 测试创建 stock_quote_agent ==="
|
||||
RESULT=$(curl -s -X POST "${AGENT_MANAGER_URL}/agents" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"name": "test-stock-quote",
|
||||
"template": "stock_quote_agent",
|
||||
"config": {
|
||||
"user_id": "test-user",
|
||||
"cpu_request": "100m",
|
||||
"cpu_limit": "500m",
|
||||
"memory_request": "128Mi",
|
||||
"memory_limit": "512Mi",
|
||||
"replicas": 1
|
||||
}
|
||||
}')
|
||||
|
||||
echo "$RESULT" | jq .
|
||||
|
||||
# 提取域名
|
||||
DOMAIN=$(echo "$RESULT" | jq -r '.access_info.domain // empty')
|
||||
if [ -n "$DOMAIN" ]; then
|
||||
echo ""
|
||||
echo "✅ Agent 创建成功!"
|
||||
echo "域名: $DOMAIN"
|
||||
echo ""
|
||||
echo "等待 30 秒后测试 API..."
|
||||
sleep 30
|
||||
|
||||
echo ""
|
||||
echo "=== 步骤 8: 测试 Agent API ==="
|
||||
echo "测试获取 AAPL 行情..."
|
||||
curl -s "http://${DOMAIN}/quote?symbol=AAPL" | jq .
|
||||
else
|
||||
echo ""
|
||||
echo "⚠️ Agent 创建中,请稍后检查状态"
|
||||
fi
|
||||
|
||||
echo ""
|
||||
echo "=========================================="
|
||||
echo "✅ 部署完成!"
|
||||
echo "=========================================="
|
||||
echo ""
|
||||
echo "新增模板:"
|
||||
echo " - stock_quote_agent: 美股实时行情查询"
|
||||
echo " - stock_news_agent: 美股新闻资讯"
|
||||
echo " - stock_analysis_agent: 美股技术分析"
|
||||
echo ""
|
||||
echo "测试命令:"
|
||||
echo " curl \"http://\${DOMAIN}/quote?symbol=AAPL\""
|
||||
echo " curl \"http://\${DOMAIN}/quote?symbol=TSLA\""
|
||||
@@ -0,0 +1,646 @@
|
||||
# 链上数据分析 AI Agent 文档
|
||||
|
||||
---
|
||||
|
||||
本文档详细介绍了两个链上数据分析 Agent 的功能、API 接口和使用方法。
|
||||
|
||||
## 概述
|
||||
|
||||
| Agent | 功能 | 端口 |
|
||||
|-------|------|------|
|
||||
| Chain Explorer Agent | 链上数据查询 - 余额、交易、代币 | 8000 |
|
||||
| Chain Analysis Agent | 链上数据分析 - 活动分析、交易模式、资金流向 | 8000 |
|
||||
|
||||
## 支持的区块链
|
||||
|
||||
| 网络 | Chain ID | 符号 | 说明 |
|
||||
|------|----------|------|------|
|
||||
| Ethereum | ethereum | ETH | 以太坊主网 |
|
||||
| BSC | bsc | BNB | 币安智能链 |
|
||||
| Polygon | polygon | POL | Polygon 网络 |
|
||||
| Arbitrum | arbitrum | ETH | Arbitrum L2 |
|
||||
| Optimism | optimism | ETH | Optimism L2 |
|
||||
| Base | base | ETH | Coinbase L2 |
|
||||
|
||||
---
|
||||
|
||||
## 认证方式
|
||||
|
||||
所有 API 调用都需要通过请求头传递 API Key:
|
||||
|
||||
| Header | 说明 | 必需 |
|
||||
|--------|------|------|
|
||||
| `etherscan-key` | Etherscan API Key(区块链浏览器) | ✅ |
|
||||
| `api-key` | 备选的区块链浏览器 API Key | ⭕ |
|
||||
| `llm-key` | LLM API Key(用于 Chat 功能) | Chat 时必需 |
|
||||
| `Authorization` | Bearer Token(LLM API Key) | Chat 时备选 |
|
||||
|
||||
### 示例
|
||||
|
||||
```bash
|
||||
curl -X POST "http://agent-url/balance" \
|
||||
-H "Content-Type: application/json" \
|
||||
-H "etherscan-key: YOUR_ETHERSCAN_API_KEY" \
|
||||
-d '{"address": "0x...", "chain": "ethereum"}'
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
# 1. Chain Explorer Agent - 链上数据查询
|
||||
|
||||
## 功能概览
|
||||
|
||||
| 端点 | 方法 | 功能 |
|
||||
|------|------|------|
|
||||
| `/` | GET | 服务状态 |
|
||||
| `/health` | GET | 健康检查 |
|
||||
| `/chains` | GET | 支持的区块链列表 |
|
||||
| `/balance` | POST | 查询地址余额 |
|
||||
| `/transactions` | POST | 查询交易记录 |
|
||||
| `/tokens` | POST | 查询代币信息 |
|
||||
| `/chat` | POST | 智能对话 |
|
||||
|
||||
---
|
||||
|
||||
## 1.1 查询地址余额
|
||||
|
||||
### 请求
|
||||
|
||||
```bash
|
||||
POST /balance
|
||||
Content-Type: application/json
|
||||
etherscan-key: YOUR_API_KEY
|
||||
```
|
||||
|
||||
### 参数
|
||||
|
||||
| 参数 | 类型 | 必需 | 默认值 | 说明 |
|
||||
|------|------|------|--------|------|
|
||||
| address | string | ✅ | - | 钱包地址 (0x开头) |
|
||||
| chain | string | ❌ | ethereum | 区块链网络 |
|
||||
|
||||
### 示例
|
||||
|
||||
```bash
|
||||
curl -X POST "http://localhost:8000/balance" \
|
||||
-H "Content-Type: application/json" \
|
||||
-H "etherscan-key: F71AZ6XW2WN6GK3D63HJ7AVAEDC9M42EZY" \
|
||||
-d '{
|
||||
"address": "0xd8dA6BF26964aF9D7eEd9e03E53415D37aA96045",
|
||||
"chain": "ethereum"
|
||||
}'
|
||||
```
|
||||
|
||||
### 响应
|
||||
|
||||
```json
|
||||
{
|
||||
"success": true,
|
||||
"address": "0xd8dA6BF26964aF9D7eEd9e03E53415D37aA96045",
|
||||
"chain": "ethereum",
|
||||
"chain_name": "Ethereum",
|
||||
"balance_wei": "32116130289281011210",
|
||||
"balance": 32.11613029,
|
||||
"symbol": "ETH",
|
||||
"explorer_url": "https://etherscan.io/address/0xd8dA6BF26964aF9D7eEd9e03E53415D37aA96045",
|
||||
"timestamp": "2026-02-05T17:00:28.539769"
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 1.2 查询交易记录
|
||||
|
||||
### 请求
|
||||
|
||||
```bash
|
||||
POST /transactions
|
||||
Content-Type: application/json
|
||||
etherscan-key: YOUR_API_KEY
|
||||
```
|
||||
|
||||
### 参数
|
||||
|
||||
| 参数 | 类型 | 必需 | 默认值 | 说明 |
|
||||
|------|------|------|--------|------|
|
||||
| address | string | ✅ | - | 钱包地址 |
|
||||
| chain | string | ❌ | ethereum | 区块链网络 |
|
||||
| page | int | ❌ | 1 | 页码 |
|
||||
| limit | int | ❌ | 10 | 每页数量 (1-100) |
|
||||
|
||||
### 示例
|
||||
|
||||
```bash
|
||||
curl -X POST "http://localhost:8000/transactions" \
|
||||
-H "Content-Type: application/json" \
|
||||
-H "etherscan-key: F71AZ6XW2WN6GK3D63HJ7AVAEDC9M42EZY" \
|
||||
-d '{
|
||||
"address": "0xd8dA6BF26964aF9D7eEd9e03E53415D37aA96045",
|
||||
"chain": "ethereum",
|
||||
"limit": 5
|
||||
}'
|
||||
```
|
||||
|
||||
### 响应
|
||||
|
||||
```json
|
||||
{
|
||||
"success": true,
|
||||
"address": "0x...",
|
||||
"chain": "ethereum",
|
||||
"transactions": [
|
||||
{
|
||||
"hash": "0x5b0d81bab...",
|
||||
"block": "21780123",
|
||||
"timestamp": "2026-02-05T13:43:47",
|
||||
"from": "0x...",
|
||||
"to": "0x...",
|
||||
"value": 0.000505,
|
||||
"symbol": "ETH",
|
||||
"gas_used": "21000",
|
||||
"gas_price": "5000000000",
|
||||
"is_error": false,
|
||||
"tx_url": "https://etherscan.io/tx/0x..."
|
||||
}
|
||||
],
|
||||
"count": 5,
|
||||
"page": 1,
|
||||
"timestamp": "2026-02-05T17:00:30.123456"
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 1.3 查询代币信息
|
||||
|
||||
### 请求
|
||||
|
||||
```bash
|
||||
POST /tokens
|
||||
Content-Type: application/json
|
||||
etherscan-key: YOUR_API_KEY
|
||||
```
|
||||
|
||||
### 参数
|
||||
|
||||
| 参数 | 类型 | 必需 | 默认值 | 说明 |
|
||||
|------|------|------|--------|------|
|
||||
| address | string | ✅ | - | 钱包地址 |
|
||||
| chain | string | ❌ | ethereum | 区块链网络 |
|
||||
|
||||
### 示例
|
||||
|
||||
```bash
|
||||
curl -X POST "http://localhost:8000/tokens" \
|
||||
-H "Content-Type: application/json" \
|
||||
-H "etherscan-key: F71AZ6XW2WN6GK3D63HJ7AVAEDC9M42EZY" \
|
||||
-d '{
|
||||
"address": "0xd8dA6BF26964aF9D7eEd9e03E53415D37aA96045",
|
||||
"chain": "ethereum"
|
||||
}'
|
||||
```
|
||||
|
||||
### 响应
|
||||
|
||||
```json
|
||||
{
|
||||
"success": true,
|
||||
"address": "0x...",
|
||||
"chain": "ethereum",
|
||||
"tokens": [
|
||||
{
|
||||
"contract": "0x...",
|
||||
"name": "Dogelon",
|
||||
"symbol": "ELON",
|
||||
"decimals": 18,
|
||||
"tx_count": 5
|
||||
}
|
||||
],
|
||||
"token_count": 49,
|
||||
"timestamp": "2026-02-05T17:00:35.123456"
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 1.4 智能对话 (Chat)
|
||||
|
||||
### 请求
|
||||
|
||||
```bash
|
||||
POST /chat
|
||||
Content-Type: application/json
|
||||
etherscan-key: YOUR_ETHERSCAN_KEY
|
||||
llm-key: YOUR_LLM_API_KEY
|
||||
```
|
||||
|
||||
### 参数
|
||||
|
||||
| 参数 | 类型 | 必需 | 默认值 | 说明 |
|
||||
|------|------|------|--------|------|
|
||||
| message | string | ✅ | - | 用户消息(包含地址) |
|
||||
| chain | string | ❌ | ethereum | 默认区块链网络 |
|
||||
|
||||
### 示例
|
||||
|
||||
```bash
|
||||
curl -X POST "http://localhost:8000/chat" \
|
||||
-H "Content-Type: application/json" \
|
||||
-H "etherscan-key: F71AZ6XW2WN6GK3D63HJ7AVAEDC9M42EZY" \
|
||||
-H "llm-key: sk-xxx" \
|
||||
-d '{
|
||||
"message": "帮我查看 0xd8dA6BF26964aF9D7eEd9e03E53415D37aA96045 的余额和交易",
|
||||
"chain": "ethereum"
|
||||
}'
|
||||
```
|
||||
|
||||
### 响应
|
||||
|
||||
```json
|
||||
{
|
||||
"response": "地址 0xd8dA6BF26964aF9D7eEd9e03E53415D37aA96045 目前的余额为 32.12 ETH...",
|
||||
"data": {
|
||||
"balance": { ... },
|
||||
"recent_transactions": [ ... ],
|
||||
"tokens": [ ... ]
|
||||
},
|
||||
"timestamp": "2026-02-05T17:01:00.123456"
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
# 2. Chain Analysis Agent - 链上数据分析
|
||||
|
||||
## 功能概览
|
||||
|
||||
| 端点 | 方法 | 功能 |
|
||||
|------|------|------|
|
||||
| `/` | GET | 服务状态 |
|
||||
| `/health` | GET | 健康检查 |
|
||||
| `/chains` | GET | 支持的区块链列表 |
|
||||
| `/address-analysis` | POST | 地址活动分析 |
|
||||
| `/transaction-patterns` | POST | 交易模式分析 |
|
||||
| `/fund-flow` | POST | 资金流向分析 |
|
||||
| `/contract-interactions` | POST | 合约交互分析 |
|
||||
| `/chat` | POST | 智能分析对话 |
|
||||
|
||||
---
|
||||
|
||||
## 2.1 地址活动分析
|
||||
|
||||
分析指定时间段内的地址活动,包括收支统计、活跃度等。
|
||||
|
||||
### 请求
|
||||
|
||||
```bash
|
||||
POST /address-analysis
|
||||
Content-Type: application/json
|
||||
etherscan-key: YOUR_API_KEY
|
||||
```
|
||||
|
||||
### 参数
|
||||
|
||||
| 参数 | 类型 | 必需 | 默认值 | 说明 |
|
||||
|------|------|------|--------|------|
|
||||
| address | string | ✅ | - | 钱包地址 |
|
||||
| chain | string | ❌ | ethereum | 区块链网络 |
|
||||
| days | int | ❌ | 30 | 分析天数 (1-365) |
|
||||
|
||||
### 示例
|
||||
|
||||
```bash
|
||||
curl -X POST "http://localhost:8000/address-analysis" \
|
||||
-H "Content-Type: application/json" \
|
||||
-H "etherscan-key: F71AZ6XW2WN6GK3D63HJ7AVAEDC9M42EZY" \
|
||||
-d '{
|
||||
"address": "0xd8dA6BF26964aF9D7eEd9e03E53415D37aA96045",
|
||||
"chain": "ethereum",
|
||||
"days": 30
|
||||
}'
|
||||
```
|
||||
|
||||
### 响应
|
||||
|
||||
```json
|
||||
{
|
||||
"address": "0x...",
|
||||
"chain": "ethereum",
|
||||
"period_days": 30,
|
||||
"summary": {
|
||||
"total_sent": 1.0,
|
||||
"total_received": 0.004591,
|
||||
"net_flow": -0.995409,
|
||||
"tx_count_in": 52,
|
||||
"tx_count_out": 11,
|
||||
"total_tx": 63,
|
||||
"failed_tx": 2,
|
||||
"unique_addresses": 31,
|
||||
"active_days": 18
|
||||
},
|
||||
"symbol": "ETH",
|
||||
"current_balance": 32.11613029,
|
||||
"daily_activity": { ... },
|
||||
"timestamp": "2026-02-05T17:02:00.123456"
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 2.2 交易模式分析
|
||||
|
||||
分析地址的交易行为模式,包括时间分布、金额分布、高频交互对手等。
|
||||
|
||||
### 请求
|
||||
|
||||
```bash
|
||||
POST /transaction-patterns
|
||||
Content-Type: application/json
|
||||
etherscan-key: YOUR_API_KEY
|
||||
```
|
||||
|
||||
### 参数
|
||||
|
||||
| 参数 | 类型 | 必需 | 默认值 | 说明 |
|
||||
|------|------|------|--------|------|
|
||||
| address | string | ✅ | - | 钱包地址 |
|
||||
| chain | string | ❌ | ethereum | 区块链网络 |
|
||||
|
||||
### 示例
|
||||
|
||||
```bash
|
||||
curl -X POST "http://localhost:8000/transaction-patterns" \
|
||||
-H "Content-Type: application/json" \
|
||||
-H "etherscan-key: F71AZ6XW2WN6GK3D63HJ7AVAEDC9M42EZY" \
|
||||
-d '{
|
||||
"address": "0xd8dA6BF26964aF9D7eEd9e03E53415D37aA96045",
|
||||
"chain": "ethereum"
|
||||
}'
|
||||
```
|
||||
|
||||
### 响应
|
||||
|
||||
```json
|
||||
{
|
||||
"address": "0x...",
|
||||
"chain": "ethereum",
|
||||
"patterns": {
|
||||
"hourly_distribution": { "0": 5, "14": 20, ... },
|
||||
"daily_distribution": { "Monday": 10, "Tuesday": 15, ... },
|
||||
"value_distribution": {
|
||||
"micro": 198, // < 0.01 ETH
|
||||
"small": 1, // 0.01 - 0.1 ETH
|
||||
"medium": 0, // 0.1 - 1 ETH
|
||||
"large": 1, // 1 - 10 ETH
|
||||
"whale": 0 // > 10 ETH
|
||||
},
|
||||
"avg_interval_hours": 12.5,
|
||||
"top_counterparties": [
|
||||
{ "address": "0x...", "tx_count": 15 }
|
||||
]
|
||||
},
|
||||
"behavior_summary": "活跃高峰时段: 14:00 UTC; 以小额交易为主(可能是频繁交易者或机器人)",
|
||||
"timestamp": "2026-02-05T17:02:30.123456"
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 2.3 资金流向分析
|
||||
|
||||
分析资金来源和去向,识别主要入金/出金地址。
|
||||
|
||||
### 请求
|
||||
|
||||
```bash
|
||||
POST /fund-flow
|
||||
Content-Type: application/json
|
||||
etherscan-key: YOUR_API_KEY
|
||||
```
|
||||
|
||||
### 参数
|
||||
|
||||
| 参数 | 类型 | 必需 | 默认值 | 说明 |
|
||||
|------|------|------|--------|------|
|
||||
| address | string | ✅ | - | 钱包地址 |
|
||||
| chain | string | ❌ | ethereum | 区块链网络 |
|
||||
| limit | int | ❌ | 100 | 分析交易数量 (10-500) |
|
||||
|
||||
### 示例
|
||||
|
||||
```bash
|
||||
curl -X POST "http://localhost:8000/fund-flow" \
|
||||
-H "Content-Type: application/json" \
|
||||
-H "etherscan-key: F71AZ6XW2WN6GK3D63HJ7AVAEDC9M42EZY" \
|
||||
-d '{
|
||||
"address": "0xd8dA6BF26964aF9D7eEd9e03E53415D37aA96045",
|
||||
"chain": "ethereum",
|
||||
"limit": 100
|
||||
}'
|
||||
```
|
||||
|
||||
### 响应
|
||||
|
||||
```json
|
||||
{
|
||||
"address": "0x...",
|
||||
"chain": "ethereum",
|
||||
"fund_flow": {
|
||||
"total_inflow": 5.234,
|
||||
"total_outflow": 3.156,
|
||||
"net_flow": 2.078,
|
||||
"inflow_sources": 42,
|
||||
"outflow_destinations": 8,
|
||||
"top_inflow": [
|
||||
{ "address": "0x...", "amount": 2.5, "symbol": "ETH" }
|
||||
],
|
||||
"top_outflow": [
|
||||
{ "address": "0x...", "amount": 1.0, "symbol": "ETH" }
|
||||
]
|
||||
},
|
||||
"symbol": "ETH",
|
||||
"timestamp": "2026-02-05T17:03:00.123456"
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 2.4 合约交互分析
|
||||
|
||||
分析地址与智能合约的交互情况。
|
||||
|
||||
### 请求
|
||||
|
||||
```bash
|
||||
POST /contract-interactions
|
||||
Content-Type: application/json
|
||||
etherscan-key: YOUR_API_KEY
|
||||
```
|
||||
|
||||
### 参数
|
||||
|
||||
| 参数 | 类型 | 必需 | 默认值 | 说明 |
|
||||
|------|------|------|--------|------|
|
||||
| address | string | ✅ | - | 钱包地址 |
|
||||
| chain | string | ❌ | ethereum | 区块链网络 |
|
||||
|
||||
### 示例
|
||||
|
||||
```bash
|
||||
curl -X POST "http://localhost:8000/contract-interactions" \
|
||||
-H "Content-Type: application/json" \
|
||||
-H "etherscan-key: F71AZ6XW2WN6GK3D63HJ7AVAEDC9M42EZY" \
|
||||
-d '{
|
||||
"address": "0xd8dA6BF26964aF9D7eEd9e03E53415D37aA96045",
|
||||
"chain": "ethereum"
|
||||
}'
|
||||
```
|
||||
|
||||
### 响应
|
||||
|
||||
```json
|
||||
{
|
||||
"address": "0x...",
|
||||
"chain": "ethereum",
|
||||
"contract_interactions": {
|
||||
"total_contracts": 10,
|
||||
"top_contracts": [
|
||||
{
|
||||
"contract": "0x...",
|
||||
"interaction_count": 9,
|
||||
"unique_methods": 1,
|
||||
"total_value": 0.5,
|
||||
"symbol": "ETH",
|
||||
"explorer_url": "https://etherscan.io/address/0x..."
|
||||
}
|
||||
]
|
||||
},
|
||||
"timestamp": "2026-02-05T17:03:30.123456"
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 2.5 智能分析对话 (Chat)
|
||||
|
||||
### 请求
|
||||
|
||||
```bash
|
||||
POST /chat
|
||||
Content-Type: application/json
|
||||
etherscan-key: YOUR_ETHERSCAN_KEY
|
||||
llm-key: YOUR_LLM_API_KEY
|
||||
```
|
||||
|
||||
### 参数
|
||||
|
||||
| 参数 | 类型 | 必需 | 默认值 | 说明 |
|
||||
|------|------|------|--------|------|
|
||||
| message | string | ✅ | - | 分析请求(包含地址) |
|
||||
| chain | string | ❌ | ethereum | 默认区块链网络 |
|
||||
|
||||
### 示例
|
||||
|
||||
```bash
|
||||
curl -X POST "http://localhost:8000/chat" \
|
||||
-H "Content-Type: application/json" \
|
||||
-H "etherscan-key: F71AZ6XW2WN6GK3D63HJ7AVAEDC9M42EZY" \
|
||||
-H "llm-key: sk-xxx" \
|
||||
-d '{
|
||||
"message": "分析 0xd8dA6BF26964aF9D7eEd9e03E53415D37aA96045 是不是巨鲸或机器人",
|
||||
"chain": "ethereum"
|
||||
}'
|
||||
```
|
||||
|
||||
### 响应
|
||||
|
||||
```json
|
||||
{
|
||||
"response": "### 地址分析报告\n\n#### 一、基本信息\n- **当前余额**: 32.12 ETH\n...",
|
||||
"analysis": {
|
||||
"activity": { ... },
|
||||
"patterns": { ... },
|
||||
"fund_flow": { ... },
|
||||
"contracts": { ... },
|
||||
"balance": 32.11613029
|
||||
},
|
||||
"timestamp": "2026-02-05T17:04:00.123456"
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 统一错误格式
|
||||
|
||||
### 成功响应
|
||||
|
||||
```json
|
||||
{
|
||||
"success": true,
|
||||
"data": { ... },
|
||||
"timestamp": "2026-02-05T17:00:00.000000"
|
||||
}
|
||||
```
|
||||
|
||||
### 错误响应
|
||||
|
||||
```json
|
||||
{
|
||||
"detail": "错误信息描述"
|
||||
}
|
||||
```
|
||||
|
||||
### HTTP 状态码
|
||||
|
||||
| 状态码 | 说明 |
|
||||
|--------|------|
|
||||
| 200 | 成功 |
|
||||
| 400 | 请求参数错误 |
|
||||
| 401 | 未提供 API Key |
|
||||
| 404 | 未找到数据 |
|
||||
| 500 | 服务器错误 |
|
||||
|
||||
---
|
||||
|
||||
## 部署信息
|
||||
|
||||
| Agent | 镜像地址 | 端口 |
|
||||
|-------|----------|------|
|
||||
| Chain Explorer | `agnettaiji.azurecr.io/ai-agents/chain-explorer-agent:latest` | 8000 |
|
||||
| Chain Analysis | `agnettaiji.azurecr.io/ai-agents/chain-analysis-agent:latest` | 8000 |
|
||||
|
||||
### 环境变量
|
||||
|
||||
| 变量 | 默认值 | 说明 |
|
||||
|------|--------|------|
|
||||
| `SERVICE_HOST` | 0.0.0.0 | 服务绑定地址 |
|
||||
| `SERVICE_PORT` | 8000 | 服务端口 |
|
||||
| `LLM_BASE_URL` | https://litellm.xxx | LLM 服务地址 |
|
||||
| `LLM_MODEL` | taiji/gpt-4o-mini | LLM 模型 |
|
||||
|
||||
---
|
||||
|
||||
## 测试用地址
|
||||
|
||||
| 地址 | 说明 |
|
||||
|------|------|
|
||||
| `0xd8dA6BF26964aF9D7eEd9e03E53415D37aA96045` | Vitalik Buterin |
|
||||
| `0xBE0eB53F46cd790Cd13851d5EFf43D12404d33E8` | Binance Cold Wallet |
|
||||
| `0x28C6c06298d514Db089934071355E5743bf21d60` | Binance Hot Wallet |
|
||||
|
||||
---
|
||||
|
||||
## 最佳实践
|
||||
|
||||
1. **API Key 管理**:Etherscan API 有速率限制,建议申请付费 API Key
|
||||
2. **缓存策略**:对于不常变化的数据(如历史交易),建议本地缓存
|
||||
3. **并发控制**:避免短时间内大量请求,建议间隔 200ms
|
||||
4. **多链支持**:使用统一的 Etherscan V2 API,通过 chainid 区分网络
|
||||
|
||||
---
|
||||
|
||||
## 版本历史
|
||||
|
||||
| 版本 | 日期 | 更新内容 |
|
||||
|------|------|----------|
|
||||
| 1.0.0 | 2026-02-05 | 初始版本,支持 Etherscan V2 API |
|
||||
@@ -0,0 +1,563 @@
|
||||
# 美股 AI Agent 文档
|
||||
|
||||
本项目包含 **三个独立的美股 AI Agent 服务**,均通过 **HTTP API** 对外提供能力,支持智能对话分析。
|
||||
|
||||
- **Stock Quote Agent**:美股实时行情查询与分析
|
||||
- **Stock News Agent**:美股新闻资讯获取与解读
|
||||
- **Stock Analysis Agent**:美股技术分析与投资建议
|
||||
|
||||
---
|
||||
|
||||
## 认证方式
|
||||
|
||||
所有 Chat API 需要在请求头中提供 API Key,支持两种方式:
|
||||
|
||||
| 方式 | Header | 示例 |
|
||||
|------|--------|------|
|
||||
| api-key | `api-key` | `api-key: sk-xxxxx` |
|
||||
| Bearer Token | `Authorization` | `Authorization: Bearer sk-xxxxx` |
|
||||
|
||||
> ⚠️ 未提供认证信息将返回 `401 Unauthorized`
|
||||
|
||||
---
|
||||
|
||||
## Agent 1:Stock Quote Agent
|
||||
|
||||
### 功能概览
|
||||
|
||||
提供美股 **实时行情查询** 能力,支持自然语言交互,返回股票价格、涨跌幅、成交量等数据。
|
||||
|
||||
支持能力:
|
||||
|
||||
- 实时股票行情查询
|
||||
- 多股票批量查询
|
||||
- 热门股票行情
|
||||
- **AI 智能对话分析**
|
||||
|
||||
---
|
||||
|
||||
### 1️⃣ /chat — 智能对话
|
||||
|
||||
#### 功能说明
|
||||
|
||||
通过自然语言与 AI 交互,自动识别股票代码并返回行情数据及投资建议。
|
||||
|
||||
---
|
||||
|
||||
#### REST API 调用
|
||||
|
||||
```
|
||||
POST /chat
|
||||
Content-Type: application/json
|
||||
api-key: your-api-key
|
||||
```
|
||||
|
||||
```json
|
||||
{
|
||||
"message": "AAPL 和 TSLA 今天表现如何?"
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
#### 参数说明
|
||||
|
||||
| 参数 | 类型 | 必需 | 默认值 | 说明 |
|
||||
|------|------|------|--------|------|
|
||||
| message | string | ✅ | - | 用户消息(支持自然语言) |
|
||||
| user_id | string | ❌ | null | 用户ID(用于计费回调) |
|
||||
|
||||
**Header 参数:**
|
||||
|
||||
| 参数 | 类型 | 必需 | 说明 |
|
||||
|------|------|------|------|
|
||||
| api-key | string | ⚠ | API Key(二选一) |
|
||||
| Authorization | string | ⚠ | Bearer Token(二选一) |
|
||||
|
||||
---
|
||||
|
||||
#### 返回结果
|
||||
|
||||
```json
|
||||
{
|
||||
"response": "今天AAPL的股价为$274.43,涨跌幅为0.00%。TSLA的股价为$398.78...",
|
||||
"data": {
|
||||
"stocks": [
|
||||
{
|
||||
"success": true,
|
||||
"symbol": "AAPL",
|
||||
"name": "Apple Inc.",
|
||||
"price": 274.43,
|
||||
"change": 0,
|
||||
"change_percent": 0,
|
||||
"volume": 5212932,
|
||||
"high_52week": 288.62,
|
||||
"low_52week": 169.21
|
||||
},
|
||||
{
|
||||
"success": true,
|
||||
"symbol": "TSLA",
|
||||
"name": "Tesla, Inc.",
|
||||
"price": 398.78,
|
||||
"change": 0,
|
||||
"change_percent": 0,
|
||||
"volume": 7867089,
|
||||
"high_52week": 498.83,
|
||||
"low_52week": 214.25
|
||||
}
|
||||
],
|
||||
"symbols_detected": ["AAPL", "TSLA"]
|
||||
},
|
||||
"timestamp": "2026-02-05T15:30:13.347583"
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### 2️⃣ /quote — 单股行情查询
|
||||
|
||||
#### REST API 调用
|
||||
|
||||
```
|
||||
GET /quote?symbol=AAPL
|
||||
```
|
||||
|
||||
或
|
||||
|
||||
```
|
||||
POST /quote
|
||||
Content-Type: application/json
|
||||
```
|
||||
|
||||
```json
|
||||
{
|
||||
"symbol": "AAPL"
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
#### 参数说明
|
||||
|
||||
| 参数 | 类型 | 必需 | 说明 |
|
||||
|------|------|------|------|
|
||||
| symbol | string | ✅ | 股票代码(如 AAPL, TSLA) |
|
||||
|
||||
---
|
||||
|
||||
#### 返回结果
|
||||
|
||||
```json
|
||||
{
|
||||
"symbol": "AAPL",
|
||||
"name": "Apple Inc.",
|
||||
"price": 274.43,
|
||||
"change": 2.31,
|
||||
"change_percent": 0.85,
|
||||
"volume": 52129320,
|
||||
"market_cap": 4200000000000,
|
||||
"high_52week": 288.62,
|
||||
"low_52week": 169.21,
|
||||
"timestamp": "2026-02-05T15:30:00.000000"
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### 3️⃣ /batch — 批量行情查询
|
||||
|
||||
```
|
||||
POST /batch
|
||||
Content-Type: application/json
|
||||
```
|
||||
|
||||
```json
|
||||
{
|
||||
"symbols": ["AAPL", "TSLA", "NVDA", "MSFT", "GOOGL"]
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### 4️⃣ /popular — 热门股票行情
|
||||
|
||||
```
|
||||
GET /popular
|
||||
```
|
||||
|
||||
返回 AAPL, MSFT, GOOGL, AMZN, TSLA, NVDA, META 等热门股票行情。
|
||||
|
||||
---
|
||||
|
||||
## Agent 2:Stock News Agent
|
||||
|
||||
### 功能概览
|
||||
|
||||
提供美股 **新闻资讯获取与分析** 能力,支持按股票代码或关键词搜索新闻。
|
||||
|
||||
支持能力:
|
||||
|
||||
- 股票相关新闻查询
|
||||
- 市场动态获取
|
||||
- 热门财经新闻
|
||||
- **AI 新闻解读与影响分析**
|
||||
|
||||
---
|
||||
|
||||
### 1️⃣ /chat — 智能对话
|
||||
|
||||
#### 功能说明
|
||||
|
||||
通过自然语言获取股票新闻并进行 AI 分析解读。
|
||||
|
||||
---
|
||||
|
||||
#### REST API 调用
|
||||
|
||||
```
|
||||
POST /chat
|
||||
Content-Type: application/json
|
||||
api-key: your-api-key
|
||||
```
|
||||
|
||||
```json
|
||||
{
|
||||
"message": "NVDA 最近有什么重要新闻?"
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
#### 返回结果
|
||||
|
||||
```json
|
||||
{
|
||||
"response": "最近关于NVDA的新闻显示出其股票在盘前交易中表现强劲,市场对其AI芯片业务的前景保持乐观...",
|
||||
"data": {
|
||||
"news_count": 5,
|
||||
"symbols": ["NVDA"]
|
||||
},
|
||||
"timestamp": "2026-02-05T15:31:29.298428"
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### 2️⃣ /news — 获取新闻
|
||||
|
||||
```
|
||||
POST /news
|
||||
Content-Type: application/json
|
||||
```
|
||||
|
||||
```json
|
||||
{
|
||||
"symbol": "AAPL",
|
||||
"limit": 10
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
#### 参数说明
|
||||
|
||||
| 参数 | 类型 | 必需 | 默认值 | 说明 |
|
||||
|------|------|------|--------|------|
|
||||
| symbol | string | ⚠ | null | 股票代码 |
|
||||
| query | string | ⚠ | null | 搜索关键词 |
|
||||
| limit | integer | ❌ | 10 | 返回新闻数量(1-50) |
|
||||
|
||||
> `symbol` 与 `query` 二选一
|
||||
|
||||
---
|
||||
|
||||
### 3️⃣ /market — 市场动态
|
||||
|
||||
```
|
||||
GET /market
|
||||
```
|
||||
|
||||
返回市场涨跌排行、活跃股票等信息。
|
||||
|
||||
---
|
||||
|
||||
### 4️⃣ /trending — 热门新闻
|
||||
|
||||
```
|
||||
GET /trending
|
||||
```
|
||||
|
||||
返回当前热门财经新闻。
|
||||
|
||||
---
|
||||
|
||||
## Agent 3:Stock Analysis Agent
|
||||
|
||||
### 功能概览
|
||||
|
||||
提供美股 **技术分析** 能力,计算技术指标并给出交易信号与投资建议。
|
||||
|
||||
支持能力:
|
||||
|
||||
- 技术指标计算(SMA, RSI, MACD, 布林带)
|
||||
- 趋势判断(看涨/看跌/中性)
|
||||
- 买卖信号生成
|
||||
- 支撑位/阻力位计算
|
||||
- **AI 综合分析与投资建议**
|
||||
|
||||
---
|
||||
|
||||
### 1️⃣ /chat — 智能对话
|
||||
|
||||
#### 功能说明
|
||||
|
||||
通过自然语言获取股票技术分析并由 AI 提供投资建议。
|
||||
|
||||
---
|
||||
|
||||
#### REST API 调用
|
||||
|
||||
```
|
||||
POST /chat
|
||||
Content-Type: application/json
|
||||
Authorization: Bearer your-api-key
|
||||
```
|
||||
|
||||
```json
|
||||
{
|
||||
"message": "帮我分析 AAPL,现在适合买入吗?"
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
#### 返回结果
|
||||
|
||||
```json
|
||||
{
|
||||
"response": "根据技术分析数据,AAPL当前价格为$274.31,趋势为中性,信号为持有。RSI值为66.49,接近超买区域...",
|
||||
"data": {
|
||||
"analysis": [
|
||||
{
|
||||
"symbol": "AAPL",
|
||||
"current_price": 274.31,
|
||||
"indicators": {
|
||||
"sma_20": 259.12,
|
||||
"sma_50": 268.63,
|
||||
"sma_200": null,
|
||||
"rsi_14": 66.49,
|
||||
"macd": -0.9012,
|
||||
"macd_signal": -0.8111,
|
||||
"bollinger_upper": 275.39,
|
||||
"bollinger_lower": 242.84
|
||||
},
|
||||
"trend": "neutral",
|
||||
"signal": "hold",
|
||||
"support_level": 243.42,
|
||||
"resistance_level": 279.5,
|
||||
"risk_level": "medium"
|
||||
}
|
||||
],
|
||||
"symbols": ["AAPL"]
|
||||
},
|
||||
"timestamp": "2026-02-05T15:58:50.578048"
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### 2️⃣ /analyze — 技术分析
|
||||
|
||||
```
|
||||
POST /analyze
|
||||
Content-Type: application/json
|
||||
```
|
||||
|
||||
```json
|
||||
{
|
||||
"symbol": "AAPL"
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
#### 返回字段说明
|
||||
|
||||
| 字段 | 类型 | 说明 |
|
||||
|------|------|------|
|
||||
| current_price | float | 当前价格 |
|
||||
| sma_20 | float | 20日简单移动平均线 |
|
||||
| sma_50 | float | 50日简单移动平均线 |
|
||||
| sma_200 | float | 200日简单移动平均线 |
|
||||
| rsi_14 | float | 14日相对强弱指数 |
|
||||
| macd | float | MACD 值 |
|
||||
| macd_signal | float | MACD 信号线 |
|
||||
| bollinger_upper | float | 布林带上轨 |
|
||||
| bollinger_lower | float | 布林带下轨 |
|
||||
| trend | string | 趋势:bullish / bearish / neutral |
|
||||
| signal | string | 信号:buy / sell / hold |
|
||||
| support_level | float | 支撑位 |
|
||||
| resistance_level | float | 阻力位 |
|
||||
| risk_level | string | 风险等级:low / medium / high |
|
||||
|
||||
---
|
||||
|
||||
### 3️⃣ /compare — 多股对比
|
||||
|
||||
```
|
||||
POST /compare
|
||||
Content-Type: application/json
|
||||
```
|
||||
|
||||
```json
|
||||
{
|
||||
"symbols": ["AAPL", "MSFT", "GOOGL"]
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### 4️⃣ /screen — 股票筛选
|
||||
|
||||
```
|
||||
GET /screen?trend=bullish&signal=buy
|
||||
```
|
||||
|
||||
根据技术指标筛选符合条件的股票。
|
||||
|
||||
---
|
||||
|
||||
## 统一错误格式
|
||||
|
||||
**成功:**
|
||||
|
||||
```json
|
||||
{
|
||||
"response": "AI 分析结果...",
|
||||
"data": {},
|
||||
"timestamp": "2026-02-05T15:30:00.000000"
|
||||
}
|
||||
```
|
||||
|
||||
**认证失败(401):**
|
||||
|
||||
```json
|
||||
{
|
||||
"detail": "请在请求头中提供 api-key 或 Authorization"
|
||||
}
|
||||
```
|
||||
|
||||
**请求错误(400):**
|
||||
|
||||
```json
|
||||
{
|
||||
"detail": "错误描述"
|
||||
}
|
||||
```
|
||||
|
||||
**服务器错误(500):**
|
||||
|
||||
```json
|
||||
{
|
||||
"detail": "Internal Server Error"
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 调用示例
|
||||
|
||||
### cURL 示例
|
||||
|
||||
**使用 api-key Header:**
|
||||
|
||||
```bash
|
||||
curl -X POST "http://test-stock-quote.taijiagnet.com/chat" \
|
||||
-H "Content-Type: application/json" \
|
||||
-H "api-key: sk-mPV5MVVVVvfGSkXA-ASQXQ" \
|
||||
-d '{"message": "AAPL 和 TSLA 今天表现如何?"}'
|
||||
```
|
||||
|
||||
**使用 Authorization Bearer:**
|
||||
|
||||
```bash
|
||||
curl -X POST "http://test-stock-analysis.taijiagnet.com/chat" \
|
||||
-H "Content-Type: application/json" \
|
||||
-H "Authorization: Bearer sk-mPV5MVVVVvfGSkXA-ASQXQ" \
|
||||
-d '{"message": "帮我分析 NVDA,现在适合买入吗?"}'
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### Python 示例
|
||||
|
||||
```python
|
||||
import requests
|
||||
|
||||
API_KEY = "sk-mPV5MVVVVvfGSkXA-ASQXQ"
|
||||
|
||||
# Stock Quote Agent
|
||||
response = requests.post(
|
||||
"http://test-stock-quote.taijiagnet.com/chat",
|
||||
headers={
|
||||
"Content-Type": "application/json",
|
||||
"api-key": API_KEY
|
||||
},
|
||||
json={"message": "AAPL 现在多少钱?"}
|
||||
)
|
||||
print(response.json())
|
||||
|
||||
# Stock News Agent
|
||||
response = requests.post(
|
||||
"http://test-stock-news.taijiagnet.com/chat",
|
||||
headers={
|
||||
"Content-Type": "application/json",
|
||||
"Authorization": f"Bearer {API_KEY}"
|
||||
},
|
||||
json={"message": "TSLA 最近有什么新闻?"}
|
||||
)
|
||||
print(response.json())
|
||||
|
||||
# Stock Analysis Agent
|
||||
response = requests.post(
|
||||
"http://test-stock-analysis.taijiagnet.com/chat",
|
||||
headers={
|
||||
"Content-Type": "application/json",
|
||||
"api-key": API_KEY
|
||||
},
|
||||
json={"message": "帮我技术分析 NVDA"}
|
||||
)
|
||||
print(response.json())
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 部署信息
|
||||
|
||||
| Agent | 模板名称 | 镜像 | 端口 |
|
||||
|-------|----------|------|------|
|
||||
| Stock Quote | stock_quote_agent | agnettaiji.azurecr.io/ai-agents/stock-quote-agent:latest | 8080 |
|
||||
| Stock News | stock_news_agent | agnettaiji.azurecr.io/ai-agents/stock-news-agent:latest | 8080 |
|
||||
| Stock Analysis | stock_analysis_agent | agnettaiji.azurecr.io/ai-agents/stock-analysis-agent:latest | 8080 |
|
||||
|
||||
---
|
||||
|
||||
## 环境变量配置
|
||||
|
||||
| 变量名 | 说明 | 默认值 |
|
||||
|--------|------|--------|
|
||||
| LLM_BASE_URL | LLM 服务地址 | https://litellm.graystone-fb459c5d.southeastasia.azurecontainerapps.io/v1 |
|
||||
| LLM_MODEL | 模型名称 | taiji/gpt-4o-mini |
|
||||
| SERVICE_HOST | 服务监听地址 | 0.0.0.0 |
|
||||
| SERVICE_PORT | 服务端口 | 8080 |
|
||||
|
||||
---
|
||||
|
||||
## 免责声明
|
||||
|
||||
> ⚠️ **投资有风险,入市需谨慎。** 本 Agent 提供的分析和建议仅供参考,不构成任何投资建议。用户应自行判断并承担投资风险。
|
||||
|
||||
---
|
||||
|
||||
*文档版本:v1.0.0*
|
||||
*更新日期:2026-02-05*
|
||||
@@ -0,0 +1,125 @@
|
||||
# Agent 端口统一为 8000 改动方案(除 search_agent 外)
|
||||
|
||||
目标:除 **search_agent / search_agent_a2a / search_agent_mcp** 外,所有 agent 统一使用端口 **8000**。
|
||||
|
||||
---
|
||||
|
||||
## 一、当前状态汇总
|
||||
|
||||
| 模板名 | template_manager port | k8s TEMPLATE_PORTS | k8s container_port | 代码/Dockerfile 默认 | 需改动 |
|
||||
|--------|------------------------|--------------------|---------------------|-----------------------|--------|
|
||||
| search_agent | 8000 | 8080 | 8080 | 8080 | 否(保持 8080) |
|
||||
| search_agent_a2a | 8000 | 8080 | 8080 | 8080 | 否(保持 8080) |
|
||||
| search_agent_mcp | 8000 | 8080 | 8080 | 8080 | 否(保持 8080) |
|
||||
| jina_search_agent | 8080 | 8080 | 8080 | 8080 | **是 → 8000** |
|
||||
| azure_blob_agent | 8080 | 8080 | 8080 | 8080 | **是 → 8000** |
|
||||
| azure_blob_agent_mcp | 8080 | 8080 | 8080 | 8080 | **是 → 8000** |
|
||||
| azure_blob_agent_a2a | 8080 | 8080 | 8080 | 8080 | **是 → 8000** |
|
||||
| a2a_litellm_agent | 8080 | 8080 | 8080 | 8080 | **是 → 8000** |
|
||||
| echo_agent / mysql_agent / postgresql_agent / code_ai_agent / facebook_agent / media_downloader / content_analyzer / huoke / microsoft_learn_agent / aws_docs_mcp / google_mcp | 8000 | 8000 | 8000 或未显式 | 8000 | 否 |
|
||||
|
||||
说明:search_agent 系列在 template_manager 中为 8000 与容器实际 8080 不一致,若需与现状一致可单独将 template_manager 中 search_agent* 改为 8080(本次方案不包含,仅统一「非 search_agent」为 8000)。
|
||||
|
||||
---
|
||||
|
||||
## 二、需改动的 5 个 Agent
|
||||
|
||||
1. **jina_search_agent**
|
||||
2. **azure_blob_agent**
|
||||
3. **azure_blob_agent_mcp**
|
||||
4. **azure_blob_agent_a2a**
|
||||
5. **a2a_litellm_agent**
|
||||
|
||||
---
|
||||
|
||||
## 三、具体改动清单
|
||||
|
||||
### 1. template_manager.py
|
||||
|
||||
- **jina_search_agent**:`"port": 8080` → `"port": 8000`
|
||||
- **azure_blob_agent**:`"port": 8080` → `"port": 8000`
|
||||
- **azure_blob_agent_mcp**:`"port": 8080` → `"port": 8000`
|
||||
- **azure_blob_agent_a2a**:`"port": 8080` → `"port": 8000`
|
||||
- **a2a_litellm_agent**:`"port": 8080` → `"port": 8000`
|
||||
|
||||
### 2. k8s_manager.py
|
||||
|
||||
- **TEMPLATE_PORTS**:上述 5 个模板的端口由 `8080` 改为 `8000`(search_agent / search_agent_a2a / search_agent_mcp 保持 8080)。
|
||||
- **container_ports**(约 746–750 行):
|
||||
- 当前 8080 列表:`jina_search_agent, azure_blob_agent, azure_blob_agent_mcp, azure_blob_agent_a2a, search_agent, search_agent_a2a, search_agent_mcp, a2a_litellm_agent`
|
||||
- 修改为:仅 **search_agent, search_agent_a2a, search_agent_mcp** 使用 `container_port=8080`;**jina_search_agent, azure_blob_agent, azure_blob_agent_mcp, azure_blob_agent_a2a, a2a_litellm_agent** 移到 8000 分支,使用 `container_port=8000`。
|
||||
- **TEMPLATE_ENV_INFO** 中「SERVICE_PORT」说明:
|
||||
- **jina_search_agent**:`"默认8080"` → `"默认8000"`
|
||||
- **azure_blob_agent**:`"默认8080"` → `"默认8000"`
|
||||
- **azure_blob_agent_mcp**:无 SERVICE_PORT 时可补充 `"SERVICE_PORT": "HTTP服务端口,默认8000"`(若有则改为 8000)
|
||||
- **azure_blob_agent_a2a**:同上
|
||||
- **a2a_litellm_agent**:`"默认8080"` → `"默认8000"`
|
||||
|
||||
### 3. Agent 代码与 Dockerfile(每个 agent 内默认端口 8080 → 8000)
|
||||
|
||||
#### 3.1 jina_search_agent
|
||||
|
||||
| 文件 | 改动 |
|
||||
|------|------|
|
||||
| `agent_templates/agents/jina_search_agent/jina_search_agent.py` | `SERVICE_PORT = int(os.getenv("SERVICE_PORT", "8080"))` → `"8000"` |
|
||||
| `agent_templates/agents/jina_search_agent/jina_search_agent.Dockerfile` | `ENV SERVICE_PORT=8080` → `8000`;健康检查与 `EXPOSE` 中 `8080` → `8000` |
|
||||
|
||||
#### 3.2 azure_blob_agent
|
||||
|
||||
| 文件 | 改动 |
|
||||
|------|------|
|
||||
| `agent_templates/agents/azure_blob_agent/azure_blob_agent.py` | `SERVICE_PORT = int(os.getenv("SERVICE_PORT", "8080"))` → `"8000"` |
|
||||
| `agent_templates/agents/azure_blob_agent/azure_blob_agent.Dockerfile` | `ENV SERVICE_PORT=8080` → `8000`;健康检查 URL 中 `8080` → `8000` |
|
||||
|
||||
#### 3.3 azure_blob_agent_mcp
|
||||
|
||||
| 文件 | 改动 |
|
||||
|------|------|
|
||||
| `agent_templates/agents/azure_blob_agent_mcp/azure_blob_agent_mcp.py` | `SERVICE_PORT = int(os.getenv("SERVICE_PORT", "8080"))` → `"8000"` |
|
||||
| `agent_templates/agents/azure_blob_agent_mcp/azure_blob_agent_mcp.Dockerfile` | `EXPOSE 8080` → `8000`;健康检查 URL 中 `8080` → `8000` |
|
||||
|
||||
#### 3.4 azure_blob_agent_a2a
|
||||
|
||||
| 文件 | 改动 |
|
||||
|------|------|
|
||||
| `agent_templates/agents/azure_blob_agent_a2a/azure_blob_agent_a2a.py` | `SERVICE_PORT = int(os.getenv("SERVICE_PORT", "8080"))` → `"8000"` |
|
||||
| `agent_templates/agents/azure_blob_agent_a2a/azure_blob_agent_a2a.Dockerfile` | `EXPOSE 8080` → `8000`;健康检查 URL 中 `8080` → `8000` |
|
||||
|
||||
#### 3.5 a2a_litellm_agent
|
||||
|
||||
| 文件 | 改动 |
|
||||
|------|------|
|
||||
| `agent_templates/agents/a2a_litellm_agent/main.py` | `SERVICE_PORT = int(os.getenv("SERVICE_PORT", "8080"))` → `"8000"` |
|
||||
| `agent_templates/agents/a2a_litellm_agent/config.py` | `port: int = 8080` → `port: int = 8000` |
|
||||
| `agent_templates/agents/a2a_litellm_agent/a2a_server.py` | `SERVICE_PORT = int(os.getenv("SERVICE_PORT", "8080"))` → `"8000"`;文档/示例中的 `--port 8080` → `--port 8000` |
|
||||
| `agent_templates/agents/a2a_litellm_agent/a2a_litellm_agent.Dockerfile` | `ENV SERVICE_PORT=8080` → `8000`;健康检查与 `EXPOSE` 中 `8080` → `8000` |
|
||||
|
||||
---
|
||||
|
||||
## 四、未在 Agent Manager 默认模板中的 Agent(可选统一)
|
||||
|
||||
以下 agent 在仓库中存在且默认使用 8080,但未出现在 `template_manager.DEFAULT_TEMPLATES` / `k8s_manager.TEMPLATE_PORTS` 中。若希望「全仓库除 search_agent 外一律 8000」,可一并改:
|
||||
|
||||
- **chain_explorer_agent**:`chain_explorer_agent.py` 默认 8080 → 8000;Dockerfile 已是 8000,无需改。
|
||||
- **chain_analysis_agent**:同上(代码 8080 → 8000,Dockerfile 已是 8000)。
|
||||
- **stock_analysis_agent**:`stock_analysis_agent.py` 默认 8080 → 8000;Dockerfile 中 8080 → 8000。
|
||||
- **stock_news_agent**:同上。
|
||||
- **stock_quote_agent**:同上。
|
||||
|
||||
若这些模板后续加入 Agent Manager,建议直接以 8000 登记。
|
||||
|
||||
---
|
||||
|
||||
## 五、实施顺序建议
|
||||
|
||||
1. 改 **template_manager.py**、**k8s_manager.py**(端口与 K8s 一致)。
|
||||
2. 改上述 5 个 agent 的 **代码 + Dockerfile**。
|
||||
3. 重新构建并推送对应镜像;已运行中的 Pod 需用新镜像重启或重新部署。
|
||||
4. 若数据库已初始化过模板,需将上述 5 个模板的 `port` 字段更新为 8000(或重新从默认模板初始化)。
|
||||
|
||||
---
|
||||
|
||||
## 六、改动后端口约定(总结)
|
||||
|
||||
- **8000**:除 search_agent 系列外的所有 agent(含 jina_search_agent、azure_blob_agent 系列、a2a_litellm_agent、echo/mysql/postgresql/code_ai/facebook 等)。
|
||||
- **8080**:仅 **search_agent**、**search_agent_a2a**、**search_agent_mcp**。
|
||||
@@ -21,6 +21,7 @@ from fastapi import APIRouter, HTTPException, Query
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from agent_code_generator import agent_code_generator
|
||||
from k8s_manager import sanitize_k8s_name
|
||||
from tool_storage import tool_storage
|
||||
from gitee_manager import gitee_manager
|
||||
|
||||
@@ -860,10 +861,10 @@ async def create_agent_with_tools(request: CreateAgentWithToolsRequest):
|
||||
|
||||
# 生成唯一后缀,确保不同用户创建同名 Agent 不会冲突
|
||||
unique_suffix = uuid.uuid4().hex[:6]
|
||||
base_name = request.name.lower().replace("_", "-").replace(" ", "-")
|
||||
base_name = sanitize_k8s_name(request.name)
|
||||
|
||||
# k8s_name 带唯一后缀,避免域名/namespace 冲突
|
||||
k8s_name = f"{base_name}-{unique_suffix}"
|
||||
# k8s_name 带唯一后缀,避免域名/namespace 冲突(确保 DNS-1035 合规)
|
||||
k8s_name = sanitize_k8s_name(f"{base_name}-{unique_suffix}")
|
||||
repo_name = f"agent-{k8s_name}"
|
||||
agent_ref_id = f"agent-{repo_name}"
|
||||
|
||||
@@ -905,7 +906,7 @@ async def create_agent_with_tools(request: CreateAgentWithToolsRequest):
|
||||
"ACR_USERNAME": "agnettaiji",
|
||||
"ACR_PASSWORD": "hDpX5t34N5ZmnKdtqyjYL5co/SnXJrmD20CRpGpWaG+ACRCw2wGM",
|
||||
"AZ_CLIENT_ID": "f2dd1cb2-02f6-4efb-bc72-d148f6e01545",
|
||||
"AZ_CLIENT_SECRET": "S0J8Q~DE.DEu29nreaBn2EbeuGOg7GEIkonRMbYj",
|
||||
"AZ_CLIENT_SECRET": "UVU8Q~Hcrf5KeLi2RvUXB2rcuKFEjRCCrf_JrbwA",
|
||||
"AZ_TENANT_ID": "263c3ff6-1be5-4141-8308-b188464fb297",
|
||||
"AZ_SUBSCRIPTION_ID": "45d7a360-af09-40fc-9afc-56dc475245ec",
|
||||
"AZ_RG": "taiji-ai-pda",
|
||||
@@ -1102,12 +1103,12 @@ async def get_agent_build_status(agent_ref_id: str):
|
||||
if agent_info and agent_info.get("k8s_name"):
|
||||
k8s_name = agent_info.get("k8s_name")
|
||||
else:
|
||||
# 无法从 repo_name 准确推断,使用 agent 名称
|
||||
# 无法从 repo_name 准确推断,使用 agent 名称(确保 DNS-1035 合规)
|
||||
agent_name = agent_info.get("name") if agent_info else None
|
||||
if agent_name:
|
||||
k8s_name = agent_name.lower().replace("_", "-").replace(" ", "-")
|
||||
k8s_name = sanitize_k8s_name(agent_name)
|
||||
else:
|
||||
k8s_name = repo_name.lower().replace("_", "-").replace(" ", "-")
|
||||
k8s_name = sanitize_k8s_name(repo_name)
|
||||
|
||||
expected_domain = f"{k8s_name}.taijiagnet.com"
|
||||
namespace = f"agent-{k8s_name}"
|
||||
@@ -1172,7 +1173,7 @@ async def get_agent_deployment_info(agent_ref_id: str):
|
||||
repo_name = agent_ref_id.replace("agent-", "", 1)
|
||||
else:
|
||||
repo_name = agent_ref_id
|
||||
k8s_name = repo_name.lower().replace("_", "-").replace(" ", "-")
|
||||
k8s_name = sanitize_k8s_name(repo_name)
|
||||
namespace = f"agent-{k8s_name}"
|
||||
|
||||
# 导入 K8sManager 查询实际状态
|
||||
|
||||
@@ -10,3 +10,14 @@ data:
|
||||
AZURE_DNS_ZONE: "taijiagnet.com"
|
||||
AZURE_SUBSCRIPTION_ID: "your-subscription-id"
|
||||
AZURE_RESOURCE_GROUP: "your-resource-group"
|
||||
|
||||
# Gitee 配置(非敏感信息)
|
||||
GITEE_API_URL: "http://gitee.ath.cx:3000/api/v1"
|
||||
GITEE_BASE_URL: "http://gitee.ath.cx:3000"
|
||||
GITEE_OWNER: "xiaohei"
|
||||
GITEE_USERNAME: "zhanggangyong"
|
||||
GITEE_TEMPLATE_REPO: "cicd-AKS"
|
||||
|
||||
# ACR 配置
|
||||
ACR_REGISTRY: "agnettaiji.azurecr.io"
|
||||
ACR_NAMESPACE: "ai-agents"
|
||||
@@ -61,6 +61,17 @@ spec:
|
||||
secretKeyRef:
|
||||
name: agent-manager-secret
|
||||
key: AZURE_CLIENT_SECRET
|
||||
# Gitee 凭据
|
||||
- name: GITEE_TOKEN
|
||||
valueFrom:
|
||||
secretKeyRef:
|
||||
name: agent-manager-secret
|
||||
key: GITEE_TOKEN
|
||||
- name: GITEE_PASSWORD
|
||||
valueFrom:
|
||||
secretKeyRef:
|
||||
name: agent-manager-secret
|
||||
key: GITEE_PASSWORD
|
||||
|
||||
# 挂载 kubeconfig(用于管理其他 Agent)
|
||||
volumeMounts:
|
||||
|
||||
@@ -10,5 +10,9 @@ stringData:
|
||||
AZURE_CLIENT_ID: "your-client-id"
|
||||
AZURE_CLIENT_SECRET: "your-client-secret"
|
||||
|
||||
# Gitee 凭据(敏感信息)
|
||||
GITEE_TOKEN: "your-gitee-token"
|
||||
GITEE_PASSWORD: "your-gitee-password"
|
||||
|
||||
# 数据库密码(如果需要单独管理)
|
||||
# DB_PASSWORD: "By@123456."
|
||||
|
||||
+96
-48
@@ -6,6 +6,7 @@ from kubernetes.client.rest import ApiException
|
||||
from typing import Dict, List, Optional
|
||||
import logging
|
||||
import os
|
||||
import re
|
||||
import requests
|
||||
import time
|
||||
|
||||
@@ -13,6 +14,45 @@ logging.basicConfig(level=logging.INFO)
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def sanitize_k8s_name(name: str, max_length: int = 63) -> str:
|
||||
"""将名称转换为 DNS-1035 合规格式
|
||||
|
||||
Kubernetes Service、Deployment 等资源名称必须符合 DNS-1035 标准:
|
||||
- 只能包含小写字母、数字和连字符 '-'
|
||||
- 必须以字母开头
|
||||
- 必须以字母或数字结尾
|
||||
- 最长 63 个字符
|
||||
|
||||
Args:
|
||||
name: 原始名称
|
||||
max_length: 最大长度(默认 63)
|
||||
|
||||
Returns:
|
||||
合规的 K8s 资源名称
|
||||
"""
|
||||
# 转小写
|
||||
name = name.lower()
|
||||
# 将下划线和空格替换为连字符
|
||||
name = name.replace("_", "-").replace(" ", "-")
|
||||
# 移除非字母、数字、连字符的字符
|
||||
name = re.sub(r'[^a-z0-9-]', '', name)
|
||||
# 合并连续的连字符
|
||||
name = re.sub(r'-+', '-', name)
|
||||
# 如果以数字开头,添加 'a' 前缀
|
||||
if name and name[0].isdigit():
|
||||
name = 'a' + name
|
||||
# 如果以连字符开头,去掉
|
||||
name = name.lstrip('-')
|
||||
# 截断到最大长度
|
||||
name = name[:max_length]
|
||||
# 去掉末尾的连字符
|
||||
name = name.rstrip('-')
|
||||
# 最终兜底:如果名称为空
|
||||
if not name:
|
||||
name = 'agent'
|
||||
return name
|
||||
|
||||
|
||||
class K8sManager:
|
||||
"""Kubernetes资源管理器"""
|
||||
|
||||
@@ -94,11 +134,9 @@ class K8sManager:
|
||||
Returns:
|
||||
创建的命名空间名称
|
||||
"""
|
||||
# 生成命名空间名称(使用 agent-{agent_name} 格式)
|
||||
namespace_name = f"agent-{agent_name}"
|
||||
|
||||
# 确保命名空间名称符合 DNS 标准(最多 63 个字符,只能包含小写字母、数字和连字符)
|
||||
namespace_name = namespace_name[:63].lower().strip('-')
|
||||
# 生成命名空间名称(使用 agent-{sanitized_name} 格式,确保 DNS 合规)
|
||||
sanitized_name = sanitize_k8s_name(agent_name)
|
||||
namespace_name = f"agent-{sanitized_name}"[:63].rstrip('-')
|
||||
|
||||
try:
|
||||
# 检查命名空间是否已存在
|
||||
@@ -132,41 +170,45 @@ class K8sManager:
|
||||
raise
|
||||
|
||||
def _copy_acr_secret_to_namespace(self, target_namespace: str):
|
||||
"""复制 ACR secret 到目标命名空间
|
||||
"""复制所有 ACR secrets 到目标命名空间
|
||||
|
||||
Args:
|
||||
target_namespace: 目标命名空间
|
||||
"""
|
||||
try:
|
||||
# 从 agent-manager 命名空间读取 acr-secret
|
||||
source_secret = self.v1.read_namespaced_secret(
|
||||
name="acr-secret",
|
||||
namespace="agent-manager"
|
||||
)
|
||||
|
||||
# 创建新的 secret(去除自动生成的字段)
|
||||
new_secret = client.V1Secret(
|
||||
metadata=client.V1ObjectMeta(
|
||||
name="acr-secret",
|
||||
namespace=target_namespace
|
||||
),
|
||||
data=source_secret.data,
|
||||
type=source_secret.type
|
||||
)
|
||||
|
||||
# 在目标命名空间创建 secret
|
||||
self.v1.create_namespaced_secret(
|
||||
namespace=target_namespace,
|
||||
body=new_secret
|
||||
)
|
||||
logger.info(f"✅ 已复制 ACR secret 到命名空间 {target_namespace}")
|
||||
except ApiException as e:
|
||||
if e.status == 404:
|
||||
logger.warning(f"⚠️ 源 ACR secret 不存在,跳过复制")
|
||||
elif e.status == 409:
|
||||
logger.info(f"ACR secret 已存在于命名空间 {target_namespace}")
|
||||
else:
|
||||
logger.error(f"❌ 复制 ACR secret 失败: {e}")
|
||||
# 需要复制的 ACR secrets 列表
|
||||
acr_secrets = ["acr-secret", "openclaw-acr-secret"]
|
||||
|
||||
for secret_name in acr_secrets:
|
||||
try:
|
||||
# 从 agent-manager 命名空间读取 secret
|
||||
source_secret = self.v1.read_namespaced_secret(
|
||||
name=secret_name,
|
||||
namespace="agent-manager"
|
||||
)
|
||||
|
||||
# 创建新的 secret(去除自动生成的字段)
|
||||
new_secret = client.V1Secret(
|
||||
metadata=client.V1ObjectMeta(
|
||||
name=secret_name,
|
||||
namespace=target_namespace
|
||||
),
|
||||
data=source_secret.data,
|
||||
type=source_secret.type
|
||||
)
|
||||
|
||||
# 在目标命名空间创建 secret
|
||||
self.v1.create_namespaced_secret(
|
||||
namespace=target_namespace,
|
||||
body=new_secret
|
||||
)
|
||||
logger.info(f"✅ 已复制 {secret_name} 到命名空间 {target_namespace}")
|
||||
except ApiException as e:
|
||||
if e.status == 404:
|
||||
logger.warning(f"⚠️ 源 {secret_name} 不存在,跳过复制")
|
||||
elif e.status == 409:
|
||||
logger.info(f"{secret_name} 已存在于命名空间 {target_namespace}")
|
||||
else:
|
||||
logger.error(f"❌ 复制 {secret_name} 失败: {e}")
|
||||
|
||||
def create_service(self, service_name: str, namespace: str, pod_selector: Dict[str, str],
|
||||
service_port: int, target_port: int) -> Dict:
|
||||
@@ -318,11 +360,11 @@ class K8sManager:
|
||||
"search_agent_mcp": 8080,
|
||||
"mysql_agent": 8000,
|
||||
"postgresql_agent": 8000,
|
||||
"jina_search_agent": 8080,
|
||||
"azure_blob_agent": 8080,
|
||||
"azure_blob_agent_mcp": 8080,
|
||||
"azure_blob_agent_a2a": 8080,
|
||||
"a2a_litellm_agent": 8080,
|
||||
"jina_search_agent": 8000,
|
||||
"azure_blob_agent": 8000,
|
||||
"azure_blob_agent_mcp": 8000,
|
||||
"azure_blob_agent_a2a": 8000,
|
||||
"a2a_litellm_agent": 8000,
|
||||
"code_ai_agent": 8000,
|
||||
"facebook_agent": 8000,
|
||||
"media_downloader": 8000,
|
||||
@@ -340,7 +382,7 @@ class K8sManager:
|
||||
"JINA_API_KEY": "Jina API密钥,从 https://jina.ai/ 获取"
|
||||
},
|
||||
"optional": {
|
||||
"SERVICE_PORT": "HTTP服务端口,默认8080",
|
||||
"SERVICE_PORT": "HTTP服务端口,默认8000",
|
||||
"SERVICE_HOST": "HTTP服务监听地址,默认0.0.0.0"
|
||||
}
|
||||
},
|
||||
@@ -376,7 +418,7 @@ class K8sManager:
|
||||
},
|
||||
"optional": {
|
||||
"AZURE_STORAGE_CONNECTION_STRING": "Azure Storage连接字符串(可选,也可通过 /connect API 动态传入)",
|
||||
"SERVICE_PORT": "HTTP服务端口,默认8080",
|
||||
"SERVICE_PORT": "HTTP服务端口,默认8000",
|
||||
"SERVICE_HOST": "HTTP服务监听地址,默认0.0.0.0"
|
||||
}
|
||||
},
|
||||
@@ -479,7 +521,7 @@ class K8sManager:
|
||||
"LITELLM_API_KEY": "LiteLLM API 密钥(可在请求中传入)",
|
||||
"AGENT_NAME": "Agent 名称",
|
||||
"AGENT_DESCRIPTION": "Agent 描述",
|
||||
"SERVICE_PORT": "HTTP服务端口,默认 8080",
|
||||
"SERVICE_PORT": "HTTP服务端口,默认 8000",
|
||||
"SERVICE_HOST": "HTTP服务监听地址,默认 0.0.0.0"
|
||||
}
|
||||
},
|
||||
@@ -706,9 +748,9 @@ class K8sManager:
|
||||
|
||||
# 设置容器端口(如果是HTTP服务类型的agent)
|
||||
container_ports = None
|
||||
if template in ["jina_search_agent", "azure_blob_agent", "azure_blob_agent_mcp", "azure_blob_agent_a2a", "search_agent", "search_agent_a2a", "search_agent_mcp", "a2a_litellm_agent"]:
|
||||
if template in ["search_agent", "search_agent_a2a", "search_agent_mcp"]:
|
||||
container_ports = [client.V1ContainerPort(container_port=8080)]
|
||||
elif template in ["code_ai_agent", "facebook_agent", "echo_agent", "mysql_agent", "postgresql_agent"]:
|
||||
elif template in ["jina_search_agent", "azure_blob_agent", "azure_blob_agent_mcp", "azure_blob_agent_a2a", "a2a_litellm_agent", "code_ai_agent", "facebook_agent", "echo_agent", "mysql_agent", "postgresql_agent"]:
|
||||
container_ports = [client.V1ContainerPort(container_port=8000)]
|
||||
|
||||
# 创建Pod规格
|
||||
@@ -723,10 +765,15 @@ class K8sManager:
|
||||
ports=container_ports
|
||||
)
|
||||
|
||||
# 根据镜像来源选择合适的 imagePullSecrets
|
||||
image_pull_secrets = [client.V1LocalObjectReference(name="acr-secret")]
|
||||
if "openclawacr" in image:
|
||||
image_pull_secrets.append(client.V1LocalObjectReference(name="openclaw-acr-secret"))
|
||||
|
||||
pod_spec = client.V1PodSpec(
|
||||
containers=[container],
|
||||
restart_policy="Always",
|
||||
image_pull_secrets=[client.V1LocalObjectReference(name="acr-secret")]
|
||||
image_pull_secrets=image_pull_secrets
|
||||
)
|
||||
|
||||
# 构建标签(合并默认标签和用户自定义标签)
|
||||
@@ -785,7 +832,8 @@ class K8sManager:
|
||||
Returns:
|
||||
删除结果
|
||||
"""
|
||||
namespace_name = f"agent-{agent_name}".lower().strip('-')[:63]
|
||||
sanitized_name = sanitize_k8s_name(agent_name)
|
||||
namespace_name = f"agent-{sanitized_name}"[:63].rstrip('-')
|
||||
# region agent log
|
||||
try:
|
||||
import json, time
|
||||
|
||||
@@ -30,6 +30,10 @@ AZURE_CLIENT_SECRET="${AZURE_CLIENT_SECRET:-your-client-secret}"
|
||||
AZURE_SUBSCRIPTION_ID="${AZURE_SUBSCRIPTION_ID:-your-subscription-id}"
|
||||
AZURE_RESOURCE_GROUP="${AZURE_RESOURCE_GROUP:-your-resource-group}"
|
||||
|
||||
# Gitee 配置(需要替换为实际值)
|
||||
GITEE_TOKEN="${GITEE_TOKEN:-your-gitee-token}"
|
||||
GITEE_PASSWORD="${GITEE_PASSWORD:-your-gitee-password}"
|
||||
|
||||
echo -e "${BLUE}========================================${NC}"
|
||||
echo -e "${BLUE} Agent Manager K8s 部署 (ARM64)${NC}"
|
||||
echo -e "${BLUE}========================================${NC}"
|
||||
@@ -154,7 +158,18 @@ update_config() {
|
||||
fi
|
||||
fi
|
||||
|
||||
# 创建临时 secret 文件
|
||||
# 检查 Gitee 凭据
|
||||
if [ "$GITEE_TOKEN" = "your-gitee-token" ]; then
|
||||
print_warning "请设置 GITEE_TOKEN 环境变量(创建 Agent 仓库必需)"
|
||||
print_warning " export GITEE_TOKEN=your-actual-token"
|
||||
read -p "是否继续部署(不含 Gitee 仓库功能)?[y/N] " -n 1 -r
|
||||
echo
|
||||
if [[ ! $REPLY =~ ^[Yy]$ ]]; then
|
||||
exit 1
|
||||
fi
|
||||
fi
|
||||
|
||||
# 创建临时 secret 文件(包含 Azure 和 Gitee 凭据)
|
||||
cat > /tmp/agent-manager-secret.yaml <<EOF
|
||||
apiVersion: v1
|
||||
kind: Secret
|
||||
@@ -166,6 +181,8 @@ stringData:
|
||||
AZURE_TENANT_ID: "${AZURE_TENANT_ID}"
|
||||
AZURE_CLIENT_ID: "${AZURE_CLIENT_ID}"
|
||||
AZURE_CLIENT_SECRET: "${AZURE_CLIENT_SECRET}"
|
||||
GITEE_TOKEN: "${GITEE_TOKEN}"
|
||||
GITEE_PASSWORD: "${GITEE_PASSWORD}"
|
||||
EOF
|
||||
|
||||
kubectl apply -f ${K8S_DIR}/agent-manager-configmap.yaml
|
||||
|
||||
Executable
+107
@@ -0,0 +1,107 @@
|
||||
#!/bin/bash
|
||||
# 创建 Agent 测试是否正常工作,然后删除
|
||||
# 使用: AGENT_MANAGER_URL=http://localhost:8000 ./scripts/test_create_delete_agent.sh
|
||||
# 或: ./scripts/test_create_delete_agent.sh http://your-manager:8000
|
||||
|
||||
set -e
|
||||
BASE_URL="${1:-${AGENT_MANAGER_URL:-http://localhost:8000}}"
|
||||
AGENT_NAME="test-echo-$(date +%s)"
|
||||
|
||||
echo "=========================================="
|
||||
echo "Agent 创建/删除测试"
|
||||
echo "=========================================="
|
||||
echo "Manager URL: $BASE_URL"
|
||||
echo "Agent 名称: $AGENT_NAME"
|
||||
echo ""
|
||||
|
||||
# 1. 创建 Agent (echo_agent 无需额外 config)
|
||||
echo ">>> 1. 创建 Agent (template=echo_agent)..."
|
||||
CREATE_RESP=$(curl -s -w "\n%{http_code}" -X POST "${BASE_URL}/agents" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d "{
|
||||
\"name\": \"${AGENT_NAME}\",
|
||||
\"template\": \"echo_agent\",
|
||||
\"config\": {}
|
||||
}")
|
||||
HTTP_BODY=$(echo "$CREATE_RESP" | head -n -1)
|
||||
HTTP_CODE=$(echo "$CREATE_RESP" | tail -n 1)
|
||||
|
||||
if [ "$HTTP_CODE" != "200" ]; then
|
||||
echo "创建失败 HTTP $HTTP_CODE"
|
||||
echo "$HTTP_BODY" | python3 -m json.tool 2>/dev/null || echo "$HTTP_BODY"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
echo "创建成功"
|
||||
echo "$HTTP_BODY" | python3 -m json.tool 2>/dev/null || echo "$HTTP_BODY"
|
||||
echo ""
|
||||
|
||||
# 从响应中取 pod_ip 或 access_info
|
||||
POD_IP=$(echo "$HTTP_BODY" | python3 -c "import sys,json; d=json.load(sys.stdin); print(d.get('pod_ip','') or (d.get('access_info',{}) or {}).get('pod_url','').split('//')[-1].split(':')[0])" 2>/dev/null || true)
|
||||
EXTERNAL_IP=$(echo "$HTTP_BODY" | python3 -c "import sys,json; d=json.load(sys.stdin); a=d.get('access_info',{}); print(a.get('external_ip','') or a.get('ip_url','').split('//')[-1].split(':')[0] if isinstance(a,dict) else '')" 2>/dev/null || true)
|
||||
SERVICE_PORT=$(echo "$HTTP_BODY" | python3 -c "import sys,json; d=json.load(sys.stdin); print(d.get('service_port', 8000) or 8000)" 2>/dev/null || echo "8000")
|
||||
|
||||
# 2. 等待 Pod 就绪并测活
|
||||
echo ">>> 2. 等待 Pod 就绪并测试健康..."
|
||||
for i in 1 2 3 4 5 6 7 8 9 10; do
|
||||
STATUS_RESP=$(curl -s -w "\n%{http_code}" "${BASE_URL}/agents/${AGENT_NAME}/status")
|
||||
STATUS_BODY=$(echo "$STATUS_RESP" | head -n -1)
|
||||
STATUS_CODE=$(echo "$STATUS_RESP" | tail -n 1)
|
||||
if [ "$STATUS_CODE" != "200" ]; then
|
||||
sleep 3
|
||||
continue
|
||||
fi
|
||||
STATUS=$(echo "$STATUS_BODY" | python3 -c "import sys,json; d=json.load(sys.stdin); print(d.get('status',''))" 2>/dev/null || true)
|
||||
if [ "$STATUS" = "Running" ]; then
|
||||
break
|
||||
fi
|
||||
sleep 3
|
||||
done
|
||||
|
||||
if [ "$STATUS" != "Running" ]; then
|
||||
echo "Pod 未在预期内变为 Running,当前 status: $STATUS"
|
||||
echo "继续尝试访问 Agent 健康端点..."
|
||||
fi
|
||||
|
||||
# 尝试访问 Agent:优先外网 IP:80,否则 pod_ip:service_port
|
||||
AGENT_URL=""
|
||||
if [ -n "$EXTERNAL_IP" ]; then
|
||||
AGENT_URL="http://${EXTERNAL_IP}:80"
|
||||
elif [ -n "$POD_IP" ]; then
|
||||
AGENT_URL="http://${POD_IP}:${SERVICE_PORT}"
|
||||
fi
|
||||
|
||||
if [ -n "$AGENT_URL" ]; then
|
||||
echo "访问 Agent: $AGENT_URL/health"
|
||||
if curl -sf --connect-timeout 10 "${AGENT_URL}/health" > /dev/null; then
|
||||
echo "Agent 健康检查通过"
|
||||
else
|
||||
echo "健康检查失败(可能 LoadBalancer 未就绪或网络不可达)"
|
||||
fi
|
||||
# 尝试根路径
|
||||
if curl -sf --connect-timeout 5 "${AGENT_URL}/" > /dev/null; then
|
||||
echo "Agent 根路径可访问"
|
||||
fi
|
||||
else
|
||||
echo "未获取到 Pod IP 或外网 IP,跳过 Agent 端点测试"
|
||||
fi
|
||||
|
||||
echo ""
|
||||
echo ">>> 3. 删除 Agent..."
|
||||
DEL_RESP=$(curl -s -w "\n%{http_code}" -X DELETE "${BASE_URL}/agents/${AGENT_NAME}")
|
||||
DEL_BODY=$(echo "$DEL_RESP" | head -n -1)
|
||||
DEL_CODE=$(echo "$DEL_RESP" | tail -n 1)
|
||||
|
||||
if [ "$DEL_CODE" = "200" ] || [ "$DEL_CODE" = "204" ]; then
|
||||
echo "删除成功"
|
||||
echo "$DEL_BODY" | python3 -m json.tool 2>/dev/null || echo "$DEL_BODY"
|
||||
else
|
||||
echo "删除返回 HTTP $DEL_CODE"
|
||||
echo "$DEL_BODY"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
echo ""
|
||||
echo "=========================================="
|
||||
echo "测试完成: 创建 -> 检查状态 -> 删除 均成功"
|
||||
echo "=========================================="
|
||||
Executable
+148
@@ -0,0 +1,148 @@
|
||||
#!/bin/bash
|
||||
# 创建 Search Agent、测试健康与状态、可选调用 /search、然后删除
|
||||
# 使用: AGENT_MANAGER_URL=http://20.212.121.126 ./scripts/test_search_agent.sh
|
||||
# 可选环境变量: LLM_BASE_URL, SERPER_API_KEY, JINA_API_KEY (不设则用默认/镜像内建)
|
||||
|
||||
set -e
|
||||
BASE_URL="${1:-${AGENT_MANAGER_URL:-http://localhost:8000}}"
|
||||
AGENT_NAME="test-search-$(date +%s)"
|
||||
|
||||
# 默认 LLM 地址(与 code_ai_agent 等一致),可通过环境变量覆盖
|
||||
LLM_BASE_URL="${LLM_BASE_URL:-https://litellm.graystone-fb459c5d.southeastasia.azurecontainerapps.io/v1}"
|
||||
SERPER_API_KEY="${SERPER_API_KEY:-8253b4f240b520194065312f90e85f9be0fa205f}"
|
||||
JINA_API_KEY="${JINA_API_KEY:-jina_e26dc30420a44a1e859216528065b203TkMRmsoz-FgMDQC5FZX9jr5oF2CI}"
|
||||
|
||||
echo "=========================================="
|
||||
echo "Search Agent 测试"
|
||||
echo "=========================================="
|
||||
echo "Manager URL: $BASE_URL"
|
||||
echo "Agent 名称: $AGENT_NAME"
|
||||
echo "LLM_BASE_URL: $LLM_BASE_URL"
|
||||
echo ""
|
||||
|
||||
# 1. 创建 Search Agent(传入 env 以通过 K8s 注入)
|
||||
echo ">>> 1. 创建 Search Agent (template=search_agent)..."
|
||||
CREATE_RESP=$(curl -s -w "\n%{http_code}" -X POST "${BASE_URL}/agents" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d "{
|
||||
\"name\": \"${AGENT_NAME}\",
|
||||
\"template\": \"search_agent\",
|
||||
\"framework\": \"API\",
|
||||
\"config\": {
|
||||
\"env\": {
|
||||
\"LLM_BASE_URL\": \"${LLM_BASE_URL}\",
|
||||
\"SERPER_API_KEY\": \"${SERPER_API_KEY}\",
|
||||
\"JINA_API_KEY\": \"${JINA_API_KEY}\"
|
||||
}
|
||||
}
|
||||
}")
|
||||
HTTP_BODY=$(echo "$CREATE_RESP" | head -n -1)
|
||||
HTTP_CODE=$(echo "$CREATE_RESP" | tail -n 1)
|
||||
|
||||
if [ "$HTTP_CODE" != "200" ]; then
|
||||
echo "创建失败 HTTP $HTTP_CODE"
|
||||
echo "$HTTP_BODY" | python3 -m json.tool 2>/dev/null || echo "$HTTP_BODY"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
echo "创建成功"
|
||||
echo "$HTTP_BODY" | python3 -m json.tool 2>/dev/null || echo "$HTTP_BODY"
|
||||
echo ""
|
||||
|
||||
POD_IP=$(echo "$HTTP_BODY" | python3 -c "import sys,json; d=json.load(sys.stdin); print(d.get('pod_ip','') or (d.get('access_info',{}) or {}).get('pod_url','').split('//')[-1].split(':')[0])" 2>/dev/null || true)
|
||||
EXTERNAL_IP=$(echo "$HTTP_BODY" | python3 -c "import sys,json; d=json.load(sys.stdin); a=d.get('access_info',{}); print(a.get('external_ip','') or (a.get('ip_url') or '').split('//')[-1].split(':')[0] if isinstance(a,dict) else '')" 2>/dev/null || true)
|
||||
SERVICE_PORT=$(echo "$HTTP_BODY" | python3 -c "import sys,json; d=json.load(sys.stdin); print(d.get('service_port', 8080) or 8080)" 2>/dev/null || echo "8080")
|
||||
|
||||
# 2. 等待 Pod 就绪
|
||||
echo ">>> 2. 等待 Pod 就绪..."
|
||||
for i in 1 2 3 4 5 6 7 8 9 10 11 12; do
|
||||
STATUS_RESP=$(curl -s -w "\n%{http_code}" "${BASE_URL}/agents/${AGENT_NAME}/status")
|
||||
STATUS_BODY=$(echo "$STATUS_RESP" | head -n -1)
|
||||
STATUS_CODE=$(echo "$STATUS_RESP" | tail -n 1)
|
||||
if [ "$STATUS_CODE" != "200" ]; then
|
||||
sleep 5
|
||||
continue
|
||||
fi
|
||||
STATUS=$(echo "$STATUS_BODY" | python3 -c "import sys,json; d=json.load(sys.stdin); print(d.get('status',''))" 2>/dev/null || true)
|
||||
if [ "$STATUS" = "Running" ]; then
|
||||
echo "Pod 状态: Running"
|
||||
break
|
||||
fi
|
||||
echo " 等待中... status=$STATUS (${i}/12)"
|
||||
sleep 5
|
||||
done
|
||||
|
||||
if [ "$STATUS" != "Running" ]; then
|
||||
echo "Pod 未在预期内变为 Running,当前: $STATUS"
|
||||
fi
|
||||
|
||||
# 3. 测试 Agent 端点(Search Agent 端口 8080,经 Service 暴露为 80)
|
||||
AGENT_URL=""
|
||||
if [ -n "$EXTERNAL_IP" ]; then
|
||||
AGENT_URL="http://${EXTERNAL_IP}:80"
|
||||
elif [ -n "$POD_IP" ]; then
|
||||
AGENT_URL="http://${POD_IP}:${SERVICE_PORT}"
|
||||
fi
|
||||
|
||||
if [ -n "$AGENT_URL" ]; then
|
||||
echo ""
|
||||
echo ">>> 3. 测试 Search Agent 端点..."
|
||||
echo " Base URL: $AGENT_URL"
|
||||
|
||||
if curl -sf --connect-timeout 15 "${AGENT_URL}/health" > /tmp/search_health.json 2>/dev/null; then
|
||||
echo " GET /health: 成功"
|
||||
cat /tmp/search_health.json | python3 -m json.tool 2>/dev/null || cat /tmp/search_health.json
|
||||
else
|
||||
echo " GET /health: 失败或超时"
|
||||
fi
|
||||
|
||||
if curl -sf --connect-timeout 10 "${AGENT_URL}/status" > /tmp/search_status.json 2>/dev/null; then
|
||||
echo " GET /status: 成功"
|
||||
cat /tmp/search_status.json | python3 -m json.tool 2>/dev/null || cat /tmp/search_status.json
|
||||
else
|
||||
echo " GET /status: 失败或超时"
|
||||
fi
|
||||
|
||||
if curl -sf --connect-timeout 5 "${AGENT_URL}/" > /dev/null; then
|
||||
echo " GET /: 成功"
|
||||
fi
|
||||
|
||||
# 可选:调用 /search(需要有效 llm_api_key,否则可能 400/500)
|
||||
if [ -n "${LLM_API_KEY_FOR_TEST}" ]; then
|
||||
echo " 调用 POST /search (简短查询)..."
|
||||
SEARCH_RESP=$(curl -s -w "\n%{http_code}" -X POST "${AGENT_URL}/search" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d "{\"query\": \"What is 2+2?\", \"llm_api_key\": \"${LLM_API_KEY_FOR_TEST}\"}")
|
||||
SEARCH_CODE=$(echo "$SEARCH_RESP" | tail -n 1)
|
||||
if [ "$SEARCH_CODE" = "200" ]; then
|
||||
echo " POST /search: 成功 (HTTP 200)"
|
||||
else
|
||||
echo " POST /search: HTTP $SEARCH_CODE"
|
||||
fi
|
||||
else
|
||||
echo " (跳过 /search:设置 LLM_API_KEY_FOR_TEST 可测试搜索)"
|
||||
fi
|
||||
else
|
||||
echo "未获取到 Agent 地址,跳过端点测试"
|
||||
fi
|
||||
|
||||
# 4. 删除 Agent
|
||||
echo ""
|
||||
echo ">>> 4. 删除 Agent..."
|
||||
DEL_RESP=$(curl -s -w "\n%{http_code}" -X DELETE "${BASE_URL}/agents/${AGENT_NAME}")
|
||||
DEL_BODY=$(echo "$DEL_RESP" | head -n -1)
|
||||
DEL_CODE=$(echo "$DEL_RESP" | tail -n 1)
|
||||
|
||||
if [ "$DEL_CODE" = "200" ] || [ "$DEL_CODE" = "204" ]; then
|
||||
echo "删除成功"
|
||||
echo "$DEL_BODY" | python3 -m json.tool 2>/dev/null || echo "$DEL_BODY"
|
||||
else
|
||||
echo "删除返回 HTTP $DEL_CODE"
|
||||
echo "$DEL_BODY"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
echo ""
|
||||
echo "=========================================="
|
||||
echo "Search Agent 测试完成"
|
||||
echo "=========================================="
|
||||
Executable
+116
@@ -0,0 +1,116 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
更新数据库中模板的端口。
|
||||
|
||||
将 search_agent / search_agent_a2a / search_agent_mcp 的 port 改为 8080,
|
||||
与容器实际监听端口一致,使 Service target_port 正确。
|
||||
|
||||
用法:
|
||||
# 使用项目 database 配置,仅修改 search_agent* 为 8080
|
||||
python scripts/update_template_ports.py
|
||||
|
||||
# 指定要改的模板和端口
|
||||
python scripts/update_template_ports.py --names search_agent,search_agent_a2a,search_agent_mcp --port 8080
|
||||
|
||||
# 仅打印当前端口,不修改(dry-run)
|
||||
python scripts/update_template_ports.py --dry-run
|
||||
"""
|
||||
|
||||
import os
|
||||
import sys
|
||||
import argparse
|
||||
|
||||
# 确保项目根在 path 中
|
||||
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
|
||||
from database import SessionLocal, Template
|
||||
|
||||
|
||||
# 默认:需要改为 8080 的模板(与 search_agent 镜像一致)
|
||||
DEFAULT_TEMPLATES_8080 = ["search_agent", "search_agent_a2a", "search_agent_mcp"]
|
||||
|
||||
|
||||
def get_current_ports(db, names):
|
||||
"""返回 {name: port}"""
|
||||
rows = db.query(Template).filter(Template.name.in_(names)).all()
|
||||
return {r.name: r.port for r in rows}
|
||||
|
||||
|
||||
def update_ports(names: list, port: int, dry_run: bool = False):
|
||||
db = SessionLocal()
|
||||
try:
|
||||
current = get_current_ports(db, names)
|
||||
missing = [n for n in names if n not in current]
|
||||
if missing:
|
||||
print(f"未找到模板: {missing}")
|
||||
names = [n for n in names if n in current]
|
||||
if not names:
|
||||
return False
|
||||
|
||||
print("当前端口:")
|
||||
for n in names:
|
||||
print(f" {n}: {current.get(n)}")
|
||||
|
||||
if dry_run:
|
||||
print("\n[DRY-RUN] 未执行修改。去掉 --dry-run 将执行更新。")
|
||||
return True
|
||||
|
||||
updated = 0
|
||||
for name in names:
|
||||
row = db.query(Template).filter(Template.name == name).first()
|
||||
if row is not None and row.port != port:
|
||||
row.port = port
|
||||
updated += 1
|
||||
print(f" 更新 {name} -> port={port}")
|
||||
|
||||
if updated:
|
||||
db.commit()
|
||||
print(f"\n已提交: {updated} 条记录 port 已改为 {port}")
|
||||
else:
|
||||
print("\n无需更新(端口已是目标值)")
|
||||
|
||||
# 再次查询确认
|
||||
after = get_current_ports(db, names)
|
||||
print("更新后端口:")
|
||||
for n in names:
|
||||
print(f" {n}: {after.get(n)}")
|
||||
return True
|
||||
except Exception as e:
|
||||
db.rollback()
|
||||
print(f"错误: {e}", file=sys.stderr)
|
||||
return False
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="更新模板端口")
|
||||
parser.add_argument(
|
||||
"--names",
|
||||
type=str,
|
||||
default=",".join(DEFAULT_TEMPLATES_8080),
|
||||
help="模板名称,逗号分隔,默认: search_agent,search_agent_a2a,search_agent_mcp",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--port",
|
||||
type=int,
|
||||
default=8080,
|
||||
help="目标端口,默认 8080",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--dry-run",
|
||||
action="store_true",
|
||||
help="仅打印当前端口,不修改",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
names = [n.strip() for n in args.names.split(",") if n.strip()]
|
||||
if not names:
|
||||
print("请至少指定一个模板名 (--names)")
|
||||
sys.exit(1)
|
||||
|
||||
ok = update_ports(names, args.port, dry_run=args.dry_run)
|
||||
sys.exit(0 if ok else 1)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
+8
-8
@@ -28,7 +28,7 @@ DEFAULT_TEMPLATES = {
|
||||
"display_name": "Search Agent",
|
||||
"description": "搜索 Agent (LangChain)",
|
||||
"image": "agnettaiji.azurecr.io/ai-agents/search-agent:latest",
|
||||
"port": 8000,
|
||||
"port": 8080,
|
||||
"agent_framework": "langchain",
|
||||
"env_requirements": {},
|
||||
},
|
||||
@@ -36,7 +36,7 @@ DEFAULT_TEMPLATES = {
|
||||
"display_name": "Search Agent A2A",
|
||||
"description": "搜索 Agent (A2A 协议)",
|
||||
"image": "agnettaiji.azurecr.io/ai-agents/search-agent-a2a:latest",
|
||||
"port": 8000,
|
||||
"port": 8080,
|
||||
"agent_framework": "a2a",
|
||||
"env_requirements": {},
|
||||
},
|
||||
@@ -44,7 +44,7 @@ DEFAULT_TEMPLATES = {
|
||||
"display_name": "Search Agent MCP",
|
||||
"description": "搜索 Agent (MCP 协议)",
|
||||
"image": "agnettaiji.azurecr.io/ai-agents/search-agent-mcp:latest",
|
||||
"port": 8000,
|
||||
"port": 8080,
|
||||
"agent_framework": "mcp",
|
||||
"env_requirements": {},
|
||||
},
|
||||
@@ -82,7 +82,7 @@ DEFAULT_TEMPLATES = {
|
||||
"display_name": "Jina Search Agent",
|
||||
"description": "Jina AI 搜索 Agent",
|
||||
"image": "agnettaiji.azurecr.io/ai-agents/jina-search-agent:latest",
|
||||
"port": 8080,
|
||||
"port": 8000,
|
||||
"agent_framework": "api",
|
||||
"env_requirements": {
|
||||
"required": {
|
||||
@@ -94,7 +94,7 @@ DEFAULT_TEMPLATES = {
|
||||
"display_name": "Azure Blob Agent",
|
||||
"description": "Azure Blob 存储 Agent",
|
||||
"image": "agnettaiji.azurecr.io/ai-agents/azure-blob-agent:latest",
|
||||
"port": 8080,
|
||||
"port": 8000,
|
||||
"agent_framework": "api",
|
||||
"env_requirements": {
|
||||
"required": {
|
||||
@@ -106,7 +106,7 @@ DEFAULT_TEMPLATES = {
|
||||
"display_name": "Azure Blob Agent MCP",
|
||||
"description": "Azure Blob 存储 Agent (MCP 协议)",
|
||||
"image": "agnettaiji.azurecr.io/ai-agents/azure-blob-agent-mcp:latest",
|
||||
"port": 8080,
|
||||
"port": 8000,
|
||||
"agent_framework": "mcp",
|
||||
"env_requirements": {
|
||||
"required": {
|
||||
@@ -118,7 +118,7 @@ DEFAULT_TEMPLATES = {
|
||||
"display_name": "Azure Blob Agent A2A",
|
||||
"description": "Azure Blob 存储 Agent (A2A 协议)",
|
||||
"image": "agnettaiji.azurecr.io/ai-agents/azure-blob-agent-a2a:latest",
|
||||
"port": 8080,
|
||||
"port": 8000,
|
||||
"agent_framework": "a2a",
|
||||
"env_requirements": {
|
||||
"required": {
|
||||
@@ -130,7 +130,7 @@ DEFAULT_TEMPLATES = {
|
||||
"display_name": "A2A LiteLLM Agent",
|
||||
"description": "LiteLLM A2A 协议 Agent",
|
||||
"image": "agnettaiji.azurecr.io/ai-agents/a2a-litellm-agent:latest",
|
||||
"port": 8080,
|
||||
"port": 8000,
|
||||
"agent_framework": "a2a",
|
||||
"env_requirements": {},
|
||||
},
|
||||
|
||||
@@ -0,0 +1,2 @@
|
||||
# CI/CD Test
|
||||
Written by code_ai_agent.
|
||||
@@ -0,0 +1,242 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
测试美股 Agents 功能
|
||||
"""
|
||||
import asyncio
|
||||
import aiohttp
|
||||
import os
|
||||
import sys
|
||||
|
||||
# LLM 配置
|
||||
LLM_BASE_URL = "https://litellm.graystone-fb459c5d.southeastasia.azurecontainerapps.io/v1"
|
||||
LLM_API_KEY = "sk-mPV5MVVVVvfGSkXA-ASQXQ"
|
||||
LLM_MODEL = "taiji/gpt-4o-mini"
|
||||
|
||||
async def fetch_stock_quote(symbol: str):
|
||||
"""获取股票行情数据 (Stock Quote Agent)"""
|
||||
url = f"https://query1.finance.yahoo.com/v8/finance/chart/{symbol}"
|
||||
params = {"interval": "1d", "range": "1d"}
|
||||
headers = {"User-Agent": "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36"}
|
||||
|
||||
try:
|
||||
async with aiohttp.ClientSession() as session:
|
||||
async with session.get(url, params=params, headers=headers, timeout=aiohttp.ClientTimeout(total=15)) as response:
|
||||
if response.status == 200:
|
||||
data = await response.json()
|
||||
result = data.get("chart", {}).get("result", [])
|
||||
|
||||
if not result:
|
||||
return {"success": False, "error": f"未找到股票: {symbol}"}
|
||||
|
||||
quote_data = result[0]
|
||||
meta = quote_data.get("meta", {})
|
||||
|
||||
current_price = meta.get("regularMarketPrice", 0)
|
||||
previous_close = meta.get("previousClose", 0)
|
||||
change = current_price - previous_close if previous_close else 0
|
||||
change_percent = (change / previous_close * 100) if previous_close else 0
|
||||
|
||||
return {
|
||||
"success": True,
|
||||
"symbol": symbol.upper(),
|
||||
"name": meta.get("shortName", symbol),
|
||||
"price": current_price,
|
||||
"change": round(change, 2),
|
||||
"change_percent": round(change_percent, 2),
|
||||
"high_52week": meta.get("fiftyTwoWeekHigh"),
|
||||
"low_52week": meta.get("fiftyTwoWeekLow"),
|
||||
"market_cap": meta.get("marketCap"),
|
||||
}
|
||||
else:
|
||||
return {"success": False, "error": f"API 请求失败: HTTP {response.status}"}
|
||||
except Exception as e:
|
||||
return {"success": False, "error": str(e)}
|
||||
|
||||
|
||||
async def fetch_stock_news(symbol: str):
|
||||
"""获取股票新闻 (Stock News Agent)"""
|
||||
url = f"https://query1.finance.yahoo.com/v1/finance/search"
|
||||
params = {"q": symbol, "newsCount": 5}
|
||||
headers = {"User-Agent": "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36"}
|
||||
|
||||
try:
|
||||
async with aiohttp.ClientSession() as session:
|
||||
async with session.get(url, params=params, headers=headers, timeout=aiohttp.ClientTimeout(total=15)) as response:
|
||||
if response.status == 200:
|
||||
data = await response.json()
|
||||
news = data.get("news", [])
|
||||
|
||||
return {
|
||||
"success": True,
|
||||
"symbol": symbol.upper(),
|
||||
"news_count": len(news),
|
||||
"news": [{"title": n.get("title", ""), "publisher": n.get("publisher", "")} for n in news[:5]]
|
||||
}
|
||||
else:
|
||||
return {"success": False, "error": f"API 请求失败: HTTP {response.status}"}
|
||||
except Exception as e:
|
||||
return {"success": False, "error": str(e)}
|
||||
|
||||
|
||||
async def fetch_historical_data(symbol: str, period: str = "1mo"):
|
||||
"""获取历史数据用于技术分析 (Stock Analysis Agent)"""
|
||||
url = f"https://query1.finance.yahoo.com/v8/finance/chart/{symbol}"
|
||||
params = {"interval": "1d", "range": period}
|
||||
headers = {"User-Agent": "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36"}
|
||||
|
||||
try:
|
||||
async with aiohttp.ClientSession() as session:
|
||||
async with session.get(url, params=params, headers=headers, timeout=aiohttp.ClientTimeout(total=15)) as response:
|
||||
if response.status == 200:
|
||||
data = await response.json()
|
||||
result = data.get("chart", {}).get("result", [])
|
||||
|
||||
if not result:
|
||||
return {"success": False, "error": f"未找到股票: {symbol}"}
|
||||
|
||||
quote_data = result[0]
|
||||
indicators = quote_data.get("indicators", {}).get("quote", [{}])[0]
|
||||
timestamps = quote_data.get("timestamp", [])
|
||||
|
||||
closes = indicators.get("close", [])
|
||||
closes = [c for c in closes if c is not None]
|
||||
|
||||
if len(closes) >= 5:
|
||||
# 简单技术分析
|
||||
sma_5 = sum(closes[-5:]) / 5
|
||||
current = closes[-1]
|
||||
trend = "看涨" if current > sma_5 else "看跌"
|
||||
|
||||
return {
|
||||
"success": True,
|
||||
"symbol": symbol.upper(),
|
||||
"data_points": len(closes),
|
||||
"current_price": round(current, 2),
|
||||
"sma_5": round(sma_5, 2),
|
||||
"trend": trend,
|
||||
"high": round(max(closes), 2),
|
||||
"low": round(min(closes), 2),
|
||||
}
|
||||
else:
|
||||
return {"success": False, "error": "数据点不足"}
|
||||
else:
|
||||
return {"success": False, "error": f"API 请求失败: HTTP {response.status}"}
|
||||
except Exception as e:
|
||||
return {"success": False, "error": str(e)}
|
||||
|
||||
|
||||
async def test_with_llm(symbol: str, context: str):
|
||||
"""使用 LLM 生成分析报告"""
|
||||
try:
|
||||
async with aiohttp.ClientSession() as session:
|
||||
payload = {
|
||||
"model": LLM_MODEL,
|
||||
"messages": [
|
||||
{"role": "system", "content": "你是一个专业的美股分析师,请根据提供的数据给出简洁的分析。"},
|
||||
{"role": "user", "content": f"请分析以下 {symbol} 股票数据并给出简要建议(50字以内):\n{context}"}
|
||||
],
|
||||
"max_tokens": 200
|
||||
}
|
||||
headers = {
|
||||
"Authorization": f"Bearer {LLM_API_KEY}",
|
||||
"Content-Type": "application/json"
|
||||
}
|
||||
|
||||
async with session.post(
|
||||
f"{LLM_BASE_URL}/chat/completions",
|
||||
json=payload,
|
||||
headers=headers,
|
||||
timeout=aiohttp.ClientTimeout(total=30)
|
||||
) as response:
|
||||
if response.status == 200:
|
||||
data = await response.json()
|
||||
content = data.get("choices", [{}])[0].get("message", {}).get("content", "")
|
||||
return {"success": True, "analysis": content}
|
||||
else:
|
||||
error_text = await response.text()
|
||||
return {"success": False, "error": f"LLM 请求失败: {response.status} - {error_text[:100]}"}
|
||||
except Exception as e:
|
||||
return {"success": False, "error": str(e)}
|
||||
|
||||
|
||||
async def main():
|
||||
print("=" * 70)
|
||||
print("美股 AI Agents 本地测试")
|
||||
print("=" * 70)
|
||||
print(f"LLM: {LLM_MODEL}")
|
||||
print(f"API: {LLM_BASE_URL}")
|
||||
print("=" * 70)
|
||||
|
||||
symbols = ["AAPL", "TSLA", "NVDA", "MSFT", "GOOGL"]
|
||||
|
||||
# 测试 1: Stock Quote Agent
|
||||
print("\n📈 【测试 1: Stock Quote Agent - 实时行情】")
|
||||
print("-" * 70)
|
||||
for symbol in symbols:
|
||||
result = await fetch_stock_quote(symbol)
|
||||
if result.get("success"):
|
||||
print(f"✅ {result['symbol']:6} | {result['name'][:25]:25} | ${result['price']:>10.2f} | {result['change_percent']:+6.2f}%")
|
||||
else:
|
||||
print(f"❌ {symbol}: {result.get('error')}")
|
||||
|
||||
# 测试 2: Stock News Agent
|
||||
print("\n📰 【测试 2: Stock News Agent - 新闻资讯】")
|
||||
print("-" * 70)
|
||||
for symbol in ["AAPL", "TSLA"]:
|
||||
result = await fetch_stock_news(symbol)
|
||||
if result.get("success"):
|
||||
print(f"✅ {result['symbol']} - 找到 {result['news_count']} 条新闻:")
|
||||
for news in result['news'][:2]:
|
||||
print(f" • {news['title'][:60]}...")
|
||||
else:
|
||||
print(f"❌ {symbol}: {result.get('error')}")
|
||||
|
||||
# 测试 3: Stock Analysis Agent
|
||||
print("\n📊 【测试 3: Stock Analysis Agent - 技术分析】")
|
||||
print("-" * 70)
|
||||
for symbol in ["AAPL", "NVDA", "TSLA"]:
|
||||
result = await fetch_historical_data(symbol)
|
||||
if result.get("success"):
|
||||
print(f"✅ {result['symbol']:6} | 当前: ${result['current_price']:>8.2f} | SMA5: ${result['sma_5']:>8.2f} | 趋势: {result['trend']} | 区间: ${result['low']:.2f}-${result['high']:.2f}")
|
||||
else:
|
||||
print(f"❌ {symbol}: {result.get('error')}")
|
||||
|
||||
# 测试 4: LLM 综合分析
|
||||
print("\n🤖 【测试 4: LLM 综合分析】")
|
||||
print("-" * 70)
|
||||
|
||||
# 获取一只股票的完整数据
|
||||
symbol = "AAPL"
|
||||
quote = await fetch_stock_quote(symbol)
|
||||
analysis = await fetch_historical_data(symbol)
|
||||
|
||||
if quote.get("success") and analysis.get("success"):
|
||||
context = f"""
|
||||
股票: {symbol} ({quote['name']})
|
||||
当前价格: ${quote['price']:.2f}
|
||||
涨跌幅: {quote['change_percent']:+.2f}%
|
||||
52周范围: ${quote.get('low_52week', 0):.2f} - ${quote.get('high_52week', 0):.2f}
|
||||
5日均线: ${analysis['sma_5']:.2f}
|
||||
技术趋势: {analysis['trend']}
|
||||
近期区间: ${analysis['low']:.2f} - ${analysis['high']:.2f}
|
||||
"""
|
||||
print(f"📋 {symbol} 数据汇总:")
|
||||
print(context)
|
||||
|
||||
print("🔄 调用 LLM 生成分析报告...")
|
||||
llm_result = await test_with_llm(symbol, context)
|
||||
if llm_result.get("success"):
|
||||
print(f"\n💡 AI 分析建议:")
|
||||
print(f" {llm_result['analysis']}")
|
||||
else:
|
||||
print(f"❌ LLM 调用失败: {llm_result.get('error')}")
|
||||
else:
|
||||
print(f"❌ 获取数据失败")
|
||||
|
||||
print("\n" + "=" * 70)
|
||||
print("✅ 测试完成!")
|
||||
print("=" * 70)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
asyncio.run(main())
|
||||
+11
-10
@@ -12,6 +12,7 @@ from datetime import datetime
|
||||
from typing import Dict, List, Optional, Any
|
||||
from fastapi import APIRouter, HTTPException, BackgroundTasks
|
||||
from pydantic import BaseModel, Field
|
||||
from k8s_manager import sanitize_k8s_name
|
||||
|
||||
from gitee_manager import gitee_manager
|
||||
from agent_code_generator import agent_code_generator
|
||||
@@ -185,9 +186,9 @@ async def generate_agent(request: GenerateAgentRequest, background_tasks: Backgr
|
||||
"timeout": tool.timeout
|
||||
})
|
||||
|
||||
# 生成完整项目文件(参考 http://gitee.ath.cx:3000/xiaohei/cicd-AKS)
|
||||
# 生成完整项目文件(k8s_name 必须合规,否则 Service 创建会报 DNS-1035)
|
||||
project_files = agent_code_generator.generate_full_project(
|
||||
agent_name=request.agent_name,
|
||||
agent_name=k8s_name_for_cicd,
|
||||
description=request.description,
|
||||
tools_config=tools_config,
|
||||
auto_deploy=request.auto_deploy
|
||||
@@ -225,10 +226,10 @@ async def generate_agent(request: GenerateAgentRequest, background_tasks: Backgr
|
||||
"ACR_LOGIN_SERVER": "agnettaiji.azurecr.io",
|
||||
"ACR_USERNAME": "agnettaiji",
|
||||
"ACR_PASSWORD": "hDpX5t34N5ZmnKdtqyjYL5co/SnXJrmD20CRpGpWaG+ACRCw2wGM",
|
||||
"AZ_CLIENT_ID": "fb306798-2cfe-4ac9-ba48-eab7bc71bcfe",
|
||||
"AZ_CLIENT_SECRET": "cVK8Q~xlfBwm2_t2TC24yrTukWV4F3G~eIjBBa0D",
|
||||
"AZ_CLIENT_ID": "f2dd1cb2-02f6-4efb-bc72-d148f6e01545",
|
||||
"AZ_CLIENT_SECRET": "UVU8Q~Hcrf5KeLi2RvUXB2rcuKFEjRCCrf_JrbwA",
|
||||
"AZ_TENANT_ID": "263c3ff6-1be5-4141-8308-b188464fb297",
|
||||
"AZ_SUBSCRIPTION_ID": "c6c47e4c-f5f4-49f8-b26f-7728862c17d6",
|
||||
"AZ_SUBSCRIPTION_ID": "45d7a360-af09-40fc-9afc-56dc475245ec",
|
||||
"AZ_RG": "taiji-ai-pda",
|
||||
"AZ_AKS": "taiji-ai-pda",
|
||||
"AZURE_DNS_ZONE": "taijiagnet.com"
|
||||
@@ -261,9 +262,9 @@ async def generate_agent(request: GenerateAgentRequest, background_tasks: Backgr
|
||||
"created_at": datetime.utcnow().isoformat(),
|
||||
"status": "building",
|
||||
"auto_deploy": request.auto_deploy,
|
||||
"image_name": f"agnettaiji.azurecr.io/ai-agents/{repo_name}:latest",
|
||||
"expected_domain": f"{k8s_name}.taijiagnet.com",
|
||||
"expected_namespace": f"agent-{k8s_name}"
|
||||
"image_name": f"agnettaiji.azurecr.io/ai-agents/{k8s_name_for_cicd}:latest",
|
||||
"expected_domain": f"{k8s_name_for_cicd}.taijiagnet.com",
|
||||
"expected_namespace": f"agent-{k8s_name_for_cicd}"
|
||||
}
|
||||
|
||||
logger.info(f"✅ Agent 项目创建成功: {repo_name}")
|
||||
@@ -278,12 +279,12 @@ async def generate_agent(request: GenerateAgentRequest, background_tasks: Backgr
|
||||
"agent_ref_id": agent_ref_id,
|
||||
"repo_name": repo_name,
|
||||
"repo_url": repo_result.get("html_url"),
|
||||
"image_name": f"agnettaiji.azurecr.io/ai-agents/{repo_name}:latest",
|
||||
"image_name": f"agnettaiji.azurecr.io/ai-agents/{k8s_name_for_cicd}:latest",
|
||||
"status": "building",
|
||||
"files_pushed": len(project_files),
|
||||
"tools_count": len(request.tools),
|
||||
"expected_domain": expected_domain,
|
||||
"expected_namespace": f"agent-{k8s_name}"
|
||||
"expected_namespace": f"agent-{k8s_name_for_cicd}"
|
||||
},
|
||||
"message": f"Agent 项目已创建并推送到 Gitee,CI/CD 正在构建中。部署后访问: http://{expected_domain}"
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user