diff --git a/docs/references/data/api-design-guidelines.md b/docs/references/data/api-design-guidelines.md index 12809a53981..0e64b5c16ff 100644 --- a/docs/references/data/api-design-guidelines.md +++ b/docs/references/data/api-design-guidelines.md @@ -188,6 +188,8 @@ Use verb-based paths for operations that don't fit CRUD semantics: > For sortable resources (drag-and-drop ordering), do not invent ad-hoc endpoints — follow the canonical `PATCH /{resource}/:id/order` pattern documented in the [Ordering Guide](./data-ordering-guide.md). +Provider enablement is a narrow exception: `PATCH /providers/:providerId` atomically moves that provider to the first position only when `isEnabled` transitions from `false` to `true`. This provider-specific invariant does not establish a general permission for resource updates to mutate ordering. Redundant provider `true` updates preserve the user's existing order, and explicit reorder requests still use the canonical order routes. + ```typescript // Search diff --git a/docs/references/data/data-ordering-guide.md b/docs/references/data/data-ordering-guide.md index 289fe129872..e3e02d3ea20 100644 --- a/docs/references/data/data-ordering-guide.md +++ b/docs/references/data/data-ordering-guide.md @@ -2,6 +2,8 @@ Canonical spec for any sortable resource in the DataApi system. Uses a single fractional-indexing design ([fractional-indexing](https://www.npmjs.com/package/fractional-indexing), Rocicorp, ~2 KB gzip) — `PATCH /{resource}/:id/order` with an anchor body. Scales from tens to thousands of rows without background rebalancing; applies uniformly whether the view is paginated or not. Replaces the two incompatible predecessors (`PATCH /mini-apps` absolute `sortOrder` integers and `PATCH /mcp-servers` full `orderedIds` list). +Provider enablement is transition-aware: `PATCH /providers/:providerId` moves a provider to the first position in the same transaction only when `isEnabled` changes from `false` to `true`. Redundant `true` updates preserve the user's existing order, while explicit reorder requests continue to use the canonical order routes. + Every sortable resource stores its position as a string `order_key` column. A reorder is always **relative** against an anchor (another row's id, or a `first` / `last` sentinel), never an absolute index. The server computes a new key between neighbours in one transaction; the renderer optimistically reorders its local cache and revalidates on completion. ## Quickstart — The Four Layers diff --git a/src/main/data/api/handlers/__tests__/providers.test.ts b/src/main/data/api/handlers/__tests__/providers.test.ts index 8ef2340d286..a3b5126ddb2 100644 --- a/src/main/data/api/handlers/__tests__/providers.test.ts +++ b/src/main/data/api/handlers/__tests__/providers.test.ts @@ -100,15 +100,15 @@ describe('providerHandlers', () => { describe('/providers/:providerId', () => { it('delegates PATCH to providerService.update with parsed body', async () => { - const updated = { id: 'openai', isEnabled: false } + const updated = { id: 'openai', isEnabled: true } updateMock.mockReturnValueOnce(updated) const result = await providerHandlers['/providers/:providerId'].PATCH({ params: { providerId: 'openai' }, - body: { isEnabled: false } + body: { isEnabled: true } } as never) - expect(updateMock).toHaveBeenCalledWith('openai', { isEnabled: false }) + expect(updateMock).toHaveBeenCalledWith('openai', { isEnabled: true }) expect(result).toBe(updated) }) diff --git a/src/main/data/services/ProviderService.ts b/src/main/data/services/ProviderService.ts index c76538114b8..25a4bb1a18e 100644 --- a/src/main/data/services/ProviderService.ts +++ b/src/main/data/services/ProviderService.ts @@ -230,7 +230,9 @@ class ProviderService { } /** - * Update an existing provider + * Update an existing provider. A false-to-true enabled transition moves the + * provider to the first position in the same transaction; redundant enabled + * writes preserve the user's current order. */ update(providerId: string, dto: UpdateProviderDto): Provider { assertManagedCherryAiProviderPatchAllowed(providerId, dto) @@ -246,7 +248,10 @@ class ProviderService { // defaults — otherwise DEFAULT_PROVIDER_SETTINGS would be persisted // into the row and break the "row stores only overrides" contract. const [current] = tx - .select({ providerSettings: userProviderTable.providerSettings }) + .select({ + providerSettings: userProviderTable.providerSettings, + isEnabled: userProviderTable.isEnabled + }) .from(userProviderTable) .where(eq(userProviderTable.providerId, providerId)) .limit(1) @@ -269,6 +274,16 @@ class ProviderService { ...dto.providerSettings } } + + if (dto.isEnabled === true && !current.isEnabled) { + try { + applyMoves(tx, userProviderTable, [{ id: providerId, anchor: { position: 'first' } }], { + pkColumn: userProviderTable.providerId + }) + } catch (error) { + this.rethrowOrderError(error) + } + } if (dto.isEnabled !== undefined) updates.isEnabled = dto.isEnabled const [updated] = tx diff --git a/src/main/data/services/__tests__/ProviderService.reorder.test.ts b/src/main/data/services/__tests__/ProviderService.reorder.test.ts index ed3783d4cf4..a92c53789c3 100644 --- a/src/main/data/services/__tests__/ProviderService.reorder.test.ts +++ b/src/main/data/services/__tests__/ProviderService.reorder.test.ts @@ -1,14 +1,15 @@ // Load the sibling so it self-registers in the data-service registry (prod loads it via its DataApi handler). import '@data/services/ProviderRegistryService' +import { application } from '@application' import { userProviderTable } from '@data/db/schemas/userProvider' import { providerService } from '@data/services/ProviderService' import { generateOrderKeySequence } from '@data/services/utils/orderKey' import { ErrorCode } from '@shared/data/api/errors' import { CHERRYAI_PROVIDER_ID } from '@shared/data/presets/cherryai' import { setupTestDatabase } from '@test-helpers/db' -import { asc, eq } from 'drizzle-orm' -import { describe, expect, it } from 'vitest' +import { asc, eq, sql } from 'drizzle-orm' +import { describe, expect, it, type Mock } from 'vitest' describe('ProviderService reorder', () => { const dbh = setupTestDatabase() @@ -61,6 +62,69 @@ describe('ProviderService reorder', () => { expect(await readOrder()).toEqual(['gemini', 'openai', 'anthropic']) }) + it('moves a provider to the first position when update enables it', async () => { + await seedProviders() + + const updated = providerService.update('gemini', { isEnabled: true, name: 'Gemini OAuth' }) + + expect(updated.isEnabled).toBe(true) + expect(updated.name).toBe('Gemini OAuth') + expect(await readOrder()).toEqual(['gemini', 'openai', 'anthropic']) + + const [row] = await dbh.db.select().from(userProviderTable).where(eq(userProviderTable.providerId, 'gemini')) + expect(row.isEnabled).toBe(true) + expect(row.name).toBe('Gemini OAuth') + }) + + it('rolls back the order move when the enable update fails', async () => { + await seedProviders() + + const withWriteTx = application.get('DbService').withWriteTx as Mock + withWriteTx.mockImplementationOnce((fn: (tx: unknown) => unknown) => dbh.db.transaction(fn as never)) + dbh.db.run( + sql.raw(` + CREATE TRIGGER fail_provider_enable_update + BEFORE UPDATE OF is_enabled ON user_provider + WHEN NEW.provider_id = 'gemini' + BEGIN + SELECT RAISE(ABORT, 'forced provider enable update failure'); + END; + `) + ) + + try { + expect(() => providerService.update('gemini', { isEnabled: true })).toThrow( + 'forced provider enable update failure' + ) + } finally { + dbh.db.run(sql.raw('DROP TRIGGER IF EXISTS fail_provider_enable_update')) + } + + const [row] = await dbh.db.select().from(userProviderTable).where(eq(userProviderTable.providerId, 'gemini')) + expect(row.isEnabled).toBe(false) + expect(await readOrder()).toEqual(['openai', 'anthropic', 'gemini']) + }) + + it('preserves user order when update receives a redundant enabled state', async () => { + await seedProviders() + await dbh.db.update(userProviderTable).set({ isEnabled: true }).where(eq(userProviderTable.providerId, 'gemini')) + + const updated = providerService.update('gemini', { isEnabled: true }) + + expect(updated.isEnabled).toBe(true) + expect(await readOrder()).toEqual(['openai', 'anthropic', 'gemini']) + }) + + it('preserves user order when update disables a provider', async () => { + await seedProviders() + await dbh.db.update(userProviderTable).set({ isEnabled: true }).where(eq(userProviderTable.providerId, 'gemini')) + + const updated = providerService.update('gemini', { isEnabled: false }) + + expect(updated.isEnabled).toBe(false) + expect(await readOrder()).toEqual(['openai', 'anthropic', 'gemini']) + }) + it('moves a provider after an anchor', async () => { await seedProviders() diff --git a/src/main/services/oauth/runtime/__tests__/OAuthRuntimeService.test.ts b/src/main/services/oauth/runtime/__tests__/OAuthRuntimeService.test.ts index 51a129d12f6..70fd574125d 100644 --- a/src/main/services/oauth/runtime/__tests__/OAuthRuntimeService.test.ts +++ b/src/main/services/oauth/runtime/__tests__/OAuthRuntimeService.test.ts @@ -349,6 +349,7 @@ describe('OAuthRuntimeService', () => { const stored = h.providerStore.get('codex') expect(stored?.authConfig).toMatchObject({ accessToken: 'at' }) expect(stored?.isEnabled).toBe(true) + expect(h.providerServiceMock.update).toHaveBeenCalledWith('codex', { isEnabled: true }) expect(account).toEqual({ accountId: null }) expect(h.transportMock.close).toHaveBeenCalled() }) @@ -373,6 +374,7 @@ describe('OAuthRuntimeService', () => { await service.handleDeepLinkCallback(new URL('app://cb?state=st&code=c')) expect(h.providerStore.get('cherryin')?.authConfig).toMatchObject({ accessToken: 'at' }) + expect(h.providerServiceMock.update).toHaveBeenCalledWith('cherryin', { isEnabled: true }) expect(h.deepLinkTransportMock.sendConsumedResult).toHaveBeenCalledWith('st', 'win-1', { apiKeys: '' }) }) diff --git a/src/renderer/hooks/__tests__/useProvider.test.ts b/src/renderer/hooks/__tests__/useProvider.test.ts index e85461042d4..ad32495a4da 100644 --- a/src/renderer/hooks/__tests__/useProvider.test.ts +++ b/src/renderer/hooks/__tests__/useProvider.test.ts @@ -257,6 +257,7 @@ describe('useProvider', () => { const { result } = renderHook(() => useProvider('openai')) expect(result.current.updateProvider).toBeDefined() + expect(result.current.enableProvider).toBeDefined() expect(result.current.deleteProvider).toBeDefined() expect(result.current.updateAuthConfig).toBeDefined() expect(result.current.updateApiKeys).toBeDefined() @@ -371,6 +372,24 @@ describe('useProviderMutations', () => { expect(mockTrigger).toHaveBeenCalledWith({ params: { providerId: 'openai' }, body: { isEnabled: false } }) }) + it('should enable a provider through the generic PATCH mutation', async () => { + const patchTrigger = vi.fn().mockResolvedValue({}) + mockUseMutation.mockImplementation((_method: string, path: string) => ({ + trigger: + _method === 'PATCH' && path === '/providers/:providerId' ? patchTrigger : vi.fn().mockResolvedValue(undefined), + isLoading: false, + error: undefined + })) + + const { result } = renderHook(() => useProviderMutations('openai')) + + await act(async () => { + await result.current.enableProvider() + }) + + expect(patchTrigger).toHaveBeenCalledWith({ params: { providerId: 'openai' }, body: { isEnabled: true } }) + }) + it('should call deleteTrigger with providerId param when deleteProvider is invoked', async () => { const mockTrigger = vi.fn().mockResolvedValue(undefined) mockUseMutation.mockImplementation(() => ({ diff --git a/src/renderer/hooks/useProvider.ts b/src/renderer/hooks/useProvider.ts index 0877a12001f..3607beef2e2 100644 --- a/src/renderer/hooks/useProvider.ts +++ b/src/renderer/hooks/useProvider.ts @@ -155,6 +155,8 @@ export function useProviderMutations(providerId: string) { } }, [deleteTrigger, providerId]) + const enableProvider = useCallback(() => updateProvider({ isEnabled: true }), [updateProvider]) + const updateAuthConfig = useCallback( async (authConfig: AuthConfig) => { try { @@ -222,6 +224,7 @@ export function useProviderMutations(providerId: string) { deleteProvider, isDeleting, deleteError, + enableProvider, updateAuthConfig, addApiKey, isAddingApiKey, diff --git a/src/renderer/i18n/locales/en-us.json b/src/renderer/i18n/locales/en-us.json index 21cff8f4ea8..88da0e8cd92 100644 --- a/src/renderer/i18n/locales/en-us.json +++ b/src/renderer/i18n/locales/en-us.json @@ -6797,6 +6797,7 @@ "confirm": "Are you sure you want to add all models to the list?", "label": "Add all models" }, + "add_success_enable_failed": "Models were added, but the provider could not be enabled.", "add_whole_group": "Add the whole group", "clean_stale_models": "Clean stale models", "clean_stale_success": "Cleaned {{count}} stale model(s)", @@ -7208,6 +7209,7 @@ "fill_after_create": "Authentication fields can be filled after create", "menu_label": "Add instance" }, + "enable_failed_after_connection": "Connection succeeded, but the provider could not be enabled.", "filter": { "agent": "Agent Supported", "all": "All Providers", diff --git a/src/renderer/i18n/locales/zh-cn.json b/src/renderer/i18n/locales/zh-cn.json index 4e88fc8f69b..ac28a243a3c 100644 --- a/src/renderer/i18n/locales/zh-cn.json +++ b/src/renderer/i18n/locales/zh-cn.json @@ -6797,6 +6797,7 @@ "confirm": "确定要添加所有模型到列表吗?", "label": "添加全部模型" }, + "add_success_enable_failed": "模型已添加,但服务商启用失败。", "add_whole_group": "添加整个分组", "clean_stale_models": "清理失效模型", "clean_stale_success": "已清理 {{count}} 个失效模型", @@ -7208,6 +7209,7 @@ "fill_after_create": "创建后请在详情页填写认证字段", "menu_label": "添加实例" }, + "enable_failed_after_connection": "连接成功,但服务商启用失败。", "filter": { "agent": "支持 Agent", "all": "全部服务商", diff --git a/src/renderer/i18n/locales/zh-tw.json b/src/renderer/i18n/locales/zh-tw.json index 354cd755107..b942eefe7c9 100644 --- a/src/renderer/i18n/locales/zh-tw.json +++ b/src/renderer/i18n/locales/zh-tw.json @@ -6797,6 +6797,7 @@ "confirm": "確定要新增所有模型到列表嗎?", "label": "新增全部模型" }, + "add_success_enable_failed": "模型已新增,但服務商啟用失敗。", "add_whole_group": "新增整個分組", "clean_stale_models": "清理失效模型", "clean_stale_success": "已清理 {{count}} 個失效模型", @@ -7208,6 +7209,7 @@ "fill_after_create": "建立後請在詳情頁填寫認證欄位", "menu_label": "新增實例" }, + "enable_failed_after_connection": "連線成功,但服務商啟用失敗。", "filter": { "agent": "支援 Agent", "all": "全部服務商", diff --git a/src/renderer/pages/settings/ProviderSettings/ModelList/__tests__/useProviderModelPullReconcile.test.ts b/src/renderer/pages/settings/ProviderSettings/ModelList/__tests__/useProviderModelPullReconcile.test.ts index 66b19720d2f..f272ae66547 100644 --- a/src/renderer/pages/settings/ProviderSettings/ModelList/__tests__/useProviderModelPullReconcile.test.ts +++ b/src/renderer/pages/settings/ProviderSettings/ModelList/__tests__/useProviderModelPullReconcile.test.ts @@ -24,7 +24,7 @@ const toCreateModelDtoMock = vi.fn((providerId, model, endpointTypes) => ({ group: model.group, endpointTypes })) -const updateProviderMock = vi.fn() +const enableProviderMock = vi.fn() const useModelsMock = vi.fn() const useProviderMock = vi.fn() @@ -104,14 +104,14 @@ describe('useProviderModelPullReconcile', () => { createModelsMock.mockResolvedValue([]) deleteModelsMock.mockResolvedValue(undefined) reconcileTriggerMock.mockResolvedValue([]) - enableProviderWhenModelsAvailableMock.mockResolvedValue(false) + enableProviderWhenModelsAvailableMock.mockResolvedValue(undefined) fetchProviderCatalogModelsMock.mockResolvedValue([catalogModel]) fetchResolvedProviderModelsMock.mockResolvedValue([fetchedModel]) resolveCreateModelEndpointTypesMock.mockReturnValue(undefined) useModelsMock.mockReturnValue({ models: [localModel] }) useProviderMock.mockReturnValue({ provider: { id: 'openai', isEnabled: false }, - updateProvider: updateProviderMock + enableProvider: enableProviderMock }) }) @@ -265,7 +265,7 @@ describe('useProviderModelPullReconcile', () => { expect(resolveCreateModelEndpointTypesMock).toHaveBeenCalledWith({ id: 'openai', isEnabled: false }, fetchedModel) expect(enableProviderWhenModelsAvailableMock).toHaveBeenCalledWith( { id: 'openai', isEnabled: false }, - updateProviderMock, + enableProviderMock, 2, 'model_manage_add' ) @@ -282,6 +282,19 @@ describe('useProviderModelPullReconcile', () => { expect(toast.error).toHaveBeenCalledWith('settings.models.manage.operation_failed') }) + it('warns that models were added when provider enablement fails', async () => { + enableProviderWhenModelsAvailableMock.mockRejectedValueOnce(new Error('enable failed')) + const { result } = renderHook(() => useProviderModelPullReconcile('openai')) + + await act(async () => { + await result.current.addModels([fetchedModel as any]) + }) + + expect(createModelsMock).toHaveBeenCalledTimes(1) + expect(toast.warning).toHaveBeenCalledWith('settings.models.manage.add_success_enable_failed') + expect(toast.error).not.toHaveBeenCalledWith('settings.models.manage.operation_failed') + }) + it('removes unique local model ids', async () => { const { result } = renderHook(() => useProviderModelPullReconcile('openai')) diff --git a/src/renderer/pages/settings/ProviderSettings/ModelList/useProviderModelPullReconcile.ts b/src/renderer/pages/settings/ProviderSettings/ModelList/useProviderModelPullReconcile.ts index 543a7b4b1bb..dbdfc30fbcb 100644 --- a/src/renderer/pages/settings/ProviderSettings/ModelList/useProviderModelPullReconcile.ts +++ b/src/renderer/pages/settings/ProviderSettings/ModelList/useProviderModelPullReconcile.ts @@ -72,7 +72,7 @@ export function useProviderModelPullReconcile(providerId: string) { const [defaultModelId] = usePreference('chat.default_model_id') const [quickAssistantModelId] = usePreference('feature.quick_assistant.model_id') const [translateModelId] = usePreference('feature.translate.model_id') - const { provider, updateProvider } = useProvider(providerId) + const { provider, enableProvider } = useProvider(providerId) const { models } = useModels({ providerId }) const { createModels, deleteModels, isCreating, isDeleting, isBulkDeleting } = useModelMutations() const { trigger: reconcileModels, isLoading: isReconciling } = useMutation( @@ -179,18 +179,29 @@ export function useProviderModelPullReconcile(providerId: string) { await createModels( toAdd.map((model) => toCreateModelDto(providerId, model, resolveCreateModelEndpointTypes(provider, model))) ) + } catch (error) { + logger.error('Failed to add provider models from manage drawer', { providerId, count: toAdd.length, error }) + toast.error(t('settings.models.manage.operation_failed')) + return + } + + try { await enableProviderWhenModelsAvailable( provider, - updateProvider, + enableProvider, models.length + toAdd.length, 'model_manage_add' ) } catch (error) { - logger.error('Failed to add provider models from manage drawer', { providerId, count: toAdd.length, error }) - toast.error(t('settings.models.manage.operation_failed')) + logger.error('Models were added but provider enablement failed', { + providerId, + count: toAdd.length, + error + }) + toast.warning(t('settings.models.manage.add_success_enable_failed')) } }, - [createModels, models, provider, providerId, t, updateProvider] + [createModels, enableProvider, models, provider, providerId, t] ) const removeModels = useCallback( diff --git a/src/renderer/pages/settings/ProviderSettings/hooks/providerSetting/__tests__/useProviderConnectionCheck.test.tsx b/src/renderer/pages/settings/ProviderSettings/hooks/providerSetting/__tests__/useProviderConnectionCheck.test.tsx index bd2facaf03e..155d03e1e7d 100644 --- a/src/renderer/pages/settings/ProviderSettings/hooks/providerSetting/__tests__/useProviderConnectionCheck.test.tsx +++ b/src/renderer/pages/settings/ProviderSettings/hooks/providerSetting/__tests__/useProviderConnectionCheck.test.tsx @@ -1,3 +1,4 @@ +import { HealthStatus } from '@renderer/pages/settings/ProviderSettings/types/healthCheck' import { toast } from '@renderer/services/toast' import { ENDPOINT_TYPE, MODEL_CAPABILITY } from '@shared/data/types/model' import { act, renderHook } from '@testing-library/react' @@ -11,11 +12,12 @@ const useTimerMock = vi.fn() const useAuthenticationApiKeyMock = vi.fn() const useProviderEndpointsMock = vi.fn() const checkApiMock = vi.fn() -const updateProviderMock = vi.fn() +const enableProviderMock = vi.fn() const commitInputApiKeyNowMock = vi.fn() const { loggerErrorMock } = vi.hoisted(() => ({ loggerErrorMock: vi.fn() })) +let inputApiKey = 'sk-a,sk-b' vi.mock('react-i18next', async (importOriginal) => { const actual = await importOriginal() @@ -66,10 +68,11 @@ describe('useProviderConnectionCheck', () => { beforeEach(() => { vi.clearAllMocks() + inputApiKey = 'sk-a,sk-b' useProviderMock.mockReturnValue({ provider: { id: 'cherryin', name: 'CherryIN', isEnabled: false }, - updateProvider: updateProviderMock + enableProvider: enableProviderMock }) useModelsMock.mockReturnValue({ models: [ @@ -91,10 +94,10 @@ describe('useProviderConnectionCheck', () => { }) useTimerMock.mockReturnValue({ setTimeoutTimer }) commitInputApiKeyNowMock.mockResolvedValue(undefined) - useAuthenticationApiKeyMock.mockReturnValue({ - inputApiKey: 'sk-a,sk-b', + useAuthenticationApiKeyMock.mockImplementation(() => ({ + inputApiKey, commitInputApiKeyNow: commitInputApiKeyNowMock - }) + })) useProviderEndpointsMock.mockReturnValue({ apiHost: 'https://open.cherryin.net', anthropicApiHost: 'https://open.cherryin.net' @@ -119,7 +122,7 @@ describe('useProviderConnectionCheck', () => { it('opens the connection drawer without API keys for no-key providers', () => { useProviderMock.mockReturnValue({ provider: { id: 'ollama', name: 'Ollama', isEnabled: false }, - updateProvider: updateProviderMock + enableProvider: enableProviderMock }) useAuthenticationApiKeyMock.mockReturnValue({ inputApiKey: '', @@ -139,7 +142,7 @@ describe('useProviderConnectionCheck', () => { it('opens the connection drawer without API keys for providers derived from no-key presets', () => { useProviderMock.mockReturnValue({ provider: { id: 'custom-ollama', presetProviderId: 'ollama', name: 'Custom Ollama', isEnabled: false }, - updateProvider: updateProviderMock + enableProvider: enableProviderMock }) useAuthenticationApiKeyMock.mockReturnValue({ inputApiKey: '', @@ -181,7 +184,7 @@ describe('useProviderConnectionCheck', () => { it('runs no-key provider checks without an API key override', async () => { useProviderMock.mockReturnValue({ provider: { id: 'ollama', name: 'Ollama', isEnabled: false }, - updateProvider: updateProviderMock + enableProvider: enableProviderMock }) useAuthenticationApiKeyMock.mockReturnValue({ inputApiKey: '', @@ -213,7 +216,7 @@ describe('useProviderConnectionCheck', () => { }) }) - expect(updateProviderMock).toHaveBeenCalledWith({ isEnabled: true }) + expect(enableProviderMock).toHaveBeenCalledTimes(1) }) it('persists the pending API key before running the check and before enabling the provider', async () => { @@ -227,11 +230,11 @@ describe('useProviderConnectionCheck', () => { }) expect(commitInputApiKeyNowMock).toHaveBeenCalledTimes(1) - expect(updateProviderMock).toHaveBeenCalledWith({ isEnabled: true }) + expect(enableProviderMock).toHaveBeenCalledTimes(1) // commit must run before the check so a freshly typed key is saved before // provider enablement, while the check still uses the selected key override. expect(commitInputApiKeyNowMock.mock.invocationCallOrder[0]).toBeLessThan(checkApiMock.mock.invocationCallOrder[0]) - expect(checkApiMock.mock.invocationCallOrder[0]).toBeLessThan(updateProviderMock.mock.invocationCallOrder[0]) + expect(checkApiMock.mock.invocationCallOrder[0]).toBeLessThan(enableProviderMock.mock.invocationCallOrder[0]) }) it('does not run the check or enable the provider when saving the pending API key fails', async () => { @@ -251,7 +254,7 @@ describe('useProviderConnectionCheck', () => { // enabling, surfacing only the failure path — never success-then-failure. // The toast must name the save failure, not the connection: nothing was probed. expect(checkApiMock).not.toHaveBeenCalled() - expect(updateProviderMock).not.toHaveBeenCalled() + expect(enableProviderMock).not.toHaveBeenCalled() expect(loggerErrorMock).toHaveBeenCalledWith('Failed to persist pending API key before connection check', { providerId: 'cherryin', modelId: 'cherryin::claude-4-sonnet', @@ -266,8 +269,23 @@ describe('useProviderConnectionCheck', () => { it('does not patch an already enabled provider after a successful model connection check', async () => { useProviderMock.mockReturnValue({ provider: { id: 'cherryin', name: 'CherryIN', isEnabled: true }, - updateProvider: updateProviderMock + enableProvider: enableProviderMock + }) + const { result } = renderHook(() => useProviderConnectionCheck('cherryin')) + + await act(async () => { + await result.current.startConnectionCheck({ + model: result.current.checkableModels[0], + apiKey: 'sk-a' + }) }) + + expect(enableProviderMock).not.toHaveBeenCalled() + }) + + it('preserves connection success and warns when provider enablement fails', async () => { + const enableError = new Error('enable and pin failed') + enableProviderMock.mockRejectedValueOnce(enableError) const { result } = renderHook(() => useProviderConnectionCheck('cherryin')) await act(async () => { @@ -277,7 +295,44 @@ describe('useProviderConnectionCheck', () => { }) }) - expect(updateProviderMock).not.toHaveBeenCalled() + expect(loggerErrorMock).toHaveBeenCalledWith('Provider connection succeeded but enablement failed', { + providerId: 'cherryin', + modelId: 'cherryin::claude-4-sonnet', + error: enableError + }) + expect(toast.warning).toHaveBeenCalledWith('settings.provider.enable_failed_after_connection') + expect(toast.success).toHaveBeenCalled() + expect(result.current.apiKeyConnectivity.status).toBe(HealthStatus.SUCCESS) + }) + + it('ignores an enablement failure from a superseded connection check', async () => { + let rejectEnable: ((error: Error) => void) | undefined + enableProviderMock.mockImplementationOnce( + () => + new Promise((_, reject) => { + rejectEnable = reject + }) + ) + const { result, rerender } = renderHook(() => useProviderConnectionCheck('cherryin')) + + act(() => { + void result.current.startConnectionCheck({ + model: result.current.checkableModels[0], + apiKey: 'sk-a' + }) + }) + await vi.waitFor(() => expect(enableProviderMock).toHaveBeenCalledTimes(1)) + + inputApiKey = 'sk-new' + rerender() + + await act(async () => { + rejectEnable?.(new Error('stale enable failure')) + await Promise.resolve() + }) + + expect(toast.warning).not.toHaveBeenCalled() + expect(toast.success).not.toHaveBeenCalled() }) it('logs provider/model context when the connection check fails', async () => { diff --git a/src/renderer/pages/settings/ProviderSettings/hooks/providerSetting/__tests__/useProviderEnable.test.tsx b/src/renderer/pages/settings/ProviderSettings/hooks/providerSetting/__tests__/useProviderEnable.test.tsx index c231f60196d..936350afbe1 100644 --- a/src/renderer/pages/settings/ProviderSettings/hooks/providerSetting/__tests__/useProviderEnable.test.tsx +++ b/src/renderer/pages/settings/ProviderSettings/hooks/providerSetting/__tests__/useProviderEnable.test.tsx @@ -5,19 +5,14 @@ import { useProviderEnable } from '../useProviderEnable' const useProviderMock = vi.fn() const useProviderMutationsMock = vi.fn() -const useReorderMock = vi.fn() const updateProviderMock = vi.fn().mockResolvedValue(undefined) -const moveMock = vi.fn().mockResolvedValue(undefined) +const enableProviderMock = vi.fn().mockResolvedValue(undefined) vi.mock('@renderer/hooks/useProvider', () => ({ useProvider: (...args: any[]) => useProviderMock(...args), useProviderMutations: (...args: any[]) => useProviderMutationsMock(...args) })) -vi.mock('@data/hooks/useReorder', () => ({ - useReorder: (...args: any[]) => useReorderMock(...args) -})) - describe('useProviderEnable', () => { beforeEach(() => { vi.clearAllMocks() @@ -25,10 +20,8 @@ describe('useProviderEnable', () => { provider: { id: 'openai', isEnabled: true } }) useProviderMutationsMock.mockReturnValue({ - updateProvider: updateProviderMock - }) - useReorderMock.mockReturnValue({ - move: moveMock + updateProvider: updateProviderMock, + enableProvider: enableProviderMock }) }) @@ -40,18 +33,18 @@ describe('useProviderEnable', () => { }) expect(updateProviderMock).toHaveBeenCalledWith({ isEnabled: false }) - expect(moveMock).not.toHaveBeenCalled() + expect(enableProviderMock).not.toHaveBeenCalled() }) - it('moves the provider to the top after enabling it', async () => { + it('enables and moves the provider to the top through the atomic mutation', async () => { const { result } = renderHook(() => useProviderEnable('openai')) await act(async () => { await result.current.toggleProviderEnabled(true) }) - expect(updateProviderMock).toHaveBeenCalledWith({ isEnabled: true }) - expect(moveMock).toHaveBeenCalledWith('openai', { position: 'first' }) + expect(enableProviderMock).toHaveBeenCalledTimes(1) + expect(updateProviderMock).not.toHaveBeenCalled() }) it('does nothing when the provider is missing', async () => { @@ -66,15 +59,15 @@ describe('useProviderEnable', () => { }) expect(updateProviderMock).not.toHaveBeenCalled() - expect(moveMock).not.toHaveBeenCalled() + expect(enableProviderMock).not.toHaveBeenCalled() }) - it('rolls the enable state back when pin-to-top fails after enabling', async () => { + it('surfaces atomic enable-and-pin failures without stale rollback', async () => { useProviderMock.mockReturnValue({ provider: { id: 'openai', isEnabled: false } }) - const moveError = new Error('move failed') - moveMock.mockRejectedValueOnce(moveError) + const enableError = new Error('enable and pin failed') + enableProviderMock.mockRejectedValueOnce(enableError) const { result } = renderHook(() => useProviderEnable('openai')) @@ -87,10 +80,8 @@ describe('useProviderEnable', () => { } }) - expect(thrown).toBe(moveError) - expect(updateProviderMock).toHaveBeenCalledTimes(2) - expect(updateProviderMock).toHaveBeenNthCalledWith(1, { isEnabled: true }) - expect(moveMock).toHaveBeenCalledWith('openai', { position: 'first' }) - expect(updateProviderMock).toHaveBeenNthCalledWith(2, { isEnabled: false }) + expect(thrown).toBe(enableError) + expect(enableProviderMock).toHaveBeenCalledTimes(1) + expect(updateProviderMock).not.toHaveBeenCalled() }) }) diff --git a/src/renderer/pages/settings/ProviderSettings/hooks/providerSetting/useProviderConnectionCheck.ts b/src/renderer/pages/settings/ProviderSettings/hooks/providerSetting/useProviderConnectionCheck.ts index 9c4dbafb9cb..221747f23c4 100644 --- a/src/renderer/pages/settings/ProviderSettings/hooks/providerSetting/useProviderConnectionCheck.ts +++ b/src/renderer/pages/settings/ProviderSettings/hooks/providerSetting/useProviderConnectionCheck.ts @@ -22,7 +22,7 @@ import { useProviderEndpoints } from './useProviderEndpoints' const logger = loggerService.withContext('ProviderSettings:ConnectionCheck') export function useProviderConnectionCheck(providerId: string) { - const { provider, updateProvider } = useProvider(providerId) + const { provider, enableProvider } = useProvider(providerId) const [connectionCheckOpen, setConnectionCheckOpen] = useState(false) const { models } = useModels( { providerId }, @@ -113,9 +113,20 @@ export function useProviderConnectionCheck(providerId: string) { if (runId !== runIdRef.current) return - // Enable the provider (if disabled) only after a successful check. Enable - // swallows its own errors, so it never diverts to the failure path. - await enableProviderWhenModelsAvailable(provider, updateProvider, checkableModels.length, 'connection_check') + // Connectivity has already succeeded. Provider enablement is a follow-up + // action, so report its failure separately without marking the probe failed. + try { + await enableProviderWhenModelsAvailable(provider, enableProvider, checkableModels.length, 'connection_check') + } catch (error) { + if (runId !== runIdRef.current || controller.signal.aborted) return + + logger.error('Provider connection succeeded but enablement failed', { + providerId: provider.id, + modelId: model.id, + error + }) + toast.warning(i18n.t('settings.provider.enable_failed_after_connection')) + } // The enable await can interleave with a newer check; drop this run if it // was superseded or aborted before touching success state. @@ -169,7 +180,7 @@ export function useProviderConnectionCheck(providerId: string) { provider, requiresApiKey, setTimeoutTimer, - updateProvider + enableProvider ] ) diff --git a/src/renderer/pages/settings/ProviderSettings/hooks/providerSetting/useProviderEnable.ts b/src/renderer/pages/settings/ProviderSettings/hooks/providerSetting/useProviderEnable.ts index e8f699653c1..d1da52038e9 100644 --- a/src/renderer/pages/settings/ProviderSettings/hooks/providerSetting/useProviderEnable.ts +++ b/src/renderer/pages/settings/ProviderSettings/hooks/providerSetting/useProviderEnable.ts @@ -1,12 +1,10 @@ -import { useReorder } from '@data/hooks/useReorder' import { useProvider, useProviderMutations } from '@renderer/hooks/useProvider' import { useCallback } from 'react' /** Persists provider enable changes and moves newly enabled providers to the top. */ export function useProviderEnable(providerId: string) { const { provider } = useProvider(providerId) - const { updateProvider } = useProviderMutations(providerId) - const { move } = useReorder('/providers') + const { updateProvider, enableProvider } = useProviderMutations(providerId) const toggleProviderEnabled = useCallback( async (enabled: boolean) => { @@ -14,25 +12,14 @@ export function useProviderEnable(providerId: string) { return } - const previousEnabled = provider.isEnabled - await updateProvider({ isEnabled: enabled }) - - if (!enabled) { + if (enabled) { + await enableProvider() return } - try { - await move(providerId, { position: 'first' }) - } catch (error) { - // Enable + pin-to-top is one user-facing action. If pinning fails after the - // enable already committed, roll the enable state back so we don't leave a - // half-success ("enabled but not pinned") with no path back. Best-effort — - // if the rollback also fails the original error still surfaces. - await updateProvider({ isEnabled: previousEnabled }).catch(() => undefined) - throw error - } + await updateProvider({ isEnabled: false }) }, - [move, provider, providerId, updateProvider] + [enableProvider, provider, updateProvider] ) return { diff --git a/src/renderer/pages/settings/ProviderSettings/utils/__tests__/providerEnablement.test.ts b/src/renderer/pages/settings/ProviderSettings/utils/__tests__/providerEnablement.test.ts index 7faf112f3d7..6586efb3b08 100644 --- a/src/renderer/pages/settings/ProviderSettings/utils/__tests__/providerEnablement.test.ts +++ b/src/renderer/pages/settings/ProviderSettings/utils/__tests__/providerEnablement.test.ts @@ -14,53 +14,48 @@ describe('enableProviderWhenModelsAvailable', () => { loggerErrorSpy = vi.spyOn(mockRendererLoggerService, 'error').mockImplementation(() => {}) }) - it('enables a disabled provider when at least one model is available', async () => { - const updateProvider = vi.fn().mockResolvedValue(undefined) + it('enables a disabled provider with pin-to-top when at least one model is available', async () => { + const enableProvider = vi.fn().mockResolvedValue(undefined) - const enabled = await enableProviderWhenModelsAvailable(disabledProvider, updateProvider, 2, 'test') + await enableProviderWhenModelsAvailable(disabledProvider, enableProvider, 2, 'test') - expect(enabled).toBe(true) - expect(updateProvider).toHaveBeenCalledWith({ isEnabled: true }) + expect(enableProvider).toHaveBeenCalledTimes(1) }) - it('no-ops when the provider is already enabled', async () => { - const updateProvider = vi.fn().mockResolvedValue(undefined) + it('skips when the provider is already enabled', async () => { + const enableProvider = vi.fn().mockResolvedValue(undefined) - const enabled = await enableProviderWhenModelsAvailable(enabledProvider, updateProvider, 2, 'test') + await enableProviderWhenModelsAvailable(enabledProvider, enableProvider, 2, 'test') - expect(enabled).toBe(false) - expect(updateProvider).not.toHaveBeenCalled() + expect(enableProvider).not.toHaveBeenCalled() }) - it('no-ops when no models are available', async () => { - const updateProvider = vi.fn().mockResolvedValue(undefined) + it('skips when no models are available', async () => { + const enableProvider = vi.fn().mockResolvedValue(undefined) - const enabled = await enableProviderWhenModelsAvailable(disabledProvider, updateProvider, 0, 'test') + await enableProviderWhenModelsAvailable(disabledProvider, enableProvider, 0, 'test') - expect(enabled).toBe(false) - expect(updateProvider).not.toHaveBeenCalled() + expect(enableProvider).not.toHaveBeenCalled() }) - it('no-ops when the provider has not resolved yet', async () => { - const updateProvider = vi.fn().mockResolvedValue(undefined) + it('skips when the provider has not resolved yet', async () => { + const enableProvider = vi.fn().mockResolvedValue(undefined) - const enabled = await enableProviderWhenModelsAvailable(undefined, updateProvider, 2, 'test') + await enableProviderWhenModelsAvailable(undefined, enableProvider, 2, 'test') - expect(enabled).toBe(false) - expect(updateProvider).not.toHaveBeenCalled() + expect(enableProvider).not.toHaveBeenCalled() }) - it('returns false and logs without throwing when the update fails', async () => { - const updateError = new Error('patch failed') - const updateProvider = vi.fn().mockRejectedValue(updateError) + it('throws and logs when the atomic enable-and-pin action rejects', async () => { + const enableError = new Error('enable and pin failed') + const enableProvider = vi.fn().mockRejectedValue(enableError) - const enabled = await enableProviderWhenModelsAvailable(disabledProvider, updateProvider, 2, 'test') - - expect(enabled).toBe(false) - expect(updateProvider).toHaveBeenCalledWith({ isEnabled: true }) + await expect(enableProviderWhenModelsAvailable(disabledProvider, enableProvider, 2, 'test')).rejects.toBe( + enableError + ) expect(loggerErrorSpy).toHaveBeenCalledWith( - 'Failed to enable provider when models are available', - expect.objectContaining({ providerId: 'cherryin', modelCount: 2, source: 'test', error: updateError }) + 'Failed to enable provider with pin-to-top when models are available', + expect.objectContaining({ providerId: 'cherryin', modelCount: 2, source: 'test', error: enableError }) ) }) }) diff --git a/src/renderer/pages/settings/ProviderSettings/utils/providerEnablement.ts b/src/renderer/pages/settings/ProviderSettings/utils/providerEnablement.ts index ae7ef40a6f3..005d7145403 100644 --- a/src/renderer/pages/settings/ProviderSettings/utils/providerEnablement.ts +++ b/src/renderer/pages/settings/ProviderSettings/utils/providerEnablement.ts @@ -1,30 +1,28 @@ import { loggerService } from '@logger' -import type { UpdateProviderDto } from '@shared/data/api/schemas/providers' import type { Provider } from '@shared/data/types/provider' const logger = loggerService.withContext('ProviderSettings:EnableProviderWhenModelsAvailable') -/** Enables a disabled provider once a flow has confirmed it has usable models. */ +/** Enables a disabled provider once a flow has confirmed it has usable models, then moves it to the top. */ export async function enableProviderWhenModelsAvailable( provider: Pick | undefined, - updateProvider: (updates: UpdateProviderDto) => Promise, + enableProvider: () => Promise, modelCount: number, source: string -): Promise { +): Promise { if (!provider || provider.isEnabled || modelCount <= 0) { - return false + return } try { - await updateProvider({ isEnabled: true }) - return true + await enableProvider() } catch (error) { - logger.error('Failed to enable provider when models are available', { + logger.error('Failed to enable provider with pin-to-top when models are available', { providerId: provider.id, modelCount, source, error }) - return false + throw error } } diff --git a/src/shared/data/api/schemas/providers.ts b/src/shared/data/api/schemas/providers.ts index 63b556dc17a..b8aeae096ad 100644 --- a/src/shared/data/api/schemas/providers.ts +++ b/src/shared/data/api/schemas/providers.ts @@ -90,7 +90,11 @@ const ProviderMutableFieldsSchema = CreateProviderSchema.pick({ }) export const UpdateProviderSchema = ProviderMutableFieldsSchema.partial().extend({ - /** Whether this provider is enabled */ + /** + * Whether this provider is enabled. A persisted false-to-true transition also + * moves the provider to the first position atomically; redundant true updates + * preserve the existing order. + */ isEnabled: z.boolean().optional() }) export type UpdateProviderDto = z.infer