fix(provider): support multi-key selection in model tests

This commit is contained in:
zhefox
2026-05-26 13:08:35 +08:00
parent 5fc6dc8019
commit 0bf63cc80e
7 changed files with 253 additions and 26 deletions
@@ -1,5 +1,5 @@
use super::super::payload::{
provider_query_extract_api_key_id, provider_query_extract_force_refresh,
provider_query_extract_api_key_ids, provider_query_extract_force_refresh,
provider_query_extract_model, provider_query_extract_provider_id,
provider_query_extract_request_id,
};
@@ -860,7 +860,7 @@ async fn provider_query_select_preferred_non_kiro_endpoint(
provider: &StoredProviderCatalogProvider,
endpoints: &[StoredProviderCatalogEndpoint],
keys: &[StoredProviderCatalogKey],
selected_key_id: Option<&str>,
selected_key_ids: Option<&BTreeSet<String>>,
) -> Option<StoredProviderCatalogEndpoint> {
for priority in 0..=2 {
for endpoint in endpoints.iter().filter(|endpoint| endpoint.is_active) {
@@ -873,7 +873,7 @@ async fn provider_query_select_preferred_non_kiro_endpoint(
}
for key in keys {
if !key.is_active
|| selected_key_id.is_some_and(|value| value != key.id.as_str())
|| !provider_query_selected_key_ids_allow_key(selected_key_ids, &key.id)
|| !provider_query_key_supports_endpoint(
key,
&provider.provider_type,
@@ -905,7 +905,7 @@ async fn provider_query_select_preferred_non_kiro_endpoint(
endpoint.is_active
&& keys.iter().any(|key| {
key.is_active
&& selected_key_id.is_none_or(|value| value == key.id.as_str())
&& provider_query_selected_key_ids_allow_key(selected_key_ids, &key.id)
&& provider_query_key_supports_endpoint(
key,
&provider.provider_type,
@@ -917,6 +917,22 @@ async fn provider_query_select_preferred_non_kiro_endpoint(
.cloned()
}
fn provider_query_selected_key_ids_allow_key(
selected_key_ids: Option<&BTreeSet<String>>,
key_id: &str,
) -> bool {
selected_key_ids.is_none_or(|ids| ids.contains(key_id))
}
fn provider_query_selected_key_ids_all_exist(
selected_key_ids: &BTreeSet<String>,
keys: &[StoredProviderCatalogKey],
) -> bool {
selected_key_ids
.iter()
.all(|id| keys.iter().any(|key| key.id == *id))
}
fn provider_query_test_key_sort_key(
provider_type: &str,
key: &StoredProviderCatalogKey,
@@ -1262,7 +1278,7 @@ async fn provider_query_build_kiro_test_candidates(
ADMIN_PROVIDER_QUERY_NO_ACTIVE_API_KEY_DETAIL,
)
})?;
let selected_key_id = provider_query_extract_api_key_id(payload);
let selected_key_ids = provider_query_extract_api_key_ids(payload);
let requested_endpoint_id = provider_query_extract_endpoint_id(payload);
let requested_api_format = provider_query_extract_api_format(payload);
let endpoint = if requested_endpoint_id.is_none()
@@ -1274,7 +1290,7 @@ async fn provider_query_build_kiro_test_candidates(
provider,
&endpoints,
&all_keys,
selected_key_id.as_deref(),
selected_key_ids.as_ref(),
)
.await
.ok_or_else(|| {
@@ -1307,22 +1323,11 @@ async fn provider_query_build_kiro_test_candidates(
}
};
if let Some(api_key_id) = selected_key_id.as_deref() {
let Some(key) = all_keys.iter().find(|key| key.id == api_key_id) else {
if let Some(selected_key_ids) = selected_key_ids.as_ref() {
if !provider_query_selected_key_ids_all_exist(selected_key_ids, &all_keys) {
return Err(build_admin_provider_query_not_found_response(
ADMIN_PROVIDER_QUERY_API_KEY_NOT_FOUND_DETAIL,
));
};
if !key.is_active
|| !provider_query_key_supports_endpoint(
key,
&provider.provider_type,
&endpoint.api_format,
)
{
return Err(build_admin_provider_query_not_found_response(
ADMIN_PROVIDER_QUERY_NO_ACTIVE_TEST_CANDIDATE_DETAIL,
));
}
}
@@ -1380,11 +1385,7 @@ async fn provider_query_build_kiro_test_candidates(
let mut keys = all_keys
.into_iter()
.filter(|key| key.is_active)
.filter(|key| {
selected_key_id
.as_deref()
.is_none_or(|value| value == key.id.as_str())
})
.filter(|key| provider_query_selected_key_ids_allow_key(selected_key_ids.as_ref(), &key.id))
.filter(|key| {
provider_query_key_supports_endpoint(key, &provider.provider_type, &endpoint.api_format)
})
@@ -231,6 +231,27 @@ fn provider_query_request_body_model_uses_non_empty_string_only() {
);
}
#[test]
fn provider_query_model_test_extracts_multiple_selected_key_ids() {
let payload = json!({
"api_key_ids": [" key-b ", "", "key-a", "key-b"],
"api_key_id": "key-c"
});
let ids = provider_query_extract_api_key_ids(&payload)
.expect("non-empty key selection should be extracted")
.into_iter()
.collect::<Vec<_>>();
assert_eq!(ids, vec!["key-a", "key-b", "key-c"]);
}
#[test]
fn provider_query_model_test_empty_selected_key_ids_keep_default_selection() {
assert!(provider_query_extract_api_key_ids(&json!({})).is_none());
assert!(provider_query_extract_api_key_ids(&json!({ "api_key_ids": [] })).is_none());
}
#[test]
fn provider_query_standard_test_resolves_codex_responses_upstream_streaming() {
assert!(provider_query_resolve_standard_test_upstream_is_stream(
@@ -1,6 +1,7 @@
use axum::body::Bytes;
use axum::response::{IntoResponse, Response};
use serde_json::json;
use std::collections::BTreeSet;
pub(crate) fn parse_admin_provider_query_body(
request_body: Option<&Bytes>,
@@ -36,6 +37,47 @@ pub(crate) fn provider_query_extract_api_key_id(payload: &serde_json::Value) ->
.map(ToOwned::to_owned)
}
fn provider_query_insert_api_key_id(ids: &mut BTreeSet<String>, value: &str) {
let value = value.trim();
if !value.is_empty() {
ids.insert(value.to_string());
}
}
pub(crate) fn provider_query_extract_api_key_ids(
payload: &serde_json::Value,
) -> Option<BTreeSet<String>> {
let mut ids = BTreeSet::new();
if let Some(value) = payload
.get("api_key_ids")
.or_else(|| payload.get("provider_key_ids"))
.or_else(|| payload.get("key_ids"))
{
match value {
serde_json::Value::Array(items) => {
for item in items {
if let Some(value) = item.as_str() {
provider_query_insert_api_key_id(&mut ids, value);
}
}
}
serde_json::Value::String(value) => {
for item in value.split(',') {
provider_query_insert_api_key_id(&mut ids, item);
}
}
_ => {}
}
}
if let Some(api_key_id) = provider_query_extract_api_key_id(payload) {
ids.insert(api_key_id);
}
(!ids.is_empty()).then_some(ids)
}
pub(crate) fn provider_query_extract_force_refresh(payload: &serde_json::Value) -> bool {
payload
.get("force_refresh")
+2
View File
@@ -196,6 +196,7 @@ export interface TestModelRequest {
provider_id: string
model_name: string
api_key_id?: string
api_key_ids?: string[]
endpoint_id?: string
message?: string
api_format?: string
@@ -249,6 +250,7 @@ export interface TestModelFailoverRequest {
mode: 'global' | 'direct' | 'pool'
model_name: string
failover_models?: string[]
api_key_ids?: string[]
api_format?: string
endpoint_id?: string
message?: string
+15 -1
View File
@@ -20,6 +20,7 @@ export interface StartTestParams {
endpointId?: string
endpointBaseUrl?: string
message?: string
apiKeyIds?: string[]
applyModelMapping?: boolean
mappedModelName?: string
requestHeaders?: Record<string, unknown>
@@ -150,13 +151,17 @@ export function useModelTest(options: UseModelTestOptions) {
reqId: string,
signal?: AbortSignal,
): Promise<TestModelFailoverResponse> {
const message = normalizedMessage(params.message)
const apiKeyIds = normalizedApiKeyIds(params.apiKeyIds)
return normalizeDirectTestResult(params, await testModel({
provider_id: providerId(),
model_name: params.modelName,
mode: params.mode,
api_format: params.apiFormat,
endpoint_id: params.endpointId,
...(normalizedMessage(params.message) ? { message: normalizedMessage(params.message) } : {}),
...(apiKeyIds ? { api_key_ids: apiKeyIds } : {}),
...(message ? { message } : {}),
...(typeof params.applyModelMapping === 'boolean' ? { apply_model_mapping: params.applyModelMapping } : {}),
...(params.mappedModelName ? { mapped_model_name: params.mappedModelName } : {}),
...(params.requestHeaders ? { request_headers: params.requestHeaders } : {}),
@@ -173,6 +178,13 @@ export function useModelTest(options: UseModelTestOptions) {
: undefined
}
function normalizedApiKeyIds(apiKeyIds?: string[]): string[] | undefined {
const ids = Array.isArray(apiKeyIds)
? apiKeyIds.map(item => item.trim()).filter(Boolean)
: []
return ids.length > 0 ? [...new Set(ids)] : undefined
}
async function pollTestTrace(reqId: string, token: number) {
try {
const trace = await requestTraceApi.getRequestTrace(reqId, { attemptedOnly: false })
@@ -249,6 +261,7 @@ export function useModelTest(options: UseModelTestOptions) {
try {
const message = normalizedMessage(params.message)
const apiKeyIds = normalizedApiKeyIds(params.apiKeyIds)
let result = params.mode === 'direct'
? await runDirectTest(params, reqId, abortController.signal)
@@ -257,6 +270,7 @@ export function useModelTest(options: UseModelTestOptions) {
mode: params.mode,
model_name: params.modelName,
failover_models: [params.modelName],
...(apiKeyIds ? { api_key_ids: apiKeyIds } : {}),
api_format: params.apiFormat,
endpoint_id: params.endpointId,
...(message ? { message } : {}),
@@ -99,6 +99,33 @@
</div>
</div>
<div
v-if="showKeySelector"
class="space-y-2"
>
<div class="flex items-center justify-between gap-3">
<div class="text-sm font-medium text-foreground">
测试 Key
</div>
<div class="text-xs text-muted-foreground">
{{ keySelectionStatus }}
</div>
</div>
<MultiSelect
:model-value="selectedKeyIds"
:options="keyOptions"
:placeholder="keySelectorPlaceholder"
search-placeholder="搜索 Key"
empty-text="暂无可选 Key"
no-results-text="未找到匹配 Key"
trigger-class="h-9 min-h-9 rounded-md border-border/60 text-xs"
dropdown-min-width="24rem"
:search-threshold="0"
:disabled="keyOptionsLoading && keyOptions.length === 0"
@update:model-value="emit('update:selectedKeyIds', $event)"
/>
</div>
<div class="grid gap-4 lg:grid-cols-2 lg:items-start">
<div class="space-y-2">
<div class="flex items-center justify-between gap-3">
@@ -772,9 +799,11 @@ import {
} from '@/components/ui'
import Button from '@/components/ui/button.vue'
import Textarea from '@/components/ui/textarea.vue'
import MultiSelect from '@/components/common/MultiSelect.vue'
import { formatApiFormat } from '@/api/endpoints/types/api-format'
import type { TestAttemptDetail, TestCandidateSummary, TestModelFailoverResponse } from '@/api/endpoints/providers'
import type { CandidateRecord, RequestTrace } from '@/api/requestTrace'
import type { MultiSelectOption } from '@/components/common/MultiSelect.vue'
import JsonContent from '@/features/usage/components/RequestDetailDrawer/JsonContent.vue'
import { useClipboard } from '@/composables/useClipboard'
import { useDarkMode } from '@/composables/useDarkMode'
@@ -797,6 +826,8 @@ type TestModelMappingOption = {
priority?: number
}
type TestKeyOption = MultiSelectOption
const props = defineProps<{
open: boolean
result: TestModelFailoverResponse | null
@@ -818,6 +849,9 @@ const props = defineProps<{
modelMappingAvailable?: boolean
modelMappingOptions?: TestModelMappingOption[]
selectedModelMapping?: string | null
keyOptions?: TestKeyOption[]
selectedKeyIds?: string[]
keyOptionsLoading?: boolean
startDisabled?: boolean
}>()
@@ -827,15 +861,30 @@ const emit = defineEmits<{
start: []
selectEndpoint: [endpointId: string]
selectModelMapping: [modelName: string]
'update:selectedKeyIds': [value: string[]]
'update:requestHeadersDraft': [value: string]
'update:requestBodyDraft': [value: string]
}>()
const endpoints = computed(() => props.endpoints ?? [])
const modelMappingOptions = computed(() => props.modelMappingOptions ?? [])
const keyOptions = computed(() => props.keyOptions ?? [])
const selectedKeyIds = computed(() => props.selectedKeyIds ?? [])
const keyOptionsLoading = computed(() => props.keyOptionsLoading === true)
const modelMappingAvailable = computed(
() => props.modelMappingAvailable === true && modelMappingOptions.value.length > 0,
)
const showKeySelector = computed(() => (
keyOptionsLoading.value || keyOptions.value.length > 0 || selectedKeyIds.value.length > 0
))
const keySelectorPlaceholder = computed(() => (
keyOptionsLoading.value && keyOptions.value.length === 0 ? '正在加载 Key' : '默认调度(不指定 Key'
))
const keySelectionStatus = computed(() => {
if (selectedKeyIds.value.length > 0) return `已选 ${selectedKeyIds.value.length}`
if (keyOptionsLoading.value) return '加载中'
return '默认'
})
const requestedModelName = computed(() => props.requestedModelName?.trim() || '')
const selectedModelMapping = computed(() => props.selectedModelMapping?.trim() || '')
const selectedModelMappingValue = computed(() => (
@@ -232,12 +232,16 @@
:model-mapping-available="testModelMappingAvailable"
:model-mapping-options="testModelMappingOptions"
:selected-model-mapping="selectedTestMappedModelName"
:key-options="testKeyOptions"
:selected-key-ids="selectedTestKeyIds"
:key-options-loading="loadingModelTestKeys"
:start-disabled="!selectedTestEndpoint || !!testRequestHeadersError || !!testRequestBodyError"
@close="handleTestDialogClose"
@back="handleTestDialogBack"
@start="handleStartPendingTest"
@select-endpoint="handleSelectTestEndpoint"
@select-model-mapping="handleSelectModelMapping"
@update:selected-key-ids="handleSelectTestKeyIds"
@update:request-headers-draft="testRequestHeadersDraft = $event"
@update:request-body-draft="testRequestBodyDraft = $event"
/>
@@ -257,7 +261,7 @@ import {
type Model,
type ProviderEndpoint,
} from '@/api/endpoints'
import { type EndpointAPIKey } from '@/api/endpoints/keys'
import { getProviderKeys, type EndpointAPIKey } from '@/api/endpoints/keys'
import { updateModel } from '@/api/endpoints/models'
import { parseApiError } from '@/utils/errorParser'
import { formatApiFormat } from '@/api/endpoints/types/api-format'
@@ -269,6 +273,7 @@ import {
isModelTestableApiFormat,
isModelTestableEndpoint,
listModelTestMappedModelOptions,
modelTestKeySupportsEndpoint,
normalizeModelTestMappedModelSelection,
parseModelTestRequestHeadersDraft,
parseModelTestRequestBodyDraft,
@@ -307,6 +312,10 @@ const testRequestHeadersResetValue = ref('')
const testRequestBodyDraft = ref('')
const testRequestBodyResetValue = ref('')
const selectedTestMappedModelName = ref<string | null>(null)
const selectedTestKeyIds = ref<string[]>([])
const modelTestProviderKeys = ref<EndpointAPIKey[]>([])
const modelTestKeysLoadedProviderId = ref<string | null>(null)
const loadingModelTestKeys = ref(false)
const isPoolManagedProvider = computed(() => Boolean(props.provider.pool_advanced))
const activeEndpoints = computed(() => (props.endpoints ?? [])
.filter(endpoint => {
@@ -336,6 +345,32 @@ const mappedTestModelName = computed(() => {
: null
})
const testModelMappingAvailable = computed(() => testModelMappingOptions.value.length > 0)
const providerKeysForModelTest = computed(() => (
modelTestKeysLoadedProviderId.value === props.provider.id
? modelTestProviderKeys.value
: props.providerKeys ?? []
))
const testKeyOptions = computed(() => {
const endpoint = selectedTestEndpoint.value
if (!endpoint) return []
const seen = new Set<string>()
return [...providerKeysForModelTest.value]
.filter((key) => {
if (seen.has(key.id)) return false
seen.add(key.id)
return modelTestKeySupportsEndpoint(key, endpoint, props.provider.provider_type)
})
.sort((left, right) => {
const priority = left.internal_priority - right.internal_priority
if (priority !== 0) return priority
return formatTestKeyOptionLabel(left).localeCompare(formatTestKeyOptionLabel(right))
})
.map(key => ({
value: key.id,
label: formatTestKeyOptionLabel(key),
}))
})
const effectiveTestRequestModelName = computed(() => (
mappedTestModelName.value || pendingRequestedModelName.value
))
@@ -506,6 +541,7 @@ function handleTestDialogClose() {
pendingTestModel.value = null
selectedTestEndpoint.value = null
selectedTestMappedModelName.value = null
selectedTestKeyIds.value = []
testRequestHeadersDraft.value = ''
testRequestHeadersResetValue.value = ''
testRequestBodyDraft.value = ''
@@ -524,6 +560,7 @@ function handleSelectTestEndpoint(endpointId: string) {
selectedTestEndpoint.value = endpoint
syncSelectedTestModelMapping()
resetTestRequestBodyForSelectedEndpoint()
pruneSelectedTestKeyIds()
}
function handleSelectModelMapping(modelName: string) {
@@ -534,6 +571,10 @@ function handleSelectModelMapping(modelName: string) {
syncTestRequestBodyModel()
}
function handleSelectTestKeyIds(ids: string[]) {
selectedTestKeyIds.value = normalizeSelectedTestKeyIds(ids)
}
async function handleStartPendingTest() {
if (modelTest.testing.value) return
if (!pendingTestModel.value) return
@@ -557,6 +598,7 @@ async function handleStartPendingTest() {
}
selectedTestEndpoint.value = endpoint
pruneSelectedTestKeyIds()
const model = pendingTestModel.value
const modelName = model.global_model_name || model.provider_model_name
const endpointPrefix = `[${formatApiFormat(endpoint.api_format)}] `
@@ -567,6 +609,7 @@ async function handleStartPendingTest() {
apiFormat: endpoint.api_format,
endpointId: endpoint.id,
endpointBaseUrl: endpoint.base_url,
apiKeyIds: selectedTestKeyIds.value,
applyModelMapping: Boolean(mappedTestModelName.value),
mappedModelName: mappedTestModelName.value ?? undefined,
requestHeaders,
@@ -591,6 +634,7 @@ async function testModelConnection(model: Model) {
selectedTestEndpoint.value = selectPreferredModelTestEndpoint(model, activeEndpoints.value)
const requestedModelName = getModelTestRequestedModelName(model)
selectedTestMappedModelName.value = null
selectedTestKeyIds.value = []
testRequestHeadersResetValue.value = buildDefaultModelTestRequestHeaders()
testRequestHeadersDraft.value = testRequestHeadersResetValue.value
testRequestBodyResetValue.value = buildDefaultModelTestRequestBody(
@@ -601,6 +645,49 @@ async function testModelConnection(model: Model) {
testRequestBodyDraft.value = testRequestBodyResetValue.value
modelTest.testResult.value = null
modelTest.dialogOpen.value = true
void ensureModelTestKeysLoaded()
}
function normalizeSelectedTestKeyIds(ids: string[]): string[] {
const allowed = new Set(testKeyOptions.value.map(option => option.value))
const selected = ids
.map(id => id.trim())
.filter(id => id && allowed.has(id))
return [...new Set(selected)]
}
function pruneSelectedTestKeyIds() {
if (selectedTestKeyIds.value.length === 0) return
selectedTestKeyIds.value = normalizeSelectedTestKeyIds(selectedTestKeyIds.value)
}
async function ensureModelTestKeysLoaded() {
if (modelTestKeysLoadedProviderId.value === props.provider.id || loadingModelTestKeys.value) {
return
}
loadingModelTestKeys.value = true
try {
modelTestProviderKeys.value = await getProviderKeys(props.provider.id)
modelTestKeysLoadedProviderId.value = props.provider.id
pruneSelectedTestKeyIds()
} catch (err: unknown) {
showError(parseApiError(err, '加载测试 Key 失败'), '错误')
} finally {
loadingModelTestKeys.value = false
}
}
function formatTestKeyOptionLabel(key: EndpointAPIKey): string {
const name = key.name?.trim()
const masked = key.api_key_masked?.trim()
const authType = key.auth_type?.trim()
const primary = name || masked || key.id
const suffix = [
masked && masked !== primary ? masked : '',
authType || '',
].filter(Boolean)
return suffix.length > 0 ? `${primary} · ${suffix.join(' · ')}` : primary
}
function getModelTestRequestedModelName(model: Model | null): string {
@@ -661,6 +748,17 @@ watch(
() => syncTestRequestBodyModel(),
)
watch(testKeyOptions, () => pruneSelectedTestKeyIds())
watch(
() => props.provider.id,
() => {
modelTestProviderKeys.value = []
modelTestKeysLoadedProviderId.value = null
selectedTestKeyIds.value = []
},
)
// 暴露给父组件
defineExpose({
reload: refresh