feat(provider): support custom prompts for model tests (#242)

This commit is contained in:
RWDai
2026-03-19 13:11:54 +08:00
committed by GitHub
parent b570aaac48
commit ddd6adbcf7
8 changed files with 164 additions and 46 deletions

View File

@@ -117,13 +117,17 @@ export function useModelTest(options: UseModelTestOptions) {
startPolling(reqId) startPolling(reqId)
try { try {
const normalizedMessage = typeof params.message === 'string' && params.message.trim()
? params.message.trim()
: undefined
const result = await testModelFailover({ const result = await testModelFailover({
provider_id: providerId(), provider_id: providerId(),
mode: params.mode, mode: params.mode,
model_name: params.modelName, model_name: params.modelName,
api_format: params.apiFormat, api_format: params.apiFormat,
endpoint_id: params.endpointId, endpoint_id: params.endpointId,
message: params.message ?? 'hello', ...(normalizedMessage ? { message: normalizedMessage } : {}),
request_id: reqId, request_id: reqId,
concurrency: params.concurrency, concurrency: params.concurrency,
}, { }, {

View File

@@ -192,7 +192,6 @@ async function handleTestModel(modelName: string) {
model_name: modelName, model_name: modelName,
api_key_id: props.keyId, api_key_id: props.keyId,
api_format: 'gemini:chat', api_format: 'gemini:chat',
message: 'hello',
}) })
if (result.success) { if (result.success) {

View File

@@ -391,7 +391,6 @@ async function testMapping(group: AliasGroup, mapping: ProviderModelAlias) {
const result = await testModel({ const result = await testModel({
provider_id: props.provider.id, provider_id: props.provider.id,
model_name: mapping.name, // 使用映射名称进行测试 model_name: mapping.name, // 使用映射名称进行测试
message: "hello",
api_format: apiFormat api_format: apiFormat
}) })

View File

@@ -317,7 +317,10 @@
:testing="modelTest.testing.value" :testing="modelTest.testing.value"
:trace="modelTest.testTrace.value" :trace="modelTest.testTrace.value"
:request-id="modelTest.requestId.value" :request-id="modelTest.requestId.value"
:message-draft="testMessageDraft"
@close="handleTestDialogClose" @close="handleTestDialogClose"
@start="handleStartMappingTest"
@update:message-draft="testMessageDraft = $event"
/> />
</template> </template>
@@ -390,8 +393,10 @@ const deleteConfirmOpen = ref(false)
const editingGroup = ref<AliasGroup | null>(null) const editingGroup = ref<AliasGroup | null>(null)
const deletingGroup = ref<AliasGroup | null>(null) const deletingGroup = ref<AliasGroup | null>(null)
const testingMapping = ref<string | null>(null) const testingMapping = ref<string | null>(null)
const pendingMappingKey = ref<string | null>(null)
const testingModelName = ref<string | null>(null) const testingModelName = ref<string | null>(null)
const preselectedModelId = ref<string | null>(null) const preselectedModelId = ref<string | null>(null)
const testMessageDraft = ref('')
// 使用 props 传入的数据 // 使用 props 传入的数据
const models = computed(() => props.models ?? []) const models = computed(() => props.models ?? [])
@@ -626,22 +631,38 @@ async function onDialogSaved() {
function handleTestDialogClose() { function handleTestDialogClose() {
modelTest.resetState() modelTest.resetState()
pendingMappingKey.value = null
testingModelName.value = null testingModelName.value = null
testingMapping.value = null
} }
// 测试映射(直连测试,带故障转移和实时进度) // 测试映射(直连测试,带故障转移和实时进度)
async function runMappingTest(testingKey: string, modelName: string) { function runMappingTest(testingKey: string, modelName: string) {
testingMapping.value = testingKey pendingMappingKey.value = testingKey
modelTest.testResult.value = null
modelTest.dialogOpen.value = true
testingMapping.value = null
testingModelName.value = modelName testingModelName.value = modelName
}
async function handleStartMappingTest() {
if (modelTest.testing.value || !testingModelName.value || !pendingMappingKey.value) return
const currentMappingKey = pendingMappingKey.value
testingMapping.value = currentMappingKey
await modelTest.startTest({ await modelTest.startTest({
mode: 'direct', mode: 'direct',
modelName, modelName: testingModelName.value,
displayLabel: `映射 "${modelName}"`, displayLabel: `映射 "${testingModelName.value}"`,
message: 'hello', message: testMessageDraft.value,
onSuccess: () => { onSuccess: () => {
pendingMappingKey.value = null
testingModelName.value = null testingModelName.value = null
}, },
}) })
if (pendingMappingKey.value === currentMappingKey) {
pendingMappingKey.value = null
}
testingMapping.value = null testingMapping.value = null
} }

View File

@@ -7,35 +7,91 @@
@update:open="(val: boolean) => { if (!val) emit('close') }" @update:open="(val: boolean) => { if (!val) emit('close') }"
> >
<div <div
v-if="showSelection" v-if="showSetup"
class="space-y-2" class="space-y-4"
> >
<button <div class="space-y-2">
v-for="endpoint in endpoints"
:key="endpoint.id"
type="button"
class="w-full rounded-lg border border-border/60 px-3 py-3 text-left transition-colors hover:bg-muted/40"
@click="emit('select-endpoint', endpoint.id)"
>
<div class="flex items-center justify-between gap-3"> <div class="flex items-center justify-between gap-3">
<div class="min-w-0"> <div class="text-sm font-medium">
<div class="text-sm font-medium"> 测试内容
{{ formatApiFormat(endpoint.api_format) }} </div>
</div> <div class="text-[11px] text-muted-foreground">
<div class="mt-1 text-xs text-muted-foreground truncate"> 留空时使用默认提示词
{{ endpoint.base_url }}
</div>
</div> </div>
<Badge variant="outline">
{{ endpoint.is_active ? '已启用' : '已禁用' }}
</Badge>
</div> </div>
</button> <Textarea
:model-value="messageDraft"
class="min-h-[132px]"
placeholder="输入自定义测试内容;留空时使用系统默认测试提示词"
@update:model-value="emit('update:messageDraft', $event)"
/>
</div>
<div <div
v-if="endpoints.length === 0" v-if="showEndpointChoices"
class="rounded-lg border border-dashed border-border/60 px-3 py-6 text-center text-sm text-muted-foreground" class="space-y-2"
> >
暂无可用于测试的活跃端点 <div class="text-sm font-medium">
选择测试端点
</div>
<button
v-for="endpoint in endpoints"
:key="endpoint.id"
type="button"
class="w-full rounded-lg border border-border/60 px-3 py-3 text-left transition-colors hover:bg-muted/40"
@click="emit('selectEndpoint', endpoint.id)"
>
<div class="flex items-center justify-between gap-3">
<div class="min-w-0">
<div class="text-sm font-medium">
{{ formatApiFormat(endpoint.api_format) }}
</div>
<div class="mt-1 text-xs text-muted-foreground truncate">
{{ endpoint.base_url }}
</div>
</div>
<Badge variant="outline">
{{ endpoint.is_active ? '已启用' : '已禁用' }}
</Badge>
</div>
</button>
<div class="text-[11px] text-muted-foreground">
选择端点后会立即开始测试
</div>
</div>
<div
v-else
class="space-y-4"
>
<div
v-if="setupEndpoint"
class="rounded-lg border border-border/60 bg-muted/20 p-4 space-y-2"
>
<div class="text-xs text-muted-foreground">
本次测试端点
</div>
<div class="text-sm font-medium">
{{ formatApiFormat(setupEndpoint.api_format) }}
</div>
<div class="text-xs text-muted-foreground break-all">
{{ setupEndpoint.base_url }}
</div>
</div>
<div
v-else-if="mode === 'direct'"
class="rounded-lg border border-border/60 bg-muted/20 px-3 py-2 text-xs text-muted-foreground"
>
该测试会直接在当前 Provider 内执行故障转移并使用上方内容作为测试消息
</div>
<Button
class="w-full"
@click="emit('start')"
>
开始测试
</Button>
</div> </div>
</div> </div>
@@ -453,7 +509,7 @@
size="sm" size="sm"
@click="emit('close')" @click="emit('close')"
> >
{{ showSelection ? '取消' : '关闭' }} {{ showSetup ? '取消' : '关闭' }}
</Button> </Button>
</template> </template>
</Dialog> </Dialog>
@@ -464,6 +520,7 @@ import { computed, ref, watch } from 'vue'
import { Loader2 } from 'lucide-vue-next' import { Loader2 } from 'lucide-vue-next'
import { Dialog, Badge } from '@/components/ui' import { Dialog, Badge } from '@/components/ui'
import Button from '@/components/ui/button.vue' import Button from '@/components/ui/button.vue'
import Textarea from '@/components/ui/textarea.vue'
import { formatApiFormat } from '@/api/endpoints/types/api-format' import { formatApiFormat } from '@/api/endpoints/types/api-format'
import type { TestModelFailoverResponse, TestAttemptDetail } from '@/api/endpoints/providers' import type { TestModelFailoverResponse, TestAttemptDetail } from '@/api/endpoints/providers'
import type { CandidateRecord, RequestTrace } from '@/api/requestTrace' import type { CandidateRecord, RequestTrace } from '@/api/requestTrace'
@@ -486,19 +543,29 @@ const props = defineProps<{
trace?: RequestTrace | null trace?: RequestTrace | null
requestId?: string | null requestId?: string | null
showEndpointSelector?: boolean showEndpointSelector?: boolean
messageDraft?: string
}>() }>()
const emit = defineEmits<{ const emit = defineEmits<{
close: [] close: []
back: [] back: []
'select-endpoint': [endpointId: string] start: []
selectEndpoint: [endpointId: string]
'update:messageDraft': [value: string]
}>() }>()
const endpoints = computed(() => props.endpoints ?? []) const endpoints = computed(() => props.endpoints ?? [])
const traceCandidates = computed(() => props.trace?.candidates ?? []) const traceCandidates = computed(() => props.trace?.candidates ?? [])
const showSelection = computed(() => props.open && !!props.showEndpointSelector && !props.testing && !props.result) const messageDraft = computed(() => props.messageDraft ?? '')
const showSetup = computed(() => props.open && !props.testing && !props.result)
const showEndpointChoices = computed(() => showSetup.value && !!props.showEndpointSelector && endpoints.value.length > 1)
const showResult = computed(() => !!props.result) const showResult = computed(() => !!props.result)
const canReselect = computed(() => !!props.showEndpointSelector && endpoints.value.length > 1) const canReselect = computed(() => !!props.showEndpointSelector && endpoints.value.length > 1)
const setupEndpoint = computed(() => {
if (props.selectedEndpoint) return props.selectedEndpoint
if (endpoints.value.length === 1) return endpoints.value[0]
return null
})
const dialogTitle = computed(() => { const dialogTitle = computed(() => {
if (props.result) return '模型测试结果' if (props.result) return '模型测试结果'
@@ -506,9 +573,12 @@ const dialogTitle = computed(() => {
}) })
const dialogDescription = computed(() => { const dialogDescription = computed(() => {
if (showSelection.value && props.selectingModelName) { if (showEndpointChoices.value && props.selectingModelName) {
return `${props.selectingModelName} 选择端点` return `${props.selectingModelName} 选择端点`
} }
if (showSetup.value && props.selectingModelName) {
return `${props.selectingModelName} 设置测试内容`
}
if (props.testing && props.selectedEndpoint) { if (props.testing && props.selectedEndpoint) {
return `正在通过 ${formatApiFormat(props.selectedEndpoint.api_format)} 测试 ${props.selectingModelName || '模型'}` return `正在通过 ${formatApiFormat(props.selectedEndpoint.api_format)} 测试 ${props.selectingModelName || '模型'}`
} }

View File

@@ -222,9 +222,12 @@
:trace="modelTest.testTrace.value" :trace="modelTest.testTrace.value"
:request-id="modelTest.requestId.value" :request-id="modelTest.requestId.value"
:show-endpoint-selector="activeEndpoints.length > 1" :show-endpoint-selector="activeEndpoints.length > 1"
:message-draft="testMessageDraft"
@close="handleTestDialogClose" @close="handleTestDialogClose"
@back="handleTestDialogBack" @back="handleTestDialogBack"
@start="handleStartPendingTest"
@select-endpoint="handleSelectTestEndpoint" @select-endpoint="handleSelectTestEndpoint"
@update:message-draft="testMessageDraft = $event"
/> />
</template> </template>
@@ -272,6 +275,7 @@ const localModels = ref<Model[]>([])
const togglingModelId = ref<string | null>(null) const togglingModelId = ref<string | null>(null)
const pendingTestModel = ref<Model | null>(null) const pendingTestModel = ref<Model | null>(null)
const selectedTestEndpoint = ref<ProviderEndpoint | null>(null) const selectedTestEndpoint = ref<ProviderEndpoint | null>(null)
const testMessageDraft = ref('')
const activeEndpoints = computed(() => (props.endpoints ?? []).filter(endpoint => endpoint.is_active)) const activeEndpoints = computed(() => (props.endpoints ?? []).filter(endpoint => endpoint.is_active))
const models = computed(() => props.models ?? localModels.value) const models = computed(() => props.models ?? localModels.value)
// 按名称排序的模型列表 // 按名称排序的模型列表
@@ -460,7 +464,7 @@ async function handleSelectTestEndpoint(endpointId: string) {
displayLabel: `${endpointPrefix}${modelName}`, displayLabel: `${endpointPrefix}${modelName}`,
apiFormat: endpoint.api_format, apiFormat: endpoint.api_format,
endpointId: endpoint.id, endpointId: endpoint.id,
message: 'hello', message: testMessageDraft.value,
concurrency: 5, concurrency: 5,
onSuccess: () => { onSuccess: () => {
pendingTestModel.value = null pendingTestModel.value = null
@@ -475,6 +479,13 @@ async function handleSelectTestEndpoint(endpointId: string) {
}) })
} }
async function handleStartPendingTest() {
if (modelTest.testing.value) return
const endpoint = activeEndpoints.value[0]
if (!endpoint) return
await handleSelectTestEndpoint(endpoint.id)
}
async function testModelConnection(model: Model) { async function testModelConnection(model: Model) {
if (modelTest.testing.value) return if (modelTest.testing.value) return
@@ -487,10 +498,6 @@ async function testModelConnection(model: Model) {
selectedTestEndpoint.value = null selectedTestEndpoint.value = null
modelTest.testResult.value = null modelTest.testResult.value = null
modelTest.dialogOpen.value = true modelTest.dialogOpen.value = true
if (activeEndpoints.value.length === 1) {
await handleSelectTestEndpoint(activeEndpoints.value[0].id)
}
} }
// 暴露给父组件 // 暴露给父组件

View File

@@ -226,7 +226,7 @@ class TestModelRequest(BaseModel):
api_key_id: str | None = None api_key_id: str | None = None
endpoint_id: str | None = None # 指定使用的端点ID endpoint_id: str | None = None # 指定使用的端点ID
stream: bool = False stream: bool = False
message: str | None = "你好" message: str | None = None
api_format: str | None = None # 指定使用的API格式如果不指定则使用端点的默认格式 api_format: str | None = None # 指定使用的API格式如果不指定则使用端点的默认格式
@@ -238,7 +238,7 @@ class TestModelFailoverRequest(BaseModel):
model_name: str # global 模式传 global_model_name, direct 模式传 provider_model_name model_name: str # global 模式传 global_model_name, direct 模式传 provider_model_name
api_format: str | None = None # 指定 API 格式endpoint signature api_format: str | None = None # 指定 API 格式endpoint signature
endpoint_id: str | None = None # 指定仅使用该端点测试 endpoint_id: str | None = None # 指定仅使用该端点测试
message: str | None = "Hello" message: str | None = None
request_id: str | None = None request_id: str | None = None
concurrency: int = Field(default=1, ge=1, le=20) concurrency: int = Field(default=1, ge=1, le=20)
@@ -277,6 +277,14 @@ class TestModelFailoverResponse(BaseModel):
# ============ Internal helpers ============ # ============ Internal helpers ============
DEFAULT_MODEL_TEST_MESSAGE = "Hello! This is a test message."
def _resolve_test_message(message: str | None) -> str:
normalized = str(message or "").strip()
return normalized or DEFAULT_MODEL_TEST_MESSAGE
def _test_check_response_has_error(resp: dict[str, Any]) -> bool: def _test_check_response_has_error(resp: dict[str, Any]) -> bool:
"""快速判断 check_endpoint 结果是否失败。""" """快速判断 check_endpoint 结果是否失败。"""
if resp.get("error"): if resp.get("error"):
@@ -984,9 +992,7 @@ async def test_model(
# 准备测试请求数据(优先使用流式) # 准备测试请求数据(优先使用流式)
check_request = { check_request = {
"model": request.model_name, "model": request.model_name,
"messages": [ "messages": [{"role": "user", "content": _resolve_test_message(request.message)}],
{"role": "user", "content": request.message or "Hello! This is a test message."}
],
"max_tokens": 30, "max_tokens": 30,
"temperature": 0.7, "temperature": 0.7,
"stream": True, "stream": True,
@@ -2281,7 +2287,7 @@ async def test_model_failover(
request_payload = { request_payload = {
"model": request.model_name, "model": request.model_name,
"messages": [{"role": "user", "content": request.message or "Hello"}], "messages": [{"role": "user", "content": _resolve_test_message(request.message)}],
"max_tokens": 30, "max_tokens": 30,
"temperature": 0.7, "temperature": 0.7,
"stream": True, "stream": True,

View File

@@ -5,11 +5,13 @@ from pydantic import ValidationError
from src.api.admin.provider_query import TestModelFailoverRequest as FailoverRequestModel from src.api.admin.provider_query import TestModelFailoverRequest as FailoverRequestModel
from src.api.admin.provider_query import ( from src.api.admin.provider_query import (
DEFAULT_MODEL_TEST_MESSAGE,
_build_direct_test_candidates, _build_direct_test_candidates,
_build_test_attempts_from_candidate_keys, _build_test_attempts_from_candidate_keys,
_filter_test_candidates_by_endpoint, _filter_test_candidates_by_endpoint,
_flatten_test_candidates_for_concurrency, _flatten_test_candidates_for_concurrency,
_require_test_endpoint_base_url, _require_test_endpoint_base_url,
_resolve_test_message,
_resolve_test_effective_model, _resolve_test_effective_model,
) )
from src.services.scheduling.schemas import PoolCandidate from src.services.scheduling.schemas import PoolCandidate
@@ -151,6 +153,16 @@ def test_test_model_failover_request_validates_concurrency_range() -> None:
) )
def test_resolve_test_message_uses_default_for_blank_input() -> None:
assert _resolve_test_message(None) == DEFAULT_MODEL_TEST_MESSAGE
assert _resolve_test_message("") == DEFAULT_MODEL_TEST_MESSAGE
assert _resolve_test_message(" ") == DEFAULT_MODEL_TEST_MESSAGE
def test_resolve_test_message_preserves_custom_input() -> None:
assert _resolve_test_message(" custom prompt ") == "custom prompt"
def test_require_test_endpoint_base_url_rejects_non_string() -> None: def test_require_test_endpoint_base_url_rejects_non_string() -> None:
endpoint = SimpleNamespace(id="ep-bad", api_format="claude:chat", base_url={"url": "https://x"}) endpoint = SimpleNamespace(id="ep-bad", api_format="claude:chat", base_url={"url": "https://x"})