diff --git a/packages/pi-plugin/src/commands/ctx-wrapup.test.ts b/packages/pi-plugin/src/commands/ctx-wrapup.test.ts index be34a02bc..2d6a361be 100644 --- a/packages/pi-plugin/src/commands/ctx-wrapup.test.ts +++ b/packages/pi-plugin/src/commands/ctx-wrapup.test.ts @@ -39,6 +39,7 @@ import { clearProducerModelObservations, observeProducerModelsForTest, } from "@magic-context/core/hooks/magic-context/producer-window-test-support"; +import { hasRawMessageProvider } from "@magic-context/core/hooks/magic-context/read-session-chunk"; import * as logger from "@magic-context/core/shared/logger"; import { Database } from "@magic-context/core/shared/sqlite"; import { closeQuietly } from "@magic-context/core/shared/sqlite-helpers"; @@ -485,10 +486,13 @@ describe("Pi /ctx-wrapup", () => { } }); - it("persists tokens when wrapup uses the real Pi historian", async () => { + it.each([ + [8, 100_000], + [24, 100], + ])("persists tokens and covers the full target with the real Pi historian (%i messages, %i tokens)", async (messageCount, historianChunkTokens) => { const db = createDb(); try { - const sessionId = "pi-wrapup-persisted-tokens"; + const sessionId = `pi-wrapup-persisted-tokens-${messageCount}`; const runner = { harness: "pi", run: mock(async (options: SubagentRunOptions) => { @@ -510,9 +514,16 @@ describe("Pi /ctx-wrapup", () => { const range = ranges.at(-1); if (!range) throw new Error("historian prompt did not include a message range"); + const start = Number(range[1]); + const end = Number(range[2]); + // Non-final chunks need lookahead so the runner can discard only the last compartment. + const head = + start < end + ? `Summarized the eligible Pi history.` + : ""; return { ok: true as const, - assistantText: `Summarized the eligible Pi history.`, + assistantText: `${head}Summarized the last message.`, durationMs: 1, }; }), @@ -523,17 +534,42 @@ describe("Pi /ctx-wrapup", () => { deps(db, { runner, runPiHistorianForWrapup: undefined, - historianChunkTokens: 100_000, + historianChunkTokens, }), - ctx(sessionId, 8), + ctx( + sessionId, + branch(messageCount).map((entry, index) => ({ + ...entry, + message: { + role: index % 2 === 0 ? "user" : "assistant", + content: [{ type: "text", text: entry.message.content }], + }, + })), + ), sessionId, 2, ); expect(result).toContain("## Magic Wrapup"); expect(result).not.toContain("## Magic Wrapup — Partial"); + expect(getLastCompartmentEndMessage(db, sessionId)).toBe( + messageCount - 2, + ); + expect(getPendingPiCompactionMarkerState(db, sessionId)?.ordinal).toBe( + messageCount - 2, + ); + expect(getWrapupInProgressState(db, sessionId)).toBeNull(); + expect(hasRawMessageProvider(sessionId)).toBe(false); + const compartments = getCompartments(db, sessionId); + expect(compartments[0].startMessage).toBe(1); + for (let i = 1; i < compartments.length; i++) { + expect(compartments[i].startMessage).toBe( + compartments[i - 1].endMessage + 1, + ); + } const rows = getSubagentInvocations(db, sessionId); - expect(rows).toHaveLength(1); + if (messageCount === 8) expect(rows).toHaveLength(1); + else expect(rows.length).toBeGreaterThan(1); expect(rows[0]).toMatchObject({ harness: "pi", subagent: "historian", diff --git a/packages/plugin/src/hooks/magic-context/read-session-chunk.test.ts b/packages/plugin/src/hooks/magic-context/read-session-chunk.test.ts index 7a748355d..1c707bbde 100644 --- a/packages/plugin/src/hooks/magic-context/read-session-chunk.test.ts +++ b/packages/plugin/src/hooks/magic-context/read-session-chunk.test.ts @@ -11,11 +11,14 @@ import { v2NonNarrativeStoredGapRanges } from "./compartment-runner-incremental" import { validateHistorianOutput } from "./compartment-runner-validation"; import { getProtectedTailStartOrdinal, + getRawSessionMessageCount, getRawSessionMessageIdsThrough, + hasRawMessageProvider, primeTailRawMessageCache, readRawSessionMessageRange, readRawSessionMessages, readSessionChunk, + setRawMessageProvider, withRawMessageProvider, withRawSessionMessageCache, } from "./read-session-chunk"; @@ -181,6 +184,120 @@ function appendOpenCodeMessage( } } +describe("raw message provider lifecycle", () => { + const provider = { readMessages: () => [], getMessageCount: () => 7 }; + + it("keeps the outer provider after nested synchronous return and throw", () => { + const sessionId = "provider-nested-sync"; + const cleanup = setRawMessageProvider(sessionId, provider); + try { + expect(withRawMessageProvider(sessionId, provider, () => 42)).toBe(42); + expect(getRawSessionMessageCount(sessionId)).toBe(7); + expect(() => + withRawMessageProvider(sessionId, provider, () => { + throw new Error("nested failure"); + }), + ).toThrow("nested failure"); + expect(getRawSessionMessageCount(sessionId)).toBe(7); + } finally { + cleanup(); + } + expect(hasRawMessageProvider(sessionId)).toBe(false); + }); + + it("keeps the outer provider after nested async settlement and rejection", async () => { + const sessionId = "provider-nested-async"; + await withRawMessageProvider(sessionId, provider, async () => { + await withRawMessageProvider(sessionId, provider, async () => { + await Promise.resolve(); + expect(getRawSessionMessageCount(sessionId)).toBe(7); + }); + expect(getRawSessionMessageCount(sessionId)).toBe(7); + await expect( + withRawMessageProvider(sessionId, provider, async () => { + await Promise.resolve(); + throw new Error("nested rejection"); + }), + ).rejects.toThrow("nested rejection"); + expect(getRawSessionMessageCount(sessionId)).toBe(7); + }); + expect(hasRawMessageProvider(sessionId)).toBe(false); + }); + + it("keeps a shared registration until every scope cleans up, exactly once", () => { + const sessionId = "provider-shared-cleanup"; + const outerCleanup = setRawMessageProvider(sessionId, provider); + const innerCleanup = setRawMessageProvider(sessionId, provider); + try { + outerCleanup(); + outerCleanup(); + expect(getRawSessionMessageCount(sessionId)).toBe(7); + } finally { + innerCleanup(); + outerCleanup(); + } + expect(hasRawMessageProvider(sessionId)).toBe(false); + }); + + it.each(["old-first", "new-first"])( + "never overwrites or restores a replaced provider (%s cleanup)", + (order) => { + const sessionId = `provider-replaced-${order}`; + const oldCleanup = setRawMessageProvider(sessionId, provider); + const newCleanup = setRawMessageProvider(sessionId, { + readMessages: () => [], + getMessageCount: () => 11, + }); + try { + if (order === "old-first") { + oldCleanup(); + expect(getRawSessionMessageCount(sessionId)).toBe(11); + newCleanup(); + } else { + newCleanup(); + expect(hasRawMessageProvider(sessionId)).toBe(false); + oldCleanup(); + } + expect(hasRawMessageProvider(sessionId)).toBe(false); + } finally { + newCleanup(); + oldCleanup(); + } + }, + ); + + it("does not let an old scope remove a later registration of the same object", () => { + const sessionId = "provider-reregistered"; + const oldCleanup = setRawMessageProvider(sessionId, provider); + const replacementCleanup = setRawMessageProvider(sessionId, { readMessages: () => [] }); + const latestCleanup = setRawMessageProvider(sessionId, provider); + try { + oldCleanup(); + replacementCleanup(); + expect(getRawSessionMessageCount(sessionId)).toBe(7); + } finally { + latestCleanup(); + replacementCleanup(); + oldCleanup(); + } + expect(hasRawMessageProvider(sessionId)).toBe(false); + }); + + it("owns registrations separately for each session", () => { + const firstCleanup = setRawMessageProvider("provider-session-one", provider); + const secondCleanup = setRawMessageProvider("provider-session-two", provider); + try { + firstCleanup(); + expect(hasRawMessageProvider("provider-session-one")).toBe(false); + expect(getRawSessionMessageCount("provider-session-two")).toBe(7); + } finally { + secondCleanup(); + firstCleanup(); + } + expect(hasRawMessageProvider("provider-session-two")).toBe(false); + }); +}); + describe("readSessionChunk", () => { it("reads raw OpenCode messages with stable ordinals and ids", () => { useTempDataHome("read-session-chunk-"); diff --git a/packages/plugin/src/hooks/magic-context/read-session-chunk.ts b/packages/plugin/src/hooks/magic-context/read-session-chunk.ts index e330c09d3..a7941e53f 100644 --- a/packages/plugin/src/hooks/magic-context/read-session-chunk.ts +++ b/packages/plugin/src/hooks/magic-context/read-session-chunk.ts @@ -181,7 +181,7 @@ export interface BoundedRawMessageProvider { readServedBoundaryId?: (messageId: string) => string | null; } -const sessionProviders = new Map(); +const sessionProviders = new Map(); /** * Map a stored compartment boundary id to the message id a request actually @@ -189,7 +189,7 @@ const sessionProviders = new Map(); */ export function resolveHostServedBoundaryId(sessionId: string, messageId: string): string { if (messageId.length === 0) return messageId; - return sessionProviders.get(sessionId)?.readServedBoundaryId?.(messageId) ?? messageId; + return sessionProviders.get(sessionId)?.provider.readServedBoundaryId?.(messageId) ?? messageId; } /** Whether this session has an explicit non-OpenCode raw-history source. */ @@ -202,12 +202,22 @@ export function hasRawMessageProvider(sessionId: string): boolean { * unregister function. Pass-through harnesses (OpenCode) never call * this; only Pi/future harnesses install themselves before triggering * historian. + * Re-registering the current provider shares its lifetime across scopes. + * A different provider replaces it; cleanup never restores an older source. */ export function setRawMessageProvider(sessionId: string, provider: RawMessageProvider): () => void { - sessionProviders.set(sessionId, provider); + const current = sessionProviders.get(sessionId); + const registration = current?.provider === provider ? current : { provider, scopes: 0 }; + registration.scopes += 1; + sessionProviders.set(sessionId, registration); + let active = true; return () => { - const current = sessionProviders.get(sessionId); - if (current === provider) sessionProviders.delete(sessionId); + if (!active) return; + active = false; + registration.scopes -= 1; + if (registration.scopes === 0 && sessionProviders.get(sessionId) === registration) { + sessionProviders.delete(sessionId); + } }; } @@ -346,7 +356,7 @@ export function readRawSessionMessagePage( limit: number, finalWatermark: number, ): RawMessage[] { - const provider = sessionProviders.get(sessionId); + const provider = sessionProviders.get(sessionId)?.provider; if (provider?.readMessagePage) { return provider.readMessagePage(afterOrdinal, limit, finalWatermark); } @@ -365,7 +375,7 @@ export function readRawSessionMessagePage( } export function getRawSessionMessageOrdinalCount(sessionId: string): number { - const provider = sessionProviders.get(sessionId); + const provider = sessionProviders.get(sessionId)?.provider; if (provider) { if (provider.getMessageCount) return provider.getMessageCount(); const messages = provider.readMessages(); @@ -385,7 +395,7 @@ function readRawSessionMessageRangeFromSource( fromOrdinal: number, toOrdinal: number, ): RawMessage[] { - const provider = sessionProviders.get(sessionId); + const provider = sessionProviders.get(sessionId)?.provider; if (provider?.iterateMessageRange) return [...provider.iterateMessageRange(fromOrdinal, toOrdinal)]; if (provider && !provider.readMessagePage) { @@ -450,7 +460,7 @@ export function visitRawSessionMessages( const from = Math.max(1, Math.floor(fromOrdinal)); const to = Math.floor(toOrdinal); if (to < from) return; - const provider = sessionProviders.get(sessionId); + const provider = sessionProviders.get(sessionId)?.provider; if (provider?.iterateMessageRange) { for (const message of provider.iterateMessageRange(from, to)) { if (!visit(message)) return; @@ -571,7 +581,7 @@ export function primeTailRawMessageCache(args: { // the full read (correct for the no-compartment / #132 case). if (lastCompartmentEnd < 1 || !anchorMessageId) return false; - const provider = sessionProviders.get(sessionId); + const provider = sessionProviders.get(sessionId)?.provider; if (provider) { if (!provider.readMessagePage || !provider.getMessageCount) return false; const absoluteMessageCount = provider.getMessageCount(); @@ -656,7 +666,7 @@ export function readRawSessionMessageOrdinalPage( after: RawMessageOrdinalAnchor | null, limit: number, ): RawMessageOrdinalEntry[] { - const provider = sessionProviders.get(sessionId); + const provider = sessionProviders.get(sessionId)?.provider; if (provider?.readMessageOrdinalPage) return provider.readMessageOrdinalPage(after, limit); if (provider) { const rows = provider @@ -686,7 +696,7 @@ export function readRawSessionMessageOrdinalPage( } export function getRawSessionStoredMessageCount(sessionId: string): number { - const provider = sessionProviders.get(sessionId); + const provider = sessionProviders.get(sessionId)?.provider; if (provider?.getStoredMessageCount) return provider.getStoredMessageCount(); if (provider) return provider.readMessages().length; if (!openCodeDbExists()) return 0; @@ -701,7 +711,7 @@ export function readRawSessionMessageIdOrdinalsForRange( const from = Math.max(1, Math.floor(fromOrdinal)); const to = Math.floor(toOrdinal); if (to < from) return new Map(); - const provider = sessionProviders.get(sessionId); + const provider = sessionProviders.get(sessionId)?.provider; if (provider?.readMessageIdOrdinalsForRange) { return provider.readMessageIdOrdinalsForRange(from, to); } @@ -725,7 +735,7 @@ export function readRawSessionMessagePartsById( messageId: string, onQuery?: () => void, ): RawMessageParts | null { - const provider = sessionProviders.get(sessionId); + const provider = sessionProviders.get(sessionId)?.provider; if (provider?.readMessagePartsById) return provider.readMessagePartsById(messageId); if (provider?.readMessageById) return provider.readMessageById(messageId); if (provider) { @@ -738,7 +748,7 @@ export function readRawSessionMessagePartsById( } export function hasRawSessionMessageById(sessionId: string, messageId: string): boolean { - const provider = sessionProviders.get(sessionId); + const provider = sessionProviders.get(sessionId)?.provider; if (provider?.hasMessageById) return provider.hasMessageById(messageId); return readRawSessionMessageById(sessionId, messageId) !== null; } @@ -747,7 +757,7 @@ export function readRawSessionMessageOrdinalById( sessionId: string, messageId: string, ): number | null { - const provider = sessionProviders.get(sessionId); + const provider = sessionProviders.get(sessionId)?.provider; if (provider?.readMessageOrdinalById) { return provider.readMessageOrdinalById(messageId); } @@ -797,7 +807,7 @@ export function compareRawSessionMessageOrder( leftId: string, rightId: string, ): number | null { - const provider = sessionProviders.get(sessionId); + const provider = sessionProviders.get(sessionId)?.provider; if (provider) { if (!provider.readMessageOrdinalById) return null; const left = provider.readMessageOrdinalById(leftId); @@ -833,7 +843,7 @@ export function compareRawSessionMessageOrder( } export function readRawSessionMessageById(sessionId: string, messageId: string): RawMessage | null { - const provider = sessionProviders.get(sessionId); + const provider = sessionProviders.get(sessionId)?.provider; if (provider?.readMessageById) { return provider.readMessageById(messageId); } @@ -845,7 +855,7 @@ export function readRawSessionMessageById(sessionId: string, messageId: string): } function readRawSessionMessagesFromSource(sessionId: string): RawMessage[] { - const provider = sessionProviders.get(sessionId); + const provider = sessionProviders.get(sessionId)?.provider; if (provider) return provider.readMessages(); // No provider: fall back to OpenCode's session DB — but only if it exists. // A Pi-only install has no opencode.db, and a Pi transform whose provider @@ -857,7 +867,7 @@ function readRawSessionMessagesFromSource(sessionId: string): RawMessage[] { } export function getRawSessionMessageCount(sessionId: string): number { - const provider = sessionProviders.get(sessionId); + const provider = sessionProviders.get(sessionId)?.provider; if (provider) { if (provider.getMessageCount) return provider.getMessageCount(); const messages = provider.readMessages(); @@ -1367,7 +1377,7 @@ export function readRawSessionSeedTail( boundaryId: string | null, onQuery?: () => void, ): Map { - const provider = sessionProviders.get(sessionId); + const provider = sessionProviders.get(sessionId)?.provider; if (provider) { const boundaryOrdinal = boundaryId === null ? 1 : readRawSessionMessageOrdinalById(sessionId, boundaryId);