fix(session): harden cross-session references

This commit is contained in:
Yichen Jiang
2026-07-21 17:53:30 +08:00
parent 8394898ef5
commit ebb62c482c
19 changed files with 213 additions and 55 deletions

View File

@@ -397,25 +397,25 @@ describe('acp bridge', () => {
it('cancels reference preparation before a turn is created', async () => {
harness = await makeBridgeHarness({ storageDir, withSessionReferences: true, script: [] })
const source = harness.ctx.sessions.create(SessionId('source'))
const snapshot = await harness.ctx.sessionQuery.readSurface(source.id)
await harness.client.initialize({ protocolVersion: PROTOCOL_VERSION, clientCapabilities: {} })
const { sessionId } = await harness.client.newSession({ cwd: process.cwd(), mcpServers: [] })
const prepare = vi.spyOn(harness.ctx.sessionReferences, 'prepare').mockImplementation(
(_agent, _content, _references, signal) => new Promise((_resolve, reject) => {
if (signal?.aborted === true) {
reject(new Error('already aborted'))
return
}
signal?.addEventListener('abort', () => { reject(new Error('aborted')) }, { once: true })
}),
)
let releaseRead: (() => void) | undefined
const readSurface = vi.spyOn(harness.ctx.sessionQuery, 'readSurface').mockImplementationOnce(async () => {
await new Promise<void>((resolve) => { releaseRead = resolve })
return snapshot
})
const pending = harness.client.prompt({
sessionId,
prompt: [{ type: 'resource_link', uri: encodeSessionReferenceUri(source.id), name: 'source' }],
})
await vi.waitFor(() => { expect(prepare).toHaveBeenCalledOnce() })
await vi.waitFor(() => { expect(releaseRead).toBeTypeOf('function') })
await harness.client.cancel({ sessionId })
await expect(pending).resolves.toEqual({ stopReason: 'cancelled' })
expect(harness.ctx.agents.get(SessionId(sessionId))?.session.events).toHaveLength(0)
releaseRead?.()
await Promise.resolve()
readSurface.mockRestore()
})
it('rejects a prompt for an unknown session', async () => {

View File

@@ -202,6 +202,11 @@ function displayText(text: string): string {
`\\x${control.charCodeAt(0).toString(16).padStart(2, '0')}`)
}
/** Escape external controls for terminal fields that must remain on one line. */
function displayInlineText(text: string): string {
return displayText(text).replaceAll('\n', '\\x0a')
}
/**
* Theme-agnostic palette built from the standard 16-color ANSI set plus SGR
* attributes, which every terminal remaps to its active color scheme. Body
@@ -827,17 +832,20 @@ class SessionAutocompleteProvider implements AutocompleteProvider {
if (token === undefined) return basePromise
let candidates
try {
candidates = await this.sessions.listCandidates(this.agent, token.slice(1))
candidates = await this.sessions.listCandidates(this.agent, token.slice(1), undefined, options.signal)
} catch {
return basePromise
}
const base = await basePromise
if (options.signal.aborted) return base
const items: AutocompleteItem[] = candidates.map(candidate => ({
value: formatSessionReferenceMention({ sessionId: candidate.sessionId, label: candidate.label }),
label: `Session · ${candidate.sessionId}`,
description: `${candidate.cwd ?? '(no cwd)'} · ${new Date(candidate.createdAt).toISOString()}`,
}))
const items: AutocompleteItem[] = candidates.map((candidate) => {
const mentionLabel = displayInlineText(candidate.label)
return {
value: formatSessionReferenceMention({ sessionId: candidate.sessionId, label: mentionLabel }),
label: `Session · ${displayInlineText(candidate.sessionId)}`,
description: `${candidate.cwd === undefined ? '(no cwd)' : displayInlineText(candidate.cwd)} · ${new Date(candidate.createdAt).toISOString()}`,
}
})
if (items.length === 0) return base
return { items: [...items, ...(base?.items ?? [])], prefix: token }
}

View File

@@ -2,7 +2,7 @@ import { homedir } from 'node:os'
import { join } from 'node:path'
import { describe, expect, it, vi } from 'vitest'
import { Context } from 'cordis'
import type { Terminal } from '@earendil-works/pi-tui'
import { CombinedAutocompleteProvider, type Terminal } from '@earendil-works/pi-tui'
import AgentRegistry, { type Agent } from '@deepseek-ai/dsh-agent'
import CommandService, { type CommandInvocation } from '@deepseek-ai/dsh-commands'
import SessionStore, { SessionId, type JsonValue } from '@deepseek-ai/dsh-session'
@@ -545,6 +545,40 @@ describe('pi-tui chat lifecycle and transcript', () => {
await dispose(result)
})
it('escapes session autocomplete metadata while preserving the referenced session id', async () => {
const unsafeId = SessionId('evil\x1b\x07\u009b\ns')
const unsafeCwd = '/x/\x1b\x07\u009b\nf'
const result = await setup({
async configureContext(ctx) {
ctx.provide('tools', { get: () => undefined } as never)
await ctx.plugin(SessionQueryService)
await ctx.plugin(SessionReferenceService)
const source = ctx.sessions.create(unsafeId, { meta: { cwd: unsafeCwd, createdAt: 1 } })
appendUser(source, 'safe background')
},
})
result.terminal.send('@evil')
await vi.waitFor(() => {
expect(result.terminal.output).toContain('Session · evil\\x1b\\x07\\x9b\\x0a')
})
expect(result.terminal.output).toContain('/x/\\x1b\\x07\\x9b\\x0af')
expect(result.terminal.output).not.toContain('evil\x1b\x07')
expect(result.terminal.output).not.toContain('/x/\x1b\x07')
result.terminal.send('\t')
await tick()
result.terminal.send('\r')
await vi.waitFor(() => { expect(result.agent.sent).toHaveLength(1) })
expect(result.agent.sent).toEqual([[
{ type: 'text', text: '@evil\\x1b\\x07\\x9b\\x0as' },
]])
expect(result.agent.sentOptions[0]?.contexts).toMatchObject([{
meta: { references: [{ sessionId: unsafeId }] },
}])
await dispose(result)
})
it('falls back cleanly for non-session, empty, failed, and superseded autocomplete requests', async () => {
const result = await setup({
async configureContext(ctx) {
@@ -575,19 +609,38 @@ describe('pi-tui chat lifecycle and transcript', () => {
await tick()
result.terminal.send('\x03')
let releaseFirst: (() => void) | undefined
let releaseBase: (() => void) | undefined
const baseSuggestions = vi.spyOn(CombinedAutocompleteProvider.prototype, 'getSuggestions')
.mockImplementationOnce(async () => {
await new Promise<void>((resolve) => { releaseBase = resolve })
return null
})
listCandidates.mockResolvedValueOnce([])
result.terminal.send('@base-slow')
await vi.waitFor(() => { expect(releaseBase).toBeTypeOf('function') })
const baseWaitSignal = listCandidates.mock.calls.at(-1)?.[3]
result.terminal.send('x')
await vi.waitFor(() => { expect(baseWaitSignal?.aborted).toBe(true) })
releaseBase?.()
await tick()
baseSuggestions.mockRestore()
let delayedSignal: AbortSignal | undefined
let delayed = true
listCandidates.mockImplementation(async (...args) => {
if (!delayed) return originalListCandidates(...args)
delayed = false
await new Promise<void>((resolve) => { releaseFirst = resolve })
delayedSignal = args[3]
if (delayedSignal === undefined) throw new Error('expected autocomplete cancellation signal')
await new Promise<void>((_resolve, reject) => {
delayedSignal?.addEventListener('abort', () => { reject(new Error('superseded')) }, { once: true })
})
return []
})
result.terminal.send('@slow')
await vi.waitFor(() => { expect(releaseFirst).toBeTypeOf('function') })
await vi.waitFor(() => { expect(delayedSignal).toBeDefined() })
result.terminal.send('x')
releaseFirst?.()
await tick()
await vi.waitFor(() => { expect(delayedSignal?.aborted).toBe(true) })
await dispose(result)
})