lightpdf-aipdf-backend 0.1.1__tar.gz
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- lightpdf_aipdf_backend-0.1.1/.env +7 -0
- lightpdf_aipdf_backend-0.1.1/.gitignore +20 -0
- lightpdf_aipdf_backend-0.1.1/PKG-INFO +87 -0
- lightpdf_aipdf_backend-0.1.1/README.md +68 -0
- lightpdf_aipdf_backend-0.1.1/publish.sh +160 -0
- lightpdf_aipdf_backend-0.1.1/pyproject.toml +40 -0
- lightpdf_aipdf_backend-0.1.1/run.py +10 -0
- lightpdf_aipdf_backend-0.1.1/setup.py +7 -0
- lightpdf_aipdf_backend-0.1.1/src/lightpdf_aipdf_backend/__init__.py +3 -0
- lightpdf_aipdf_backend-0.1.1/src/lightpdf_aipdf_backend/__main__.py +8 -0
- lightpdf_aipdf_backend-0.1.1/src/lightpdf_aipdf_backend/app.py +58 -0
- lightpdf_aipdf_backend-0.1.1/src/lightpdf_aipdf_backend/chat_handler.py +298 -0
- lightpdf_aipdf_backend-0.1.1/src/lightpdf_aipdf_backend/cli.py +9 -0
- lightpdf_aipdf_backend-0.1.1/src/lightpdf_aipdf_backend/config.py +22 -0
- lightpdf_aipdf_backend-0.1.1/src/lightpdf_aipdf_backend/file_handler.py +81 -0
- lightpdf_aipdf_backend-0.1.1/src/lightpdf_aipdf_backend/main.py +45 -0
- lightpdf_aipdf_backend-0.1.1/src/lightpdf_aipdf_backend/models.py +22 -0
- lightpdf_aipdf_backend-0.1.1/src/lightpdf_aipdf_backend/state.py +53 -0
- lightpdf_aipdf_backend-0.1.1/src/lightpdf_aipdf_backend/tools.py +96 -0
- lightpdf_aipdf_backend-0.1.1/src/lightpdf_aipdf_backend/utils.py +52 -0
|
@@ -0,0 +1,7 @@
|
|
|
1
|
+
# OPENAI_API_KEY=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
|
|
2
|
+
# OPENAI_BASE_URL=https://openai-api.aoscdn.com/v1
|
|
3
|
+
# OPENAI_MODEL=gpt-4o-mini
|
|
4
|
+
OPENAI_API_KEY=sk-04c4fd40cbc34de5a7858cb723b89067
|
|
5
|
+
OPENAI_BASE_URL=https://dashscope.aliyuncs.com/compatible-mode/v1
|
|
6
|
+
OPENAI_MODEL=qwen-turbo
|
|
7
|
+
API_KEY=wx6ce6ckkey17wany
|
|
@@ -0,0 +1,20 @@
|
|
|
1
|
+
# Python-generated files
|
|
2
|
+
__pycache__/
|
|
3
|
+
*.py[oc]
|
|
4
|
+
build/
|
|
5
|
+
dist/
|
|
6
|
+
wheels/
|
|
7
|
+
*.egg-info
|
|
8
|
+
|
|
9
|
+
# Virtual environments
|
|
10
|
+
.venv/
|
|
11
|
+
node_modules/
|
|
12
|
+
|
|
13
|
+
# macOS system files
|
|
14
|
+
.DS_Store
|
|
15
|
+
|
|
16
|
+
mcp_server/demo/
|
|
17
|
+
mcp_server/docs/
|
|
18
|
+
homepage/backend/uploads
|
|
19
|
+
.version
|
|
20
|
+
lightpdf-mcp.zip
|
|
@@ -0,0 +1,87 @@
|
|
|
1
|
+
Metadata-Version: 2.4
|
|
2
|
+
Name: lightpdf-aipdf-backend
|
|
3
|
+
Version: 0.1.1
|
|
4
|
+
Summary: Backend API for LightPDF AI-PDF
|
|
5
|
+
License: Proprietary
|
|
6
|
+
Requires-Python: >=3.9
|
|
7
|
+
Requires-Dist: fastapi
|
|
8
|
+
Requires-Dist: mcp[cli]
|
|
9
|
+
Requires-Dist: openai
|
|
10
|
+
Requires-Dist: pydantic
|
|
11
|
+
Requires-Dist: python-dotenv
|
|
12
|
+
Requires-Dist: python-multipart
|
|
13
|
+
Requires-Dist: uvicorn
|
|
14
|
+
Provides-Extra: dev
|
|
15
|
+
Requires-Dist: black; extra == 'dev'
|
|
16
|
+
Requires-Dist: isort; extra == 'dev'
|
|
17
|
+
Requires-Dist: pytest; extra == 'dev'
|
|
18
|
+
Description-Content-Type: text/markdown
|
|
19
|
+
|
|
20
|
+
# LightPDF AI-PDF Backend
|
|
21
|
+
|
|
22
|
+
这是LightPDF AI-PDF的后端API服务。
|
|
23
|
+
|
|
24
|
+
## 安装
|
|
25
|
+
|
|
26
|
+
使用uv安装:
|
|
27
|
+
|
|
28
|
+
```bash
|
|
29
|
+
uv pip install -e .
|
|
30
|
+
```
|
|
31
|
+
|
|
32
|
+
或者构建并安装:
|
|
33
|
+
|
|
34
|
+
```bash
|
|
35
|
+
uv build
|
|
36
|
+
uv pip install dist/*.whl
|
|
37
|
+
```
|
|
38
|
+
|
|
39
|
+
## 配置
|
|
40
|
+
|
|
41
|
+
有两种配置方式:
|
|
42
|
+
|
|
43
|
+
### 1. 使用.env文件(开发环境)
|
|
44
|
+
|
|
45
|
+
创建一个`.env`文件,包含以下环境变量:
|
|
46
|
+
|
|
47
|
+
```
|
|
48
|
+
OPENAI_API_KEY=your_openai_api_key
|
|
49
|
+
OPENAI_MODEL=gpt-4o-mini
|
|
50
|
+
OPENAI_BASE_URL=https://api.openai.com/v1
|
|
51
|
+
# 其他必要的环境变量
|
|
52
|
+
```
|
|
53
|
+
|
|
54
|
+
注意:`.env`文件不会包含在构建包中,需要在部署环境中手动创建。
|
|
55
|
+
|
|
56
|
+
### 2. 使用环境变量(生产环境)
|
|
57
|
+
|
|
58
|
+
直接设置环境变量:
|
|
59
|
+
|
|
60
|
+
```bash
|
|
61
|
+
export OPENAI_API_KEY=your_openai_api_key
|
|
62
|
+
export OPENAI_MODEL=gpt-4o-mini
|
|
63
|
+
export OPENAI_BASE_URL=https://api.openai.com/v1
|
|
64
|
+
# 设置其他必要的环境变量
|
|
65
|
+
```
|
|
66
|
+
|
|
67
|
+
## 使用方法
|
|
68
|
+
|
|
69
|
+
安装后,可以通过以下命令启动服务器:
|
|
70
|
+
|
|
71
|
+
```bash
|
|
72
|
+
lightpdf-aipdf-server
|
|
73
|
+
```
|
|
74
|
+
|
|
75
|
+
也可以作为模块运行:
|
|
76
|
+
|
|
77
|
+
```bash
|
|
78
|
+
python -m lightpdf_aipdf_backend
|
|
79
|
+
```
|
|
80
|
+
|
|
81
|
+
## 开发
|
|
82
|
+
|
|
83
|
+
安装开发依赖:
|
|
84
|
+
|
|
85
|
+
```bash
|
|
86
|
+
uv pip install -e ".[dev]"
|
|
87
|
+
```
|
|
@@ -0,0 +1,68 @@
|
|
|
1
|
+
# LightPDF AI-PDF Backend
|
|
2
|
+
|
|
3
|
+
这是LightPDF AI-PDF的后端API服务。
|
|
4
|
+
|
|
5
|
+
## 安装
|
|
6
|
+
|
|
7
|
+
使用uv安装:
|
|
8
|
+
|
|
9
|
+
```bash
|
|
10
|
+
uv pip install -e .
|
|
11
|
+
```
|
|
12
|
+
|
|
13
|
+
或者构建并安装:
|
|
14
|
+
|
|
15
|
+
```bash
|
|
16
|
+
uv build
|
|
17
|
+
uv pip install dist/*.whl
|
|
18
|
+
```
|
|
19
|
+
|
|
20
|
+
## 配置
|
|
21
|
+
|
|
22
|
+
有两种配置方式:
|
|
23
|
+
|
|
24
|
+
### 1. 使用.env文件(开发环境)
|
|
25
|
+
|
|
26
|
+
创建一个`.env`文件,包含以下环境变量:
|
|
27
|
+
|
|
28
|
+
```
|
|
29
|
+
OPENAI_API_KEY=your_openai_api_key
|
|
30
|
+
OPENAI_MODEL=gpt-4o-mini
|
|
31
|
+
OPENAI_BASE_URL=https://api.openai.com/v1
|
|
32
|
+
# 其他必要的环境变量
|
|
33
|
+
```
|
|
34
|
+
|
|
35
|
+
注意:`.env`文件不会包含在构建包中,需要在部署环境中手动创建。
|
|
36
|
+
|
|
37
|
+
### 2. 使用环境变量(生产环境)
|
|
38
|
+
|
|
39
|
+
直接设置环境变量:
|
|
40
|
+
|
|
41
|
+
```bash
|
|
42
|
+
export OPENAI_API_KEY=your_openai_api_key
|
|
43
|
+
export OPENAI_MODEL=gpt-4o-mini
|
|
44
|
+
export OPENAI_BASE_URL=https://api.openai.com/v1
|
|
45
|
+
# 设置其他必要的环境变量
|
|
46
|
+
```
|
|
47
|
+
|
|
48
|
+
## 使用方法
|
|
49
|
+
|
|
50
|
+
安装后,可以通过以下命令启动服务器:
|
|
51
|
+
|
|
52
|
+
```bash
|
|
53
|
+
lightpdf-aipdf-server
|
|
54
|
+
```
|
|
55
|
+
|
|
56
|
+
也可以作为模块运行:
|
|
57
|
+
|
|
58
|
+
```bash
|
|
59
|
+
python -m lightpdf_aipdf_backend
|
|
60
|
+
```
|
|
61
|
+
|
|
62
|
+
## 开发
|
|
63
|
+
|
|
64
|
+
安装开发依赖:
|
|
65
|
+
|
|
66
|
+
```bash
|
|
67
|
+
uv pip install -e ".[dev]"
|
|
68
|
+
```
|
|
@@ -0,0 +1,160 @@
|
|
|
1
|
+
#!/bin/bash
|
|
2
|
+
|
|
3
|
+
# 显示帮助信息
|
|
4
|
+
show_help() {
|
|
5
|
+
echo "用法: $0 [-p|--publish] [版本号]"
|
|
6
|
+
echo
|
|
7
|
+
echo "参数:"
|
|
8
|
+
echo " -p, --publish 执行发布操作(同时会更新版本号)"
|
|
9
|
+
echo " 版本号 指定新的版本号(可选,格式:x.y.z,仅在发布时有效)"
|
|
10
|
+
echo
|
|
11
|
+
echo "示例:"
|
|
12
|
+
echo " $0 # 仅构建,使用版本号0.0.1"
|
|
13
|
+
echo " $0 1.0.0 # 仅构建,使用版本号0.0.1"
|
|
14
|
+
echo " $0 -p # 构建并发布,自动递增版本号"
|
|
15
|
+
echo " $0 -p 1.0.0 # 构建并发布,使用指定版本号"
|
|
16
|
+
}
|
|
17
|
+
|
|
18
|
+
# 解析命令行参数
|
|
19
|
+
PUBLISH=false
|
|
20
|
+
VERSION=""
|
|
21
|
+
VERSION_FILE=".version"
|
|
22
|
+
|
|
23
|
+
while [[ $# -gt 0 ]]; do
|
|
24
|
+
case $1 in
|
|
25
|
+
-p|--publish)
|
|
26
|
+
PUBLISH=true
|
|
27
|
+
shift
|
|
28
|
+
;;
|
|
29
|
+
-h|--help)
|
|
30
|
+
show_help
|
|
31
|
+
exit 0
|
|
32
|
+
;;
|
|
33
|
+
*)
|
|
34
|
+
if [[ -z "$VERSION" ]]; then
|
|
35
|
+
VERSION="$1"
|
|
36
|
+
else
|
|
37
|
+
echo "错误: 无效的参数 '$1'"
|
|
38
|
+
show_help
|
|
39
|
+
exit 1
|
|
40
|
+
fi
|
|
41
|
+
shift
|
|
42
|
+
;;
|
|
43
|
+
esac
|
|
44
|
+
done
|
|
45
|
+
|
|
46
|
+
# 检查是否安装了必要的工具
|
|
47
|
+
if ! command -v uv &> /dev/null; then
|
|
48
|
+
echo "错误: 未找到 uv 命令,请先安装 uv"
|
|
49
|
+
exit 1
|
|
50
|
+
fi
|
|
51
|
+
|
|
52
|
+
if ! command -v sed &> /dev/null; then
|
|
53
|
+
echo "错误: 未找到 sed 命令"
|
|
54
|
+
exit 1
|
|
55
|
+
fi
|
|
56
|
+
|
|
57
|
+
# 激活虚拟环境
|
|
58
|
+
# if [ -d ".venv" ]; then
|
|
59
|
+
# echo "正在激活虚拟环境..."
|
|
60
|
+
# source .venv/bin/activate
|
|
61
|
+
# if [ $? -ne 0 ]; then
|
|
62
|
+
# echo "错误: 虚拟环境激活失败"
|
|
63
|
+
# exit 1
|
|
64
|
+
# fi
|
|
65
|
+
# else
|
|
66
|
+
# echo "错误: 未找到虚拟环境(.venv目录)"
|
|
67
|
+
# exit 1
|
|
68
|
+
# fi
|
|
69
|
+
|
|
70
|
+
if [ "$PUBLISH" = true ]; then
|
|
71
|
+
# 读取或创建版本文件
|
|
72
|
+
if [ ! -f "$VERSION_FILE" ]; then
|
|
73
|
+
echo "0.1.0" > "$VERSION_FILE"
|
|
74
|
+
fi
|
|
75
|
+
current_version=$(cat "$VERSION_FILE")
|
|
76
|
+
echo "当前版本: $current_version"
|
|
77
|
+
|
|
78
|
+
# 确定新版本号
|
|
79
|
+
if [ -n "$VERSION" ]; then
|
|
80
|
+
# 如果提供了版本号参数,使用该版本号
|
|
81
|
+
new_version="$VERSION"
|
|
82
|
+
echo "使用指定版本: $new_version"
|
|
83
|
+
else
|
|
84
|
+
# 如果没有提供版本号,递增当前版本的修订号
|
|
85
|
+
IFS='.' read -r major minor patch <<< "$current_version"
|
|
86
|
+
new_patch=$((patch + 1))
|
|
87
|
+
new_version="$major.$minor.$new_patch"
|
|
88
|
+
echo "自动递增版本: $new_version"
|
|
89
|
+
fi
|
|
90
|
+
|
|
91
|
+
# 验证版本号格式
|
|
92
|
+
if ! [[ $new_version =~ ^[0-9]+\.[0-9]+\.[0-9]+$ ]]; then
|
|
93
|
+
echo "错误: 无效的版本号格式。版本号必须是 x.y.z 格式,其中 x、y、z 都是数字"
|
|
94
|
+
exit 1
|
|
95
|
+
fi
|
|
96
|
+
|
|
97
|
+
# 更新版本文件
|
|
98
|
+
echo "$new_version" > "$VERSION_FILE"
|
|
99
|
+
echo "版本文件已更新"
|
|
100
|
+
|
|
101
|
+
# 更新pyproject.toml中的版本号
|
|
102
|
+
if [[ "$OSTYPE" == "darwin"* ]]; then
|
|
103
|
+
# macOS
|
|
104
|
+
sed -i '' "s/^version = \".*\"/version = \"$new_version\"/" pyproject.toml
|
|
105
|
+
else
|
|
106
|
+
# Linux
|
|
107
|
+
sed -i "s/^version = \".*\"/version = \"$new_version\"/" pyproject.toml
|
|
108
|
+
fi
|
|
109
|
+
else
|
|
110
|
+
# 非发布模式,使用固定版本号0.0.1
|
|
111
|
+
new_version="0.0.1"
|
|
112
|
+
echo "使用开发版本: $new_version"
|
|
113
|
+
|
|
114
|
+
# 更新pyproject.toml中的版本号(不更新版本文件)
|
|
115
|
+
if [[ "$OSTYPE" == "darwin"* ]]; then
|
|
116
|
+
# macOS
|
|
117
|
+
sed -i '' "s/^version = \".*\"/version = \"$new_version\"/" pyproject.toml
|
|
118
|
+
else
|
|
119
|
+
# Linux
|
|
120
|
+
sed -i "s/^version = \".*\"/version = \"$new_version\"/" pyproject.toml
|
|
121
|
+
fi
|
|
122
|
+
fi
|
|
123
|
+
|
|
124
|
+
# 删除旧包
|
|
125
|
+
echo "正在删除旧包..."
|
|
126
|
+
rm -rf dist/*
|
|
127
|
+
rm -rf build/*
|
|
128
|
+
|
|
129
|
+
# 构建包
|
|
130
|
+
echo "正在构建包..."
|
|
131
|
+
uv build
|
|
132
|
+
|
|
133
|
+
# 检查构建是否成功
|
|
134
|
+
if [ $? -ne 0 ]; then
|
|
135
|
+
echo "错误: 构建失败"
|
|
136
|
+
exit 1
|
|
137
|
+
fi
|
|
138
|
+
|
|
139
|
+
# 如果指定了发布标志,则执行发布
|
|
140
|
+
if [ "$PUBLISH" = true ]; then
|
|
141
|
+
# 设置环境变量
|
|
142
|
+
export UV_PUBLISH_URL=https://upload.pypi.org/legacy/
|
|
143
|
+
export UV_PUBLISH_TOKEN=pypi-AgEIcHlwaS5vcmcCJDNhZjM3NGUzLTc0OTItNDU5MC1iNWYxLTQ2Yjc5MDNmOTkzMgACKlszLCJmOGZmOWJlOS02M2M1LTRiZmYtOTllZS1mMjRkYjEyMGE0ZGUiXQAABiCcCvNSctcaY_DWzZx9ShYEjtMqIdiyHkWBvGhc9vmPAg
|
|
144
|
+
|
|
145
|
+
echo "正在发布包..."
|
|
146
|
+
uv publish
|
|
147
|
+
|
|
148
|
+
# 检查发布是否成功
|
|
149
|
+
if [ $? -ne 0 ]; then
|
|
150
|
+
echo "错误: 发布失败"
|
|
151
|
+
exit 1
|
|
152
|
+
fi
|
|
153
|
+
echo "发布成功!新版本: $new_version"
|
|
154
|
+
else
|
|
155
|
+
echo "构建成功!新版本: $new_version"
|
|
156
|
+
echo "提示: 使用 -p 或 --publish 参数来执行发布"
|
|
157
|
+
fi
|
|
158
|
+
|
|
159
|
+
# 退出虚拟环境
|
|
160
|
+
# deactivate
|
|
@@ -0,0 +1,40 @@
|
|
|
1
|
+
[build-system]
|
|
2
|
+
requires = ["hatchling"]
|
|
3
|
+
build-backend = "hatchling.build"
|
|
4
|
+
|
|
5
|
+
[project]
|
|
6
|
+
name = "lightpdf-aipdf-backend"
|
|
7
|
+
version = "0.1.1"
|
|
8
|
+
description = "Backend API for LightPDF AI-PDF"
|
|
9
|
+
readme = "README.md"
|
|
10
|
+
requires-python = ">=3.9"
|
|
11
|
+
license = { text = "Proprietary" }
|
|
12
|
+
dependencies = [
|
|
13
|
+
"fastapi",
|
|
14
|
+
"uvicorn",
|
|
15
|
+
"python-dotenv",
|
|
16
|
+
"openai",
|
|
17
|
+
"mcp[cli]",
|
|
18
|
+
"pydantic",
|
|
19
|
+
"python-multipart",
|
|
20
|
+
]
|
|
21
|
+
|
|
22
|
+
[project.optional-dependencies]
|
|
23
|
+
dev = [
|
|
24
|
+
"pytest",
|
|
25
|
+
"black",
|
|
26
|
+
"isort",
|
|
27
|
+
]
|
|
28
|
+
|
|
29
|
+
[project.scripts]
|
|
30
|
+
lightpdf-aipdf-server = "lightpdf_aipdf_backend.cli:run_server"
|
|
31
|
+
|
|
32
|
+
[tool.hatch.build.targets.wheel]
|
|
33
|
+
packages = ["src/lightpdf_aipdf_backend"]
|
|
34
|
+
exclude = [".env", "*.log", "uploads/*"]
|
|
35
|
+
|
|
36
|
+
[tool.black]
|
|
37
|
+
line-length = 88
|
|
38
|
+
|
|
39
|
+
[tool.isort]
|
|
40
|
+
profile = "black"
|
|
@@ -0,0 +1,58 @@
|
|
|
1
|
+
from fastapi import FastAPI, Depends, HTTPException, File, UploadFile
|
|
2
|
+
from fastapi.middleware.cors import CORSMiddleware
|
|
3
|
+
from fastapi.responses import StreamingResponse
|
|
4
|
+
from contextlib import asynccontextmanager
|
|
5
|
+
|
|
6
|
+
from .models import ChatRequest
|
|
7
|
+
from .file_handler import handle_file_upload
|
|
8
|
+
from .chat_handler import process_messages, generate_chat_response
|
|
9
|
+
from .utils import async_generator_to_json_stream
|
|
10
|
+
|
|
11
|
+
@asynccontextmanager
|
|
12
|
+
async def lifespan(app: FastAPI):
|
|
13
|
+
"""应用生命周期管理"""
|
|
14
|
+
# 应用启动时的初始化
|
|
15
|
+
yield
|
|
16
|
+
# 应用关闭时的清理
|
|
17
|
+
|
|
18
|
+
# 创建FastAPI应用
|
|
19
|
+
app = FastAPI(title="LightPDF AI助手API", lifespan=lifespan)
|
|
20
|
+
|
|
21
|
+
# 配置CORS
|
|
22
|
+
app.add_middleware(
|
|
23
|
+
CORSMiddleware,
|
|
24
|
+
allow_origins=["*"], # 在生产环境中应该限制为特定域名
|
|
25
|
+
allow_credentials=True,
|
|
26
|
+
allow_methods=["*"],
|
|
27
|
+
allow_headers=["*"],
|
|
28
|
+
)
|
|
29
|
+
|
|
30
|
+
@app.get("/")
|
|
31
|
+
async def root():
|
|
32
|
+
"""API根路径"""
|
|
33
|
+
return {"message": "LightPDF AI助手 API已启动"}
|
|
34
|
+
|
|
35
|
+
@app.post("/api/upload")
|
|
36
|
+
async def upload_file(file: UploadFile = File(...)):
|
|
37
|
+
"""处理文件上传请求"""
|
|
38
|
+
return await handle_file_upload(file)
|
|
39
|
+
|
|
40
|
+
@app.post("/api/chat")
|
|
41
|
+
async def chat(request: ChatRequest):
|
|
42
|
+
"""处理聊天请求"""
|
|
43
|
+
try:
|
|
44
|
+
# 处理消息
|
|
45
|
+
processed_messages = await process_messages(request.messages)
|
|
46
|
+
|
|
47
|
+
# 生成响应 - 修改响应头部,确保流式内容不被缓存或批处理
|
|
48
|
+
return StreamingResponse(
|
|
49
|
+
async_generator_to_json_stream(generate_chat_response(processed_messages)),
|
|
50
|
+
media_type="application/x-ndjson",
|
|
51
|
+
headers={
|
|
52
|
+
"Cache-Control": "no-cache, no-transform",
|
|
53
|
+
"X-Accel-Buffering": "no" # 禁用Nginx缓冲
|
|
54
|
+
}
|
|
55
|
+
)
|
|
56
|
+
except Exception as e:
|
|
57
|
+
error_msg = str(e)
|
|
58
|
+
raise HTTPException(status_code=500, detail=f"LightPDF AI助手处理错误: {error_msg}")
|
|
@@ -0,0 +1,298 @@
|
|
|
1
|
+
import json
|
|
2
|
+
from typing import List, Dict, AsyncGenerator, Any, Tuple
|
|
3
|
+
from fastapi import HTTPException
|
|
4
|
+
|
|
5
|
+
from .state import get_openai_client, get_session_history, update_session_history, reset_session_history, get_mcp_session
|
|
6
|
+
from .models import Message
|
|
7
|
+
from .file_handler import get_file_references
|
|
8
|
+
from .tools import get_tools, format_tool_response, process_tool_path
|
|
9
|
+
from .utils import validate_and_fix_messages, async_generator_to_json_stream
|
|
10
|
+
from .config import Config
|
|
11
|
+
|
|
12
|
+
async def handle_tool_call(tool_call: Any, api_messages: List[Dict]) -> Tuple[Dict, Dict]:
|
|
13
|
+
"""处理单个工具调用
|
|
14
|
+
|
|
15
|
+
Args:
|
|
16
|
+
tool_call: 工具调用对象
|
|
17
|
+
api_messages: API消息列表
|
|
18
|
+
|
|
19
|
+
Returns:
|
|
20
|
+
Tuple[Dict, Dict]: 工具响应消息和yield消息
|
|
21
|
+
"""
|
|
22
|
+
tool_name = tool_call.function.name
|
|
23
|
+
tool_args = json.loads(tool_call.function.arguments)
|
|
24
|
+
tool_args = await process_tool_path(tool_args)
|
|
25
|
+
|
|
26
|
+
try:
|
|
27
|
+
mcp_session = get_mcp_session()
|
|
28
|
+
tool_response = await mcp_session.call_tool(tool_name, tool_args)
|
|
29
|
+
formatted_response = format_tool_response(tool_response.model_dump())
|
|
30
|
+
|
|
31
|
+
tool_response_message = {
|
|
32
|
+
"role": "tool",
|
|
33
|
+
"tool_call_id": tool_call.id,
|
|
34
|
+
"name": tool_name,
|
|
35
|
+
"content": formatted_response
|
|
36
|
+
}
|
|
37
|
+
|
|
38
|
+
yield_message = {
|
|
39
|
+
"type": "step",
|
|
40
|
+
"step_type": "result",
|
|
41
|
+
"content": formatted_response
|
|
42
|
+
}
|
|
43
|
+
|
|
44
|
+
except Exception as e:
|
|
45
|
+
error_msg = str(e)
|
|
46
|
+
tool_response_message = {
|
|
47
|
+
"role": "tool",
|
|
48
|
+
"tool_call_id": tool_call.id,
|
|
49
|
+
"name": tool_name,
|
|
50
|
+
"content": f"Error: {error_msg}"
|
|
51
|
+
}
|
|
52
|
+
|
|
53
|
+
yield_message = {
|
|
54
|
+
"type": "error",
|
|
55
|
+
"content": f"工具 {tool_name} 调用失败:{error_msg}"
|
|
56
|
+
}
|
|
57
|
+
|
|
58
|
+
return tool_response_message, yield_message
|
|
59
|
+
|
|
60
|
+
async def collect_stream_content(response: Any) -> AsyncGenerator[Dict, Tuple[str, str, List[Dict]]]:
|
|
61
|
+
"""收集流式响应内容
|
|
62
|
+
|
|
63
|
+
Args:
|
|
64
|
+
response: OpenAI流式响应对象
|
|
65
|
+
|
|
66
|
+
Yields:
|
|
67
|
+
Dict: 处理后的消息
|
|
68
|
+
|
|
69
|
+
Returns:
|
|
70
|
+
Tuple[str, str, List[Dict]]: 完整内容、完成原因和工具调用数据
|
|
71
|
+
"""
|
|
72
|
+
full_content = ""
|
|
73
|
+
finish_reason = None
|
|
74
|
+
tool_calls_data = []
|
|
75
|
+
|
|
76
|
+
for chunk in response:
|
|
77
|
+
choice = chunk.choices[0]
|
|
78
|
+
delta = choice.delta
|
|
79
|
+
|
|
80
|
+
if delta.content:
|
|
81
|
+
full_content += delta.content
|
|
82
|
+
yield {
|
|
83
|
+
"type": "stream_chunk",
|
|
84
|
+
"content": delta.content
|
|
85
|
+
}
|
|
86
|
+
|
|
87
|
+
if delta.tool_calls:
|
|
88
|
+
for tool_call_delta in delta.tool_calls:
|
|
89
|
+
tc_index = tool_call_delta.index
|
|
90
|
+
while len(tool_calls_data) <= (tc_index or 0):
|
|
91
|
+
tool_calls_data.append({"id": None, "type": "function", "function": {"name": "", "arguments": ""}})
|
|
92
|
+
|
|
93
|
+
tc_data = tool_calls_data[tc_index]
|
|
94
|
+
if tool_call_delta.id:
|
|
95
|
+
tc_data.update({"id": tool_call_delta.id})
|
|
96
|
+
|
|
97
|
+
tc_function = tool_call_delta.function
|
|
98
|
+
if tc_function:
|
|
99
|
+
current_tool = tc_data["function"]
|
|
100
|
+
if tc_function.name:
|
|
101
|
+
current_tool["name"] = tc_function.name
|
|
102
|
+
if tc_function.arguments:
|
|
103
|
+
current_tool["arguments"] = current_tool.get("arguments", "") + tc_function.arguments
|
|
104
|
+
|
|
105
|
+
if choice.finish_reason:
|
|
106
|
+
finish_reason = choice.finish_reason
|
|
107
|
+
|
|
108
|
+
yield (full_content, finish_reason, tool_calls_data) # 替换return,作为最后一个yield值
|
|
109
|
+
|
|
110
|
+
async def process_messages(messages: List[Message]) -> List[Dict]:
|
|
111
|
+
"""处理消息列表,包括文件处理和格式转换
|
|
112
|
+
|
|
113
|
+
Args:
|
|
114
|
+
messages: 消息列表
|
|
115
|
+
|
|
116
|
+
Returns:
|
|
117
|
+
List[Dict]: 处理后的消息列表
|
|
118
|
+
"""
|
|
119
|
+
# 检查是否是新会话
|
|
120
|
+
is_new_session = len(messages) <= 2 # 通常新会话只有系统消息和用户消息
|
|
121
|
+
session_history = get_session_history()
|
|
122
|
+
|
|
123
|
+
if is_new_session:
|
|
124
|
+
reset_session_history()
|
|
125
|
+
|
|
126
|
+
# 转换消息并处理文件
|
|
127
|
+
processed_messages = []
|
|
128
|
+
|
|
129
|
+
# 如果有会话历史且不是新会话,优先使用会话历史
|
|
130
|
+
if session_history and not is_new_session:
|
|
131
|
+
# 添加会话历史中的消息
|
|
132
|
+
processed_messages.extend(session_history)
|
|
133
|
+
|
|
134
|
+
# 只处理会话历史之后的新消息
|
|
135
|
+
new_messages = messages[len(session_history):]
|
|
136
|
+
else:
|
|
137
|
+
new_messages = messages
|
|
138
|
+
|
|
139
|
+
for msg in new_messages:
|
|
140
|
+
# 基本消息结构
|
|
141
|
+
message_dict = {"role": msg.role, "content": msg.content}
|
|
142
|
+
|
|
143
|
+
# 处理assistant角色的工具调用
|
|
144
|
+
if msg.role == "assistant" and msg.tool_calls:
|
|
145
|
+
message_dict["tool_calls"] = msg.tool_calls
|
|
146
|
+
# 当有工具调用时,content可以为空字符串
|
|
147
|
+
if not message_dict["content"]:
|
|
148
|
+
message_dict["content"] = ""
|
|
149
|
+
|
|
150
|
+
# 处理tool角色的工具响应
|
|
151
|
+
if msg.role == "tool":
|
|
152
|
+
if msg.tool_call_id:
|
|
153
|
+
message_dict["tool_call_id"] = msg.tool_call_id
|
|
154
|
+
if msg.name:
|
|
155
|
+
message_dict["name"] = msg.name
|
|
156
|
+
# 确保tool角色必须有tool_call_id和name
|
|
157
|
+
if "tool_call_id" not in message_dict or "name" not in message_dict:
|
|
158
|
+
continue
|
|
159
|
+
|
|
160
|
+
# 如果消息包含文件ID
|
|
161
|
+
if msg.file_ids:
|
|
162
|
+
file_urls = get_file_references(msg.file_ids)
|
|
163
|
+
|
|
164
|
+
# 将文件引用添加到消息内容
|
|
165
|
+
if file_urls:
|
|
166
|
+
message_dict["content"] += "\n\n以下是文件附件的Markdown链接:\n\n" + "\n".join(file_urls)
|
|
167
|
+
|
|
168
|
+
processed_messages.append(message_dict)
|
|
169
|
+
|
|
170
|
+
# 限制消息历史长度
|
|
171
|
+
MAX_MESSAGES = 50 # 增加消息历史长度以保留工具调用上下文
|
|
172
|
+
if len(processed_messages) > MAX_MESSAGES:
|
|
173
|
+
processed_messages = processed_messages[-MAX_MESSAGES:]
|
|
174
|
+
|
|
175
|
+
# 检查工具调用和响应的匹配
|
|
176
|
+
return validate_and_fix_messages(processed_messages)
|
|
177
|
+
|
|
178
|
+
async def generate_chat_response(messages: List[Dict]) -> AsyncGenerator:
|
|
179
|
+
"""生成聊天响应
|
|
180
|
+
|
|
181
|
+
Args:
|
|
182
|
+
messages: 处理后的消息列表
|
|
183
|
+
|
|
184
|
+
Yields:
|
|
185
|
+
Dict: 响应内容
|
|
186
|
+
"""
|
|
187
|
+
openai_client = get_openai_client()
|
|
188
|
+
|
|
189
|
+
try:
|
|
190
|
+
tools = await get_tools()
|
|
191
|
+
|
|
192
|
+
api_messages = messages.copy()
|
|
193
|
+
while True:
|
|
194
|
+
# 第一次调用API使用流式请求
|
|
195
|
+
response = openai_client.chat.completions.create(
|
|
196
|
+
model=Config.MODEL_NAME,
|
|
197
|
+
temperature=0.4,
|
|
198
|
+
messages=api_messages,
|
|
199
|
+
tools=tools,
|
|
200
|
+
parallel_tool_calls=False,
|
|
201
|
+
stream=True
|
|
202
|
+
)
|
|
203
|
+
|
|
204
|
+
full_content = ""
|
|
205
|
+
finish_reason = None
|
|
206
|
+
tool_calls_data = []
|
|
207
|
+
yield_content = False
|
|
208
|
+
# 收集流式响应内容
|
|
209
|
+
async for item in collect_stream_content(response):
|
|
210
|
+
if isinstance(item, tuple):
|
|
211
|
+
full_content, finish_reason, tool_calls_data = item
|
|
212
|
+
else:
|
|
213
|
+
if not yield_content:
|
|
214
|
+
# 第一次消息通知
|
|
215
|
+
yield {
|
|
216
|
+
"type": "stream_start",
|
|
217
|
+
"content": ""
|
|
218
|
+
}
|
|
219
|
+
yield_content = True
|
|
220
|
+
|
|
221
|
+
yield item
|
|
222
|
+
|
|
223
|
+
# 如果需要调用工具
|
|
224
|
+
if finish_reason == 'tool_calls' and tool_calls_data:
|
|
225
|
+
if yield_content:
|
|
226
|
+
# 停止当前流
|
|
227
|
+
yield {
|
|
228
|
+
"type": "stream_end",
|
|
229
|
+
"content": full_content
|
|
230
|
+
}
|
|
231
|
+
|
|
232
|
+
# 转换工具调用格式并创建消息
|
|
233
|
+
tool_calls = []
|
|
234
|
+
for tc_data in tool_calls_data:
|
|
235
|
+
tc_function = tc_data["function"]
|
|
236
|
+
if tc_data["id"] and tc_function["name"]:
|
|
237
|
+
tool_calls.append(type('ToolCall', (), {
|
|
238
|
+
'id': tc_data["id"],
|
|
239
|
+
'function': type('Function', (), {
|
|
240
|
+
'name': tc_function["name"],
|
|
241
|
+
'arguments': tc_function["arguments"]
|
|
242
|
+
})
|
|
243
|
+
}))
|
|
244
|
+
|
|
245
|
+
# 创建助手消息
|
|
246
|
+
assistant_message = {
|
|
247
|
+
"role": "assistant",
|
|
248
|
+
"content": full_content or "",
|
|
249
|
+
"tool_calls": [
|
|
250
|
+
{
|
|
251
|
+
"id": tc.id,
|
|
252
|
+
"type": "function",
|
|
253
|
+
"function": {
|
|
254
|
+
"name": tc.function.name,
|
|
255
|
+
"arguments": tc.function.arguments
|
|
256
|
+
}
|
|
257
|
+
} for tc in tool_calls
|
|
258
|
+
]
|
|
259
|
+
}
|
|
260
|
+
|
|
261
|
+
# 准备API消息列表
|
|
262
|
+
api_messages.append(assistant_message)
|
|
263
|
+
|
|
264
|
+
# 处理每个工具调用
|
|
265
|
+
for tool_call in tool_calls:
|
|
266
|
+
# 发送工具调用信息
|
|
267
|
+
yield {
|
|
268
|
+
"type": "step",
|
|
269
|
+
"step_type": "call",
|
|
270
|
+
"content": f"调用工具 {tool_call.function.name}\n参数:{tool_call.function.arguments}"
|
|
271
|
+
}
|
|
272
|
+
|
|
273
|
+
# 处理工具调用
|
|
274
|
+
tool_response_message, yield_message = await handle_tool_call(tool_call, api_messages)
|
|
275
|
+
api_messages.append(tool_response_message)
|
|
276
|
+
yield yield_message
|
|
277
|
+
|
|
278
|
+
else:
|
|
279
|
+
# 如果不需要调用工具,直接结束流
|
|
280
|
+
yield {
|
|
281
|
+
"type": "stream_end",
|
|
282
|
+
"content": full_content
|
|
283
|
+
}
|
|
284
|
+
|
|
285
|
+
# 更新会话历史
|
|
286
|
+
api_messages.append({
|
|
287
|
+
"role": "assistant",
|
|
288
|
+
"content": full_content
|
|
289
|
+
})
|
|
290
|
+
update_session_history(api_messages)
|
|
291
|
+
|
|
292
|
+
break
|
|
293
|
+
|
|
294
|
+
except Exception as e:
|
|
295
|
+
yield {
|
|
296
|
+
"type": "error",
|
|
297
|
+
"content": str(e)
|
|
298
|
+
}
|
|
@@ -0,0 +1,22 @@
|
|
|
1
|
+
import os
|
|
2
|
+
from pathlib import Path
|
|
3
|
+
from dotenv import load_dotenv
|
|
4
|
+
|
|
5
|
+
# 加载环境变量
|
|
6
|
+
load_dotenv()
|
|
7
|
+
|
|
8
|
+
# 文件上传目录
|
|
9
|
+
UPLOADS_DIR = Path("uploads")
|
|
10
|
+
UPLOADS_DIR.mkdir(exist_ok=True)
|
|
11
|
+
|
|
12
|
+
class Config:
|
|
13
|
+
"""应用配置类"""
|
|
14
|
+
MODEL_NAME = os.getenv("OPENAI_MODEL", "gpt-4o-mini")
|
|
15
|
+
BASE_URL = os.getenv("OPENAI_BASE_URL")
|
|
16
|
+
API_KEY = os.getenv("OPENAI_API_KEY")
|
|
17
|
+
|
|
18
|
+
@classmethod
|
|
19
|
+
def validate(cls):
|
|
20
|
+
"""验证关键配置项"""
|
|
21
|
+
if not cls.API_KEY:
|
|
22
|
+
raise ValueError("Missing OpenAI API Key")
|
|
@@ -0,0 +1,81 @@
|
|
|
1
|
+
import os
|
|
2
|
+
import shutil
|
|
3
|
+
from uuid import uuid4
|
|
4
|
+
from fastapi import UploadFile, HTTPException
|
|
5
|
+
from pathlib import Path
|
|
6
|
+
|
|
7
|
+
from .config import UPLOADS_DIR
|
|
8
|
+
from .models import FileInfo
|
|
9
|
+
from .state import uploaded_files
|
|
10
|
+
|
|
11
|
+
async def handle_file_upload(file: UploadFile) -> FileInfo:
|
|
12
|
+
"""处理文件上传请求
|
|
13
|
+
|
|
14
|
+
Args:
|
|
15
|
+
file: 上传的文件
|
|
16
|
+
|
|
17
|
+
Returns:
|
|
18
|
+
FileInfo: 文件信息对象
|
|
19
|
+
|
|
20
|
+
Raises:
|
|
21
|
+
HTTPException: 文件上传失败时抛出
|
|
22
|
+
"""
|
|
23
|
+
# 生成唯一文件名
|
|
24
|
+
file_id = str(uuid4())
|
|
25
|
+
filename = file.filename or "unknown_file"
|
|
26
|
+
content_type = file.content_type or "application/octet-stream"
|
|
27
|
+
|
|
28
|
+
# 确保文件名安全
|
|
29
|
+
safe_filename = f"{file_id}_{filename}"
|
|
30
|
+
file_path = UPLOADS_DIR / safe_filename
|
|
31
|
+
|
|
32
|
+
# 确保UPLOADS_DIR存在
|
|
33
|
+
UPLOADS_DIR.mkdir(exist_ok=True)
|
|
34
|
+
|
|
35
|
+
# 保存文件
|
|
36
|
+
try:
|
|
37
|
+
with open(file_path, "wb") as buffer:
|
|
38
|
+
shutil.copyfileobj(file.file, buffer)
|
|
39
|
+
|
|
40
|
+
# 使用相对路径(仅文件名)
|
|
41
|
+
relative_path = safe_filename
|
|
42
|
+
|
|
43
|
+
# 创建文件信息
|
|
44
|
+
file_info = FileInfo(
|
|
45
|
+
file_id=file_id,
|
|
46
|
+
filename=filename,
|
|
47
|
+
content_type=content_type,
|
|
48
|
+
path=relative_path
|
|
49
|
+
)
|
|
50
|
+
|
|
51
|
+
# 存储文件信息
|
|
52
|
+
uploaded_files[file_id] = file_info
|
|
53
|
+
|
|
54
|
+
return file_info
|
|
55
|
+
except Exception as e:
|
|
56
|
+
raise HTTPException(status_code=500, detail=f"文件上传失败: {str(e)}")
|
|
57
|
+
|
|
58
|
+
def get_file_references(file_ids: list[str]) -> list[str]:
|
|
59
|
+
"""获取文件引用链接
|
|
60
|
+
|
|
61
|
+
Args:
|
|
62
|
+
file_ids: 文件ID列表
|
|
63
|
+
|
|
64
|
+
Returns:
|
|
65
|
+
list[str]: 文件引用链接列表
|
|
66
|
+
"""
|
|
67
|
+
file_urls = []
|
|
68
|
+
|
|
69
|
+
for file_id in file_ids:
|
|
70
|
+
if file_id in uploaded_files:
|
|
71
|
+
file_info = uploaded_files[file_id]
|
|
72
|
+
# 检查文件是否存在
|
|
73
|
+
file_path = os.path.join(UPLOADS_DIR, file_info.path) if not os.path.isabs(file_info.path) else file_info.path
|
|
74
|
+
if os.path.exists(file_path):
|
|
75
|
+
# 根据文件类型添加不同的引用格式
|
|
76
|
+
if file_info.content_type.startswith("image/"):
|
|
77
|
+
file_urls.append(f"")
|
|
78
|
+
else:
|
|
79
|
+
file_urls.append(f"[文件: {file_info.filename}]({file_info.path})")
|
|
80
|
+
|
|
81
|
+
return file_urls
|
|
@@ -0,0 +1,45 @@
|
|
|
1
|
+
import os
|
|
2
|
+
import asyncio
|
|
3
|
+
import uvicorn
|
|
4
|
+
from mcp import ClientSession, StdioServerParameters
|
|
5
|
+
from mcp.client.stdio import stdio_client
|
|
6
|
+
|
|
7
|
+
from .config import Config, UPLOADS_DIR
|
|
8
|
+
from .state import set_mcp_session, init_openai_client
|
|
9
|
+
from .app import app
|
|
10
|
+
from .tools import get_tools
|
|
11
|
+
|
|
12
|
+
async def main():
|
|
13
|
+
"""应用主入口"""
|
|
14
|
+
# 初始化 OpenAI 客户端
|
|
15
|
+
init_openai_client()
|
|
16
|
+
|
|
17
|
+
# 准备 MCP 服务参数
|
|
18
|
+
server_params = StdioServerParameters(
|
|
19
|
+
command="uvx",
|
|
20
|
+
args=["mcp_lightpdf@latest"],
|
|
21
|
+
# args=["-n", "--from", "../../../mcp_server/dist/mcp_lightpdf-0.0.1-py3-none-any.whl", "mcp_lightpdf"],
|
|
22
|
+
cwd=UPLOADS_DIR,
|
|
23
|
+
env={
|
|
24
|
+
"API_KEY": os.getenv("API_KEY"),
|
|
25
|
+
}
|
|
26
|
+
)
|
|
27
|
+
|
|
28
|
+
# 启动 MCP 会话
|
|
29
|
+
async with stdio_client(server_params) as (read, write):
|
|
30
|
+
async with ClientSession(read, write) as session:
|
|
31
|
+
await session.initialize()
|
|
32
|
+
# 设置全局 MCP 会话
|
|
33
|
+
set_mcp_session(session)
|
|
34
|
+
|
|
35
|
+
tools = await get_tools()
|
|
36
|
+
print(tools)
|
|
37
|
+
|
|
38
|
+
# 启动 FastAPI 服务器
|
|
39
|
+
config = uvicorn.Config(app, port=3300)
|
|
40
|
+
server = uvicorn.Server(config)
|
|
41
|
+
await server.serve()
|
|
42
|
+
|
|
43
|
+
# 确保与原始 main 函数相同
|
|
44
|
+
if __name__ == "__main__":
|
|
45
|
+
asyncio.run(main())
|
|
@@ -0,0 +1,22 @@
|
|
|
1
|
+
from typing import List, Optional, Dict
|
|
2
|
+
from pydantic import BaseModel
|
|
3
|
+
|
|
4
|
+
class FileInfo(BaseModel):
|
|
5
|
+
"""文件信息模型"""
|
|
6
|
+
file_id: str
|
|
7
|
+
filename: str
|
|
8
|
+
content_type: str
|
|
9
|
+
path: str
|
|
10
|
+
|
|
11
|
+
class Message(BaseModel):
|
|
12
|
+
"""聊天消息模型"""
|
|
13
|
+
role: str
|
|
14
|
+
content: str
|
|
15
|
+
file_ids: Optional[List[str]] = None
|
|
16
|
+
tool_calls: Optional[List[Dict]] = None
|
|
17
|
+
tool_call_id: Optional[str] = None
|
|
18
|
+
name: Optional[str] = None
|
|
19
|
+
|
|
20
|
+
class ChatRequest(BaseModel):
|
|
21
|
+
"""聊天请求模型"""
|
|
22
|
+
messages: List[Message]
|
|
@@ -0,0 +1,53 @@
|
|
|
1
|
+
from typing import Dict, List, Optional, Any
|
|
2
|
+
from collections import defaultdict
|
|
3
|
+
from openai import OpenAI
|
|
4
|
+
from mcp import ClientSession
|
|
5
|
+
|
|
6
|
+
from .config import Config
|
|
7
|
+
|
|
8
|
+
# 全局变量
|
|
9
|
+
mcp_session: Optional[ClientSession] = None
|
|
10
|
+
openai_client: Optional[OpenAI] = None
|
|
11
|
+
active_conversations = defaultdict(lambda: {"messages": [], "generator": None})
|
|
12
|
+
session_history: List[Dict] = []
|
|
13
|
+
uploaded_files: Dict[str, Any] = {}
|
|
14
|
+
|
|
15
|
+
def init_openai_client():
|
|
16
|
+
"""初始化OpenAI客户端"""
|
|
17
|
+
global openai_client
|
|
18
|
+
Config.validate()
|
|
19
|
+
openai_client = OpenAI(
|
|
20
|
+
api_key=Config.API_KEY,
|
|
21
|
+
base_url=Config.BASE_URL
|
|
22
|
+
)
|
|
23
|
+
return openai_client
|
|
24
|
+
|
|
25
|
+
def get_openai_client() -> OpenAI:
|
|
26
|
+
"""获取OpenAI客户端实例"""
|
|
27
|
+
global openai_client
|
|
28
|
+
if openai_client is None:
|
|
29
|
+
openai_client = init_openai_client()
|
|
30
|
+
return openai_client
|
|
31
|
+
|
|
32
|
+
def set_mcp_session(session: ClientSession):
|
|
33
|
+
"""设置MCP会话"""
|
|
34
|
+
global mcp_session
|
|
35
|
+
mcp_session = session
|
|
36
|
+
|
|
37
|
+
def get_mcp_session() -> Optional[ClientSession]:
|
|
38
|
+
"""获取MCP会话"""
|
|
39
|
+
return mcp_session
|
|
40
|
+
|
|
41
|
+
def reset_session_history():
|
|
42
|
+
"""重置会话历史"""
|
|
43
|
+
global session_history
|
|
44
|
+
session_history = []
|
|
45
|
+
|
|
46
|
+
def get_session_history() -> List[Dict]:
|
|
47
|
+
"""获取会话历史"""
|
|
48
|
+
return session_history
|
|
49
|
+
|
|
50
|
+
def update_session_history(history: List[Dict]):
|
|
51
|
+
"""更新会话历史"""
|
|
52
|
+
global session_history
|
|
53
|
+
session_history = history
|
|
@@ -0,0 +1,96 @@
|
|
|
1
|
+
import os
|
|
2
|
+
import json
|
|
3
|
+
from typing import List, Dict, Any
|
|
4
|
+
from fastapi import HTTPException
|
|
5
|
+
|
|
6
|
+
from .state import get_mcp_session
|
|
7
|
+
|
|
8
|
+
async def get_tools() -> List[Dict]:
|
|
9
|
+
"""从MCP client获取工具列表
|
|
10
|
+
|
|
11
|
+
Returns:
|
|
12
|
+
List[Dict]: 工具列表
|
|
13
|
+
|
|
14
|
+
Raises:
|
|
15
|
+
HTTPException: MCP会话未初始化时抛出
|
|
16
|
+
"""
|
|
17
|
+
mcp_session = get_mcp_session()
|
|
18
|
+
if not mcp_session:
|
|
19
|
+
raise HTTPException(status_code=500, detail="MCP session not initialized")
|
|
20
|
+
|
|
21
|
+
tools_result = await mcp_session.list_tools()
|
|
22
|
+
tools = []
|
|
23
|
+
|
|
24
|
+
for tool in tools_result.tools:
|
|
25
|
+
tool_dict = {
|
|
26
|
+
"type": "function",
|
|
27
|
+
"function": {
|
|
28
|
+
"name": tool.name,
|
|
29
|
+
"description": tool.description,
|
|
30
|
+
"parameters": tool.inputSchema
|
|
31
|
+
}
|
|
32
|
+
}
|
|
33
|
+
tools.append(tool_dict)
|
|
34
|
+
return tools
|
|
35
|
+
|
|
36
|
+
def format_tool_response(response_dict: Dict) -> str:
|
|
37
|
+
"""格式化工具响应,提取可能包含的Markdown内容
|
|
38
|
+
|
|
39
|
+
Args:
|
|
40
|
+
response_dict: 工具响应字典
|
|
41
|
+
|
|
42
|
+
Returns:
|
|
43
|
+
str: 格式化后的响应
|
|
44
|
+
"""
|
|
45
|
+
print(response_dict)
|
|
46
|
+
# 如果有 content 字段,优先使用它
|
|
47
|
+
if 'content' in response_dict:
|
|
48
|
+
return response_dict['content']
|
|
49
|
+
|
|
50
|
+
# 如果有 markdown 字段,优先使用它
|
|
51
|
+
if 'markdown' in response_dict:
|
|
52
|
+
return response_dict['markdown']
|
|
53
|
+
|
|
54
|
+
# 检查是否有返回文本字段
|
|
55
|
+
if 'text' in response_dict:
|
|
56
|
+
return response_dict['text']
|
|
57
|
+
|
|
58
|
+
# 如果有结果字段,检查是否包含Markdown内容
|
|
59
|
+
if 'result' in response_dict:
|
|
60
|
+
result = response_dict['result']
|
|
61
|
+
if isinstance(result, dict):
|
|
62
|
+
# 优先查找markdown字段
|
|
63
|
+
if 'markdown' in result:
|
|
64
|
+
return result['markdown']
|
|
65
|
+
elif 'text' in result:
|
|
66
|
+
return result['text']
|
|
67
|
+
elif 'content' in result:
|
|
68
|
+
return result['content']
|
|
69
|
+
|
|
70
|
+
# 否则返回格式化的JSON
|
|
71
|
+
return f"工具返回结果:\n```json\n{json.dumps(response_dict, ensure_ascii=False, indent=2)}\n```"
|
|
72
|
+
|
|
73
|
+
async def process_tool_path(tool_args: Dict) -> Dict:
|
|
74
|
+
"""处理工具调用参数中的文件路径
|
|
75
|
+
|
|
76
|
+
Args:
|
|
77
|
+
tool_args: 工具调用参数
|
|
78
|
+
|
|
79
|
+
Returns:
|
|
80
|
+
Dict: 处理后的参数
|
|
81
|
+
"""
|
|
82
|
+
if 'file_path' in tool_args:
|
|
83
|
+
from .config import UPLOADS_DIR # 避免循环导入
|
|
84
|
+
|
|
85
|
+
file_path = tool_args['file_path']
|
|
86
|
+
# 尝试确保文件路径正确
|
|
87
|
+
if not os.path.exists(file_path):
|
|
88
|
+
file_name = os.path.basename(file_path)
|
|
89
|
+
# 在uploads目录中查找
|
|
90
|
+
for uploaded_file in os.listdir(UPLOADS_DIR):
|
|
91
|
+
if file_name in uploaded_file:
|
|
92
|
+
full_path = os.path.join(UPLOADS_DIR, uploaded_file)
|
|
93
|
+
tool_args['file_path'] = full_path
|
|
94
|
+
break
|
|
95
|
+
|
|
96
|
+
return tool_args
|
|
@@ -0,0 +1,52 @@
|
|
|
1
|
+
import json
|
|
2
|
+
from typing import Dict, List, AsyncGenerator, Any
|
|
3
|
+
|
|
4
|
+
def validate_and_fix_messages(messages: List[Dict]) -> List[Dict]:
|
|
5
|
+
"""验证并修复消息格式,确保工具调用和响应的匹配
|
|
6
|
+
|
|
7
|
+
Args:
|
|
8
|
+
messages: 消息列表
|
|
9
|
+
|
|
10
|
+
Returns:
|
|
11
|
+
List[Dict]: 修复后的消息列表
|
|
12
|
+
"""
|
|
13
|
+
if not messages:
|
|
14
|
+
return []
|
|
15
|
+
|
|
16
|
+
fixed_messages = []
|
|
17
|
+
tool_call_ids = {} # 用于跟踪工具调用ID及其对应的工具名称
|
|
18
|
+
|
|
19
|
+
for msg in messages:
|
|
20
|
+
# 检查assistant角色的工具调用
|
|
21
|
+
if msg["role"] == "assistant" and "tool_calls" in msg:
|
|
22
|
+
# 存储工具调用ID和对应的工具名称
|
|
23
|
+
for tool_call in msg["tool_calls"]:
|
|
24
|
+
tool_call_ids[tool_call["id"]] = tool_call["function"]["name"]
|
|
25
|
+
|
|
26
|
+
# 处理tool角色的响应,确保有对应的tool_call_id和name
|
|
27
|
+
if msg["role"] == "tool":
|
|
28
|
+
if "tool_call_id" not in msg or not msg["tool_call_id"]:
|
|
29
|
+
continue
|
|
30
|
+
|
|
31
|
+
if "name" not in msg or not msg["name"]:
|
|
32
|
+
# 如果能找到对应的工具调用,自动填充name
|
|
33
|
+
if msg["tool_call_id"] in tool_call_ids:
|
|
34
|
+
msg["name"] = tool_call_ids[msg["tool_call_id"]]
|
|
35
|
+
else:
|
|
36
|
+
continue
|
|
37
|
+
|
|
38
|
+
fixed_messages.append(msg)
|
|
39
|
+
|
|
40
|
+
return fixed_messages
|
|
41
|
+
|
|
42
|
+
async def async_generator_to_json_stream(generator: AsyncGenerator) -> AsyncGenerator:
|
|
43
|
+
"""将异步生成器转换为JSON流
|
|
44
|
+
|
|
45
|
+
Args:
|
|
46
|
+
generator: 异步生成器
|
|
47
|
+
|
|
48
|
+
Yields:
|
|
49
|
+
bytes: JSON流
|
|
50
|
+
"""
|
|
51
|
+
async for item in generator:
|
|
52
|
+
yield json.dumps(item, ensure_ascii=False).encode('utf-8') + b'\n'
|