feat(frontend): 支持提供商预设与自动获取模型
This commit is contained in:
@@ -0,0 +1,79 @@
|
||||
// @vitest-environment happy-dom
|
||||
import { beforeEach, describe, expect, it, vi } from 'vitest'
|
||||
import { createPinia, setActivePinia } from 'pinia'
|
||||
import type { ProviderConfig, ProviderPreset } from '@/contracts'
|
||||
|
||||
vi.mock('@/services/providerService', () => ({
|
||||
mockProviders: [],
|
||||
mockModels: {},
|
||||
listProviders: vi.fn(),
|
||||
listProviderPresets: vi.fn(),
|
||||
listModels: vi.fn(),
|
||||
createProvider: vi.fn(),
|
||||
updateProvider: vi.fn(),
|
||||
deleteProvider: vi.fn(),
|
||||
testProvider: vi.fn(),
|
||||
}))
|
||||
|
||||
import { useProviderStore } from './provider'
|
||||
import { listModels, listProviderPresets, listProviders } from '@/services/providerService'
|
||||
|
||||
const providers: ProviderConfig[] = [
|
||||
{
|
||||
provider_id: 'openai',
|
||||
provider_type: 'openai_chat',
|
||||
name: 'OpenAI',
|
||||
base_url: 'https://api.openai.com/v1',
|
||||
default_model: '',
|
||||
enabled: true,
|
||||
capabilities: { chat: true },
|
||||
has_credential: true,
|
||||
},
|
||||
]
|
||||
|
||||
const presets: ProviderPreset[] = [
|
||||
{
|
||||
preset_id: 'deepseek',
|
||||
name: 'DeepSeek',
|
||||
provider_type: 'openai_compatible',
|
||||
base_url: 'https://api.deepseek.com',
|
||||
requires_credential: true,
|
||||
},
|
||||
]
|
||||
|
||||
beforeEach(() => {
|
||||
setActivePinia(createPinia())
|
||||
vi.clearAllMocks()
|
||||
vi.mocked(listProviders).mockResolvedValue(providers)
|
||||
vi.mocked(listProviderPresets).mockResolvedValue(presets)
|
||||
})
|
||||
|
||||
describe('provider store model discovery', () => {
|
||||
it('loads provider presets and automatically refreshes enabled providers', async () => {
|
||||
vi.mocked(listModels).mockResolvedValue([
|
||||
{ model_id: 'model-z', name: 'Zulu', capabilities: { chat: true } },
|
||||
{ model_id: 'model-a', name: 'Alpha', capabilities: { chat: true } },
|
||||
{ model_id: 'model-a', name: 'Alpha duplicate', capabilities: { chat: true } },
|
||||
])
|
||||
const store = useProviderStore()
|
||||
|
||||
await store.loadProviders()
|
||||
await store.loadPresets()
|
||||
await store.refreshEnabledModels()
|
||||
|
||||
expect(store.presets[0].preset_id).toBe('deepseek')
|
||||
expect(listModels).toHaveBeenCalledWith('openai')
|
||||
expect(store.modelsByProvider.openai.map((model) => model.model_id)).toEqual(['model-a', 'model-z'])
|
||||
expect(store.modelErrorsByProvider.openai).toBeUndefined()
|
||||
})
|
||||
|
||||
it('records a provider-specific error when model discovery fails', async () => {
|
||||
vi.mocked(listModels).mockRejectedValue(new Error('认证失败'))
|
||||
const store = useProviderStore()
|
||||
|
||||
await expect(store.loadModels('openai')).rejects.toThrow('认证失败')
|
||||
|
||||
expect(store.modelLoadingByProvider.openai).toBe(false)
|
||||
expect(store.modelErrorsByProvider.openai).toBe('认证失败')
|
||||
})
|
||||
})
|
||||
@@ -1,11 +1,14 @@
|
||||
import { defineStore } from 'pinia'
|
||||
import { ref, computed } from 'vue'
|
||||
import type { ProviderConfig, ModelInfo } from '@/contracts'
|
||||
import { createProvider, deleteProvider as deleteProviderRequest, listModels, listProviders, mockProviders, mockModels, testProvider as testProviderRequest, updateProvider as updateProviderRequest } from '@/services/providerService'
|
||||
import type { ProviderConfig, ModelInfo, ProviderPreset } from '@/contracts'
|
||||
import { createProvider, deleteProvider as deleteProviderRequest, listModels, listProviderPresets, listProviders, mockProviders, mockModels, testProvider as testProviderRequest, updateProvider as updateProviderRequest } from '@/services/providerService'
|
||||
|
||||
export const useProviderStore = defineStore('provider', () => {
|
||||
const providers = ref<ProviderConfig[]>(mockProviders)
|
||||
const presets = ref<ProviderPreset[]>([])
|
||||
const modelsByProvider = ref<Record<string, ModelInfo[]>>(mockModels)
|
||||
const modelLoadingByProvider = ref<Record<string, boolean>>({})
|
||||
const modelErrorsByProvider = ref<Record<string, string>>({})
|
||||
const defaultProviderId = ref('mock')
|
||||
const isLoading = ref(false)
|
||||
const error = ref<string | null>(null)
|
||||
@@ -27,8 +30,36 @@ export const useProviderStore = defineStore('provider', () => {
|
||||
}
|
||||
}
|
||||
|
||||
async function loadModels(providerId: string) {
|
||||
modelsByProvider.value[providerId] = await listModels(providerId)
|
||||
async function loadPresets() {
|
||||
try {
|
||||
presets.value = await listProviderPresets()
|
||||
} catch (reason) {
|
||||
error.value = reason instanceof Error ? reason.message : 'Provider 预设加载失败'
|
||||
}
|
||||
}
|
||||
|
||||
async function loadModels(providerId: string): Promise<ModelInfo[]> {
|
||||
modelLoadingByProvider.value[providerId] = true
|
||||
delete modelErrorsByProvider.value[providerId]
|
||||
try {
|
||||
const models = await listModels(providerId)
|
||||
const uniqueModels = [...new Map(models.map((model) => [model.model_id, model])).values()]
|
||||
.sort((left, right) => left.name.localeCompare(right.name))
|
||||
modelsByProvider.value[providerId] = uniqueModels
|
||||
return uniqueModels
|
||||
} catch (reason) {
|
||||
const message = reason instanceof Error ? reason.message : '模型列表获取失败'
|
||||
modelErrorsByProvider.value[providerId] = message
|
||||
throw reason
|
||||
} finally {
|
||||
modelLoadingByProvider.value[providerId] = false
|
||||
}
|
||||
}
|
||||
|
||||
async function refreshEnabledModels() {
|
||||
await Promise.allSettled(
|
||||
providers.value.filter((provider) => provider.enabled).map((provider) => loadModels(provider.provider_id))
|
||||
)
|
||||
}
|
||||
|
||||
async function addProvider(data: Omit<ProviderConfig, 'provider_id'>) {
|
||||
@@ -41,6 +72,7 @@ export const useProviderStore = defineStore('provider', () => {
|
||||
const updated = await updateProviderRequest(providerId, data)
|
||||
const index = providers.value.findIndex((provider) => provider.provider_id === providerId)
|
||||
if (index >= 0) providers.value[index] = updated
|
||||
return updated
|
||||
}
|
||||
|
||||
async function deleteProvider(providerId: string) {
|
||||
@@ -61,14 +93,19 @@ export const useProviderStore = defineStore('provider', () => {
|
||||
|
||||
return {
|
||||
providers,
|
||||
presets,
|
||||
modelsByProvider,
|
||||
modelLoadingByProvider,
|
||||
modelErrorsByProvider,
|
||||
defaultProviderId,
|
||||
enabledProviders,
|
||||
defaultProvider,
|
||||
isLoading,
|
||||
error,
|
||||
loadProviders,
|
||||
loadPresets,
|
||||
loadModels,
|
||||
refreshEnabledModels,
|
||||
addProvider,
|
||||
updateProvider,
|
||||
deleteProvider,
|
||||
|
||||
Reference in New Issue
Block a user