refactor: 重构 config 模块
This commit is contained in:
parent
48cf8e2f21
commit
0ee327cb97
@ -33,6 +33,7 @@ export default defineConfig({
|
|||||||
{
|
{
|
||||||
text: '高级配置',
|
text: '高级配置',
|
||||||
items: [
|
items: [
|
||||||
|
{ text: '配置系统详解', link: '/advanced/configuration' },
|
||||||
{ text: '文档解析', link: '/advanced/document-processing' },
|
{ text: '文档解析', link: '/advanced/document-processing' },
|
||||||
{ text: '智能体', link: '/advanced/agents' },
|
{ text: '智能体', link: '/advanced/agents' },
|
||||||
{ text: '品牌自定义', link: '/advanced/branding' },
|
{ text: '品牌自定义', link: '/advanced/branding' },
|
||||||
@ -42,10 +43,10 @@ export default defineConfig({
|
|||||||
{
|
{
|
||||||
text: '更新日志',
|
text: '更新日志',
|
||||||
items: [
|
items: [
|
||||||
|
{ text: '版本说明 v0.3', link: '/changelog/0.3-release-notes' },
|
||||||
{ text: '路线图', link: '/changelog/roadmap' },
|
{ text: '路线图', link: '/changelog/roadmap' },
|
||||||
{ text: '参与贡献', link: '/changelog/contributing' },
|
{ text: '参与贡献', link: '/changelog/contributing' },
|
||||||
{ text: '常见问题', link: '/changelog/faq' },
|
{ text: '常见问题', link: '/changelog/faq' }
|
||||||
{ text: '版本说明 v0.3', link: '/changelog/0.3-release-notes' }
|
|
||||||
]
|
]
|
||||||
}
|
}
|
||||||
],
|
],
|
||||||
|
|||||||
185
docs/advanced/configuration.md
Normal file
185
docs/advanced/configuration.md
Normal file
@ -0,0 +1,185 @@
|
|||||||
|
# 配置系统详解
|
||||||
|
|
||||||
|
## 概述
|
||||||
|
|
||||||
|
Yuxi-Know 从 v0.3.x 版本开始采用了全新的配置系统,基于 Pydantic BaseModel 和 TOML 格式,提供了类型安全、智能提示和选择性持久化等现代化特性。
|
||||||
|
|
||||||
|
## 架构设计
|
||||||
|
|
||||||
|
### 配置层次结构
|
||||||
|
|
||||||
|
```
|
||||||
|
配置系统架构
|
||||||
|
├── 默认配置 (代码定义)
|
||||||
|
│ ├── src/config/static/models.py (模型配置)
|
||||||
|
│ └── src/config/app.py (应用配置)
|
||||||
|
├── 用户配置 (TOML 文件)
|
||||||
|
│ └── saves/config/base.toml (仅保存用户修改)
|
||||||
|
└── 环境变量 (运行时覆盖)
|
||||||
|
└── .env 文件
|
||||||
|
```
|
||||||
|
|
||||||
|
### 核心组件
|
||||||
|
|
||||||
|
#### 1. Config 类 (`src/config/app.py`)
|
||||||
|
|
||||||
|
主配置类,继承自 Pydantic BaseModel,提供:
|
||||||
|
|
||||||
|
- **类型验证**: 自动检查配置项类型
|
||||||
|
- **默认值管理**: 内置合理的默认配置
|
||||||
|
- **选择性持久化**: 仅保存用户修改的配置项
|
||||||
|
- **向后兼容**: 支持旧的字典式访问方式
|
||||||
|
|
||||||
|
```python
|
||||||
|
class Config(BaseModel):
|
||||||
|
# 功能开关
|
||||||
|
enable_reranker: bool = Field(default=False, description="是否开启重排序")
|
||||||
|
enable_content_guard: bool = Field(default=False, description="是否启用内容审查")
|
||||||
|
|
||||||
|
# 模型配置
|
||||||
|
default_model: str = Field(default="siliconflow/deepseek-ai/DeepSeek-V3.2-Exp")
|
||||||
|
embed_model: str = Field(default="siliconflow/BAAI/bge-m3")
|
||||||
|
|
||||||
|
# 运行时状态 (不持久化)
|
||||||
|
model_provider_status: dict[str, bool] = Field(exclude=True)
|
||||||
|
```
|
||||||
|
|
||||||
|
#### 2. 模型配置类 (`src/config/static/models.py`)
|
||||||
|
|
||||||
|
定义了三种类型的模型配置:
|
||||||
|
|
||||||
|
- **ChatModelProvider**: 聊天模型提供商
|
||||||
|
- **EmbedModelInfo**: 嵌入模型信息
|
||||||
|
- **RerankerInfo**: 重排序模型信息
|
||||||
|
|
||||||
|
```python
|
||||||
|
class ChatModelProvider(BaseModel):
|
||||||
|
name: str = Field(..., description="提供商显示名称")
|
||||||
|
url: str = Field(..., description="提供商文档或模型列表 URL")
|
||||||
|
base_url: str = Field(..., description="API 基础 URL")
|
||||||
|
default: str = Field(..., description="默认模型名称")
|
||||||
|
env: str = Field(..., description="API Key 环境变量名")
|
||||||
|
models: list[str] = Field(default_factory=list, description="支持的模型列表")
|
||||||
|
```
|
||||||
|
|
||||||
|
添加配置:
|
||||||
|
|
||||||
|
|
||||||
|
```python
|
||||||
|
# 1. 在 DEFAULT_CHAT_MODEL_PROVIDERS 中添加
|
||||||
|
"new-provider": ChatModelProvider(
|
||||||
|
name="新提供商",
|
||||||
|
url="https://provider.com/docs",
|
||||||
|
base_url="https://api.provider.com/v1",
|
||||||
|
default="default-model",
|
||||||
|
env="NEW_PROVIDER_API_KEY",
|
||||||
|
models=["model1", "model2"],
|
||||||
|
),
|
||||||
|
|
||||||
|
# 2. 在 .env 中配置 API Key
|
||||||
|
# NEW_PROVIDER_API_KEY=your_api_key
|
||||||
|
|
||||||
|
# 3. 重启服务或重新加载配置
|
||||||
|
```
|
||||||
|
|
||||||
|
## 配置管理特性
|
||||||
|
|
||||||
|
系统只会保存用户修改过的配置项:
|
||||||
|
|
||||||
|
```python
|
||||||
|
# 用户只修改了 enable_reranker
|
||||||
|
config.enable_reranker = True
|
||||||
|
config.save() # 只保存 enable_reranker 到 TOML 文件
|
||||||
|
|
||||||
|
# TOML 文件内容
|
||||||
|
# enable_reranker = true
|
||||||
|
```
|
||||||
|
|
||||||
|
### 默认模型配置 (`src/config/static/models.py`)
|
||||||
|
|
||||||
|
包含所有支持的模型提供商的默认配置,开发者可以直接修改此文件添加新的模型:
|
||||||
|
|
||||||
|
```python
|
||||||
|
DEFAULT_CHAT_MODEL_PROVIDERS: dict[str, ChatModelProvider] = {
|
||||||
|
"siliconflow": ChatModelProvider(
|
||||||
|
name="SiliconFlow",
|
||||||
|
url="https://cloud.siliconflow.cn/models",
|
||||||
|
base_url="https://api.siliconflow.cn/v1",
|
||||||
|
default="deepseek-ai/DeepSeek-V3.2-Exp",
|
||||||
|
env="SILICONFLOW_API_KEY",
|
||||||
|
models=[
|
||||||
|
"deepseek-ai/DeepSeek-V3.2-Exp",
|
||||||
|
"Qwen/Qwen3-235B-A22B-Instruct-2507",
|
||||||
|
# ...
|
||||||
|
],
|
||||||
|
),
|
||||||
|
# 更多提供商...
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### 用户配置 (`saves/config/base.toml`)
|
||||||
|
|
||||||
|
只包含用户修改过的配置项,使用 TOML 格式:
|
||||||
|
|
||||||
|
```toml
|
||||||
|
# 用户只修改了这些配置项
|
||||||
|
enable_reranker = true
|
||||||
|
default_agent_id = "MyCustomAgent"
|
||||||
|
enable_content_guard = true
|
||||||
|
|
||||||
|
# 模型配置修改
|
||||||
|
[model_names.siliconflow]
|
||||||
|
models = [
|
||||||
|
"deepseek-ai/DeepSeek-V3.2-Exp",
|
||||||
|
"custom-model-name",
|
||||||
|
]
|
||||||
|
```
|
||||||
|
|
||||||
|
## 高级配置
|
||||||
|
|
||||||
|
### 动态配置更新
|
||||||
|
|
||||||
|
```python
|
||||||
|
from src.config import config
|
||||||
|
|
||||||
|
# 更新配置
|
||||||
|
config.enable_reranker = True
|
||||||
|
config.default_agent_id = "CustomAgent"
|
||||||
|
|
||||||
|
# 更新模型列表
|
||||||
|
config.model_names["siliconflow"].models.append("new-model")
|
||||||
|
|
||||||
|
# 保存配置
|
||||||
|
config.save()
|
||||||
|
|
||||||
|
# 或者只保存特定提供商的模型配置
|
||||||
|
config._save_models_to_file("siliconflow")
|
||||||
|
```
|
||||||
|
|
||||||
|
### 配置验证
|
||||||
|
|
||||||
|
```python
|
||||||
|
# 验证配置
|
||||||
|
from src.config import config
|
||||||
|
|
||||||
|
# 检查模型提供商可用性
|
||||||
|
for provider, status in config.model_provider_status.items():
|
||||||
|
print(f"{provider}: {'✅' if status else '❌'}")
|
||||||
|
|
||||||
|
# 获取可用模型列表
|
||||||
|
available_models = config.get_model_choices()
|
||||||
|
available_embed_models = config.get_embed_model_choices()
|
||||||
|
available_rerankers = config.get_reranker_choices()
|
||||||
|
```
|
||||||
|
|
||||||
|
### 配置导出
|
||||||
|
|
||||||
|
```python
|
||||||
|
# 导出完整配置(包含运行时状态)
|
||||||
|
full_config = config.dump_config()
|
||||||
|
|
||||||
|
# 导出用户配置(仅保存到文件的部分)
|
||||||
|
user_config = {
|
||||||
|
field: getattr(config, field)
|
||||||
|
for field in config._user_modified_fields
|
||||||
|
}
|
||||||
@ -18,6 +18,56 @@ Yuxi-Know v0.3 是一个重要的里程碑版本,包含了多项架构重构
|
|||||||
cp .env.template .env
|
cp .env.template .env
|
||||||
```
|
```
|
||||||
|
|
||||||
|
### 配置文件管理调整
|
||||||
|
|
||||||
|
配置系统从 YAML 格式迁移到了基于 Pydantic BaseModel + TOML 的现代化配置系统。
|
||||||
|
|
||||||
|
| 项目 | v0.2.x | v0.3.x |
|
||||||
|
|------|--------|--------|
|
||||||
|
| 默认模型配置 | `src/config/static/models.yaml` | `src/config/static/models.py` |
|
||||||
|
| 用户配置 | `saves/config/base.yaml` | `saves/config/base.toml` |
|
||||||
|
| 配置格式 | YAML | Python 代码 + TOML |
|
||||||
|
| 类型安全 | ❌ 无 | ✅ Pydantic 验证 |
|
||||||
|
| IDE 支持 | ❌ 基础 | ✅ 完整智能提示 |
|
||||||
|
| 持久化策略 | 全量保存 | 选择性保存 |
|
||||||
|
|
||||||
|
|
||||||
|
示例:迁移自定义模型提供商
|
||||||
|
|
||||||
|
假设你的旧 `models.yaml` 中有:
|
||||||
|
|
||||||
|
```yaml
|
||||||
|
# 旧的 models.yaml
|
||||||
|
MODEL_NAMES:
|
||||||
|
custom-provider:
|
||||||
|
name: "My Custom Provider"
|
||||||
|
base_url: "https://api.custom.com/v1"
|
||||||
|
default: "custom-model"
|
||||||
|
env: "CUSTOM_API_KEY"
|
||||||
|
models:
|
||||||
|
- "custom-model"
|
||||||
|
- "another-model"
|
||||||
|
```
|
||||||
|
|
||||||
|
需要在新的 `src/config/static/models.py` 中添加:
|
||||||
|
|
||||||
|
```python
|
||||||
|
# 新的 models.py
|
||||||
|
DEFAULT_CHAT_MODEL_PROVIDERS: dict[str, ChatModelProvider] = {
|
||||||
|
# ... 现有配置 ...
|
||||||
|
|
||||||
|
"custom-provider": ChatModelProvider(
|
||||||
|
name="My Custom Provider",
|
||||||
|
url="https://custom.com/docs",
|
||||||
|
base_url="https://api.custom.com/v1",
|
||||||
|
default="custom-model",
|
||||||
|
env="CUSTOM_API_KEY",
|
||||||
|
models=["custom-model", "another-model"],
|
||||||
|
),
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
|
||||||
### 数据库存储架构重构
|
### 数据库存储架构重构
|
||||||
- **变更内容**: 重新实现了对话管理的存储与管理,不再依赖于 MemorySaver
|
- **变更内容**: 重新实现了对话管理的存储与管理,不再依赖于 MemorySaver
|
||||||
- **影响**: 使用新的存储结构,之前存储的历史记录无法直接迁移
|
- **影响**: 使用新的存储结构,之前存储的历史记录无法直接迁移
|
||||||
|
|||||||
@ -7,8 +7,7 @@
|
|||||||
|
|
||||||
## Bugs
|
## Bugs
|
||||||
|
|
||||||
- [ ] 当前 ReAct 智能体有消息顺序错乱的 bug,且不会默认调用工具
|
-
|
||||||
|
|
||||||
|
|
||||||
## Next
|
## Next
|
||||||
|
|
||||||
@ -20,7 +19,6 @@
|
|||||||
- [ ] 集成智能体评估,首先使用命令行来实现,然后考虑放在 UI 里面展示
|
- [ ] 集成智能体评估,首先使用命令行来实现,然后考虑放在 UI 里面展示
|
||||||
- [ ] 开发与生产环境隔离,构建生产镜像 <Badge type="info" text="0.4" />
|
- [ ] 开发与生产环境隔离,构建生产镜像 <Badge type="info" text="0.4" />
|
||||||
- [ ] 支持 MinerU 2.5 的解析方法 <Badge type="info" text="0.3.5" />
|
- [ ] 支持 MinerU 2.5 的解析方法 <Badge type="info" text="0.3.5" />
|
||||||
- [ ] 优化全局配置的管理模型,优化配置管理
|
|
||||||
|
|
||||||
## Later
|
## Later
|
||||||
|
|
||||||
@ -34,3 +32,5 @@
|
|||||||
- [x] 添加测试脚本,覆盖最常见的功能(已覆盖API)
|
- [x] 添加测试脚本,覆盖最常见的功能(已覆盖API)
|
||||||
- [x] 新建 tasker 模块,用来管理所有的后台任务,UI 上使用侧边栏管理。
|
- [x] 新建 tasker 模块,用来管理所有的后台任务,UI 上使用侧边栏管理。
|
||||||
- [x] 优化对文档信息的检索展示(检索结果页、详情页)
|
- [x] 优化对文档信息的检索展示(检索结果页、详情页)
|
||||||
|
- [x] 当前 ReAct 智能体有消息顺序错乱的 bug,且不会默认调用工具
|
||||||
|
- [x] 优化全局配置的管理模型,优化配置管理
|
||||||
|
|||||||
@ -39,7 +39,13 @@ default_model: siliconflow/deepseek-ai/DeepSeek-V3.2-Exp
|
|||||||
## 自定义模型供应商
|
## 自定义模型供应商
|
||||||
|
|
||||||
::: warning
|
::: warning
|
||||||
原本网页中的自定义模型已在 `0.3.x` 版本移除,请在 `src/config/static/models.yaml` 中按如下方式配置,并重启服务后选择并使用。此外,这里也推荐一下团队的另外一个小工具 [mvllm (Manage and Route vLLM Servers)](https://github.com/xerrors/mvllm)。
|
原本网页中的自定义模型已在 `0.3.x` 版本移除,请在 `src/config/static/models.py` 中按如下方式配置,并重启服务后选择并使用。此外,这里也推荐一下团队的另外一个小工具 [mvllm (Manage and Route vLLM Servers)](https://github.com/xerrors/mvllm)。
|
||||||
|
:::
|
||||||
|
|
||||||
|
::: tip 配置系统升级 (v0.3.x)
|
||||||
|
从 `v0.3.x` 版本开始,模型配置系统已升级为基于 Pydantic BaseModel 的类型安全配置,支持 TOML 格式的用户配置文件。
|
||||||
|
- **默认配置**: `src/config/static/models.py` (Python 代码)
|
||||||
|
- **用户配置**: `saves/config/base.toml` (TOML 格式,仅保存用户修改)
|
||||||
:::
|
:::
|
||||||
|
|
||||||
系统理论上兼容任何 OpenAI 兼容的模型服务,包括:
|
系统理论上兼容任何 OpenAI 兼容的模型服务,包括:
|
||||||
@ -52,52 +58,50 @@ default_model: siliconflow/deepseek-ai/DeepSeek-V3.2-Exp
|
|||||||
|
|
||||||
### 1. 编辑模型配置文件
|
### 1. 编辑模型配置文件
|
||||||
|
|
||||||
**方式一:修改默认配置**
|
**方式一:修改默认配置(推荐)**
|
||||||
编辑 `src/config/static/models.yaml` 文件
|
编辑 `src/config/static/models.py` 文件中的 `DEFAULT_CHAT_MODEL_PROVIDERS` 字典
|
||||||
|
|
||||||
**方式二:使用覆盖配置**
|
在 `src/config/static/models.py` 中添加新的模型供应商:
|
||||||
创建自定义配置文件并通过环境变量指定:
|
|
||||||
```bash
|
|
||||||
# 创建自定义配置文件
|
|
||||||
cp src/config/static/models.yaml /path/to/your/custom-models.yaml
|
|
||||||
|
|
||||||
# 设置环境变量
|
```python
|
||||||
export OVERRIDE_DEFAULT_MODELS_CONFIG_WITH=/path/to/your/custom-models.yaml
|
DEFAULT_CHAT_MODEL_PROVIDERS: dict[str, ChatModelProvider] = {
|
||||||
```
|
# ... 现有配置 ...
|
||||||
|
|
||||||
### 2. 添加模型配置
|
"custom-provider": ChatModelProvider(
|
||||||
|
name="自定义提供商",
|
||||||
|
url="https://your-provider.com/docs",
|
||||||
|
base_url="https://api.your-provider.com/v1",
|
||||||
|
default="custom-model-name",
|
||||||
|
env="CUSTOM_API_KEY_ENV_NAME",
|
||||||
|
models=[
|
||||||
|
"supported-model-name",
|
||||||
|
"another-model-name",
|
||||||
|
],
|
||||||
|
),
|
||||||
|
|
||||||
在配置文件中添加新的模型供应商:
|
# 本地 Ollama 服务
|
||||||
|
"local-ollama": ChatModelProvider(
|
||||||
|
name="Local Ollama",
|
||||||
|
url="https://ollama.com",
|
||||||
|
base_url="http://localhost:11434/v1",
|
||||||
|
default="llama3.2",
|
||||||
|
env="NO_API_KEY", # 对于不需要API Key的服务,使用NO_API_KEY
|
||||||
|
models=["llama3.2", "qwen2.5"],
|
||||||
|
),
|
||||||
|
|
||||||
```yaml
|
# 本地 vLLM 服务
|
||||||
custom-provider-name:
|
"local-vllm": ChatModelProvider(
|
||||||
name: custom-provider-name
|
name="Local vLLM",
|
||||||
default: custom-model-name
|
url="https://docs.vllm.ai",
|
||||||
base_url: "https://api.your-provider.com/v1"
|
base_url="http://localhost:8000/v1",
|
||||||
env: CUSTOM_API_KEY_ENV_NAME # 注意:现在是单个环境变量
|
default="Qwen/Qwen2.5-7B-Instruct",
|
||||||
models:
|
env="NO_API_KEY",
|
||||||
- supported-model-name
|
models=[
|
||||||
- another-model-name
|
"Qwen/Qwen2.5-7B-Instruct",
|
||||||
|
"Qwen/Qwen2.5-14B-Instruct",
|
||||||
# 本地 Ollama 服务
|
],
|
||||||
local-ollama:
|
),
|
||||||
name: Local Ollama
|
}
|
||||||
base_url: "http://localhost:11434/v1"
|
|
||||||
default: llama3.2
|
|
||||||
env: NO_API_KEY # 对于不需要API Key的服务,使用NO_API_KEY
|
|
||||||
models:
|
|
||||||
- llama3.2
|
|
||||||
- qwen2.5
|
|
||||||
|
|
||||||
# 本地 vLLM 服务
|
|
||||||
local-vllm:
|
|
||||||
name: Local vLLM
|
|
||||||
base_url: "http://localhost:8000/v1"
|
|
||||||
default: Qwen/Qwen2.5-7B-Instruct
|
|
||||||
env: NO_API_KEY
|
|
||||||
models:
|
|
||||||
- Qwen/Qwen2.5-7B-Instruct
|
|
||||||
- Qwen/Qwen2.5-14B-Instruct
|
|
||||||
```
|
```
|
||||||
|
|
||||||
### 3. 配置环境变量
|
### 3. 配置环境变量
|
||||||
@ -123,21 +127,58 @@ docker compose restart api-dev
|
|||||||
|
|
||||||
#### 1. 配置模型信息
|
#### 1. 配置模型信息
|
||||||
|
|
||||||
在 `src/config/static/models.yaml` 中或通过覆盖配置文件添加配置:
|
在 `src/config/static/models.py` 中的默认配置部分添加:
|
||||||
|
|
||||||
```yaml
|
```python
|
||||||
EMBED_MODEL_INFO:
|
# 默认嵌入模型配置
|
||||||
vllm/Qwen/Qwen3-Embedding-0.6B:
|
DEFAULT_EMBED_MODELS: dict[str, EmbedModelInfo] = {
|
||||||
name: Qwen/Qwen3-Embedding-0.6B
|
# ... 现有配置 ...
|
||||||
dimension: 1024
|
|
||||||
base_url: http://localhost:8000/v1/embeddings
|
|
||||||
api_key: no_api_key
|
|
||||||
|
|
||||||
RERANKER_LIST:
|
"vllm/Qwen/Qwen3-Embedding-0.6B": EmbedModelInfo(
|
||||||
vllm/BAAI/bge-reranker-v2-m3:
|
name="Qwen/Qwen3-Embedding-0.6B",
|
||||||
name: BAAI/bge-reranker-v2-m3
|
dimension=1024,
|
||||||
base_url: http://localhost:8000/v1/rerank
|
base_url="http://localhost:8000/v1/embeddings",
|
||||||
api_key: no_api_key
|
api_key="no_api_key",
|
||||||
|
),
|
||||||
|
}
|
||||||
|
|
||||||
|
# 默认重排序模型配置
|
||||||
|
DEFAULT_RERANKERS: dict[str, RerankerInfo] = {
|
||||||
|
# ... 现有配置 ...
|
||||||
|
|
||||||
|
"vllm/BAAI/bge-reranker-v2-m3": RerankerInfo(
|
||||||
|
name="BAAI/bge-reranker-v2-m3",
|
||||||
|
base_url="http://localhost:8000/v1/rerank",
|
||||||
|
api_key="no_api_key",
|
||||||
|
),
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
#### 2. 动态配置(可选)
|
||||||
|
|
||||||
|
你也可以通过代码动态添加本地模型:
|
||||||
|
|
||||||
|
```python
|
||||||
|
from src.config import config
|
||||||
|
from src.config.static.models import EmbedModelInfo, RerankerInfo
|
||||||
|
|
||||||
|
# 添加本地嵌入模型
|
||||||
|
config.embed_model_names["local/embed-model"] = EmbedModelInfo(
|
||||||
|
name="local-embed-model",
|
||||||
|
dimension=1024,
|
||||||
|
base_url="http://localhost:8000/v1/embeddings",
|
||||||
|
api_key="no_api_key",
|
||||||
|
)
|
||||||
|
|
||||||
|
# 添加本地重排序模型
|
||||||
|
config.reranker_names["local/reranker-model"] = RerankerInfo(
|
||||||
|
name="local-reranker-model",
|
||||||
|
base_url="http://localhost:8000/v1/rerank",
|
||||||
|
api_key="no_api_key",
|
||||||
|
)
|
||||||
|
|
||||||
|
# 保存配置
|
||||||
|
config.save()
|
||||||
```
|
```
|
||||||
|
|
||||||
#### 2. 启动模型服务
|
#### 2. 启动模型服务
|
||||||
@ -154,20 +195,4 @@ vllm serve BAAI/bge-reranker-v2-m3 \
|
|||||||
--task score \
|
--task score \
|
||||||
--dtype fp16 \
|
--dtype fp16 \
|
||||||
--port 8000
|
--port 8000
|
||||||
```
|
```
|
||||||
|
|
||||||
## 常见问题
|
|
||||||
|
|
||||||
**Q: 如何查看当前可用的模型?**
|
|
||||||
|
|
||||||
在 Web 界面的"设置"页面可以查看所有已配置的模型。
|
|
||||||
|
|
||||||
**Q: 模型配置不生效?**
|
|
||||||
|
|
||||||
1. 检查环境变量是否正确设置
|
|
||||||
2. 确认 API 密钥有效
|
|
||||||
3. 重启服务:`docker compose restart api-dev`
|
|
||||||
|
|
||||||
**Q: 如何测试模型连接?**
|
|
||||||
|
|
||||||
在 Web 界面的对话页面选择对应模型进行测试。
|
|
||||||
@ -57,6 +57,8 @@ dependencies = [
|
|||||||
"pymysql>=1.1.0",
|
"pymysql>=1.1.0",
|
||||||
"tenacity>=8.0.0",
|
"tenacity>=8.0.0",
|
||||||
"pypinyin>=0.55.0",
|
"pypinyin>=0.55.0",
|
||||||
|
"tomli",
|
||||||
|
"tomli-w",
|
||||||
]
|
]
|
||||||
[tool.ruff]
|
[tool.ruff]
|
||||||
line-length = 120 # 代码最大行宽
|
line-length = 120 # 代码最大行宽
|
||||||
|
|||||||
@ -393,9 +393,9 @@ async def get_chat_models(model_provider: str, current_user: User = Depends(get_
|
|||||||
@chat.post("/models/update")
|
@chat.post("/models/update")
|
||||||
async def update_chat_models(model_provider: str, model_names: list[str], current_user=Depends(get_admin_user)):
|
async def update_chat_models(model_provider: str, model_names: list[str], current_user=Depends(get_admin_user)):
|
||||||
"""更新指定模型提供商的模型列表 (仅管理员)"""
|
"""更新指定模型提供商的模型列表 (仅管理员)"""
|
||||||
conf.model_names[model_provider]["models"] = model_names
|
conf.model_names[model_provider].models = model_names
|
||||||
conf._save_models_to_file()
|
conf._save_models_to_file(model_provider)
|
||||||
return {"models": conf.model_names[model_provider]["models"]}
|
return {"models": conf.model_names[model_provider].models}
|
||||||
|
|
||||||
|
|
||||||
@chat.get("/tools")
|
@chat.get("/tools")
|
||||||
|
|||||||
@ -14,14 +14,17 @@ def load_chat_model(fully_specified_name: str, **kwargs) -> BaseChatModel:
|
|||||||
"""
|
"""
|
||||||
provider, model = fully_specified_name.split("/", maxsplit=1)
|
provider, model = fully_specified_name.split("/", maxsplit=1)
|
||||||
|
|
||||||
assert provider != "custom", "[弃用] 自定义模型已移除,请在 src/config/static/models.yaml 中配置"
|
assert provider != "custom", "[弃用] 自定义模型已移除,请在 src/config/static/models.py 中配置"
|
||||||
|
|
||||||
model_info = config.model_names.get(provider, {})
|
model_info = config.model_names.get(provider)
|
||||||
env_var = model_info["env"]
|
if not model_info:
|
||||||
|
raise ValueError(f"Unknown model provider: {provider}")
|
||||||
|
|
||||||
|
env_var = model_info.env
|
||||||
|
|
||||||
api_key = os.getenv(env_var, env_var)
|
api_key = os.getenv(env_var, env_var)
|
||||||
|
|
||||||
base_url = get_docker_safe_url(model_info["base_url"])
|
base_url = get_docker_safe_url(model_info.base_url)
|
||||||
|
|
||||||
if provider in ["deepseek", "dashscope"]:
|
if provider in ["deepseek", "dashscope"]:
|
||||||
from langchain_deepseek import ChatDeepSeek
|
from langchain_deepseek import ChatDeepSeek
|
||||||
|
|||||||
@ -1,246 +1,371 @@
|
|||||||
import json
|
"""
|
||||||
|
应用配置模块
|
||||||
|
|
||||||
|
使用 Pydantic BaseModel 实现配置管理,支持:
|
||||||
|
- 从 TOML 文件加载用户配置
|
||||||
|
- 仅保存用户修改过的配置项
|
||||||
|
- 默认配置定义在代码中
|
||||||
|
"""
|
||||||
|
|
||||||
import os
|
import os
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
import yaml
|
import tomli
|
||||||
|
import tomli_w
|
||||||
|
from pydantic import BaseModel, Field, field_validator
|
||||||
|
|
||||||
|
from src.config.static.models import (
|
||||||
|
DEFAULT_CHAT_MODEL_PROVIDERS,
|
||||||
|
DEFAULT_EMBED_MODELS,
|
||||||
|
DEFAULT_RERANKERS,
|
||||||
|
ChatModelProvider,
|
||||||
|
EmbedModelInfo,
|
||||||
|
RerankerInfo,
|
||||||
|
)
|
||||||
from src.utils.logging_config import logger
|
from src.utils.logging_config import logger
|
||||||
|
|
||||||
|
|
||||||
class SimpleConfig(dict):
|
class Config(BaseModel):
|
||||||
def __key(self, key):
|
"""应用配置类"""
|
||||||
return "" if key is None else key # 目前忘记了这里为什么要 lower 了,只能说配置项最好不要有大写的
|
|
||||||
|
|
||||||
def __str__(self):
|
# ============================================================
|
||||||
return json.dumps(self)
|
# 基础配置
|
||||||
|
# ============================================================
|
||||||
|
save_dir: str = Field(default="saves", description="保存目录")
|
||||||
|
model_dir: str = Field(default="", description="本地模型目录")
|
||||||
|
|
||||||
def __setattr__(self, key, value):
|
# ============================================================
|
||||||
self[self.__key(key)] = value
|
# 功能开关
|
||||||
|
# ============================================================
|
||||||
|
enable_reranker: bool = Field(default=False, description="是否开启重排序")
|
||||||
|
enable_content_guard: bool = Field(default=False, description="是否启用内容审查")
|
||||||
|
enable_content_guard_llm: bool = Field(default=False, description="是否启用LLM内容审查")
|
||||||
|
enable_web_search: bool = Field(default=False, description="是否启用网络搜索")
|
||||||
|
|
||||||
def __getattr__(self, key):
|
# ============================================================
|
||||||
return self.get(self.__key(key))
|
# 模型配置
|
||||||
|
# ============================================================
|
||||||
|
default_model: str = Field(
|
||||||
|
default="siliconflow/deepseek-ai/DeepSeek-V3.2-Exp",
|
||||||
|
description="默认对话模型",
|
||||||
|
)
|
||||||
|
fast_model: str = Field(
|
||||||
|
default="siliconflow/THUDM/GLM-4-9B-0414",
|
||||||
|
description="快速响应模型",
|
||||||
|
)
|
||||||
|
embed_model: str = Field(
|
||||||
|
default="siliconflow/BAAI/bge-m3",
|
||||||
|
description="Embedding 模型",
|
||||||
|
)
|
||||||
|
reranker: str = Field(
|
||||||
|
default="siliconflow/BAAI/bge-reranker-v2-m3",
|
||||||
|
description="Re-Ranker 模型",
|
||||||
|
)
|
||||||
|
content_guard_llm_model: str = Field(
|
||||||
|
default="siliconflow/Qwen/Qwen3-235B-A22B-Instruct-2507",
|
||||||
|
description="内容审查LLM模型",
|
||||||
|
)
|
||||||
|
|
||||||
def __getitem__(self, key):
|
# ============================================================
|
||||||
return self.get(self.__key(key))
|
# 智能体配置
|
||||||
|
# ============================================================
|
||||||
|
default_agent_id: str = Field(default="", description="默认智能体ID")
|
||||||
|
|
||||||
def __setitem__(self, key, value):
|
# ============================================================
|
||||||
return super().__setitem__(self.__key(key), value)
|
# 模型信息(只读,不持久化)
|
||||||
|
# ============================================================
|
||||||
|
model_names: dict[str, ChatModelProvider] = Field(
|
||||||
|
default_factory=lambda: DEFAULT_CHAT_MODEL_PROVIDERS.copy(),
|
||||||
|
description="聊天模型提供商配置",
|
||||||
|
exclude=True,
|
||||||
|
)
|
||||||
|
embed_model_names: dict[str, EmbedModelInfo] = Field(
|
||||||
|
default_factory=lambda: DEFAULT_EMBED_MODELS.copy(),
|
||||||
|
description="嵌入模型配置",
|
||||||
|
exclude=True,
|
||||||
|
)
|
||||||
|
reranker_names: dict[str, RerankerInfo] = Field(
|
||||||
|
default_factory=lambda: DEFAULT_RERANKERS.copy(),
|
||||||
|
description="重排序模型配置",
|
||||||
|
exclude=True,
|
||||||
|
)
|
||||||
|
|
||||||
def __dict__(self):
|
# ============================================================
|
||||||
return {k: v for k, v in self.items()}
|
# 运行时状态(不持久化)
|
||||||
|
# ============================================================
|
||||||
|
model_provider_status: dict[str, bool] = Field(
|
||||||
|
default_factory=dict,
|
||||||
|
description="模型提供商可用状态",
|
||||||
|
exclude=True,
|
||||||
|
)
|
||||||
|
valuable_model_provider: list[str] = Field(
|
||||||
|
default_factory=list,
|
||||||
|
description="可用的模型提供商列表",
|
||||||
|
exclude=True,
|
||||||
|
)
|
||||||
|
|
||||||
def update(self, other):
|
# 内部状态
|
||||||
for key, value in other.items():
|
_config_file: Path | None = None
|
||||||
self[key] = value
|
_user_modified_fields: set[str] = set()
|
||||||
|
_modified_providers: set[str] = set() # 记录具体修改的模型提供商
|
||||||
|
|
||||||
|
model_config = {"arbitrary_types_allowed": True, "extra": "allow"}
|
||||||
|
|
||||||
class Config(SimpleConfig):
|
def __init__(self, **data):
|
||||||
def __init__(self):
|
super().__init__(**data)
|
||||||
super().__init__()
|
self._setup_paths()
|
||||||
self._config_items = {}
|
self._load_user_config()
|
||||||
self.save_dir = os.getenv("SAVE_DIR", "saves")
|
self._handle_environment()
|
||||||
self.filename = str(Path(f"{self.save_dir}/config/base.yaml"))
|
|
||||||
os.makedirs(os.path.dirname(self.filename), exist_ok=True)
|
|
||||||
|
|
||||||
self._models_config_path: Path | None = os.getenv("OVERRIDE_DEFAULT_MODELS_CONFIG_WITH")
|
def _setup_paths(self):
|
||||||
self._update_models_from_file()
|
"""设置配置文件路径"""
|
||||||
|
self.save_dir = os.getenv("SAVE_DIR", self.save_dir)
|
||||||
|
self._config_file = Path(self.save_dir) / "config" / "base.toml"
|
||||||
|
self._config_file.parent.mkdir(parents=True, exist_ok=True)
|
||||||
|
|
||||||
### >>> 默认配置
|
def _load_user_config(self):
|
||||||
# 功能选项
|
"""从 TOML 文件加载用户配置"""
|
||||||
self.add_item("enable_reranker", default=False, des="是否开启重排序")
|
if not self._config_file or not self._config_file.exists():
|
||||||
self.add_item("enable_content_guard", default=False, des="是否启用内容审查")
|
logger.info(f"Config file not found, using defaults: {self._config_file}")
|
||||||
self.add_item("enable_content_guard_llm", default=False, des="是否启用LLM内容审查")
|
return
|
||||||
self.add_item(
|
|
||||||
"content_guard_llm_model", default="siliconflow/Qwen/Qwen3-235B-A22B-Instruct-2507", des="内容审查LLM模型"
|
|
||||||
)
|
|
||||||
# 默认智能体配置
|
|
||||||
self.add_item("default_agent_id", default="", des="默认智能体ID")
|
|
||||||
# 模型配置
|
|
||||||
## 注意这里是模型名,而不是具体的模型路径,默认使用 HuggingFace 的路径
|
|
||||||
## 如果需要自定义本地模型路径,则在 .env 中配置 MODEL_DIR
|
|
||||||
self.add_item(
|
|
||||||
"default_model",
|
|
||||||
default=self._get_default_chat_model_spec(),
|
|
||||||
des="默认对话模型",
|
|
||||||
)
|
|
||||||
self.add_item(
|
|
||||||
"fast_model",
|
|
||||||
default="siliconflow/THUDM/GLM-4-9B-0414",
|
|
||||||
des="快速响应模型",
|
|
||||||
)
|
|
||||||
|
|
||||||
self.add_item(
|
logger.info(f"Loading config from {self._config_file}")
|
||||||
"embed_model",
|
try:
|
||||||
default="siliconflow/BAAI/bge-m3",
|
with open(self._config_file, "rb") as f:
|
||||||
des="Embedding 模型",
|
user_config = tomli.load(f)
|
||||||
choices=list(self.embed_model_names.keys()),
|
|
||||||
)
|
|
||||||
self.add_item(
|
|
||||||
"reranker",
|
|
||||||
default="siliconflow/BAAI/bge-reranker-v2-m3",
|
|
||||||
des="Re-Ranker 模型",
|
|
||||||
choices=list(self.reranker_names.keys()),
|
|
||||||
) # noqa: E501
|
|
||||||
### <<< 默认配置结束
|
|
||||||
|
|
||||||
self.load()
|
# 记录用户修改的字段
|
||||||
# 清理已废弃的配置项
|
self._user_modified_fields = set(user_config.keys())
|
||||||
self.pop("model_provider", None)
|
|
||||||
self.pop("model_name", None)
|
|
||||||
self.handle_self()
|
|
||||||
|
|
||||||
def add_item(self, key, default, des=None, choices=None):
|
# 更新配置
|
||||||
self.__setattr__(key, default)
|
for key, value in user_config.items():
|
||||||
self._config_items[key] = {"default": default, "des": des, "choices": choices}
|
if key == "model_names":
|
||||||
|
# 特殊处理模型配置
|
||||||
|
self._load_model_names(value)
|
||||||
|
elif hasattr(self, key):
|
||||||
|
setattr(self, key, value)
|
||||||
|
else:
|
||||||
|
logger.warning(f"Unknown config key: {key}")
|
||||||
|
|
||||||
def __dict__(self):
|
except Exception as e:
|
||||||
blocklist = [
|
logger.error(f"Failed to load config from {self._config_file}: {e}")
|
||||||
"_config_items",
|
|
||||||
"model_names",
|
|
||||||
"model_provider_status",
|
|
||||||
"embed_model_names",
|
|
||||||
"reranker_names",
|
|
||||||
"_models_config_path",
|
|
||||||
]
|
|
||||||
return {k: v for k, v in self.items() if k not in blocklist}
|
|
||||||
|
|
||||||
def _update_models_from_file(self):
|
def _load_model_names(self, model_names_data):
|
||||||
"""
|
"""加载用户自定义的模型配置"""
|
||||||
从 models.yaml 或覆盖配置文件中更新 MODEL_NAMES
|
try:
|
||||||
"""
|
for provider, provider_data in model_names_data.items():
|
||||||
# 检查是否设置了覆盖配置文件的环境变量
|
if provider in self.model_names:
|
||||||
override_config_path = os.getenv("OVERRIDE_DEFAULT_MODELS_CONFIG_WITH")
|
# 更新现有提供商的模型列表
|
||||||
|
if "models" in provider_data:
|
||||||
if override_config_path and os.path.exists(override_config_path):
|
self.model_names[provider].models = provider_data["models"]
|
||||||
config_file = Path(override_config_path)
|
else:
|
||||||
logger.info(f"Using override models config from: {override_config_path}")
|
# 添加新的提供商
|
||||||
else:
|
self.model_names[provider] = ChatModelProvider(**provider_data)
|
||||||
config_file = Path("src/config/static/models.yaml")
|
logger.info(f"Loaded custom model configurations for {len(model_names_data)} providers")
|
||||||
logger.info("Using default models config")
|
except Exception as e:
|
||||||
|
logger.error(f"Failed to load model names: {e}")
|
||||||
self._models_config_path = str(config_file)
|
|
||||||
|
|
||||||
with open(self._models_config_path, encoding="utf-8") as f:
|
|
||||||
_models = yaml.safe_load(f)
|
|
||||||
|
|
||||||
self.model_names = _models["MODEL_NAMES"]
|
|
||||||
self.embed_model_names = _models["EMBED_MODEL_INFO"]
|
|
||||||
self.reranker_names = _models["RERANKER_LIST"]
|
|
||||||
|
|
||||||
def _save_models_to_file(self):
|
|
||||||
"""
|
|
||||||
将当前模型配置写回模型配置文件
|
|
||||||
"""
|
|
||||||
if self._models_config_path is None:
|
|
||||||
self._models_config_path = str(Path("src/config/static/models.yaml"))
|
|
||||||
|
|
||||||
models_payload = {
|
|
||||||
"MODEL_NAMES": self.model_names,
|
|
||||||
"EMBED_MODEL_INFO": self.embed_model_names,
|
|
||||||
"RERANKER_LIST": self.reranker_names,
|
|
||||||
}
|
|
||||||
|
|
||||||
with open(self._models_config_path, "w", encoding="utf-8") as f:
|
|
||||||
yaml.safe_dump(models_payload, f, indent=2, allow_unicode=True, sort_keys=False)
|
|
||||||
|
|
||||||
def _get_default_chat_model_spec(self):
|
|
||||||
"""选择一个默认的聊天模型,优先使用 siliconflow 的默认模型"""
|
|
||||||
preferred_provider = "siliconflow"
|
|
||||||
provider_info = (self.model_names or {}).get(preferred_provider)
|
|
||||||
if provider_info:
|
|
||||||
default_model = provider_info.get("default")
|
|
||||||
if default_model:
|
|
||||||
return f"{preferred_provider}/{default_model}"
|
|
||||||
|
|
||||||
for provider, info in (self.model_names or {}).items():
|
|
||||||
default_model = info.get("default")
|
|
||||||
if default_model:
|
|
||||||
return f"{provider}/{default_model}"
|
|
||||||
|
|
||||||
return ""
|
|
||||||
|
|
||||||
def handle_self(self):
|
|
||||||
"""
|
|
||||||
处理配置
|
|
||||||
"""
|
|
||||||
self.model_dir = os.environ.get("MODEL_DIR", "")
|
|
||||||
|
|
||||||
|
def _handle_environment(self):
|
||||||
|
"""处理环境变量和运行时状态"""
|
||||||
|
# 处理模型目录
|
||||||
|
self.model_dir = os.environ.get("MODEL_DIR", self.model_dir)
|
||||||
if self.model_dir:
|
if self.model_dir:
|
||||||
if os.path.exists(self.model_dir):
|
if os.path.exists(self.model_dir):
|
||||||
logger.debug(
|
logger.debug(
|
||||||
f"The model directory ({self.model_dir}) "
|
f"Model directory ({self.model_dir}) contains: {os.listdir(self.model_dir)}"
|
||||||
f"contains the following folders: {os.listdir(self.model_dir)}"
|
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
f"Warning: The model directory ({self.model_dir}) does not exist. "
|
f"Model directory ({self.model_dir}) does not exist. "
|
||||||
"If not configured, please ignore it. "
|
"If not configured, please ignore it."
|
||||||
"If configured, please check if the configuration is correct; "
|
|
||||||
"For example, the mapping in the docker-compose file"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
# 检查模型提供商的环境变量
|
# 检查模型提供商的环境变量
|
||||||
self.model_provider_status = {}
|
self.model_provider_status = {}
|
||||||
for provider in self.model_names:
|
for provider, info in self.model_names.items():
|
||||||
env_var = self.model_names[provider]["env"]
|
env_var = info.env
|
||||||
# 如果环境变量名为 NO_API_KEY,则认为总是可用
|
|
||||||
if env_var == "NO_API_KEY":
|
if env_var == "NO_API_KEY":
|
||||||
self.model_provider_status[provider] = True
|
self.model_provider_status[provider] = True
|
||||||
else:
|
else:
|
||||||
self.model_provider_status[provider] = bool(os.getenv(env_var))
|
self.model_provider_status[provider] = bool(os.getenv(env_var))
|
||||||
|
|
||||||
|
# 检查网络搜索
|
||||||
if os.getenv("TAVILY_API_KEY"):
|
if os.getenv("TAVILY_API_KEY"):
|
||||||
self.enable_web_search = True
|
self.enable_web_search = True
|
||||||
|
|
||||||
self.valuable_model_provider = [k for k, v in self.model_provider_status.items() if v]
|
# 获取可用的模型提供商
|
||||||
assert len(self.valuable_model_provider) > 0, "No model provider available, please check your `.env` file."
|
self.valuable_model_provider = [
|
||||||
|
k for k, v in self.model_provider_status.items() if v
|
||||||
|
]
|
||||||
|
|
||||||
def load(self):
|
if not self.valuable_model_provider:
|
||||||
"""根据传入的文件覆盖掉默认配置"""
|
raise ValueError(
|
||||||
logger.info(f"Loading config from {self.filename}")
|
"No model provider available, please check your `.env` file."
|
||||||
if self.filename is not None and os.path.exists(self.filename):
|
)
|
||||||
if self.filename.endswith(".json"):
|
|
||||||
with open(self.filename) as f:
|
|
||||||
content = f.read()
|
|
||||||
if content:
|
|
||||||
local_config = json.loads(content)
|
|
||||||
self.update(local_config)
|
|
||||||
else:
|
|
||||||
print(f"{self.filename} is empty.")
|
|
||||||
|
|
||||||
elif self.filename.endswith(".yaml"):
|
|
||||||
with open(self.filename) as f:
|
|
||||||
content = f.read()
|
|
||||||
if content:
|
|
||||||
local_config = yaml.safe_load(content)
|
|
||||||
self.update(local_config)
|
|
||||||
else:
|
|
||||||
print(f"{self.filename} is empty.")
|
|
||||||
else:
|
|
||||||
logger.warning(f"Unknown config file type {self.filename}")
|
|
||||||
|
|
||||||
def save(self):
|
def save(self):
|
||||||
logger.info(f"Saving config to {self.filename}")
|
"""保存配置到 TOML 文件(仅保存用户修改的字段)"""
|
||||||
if self.filename is None:
|
if not self._config_file:
|
||||||
logger.warning("Config file is not specified, save to default config/base.yaml")
|
logger.warning("Config file path not set")
|
||||||
self.filename = os.path.join(self.save_dir, "config", "base.yaml")
|
return
|
||||||
os.makedirs(os.path.dirname(self.filename), exist_ok=True)
|
|
||||||
|
|
||||||
if self.filename.endswith(".json"):
|
logger.info(f"Saving config to {self._config_file}")
|
||||||
with open(self.filename, "w+") as f:
|
|
||||||
json.dump(self.__dict__(), f, indent=4, ensure_ascii=False)
|
|
||||||
elif self.filename.endswith(".yaml"):
|
|
||||||
with open(self.filename, "w+") as f:
|
|
||||||
yaml.dump(self.__dict__(), f, indent=2, allow_unicode=True)
|
|
||||||
else:
|
|
||||||
logger.warning(f"Unknown config file type {self.filename}, save as json")
|
|
||||||
with open(self.filename, "w+") as f:
|
|
||||||
json.dump(self, f, indent=4)
|
|
||||||
|
|
||||||
logger.info(f"Config file {self.filename} saved")
|
# 获取默认配置
|
||||||
|
default_config = Config.model_construct()
|
||||||
|
|
||||||
def dump_config(self):
|
# 对比当前配置和默认配置,找出用户修改的字段
|
||||||
return json.loads(str(self))
|
user_modified = {}
|
||||||
|
for field_name in self.model_fields.keys():
|
||||||
|
# 跳过 exclude=True 的字段
|
||||||
|
field_info = self.model_fields[field_name]
|
||||||
|
if field_info.exclude:
|
||||||
|
continue
|
||||||
|
|
||||||
|
current_value = getattr(self, field_name)
|
||||||
|
default_value = getattr(default_config, field_name)
|
||||||
|
|
||||||
|
# 如果值不同,说明用户修改了
|
||||||
|
if current_value != default_value:
|
||||||
|
user_modified[field_name] = current_value
|
||||||
|
|
||||||
|
# 写入 TOML 文件
|
||||||
|
try:
|
||||||
|
with open(self._config_file, "wb") as f:
|
||||||
|
tomli_w.dump(user_modified, f)
|
||||||
|
logger.info(f"Config saved to {self._config_file}")
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Failed to save config to {self._config_file}: {e}")
|
||||||
|
|
||||||
|
def dump_config(self) -> dict[str, Any]:
|
||||||
|
"""导出配置为字典(用于 API 返回)"""
|
||||||
|
config_dict = self.model_dump(
|
||||||
|
exclude={
|
||||||
|
"model_names",
|
||||||
|
"embed_model_names",
|
||||||
|
"reranker_names",
|
||||||
|
"model_provider_status",
|
||||||
|
"valuable_model_provider",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
# 添加模型信息(转换为字典格式供前端使用)
|
||||||
|
config_dict["model_names"] = {
|
||||||
|
provider: info.model_dump() for provider, info in self.model_names.items()
|
||||||
|
}
|
||||||
|
config_dict["embed_model_names"] = {
|
||||||
|
model_id: info.model_dump() for model_id, info in self.embed_model_names.items()
|
||||||
|
}
|
||||||
|
config_dict["reranker_names"] = {
|
||||||
|
model_id: info.model_dump() for model_id, info in self.reranker_names.items()
|
||||||
|
}
|
||||||
|
|
||||||
|
# 添加运行时状态信息
|
||||||
|
config_dict["model_provider_status"] = self.model_provider_status
|
||||||
|
config_dict["valuable_model_provider"] = self.valuable_model_provider
|
||||||
|
|
||||||
|
fields_info = {}
|
||||||
|
for field_name, field_info in Config.model_fields.items():
|
||||||
|
if not field_info.exclude: # 排除内部字段
|
||||||
|
fields_info[field_name] = {
|
||||||
|
'des': field_info.description,
|
||||||
|
'default': field_info.default,
|
||||||
|
'type': field_info.annotation.__name__ if hasattr(field_info.annotation, '__name__') else str(field_info.annotation),
|
||||||
|
'exclude': field_info.exclude if hasattr(field_info, 'exclude') else False,
|
||||||
|
}
|
||||||
|
config_dict["_config_items"] = fields_info
|
||||||
|
|
||||||
|
return config_dict
|
||||||
|
|
||||||
|
def get_model_choices(self) -> list[str]:
|
||||||
|
"""获取所有可用的聊天模型列表"""
|
||||||
|
choices = []
|
||||||
|
for provider, info in self.model_names.items():
|
||||||
|
if self.model_provider_status.get(provider, False):
|
||||||
|
for model in info.models:
|
||||||
|
choices.append(f"{provider}/{model}")
|
||||||
|
return choices
|
||||||
|
|
||||||
|
def get_embed_model_choices(self) -> list[str]:
|
||||||
|
"""获取所有可用的嵌入模型列表"""
|
||||||
|
return list(self.embed_model_names.keys())
|
||||||
|
|
||||||
|
def get_reranker_choices(self) -> list[str]:
|
||||||
|
"""获取所有可用的重排序模型列表"""
|
||||||
|
return list(self.reranker_names.keys())
|
||||||
|
|
||||||
|
# ============================================================
|
||||||
|
# 兼容旧代码的方法
|
||||||
|
# ============================================================
|
||||||
|
|
||||||
|
def __getitem__(self, key: str) -> Any:
|
||||||
|
"""支持字典式访问 config[key]"""
|
||||||
|
logger.warning("Using deprecated dict-style access for Config. "
|
||||||
|
"Please use attribute access instead.")
|
||||||
|
return getattr(self, key, None)
|
||||||
|
|
||||||
|
def __setitem__(self, key: str, value: Any):
|
||||||
|
"""支持字典式赋值 config[key] = value"""
|
||||||
|
logger.warning("Using deprecated dict-style assignment for Config. "
|
||||||
|
"Please use attribute access instead.")
|
||||||
|
setattr(self, key, value)
|
||||||
|
|
||||||
|
def update(self, other: dict):
|
||||||
|
"""批量更新配置(兼容旧代码)"""
|
||||||
|
for key, value in other.items():
|
||||||
|
if hasattr(self, key):
|
||||||
|
setattr(self, key, value)
|
||||||
|
else:
|
||||||
|
logger.warning(f"Unknown config key: {key}")
|
||||||
|
|
||||||
|
def _save_models_to_file(self, provider_name: str = None):
|
||||||
|
"""保存模型配置到主配置文件
|
||||||
|
|
||||||
|
Args:
|
||||||
|
provider_name: 如果提供,只保存特定provider的修改;否则保存所有model_names
|
||||||
|
"""
|
||||||
|
if not self._config_file:
|
||||||
|
logger.warning("Config file path not set")
|
||||||
|
return
|
||||||
|
|
||||||
|
logger.info(f"Saving models config to {self._config_file}")
|
||||||
|
|
||||||
|
try:
|
||||||
|
# 读取现有配置
|
||||||
|
user_config = {}
|
||||||
|
if self._config_file.exists():
|
||||||
|
with open(self._config_file, "rb") as f:
|
||||||
|
user_config = tomli.load(f)
|
||||||
|
|
||||||
|
# 初始化 model_names 配置(如果不存在)
|
||||||
|
if "model_names" not in user_config:
|
||||||
|
user_config["model_names"] = {}
|
||||||
|
|
||||||
|
if provider_name:
|
||||||
|
# 只保存特定 provider 的修改
|
||||||
|
if provider_name in self.model_names:
|
||||||
|
user_config["model_names"][provider_name] = self.model_names[provider_name].model_dump()
|
||||||
|
# 记录具体修改的 provider
|
||||||
|
self._modified_providers.add(provider_name)
|
||||||
|
logger.info(f"Saved models config for provider: {provider_name}")
|
||||||
|
else:
|
||||||
|
# 保存所有 model_names
|
||||||
|
user_config["model_names"] = {
|
||||||
|
provider: info.model_dump()
|
||||||
|
for provider, info in self.model_names.items()
|
||||||
|
}
|
||||||
|
# 记录整个 model_names 字段的修改
|
||||||
|
self._user_modified_fields.add("model_names")
|
||||||
|
logger.info("Saved all models config")
|
||||||
|
|
||||||
|
# 写入配置文件
|
||||||
|
with open(self._config_file, "wb") as f:
|
||||||
|
tomli_w.dump(user_config, f)
|
||||||
|
logger.info(f"Models config saved to {self._config_file}")
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Failed to save models config to {self._config_file}: {e}")
|
||||||
|
|
||||||
|
|
||||||
|
# 全局配置实例
|
||||||
config = Config()
|
config = Config()
|
||||||
|
|||||||
187
src/config/static/models.py
Normal file
187
src/config/static/models.py
Normal file
@ -0,0 +1,187 @@
|
|||||||
|
"""
|
||||||
|
默认模型配置
|
||||||
|
|
||||||
|
该文件定义了系统支持的所有默认模型配置,包括:
|
||||||
|
- 聊天模型(LLM)
|
||||||
|
- 嵌入模型(Embedding)
|
||||||
|
- 重排序模型(Reranker)
|
||||||
|
"""
|
||||||
|
|
||||||
|
from pydantic import BaseModel, Field
|
||||||
|
|
||||||
|
|
||||||
|
class ChatModelProvider(BaseModel):
|
||||||
|
"""聊天模型提供商配置"""
|
||||||
|
|
||||||
|
name: str = Field(..., description="提供商显示名称")
|
||||||
|
url: str = Field(..., description="提供商文档或模型列表 URL")
|
||||||
|
base_url: str = Field(..., description="API 基础 URL")
|
||||||
|
default: str = Field(..., description="默认模型名称")
|
||||||
|
env: str = Field(..., description="API Key 环境变量名")
|
||||||
|
models: list[str] = Field(default_factory=list, description="支持的模型列表")
|
||||||
|
|
||||||
|
|
||||||
|
class EmbedModelInfo(BaseModel):
|
||||||
|
"""嵌入模型配置"""
|
||||||
|
|
||||||
|
name: str = Field(..., description="模型名称")
|
||||||
|
dimension: int = Field(..., description="向量维度")
|
||||||
|
base_url: str = Field(..., description="API 基础 URL")
|
||||||
|
api_key: str = Field(..., description="API Key 或环境变量名")
|
||||||
|
|
||||||
|
|
||||||
|
class RerankerInfo(BaseModel):
|
||||||
|
"""重排序模型配置"""
|
||||||
|
|
||||||
|
name: str = Field(..., description="模型名称")
|
||||||
|
base_url: str = Field(..., description="API 基础 URL")
|
||||||
|
api_key: str = Field(..., description="API Key 或环境变量名")
|
||||||
|
|
||||||
|
|
||||||
|
# ============================================================
|
||||||
|
# 默认聊天模型配置
|
||||||
|
# ============================================================
|
||||||
|
|
||||||
|
DEFAULT_CHAT_MODEL_PROVIDERS: dict[str, ChatModelProvider] = {
|
||||||
|
"openai": ChatModelProvider(
|
||||||
|
name="OpenAI",
|
||||||
|
url="https://platform.openai.com/docs/models",
|
||||||
|
base_url="https://api.openai.com/v1",
|
||||||
|
default="gpt-4o-mini",
|
||||||
|
env="OPENAI_API_KEY",
|
||||||
|
models=["gpt-4", "gpt-4o", "gpt-4o-mini"],
|
||||||
|
),
|
||||||
|
"deepseek": ChatModelProvider(
|
||||||
|
name="DeepSeek",
|
||||||
|
url="https://platform.deepseek.com/api-docs/zh-cn/pricing",
|
||||||
|
base_url="https://api.deepseek.com/v1",
|
||||||
|
default="deepseek-chat",
|
||||||
|
env="DEEPSEEK_API_KEY",
|
||||||
|
models=["deepseek-chat", "deepseek-reasoner"],
|
||||||
|
),
|
||||||
|
"zhipu": ChatModelProvider(
|
||||||
|
name="智谱AI (Zhipu)",
|
||||||
|
url="https://open.bigmodel.cn/dev/api",
|
||||||
|
base_url="https://open.bigmodel.cn/api/paas/v4/",
|
||||||
|
default="glm-4.5-flash",
|
||||||
|
env="ZHIPUAI_API_KEY",
|
||||||
|
models=["glm-4.6", "glm-4.5-air", "glm-4.5-flash"],
|
||||||
|
),
|
||||||
|
"siliconflow": ChatModelProvider(
|
||||||
|
name="SiliconFlow",
|
||||||
|
url="https://cloud.siliconflow.cn/models",
|
||||||
|
base_url="https://api.siliconflow.cn/v1",
|
||||||
|
default="deepseek-ai/DeepSeek-V3.2-Exp",
|
||||||
|
env="SILICONFLOW_API_KEY",
|
||||||
|
models=[
|
||||||
|
"deepseek-ai/DeepSeek-V3.2-Exp",
|
||||||
|
"Qwen/Qwen3-235B-A22B-Thinking-2507",
|
||||||
|
"Qwen/Qwen3-235B-A22B-Instruct-2507",
|
||||||
|
"moonshotai/Kimi-K2-Instruct-0905",
|
||||||
|
"zai-org/GLM-4.6",
|
||||||
|
],
|
||||||
|
),
|
||||||
|
"together.ai": ChatModelProvider(
|
||||||
|
name="Together.ai",
|
||||||
|
url="https://api.together.ai/models",
|
||||||
|
base_url="https://api.together.xyz/v1/",
|
||||||
|
default="meta-llama/Llama-3.3-70B-Instruct-Turbo-Free",
|
||||||
|
env="TOGETHER_API_KEY",
|
||||||
|
models=["meta-llama/Llama-3.3-70B-Instruct-Turbo-Free"],
|
||||||
|
),
|
||||||
|
"dashscope": ChatModelProvider(
|
||||||
|
name="阿里百炼 (DashScope)",
|
||||||
|
url="https://bailian.console.aliyun.com/?switchAgent=10226727&productCode=p_efm#/model-market",
|
||||||
|
base_url="https://dashscope.aliyuncs.com/compatible-mode/v1",
|
||||||
|
default="qwen-max-latest",
|
||||||
|
env="DASHSCOPE_API_KEY",
|
||||||
|
models=[
|
||||||
|
"qwen-max-latest",
|
||||||
|
"qwen-plus-latest",
|
||||||
|
"qwen-turbo-latest",
|
||||||
|
"qwen3-235b-a22b-thinking-2507",
|
||||||
|
"qwen3-235b-a22b-instruct-2507",
|
||||||
|
],
|
||||||
|
),
|
||||||
|
"ark": ChatModelProvider(
|
||||||
|
name="豆包(Ark)",
|
||||||
|
url="https://console.volcengine.com/ark/region:ark+cn-beijing/model",
|
||||||
|
base_url="https://ark.cn-beijing.volces.com/api/v3",
|
||||||
|
default="doubao-seed-1-6-250615",
|
||||||
|
env="ARK_API_KEY",
|
||||||
|
models=[
|
||||||
|
"doubao-seed-1-6-250615",
|
||||||
|
"doubao-seed-1-6-thinking-250715",
|
||||||
|
"doubao-seed-1-6-flash-250715",
|
||||||
|
],
|
||||||
|
),
|
||||||
|
"openrouter": ChatModelProvider(
|
||||||
|
name="OpenRouter",
|
||||||
|
url="https://openrouter.ai/models",
|
||||||
|
base_url="https://openrouter.ai/api/v1",
|
||||||
|
default="openai/gpt-4o",
|
||||||
|
env="OPENROUTER_API_KEY",
|
||||||
|
models=[
|
||||||
|
"openai/gpt-4o",
|
||||||
|
"x-ai/grok-4",
|
||||||
|
"google/gemini-2.5-pro",
|
||||||
|
"anthropic/claude-sonnet-4",
|
||||||
|
],
|
||||||
|
),
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
# ============================================================
|
||||||
|
# 默认嵌入模型配置
|
||||||
|
# ============================================================
|
||||||
|
|
||||||
|
DEFAULT_EMBED_MODELS: dict[str, EmbedModelInfo] = {
|
||||||
|
"siliconflow/BAAI/bge-m3": EmbedModelInfo(
|
||||||
|
name="BAAI/bge-m3",
|
||||||
|
dimension=1024,
|
||||||
|
base_url="https://api.siliconflow.cn/v1/embeddings",
|
||||||
|
api_key="SILICONFLOW_API_KEY",
|
||||||
|
),
|
||||||
|
"siliconflow/Qwen/Qwen3-Embedding-0.6B": EmbedModelInfo(
|
||||||
|
name="Qwen/Qwen3-Embedding-0.6B",
|
||||||
|
dimension=1024,
|
||||||
|
base_url="https://api.siliconflow.cn/v1/embeddings",
|
||||||
|
api_key="SILICONFLOW_API_KEY",
|
||||||
|
),
|
||||||
|
"vllm/Qwen/Qwen3-Embedding-0.6B": EmbedModelInfo(
|
||||||
|
name="Qwen3-Embedding-0.6B",
|
||||||
|
dimension=1024,
|
||||||
|
base_url="http://localhost:8000/v1/embeddings",
|
||||||
|
api_key="no_api_key",
|
||||||
|
),
|
||||||
|
"ollama/nomic-embed-text": EmbedModelInfo(
|
||||||
|
name="nomic-embed-text",
|
||||||
|
dimension=768,
|
||||||
|
base_url="http://localhost:11434/api/embed",
|
||||||
|
api_key="no_api_key",
|
||||||
|
),
|
||||||
|
"ollama/bge-m3": EmbedModelInfo(
|
||||||
|
name="bge-m3",
|
||||||
|
dimension=1024,
|
||||||
|
base_url="http://localhost:11434/api/embed",
|
||||||
|
api_key="no_api_key",
|
||||||
|
),
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
# ============================================================
|
||||||
|
# 默认重排序模型配置
|
||||||
|
# ============================================================
|
||||||
|
|
||||||
|
DEFAULT_RERANKERS: dict[str, RerankerInfo] = {
|
||||||
|
"siliconflow/BAAI/bge-reranker-v2-m3": RerankerInfo(
|
||||||
|
name="BAAI/bge-reranker-v2-m3",
|
||||||
|
base_url="https://api.siliconflow.cn/v1/rerank",
|
||||||
|
api_key="SILICONFLOW_API_KEY",
|
||||||
|
),
|
||||||
|
"vllm/BAAI/bge-reranker-v2-m3": RerankerInfo(
|
||||||
|
name="BAAI/bge-reranker-v2-m3",
|
||||||
|
base_url="http://localhost:8000/v1/rerank",
|
||||||
|
api_key="no_api_key",
|
||||||
|
),
|
||||||
|
}
|
||||||
@ -114,8 +114,11 @@ def select_model(model_provider=None, model_name=None, model_spec=None):
|
|||||||
|
|
||||||
assert model_provider, "Model provider not specified"
|
assert model_provider, "Model provider not specified"
|
||||||
|
|
||||||
model_info = config.model_names.get(model_provider, {})
|
model_info = config.model_names.get(model_provider)
|
||||||
model_name = model_name or model_info.get("default", "")
|
if not model_info:
|
||||||
|
raise ValueError(f"Unknown model provider: {model_provider}")
|
||||||
|
|
||||||
|
model_name = model_name or model_info.default
|
||||||
|
|
||||||
if not model_name:
|
if not model_name:
|
||||||
raise ValueError(f"Model name not specified for provider {model_provider}")
|
raise ValueError(f"Model name not specified for provider {model_provider}")
|
||||||
@ -128,8 +131,8 @@ def select_model(model_provider=None, model_name=None, model_spec=None):
|
|||||||
# 其他模型,默认使用OpenAIBase
|
# 其他模型,默认使用OpenAIBase
|
||||||
try:
|
try:
|
||||||
model = OpenAIBase(
|
model = OpenAIBase(
|
||||||
api_key=os.getenv(model_info["env"]),
|
api_key=os.getenv(model_info.env),
|
||||||
base_url=model_info["base_url"],
|
base_url=model_info.base_url,
|
||||||
model_name=model_name,
|
model_name=model_name,
|
||||||
)
|
)
|
||||||
return model
|
return model
|
||||||
|
|||||||
@ -254,10 +254,12 @@ def select_embedding_model(model_id):
|
|||||||
if provider == "local":
|
if provider == "local":
|
||||||
raise ValueError("Local embedding model is not supported, please use other embedding models")
|
raise ValueError("Local embedding model is not supported, please use other embedding models")
|
||||||
|
|
||||||
elif provider == "ollama":
|
# 获取嵌入模型配置并转换为字典
|
||||||
model = OllamaEmbedding(**config.embed_model_names[model_id])
|
embed_config = config.embed_model_names[model_id].model_dump()
|
||||||
|
|
||||||
|
if provider == "ollama":
|
||||||
|
model = OllamaEmbedding(**embed_config)
|
||||||
else:
|
else:
|
||||||
model = OtherEmbedding(**config.embed_model_names[model_id])
|
model = OtherEmbedding(**embed_config)
|
||||||
|
|
||||||
return model
|
return model
|
||||||
|
|||||||
@ -49,7 +49,7 @@ def get_reranker(model_id, **kwargs):
|
|||||||
assert model_id in support_rerankers, f"Unsupported Reranker: {model_id}, only support {support_rerankers}"
|
assert model_id in support_rerankers, f"Unsupported Reranker: {model_id}, only support {support_rerankers}"
|
||||||
|
|
||||||
model_info = config.reranker_names[model_id]
|
model_info = config.reranker_names[model_id]
|
||||||
base_url = model_info["base_url"]
|
base_url = model_info.base_url
|
||||||
api_key = os.getenv(model_info["api_key"], model_info["api_key"])
|
api_key = os.getenv(model_info.api_key, model_info.api_key)
|
||||||
assert api_key, f"{model_info['name']} api_key is required"
|
assert api_key, f"{model_info.name} api_key is required"
|
||||||
return OnlineReranker(model_name=model_info["name"], api_key=api_key, base_url=base_url, **kwargs)
|
return OnlineReranker(model_name=model_info.name, api_key=api_key, base_url=base_url, **kwargs)
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user