fix(chat): stabilize context usage across model switches

This commit is contained in:
程序员阿江(Relakkes)
2026-08-03 03:03:08 +08:00
parent 0f1ebf8d1f
commit 2262973a48
11 changed files with 172 additions and 12 deletions
@@ -188,6 +188,61 @@ describe('ContextUsageIndicator request behavior', () => {
expect(sessionsApiMock.getInspection).toHaveBeenCalledTimes(2)
})
it('keeps the last context visible while the switched runtime is still starting', async () => {
const nextInspection = deferred<typeof baseInspection>()
sessionsApiMock.getInspection
.mockResolvedValueOnce(baseInspection)
.mockReturnValueOnce(nextInspection.promise)
const { rerender } = render(
<ContextUsageIndicator
sessionId="session-1"
chatState="idle"
messageCount={1}
runtimeSelectionKey="deepseek:deepseek-chat"
fallbackModelLabel="deepseek-chat"
/>,
)
await waitFor(() => {
expect(screen.getByTestId('context-usage-indicator')).toHaveTextContent('21%')
})
rerender(
<ContextUsageIndicator
sessionId="session-1"
chatState="idle"
messageCount={1}
runtimeSelectionKey="deepseek:deepseek-reasoner"
fallbackModelLabel="deepseek-reasoner"
/>,
)
// A runtime switch restarts the CLI. Keep the last confirmed percentage
// in place instead of replacing it with a long-running spinner while the
// replacement control channel comes online.
expect(screen.getByTestId('context-usage-indicator')).toHaveTextContent('21%')
expect(screen.queryByLabelText('Context usage loading')).not.toBeInTheDocument()
expect(screen.getByText('deepseek-reasoner')).toBeInTheDocument()
await act(async () => {
nextInspection.resolve({
...baseInspection,
status: { ...baseInspection.status, model: 'deepseek-reasoner' },
context: {
...baseInspection.context,
model: 'deepseek-reasoner',
percentage: 12,
},
})
await nextInspection.promise
})
await waitFor(() => {
expect(screen.getByTestId('context-usage-indicator')).toHaveTextContent('12%')
})
})
it('ignores a stale inspection response after the runtime identity changes', async () => {
const first = deferred<typeof baseInspection>()
sessionsApiMock.getInspection
@@ -98,6 +98,7 @@ export function ContextUsageIndicator({
const [mobileDetailsOpen, setMobileDetailsOpen] = useState(false)
const requestSeq = useRef(0)
const contextIdentityRef = useRef('')
const contextDataIdentityRef = useRef('')
const inFlightRequestRef = useRef<Promise<boolean> | null>(null)
const inFlightIdentityRef = useRef<string | null>(null)
const lastAutoRefreshAtRef = useRef(0)
@@ -140,6 +141,7 @@ export function ContextUsageIndicator({
const nextContext = inspection.context ?? inspection.contextEstimate ?? null
const nextSource = inspection.context ? 'live' : inspection.contextEstimate ? 'estimate' : null
const usageModel = inspection.usage?.models.find((model) => firstNonEmpty(model.displayName, model.model)) ?? null
contextDataIdentityRef.current = nextContext ? activeContextIdentity : ''
setContext(nextContext)
setContextSource(nextSource)
setInspectionModel(firstNonEmpty(
@@ -155,6 +157,12 @@ export function ContextUsageIndicator({
})
.catch((err) => {
if (seq !== requestSeq.current || activeContextIdentity !== contextIdentityRef.current) return false
if (contextDataIdentityRef.current !== activeContextIdentity) {
contextDataIdentityRef.current = ''
setContext(null)
setContextSource(null)
setUpdatedAt(null)
}
setError(err instanceof Error ? err.message : String(err))
return false
})
@@ -199,10 +207,7 @@ export function ContextUsageIndicator({
if (identityChanged) {
requestSeq.current += 1
lastAutoRefreshAtRef.current = 0
setContext(null)
setContextSource(null)
setError(null)
setUpdatedAt(null)
setInspectionModel(null)
}
void refresh('auto')
@@ -251,7 +256,12 @@ export function ContextUsageIndicator({
: 'var(--color-surface-container-high)',
}
const displayPercent = displayContext ? formatPercent(percentage) : '--'
const displayModel = firstNonEmpty(context?.model, inspectionModel, fallbackModelLabel)
const currentContextIdentity = `${sessionId}:${runtimeSelectionKey}`
const isContextFromPreviousRuntime = Boolean(displayContext) &&
contextDataIdentityRef.current !== currentContextIdentity
const displayModel = isContextFromPreviousRuntime
? firstNonEmpty(fallbackModelLabel, inspectionModel, context?.model)
: firstNonEmpty(context?.model, inspectionModel, fallbackModelLabel)
const ariaLabel = displayContext
? t('contextIndicator.ariaLabel', { percent: formatPercent(percentage) })
: isPendingContext
+2
View File
@@ -2301,6 +2301,8 @@ Row 9, all 8 cells: continuing from straight down, turning left through lower-le
'session.active': 'session active',
'session.lastUpdated': 'last updated {time}',
'session.messages': '{count} messages',
'session.apiTokens': '{count} API tokens',
'session.apiTokenBreakdown': 'API usage: {total} tokens · input {input} · output {output} · cache {cache}',
'session.historyLoadFailed': 'Failed to load session history.',
'session.workspaceUnavailable': 'Workspace unavailable: {dir}',
'session.worktreeRemoved': 'The temporary workspace was cleaned up. History is still available; start a new session in {dir} to continue.',
+2
View File
@@ -2303,6 +2303,8 @@ export const jp: Record<TranslationKey, string> = {
'session.active': 'セッションがアクティブ',
'session.lastUpdated': '最終更新 {time}',
'session.messages': '{count} 件のメッセージ',
'session.apiTokens': '{count} API tokens',
'session.apiTokenBreakdown': 'API 使用量: {total} tokens · 入力 {input} · 出力 {output} · キャッシュ {cache}',
'session.historyLoadFailed': 'セッション履歴の読み込みに失敗しました。',
'session.workspaceUnavailable': 'ワークスペースが利用できません: {dir}',
'session.worktreeRemoved': '一時ワークスペースはクリーンアップされました。履歴は引き続き閲覧できます。{dir} で新しいセッションを開始してください。',
+2
View File
@@ -2303,6 +2303,8 @@ export const kr: Record<TranslationKey, string> = {
'session.active': '세션 활성',
'session.lastUpdated': '마지막 업데이트 {time}',
'session.messages': '{count}개의 메시지',
'session.apiTokens': '{count} API tokens',
'session.apiTokenBreakdown': 'API 사용량: {total} tokens · 입력 {input} · 출력 {output} · 캐시 {cache}',
'session.historyLoadFailed': '세션 기록을 불러오지 못했습니다.',
'session.workspaceUnavailable': '작업 공간을 사용할 수 없습니다: {dir}',
'session.worktreeRemoved': '임시 작업 공간이 정리되었습니다. 기록은 계속 볼 수 있습니다. {dir}에서 새 세션을 시작해 계속하세요.',
+2
View File
@@ -2302,6 +2302,8 @@ export const zh: Record<TranslationKey, string> = {
'session.active': '會話活躍中',
'session.lastUpdated': '最後更新 {time}',
'session.messages': '{count} 條訊息',
'session.apiTokens': '{count} API tokens',
'session.apiTokenBreakdown': 'API 用量:{total} tokens · 輸入 {input} · 輸出 {output} · 快取 {cache}',
'session.historyLoadFailed': '歷史會話載入失敗。',
'session.workspaceUnavailable': '工作目錄不可用: {dir}',
'session.worktreeRemoved': '臨時工作區已清理,歷史記錄仍可查看。請在原專案 {dir} 中新建會話繼續。',
+2
View File
@@ -2302,6 +2302,8 @@ export const zh: Record<TranslationKey, string> = {
'session.active': '会话活跃中',
'session.lastUpdated': '最后更新 {time}',
'session.messages': '{count} 条消息',
'session.apiTokens': '{count} API tokens',
'session.apiTokenBreakdown': 'API 用量:{total} tokens · 输入 {input} · 输出 {output} · 缓存 {cache}',
'session.historyLoadFailed': '历史会话加载失败。',
'session.workspaceUnavailable': '工作目录不可用: {dir}',
'session.worktreeRemoved': '临时工作区已清理,历史记录仍可查看。请在原项目 {dir} 中新建会话继续。',
+3 -3
View File
@@ -266,8 +266,8 @@ describe('ActiveSession task polling', () => {
render(<ActiveSession />)
const tokenBadge = screen.getByTitle(/1,500/)
expect(tokenBadge).toHaveTextContent('1.5k')
const tokenBadge = screen.getByTitle(/cache 1,500/i)
expect(tokenBadge).toHaveTextContent('1.5k API tokens')
})
it('shows a loading state for historical sessions while messages are loading', () => {
@@ -2034,7 +2034,7 @@ describe('ActiveSession header', () => {
// 元数据挤在标题右侧时会离标题很远,读起来像飘在角落的另一块内容。
expect(within(titleRow).queryByText('2 messages')).not.toBeInTheDocument()
expect(within(meta).getByText('2 messages')).toBeInTheDocument()
expect(within(meta).getByText('15k tokens')).toBeInTheDocument()
expect(within(meta).getByText('15k API tokens')).toBeInTheDocument()
expect(header).toHaveClass('py-3')
})
+13 -2
View File
@@ -408,6 +408,8 @@ export function ActiveSession() {
(trackedTaskSessionId === activeTabId && hasRunningTasks) ||
hasRunningBackgroundTasks
const totalTokens = getTokenUsageTotal(tokenUsage)
const cachedTokens = (tokenUsage.cache_read_tokens ?? 0) +
(tokenUsage.cache_creation_tokens ?? 0)
const activityTeamMembers = useMemo(() => {
if (!activeTeam || activeTeam.leadSessionId !== activeTabId) return []
return activeTeam.members.filter((member) =>
@@ -661,8 +663,17 @@ export function ActiveSession() {
</span>
),
totalTokens > 0 && (
<span key="tokens" className="shrink-0" title={t('common.tokens', { count: totalTokens.toLocaleString() })}>
{t('common.tokens', { count: formatTokenCount(totalTokens) })}
<span
key="tokens"
className="shrink-0"
title={t('session.apiTokenBreakdown', {
total: totalTokens.toLocaleString(),
input: tokenUsage.input_tokens.toLocaleString(),
output: tokenUsage.output_tokens.toLocaleString(),
cache: cachedTokens.toLocaleString(),
})}
>
{t('session.apiTokens', { count: formatTokenCount(totalTokens) })}
</span>
),
lastUpdated && (
@@ -226,6 +226,47 @@ describe('ConversationService', () => {
expect(removeAbortListener).toHaveBeenCalledWith('abort', expect.any(Function))
})
it('should reject an in-flight control request when its CLI session is stopped', async () => {
const svc = new ConversationService()
const sid = crypto.randomUUID()
const sent: unknown[] = []
const session: any = {
proc: { kill() {}, exited: Promise.resolve(0) },
outputCallbacks: [],
workDir: process.cwd(),
permissionMode: 'default',
sdkToken: 'token',
sdkSocket: {
send(data: string) {
sent.push(JSON.parse(data))
},
},
pendingOutbound: [],
startupPending: false,
startupExitCode: null,
stdoutLines: [],
stderrLines: [],
outputDrain: Promise.resolve(),
sdkMessages: [],
initMessage: null,
pendingPermissionRequests: new Map(),
}
;(svc as any).sessions.set(sid, session)
const request = svc.requestControl(
sid,
{ subtype: 'get_context_usage' },
50,
)
await new Promise((resolve) => setTimeout(resolve, 0))
expect(sent).toHaveLength(1)
svc.stopSession(sid)
await expect(request).rejects.toThrow('CLI session stopped')
expect(session.outputCallbacks).toHaveLength(0)
})
it('should ignore a stale SDK disconnect after a replacement socket attaches', () => {
const svc = new ConversationService()
const sessionId = crypto.randomUUID()
+36 -3
View File
@@ -223,6 +223,7 @@ type SessionProcess = {
permissionSuggestions?: unknown[]
}
>
pendingControlRequests: Map<string, (reason: Error) => void>
}
export type PendingPermissionRequest = {
@@ -480,6 +481,7 @@ export class ConversationService {
usesOfficialOAuth,
officialOAuthToken: childEnv.CLAUDE_CODE_OAUTH_TOKEN ?? null,
pendingPermissionRequests: new Map(),
pendingControlRequests: new Map(),
}
this.sessions.set(sessionId, session)
@@ -856,6 +858,10 @@ export class ConversationService {
const startedAt = Date.now()
await this.waitForControlChannelReady(sessionId, timeoutMs, signal)
const session = this.sessions.get(sessionId)
if (!session) {
throw new Error('CLI session is not running')
}
const responseTimeoutMs = Math.max(1, timeoutMs - (Date.now() - startedAt))
const requestId = crypto.randomUUID()
return new Promise((resolve, reject) => {
@@ -867,7 +873,8 @@ export class ConversationService {
settled = true
clearTimeout(timeout)
signal?.removeEventListener('abort', handleAbort)
this.removeOutputCallback(sessionId, handleOutput)
session.outputCallbacks = session.outputCallbacks.filter((entry) => entry !== handleOutput)
session.pendingControlRequests?.delete(requestId)
fn()
}
@@ -900,13 +907,18 @@ export class ConversationService {
`Timed out waiting for ${String(request.subtype ?? 'control')} response`,
)))
}, responseTimeoutMs)
this.onOutput(sessionId, handleOutput)
session.outputCallbacks.push(handleOutput)
const pendingControlRequests = session.pendingControlRequests ?? new Map<string, (reason: Error) => void>()
session.pendingControlRequests = pendingControlRequests
pendingControlRequests.set(requestId, (reason) => {
finish(() => reject(reason))
})
signal?.addEventListener('abort', handleAbort, { once: true })
if (signal?.aborted) {
handleAbort()
return
}
const sent = this.sendSdkMessage(sessionId, {
const sent = this.sessions.get(sessionId) === session && this.sendSdkMessage(sessionId, {
type: 'control_request',
request_id: requestId,
request,
@@ -1214,10 +1226,23 @@ export class ConversationService {
return `${truncated}\n[truncated]`
}
private cancelPendingControlRequests(
session: SessionProcess,
reason = new Error('CLI session stopped'),
): void {
const pending = session.pendingControlRequests
if (!pending || pending.size === 0) return
for (const cancel of [...pending.values()]) {
cancel(reason)
}
pending.clear()
}
stopSession(sessionId: string): void {
const session = this.sessions.get(sessionId)
if (!session) return
this.cancelPendingControlRequests(session)
this.sessions.delete(sessionId)
this.killProcess(sessionId, session)
}
@@ -1229,6 +1254,7 @@ export class ConversationService {
const session = this.sessions.get(sessionId)
if (!session) return
this.cancelPendingControlRequests(session)
this.sessions.delete(sessionId)
await this.stopProcessAndWait(sessionId, session, timeoutMs)
}
@@ -1245,6 +1271,9 @@ export class ConversationService {
const activeSessions = Array.from(this.sessions.entries())
if (activeSessions.length === 0) return
for (const [, session] of activeSessions) {
this.cancelPendingControlRequests(session)
}
this.sessions.clear()
await Promise.all(
activeSessions.map(([sessionId, session]) =>
@@ -1398,6 +1427,10 @@ export class ConversationService {
const activeSession = this.sessions.get(sessionId)
if (activeSession?.proc === proc) {
this.cancelPendingControlRequests(
activeSession,
new Error('CLI session exited before the control request completed'),
)
if (activeSession.startupPending) {
activeSession.startupExitCode = code
return