diff --git a/desktop/src/components/chat/ChatInput.test.tsx b/desktop/src/components/chat/ChatInput.test.tsx index 9ff02252..8ce47eca 100644 --- a/desktop/src/components/chat/ChatInput.test.tsx +++ b/desktop/src/components/chat/ChatInput.test.tsx @@ -1093,6 +1093,16 @@ describe('ChatInput file mentions', () => { // moved. The button has no label to shed any more — it is one round icon at // every width — so what needs pinning is that it does *not* change with the // column, leaving the location as the only thing that degrades (next test). + it.each(['mixed', 'unknown'] as const)('blocks sending in a %s protocol session with new-session guidance', async (sessionApiFormat) => { + useSessionStore.setState((state) => ({ sessions: state.sessions.map((session) => ({ ...session, sessionApiFormat })) })) + render() + setComposerText('Continue please', 15) + expect(screen.getByRole('alert')).toHaveTextContent('Start a new session') + expect(screen.getByRole('button', { name: 'Run' })).toBeDisabled() + fireEvent.keyDown(screen.getByRole('textbox'), { key: 'Enter' }) + expect(mocks.wsSend).not.toHaveBeenCalledWith(sessionId, expect.objectContaining({type:'user_message'})) + }) + it('keeps the send button a fixed circle as the column narrows', async () => { const column = stubComposerColumnWidth(700) diff --git a/desktop/src/components/chat/ChatInput.tsx b/desktop/src/components/chat/ChatInput.tsx index 663f1a1d..e3fb4ad2 100644 --- a/desktop/src/components/chat/ChatInput.tsx +++ b/desktop/src/components/chat/ChatInput.tsx @@ -212,6 +212,8 @@ export function ChatInput({ variant = 'default', compact = false }: ChatInputPro : undefined const runtimeModelLabel = runtimeSelection?.modelId ?? currentModel?.name ?? currentModel?.id const activeSession = useSessionStore((state) => activeTabId ? state.sessions.find((session) => session.id === activeTabId) ?? null : null) + const sessionProtocol = sessionState?.sessionApiFormat ?? activeSession?.sessionApiFormat + const isProtocolBlocked = sessionProtocol === 'mixed' || sessionProtocol === 'unknown' const loadedMessageCount = sessionState?.messages?.length ?? 0 const messageCount = Math.max(loadedMessageCount, activeSession?.messageCount ?? 0) const memberInfo = useTeamStore((s) => activeTabId ? s.getMemberBySessionId(activeTabId) : null) @@ -298,7 +300,7 @@ export function ChatInput({ variant = 'default', compact = false }: ChatInputPro const pendingSlashUiAction = !isMemberSession && input.trim().startsWith('/') ? resolveSlashUiAction(input.trim().slice(1)) : null - const canSubmit = !isWorkspaceMissing && + const canSubmit = !isWorkspaceMissing && !isProtocolBlocked && !launchTransitioning && !isPreparingTurn && (!showLaunchControls || launchReady || !!pendingSlashUiAction) && @@ -705,6 +707,7 @@ export function ChatInput({ variant = 'default', compact = false }: ChatInputPro const handleSubmit = async () => { const text = input.trim() + if (isProtocolBlocked) return if ((!text && ((!attachments.length && !hasWorkspaceReferences) || isMemberSession)) || isWorkspaceMissing) return if (pendingSlashUiAction?.type === 'panel') { @@ -1086,6 +1089,11 @@ export function ChatInput({ variant = 'default', compact = false }: ChatInputPro : `${isMobileComposer ? 'mx-0 max-w-none' : 'mx-auto max-w-[900px]'}` } > + {isProtocolBlocked && ( +

+ {t(sessionProtocol === 'mixed' ? 'model.protocolMixed' : 'model.protocolUnknown')} +

+ )}
{ }) import { ModelSelector } from './ModelSelector' +import { useSessionStore } from '@/stores/sessionStore' +import type { SessionListItem } from '@/types/session' +import type { SessionProtocolState, SessionApiFormat } from '../../../../src/shared/sessionProtocol' import { useChatStore } from '../../stores/chatStore' import { useHahaOAuthStore } from '../../stores/hahaOAuthStore' import { useHahaOpenAIOAuthStore } from '../../stores/hahaOpenAIOAuthStore' @@ -47,6 +50,7 @@ afterEach(() => { useProviderStore.setState(useProviderStore.getInitialState(), true) useSessionRuntimeStore.setState(useSessionRuntimeStore.getInitialState(), true) useChatStore.setState(useChatStore.getInitialState(), true) + useSessionStore.setState(useSessionStore.getInitialState(), true) useHahaOAuthStore.setState(useHahaOAuthStore.getInitialState(), true) useHahaOpenAIOAuthStore.setState(useHahaOpenAIOAuthStore.getInitialState(), true) useHahaGrokOAuthStore.setState(useHahaGrokOAuthStore.getInitialState(), true) @@ -61,6 +65,88 @@ beforeEach(() => { }) describe('ModelSelector', () => { + function prepareProtocolSession(protocol: SessionProtocolState | undefined, apiFormat: SessionApiFormat = 'anthropic') { + useSettingsStore.setState({ locale: 'en' }) + useSessionStore.setState({ sessions: [{ + id: 'protocol-session', title: 'Protocol session', messageCount: protocol ? 2 : 0, + sessionApiFormat: protocol, workDir: '/fixture/project', + } as SessionListItem] }) + useProviderStore.setState({ + activeId: 'provider-a', hasLoadedProviders: true, isLoading: false, + providers: [ + { id: 'provider-a', name: 'Provider A', apiFormat, model: 'model-a' }, + { id: 'provider-b', name: 'Provider B', apiFormat, model: 'model-b' }, + { id: 'provider-c', name: 'Provider C', apiFormat: apiFormat === 'anthropic' ? 'openai_responses' : 'anthropic', model: 'model-c' }, + ].map(({ model, ...provider }) => ({ + ...provider, apiFormat: provider.apiFormat as SessionApiFormat, + presetId: 'custom', apiKey: 'fixture', baseUrl: 'http://127.0.0.1:9999', + models: { main: model, haiku: model, sonnet: model, opus: model }, + })), + }) + useSessionRuntimeStore.getState().setSelection('protocol-session', { providerId: 'provider-a', modelId: 'model-a' }) + } + + it.each(['anthropic', 'openai_chat', 'openai_responses'] as const)( + 'allows other providers using %s and disables a different protocol', async (protocol) => { + prepareProtocolSession(protocol, protocol) + const runtimeChange = vi.fn() + render() + await clickByRole('model-a, Provider A') + expect(screen.getByRole('button', { name: /model-c/ })).toBeDisabled() + expect(screen.getByText('Different API protocol — start a new session')).toBeVisible() + expect(screen.getByRole('button', { name: /model-b/ })).toBeEnabled() + await clickByRole(/model-b/) + expect(runtimeChange).toHaveBeenCalledWith(expect.objectContaining({providerId:'provider-b',modelId:'model-b'})) + expect(useSessionStore.getState().sessions[0]?.sessionApiFormat).toBe(protocol) + }, + ) + + it.each(['mixed', 'unknown'] as const)('explains %s legacy history and requires a new session', async (protocol) => { + prepareProtocolSession(protocol) + render() + await clickByRole('model-a, Provider A') + for (const model of ['model-a', 'model-b', 'model-c']) { + expect(screen.getAllByRole('button', { name: new RegExp(model) }).at(-1)).toBeDisabled() + } + expect(screen.getByRole('status')).toHaveTextContent('Start a new session') + expect(screen.getByRole('button', { name: 'New session' })).toBeEnabled() + }) + + it('leaves a new session unlocked when its picker is opened or changed', async () => { + prepareProtocolSession(undefined) + render() + await clickByRole('model-a, Provider A') + await clickByRole(/model-c/) + expect(useSessionStore.getState().sessions[0]?.sessionApiFormat).toBeUndefined() + expect(useChatStore.getState().sessions['protocol-session']?.sessionApiFormat).toBeUndefined() + }) + + it('updates an open picker when the server confirms a protocol lock', async () => { + prepareProtocolSession(undefined) + useChatStore.setState({ sessions: { 'protocol-session': useChatStore.getState().getSession('protocol-session') } }) + render() + await clickByRole('model-a, Provider A') + expect(screen.getByRole('button', { name: /model-c/ })).toBeEnabled() + act(() => useChatStore.getState().handleServerMessage('protocol-session', {type:'session_protocol',sessionApiFormat:'anthropic'})) + expect(screen.getByRole('button', { name: /model-c/ })).toBeDisabled() + expect(screen.getByRole('status')).toHaveTextContent('Messages') + }) + + it('starts an unlocked session in the same directory without altering the old session', async () => { + prepareProtocolSession('mixed') + const createSession = vi.fn(async () => 'protocol-new') + const connectToSession = vi.fn() + useSessionStore.setState({ createSession }) + useChatStore.setState({ connectToSession }) + render() + await clickByRole('model-a, Provider A') + await clickByRole('New session') + expect(createSession).toHaveBeenCalledWith('/fixture/project') + expect(useTabStore.getState().activeTabId).toBe('protocol-new') + expect(connectToSession).toHaveBeenCalledWith('protocol-new') + expect(useSessionStore.getState().sessions[0]?.sessionApiFormat).toBe('mixed') + }) + it('keeps the current Claude Official catalog visible when the API returns legacy settings models', async () => { const legacyModels: ModelInfo[] = [ { id: 'claude-opus-4-7', name: 'Opus 4.7', description: 'Legacy Opus', context: '1m' }, diff --git a/desktop/src/components/controls/ModelSelector.tsx b/desktop/src/components/controls/ModelSelector.tsx index 177510a5..8d4ec03f 100644 --- a/desktop/src/components/controls/ModelSelector.tsx +++ b/desktop/src/components/controls/ModelSelector.tsx @@ -10,6 +10,9 @@ import { OPENAI_OFFICIAL_PROVIDER_ID, } from '../../constants/openaiOfficialProvider' import { useTranslation } from '../../i18n' +import { Button } from '@/components/ui/Button' +import { resolveProviderApiFormat, type SessionApiFormat } from '../../../../src/shared/sessionProtocol' +import { useSessionStore } from '@/stores/sessionStore' import { useChatStore } from '../../stores/chatStore' import { useProviderStore } from '../../stores/providerStore' import { DRAFT_RUNTIME_SELECTION_KEY, useSessionRuntimeStore } from '../../stores/sessionRuntimeStore' @@ -48,6 +51,7 @@ type ProviderChoice = { providerId: string | null providerName: string isDefault: boolean + apiFormat?: SessionApiFormat models: ModelInfo[] } @@ -112,6 +116,7 @@ function officialChoices( providerId, providerName: officialName, isDefault, + apiFormat: resolveProviderApiFormat(providerId), models, } } @@ -221,6 +226,7 @@ function buildProviderChoices( providerId: provider.id, providerName: provider.name, isDefault: activeId === provider.id, + apiFormat: resolveProviderApiFormat(provider.id, provider), models: buildProviderModels(provider, labels), }) } @@ -271,6 +277,16 @@ export const ModelSelector = forwardRef(function Mod const runtimeSelection = useSessionRuntimeStore((state) => runtimeKey ? state.selections[runtimeKey] : undefined, ) + const listedSession = useSessionStore((state) => + runtimeKey ? state.sessions.find((session) => session.id === runtimeKey) : undefined, + ) + const liveProtocol = useChatStore((state) => + runtimeKey ? state.sessions[runtimeKey]?.sessionApiFormat : undefined, + ) + const sessionProtocol = runtimeKey && runtimeKey !== DRAFT_RUNTIME_SELECTION_KEY + ? liveProtocol ?? listedSession?.sessionApiFormat + : undefined + const [creatingSession, setCreatingSession] = useState(false) const [open, setOpen] = useState(false) const [effortOpen, setEffortOpen] = useState(false) const [searchQuery, setSearchQuery] = useState('') @@ -565,8 +581,47 @@ export const ModelSelector = forwardRef(function Mod open: openSelector, }), [openSelector]) + const protocolLabel = (protocol: SessionApiFormat) => ({ + anthropic: 'Messages', + openai_chat: 'Chat Completions', + openai_responses: 'Responses', + })[protocol] + const protocolNotice = sessionProtocol === 'mixed' + ? t('model.protocolMixed') + : sessionProtocol === 'unknown' + ? t('model.protocolUnknown') + : sessionProtocol + ? t('model.protocolLocked', { protocol: protocolLabel(sessionProtocol) }) + : null + const blocksProtocol = (apiFormat: SessionApiFormat | undefined) => Boolean( + sessionProtocol && sessionProtocol !== apiFormat, + ) + const createNewSession = async () => { + if (creatingSession) return + setCreatingSession(true) + try { + const sessionId = await useSessionStore.getState().createSession( + listedSession?.workDir || listedSession?.projectRoot || undefined, + ) + if (requestedRuntimeSelection) { + useSessionRuntimeStore.getState().setSelection(sessionId, requestedRuntimeSelection) + } + useTabStore.getState().openTab(sessionId, t('sidebar.newSession')) + useChatStore.getState().connectToSession(sessionId) + setOpen(false) + } catch (error) { + useUIStore.getState().addToast({ + type: 'error', + message: error instanceof Error ? error.message : t('empty.failedToCreate'), + }) + } finally { + setCreatingSession(false) + } + } + const handleRuntimeSelect = (selection: RuntimeSelection) => { const provider = providers.find((entry) => entry.id === selection.providerId) + if (blocksProtocol(resolveProviderApiFormat(selection.providerId, provider))) return const normalizedSelection = normalizeRuntimeSelection( selection, provider?.apiFormat, @@ -607,6 +662,14 @@ export const ModelSelector = forwardRef(function Mod const dropdownContent = ( <> + {protocolNotice && ( +
+

{protocolNotice}

+ +
+ )} {/* The header stays OUTSIDE the scroll region: a sticky header inside `overflow-y-auto` depends on the engine compositing it above the scrolling layer, and on the desktop shell scrolled items paint @@ -650,10 +713,14 @@ export const ModelSelector = forwardRef(function Mod const isSelected = activeRuntimeSelection?.providerId === choice.providerId && activeRuntimeSelection.modelId === model.id + const protocolBlocked = blocksProtocol(choice.apiFormat) return (