mirror of
https://github.com/NanmiCoder/claude-code-haha.git
synced 2026-10-10 03:43:11 +08:00
fix(cli): backport Claude Code stability fixes
This commit is contained in:
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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':
|
||||
|
||||
@@ -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
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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 = () =>
|
||||
|
||||
@@ -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,
|
||||
});
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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
@@ -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 })
|
||||
}
|
||||
})
|
||||
})
|
||||
|
||||
@@ -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([])
|
||||
})
|
||||
})
|
||||
@@ -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
@@ -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
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
|
||||
@@ -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
@@ -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}`,
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
|
||||
@@ -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])
|
||||
})
|
||||
})
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
}
|
||||
@@ -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
@@ -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')
|
||||
})
|
||||
})
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user