mirror of
https://github.com/NanmiCoder/claude-code-haha.git
synced 2026-10-10 20:03:13 +08:00
fix(session): lock model switching to the session API protocol
This commit is contained in:
@@ -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(<ChatInput compact />)
|
||||
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)
|
||||
|
||||
|
||||
@@ -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 && (
|
||||
<p role="alert" className="mb-2 text-sm text-[var(--color-text-secondary)]">
|
||||
{t(sessionProtocol === 'mixed' ? 'model.protocolMixed' : 'model.protocolUnknown')}
|
||||
</p>
|
||||
)}
|
||||
<div
|
||||
ref={panelRef}
|
||||
data-testid="chat-input-panel"
|
||||
|
||||
@@ -16,6 +16,9 @@ vi.mock('../../lib/desktopRuntime', async (importOriginal) => {
|
||||
})
|
||||
|
||||
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(<ModelSelector runtimeKey="protocol-session" onRuntimeSelectionChange={runtimeChange} />)
|
||||
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(<ModelSelector runtimeKey="protocol-session" />)
|
||||
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(<ModelSelector runtimeKey="protocol-session" />)
|
||||
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(<ModelSelector runtimeKey="protocol-session" />)
|
||||
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(<ModelSelector runtimeKey="protocol-session" />)
|
||||
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' },
|
||||
|
||||
@@ -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<ModelSelectorHandle, Props>(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<ModelSelectorHandle, Props>(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<ModelSelectorHandle, Props>(function Mod
|
||||
|
||||
const dropdownContent = (
|
||||
<>
|
||||
{protocolNotice && (
|
||||
<div className="flex-none border-b border-[var(--color-border)] px-4 py-3">
|
||||
<p role="status" className="text-xs text-[var(--color-text-secondary)]">{protocolNotice}</p>
|
||||
<Button size="sm" variant="ghost" disabled={creatingSession} onClick={() => void createNewSession()}>
|
||||
{t('sidebar.newSession')}
|
||||
</Button>
|
||||
</div>
|
||||
)}
|
||||
{/* 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<ModelSelectorHandle, Props>(function Mod
|
||||
const isSelected =
|
||||
activeRuntimeSelection?.providerId === choice.providerId &&
|
||||
activeRuntimeSelection.modelId === model.id
|
||||
const protocolBlocked = blocksProtocol(choice.apiFormat)
|
||||
return (
|
||||
<button
|
||||
key={`${choice.providerId ?? 'official'}:${model.id}`}
|
||||
disabled={protocolBlocked}
|
||||
title={protocolBlocked ? t('model.protocolRequiresNewSession') : undefined}
|
||||
onClick={() => {
|
||||
if (protocolBlocked) return
|
||||
const supportedEfforts = model.supportedReasoningEfforts
|
||||
const explicitEffort = activeRuntimeSelection?.effortLevel
|
||||
const selectedProvider = providers.find(
|
||||
@@ -696,7 +763,7 @@ export const ModelSelector = forwardRef<ModelSelectorHandle, Props>(function Mod
|
||||
})
|
||||
}}
|
||||
className={`
|
||||
w-full rounded-[var(--radius-md)] border px-3 text-left transition-colors
|
||||
w-full rounded-[var(--radius-md)] border px-3 text-left transition-colors disabled:cursor-not-allowed disabled:opacity-50
|
||||
${isMobileBrowser ? 'min-h-[56px] py-3' : 'py-2'}
|
||||
${isSelected
|
||||
? 'border-[var(--color-model-option-selected-border)] bg-[var(--color-model-option-selected-bg)]'
|
||||
@@ -711,6 +778,11 @@ export const ModelSelector = forwardRef<ModelSelectorHandle, Props>(function Mod
|
||||
<div className="truncate font-mono text-[13px] font-medium text-[var(--color-text-primary)]">
|
||||
{model.name}
|
||||
</div>
|
||||
{protocolBlocked && (
|
||||
<div aria-hidden="true" className="mt-0.5 text-[11px] text-[var(--color-text-tertiary)]">
|
||||
{t('model.protocolRequiresNewSession')}
|
||||
</div>
|
||||
)}
|
||||
{model.description && (
|
||||
<div className="mt-0.5 truncate pr-[6px] text-[11px] text-[var(--color-text-tertiary)]">
|
||||
{model.description}
|
||||
|
||||
@@ -2342,6 +2342,10 @@ Row 9, all 8 cells: continuing from straight down, turning left through lower-le
|
||||
// ─── Model Selector ──────────────────────────────────────
|
||||
'model.selectModel': 'Select model',
|
||||
'model.configureProvider': 'Configure model provider',
|
||||
'model.protocolLocked': 'This session uses {protocol}. To use a different API protocol, start a new session.',
|
||||
'model.protocolMixed': 'This session contains history from multiple API protocols and cannot continue. Start a new session from the model menu.',
|
||||
'model.protocolUnknown': 'The API protocol of this older session could not be determined. Start a new session from the model menu.',
|
||||
'model.protocolRequiresNewSession': 'Different API protocol — start a new session',
|
||||
'model.configuration': 'Model Configuration',
|
||||
'model.searchPlaceholder': 'Search models',
|
||||
'model.clearSearch': 'Clear model search',
|
||||
|
||||
@@ -2344,6 +2344,10 @@ export const jp: Record<TranslationKey, string> = {
|
||||
// ─── Model Selector ──────────────────────────────────────
|
||||
'model.selectModel': 'モデルを選択',
|
||||
'model.configureProvider': 'モデルプロバイダーを設定',
|
||||
'model.protocolLocked': 'このセッションは {protocol} を使用します。別の API プロトコルを使うには、新しいセッションを作成してください。',
|
||||
'model.protocolMixed': '複数の API プロトコルの履歴が混在しているため、このセッションは続行できません。モデルメニューから新しいセッションを作成してください。',
|
||||
'model.protocolUnknown': 'この古いセッションの API プロトコルを特定できません。モデルメニューから新しいセッションを作成してください。',
|
||||
'model.protocolRequiresNewSession': 'API プロトコルが異なります。新しいセッションを作成してください',
|
||||
'model.configuration': 'モデル設定',
|
||||
'model.searchPlaceholder': 'モデルを検索',
|
||||
'model.clearSearch': 'モデル検索をクリア',
|
||||
|
||||
@@ -2344,6 +2344,10 @@ export const kr: Record<TranslationKey, string> = {
|
||||
// ─── Model Selector ──────────────────────────────────────
|
||||
'model.selectModel': '모델 선택',
|
||||
'model.configureProvider': '모델 공급자 설정',
|
||||
'model.protocolLocked': '이 세션은 {protocol} 프로토콜을 사용합니다. 다른 API 프로토콜을 사용하려면 새 세션을 시작하세요.',
|
||||
'model.protocolMixed': '여러 API 프로토콜의 기록이 섞여 있어 이 세션을 계속할 수 없습니다. 모델 메뉴에서 새 세션을 시작하세요.',
|
||||
'model.protocolUnknown': '이전 세션의 API 프로토콜을 확인할 수 없습니다. 모델 메뉴에서 새 세션을 시작하세요.',
|
||||
'model.protocolRequiresNewSession': 'API 프로토콜이 다릅니다. 새 세션을 시작하세요',
|
||||
'model.configuration': '모델 구성',
|
||||
'model.searchPlaceholder': '모델 검색',
|
||||
'model.clearSearch': '모델 검색 지우기',
|
||||
|
||||
@@ -2343,6 +2343,10 @@ export const zh: Record<TranslationKey, string> = {
|
||||
// ─── Model Selector ──────────────────────────────────────
|
||||
'model.selectModel': '選擇模型',
|
||||
'model.configureProvider': '設定模型服務商',
|
||||
'model.protocolLocked': '此對話使用 {protocol} 協定。如需使用其他 API 協定,請建立新對話。',
|
||||
'model.protocolMixed': '此對話混用了多種 API 協定的歷史,無法繼續。請在模型選單中建立新對話。',
|
||||
'model.protocolUnknown': '無法判定此舊對話的 API 協定。請在模型選單中建立新對話。',
|
||||
'model.protocolRequiresNewSession': 'API 協定不同,請建立新對話',
|
||||
'model.configuration': '模型配置',
|
||||
'model.searchPlaceholder': '搜尋模型',
|
||||
'model.clearSearch': '清除模型搜尋',
|
||||
|
||||
@@ -2343,6 +2343,10 @@ export const zh: Record<TranslationKey, string> = {
|
||||
// ─── Model Selector ──────────────────────────────────────
|
||||
'model.selectModel': '选择模型',
|
||||
'model.configureProvider': '配置模型供应商',
|
||||
'model.protocolLocked': '此会话使用 {protocol} 协议。如需使用其他 API 协议,请新建会话。',
|
||||
'model.protocolMixed': '此会话混用了多种 API 协议的历史,无法继续。请在模型菜单中新建会话。',
|
||||
'model.protocolUnknown': '无法确定此旧会话的 API 协议。请在模型菜单中新建会话。',
|
||||
'model.protocolRequiresNewSession': 'API 协议不同,请新建会话',
|
||||
'model.configuration': '模型配置',
|
||||
'model.searchPlaceholder': '搜索模型',
|
||||
'model.clearSearch': '清除模型搜索',
|
||||
|
||||
@@ -29,6 +29,7 @@ const {
|
||||
updateSessionTitleMock,
|
||||
updateSessionMessageCountMock,
|
||||
updateSessionPermissionModeMock,
|
||||
updateSessionApiFormatMock,
|
||||
sessionStoreSnapshot,
|
||||
cliTaskStoreSnapshot,
|
||||
connectionStateHandlers,
|
||||
@@ -55,6 +56,7 @@ const {
|
||||
updateSessionTitleMock: vi.fn(),
|
||||
updateSessionMessageCountMock: vi.fn(),
|
||||
updateSessionPermissionModeMock: vi.fn(),
|
||||
updateSessionApiFormatMock: vi.fn(),
|
||||
sessionStoreSnapshot: {
|
||||
sessions: [] as Array<{
|
||||
id: string
|
||||
@@ -139,6 +141,7 @@ vi.mock('./sessionStore', () => ({
|
||||
updateSessionTitle: updateSessionTitleMock,
|
||||
updateSessionMessageCount: updateSessionMessageCountMock,
|
||||
updateSessionPermissionMode: updateSessionPermissionModeMock,
|
||||
updateSessionApiFormat: updateSessionApiFormatMock,
|
||||
}),
|
||||
},
|
||||
}))
|
||||
@@ -8880,6 +8883,22 @@ describe('chatStore history mapping', () => {
|
||||
expect(sendMock).not.toHaveBeenCalledWith(TEST_SESSION_ID, { type: 'prewarm_session' })
|
||||
})
|
||||
|
||||
it('hydrates the authoritative protocol before reconnecting and mirrors it to session metadata', () => {
|
||||
useChatStore.setState({ sessions: { [TEST_SESSION_ID]: makeSession() } })
|
||||
useChatStore.getState().handleServerMessage(TEST_SESSION_ID, {
|
||||
type: 'session_protocol', sessionApiFormat: 'openai_responses',
|
||||
})
|
||||
expect(useChatStore.getState().sessions[TEST_SESSION_ID]?.sessionApiFormat).toBe('openai_responses')
|
||||
expect(updateSessionApiFormatMock).toHaveBeenCalledWith(TEST_SESSION_ID, 'openai_responses')
|
||||
useChatStore.setState((state) => ({ sessions: {
|
||||
...state.sessions, [TEST_SESSION_ID]: { ...state.sessions[TEST_SESSION_ID]!, connectionState: 'disconnected' },
|
||||
} }))
|
||||
useChatStore.getState().connectToSession(TEST_SESSION_ID, { minimalBootstrap: true, prewarm: false })
|
||||
expect(useChatStore.getState().sessions[TEST_SESSION_ID]?.sessionApiFormat).toBe('openai_responses')
|
||||
useChatStore.getState().handleServerMessage(TEST_SESSION_ID, { type: 'session_protocol', sessionApiFormat: 'mixed' })
|
||||
expect(useChatStore.getState().sessions[TEST_SESSION_ID]?.sessionApiFormat).toBe('mixed')
|
||||
})
|
||||
|
||||
it('sends explicit runtime overrides over websocket', () => {
|
||||
useChatStore.getState().setSessionRuntime(TEST_SESSION_ID, {
|
||||
providerId: null,
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
import type { SessionProtocolState } from '../../../src/shared/sessionProtocol'
|
||||
import { create } from 'zustand'
|
||||
import { wsManager } from '../api/websocket'
|
||||
import { sessionsApi } from '../api/sessions'
|
||||
@@ -106,6 +107,7 @@ type PendingComputerUsePermission = {
|
||||
type PendingComputerUsePermissions = Record<string, PendingComputerUsePermission>
|
||||
|
||||
export type PerSessionState = {
|
||||
sessionApiFormat?: SessionProtocolState
|
||||
messages: UIMessage[]
|
||||
chatState: ChatState
|
||||
/**
|
||||
@@ -2596,6 +2598,7 @@ export const useChatStore = create<ChatStore>((set, get) => ({
|
||||
[sessionId]: {
|
||||
...createDefaultSessionState(),
|
||||
connectionState: 'connecting',
|
||||
sessionApiFormat: existing?.sessionApiFormat,
|
||||
connectionSnapshotReady: false,
|
||||
messages: existing?.messages ?? [],
|
||||
// A new connection lifecycle may have durable transcript rows that
|
||||
@@ -4337,6 +4340,11 @@ export const useChatStore = create<ChatStore>((set, get) => ({
|
||||
)
|
||||
break
|
||||
|
||||
case 'session_protocol':
|
||||
update(() => ({ sessionApiFormat: msg.sessionApiFormat }))
|
||||
useSessionStore.getState().updateSessionApiFormat(sessionId, msg.sessionApiFormat)
|
||||
break
|
||||
|
||||
case 'runtime_config_applied': {
|
||||
const selected = useSessionRuntimeStore.getState().selections[sessionId]
|
||||
const matchesCurrentSelection = Boolean(selected) &&
|
||||
|
||||
@@ -24,7 +24,10 @@ type SessionRuntimeStore = {
|
||||
setSelection: (key: string, selection: RuntimeSelection) => void
|
||||
clearSelection: (key: string) => void
|
||||
moveSelection: (fromKey: string, toKey: string) => void
|
||||
syncFromSessions: (sessions: SessionListItem[]) => void
|
||||
syncFromSessions: (
|
||||
sessions: SessionListItem[],
|
||||
expectedSelections?: Record<string, RuntimeSelection>,
|
||||
) => void
|
||||
}
|
||||
|
||||
function normalizeSelection(selection: RuntimeSelection): RuntimeSelection | null {
|
||||
@@ -132,10 +135,13 @@ export const useSessionRuntimeStore = create<SessionRuntimeStore>((set) => ({
|
||||
return { selections }
|
||||
}),
|
||||
|
||||
syncFromSessions: (sessions) =>
|
||||
syncFromSessions: (sessions, expectedSelections) =>
|
||||
set((state) => {
|
||||
let selections = state.selections
|
||||
for (const session of sessions) {
|
||||
// A list response can describe the runtime before a user switched
|
||||
// models. Never replay that snapshot over a selection made in flight.
|
||||
if (expectedSelections && state.selections[session.id] !== expectedSelections[session.id]) continue
|
||||
if (!session.runtimeModelId || session.runtimeProviderId === undefined) continue
|
||||
const selection = normalizeSelection({
|
||||
providerId: session.runtimeProviderId,
|
||||
|
||||
@@ -244,6 +244,53 @@ describe('sessionStore', () => {
|
||||
})
|
||||
})
|
||||
|
||||
it('keeps a model chosen while an older session refresh is in flight', async () => {
|
||||
const sessionId = 'session-runtime-switch'
|
||||
const oldSelection = { providerId: 'openai-official', modelId: 'gpt-5.6' }
|
||||
const nextSelection = { providerId: 'deepseek-provider', modelId: 'deepseek-v4.1' }
|
||||
useSessionRuntimeStore.getState().setSelection(sessionId, oldSelection)
|
||||
const response = {
|
||||
sessions: [{
|
||||
...makeSession(sessionId, '2026-09-09T00:00:00.000Z'),
|
||||
runtimeProviderId: oldSelection.providerId,
|
||||
runtimeModelId: oldSelection.modelId,
|
||||
}],
|
||||
total: 1,
|
||||
}
|
||||
const refresh = createDeferred<typeof response>()
|
||||
listMock.mockReturnValueOnce(refresh.promise)
|
||||
|
||||
const refreshing = useSessionStore.getState().fetchSessions()
|
||||
useSessionRuntimeStore.getState().setSelection(sessionId, nextSelection)
|
||||
refresh.resolve(response)
|
||||
await refreshing
|
||||
|
||||
expect(useSessionRuntimeStore.getState().selections[sessionId]).toEqual(nextSelection)
|
||||
expect(JSON.parse(localStorage.getItem('cc-haha-session-runtime')!)[sessionId])
|
||||
.toEqual(nextSelection)
|
||||
})
|
||||
|
||||
it('does not erase a confirmed protocol when an earlier unlocked refresh arrives', async () => {
|
||||
const session = makeSession('locked-during-refresh', '2026-09-09T00:00:00.000Z')
|
||||
useSessionStore.setState({ sessions: [session] })
|
||||
const response = { sessions: [session], total: 1 }
|
||||
const refresh = createDeferred<typeof response>()
|
||||
listMock.mockReturnValueOnce(refresh.promise)
|
||||
const refreshing = useSessionStore.getState().fetchSessions()
|
||||
useSessionStore.getState().updateSessionApiFormat(session.id, 'anthropic')
|
||||
refresh.resolve(response)
|
||||
await refreshing
|
||||
expect(useSessionStore.getState().sessions[0]?.sessionApiFormat).toBe('anthropic')
|
||||
})
|
||||
|
||||
it('preserves the confirmed protocol when hydrating an older session snapshot', () => {
|
||||
const session = { ...makeSession('locked-session', '2026-09-09T00:00:00.000Z'), sessionApiFormat: 'anthropic' as const }
|
||||
useSessionStore.setState({ sessions: [session] })
|
||||
useSessionStore.getState().updateSessionApiFormat(session.id, 'openai_responses')
|
||||
useSessionStore.getState().hydrateHistoricalSessions([session])
|
||||
expect(useSessionStore.getState().sessions[0]?.sessionApiFormat).toBe('openai_responses')
|
||||
})
|
||||
|
||||
it('updates a session message count without changing other metadata', () => {
|
||||
useSessionStore.setState({
|
||||
sessions: [makeSession('session-count-1', '2026-05-07T00:00:00.000Z', 'Working session')],
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
import type { SessionProtocolState } from '../../../src/shared/sessionProtocol'
|
||||
import { create } from 'zustand'
|
||||
import {
|
||||
sessionsApi,
|
||||
@@ -13,6 +14,7 @@ import { useSettingsStore } from './settingsStore'
|
||||
import { useTabStore } from './tabStore'
|
||||
import type { LocalIndexStatus, SessionListItem } from '../types/session'
|
||||
import type { PermissionMode } from '../types/settings'
|
||||
import type { RuntimeSelection } from '../types/runtime'
|
||||
import { isPlaceholderSessionTitle } from '../lib/sessionTitle'
|
||||
import { invalidateRecentProjectsCache } from '../lib/recentProjectsCache'
|
||||
|
||||
@@ -55,7 +57,10 @@ type SessionStore = {
|
||||
fetchSessions: (project?: string) => Promise<void>
|
||||
loadMoreProjectSessions: (projectRoot: string) => Promise<void>
|
||||
releaseProjectHistory: (projectRoot: string) => void
|
||||
hydrateHistoricalSessions: (sessions: SessionListItem[]) => SessionListItem[]
|
||||
hydrateHistoricalSessions: (
|
||||
sessions: SessionListItem[],
|
||||
expectedSelections?: Record<string, RuntimeSelection>,
|
||||
) => SessionListItem[]
|
||||
openHistoricalSession: (session: SessionListItem) => void
|
||||
createSession: (workDir?: string, options?: CreateSessionOptions) => Promise<string>
|
||||
branchSession: (
|
||||
@@ -75,6 +80,7 @@ type SessionStore = {
|
||||
updateSessionTitle: (id: string, title: string) => void
|
||||
updateSessionMessageCount: (id: string, messageCount: number) => void
|
||||
updateSessionPermissionMode: (id: string, mode: PermissionMode) => void
|
||||
updateSessionApiFormat: (id: string, sessionApiFormat: SessionProtocolState) => void
|
||||
setActiveSession: (id: string | null) => void
|
||||
}
|
||||
|
||||
@@ -105,13 +111,14 @@ export const useSessionStore = create<SessionStore>((set, get) => ({
|
||||
|
||||
fetchSessions: async (project?: string) => {
|
||||
const requestId = ++fetchSessionsRequestId
|
||||
const expectedSelections = useSessionRuntimeStore.getState().selections
|
||||
set({ isLoading: true, error: null, sessionListRequestId: requestId })
|
||||
try {
|
||||
const response = await sessionsApi.list(buildSessionListParams(project))
|
||||
if (requestId !== get().sessionListRequestId) return
|
||||
const raw = response.sessions
|
||||
const indexStatus = response.index ?? null
|
||||
useSessionRuntimeStore.getState().syncFromSessions(raw)
|
||||
useSessionRuntimeStore.getState().syncFromSessions(raw, expectedSelections)
|
||||
let syncedSessions: SessionListItem[] = []
|
||||
set((state) => {
|
||||
if (requestId !== state.sessionListRequestId) return state
|
||||
@@ -266,13 +273,13 @@ export const useSessionStore = create<SessionStore>((set, get) => ({
|
||||
})
|
||||
},
|
||||
|
||||
hydrateHistoricalSessions: (snapshots) => {
|
||||
hydrateHistoricalSessions: (snapshots, expectedSelections) => {
|
||||
if (snapshots.length === 0) return []
|
||||
const selected = reconcileSessionSnapshots(snapshots, get().sessions)
|
||||
const selectedIds = new Set(selected.map((session) => session.id))
|
||||
// Hydrate before activating the tab: connecting immediately applies its
|
||||
// runtime selection and the composer reads workspace/permission metadata.
|
||||
useSessionRuntimeStore.getState().syncFromSessions(selected)
|
||||
useSessionRuntimeStore.getState().syncFromSessions(selected, expectedSelections)
|
||||
set((state) => ({
|
||||
sessions: mergeSessionList([
|
||||
...selected,
|
||||
@@ -462,6 +469,14 @@ export const useSessionStore = create<SessionStore>((set, get) => ({
|
||||
}))
|
||||
},
|
||||
|
||||
updateSessionApiFormat: (id, sessionApiFormat) => {
|
||||
set((state) => ({
|
||||
sessions: state.sessions.map((session) =>
|
||||
session.id === id ? { ...session, sessionApiFormat } : session,
|
||||
),
|
||||
}))
|
||||
},
|
||||
|
||||
setActiveSession: (id) => set({ activeSessionId: id }),
|
||||
}))
|
||||
|
||||
@@ -566,7 +581,7 @@ function mergeSessionList(
|
||||
|
||||
for (const item of incoming) {
|
||||
const current = currentById.get(item.id)
|
||||
const candidate = preserveLocalTitle(current, item)
|
||||
const candidate = preserveLocalMetadata(current, item)
|
||||
const existing = byId.get(candidate.id)
|
||||
if (!existing || sessionModifiedTime(candidate) > sessionModifiedTime(existing)) {
|
||||
byId.set(candidate.id, candidate)
|
||||
@@ -598,11 +613,16 @@ function sessionModifiedTime(session: SessionListItem): number {
|
||||
return Number.isFinite(timestamp) ? timestamp : 0
|
||||
}
|
||||
|
||||
function preserveLocalTitle(
|
||||
function preserveLocalMetadata(
|
||||
current: SessionListItem | undefined,
|
||||
incoming: SessionListItem,
|
||||
): SessionListItem {
|
||||
if (!current) return incoming
|
||||
// Protocol locks cannot be removed by a session-list snapshot taken before
|
||||
// the first accepted turn. The live event is mirrored into this metadata.
|
||||
if (current.sessionApiFormat && incoming.sessionApiFormat === undefined) {
|
||||
incoming = { ...incoming, sessionApiFormat: current.sessionApiFormat }
|
||||
}
|
||||
if (isPlaceholderSessionTitle(incoming.title) && !isPlaceholderSessionTitle(current.title)) {
|
||||
return { ...incoming, title: current.title }
|
||||
}
|
||||
|
||||
@@ -511,6 +511,35 @@ describe('tabStore', () => {
|
||||
},
|
||||
)
|
||||
|
||||
it.each([false, true])('preserves a runtime chosen during tab restoration (historical: %s)', async (historical) => {
|
||||
const session = historicalSummary('runtime-restore-race')
|
||||
const nextSelection = { providerId: 'deepseek-provider', modelId: 'deepseek-v4.1' }
|
||||
localStorage.setItem('cc-haha-open-tabs', JSON.stringify({
|
||||
openTabs: [{ sessionId: session.id, title: session.title, type: 'session' }],
|
||||
activeTabId: session.id,
|
||||
}))
|
||||
let resolveResponse!: () => void
|
||||
if (historical) {
|
||||
vi.mocked(sessionsApi.list).mockResolvedValueOnce({ sessions: [], total: 1 })
|
||||
vi.mocked(sessionsApi.getSummary).mockReturnValueOnce(new Promise((resolve) => {
|
||||
resolveResponse = () => resolve(session)
|
||||
}))
|
||||
} else {
|
||||
vi.mocked(sessionsApi.list).mockReturnValueOnce(new Promise((resolve) => {
|
||||
resolveResponse = () => resolve({ sessions: [session], total: 1 })
|
||||
}))
|
||||
}
|
||||
|
||||
const restoring = useTabStore.getState().restoreTabs()
|
||||
if (historical) await vi.waitFor(() => expect(sessionsApi.getSummary).toHaveBeenCalledOnce())
|
||||
useSessionRuntimeStore.getState().setSelection(session.id, nextSelection)
|
||||
resolveResponse()
|
||||
await restoring
|
||||
|
||||
expect(useSessionRuntimeStore.getState().selections[session.id]).toEqual(nextSelection)
|
||||
expect(useTabStore.getState().activeTabId).toBe(session.id)
|
||||
})
|
||||
|
||||
it('hydrates restored tabs with authoritative transcript runtime metadata', async () => {
|
||||
useSessionRuntimeStore.getState().setSelection('session-1', {
|
||||
providerId: null,
|
||||
|
||||
@@ -474,6 +474,7 @@ export const useTabStore = create<TabStore>((set, get) => ({
|
||||
restoreTabs: async () => {
|
||||
try {
|
||||
const restoreStartedWith = get()
|
||||
const expectedSelections = useSessionRuntimeStore.getState().selections
|
||||
const restoreStillCurrent = () => {
|
||||
const current = get()
|
||||
return current.tabs === restoreStartedWith.tabs &&
|
||||
@@ -522,10 +523,10 @@ export const useTabStore = create<TabStore>((set, get) => ({
|
||||
const recentSessions = reconcileSessionSnapshots(sessions, useSessionStore.getState().sessions)
|
||||
for (const session of recentSessions) sessionsById.set(session.id, session)
|
||||
if (historicalSessions.length > 0) {
|
||||
const hydrated = useSessionStore.getState().hydrateHistoricalSessions(historicalSessions)
|
||||
const hydrated = useSessionStore.getState().hydrateHistoricalSessions(historicalSessions, expectedSelections)
|
||||
for (const session of hydrated) sessionsById.set(session.id, session)
|
||||
}
|
||||
useSessionRuntimeStore.getState().syncFromSessions(recentSessions)
|
||||
useSessionRuntimeStore.getState().syncFromSessions(recentSessions, expectedSelections)
|
||||
|
||||
const validTabs: Tab[] = data.openTabs
|
||||
.filter((t) => {
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
import type { SessionProtocolState } from '../../../src/shared/sessionProtocol'
|
||||
import type { PermissionMode } from './settings'
|
||||
import type { RuntimeSelection } from './runtime'
|
||||
|
||||
@@ -83,6 +84,7 @@ export type UIAttachment = {
|
||||
|
||||
export type ServerMessage =
|
||||
| { type: 'connected'; sessionId: string }
|
||||
| { type: 'session_protocol'; sessionApiFormat: SessionProtocolState }
|
||||
| {
|
||||
type: 'session_state'
|
||||
turnState: 'running' | 'idle'
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
// Source: src/server/services/sessionService.ts
|
||||
|
||||
import type { SessionProtocolState } from '../../../src/shared/sessionProtocol'
|
||||
import type { ReasoningEffortLevel } from './settings'
|
||||
|
||||
export type LocalIndexMode = 'off' | 'shadow' | 'on'
|
||||
@@ -29,6 +30,7 @@ export type SessionListItem = {
|
||||
workDirExists: boolean
|
||||
workspaceState?: SessionWorkspaceState
|
||||
permissionMode?: string
|
||||
sessionApiFormat?: SessionProtocolState
|
||||
runtimeProviderId?: string | null
|
||||
runtimeModelId?: string
|
||||
effortLevel?: ReasoningEffortLevel
|
||||
|
||||
@@ -12,6 +12,16 @@ const checks: Check[] = [
|
||||
title: 'Server persistent JSON migrations',
|
||||
command: ['bun', 'test', './src/server/__tests__/persistence-upgrade.test.ts'],
|
||||
},
|
||||
{
|
||||
title: 'Session protocol history and local index migrations',
|
||||
command: [
|
||||
'bun', 'test',
|
||||
'./src/server/__tests__/session-protocol.test.ts',
|
||||
'./src/server/services/sessionProtocolHistory.test.ts',
|
||||
'./src/server/services/localIndex/database.test.ts',
|
||||
'./src/server/services/localIndex/sessionProjector.test.ts',
|
||||
],
|
||||
},
|
||||
{
|
||||
title: 'Desktop UI preference migrations',
|
||||
command: [
|
||||
|
||||
@@ -72,6 +72,9 @@ describe('ConversationService attachment materialization', () => {
|
||||
},
|
||||
},
|
||||
pendingOutbound: [],
|
||||
workDir: tmpDir,
|
||||
providerId: null,
|
||||
apiFormat: 'anthropic',
|
||||
})
|
||||
|
||||
const ok = await svc.sendMessage(sessionId, '这张图说了什么?', [
|
||||
@@ -119,6 +122,9 @@ describe('ConversationService attachment materialization', () => {
|
||||
},
|
||||
},
|
||||
pendingOutbound: [],
|
||||
workDir: tmpDir,
|
||||
providerId: null,
|
||||
apiFormat: 'anthropic',
|
||||
})
|
||||
|
||||
const ok = await svc.sendMessage(sessionId, '', [
|
||||
@@ -155,6 +161,9 @@ describe('ConversationService attachment materialization', () => {
|
||||
},
|
||||
},
|
||||
pendingOutbound: [],
|
||||
workDir: tmpDir,
|
||||
providerId: null,
|
||||
apiFormat: 'anthropic',
|
||||
})
|
||||
|
||||
const ok = await svc.sendMessage(sessionId, '看这个截图', [
|
||||
|
||||
@@ -187,6 +187,9 @@ describe('ConversationService', () => {
|
||||
) {
|
||||
const session = {
|
||||
outputCallbacks: [],
|
||||
workDir: tmpDir,
|
||||
providerId: null,
|
||||
apiFormat: 'anthropic',
|
||||
networkRoutingFingerprint: '',
|
||||
networkDerivedFirstTokenTimeout,
|
||||
sdkSocket: {
|
||||
@@ -525,6 +528,8 @@ describe('ConversationService', () => {
|
||||
const sent: string[] = []
|
||||
service.sessions.set('sleep-wake-session', {
|
||||
proc: {},
|
||||
providerId: null,
|
||||
apiFormat: 'anthropic',
|
||||
outputCallbacks: [],
|
||||
workDir: tmpDir,
|
||||
permissionMode: 'default',
|
||||
|
||||
@@ -403,6 +403,7 @@ describe('SessionService local-index routing parity', () => {
|
||||
projectPath: fileSession.projectPath,
|
||||
workDir: fileSession.workDir,
|
||||
runtimeProviderId: null,
|
||||
sessionApiFormat: fileSession.sessionApiFormat,
|
||||
}
|
||||
gateway.page = { sessions: [explicitNullRow], total: 1 }
|
||||
await service.listSessions()
|
||||
|
||||
@@ -0,0 +1,10 @@
|
||||
import { expect, test } from 'bun:test'
|
||||
import { errorResponse } from '../middleware/errorHandler.js'
|
||||
import { SessionProtocolError } from '../services/sessionProtocolHistory.js'
|
||||
|
||||
test.each(['anthropic', 'mixed', 'unknown'] as const)('API reports protocol conflict for %s without a generic server failure', async current => {
|
||||
const error = new SessionProtocolError(current, 'openai_responses')
|
||||
const response = errorResponse(error)
|
||||
expect(response.status).toBe(409)
|
||||
expect(await response.json()).toEqual({ error: error.code, message: error.message })
|
||||
})
|
||||
@@ -0,0 +1,285 @@
|
||||
import { afterAll, afterEach, beforeAll, describe, expect, it, spyOn } from 'bun:test'
|
||||
import { appendFile, mkdir, readFile } from 'node:fs/promises'
|
||||
import { join } from 'node:path'
|
||||
import { fileURLToPath } from 'node:url'
|
||||
import { createQualityGateSandbox, type QualityGateSandbox } from '../../../scripts/quality-gate/sandbox.js'
|
||||
import { createOfflineTestEnvironment } from '../../../scripts/pr/test-environment.js'
|
||||
import type { SessionApiFormat } from '../../shared/sessionProtocol.js'
|
||||
import { conversationService } from '../services/conversationService.js'
|
||||
import { ProviderService } from '../services/providerService.js'
|
||||
import { SessionService, sessionService } from '../services/sessionService.js'
|
||||
import { resetTerminalShellEnvironmentCacheForTests } from '../../utils/terminalShellEnvironment.js'
|
||||
|
||||
type Event = { type: string; [key: string]: any }
|
||||
type Client = {
|
||||
socket: WebSocket
|
||||
events: Event[]
|
||||
send(message: Record<string, unknown>): void
|
||||
wait(predicate: (event: Event) => boolean, after?: number): Promise<Event>
|
||||
}
|
||||
|
||||
describe('session protocol routing over WebSocket', () => {
|
||||
const originalEnv = { ...process.env }
|
||||
const sockets = new Set<WebSocket>()
|
||||
const providerService = new ProviderService()
|
||||
let sandbox: QualityGateSandbox
|
||||
let server: ReturnType<typeof Bun.serve>
|
||||
let baseUrl: string
|
||||
let workDir: string
|
||||
|
||||
beforeAll(async () => {
|
||||
sandbox = createQualityGateSandbox({
|
||||
label: 'session-protocol',
|
||||
seedProviders: false,
|
||||
// Do not inherit proxy/provider credentials or access the login shell.
|
||||
source: createOfflineTestEnvironment({}, originalEnv),
|
||||
sourceConfigDir: originalEnv.CLAUDE_CONFIG_DIR,
|
||||
envOverrides: {
|
||||
NODE_ENV: 'test',
|
||||
CLAUDE_CLI_PATH: fileURLToPath(new URL('./fixtures/mock-sdk-cli.ts', import.meta.url)),
|
||||
},
|
||||
})
|
||||
for (const name of Object.keys(process.env)) delete process.env[name]
|
||||
Object.assign(process.env, sandbox.env)
|
||||
resetTerminalShellEnvironmentCacheForTests()
|
||||
workDir = join(sandbox.home, 'workspace')
|
||||
await mkdir(workDir, { recursive: true })
|
||||
await mkdir(join(sandbox.configDir, 'projects'), { recursive: true })
|
||||
const { startServer } = await import('../index.js')
|
||||
server = startServer(0, '127.0.0.1')
|
||||
baseUrl = `http://127.0.0.1:${server.port}`
|
||||
})
|
||||
|
||||
afterEach(async () => {
|
||||
for (const socket of sockets) socket.close()
|
||||
sockets.clear()
|
||||
await conversationService.stopAllSessionsAndWait(1_000)
|
||||
})
|
||||
|
||||
afterAll(async () => {
|
||||
try {
|
||||
server?.stop(true)
|
||||
const { stopServerRuntimeForShutdown } = await import('../index.js')
|
||||
await stopServerRuntimeForShutdown()
|
||||
expect(sandbox.detectUserStateMutations()).toEqual([])
|
||||
} finally {
|
||||
sandbox?.cleanup()
|
||||
for (const name of Object.keys(process.env)) delete process.env[name]
|
||||
Object.assign(process.env, originalEnv)
|
||||
resetTerminalShellEnvironmentCacheForTests()
|
||||
}
|
||||
})
|
||||
|
||||
async function addProvider(apiFormat: SessionApiFormat) {
|
||||
return providerService.addProvider({
|
||||
presetId: 'custom',
|
||||
name: `Protocol ${apiFormat} ${crypto.randomUUID()}`,
|
||||
apiFormat,
|
||||
apiKey: 'fixture-protocol-key',
|
||||
baseUrl: 'http://127.0.0.1:1',
|
||||
// Identical model names prove that protocol selection uses provider routing.
|
||||
models: { main: 'fixture-main', haiku: 'fixture-small', sonnet: 'fixture-main', opus: 'fixture-large' },
|
||||
})
|
||||
}
|
||||
|
||||
async function createSession(): Promise<string> {
|
||||
const response = await fetch(`${baseUrl}/api/sessions`, {
|
||||
method: 'POST',
|
||||
headers: { 'Content-Type': 'application/json' },
|
||||
body: JSON.stringify({ workDir }),
|
||||
})
|
||||
expect(response.status).toBe(201)
|
||||
const body = await response.json() as { sessionId: string }
|
||||
return body.sessionId
|
||||
}
|
||||
|
||||
async function connect(sessionId: string): Promise<Client> {
|
||||
const socket = new WebSocket(`ws://127.0.0.1:${server.port}/ws/${sessionId}`)
|
||||
sockets.add(socket)
|
||||
const events: Event[] = []
|
||||
const listeners = new Set<() => void>()
|
||||
let failed = false
|
||||
socket.onmessage = event => {
|
||||
events.push(JSON.parse(event.data as string))
|
||||
for (const listener of listeners) listener()
|
||||
}
|
||||
socket.onerror = () => {
|
||||
failed = true
|
||||
for (const listener of listeners) listener()
|
||||
}
|
||||
const client: Client = {
|
||||
socket,
|
||||
events,
|
||||
send(message) { socket.send(JSON.stringify(message)) },
|
||||
wait(predicate, after = 0) {
|
||||
return new Promise((resolve, reject) => {
|
||||
const timer = setTimeout(() => {
|
||||
listeners.delete(check)
|
||||
reject(new Error(`Timed out waiting for protocol event: ${JSON.stringify(events.slice(after))}`))
|
||||
}, 8_000)
|
||||
function check() {
|
||||
const event = events.slice(after).find(predicate)
|
||||
if (!event && !failed) return
|
||||
clearTimeout(timer)
|
||||
listeners.delete(check)
|
||||
if (failed) reject(new Error('Protocol fixture WebSocket failed'))
|
||||
else resolve(event!)
|
||||
}
|
||||
listeners.add(check)
|
||||
check()
|
||||
})
|
||||
},
|
||||
}
|
||||
await client.wait(event => event.type === 'connected')
|
||||
return client
|
||||
}
|
||||
|
||||
async function select(client: Client, providerId: string, modelId = 'fixture-main') {
|
||||
const after = client.events.length
|
||||
client.send({ type: 'set_runtime_config', providerId, modelId })
|
||||
const result = await client.wait(event => event.type === 'error' || (
|
||||
event.type === 'runtime_config_applied' && event.providerId === providerId && event.modelId === modelId
|
||||
), after)
|
||||
expect(result.type).toBe('runtime_config_applied')
|
||||
}
|
||||
|
||||
async function sendTurn(client: Client, content = 'hello fixture') {
|
||||
const after = client.events.length
|
||||
client.send({ type: 'user_message', content })
|
||||
const result = await client.wait(event => event.type === 'message_complete' || event.type === 'error', after)
|
||||
expect(result.type).toBe('message_complete')
|
||||
return after
|
||||
}
|
||||
|
||||
for (const apiFormat of ['anthropic', 'openai_chat', 'openai_responses'] as const) {
|
||||
it(`locks ${apiFormat} only on first send and rejects other protocols before restart or persistence`, async () => {
|
||||
const providers = []
|
||||
for (const format of ['anthropic', 'openai_chat', 'openai_responses'] as const) {
|
||||
providers.push(await addProvider(format))
|
||||
}
|
||||
const selected = providers.find(provider => provider.apiFormat === apiFormat)!
|
||||
const others = providers.filter(provider => provider.id !== selected.id)
|
||||
const sessionId = await createSession()
|
||||
const client = await connect(sessionId)
|
||||
// Changing choices (including protocol) before sending must remain possible.
|
||||
await select(client, others[0]!.id)
|
||||
expect(await sessionService.getSessionApiFormat(sessionId)).toBeUndefined()
|
||||
await select(client, selected.id)
|
||||
expect(await sessionService.getSessionApiFormat(sessionId)).toBeUndefined()
|
||||
const after = await sendTurn(client)
|
||||
await client.wait(event => event.type === 'session_protocol' && event.sessionApiFormat === apiFormat, after)
|
||||
expect(await new SessionService().getSessionApiFormat(sessionId)).toBe(apiFormat)
|
||||
const before = await sessionService.getSessionLaunchInfo(sessionId)
|
||||
const start = spyOn(conversationService, 'startSession')
|
||||
const stop = spyOn(conversationService, 'stopSession')
|
||||
try {
|
||||
for (const other of others) {
|
||||
const index = client.events.length
|
||||
client.send({ type: 'set_runtime_config', providerId: other.id, modelId: 'fixture-main' })
|
||||
const error = await client.wait(event => event.type === 'error', index)
|
||||
expect(error.code).toBe('SESSION_PROTOCOL_MISMATCH')
|
||||
await client.wait(event => event.type === 'runtime_config_applied' && event.providerId === selected.id, index)
|
||||
expect(await sessionService.getSessionLaunchInfo(sessionId)).toMatchObject({
|
||||
sessionApiFormat: apiFormat,
|
||||
runtimeProviderId: before!.runtimeProviderId,
|
||||
runtimeModelId: before!.runtimeModelId,
|
||||
})
|
||||
}
|
||||
expect(start).not.toHaveBeenCalled()
|
||||
expect(stop).not.toHaveBeenCalled()
|
||||
// Rejection must leave the original session usable.
|
||||
await sendTurn(client, 'continue on the original protocol')
|
||||
} finally {
|
||||
start.mockRestore()
|
||||
stop.mockRestore()
|
||||
}
|
||||
}, 25_000)
|
||||
}
|
||||
|
||||
it('allows another provider and model on the locked protocol', async () => {
|
||||
const first = await addProvider('openai_chat')
|
||||
const second = await addProvider('openai_chat')
|
||||
const sessionId = await createSession()
|
||||
const client = await connect(sessionId)
|
||||
await select(client, first.id)
|
||||
await sendTurn(client)
|
||||
await select(client, second.id, 'fixture-large')
|
||||
await sendTurn(client, 'continue with a different model')
|
||||
expect(await sessionService.getSessionLaunchInfo(sessionId)).toMatchObject({
|
||||
sessionApiFormat: 'openai_chat',
|
||||
runtimeProviderId: second.id,
|
||||
runtimeModelId: 'fixture-large',
|
||||
})
|
||||
}, 20_000)
|
||||
|
||||
it('restores the durable protocol in the API and on reconnect after the CLI stops', async () => {
|
||||
const provider = await addProvider('openai_responses')
|
||||
const sessionId = await createSession()
|
||||
const client = await connect(sessionId)
|
||||
await select(client, provider.id)
|
||||
await sendTurn(client)
|
||||
client.socket.close()
|
||||
await conversationService.stopAllSessionsAndWait(1_000)
|
||||
const response = await fetch(`${baseUrl}/api/sessions/${sessionId}`)
|
||||
expect(response.status).toBe(200)
|
||||
expect(await response.json()).toMatchObject({ sessionApiFormat: 'openai_responses' })
|
||||
expect(await new SessionService().getSessionApiFormat(sessionId)).toBe('openai_responses')
|
||||
const reconnected = await connect(sessionId)
|
||||
await reconnected.wait(event => event.type === 'session_protocol' && event.sessionApiFormat === 'openai_responses')
|
||||
await sendTurn(reconnected, 'continue after CLI restart')
|
||||
expect(await sessionService.getSessionApiFormat(sessionId)).toBe('openai_responses')
|
||||
}, 20_000)
|
||||
|
||||
it('rejects sending when a saved provider changes protocol under an active CLI', async () => {
|
||||
const provider = await addProvider('anthropic')
|
||||
const sessionId = await createSession()
|
||||
const client = await connect(sessionId)
|
||||
await select(client, provider.id)
|
||||
await sendTurn(client)
|
||||
await providerService.updateProvider(provider.id, { apiFormat: 'openai_responses' })
|
||||
const before = conversationService.getRecentSdkMessages(sessionId).length
|
||||
const after = client.events.length
|
||||
client.send({ type: 'user_message', content: 'must not reach the SDK' })
|
||||
const error = await client.wait(event => event.type === 'error', after)
|
||||
expect(error.code).toBe('SESSION_PROTOCOL_MISMATCH')
|
||||
await client.wait(event => event.type === 'status' && event.state === 'idle', after)
|
||||
expect(await sessionService.getSessionApiFormat(sessionId)).toBe('anthropic')
|
||||
expect(conversationService.getRecentSdkMessages(sessionId).slice(before)
|
||||
.some(event => event.type === 'assistant' || event.type === 'result')).toBe(false)
|
||||
// Restoring the provider makes the original route usable without a stuck turn.
|
||||
await providerService.updateProvider(provider.id, { apiFormat: 'anthropic' })
|
||||
await sendTurn(client, 'continue after restoring the provider')
|
||||
}, 20_000)
|
||||
|
||||
for (const state of ['unknown', 'mixed'] as const) {
|
||||
it(`blocks old ${state} history without mutating it or starting a CLI`, async () => {
|
||||
const provider = await addProvider('anthropic')
|
||||
const sessionId = await createSession()
|
||||
const launch = await sessionService.getSessionLaunchInfo(sessionId)
|
||||
const oldEntries = state === 'unknown'
|
||||
? [{ type: 'assistant', message: { role: 'assistant', model: 'legacy-model', content: 'old reply' } }]
|
||||
: [
|
||||
{ type: 'session-meta', runtimeProviderId: null },
|
||||
{ type: 'assistant', message: { role: 'assistant', content: 'Messages reply' } },
|
||||
{ type: 'session-meta', runtimeProviderId: 'openai-official' },
|
||||
{ type: 'assistant', message: { role: 'assistant', content: 'Responses reply' } },
|
||||
]
|
||||
await appendFile(launch!.filePath, oldEntries.map(entry => JSON.stringify(entry)).join('\n') + '\n')
|
||||
const original = await readFile(launch!.filePath, 'utf8')
|
||||
const client = await connect(sessionId)
|
||||
await client.wait(event => event.type === 'session_protocol' && event.sessionApiFormat === state)
|
||||
const after = client.events.length
|
||||
client.send({ type: 'set_runtime_config', providerId: provider.id, modelId: 'fixture-main' })
|
||||
expect(await client.wait(event => event.type === 'error', after)).toMatchObject({
|
||||
code: 'SESSION_PROTOCOL_UNRESOLVED', retryable: false,
|
||||
})
|
||||
const sendIndex = client.events.length
|
||||
client.send({ type: 'user_message', content: 'must not replay unresolved history' })
|
||||
expect(await client.wait(event => event.type === 'error', sendIndex)).toMatchObject({
|
||||
code: 'SESSION_PROTOCOL_UNRESOLVED', retryable: false,
|
||||
})
|
||||
expect(conversationService.hasSession(sessionId)).toBe(false)
|
||||
expect(await readFile(launch!.filePath, 'utf8')).toBe(original)
|
||||
}, 15_000)
|
||||
}
|
||||
})
|
||||
@@ -0,0 +1,125 @@
|
||||
import { afterEach, beforeEach, describe, expect, it } from 'bun:test'
|
||||
import * as fs from 'node:fs/promises'
|
||||
import * as os from 'node:os'
|
||||
import * as path from 'node:path'
|
||||
import { SessionService } from '../services/sessionService.js'
|
||||
import { SessionProtocolError } from '../services/sessionProtocolHistory.js'
|
||||
import { sanitizePath } from '../../utils/sessionStoragePortable.js'
|
||||
|
||||
const sessionId = '12870000-bbbb-cccc-dddd-eeeeeeeeeeee'
|
||||
let configDir: string
|
||||
let previousConfigDir: string | undefined
|
||||
let previousIndexMode: string | undefined
|
||||
let workDir: string
|
||||
let filePath: string
|
||||
|
||||
beforeEach(async () => {
|
||||
previousConfigDir = process.env.CLAUDE_CONFIG_DIR
|
||||
previousIndexMode = process.env.CC_HAHA_LOCAL_INDEX
|
||||
configDir = await fs.mkdtemp(path.join(os.tmpdir(), 'session-protocol-'))
|
||||
process.env.CLAUDE_CONFIG_DIR = configDir
|
||||
process.env.CC_HAHA_LOCAL_INDEX = 'off'
|
||||
workDir = path.join(configDir, 'repo')
|
||||
await fs.mkdir(workDir)
|
||||
filePath = path.join(configDir, 'projects', sanitizePath(workDir), `${sessionId}.jsonl`)
|
||||
await fs.mkdir(path.dirname(filePath), { recursive: true })
|
||||
})
|
||||
|
||||
afterEach(async () => {
|
||||
if (previousConfigDir === undefined) delete process.env.CLAUDE_CONFIG_DIR
|
||||
else process.env.CLAUDE_CONFIG_DIR = previousConfigDir
|
||||
if (previousIndexMode === undefined) delete process.env.CC_HAHA_LOCAL_INDEX
|
||||
else process.env.CC_HAHA_LOCAL_INDEX = previousIndexMode
|
||||
await fs.rm(configDir, { recursive: true, force: true })
|
||||
})
|
||||
|
||||
async function seed(entries: object[]) {
|
||||
const original = [
|
||||
{ type: 'session-meta', workDir, customFutureField: { keep: true } }, ...entries,
|
||||
].map(entry => JSON.stringify(entry)).join('\n') + '\n'
|
||||
await fs.writeFile(filePath, original)
|
||||
return original
|
||||
}
|
||||
const user = { type: 'user', uuid: 'user-1', message: { role: 'user', content: 'hello' } }
|
||||
const assistant = { type: 'assistant', uuid: 'assistant-1', message: { role: 'assistant', model: 'fixture-model', content: 'reply' } }
|
||||
|
||||
describe('session API protocol persistence', () => {
|
||||
it('locks on first send, survives restart, rejects a different protocol without writes', async () => {
|
||||
const original = await seed([])
|
||||
const service = new SessionService()
|
||||
expect(await service.getSessionApiFormat(sessionId)).toBeUndefined()
|
||||
await service.lockSessionApiFormat(sessionId, 'openai_chat')
|
||||
const locked = await fs.readFile(filePath, 'utf8')
|
||||
expect(locked.startsWith(original)).toBe(true)
|
||||
const restarted = new SessionService()
|
||||
expect(await restarted.getSessionApiFormat(sessionId)).toBe('openai_chat')
|
||||
await restarted.lockSessionApiFormat(sessionId, 'openai_chat')
|
||||
await expect(restarted.lockSessionApiFormat(sessionId, 'anthropic')).rejects.toBeInstanceOf(SessionProtocolError)
|
||||
expect(await fs.readFile(filePath, 'utf8')).toBe(locked)
|
||||
})
|
||||
|
||||
it('creates protocol metadata for a live SDK session before its first transcript write', async () => {
|
||||
const service = new SessionService()
|
||||
await service.lockSessionApiFormat('live-sdk-session', 'openai_responses', workDir)
|
||||
expect(await service.getSessionApiFormat('live-sdk-session')).toBe('openai_responses')
|
||||
await expect(service.lockSessionApiFormat('live-sdk-session', 'anthropic', workDir)).rejects.toBeInstanceOf(SessionProtocolError)
|
||||
await expect(service.lockSessionApiFormat('../outside', 'anthropic', workDir)).rejects.toThrow('Invalid session ID')
|
||||
await expect(service.lockSessionApiFormat(sessionId, 'anthropic')).rejects.toThrow('Session not found')
|
||||
})
|
||||
|
||||
it('serializes competing first sends across service instances', async () => {
|
||||
await seed([])
|
||||
const results = await Promise.allSettled([
|
||||
new SessionService().lockSessionApiFormat(sessionId, 'anthropic'),
|
||||
new SessionService().lockSessionApiFormat(sessionId, 'openai_responses'),
|
||||
])
|
||||
expect(results.map(result => result.status)).toEqual(['fulfilled', 'rejected'])
|
||||
expect(await new SessionService().getSessionApiFormat(sessionId)).toBe('anthropic')
|
||||
expect((await fs.readFile(filePath, 'utf8')).match(/sessionApiFormat/g)).toHaveLength(1)
|
||||
})
|
||||
|
||||
it('upgrades old official-route history additively only on send and exposes every read surface', async () => {
|
||||
const original = await seed([{ type: 'session-meta', runtimeProviderId: 'openai-official' }, user, assistant])
|
||||
const service = new SessionService()
|
||||
expect(await service.getSessionApiFormat(sessionId)).toBe('openai_responses')
|
||||
expect((await service.getSessionLaunchInfo(sessionId))?.sessionApiFormat).toBe('openai_responses')
|
||||
expect((await service.getSession(sessionId))?.sessionApiFormat).toBe('openai_responses')
|
||||
expect((await service.getSessionSummary(sessionId))?.sessionApiFormat).toBe('openai_responses')
|
||||
expect((await service.listSessions()).sessions[0]?.sessionApiFormat).toBe('openai_responses')
|
||||
expect((await service.getInspectionTranscriptSnapshot(sessionId))?.launchInfo.sessionApiFormat).toBe('openai_responses')
|
||||
expect(await fs.readFile(filePath, 'utf8')).toBe(original)
|
||||
await service.lockSessionApiFormat(sessionId, 'openai_responses')
|
||||
expect((await fs.readFile(filePath, 'utf8')).startsWith(original)).toBe(true)
|
||||
})
|
||||
|
||||
it.each(['unknown', 'mixed'] as const)('blocks %s old histories without rewriting them', async state => {
|
||||
const entries = state === 'mixed'
|
||||
? [{ type: 'session-meta', runtimeProviderId: null }, user, assistant,
|
||||
{ type: 'session-meta', runtimeProviderId: 'openai-official' }, user, assistant]
|
||||
: [{ type: 'session-meta', runtimeProviderId: 'mutable-saved-provider' }, user, assistant]
|
||||
const original = await seed(entries)
|
||||
const service = new SessionService()
|
||||
expect(await service.getSessionApiFormat(sessionId)).toBe(state)
|
||||
await expect(service.lockSessionApiFormat(sessionId, 'anthropic')).rejects.toMatchObject({ code: 'SESSION_PROTOCOL_UNRESOLVED' })
|
||||
expect(await fs.readFile(filePath, 'utf8')).toBe(original)
|
||||
})
|
||||
|
||||
it('preserves protocol when runtime metadata moves to another workspace', async () => {
|
||||
await seed([{ type: 'session-meta', sessionApiFormat: 'openai_chat' }, user, assistant])
|
||||
const destination = path.join(configDir, 'second-repo')
|
||||
await fs.mkdir(destination)
|
||||
const service = new SessionService()
|
||||
await service.appendSessionMetadata(sessionId, { workDir: destination, runtimeProviderId: 'different-saved-provider' })
|
||||
expect(await service.getSessionApiFormat(sessionId)).toBe('openai_chat')
|
||||
expect((await service.getSessionLaunchInfo(sessionId))?.sessionApiFormat).toBe('openai_chat')
|
||||
})
|
||||
|
||||
it('keeps the lock when clearing or rewinding the same session', async () => {
|
||||
await seed([{ type: 'session-meta', sessionApiFormat: 'anthropic' }, user, assistant])
|
||||
const service = new SessionService()
|
||||
await service.trimSessionMessagesFrom(sessionId, 'user-1')
|
||||
expect(await service.getSessionApiFormat(sessionId)).toBe('anthropic')
|
||||
await service.clearSessionTranscript(sessionId, workDir)
|
||||
expect(await new SessionService().getSessionApiFormat(sessionId)).toBe('anthropic')
|
||||
})
|
||||
})
|
||||
@@ -853,6 +853,8 @@ describe('SessionService', () => {
|
||||
|
||||
expect(scanned).toEqual(reduced.summary)
|
||||
expect(scanned).toEqual({
|
||||
// The completed legacy messages predate any provable protocol snapshot.
|
||||
sessionApiFormat: 'unknown',
|
||||
title: 'Canonical parity title',
|
||||
createdAt: '2026-07-01T01:00:00.000Z',
|
||||
modifiedAt: '2026-07-01T02:05:00.000Z',
|
||||
|
||||
@@ -3,6 +3,7 @@
|
||||
*/
|
||||
|
||||
import { diagnosticsService } from '../services/diagnosticsService.js'
|
||||
import { SessionProtocolError } from '../services/sessionProtocolHistory.js'
|
||||
|
||||
export class ApiError extends Error {
|
||||
constructor(
|
||||
@@ -32,7 +33,7 @@ export class ApiError extends Error {
|
||||
}
|
||||
|
||||
export function errorResponse(error: unknown): Response {
|
||||
if (error instanceof ApiError) {
|
||||
if (error instanceof ApiError || error instanceof SessionProtocolError) {
|
||||
return Response.json(
|
||||
{ error: error.code || 'ERROR', message: error.message },
|
||||
{ status: error.statusCode }
|
||||
|
||||
@@ -31,6 +31,8 @@ import {
|
||||
IMAGE_GENERATION_PROVIDER_KIND_ENV_KEY,
|
||||
} from '../../services/imageGeneration/config.js'
|
||||
import { sessionService } from './sessionService.js'
|
||||
import { assertSessionApiFormat } from './sessionProtocolHistory.js'
|
||||
import { resolveProviderApiFormat, type SessionApiFormat } from '../../shared/sessionProtocol.js'
|
||||
import { diagnosticsService } from './diagnosticsService.js'
|
||||
import {
|
||||
isMaterializedWorktreeLaunch,
|
||||
@@ -163,6 +165,7 @@ type SendMessageOptions = {
|
||||
canSend?: () => boolean
|
||||
messageUuid?: string
|
||||
onCommitted?: () => void
|
||||
onProtocolLocked?: (apiFormat: SessionApiFormat) => void
|
||||
}
|
||||
|
||||
type HandleSdkPayloadOptions = {
|
||||
@@ -198,6 +201,8 @@ type SessionProcess = {
|
||||
outputCallbacks: SessionOutputCallback[]
|
||||
workDir: string
|
||||
permissionMode: string
|
||||
providerId: string | null
|
||||
apiFormat: SessionApiFormat
|
||||
networkRoutingFingerprint: string
|
||||
networkDerivedFirstTokenTimeout: boolean
|
||||
sdkToken: string
|
||||
@@ -274,6 +279,21 @@ export class ConversationService {
|
||||
private providerService = new ProviderService()
|
||||
private pendingPermissionModeChanges = new Map<string, Map<string, number>>()
|
||||
|
||||
async resolveRuntimeApiFormat(providerId: string | null = null): Promise<SessionApiFormat> {
|
||||
const provider = providerId ? await this.providerService.getProvider(providerId) : undefined
|
||||
const format = resolveProviderApiFormat(providerId, provider ?? undefined)
|
||||
if (!format) {
|
||||
throw new ConversationStartupError('The selected provider is unavailable or has an unsupported API protocol.', 'CLI_START_FAILED')
|
||||
}
|
||||
return format
|
||||
}
|
||||
|
||||
async validateSessionProtocol(sessionId: string, providerId: string | null = null): Promise<SessionApiFormat> {
|
||||
const format = await this.resolveRuntimeApiFormat(providerId)
|
||||
assertSessionApiFormat(await sessionService.getSessionApiFormat(sessionId), format)
|
||||
return format
|
||||
}
|
||||
|
||||
private trackPendingPermissionModeChange(sessionId: string, mode: string, delta: 1 | -1): void {
|
||||
const sessionChanges = this.pendingPermissionModeChanges.get(sessionId) ?? new Map<string, number>()
|
||||
const nextCount = (sessionChanges.get(mode) ?? 0) + delta
|
||||
@@ -346,6 +366,7 @@ export class ConversationService {
|
||||
if (this.sessions.has(sessionId)) return
|
||||
|
||||
const launchInfo = await sessionService.getSessionLaunchInfo(sessionId)
|
||||
const apiFormat = await this.validateSessionProtocol(sessionId, options?.providerId ?? null)
|
||||
const shouldResume = !!launchInfo && launchInfo.transcriptMessageCount > 0
|
||||
const shouldReplacePlaceholder =
|
||||
!!launchInfo && launchInfo.transcriptMessageCount === 0
|
||||
@@ -368,7 +389,7 @@ export class ConversationService {
|
||||
)
|
||||
}
|
||||
|
||||
if (shouldReplacePlaceholder) {
|
||||
if (shouldReplacePlaceholder && !launchInfo?.sessionApiFormat) {
|
||||
await sessionService.clearSessionTranscript(sessionId, workDir)
|
||||
}
|
||||
|
||||
@@ -467,6 +488,8 @@ export class ConversationService {
|
||||
})
|
||||
const session: SessionProcess = {
|
||||
proc,
|
||||
providerId: options?.providerId ?? null,
|
||||
apiFormat,
|
||||
outputCallbacks: [],
|
||||
workDir: launchWorkDir,
|
||||
permissionMode: options?.permissionMode || 'default',
|
||||
@@ -606,8 +629,11 @@ export class ConversationService {
|
||||
attachments?: AttachmentRef[],
|
||||
options?: SendMessageOptions,
|
||||
): Promise<boolean> {
|
||||
const userContent = await this.buildUserContent(content, sessionId, attachments)
|
||||
let session = this.sessions.get(sessionId)
|
||||
if (!session) return false
|
||||
const selectedFormat = await this.validateSessionProtocol(sessionId, session.providerId)
|
||||
assertSessionApiFormat(session.apiFormat, selectedFormat)
|
||||
const userContent = await this.buildUserContent(content, sessionId, attachments)
|
||||
if (session && !await this.refreshNetworkEnvironmentBeforeTurn(sessionId, session)) {
|
||||
return false
|
||||
}
|
||||
@@ -620,6 +646,15 @@ export class ConversationService {
|
||||
// one of those awaits is pending, so check ownership at the last possible
|
||||
// point before writing the user message to the SDK socket.
|
||||
if (options?.canSend && !options.canSend()) return false
|
||||
if (!session || this.sessions.get(sessionId) !== session) return false
|
||||
// Saved provider configuration can change while its CLI is alive. Check
|
||||
// both the selected route and the transport captured when it was started.
|
||||
const currentFormat = await this.resolveRuntimeApiFormat(session.providerId)
|
||||
assertSessionApiFormat(session.apiFormat, currentFormat)
|
||||
await sessionService.lockSessionApiFormat(sessionId, currentFormat, session.workDir)
|
||||
if (options?.canSend && !options.canSend()) return false
|
||||
if (this.sessions.get(sessionId) !== session) return false
|
||||
options?.onProtocolLocked?.(currentFormat)
|
||||
const sent = this.sendSdkMessage(sessionId, {
|
||||
type: 'user',
|
||||
...(options?.messageUuid ? { uuid: options.messageUuid } : {}),
|
||||
|
||||
@@ -604,6 +604,37 @@ describe('local index database', () => {
|
||||
}
|
||||
})
|
||||
|
||||
it('upgrades a frozen v4 cache additively and preserves its existing rows', async () => {
|
||||
const databasePath = join(process.env.CLAUDE_CONFIG_DIR!, 'frozen-v4.sqlite')
|
||||
await mkdir(dirname(databasePath), { recursive: true })
|
||||
const seed = await openRawDatabase(databasePath)
|
||||
seedFrozenV3(seed)
|
||||
seed.exec('ALTER TABLE activity_sessions ADD COLUMN active_duration_ms INTEGER NOT NULL DEFAULT 0')
|
||||
seed.exec('PRAGMA user_version = 4')
|
||||
seed.exec("INSERT INTO schema_meta (key, value) VALUES ('future-extension', 'keep-me')")
|
||||
const originalSessions = queryAll<{ transcript_path: string; title: string }>(seed,
|
||||
'SELECT transcript_path, title FROM sessions ORDER BY transcript_path')
|
||||
seed.close(true)
|
||||
const { openLocalIndexDatabase } = await loadDatabase()
|
||||
const upgraded = openLocalIndexDatabase({ path: databasePath })
|
||||
try {
|
||||
expect(upgraded.read(operation => operation.all<{ name: string }>(
|
||||
'PRAGMA table_info(sessions)',
|
||||
).map(row => row.name))).toContain('session_api_format')
|
||||
expect(upgraded.read(operation => operation.all<{ transcript_path: string; title: string }>(
|
||||
'SELECT transcript_path, title FROM sessions ORDER BY transcript_path',
|
||||
))).toEqual(originalSessions)
|
||||
expect(upgraded.read(operation => operation.get<{ value: string }>(
|
||||
"SELECT value FROM schema_meta WHERE key = 'future-extension'",
|
||||
)?.value)).toBe('keep-me')
|
||||
expect(upgraded.read(operation => operation.get<{ count: number }>(
|
||||
'SELECT COUNT(*) AS count FROM sessions WHERE session_api_format IS NOT NULL',
|
||||
)?.count)).toBe(0)
|
||||
} finally {
|
||||
upgraded.close()
|
||||
}
|
||||
})
|
||||
|
||||
it('rolls back an interrupted v2 to v3 migration without changing v2 data', async () => {
|
||||
const databasePath = join(process.env.CLAUDE_CONFIG_DIR!, 'blocked-v3.sqlite')
|
||||
await mkdir(dirname(databasePath), { recursive: true })
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
import type { Database } from 'bun:sqlite'
|
||||
|
||||
export const LOCAL_INDEX_SCHEMA_VERSION = 4
|
||||
export const LOCAL_INDEX_SCHEMA_VERSION = 5
|
||||
export const LOCAL_INDEX_SCHEMA_UNSUPPORTED =
|
||||
'LOCAL_INDEX_SCHEMA_UNSUPPORTED' as const
|
||||
|
||||
@@ -182,11 +182,17 @@ const SCHEMA_V4 = `
|
||||
ALTER TABLE activity_sessions ADD COLUMN active_duration_ms INTEGER NOT NULL DEFAULT 0;
|
||||
`
|
||||
|
||||
// Additive cache migration. Parser v5 reprojects old rows from unchanged transcripts.
|
||||
const SCHEMA_V5 = `
|
||||
ALTER TABLE sessions ADD COLUMN session_api_format TEXT;
|
||||
`
|
||||
|
||||
const MIGRATIONS = [
|
||||
{ version: 1, sql: SCHEMA_V1 },
|
||||
{ version: 2, sql: SCHEMA_V2 },
|
||||
{ version: 3, sql: SCHEMA_V3 },
|
||||
{ version: 4, sql: SCHEMA_V4 },
|
||||
{ version: 5, sql: SCHEMA_V5 },
|
||||
] as const
|
||||
|
||||
export class UnsupportedLocalIndexSchemaError extends Error {
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
import type { SessionProtocolState } from '../../../shared/sessionProtocol.js'
|
||||
import type { LocalIndexDatabase } from './database.js'
|
||||
import { createActivityIndex, type ActivityIndex } from './activityIndex.js'
|
||||
import type {
|
||||
@@ -26,6 +27,7 @@ export type IndexedSessionRow = {
|
||||
runtimeProviderId?: string | null
|
||||
runtimeModelId?: string
|
||||
effortLevel?: string
|
||||
sessionApiFormat?: SessionProtocolState
|
||||
repository?: PersistedRepositorySession
|
||||
worktreeSession?: PersistedWorktreeSession | null
|
||||
}
|
||||
@@ -134,6 +136,7 @@ type SessionRow = {
|
||||
runtime_provider_present: number
|
||||
runtime_model_id: string | null
|
||||
effort_level: string | null
|
||||
session_api_format: SessionProtocolState | null
|
||||
repository_json: string | null
|
||||
worktree_session_json: string | null
|
||||
}
|
||||
@@ -199,6 +202,7 @@ function sessionFromRow(row: SessionRow): IndexedSessionRow {
|
||||
: {}),
|
||||
...(row.runtime_model_id ? { runtimeModelId: row.runtime_model_id } : {}),
|
||||
...(row.effort_level ? { effortLevel: row.effort_level } : {}),
|
||||
...(row.session_api_format ? { sessionApiFormat: row.session_api_format } : {}),
|
||||
...(repository ? { repository } : {}),
|
||||
...(row.worktree_session_json !== null ? { worktreeSession } : {}),
|
||||
}
|
||||
@@ -276,7 +280,7 @@ export function createSessionIndex(database: LocalIndexDatabase): SessionIndex {
|
||||
SELECT transcript_path, session_id, project_path, title, created_at,
|
||||
modified_at, message_count, work_dir, permission_mode,
|
||||
runtime_provider_id, runtime_provider_present,
|
||||
runtime_model_id, effort_level,
|
||||
runtime_model_id, effort_level, session_api_format,
|
||||
repository_json, worktree_session_json
|
||||
FROM sessions
|
||||
ORDER BY modified_at_ms DESC, session_id ASC, transcript_path ASC
|
||||
@@ -286,7 +290,7 @@ export function createSessionIndex(database: LocalIndexDatabase): SessionIndex {
|
||||
SELECT transcript_path, session_id, project_path, title, created_at,
|
||||
modified_at, message_count, work_dir, permission_mode,
|
||||
runtime_provider_id, runtime_provider_present,
|
||||
runtime_model_id, effort_level,
|
||||
runtime_model_id, effort_level, session_api_format,
|
||||
repository_json, worktree_session_json
|
||||
FROM sessions
|
||||
WHERE project_path = ?
|
||||
@@ -394,7 +398,7 @@ export function createSessionIndex(database: LocalIndexDatabase): SessionIndex {
|
||||
sessions.message_count, sessions.work_dir,
|
||||
sessions.permission_mode, sessions.runtime_provider_id,
|
||||
sessions.runtime_provider_present,
|
||||
sessions.runtime_model_id, sessions.effort_level,
|
||||
sessions.runtime_model_id, sessions.effort_level, sessions.session_api_format,
|
||||
sessions.repository_json, sessions.worktree_session_json,
|
||||
source_files.indexed_bytes, source_files.size_bytes
|
||||
FROM sessions
|
||||
@@ -423,6 +427,7 @@ export function createSessionIndex(database: LocalIndexDatabase): SessionIndex {
|
||||
: {}),
|
||||
...(row.runtime_model_id ? { runtimeModelId: row.runtime_model_id } : {}),
|
||||
...(row.effort_level ? { effortLevel: row.effort_level } : {}),
|
||||
...(row.session_api_format ? { sessionApiFormat: row.session_api_format } : {}),
|
||||
...(repository ? { repository } : {}),
|
||||
...(row.worktree_session_json !== null
|
||||
? { worktreeSession }
|
||||
|
||||
@@ -431,6 +431,35 @@ describe('session projector', () => {
|
||||
}
|
||||
})
|
||||
|
||||
it('round-trips inferred and locked protocols through projection, SQLite reads and restart seeds', async () => {
|
||||
const root = await createTempDir('projector-protocol')
|
||||
const database = openLocalIndexDatabase({ path: join(root, 'index.sqlite') })
|
||||
const index = createSessionIndex(database)
|
||||
const projector = createSessionProjector({ database, index, scope: root })
|
||||
try {
|
||||
for (const protocol of ['anthropic', 'openai_chat', 'openai_responses', 'mixed', 'unknown'] as const) {
|
||||
const candidate = await createCandidate({
|
||||
root, projectPath: '-repo-a', sessionId: protocol,
|
||||
content: line({ type: 'session-meta', sessionApiFormat: protocol }),
|
||||
})
|
||||
const original = await sourceHash(candidate.path)
|
||||
await projector.projectSource(candidate)
|
||||
expect(index.listSessions({ limit: 20 }).sessions.find(row => row.id === protocol)?.sessionApiFormat).toBe(protocol)
|
||||
expect(index.getProjectionSeed(candidate.path)?.summary.sessionApiFormat).toBe(protocol)
|
||||
expect(await sourceHash(candidate.path)).toBe(original)
|
||||
}
|
||||
const legacy = await createCandidate({
|
||||
root, projectPath: '-repo-a', sessionId: 'legacy',
|
||||
content: line({ type: 'session-meta', runtimeProviderId: 'openai-official' }) +
|
||||
line(user('Legacy', '2026-01-01T00:00:00.000Z')) + line(assistant('2026-01-01T00:00:01.000Z')),
|
||||
})
|
||||
await projector.projectSource(legacy)
|
||||
expect(index.getProjectionSeed(legacy.path)?.summary.sessionApiFormat).toBe('openai_responses')
|
||||
} finally {
|
||||
database.close()
|
||||
}
|
||||
})
|
||||
|
||||
it('removes only a confirmed missing source projection', async () => {
|
||||
const root = await createTempDir('projector-delete')
|
||||
const candidate = await createCandidate({
|
||||
|
||||
@@ -34,7 +34,7 @@ import type {
|
||||
// that refreshes already-indexed transcripts.
|
||||
// 3: usage is deduplicated per (message.id, requestId), and sessions carry active working time.
|
||||
// 4: usage copied into a fork is excluded from the fork's activity projection.
|
||||
export const SESSION_SUMMARY_PARSER_VERSION = 4
|
||||
export const SESSION_SUMMARY_PARSER_VERSION = 5
|
||||
|
||||
export type SessionSourceCandidate = {
|
||||
path: string
|
||||
@@ -573,8 +573,8 @@ export function createSessionProjector(options: SessionProjectorOptions): Sessio
|
||||
transcript_path, session_id, project_path, title, created_at,
|
||||
modified_at, modified_at_ms, message_count, work_dir, repository_json,
|
||||
worktree_session_json, permission_mode, runtime_provider_id,
|
||||
runtime_provider_present, runtime_model_id, effort_level
|
||||
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
runtime_provider_present, runtime_model_id, effort_level, session_api_format
|
||||
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
ON CONFLICT(transcript_path) DO UPDATE SET
|
||||
session_id = excluded.session_id,
|
||||
project_path = excluded.project_path,
|
||||
@@ -590,7 +590,8 @@ export function createSessionProjector(options: SessionProjectorOptions): Sessio
|
||||
runtime_provider_id = excluded.runtime_provider_id,
|
||||
runtime_provider_present = excluded.runtime_provider_present,
|
||||
runtime_model_id = excluded.runtime_model_id,
|
||||
effort_level = excluded.effort_level
|
||||
effort_level = excluded.effort_level,
|
||||
session_api_format = excluded.session_api_format
|
||||
`,
|
||||
bundle.candidate.path,
|
||||
bundle.candidate.sessionId,
|
||||
@@ -609,7 +610,8 @@ export function createSessionProjector(options: SessionProjectorOptions): Sessio
|
||||
summary.runtimeProviderId ?? null,
|
||||
runtimeProviderPresent,
|
||||
summary.runtimeModelId ?? null,
|
||||
summary.effortLevel ?? null)
|
||||
summary.effortLevel ?? null,
|
||||
summary.sessionApiFormat ?? null)
|
||||
|
||||
writeBackfillState(
|
||||
writer,
|
||||
|
||||
@@ -116,6 +116,7 @@ describe('reduceTranscript', () => {
|
||||
const result = reduceTranscript(chunks, initialProjection())
|
||||
|
||||
expect(result.summary).toEqual({
|
||||
sessionApiFormat: 'unknown',
|
||||
title: 'Pinned title',
|
||||
createdAt: '2026-01-01T00:00:00.000Z',
|
||||
modifiedAt: '2026-01-01T00:02:00.000Z',
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
import { createSessionProtocolAccumulator } from '../sessionProtocolHistory.js'
|
||||
import { cleanSessionTitleSource } from '../../../utils/sessionTitleText.js'
|
||||
import { SYNTHETIC_MODEL } from '../../../utils/messages.js'
|
||||
import { extractShotCountFromAssistantContent } from '../../../utils/shotStats.js'
|
||||
@@ -89,6 +90,7 @@ type ReducerState = {
|
||||
runtimeProviderId: string | null | undefined
|
||||
runtimeModelId: string | undefined
|
||||
effortLevel: string | undefined
|
||||
sessionProtocol: ReturnType<typeof createSessionProtocolAccumulator>
|
||||
repository: PersistedRepositorySession | undefined
|
||||
worktreeSession: PersistedWorktreeSession | null | undefined
|
||||
nextOrdinal: number
|
||||
@@ -204,6 +206,7 @@ function latestTimestamp(current: string | null, candidate: unknown): string | n
|
||||
function cloneState(state: ReducerState): ReducerState {
|
||||
return {
|
||||
...state,
|
||||
sessionProtocol: state.sessionProtocol.clone(),
|
||||
repository: state.repository ? { ...state.repository } : undefined,
|
||||
worktreeSession: state.worktreeSession
|
||||
? { ...state.worktreeSession }
|
||||
@@ -248,6 +251,7 @@ function createInitialState(
|
||||
runtimeProviderId: undefined,
|
||||
runtimeModelId: undefined,
|
||||
effortLevel: undefined,
|
||||
sessionProtocol: createSessionProtocolAccumulator(),
|
||||
repository: undefined,
|
||||
worktreeSession: undefined,
|
||||
nextOrdinal: 0,
|
||||
@@ -506,6 +510,7 @@ function applyActivityEntry(state: ReducerState, entry: ReducerEntry): void {
|
||||
}
|
||||
|
||||
function applyEntry(state: ReducerState, entry: ReducerEntry): void {
|
||||
state.sessionProtocol.add(entry)
|
||||
applyActivityEntry(state, entry)
|
||||
if (!state.hasCreatedAt && entry.timestamp) {
|
||||
state.createdAt = entry.timestamp
|
||||
@@ -612,6 +617,7 @@ function summaryFromState(state: ReducerState): SessionListSummary {
|
||||
: {}),
|
||||
...(state.runtimeModelId ? { runtimeModelId: state.runtimeModelId } : {}),
|
||||
...(state.effortLevel ? { effortLevel: state.effortLevel } : {}),
|
||||
...(state.sessionProtocol.get() ? { sessionApiFormat: state.sessionProtocol.get() } : {}),
|
||||
...(state.repository ? { repository: { ...state.repository } } : {}),
|
||||
...(state.worktreeSession !== undefined
|
||||
? {
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
import type { SessionProtocolState } from '../../../shared/sessionProtocol.js'
|
||||
export type LocalIndexMode = 'off' | 'shadow' | 'on'
|
||||
export type LocalIndexState = 'off' | 'building' | 'ready' | 'degraded'
|
||||
|
||||
@@ -46,6 +47,7 @@ export type SessionListSummary = {
|
||||
runtimeProviderId?: string | null
|
||||
runtimeModelId?: string
|
||||
effortLevel?: string
|
||||
sessionApiFormat?: SessionProtocolState
|
||||
repository?: PersistedRepositorySession
|
||||
worktreeSession?: PersistedWorktreeSession | null
|
||||
}
|
||||
|
||||
@@ -0,0 +1,65 @@
|
||||
import { describe, expect, it } from 'bun:test'
|
||||
import { encodeOpenAIReasoningEnvelope } from '../proxy/transform/openaiReasoning.js'
|
||||
import { createSessionProtocolAccumulator, inferSessionApiFormat } from './sessionProtocolHistory.js'
|
||||
|
||||
const user = { type: 'user', message: { role: 'user', content: 'hello' } }
|
||||
const assistant = { type: 'assistant', message: { id: 'reply', role: 'assistant', model: 'same-model', content: 'hello' } }
|
||||
const meta = (runtimeProviderId: string | null) => ({ type: 'session-meta', runtimeProviderId })
|
||||
|
||||
describe('legacy session protocol inference', () => {
|
||||
it('leaves selection-only and synthetic placeholder sessions unlocked', () => {
|
||||
expect(inferSessionApiFormat([meta('openai-official')])).toBeUndefined()
|
||||
expect(inferSessionApiFormat([{ ...assistant, message: { ...assistant.message, model: '<synthetic>' } }])).toBeUndefined()
|
||||
})
|
||||
|
||||
it.each([
|
||||
[null, 'anthropic'], ['claude-official', 'anthropic'],
|
||||
['openai-official', 'openai_responses'], ['grok-official', 'openai_responses'],
|
||||
] as const)('uses the recorded immutable route %s', (provider, protocol) => {
|
||||
expect(inferSessionApiFormat([meta(provider), user, assistant])).toBe(protocol)
|
||||
})
|
||||
|
||||
it('never guesses from model names or mutable saved-provider ids', () => {
|
||||
expect(inferSessionApiFormat([user, assistant])).toBe('unknown')
|
||||
expect(inferSessionApiFormat([meta('saved-provider'), user, assistant])).toBe('unknown')
|
||||
})
|
||||
|
||||
it('does not count unused model selections or assign a later route to earlier messages', () => {
|
||||
expect(inferSessionApiFormat([meta(null), user, assistant, meta('openai-official')])).toBe('anthropic')
|
||||
expect(inferSessionApiFormat([user, assistant, meta('openai-official')])).toBe('unknown')
|
||||
expect(inferSessionApiFormat([meta(null), user, meta('openai-official')])).toBe('anthropic')
|
||||
})
|
||||
|
||||
it('marks actual cross-protocol history mixed and incomplete evidence unknown', () => {
|
||||
expect(inferSessionApiFormat([meta(null), user, assistant, meta('openai-official'), user, assistant])).toBe('mixed')
|
||||
expect(inferSessionApiFormat([meta('saved'), user, assistant, meta('openai-official'), user, {
|
||||
...assistant, message: { ...assistant.message, id: 'reply-2' },
|
||||
}])).toBe('unknown')
|
||||
})
|
||||
|
||||
it('recognizes application-owned Responses evidence even across streamed content blocks', () => {
|
||||
const data = encodeOpenAIReasoningEnvelope({ type: 'reasoning', summary: [], encrypted_content: 'fixture-only' })!
|
||||
const reasoning = { ...assistant, message: { ...assistant.message, content: [{ type: 'redacted_thinking', data }] } }
|
||||
expect(inferSessionApiFormat([user, assistant, reasoning])).toBe('openai_responses')
|
||||
expect(inferSessionApiFormat([user, reasoning, assistant])).toBe('openai_responses')
|
||||
expect(inferSessionApiFormat([user, { ...reasoning, message: { ...reasoning.message, content: [{ type: 'redacted_thinking', data: 'opaque' }] } }])).toBe('unknown')
|
||||
})
|
||||
|
||||
it('preserves explicit locks and detects conflicting locks', () => {
|
||||
const lock = { type: 'session-meta', sessionApiFormat: 'openai_chat' }
|
||||
expect(inferSessionApiFormat([lock, meta('saved-provider'), user, assistant])).toBe('openai_chat')
|
||||
expect(inferSessionApiFormat([lock, { ...lock, sessionApiFormat: 'anthropic' }])).toBe('mixed')
|
||||
})
|
||||
|
||||
it('keeps incremental projection branches independent', () => {
|
||||
const base = createSessionProtocolAccumulator()
|
||||
base.add(meta(null))
|
||||
base.add(user)
|
||||
base.add(assistant)
|
||||
const changed = base.clone()
|
||||
changed.add(meta('openai-official'))
|
||||
changed.add(assistant)
|
||||
expect(changed.get()).toBe('mixed')
|
||||
expect(base.get()).toBe('anthropic')
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,135 @@
|
||||
import type { SessionApiFormat, SessionProtocolState } from '../../shared/sessionProtocol.js'
|
||||
import { resolveProviderApiFormat, isSessionApiFormat } from '../../shared/sessionProtocol.js'
|
||||
import { parseOpenAIReasoningEnvelope } from '../../utils/openAIReasoningEnvelope.js'
|
||||
|
||||
export class SessionProtocolError extends Error {
|
||||
readonly statusCode = 409
|
||||
readonly code: 'SESSION_PROTOCOL_UNRESOLVED' | 'SESSION_PROTOCOL_MISMATCH'
|
||||
constructor(
|
||||
public readonly currentFormat: SessionProtocolState,
|
||||
public readonly requestedFormat: SessionApiFormat,
|
||||
) {
|
||||
super(
|
||||
currentFormat === 'mixed' || currentFormat === 'unknown'
|
||||
? `This session's API protocol is ${currentFormat}. Start a new session to use ${requestedFormat}.`
|
||||
: `This session uses ${currentFormat}; switching to ${requestedFormat} requires a new session.`)
|
||||
this.name = 'SessionProtocolError'
|
||||
this.code = currentFormat === 'mixed' || currentFormat === 'unknown'
|
||||
? 'SESSION_PROTOCOL_UNRESOLVED'
|
||||
: 'SESSION_PROTOCOL_MISMATCH'
|
||||
}
|
||||
}
|
||||
|
||||
export function assertSessionApiFormat(
|
||||
current: SessionProtocolState | undefined,
|
||||
requested: SessionApiFormat,
|
||||
): void {
|
||||
if (current && current !== requested) throw new SessionProtocolError(current, requested)
|
||||
}
|
||||
|
||||
type HistoryEntry = {
|
||||
type?: string
|
||||
isMeta?: boolean
|
||||
message?: { id?: unknown; role?: string; model?: string; content?: unknown }
|
||||
[key: string]: unknown
|
||||
}
|
||||
|
||||
/** Derives a read-only forward view of old transcripts; never guesses from model names
|
||||
* or today's mutable saved-provider configuration. The durable lock is additive metadata. */
|
||||
type ProtocolAccumulatorSeed = {
|
||||
formats: Set<SessionApiFormat>
|
||||
lockedFormat?: SessionProtocolState
|
||||
runtimeFormat?: SessionApiFormat
|
||||
hasMessages: boolean
|
||||
userFormat?: SessionApiFormat
|
||||
hasAssistant: boolean
|
||||
unknownResponses: Set<string>
|
||||
knownResponses: Set<string>
|
||||
anonymousResponseCount: number
|
||||
}
|
||||
|
||||
type SessionProtocolAccumulator = {
|
||||
clone(): SessionProtocolAccumulator
|
||||
add(entry: HistoryEntry): void
|
||||
get(): SessionProtocolState | undefined
|
||||
}
|
||||
|
||||
export function createSessionProtocolAccumulator(seed?: ProtocolAccumulatorSeed): SessionProtocolAccumulator {
|
||||
const formats = new Set<SessionApiFormat>(seed?.formats)
|
||||
let lockedFormat = seed?.lockedFormat
|
||||
let runtimeFormat = seed?.runtimeFormat
|
||||
let hasMessages = seed?.hasMessages ?? false
|
||||
let userFormat = seed?.userFormat
|
||||
let hasAssistant = seed?.hasAssistant ?? false
|
||||
const unknownResponses = new Set(seed?.unknownResponses)
|
||||
const knownResponses = new Set(seed?.knownResponses)
|
||||
let anonymousResponseCount = seed?.anonymousResponseCount ?? 0
|
||||
|
||||
return {
|
||||
clone() {
|
||||
return createSessionProtocolAccumulator({
|
||||
formats, lockedFormat, runtimeFormat, hasMessages, userFormat, hasAssistant,
|
||||
unknownResponses, knownResponses, anonymousResponseCount,
|
||||
})
|
||||
},
|
||||
add(entry: HistoryEntry): void {
|
||||
if (entry.type === 'session-meta') {
|
||||
const format = entry.sessionApiFormat
|
||||
if (isSessionApiFormat(format)) {
|
||||
formats.add(format)
|
||||
lockedFormat = format
|
||||
} else if (format === 'mixed' || format === 'unknown') {
|
||||
lockedFormat = format
|
||||
} else if (format !== undefined) {
|
||||
lockedFormat = 'unknown'
|
||||
}
|
||||
// These explicit routing snapshots are usable when present in imported
|
||||
// history. Merely selecting a route does not count as having used it.
|
||||
if (isSessionApiFormat(entry.runtimeApiFormat)) {
|
||||
runtimeFormat = entry.runtimeApiFormat
|
||||
} else if (Object.prototype.hasOwnProperty.call(entry, 'runtimeProviderId')) {
|
||||
runtimeFormat = resolveProviderApiFormat(entry.runtimeProviderId as string | null)
|
||||
}
|
||||
return
|
||||
}
|
||||
if (entry.isMeta || (entry.type !== 'user' && entry.type !== 'assistant') || !entry.message?.role) return
|
||||
if (entry.type === 'assistant' && entry.message.model === '<synthetic>') return
|
||||
hasMessages = true
|
||||
if (entry.type !== 'assistant') {
|
||||
userFormat = runtimeFormat
|
||||
return
|
||||
}
|
||||
hasAssistant = true
|
||||
const hasResponsesEnvelope = Array.isArray(entry.message.content) && entry.message.content.some(block => (
|
||||
block && typeof block === 'object' && block.type === 'redacted_thinking' &&
|
||||
typeof block.data === 'string' && parseOpenAIReasoningEnvelope(block.data) !== null
|
||||
))
|
||||
const format = hasResponsesEnvelope ? 'openai_responses' : runtimeFormat
|
||||
// One streamed reply may occupy several JSONL content-block entries.
|
||||
const responseId = typeof entry.message.id === 'string'
|
||||
? entry.message.id
|
||||
: `anonymous:${++anonymousResponseCount}`
|
||||
if (format) {
|
||||
formats.add(format)
|
||||
knownResponses.add(responseId)
|
||||
unknownResponses.delete(responseId)
|
||||
} else if (!knownResponses.has(responseId)) {
|
||||
unknownResponses.add(responseId)
|
||||
}
|
||||
},
|
||||
get(): SessionProtocolState | undefined {
|
||||
if (formats.size > 1 || lockedFormat === 'mixed') return 'mixed'
|
||||
if (lockedFormat) return lockedFormat
|
||||
if (!hasMessages) return undefined
|
||||
if (unknownResponses.size > 0) return 'unknown'
|
||||
if (!hasAssistant) return userFormat ?? 'unknown'
|
||||
return formats.values().next().value ?? 'unknown'
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
export function inferSessionApiFormat(entries: HistoryEntry[]): SessionProtocolState | undefined {
|
||||
const accumulator = createSessionProtocolAccumulator()
|
||||
for (const entry of entries) accumulator.add(entry)
|
||||
return accumulator.get()
|
||||
}
|
||||
@@ -76,6 +76,12 @@ import {
|
||||
type ProjectHistoryRow,
|
||||
} from './projectSessionHistory.js'
|
||||
|
||||
import { isSessionApiFormat, type SessionApiFormat, type SessionProtocolState } from '../../shared/sessionProtocol.js'
|
||||
import { createSessionProtocolAccumulator, inferSessionApiFormat, SessionProtocolError } from './sessionProtocolHistory.js'
|
||||
|
||||
// Shared across service instances: two first sends cannot establish different locks.
|
||||
const sessionProtocolWrites = new Map<string, Promise<void>>()
|
||||
|
||||
// ============================================================================
|
||||
// Types
|
||||
// ============================================================================
|
||||
@@ -95,6 +101,7 @@ export type SessionListItem = {
|
||||
runtimeProviderId?: string | null
|
||||
runtimeModelId?: string
|
||||
effortLevel?: string
|
||||
sessionApiFormat?: SessionProtocolState
|
||||
}
|
||||
|
||||
export type SubagentTranscriptFragment = {
|
||||
@@ -182,6 +189,7 @@ export type SessionLaunchInfo = {
|
||||
runtimeProviderId?: string | null
|
||||
runtimeModelId?: string
|
||||
effortLevel?: string
|
||||
sessionApiFormat?: SessionProtocolState
|
||||
}
|
||||
|
||||
type ProviderContextWindowHint = Pick<SessionLaunchInfo, 'runtimeProviderId' | 'runtimeModelId'>
|
||||
@@ -2954,6 +2962,7 @@ export class SessionService {
|
||||
let runtimeModelId: string | undefined
|
||||
let effortLevel: string | undefined
|
||||
let customTitle: string | null = null
|
||||
const sessionProtocol = createSessionProtocolAccumulator()
|
||||
let transcriptMessageCount = 0
|
||||
const metadata: TranscriptMetadataSnapshot = {}
|
||||
|
||||
@@ -2971,6 +2980,7 @@ export class SessionService {
|
||||
const contextState = createTranscriptContextAccumulator()
|
||||
|
||||
await this.streamJsonlFile(found.filePath, (entry) => {
|
||||
sessionProtocol.add(entry)
|
||||
if (typeof entry.message?.model === 'string') {
|
||||
metadata.model = entry.message.model
|
||||
}
|
||||
@@ -3124,6 +3134,7 @@ export class SessionService {
|
||||
|
||||
const workDir = latestWorkDir || latestCwd || this.desanitizePath(found.projectDir) || process.cwd()
|
||||
const launchInfo: SessionLaunchInfo = {
|
||||
sessionApiFormat: sessionProtocol.get(),
|
||||
filePath: found.filePath,
|
||||
projectDir: found.projectDir,
|
||||
workDir,
|
||||
@@ -3432,6 +3443,7 @@ export class SessionService {
|
||||
workDir,
|
||||
workDirExists,
|
||||
workspaceState,
|
||||
sessionApiFormat: row.sessionApiFormat,
|
||||
permissionMode: row.permissionMode,
|
||||
...(row.runtimeProviderId !== undefined
|
||||
? { runtimeProviderId: row.runtimeProviderId }
|
||||
@@ -3447,6 +3459,7 @@ export class SessionService {
|
||||
): SessionListShadowComparison {
|
||||
const fieldHashes: SessionListShadowComparison['fieldHashes'] = []
|
||||
const fields: Array<keyof SessionListItem> = [
|
||||
'sessionApiFormat',
|
||||
'id',
|
||||
'title',
|
||||
'createdAt',
|
||||
@@ -3611,6 +3624,7 @@ export class SessionService {
|
||||
workDir,
|
||||
workDirExists,
|
||||
workspaceState,
|
||||
sessionApiFormat: summary.sessionApiFormat,
|
||||
permissionMode: summary.permissionMode,
|
||||
...(summary.runtimeProviderId !== undefined
|
||||
? { runtimeProviderId: summary.runtimeProviderId }
|
||||
@@ -3711,6 +3725,7 @@ export class SessionService {
|
||||
workDirExists,
|
||||
workspaceState,
|
||||
permissionMode,
|
||||
sessionApiFormat: inferSessionApiFormat(entries),
|
||||
messages,
|
||||
}
|
||||
}
|
||||
@@ -4149,6 +4164,62 @@ export class SessionService {
|
||||
return typeof entry?.cwd === 'string' && entry.cwd.trim() ? entry.cwd : null
|
||||
}
|
||||
|
||||
private async findProtocolSessionFile(sessionId: string): Promise<{ filePath: string; projectDir: string } | null> {
|
||||
// SDK/WebSocket callers can use safe ad-hoc IDs before a UUID is assigned.
|
||||
if (!/^[a-zA-Z0-9_-]+$/.test(sessionId)) throw ApiError.badRequest('Invalid session ID')
|
||||
if (this.isValidSessionId(sessionId)) return this.findSessionFile(sessionId)
|
||||
return (await this.findSessionFilesFromFiles(sessionId))[0] ?? null
|
||||
}
|
||||
|
||||
async getSessionApiFormat(sessionId: string): Promise<SessionProtocolState | undefined> {
|
||||
const found = await this.findProtocolSessionFile(sessionId)
|
||||
if (!found) return undefined
|
||||
const protocol = createSessionProtocolAccumulator()
|
||||
await this.streamJsonlFile(found.filePath, entry => protocol.add(entry))
|
||||
return protocol.get()
|
||||
}
|
||||
|
||||
/** Establish the immutable protocol immediately before the first real send.
|
||||
* Legacy histories are upgraded by appending metadata only after unambiguous inference. */
|
||||
async lockSessionApiFormat(sessionId: string, apiFormat: SessionApiFormat, workDir?: string): Promise<void> {
|
||||
if (!isSessionApiFormat(apiFormat)) throw ApiError.badRequest('Invalid API protocol')
|
||||
const key = `${this.getConfigDir()}\0${sessionId}`
|
||||
const previous = sessionProtocolWrites.get(key) ?? Promise.resolve()
|
||||
const write = previous.catch(() => {}).then(async () => {
|
||||
let found = await this.findProtocolSessionFile(sessionId)
|
||||
if (!found && workDir) {
|
||||
const absoluteWorkDir = await fs.realpath(path.resolve(normalizeDriveRootPathForPlatform(workDir)))
|
||||
const projectDir = this.sanitizePath(absoluteWorkDir)
|
||||
const directory = path.join(this.getProjectsDir(), projectDir)
|
||||
await fs.mkdir(directory, { recursive: true })
|
||||
found = { filePath: path.join(directory, `${sessionId}.jsonl`), projectDir }
|
||||
await this.appendJsonlEntry(found.filePath, {
|
||||
type: 'session-meta', isMeta: true, workDir: absoluteWorkDir,
|
||||
timestamp: new Date().toISOString(),
|
||||
})
|
||||
}
|
||||
if (!found) throw ApiError.notFound(`Session not found: ${sessionId}`)
|
||||
const entries = await this.readJsonlFile(found.filePath)
|
||||
const existing = inferSessionApiFormat(entries)
|
||||
if (existing && existing !== apiFormat) throw new SessionProtocolError(existing, apiFormat)
|
||||
if (entries.some(entry => entry.type === 'session-meta' &&
|
||||
(entry as Record<string, unknown>).sessionApiFormat === apiFormat)) return
|
||||
await this.appendJsonlEntry(found.filePath, {
|
||||
type: 'session-meta',
|
||||
isMeta: true,
|
||||
sessionApiFormat: apiFormat,
|
||||
timestamp: new Date().toISOString(),
|
||||
})
|
||||
this.invalidateSessionListCache()
|
||||
})
|
||||
sessionProtocolWrites.set(key, write)
|
||||
try {
|
||||
await write
|
||||
} finally {
|
||||
if (sessionProtocolWrites.get(key) === write) sessionProtocolWrites.delete(key)
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Inspect how a session should be launched.
|
||||
* Placeholder desktop-created sessions have zero transcript messages.
|
||||
@@ -4190,6 +4261,7 @@ export class SessionService {
|
||||
const transcriptMessageCount = this.countTranscriptMessages(entries)
|
||||
|
||||
return {
|
||||
sessionApiFormat: inferSessionApiFormat(entries),
|
||||
filePath: found.filePath,
|
||||
projectDir: found.projectDir,
|
||||
workDir,
|
||||
@@ -4250,6 +4322,7 @@ export class SessionService {
|
||||
? preservedPermissionMode
|
||||
: this.resolvePermissionModeFromEntries(entries)
|
||||
const now = new Date().toISOString()
|
||||
const sessionApiFormat = inferSessionApiFormat(entries)
|
||||
|
||||
const initialEntry = {
|
||||
type: 'file-history-snapshot',
|
||||
@@ -4265,6 +4338,7 @@ export class SessionService {
|
||||
const metaEntry = {
|
||||
type: 'session-meta',
|
||||
isMeta: true,
|
||||
...(sessionApiFormat ? { sessionApiFormat } : {}),
|
||||
workDir,
|
||||
repository,
|
||||
...(permissionMode ? { permissionMode } : {}),
|
||||
@@ -4336,11 +4410,18 @@ export class SessionService {
|
||||
}
|
||||
}
|
||||
|
||||
// Moving the session metadata to a new workspace must not turn an existing
|
||||
// conversation into an unlocked placeholder. Same-file updates retain the old record.
|
||||
const sessionApiFormat = matches[0]?.filePath !== targetFilePath
|
||||
? inferSessionApiFormat(await this.readJsonlFile(matches[0]!.filePath))
|
||||
: undefined
|
||||
|
||||
await fs.mkdir(path.dirname(targetFilePath), { recursive: true })
|
||||
|
||||
await this.appendJsonlEntry(targetFilePath, {
|
||||
type: 'session-meta',
|
||||
isMeta: true,
|
||||
...(sessionApiFormat ? { sessionApiFormat } : {}),
|
||||
workDir: normalizedWorkDir,
|
||||
repository,
|
||||
...(metadata.permissionMode && VALID_SESSION_PERMISSION_MODES.has(metadata.permissionMode)
|
||||
|
||||
@@ -109,6 +109,7 @@ export type ServerMessage =
|
||||
*/
|
||||
| { type: 'thinking'; text: string; complete?: boolean }
|
||||
| { type: 'status'; state: ChatState; verb?: string; attemptStart?: boolean }
|
||||
| { type: 'session_protocol'; sessionApiFormat: import('../../shared/sessionProtocol.js').SessionProtocolState }
|
||||
| {
|
||||
type: typeof RUNTIME_CONFIG_APPLIED_EVENT
|
||||
providerId: string | null
|
||||
|
||||
@@ -15,6 +15,7 @@ import type {
|
||||
TokenUsage,
|
||||
} from './events.js'
|
||||
import { RUNTIME_CONFIG_APPLIED_EVENT } from './events.js'
|
||||
import { SessionProtocolError } from '../services/sessionProtocolHistory.js'
|
||||
import * as os from 'node:os'
|
||||
import {
|
||||
ConversationStartupError,
|
||||
@@ -577,6 +578,9 @@ export const handleWebSocket = {
|
||||
|
||||
const msg: ServerMessage = { type: 'connected', sessionId }
|
||||
sendMessage(ws, msg)
|
||||
void sessionService.getSessionApiFormat(sessionId).then(sessionApiFormat => {
|
||||
if (sessionApiFormat) sendMessage(ws, { type: 'session_protocol', sessionApiFormat })
|
||||
}).catch(err => console.warn('[WS] Failed to restore session protocol:', err))
|
||||
const toolRequestIds = replayPendingPermissionRequests(ws, sessionId)
|
||||
const computerUseRequestIds = replayPendingComputerUsePermissionRequests(ws, sessionId)
|
||||
sendMessage(ws, {
|
||||
@@ -642,11 +646,14 @@ export const handleWebSocket = {
|
||||
clearActiveUserTurn(sessionId, activeTurn)
|
||||
const titleState = sessionTitleState.get(sessionId)
|
||||
if (titleState) titleState.activeTurn = undefined
|
||||
if (err instanceof SessionProtocolError) {
|
||||
sendToSession(sessionId, { type: 'session_protocol', sessionApiFormat: err.currentFormat })
|
||||
}
|
||||
sendMessage(ws, {
|
||||
type: 'error',
|
||||
message: 'The request could not be started. Please retry.',
|
||||
code: 'USER_TURN_FAILED',
|
||||
retryable: true,
|
||||
message: err instanceof SessionProtocolError ? err.message : 'The request could not be started. Please retry.',
|
||||
code: err instanceof SessionProtocolError ? err.code : 'USER_TURN_FAILED',
|
||||
retryable: !(err instanceof SessionProtocolError),
|
||||
})
|
||||
sendMessage(ws, { type: 'status', state: 'idle' })
|
||||
}
|
||||
@@ -848,6 +855,15 @@ async function handleUserMessage(
|
||||
await ensureCliSessionStarted(ws, sessionId, 'user_message')
|
||||
} catch (err) {
|
||||
if (activeUserTurns.get(sessionId) !== activeTurn || activeTurn.cancelled) return
|
||||
if (err instanceof SessionProtocolError) {
|
||||
const sessionApiFormat = await sessionService.getSessionApiFormat(sessionId)
|
||||
if (sessionApiFormat) sendMessage(ws, { type: 'session_protocol', sessionApiFormat })
|
||||
sendMessage(ws, { type: 'error', message: err.message, code: err.code, retryable: false })
|
||||
sendMessage(ws, { type: 'status', state: 'idle' })
|
||||
failSessionChatActivity(sessionId)
|
||||
clearActiveUserTurn(sessionId, activeTurn)
|
||||
return
|
||||
}
|
||||
const errMsg = err instanceof Error ? err.message : String(err)
|
||||
const code =
|
||||
err instanceof ConversationStartupError ? err.code : 'CLI_START_FAILED'
|
||||
@@ -926,6 +942,9 @@ async function handleUserMessage(
|
||||
canSend: () =>
|
||||
activeUserTurns.get(sessionId) === activeTurn && !activeTurn.cancelled,
|
||||
messageUuid: activeTurn.expectedReplayUuid,
|
||||
onProtocolLocked: sessionApiFormat => {
|
||||
sendToSession(sessionId, { type: 'session_protocol', sessionApiFormat })
|
||||
},
|
||||
onCommitted: () => {
|
||||
activeTurn.messageSent = true
|
||||
},
|
||||
@@ -1466,6 +1485,24 @@ async function handleSetRuntimeConfig(
|
||||
// A user message arriving in that async admission window must wait for the
|
||||
// selected runtime instead of entering the previous provider's CLI process.
|
||||
await enqueueRuntimeTransition(sessionId, async () => {
|
||||
try {
|
||||
// Empty sessions retain the existing stale-provider fallback. Their
|
||||
// actual resolved transport is checked and locked when sending.
|
||||
if (await sessionService.getSessionApiFormat(sessionId)) {
|
||||
await conversationService.validateSessionProtocol(sessionId, message.providerId ?? null)
|
||||
}
|
||||
} catch (err) {
|
||||
const sessionApiFormat = await sessionService.getSessionApiFormat(sessionId)
|
||||
if (sessionApiFormat) sendMessage(ws, { type: 'session_protocol', sessionApiFormat })
|
||||
sendMessage(ws, {
|
||||
type: 'error',
|
||||
message: err instanceof Error ? err.message : String(err),
|
||||
code: err instanceof SessionProtocolError ? err.code : 'RUNTIME_CONFIG_INVALID',
|
||||
retryable: false,
|
||||
})
|
||||
broadcastAppliedRuntimeConfig(sessionId)
|
||||
return
|
||||
}
|
||||
let modelId = requestedModelId
|
||||
if (isGrokOfficialProviderId(message.providerId)) {
|
||||
modelId = (await getGrokReasoningEfforts(modelId)).modelId
|
||||
|
||||
@@ -0,0 +1,30 @@
|
||||
import { describe, expect, test } from 'bun:test'
|
||||
import { isSessionApiFormat, resolveProviderApiFormat } from './sessionProtocol.js'
|
||||
|
||||
describe('session upstream protocol', () => {
|
||||
test.each([
|
||||
[null, 'anthropic'],
|
||||
['claude-official', 'anthropic'],
|
||||
['openai-official', 'openai_responses'],
|
||||
['grok-official', 'openai_responses'],
|
||||
] as const)('resolves the immutable built-in route %s', (id, format) => {
|
||||
expect(resolveProviderApiFormat(id)).toBe(format)
|
||||
})
|
||||
|
||||
test.each(['anthropic', 'openai_chat', 'openai_responses'] as const)('uses the saved upstream format %s', apiFormat => {
|
||||
expect(resolveProviderApiFormat('custom-provider', { apiFormat })).toBe(apiFormat)
|
||||
})
|
||||
|
||||
test.each(['openai_oauth', 'grok_oauth'])('uses the actual OAuth transport for %s', runtimeKind => {
|
||||
expect(resolveProviderApiFormat('custom-provider', { apiFormat: 'anthropic', runtimeKind })).toBe('openai_responses')
|
||||
})
|
||||
|
||||
test('does not guess protocols for missing providers or unsupported formats', () => {
|
||||
expect(resolveProviderApiFormat(undefined)).toBeUndefined()
|
||||
expect(resolveProviderApiFormat('deleted-provider')).toBeUndefined()
|
||||
expect(resolveProviderApiFormat('provider', { apiFormat: 'future-format' })).toBeUndefined()
|
||||
expect(resolveProviderApiFormat('legacy-provider', {})).toBe('anthropic')
|
||||
expect(isSessionApiFormat('mixed')).toBe(false)
|
||||
expect(isSessionApiFormat('unknown')).toBe(false)
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,21 @@
|
||||
export type SessionApiFormat = 'anthropic' | 'openai_chat' | 'openai_responses'
|
||||
export type SessionProtocolState = SessionApiFormat | 'mixed' | 'unknown'
|
||||
|
||||
export function isSessionApiFormat(value: unknown): value is SessionApiFormat {
|
||||
return value === 'anthropic' || value === 'openai_chat' || value === 'openai_responses'
|
||||
}
|
||||
|
||||
/** Resolve the upstream wire protocol, never the model name or internal envelope. */
|
||||
export function resolveProviderApiFormat(
|
||||
providerId: string | null | undefined,
|
||||
provider?: { apiFormat?: string; runtimeKind?: string },
|
||||
): SessionApiFormat | undefined {
|
||||
if (providerId === null || providerId === 'claude-official') return 'anthropic'
|
||||
if (providerId === 'openai-official' || providerId === 'grok-official') return 'openai_responses'
|
||||
if (!provider) return undefined
|
||||
if (provider.runtimeKind === 'openai_oauth' || provider.runtimeKind === 'grok_oauth') {
|
||||
return 'openai_responses'
|
||||
}
|
||||
if (provider.apiFormat === undefined) return 'anthropic'
|
||||
return isSessionApiFormat(provider.apiFormat) ? provider.apiFormat : undefined
|
||||
}
|
||||
@@ -221,6 +221,45 @@ describe('stripSignatureBlocksAfterModelChange', () => {
|
||||
).toBe(messages)
|
||||
})
|
||||
|
||||
test('cleans restored mixed history even after the selected model has replied', () => {
|
||||
const gptThinking = assistant('gpt-response', [
|
||||
{ type: 'redacted_thinking', data: 'cc-haha:openai-reasoning:v1:fixture' },
|
||||
])
|
||||
gptThinking.message.model = 'gpt-luna'
|
||||
const gptTool = assistant('gpt-response', [toolUse('read-gpt')])
|
||||
gptTool.message.model = 'gpt-luna'
|
||||
const deepseek = assistant('deepseek-response', [
|
||||
{ type: 'thinking', thinking: 'Current model reasoning', signature: 'deepseek-signature' },
|
||||
{ type: 'text', text: 'DeepSeek reply' },
|
||||
])
|
||||
deepseek.message.model = 'deepseek-v4-flash'
|
||||
const history = [
|
||||
createUserMessage({ content: 'Start' }),
|
||||
gptThinking, gptTool, toolResult('read-gpt'),
|
||||
deepseek, createUserMessage({ content: 'Continue next turn' }),
|
||||
]
|
||||
|
||||
const cleaned = stripSignatureBlocksAfterModelChange(history, 'deepseek-v4-flash')
|
||||
const normalized = normalizeMessagesForAPI(cleaned)
|
||||
const replies = normalized.filter((msg): msg is AssistantMessage => msg.type === 'assistant')
|
||||
|
||||
expect(replies.flatMap(msg => msg.message.content).some(block => block.type === 'redacted_thinking')).toBe(false)
|
||||
expect(replies.find(msg => msg.message.id === 'gpt-response')?.message.content).toEqual([toolUse('read-gpt')])
|
||||
expect(cleaned[4]).toBe(deepseek)
|
||||
expect(normalized.some(msg => msg.type === 'user' && Array.isArray(msg.message.content)
|
||||
&& msg.message.content.some(block => block.type === 'tool_result' && block.tool_use_id === 'read-gpt'))).toBe(true)
|
||||
expect(gptThinking.message.content[0]?.type).toBe('redacted_thinking')
|
||||
expect(stripSignatureBlocksAfterModelChange(cleaned, 'deepseek-v4-flash')).toBe(cleaned)
|
||||
|
||||
// Switching back also checks every historical source, preserving this model's
|
||||
// original encrypted block without replaying DeepSeek's signed thinking.
|
||||
const switchedBack = stripSignatureBlocksAfterModelChange(history, 'gpt-luna')
|
||||
expect(switchedBack[1]).toBe(gptThinking)
|
||||
expect(switchedBack[4]?.type === 'assistant' && switchedBack[4].message.content).toEqual([
|
||||
{ type: 'text', text: 'DeepSeek reply' },
|
||||
])
|
||||
})
|
||||
|
||||
test('leaves history without protected thinking untouched', () => {
|
||||
const messages = [createUserMessage({ content: 'Continue' })]
|
||||
|
||||
|
||||
+19
-13
@@ -5329,23 +5329,29 @@ export function stripSignatureBlocksAfterModelChange(
|
||||
messages: Message[],
|
||||
currentModel: string,
|
||||
): Message[] {
|
||||
const signatureSource = messages.findLast(msg => (
|
||||
msg.type === 'assistant' &&
|
||||
msg.message.model !== SYNTHETIC_MODEL &&
|
||||
msg.message.content.some(isThinkingBlock)
|
||||
))
|
||||
if (signatureSource?.type !== 'assistant' || !signatureSource.message.model) {
|
||||
return messages
|
||||
}
|
||||
|
||||
const normalize = (model: string) => normalizeModelStringForAPI(
|
||||
parseUserSpecifiedModel(model),
|
||||
).trim().toLowerCase()
|
||||
const targetModel = normalize(currentModel)
|
||||
let changed = false
|
||||
|
||||
if (normalize(signatureSource.message.model) === normalize(currentModel)) {
|
||||
return messages
|
||||
}
|
||||
return stripSignatureBlocks(messages)
|
||||
// A new turn reloads the original transcript, including blocks cleaned only
|
||||
// in memory on the previous turn. Inspect every source: the latest reply can
|
||||
// already match the target while older, incompatible blocks remain.
|
||||
const result = messages.map(msg => {
|
||||
if (
|
||||
msg.type !== 'assistant' ||
|
||||
!msg.message.model ||
|
||||
msg.message.model === SYNTHETIC_MODEL ||
|
||||
normalize(msg.message.model) === targetModel
|
||||
) return msg
|
||||
|
||||
const [cleaned] = stripSignatureBlocks([msg])
|
||||
if (cleaned !== msg) changed = true
|
||||
return cleaned!
|
||||
})
|
||||
|
||||
return changed ? result : messages
|
||||
}
|
||||
|
||||
/**
|
||||
|
||||
@@ -0,0 +1,84 @@
|
||||
import { afterEach, describe, expect, it } from 'bun:test'
|
||||
import { mkdtemp, readFile, rm, writeFile } from 'node:fs/promises'
|
||||
import { tmpdir } from 'node:os'
|
||||
import { join } from 'node:path'
|
||||
import { inferSessionApiFormat } from '../server/services/sessionProtocolHistory.js'
|
||||
import { createSessionBranch } from './sessionBranching.js'
|
||||
|
||||
describe('branch protocol inheritance', () => {
|
||||
const tempDirs: string[] = []
|
||||
|
||||
afterEach(async () => {
|
||||
for (const directory of tempDirs.splice(0)) await rm(directory, { recursive: true, force: true })
|
||||
})
|
||||
|
||||
async function branchFixture(options: {
|
||||
firstProvider?: string | null
|
||||
branchAtFirst?: boolean
|
||||
lock?: 'anthropic' | 'openai_responses'
|
||||
}) {
|
||||
const directory = await mkdtemp(join(tmpdir(), 'branch-protocol-'))
|
||||
tempDirs.push(directory)
|
||||
const sourceSessionId = crypto.randomUUID()
|
||||
const sourceTranscriptPath = join(directory, `${sourceSessionId}.jsonl`)
|
||||
let parentUuid: string | null = null
|
||||
let sequence = 0
|
||||
function message(type: 'user' | 'assistant', content: string) {
|
||||
const uuid = crypto.randomUUID()
|
||||
const result = {
|
||||
type, uuid, parentUuid, sessionId: sourceSessionId, isSidechain: false,
|
||||
timestamp: new Date(Date.UTC(2026, 0, 1, 0, 0, sequence++)).toISOString(),
|
||||
cwd: directory,
|
||||
message: { role: type, content },
|
||||
}
|
||||
parentUuid = uuid
|
||||
return result
|
||||
}
|
||||
const user = message('user', 'first prompt')
|
||||
const firstReply = message('assistant', 'first reply')
|
||||
const entries = [
|
||||
{ type: 'session-meta', workDir: directory, unknownField: { preserve: true },
|
||||
...(options.firstProvider !== undefined ? { runtimeProviderId: options.firstProvider } : {}) },
|
||||
user,
|
||||
firstReply,
|
||||
{ type: 'session-meta', runtimeProviderId: 'openai-official', runtimeModelId: 'same-model-name' },
|
||||
message('user', 'later prompt'),
|
||||
message('assistant', 'later reply'),
|
||||
...(options.lock ? [{ type: 'session-meta', sessionApiFormat: options.lock }] : []),
|
||||
]
|
||||
const original = entries.map(entry => JSON.stringify(entry)).join('\n') + '\n'
|
||||
await writeFile(sourceTranscriptPath, original)
|
||||
const branch = await createSessionBranch({
|
||||
sourceSessionId, sourceTranscriptPath,
|
||||
...(options.branchAtFirst ? { targetMessageId: firstReply.uuid } : {}),
|
||||
})
|
||||
const copied = (await readFile(branch.forkPath, 'utf8')).trim().split('\n').map(line => JSON.parse(line))
|
||||
expect(await readFile(sourceTranscriptPath, 'utf8')).toBe(original)
|
||||
expect(copied.find(entry => entry.unknownField)?.unknownField).toEqual({ preserve: true })
|
||||
return copied
|
||||
}
|
||||
|
||||
it('inherits the selected earlier protocol instead of a later provider selection', async () => {
|
||||
const entries = await branchFixture({ firstProvider: null, branchAtFirst: true })
|
||||
expect(inferSessionApiFormat(entries)).toBe('anthropic')
|
||||
expect(entries.some(entry => entry.type === 'session-meta' && entry.sessionApiFormat === 'anthropic')).toBe(true)
|
||||
expect(entries.filter(entry => entry.type === 'assistant')).toHaveLength(1)
|
||||
})
|
||||
|
||||
it('preserves mixed legacy history instead of reclassifying every reply as the last protocol', async () => {
|
||||
const entries = await branchFixture({ firstProvider: null })
|
||||
expect(inferSessionApiFormat(entries)).toBe('mixed')
|
||||
expect(entries.some(entry => entry.type === 'session-meta' && entry.sessionApiFormat === 'mixed')).toBe(true)
|
||||
})
|
||||
|
||||
it('preserves unknown legacy history instead of assigning a later known provider retroactively', async () => {
|
||||
const entries = await branchFixture({})
|
||||
expect(inferSessionApiFormat(entries)).toBe('unknown')
|
||||
expect(entries.some(entry => entry.type === 'session-meta' && entry.sessionApiFormat === 'unknown')).toBe(true)
|
||||
})
|
||||
|
||||
it('keeps an explicit parent lock when branching at an earlier message', async () => {
|
||||
const entries = await branchFixture({ branchAtFirst: true, lock: 'openai_responses' })
|
||||
expect(inferSessionApiFormat(entries)).toBe('openai_responses')
|
||||
})
|
||||
})
|
||||
@@ -12,6 +12,7 @@ import { parseJSONL } from './json.js'
|
||||
import { buildConversationChain, loadTranscriptFile } from './sessionStorage.js'
|
||||
import { jsonStringify } from './slowOperations.js'
|
||||
import { escapeRegExp } from './stringUtils.js'
|
||||
import { inferSessionApiFormat } from '../server/services/sessionProtocolHistory.js'
|
||||
|
||||
type SessionMetaEntry = {
|
||||
type: 'session-meta'
|
||||
@@ -452,6 +453,27 @@ export async function createSessionBranch(
|
||||
)
|
||||
}
|
||||
|
||||
// Provider snapshots describe the messages that follow them. Hoisting every
|
||||
// session-meta record to the header would assign the final provider to all
|
||||
// inherited replies. Keep those records next to the copied message that
|
||||
// originally followed them, while preserving the active chain's message order.
|
||||
const metadataBeforeMessage = new Map<string, SessionMetaEntry[]>()
|
||||
let trailingSessionMetadata: SessionMetaEntry[] = []
|
||||
const inheritedProtocolEntries: Parameters<typeof inferSessionApiFormat>[0] = []
|
||||
for (const entry of sourceEntries) {
|
||||
if (isSessionMetaEntry(entry)) {
|
||||
trailingSessionMetadata.push(entry)
|
||||
inheritedProtocolEntries.push(entry)
|
||||
} else if (isTranscriptEntry(entry) && sourceMessageEntriesById.get(entry.uuid) === entry) {
|
||||
metadataBeforeMessage.set(entry.uuid, trailingSessionMetadata)
|
||||
trailingSessionMetadata = []
|
||||
inheritedProtocolEntries.push(entry)
|
||||
}
|
||||
}
|
||||
// Include explicit parent locks even if they were appended after the target.
|
||||
// Later provider selections alone do not count as using that protocol.
|
||||
const inheritedApiFormat = inferSessionApiFormat(inheritedProtocolEntries)
|
||||
|
||||
const copiedToolResultIds = new Set(
|
||||
copiedMessages.flatMap((message) => extractToolResultIds(message)),
|
||||
)
|
||||
@@ -468,6 +490,7 @@ export async function createSessionBranch(
|
||||
let parentUuid: UUID | null = null
|
||||
|
||||
for (const entry of branchMessageEntries) {
|
||||
messageLines.push(...(metadataBeforeMessage.get(entry.uuid) ?? []).map(metadata => jsonStringify(metadata)))
|
||||
const forkedEntry: TranscriptEntry = {
|
||||
...entry,
|
||||
sessionId: forkSessionId,
|
||||
@@ -512,10 +535,24 @@ export async function createSessionBranch(
|
||||
)
|
||||
|
||||
const lines = [
|
||||
...metadataEntries.map((entry) => jsonStringify(entry)),
|
||||
// Only synthesized session metadata belongs in the header; source records
|
||||
// retain their relationship to the inherited messages above.
|
||||
...metadataEntries
|
||||
.filter(entry => !isSessionMetaEntry(entry) || !sourceEntries.includes(entry))
|
||||
.map((entry) => jsonStringify(entry)),
|
||||
...messageLines,
|
||||
...trailingSessionMetadata.map(entry => jsonStringify(entry)),
|
||||
]
|
||||
|
||||
if (inheritedApiFormat) {
|
||||
lines.push(jsonStringify({
|
||||
type: 'session-meta',
|
||||
isMeta: true,
|
||||
sessionApiFormat: inheritedApiFormat,
|
||||
timestamp: new Date().toISOString(),
|
||||
}))
|
||||
}
|
||||
|
||||
if (contentReplacementRecords.length > 0) {
|
||||
lines.push(jsonStringify({
|
||||
type: 'content-replacement',
|
||||
|
||||
Reference in New Issue
Block a user