fix: handle Xinference model credential context (#38348)

This commit is contained in:
AsperforMias 2026-07-03 12:12:00 +08:00 committed by GitHub
parent f4ec608ef4
commit 1bc279e7d4
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
8 changed files with 177 additions and 32 deletions

View File

@ -1188,7 +1188,17 @@ class ProviderConfiguration(BaseModel):
) )
credential_record = session.execute(stmt).scalar_one_or_none() credential_record = session.execute(stmt).scalar_one_or_none()
if not credential_record: if not credential_record:
raise ValueError("Credential record not found.") fallback_stmt = select(ProviderModelCredential).where(
ProviderModelCredential.id == credential_id,
ProviderModelCredential.tenant_id == self.tenant_id,
ProviderModelCredential.provider_name.in_(self._get_provider_names()),
)
credential_record = session.execute(fallback_stmt).scalar_one_or_none()
if not credential_record:
raise ValueError("Credential record not found.")
model = credential_record.model_name
model_type = ModelType(credential_record.model_type)
lb_stmt = select(LoadBalancingModelConfig).where( lb_stmt = select(LoadBalancingModelConfig).where(
LoadBalancingModelConfig.tenant_id == self.tenant_id, LoadBalancingModelConfig.tenant_id == self.tenant_id,

View File

@ -1337,6 +1337,29 @@ def test_create_update_delete_custom_model_credential_flow() -> None:
assert provider_model_record.credential_id is None assert provider_model_record.credential_id is None
assert mock_cache.return_value.delete.call_count == 2 assert mock_cache.return_value.delete.call_count == 2
session = Mock()
mismatched_credential_record = SimpleNamespace(
id="cred-2",
model_name="stored-model",
model_type=ModelType.TEXT_EMBEDDING,
)
provider_model_record = SimpleNamespace(id="model-2", credential_id="cred-2", updated_at=None)
session.execute.side_effect = [
_exec_result(scalar_one_or_none=None),
_exec_result(scalar_one_or_none=mismatched_credential_record),
_exec_result(scalars_all=[]),
_exec_result(scalar=1),
]
with _patched_session(session):
with patch.object(
ProviderConfiguration,
"_get_custom_model_record",
return_value=provider_model_record,
) as mock_get_model:
configuration.delete_custom_model_credential(ModelType.LLM, "request-model", "cred-2")
mock_get_model.assert_called_once_with(ModelType.TEXT_EMBEDDING, "stored-model", session=session)
session.delete.assert_any_call(mismatched_credential_record)
def test_add_model_credential_to_model_and_switch_custom_model_credential() -> None: def test_add_model_credential_to_model_and_switch_custom_model_credential() -> None:
configuration = _build_provider_configuration() configuration = _build_provider_configuration()

View File

@ -306,6 +306,14 @@ export type ModelCredential = {
current_credential_name?: string current_credential_name?: string
} }
export type ModelCredentialPayload = {
credentials: Record<string, unknown>
model: string
model_type: ModelTypeEnum
name?: string
credential_id?: string
}
export enum ModelModalModeEnum { export enum ModelModalModeEnum {
configProviderCredential = 'config-provider-credential', configProviderCredential = 'config-provider-credential',
configCustomModel = 'config-custom-model', configCustomModel = 'config-custom-model',

View File

@ -168,8 +168,6 @@ describe('useAuth', () => {
expect(mockDeleteProviderCredential).toHaveBeenCalledWith({ expect(mockDeleteProviderCredential).toHaveBeenCalledWith({
credential_id: 'cred-1', credential_id: 'cred-1',
model: 'gpt-4',
model_type: ModelTypeEnum.textGeneration,
}) })
expect(mockDeleteModelService).not.toHaveBeenCalled() expect(mockDeleteModelService).not.toHaveBeenCalled()
expect(onRemove).toHaveBeenCalledWith('cred-1') expect(onRemove).toHaveBeenCalledWith('cred-1')

View File

@ -6,6 +6,16 @@ import { useModelModalHandler, useRefreshModel } from '@/app/components/header/a
import { useDeleteModel } from '@/service/use-models' import { useDeleteModel } from '@/service/use-models'
import { useAuthService } from './use-auth-service' import { useAuthService } from './use-auth-service'
type ProviderCredentialOperationPayload = {
credential_id: string
}
type ModelCredentialOperationPayload = {
credential_id: string
model: string
model_type: CustomModel['model_type']
}
export const useAuth = (provider: ModelProvider, configurationMethod: ConfigurationMethodEnum, currentCustomConfigurationModelFixedFields?: CustomConfigurationModelFixedFields, extra: { export const useAuth = (provider: ModelProvider, configurationMethod: ConfigurationMethodEnum, currentCustomConfigurationModelFixedFields?: CustomConfigurationModelFixedFields, extra: {
isModelCredential?: boolean isModelCredential?: boolean
onUpdate?: (newPayload?: any, formValues?: Record<string, any>) => void onUpdate?: (newPayload?: any, formValues?: Record<string, any>) => void
@ -30,6 +40,18 @@ export const useAuth = (provider: ModelProvider, configurationMethod: Configurat
setDeleteModel(model) setDeleteModel(model)
pendingOperationModel.current = model pendingOperationModel.current = model
}, []) }, [])
const resolveModelContext = useCallback((model?: CustomModel | null): CustomModel | undefined => {
if (model)
return model
if (!currentCustomConfigurationModelFixedFields)
return undefined
return {
model: currentCustomConfigurationModelFixedFields.__model_name,
model_type: currentCustomConfigurationModelFixedFields.__model_type,
}
}, [currentCustomConfigurationModelFixedFields])
const openConfirmDelete = useCallback((credential?: Credential, model?: CustomModel) => { const openConfirmDelete = useCallback((credential?: Credential, model?: CustomModel) => {
if (credential) if (credential)
handleSetDeleteCredentialId(credential.credential_id) handleSetDeleteCredentialId(credential.credential_id)
@ -51,18 +73,32 @@ export const useAuth = (provider: ModelProvider, configurationMethod: Configurat
return return
try { try {
handleSetDoingAction(true) handleSetDoingAction(true)
await getActiveCredentialService(!!model)({ const modelContext = model ?? (isModelCredential ? resolveModelContext() : undefined)
credential_id: credential.credential_id, if (modelContext) {
model: model?.model, const activeModelCredential = getActiveCredentialService(true) as (
model_type: model?.model_type, payload: ModelCredentialOperationPayload,
}) ) => Promise<unknown>
await activeModelCredential({
credential_id: credential.credential_id,
model: modelContext.model,
model_type: modelContext.model_type,
})
}
else {
const activeProviderCredential = getActiveCredentialService(false) as (
payload: ProviderCredentialOperationPayload,
) => Promise<unknown>
await activeProviderCredential({
credential_id: credential.credential_id,
})
}
toast.success(t('api.actionSuccess', { ns: 'common' })) toast.success(t('api.actionSuccess', { ns: 'common' }))
handleRefreshModel(provider, undefined, true) handleRefreshModel(provider, undefined, true)
} }
finally { finally {
handleSetDoingAction(false) handleSetDoingAction(false)
} }
}, [getActiveCredentialService, t, handleSetDoingAction]) }, [getActiveCredentialService, isModelCredential, resolveModelContext, t, handleSetDoingAction])
const handleConfirmDelete = useCallback(async () => { const handleConfirmDelete = useCallback(async () => {
if (doingActionRef.current) if (doingActionRef.current)
return return
@ -74,12 +110,30 @@ export const useAuth = (provider: ModelProvider, configurationMethod: Configurat
handleSetDoingAction(true) handleSetDoingAction(true)
let payload: any = {} let payload: any = {}
if (pendingOperationCredentialId.current) { if (pendingOperationCredentialId.current) {
payload = { if (isModelCredential) {
credential_id: pendingOperationCredentialId.current, const modelContext = resolveModelContext(pendingOperationModel.current)
model: pendingOperationModel.current?.model, if (!modelContext)
model_type: pendingOperationModel.current?.model_type, return
payload = {
credential_id: pendingOperationCredentialId.current,
model: modelContext.model,
model_type: modelContext.model_type,
}
const deleteModelCredential = getDeleteCredentialService(true) as (
payload: ModelCredentialOperationPayload,
) => Promise<unknown>
await deleteModelCredential(payload)
}
else {
payload = {
credential_id: pendingOperationCredentialId.current,
}
const deleteProviderCredential = getDeleteCredentialService(false) as (
payload: ProviderCredentialOperationPayload,
) => Promise<unknown>
await deleteProviderCredential(payload)
} }
await getDeleteCredentialService(!!isModelCredential)(payload)
} }
if (!pendingOperationCredentialId.current && pendingOperationModel.current) { if (!pendingOperationCredentialId.current && pendingOperationModel.current) {
payload = { payload = {
@ -96,7 +150,7 @@ export const useAuth = (provider: ModelProvider, configurationMethod: Configurat
finally { finally {
handleSetDoingAction(false) handleSetDoingAction(false)
} }
}, [t, handleSetDoingAction, getDeleteCredentialService, isModelCredential, closeConfirmDelete, handleRefreshModel, provider, configurationMethod, deleteModelService]) }, [t, handleSetDoingAction, getDeleteCredentialService, isModelCredential, closeConfirmDelete, handleRefreshModel, provider, configurationMethod, deleteModelService, resolveModelContext])
const handleSaveCredential = useCallback(async (payload: Record<string, any>) => { const handleSaveCredential = useCallback(async (payload: Record<string, any>) => {
if (doingActionRef.current) if (doingActionRef.current)
return return

View File

@ -353,4 +353,44 @@ describe('ModelModal', () => {
expect(mockHandlers.openConfirmDelete).toHaveBeenCalledWith({ credential_id: 'remove-1' }, undefined) expect(mockHandlers.openConfirmDelete).toHaveBeenCalledWith({ credential_id: 'remove-1' }, undefined)
removable.unmount() removable.unmount()
}) })
it('should use fixed model context when saving a model credential without model prop', async () => {
mockState.formSchemas = [{ variable: 'api_key', type: 'secret-input' } as unknown as CredentialFormSchema]
mockFormState.responses = [
{ isCheckValidated: true, values: { __authorization_name__: 'Xinference Auth', api_key: 'secret' } },
]
renderModal({
mode: ModelModalModeEnum.configModelCredential,
currentCustomConfigurationModelFixedFields: {
__model_name: 'bge-m3',
__model_type: ModelTypeEnum.textEmbedding,
},
})
fireEvent.click(screen.getByRole('button', { name: 'common.operation.save' }))
await waitFor(() => {
expect(mockHandlers.handleSaveCredential).toHaveBeenCalledWith({
credential_id: undefined,
credentials: { api_key: 'secret' },
name: 'Xinference Auth',
model: 'bge-m3',
model_type: ModelTypeEnum.textEmbedding,
})
})
})
it('should not submit model credential payload when model context is missing', async () => {
mockState.formSchemas = [{ variable: 'api_key', type: 'secret-input' } as unknown as CredentialFormSchema]
mockFormState.responses = [
{ isCheckValidated: true, values: { __authorization_name__: 'Missing Model Auth', api_key: 'secret' } },
]
renderModal({ mode: ModelModalModeEnum.configModelCredential })
fireEvent.click(screen.getByRole('button', { name: 'common.operation.save' }))
await waitFor(() => {
expect(mockHandlers.handleSaveCredential).not.toHaveBeenCalled()
})
})
}) })

View File

@ -147,17 +147,22 @@ const ModelModal: FC<ModelModalProps> = ({
modelNameAndTypeValues = formResult.values modelNameAndTypeValues = formResult.values
} }
if (mode === ModelModalModeEnum.configModelCredential && model) { if (
modelNameAndTypeValues = { mode === ModelModalModeEnum.configModelCredential
__model_name: model.model, || (mode === ModelModalModeEnum.addCustomModelToModelList && selectedCredential?.addNewCredential)
__model_type: model.model_type, ) {
} const modelContext = model ?? (currentCustomConfigurationModelFixedFields
} ? {
model: currentCustomConfigurationModelFixedFields.__model_name,
model_type: currentCustomConfigurationModelFixedFields.__model_type,
}
: undefined)
if (!modelContext)
return
if (mode === ModelModalModeEnum.addCustomModelToModelList && selectedCredential?.addNewCredential && model) {
modelNameAndTypeValues = { modelNameAndTypeValues = {
__model_name: model.model, __model_name: modelContext.model,
__model_type: model.model_type, __model_type: modelContext.model_type,
} }
} }
const { const {
@ -178,7 +183,13 @@ const ModelModal: FC<ModelModalProps> = ({
__authorization_name__, __authorization_name__,
...rest ...rest
} = values } = values
if (__model_name && __model_type) { const shouldSaveModelCredential = mode === ModelModalModeEnum.configCustomModel
|| mode === ModelModalModeEnum.configModelCredential
|| (mode === ModelModalModeEnum.addCustomModelToModelList && selectedCredential?.addNewCredential)
if (shouldSaveModelCredential) {
if (!__model_name || !__model_type)
return
await handleSaveCredential({ await handleSaveCredential({
credential_id: credential?.credential_id, credential_id: credential?.credential_id,
credentials: rest, credentials: rest,
@ -195,7 +206,7 @@ const ModelModal: FC<ModelModalProps> = ({
}) })
} }
onSave(values) onSave(values)
}, [mode, selectedCredential, model, canUseCredential, canCreateCredential, canManageCredential, onSave, handleActiveCredential, onCancel, handleSaveCredential, credential]) }, [mode, selectedCredential, model, currentCustomConfigurationModelFixedFields, canUseCredential, canCreateCredential, canManageCredential, onSave, handleActiveCredential, onCancel, handleSaveCredential, credential])
const modalTitle = useMemo(() => { const modalTitle = useMemo(() => {
let label = t('modelProvider.auth.apiKeyModal.title', { ns: 'common' }) let label = t('modelProvider.auth.apiKeyModal.title', { ns: 'common' })

View File

@ -1,5 +1,6 @@
import type { import type {
ModelCredential, ModelCredential,
ModelCredentialPayload,
ModelItem, ModelItem,
ModelLoadBalancingConfig, ModelLoadBalancingConfig,
ModelTypeEnum, ModelTypeEnum,
@ -91,7 +92,7 @@ export const useGetModelCredential = (
export const useAddModelCredential = (provider: string) => { export const useAddModelCredential = (provider: string) => {
return useMutation({ return useMutation({
mutationFn: (data: ModelCredential) => post<{ result: string }>(`/workspaces/current/model-providers/${provider}/models/credentials`, { mutationFn: (data: ModelCredentialPayload) => post<{ result: string }>(`/workspaces/current/model-providers/${provider}/models/credentials`, {
body: data, body: data,
}), }),
}) })
@ -99,7 +100,7 @@ export const useAddModelCredential = (provider: string) => {
export const useEditModelCredential = (provider: string) => { export const useEditModelCredential = (provider: string) => {
return useMutation({ return useMutation({
mutationFn: (data: ModelCredential) => put<{ result: string }>(`/workspaces/current/model-providers/${provider}/models/credentials`, { mutationFn: (data: ModelCredentialPayload) => put<{ result: string }>(`/workspaces/current/model-providers/${provider}/models/credentials`, {
body: data, body: data,
}), }),
}) })
@ -109,8 +110,8 @@ export const useDeleteModelCredential = (provider: string) => {
return useMutation({ return useMutation({
mutationFn: (data: { mutationFn: (data: {
credential_id: string credential_id: string
model?: string model: string
model_type?: ModelTypeEnum model_type: ModelTypeEnum
}) => del<{ result: string }>(`/workspaces/current/model-providers/${provider}/models/credentials`, { }) => del<{ result: string }>(`/workspaces/current/model-providers/${provider}/models/credentials`, {
body: data, body: data,
}), }),
@ -132,8 +133,8 @@ export const useActiveModelCredential = (provider: string) => {
return useMutation({ return useMutation({
mutationFn: (data: { mutationFn: (data: {
credential_id: string credential_id: string
model?: string model: string
model_type?: ModelTypeEnum model_type: ModelTypeEnum
}) => post<{ result: string }>(`/workspaces/current/model-providers/${provider}/models/credentials/switch`, { }) => post<{ result: string }>(`/workspaces/current/model-providers/${provider}/models/credentials/switch`, {
body: data, body: data,
}), }),