fix(cli): backport Claude Code stability fixes

This commit is contained in:
程序员阿江(Relakkes)
2026-07-29 09:35:19 +08:00
parent 07a06227ca
commit b21c63cff3
26 changed files with 2598 additions and 163 deletions
+154
View File
@@ -0,0 +1,154 @@
import { describe, expect, test } from 'bun:test'
import type {
SDKAssistantMessage,
SDKCompactBoundaryMessage,
SDKMessage,
SDKUserMessage,
} from 'src/entrypoints/agentSdkTypes.js'
import { PrintPartialOutputTracker } from './partialOutput.js'
describe('PrintPartialOutputTracker', () => {
test('combines completed text blocks from the failed response', () => {
const tracker = new PrintPartialOutputTracker()
tracker.observe(assistant('response-a', 'first'))
tracker.observe(assistant('response-a', ' second'))
tracker.observe(assistant('synthetic-error', 'API Error: failed', 'unknown'))
expect(tracker.formatResult('API Error: failed', true)).toBe(
'first second\nAPI Error: failed',
)
})
test('does not duplicate normal results or the same error text', () => {
const tracker = new PrintPartialOutputTracker()
tracker.observe(assistant('response-a', 'complete'))
expect(tracker.formatResult('complete', false)).toBe('complete')
expect(tracker.formatResult('complete', true)).toBe('complete')
expect(tracker.formatResultLine('complete', false)).toBe('complete\n')
expect(tracker.formatResultLine('complete\n', false)).toBe('complete\n')
})
test.each([
['a user turn', user()],
['a compact boundary', compactBoundary()],
])('resets accumulated text at %s', (_name, boundary) => {
const tracker = new PrintPartialOutputTracker()
tracker.observe(assistant('response-a', 'stale'))
tracker.observe(boundary)
expect(tracker.formatResult('API Error: failed', true)).toBe(
'API Error: failed',
)
})
test('keeps only the newest assistant response', () => {
const tracker = new PrintPartialOutputTracker()
tracker.observe(assistant('response-a', 'stale'))
tracker.observe(assistant('response-b', 'current'))
tracker.observe(assistant('response-error', 'API Error: failed', 'unknown'))
expect(tracker.formatResult('API Error: failed', true)).toBe(
'current\nAPI Error: failed',
)
})
test('keeps the prior partial when a new response contains an untagged terminal error', () => {
const tracker = new PrintPartialOutputTracker()
tracker.observe(assistant('response-a', 'partial'))
tracker.observe(assistant('response-b', 'Image was too large to process.'))
expect(
tracker.formatResult('Image was too large to process.', true),
).toBe('partial\nImage was too large to process.')
})
test('does not duplicate identical prior and terminal text', () => {
const tracker = new PrintPartialOutputTracker()
tracker.observe(assistant('response-a', 'same'))
tracker.observe(assistant('response-b', 'same'))
expect(tracker.formatResult('same', true)).toBe('same')
})
test('does not combine two known error responses', () => {
const tracker = new PrintPartialOutputTracker()
tracker.observe(assistant('response-a', 'API Error: first', 'unknown'))
tracker.observe(assistant('response-b', 'API Error: final', 'unknown'))
expect(tracker.formatResult('API Error: final', true)).toBe(
'API Error: final',
)
})
test('ignores nested assistant messages', () => {
const tracker = new PrintPartialOutputTracker()
const nested = assistant('response-a', 'nested')
nested.parent_tool_use_id = 'tool-parent'
tracker.observe(nested)
expect(tracker.formatResult('API Error: failed', true)).toBe(
'API Error: failed',
)
})
})
function assistant(
id: string,
text: string,
error?: SDKAssistantMessage['error'],
): SDKAssistantMessage {
return {
type: 'assistant',
message: {
id,
type: 'message',
role: 'assistant',
model: 'test-model',
content: [{ type: 'text', text, citations: null }],
stop_reason: null,
stop_sequence: null,
usage: {
input_tokens: 0,
output_tokens: 0,
cache_creation_input_tokens: 0,
cache_read_input_tokens: 0,
},
},
parent_tool_use_id: null,
session_id: 'test-session',
uuid: crypto.randomUUID(),
error,
}
}
function user(): SDKUserMessage {
return {
type: 'user',
message: { role: 'user', content: 'next turn' },
parent_tool_use_id: null,
session_id: 'test-session',
uuid: crypto.randomUUID(),
}
}
function compactBoundary(): SDKMessage {
return {
type: 'system',
subtype: 'compact_boundary',
session_id: 'test-session',
uuid: crypto.randomUUID(),
compact_metadata: {
trigger: 'auto',
pre_tokens: 1,
},
} as SDKCompactBoundaryMessage
}
+67
View File
@@ -0,0 +1,67 @@
import type { SDKMessage } from 'src/entrypoints/agentSdkTypes.js'
import { extractTextContent } from 'src/utils/messages.js'
export class PrintPartialOutputTracker {
private responseId: string | undefined
private responseText = ''
private responseIsKnownError = false
private previousResponseText = ''
private previousResponseIsKnownError = false
observe(message: SDKMessage): void {
if (
(message.type === 'system' &&
message.subtype === 'compact_boundary') ||
(message.type === 'user' && message.parent_tool_use_id === null)
) {
this.reset()
return
}
if (
message.type !== 'assistant' ||
message.parent_tool_use_id !== null
) {
return
}
const text = extractTextContent(message.message.content)
if (this.responseId !== message.message.id) {
this.previousResponseText = this.responseText
this.previousResponseIsKnownError = this.responseIsKnownError
this.responseId = message.message.id
this.responseText = ''
this.responseIsKnownError = false
}
this.responseText += text
this.responseIsKnownError ||=
message.error !== undefined || text.startsWith('API Error:')
}
formatResult(result: string, isError: boolean): string {
if (!isError) {
return result
}
const partial =
this.responseText === result &&
!this.previousResponseIsKnownError &&
this.previousResponseText &&
this.previousResponseText !== result
? this.previousResponseText
: ''
return partial ? `${partial}\n${result}` : result
}
formatResultLine(result: string, isError: boolean): string {
const formatted = this.formatResult(result, isError)
return formatted.endsWith('\n') ? formatted : `${formatted}\n`
}
private reset(): void {
this.responseId = undefined
this.responseText = ''
this.responseIsKnownError = false
this.previousResponseText = ''
this.previousResponseIsKnownError = false
}
}
+241
View File
@@ -0,0 +1,241 @@
import { afterEach, describe, expect, test } from 'bun:test'
import { mkdtemp, rm } from 'fs/promises'
import { tmpdir } from 'os'
import { join } from 'path'
let configDir: string | null = null
afterEach(async () => {
if (configDir) {
await rm(configDir, { recursive: true, force: true })
configDir = null
}
})
describe('print mode partial output', () => {
test(
'keeps assistant text produced before a mid-stream API error',
async () => {
const server = Bun.serve({
hostname: '127.0.0.1',
port: 0,
fetch() {
return new Response(midStreamErrorResponse(), {
headers: { 'content-type': 'text/event-stream' },
})
},
})
configDir = await mkdtemp(join(tmpdir(), 'cc-haha-print-partial-'))
try {
const child = Bun.spawn(
['./bin/claude-haha', '--bare', '-p', 'Reply briefly'],
{
cwd: process.cwd(),
env: {
...process.env,
NODE_ENV: 'production',
CI: '1',
CC_HAHA_SKIP_DOTENV: '1',
CLAUDE_CONFIG_DIR: configDir,
CLAUDE_CODE_SKIP_PROMPT_HISTORY: '1',
CLAUDE_CODE_DISABLE_NONESSENTIAL_TRAFFIC: '1',
CLAUDE_CODE_DISABLE_NONSTREAMING_FALLBACK: '1',
CLAUDE_STREAM_TRANSIENT_RETRY_MAX: '0',
DISABLE_AUTOUPDATER: '1',
DISABLE_TELEMETRY: '1',
DISABLE_ERROR_REPORTING: '1',
ANTHROPIC_API_KEY: 'loopback-test-key',
ANTHROPIC_BASE_URL: `http://127.0.0.1:${server.port}`,
ANTHROPIC_MODEL: 'claude-sonnet-4-5',
CLAUDE_CODE_USE_BEDROCK: undefined,
CLAUDE_CODE_USE_VERTEX: undefined,
CLAUDE_CODE_USE_FOUNDRY: undefined,
ANTHROPIC_AUTH_TOKEN: undefined,
},
stdout: 'pipe',
stderr: 'pipe',
},
)
const [stdout, stderr, exitCode] = await Promise.all([
new Response(child.stdout).text(),
new Response(child.stderr).text(),
child.exited,
])
expect(exitCode).toBe(1)
expect(stderr).not.toContain('FIRST_PARTIAL_SENTINEL')
expect(stderr).not.toContain('SECOND_PARTIAL_SENTINEL')
expect(stdout).toContain('FIRST_PARTIAL_SENTINEL')
expect(stdout).toContain('SECOND_PARTIAL_SENTINEL')
expect(stdout).toContain('MIDSTREAM_SENTINEL_ERROR')
} finally {
server.stop(true)
}
},
30_000,
)
test(
'keeps completed text when a socket reset exhausts stream retries',
async () => {
const server = Bun.listen({
hostname: '127.0.0.1',
port: 0,
socket: {
open(socket) {
const payload = completedBlockWithoutMessageStop()
socket.write(
[
'HTTP/1.1 200 OK',
'Content-Type: text/event-stream',
'Transfer-Encoding: chunked',
'Connection: close',
'',
`${Buffer.byteLength(payload).toString(16)}\r\n${payload}\r\n`,
].join('\r\n'),
)
socket.flush()
setTimeout(() => socket.end(), 50)
},
data() {},
},
})
configDir = await mkdtemp(join(tmpdir(), 'cc-haha-print-transport-'))
try {
const child = Bun.spawn(
['./bin/claude-haha', '--bare', '-p', 'Reply briefly'],
{
cwd: process.cwd(),
env: {
...process.env,
NODE_ENV: 'production',
CI: '1',
CC_HAHA_SKIP_DOTENV: '1',
CLAUDE_CONFIG_DIR: configDir,
CLAUDE_CODE_SKIP_PROMPT_HISTORY: '1',
CLAUDE_CODE_DISABLE_NONESSENTIAL_TRAFFIC: '1',
CLAUDE_CODE_DISABLE_NONSTREAMING_FALLBACK: '1',
CLAUDE_STREAM_TRANSIENT_RETRY_MAX: '0',
DISABLE_AUTOUPDATER: '1',
DISABLE_TELEMETRY: '1',
DISABLE_ERROR_REPORTING: '1',
ANTHROPIC_API_KEY: 'loopback-test-key',
ANTHROPIC_BASE_URL: `http://127.0.0.1:${server.port}`,
ANTHROPIC_MODEL: 'claude-sonnet-4-5',
CLAUDE_CODE_USE_BEDROCK: undefined,
CLAUDE_CODE_USE_VERTEX: undefined,
CLAUDE_CODE_USE_FOUNDRY: undefined,
ANTHROPIC_AUTH_TOKEN: undefined,
},
stdout: 'pipe',
stderr: 'pipe',
},
)
const [stdout, exitCode] = await Promise.all([
new Response(child.stdout).text(),
child.exited,
])
expect(exitCode).toBe(1)
expect(stdout.match(/TRANSPORT_PARTIAL_SENTINEL/g)).toHaveLength(1)
expect(stdout).toMatch(/socket|connection|stream/i)
} finally {
server.stop(true)
}
},
30_000,
)
})
function sseEvent(name: string, data: unknown): string {
return `event: ${name}\ndata: ${JSON.stringify(data)}\n\n`
}
function midStreamErrorResponse(): string {
return [
sseEvent('message_start', {
type: 'message_start',
message: {
id: 'msg_partial_output',
type: 'message',
role: 'assistant',
model: 'claude-sonnet-4-5',
content: [],
stop_reason: null,
stop_sequence: null,
usage: { input_tokens: 1, output_tokens: 0 },
},
}),
sseEvent('content_block_start', {
type: 'content_block_start',
index: 0,
content_block: { type: 'text', text: '' },
}),
sseEvent('content_block_delta', {
type: 'content_block_delta',
index: 0,
delta: { type: 'text_delta', text: 'FIRST_PARTIAL_SENTINEL' },
}),
sseEvent('content_block_stop', {
type: 'content_block_stop',
index: 0,
}),
sseEvent('content_block_start', {
type: 'content_block_start',
index: 1,
content_block: { type: 'text', text: '' },
}),
sseEvent('content_block_delta', {
type: 'content_block_delta',
index: 1,
delta: { type: 'text_delta', text: 'SECOND_PARTIAL_SENTINEL' },
}),
sseEvent('content_block_stop', {
type: 'content_block_stop',
index: 1,
}),
sseEvent('error', {
type: 'error',
error: {
type: 'api_error',
message: 'MIDSTREAM_SENTINEL_ERROR',
},
}),
].join('')
}
function completedBlockWithoutMessageStop(): string {
return [
sseEvent('message_start', {
type: 'message_start',
message: {
id: 'msg_transport_partial',
type: 'message',
role: 'assistant',
model: 'claude-sonnet-4-5',
content: [],
stop_reason: null,
stop_sequence: null,
usage: { input_tokens: 1, output_tokens: 0 },
},
}),
sseEvent('content_block_start', {
type: 'content_block_start',
index: 0,
content_block: { type: 'text', text: '' },
}),
sseEvent('content_block_delta', {
type: 'content_block_delta',
index: 0,
delta: { type: 'text_delta', text: 'TRANSPORT_PARTIAL_SENTINEL' },
}),
sseEvent('content_block_stop', {
type: 'content_block_stop',
index: 0,
}),
].join('')
}
+8 -3
View File
@@ -262,6 +262,7 @@ import {
toSDKRateLimitInfo,
} from 'src/utils/messages/mappers.js'
import { createModelSwitchBreadcrumbs } from 'src/utils/messages.js'
import { PrintPartialOutputTracker } from './partialOutput.js'
import { collectContextData } from 'src/commands/context/context-noninteractive.js'
import { getSessionUsageSnapshot } from 'src/cost-tracker.js'
import { LOCAL_COMMAND_STDOUT_TAG } from 'src/constants/xml.js'
@@ -854,6 +855,7 @@ export async function runHeadless(
const needsFullArray = options.outputFormat === 'json' && options.verbose
const messages: SDKMessage[] = []
let lastMessage: SDKMessage | undefined
const partialOutputTracker = new PrintPartialOutputTracker()
// Streamlined mode transforms messages when CLAUDE_CODE_STREAMLINED_OUTPUT=true and using stream-json
// Build flag gates this out of external builds; env var is the runtime opt-in for ant builds
const transformToStreamlined =
@@ -878,6 +880,8 @@ export async function runHeadless(
options,
turnInterruptionState,
)) {
partialOutputTracker.observe(message)
if (transformToStreamlined) {
// Streamlined mode: transform messages and stream immediately
const transformed = transformToStreamlined(message)
@@ -938,9 +942,10 @@ export async function runHeadless(
switch (lastMessage.subtype) {
case 'success':
writeToStdout(
lastMessage.result.endsWith('\n')
? lastMessage.result
: lastMessage.result + '\n',
partialOutputTracker.formatResultLine(
lastMessage.result,
lastMessage.is_error,
),
)
break
case 'error_during_execution':
+193
View File
@@ -0,0 +1,193 @@
import { afterEach, describe, expect, mock, test } from 'bun:test'
import * as fsPromises from 'fs/promises'
import { tmpdir } from 'os'
import { join } from 'path'
const originalConfigDir = process.env.CLAUDE_CONFIG_DIR
let testConfigDir: string | null = null
afterEach(async () => {
mock.restore()
if (originalConfigDir === undefined) {
delete process.env.CLAUDE_CONFIG_DIR
} else {
process.env.CLAUDE_CONFIG_DIR = originalConfigDir
}
if (testConfigDir) {
await fsPromises.rm(testConfigDir, { recursive: true, force: true })
testConfigDir = null
}
})
describe('prompt history persistence', () => {
test('reconciles partial and fully committed append failures without duplicates', async () => {
const configDir = await fsPromises.mkdtemp(
join(tmpdir(), 'cc-haha-history-test-'),
)
testConfigDir = configDir
process.env.CLAUDE_CONFIG_DIR = configDir
let appendCalls = 0
let behaviorCalls = 0
let behavior:
| 'partial-then-success'
| 'full-then-error'
| 'rollback-fails'
| 'read-fails'
| 'unexpected-tail' = 'partial-then-success'
const realAppendFile = fsPromises.appendFile
const realReadFile = fsPromises.readFile
const realTruncate = fsPromises.truncate
mock.module('fs/promises', () => ({
...fsPromises,
appendFile: async (...args: Parameters<typeof fsPromises.appendFile>) => {
appendCalls += 1
behaviorCalls += 1
if (behavior === 'partial-then-success' && behaviorCalls === 1) {
const payload = Buffer.from(String(args[1]))
await realAppendFile(
args[0],
payload.subarray(0, Math.floor(payload.length / 2)),
{ mode: 0o600 },
)
throw new Error('injected partial append failure')
}
if (behavior === 'full-then-error' && behaviorCalls === 1) {
await realAppendFile(...args)
throw new Error('injected post-commit append failure')
}
if (behavior === 'rollback-fails' && behaviorCalls === 1) {
const payload = Buffer.from(String(args[1]))
await realAppendFile(
args[0],
payload.subarray(0, Math.floor(payload.length / 2)),
{ mode: 0o600 },
)
throw new Error('injected partial append failure')
}
if (behavior === 'read-fails' && behaviorCalls === 1) {
throw new Error('injected append failure before reconciliation read')
}
if (behavior === 'unexpected-tail' && behaviorCalls === 1) {
await realAppendFile(args[0], 'not-a-payload-prefix', {
mode: 0o600,
})
throw new Error('injected append with unexpected tail')
}
return realAppendFile(...args)
},
readFile: async (...args: Parameters<typeof fsPromises.readFile>) => {
if (behavior === 'read-fails' && behaviorCalls === 1) {
throw new Error('injected reconciliation read failure')
}
return realReadFile(...args)
},
truncate: async (...args: Parameters<typeof fsPromises.truncate>) => {
if (behavior === 'rollback-fails') {
throw new Error('injected rollback failure')
}
return realTruncate(...args)
},
}))
const history = await import('./history.js')
history.clearPendingHistoryEntries()
history.addToHistory('FIRST_SENTINEL_你好😀')
await waitFor(() => appendCalls === 1)
const historyPath = join(configDir, 'history.jsonl')
await waitFor(async () => {
const contents = await fsPromises
.readFile(historyPath, 'utf8')
.catch(() => '')
return contents.includes('FIRST_SENTINEL')
})
const contents = await fsPromises.readFile(historyPath, 'utf8')
expect(contents.match(/FIRST_SENTINEL/g)).toHaveLength(1)
behavior = 'full-then-error'
behaviorCalls = 0
history.addToHistory('SECOND_SENTINEL')
await waitFor(() => behaviorCalls === 1)
await Bun.sleep(600)
const reconciled = await fsPromises.readFile(historyPath, 'utf8')
expect(reconciled.match(/FIRST_SENTINEL/g)).toHaveLength(1)
expect(reconciled.match(/SECOND_SENTINEL/g)).toHaveLength(1)
for (const line of reconciled.trim().split('\n')) {
expect(() => JSON.parse(line)).not.toThrow()
}
const realDateNow = Date.now
Date.now = () => 1_234_567_890
behavior = 'partial-then-success'
behaviorCalls = 1
try {
history.addToHistory('SAME_TIME_A')
history.addToHistory('SAME_TIME_B')
} finally {
Date.now = realDateNow
}
await waitFor(async () => {
const current = await fsPromises.readFile(historyPath, 'utf8')
return (
current.includes('SAME_TIME_A') && current.includes('SAME_TIME_B')
)
})
history.removeLastFromHistory()
const visible: string[] = []
for await (const entry of history.makeHistoryReader()) {
if (entry.display.startsWith('SAME_TIME_')) {
visible.push(entry.display)
}
}
expect(visible).toEqual(['SAME_TIME_A'])
history.clearPendingHistoryEntries()
behavior = 'rollback-fails'
behaviorCalls = 0
history.addToHistory('POISONED_SENTINEL')
await waitFor(() => behaviorCalls === 1)
await Bun.sleep(600)
expect(behaviorCalls).toBe(1)
const pending: string[] = []
for await (const entry of history.makeHistoryReader()) {
if (entry.display === 'POISONED_SENTINEL') {
pending.push(entry.display)
}
}
expect(pending).toEqual(['POISONED_SENTINEL'])
history.clearPendingHistoryEntries()
behavior = 'read-fails'
behaviorCalls = 0
history.addToHistory('READ_FAILURE_SENTINEL')
await waitFor(() => behaviorCalls === 1)
await Bun.sleep(600)
expect(behaviorCalls).toBe(1)
history.clearPendingHistoryEntries()
behavior = 'unexpected-tail'
behaviorCalls = 0
history.addToHistory('UNEXPECTED_TAIL_SENTINEL')
await waitFor(() => behaviorCalls === 1)
await Bun.sleep(600)
expect(behaviorCalls).toBe(1)
history.clearPendingHistoryEntries()
})
})
async function waitFor(
predicate: () => boolean | Promise<boolean>,
timeoutMs = 2_000,
): Promise<void> {
const startedAt = Date.now()
while (!(await predicate())) {
if (Date.now() - startedAt > timeoutMs) {
throw new Error(`condition not met within ${timeoutMs}ms`)
}
await Bun.sleep(10)
}
}
+114 -14
View File
@@ -1,4 +1,10 @@
import { appendFile, writeFile } from 'fs/promises'
import {
appendFile,
readFile,
stat,
truncate,
writeFile,
} from 'fs/promises'
import { join } from 'path'
import { getProjectRoot, getSessionId } from './bootstrap/state.js'
import { registerCleanup } from './utils/cleanupRegistry.js'
@@ -105,6 +111,7 @@ function deserializeLogEntry(line: string): LogEntry {
async function* makeLogEntryReader(): AsyncGenerator<LogEntry> {
const currentSession = getSessionId()
const remainingSkippedTimestamps = new Map(skippedTimestampCounts)
// Start with entries that have yet to be flushed to disk
for (let i = pendingEntries.length - 1; i >= 0; i--) {
@@ -121,11 +128,17 @@ async function* makeLogEntryReader(): AsyncGenerator<LogEntry> {
// removeLastFromHistory slow path: entry was flushed before removal,
// so filter here so both getHistory (Up-arrow) and makeHistoryReader
// (ctrl+r search) skip it consistently.
if (
entry.sessionId === currentSession &&
skippedTimestamps.has(entry.timestamp)
) {
continue
if (entry.sessionId === currentSession) {
const remaining =
remainingSkippedTimestamps.get(entry.timestamp) ?? 0
if (remaining > 0) {
if (remaining === 1) {
remainingSkippedTimestamps.delete(entry.timestamp)
} else {
remainingSkippedTimestamps.set(entry.timestamp, remaining - 1)
}
continue
}
}
yield entry
} catch (error) {
@@ -286,11 +299,36 @@ let lastAddedEntry: LogEntry | null = null
// Timestamps of entries already flushed to disk that should be skipped when
// reading. Used by removeLastFromHistory when the entry has raced past the
// pending buffer. Session-scoped (module state resets on process restart).
const skippedTimestamps = new Set<number>()
const skippedTimestampCounts = new Map<number, number>()
let inFlightEntries = new Set<LogEntry>()
const removedInFlightEntries = new Set<LogEntry>()
let historyWriterPoisoned = false
function markTimestampSkipped(timestamp: number): void {
skippedTimestampCounts.set(
timestamp,
(skippedTimestampCounts.get(timestamp) ?? 0) + 1,
)
}
function commitHistoryEntries(entries: LogEntry[]): void {
for (const entry of entries) {
if (removedInFlightEntries.delete(entry)) {
markTimestampSkipped(entry.timestamp)
}
}
pendingEntries = pendingEntries.filter(entry => !inFlightEntries.has(entry))
}
function discardRemovedMarkers(entries: LogEntry[]): void {
for (const entry of entries) {
removedInFlightEntries.delete(entry)
}
}
// Core flush logic - writes pending entries to disk
async function immediateFlushHistory(): Promise<void> {
if (pendingEntries.length === 0) {
if (pendingEntries.length === 0 || historyWriterPoisoned) {
return
}
@@ -313,13 +351,65 @@ async function immediateFlushHistory(): Promise<void> {
},
})
const jsonLines = pendingEntries.map(entry => jsonStringify(entry) + '\n')
pendingEntries = []
const entriesToWrite = [...pendingEntries]
inFlightEntries = new Set(entriesToWrite)
const payload = entriesToWrite
.map(entry => jsonStringify(entry) + '\n')
.join('')
const payloadBytes = Buffer.from(payload)
const originalSize = (await stat(historyPath)).size
await appendFile(historyPath, jsonLines.join(''), { mode: 0o600 })
try {
await appendFile(historyPath, payload, { mode: 0o600 })
commitHistoryEntries(entriesToWrite)
} catch (appendError) {
let tail: Buffer
try {
const contents = await readFile(historyPath)
tail = contents.subarray(originalSize)
} catch (reconcileError) {
historyWriterPoisoned = true
logForDebugging(
`Prompt history writer disabled after an indeterminate append failure: ${reconcileError}`,
)
return
}
if (
tail.length >= payloadBytes.length &&
tail.subarray(0, payloadBytes.length).equals(payloadBytes)
) {
commitHistoryEntries(entriesToWrite)
return
}
if (
tail.length <= payloadBytes.length &&
payloadBytes.subarray(0, tail.length).equals(tail)
) {
try {
await truncate(historyPath, originalSize)
discardRemovedMarkers(entriesToWrite)
} catch (rollbackError) {
historyWriterPoisoned = true
logForDebugging(
`Prompt history writer disabled after rollback failed: ${rollbackError}`,
)
return
}
} else {
historyWriterPoisoned = true
logForDebugging(
'Prompt history writer disabled after append produced an unexpected file tail',
)
return
}
throw appendError
}
} catch (error) {
logForDebugging(`Failed to write prompt history: ${error}`)
} finally {
inFlightEntries.clear()
if (release) {
await release()
}
@@ -327,7 +417,11 @@ async function immediateFlushHistory(): Promise<void> {
}
async function flushPromptHistory(retries: number): Promise<void> {
if (isWriting || pendingEntries.length === 0) {
if (
isWriting ||
pendingEntries.length === 0 ||
historyWriterPoisoned
) {
return
}
@@ -435,8 +529,11 @@ export function addToHistory(command: HistoryEntry | string): void {
export function clearPendingHistoryEntries(): void {
pendingEntries = []
inFlightEntries.clear()
removedInFlightEntries.clear()
lastAddedEntry = null
skippedTimestamps.clear()
skippedTimestampCounts.clear()
historyWriterPoisoned = false
}
/**
@@ -458,7 +555,10 @@ export function removeLastFromHistory(): void {
const idx = pendingEntries.lastIndexOf(entry)
if (idx !== -1) {
pendingEntries.splice(idx, 1)
if (inFlightEntries.has(entry)) {
removedInFlightEntries.add(entry)
}
} else {
skippedTimestamps.add(entry.timestamp)
markTimestampSkipped(entry.timestamp)
}
}
+5 -3
View File
@@ -2774,7 +2774,7 @@ async function* queryModel(
)}`,
{ level: "warn" },
);
throw new RetriableStreamError(streamingError);
throw new RetriableStreamError(streamingError, assistantCommitBuffer.flush());
}
// The socket under the stream died mid-response (stale pooled keep-alive
@@ -2806,7 +2806,7 @@ async function* queryModel(
)}`,
{ level: "warn" },
);
throw new RetriableStreamError(streamingError);
throw new RetriableStreamError(streamingError, assistantCommitBuffer.flush());
}
if (
@@ -2821,7 +2821,7 @@ async function* queryModel(
)}`,
{ level: "warn" },
);
throw new RetriableStreamError(streamingError);
throw new RetriableStreamError(streamingError, assistantCommitBuffer.flush());
}
// When the flag is enabled, skip the non-streaming fallback and let the
@@ -3121,6 +3121,7 @@ async function* queryModel(
return;
}
yield* assistantCommitBuffer.flush();
yield getAssistantMessageFromError(error, errorModel, {
messages,
messagesForAPI,
@@ -3179,6 +3180,7 @@ async function* queryModel(
return;
}
yield* assistantCommitBuffer.flush();
yield getAssistantMessageFromError(error, errorModel, {
messages,
messagesForAPI,
+46
View File
@@ -173,6 +173,52 @@ describe('withStreamRetry', () => {
delete process.env[RETRY_ENV]
})
test('yields completed text from only the final exhausted attempt', async () => {
process.env[RETRY_ENV] = '1'
let calls = 0
const attempt = () =>
// biome-ignore lint/suspicious/noExplicitAny: mock stream messages
(async function* (): AsyncGenerator<any, void> {
calls++
throw new RetriableStreamError(
new Error('socket reset'),
[
{
type: 'assistant',
message: {
id: `response-${calls}`,
type: 'message',
role: 'assistant',
model: 'test-model',
content: [{ type: 'text', text: `partial-${calls}` }],
stop_reason: null,
stop_sequence: null,
usage: {
input_tokens: 0,
output_tokens: 0,
},
},
uuid: `partial-${calls}`,
timestamp: new Date().toISOString(),
},
],
)
})()
const out = await collect(withStreamRetry(attempt, 'test-model', []))
const partials = out.filter(
message =>
message.type === 'assistant' &&
typeof message.uuid === 'string' &&
message.uuid.startsWith('partial-'),
)
expect(calls).toBe(2)
expect(partials.map(message => message.uuid)).toEqual(['partial-2'])
expect(out.at(-1)?.isApiErrorMessage).toBe(true)
delete process.env[RETRY_ENV]
})
test('passes through a clean attempt without retrying', async () => {
let calls = 0
const attempt = () =>
+3
View File
@@ -72,6 +72,9 @@ export async function* withStreamRetry(
model:
model as AnalyticsMetadata_I_VERIFIED_THIS_IS_NOT_CODE_OR_FILEPATHS,
});
for (const bufferedMessage of error.bufferedMessages) {
yield bufferedMessage;
}
yield getAssistantMessageFromError(error.originalError, model, {
messages,
});
+96
View File
@@ -3,6 +3,7 @@ import type Anthropic from '@anthropic-ai/sdk'
import { APIConnectionError, APIError } from '@anthropic-ai/sdk'
import { _resetKeepAliveForTesting, getProxyFetchOptions } from '../../utils/proxy.js'
import {
CannotRetryError,
getMaxStreamTransientRetries,
isRetryableStreamError,
isRetryableStreamTransportError,
@@ -55,6 +56,101 @@ describe('withRetry stale connections', () => {
})
})
describe('withRetry context overflow recovery', () => {
test('uses the available context even when the thinking budget is larger', async () => {
const overrides: Array<number | undefined> = []
const overflowMessage =
'input length and `max_tokens` exceed context limit: 190000 + 20000 > 200000'
const overflow = new APIError(
400,
{
type: 'error',
error: {
type: 'invalid_request_error',
message: overflowMessage,
},
},
overflowMessage,
undefined,
)
const generator = withRetry(
async () => ({} as Anthropic),
async (_client, attempt, context) => {
overrides.push(context.maxTokensOverride)
if (attempt === 1) {
throw overflow
}
return 'ok'
},
{
model: 'claude-opus-4-7',
thinkingConfig: { type: 'enabled', budgetTokens: 20_000 },
maxRetries: 1,
},
)
let finalValue: string | undefined
for (;;) {
const next = await generator.next()
if (next.done) {
finalValue = next.value
break
}
}
expect(finalValue).toBe('ok')
expect(overrides).toEqual([undefined, 9_000])
})
test('stops when the provider repeats the same overflow after adjustment', async () => {
let attempts = 0
const overflowMessage =
'input length and `max_tokens` exceed context limit: 190000 + 20000 > 200000'
const overflow = new APIError(
400,
{
type: 'error',
error: {
type: 'invalid_request_error',
message: overflowMessage,
},
},
overflowMessage,
undefined,
)
const generator = withRetry(
async () => ({} as Anthropic),
async () => {
attempts += 1
throw overflow
},
{
model: 'claude-opus-4-7',
thinkingConfig: { type: 'disabled' },
maxRetries: 5,
},
)
let thrown: unknown
try {
while (!(await generator.next()).done) {
// Drain retry status messages until the generator terminates.
}
} catch (error) {
thrown = error
}
expect(thrown).toBeInstanceOf(CannotRetryError)
expect((thrown as CannotRetryError).originalError).toBe(overflow)
expect((thrown as CannotRetryError).retryContext.maxTokensOverride).toBe(
9_000,
)
expect(attempts).toBe(2)
})
})
describe('isRetryableStreamError', () => {
// The SDK embeds the serialized error body in `error.message`; mirror that so
// the matcher sees the same shape it does in production.
+19 -12
View File
@@ -6,7 +6,10 @@ import {
APIUserAbortError,
} from '@anthropic-ai/sdk'
import type { QuerySource } from 'src/constants/querySource.js'
import type { SystemAPIErrorMessage } from 'src/types/message.js'
import type {
AssistantMessage,
SystemAPIErrorMessage,
} from 'src/types/message.js'
import { isAwsCredentialsProviderError } from 'src/utils/aws.js'
import { logForDebugging } from 'src/utils/debug.js'
import { logError } from 'src/utils/log.js'
@@ -175,9 +178,16 @@ export class FallbackTriggeredError extends Error {
* path can surface a faithful API-error message.
*/
export class RetriableStreamError extends Error {
constructor(public readonly originalError: unknown) {
public readonly bufferedMessages: readonly AssistantMessage[]
constructor(
public readonly originalError: unknown,
bufferedMessages: readonly AssistantMessage[] = [],
) {
super(errorMessage(originalError))
this.name = 'RetriableStreamError'
this.bufferedMessages = bufferedMessages
Object.defineProperty(this, 'bufferedMessages', { enumerable: false })
if (originalError instanceof Error && originalError.stack) {
this.stack = originalError.stack
}
@@ -515,16 +525,13 @@ export async function* withRetry<T>(
)
throw error
}
// Ensure we have enough tokens for thinking + at least 1 output token
const minRequired =
(retryContext.thinkingConfig.type === 'enabled'
? retryContext.thinkingConfig.budgetTokens
: 0) + 1
const adjustedMaxTokens = Math.max(
FLOOR_OUTPUT_TOKENS,
availableContext,
minRequired,
)
const adjustedMaxTokens = availableContext
if (
retryContext.maxTokensOverride !== undefined &&
adjustedMaxTokens >= retryContext.maxTokensOverride
) {
throw new CannotRetryError(error, retryContext)
}
retryContext.maxTokensOverride = adjustedMaxTokens
logEvent('tengu_max_tokens_context_overflow_adjustment', {
@@ -0,0 +1,72 @@
import { describe, expect, test } from 'bun:test'
import { AskUserQuestionTool, type Output } from './AskUserQuestionTool.js'
const question = {
question: 'What should I do next?',
header: 'Next step',
options: [
{ label: 'Continue', description: 'Proceed with the task' },
{ label: 'Pause', description: 'Stop and wait' },
],
multiSelect: false,
}
describe('AskUserQuestion tool result guidance', () => {
test('keeps the concise continuation guidance for a predefined option', () => {
const result = mapResult({
questions: [question],
answers: { [question.question]: 'Continue' },
})
expect(result.content).toContain('You can now continue')
})
test('uses neutral guidance when the user gives a free-text instruction', () => {
const result = mapResult({
questions: [question],
answers: {
[question.question]: 'Wait. Explain the risks before doing anything.',
},
})
expect(result.content).not.toContain('You can now continue')
expect(result.content).toContain('may require that you pause')
})
test('uses neutral guidance when a predefined option has user notes', () => {
const result = mapResult({
questions: [question],
answers: { [question.question]: 'Continue' },
annotations: {
[question.question]: {
notes: 'Explain the risks first.',
},
},
})
expect(result.content).not.toContain('You can now continue')
expect(result.content).toContain('may require that you pause')
})
test('uses neutral guidance when not every question was answered', () => {
const secondQuestion = {
...question,
question: 'Where should I continue?',
header: 'Scope',
}
const result = mapResult({
questions: [question, secondQuestion],
answers: { [question.question]: 'Continue' },
})
expect(result.content).not.toContain('You can now continue')
expect(result.content).toContain('may require that you pause')
})
})
function mapResult(output: Output) {
return AskUserQuestionTool.mapToolResultToToolResultBlockParam(
output,
'tool-use-id',
)
}
File diff suppressed because one or more lines are too long
+77 -28
View File
@@ -1,5 +1,5 @@
import { describe, expect, test, beforeEach, afterEach } from 'bun:test'
import { mkdir, writeFile, rm } from 'fs/promises'
import { mkdir, readdir, writeFile, rm, symlink } from 'fs/promises'
import { join } from 'path'
import { randomUUID } from 'crypto'
@@ -45,35 +45,33 @@ describe('updateCronTask integration', () => {
describe('CronTaskMeta type coverage', () => {
test('all UI fields are optional on CronTask', async () => {
// Verify all new fields exist on the type by creating tasks with them
const { addCronTask } = await import('../cronTasks.js')
const { addCronTask, removeCronTasks } = await import('../cronTasks.js')
// Create a task with all metadata fields (durable=true writes to disk in test dir)
const tmpDir = join('/tmp', `cron-meta-test-${randomUUID().slice(0, 8)}`)
await mkdir(join(tmpDir, '.claude'), { recursive: true })
let id: string | undefined
try {
id = await addCronTask(
'0 9 * * *',
'test prompt',
true, // recurring
false, // session-only; this type test must not write project state
undefined, // agentId
{
name: 'test-name',
description: 'test description',
folder: '/test/folder',
model: 'claude-opus-4-7',
permissionMode: 'ask',
worktree: false,
frequency: 'daily',
scheduledTime: '09:00',
},
)
const id = await addCronTask(
'0 9 * * *',
'test prompt',
true, // recurring
true, // durable (writes to disk)
undefined, // agentId
{
name: 'test-name',
description: 'test description',
folder: '/test/folder',
model: 'claude-opus-4-7',
permissionMode: 'ask',
worktree: false,
frequency: 'daily',
scheduledTime: '09:00',
},
)
expect(typeof id).toBe('string')
expect(id.length).toBe(8) // Short ID
// Clean up
await rm(tmpDir, { recursive: true, force: true })
expect(typeof id).toBe('string')
expect(id.length).toBe(8) // Short ID
} finally {
if (id) await removeCronTasks([id])
}
})
})
@@ -202,3 +200,54 @@ describe('writeCronTasks strips runtime fields', () => {
await rm(tmpDir, { recursive: true, force: true })
})
})
describe('writeCronTasks symlink safety', () => {
test('refuses to write through a project .claude directory symlink', async () => {
const { writeCronTasks } = await import('../cronTasks.js')
const tmpDir = join('/tmp', `cron-symlink-${randomUUID().slice(0, 8)}`)
const projectDir = join(tmpDir, 'project')
const outsideDir = join(tmpDir, 'outside')
try {
await mkdir(projectDir, { recursive: true })
await mkdir(outsideDir, { recursive: true })
await symlink(outsideDir, join(projectDir, '.claude'), 'dir')
const task = {
id: 'abcd1234',
cron: '0 9 * * *',
prompt: 'must stay inside the project',
createdAt: Date.now(),
}
await expect(writeCronTasks([task], projectDir)).rejects.toThrow(
'symbolic link',
)
expect(
await Bun.file(join(outsideDir, 'scheduled_tasks.json')).exists(),
).toBe(false)
} finally {
await rm(tmpDir, { recursive: true, force: true })
}
})
test('cleans up the temporary file when atomic replacement fails', async () => {
const { writeCronTasks } = await import('../cronTasks.js')
const tmpDir = join('/tmp', `cron-atomic-${randomUUID().slice(0, 8)}`)
const claudeDir = join(tmpDir, '.claude')
try {
await mkdir(join(claudeDir, 'scheduled_tasks.json'), { recursive: true })
const task = {
id: 'abcd1234',
cron: '0 9 * * *',
prompt: 'test',
createdAt: Date.now(),
}
await expect(writeCronTasks([task], tmpDir)).rejects.toThrow()
expect(await readdir(claudeDir)).toEqual(['scheduled_tasks.json'])
} finally {
await rm(tmpDir, { recursive: true, force: true })
}
})
})
+172
View File
@@ -0,0 +1,172 @@
import { describe, expect, test } from 'bun:test'
import type { Message } from '../types/message.js'
import { deserializeMessages } from './conversationRecovery.js'
import { normalizeMessagesForAPI } from './messages.js'
describe('deserializeMessages malformed attachments', () => {
test.each([
['missing payload', { type: 'attachment' }],
['null payload', { type: 'attachment', attachment: null }],
[
'invalid hook context',
{
type: 'attachment',
attachment: {
type: 'hook_additional_context',
content: null,
},
},
],
[
'new-file attachment with a non-string filename',
{
type: 'attachment',
attachment: { type: 'new_file', filename: null },
},
],
[
'new-directory attachment with a non-string path',
{
type: 'attachment',
attachment: { type: 'new_directory', path: null },
},
],
[
'current file attachment with a non-string filename',
{
type: 'attachment',
attachment: { type: 'file', filename: 42 },
},
],
[
'current directory attachment with a non-string path',
{
type: 'attachment',
attachment: { type: 'directory', path: 42 },
},
],
[
'IDE selection with non-string content',
{
type: 'attachment',
attachment: {
type: 'selected_lines_in_ide',
content: null,
},
},
],
[
'invoked-skills attachment with a non-array payload',
{
type: 'attachment',
attachment: { type: 'invoked_skills', skills: null },
},
],
[
'hook-success attachment with a non-string payload',
{
type: 'attachment',
attachment: { type: 'hook_success', content: null },
},
],
[
'skill-listing attachment with a non-string payload',
{
type: 'attachment',
attachment: { type: 'skill_listing', content: null },
},
],
[
'deferred-tools delta without rendered lines',
{
type: 'attachment',
attachment: {
type: 'deferred_tools_delta',
addedNames: ['Read'],
removedNames: [],
},
},
],
[
'MCP instructions delta without rendered blocks',
{
type: 'attachment',
attachment: {
type: 'mcp_instructions_delta',
addedNames: ['server'],
removedNames: [],
},
},
],
[
'agent-listing delta without rendered lines',
{
type: 'attachment',
attachment: {
type: 'agent_listing_delta',
addedTypes: ['Explore'],
removedTypes: [],
},
},
],
])('drops a %s instead of crashing resume', (_name, malformed) => {
const messages = deserializeMessages([
malformed as unknown as Message,
])
expect(messages).toEqual([])
})
test('keeps non-attachment messages and unknown forward-compatible attachments', () => {
const messages = deserializeMessages([
{
type: 'system',
subtype: 'local_command',
content: 'status',
level: 'info',
uuid: crypto.randomUUID(),
timestamp: new Date().toISOString(),
} as Message,
{
type: 'attachment',
attachment: { type: 'future_attachment' },
} as unknown as Message,
])
expect(
messages.some(
message =>
message.type === 'system' && message.content === 'status',
),
).toBe(true)
expect(
messages.some(
message =>
message.type === 'attachment' &&
(message.attachment as { type: string }).type ===
'future_attachment',
),
).toBe(true)
})
test('drops a malformed known attachment at API normalization', () => {
const malformed = {
type: 'attachment',
attachment: {
type: 'todo_reminder',
content: null,
itemCount: 0,
},
} as unknown as Message
expect(normalizeMessagesForAPI([malformed])).toEqual([])
})
test('drops a missing attachment envelope at API normalization', () => {
expect(
normalizeMessagesForAPI([
{ type: 'attachment' } as unknown as Message,
]),
).toEqual([])
})
})
+101 -4
View File
@@ -19,6 +19,7 @@ import type {
} from '../types/message.js'
import { PERMISSION_MODES } from '../types/permissions.js'
import { suppressNextSkillListing } from './attachments.js'
import { logForDebugging } from './debug.js'
import {
copyFileHistoryForResume,
type FileHistorySnapshot,
@@ -117,7 +118,7 @@ function migrateLegacyAttachmentTypes(message: Message): Message {
: 'skillDir' in attachment
? (attachment.skillDir as string)
: undefined
if (path) {
if (typeof path === 'string') {
return {
...message,
attachment: {
@@ -131,6 +132,99 @@ function migrateLegacyAttachmentTypes(message: Message): Message {
return message
}
function isStringArray(value: unknown): value is string[] {
return Array.isArray(value) && value.every(item => typeof item === 'string')
}
function isWellFormedAttachmentPayload(message: Message): boolean {
if (message.type !== 'attachment') {
return true
}
const attachment = (message as { attachment?: unknown }).attachment
if (
typeof attachment !== 'object' ||
attachment === null ||
!('type' in attachment) ||
typeof attachment.type !== 'string'
) {
return false
}
switch (attachment.type) {
case 'new_file':
return 'filename' in attachment && typeof attachment.filename === 'string'
case 'new_directory':
return 'path' in attachment && typeof attachment.path === 'string'
case 'invoked_skills':
return (
'skills' in attachment &&
Array.isArray(attachment.skills) &&
attachment.skills.every(
skill => typeof skill === 'object' && skill !== null,
)
)
case 'file':
return (
'filename' in attachment && typeof attachment.filename === 'string'
)
case 'directory':
return 'path' in attachment && typeof attachment.path === 'string'
case 'selected_lines_in_ide':
return (
'content' in attachment && typeof attachment.content === 'string'
)
case 'hook_success':
case 'skill_listing':
return 'content' in attachment && typeof attachment.content === 'string'
case 'hook_additional_context':
return (
'content' in attachment && isStringArray(attachment.content)
)
case 'deferred_tools_delta':
return (
'addedNames' in attachment &&
isStringArray(attachment.addedNames) &&
'addedLines' in attachment &&
isStringArray(attachment.addedLines) &&
'removedNames' in attachment &&
isStringArray(attachment.removedNames)
)
case 'mcp_instructions_delta':
return (
'addedNames' in attachment &&
isStringArray(attachment.addedNames) &&
'addedBlocks' in attachment &&
isStringArray(attachment.addedBlocks) &&
'removedNames' in attachment &&
isStringArray(attachment.removedNames)
)
case 'agent_listing_delta':
return (
'addedTypes' in attachment &&
isStringArray(attachment.addedTypes) &&
'addedLines' in attachment &&
isStringArray(attachment.addedLines) &&
'removedTypes' in attachment &&
isStringArray(attachment.removedTypes)
)
default:
return true
}
}
function dropMalformedAttachments(messages: Message[]): Message[] {
const filtered = messages.filter(isWellFormedAttachmentPayload)
const droppedCount = messages.length - filtered.length
if (droppedCount > 0) {
logForDebugging(
`resume: dropped ${droppedCount} attachment entries with a missing or malformed payload`,
{ level: 'warn' },
)
}
return filtered
}
export type TeleportRemoteResponse = {
log: Message[]
branch?: string
@@ -166,7 +260,7 @@ export function deserializeMessagesWithInterruptDetection(
): DeserializeResult {
try {
// Transform legacy attachment types before processing
const migratedMessages = serializedMessages.map(
const migratedMessages = dropMalformedAttachments(serializedMessages).map(
migrateLegacyAttachmentTypes,
)
@@ -553,12 +647,15 @@ export async function loadConversationForResume(
checkResumeConsistency(messages)
}
// Filter unsafe persisted payloads before any resume-side effects.
messages = dropMalformedAttachments(messages!)
// Restore skill state from invoked_skills attachments before deserialization.
// This ensures skills survive multiple compaction cycles after resume.
restoreSkillStateFromMessages(messages!)
restoreSkillStateFromMessages(messages)
// Deserialize messages to handle unresolved tool uses and ensure proper format
const deserialized = deserializeMessagesWithInterruptDetection(messages!)
const deserialized = deserializeMessagesWithInterruptDetection(messages)
messages = deserialized.messages
// Process session start hooks for resume
+65 -7
View File
@@ -10,8 +10,8 @@
// { "tasks": [{ id, cron, prompt, createdAt, recurring?, permanent? }] }
import { randomUUID } from 'crypto'
import { readFileSync } from 'fs'
import { mkdir, writeFile } from 'fs/promises'
import { readFileSync, type Stats } from 'fs'
import { lstat, mkdir, rename, unlink, writeFile } from 'fs/promises'
import { join } from 'path'
import {
addSessionCronTask,
@@ -216,17 +216,75 @@ export async function writeCronTasks(
dir?: string,
): Promise<void> {
const root = dir ?? getProjectRoot()
await mkdir(join(root, '.claude'), { recursive: true })
const claudeDir = join(root, '.claude')
await mkdir(claudeDir, { recursive: true })
const originalDirectoryStats = await assertSafeCronDirectory(claudeDir)
// Strip runtime-only flags — everything on disk is durable by definition,
// and agentId is session-scoped (teammates don't persist across sessions).
const body: CronFile = {
tasks: tasks.map(({ durable: _durable, agentId: _agentId, ...rest }) => rest),
}
await writeFile(
getCronFilePath(root),
jsonStringify(body, null, 2) + '\n',
'utf-8',
const targetPath = getCronFilePath(root)
const temporaryPath = join(
claudeDir,
`.scheduled_tasks.${process.pid}.${randomUUID()}.tmp`,
)
try {
await writeFile(temporaryPath, jsonStringify(body, null, 2) + '\n', {
encoding: 'utf-8',
flag: 'wx',
mode: 0o600,
})
await assertSameCronDirectory(claudeDir, originalDirectoryStats)
await rename(temporaryPath, targetPath)
await assertSameCronDirectory(claudeDir, originalDirectoryStats)
} catch (error) {
if (
await isSameCronDirectory(claudeDir, originalDirectoryStats)
) {
await unlink(temporaryPath).catch(() => {})
}
throw error
}
}
async function assertSafeCronDirectory(claudeDir: string): Promise<Stats> {
const stats = await lstat(claudeDir)
if (stats.isSymbolicLink()) {
throw new Error(
`Refusing to write scheduled tasks through symbolic link: ${claudeDir}`,
)
}
if (!stats.isDirectory()) {
throw new Error(
`Refusing to write scheduled tasks because the path is not a directory: ${claudeDir}`,
)
}
return stats
}
async function assertSameCronDirectory(
claudeDir: string,
expected: Stats,
): Promise<void> {
const current = await assertSafeCronDirectory(claudeDir)
if (current.dev !== expected.dev || current.ino !== expected.ino) {
throw new Error(
`Refusing to write scheduled tasks because the directory changed: ${claudeDir}`,
)
}
}
async function isSameCronDirectory(
claudeDir: string,
expected: Stats,
): Promise<boolean> {
try {
await assertSameCronDirectory(claudeDir, expected)
return true
} catch {
return false
}
}
/**
+389
View File
@@ -0,0 +1,389 @@
import { afterEach, beforeEach, describe, expect, test } from 'bun:test'
import { randomUUID, type UUID } from 'crypto'
import {
link,
mkdir,
mkdtemp,
open,
readdir,
readFile,
rm,
symlink,
unlink,
writeFile,
} from 'fs/promises'
import { tmpdir } from 'os'
import { join } from 'path'
import {
getIsInteractive,
getOriginalCwd,
getSessionId,
setIsInteractive,
setOriginalCwd,
} from '../bootstrap/state.js'
import {
fileHistoryMakeSnapshot,
fileHistoryRewind,
fileHistoryTrackEdit,
type FileHistoryState,
} from './fileHistory.js'
const originalConfigDir = process.env.CLAUDE_CONFIG_DIR
const originalDisableCheckpointing =
process.env.CLAUDE_CODE_DISABLE_FILE_CHECKPOINTING
const originalCwd = getOriginalCwd()
const originalInteractive = getIsInteractive()
let testRoot: string | null = null
beforeEach(async () => {
testRoot = await mkdtemp(join(tmpdir(), 'cc-haha-file-history-security-'))
process.env.CLAUDE_CONFIG_DIR = join(testRoot, 'config')
delete process.env.CLAUDE_CODE_DISABLE_FILE_CHECKPOINTING
setOriginalCwd(join(testRoot, 'project'))
setIsInteractive(true)
await mkdir(getOriginalCwd(), { recursive: true })
})
afterEach(async () => {
setOriginalCwd(originalCwd)
setIsInteractive(originalInteractive)
restoreEnv('CLAUDE_CONFIG_DIR', originalConfigDir)
restoreEnv(
'CLAUDE_CODE_DISABLE_FILE_CHECKPOINTING',
originalDisableCheckpointing,
)
if (testRoot) {
await Bun.sleep(50)
await rm(testRoot, { recursive: true, force: true })
testRoot = null
}
})
describe('file history rewind link safety', () => {
test('refuses to create a backup through a symlinked backup directory', async () => {
const targetMessageId = randomUUID() as UUID
const trackedPath = join(getOriginalCwd(), 'tracked.txt')
const outsideBackupDirectory = join(testRoot!, 'outside-backups')
const backupRoot = join(
process.env.CLAUDE_CONFIG_DIR!,
'file-history',
)
await writeFile(trackedPath, 'snapshot content')
await mkdir(backupRoot, { recursive: true })
await mkdir(outsideBackupDirectory)
await symlink(
outsideBackupDirectory,
join(backupRoot, getSessionId()),
)
const { getState, updateState } = createHistoryState(targetMessageId)
await fileHistoryTrackEdit(updateState, trackedPath, targetMessageId)
expect(await readdir(outsideBackupDirectory)).toEqual([])
expect(getState().trackedFiles.size).toBe(0)
})
test('cleans up a partial backup when snapshot copying fails', async () => {
const targetMessageId = randomUUID() as UUID
const trackedPath = join(getOriginalCwd(), 'tracked.txt')
await writeFile(trackedPath, 'snapshot content')
const { getState, updateState } = createHistoryState(targetMessageId)
const probe = await open(trackedPath, 'r')
const fileHandlePrototype = Object.getPrototypeOf(probe) as {
write: (...args: unknown[]) => Promise<unknown>
}
await probe.close()
const originalWrite = fileHandlePrototype.write
fileHandlePrototype.write = async function () {
throw new Error('injected backup write failure')
}
try {
await fileHistoryTrackEdit(updateState, trackedPath, targetMessageId)
} finally {
fileHandlePrototype.write = originalWrite
}
const backupDirectory = join(
process.env.CLAUDE_CONFIG_DIR!,
'file-history',
getSessionId(),
)
expect(await readdir(backupDirectory)).toEqual([])
expect(getState().trackedFiles.size).toBe(0)
})
test.each(['symlink', 'hardlink'] as const)(
'does not overwrite an external victim through a %s',
async linkType => {
const targetMessageId = randomUUID() as UUID
const trackedPath = join(getOriginalCwd(), 'tracked.txt')
const victimPath = join(testRoot!, 'outside-victim.txt')
await writeFile(trackedPath, 'snapshot content')
await writeFile(victimPath, 'outside content')
let state: FileHistoryState = {
snapshots: [
{
messageId: targetMessageId,
trackedFileBackups: {},
timestamp: new Date(),
},
],
trackedFiles: new Set(),
snapshotSequence: 1,
}
const updateState = (
updater: (previous: FileHistoryState) => FileHistoryState,
) => {
state = updater(state)
}
await fileHistoryTrackEdit(updateState, trackedPath, targetMessageId)
await unlink(trackedPath)
if (linkType === 'symlink') {
await symlink(victimPath, trackedPath)
} else {
await link(victimPath, trackedPath)
}
await fileHistoryRewind(updateState, targetMessageId)
expect(await readFile(victimPath, 'utf8')).toBe('outside content')
},
)
test('restores a regular file after its parent directory was deleted', async () => {
const targetMessageId = randomUUID() as UUID
const trackedPath = join(getOriginalCwd(), 'nested', 'tracked.txt')
await mkdir(join(getOriginalCwd(), 'nested'))
await writeFile(trackedPath, 'snapshot content')
const { updateState } = createHistoryState(targetMessageId)
await fileHistoryTrackEdit(updateState, trackedPath, targetMessageId)
await rm(join(getOriginalCwd(), 'nested'), {
recursive: true,
force: true,
})
await fileHistoryRewind(updateState, targetMessageId)
expect(await readFile(trackedPath, 'utf8')).toBe('snapshot content')
})
test('keeps the new-file marker when the future parent does not exist', async () => {
const targetMessageId = randomUUID() as UUID
const trackedPath = join(getOriginalCwd(), 'future', 'tracked.txt')
const { updateState } = createHistoryState(targetMessageId)
await fileHistoryTrackEdit(updateState, trackedPath, targetMessageId)
await mkdir(join(getOriginalCwd(), 'future'))
await writeFile(trackedPath, 'created later')
await fileHistoryRewind(updateState, targetMessageId)
await expect(Bun.file(trackedPath).exists()).resolves.toBe(false)
})
test.each(['symlink', 'hardlink'] as const)(
'refuses to snapshot a tracked file replaced by a %s',
async linkType => {
const targetMessageId = randomUUID() as UUID
const trackedPath = join(getOriginalCwd(), 'tracked.txt')
const victimPath = join(testRoot!, 'snapshot-victim.txt')
await writeFile(trackedPath, 'snapshot content')
await writeFile(victimPath, 'outside content')
const { getState, updateState } = createHistoryState(targetMessageId)
await fileHistoryTrackEdit(updateState, trackedPath, targetMessageId)
const trackingPath = 'tracked.txt'
const before =
getState().snapshots.at(-1)!.trackedFileBackups[trackingPath]!
await unlink(trackedPath)
if (linkType === 'symlink') {
await symlink(victimPath, trackedPath)
} else {
await link(victimPath, trackedPath)
}
await fileHistoryMakeSnapshot(updateState, randomUUID() as UUID)
expect(await readFile(victimPath, 'utf8')).toBe('outside content')
expect(
getState().snapshots.at(-1)!.trackedFileBackups[trackingPath],
).toEqual(before)
},
)
test('does not begin tracking an existing hardlink', async () => {
const targetMessageId = randomUUID() as UUID
const trackedPath = join(getOriginalCwd(), 'tracked.txt')
const victimPath = join(testRoot!, 'track-victim.txt')
await writeFile(victimPath, 'outside content')
await link(victimPath, trackedPath)
const { getState, updateState } = createHistoryState(targetMessageId)
await fileHistoryTrackEdit(updateState, trackedPath, targetMessageId)
expect(getState().trackedFiles.size).toBe(0)
})
test('refuses to restore through a parent directory symlink', async () => {
const targetMessageId = randomUUID() as UUID
const nestedDir = join(getOriginalCwd(), 'nested')
const trackedPath = join(nestedDir, 'tracked.txt')
const outsideDir = join(testRoot!, 'outside-directory')
const victimPath = join(outsideDir, 'tracked.txt')
await mkdir(nestedDir)
await mkdir(outsideDir)
await writeFile(trackedPath, 'snapshot content')
await writeFile(victimPath, 'outside content')
const { updateState } = createHistoryState(targetMessageId)
await fileHistoryTrackEdit(updateState, trackedPath, targetMessageId)
await rm(nestedDir, { recursive: true, force: true })
await symlink(outsideDir, nestedDir)
await fileHistoryRewind(updateState, targetMessageId)
expect(await readFile(victimPath, 'utf8')).toBe('outside content')
})
test('refuses to delete through a parent directory symlink', async () => {
const targetMessageId = randomUUID() as UUID
const nestedDir = join(getOriginalCwd(), 'nested')
const trackedPath = join(nestedDir, 'tracked.txt')
const outsideDir = join(testRoot!, 'outside-delete-directory')
const victimPath = join(outsideDir, 'tracked.txt')
const { updateState } = createHistoryState(targetMessageId)
await fileHistoryTrackEdit(updateState, trackedPath, targetMessageId)
await mkdir(outsideDir)
await writeFile(victimPath, 'outside content')
await symlink(outsideDir, nestedDir)
await fileHistoryRewind(updateState, targetMessageId)
expect(await readFile(victimPath, 'utf8')).toBe('outside content')
})
test('treats a missing delete parent as already absent', async () => {
const targetMessageId = randomUUID() as UUID
const trackedPath = join(getOriginalCwd(), 'missing', 'tracked.txt')
const { updateState } = createHistoryState(targetMessageId)
await fileHistoryTrackEdit(updateState, trackedPath, targetMessageId)
await fileHistoryRewind(updateState, targetMessageId)
await expect(Bun.file(trackedPath).exists()).resolves.toBe(false)
})
test('keeps the original file intact when a restore write fails midway', async () => {
const targetMessageId = randomUUID() as UUID
const trackedPath = join(getOriginalCwd(), 'tracked.txt')
await writeFile(trackedPath, 's'.repeat(3 * 1024 * 1024))
const { updateState } = createHistoryState(targetMessageId)
await fileHistoryTrackEdit(updateState, trackedPath, targetMessageId)
await writeFile(trackedPath, 'modified content')
const probe = await open(trackedPath, 'r')
const fileHandlePrototype = Object.getPrototypeOf(probe) as {
write: (...args: unknown[]) => Promise<unknown>
}
await probe.close()
const originalWrite = fileHandlePrototype.write
let writes = 0
fileHandlePrototype.write = async function (...args: unknown[]) {
writes += 1
if (writes === 2) {
throw new Error('injected restore write failure')
}
return Reflect.apply(originalWrite, this, args) as Promise<unknown>
}
try {
await fileHistoryRewind(updateState, targetMessageId)
} finally {
fileHandlePrototype.write = originalWrite
}
expect(await readFile(trackedPath, 'utf8')).toBe('modified content')
})
test.each(['missing', 'symlink'] as const)(
'keeps the current file when its backup is %s',
async backupState => {
const targetMessageId = randomUUID() as UUID
const trackedPath = join(getOriginalCwd(), 'tracked.txt')
await writeFile(trackedPath, 'snapshot content')
const { getState, updateState } = createHistoryState(targetMessageId)
await fileHistoryTrackEdit(updateState, trackedPath, targetMessageId)
const backupName =
getState().snapshots.at(-1)!.trackedFileBackups['tracked.txt']!
.backupFileName
if (!backupName) throw new Error('expected a file backup')
const backupPath = join(
process.env.CLAUDE_CONFIG_DIR!,
'file-history',
getSessionId(),
backupName,
)
await unlink(backupPath)
if (backupState === 'symlink') {
const outsideBackup = join(testRoot!, 'outside-backup.txt')
await writeFile(outsideBackup, 'outside backup content')
await symlink(outsideBackup, backupPath)
}
await writeFile(trackedPath, 'modified content')
await fileHistoryRewind(updateState, targetMessageId)
expect(await readFile(trackedPath, 'utf8')).toBe('modified content')
},
)
test('preserves explicitly tracked absolute paths outside the project', async () => {
const targetMessageId = randomUUID() as UUID
const trackedPath = join(testRoot!, 'explicit-external.txt')
await writeFile(trackedPath, 'snapshot content')
const { updateState } = createHistoryState(targetMessageId)
await fileHistoryTrackEdit(updateState, trackedPath, targetMessageId)
await writeFile(trackedPath, 'modified content')
await fileHistoryRewind(updateState, targetMessageId)
expect(await readFile(trackedPath, 'utf8')).toBe('snapshot content')
})
})
function createHistoryState(targetMessageId: UUID): {
getState: () => FileHistoryState
updateState: (
updater: (previous: FileHistoryState) => FileHistoryState,
) => void
} {
let state: FileHistoryState = {
snapshots: [
{
messageId: targetMessageId,
trackedFileBackups: {},
timestamp: new Date(),
},
],
trackedFiles: new Set(),
snapshotSequence: 1,
}
return {
getState() {
return state
},
updateState(updater) {
state = updater(state)
},
}
}
function restoreEnv(name: string, value: string | undefined): void {
if (value === undefined) {
delete process.env[name]
} else {
process.env[name] = value
}
}
+439 -58
View File
@@ -1,16 +1,20 @@
import { createHash, type UUID } from 'crypto'
import { createHash, randomUUID, type UUID } from 'crypto'
import { diffLines } from 'diff'
import type { Stats } from 'fs'
import { constants as fsConstants, type Stats } from 'fs'
import {
chmod,
copyFile,
type FileHandle,
link,
lstat,
mkdir,
open,
readFile,
realpath,
rename,
stat,
unlink,
} from 'fs/promises'
import { dirname, isAbsolute, join, relative } from 'path'
import { dirname, isAbsolute, join, relative, resolve } from 'path'
import {
getIsNonInteractiveSession,
getOriginalCwd,
@@ -52,6 +56,8 @@ export type FileHistoryState = {
}
const MAX_SNAPSHOTS = 100
const O_NOFOLLOW = fsConstants.O_NOFOLLOW ?? 0
const COPY_BUFFER_SIZE = 1024 * 1024
export type DiffStats =
| {
filesChanged?: string[]
@@ -233,7 +239,7 @@ export async function fileHistoryMakeSnapshot(
// Stat the file once; ENOENT means the tracked file was deleted.
let fileStats: Stats | undefined
try {
fileStats = await stat(filePath)
fileStats = await lstat(filePath)
} catch (e: unknown) {
if (!isENOENT(e)) throw e
}
@@ -252,6 +258,15 @@ export async function fileHistoryMakeSnapshot(
)
return
}
if (
fileStats.isSymbolicLink() ||
!fileStats.isFile() ||
fileStats.nlink > 1
) {
throw new Error(
`FileHistory: Refusing to snapshot unsafe linked file: ${filePath}`,
)
}
// File exists - check if it needs to be backed up
if (
@@ -561,24 +576,32 @@ async function applySnapshot(
if (backupFileName === null) {
// File did not exist at the target version; delete it if present.
try {
await unlink(filePath)
if (
await deleteTrackedFile(
filePath,
!isAbsolute(trackingPath),
)
) {
logForDebugging(`FileHistory: [Rewind] Deleted ${filePath}`)
filesChanged.push(filePath)
} catch (e: unknown) {
if (!isENOENT(e)) throw e
// Already absent; nothing to do.
}
continue
}
// File should exist at a specific version. Restore only if it differs.
if (await checkOriginFileChanged(filePath, backupFileName)) {
await restoreBackup(filePath, backupFileName)
logForDebugging(
`FileHistory: [Rewind] Restored ${filePath} from ${backupFileName}`,
)
filesChanged.push(filePath)
if (
await restoreBackup(
filePath,
backupFileName,
!isAbsolute(trackingPath),
)
) {
logForDebugging(
`FileHistory: [Rewind] Restored ${filePath} from ${backupFileName}`,
)
filesChanged.push(filePath)
}
}
} catch (error) {
logError(error)
@@ -607,11 +630,19 @@ export async function checkOriginFileChanged(
let originalStats: Stats | null = originalStatsHint ?? null
if (!originalStats) {
try {
originalStats = await stat(originalFile)
originalStats = await lstat(originalFile)
} catch (e: unknown) {
if (!isENOENT(e)) return true
}
}
if (
originalStats &&
(originalStats.isSymbolicLink() ||
!originalStats.isFile() ||
originalStats.nlink > 1)
) {
return true
}
let backupStats: Stats | null = null
try {
backupStats = await stat(backupPath)
@@ -731,13 +762,100 @@ function getBackupFileName(filePath: string, version: number): string {
}
function resolveBackupPath(backupFileName: string, sessionId?: string): string {
assertSafePathSegment(backupFileName, 'backup file name')
return join(resolveBackupDirectory(sessionId), backupFileName)
}
function resolveBackupDirectory(sessionId?: string): string {
const resolvedSessionId = sessionId || getSessionId()
assertSafePathSegment(resolvedSessionId, 'session ID')
return join(getClaudeConfigHomeDir(), 'file-history', resolvedSessionId)
}
function assertSafePathSegment(segment: string, label: string): void {
if (
!segment ||
segment === '.' ||
segment === '..' ||
segment.includes('/') ||
segment.includes('\\')
) {
throw new Error(`FileHistory: Refusing unsafe ${label}: ${segment}`)
}
}
type SafeDirectoryEntry = {
path: string
stats: Stats
}
async function ensureSafeBackupDirectory(
sessionId?: string,
): Promise<SafeDirectoryEntry[]> {
const configDir = getClaudeConfigHomeDir()
return join(
configDir,
'file-history',
sessionId || getSessionId(),
backupFileName,
)
await mkdir(configDir, { recursive: true })
const backupDirectory = resolveBackupDirectory(sessionId)
const relativeBackupDirectory = relative(configDir, backupDirectory)
const entries: SafeDirectoryEntry[] = []
let currentPath = configDir
for (const segment of relativeBackupDirectory.split(/[\\/]/)) {
assertSafePathSegment(segment, 'backup directory segment')
currentPath = join(currentPath, segment)
try {
await mkdir(currentPath)
} catch (error) {
if (getErrnoCode(error) !== 'EEXIST') throw error
}
const stats = await lstat(currentPath)
assertSafeDirectory(stats, currentPath)
entries.push({ path: currentPath, stats })
}
return entries
}
async function inspectSafeBackupDirectory(
sessionId?: string,
): Promise<SafeDirectoryEntry[]> {
const configDir = getClaudeConfigHomeDir()
const backupDirectory = resolveBackupDirectory(sessionId)
const relativeBackupDirectory = relative(configDir, backupDirectory)
const entries: SafeDirectoryEntry[] = []
let currentPath = configDir
for (const segment of relativeBackupDirectory.split(/[\\/]/)) {
assertSafePathSegment(segment, 'backup directory segment')
currentPath = join(currentPath, segment)
const stats = await lstat(currentPath)
assertSafeDirectory(stats, currentPath)
entries.push({ path: currentPath, stats })
}
return entries
}
async function assertSafeDirectoryEntriesUnchanged(
entries: SafeDirectoryEntry[],
): Promise<void> {
for (const entry of entries) {
const currentStats = await lstat(entry.path)
assertSafeDirectory(currentStats, entry.path)
if (!sameFileIdentity(entry.stats, currentStats)) {
throw new Error(
`FileHistory: Refusing a backup directory that changed: ${entry.path}`,
)
}
}
}
async function areSafeDirectoryEntriesUnchanged(
entries: SafeDirectoryEntry[],
): Promise<boolean> {
try {
await assertSafeDirectoryEntriesUnchanged(entries)
return true
} catch {
return false
}
}
/**
@@ -756,13 +874,25 @@ async function createBackup(
const backupFileName = getBackupFileName(filePath, version)
const backupPath = resolveBackupPath(backupFileName)
// Stat first: if the source is missing, record a null backup and skip the
// copy. Separates "source missing" from "backup dir missing" cleanly —
// sharing a catch for both meant a file deleted between copyFile-success
// and stat would leave an orphaned backup with a null state record.
let srcStats: Stats
let pathStats: Stats
try {
srcStats = await stat(filePath)
pathStats = await lstat(filePath)
} catch (e: unknown) {
if (isENOENT(e)) {
return { backupFileName: null, version, backupTime: new Date() }
}
throw e
}
assertSafeRegularFile(pathStats, filePath, 'snapshot')
let source: FileHandle
try {
source = await open(
filePath,
process.platform === 'win32'
? fsConstants.O_RDONLY
: fsConstants.O_RDONLY | O_NOFOLLOW,
)
} catch (e: unknown) {
if (isENOENT(e)) {
return { backupFileName: null, version, backupTime: new Date() }
@@ -770,26 +900,56 @@ async function createBackup(
throw e
}
// copyFile preserves content and avoids reading the whole file into the JS
// heap (which the previous readFileSync+writeFileSync pipeline did, OOMing
// on large tracked files). Lazy mkdir: 99% of calls hit the fast path
// (directory already exists); on ENOENT, mkdir then retry.
let temporaryPath: string | undefined
let backupDirectoryEntries: SafeDirectoryEntry[] | undefined
try {
await copyFile(filePath, backupPath)
} catch (e: unknown) {
if (!isENOENT(e)) throw e
await mkdir(dirname(backupPath), { recursive: true })
await copyFile(filePath, backupPath)
await assertTrackedPathStaysWithinProject(filePath)
const srcStats = await source.stat()
assertSafeRegularFile(srcStats, filePath, 'snapshot')
if (!sameFileIdentity(pathStats, srcStats)) {
throw new Error(
`FileHistory: Refusing to snapshot a file that changed while opening: ${filePath}`,
)
}
backupDirectoryEntries = await ensureSafeBackupDirectory()
temporaryPath = `${backupPath}.${randomUUID()}.tmp`
const destination = await open(
temporaryPath,
fsConstants.O_WRONLY |
fsConstants.O_CREAT |
fsConstants.O_EXCL |
(process.platform === 'win32' ? 0 : O_NOFOLLOW),
srcStats.mode,
)
try {
await copyBetweenFileHandles(source, destination)
await destination.chmod(srcStats.mode)
await destination.sync()
} finally {
await destination.close()
}
await assertSafeDirectoryEntriesUnchanged(backupDirectoryEntries)
const temporaryStats = await lstat(temporaryPath)
assertSafeRegularFile(temporaryStats, temporaryPath, 'snapshot')
await rename(temporaryPath, backupPath)
temporaryPath = undefined
logEvent('tengu_file_history_backup_file_created', {
version: version,
fileSize: srcStats.size,
})
} finally {
await source.close()
if (
temporaryPath &&
backupDirectoryEntries &&
(await areSafeDirectoryEntriesUnchanged(backupDirectoryEntries))
) {
await unlink(temporaryPath).catch(() => {})
}
}
// Preserve file permissions on the backup.
await chmod(backupPath, srcStats.mode)
logEvent('tengu_file_history_backup_file_created', {
version: version,
fileSize: srcStats.size,
})
return {
backupFileName,
version,
@@ -804,36 +964,257 @@ async function createBackup(
async function restoreBackup(
filePath: string,
backupFileName: string,
): Promise<void> {
enforceProjectBoundary: boolean,
): Promise<boolean> {
const backupPath = resolveBackupPath(backupFileName)
const backupDirectoryEntries = await inspectSafeBackupDirectory()
// Stat first: if the backup is missing, log and bail before attempting
// the copy. Separates "backup missing" from "destination dir missing".
let backupStats: Stats
let backupPathStats: Stats
try {
backupStats = await stat(backupPath)
backupPathStats = await lstat(backupPath)
assertSafeRegularFile(backupPathStats, backupPath, 'restore')
} catch (e: unknown) {
if (isENOENT(e)) {
logEvent('tengu_file_history_rewind_restore_file_failed', {})
logError(
new Error(`FileHistory: [Rewind] Backup file not found: ${backupPath}`),
)
return
return false
}
throw e
}
// Lazy mkdir: 99% of calls hit the fast path (destination dir exists).
let source: FileHandle
try {
await copyFile(backupPath, filePath)
source = await open(
backupPath,
process.platform === 'win32'
? fsConstants.O_RDONLY
: fsConstants.O_RDONLY | O_NOFOLLOW,
)
} catch (e: unknown) {
if (!isENOENT(e)) throw e
await mkdir(dirname(filePath), { recursive: true })
await copyFile(backupPath, filePath)
if (isENOENT(e)) {
logEvent('tengu_file_history_rewind_restore_file_failed', {})
logError(
new Error(`FileHistory: [Rewind] Backup file not found: ${backupPath}`),
)
return false
}
throw e
}
// Restore the file permissions
await chmod(filePath, backupStats.mode)
try {
const backupStats = await source.stat()
if (!backupStats.isFile()) {
throw new Error(
`FileHistory: Refusing to restore from non-file backup: ${backupPath}`,
)
}
if (!sameFileIdentity(backupPathStats, backupStats)) {
throw new Error(
`FileHistory: Refusing a backup that changed while opening: ${backupPath}`,
)
}
await assertSafeDirectoryEntriesUnchanged(backupDirectoryEntries)
if (enforceProjectBoundary) {
await assertTrackedPathStaysWithinProject(filePath)
}
const parentPath = dirname(filePath)
await mkdir(parentPath, { recursive: true })
const parentStats = await lstat(parentPath)
assertSafeDirectory(parentStats, parentPath)
const originalTargetStats = await getSafeTargetStats(filePath, 'restore')
const temporaryPath = join(
parentPath,
`.${randomUUID()}.file-history-restore.tmp`,
)
let destination: FileHandle | undefined
try {
destination = await open(
temporaryPath,
fsConstants.O_WRONLY |
fsConstants.O_CREAT |
fsConstants.O_EXCL |
(process.platform === 'win32' ? 0 : O_NOFOLLOW),
backupStats.mode,
)
await copyBetweenFileHandles(source, destination)
await destination.chmod(backupStats.mode)
await destination.sync()
await destination.close()
destination = undefined
if (enforceProjectBoundary) {
await assertTrackedPathStaysWithinProject(filePath)
}
const currentParentStats = await lstat(parentPath)
assertSafeDirectory(currentParentStats, parentPath)
if (!sameFileIdentity(parentStats, currentParentStats)) {
throw new Error(
`FileHistory: Refusing to restore after the parent directory changed: ${filePath}`,
)
}
const currentTargetStats = await getSafeTargetStats(filePath, 'restore')
if (!sameOptionalFileIdentity(originalTargetStats, currentTargetStats)) {
throw new Error(
`FileHistory: Refusing to restore a file that changed concurrently: ${filePath}`,
)
}
const temporaryStats = await lstat(temporaryPath)
assertSafeRegularFile(temporaryStats, temporaryPath, 'restore')
await rename(temporaryPath, filePath)
return true
} finally {
await destination?.close().catch(() => {})
await unlink(temporaryPath).catch(() => {})
}
} finally {
await source.close()
}
}
async function deleteTrackedFile(
filePath: string,
enforceProjectBoundary: boolean,
): Promise<boolean> {
if (enforceProjectBoundary) {
await assertTrackedPathStaysWithinProject(filePath)
}
const parentPath = dirname(filePath)
let parentStats: Stats
try {
parentStats = await lstat(parentPath)
} catch (error) {
if (isENOENT(error)) return false
throw error
}
assertSafeDirectory(parentStats, parentPath)
const targetStats = await getSafeTargetStats(filePath, 'delete')
if (!targetStats) return false
if (enforceProjectBoundary) {
await assertTrackedPathStaysWithinProject(filePath)
}
const currentParentStats = await lstat(parentPath)
assertSafeDirectory(currentParentStats, parentPath)
if (!sameFileIdentity(parentStats, currentParentStats)) {
throw new Error(
`FileHistory: Refusing to delete after the parent directory changed: ${filePath}`,
)
}
const currentTargetStats = await getSafeTargetStats(filePath, 'delete')
if (!sameOptionalFileIdentity(targetStats, currentTargetStats)) {
throw new Error(
`FileHistory: Refusing to delete a file that changed concurrently: ${filePath}`,
)
}
await unlink(filePath)
return true
}
function assertSafeDirectory(stats: Stats, path: string): void {
if (stats.isSymbolicLink() || !stats.isDirectory()) {
throw new Error(`FileHistory: Refusing unsafe directory: ${path}`)
}
}
function assertSafeRegularFile(
stats: Stats,
path: string,
operation: 'snapshot' | 'restore' | 'delete',
): void {
if (stats.isSymbolicLink() || !stats.isFile() || stats.nlink > 1) {
throw new Error(
`FileHistory: Refusing to ${operation} unsafe linked file: ${path}`,
)
}
}
async function getSafeTargetStats(
path: string,
operation: 'restore' | 'delete',
): Promise<Stats | null> {
try {
const stats = await lstat(path)
assertSafeRegularFile(stats, path, operation)
return stats
} catch (error) {
if (isENOENT(error)) return null
throw error
}
}
function sameFileIdentity(left: Stats, right: Stats): boolean {
return left.dev === right.dev && left.ino === right.ino
}
function sameOptionalFileIdentity(
left: Stats | null,
right: Stats | null,
): boolean {
if (!left || !right) return left === right
return sameFileIdentity(left, right)
}
async function copyBetweenFileHandles(
source: FileHandle,
destination: FileHandle,
): Promise<void> {
const buffer = Buffer.allocUnsafe(COPY_BUFFER_SIZE)
let position = 0
while (true) {
const { bytesRead } = await source.read(
buffer,
0,
buffer.length,
position,
)
if (bytesRead === 0) return
let written = 0
while (written < bytesRead) {
const { bytesWritten } = await destination.write(
buffer,
written,
bytesRead - written,
position + written,
)
written += bytesWritten
}
position += bytesRead
}
}
async function assertTrackedPathStaysWithinProject(
filePath: string,
): Promise<void> {
const projectPath = resolve(getOriginalCwd())
const resolvedFilePath = resolve(filePath)
const lexicalRelative = relative(projectPath, resolvedFilePath)
if (lexicalRelative.startsWith('..') || isAbsolute(lexicalRelative)) {
return
}
const realProjectPath = await realpath(projectPath)
let existingParentPath = dirname(resolvedFilePath)
let realParentPath: string
while (true) {
try {
realParentPath = await realpath(existingParentPath)
break
} catch (error) {
if (!isENOENT(error)) throw error
const nextParent = dirname(existingParentPath)
if (nextParent === existingParentPath) throw error
existingParentPath = nextParent
}
}
const realRelative = relative(realProjectPath, realParentPath)
if (realRelative.startsWith('..') || isAbsolute(realRelative)) {
throw new Error(
`FileHistory: Refusing path whose parent escapes the project through a symbolic link: ${filePath}`,
)
}
}
/**
+20 -1
View File
@@ -1,5 +1,8 @@
import { describe, expect, test } from 'bun:test'
import { parseFrontmatter } from './frontmatterParser.js'
import {
parseFrontmatter,
splitPathInFrontmatter,
} from './frontmatterParser.js'
describe('parseFrontmatter delimiters', () => {
test('keeps indented delimiter text inside YAML block scalars', () => {
@@ -51,3 +54,19 @@ describe('parseFrontmatter delimiters', () => {
expect(content).toBe('Keep the body intact.')
})
})
describe('splitPathInFrontmatter brace expansion', () => {
test('keeps an exponentially large pattern unexpanded', () => {
const pattern = Array.from({ length: 11 }, () => '{a,b}').join('/')
expect(splitPathInFrontmatter(pattern)).toEqual([pattern])
})
test('shares the expansion budget across comma-separated patterns', () => {
const pattern = Array.from({ length: 9 }, () => '{a,b}').join('/')
const expanded = splitPathInFrontmatter([pattern, pattern, pattern])
expect(expanded).toHaveLength(514)
expect(expanded.slice(-2)).toEqual([pattern, pattern])
})
})
+63 -22
View File
@@ -190,8 +190,20 @@ export function parseFrontmatter(
* splitPathInFrontmatter(["a", "src/*.{ts,tsx}"]) // returns ["a", "src/*.ts", "src/*.tsx"]
*/
export function splitPathInFrontmatter(input: string | string[]): string[] {
const budget = {
results: 0,
bytes: 0,
exhausted: false,
}
return splitPathInFrontmatterWithBudget(input, budget)
}
function splitPathInFrontmatterWithBudget(
input: string | string[],
budget: BraceExpansionBudget,
): string[] {
if (Array.isArray(input)) {
return input.flatMap(splitPathInFrontmatter)
return input.flatMap(value => splitPathInFrontmatterWithBudget(value, budget))
}
if (typeof input !== 'string') {
return []
@@ -231,7 +243,13 @@ export function splitPathInFrontmatter(input: string | string[]): string[] {
// Expand brace patterns in each part
return parts
.filter(p => p.length > 0)
.flatMap(pattern => expandBraces(pattern))
.flatMap(pattern => expandBraces(pattern, budget))
}
type BraceExpansionBudget = {
results: number
bytes: number
exhausted: boolean
}
/**
@@ -240,31 +258,54 @@ export function splitPathInFrontmatter(input: string | string[]): string[] {
* expandBraces("src/*.{ts,tsx}") // returns ["src/*.ts", "src/*.tsx"]
* expandBraces("{a,b}/{c,d}") // returns ["a/c", "a/d", "b/c", "b/d"]
*/
function expandBraces(pattern: string): string[] {
// Find the first brace group
const braceMatch = pattern.match(/^([^{]*)\{([^}]+)\}(.*)$/)
if (!braceMatch) {
// No braces found, return pattern as-is
function expandBraces(
pattern: string,
budget: BraceExpansionBudget,
): string[] {
const maxResults = 1_000
const maxBytes = 4 * 1024 * 1024
if (budget.exhausted) {
return [pattern]
}
const prefix = braceMatch[1] || ''
const alternatives = braceMatch[2] || ''
const suffix = braceMatch[3] || ''
// Split alternatives by comma and expand each one
const parts = alternatives.split(',').map(alt => alt.trim())
// Recursively expand remaining braces in suffix
const pending = [pattern]
const expanded: string[] = []
for (const part of parts) {
const combined = prefix + part + suffix
// Recursively handle additional brace groups
const furtherExpanded = expandBraces(combined)
expanded.push(...furtherExpanded)
while (pending.length > 0) {
const current = pending.pop()!
budget.bytes += current.length
const braceMatch = current.match(/^([^{]*)\{([^}]+)\}(.*)$/)
if (!braceMatch) {
expanded.push(current)
continue
}
const prefix = braceMatch[1] || ''
const alternatives = braceMatch[2] || ''
const suffix = braceMatch[3] || ''
const parts = alternatives.split(',').map(alt => alt.trim())
const projectedResults =
budget.results + expanded.length + pending.length + parts.length
if (
budget.bytes > maxBytes ||
projectedResults > maxResults ||
projectedResults * pattern.length > maxBytes - budget.bytes
) {
budget.exhausted = true
logForDebugging(
`Brace pattern expansion exceeds the budget; using it unexpanded: ${pattern}`,
{ level: 'warn' },
)
return [pattern]
}
for (let i = parts.length - 1; i >= 0; i--) {
pending.push(prefix + parts[i]! + suffix)
}
}
budget.results += expanded.length
return expanded
}
+107
View File
@@ -0,0 +1,107 @@
import { afterEach, describe, expect, test } from 'bun:test'
import { truncateMcpContent } from './mcpValidation.js'
const originalLimit = process.env.MAX_MCP_OUTPUT_TOKENS
afterEach(() => {
if (originalLimit === undefined) {
delete process.env.MAX_MCP_OUTPUT_TOKENS
} else {
process.env.MAX_MCP_OUTPUT_TOKENS = originalLimit
}
})
describe('truncateMcpContent memory ownership', () => {
test('does not retain the backing store of a large string', async () => {
process.env.MAX_MCP_OUTPUT_TOKENS = '1000'
Bun.gc(true)
const baselineExternal = process.memoryUsage().external
let content = Buffer.alloc(32 * 1024 * 1024, 97).toString('utf8')
const truncated = await truncateMcpContent(content)
content = ''
await waitForExternalMemoryRelease(baselineExternal)
expect(typeof truncated).toBe('string')
expect(truncated).toStartWith('a'.repeat(4_000))
})
test('does not split a surrogate pair at the truncation boundary', async () => {
process.env.MAX_MCP_OUTPUT_TOKENS = '1000'
const content = `${'a'.repeat(3_999)}😀tail`
const truncated = await truncateMcpContent(content)
expect(typeof truncated).toBe('string')
expect(truncated).not.toContain('\ud83d')
})
test('detaches a truncated text block from its backing string', async () => {
process.env.MAX_MCP_OUTPUT_TOKENS = '1000'
const content = 'x'.repeat(8_000)
const truncated = await truncateMcpContent([
{ type: 'text', text: content },
])
if (!Array.isArray(truncated)) {
throw new Error('expected content blocks')
}
expect(truncated[0]).toEqual({
type: 'text',
text: 'x'.repeat(4_000),
})
})
test('detaches fully retained text blocks when later blocks exceed the limit', async () => {
process.env.MAX_MCP_OUTPUT_TOKENS = '1000'
Bun.gc(true)
const baselineExternal = process.memoryUsage().external
const truncated = await truncateRetainedBlockFromLargeBacking()
await waitForExternalMemoryRelease(baselineExternal)
expect(Array.isArray(truncated)).toBe(true)
})
test('detaches direct strings even when truncateMcpContent receives a short slice', async () => {
process.env.MAX_MCP_OUTPUT_TOKENS = '1000'
Bun.gc(true)
const baselineExternal = process.memoryUsage().external
const truncated = await truncateDirectSliceFromLargeBacking()
await waitForExternalMemoryRelease(baselineExternal)
expect(truncated).toStartWith('a'.repeat(3_000))
})
})
async function truncateRetainedBlockFromLargeBacking() {
const backing = Buffer.alloc(32 * 1024 * 1024, 97).toString('utf8')
return truncateMcpContent([
{ type: 'text', text: backing.slice(0, 3_000) },
{ type: 'text', text: 'b'.repeat(8_000) },
])
}
async function truncateDirectSliceFromLargeBacking() {
const backing = Buffer.alloc(32 * 1024 * 1024, 97).toString('utf8')
return truncateMcpContent(backing.slice(0, 3_000))
}
async function waitForExternalMemoryRelease(
baselineExternal: number,
): Promise<void> {
const maximumRetainedBytes = 8 * 1024 * 1024
const deadline = Date.now() + 2_000
while (
process.memoryUsage().external - baselineExternal >= maximumRetainedBytes &&
Date.now() < deadline
) {
await Bun.sleep(10)
Bun.gc(true)
}
expect(process.memoryUsage().external - baselineExternal).toBeLessThan(
maximumRetainedBytes,
)
}
+15 -5
View File
@@ -84,11 +84,18 @@ function getTruncationMessage(): string {
The tool output was truncated. If this MCP server provides pagination or filtering tools, use them to retrieve specific portions of the data. If pagination is not available, inform the user that you are working with truncated output and results may be incomplete.`
}
function forceCopyString(content: string): string {
return Buffer.from(content, 'utf16le').toString('utf16le')
}
function truncateString(content: string, maxChars: number): string {
if (content.length <= maxChars) {
return content
let truncated =
content.length <= maxChars ? content : content.slice(0, maxChars)
const lastCodeUnit = truncated.charCodeAt(truncated.length - 1)
if (lastCodeUnit >= 0xd800 && lastCodeUnit <= 0xdbff) {
truncated = truncated.slice(0, -1)
}
return content.slice(0, maxChars)
return forceCopyString(truncated)
}
async function truncateContentBlocks(
@@ -104,10 +111,13 @@ async function truncateContentBlocks(
if (remainingChars <= 0) break
if (block.text.length <= remainingChars) {
result.push(block)
result.push({ ...block, text: forceCopyString(block.text) })
currentChars += block.text.length
} else {
result.push({ type: 'text', text: block.text.slice(0, remainingChars) })
result.push({
type: 'text',
text: truncateString(block.text, remainingChars),
})
break
}
} else if (isImageBlock(block)) {
+15 -3
View File
@@ -2310,9 +2310,21 @@ export function normalizeMessagesForAPI(
return
}
case 'attachment': {
const rawAttachmentMessage = normalizeAttachmentForAPI(
message.attachment,
)
let rawAttachmentMessage: UserMessage[]
try {
rawAttachmentMessage = normalizeAttachmentForAPI(message.attachment)
} catch (error) {
const attachmentType = (message as {
attachment?: { type?: unknown }
}).attachment?.type
logForDebugging(
`Dropping malformed attachment during API normalization: ${
typeof attachmentType === 'string' ? attachmentType : 'unknown'
}: ${error instanceof Error ? error.message : String(error)}`,
{ level: 'warn' },
)
return
}
const attachmentMessage = checkStatsigFeatureGate_CACHED_MAY_BE_STALE(
'tengu_chair_sermon',
)
@@ -0,0 +1,91 @@
import { afterEach, describe, expect, test } from 'bun:test'
import { mkdtemp, mkdir, rm, writeFile } from 'fs/promises'
import { tmpdir } from 'os'
import { join } from 'path'
import { setInlinePlugins } from '../../bootstrap/state.js'
import { getCommandName } from '../../types/command.js'
import { clearPluginCache } from './pluginLoader.js'
import { getPluginSkills } from './loadPluginCommands.js'
let pluginDir: string | null = null
afterEach(async () => {
setInlinePlugins([])
getPluginSkills.cache?.clear?.()
clearPluginCache('load-plugin-commands-test')
if (pluginDir) {
await rm(pluginDir, { recursive: true, force: true })
pluginDir = null
}
})
describe('plugin skill display names', () => {
test('keeps the plugin prefix when frontmatter overrides the skill name', async () => {
pluginDir = await mkdtemp(join(tmpdir(), 'cc-haha-plugin-skill-'))
await mkdir(join(pluginDir, '.claude-plugin'), { recursive: true })
await mkdir(join(pluginDir, 'skills', 'review'), { recursive: true })
await writeFile(
join(pluginDir, '.claude-plugin', 'plugin.json'),
JSON.stringify({
name: 'acme',
version: '1.0.0',
description: 'Test plugin',
}),
)
await writeFile(
join(pluginDir, 'skills', 'review', 'SKILL.md'),
[
'---',
'name: custom-review',
'description: Review a change',
'---',
'Review the requested change.',
].join('\n'),
)
setInlinePlugins([pluginDir])
getPluginSkills.cache?.clear?.()
clearPluginCache('load-plugin-commands-test-setup')
const skills = await getPluginSkills()
const skill = skills.find(command => command.name === 'acme:review')
expect(skill).toBeDefined()
expect(getCommandName(skill!)).toBe('acme:custom-review')
expect(skill?.aliases).toContain('custom-review')
})
test('keeps the complete prefix when the plugin name contains a colon', async () => {
pluginDir = await mkdtemp(join(tmpdir(), 'cc-haha-plugin-skill-'))
await mkdir(join(pluginDir, '.claude-plugin'), { recursive: true })
await mkdir(join(pluginDir, 'skills', 'review'), { recursive: true })
await writeFile(
join(pluginDir, '.claude-plugin', 'plugin.json'),
JSON.stringify({
name: 'foo:bar',
version: '1.0.0',
description: 'Test plugin',
}),
)
await writeFile(
join(pluginDir, 'skills', 'review', 'SKILL.md'),
[
'---',
'name: custom-review',
'description: Review a change',
'---',
'Review the requested change.',
].join('\n'),
)
setInlinePlugins([pluginDir])
getPluginSkills.cache?.clear?.()
clearPluginCache('load-plugin-commands-colon-test')
const skills = await getPluginSkills()
const skill = skills.find(command => command.name === 'foo:bar:review')
expect(skill).toBeDefined()
expect(getCommandName(skill!)).toBe('foo:bar:custom-review')
})
})
+10 -1
View File
@@ -267,6 +267,11 @@ function createPluginCommand(
const whenToUse = frontmatter.when_to_use as string | undefined
const version = frontmatter.version as string | undefined
const displayName = frontmatter.name as string | undefined
const pluginPrefix = `${pluginManifest.name}:`
const userFacingDisplayName =
displayName && isSkill && !displayName.startsWith(pluginPrefix)
? `${pluginPrefix}${displayName}`
: displayName
// Handle model configuration, resolving aliases like 'haiku', 'sonnet', 'opus'
const model =
@@ -300,6 +305,10 @@ function createPluginCommand(
return {
type: 'prompt',
name: commandName,
aliases:
displayName && isSkill && userFacingDisplayName !== displayName
? [displayName]
: undefined,
description,
hasUserSpecifiedDescription: validatedDescription !== null,
allowedTools,
@@ -321,7 +330,7 @@ function createPluginCommand(
isHidden: !userInvocable,
progressMessage: isSkill || config.isSkillMode ? 'loading' : 'running',
userFacingName(): string {
return displayName || commandName
return userFacingDisplayName || commandName
},
async getPromptForCommand(args, context) {
// For skills from skills/ directory, include base directory