feat: 更新模型配置和选择逻辑,修改为通过 model_spec 统一指定模型格式
This commit is contained in:
parent
be40799af9
commit
afc0664e09
@ -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 元额度,支持多种开源模型。
|
||||
|
||||
@ -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()
|
||||
|
||||
@ -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):
|
||||
"""
|
||||
处理配置
|
||||
|
||||
@ -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(
|
||||
|
||||
@ -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":
|
||||
|
||||
@ -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)
|
||||
|
||||
@ -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
|
||||
});
|
||||
};
|
||||
|
||||
|
||||
@ -3,10 +3,10 @@
|
||||
<div class="model-select" @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
|
||||
@ -48,13 +48,17 @@ 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: '请选择模型'
|
||||
}
|
||||
});
|
||||
|
||||
@ -76,6 +80,28 @@ const modelKeys = computed(() => {
|
||||
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)
|
||||
}
|
||||
|
||||
</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,11 +186,12 @@ 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;
|
||||
gap: 0.5rem;
|
||||
font-size: 13px;
|
||||
|
||||
// 修饰符类
|
||||
&.borderless {
|
||||
@ -188,7 +219,7 @@ const handleSelectModel = async (provider, name) => {
|
||||
.model-text {
|
||||
overflow: hidden;
|
||||
text-overflow: ellipsis;
|
||||
color: #000;
|
||||
color: var(--gray-1000);
|
||||
white-space: nowrap;
|
||||
}
|
||||
|
||||
@ -270,4 +301,4 @@ const handleSelectModel = async (provider, name) => {
|
||||
overflow-y: auto;
|
||||
}
|
||||
}
|
||||
</style>
|
||||
</style>
|
||||
|
||||
@ -61,8 +61,8 @@
|
||||
<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"
|
||||
style="width: 100%; height: 60px;"
|
||||
/>
|
||||
@ -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 = () => {
|
||||
|
||||
Loading…
Reference in New Issue
Block a user