diff --git a/apps/server/src/provider/Layers/OpenCodeAdapter.test.ts b/apps/server/src/provider/Layers/OpenCodeAdapter.test.ts index 1720ee736..67c6f371c 100644 --- a/apps/server/src/provider/Layers/OpenCodeAdapter.test.ts +++ b/apps/server/src/provider/Layers/OpenCodeAdapter.test.ts @@ -80,6 +80,7 @@ const runtimeMock = { | ((sessionID: string) => Promise>) | null, closeCalls: [] as string[], + revertMessageID: undefined as string | undefined, revertCalls: [] as Array<{ sessionID: string; messageID?: string }>, messageCalls: [] as Array<{ sessionID: string; messageID: string }>, messageFailures: 0, @@ -138,6 +139,7 @@ const runtimeMock = { this.state.sessionChildrenById.clear(); this.state.sessionChildrenImplementation = null; this.state.closeCalls.length = 0; + this.state.revertMessageID = undefined; this.state.revertCalls.length = 0; this.state.messageCalls.length = 0; this.state.messageFailures = 0; @@ -257,6 +259,9 @@ const OpenCodeRuntimeTestDouble: OpenCodeRuntimeShape = { return { data: { id: sessionID, + ...(runtimeMock.state.revertMessageID + ? { revert: { messageID: runtimeMock.state.revertMessageID } } + : {}), ...(directory ? { directory } : {}), ...(parentID ? { parentID } : {}), }, @@ -362,17 +367,16 @@ const OpenCodeRuntimeTestDouble: OpenCodeRuntimeShape = { ...(messageID ? { messageID } : {}), }); if (!messageID) { - runtimeMock.state.messages = []; - return; + throw new Error("Expected messageID"); + } + let lastUserID: string | undefined; + for (const entry of runtimeMock.state.messages) { + if (entry.info.role === "user") lastUserID = entry.info.id; + if (entry.info.id === messageID && entry.parts.length > 0) { + runtimeMock.state.revertMessageID = lastUserID ?? messageID; + break; + } } - - const targetIndex = runtimeMock.state.messages.findIndex( - (entry) => entry.info.id === messageID, - ); - runtimeMock.state.messages = - targetIndex >= 0 - ? runtimeMock.state.messages.slice(0, targetIndex + 1) - : runtimeMock.state.messages; }, }, event: { @@ -5530,7 +5534,7 @@ it.layer(OpenCodeAdapterTestLayer)("OpenCodeAdapterLive", (it) => { }).pipe(Effect.provide(adapterLayer)); }); - it.effect("reverts the full thread when rollback removes every assistant turn", () => + it.effect("reverts the first removed assistant message and returns only retained turns", () => Effect.gen(function* () { const adapter = yield* OpenCodeAdapter; const threadId = asThreadId("thread-rollback-all"); @@ -5541,22 +5545,62 @@ it.layer(OpenCodeAdapterTestLayer)("OpenCodeAdapterLive", (it) => { }); runtimeMock.state.messages = [ + { info: { id: "user-1", role: "user" }, parts: [] }, { info: { id: "assistant-1", role: "assistant" }, - parts: [], + parts: [{ id: "part-1", type: "text", text: "first answer" }], }, + { info: { id: "user-2", role: "user" }, parts: [] }, { info: { id: "assistant-2", role: "assistant" }, - parts: [], + parts: [{ id: "part-2", type: "text", text: "second answer" }], }, ]; - const snapshot = yield* adapter.rollbackThread(threadId, 2); - - NodeAssert.deepEqual(runtimeMock.state.revertCalls, [ - { sessionID: "http://127.0.0.1:9999/session" }, - ]); - NodeAssert.deepEqual(snapshot.turns, []); + for (const numTurns of [0, 1, 2, 3]) { + runtimeMock.state.revertMessageID = undefined; + runtimeMock.state.revertCalls.length = 0; + const snapshot = yield* adapter.rollbackThread(threadId, numTurns); + NodeAssert.deepEqual( + runtimeMock.state.revertCalls, + numTurns === 0 + ? [] + : [ + { + sessionID: "http://127.0.0.1:9999/session", + messageID: numTurns === 1 ? "assistant-2" : "assistant-1", + }, + ], + ); + NodeAssert.deepEqual( + snapshot.turns.map((turn) => turn.id), + ["assistant-1", "assistant-2"].slice(0, Math.max(0, 2 - numTurns)), + ); + } + runtimeMock.state.revertMessageID = undefined; + for (const remaining of [1, 0]) { + const snapshot = yield* adapter.rollbackThread(threadId, 1); + NodeAssert.equal(snapshot.turns.length, remaining); + NodeAssert.deepEqual((yield* adapter.readThread(threadId)).turns, snapshot.turns); + } + NodeAssert.deepEqual( + runtimeMock.state.revertCalls.slice(-2).map((call) => call.messageID), + ["assistant-2", "assistant-1"], + ); + runtimeMock.state.revertMessageID = undefined; + runtimeMock.state.messages = runtimeMock.state.messages.filter( + (entry) => entry.info.id !== "user-2", + ); + const sharedUserSnapshot = yield* adapter.rollbackThread(threadId, 1); + NodeAssert.equal(runtimeMock.state.revertMessageID, "user-1"); + NodeAssert.deepEqual(sharedUserSnapshot.turns, []); + NodeAssert.deepEqual((yield* adapter.readThread(threadId)).turns, []); + + runtimeMock.state.messages = []; + runtimeMock.state.revertCalls.length = 0; + const emptySnapshot = yield* adapter.rollbackThread(threadId, 1); + NodeAssert.deepEqual(runtimeMock.state.revertCalls, []); + NodeAssert.deepEqual(emptySnapshot.turns, []); }), ); diff --git a/apps/server/src/provider/Layers/OpenCodeAdapter.ts b/apps/server/src/provider/Layers/OpenCodeAdapter.ts index 6d7a3cd8a..024b00b9a 100644 --- a/apps/server/src/provider/Layers/OpenCodeAdapter.ts +++ b/apps/server/src/provider/Layers/OpenCodeAdapter.ts @@ -3605,6 +3605,9 @@ export function makeOpenCodeAdapter( const readThread: OpenCodeAdapterShape["readThread"] = Effect.fn("readThread")( function* (threadId) { const context = yield* ensureSessionContext(sessions, threadId); + const session = yield* runOpenCodeSdk("session.get", () => + context.client.session.get({ sessionID: context.openCodeSessionId }), + ).pipe(Effect.mapError(toRequestError)); const messages = yield* runOpenCodeSdk("session.messages", () => context.client.session.messages({ sessionID: context.openCodeSessionId, @@ -3613,6 +3616,7 @@ export function makeOpenCodeAdapter( const turns: Array = []; for (const entry of messages.data ?? []) { + if (entry.info.id === session.data?.revert?.messageID) break; if (entry.info.role === "assistant") { turns.push({ id: TurnId.make(entry.info.id), @@ -3631,25 +3635,21 @@ export function makeOpenCodeAdapter( const rollbackThread: OpenCodeAdapterShape["rollbackThread"] = Effect.fn("rollbackThread")( function* (threadId, numTurns) { const context = yield* ensureSessionContext(sessions, threadId); - const messages = yield* runOpenCodeSdk("session.messages", () => - context.client.session.messages({ - sessionID: context.openCodeSessionId, - }), - ).pipe(Effect.mapError(toRequestError)); - - const assistantMessages = (messages.data ?? []).filter( - (entry) => entry.info.role === "assistant", - ); - const targetIndex = assistantMessages.length - numTurns - 1; - const target = targetIndex >= 0 ? assistantMessages[targetIndex] : null; - yield* runOpenCodeSdk("session.revert", () => - context.client.session.revert({ - sessionID: context.openCodeSessionId, - ...(target ? { messageID: target.info.id } : {}), - }), - ).pipe(Effect.mapError(toRequestError)); + const snapshot = yield* readThread(threadId); + const targetIndex = Math.max(0, snapshot.turns.length - numTurns); + const target = snapshot.turns[targetIndex]; + if (target) { + yield* runOpenCodeSdk("session.revert", () => + context.client.session.revert({ + sessionID: context.openCodeSessionId, + messageID: target.id, + }), + ).pipe(Effect.mapError(toRequestError)); + // Native revert can move the boundary to the preceding user message. + return yield* readThread(threadId); + } - return yield* readThread(threadId); + return snapshot; }, );