diff --git a/server/routers/chat_router.py b/server/routers/chat_router.py index 4c089d7d..50fb73f3 100644 --- a/server/routers/chat_router.py +++ b/server/routers/chat_router.py @@ -203,6 +203,13 @@ async def get_chat_models(model_provider: str): model = select_model(model_provider=model_provider) return {"models": model.get_models()} +@chat.post("/models/update") +async def update_chat_models(model_provider: str, model_names: list[str]): + """更新指定模型提供商的模型列表""" + config.model_names[model_provider]["models"] = model_names + config._save_models_to_file() + return {"models": config.model_names[model_provider]["models"]} + @chat.get("/tools") async def get_tools(): """获取所有工具""" diff --git a/src/config/__init__.py b/src/config/__init__.py index b9fc17e2..5da6842c 100644 --- a/src/config/__init__.py +++ b/src/config/__init__.py @@ -4,13 +4,6 @@ import yaml from pathlib import Path from src.utils.logging_config import logger -with open(Path("src/static/models.yaml"), 'r', encoding='utf-8') as f: - _models = yaml.safe_load(f) - -MODEL_NAMES = _models["MODEL_NAMES"] -EMBED_MODEL_INFO = _models["EMBED_MODEL_INFO"] -RERANKER_LIST = _models["RERANKER_LIST"] - DEFAULT_MOCK_API = 'this_is_mock_api_key_in_frontend' class SimpleConfig(dict): @@ -46,6 +39,8 @@ class Config(SimpleConfig): self.filename = str(Path("saves/config/base.yaml")) os.makedirs(os.path.dirname(self.filename), exist_ok=True) + self._update_models_from_file() + ### >>> 默认配置 self.add_item("stream", default=True, des="是否开启流式输出") # 功能选项 @@ -56,13 +51,13 @@ class Config(SimpleConfig): # 模型配置 ## 注意这里是模型名,而不是具体的模型路径,默认使用 HuggingFace 的路径 ## 如果需要自定义本地模型路径,则在 src/.env 中配置 MODEL_DIR - self.add_item("model_provider", default="siliconflow", des="模型提供商", choices=list(MODEL_NAMES.keys())) - self.add_item("model_provider_lite", default="siliconflow", des="模型提供商(用于轻量任务)", choices=list(MODEL_NAMES.keys())) + self.add_item("model_provider", default="siliconflow", des="模型提供商", choices=list(self.model_names.keys())) + self.add_item("model_provider_lite", default="siliconflow", des="模型提供商(用于轻量任务)", choices=list(self.model_names.keys())) self.add_item("model_name", default="Qwen/Qwen2.5-7B-Instruct", des="模型名称") self.add_item("model_name_lite", default="Qwen/Qwen2.5-7B-Instruct", des="模型名称(用于轻量任务)") - self.add_item("embed_model", default="siliconflow/BAAI/bge-m3", des="Embedding 模型", choices=list(EMBED_MODEL_INFO.keys())) - self.add_item("reranker", default="siliconflow/BAAI/bge-reranker-v2-m3", des="Re-Ranker 模型", choices=list(RERANKER_LIST.keys())) + 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())) self.add_item("model_local_paths", default={}, des="本地模型路径") self.add_item("use_rewrite_query", default="off", des="重写查询", choices=["off", "on", "hyde"]) self.add_item("device", default="cuda", des="运行本地模型的设备", choices=["cpu", "cuda"]) @@ -89,11 +84,40 @@ class Config(SimpleConfig): ] return {k: v for k, v in self.items() if k not in blocklist} - def handle_self(self): - self.model_names = MODEL_NAMES - self.embed_model_names = EMBED_MODEL_INFO - self.reranker_names = RERANKER_LIST + def _update_models_from_file(self): + """ + 从 models.yaml 和 models.private.yml 中更新 MODEL_NAMES + """ + with open(Path("src/static/models.yaml"), 'r', encoding='utf-8') as f: + _models = yaml.safe_load(f) + + # 尝试打开一个 models.private.yml 文件,用来覆盖 models.yaml 中的配置 + try: + with open(Path("src/static/models.private.yml"), 'r', encoding='utf-8') as f: + _models_private = yaml.safe_load(f) + except FileNotFoundError: + _models_private = {} + + _models = {**_models, **_models_private} + + 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): + _models = { + "MODEL_NAMES": self.model_names, + "EMBED_MODEL_INFO": self.embed_model_names, + "RERANKER_LIST": self.reranker_names, + } + with open(Path("src/static/models.private.yml"), 'w', encoding='utf-8') as f: + yaml.dump(_models, f, indent=2, allow_unicode=True) + + def handle_self(self): + """ + 处理配置 + """ model_provider_info = self.model_names.get(self.model_provider, {}) self.model_dir = os.environ.get("MODEL_DIR", "") logger.info(f"MODEL_DIR: {self.model_dir}; 如果是在 docker 中运行,会自动挂载 MODEL_DIR 到 /models 目录,请检查 docker compose 文件") diff --git a/src/core/graphbase.py b/src/core/graphbase.py index 6f0af518..9c61063b 100644 --- a/src/core/graphbase.py +++ b/src/core/graphbase.py @@ -141,8 +141,7 @@ class GraphDatabase: """) # 判断模型名称是否匹配 - from src.config import EMBED_MODEL_INFO - cur_embed_info = EMBED_MODEL_INFO[config.embed_model] + cur_embed_info = config.embed_model_names[config.embed_model] self.embed_model_name = self.embed_model_name or cur_embed_info.get('name') assert self.embed_model_name == cur_embed_info.get('name') or self.embed_model_name is None, \ f"embed_model_name={self.embed_model_name}, {cur_embed_info.get('name')=}" diff --git a/src/models/embedding.py b/src/models/embedding.py index 558de9be..ff1a3bcf 100644 --- a/src/models/embedding.py +++ b/src/models/embedding.py @@ -5,7 +5,6 @@ from FlagEmbedding import FlagModel from zhipuai import ZhipuAI from src import config -from src.config import EMBED_MODEL_INFO from src.utils import hashstr, logger, get_docker_safe_url @@ -17,9 +16,9 @@ class BaseEmbeddingModel: return self.dimension if hasattr("embed_model_fullname"): - return EMBED_MODEL_INFO[self.embed_model_fullname].get("dimension", None) + return config.embed_model_names[self.embed_model_fullname].get("dimension", None) - return EMBED_MODEL_INFO[self.model].get("dimension", None) + return config.embed_model_names[self.model].get("dimension", None) def encode(self, message): return self.predict(message) @@ -54,7 +53,7 @@ class BaseEmbeddingModel: class LocalEmbeddingModel(FlagModel, BaseEmbeddingModel): def __init__(self, config, **kwargs): - info = EMBED_MODEL_INFO[config.embed_model] + info = config.embed_model_names[config.embed_model] self.model = config.model_local_paths.get(info["name"], info.get("local_path")) self.model = self.model or info["name"] @@ -79,8 +78,8 @@ class ZhipuEmbedding(BaseEmbeddingModel): def __init__(self, config) -> None: self.config = config - self.model = EMBED_MODEL_INFO[config.embed_model]["name"] - self.dimension = EMBED_MODEL_INFO[config.embed_model]["dimension"] + self.model = config.embed_model_names[config.embed_model]["name"] + self.dimension = config.embed_model_names[config.embed_model]["dimension"] self.client = ZhipuAI(api_key=os.getenv("ZHIPUAI_API_KEY")) self.embed_model_fullname = config.embed_model @@ -95,7 +94,7 @@ class ZhipuEmbedding(BaseEmbeddingModel): class OllamaEmbedding(BaseEmbeddingModel): def __init__(self, config) -> None: - self.info = EMBED_MODEL_INFO[config.embed_model] + self.info = config.embed_model_names[config.embed_model] self.model = self.info["name"] self.url = self.info.get("url", "http://localhost:11434/api/embed") self.url = get_docker_safe_url(self.url) @@ -119,7 +118,7 @@ class OllamaEmbedding(BaseEmbeddingModel): class OtherEmbedding(BaseEmbeddingModel): def __init__(self, config) -> None: - self.info = EMBED_MODEL_INFO[config.embed_model] + self.info = config.embed_model_names[config.embed_model] self.embed_model_fullname = config.embed_model self.dimension = self.info.get("dimension", None) self.model = self.info["name"] @@ -150,7 +149,7 @@ def get_embedding_model(config): return None provider, model_name = config.embed_model.split('/', 1) - assert config.embed_model in EMBED_MODEL_INFO.keys(), f"Unsupported embed model: {config.embed_model}, only support {EMBED_MODEL_INFO.keys()}" + assert config.embed_model in config.embed_model_names.keys(), f"Unsupported embed model: {config.embed_model}, only support {config.embed_model_names.keys()}" logger.debug(f"Loading embedding model {config.embed_model}") if provider == "local": model = LocalEmbeddingModel(config) diff --git a/src/models/rerank_model.py b/src/models/rerank_model.py index afe3e124..a2fb7d37 100644 --- a/src/models/rerank_model.py +++ b/src/models/rerank_model.py @@ -4,13 +4,13 @@ import requests import numpy as np from FlagEmbedding import FlagReranker -from src.config import RERANKER_LIST +from src import config from src.utils.logging_config import logger class LocalReranker(FlagReranker): def __init__(self, config, **kwargs): - model_info = RERANKER_LIST[config.reranker] + model_info = config.reranker_names[config.reranker] model_name_or_path = config.model_local_paths.get(model_info["name"], model_info.get("local_path")) model_name_or_path = model_name_or_path or model_info["name"] logger.info(f"Loading Reranker model {config.reranker} from {model_name_or_path}") @@ -25,7 +25,7 @@ def sigmoid(x): class SilconFlowReranker(): def __init__(self, config, **kwargs): self.url = "https://api.siliconflow.cn/v1/rerank" - self.model = RERANKER_LIST[config.reranker]["name"] + self.model = config.reranker_names[config.reranker]["name"] api_key = os.getenv("SILICONFLOW_API_KEY") assert api_key, "SILICONFLOW_API_KEY is required" @@ -59,12 +59,12 @@ class SilconFlowReranker(): } def get_reranker(config): - assert config.reranker in RERANKER_LIST.keys(), f"Unsupported Reranker: {config.reranker}, only support {RERANKER_LIST.keys()}" + assert config.reranker in config.reranker_names.keys(), f"Unsupported Reranker: {config.reranker}, only support {config.reranker_names.keys()}" provider, model_name = config.reranker.split('/', 1) if provider == "local": return LocalReranker(config) elif provider == "siliconflow": return SilconFlowReranker(config) else: - raise ValueError(f"Unsupported Reranker: {config.reranker}, only support {RERANKER_LIST.keys()}") + raise ValueError(f"Unsupported Reranker: {config.reranker}, only support {config.reranker_names.keys()}") diff --git a/src/static/models.yaml b/src/static/models.yaml index a1f1f371..9a41c7a9 100644 --- a/src/static/models.yaml +++ b/src/static/models.yaml @@ -1,3 +1,13 @@ +#################################################### +# +# 不要直接修改这个里面的文件,可能会有被覆盖的风险, +# 建议 复制一份 在 models.private.yml 中修改, +# 会自动加载 +# +##################################################### + + + MODEL_NAMES: openai: name: OpenAI diff --git a/web/src/views/SettingView.vue b/web/src/views/SettingView.vue index d5b56e45..6af88d50 100644 --- a/web/src/views/SettingView.vue +++ b/web/src/views/SettingView.vue @@ -169,6 +169,14 @@ + + + + + +
需配置{{ key }}
@@ -217,6 +233,43 @@ /> + + + + @@ -235,6 +288,7 @@ import { InfoCircleOutlined, DownOutlined, UpOutlined, + LoadingOutlined, } from '@ant-design/icons-vue'; import HeaderComponent from '@/components/HeaderComponent.vue'; import TableConfigComponent from '@/components/TableConfigComponent.vue'; @@ -257,6 +311,15 @@ const customModel = reactive({ api_base: '', edit_type: 'add', }) +const providerConfig = reactive({ + visible: false, + provider: '', + providerName: '', + models: [], + allModels: [], // 所有可用的模型 + selectedModels: [], // 用户选择的模型 + loading: false, +}) const state = reactive({ loading: false, section: 'base', @@ -427,6 +490,99 @@ const sendRestart = () => { }, 200) }) } + +// 获取模型提供商的模型列表 +const fetchProviderModels = (provider) => { + providerConfig.loading = true; + fetch(`/api/chat/models?model_provider=${provider}`) + .then(response => response.json()) + .then(data => { + console.log(`${provider} 模型列表:`, data); + + // 处理各种可能的API返回格式 + let modelsList = []; + + // 情况1: { data: [...] } + if (data.data && Array.isArray(data.data)) { + modelsList = data.data; + } + // 情况2: { models: [...] } (字符串数组) + else if (data.models && Array.isArray(data.models)) { + modelsList = data.models.map(model => typeof model === 'string' ? { id: model } : model); + } + // 情况3: { models: { data: [...] } } + else if (data.models && data.models.data && Array.isArray(data.models.data)) { + modelsList = data.models.data; + } + + console.log("处理后的模型列表:", modelsList); + providerConfig.allModels = modelsList; + providerConfig.loading = false; + }) + .catch(error => { + console.error(`获取${provider}模型列表失败:`, error); + message.error({ content: `获取${modelNames.value[provider].name}模型列表失败`, duration: 2 }); + providerConfig.loading = false; + }); +} + +const openProviderConfig = (provider) => { + providerConfig.provider = provider; + providerConfig.providerName = modelNames.value[provider].name; + providerConfig.allModels = []; + providerConfig.visible = true; + providerConfig.loading = true; + + // 获取当前选择的模型作为初始选中值 + const currentModels = modelNames.value[provider]?.models || []; + providerConfig.selectedModels = [...currentModels]; + + // 获取所有可用模型 + fetchProviderModels(provider); +} + +const saveProviderConfig = async () => { + if (!modelStatus.value[providerConfig.provider]) { + message.error('请在 src/.env 中配置对应的 APIKEY') + return + } + + message.loading({ content: '保存配置中...', key: 'save-config', duration: 0 }); + + try { + // 发送选择的模型列表到后端 + const response = await fetch(`/api/chat/models/update?model_provider=${providerConfig.provider}`, { + method: 'POST', + headers: { + 'Content-Type': 'application/json', + }, + body: JSON.stringify(providerConfig.selectedModels), + }); + + if (!response.ok) { + throw new Error('保存模型配置失败'); + } + + const data = await response.json(); + console.log('更新后的模型列表:', data.models); + + message.success({ content: '模型配置已保存!', key: 'save-config', duration: 2 }); + + // 关闭弹窗 + providerConfig.visible = false; + + // 刷新配置 + configStore.refreshConfig(); + + } catch (error) { + console.error('保存配置失败:', error); + message.error({ content: '保存配置失败: ' + error.message, key: 'save-config', duration: 2 }); + } +} + +const cancelProviderConfig = () => { + providerConfig.visible = false; +} +