fix(compaction): restore tail turns after summarization (#27145)

This commit is contained in:
Aiden Cline 2026-05-12 15:33:16 -05:00 committed by GitHub
commit ca28dd02ec
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 95 additions and 77 deletions

View file

@ -261,10 +261,10 @@ export const Info = Schema.Struct({
}), }),
tail_turns: Schema.optional(NonNegativeInt).annotate({ tail_turns: Schema.optional(NonNegativeInt).annotate({
description: description:
"Number of recent user turns, including their following assistant/tool responses, to serialize into the compaction summary (default: 2)", "Number of recent user turns, including their following assistant/tool responses, to keep verbatim during compaction (default: 2)",
}), }),
preserve_recent_tokens: Schema.optional(NonNegativeInt).annotate({ preserve_recent_tokens: Schema.optional(NonNegativeInt).annotate({
description: "Maximum number of tokens from recent turns to serialize into the compaction summary", description: "Maximum number of tokens from recent turns to preserve verbatim after compaction",
}), }),
reserved: Schema.optional(NonNegativeInt).annotate({ reserved: Schema.optional(NonNegativeInt).annotate({
description: "Token buffer for compaction. Leaves enough window to avoid overflow during compaction.", description: "Token buffer for compaction. Leaves enough window to avoid overflow during compaction.",

View file

@ -79,10 +79,12 @@ Rules:
type Turn = { type Turn = {
start: number start: number
end: number end: number
id: MessageID
} }
type Tail = { type Tail = {
start: number start: number
id: MessageID
} }
type CompletedCompaction = { type CompletedCompaction = {
@ -119,41 +121,19 @@ function completedCompactions(messages: MessageV2.WithParts[]) {
}) })
} }
function buildPrompt(input: { previousSummary?: string; context: string[]; tail?: string }) { function buildPrompt(input: { previousSummary?: string; context: string[] }) {
const source = input.tail
? "the conversation history above and the serialized recent conversation tail below"
: "the conversation history above"
const anchor = input.previousSummary const anchor = input.previousSummary
? [ ? [
`Update the anchored summary below using ${source}.`, "Update the anchored summary below using the conversation history above.",
"Preserve still-true details, remove stale details, and merge in the new facts.", "Preserve still-true details, remove stale details, and merge in the new facts.",
"<previous-summary>", "<previous-summary>",
input.previousSummary, input.previousSummary,
"</previous-summary>", "</previous-summary>",
].join("\n") ].join("\n")
: `Create a new anchored summary from ${source}.` : "Create a new anchored summary from the conversation history above."
const tail = input.tail return [anchor, SUMMARY_TEMPLATE, ...input.context].join("\n\n")
? [
"Fold this serialized recent conversation tail into the summary; it is not provider message history.",
"<recent-conversation-tail>",
input.tail,
"</recent-conversation-tail>",
].join("\n")
: undefined
return [anchor, ...(tail ? [tail] : []), SUMMARY_TEMPLATE, ...input.context].join("\n\n")
} }
const serialize = Effect.fn("SessionCompaction.serialize")(function* (input: {
messages: MessageV2.WithParts[]
model: Provider.Model
}) {
const messages = yield* MessageV2.toModelMessagesEffect(input.messages, input.model, {
stripMedia: true,
toolOutputMaxChars: TOOL_OUTPUT_MAX_CHARS,
})
return messages.length ? JSON.stringify(messages, null, 2) : undefined
})
function preserveRecentBudget(input: { cfg: Config.Info; model: Provider.Model }) { function preserveRecentBudget(input: { cfg: Config.Info; model: Provider.Model }) {
return ( return (
input.cfg.compaction?.preserve_recent_tokens ?? input.cfg.compaction?.preserve_recent_tokens ??
@ -170,6 +150,7 @@ function turns(messages: MessageV2.WithParts[]) {
result.push({ result.push({
start: i, start: i,
end: messages.length, end: messages.length,
id: msg.info.id,
}) })
} }
for (let i = 0; i < result.length - 1; i++) { for (let i = 0; i < result.length - 1; i++) {
@ -196,6 +177,7 @@ function splitTurn(input: {
if (size > input.budget) continue if (size > input.budget) continue
return { return {
start, start,
id: input.messages[start]!.info.id,
} satisfies Tail } satisfies Tail
} }
return undefined return undefined
@ -262,7 +244,8 @@ export const layer: Layer.Layer<
messages: MessageV2.WithParts[] messages: MessageV2.WithParts[]
model: Provider.Model model: Provider.Model
}) { }) {
return Token.estimate((yield* serialize(input)) ?? "") const msgs = yield* MessageV2.toModelMessagesEffect(input.messages, input.model)
return Token.estimate(JSON.stringify(msgs))
}) })
const select = Effect.fn("SessionCompaction.select")(function* (input: { const select = Effect.fn("SessionCompaction.select")(function* (input: {
@ -271,10 +254,10 @@ export const layer: Layer.Layer<
model: Provider.Model model: Provider.Model
}) { }) {
const limit = input.cfg.compaction?.tail_turns ?? DEFAULT_TAIL_TURNS const limit = input.cfg.compaction?.tail_turns ?? DEFAULT_TAIL_TURNS
if (limit <= 0) return { head: input.messages, tail: [] } if (limit <= 0) return { head: input.messages, tail_start_id: undefined }
const budget = preserveRecentBudget({ cfg: input.cfg, model: input.model }) const budget = preserveRecentBudget({ cfg: input.cfg, model: input.model })
const all = turns(input.messages) const all = turns(input.messages)
if (!all.length) return { head: input.messages, tail: [] } if (!all.length) return { head: input.messages, tail_start_id: undefined }
const recent = all.slice(-limit) const recent = all.slice(-limit)
const sizes = yield* Effect.forEach( const sizes = yield* Effect.forEach(
recent, recent,
@ -293,7 +276,7 @@ export const layer: Layer.Layer<
const size = sizes[i] const size = sizes[i]
if (total + size <= budget) { if (total + size <= budget) {
total += size total += size
keep = { start: turn.start } keep = { start: turn.start, id: turn.id }
continue continue
} }
const remaining = budget - total const remaining = budget - total
@ -309,10 +292,10 @@ export const layer: Layer.Layer<
break break
} }
if (!keep) return { head: input.messages, tail: [] } if (!keep || keep.start === 0) return { head: input.messages, tail_start_id: undefined }
return { return {
head: input.messages.slice(0, keep.start), head: input.messages.slice(0, keep.start),
tail: input.messages.slice(keep.start), tail_start_id: keep.id,
} }
}) })
@ -423,10 +406,7 @@ export const layer: Layer.Layer<
{ sessionID: input.sessionID }, { sessionID: input.sessionID },
{ context: [], prompt: undefined }, { context: [], prompt: undefined },
) )
const tailMessages = structuredClone(selected.tail) const nextPrompt = compacting.prompt ?? buildPrompt({ previousSummary, context: compacting.context })
yield* plugin.trigger("experimental.chat.messages.transform", {}, { messages: tailMessages })
const tail = yield* serialize({ messages: tailMessages, model })
const nextPrompt = compacting.prompt ?? buildPrompt({ previousSummary, context: compacting.context, tail })
const msgs = structuredClone(selected.head) const msgs = structuredClone(selected.head)
yield* plugin.trigger("experimental.chat.messages.transform", {}, { messages: msgs }) yield* plugin.trigger("experimental.chat.messages.transform", {}, { messages: msgs })
const modelMessages = yield* MessageV2.toModelMessagesEffect(msgs, model, { const modelMessages = yield* MessageV2.toModelMessagesEffect(msgs, model, {
@ -493,6 +473,13 @@ export const layer: Layer.Layer<
return "stop" return "stop"
} }
if (compactionPart && selected.tail_start_id && compactionPart.tail_start_id !== selected.tail_start_id) {
yield* session.updatePart({
...compactionPart,
tail_start_id: selected.tail_start_id,
})
}
if (result === "continue" && input.auto) { if (result === "continue" && input.auto) {
if (replay) { if (replay) {
const original = replay.info const original = replay.info
@ -588,6 +575,7 @@ export const layer: Layer.Layer<
sessionID: input.sessionID, sessionID: input.sessionID,
timestamp: DateTime.makeUnsafe(Date.now()), timestamp: DateTime.makeUnsafe(Date.now()),
text: summary ?? "", text: summary ?? "",
include: selected.tail_start_id,
}) })
} }
yield* bus.publish(Event.Compacted, { sessionID: input.sessionID }) yield* bus.publish(Event.Compacted, { sessionID: input.sessionID })

View file

@ -772,13 +772,12 @@ export const toModelMessagesEffect = Effect.fnUntraced(function* (
return part.metadata?.anthropic?.signature != null return part.metadata?.anthropic?.signature != null
}) })
for (const part of msg.parts) { for (const part of msg.parts) {
if (msg.info.summary && part.type !== "text") continue
if (part.type === "text") { if (part.type === "text") {
const text = part.text === "" && hasSignedReasoning ? " " : part.text const text = part.text === "" && hasSignedReasoning ? " " : part.text
assistantMessage.parts.push({ assistantMessage.parts.push({
type: "text", type: "text",
text, text,
...(differentModel || msg.info.summary ? {} : { providerMetadata: part.metadata }), ...(differentModel ? {} : { providerMetadata: part.metadata }),
}) })
} }
if (part.type === "step-start") if (part.type === "step-start")
@ -1004,16 +1003,53 @@ export function get(input: { sessionID: SessionID; messageID: MessageID }): With
export function filterCompacted(msgs: Iterable<WithParts>) { export function filterCompacted(msgs: Iterable<WithParts>) {
const result = [] as WithParts[] const result = [] as WithParts[]
const completed = new Set<string>() const completed = new Set<string>()
let retain: MessageID | undefined
for (const msg of msgs) { for (const msg of msgs) {
result.push(msg) result.push(msg)
if (msg.info.role === "user" && completed.has(msg.info.id)) { if (retain) {
if (msg.parts.some((item): item is CompactionPart => item.type === "compaction")) break if (msg.info.id === retain) break
continue continue
} }
if (msg.info.role === "user" && completed.has(msg.info.id)) {
const part = msg.parts.find((item): item is CompactionPart => item.type === "compaction")
if (!part) continue
if (!part.tail_start_id) break
retain = part.tail_start_id
if (msg.info.id === retain) break
continue
}
if (msg.info.role === "user" && completed.has(msg.info.id) && msg.parts.some((part) => part.type === "compaction"))
break
if (msg.info.role === "assistant" && msg.info.summary && msg.info.finish && !msg.info.error) if (msg.info.role === "assistant" && msg.info.summary && msg.info.finish && !msg.info.error)
completed.add(msg.info.parentID) completed.add(msg.info.parentID)
} }
result.reverse() result.reverse()
const compactionIndex = result.findLastIndex(
(msg) =>
msg.info.role === "user" &&
msg.parts.some((item): item is CompactionPart => item.type === "compaction" && item.tail_start_id !== undefined),
)
const compaction = result[compactionIndex]
const part = compaction?.parts.find(
(item): item is CompactionPart => item.type === "compaction" && item.tail_start_id !== undefined,
)
const summaryIndex = compaction
? result.findIndex(
(msg, index) =>
index > compactionIndex &&
msg.info.role === "assistant" &&
msg.info.summary &&
msg.info.parentID === compaction.info.id,
)
: -1
const tailIndex = part?.tail_start_id ? result.findIndex((msg) => msg.info.id === part.tail_start_id) : -1
if (tailIndex >= 0 && tailIndex < compactionIndex && summaryIndex > compactionIndex) {
return [
...result.slice(compactionIndex, summaryIndex + 1),
...result.slice(tailIndex, compactionIndex),
...result.slice(summaryIndex + 1),
]
}
return result return result
} }

View file

@ -926,12 +926,12 @@ describe("session.compaction.process", () => {
) )
itCompaction.instance( itCompaction.instance(
"does not persist tail_start_id for serialized recent turns", "persists tail_start_id for retained recent turns",
Effect.gen(function* () { Effect.gen(function* () {
const ssn = yield* SessionNs.Service const ssn = yield* SessionNs.Service
const session = yield* ssn.create({}) const session = yield* ssn.create({})
yield* createUserMessage(session.id, "first") yield* createUserMessage(session.id, "first")
yield* createUserMessage(session.id, "second") const keep = yield* createUserMessage(session.id, "second")
yield* createUserMessage(session.id, "third") yield* createUserMessage(session.id, "third")
yield* createSummaryCompaction(session.id) yield* createSummaryCompaction(session.id)
@ -947,18 +947,18 @@ describe("session.compaction.process", () => {
const part = yield* readCompactionPart(session.id) const part = yield* readCompactionPart(session.id)
expect(part?.type).toBe("compaction") expect(part?.type).toBe("compaction")
expect(part?.tail_start_id).toBeUndefined() expect(part?.tail_start_id).toBe(keep.id)
}).pipe(withCompaction({ config: cfg({ tail_turns: 2, preserve_recent_tokens: 10_000 }) })), }).pipe(withCompaction({ config: cfg({ tail_turns: 2, preserve_recent_tokens: 10_000 }) })),
) )
itCompaction.instance( itCompaction.instance(
"does not persist tail_start_id when shrinking serialized tail", "shrinks retained tail to fit preserve token budget",
Effect.gen(function* () { Effect.gen(function* () {
const ssn = yield* SessionNs.Service const ssn = yield* SessionNs.Service
const session = yield* ssn.create({}) const session = yield* ssn.create({})
yield* createUserMessage(session.id, "first") yield* createUserMessage(session.id, "first")
yield* createUserMessage(session.id, "x".repeat(2_000)) yield* createUserMessage(session.id, "x".repeat(2_000))
yield* createUserMessage(session.id, "tiny") const keep = yield* createUserMessage(session.id, "tiny")
yield* createSummaryCompaction(session.id) yield* createSummaryCompaction(session.id)
const msgs = yield* ssn.messages({ sessionID: session.id }) const msgs = yield* ssn.messages({ sessionID: session.id })
@ -973,7 +973,7 @@ describe("session.compaction.process", () => {
const part = yield* readCompactionPart(session.id) const part = yield* readCompactionPart(session.id)
expect(part?.type).toBe("compaction") expect(part?.type).toBe("compaction")
expect(part?.tail_start_id).toBeUndefined() expect(part?.tail_start_id).toBe(keep.id)
}).pipe(withCompaction({ config: cfg({ tail_turns: 2, preserve_recent_tokens: 100 }) })), }).pipe(withCompaction({ config: cfg({ tail_turns: 2, preserve_recent_tokens: 100 }) })),
) )
@ -1005,7 +1005,7 @@ describe("session.compaction.process", () => {
) )
itCompaction.instance( itCompaction.instance(
"serializes retained tail media as text in the summary input", "falls back to full summary when retained tail media exceeds preserve token budget",
() => { () => {
const stub = llm() const stub = llm()
let captured = "" let captured = ""
@ -1078,16 +1078,15 @@ describe("session.compaction.process", () => {
const part = yield* readCompactionPart(session.id) const part = yield* readCompactionPart(session.id)
expect(part?.type).toBe("compaction") expect(part?.type).toBe("compaction")
expect(part?.tail_start_id).toBeUndefined() expect(part?.tail_start_id).toBe(keep.id)
expect(captured).toContain("zzzz") expect(captured).toContain("zzzz")
expect(captured).toContain("keep tail") expect(captured).not.toContain("keep tail")
const filtered = MessageV2.filterCompacted(MessageV2.stream(session.id)) const filtered = MessageV2.filterCompacted(MessageV2.stream(session.id))
expect(filtered.map((msg) => msg.info.id)).toEqual([parent!, expect.any(String)]) expect(filtered.map((msg) => msg.info.id).slice(0, 3)).toEqual([parent!, expect.any(String), keep.id])
expect(filtered[1]?.info.role).toBe("assistant") expect(filtered[1]?.info.role).toBe("assistant")
expect(filtered[1]?.info.role === "assistant" ? filtered[1].info.summary : false).toBe(true) expect(filtered[1]?.info.role === "assistant" ? filtered[1].info.summary : false).toBe(true)
expect(filtered.map((msg) => msg.info.id)).not.toContain(large.id) expect(filtered.map((msg) => msg.info.id)).not.toContain(large.id)
expect(filtered.map((msg) => msg.info.id)).not.toContain(keep.id)
}).pipe(withCompaction({ llm: stub.layer, config: cfg({ tail_turns: 1, preserve_recent_tokens: 100 }) })) }).pipe(withCompaction({ llm: stub.layer, config: cfg({ tail_turns: 1, preserve_recent_tokens: 100 }) }))
}, },
{ git: true }, { git: true },
@ -1354,13 +1353,13 @@ describe("session.compaction.process", () => {
) )
itCompaction.instance( itCompaction.instance(
"summarizes the head while serializing recent tail into summary input", "summarizes only the head while keeping recent tail out of summary input",
() => { () => {
const stub = llm() const stub = llm()
let captured: LLM.StreamInput["messages"] = [] let captured = ""
stub.push( stub.push(
reply("summary", (input) => { reply("summary", (input) => {
captured = input.messages captured = JSON.stringify(input.messages)
}), }),
) )
return Effect.gen(function* () { return Effect.gen(function* () {
@ -1381,15 +1380,10 @@ describe("session.compaction.process", () => {
auto: false, auto: false,
}) })
const head = JSON.stringify(captured.slice(0, -1)) expect(captured).toContain("older context")
const prompt = JSON.stringify(captured.at(-1)) expect(captured).not.toContain("keep this turn")
expect(head).toContain("older context") expect(captured).not.toContain("and this one too")
expect(head).not.toContain("keep this turn") expect(captured).not.toContain("What did we do so far?")
expect(head).not.toContain("and this one too")
expect(prompt).toContain("keep this turn")
expect(prompt).toContain("and this one too")
expect(prompt).toContain("recent-conversation-tail")
expect(prompt).not.toContain("What did we do so far?")
}).pipe(withCompaction({ llm: stub.layer })) }).pipe(withCompaction({ llm: stub.layer }))
}, },
{ git: true }, { git: true },
@ -1437,7 +1431,7 @@ describe("session.compaction.process", () => {
{ git: true }, { git: true },
) )
itCompaction.instance("does not replay recent pre-compaction turns across repeated compactions", () => { itCompaction.instance("keeps recent pre-compaction turns across repeated compactions", () => {
const stub = llm() const stub = llm()
stub.push(reply("summary one")) stub.push(reply("summary one"))
stub.push(reply("summary two")) stub.push(reply("summary two"))
@ -1468,8 +1462,8 @@ describe("session.compaction.process", () => {
expect(ids).not.toContain(u1.id) expect(ids).not.toContain(u1.id)
expect(ids).not.toContain(u2.id) expect(ids).not.toContain(u2.id)
expect(ids).not.toContain(u3.id) expect(ids).toContain(u3.id)
expect(ids).not.toContain(u4.id) expect(ids).toContain(u4.id)
expect(filtered.some((msg) => msg.info.role === "assistant" && msg.info.summary)).toBe(true) expect(filtered.some((msg) => msg.info.role === "assistant" && msg.info.summary)).toBe(true)
expect( expect(
filtered.some((msg) => msg.info.role === "user" && msg.parts.some((part) => part.type === "compaction")), filtered.some((msg) => msg.info.role === "user" && msg.parts.some((part) => part.type === "compaction")),
@ -1478,7 +1472,7 @@ describe("session.compaction.process", () => {
}) })
itCompaction.instance( itCompaction.instance(
"ignores previous summaries when sizing the serialized tail", "ignores previous summaries when sizing the retained tail",
Effect.gen(function* () { Effect.gen(function* () {
const ssn = yield* SessionNs.Service const ssn = yield* SessionNs.Service
const test = yield* TestInstance const test = yield* TestInstance
@ -1517,7 +1511,7 @@ describe("session.compaction.process", () => {
const part = yield* readCompactionPart(session.id) const part = yield* readCompactionPart(session.id)
expect(part?.type).toBe("compaction") expect(part?.type).toBe("compaction")
expect(part?.tail_start_id).toBeUndefined() expect(part?.tail_start_id).toBe(keep.id)
}).pipe(withCompaction({ config: cfg({ tail_turns: 2, preserve_recent_tokens: 500 }) })), }).pipe(withCompaction({ config: cfg({ tail_turns: 2, preserve_recent_tokens: 500 }) })),
) )
}) })

View file

@ -650,7 +650,7 @@ describe("MessageV2.filterCompacted", () => {
), ),
) )
it.instance("ignores original tail when compaction stores tail_start_id", () => it.instance("retains original tail when compaction stores tail_start_id", () =>
withSession(({ session, sessionID }) => withSession(({ session, sessionID }) =>
Effect.gen(function* () { Effect.gen(function* () {
const u1 = yield* addUser(sessionID, "first") const u1 = yield* addUser(sessionID, "first")
@ -696,12 +696,12 @@ describe("MessageV2.filterCompacted", () => {
const result = MessageV2.filterCompacted(MessageV2.stream(sessionID)) const result = MessageV2.filterCompacted(MessageV2.stream(sessionID))
expect(result.map((item) => item.info.id)).toEqual([c1, s1, u3, a3]) expect(result.map((item) => item.info.id)).toEqual([c1, s1, u2, a2, u3, a3])
}), }),
), ),
) )
it.instance("fork keeps legacy tail_start_id without replaying the tail", () => it.instance("fork remaps compaction tail_start_id for filterCompacted", () =>
Effect.gen(function* () { Effect.gen(function* () {
const session = yield* SessionNs.Service const session = yield* SessionNs.Service
const created = yield* session.create({}) const created = yield* session.create({})
@ -748,7 +748,7 @@ describe("MessageV2.filterCompacted", () => {
}) })
const parentFiltered = MessageV2.filterCompacted(MessageV2.stream(created.id)) const parentFiltered = MessageV2.filterCompacted(MessageV2.stream(created.id))
expect(parentFiltered.map((item) => item.info.id)).toEqual([c1, s1, u3, a3]) expect(parentFiltered.map((item) => item.info.id)).toEqual([c1, s1, u2, a2, u3, a3])
const forked = yield* session.fork({ sessionID: created.id }) const forked = yield* session.fork({ sessionID: created.id })
const childFiltered = MessageV2.filterCompacted(MessageV2.stream(forked.id)) const childFiltered = MessageV2.filterCompacted(MessageV2.stream(forked.id))
@ -758,14 +758,14 @@ describe("MessageV2.filterCompacted", () => {
expect(tailPart?.type).toBe("compaction") expect(tailPart?.type).toBe("compaction")
if (!tailPart || tailPart.type !== "compaction") throw new Error("Expected forked compaction part") if (!tailPart || tailPart.type !== "compaction") throw new Error("Expected forked compaction part")
expect(tailPart.tail_start_id).toBeDefined() expect(tailPart.tail_start_id).toBeDefined()
expect(childFiltered.some((m) => m.info.id === tailPart.tail_start_id)).toBe(false) expect(childFiltered.some((m) => m.info.id === tailPart.tail_start_id)).toBe(true)
yield* session.remove(forked.id) yield* session.remove(forked.id)
yield* session.remove(created.id) yield* session.remove(created.id)
}), }),
) )
it.instance("does not replay an assistant tail when compaction starts inside a turn", () => it.instance("retains an assistant tail when compaction starts inside a turn", () =>
withSession(({ session, sessionID }) => withSession(({ session, sessionID }) =>
Effect.gen(function* () { Effect.gen(function* () {
const u1 = yield* addUser(sessionID, "first") const u1 = yield* addUser(sessionID, "first")
@ -819,7 +819,7 @@ describe("MessageV2.filterCompacted", () => {
const result = MessageV2.filterCompacted(MessageV2.stream(sessionID)) const result = MessageV2.filterCompacted(MessageV2.stream(sessionID))
expect(result.map((item) => item.info.id)).toEqual([c1, s1, u3, a4]) expect(result.map((item) => item.info.id)).toEqual([c1, s1, a3, u3, a4])
}), }),
), ),
) )
@ -891,7 +891,7 @@ describe("MessageV2.filterCompacted", () => {
const result = MessageV2.filterCompacted(MessageV2.stream(sessionID)) const result = MessageV2.filterCompacted(MessageV2.stream(sessionID))
expect(result.map((item) => item.info.id)).toEqual([c2, s2, u4, a4]) expect(result.map((item) => item.info.id)).toEqual([c2, s2, u3, a3, u4, a4])
}), }),
), ),
) )