Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
82 changes: 63 additions & 19 deletions apps/server/src/provider/Layers/OpenCodeAdapter.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -84,6 +84,7 @@ const runtimeMock = {
| ((sessionID: string) => Promise<Array<{ id: string }>>)
| 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,
Expand Down Expand Up @@ -143,6 +144,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;
Expand Down Expand Up @@ -263,6 +265,9 @@ const OpenCodeRuntimeTestDouble: OpenCodeRuntimeShape = {
return {
data: {
id: sessionID,
...(runtimeMock.state.revertMessageID
? { revert: { messageID: runtimeMock.state.revertMessageID } }
: {}),
...(directory ? { directory } : {}),
...(parentID ? { parentID } : {}),
},
Expand Down Expand Up @@ -372,17 +377,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: {
Expand Down Expand Up @@ -6316,7 +6320,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");
Expand All @@ -6327,22 +6331,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, []);
}),
);

Expand Down
36 changes: 18 additions & 18 deletions apps/server/src/provider/Layers/OpenCodeAdapter.ts
Original file line number Diff line number Diff line change
Expand Up @@ -3750,6 +3750,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,
Expand All @@ -3758,6 +3761,7 @@ export function makeOpenCodeAdapter(

const turns: Array<OpenCodeTurnSnapshot> = [];
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),
Expand All @@ -3776,25 +3780,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;
},
);

Expand Down
Loading