feat: 更新模型配置和选择逻辑,修改为通过 model_spec 统一指定模型格式

This commit is contained in:
Wenjie Zhang 2025-10-14 02:27:15 +08:00
parent be40799af9
commit 50d2cf850f
10 changed files with 186 additions and 50 deletions

View File

@ -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 元额度,支持多种开源模型。

View File

@ -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()

View File

@ -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):
"""
处理配置

View File

@ -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"]
logger.info(f"Using user-selected LLM: {provider}/{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: {model_spec}")
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(

View File

@ -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":

View File

@ -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)

View File

@ -56,8 +56,7 @@
<div v-if="value.template_metadata.kind === 'llm'" class="model-selector">
<ModelSelectorComponent
@select-model="handleModelChange"
:model_name="agentConfig[key] ? agentConfig[key].split('/').slice(1).join('/') : ''"
:model_provider="agentConfig[key] ? agentConfig[key].split('/')[0] : ''"
:model_spec="agentConfig[key] || ''"
/>
</div>
@ -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
});
};

View File

@ -1,12 +1,12 @@
<template>
<a-dropdown trigger="click">
<div class="model-select" @click.prevent>
<div class="model-select" :class="modelSelectClasses" @click.prevent>
<div class="model-select-content">
<div class="model-info">
<a-tooltip :title="model_name" placement="right">
<span class="model-text text"> {{ model_name }} </span>
<a-tooltip :title="displayModelText" placement="right">
<span class="model-text text"> {{ displayModelText }} </span>
</a-tooltip>
<span class="model-provider">{{ model_provider }}</span>
<span class="model-provider">{{ displayModelProvider }}</span>
</div>
<div class="model-status-controls">
<span
@ -18,7 +18,7 @@
{{ modelStatusIcon }}
</span>
<a-button
size="small"
:size="buttonSize"
type="text"
:loading="state.checkingStatus"
@click.stop="checkCurrentModelStatus"
@ -48,13 +48,22 @@ import { useConfigStore } from '@/stores/config'
import { chatModelApi } from '@/apis/system_api'
const props = defineProps({
model_name: {
model_spec: {
type: String,
default: ''
},
model_provider: {
sep: {
type: String,
default: ''
default: '/'
},
placeholder: {
type: String,
default: '请选择模型'
},
size: {
type: String,
default: 'small',
validator: value => ['small', 'middle', 'large'].includes(value)
}
});
@ -76,6 +85,38 @@ const modelKeys = computed(() => {
return Object.keys(modelStatus.value || {}).filter(key => modelStatus.value?.[key])
})
const resolvedSep = computed(() => props.sep || '/')
const resolvedSize = computed(() => props.size || 'small')
const modelSelectClasses = computed(() => ({
'model-select--middle': resolvedSize.value === 'middle',
'model-select--large': resolvedSize.value === 'large'
}))
const buttonSize = computed(() => {
if (resolvedSize.value === 'large') return 'large'
if (resolvedSize.value === 'middle') return 'middle'
return 'small'
})
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 +124,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 +170,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)
}
</script>
<style lang="less" scoped>
//
@status-success: #52c41a;
@status-error: #ff4d4f;
@status-warning: #faad14;
@status-default: #999;
@status-success: var(--color-success);
@status-error: var(--color-error);
@status-warning: var(--chart-warning);
@status-default: var(--gray-500);
@border-radius: 8px;
@scrollbar-width: 6px;
@status-indicator-padding: 2px 4px;
@status-check-button-padding: 0 4px;
@status-check-button-font-size: 12px;
@status-indicator-font-size: 11px;
@model-provider-color: #aaa;
@model-provider-color: var(--gray-500);
//
.model-select {
@ -156,7 +201,7 @@ const handleSelectModel = async (provider, name) => {
cursor: pointer;
border: 1px solid var(--gray-200);
border-radius: @border-radius;
background-color: white;
background-color: var(--gray-0);
min-width: 0;
display: flex;
align-items: center;
@ -171,6 +216,14 @@ const handleSelectModel = async (provider, name) => {
max-width: 380px;
}
&.model-select--middle {
font-size: 15px;
}
&.model-select--large {
font-size: 16px;
}
//
.model-select-content {
display: flex;
@ -188,7 +241,7 @@ const handleSelectModel = async (provider, name) => {
.model-text {
overflow: hidden;
text-overflow: ellipsis;
color: #000;
color: var(--gray-1000);
white-space: nowrap;
}
@ -270,4 +323,4 @@ const handleSelectModel = async (provider, name) => {
overflow-y: auto;
}
}
</style>
</style>

View File

@ -61,7 +61,7 @@
:status="progressStatus(task.status)"
stroke-width="6"
/>
<span class="task-card-progress-value">{{ Math.round(task.progress || 0) }}%</span>
<!-- <span class="task-card-progress-value">{{ Math.round(task.progress || 0) }}%</span> -->
</div>
<div v-if="task.message && !isTaskCompleted(task)" class="task-card-message">

View File

@ -61,9 +61,10 @@
<h3 style="margin-top: 20px;">语言模型 (LLM)</h3>
<p style="color: var(--gray-700); font-size: 14px;">可以在设置中配置语言模型</p>
<ModelSelectorComponent
:model_name="newDatabase.llm_info.model_name || '请选择模型'"
:model_provider="newDatabase.llm_info.provider || ''"
:model_spec="llmModelSpec"
placeholder="请选择模型"
@select-model="handleLLMSelect"
size="large"
style="width: 100%; height: 60px;"
/>
</div>
@ -202,6 +203,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 +361,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 = () => {