From afc0664e09ee3c0dc5f073e941ae452f40b14736 Mon Sep 17 00:00:00 2001 From: Wenjie Zhang Date: Tue, 14 Oct 2025 01:54:02 +0800 Subject: [PATCH] =?UTF-8?q?feat:=20=E6=9B=B4=E6=96=B0=E6=A8=A1=E5=9E=8B?= =?UTF-8?q?=E9=85=8D=E7=BD=AE=E5=92=8C=E9=80=89=E6=8B=A9=E9=80=BB=E8=BE=91?= =?UTF-8?q?=EF=BC=8C=E4=BF=AE=E6=94=B9=E4=B8=BA=E9=80=9A=E8=BF=87=20model?= =?UTF-8?q?=5Fspec=20=E7=BB=9F=E4=B8=80=E6=8C=87=E5=AE=9A=E6=A8=A1?= =?UTF-8?q?=E5=9E=8B=E6=A0=BC=E5=BC=8F?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- docs/intro/model-config.md | 10 +++ server/routers/chat_router.py | 6 +- src/config/app.py | 26 ++++++- src/knowledge/implementations/lightrag.py | 11 +-- src/models/chat.py | 33 ++++++++- test/bruteforce_simulation.py | 17 +++-- web/src/components/AgentConfigSidebar.vue | 8 +-- web/src/components/ModelSelectorComponent.vue | 67 ++++++++++++++----- web/src/views/DataBaseView.vue | 27 ++++++-- 9 files changed, 159 insertions(+), 46 deletions(-) diff --git a/docs/intro/model-config.md b/docs/intro/model-config.md index e06decc6..1a40e880 100644 --- a/docs/intro/model-config.md +++ b/docs/intro/model-config.md @@ -21,6 +21,16 @@ <<< @/../.env.template#model_provider{bash 2} +### 默认对话模型格式 + +系统的默认对话模型通过配置项 `default_model` 指定,格式统一为 `模型提供商/模型名称`,例如: + +```yaml +default_model: siliconflow/deepseek-ai/DeepSeek-V3.2-Exp +``` + +在 Web 界面中选择模型时也会自动按照这一格式保存,无需手动拆分提供商和模型名称。 + ::: tip 免费获取 API Key [硅基流动](https://cloud.siliconflow.cn/i/Eo5yTHGJ) 注册即送 14 元额度,支持多种开源模型。 diff --git a/server/routers/chat_router.py b/server/routers/chat_router.py index d3ab82f5..9954459a 100644 --- a/server/routers/chat_router.py +++ b/server/routers/chat_router.py @@ -85,7 +85,11 @@ async def set_default_agent(request_data: dict = Body(...), current_user=Depends async def call(query: str = Body(...), meta: dict = Body(None), current_user: User = Depends(get_required_user)): """调用模型进行简单问答(需要登录)""" meta = meta or {} - model = select_model(model_provider=meta.get("model_provider"), model_name=meta.get("model_name")) + model = select_model( + model_provider=meta.get("model_provider"), + model_name=meta.get("model_name"), + model_spec=meta.get("model_spec") or meta.get("model"), + ) async def call_async(query): loop = asyncio.get_event_loop() diff --git a/src/config/app.py b/src/config/app.py index 9624011b..f476f226 100644 --- a/src/config/app.py +++ b/src/config/app.py @@ -58,8 +58,11 @@ class Config(SimpleConfig): # 模型配置 ## 注意这里是模型名,而不是具体的模型路径,默认使用 HuggingFace 的路径 ## 如果需要自定义本地模型路径,则在 .env 中配置 MODEL_DIR - self.add_item("model_provider", default="siliconflow", des="模型提供商", choices=list(self.model_names.keys())) - self.add_item("model_name", default="zai-org/GLM-4.5", des="模型名称") + 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", @@ -81,6 +84,9 @@ class Config(SimpleConfig): ### <<< 默认配置结束 self.load() + # 清理已废弃的配置项 + self.pop("model_provider", None) + self.pop("model_name", None) self.handle_self() def add_item(self, key, default, des=None, choices=None): @@ -137,6 +143,22 @@ class Config(SimpleConfig): 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): """ 处理配置 diff --git a/src/knowledge/implementations/lightrag.py b/src/knowledge/implementations/lightrag.py index c2143234..3bf22536 100644 --- a/src/knowledge/implementations/lightrag.py +++ b/src/knowledge/implementations/lightrag.py @@ -174,16 +174,19 @@ class LightRagKB(KnowledgeBase): from src.models import select_model # 如果用户选择了LLM,使用用户选择的;否则使用环境变量默认值 - if llm_info and llm_info.get("provider") and llm_info.get("model_name"): - provider = llm_info["provider"] - model_name = llm_info["model_name"] + if llm_info and llm_info.get("model_spec"): + model_spec = llm_info["model_spec"] + logger.info(f"Using user-selected LLM spec: {model_spec}") + elif llm_info and llm_info.get("provider") and llm_info.get("model_name"): + model_spec = f"{llm_info['provider']}/{llm_info['model_name']}" logger.info(f"Using user-selected LLM: {provider}/{model_name}") else: provider = LIGHTRAG_LLM_PROVIDER model_name = LIGHTRAG_LLM_NAME + model_spec = f"{provider}/{model_name}" logger.info(f"Using default LLM from environment: {provider}/{model_name}") - model = select_model(provider, model_name) + model = select_model(model_spec=model_spec) async def llm_model_func(prompt, system_prompt=None, history_messages=[], **kwargs): return await openai_complete_if_cache( diff --git a/src/models/chat.py b/src/models/chat.py index aa06c7a3..87fdb5e1 100644 --- a/src/models/chat.py +++ b/src/models/chat.py @@ -8,6 +8,21 @@ from src import config from src.utils import logger +def split_model_spec(model_spec, sep="/"): + """ + 将 provider/model 形式的字符串拆分为 (provider, model) + """ + if not model_spec or not isinstance(model_spec, str): + return "", "" + if not sep: + return model_spec, "" + try: + provider, model_name = model_spec.split(sep, 1) + return provider, model_name + except ValueError: + return model_spec, "" + + class OpenAIBase: def __init__(self, api_key, base_url, model_name, **kwargs): self.api_key = api_key @@ -85,12 +100,26 @@ class GeneralResponse: self.is_full = False -def select_model(model_provider, model_name=None): +def select_model(model_provider=None, model_name=None, model_spec=None): """根据模型提供者选择模型""" - assert model_provider is not None, "Model provider not specified" + if model_spec: + spec_provider, spec_model_name = split_model_spec(model_spec) + model_provider = model_provider or spec_provider + model_name = model_name or spec_model_name + + if model_provider is None or not model_name: + default_provider, default_model = split_model_spec(getattr(config, "default_model", "")) + model_provider = model_provider or default_provider + model_name = model_name or default_model + + assert model_provider, "Model provider not specified" + model_info = config.model_names.get(model_provider, {}) model_name = model_name or model_info.get("default", "") + if not model_name: + raise ValueError(f"Model name not specified for provider {model_provider}") + logger.info(f"Selecting model from `{model_provider}` with `{model_name}`") if model_provider == "openai": diff --git a/test/bruteforce_simulation.py b/test/bruteforce_simulation.py index 74ea74de..4a367106 100644 --- a/test/bruteforce_simulation.py +++ b/test/bruteforce_simulation.py @@ -23,9 +23,7 @@ def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser(description="Simulate brute-force login attempts.") parser.add_argument("--base-url", default=os.getenv("TEST_BASE_URL", "http://localhost:5050"), help="API base URL") parser.add_argument("--username", default=os.getenv("TEST_USERNAME", "admin"), help="Login identifier to attack") - parser.add_argument( - "--attempts", type=int, default=20, help="Total number of attempts to issue (default: 20)" - ) + parser.add_argument("--attempts", type=int, default=20, help="Total number of attempts to issue (default: 20)") parser.add_argument( "--concurrency", type=int, @@ -68,10 +66,13 @@ async def attempt_login( started = time.perf_counter() response = await client.post("/api/auth/token", data=payload) elapsed = time.perf_counter() - started - detail = response.json().get("detail") if response.headers.get("content-type", "").startswith("application/json") else response.text + detail = ( + response.json().get("detail") + if response.headers.get("content-type", "").startswith("application/json") + else response.text + ) print( - f"[{attempt_no:02d}] {response.status_code} in {elapsed*1000:.1f} ms " - f"(pwd={password!r}) detail={detail!r}" + f"[{attempt_no:02d}] {response.status_code} in {elapsed * 1000:.1f} ms (pwd={password!r}) detail={detail!r}" ) return response.status_code, elapsed @@ -86,9 +87,7 @@ async def run_simulation(args: argparse.Namespace) -> int: tasks = [] for attempt_no in range(1, args.attempts + 1): tasks.append( - asyncio.create_task( - attempt_login(client, semaphore, attempt_no, args.username, args.password) - ) + asyncio.create_task(attempt_login(client, semaphore, attempt_no, args.username, args.password)) ) if args.delay: await asyncio.sleep(args.delay) diff --git a/web/src/components/AgentConfigSidebar.vue b/web/src/components/AgentConfigSidebar.vue index 91d8831c..d2f616d3 100644 --- a/web/src/components/AgentConfigSidebar.vue +++ b/web/src/components/AgentConfigSidebar.vue @@ -56,8 +56,7 @@
@@ -367,9 +366,10 @@ const getPlaceholder = (key, value) => { return `(默认: ${value.default})`; }; -const handleModelChange = (data) => { +const handleModelChange = (spec) => { + if (typeof spec !== 'string' || !spec) return; agentStore.updateAgentConfig({ - model: `${data.provider}/${data.name}` + model: spec }); }; diff --git a/web/src/components/ModelSelectorComponent.vue b/web/src/components/ModelSelectorComponent.vue index 9eb917ab..bd759c26 100644 --- a/web/src/components/ModelSelectorComponent.vue +++ b/web/src/components/ModelSelectorComponent.vue @@ -3,10 +3,10 @@
- - {{ model_name }} + + {{ displayModelText }} - {{ model_provider }} + {{ displayModelProvider }}
{ return Object.keys(modelStatus.value || {}).filter(key => modelStatus.value?.[key]) }) +const resolvedSep = computed(() => props.sep || '/') + +const resolvedModel = computed(() => { + const spec = props.model_spec || '' + const sep = resolvedSep.value + if (spec && sep) { + const index = spec.indexOf(sep) + if (index !== -1) { + const provider = spec.slice(0, index) + const name = spec.slice(index + sep.length) + if (provider && name) { + return { provider, name } + } + } + } + return { provider: '', name: '' } +}) + +const displayModelProvider = computed(() => resolvedModel.value.provider || '') +const displayModelName = computed(() => resolvedModel.value.name || '') +const displayModelText = computed(() => displayModelName.value || props.placeholder) + // 当前模型状态 const currentModelStatus = computed(() => { return state.currentModelStatus @@ -83,18 +109,19 @@ const currentModelStatus = computed(() => { // 检查当前模型状态 const checkCurrentModelStatus = async () => { - if (!props.model_provider || !props.model_name) return + const { provider, name } = resolvedModel.value + if (!provider || !name) return try { state.checkingStatus = true - const response = await chatModelApi.getModelStatus(props.model_provider, props.model_name) + const response = await chatModelApi.getModelStatus(provider, name) if (response.status) { state.currentModelStatus = response.status } else { state.currentModelStatus = null } } catch (error) { - console.error(`检查当前模型 ${props.model_provider}/${props.model_name} 状态失败:`, error) + console.error(`检查当前模型 ${provider}/${name} 状态失败:`, error) state.currentModelStatus = { status: 'error', message: error.message } } finally { state.checkingStatus = false @@ -128,24 +155,27 @@ const getCurrentModelStatusTooltip = () => { // 选择模型的方法 const handleSelectModel = async (provider, name) => { - emit('select-model', { provider, name }) + const sep = resolvedSep.value || '/' + const separator = sep || '/' + const spec = `${provider}${separator}${name}` + emit('select-model', spec) } \ No newline at end of file + diff --git a/web/src/views/DataBaseView.vue b/web/src/views/DataBaseView.vue index 593989e7..e0648c3f 100644 --- a/web/src/views/DataBaseView.vue +++ b/web/src/views/DataBaseView.vue @@ -61,8 +61,8 @@

语言模型 (LLM)

可以在设置中配置语言模型

@@ -202,6 +202,15 @@ const newDatabase = reactive({ ...emptyEmbedInfo, }) +const llmModelSpec = computed(() => { + const provider = newDatabase.llm_info?.provider || '' + const modelName = newDatabase.llm_info?.model_name || '' + if (provider && modelName) { + return `${provider}/${modelName}` + } + return '' +}) + // 支持的知识库类型 const supportedKbTypes = ref({}) @@ -351,10 +360,16 @@ const handleKbTypeChange = (type) => { } // 处理LLM选择 -const handleLLMSelect = (selection) => { - console.log('LLM选择:', selection) - newDatabase.llm_info.provider = selection.provider - newDatabase.llm_info.model_name = selection.name +const handleLLMSelect = (spec) => { + console.log('LLM选择:', spec) + if (typeof spec !== 'string' || !spec) return + + const index = spec.indexOf('/') + const provider = index !== -1 ? spec.slice(0, index) : '' + const modelName = index !== -1 ? spec.slice(index + 1) : '' + + newDatabase.llm_info.provider = provider + newDatabase.llm_info.model_name = modelName } const createDatabase = () => {