refactor: 重构 config 模块

This commit is contained in:
Wenjie Zhang 2025-10-22 11:51:32 +08:00
parent 48cf8e2f21
commit 0ee327cb97
13 changed files with 873 additions and 290 deletions

View File

@ -33,6 +33,7 @@ export default defineConfig({
{
text: '高级配置',
items: [
{ text: '配置系统详解', link: '/advanced/configuration' },
{ text: '文档解析', link: '/advanced/document-processing' },
{ text: '智能体', link: '/advanced/agents' },
{ text: '品牌自定义', link: '/advanced/branding' },
@ -42,10 +43,10 @@ export default defineConfig({
{
text: '更新日志',
items: [
{ text: '版本说明 v0.3', link: '/changelog/0.3-release-notes' },
{ text: '路线图', link: '/changelog/roadmap' },
{ text: '参与贡献', link: '/changelog/contributing' },
{ text: '常见问题', link: '/changelog/faq' },
{ text: '版本说明 v0.3', link: '/changelog/0.3-release-notes' }
{ text: '常见问题', link: '/changelog/faq' }
]
}
],

View 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
}

View File

@ -18,6 +18,56 @@ Yuxi-Know v0.3 是一个重要的里程碑版本,包含了多项架构重构
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
- **影响**: 使用新的存储结构,之前存储的历史记录无法直接迁移

View File

@ -7,8 +7,7 @@
## Bugs
- [ ] 当前 ReAct 智能体有消息顺序错乱的 bug且不会默认调用工具
-
## Next
@ -20,7 +19,6 @@
- [ ] 集成智能体评估,首先使用命令行来实现,然后考虑放在 UI 里面展示
- [ ] 开发与生产环境隔离,构建生产镜像 <Badge type="info" text="0.4" />
- [ ] 支持 MinerU 2.5 的解析方法 <Badge type="info" text="0.3.5" />
- [ ] 优化全局配置的管理模型,优化配置管理
## Later
@ -34,3 +32,5 @@
- [x] 添加测试脚本覆盖最常见的功能已覆盖API
- [x] 新建 tasker 模块用来管理所有的后台任务UI 上使用侧边栏管理。
- [x] 优化对文档信息的检索展示(检索结果页、详情页)
- [x] 当前 ReAct 智能体有消息顺序错乱的 bug且不会默认调用工具
- [x] 优化全局配置的管理模型,优化配置管理

View File

@ -39,7 +39,13 @@ default_model: siliconflow/deepseek-ai/DeepSeek-V3.2-Exp
## 自定义模型供应商
::: 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 兼容的模型服务,包括:
@ -52,52 +58,50 @@ default_model: siliconflow/deepseek-ai/DeepSeek-V3.2-Exp
### 1. 编辑模型配置文件
**方式一:修改默认配置**
编辑 `src/config/static/models.yaml` 文件
**方式一:修改默认配置(推荐)**
编辑 `src/config/static/models.py` 文件中的 `DEFAULT_CHAT_MODEL_PROVIDERS` 字典
**方式二:使用覆盖配置**
创建自定义配置文件并通过环境变量指定:
```bash
# 创建自定义配置文件
cp src/config/static/models.yaml /path/to/your/custom-models.yaml
`src/config/static/models.py` 中添加新的模型供应商:
# 设置环境变量
export OVERRIDE_DEFAULT_MODELS_CONFIG_WITH=/path/to/your/custom-models.yaml
```
```python
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
custom-provider-name:
name: custom-provider-name
default: custom-model-name
base_url: "https://api.your-provider.com/v1"
env: CUSTOM_API_KEY_ENV_NAME # 注意:现在是单个环境变量
models:
- supported-model-name
- another-model-name
# 本地 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
# 本地 vLLM 服务
"local-vllm": ChatModelProvider(
name="Local vLLM",
url="https://docs.vllm.ai",
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. 配置环境变量
@ -123,21 +127,58 @@ docker compose restart api-dev
#### 1. 配置模型信息
`src/config/static/models.yaml` 中或通过覆盖配置文件添加配置
`src/config/static/models.py` 中的默认配置部分添加
```yaml
EMBED_MODEL_INFO:
vllm/Qwen/Qwen3-Embedding-0.6B:
name: Qwen/Qwen3-Embedding-0.6B
dimension: 1024
base_url: http://localhost:8000/v1/embeddings
api_key: no_api_key
```python
# 默认嵌入模型配置
DEFAULT_EMBED_MODELS: dict[str, EmbedModelInfo] = {
# ... 现有配置 ...
RERANKER_LIST:
vllm/BAAI/bge-reranker-v2-m3:
name: BAAI/bge-reranker-v2-m3
base_url: http://localhost:8000/v1/rerank
api_key: no_api_key
"vllm/Qwen/Qwen3-Embedding-0.6B": EmbedModelInfo(
name="Qwen/Qwen3-Embedding-0.6B",
dimension=1024,
base_url="http://localhost:8000/v1/embeddings",
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. 启动模型服务
@ -154,20 +195,4 @@ vllm serve BAAI/bge-reranker-v2-m3 \
--task score \
--dtype fp16 \
--port 8000
```
## 常见问题
**Q: 如何查看当前可用的模型?**
在 Web 界面的"设置"页面可以查看所有已配置的模型。
**Q: 模型配置不生效?**
1. 检查环境变量是否正确设置
2. 确认 API 密钥有效
3. 重启服务:`docker compose restart api-dev`
**Q: 如何测试模型连接?**
在 Web 界面的对话页面选择对应模型进行测试。
```

View File

@ -57,6 +57,8 @@ dependencies = [
"pymysql>=1.1.0",
"tenacity>=8.0.0",
"pypinyin>=0.55.0",
"tomli",
"tomli-w",
]
[tool.ruff]
line-length = 120 # 代码最大行宽

View File

@ -393,9 +393,9 @@ async def get_chat_models(model_provider: str, current_user: User = Depends(get_
@chat.post("/models/update")
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._save_models_to_file()
return {"models": conf.model_names[model_provider]["models"]}
conf.model_names[model_provider].models = model_names
conf._save_models_to_file(model_provider)
return {"models": conf.model_names[model_provider].models}
@chat.get("/tools")

View File

@ -14,14 +14,17 @@ def load_chat_model(fully_specified_name: str, **kwargs) -> BaseChatModel:
"""
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, {})
env_var = model_info["env"]
model_info = config.model_names.get(provider)
if not model_info:
raise ValueError(f"Unknown model provider: {provider}")
env_var = model_info.env
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"]:
from langchain_deepseek import ChatDeepSeek

View File

@ -1,246 +1,371 @@
import json
"""
应用配置模块
使用 Pydantic BaseModel 实现配置管理支持
- TOML 文件加载用户配置
- 仅保存用户修改过的配置项
- 默认配置定义在代码中
"""
import os
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
class SimpleConfig(dict):
def __key(self, key):
return "" if key is None else key # 目前忘记了这里为什么要 lower 了,只能说配置项最好不要有大写的
class Config(BaseModel):
"""应用配置类"""
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():
self[key] = value
# 内部状态
_config_file: Path | None = None
_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):
super().__init__()
self._config_items = {}
self.save_dir = os.getenv("SAVE_DIR", "saves")
self.filename = str(Path(f"{self.save_dir}/config/base.yaml"))
os.makedirs(os.path.dirname(self.filename), exist_ok=True)
def __init__(self, **data):
super().__init__(**data)
self._setup_paths()
self._load_user_config()
self._handle_environment()
self._models_config_path: Path | None = os.getenv("OVERRIDE_DEFAULT_MODELS_CONFIG_WITH")
self._update_models_from_file()
def _setup_paths(self):
"""设置配置文件路径"""
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)
### >>> 默认配置
# 功能选项
self.add_item("enable_reranker", default=False, des="是否开启重排序")
self.add_item("enable_content_guard", default=False, des="是否启用内容审查")
self.add_item("enable_content_guard_llm", default=False, des="是否启用LLM内容审查")
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="快速响应模型",
)
def _load_user_config(self):
"""从 TOML 文件加载用户配置"""
if not self._config_file or not self._config_file.exists():
logger.info(f"Config file not found, using defaults: {self._config_file}")
return
self.add_item(
"embed_model",
default="siliconflow/BAAI/bge-m3",
des="Embedding 模型",
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
### <<< 默认配置结束
logger.info(f"Loading config from {self._config_file}")
try:
with open(self._config_file, "rb") as f:
user_config = tomli.load(f)
self.load()
# 清理已废弃的配置项
self.pop("model_provider", None)
self.pop("model_name", None)
self.handle_self()
# 记录用户修改的字段
self._user_modified_fields = set(user_config.keys())
def add_item(self, key, default, des=None, choices=None):
self.__setattr__(key, default)
self._config_items[key] = {"default": default, "des": des, "choices": choices}
# 更新配置
for key, value in user_config.items():
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):
blocklist = [
"_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}
except Exception as e:
logger.error(f"Failed to load config from {self._config_file}: {e}")
def _update_models_from_file(self):
"""
models.yaml 或覆盖配置文件中更新 MODEL_NAMES
"""
# 检查是否设置了覆盖配置文件的环境变量
override_config_path = os.getenv("OVERRIDE_DEFAULT_MODELS_CONFIG_WITH")
if override_config_path and os.path.exists(override_config_path):
config_file = Path(override_config_path)
logger.info(f"Using override models config from: {override_config_path}")
else:
config_file = Path("src/config/static/models.yaml")
logger.info("Using default models config")
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 _load_model_names(self, model_names_data):
"""加载用户自定义的模型配置"""
try:
for provider, provider_data in model_names_data.items():
if provider in self.model_names:
# 更新现有提供商的模型列表
if "models" in provider_data:
self.model_names[provider].models = provider_data["models"]
else:
# 添加新的提供商
self.model_names[provider] = ChatModelProvider(**provider_data)
logger.info(f"Loaded custom model configurations for {len(model_names_data)} providers")
except Exception as e:
logger.error(f"Failed to load model names: {e}")
def _handle_environment(self):
"""处理环境变量和运行时状态"""
# 处理模型目录
self.model_dir = os.environ.get("MODEL_DIR", self.model_dir)
if self.model_dir:
if os.path.exists(self.model_dir):
logger.debug(
f"The model directory {self.model_dir} "
f"contains the following folders: {os.listdir(self.model_dir)}"
f"Model directory ({self.model_dir}) contains: {os.listdir(self.model_dir)}"
)
else:
logger.warning(
f"Warning: The model directory {self.model_dir} does not exist. "
"If not configured, please ignore it. "
"If configured, please check if the configuration is correct; "
"For example, the mapping in the docker-compose file"
f"Model directory ({self.model_dir}) does not exist. "
"If not configured, please ignore it."
)
# 检查模型提供商的环境变量
self.model_provider_status = {}
for provider in self.model_names:
env_var = self.model_names[provider]["env"]
# 如果环境变量名为 NO_API_KEY则认为总是可用
for provider, info in self.model_names.items():
env_var = info.env
if env_var == "NO_API_KEY":
self.model_provider_status[provider] = True
else:
self.model_provider_status[provider] = bool(os.getenv(env_var))
# 检查网络搜索
if os.getenv("TAVILY_API_KEY"):
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):
"""根据传入的文件覆盖掉默认配置"""
logger.info(f"Loading config from {self.filename}")
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}")
if not self.valuable_model_provider:
raise ValueError(
"No model provider available, please check your `.env` file."
)
def save(self):
logger.info(f"Saving config to {self.filename}")
if self.filename is None:
logger.warning("Config file is not specified, save to default config/base.yaml")
self.filename = os.path.join(self.save_dir, "config", "base.yaml")
os.makedirs(os.path.dirname(self.filename), exist_ok=True)
"""保存配置到 TOML 文件(仅保存用户修改的字段)"""
if not self._config_file:
logger.warning("Config file path not set")
return
if self.filename.endswith(".json"):
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"Saving config to {self._config_file}")
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()

187
src/config/static/models.py Normal file
View 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",
),
}

View File

@ -114,8 +114,11 @@ def select_model(model_provider=None, model_name=None, model_spec=None):
assert model_provider, "Model provider not specified"
model_info = config.model_names.get(model_provider, {})
model_name = model_name or model_info.get("default", "")
model_info = config.model_names.get(model_provider)
if not model_info:
raise ValueError(f"Unknown model provider: {model_provider}")
model_name = model_name or model_info.default
if not model_name:
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
try:
model = OpenAIBase(
api_key=os.getenv(model_info["env"]),
base_url=model_info["base_url"],
api_key=os.getenv(model_info.env),
base_url=model_info.base_url,
model_name=model_name,
)
return model

View File

@ -254,10 +254,12 @@ def select_embedding_model(model_id):
if provider == "local":
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:
model = OtherEmbedding(**config.embed_model_names[model_id])
model = OtherEmbedding(**embed_config)
return model

View File

@ -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}"
model_info = config.reranker_names[model_id]
base_url = model_info["base_url"]
api_key = os.getenv(model_info["api_key"], model_info["api_key"])
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)
base_url = model_info.base_url
api_key = os.getenv(model_info.api_key, model_info.api_key)
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)