mirror of
https://github.com/NanmiCoder/claude-code-haha.git
synced 2026-10-10 11:53:10 +08:00
fix(desktop): preserve provider 1m flags in runtime selection
This commit is contained in:
@@ -63,6 +63,29 @@ beforeEach(() => {
|
||||
})
|
||||
|
||||
describe('ModelSelector', () => {
|
||||
it.each([true, false])('sends each provider slot with 1M=%s and preserves reasoning controls', async (enabled) => {
|
||||
useSettingsStore.setState({ locale: 'en', effortLevel: 'high' })
|
||||
useProviderStore.setState({
|
||||
activeId: 'provider-1m', hasLoadedProviders: true, isLoading: false,
|
||||
providers: [{
|
||||
id: 'provider-1m', presetId: 'custom', name: 'Provider 1M',
|
||||
apiFormat: 'anthropic', apiKey: 'fixture', baseUrl: 'http://127.0.0.1:9999',
|
||||
models: { main: 'main-model', haiku: 'haiku-model', sonnet: 'sonnet-model', opus: 'opus-model' },
|
||||
model1mSupport: { main: enabled, haiku: enabled, sonnet: enabled, opus: enabled },
|
||||
}],
|
||||
})
|
||||
const runtimeChange = vi.fn()
|
||||
render(<ModelSelector runtimeKey="__draft__" onRuntimeSelectionChange={runtimeChange} />)
|
||||
for (const slot of ['main', 'haiku', 'sonnet', 'opus']) {
|
||||
await clickByRole(/, Provider 1M$/)
|
||||
fireEvent.click(within(screen.getByTestId('model-selector-dropdown')).getByRole('button', { name: new RegExp(`^${slot}-model`) }))
|
||||
expect(runtimeChange).toHaveBeenLastCalledWith({
|
||||
providerId: 'provider-1m', modelId: `${slot}-model${enabled ? '[1m]' : ''}`, effortLevel: 'high',
|
||||
})
|
||||
expect(screen.getByRole('button', { name: /High/ })).toBeInTheDocument()
|
||||
}
|
||||
})
|
||||
|
||||
it.each(['unknown', 'mixed', 'anthropic'] as const)(
|
||||
'allows cross-protocol selection despite retained %s session metadata', async (sessionApiFormat) => {
|
||||
const sessionId = 'protocol-rollback-session'
|
||||
|
||||
@@ -23,6 +23,8 @@ import { isDesktopRuntime } from '../../lib/desktopRuntime'
|
||||
import {
|
||||
normalizeRuntimeSelection,
|
||||
resolveDefaultRuntimeSelection,
|
||||
resolveProviderRuntimeModelId,
|
||||
resolveProviderSlotModelId,
|
||||
} from '../../lib/runtimeSelection'
|
||||
import { useHahaOAuthStore } from '../../stores/hahaOAuthStore'
|
||||
import { useHahaOpenAIOAuthStore } from '../../stores/hahaOpenAIOAuthStore'
|
||||
@@ -97,7 +99,12 @@ function getProviderModelCapabilityOverride(
|
||||
): string | undefined {
|
||||
return getModelReasoningCapabilityOverride(
|
||||
modelId,
|
||||
provider.models,
|
||||
{
|
||||
...provider.models,
|
||||
haiku: resolveProviderSlotModelId(provider, 'haiku'),
|
||||
sonnet: resolveProviderSlotModelId(provider, 'sonnet'),
|
||||
opus: resolveProviderSlotModelId(provider, 'opus'),
|
||||
},
|
||||
PROVIDER_PRESET_DEFAULT_ENVS.get(provider.presetId) ?? {},
|
||||
)
|
||||
}
|
||||
@@ -133,10 +140,10 @@ function buildProviderModels(
|
||||
labels: Record<'main' | 'haiku' | 'sonnet' | 'opus', string>,
|
||||
): ModelInfo[] {
|
||||
const entries: Array<{ id: string; label: string }> = [
|
||||
{ id: provider.models.main.trim(), label: labels.main },
|
||||
{ id: provider.models.haiku.trim(), label: labels.haiku },
|
||||
{ id: provider.models.sonnet.trim(), label: labels.sonnet },
|
||||
{ id: provider.models.opus.trim(), label: labels.opus },
|
||||
{ id: resolveProviderSlotModelId(provider, 'main'), label: labels.main },
|
||||
{ id: resolveProviderSlotModelId(provider, 'haiku'), label: labels.haiku },
|
||||
{ id: resolveProviderSlotModelId(provider, 'sonnet'), label: labels.sonnet },
|
||||
{ id: resolveProviderSlotModelId(provider, 'opus'), label: labels.opus },
|
||||
]
|
||||
|
||||
const byId = new Map<string, { id: string; labels: string[] }>()
|
||||
@@ -442,10 +449,21 @@ export const ModelSelector = forwardRef<ModelSelectorHandle, Props>(function Mod
|
||||
storeModel?.id,
|
||||
)
|
||||
: null
|
||||
const requestedRuntimeProvider = providers.find(
|
||||
(provider) => provider.id === requestedRuntimeSelection?.providerId,
|
||||
)
|
||||
const activeRuntimeSelection = requestedRuntimeSelection && providerChoices.some(
|
||||
(choice) => choice.providerId === requestedRuntimeSelection.providerId,
|
||||
)
|
||||
? requestedRuntimeSelection
|
||||
? {
|
||||
...requestedRuntimeSelection,
|
||||
modelId: requestedRuntimeProvider
|
||||
? resolveProviderRuntimeModelId(
|
||||
requestedRuntimeProvider,
|
||||
requestedRuntimeSelection.modelId,
|
||||
)
|
||||
: requestedRuntimeSelection.modelId,
|
||||
}
|
||||
: null
|
||||
|
||||
const selectedProviderChoice = activeRuntimeSelection
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
import { describe, expect, it } from 'vitest'
|
||||
import { normalizeRuntimeSelection } from './runtimeSelection'
|
||||
import { normalizeRuntimeSelection, resolveDefaultRuntimeSelection, resolveProviderRuntimeModelId, resolveProviderSlotModelId } from './runtimeSelection'
|
||||
import type { SavedProvider } from '../types/provider'
|
||||
|
||||
describe('normalizeRuntimeSelection', () => {
|
||||
it.each([
|
||||
@@ -138,3 +139,43 @@ describe('normalizeRuntimeSelection', () => {
|
||||
})
|
||||
})
|
||||
})
|
||||
|
||||
|
||||
describe('provider 1M runtime selection', () => {
|
||||
const provider: SavedProvider = {
|
||||
id: 'provider', name: 'Provider', presetId: 'custom', apiKey: 'fixture',
|
||||
baseUrl: 'http://127.0.0.1:9999', apiFormat: 'anthropic',
|
||||
models: { main: ' main-model ', haiku: 'fast-model', sonnet: 'balanced-model', opus: 'large-model' },
|
||||
model1mSupport: { main: true, haiku: false, sonnet: true, opus: false },
|
||||
}
|
||||
|
||||
it('materializes the active provider main slot by id and by legacy name', () => {
|
||||
for (const activeId of [provider.id, null]) {
|
||||
expect(resolveDefaultRuntimeSelection(activeId, provider.name, [provider], 'stale')).toEqual({
|
||||
providerId: provider.id, modelId: 'main-model[1m]',
|
||||
})
|
||||
}
|
||||
})
|
||||
|
||||
it('reconciles restored raw and marked IDs without losing a non-main model or effort', () => {
|
||||
expect(resolveProviderRuntimeModelId(provider, 'balanced-model')).toBe('balanced-model[1m]')
|
||||
expect(resolveProviderRuntimeModelId(provider, 'large-model[1m]')).toBe('large-model')
|
||||
expect(resolveProviderRuntimeModelId(provider, 'unmapped[1m]')).toBe('unmapped[1m]')
|
||||
})
|
||||
|
||||
it('keeps distinct choices for one raw model mapped to slots with different capabilities', () => {
|
||||
const shared = { ...provider, models: { main: 'shared', haiku: 'shared', sonnet: '', opus: '' } }
|
||||
expect(resolveProviderSlotModelId(shared, 'main')).toBe('shared[1m]')
|
||||
expect(resolveProviderSlotModelId(shared, 'haiku')).toBe('shared')
|
||||
expect(resolveProviderRuntimeModelId(shared, 'shared[1m]')).toBe('shared[1m]')
|
||||
expect(resolveProviderRuntimeModelId(shared, 'shared')).toBe('shared')
|
||||
})
|
||||
|
||||
it('preserves legacy explicit suffixes when flags are absent, but obeys an explicit off', () => {
|
||||
const legacy = { ...provider, model1mSupport: undefined, models: { ...provider.models, main: 'old[1m]', haiku: 'old:1m' } }
|
||||
expect(resolveProviderSlotModelId(legacy, 'main')).toBe('old[1m]')
|
||||
expect(resolveProviderSlotModelId(legacy, 'haiku')).toBe('old:1m')
|
||||
expect(resolveProviderSlotModelId({ ...legacy, model1mSupport: provider.model1mSupport }, 'haiku')).toBe('old')
|
||||
expect(resolveProviderSlotModelId({ ...legacy, model1mSupport: provider.model1mSupport }, 'main')).toBe('old[1m]')
|
||||
})
|
||||
})
|
||||
|
||||
@@ -18,6 +18,34 @@ import {
|
||||
type ModelReasoningProviderKind,
|
||||
} from '../../../src/shared/modelReasoning'
|
||||
|
||||
const PROVIDER_MODEL_SLOTS = ['main', 'haiku', 'sonnet', 'opus', 'fable'] as const
|
||||
|
||||
function baseProviderModelId(modelId: string): string {
|
||||
return modelId.trim().replace(/\[1m\]$/i, '').replace(/:1m$/i, '').trim()
|
||||
}
|
||||
|
||||
export function resolveProviderSlotModelId(
|
||||
provider: SavedProvider,
|
||||
slot: keyof SavedProvider['models'],
|
||||
): string {
|
||||
const modelId = provider.models[slot]?.trim() ?? ''
|
||||
const enabled = slot === 'fable' ? undefined : provider.model1mSupport?.[slot]
|
||||
// Missing flags are legacy configuration: preserve explicit model suffixes.
|
||||
if (!modelId || enabled === undefined) return modelId
|
||||
const baseModelId = baseProviderModelId(modelId)
|
||||
return enabled ? `${baseModelId}[1m]` : baseModelId
|
||||
}
|
||||
|
||||
export function resolveProviderRuntimeModelId(provider: SavedProvider, modelId: string): string {
|
||||
const candidates = PROVIDER_MODEL_SLOTS
|
||||
.filter((slot) => provider.models[slot]?.trim() &&
|
||||
baseProviderModelId(provider.models[slot]!) === baseProviderModelId(modelId))
|
||||
.map((slot) => resolveProviderSlotModelId(provider, slot))
|
||||
// A provider can map one ID to slots with different capabilities. Preserve
|
||||
// an exact runtime choice; otherwise reconcile old IDs in main-first order.
|
||||
return candidates.find((candidate) => candidate === modelId.trim()) ?? candidates[0] ?? modelId
|
||||
}
|
||||
|
||||
export function resolveActiveProviderRuntimeSelection(
|
||||
activeId: string | null,
|
||||
activeProviderName: string | null,
|
||||
@@ -32,7 +60,7 @@ export function resolveActiveProviderRuntimeSelection(
|
||||
const inferredProviderId = activeId ?? activeProvider?.id ?? null
|
||||
if (!inferredProviderId) return null
|
||||
|
||||
const providerMainModelId = activeProvider?.models.main.trim()
|
||||
const providerMainModelId = activeProvider ? resolveProviderSlotModelId(activeProvider, 'main') : undefined
|
||||
|
||||
return {
|
||||
providerId: inferredProviderId,
|
||||
|
||||
@@ -647,7 +647,7 @@ describe('EmptySession', () => {
|
||||
})
|
||||
})
|
||||
|
||||
it('materializes the active provider runtime and visible default effort before the first draft message', async () => {
|
||||
it.each([true, false])('materializes raw provider models with 1M=%s before the first draft message', async (enabled) => {
|
||||
useProviderStore.setState({
|
||||
providers: [{
|
||||
id: 'provider-minimax',
|
||||
@@ -658,11 +658,12 @@ describe('EmptySession', () => {
|
||||
apiFormat: 'anthropic',
|
||||
runtimeKind: 'anthropic_compatible',
|
||||
models: {
|
||||
main: 'MiniMax-M3[1m]',
|
||||
haiku: 'MiniMax-M3[1m]',
|
||||
sonnet: 'MiniMax-M3[1m]',
|
||||
opus: 'MiniMax-M3[1m]',
|
||||
main: 'MiniMax-M3',
|
||||
haiku: 'MiniMax-M3',
|
||||
sonnet: 'MiniMax-M3',
|
||||
opus: 'MiniMax-M3',
|
||||
},
|
||||
model1mSupport: { main: enabled, haiku: enabled, sonnet: enabled, opus: enabled },
|
||||
toolSearchEnabled: true,
|
||||
}],
|
||||
activeId: 'provider-minimax',
|
||||
@@ -682,7 +683,7 @@ describe('EmptySession', () => {
|
||||
|
||||
expect(useSessionRuntimeStore.getState().selections['draft-session']).toEqual({
|
||||
providerId: 'provider-minimax',
|
||||
modelId: 'MiniMax-M3[1m]',
|
||||
modelId: enabled ? 'MiniMax-M3[1m]' : 'MiniMax-M3',
|
||||
effortLevel: 'max',
|
||||
})
|
||||
expect(mocks.wsSend.mock.calls.slice(0, 3)).toEqual([
|
||||
@@ -691,7 +692,7 @@ describe('EmptySession', () => {
|
||||
{
|
||||
type: 'set_runtime_config',
|
||||
providerId: 'provider-minimax',
|
||||
modelId: 'MiniMax-M3[1m]',
|
||||
modelId: enabled ? 'MiniMax-M3[1m]' : 'MiniMax-M3',
|
||||
effortLevel: 'max',
|
||||
},
|
||||
],
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'
|
||||
import type { AgentTaskNotification } from '../types/chat'
|
||||
import type { MessageEntry } from '../types/session'
|
||||
import type { SavedProvider } from '../types/provider'
|
||||
import {
|
||||
buildMainSessionActivityModel,
|
||||
buildSessionActivityModel,
|
||||
@@ -34,6 +35,7 @@ const {
|
||||
connectionStateHandlers,
|
||||
sendSubagentMessageMock,
|
||||
tabStoreSnapshot,
|
||||
providerStoreSnapshot,
|
||||
} = vi.hoisted(() => ({
|
||||
sendMock: vi.fn(),
|
||||
getMemberBySessionIdMock: vi.fn<(sessionId: string) => any>(() => null),
|
||||
@@ -74,6 +76,11 @@ const {
|
||||
connectionStateHandlers: new Map<string, (state: 'connecting' | 'connected' | 'reconnecting' | 'disconnected') => void>(),
|
||||
sendSubagentMessageMock: vi.fn(async () => ({ ok: true })),
|
||||
tabStoreSnapshot: { tabs: [] as Array<Record<string, unknown>> },
|
||||
providerStoreSnapshot: { providers: [] as SavedProvider[], activeId: null as string | null },
|
||||
}))
|
||||
|
||||
vi.mock('./providerStore', () => ({
|
||||
useProviderStore: { getState: () => providerStoreSnapshot },
|
||||
}))
|
||||
|
||||
vi.mock('../lib/desktopNotifications', () => ({
|
||||
@@ -462,6 +469,8 @@ describe('chatStore background agent activity interleaving', () => {
|
||||
|
||||
describe('chatStore history mapping', () => {
|
||||
beforeEach(() => {
|
||||
providerStoreSnapshot.providers = []
|
||||
providerStoreSnapshot.activeId = null
|
||||
sendMock.mockReset()
|
||||
getMemberBySessionIdMock.mockReset()
|
||||
getMemberBySessionIdMock.mockReturnValue(null)
|
||||
@@ -5149,6 +5158,46 @@ describe('chatStore history mapping', () => {
|
||||
])
|
||||
})
|
||||
|
||||
it.each([true, false])('reconciles restored raw runtime models before reconnect and the next turn (1m=%s)', (enabled) => {
|
||||
const model = 'deepseek-v4.1-flash-expires-on-0910'
|
||||
providerStoreSnapshot.providers = [{
|
||||
id: 'provider-1', presetId: 'custom', name: 'DeepSeek', apiKey: 'fixture',
|
||||
baseUrl: 'http://127.0.0.1:1', apiFormat: 'anthropic',
|
||||
models: { main: model, haiku: '', sonnet: '', opus: '' },
|
||||
model1mSupport: { main: enabled, haiku: false, sonnet: false, opus: false },
|
||||
}]
|
||||
const staleSelection = { providerId: 'provider-1', modelId: `${model}${enabled ? '' : '[1m]'}`, effortLevel: 'high' as const }
|
||||
const expectedSelection = { ...staleSelection, modelId: `${model}${enabled ? '[1m]' : ''}` }
|
||||
useSessionRuntimeStore.getState().setSelection(TEST_SESSION_ID, staleSelection)
|
||||
useChatStore.getState().connectToSession(TEST_SESSION_ID, { prewarm: false, minimalBootstrap: true })
|
||||
expect(sendMock).toHaveBeenCalledWith(TEST_SESSION_ID, { type: 'set_runtime_config', ...expectedSelection })
|
||||
|
||||
// An edit during a busy turn is deferred until the next user message.
|
||||
useSessionRuntimeStore.getState().setSelection(TEST_SESSION_ID, staleSelection)
|
||||
sendMock.mockClear()
|
||||
useChatStore.getState().sendMessage(TEST_SESSION_ID, 'continue')
|
||||
expect(sendMock.mock.calls.slice(0, 2)).toEqual([
|
||||
[TEST_SESSION_ID, { type: 'set_runtime_config', ...expectedSelection }],
|
||||
[TEST_SESSION_ID, { type: 'user_message', content: 'continue', attachments: undefined }],
|
||||
])
|
||||
expect(useSessionRuntimeStore.getState().selections[TEST_SESSION_ID]).toEqual(expectedSelection)
|
||||
})
|
||||
|
||||
it('pins current provider capabilities before sending an older implicit-default session', () => {
|
||||
providerStoreSnapshot.activeId = 'provider-1'
|
||||
providerStoreSnapshot.providers = [{
|
||||
id: 'provider-1', presetId: 'custom', name: 'DeepSeek', apiKey: 'fixture',
|
||||
baseUrl: 'http://127.0.0.1:1', apiFormat: 'anthropic',
|
||||
models: { main: 'deepseek-v4.1', haiku: '', sonnet: '', opus: '' },
|
||||
model1mSupport: { main: true, haiku: false, sonnet: false, opus: false },
|
||||
}]
|
||||
useChatStore.getState().sendMessage(TEST_SESSION_ID, 'continue')
|
||||
expect(sendMock.mock.calls.slice(0, 2)).toEqual([
|
||||
[TEST_SESSION_ID, { type: 'set_runtime_config', providerId: 'provider-1', modelId: 'deepseek-v4.1[1m]' }],
|
||||
[TEST_SESSION_ID, { type: 'user_message', content: 'continue', attachments: undefined }],
|
||||
])
|
||||
})
|
||||
|
||||
it('does not prewarm unknown desktop sessions when connecting', () => {
|
||||
useChatStore.getState().connectToSession(TEST_SESSION_ID)
|
||||
|
||||
|
||||
@@ -7,6 +7,8 @@ import { useSessionStore } from './sessionStore'
|
||||
import { useCLITaskStore } from './cliTaskStore'
|
||||
import { useWorkflowStore } from './workflowStore'
|
||||
import { useSessionRuntimeStore } from './sessionRuntimeStore'
|
||||
import { useProviderStore } from './providerStore'
|
||||
import { resolveActiveProviderRuntimeSelection, resolveProviderRuntimeModelId } from '../lib/runtimeSelection'
|
||||
import { useTabStore } from './tabStore'
|
||||
import { randomSpinnerVerb } from '../config/spinnerVerbs'
|
||||
import { notifyDesktop } from '../lib/desktopNotifications'
|
||||
@@ -49,6 +51,13 @@ import type {
|
||||
} from '../types/slashCommand'
|
||||
|
||||
type ConnectionState = 'disconnected' | 'connecting' | 'connected' | 'reconnecting'
|
||||
|
||||
function reconcileProviderRuntimeSelection(selection: RuntimeSelection): RuntimeSelection {
|
||||
const provider = useProviderStore.getState().providers.find((entry) => entry.id === selection.providerId)
|
||||
if (!provider) return selection
|
||||
const modelId = resolveProviderRuntimeModelId(provider, selection.modelId)
|
||||
return modelId === selection.modelId ? selection : { ...selection, modelId }
|
||||
}
|
||||
type ToolCall = Extract<UIMessage, { type: 'tool_use' }>
|
||||
type CompactSummaryMessage = Extract<UIMessage, { type: 'compact_summary' }>
|
||||
|
||||
@@ -2765,7 +2774,7 @@ export const useChatStore = create<ChatStore>((set, get) => ({
|
||||
|
||||
const runtimeSelection = useSessionRuntimeStore.getState().selections[sessionId]
|
||||
if (runtimeSelection && options?.applyRuntimeSelection !== false) {
|
||||
wsManager.send(sessionId, { type: 'set_runtime_config', ...runtimeSelection })
|
||||
get().setSessionRuntime(sessionId, runtimeSelection)
|
||||
}
|
||||
if (
|
||||
options?.prewarm !== false &&
|
||||
@@ -2979,6 +2988,20 @@ export const useChatStore = create<ChatStore>((set, get) => ({
|
||||
return
|
||||
}
|
||||
|
||||
const selection = useSessionRuntimeStore.getState().selections[sessionId]
|
||||
if (selection) {
|
||||
const reconciled = reconcileProviderRuntimeSelection(selection)
|
||||
if (reconciled !== selection) get().setSessionRuntime(sessionId, selection)
|
||||
} else {
|
||||
const providers = useProviderStore.getState()
|
||||
const defaultSelection = resolveActiveProviderRuntimeSelection(
|
||||
providers.activeId, null, providers.providers, undefined,
|
||||
)
|
||||
if (defaultSelection) {
|
||||
useSessionRuntimeStore.getState().setSelection(sessionId, defaultSelection)
|
||||
get().setSessionRuntime(sessionId, defaultSelection)
|
||||
}
|
||||
}
|
||||
wsManager.send(sessionId, { type: 'user_message', content, attachments })
|
||||
},
|
||||
|
||||
@@ -3038,9 +3061,13 @@ export const useChatStore = create<ChatStore>((set, get) => ({
|
||||
},
|
||||
|
||||
setSessionRuntime: (sessionId, selection) => {
|
||||
const reconciled = reconcileProviderRuntimeSelection(selection)
|
||||
if (reconciled !== selection) {
|
||||
useSessionRuntimeStore.getState().setSelection(sessionId, reconciled)
|
||||
}
|
||||
wsManager.send(sessionId, {
|
||||
type: 'set_runtime_config',
|
||||
...selection,
|
||||
...reconciled,
|
||||
})
|
||||
},
|
||||
|
||||
|
||||
@@ -151,6 +151,35 @@ describe('providerStore runtime refresh', () => {
|
||||
})
|
||||
})
|
||||
|
||||
it.each([true, false])('refreshes saved model capabilities for idle, disconnected and draft selections (1m=%s)', async (enabled) => {
|
||||
const provider = makeProvider({ model1mSupport: { main: enabled, haiku: false, sonnet: false, opus: enabled } })
|
||||
providersApiMock.update.mockResolvedValue({ provider })
|
||||
providersApiMock.list.mockResolvedValue({ providers: [provider], activeId: provider.id })
|
||||
chatStoreState.sessions = {
|
||||
idle: { connectionState: 'connected', chatState: 'idle' },
|
||||
offline: { connectionState: 'disconnected', chatState: 'idle' },
|
||||
busy: { connectionState: 'connected', chatState: 'streaming' },
|
||||
}
|
||||
const previous = `model-opus${enabled ? '' : '[1m]'}`
|
||||
runtimeStoreState.selections = Object.fromEntries(
|
||||
['idle', 'offline', 'busy', '__draft__', 'restored'].map((key) => [key, {
|
||||
providerId: provider.id,
|
||||
modelId: previous,
|
||||
effortLevel: 'high',
|
||||
}]),
|
||||
)
|
||||
|
||||
const { useProviderStore } = await import('./providerStore')
|
||||
await useProviderStore.getState().updateProvider(provider.id, { model1mSupport: provider.model1mSupport })
|
||||
|
||||
const selection = { providerId: provider.id, modelId: `model-opus${enabled ? '[1m]' : ''}`, effortLevel: 'high' }
|
||||
for (const key of ['idle', 'offline', '__draft__', 'restored']) {
|
||||
expect(setSelectionMock).toHaveBeenCalledWith(key, selection)
|
||||
}
|
||||
expect(setSelectionMock).not.toHaveBeenCalledWith('busy', expect.anything())
|
||||
expect(setSessionRuntimeMock.mock.calls).toEqual([['idle', selection]])
|
||||
})
|
||||
|
||||
it('does not restart busy sessions while a provider update is saved', async () => {
|
||||
const provider = makeProvider()
|
||||
providersApiMock.update.mockResolvedValue({ provider })
|
||||
@@ -167,6 +196,20 @@ describe('providerStore runtime refresh', () => {
|
||||
expect(setSessionRuntimeMock).not.toHaveBeenCalled()
|
||||
})
|
||||
|
||||
it('reconciles restored selections when provider capabilities arrive after connection', async () => {
|
||||
const provider = makeProvider({ model1mSupport: { main: true, haiku: false, sonnet: false, opus: false } })
|
||||
providersApiMock.list.mockResolvedValue({ providers: [provider], activeId: provider.id })
|
||||
chatStoreState.sessions = { restored: { connectionState: 'connected', chatState: 'idle' } }
|
||||
runtimeStoreState.selections = { restored: { providerId: provider.id, modelId: 'model-main' } }
|
||||
const { useProviderStore } = await import('./providerStore')
|
||||
useProviderStore.setState({ hasLoadedProviders: false })
|
||||
|
||||
await useProviderStore.getState().fetchProviders()
|
||||
|
||||
expect(setSessionRuntimeMock).toHaveBeenCalledWith('restored', { providerId: provider.id, modelId: 'model-main[1m]' })
|
||||
expect(setSelectionMock).toHaveBeenCalledWith('restored', { providerId: provider.id, modelId: 'model-main[1m]' })
|
||||
})
|
||||
|
||||
it('sets the OpenAI default model when activating built-in ChatGPT Official', async () => {
|
||||
providersApiMock.activate.mockResolvedValue({ ok: true })
|
||||
providersApiMock.list.mockResolvedValue({
|
||||
|
||||
@@ -16,6 +16,7 @@ import {
|
||||
GROK_OFFICIAL_PROVIDER_ID,
|
||||
} from '../constants/grokOfficialProvider'
|
||||
import { BUNDLED_PROVIDER_PRESETS } from '../config/providerPresets'
|
||||
import { resolveProviderRuntimeModelId, resolveProviderSlotModelId } from '../lib/runtimeSelection'
|
||||
import type {
|
||||
SavedProvider,
|
||||
CreateProviderInput,
|
||||
@@ -112,8 +113,8 @@ function mergeSavedOrderIntoProviderOrder(providerOrder: string[], savedOrder: s
|
||||
|
||||
function providerModelIds(provider: SavedProvider): Set<string> {
|
||||
return new Set(
|
||||
Object.values(provider.models)
|
||||
.map((modelId) => modelId.trim())
|
||||
(Object.keys(provider.models) as Array<keyof SavedProvider['models']>)
|
||||
.map((slot) => resolveProviderSlotModelId(provider, slot))
|
||||
.filter(Boolean),
|
||||
)
|
||||
}
|
||||
@@ -125,11 +126,12 @@ function resolveRuntimeRefreshSelection(
|
||||
): RuntimeSelection | null {
|
||||
if (currentSelection?.providerId === provider.id) {
|
||||
const modelIds = providerModelIds(provider)
|
||||
const modelId = resolveProviderRuntimeModelId(provider, currentSelection.modelId)
|
||||
return {
|
||||
providerId: provider.id,
|
||||
modelId: modelIds.has(currentSelection.modelId)
|
||||
? currentSelection.modelId
|
||||
: provider.models.main,
|
||||
modelId: modelIds.has(modelId)
|
||||
? modelId
|
||||
: resolveProviderSlotModelId(provider, 'main'),
|
||||
...(currentSelection.effortLevel ? { effortLevel: currentSelection.effortLevel } : {}),
|
||||
}
|
||||
}
|
||||
@@ -137,31 +139,38 @@ function resolveRuntimeRefreshSelection(
|
||||
if (!currentSelection && activeId === provider.id) {
|
||||
return {
|
||||
providerId: provider.id,
|
||||
modelId: provider.models.main,
|
||||
modelId: resolveProviderSlotModelId(provider, 'main'),
|
||||
}
|
||||
}
|
||||
|
||||
return null
|
||||
}
|
||||
|
||||
function refreshConnectedSessionsForProvider(provider: SavedProvider, activeId: string | null) {
|
||||
function refreshSessionsForProvider(provider: SavedProvider, activeId: string | null, onlyChanged = false) {
|
||||
const chatStore = useChatStore.getState()
|
||||
const runtimeStore = useSessionRuntimeStore.getState()
|
||||
|
||||
for (const [sessionId, session] of Object.entries(chatStore.sessions)) {
|
||||
if (session.connectionState !== 'connected' || session.chatState !== 'idle') {
|
||||
continue
|
||||
}
|
||||
const sessionIds = new Set([...Object.keys(chatStore.sessions), ...Object.keys(runtimeStore.selections)])
|
||||
for (const sessionId of sessionIds) {
|
||||
const session = chatStore.sessions[sessionId]
|
||||
// Do not replace a running turn. Its next send reconciles the stored model
|
||||
// against the current provider before submitting the prompt.
|
||||
if (session && session.chatState !== 'idle') continue
|
||||
const currentSelection = runtimeStore.selections[sessionId]
|
||||
if (!currentSelection && session?.connectionState !== 'connected') continue
|
||||
|
||||
const selection = resolveRuntimeRefreshSelection(
|
||||
provider,
|
||||
activeId,
|
||||
runtimeStore.selections[sessionId],
|
||||
currentSelection,
|
||||
)
|
||||
if (!selection) continue
|
||||
if (onlyChanged && (!currentSelection || currentSelection.modelId === selection.modelId)) continue
|
||||
|
||||
runtimeStore.setSelection(sessionId, selection)
|
||||
chatStore.setSessionRuntime(sessionId, selection)
|
||||
if (session?.connectionState === 'connected') {
|
||||
chatStore.setSessionRuntime(sessionId, selection)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -177,6 +186,7 @@ export const useProviderStore = create<ProviderStore>((set, get) => ({
|
||||
fetchProviders: async () => {
|
||||
set({ isLoading: true, error: null })
|
||||
try {
|
||||
const firstLoad = !get().hasLoadedProviders
|
||||
const { providers, activeId, providerOrder } = await providersApi.list()
|
||||
set({
|
||||
providers,
|
||||
@@ -185,6 +195,9 @@ export const useProviderStore = create<ProviderStore>((set, get) => ({
|
||||
hasLoadedProviders: true,
|
||||
isLoading: false,
|
||||
})
|
||||
if (firstLoad) {
|
||||
for (const provider of providers) refreshSessionsForProvider(provider, activeId, true)
|
||||
}
|
||||
} catch (err) {
|
||||
set({
|
||||
isLoading: false,
|
||||
@@ -211,7 +224,7 @@ export const useProviderStore = create<ProviderStore>((set, get) => ({
|
||||
await settings.fetchAll()
|
||||
}
|
||||
}
|
||||
refreshConnectedSessionsForProvider(provider, activeId)
|
||||
refreshSessionsForProvider(provider, activeId)
|
||||
return provider
|
||||
},
|
||||
|
||||
|
||||
@@ -3,6 +3,9 @@ import { mkdtemp, rm } from 'node:fs/promises'
|
||||
import { tmpdir } from 'node:os'
|
||||
import { join, resolve } from 'node:path'
|
||||
import { createSandboxedTestEnvironment } from '../../../scripts/pr/test-environment.js'
|
||||
import { resolveDefaultRuntimeSelection } from '../../../desktop/src/lib/runtimeSelection.js'
|
||||
import { buildProviderManagedEnv } from '../../server/services/providerRuntimeEnv.js'
|
||||
import type { SavedProvider } from '../../server/types/provider.js'
|
||||
|
||||
const contextBeta = 'context-1m-2025-08-07'
|
||||
const root = resolve(import.meta.dir, '../../..')
|
||||
@@ -12,9 +15,17 @@ async function runRelay(options: {
|
||||
injectHeader?: boolean
|
||||
env?: Record<string, string>
|
||||
args?: string[]
|
||||
provider?: SavedProvider
|
||||
requireContextBeta?: boolean
|
||||
} = {}) {
|
||||
const sandbox = await mkdtemp(join(tmpdir(), 'cc-haha-context-beta-'))
|
||||
const requests: { model: string, beta: string, status: number }[] = []
|
||||
const requests: {
|
||||
model: string
|
||||
beta: string
|
||||
status: number
|
||||
thinking?: unknown
|
||||
outputConfig?: unknown
|
||||
}[] = []
|
||||
const server = Bun.serve({
|
||||
hostname: '127.0.0.1',
|
||||
port: 0,
|
||||
@@ -22,14 +33,18 @@ async function runRelay(options: {
|
||||
if (!new URL(request.url).pathname.endsWith('/messages')) {
|
||||
return new Response('Unexpected route', { status: 404 })
|
||||
}
|
||||
const body = await request.json() as { model: string }
|
||||
const body = await request.json() as { model: string, thinking?: unknown, output_config?: unknown }
|
||||
const beta = request.headers.get('anthropic-beta') ?? ''
|
||||
const forwardedHeaders = new Headers(request.headers)
|
||||
if (options.injectHeader) {
|
||||
forwardedHeaders.set('anthropic-beta', [beta, contextBeta].filter(Boolean).join(','))
|
||||
}
|
||||
const accepted = forwardedHeaders.get('anthropic-beta')?.split(',').includes(contextBeta)
|
||||
requests.push({ model: body.model, beta, status: accepted ? 200 : 400 })
|
||||
const accepted = options.requireContextBeta === false ||
|
||||
forwardedHeaders.get('anthropic-beta')?.split(',').includes(contextBeta)
|
||||
requests.push({
|
||||
model: body.model, beta, status: accepted ? 200 : 400,
|
||||
thinking: body.thinking, outputConfig: body.output_config,
|
||||
})
|
||||
if (!accepted) {
|
||||
return Response.json({ type: 'error', error: {
|
||||
type: 'invalid_request_error', message: '1m 上下文已经全量可用,请启用 1m 上下文后重试',
|
||||
@@ -67,9 +82,14 @@ async function runRelay(options: {
|
||||
CALLER_DIR: sandbox,
|
||||
ANTHROPIC_API_KEY: 'loopback-test-key',
|
||||
ANTHROPIC_BASE_URL: `http://127.0.0.1:${server.port}`,
|
||||
ANTHROPIC_DEFAULT_OPUS_MODEL: 'claude-opus-5[1m]',
|
||||
ANTHROPIC_DEFAULT_OPUS_MODEL_SUPPORTED_CAPABILITIES: 'thinking,effort',
|
||||
CLAUDE_CODE_MODEL_CONTEXT_WINDOWS: '{"claude-opus-5":1000000}',
|
||||
...(options.provider ? buildProviderManagedEnv({
|
||||
...options.provider,
|
||||
baseUrl: `http://127.0.0.1:${server.port}`,
|
||||
}) : {
|
||||
ANTHROPIC_DEFAULT_OPUS_MODEL: 'claude-opus-5[1m]',
|
||||
ANTHROPIC_DEFAULT_OPUS_MODEL_SUPPORTED_CAPABILITIES: 'thinking,effort',
|
||||
CLAUDE_CODE_MODEL_CONTEXT_WINDOWS: '{"claude-opus-5":1000000}',
|
||||
}),
|
||||
CLAUDE_CODE_DISABLE_NONESSENTIAL_TRAFFIC: '1',
|
||||
CLAUDE_CODE_SKIP_UPDATE_CHECK: '1',
|
||||
...options.env,
|
||||
@@ -132,3 +152,48 @@ test('real CLI sends the opted-in 1M beta to third-party Anthropic relays', asyn
|
||||
expect(result.requests[0]?.model).toBe('claude-opus-5')
|
||||
expect(result.requests[0]?.beta.split(',')).toContain(contextBeta)
|
||||
}, 40_000)
|
||||
|
||||
for (const { enabled, disableBetas } of [
|
||||
{ enabled: true, disableBetas: false },
|
||||
{ enabled: false, disableBetas: false },
|
||||
{ enabled: true, disableBetas: true },
|
||||
]) {
|
||||
test(`desktop provider selection reaches explicit CLI model and native relay (1M: ${enabled}, disable betas: ${disableBetas})`, async () => {
|
||||
const model = 'deepseek-v4.1-flash-expires-on-0910'
|
||||
const provider: SavedProvider = {
|
||||
id: 'desktop-context-fixture',
|
||||
name: 'Desktop context fixture',
|
||||
presetId: 'custom',
|
||||
apiKey: 'loopback-test-key',
|
||||
baseUrl: 'http://127.0.0.1:1',
|
||||
apiFormat: 'anthropic',
|
||||
runtimeKind: 'anthropic_compatible',
|
||||
models: { main: model, haiku: model, sonnet: model, opus: model },
|
||||
model1mSupport: { main: enabled, haiku: enabled, sonnet: enabled, opus: enabled },
|
||||
modelContextWindows: { [model]: 1_000_000 },
|
||||
disableExperimentalBetas: disableBetas,
|
||||
}
|
||||
// Follow new-conversation selection into the same explicit --model used by
|
||||
// conversationService. The managed environment alone cannot repair a raw ID.
|
||||
const selection = resolveDefaultRuntimeSelection(provider.id, provider.name, [provider], undefined)
|
||||
const result = await runRelay({
|
||||
model: selection.modelId,
|
||||
provider,
|
||||
requireContextBeta: false,
|
||||
env: { CLAUDE_CODE_EFFORT_LEVEL: 'high' },
|
||||
})
|
||||
expect(result.exitCode, JSON.stringify(result)).toBe(0)
|
||||
expect(result.stdout).toContain('relay-ok')
|
||||
expect(result.requests).toHaveLength(1)
|
||||
const request = result.requests[0]!
|
||||
expect(request.model).toBe(model)
|
||||
expect(request.beta.split(',').includes(contextBeta)).toBe(enabled && !disableBetas)
|
||||
expect(selection.modelId).toBe(enabled ? `${model}[1m]` : model)
|
||||
// Direct Anthropic relays suppress experimental effort when betas are
|
||||
// disabled, while the model's existing thinking capability is independent.
|
||||
expect(request.outputConfig).toEqual(disableBetas ? undefined : { effort: 'high' })
|
||||
expect(request.thinking).toMatchObject({ type: 'adaptive' })
|
||||
if (disableBetas) expect(request.beta).toBe('')
|
||||
else expect(request.beta).toContain('effort-2025-11-24')
|
||||
}, 40_000)
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user