refactor: 重构 config 模块
This commit is contained in:
parent
48cf8e2f21
commit
0ee327cb97
@ -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' }
|
||||
]
|
||||
}
|
||||
],
|
||||
|
||||
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
|
||||
```
|
||||
|
||||
### 配置文件管理调整
|
||||
|
||||
配置系统从 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
|
||||
- **影响**: 使用新的存储结构,之前存储的历史记录无法直接迁移
|
||||
|
||||
@ -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] 优化全局配置的管理模型,优化配置管理
|
||||
|
||||
@ -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 界面的对话页面选择对应模型进行测试。
|
||||
```
|
||||
@ -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 # 代码最大行宽
|
||||
|
||||
@ -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")
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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
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"
|
||||
|
||||
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
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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)
|
||||
|
||||
Loading…
Reference in New Issue
Block a user