From f44ff03a662145577d68571005d30b61ebb01b38 Mon Sep 17 00:00:00 2001 From: Nassim Najjar Date: Wed, 22 Jul 2026 04:13:23 +0100 Subject: [PATCH] feat: fork from a selected assistant response --- .../Layers/ProjectionPipeline.ts | 50 +++++ .../Layers/ProjectionSnapshotQuery.ts | 23 ++ .../Layers/ProviderCommandReactor.test.ts | 19 ++ .../Layers/ProviderCommandReactor.ts | 55 ++++- apps/server/src/orchestration/Schemas.ts | 2 + .../src/orchestration/decider.fork.test.ts | 197 ++++++++++++++++++ apps/server/src/orchestration/decider.ts | 70 +++++++ apps/server/src/orchestration/projector.ts | 43 ++++ .../persistence/Layers/ProjectionThreads.ts | 10 + apps/server/src/persistence/Migrations.ts | 2 + .../033_ProjectionThreadForkLineage.test.ts | 29 +++ .../033_ProjectionThreadForkLineage.ts | 13 ++ .../persistence/Services/ProjectionThreads.ts | 2 + .../src/provider/Layers/ClaudeAdapter.test.ts | 106 +++++++++- .../src/provider/Layers/ClaudeAdapter.ts | 109 +++++++++- .../src/provider/Layers/CodexAdapter.ts | 11 + .../Layers/CodexSessionRuntime.test.ts | 58 +++++- .../provider/Layers/CodexSessionRuntime.ts | 32 ++- .../provider/Layers/OpenCodeAdapter.test.ts | 58 +++++- .../src/provider/Layers/OpenCodeAdapter.ts | 54 +++++ .../src/provider/Layers/ProviderService.ts | 63 +++++- .../src/provider/Services/ProviderAdapter.ts | 2 + .../web/src/components/ChatView.logic.test.ts | 13 ++ apps/web/src/components/ChatView.logic.ts | 14 ++ apps/web/src/components/ChatView.tsx | 180 +++++++++++----- apps/web/src/components/RightPanelTabs.tsx | 15 +- .../chat/MessagesTimeline.logic.test.ts | 31 +++ .../components/chat/MessagesTimeline.logic.ts | 15 ++ .../components/chat/MessagesTimeline.test.tsx | 58 ++++++ .../src/components/chat/MessagesTimeline.tsx | 49 +++++ apps/web/src/rightPanelStore.test.ts | 19 ++ apps/web/src/rightPanelStore.ts | 51 ++++- .../client-runtime/src/operations/commands.ts | 13 ++ .../src/state/threadCommands.ts | 9 + packages/contracts/src/orchestration.ts | 41 ++++ packages/contracts/src/provider.ts | 12 +- 36 files changed, 1438 insertions(+), 90 deletions(-) create mode 100644 apps/server/src/orchestration/decider.fork.test.ts create mode 100644 apps/server/src/persistence/Migrations/033_ProjectionThreadForkLineage.test.ts create mode 100644 apps/server/src/persistence/Migrations/033_ProjectionThreadForkLineage.ts diff --git a/apps/server/src/orchestration/Layers/ProjectionPipeline.ts b/apps/server/src/orchestration/Layers/ProjectionPipeline.ts index f12df850941..60f939c8087 100644 --- a/apps/server/src/orchestration/Layers/ProjectionPipeline.ts +++ b/apps/server/src/orchestration/Layers/ProjectionPipeline.ts @@ -603,6 +603,8 @@ const makeOrchestrationProjectionPipeline = Effect.fn("makeOrchestrationProjecti interactionMode: event.payload.interactionMode, branch: event.payload.branch, worktreePath: event.payload.worktreePath, + forkedFromThreadId: null, + forkedFromTurnId: null, latestTurnId: null, createdAt: event.payload.createdAt, updatedAt: event.payload.updatedAt, @@ -615,6 +617,33 @@ const makeOrchestrationProjectionPipeline = Effect.fn("makeOrchestrationProjecti }); return; + case "thread.forked": + yield* projectionThreadRepository.upsert({ + threadId: event.payload.threadId, + projectId: event.payload.projectId, + title: event.payload.title, + modelSelection: event.payload.modelSelection, + runtimeMode: event.payload.runtimeMode, + interactionMode: event.payload.interactionMode, + branch: event.payload.branch, + worktreePath: event.payload.worktreePath, + forkedFromThreadId: event.payload.forkedFrom.threadId, + forkedFromTurnId: event.payload.forkedFrom.turnId, + latestTurnId: null, + createdAt: event.payload.createdAt, + updatedAt: event.payload.updatedAt, + archivedAt: null, + latestUserMessageAt: + event.payload.inheritedMessages + .toReversed() + .find((message) => message.role === "user")?.createdAt ?? null, + pendingApprovalCount: 0, + pendingUserInputCount: 0, + hasActionableProposedPlan: 0, + deletedAt: null, + }); + return; + case "thread.archived": { const existingRow = yield* projectionThreadRepository.getById({ threadId: event.payload.threadId, @@ -811,6 +840,27 @@ const makeOrchestrationProjectionPipeline = Effect.fn("makeOrchestrationProjecti "applyThreadMessagesProjection", )(function* (event, attachmentSideEffects) { switch (event.type) { + case "thread.forked": + yield* Effect.forEach( + event.payload.inheritedMessages, + (message) => + projectionThreadMessageRepository.upsert({ + messageId: message.id, + threadId: event.payload.threadId, + turnId: message.turnId, + role: message.role, + text: message.text, + ...(message.attachments !== undefined + ? { attachments: [...message.attachments] } + : {}), + isStreaming: message.streaming, + createdAt: message.createdAt, + updatedAt: message.updatedAt, + }), + { concurrency: 1, discard: true }, + ); + return; + case "thread.message-sent": { const existingMessage = yield* projectionThreadMessageRepository.getByMessageId({ messageId: event.payload.messageId, diff --git a/apps/server/src/orchestration/Layers/ProjectionSnapshotQuery.ts b/apps/server/src/orchestration/Layers/ProjectionSnapshotQuery.ts index 32210436e67..9f99246bda7 100644 --- a/apps/server/src/orchestration/Layers/ProjectionSnapshotQuery.ts +++ b/apps/server/src/orchestration/Layers/ProjectionSnapshotQuery.ts @@ -224,6 +224,15 @@ function mapSessionRow( }; } +function mapForkedFrom(row: Schema.Schema.Type) { + return row.forkedFromThreadId == null + ? null + : { + threadId: row.forkedFromThreadId, + turnId: row.forkedFromTurnId ?? null, + }; +} + function mapProjectShellRow( row: Schema.Schema.Type, repositoryIdentity: OrchestrationProject["repositoryIdentity"], @@ -330,6 +339,8 @@ const makeProjectionSnapshotQuery = Effect.gen(function* () { interaction_mode AS "interactionMode", branch, worktree_path AS "worktreePath", + forked_from_thread_id AS "forkedFromThreadId", + forked_from_turn_id AS "forkedFromTurnId", latest_turn_id AS "latestTurnId", created_at AS "createdAt", updated_at AS "updatedAt", @@ -358,6 +369,8 @@ const makeProjectionSnapshotQuery = Effect.gen(function* () { interaction_mode AS "interactionMode", branch, worktree_path AS "worktreePath", + forked_from_thread_id AS "forkedFromThreadId", + forked_from_turn_id AS "forkedFromTurnId", latest_turn_id AS "latestTurnId", created_at AS "createdAt", updated_at AS "updatedAt", @@ -388,6 +401,8 @@ const makeProjectionSnapshotQuery = Effect.gen(function* () { interaction_mode AS "interactionMode", branch, worktree_path AS "worktreePath", + forked_from_thread_id AS "forkedFromThreadId", + forked_from_turn_id AS "forkedFromTurnId", latest_turn_id AS "latestTurnId", created_at AS "createdAt", updated_at AS "updatedAt", @@ -750,6 +765,8 @@ const makeProjectionSnapshotQuery = Effect.gen(function* () { interaction_mode AS "interactionMode", branch, worktree_path AS "worktreePath", + forked_from_thread_id AS "forkedFromThreadId", + forked_from_turn_id AS "forkedFromTurnId", latest_turn_id AS "latestTurnId", created_at AS "createdAt", updated_at AS "updatedAt", @@ -1182,6 +1199,7 @@ const makeProjectionSnapshotQuery = Effect.gen(function* () { interactionMode: row.interactionMode, branch: row.branch, worktreePath: row.worktreePath, + forkedFrom: mapForkedFrom(row), latestTurn: latestTurnByThread.get(row.threadId) ?? null, createdAt: row.createdAt, updatedAt: row.updatedAt, @@ -1380,6 +1398,7 @@ const makeProjectionSnapshotQuery = Effect.gen(function* () { interactionMode: row.interactionMode, branch: row.branch, worktreePath: row.worktreePath, + forkedFrom: mapForkedFrom(row), latestTurn: latestTurnByThread.get(row.threadId) ?? null, createdAt: row.createdAt, updatedAt: row.updatedAt, @@ -1509,6 +1528,7 @@ const makeProjectionSnapshotQuery = Effect.gen(function* () { interactionMode: row.interactionMode, branch: row.branch, worktreePath: row.worktreePath, + forkedFrom: mapForkedFrom(row), latestTurn: latestTurnByThread.get(row.threadId) ?? null, createdAt: row.createdAt, updatedAt: row.updatedAt, @@ -1643,6 +1663,7 @@ const makeProjectionSnapshotQuery = Effect.gen(function* () { interactionMode: row.interactionMode, branch: row.branch, worktreePath: row.worktreePath, + forkedFrom: mapForkedFrom(row), latestTurn: latestTurnByThread.get(row.threadId) ?? null, createdAt: row.createdAt, updatedAt: row.updatedAt, @@ -1883,6 +1904,7 @@ const makeProjectionSnapshotQuery = Effect.gen(function* () { interactionMode: threadRow.value.interactionMode, branch: threadRow.value.branch, worktreePath: threadRow.value.worktreePath, + forkedFrom: mapForkedFrom(threadRow.value), latestTurn: Option.isSome(latestTurnRow) ? mapLatestTurn(latestTurnRow.value) : null, createdAt: threadRow.value.createdAt, updatedAt: threadRow.value.updatedAt, @@ -1977,6 +1999,7 @@ const makeProjectionSnapshotQuery = Effect.gen(function* () { interactionMode: threadRow.value.interactionMode, branch: threadRow.value.branch, worktreePath: threadRow.value.worktreePath, + forkedFrom: mapForkedFrom(threadRow.value), latestTurn: Option.isSome(latestTurnRow) ? mapLatestTurn(latestTurnRow.value) : null, createdAt: threadRow.value.createdAt, updatedAt: threadRow.value.updatedAt, diff --git a/apps/server/src/orchestration/Layers/ProviderCommandReactor.test.ts b/apps/server/src/orchestration/Layers/ProviderCommandReactor.test.ts index ce464565dc5..bdc73d407a6 100644 --- a/apps/server/src/orchestration/Layers/ProviderCommandReactor.test.ts +++ b/apps/server/src/orchestration/Layers/ProviderCommandReactor.test.ts @@ -48,6 +48,7 @@ import { OrchestrationEngineLive } from "./OrchestrationEngine.ts"; import { OrchestrationProjectionPipelineLive } from "./ProjectionPipeline.ts"; import { OrchestrationProjectionSnapshotQueryLive } from "./ProjectionSnapshotQuery.ts"; import { + findCompletedTurnIndex, providerErrorLabel, providerErrorLabelFromInstanceHint, ProviderCommandReactorLive, @@ -66,6 +67,24 @@ const asApprovalRequestId = (value: string): ApprovalRequestId => ApprovalReques const asMessageId = (value: string): MessageId => MessageId.make(value); const asTurnId = (value: string): TurnId => TurnId.make(value); +describe("findCompletedTurnIndex", () => { + it("maps a selected T3 turn to its provider response position", () => { + const first = asTurnId("turn-1"); + const second = asTurnId("turn-2"); + expect( + findCompletedTurnIndex( + [ + { role: "assistant", turnId: first, streaming: false }, + { role: "assistant", turnId: first, streaming: false }, + { role: "assistant", turnId: second, streaming: true }, + { role: "assistant", turnId: second, streaming: false }, + ], + second, + ), + ).toBe(1); + }); +}); + const deriveServerPathsSync = (baseDir: string, devUrl: URL | undefined) => Effect.runSync(deriveServerPaths(baseDir, devUrl).pipe(Effect.provide(NodeServices.layer))); diff --git a/apps/server/src/orchestration/Layers/ProviderCommandReactor.ts b/apps/server/src/orchestration/Layers/ProviderCommandReactor.ts index 9c7a7c94bb1..9ab0affc82a 100644 --- a/apps/server/src/orchestration/Layers/ProviderCommandReactor.ts +++ b/apps/server/src/orchestration/Layers/ProviderCommandReactor.ts @@ -88,6 +88,32 @@ const HANDLED_TURN_START_KEY_TTL = Duration.minutes(30); const DEFAULT_RUNTIME_MODE: RuntimeMode = "full-access"; const DEFAULT_THREAD_TITLE = "New thread"; +export function findCompletedTurnIndex( + messages: ReadonlyArray<{ + readonly role: string; + readonly turnId: TurnId | null; + readonly streaming: boolean; + }>, + sourceTurnId: TurnId, +): number | undefined { + const completedTurnIds: TurnId[] = []; + for (const message of messages) { + if ( + message.role !== "assistant" || + message.turnId === null || + message.streaming || + completedTurnIds.some((turnId) => turnId === message.turnId) + ) { + continue; + } + completedTurnIds.push(message.turnId); + if (message.turnId === sourceTurnId) { + return completedTurnIds.length - 1; + } + } + return undefined; +} + export function providerErrorLabel(value: string | undefined): string { const normalized = value?.trim(); return normalized && normalized.length > 0 ? normalized : "unknown"; @@ -471,19 +497,40 @@ const make = Effect.gen(function* () { projects: project ? [project] : [], }); - const startProviderSession = (input?: { + const startProviderSession = Effect.fn("startProviderSession")(function* (input?: { readonly resumeCursor?: unknown; readonly provider?: ProviderDriverKind; - }) => - providerService.startSession(threadId, { + readonly includeForkSource?: boolean; + }) { + const forkedFrom = input?.includeForkSource === true ? thread.forkedFrom : null; + const forkSource = + forkedFrom != null + ? yield* Effect.gen(function* () { + const sourceThread = yield* resolveThread(forkedFrom.threadId); + const sourceTurnId = forkedFrom.turnId; + const sourceTurnIndex = + sourceThread != null && sourceTurnId !== null + ? findCompletedTurnIndex(sourceThread.messages, sourceTurnId) + : undefined; + return { + threadId: forkedFrom.threadId, + ...(sourceTurnId !== null ? { sourceTurnId } : {}), + ...(sourceTurnIndex !== undefined ? { sourceTurnIndex } : {}), + }; + }) + : undefined; + + return yield* providerService.startSession(threadId, { threadId, ...(preferredProvider ? { provider: preferredProvider } : {}), providerInstanceId: desiredInstanceId, ...(effectiveCwd ? { cwd: effectiveCwd } : {}), modelSelection: desiredModelSelection, ...(input?.resumeCursor !== undefined ? { resumeCursor: input.resumeCursor } : {}), + ...(forkSource !== undefined ? { forkFrom: forkSource } : {}), runtimeMode: desiredRuntimeMode, }); + }); const bindSessionToThread = (session: ProviderSession) => Effect.gen(function* () { @@ -578,7 +625,7 @@ const make = Effect.gen(function* () { return restartedSession.threadId; } - const startedSession = yield* startProviderSession(undefined); + const startedSession = yield* startProviderSession({ includeForkSource: true }); yield* bindSessionToThread(startedSession); return startedSession.threadId; }); diff --git a/apps/server/src/orchestration/Schemas.ts b/apps/server/src/orchestration/Schemas.ts index f7ebf693440..be74153c8cc 100644 --- a/apps/server/src/orchestration/Schemas.ts +++ b/apps/server/src/orchestration/Schemas.ts @@ -3,6 +3,7 @@ import { ProjectMetaUpdatedPayload as ContractsProjectMetaUpdatedPayloadSchema, ProjectDeletedPayload as ContractsProjectDeletedPayloadSchema, ThreadCreatedPayload as ContractsThreadCreatedPayloadSchema, + ThreadForkedPayload as ContractsThreadForkedPayloadSchema, ThreadArchivedPayload as ContractsThreadArchivedPayloadSchema, ThreadMetaUpdatedPayload as ContractsThreadMetaUpdatedPayloadSchema, ThreadRuntimeModeSetPayload as ContractsThreadRuntimeModeSetPayloadSchema, @@ -28,6 +29,7 @@ export const ProjectMetaUpdatedPayload = ContractsProjectMetaUpdatedPayloadSchem export const ProjectDeletedPayload = ContractsProjectDeletedPayloadSchema; export const ThreadCreatedPayload = ContractsThreadCreatedPayloadSchema; +export const ThreadForkedPayload = ContractsThreadForkedPayloadSchema; export const ThreadArchivedPayload = ContractsThreadArchivedPayloadSchema; export const ThreadMetaUpdatedPayload = ContractsThreadMetaUpdatedPayloadSchema; export const ThreadRuntimeModeSetPayload = ContractsThreadRuntimeModeSetPayloadSchema; diff --git a/apps/server/src/orchestration/decider.fork.test.ts b/apps/server/src/orchestration/decider.fork.test.ts new file mode 100644 index 00000000000..df7b2a8c01c --- /dev/null +++ b/apps/server/src/orchestration/decider.fork.test.ts @@ -0,0 +1,197 @@ +import { + CommandId, + DEFAULT_PROVIDER_INTERACTION_MODE, + EventId, + MessageId, + ProjectId, + ProviderInstanceId, + ThreadId, + TurnId, + type OrchestrationEvent, + type OrchestrationReadModel, +} from "@t3tools/contracts"; +import * as NodeServices from "@effect/platform-node/NodeServices"; +import { expect, it } from "@effect/vitest"; +import * as Effect from "effect/Effect"; + +import { decideOrchestrationCommand } from "./decider.ts"; +import { createEmptyReadModel, projectEvent } from "./projector.ts"; + +const now = "2026-07-21T12:00:00.000Z"; +const sourceThreadId = ThreadId.make("thread-source"); +const forkThreadId = ThreadId.make("thread-fork"); +const turnOneId = TurnId.make("turn-1"); +const turnTwoId = TurnId.make("turn-2"); + +const seedReadModel = (): OrchestrationReadModel => ({ + ...createEmptyReadModel(now), + threads: [ + { + id: sourceThreadId, + projectId: ProjectId.make("project-1"), + title: "Source thread", + modelSelection: { + instanceId: ProviderInstanceId.make("codex"), + model: "gpt-5.4", + }, + runtimeMode: "full-access", + interactionMode: DEFAULT_PROVIDER_INTERACTION_MODE, + branch: "main", + worktreePath: "/tmp/project", + forkedFrom: null, + latestTurn: { + turnId: turnTwoId, + state: "completed", + requestedAt: now, + startedAt: now, + completedAt: now, + assistantMessageId: MessageId.make("assistant-2"), + }, + createdAt: now, + updatedAt: now, + archivedAt: null, + deletedAt: null, + messages: [ + { + id: MessageId.make("user-1"), + role: "user", + text: "First question", + attachments: [], + turnId: turnOneId, + streaming: false, + createdAt: now, + updatedAt: now, + }, + { + id: MessageId.make("assistant-1"), + role: "assistant", + text: "First answer", + attachments: [], + turnId: turnOneId, + streaming: false, + createdAt: now, + updatedAt: now, + }, + { + id: MessageId.make("user-2"), + role: "user", + text: "Second question", + attachments: [], + turnId: turnTwoId, + streaming: false, + createdAt: now, + updatedAt: now, + }, + { + id: MessageId.make("assistant-2"), + role: "assistant", + text: "Second answer", + attachments: [], + turnId: turnTwoId, + streaming: false, + createdAt: now, + updatedAt: now, + }, + { + id: MessageId.make("assistant-streaming"), + role: "assistant", + text: "Unsettled answer", + attachments: [], + turnId: TurnId.make("turn-3"), + streaming: true, + createdAt: now, + updatedAt: now, + }, + ], + proposedPlans: [], + activities: [], + checkpoints: [], + session: null, + }, + ], +}); + +const forkCommand = (sourceTurnId: TurnId) => ({ + type: "thread.fork" as const, + commandId: CommandId.make("command-fork"), + threadId: forkThreadId, + sourceThreadId, + sourceTurnId, + title: "Forked thread", + createdAt: now, +}); + +type PlannedForkEvent = Omit< + Extract, + "sequence" +>; + +function requireForkEvent(result: unknown): PlannedForkEvent { + if ( + result === null || + typeof result !== "object" || + Array.isArray(result) || + !("type" in result) || + result.type !== "thread.forked" + ) { + throw new Error("Expected one thread.forked event."); + } + return result as PlannedForkEvent; +} + +it.layer(NodeServices.layer)("thread fork decider", (it) => { + it.effect("copies history through the requested turn and persists lineage", () => + Effect.gen(function* () { + const event = requireForkEvent( + yield* decideOrchestrationCommand({ + command: forkCommand(turnOneId), + readModel: seedReadModel(), + }), + ); + + expect(event.payload.forkedFrom).toEqual({ + threadId: sourceThreadId, + turnId: turnOneId, + }); + expect(event.payload.inheritedMessages.map((message) => message.text)).toEqual([ + "First question", + "First answer", + ]); + + const projected = yield* projectEvent(seedReadModel(), { + ...event, + sequence: 1, + eventId: EventId.make("event-fork"), + }); + expect(projected.threads.find((thread) => thread.id === forkThreadId)?.forkedFrom).toEqual({ + threadId: sourceThreadId, + turnId: turnOneId, + }); + }), + ); + + it.effect("rejects a fork at a running turn", () => + Effect.gen(function* () { + const readModel = seedReadModel(); + const source = readModel.threads[0]!; + const error = yield* decideOrchestrationCommand({ + command: forkCommand(turnTwoId), + readModel: { + ...readModel, + threads: [ + { + ...source, + latestTurn: { + ...source.latestTurn!, + state: "running", + completedAt: null, + }, + }, + ], + }, + }).pipe(Effect.flip); + + expect(error.message).toContain("still running and cannot be forked"); + }), + ); +}); diff --git a/apps/server/src/orchestration/decider.ts b/apps/server/src/orchestration/decider.ts index 1730494ecc6..f2908153922 100644 --- a/apps/server/src/orchestration/decider.ts +++ b/apps/server/src/orchestration/decider.ts @@ -1,5 +1,6 @@ import { EventId, + MessageId, type OrchestrationCommand, type OrchestrationEvent, type OrchestrationReadModel, @@ -260,6 +261,75 @@ export const decideOrchestrationCommand = Effect.fn("decideOrchestrationCommand" }; } + case "thread.fork": { + const sourceThread = yield* requireThread({ + readModel, + command, + threadId: command.sourceThreadId, + }); + yield* requireThreadAbsent({ + readModel, + command, + threadId: command.threadId, + }); + + const sourceTurnId = command.sourceTurnId; + if (!sourceThread.messages.some((message) => message.turnId === sourceTurnId)) { + return yield* new OrchestrationCommandInvariantError({ + commandType: command.type, + detail: `Turn '${sourceTurnId}' was not found in source thread '${sourceThread.id}'.`, + }); + } + const latestTurn = sourceThread.latestTurn; + if ( + latestTurn?.turnId === sourceTurnId && + latestTurn.state === "running" + ) { + return yield* new OrchestrationCommandInvariantError({ + commandType: command.type, + detail: `Turn '${sourceTurnId}' is still running and cannot be forked.`, + }); + } + + const cutoffIndex = sourceThread.messages.findLastIndex( + (message) => message.turnId === sourceTurnId, + ); + const inheritedMessages = sourceThread.messages + .slice(0, cutoffIndex + 1) + .filter((message) => !message.streaming) + .map((message, index) => ({ + ...message, + id: MessageId.make(`${command.threadId}:fork:${index}`), + })); + + return { + ...(yield* withEventBase({ + aggregateKind: "thread", + aggregateId: command.threadId, + occurredAt: command.createdAt, + commandId: command.commandId, + })), + type: "thread.forked", + payload: { + threadId: command.threadId, + projectId: sourceThread.projectId, + title: command.title, + modelSelection: sourceThread.modelSelection, + runtimeMode: sourceThread.runtimeMode, + interactionMode: sourceThread.interactionMode, + branch: sourceThread.branch, + worktreePath: sourceThread.worktreePath, + forkedFrom: { + threadId: sourceThread.id, + turnId: sourceTurnId, + }, + inheritedMessages, + createdAt: command.createdAt, + updatedAt: command.createdAt, + }, + }; + } + case "thread.delete": { yield* requireThread({ readModel, diff --git a/apps/server/src/orchestration/projector.ts b/apps/server/src/orchestration/projector.ts index fc6ab8f6fcf..5096adad583 100644 --- a/apps/server/src/orchestration/projector.ts +++ b/apps/server/src/orchestration/projector.ts @@ -17,6 +17,7 @@ import { ThreadActivityAppendedPayload, ThreadArchivedPayload, ThreadCreatedPayload, + ThreadForkedPayload, ThreadDeletedPayload, ThreadInteractionModeSetPayload, ThreadMetaUpdatedPayload, @@ -304,6 +305,48 @@ export function projectEvent( }; }); + case "thread.forked": + return Effect.gen(function* () { + const payload = yield* decodeForEvent( + ThreadForkedPayload, + event.payload, + event.type, + "payload", + ); + const thread: OrchestrationThread = yield* decodeForEvent( + OrchestrationThread, + { + id: payload.threadId, + projectId: payload.projectId, + title: payload.title, + modelSelection: payload.modelSelection, + runtimeMode: payload.runtimeMode, + interactionMode: payload.interactionMode, + branch: payload.branch, + worktreePath: payload.worktreePath, + forkedFrom: payload.forkedFrom, + latestTurn: null, + createdAt: payload.createdAt, + updatedAt: payload.updatedAt, + archivedAt: null, + deletedAt: null, + messages: payload.inheritedMessages, + activities: [], + checkpoints: [], + session: null, + }, + event.type, + "thread", + ); + const existing = nextBase.threads.find((entry) => entry.id === thread.id); + return { + ...nextBase, + threads: existing + ? nextBase.threads.map((entry) => (entry.id === thread.id ? thread : entry)) + : [...nextBase.threads, thread], + }; + }); + case "thread.deleted": return decodeForEvent(ThreadDeletedPayload, event.payload, event.type, "payload").pipe( Effect.map((payload) => ({ diff --git a/apps/server/src/persistence/Layers/ProjectionThreads.ts b/apps/server/src/persistence/Layers/ProjectionThreads.ts index 1baeb375c15..eb505143be2 100644 --- a/apps/server/src/persistence/Layers/ProjectionThreads.ts +++ b/apps/server/src/persistence/Layers/ProjectionThreads.ts @@ -39,6 +39,8 @@ const makeProjectionThreadRepository = Effect.gen(function* () { interaction_mode, branch, worktree_path, + forked_from_thread_id, + forked_from_turn_id, latest_turn_id, created_at, updated_at, @@ -58,6 +60,8 @@ const makeProjectionThreadRepository = Effect.gen(function* () { ${row.interactionMode}, ${row.branch}, ${row.worktreePath}, + ${row.forkedFromThreadId ?? null}, + ${row.forkedFromTurnId ?? null}, ${row.latestTurnId}, ${row.createdAt}, ${row.updatedAt}, @@ -77,6 +81,8 @@ const makeProjectionThreadRepository = Effect.gen(function* () { interaction_mode = excluded.interaction_mode, branch = excluded.branch, worktree_path = excluded.worktree_path, + forked_from_thread_id = excluded.forked_from_thread_id, + forked_from_turn_id = excluded.forked_from_turn_id, latest_turn_id = excluded.latest_turn_id, created_at = excluded.created_at, updated_at = excluded.updated_at, @@ -103,6 +109,8 @@ const makeProjectionThreadRepository = Effect.gen(function* () { interaction_mode AS "interactionMode", branch, worktree_path AS "worktreePath", + forked_from_thread_id AS "forkedFromThreadId", + forked_from_turn_id AS "forkedFromTurnId", latest_turn_id AS "latestTurnId", created_at AS "createdAt", updated_at AS "updatedAt", @@ -131,6 +139,8 @@ const makeProjectionThreadRepository = Effect.gen(function* () { interaction_mode AS "interactionMode", branch, worktree_path AS "worktreePath", + forked_from_thread_id AS "forkedFromThreadId", + forked_from_turn_id AS "forkedFromTurnId", latest_turn_id AS "latestTurnId", created_at AS "createdAt", updated_at AS "updatedAt", diff --git a/apps/server/src/persistence/Migrations.ts b/apps/server/src/persistence/Migrations.ts index d468838d5d4..b6e2b27becd 100644 --- a/apps/server/src/persistence/Migrations.ts +++ b/apps/server/src/persistence/Migrations.ts @@ -45,6 +45,7 @@ import Migration0029 from "./Migrations/029_ProjectionThreadDetailOrderingIndexe import Migration0030 from "./Migrations/030_ProjectionThreadShellArchiveIndexes.ts"; import Migration0031 from "./Migrations/031_AuthAuthorizationScopes.ts"; import Migration0032 from "./Migrations/032_AuthPairingProofKeyThumbprint.ts"; +import Migration0033 from "./Migrations/033_ProjectionThreadForkLineage.ts"; /** * Migration loader with all migrations defined inline. @@ -89,6 +90,7 @@ export const migrationEntries = [ [30, "ProjectionThreadShellArchiveIndexes", Migration0030], [31, "AuthAuthorizationScopes", Migration0031], [32, "AuthPairingProofKeyThumbprint", Migration0032], + [33, "ProjectionThreadForkLineage", Migration0033], ] as const; export const makeMigrationLoader = (throughId?: number) => diff --git a/apps/server/src/persistence/Migrations/033_ProjectionThreadForkLineage.test.ts b/apps/server/src/persistence/Migrations/033_ProjectionThreadForkLineage.test.ts new file mode 100644 index 00000000000..e0bbd6e0e1b --- /dev/null +++ b/apps/server/src/persistence/Migrations/033_ProjectionThreadForkLineage.test.ts @@ -0,0 +1,29 @@ +import { assert, it } from "@effect/vitest"; +import * as Effect from "effect/Effect"; +import * as Layer from "effect/Layer"; +import * as SqlClient from "effect/unstable/sql/SqlClient"; + +import { runMigrations } from "../Migrations.ts"; +import * as NodeSqliteClient from "../NodeSqliteClient.ts"; + +const layer = it.layer(Layer.mergeAll(NodeSqliteClient.layerMemory())); + +layer("033_ProjectionThreadForkLineage", (it) => { + it.effect("adds nullable fork lineage columns and the parent lookup index", () => + Effect.gen(function* () { + const sql = yield* SqlClient.SqlClient; + + yield* runMigrations({ toMigrationInclusive: 32 }); + const before = yield* sql<{ readonly name: string }>`PRAGMA table_info(projection_threads)`; + assert.isFalse(before.some((column) => column.name === "forked_from_thread_id")); + + yield* runMigrations({ toMigrationInclusive: 33 }); + const columns = yield* sql<{ readonly name: string }>`PRAGMA table_info(projection_threads)`; + const indexes = yield* sql<{ readonly name: string }>`PRAGMA index_list(projection_threads)`; + + assert.isTrue(columns.some((column) => column.name === "forked_from_thread_id")); + assert.isTrue(columns.some((column) => column.name === "forked_from_turn_id")); + assert.isTrue(indexes.some((index) => index.name === "idx_projection_threads_forked_from")); + }), + ); +}); diff --git a/apps/server/src/persistence/Migrations/033_ProjectionThreadForkLineage.ts b/apps/server/src/persistence/Migrations/033_ProjectionThreadForkLineage.ts new file mode 100644 index 00000000000..e75a95a6bd4 --- /dev/null +++ b/apps/server/src/persistence/Migrations/033_ProjectionThreadForkLineage.ts @@ -0,0 +1,13 @@ +import * as Effect from "effect/Effect"; +import * as SqlClient from "effect/unstable/sql/SqlClient"; + +export default Effect.gen(function* () { + const sql = yield* SqlClient.SqlClient; + + yield* sql`ALTER TABLE projection_threads ADD COLUMN forked_from_thread_id TEXT`; + yield* sql`ALTER TABLE projection_threads ADD COLUMN forked_from_turn_id TEXT`; + yield* sql` + CREATE INDEX IF NOT EXISTS idx_projection_threads_forked_from + ON projection_threads(forked_from_thread_id, created_at) + `; +}); diff --git a/apps/server/src/persistence/Services/ProjectionThreads.ts b/apps/server/src/persistence/Services/ProjectionThreads.ts index 44fdc147a4a..04c99ea957d 100644 --- a/apps/server/src/persistence/Services/ProjectionThreads.ts +++ b/apps/server/src/persistence/Services/ProjectionThreads.ts @@ -32,6 +32,8 @@ export const ProjectionThread = Schema.Struct({ interactionMode: ProviderInteractionMode, branch: Schema.NullOr(Schema.String), worktreePath: Schema.NullOr(Schema.String), + forkedFromThreadId: Schema.optional(Schema.NullOr(ThreadId)), + forkedFromTurnId: Schema.optional(Schema.NullOr(TurnId)), latestTurnId: Schema.NullOr(TurnId), createdAt: IsoDateTime, updatedAt: IsoDateTime, diff --git a/apps/server/src/provider/Layers/ClaudeAdapter.test.ts b/apps/server/src/provider/Layers/ClaudeAdapter.test.ts index 17aeff2d0e3..d621bf4dabe 100644 --- a/apps/server/src/provider/Layers/ClaudeAdapter.test.ts +++ b/apps/server/src/provider/Layers/ClaudeAdapter.test.ts @@ -9,6 +9,7 @@ import type { PermissionMode, PermissionResult, SDKMessage, + SessionMessage, SDKUserMessage, } from "@anthropic-ai/claude-agent-sdk"; import { @@ -37,7 +38,11 @@ import { ServerConfig } from "../../config.ts"; import { ServerSettingsService } from "../../serverSettings.ts"; import { ProviderAdapterProcessError, ProviderAdapterValidationError } from "../Errors.ts"; import type { ClaudeAdapterShape } from "../Services/ClaudeAdapter.ts"; -import { makeClaudeAdapter, type ClaudeAdapterLiveOptions } from "./ClaudeAdapter.ts"; +import { + makeClaudeAdapter, + resolveClaudeAssistantForkPoint, + type ClaudeAdapterLiveOptions, +} from "./ClaudeAdapter.ts"; const decodeClaudeSettings = Schema.decodeSync(ClaudeSettings); // Test-local service tag so the rest of the file can keep using `yield* ClaudeAdapter`. @@ -156,6 +161,7 @@ function makeHarness(config?: { readonly baseDir?: string; readonly claudeConfig?: Partial; readonly instanceId?: ProviderInstanceId; + readonly getSessionMessages?: (sessionId: string) => Promise; }) { const query = new FakeClaudeQuery(); let createInput: @@ -171,6 +177,7 @@ function makeHarness(config?: { createInput = input; return query; }, + ...(config?.getSessionMessages ? { getSessionMessages: config.getSessionMessages } : {}), ...(config?.nativeEventLogger ? { nativeEventLogger: config.nativeEventLogger, @@ -221,6 +228,30 @@ function makeDeterministicRandomService(seed = 0x1234_5678): { }; } +describe("resolveClaudeAssistantForkPoint", () => { + it("selects the terminal assistant message for a human turn", () => { + const message = (type: SessionMessage["type"], uuid: string, content: unknown) => ({ + type, + uuid, + session_id: "session-1", + message: { content }, + parent_tool_use_id: null, + }); + const messages = [ + message("user", "user-1", [{ type: "text", text: "first" }]), + message("assistant", "assistant-1a", [{ type: "tool_use" }]), + message("user", "tool-result-1", [{ type: "tool_result" }]), + message("assistant", "assistant-1b", [{ type: "text" }]), + message("user", "user-2", [{ type: "text", text: "second" }]), + message("assistant", "assistant-2", [{ type: "text" }]), + ] satisfies SessionMessage[]; + + assert.equal(resolveClaudeAssistantForkPoint(messages, 0), "assistant-1b"); + assert.equal(resolveClaudeAssistantForkPoint(messages, 1), "assistant-2"); + assert.equal(resolveClaudeAssistantForkPoint(messages, undefined), "assistant-2"); + }); +}); + async function readFirstPromptText( input: | { @@ -2821,6 +2852,79 @@ describe("ClaudeAdapterLive", () => { ); }); + it.effect("forks a Claude session at the selected assistant response", () => { + const sourceSessionId = "550e8400-e29b-41d4-a716-446655440000"; + const readCalls: string[] = []; + const harness = makeHarness({ + getSessionMessages: async (sessionId) => { + readCalls.push(sessionId); + return [ + { + type: "user", + uuid: "user-1", + session_id: sessionId, + message: { content: [{ type: "text", text: "first" }] }, + parent_tool_use_id: null, + }, + { + type: "assistant", + uuid: "assistant-1", + session_id: sessionId, + message: {}, + parent_tool_use_id: null, + }, + { + type: "user", + uuid: "user-2", + session_id: sessionId, + message: { content: [{ type: "text", text: "second" }] }, + parent_tool_use_id: null, + }, + { + type: "assistant", + uuid: "assistant-2", + session_id: sessionId, + message: { content: [{ type: "text", text: "response" }] }, + parent_tool_use_id: null, + }, + ]; + }, + }); + return Effect.gen(function* () { + const adapter = yield* ClaudeAdapter; + + const session = yield* adapter.startSession({ + threadId: ThreadId.make("thread-claude-fork"), + provider: ProviderDriverKind.make("claudeAgent"), + forkFrom: { + threadId: RESUME_THREAD_ID, + sourceTurnIndex: 1, + resumeCursor: { + threadId: RESUME_THREAD_ID, + resume: sourceSessionId, + turnCount: 2, + }, + }, + runtimeMode: "full-access", + }); + + const createInput = harness.getLastCreateQueryInput(); + assert.deepEqual(readCalls, [sourceSessionId]); + assert.equal(createInput?.options.resume, sourceSessionId); + assert.equal(createInput?.options.forkSession, true); + assert.equal(createInput?.options.resumeSessionAt, "assistant-2"); + assert.equal(typeof createInput?.options.sessionId, "string"); + assert.notEqual(createInput?.options.sessionId, sourceSessionId); + assert.equal( + (session.resumeCursor as { resume?: string }).resume, + createInput?.options.sessionId, + ); + }).pipe( + Effect.provideService(Random.Random, makeDeterministicRandomService()), + Effect.provide(harness.layer), + ); + }); + it.effect("preserves durable resume ids across Claude resume hooks", () => { const harness = makeHarness(); return Effect.gen(function* () { diff --git a/apps/server/src/provider/Layers/ClaudeAdapter.ts b/apps/server/src/provider/Layers/ClaudeAdapter.ts index f6e63eeffad..cc3220a1ae2 100644 --- a/apps/server/src/provider/Layers/ClaudeAdapter.ts +++ b/apps/server/src/provider/Layers/ClaudeAdapter.ts @@ -14,6 +14,8 @@ import { type PermissionResult, type PermissionUpdate, type SDKMessage, + getSessionMessages, + type SessionMessage, type SDKControlGetContextUsageResponse, type SDKResultMessage, type SettingSource, @@ -219,6 +221,10 @@ export interface ClaudeAdapterLiveOptions { readonly prompt: AsyncIterable; readonly options: ClaudeQueryOptions; }) => ClaudeQueryRuntime; + readonly getSessionMessages?: ( + sessionId: string, + options?: { readonly dir?: string }, + ) => Promise; readonly nativeEventLogPath?: string; readonly nativeEventLogger?: EventNdjsonLogger; } @@ -598,6 +604,52 @@ function readClaudeResumeState(resumeCursor: unknown): ClaudeResumeState | undef }; } +function sessionUserMessageStartsHumanTurn(message: SessionMessage): boolean { + if (message.type !== "user") { + return false; + } + const content = (message.message as { readonly content?: unknown } | null)?.content; + if (!Array.isArray(content)) { + return true; + } + return content.some( + (block) => + typeof block !== "object" || + block === null || + (block as { readonly type?: unknown }).type !== "tool_result", + ); +} + +export function resolveClaudeAssistantForkPoint( + messages: ReadonlyArray, + sourceTurnIndex: number | undefined, +): string | undefined { + const terminalAssistantUuids: string[] = []; + let currentTurnAssistantUuid: string | undefined; + let hasHumanTurn = false; + + for (const message of messages) { + if (sessionUserMessageStartsHumanTurn(message)) { + if (hasHumanTurn && currentTurnAssistantUuid !== undefined) { + terminalAssistantUuids.push(currentTurnAssistantUuid); + } + hasHumanTurn = true; + currentTurnAssistantUuid = undefined; + continue; + } + if (message.type === "assistant" && hasHumanTurn) { + currentTurnAssistantUuid = message.uuid; + } + } + if (hasHumanTurn && currentTurnAssistantUuid !== undefined) { + terminalAssistantUuids.push(currentTurnAssistantUuid); + } + + return sourceTurnIndex === undefined + ? terminalAssistantUuids.at(-1) + : terminalAssistantUuids[sourceTurnIndex]; +} + function classifyToolItemType(toolName: string): CanonicalItemType { const normalized = toolName.toLowerCase(); if (normalized.includes("agent")) { @@ -1370,6 +1422,7 @@ export const makeClaudeAdapter = Effect.fn("makeClaudeAdapter")(function* ( prompt: input.prompt, options: input.options, }) as ClaudeQueryRuntime); + const readSessionMessages = options?.getSessionMessages ?? getSessionMessages; const sessions = new Map(); const runtimeEventQueue = yield* Queue.unbounded(); @@ -3093,11 +3146,50 @@ export const makeClaudeAdapter = Effect.fn("makeClaudeAdapter")(function* ( } const startedAt = yield* nowIso; - const resumeState = readClaudeResumeState(input.resumeCursor); + const isFork = input.forkFrom !== undefined; + const resumeState = readClaudeResumeState( + isFork ? input.forkFrom?.resumeCursor : input.resumeCursor, + ); + if (isFork && resumeState?.resume === undefined) { + return yield* new ProviderAdapterValidationError({ + provider: PROVIDER, + operation: "startSession", + issue: "Claude forks require a valid persisted source session id.", + }); + } + const forkAtAssistantUuid = isFork + ? yield* Effect.tryPromise({ + try: async () => { + const messages = await readSessionMessages(resumeState!.resume!, { + ...(input.cwd ? { dir: input.cwd } : {}), + }); + const sourceTurnIndex = input.forkFrom?.sourceTurnIndex; + const assistantUuid = resolveClaudeAssistantForkPoint(messages, sourceTurnIndex); + if (!assistantUuid) { + throw new Error( + sourceTurnIndex === undefined + ? `Claude source session '${resumeState!.resume}' has no assistant response to fork.` + : `Claude source session '${resumeState!.resume}' has no assistant response at position ${sourceTurnIndex}.`, + ); + } + return assistantUuid; + }, + catch: (cause) => + new ProviderAdapterRequestError({ + provider: PROVIDER, + method: "session/fork", + detail: + cause instanceof Error ? cause.message : "Failed to resolve Claude fork point.", + cause, + }), + }) + : undefined; const threadId = input.threadId; const existingResumeSessionId = resumeState?.resume; - const newSessionId = existingResumeSessionId === undefined ? yield* randomUUIDv4 : undefined; + const newSessionId = + isFork || existingResumeSessionId === undefined ? yield* randomUUIDv4 : undefined; const sessionId = existingResumeSessionId ?? newSessionId; + const effectiveSessionId = isFork ? newSessionId : sessionId; const runtimeContext = yield* Effect.context(); const runFork = Effect.runForkWith(runtimeContext); @@ -3465,6 +3557,8 @@ export const makeClaudeAdapter = Effect.fn("makeClaudeAdapter")(function* ( ...(Object.keys(settings).length > 0 ? { settings } : {}), ...(existingResumeSessionId ? { resume: existingResumeSessionId } : {}), ...(newSessionId ? { sessionId: newSessionId } : {}), + ...(isFork ? { forkSession: true } : {}), + ...(forkAtAssistantUuid ? { resumeSessionAt: forkAtAssistantUuid } : {}), includePartialMessages: true, canUseTool, env: claudeEnvironment, @@ -3536,8 +3630,10 @@ export const makeClaudeAdapter = Effect.fn("makeClaudeAdapter")(function* ( ...(threadId ? { threadId } : {}), resumeCursor: { ...(threadId ? { threadId } : {}), - ...(sessionId ? { resume: sessionId } : {}), - ...(resumeState?.resumeSessionAt ? { resumeSessionAt: resumeState.resumeSessionAt } : {}), + ...(effectiveSessionId ? { resume: effectiveSessionId } : {}), + ...(!isFork && resumeState?.resumeSessionAt + ? { resumeSessionAt: resumeState.resumeSessionAt } + : {}), turnCount: resumeState?.turnCount ?? 0, }, createdAt: startedAt, @@ -3552,7 +3648,7 @@ export const makeClaudeAdapter = Effect.fn("makeClaudeAdapter")(function* ( startedAt, basePermissionMode: permissionMode, currentApiModelId: apiModelId, - resumeSessionId: sessionId, + resumeSessionId: effectiveSessionId, pendingApprovals, pendingUserInputs, turns: [], @@ -3562,7 +3658,7 @@ export const makeClaudeAdapter = Effect.fn("makeClaudeAdapter")(function* ( lastKnownContextWindow: initialContextWindow, lastKnownTokenUsage: undefined, lastKnownTotalProcessedTokens: undefined, - lastAssistantUuid: resumeState?.resumeSessionAt, + lastAssistantUuid: isFork ? undefined : resumeState?.resumeSessionAt, lastThreadStartedId: undefined, stopped: false, }; @@ -3855,6 +3951,7 @@ export const makeClaudeAdapter = Effect.fn("makeClaudeAdapter")(function* ( provider: PROVIDER, capabilities: { sessionModelSwitch: "in-session", + sessionFork: "native", }, startSession, sendTurn, diff --git a/apps/server/src/provider/Layers/CodexAdapter.ts b/apps/server/src/provider/Layers/CodexAdapter.ts index 38a5887cdc3..13b9da7ae8c 100644 --- a/apps/server/src/provider/Layers/CodexAdapter.ts +++ b/apps/server/src/provider/Layers/CodexAdapter.ts @@ -1406,6 +1406,16 @@ export const makeCodexAdapter = Effect.fn("makeCodexAdapter")(function* ( ...(isCodexResumeCursorSchema(input.resumeCursor) ? { resumeCursor: input.resumeCursor } : {}), + ...(input.forkFrom !== undefined && isCodexResumeCursorSchema(input.forkFrom.resumeCursor) + ? { + forkFrom: { + providerThreadId: input.forkFrom.resumeCursor.threadId, + ...(input.forkFrom.sourceTurnId !== undefined + ? { sourceTurnId: input.forkFrom.sourceTurnId } + : {}), + }, + } + : {}), runtimeMode: input.runtimeMode, ...(input.modelSelection?.instanceId === boundInstanceId ? { model: input.modelSelection.model } @@ -1703,6 +1713,7 @@ export const makeCodexAdapter = Effect.fn("makeCodexAdapter")(function* ( provider: PROVIDER, capabilities: { sessionModelSwitch: "in-session", + sessionFork: "native", }, startSession, sendTurn, diff --git a/apps/server/src/provider/Layers/CodexSessionRuntime.test.ts b/apps/server/src/provider/Layers/CodexSessionRuntime.test.ts index 119fa36303a..262f74f6b29 100644 --- a/apps/server/src/provider/Layers/CodexSessionRuntime.test.ts +++ b/apps/server/src/provider/Layers/CodexSessionRuntime.test.ts @@ -365,12 +365,64 @@ describe("isRecoverableThreadResumeError", () => { }); describe("openCodexThread", () => { + it.effect("forks the native Codex thread at the selected turn", () => + Effect.gen(function* () { + const calls: Array<{ + method: "thread/start" | "thread/resume" | "thread/fork"; + payload: unknown; + }> = []; + const forked = makeThreadOpenResponse("forked-provider-thread"); + const client = { + request: ( + method: M, + payload: CodexRpc.ClientRequestParamsByMethod[M], + ) => { + calls.push({ method, payload }); + return Effect.succeed(forked as CodexRpc.ClientRequestResponsesByMethod[M]); + }, + }; + + const opened = yield* openCodexThread({ + client, + threadId: ThreadId.make("t3-fork-thread"), + runtimeMode: "full-access", + cwd: "/tmp/project", + requestedModel: "gpt-5.4", + serviceTier: undefined, + resumeThreadId: "ignored-target-resume-thread", + forkFrom: { + providerThreadId: "source-provider-thread", + sourceTurnId: "turn-2" as import("@t3tools/contracts").TurnId, + }, + }); + + NodeAssert.equal(opened.thread.id, "forked-provider-thread"); + NodeAssert.deepStrictEqual(calls, [ + { + method: "thread/fork", + payload: { + threadId: "source-provider-thread", + cwd: "/tmp/project", + approvalPolicy: "never", + sandbox: "danger-full-access", + model: "gpt-5.4", + lastTurnId: "turn-2", + threadSource: "t3-code", + }, + }, + ]); + }), + ); + it.effect("falls back to thread/start when resume fails recoverably", () => Effect.gen(function* () { - const calls: Array<{ method: "thread/start" | "thread/resume"; payload: unknown }> = []; + const calls: Array<{ + method: "thread/start" | "thread/resume" | "thread/fork"; + payload: unknown; + }> = []; const started = makeThreadOpenResponse("fresh-thread"); const client = { - request: ( + request: ( method: M, payload: CodexRpc.ClientRequestParamsByMethod[M], ) => { @@ -408,7 +460,7 @@ describe("openCodexThread", () => { it.effect("propagates non-recoverable resume failures", () => Effect.gen(function* () { const client = { - request: ( + request: ( method: M, _payload: CodexRpc.ClientRequestParamsByMethod[M], ) => { diff --git a/apps/server/src/provider/Layers/CodexSessionRuntime.ts b/apps/server/src/provider/Layers/CodexSessionRuntime.ts index 5a81e915e34..6dde1ca8458 100644 --- a/apps/server/src/provider/Layers/CodexSessionRuntime.ts +++ b/apps/server/src/provider/Layers/CodexSessionRuntime.ts @@ -105,6 +105,10 @@ export interface CodexSessionRuntimeOptions { readonly model?: string; readonly serviceTier?: CodexServiceTier | undefined; readonly resumeCursor?: CodexResumeCursor; + readonly forkFrom?: { + readonly providerThreadId: string; + readonly sourceTurnId?: TurnId; + }; readonly appServerArgs?: ReadonlyArray; } @@ -428,9 +432,10 @@ export function isRecoverableThreadResumeError(error: unknown): boolean { type CodexThreadOpenResponse = | CodexRpc.ClientRequestResponsesByMethod["thread/start"] - | CodexRpc.ClientRequestResponsesByMethod["thread/resume"]; + | CodexRpc.ClientRequestResponsesByMethod["thread/resume"] + | CodexRpc.ClientRequestResponsesByMethod["thread/fork"]; -type CodexThreadOpenMethod = "thread/start" | "thread/resume"; +type CodexThreadOpenMethod = "thread/start" | "thread/resume" | "thread/fork"; interface CodexThreadOpenClient { readonly request: ( @@ -447,6 +452,12 @@ export const openCodexThread = (input: { readonly requestedModel: string | undefined; readonly serviceTier: CodexServiceTier | undefined; readonly resumeThreadId: string | undefined; + readonly forkFrom?: + | { + readonly providerThreadId: string; + readonly sourceTurnId?: TurnId; + } + | undefined; }): Effect.Effect => { const resumeThreadId = input.resumeThreadId; const startParams = buildThreadStartParams({ @@ -456,6 +467,22 @@ export const openCodexThread = (input: { serviceTier: input.serviceTier, }); + if (input.forkFrom !== undefined) { + const config = runtimeModeToThreadConfig(input.runtimeMode); + return input.client.request("thread/fork", { + threadId: input.forkFrom.providerThreadId, + cwd: input.cwd, + approvalPolicy: config.approvalPolicy, + sandbox: config.sandbox, + ...(startParams.model ? { model: startParams.model } : {}), + ...(startParams.serviceTier ? { serviceTier: startParams.serviceTier } : {}), + ...(input.forkFrom.sourceTurnId !== undefined + ? { lastTurnId: input.forkFrom.sourceTurnId } + : {}), + threadSource: "t3-code", + }); + } + if (resumeThreadId === undefined) { return input.client.request("thread/start", startParams); } @@ -1212,6 +1239,7 @@ export const makeCodexSessionRuntime = ( requestedModel, serviceTier: options.serviceTier, resumeThreadId: readResumeCursorThreadId(options.resumeCursor), + forkFrom: options.forkFrom, }); const providerThreadId = opened.thread.id; diff --git a/apps/server/src/provider/Layers/OpenCodeAdapter.test.ts b/apps/server/src/provider/Layers/OpenCodeAdapter.test.ts index 1385ccbaabe..9d39254a947 100644 --- a/apps/server/src/provider/Layers/OpenCodeAdapter.test.ts +++ b/apps/server/src/provider/Layers/OpenCodeAdapter.test.ts @@ -73,7 +73,7 @@ const runtimeMock = { transientErrorSessionIds: new Set(), sessionDirectoryById: new Map(), sessionUpdateCalls: [] as Array<{ sessionID: string; permission: unknown }>, - forkCalls: [] as Array<{ sessionID: string; directory?: string }>, + forkCalls: [] as Array<{ sessionID: string; directory?: string; messageID?: string }>, }, reset() { this.state.startCalls.length = 0; @@ -165,10 +165,22 @@ const OpenCodeRuntimeTestDouble: OpenCodeRuntimeShape = { runtimeMock.state.sessionUpdateCalls.push({ sessionID, permission }); return { data: { id: sessionID } }; }, - fork: async ({ sessionID, directory }: { sessionID: string; directory?: string }) => { + fork: async ({ + sessionID, + directory, + messageID, + }: { + sessionID: string; + directory?: string; + messageID?: string; + }) => { // Fork clones history into a new session bound to the directory. const forkedId = `${sessionID}_fork`; - runtimeMock.state.forkCalls.push({ sessionID, ...(directory ? { directory } : {}) }); + runtimeMock.state.forkCalls.push({ + sessionID, + ...(directory ? { directory } : {}), + ...(messageID ? { messageID } : {}), + }); if (directory) { runtimeMock.state.sessionDirectoryById.set(forkedId, directory); } @@ -353,6 +365,46 @@ it.layer(OpenCodeAdapterTestLayer)("OpenCodeAdapterLive", (it) => { }), ); + it.effect("forks the persisted OpenCode session at the selected assistant response", () => + Effect.gen(function* () { + const adapter = yield* OpenCodeAdapter; + const threadId = asThreadId("thread-opencode-fork"); + runtimeMock.state.messages = [ + { info: { id: "user-1", role: "user" }, parts: [] }, + { info: { id: "assistant-1", role: "assistant" }, parts: [] }, + { info: { id: "user-2", role: "user" }, parts: [] }, + { info: { id: "assistant-2", role: "assistant" }, parts: [] }, + ]; + + const session = yield* adapter.startSession({ + provider: ProviderDriverKind.make("opencode"), + threadId, + runtimeMode: "full-access", + forkFrom: { + threadId: asThreadId("thread-opencode-source"), + sourceTurnIndex: 1, + resumeCursor: { schemaVersion: 1, sessionId: "ses_source" }, + }, + }); + + NodeAssert.deepEqual(runtimeMock.state.forkCalls, [ + { + sessionID: "ses_source", + directory: process.cwd(), + messageID: "assistant-2", + }, + ]); + NodeAssert.deepEqual(runtimeMock.state.sessionCreateUrls, []); + NodeAssert.deepEqual(session.resumeCursor, { + schemaVersion: 1, + sessionId: "ses_source_fork", + }); + NodeAssert.equal(runtimeMock.state.sessionUpdateCalls[0]?.sessionID, "ses_source_fork"); + + yield* adapter.stopSession(threadId); + }), + ); + it.effect("sends follow-up turns to the resumed session id", () => Effect.gen(function* () { const adapter = yield* OpenCodeAdapter; diff --git a/apps/server/src/provider/Layers/OpenCodeAdapter.ts b/apps/server/src/provider/Layers/OpenCodeAdapter.ts index 73c23b77e68..5c060e04173 100644 --- a/apps/server/src/provider/Layers/OpenCodeAdapter.ts +++ b/apps/server/src/provider/Layers/OpenCodeAdapter.ts @@ -1191,6 +1191,14 @@ export function makeOpenCodeAdapter( const serverPassword = openCodeSettings.serverPassword; const directory = input.cwd ?? serverConfig.cwd; const resumeSessionId = parseOpenCodeResume(input.resumeCursor)?.sessionId; + const forkSessionId = parseOpenCodeResume(input.forkFrom?.resumeCursor)?.sessionId; + if (input.forkFrom !== undefined && forkSessionId === undefined) { + return yield* new ProviderAdapterValidationError({ + provider: PROVIDER, + operation: "startSession", + issue: "OpenCode forks require a valid persisted source session id.", + }); + } const existing = sessions.get(input.threadId); if (existing) { yield* stopOpenCodeContext(existing); @@ -1235,6 +1243,51 @@ export function makeOpenCodeAdapter( // a confirmed not-found (start fresh); transport/auth/server // errors propagate instead of masking as a new empty session. const resolved = yield* Effect.gen(function* () { + if (forkSessionId !== undefined) { + const sourceMessages = yield* runOpenCodeSdk("session.messages", () => + client.session.messages({ sessionID: forkSessionId }), + ); + const assistantMessages = (sourceMessages.data ?? []).filter( + (entry) => entry.info.role === "assistant", + ); + const sourceTurnIndex = input.forkFrom?.sourceTurnIndex; + const sourceMessage = + sourceTurnIndex === undefined + ? assistantMessages.at(-1) + : assistantMessages[sourceTurnIndex]; + if (!sourceMessage) { + return yield* new OpenCodeRuntimeError({ + operation: "session.fork", + detail: + sourceTurnIndex === undefined + ? `OpenCode source session '${forkSessionId}' has no assistant response to fork.` + : `OpenCode source session '${forkSessionId}' has no assistant response at position ${sourceTurnIndex}.`, + }); + } + + const forkedResponse = yield* runOpenCodeSdk("session.fork", () => + client.session.fork({ + sessionID: forkSessionId, + directory, + messageID: sourceMessage.info.id, + }), + ); + const forked = forkedResponse.data; + if (!forked) { + return yield* new OpenCodeRuntimeError({ + operation: "session.fork", + detail: "OpenCode session.fork returned no session payload.", + }); + } + yield* runOpenCodeSdk("session.update", () => + client.session.update({ + sessionID: forked.id, + permission: buildOpenCodePermissionRules(input.runtimeMode), + }), + ); + return { openCodeSession: forked, created: true }; + } + const adopted = resumeSessionId ? yield* runOpenCodeSdk("session.get", () => client.session.get({ sessionID: resumeSessionId }), @@ -1701,6 +1754,7 @@ export function makeOpenCodeAdapter( provider: PROVIDER, capabilities: { sessionModelSwitch: "in-session", + sessionFork: "native", }, startSession, sendTurn, diff --git a/apps/server/src/provider/Layers/ProviderService.ts b/apps/server/src/provider/Layers/ProviderService.ts index 2eaaeb8ce3c..c9e0c9d0a87 100644 --- a/apps/server/src/provider/Layers/ProviderService.ts +++ b/apps/server/src/provider/Layers/ProviderService.ts @@ -559,12 +559,64 @@ const makeProviderService = Effect.fn("makeProviderService")(function* ( `Provider instance '${resolvedInstanceId}' is disabled in T3 Code settings.`, ); } + const adapter = yield* registry.getByInstance(resolvedInstanceId); const persistedBinding = Option.getOrUndefined(yield* directory.getBinding(threadId)); + const requestedForkFrom = input.forkFrom; + const effectiveForkFrom = + requestedForkFrom === undefined + ? undefined + : yield* Effect.gen(function* () { + const sourceBinding = Option.getOrUndefined( + yield* directory.getBinding(requestedForkFrom.threadId), + ); + if (!sourceBinding) { + if (requestedForkFrom.sourceTurnId === undefined) { + return undefined; + } + return yield* toValidationError( + "ProviderService.startSession", + `Cannot fork thread '${requestedForkFrom.threadId}' because it has no persisted provider binding.`, + ); + } + if (adapter.capabilities.sessionFork !== "native") { + return yield* toValidationError( + "ProviderService.startSession", + `Provider '${resolvedProvider}' does not support native thread forks.`, + ); + } + if (sourceBinding.providerInstanceId !== resolvedInstanceId) { + return yield* toValidationError( + "ProviderService.startSession", + `Cannot fork thread '${requestedForkFrom.threadId}' into provider instance '${resolvedInstanceId}' because its provider resume state belongs to a different instance.`, + ); + } + if ( + sourceBinding.resumeCursor === null || + sourceBinding.resumeCursor === undefined + ) { + return yield* toValidationError( + "ProviderService.startSession", + `Cannot fork thread '${requestedForkFrom.threadId}' because no provider resume state is persisted.`, + ); + } + return { + threadId: requestedForkFrom.threadId, + ...(requestedForkFrom.sourceTurnId !== undefined + ? { sourceTurnId: requestedForkFrom.sourceTurnId } + : {}), + ...(requestedForkFrom.sourceTurnIndex !== undefined + ? { sourceTurnIndex: requestedForkFrom.sourceTurnIndex } + : {}), + resumeCursor: sourceBinding.resumeCursor, + }; + }); const effectiveResumeCursor = - input.resumeCursor ?? - (persistedBinding?.providerInstanceId === resolvedInstanceId - ? persistedBinding.resumeCursor - : undefined); + effectiveForkFrom === undefined + ? (input.resumeCursor ?? + (persistedBinding?.providerInstanceId === resolvedInstanceId + ? persistedBinding.resumeCursor + : undefined)) + : undefined; const effectiveCwd = input.cwd ?? (persistedBinding?.providerInstanceId === resolvedInstanceId @@ -580,6 +632,7 @@ const makeProviderService = Effect.fn("makeProviderService")(function* ( ? "persisted" : "none", "provider.resume_cursor.present": effectiveResumeCursor !== undefined, + "provider.fork.present": effectiveForkFrom !== undefined, "provider.cwd.source": input.cwd !== undefined ? "request" @@ -589,7 +642,6 @@ const makeProviderService = Effect.fn("makeProviderService")(function* ( : "none", "provider.cwd.effective": effectiveCwd ?? "", }); - const adapter = yield* registry.getByInstance(resolvedInstanceId); yield* prepareMcpSession(threadId, resolvedInstanceId); const session = yield* adapter .startSession({ @@ -597,6 +649,7 @@ const makeProviderService = Effect.fn("makeProviderService")(function* ( providerInstanceId: resolvedInstanceId, ...(effectiveCwd !== undefined ? { cwd: effectiveCwd } : {}), ...(effectiveResumeCursor !== undefined ? { resumeCursor: effectiveResumeCursor } : {}), + ...(effectiveForkFrom !== undefined ? { forkFrom: effectiveForkFrom } : {}), }) .pipe(Effect.onError(() => clearMcpSession(threadId))); diff --git a/apps/server/src/provider/Services/ProviderAdapter.ts b/apps/server/src/provider/Services/ProviderAdapter.ts index 01eeae7b7bd..941e40af3e7 100644 --- a/apps/server/src/provider/Services/ProviderAdapter.ts +++ b/apps/server/src/provider/Services/ProviderAdapter.ts @@ -24,12 +24,14 @@ import type * as Effect from "effect/Effect"; import type * as Stream from "effect/Stream"; export type ProviderSessionModelSwitchMode = "in-session" | "unsupported"; +export type ProviderSessionForkMode = "native" | "unsupported"; export interface ProviderAdapterCapabilities { /** * Declares whether changing the model on an existing session is supported. */ readonly sessionModelSwitch: ProviderSessionModelSwitchMode; + readonly sessionFork?: ProviderSessionForkMode; } export interface ProviderThreadTurnSnapshot { diff --git a/apps/web/src/components/ChatView.logic.test.ts b/apps/web/src/components/ChatView.logic.test.ts index bc7487cee29..5fad096c07b 100644 --- a/apps/web/src/components/ChatView.logic.test.ts +++ b/apps/web/src/components/ChatView.logic.test.ts @@ -3,6 +3,7 @@ import { MessageId, ProjectId, ProviderInstanceId, + ProviderDriverKind, ThreadId, TurnId, } from "@t3tools/contracts"; @@ -22,8 +23,20 @@ import { reconcileRetainedMountedThreadIds, resolveSendEnvMode, shouldWriteThreadErrorToCurrentServerThread, + supportsSelectedResponseFork, } from "./ChatView.logic"; +describe("supportsSelectedResponseFork", () => { + it("only enables providers with an exact historical fork primitive", () => { + expect(supportsSelectedResponseFork(ProviderDriverKind.make("codex"))).toBe(true); + expect(supportsSelectedResponseFork(ProviderDriverKind.make("claudeAgent"))).toBe(true); + expect(supportsSelectedResponseFork(ProviderDriverKind.make("opencode"))).toBe(true); + expect(supportsSelectedResponseFork(ProviderDriverKind.make("cursor"))).toBe(false); + expect(supportsSelectedResponseFork(ProviderDriverKind.make("grok"))).toBe(false); + expect(supportsSelectedResponseFork(null)).toBe(false); + }); +}); + const environmentId = EnvironmentId.make("environment-local"); const projectId = ProjectId.make("project-1"); const threadId = ThreadId.make("thread-1"); diff --git a/apps/web/src/components/ChatView.logic.ts b/apps/web/src/components/ChatView.logic.ts index 325f9afa90a..5f17915656f 100644 --- a/apps/web/src/components/ChatView.logic.ts +++ b/apps/web/src/components/ChatView.logic.ts @@ -27,6 +27,20 @@ export const MAX_HIDDEN_MOUNTED_PREVIEW_THREADS = 3; export const LastInvokedScriptByProjectSchema = Schema.Record(ProjectId, Schema.String); +const SELECTED_RESPONSE_FORK_DRIVERS = new Set([ + "codex" as ProviderDriverKind, + "claudeAgent" as ProviderDriverKind, + "opencode" as ProviderDriverKind, +]); + +export function supportsSelectedResponseFork( + driverKind: ProviderDriverKind | null | undefined, +): boolean { + return driverKind !== null && driverKind !== undefined + ? SELECTED_RESPONSE_FORK_DRIVERS.has(driverKind) + : false; +} + export function buildLocalDraftThread( threadId: ThreadId, draftThread: DraftThreadState, diff --git a/apps/web/src/components/ChatView.tsx b/apps/web/src/components/ChatView.tsx index c296c717066..feb0f71dc43 100644 --- a/apps/web/src/components/ChatView.tsx +++ b/apps/web/src/components/ChatView.tsx @@ -244,6 +244,7 @@ import { resolveSendEnvMode, revokeBlobPreviewUrl, revokeUserMessagePreviewUrls, + supportsSelectedResponseFork, waitForStartedServerThread, } from "./ChatView.logic"; import { useLocalStorage } from "~/hooks/useLocalStorage"; @@ -423,6 +424,7 @@ type ChatViewProps = onDiffPanelOpen?: () => void; reserveTitleBarControlInset?: boolean; forceExpandedMobileComposer?: boolean; + embedded?: boolean; routeKind: "server"; draftId?: never; } @@ -432,6 +434,7 @@ type ChatViewProps = onDiffPanelOpen?: () => void; reserveTitleBarControlInset?: boolean; forceExpandedMobileComposer?: boolean; + embedded?: boolean; routeKind: "draft"; draftId: DraftId; }; @@ -1084,6 +1087,7 @@ function ChatViewContent(props: ChatViewProps) { onDiffPanelOpen, reserveTitleBarControlInset = true, forceExpandedMobileComposer = false, + embedded = false, } = props; const draftId = routeKind === "draft" ? props.draftId : null; const routeThreadRef = useMemo( @@ -1099,6 +1103,7 @@ function ChatViewContent(props: ChatViewProps) { const writeTerminal = useAtomCommand(terminalEnvironment.write, "terminal write"); const closeTerminalMutation = useAtomCommand(terminalEnvironment.close, "terminal close"); const createThread = useAtomCommand(threadEnvironment.create, { reportFailure: false }); + const forkThread = useAtomCommand(threadEnvironment.fork, { reportFailure: false }); const deleteThread = useAtomCommand(threadEnvironment.delete, { reportFailure: false }); const updateThreadMetadata = useAtomCommand(threadEnvironment.updateMetadata, { reportFailure: false, @@ -1208,6 +1213,7 @@ function ChatViewContent(props: ChatViewProps) { >({}); const [isConnecting, _setIsConnecting] = useState(false); const [isRevertingCheckpoint, setIsRevertingCheckpoint] = useState(false); + const [isForkingToSide, setIsForkingToSide] = useState(false); const [maximizedRightPanelThreadKey, setMaximizedRightPanelThreadKey] = useState( null, ); @@ -1453,7 +1459,7 @@ function ChatViewContent(props: ChatViewProps) { [rightPanelState.surfaces], ); const previewPanelOpen = activeRightPanelKind === "preview" && isPreviewSupportedInRuntime(); - const rightPanelOpen = rightPanelState.isOpen; + const rightPanelOpen = !embedded && rightPanelState.isOpen; const canMaximizeRightPanel = rightPanelOpen && !shouldUsePlanSidebarSheet; const rightPanelMaximized = canMaximizeRightPanel && maximizedRightPanelThreadKey === routeThreadKey; @@ -4758,6 +4764,45 @@ function ChatViewContent(props: ChatViewProps) { ], ); + const onForkToSide = useCallback( + async (sourceTurnId: TurnId) => { + if (!activeThread || !activeThreadRef || !isServerThread || isForkingToSide) return; + + const nextThreadId = newThreadId(); + const title = truncate(`${activeThread.title} (fork)`); + setIsForkingToSide(true); + const result = await forkThread({ + environmentId: activeThread.environmentId, + input: { + threadId: nextThreadId, + sourceThreadId: activeThread.id, + sourceTurnId, + title, + createdAt: new Date().toISOString(), + }, + }); + setIsForkingToSide(false); + + if (result._tag === "Failure") { + if (!isAtomCommandInterrupted(result)) { + const error = squashAtomCommandFailure(result); + toastManager.add( + stackedThreadToast({ + type: "error", + title: "Could not fork side chat", + description: + error instanceof Error ? error.message : "The thread could not be forked.", + }), + ); + } + return; + } + + useRightPanelStore.getState().openThread(activeThreadRef, nextThreadId, title); + }, + [activeThread, activeThreadRef, forkThread, isForkingToSide, isServerThread], + ); + const onImplementPlanInNewThread = useCallback(async () => { if ( !activeThread || @@ -5113,7 +5158,15 @@ function ChatViewContent(props: ChatViewProps) { ); const rightPanelContent = activeThreadRef ? ( - activeRightPanelSurface?.kind === "preview" ? ( + activeRightPanelSurface?.kind === "thread" ? ( + + ) : activeRightPanelSurface?.kind === "preview" ? ( {/* Top bar */} -
- {!rightPanelOpen ? panelLayoutControls : null} - -
+ {!embedded ? ( +
+ {!rightPanelOpen ? panelLayoutControls : null} + +
+ ) : null} {/* Error banner */} @@ -5269,6 +5324,11 @@ function ChatViewContent(props: ChatViewProps) { revertTurnCountByUserMessageId={revertTurnCountByUserMessageId} onRevertUserMessage={onRevertUserMessage} isRevertingCheckpoint={isRevertingCheckpoint} + {...(isServerThread && + !embedded && + supportsSelectedResponseFork(activeProviderStatus?.driver) + ? { onForkToSide, isForkingToSide } + : {})} onImageExpand={onExpandTimelineImage} markdownCwd={gitCwd ?? undefined} resolvedTheme={resolvedTheme} @@ -5509,24 +5569,30 @@ function ChatViewContent(props: ChatViewProps) { {/* end horizontal flex container */} - {mountedTerminalThreadRefs.map(({ key: mountedThreadKey, threadRef: mountedThreadRef }) => ( - - ))} + {!embedded + ? mountedTerminalThreadRefs.map( + ({ key: mountedThreadKey, threadRef: mountedThreadRef }) => ( + + ), + ) + : null} {!shouldUsePlanSidebarSheet && rightPanelOpen && activeThreadRef ? ( diff --git a/apps/web/src/components/RightPanelTabs.tsx b/apps/web/src/components/RightPanelTabs.tsx index f5bc9880d74..960fd08fcab 100644 --- a/apps/web/src/components/RightPanelTabs.tsx +++ b/apps/web/src/components/RightPanelTabs.tsx @@ -1,6 +1,15 @@ import type { ContextMenuItem, PreviewSessionSnapshot } from "@t3tools/contracts"; import { getTerminalLabel } from "@t3tools/shared/terminalLabels"; -import { ClipboardList, FileDiff, Files, Globe2, Plus, TerminalSquare, X } from "lucide-react"; +import { + ClipboardList, + FileDiff, + Files, + GitFork, + Globe2, + Plus, + TerminalSquare, + X, +} from "lucide-react"; import { type MouseEvent as ReactMouseEvent, type ReactElement, @@ -205,6 +214,8 @@ function surfaceTitle( ); case "plan": return "Plan"; + case "thread": + return surface.title; case "preview": { const snapshot = surface.resourceId ? sessions[surface.resourceId] : null; if (!snapshot || snapshot.navStatus._tag === "Idle") return "Browser"; @@ -266,6 +277,8 @@ function SurfaceIcon({ return ; case "plan": return ; + case "thread": + return ; } } diff --git a/apps/web/src/components/chat/MessagesTimeline.logic.test.ts b/apps/web/src/components/chat/MessagesTimeline.logic.test.ts index 6d74204bc1c..430e6cd2680 100644 --- a/apps/web/src/components/chat/MessagesTimeline.logic.test.ts +++ b/apps/web/src/components/chat/MessagesTimeline.logic.test.ts @@ -1,3 +1,4 @@ +import { TurnId } from "@t3tools/contracts"; import { describe, expect, it } from "vite-plus/test"; import { computeStableMessagesTimelineRows, @@ -5,6 +6,7 @@ import { deriveMessagesTimelineRows, normalizeCompactToolLabel, resolveAssistantMessageCopyState, + resolveAssistantMessageForkState, } from "./MessagesTimeline.logic"; describe("computeMessageDurationStart", () => { @@ -260,6 +262,35 @@ describe("resolveAssistantMessageCopyState", () => { }); }); +describe("resolveAssistantMessageForkState", () => { + it("keeps the selected completed assistant turn", () => { + const turnId = TurnId.make("turn-selected"); + + expect( + resolveAssistantMessageForkState({ + turnId, + showForkButton: true, + streaming: false, + }), + ).toEqual({ turnId, visible: true }); + }); + + it("hides the action for streaming, intermediate, and unscoped messages", () => { + const turnId = TurnId.make("turn-selected"); + + expect( + resolveAssistantMessageForkState({ turnId, showForkButton: true, streaming: true }).visible, + ).toBe(false); + expect( + resolveAssistantMessageForkState({ turnId, showForkButton: false, streaming: false }).visible, + ).toBe(false); + expect( + resolveAssistantMessageForkState({ turnId: null, showForkButton: true, streaming: false }) + .visible, + ).toBe(false); + }); +}); + describe("deriveMessagesTimelineRows", () => { it("only enables assistant copy for the terminal assistant message in a turn", () => { const rows = deriveMessagesTimelineRows({ diff --git a/apps/web/src/components/chat/MessagesTimeline.logic.ts b/apps/web/src/components/chat/MessagesTimeline.logic.ts index 3227bac2413..7585d23bdcf 100644 --- a/apps/web/src/components/chat/MessagesTimeline.logic.ts +++ b/apps/web/src/components/chat/MessagesTimeline.logic.ts @@ -221,6 +221,21 @@ export function resolveAssistantMessageCopyState({ }; } +export function resolveAssistantMessageForkState({ + turnId, + showForkButton, + streaming, +}: { + turnId: TurnId | null; + showForkButton: boolean; + streaming: boolean; +}) { + return { + turnId, + visible: showForkButton && turnId !== null && !streaming, + }; +} + function deriveTerminalAssistantMessageIds(timelineEntries: ReadonlyArray) { const lastAssistantMessageIdByResponseKey = new Map(); let nullTurnResponseIndex = 0; diff --git a/apps/web/src/components/chat/MessagesTimeline.test.tsx b/apps/web/src/components/chat/MessagesTimeline.test.tsx index b340a248fbe..a45e39bd0fc 100644 --- a/apps/web/src/components/chat/MessagesTimeline.test.tsx +++ b/apps/web/src/components/chat/MessagesTimeline.test.tsx @@ -219,6 +219,64 @@ function buildUserTimelineEntry(text: string) { } describe("MessagesTimeline", () => { + it("renders the fork action beneath a completed assistant response", async () => { + const { MessagesTimeline } = await import("./MessagesTimeline"); + const turnId = TurnId.make("turn-fork-point"); + const markup = renderToStaticMarkup( + {}} + timelineEntries={[ + { + id: "entry-assistant-fork-point", + kind: "message", + createdAt: MESSAGE_CREATED_AT, + message: { + id: MessageId.make("message-assistant-fork-point"), + role: "assistant", + text: "Fork from this response.", + turnId, + createdAt: MESSAGE_CREATED_AT, + updatedAt: MESSAGE_CREATED_AT, + streaming: false, + }, + }, + ]} + />, + ); + + expect(markup).toContain('aria-label="Fork from this message"'); + expect(markup).toContain("lucide-git-fork"); + expect(markup).toContain('aria-label="Copy link"'); + }); + + it("does not render the fork action without a side-chat owner", async () => { + const { MessagesTimeline } = await import("./MessagesTimeline"); + const markup = renderToStaticMarkup( + , + ); + + expect(markup).not.toContain('aria-label="Fork from this message"'); + }); + it("keeps assistant changed-files headers sticky below the thread header", async () => { const { MessagesTimeline } = await import("./MessagesTimeline"); const assistantMessageId = MessageId.make("message-assistant-with-files"); diff --git a/apps/web/src/components/chat/MessagesTimeline.tsx b/apps/web/src/components/chat/MessagesTimeline.tsx index f759aa150be..77ec2b1aaad 100644 --- a/apps/web/src/components/chat/MessagesTimeline.tsx +++ b/apps/web/src/components/chat/MessagesTimeline.tsx @@ -49,6 +49,7 @@ import { CircleAlertIcon, EyeIcon, FileDiffIcon, + GitForkIcon, GlobeIcon, HammerIcon, MessageCircleIcon, @@ -73,6 +74,7 @@ import { deriveMessagesTimelineRows, normalizeCompactToolLabel, resolveAssistantMessageCopyState, + resolveAssistantMessageForkState, resolveTimelineIsAtEnd, resolveTimelineMinimapHasPersistentGutter, resolveTimelineMinimapHeightStyle, @@ -135,6 +137,8 @@ interface TimelineRowSharedState { skills: ReadonlyArray>; activeThreadEnvironmentId: EnvironmentId; onRevertUserMessage: (messageId: MessageId) => void; + onForkToSide: ((turnId: TurnId) => void) | null; + isForkingToSide: boolean; onImageExpand: (preview: ExpandedImagePreview) => void; onOpenTurnDiff: (turnId: TurnId, filePath?: string) => void; onToggleTurnFold: (turnId: TurnId) => void; @@ -171,6 +175,8 @@ interface MessagesTimelineProps { revertTurnCountByUserMessageId: Map; onRevertUserMessage: (messageId: MessageId) => void; isRevertingCheckpoint: boolean; + onForkToSide?: ((turnId: TurnId) => void) | undefined; + isForkingToSide?: boolean; onImageExpand: (preview: ExpandedImagePreview) => void; activeThreadEnvironmentId: EnvironmentId; markdownCwd: string | undefined; @@ -205,6 +211,8 @@ export const MessagesTimeline = memo(function MessagesTimeline({ revertTurnCountByUserMessageId, onRevertUserMessage, isRevertingCheckpoint, + onForkToSide, + isForkingToSide = false, onImageExpand, activeThreadEnvironmentId, markdownCwd, @@ -426,6 +434,8 @@ export const MessagesTimeline = memo(function MessagesTimeline({ skills, activeThreadEnvironmentId, onRevertUserMessage, + onForkToSide: onForkToSide ?? null, + isForkingToSide, onImageExpand, onOpenTurnDiff, onToggleTurnFold, @@ -440,6 +450,8 @@ export const MessagesTimeline = memo(function MessagesTimeline({ skills, activeThreadEnvironmentId, onRevertUserMessage, + onForkToSide, + isForkingToSide, onImageExpand, onOpenTurnDiff, onToggleTurnFold, @@ -1034,6 +1046,7 @@ function AssistantTimelineRow({ row }: { row: Extract + {!row.message.streaming && ( ; } +function AssistantForkButton({ row }: { row: Extract }) { + const ctx = use(TimelineRowCtx); + const forkState = resolveAssistantMessageForkState({ + turnId: row.message.turnId, + showForkButton: row.showAssistantMeta, + streaming: row.assistantCopyStreaming, + }); + + if (ctx.onForkToSide === null || !forkState.visible || forkState.turnId === null) { + return null; + } + + const turnId = forkState.turnId; + + return ( + + ctx.onForkToSide?.(turnId)} + /> + } + > + + + Fork to side chat + + ); +} + function ProposedPlanTimelineRow({ row, }: { diff --git a/apps/web/src/rightPanelStore.test.ts b/apps/web/src/rightPanelStore.test.ts index c7457cfd304..e50ed5926fd 100644 --- a/apps/web/src/rightPanelStore.test.ts +++ b/apps/web/src/rightPanelStore.test.ts @@ -292,6 +292,25 @@ describe("rightPanelStore", () => { }); }); + it("opens a forked thread as a peer surface and refreshes its title", () => { + const forkId = ThreadId.make("thread-fork"); + useRightPanelStore.getState().openThread(refA, forkId, "Initial fork"); + useRightPanelStore.getState().openThread(refA, forkId, "Renamed fork"); + + expect(selectThreadRightPanelState(useRightPanelStore.getState().byThreadKey, refA)).toEqual({ + isOpen: true, + activeSurfaceId: "thread:thread-fork", + surfaces: [ + { + id: "thread:thread-fork", + kind: "thread", + threadId: forkId, + title: "Renamed fork", + }, + ], + }); + }); + it("tracks one surface per terminal session", () => { useRightPanelStore.getState().openTerminal(refA, "term-1"); useRightPanelStore.getState().openTerminal(refA, "term-2"); diff --git a/apps/web/src/rightPanelStore.ts b/apps/web/src/rightPanelStore.ts index 70d163306cc..2e6f7c0c13a 100644 --- a/apps/web/src/rightPanelStore.ts +++ b/apps/web/src/rightPanelStore.ts @@ -8,13 +8,21 @@ * workspace paths, and diff/plan/files remain singleton surfaces. */ import { scopedThreadKey } from "@t3tools/client-runtime/environment"; -import type { ScopedThreadRef } from "@t3tools/contracts"; +import type { ScopedThreadRef, ThreadId } from "@t3tools/contracts"; import { create } from "zustand"; import { createJSONStorage, persist } from "zustand/middleware"; import { resolveStorage } from "./lib/storage"; -export const RIGHT_PANEL_KINDS = ["plan", "diff", "files", "file", "preview", "terminal"] as const; +export const RIGHT_PANEL_KINDS = [ + "plan", + "diff", + "files", + "file", + "preview", + "terminal", + "thread", +] as const; export type RightPanelKind = (typeof RIGHT_PANEL_KINDS)[number]; export type RightPanelSurface = @@ -37,10 +45,11 @@ export type RightPanelSurface = revealLine: number | null; revealRequestId: number; } - | { id: "plan"; kind: "plan" }; + | { id: "plan"; kind: "plan" } + | { id: `thread:${string}`; kind: "thread"; threadId: ThreadId; title: string }; const RIGHT_PANEL_STORAGE_KEY = "t3code:right-panel-state:v2"; -const RIGHT_PANEL_STORAGE_VERSION = 7; +const RIGHT_PANEL_STORAGE_VERSION = 8; export interface ThreadRightPanelState { isOpen: boolean; @@ -50,7 +59,11 @@ export interface ThreadRightPanelState { interface RightPanelStoreState { byThreadKey: Record; - open: (ref: ScopedThreadRef, kind: Exclude) => void; + open: ( + ref: ScopedThreadRef, + kind: Exclude, + ) => void; + openThread: (ref: ScopedThreadRef, threadId: ThreadId, title: string) => void; openBrowser: (ref: ScopedThreadRef, tabId: string | null) => void; openFile: (ref: ScopedThreadRef, relativePath: string, line?: number) => void; openTerminal: (ref: ScopedThreadRef, terminalId: string) => void; @@ -72,7 +85,10 @@ interface RightPanelStoreState { show: (ref: ScopedThreadRef) => void; close: (ref: ScopedThreadRef) => void; toggleVisibility: (ref: ScopedThreadRef) => void; - toggle: (ref: ScopedThreadRef, kind: Exclude) => void; + toggle: ( + ref: ScopedThreadRef, + kind: Exclude, + ) => void; removeThread: (ref: ScopedThreadRef) => void; } @@ -83,7 +99,7 @@ const EMPTY_THREAD_STATE: ThreadRightPanelState = { }; const singletonSurface = ( - kind: Exclude, + kind: Exclude, ): RightPanelSurface => { switch (kind) { case "diff": @@ -120,6 +136,13 @@ const terminalSurface = (terminalId: string): RightPanelSurface => ({ activeTerminalId: terminalId, }); +const threadSurface = (threadId: ThreadId, title: string): RightPanelSurface => ({ + id: `thread:${threadId}`, + kind: "thread", + threadId, + title, +}); + const upsertSurface = ( current: ThreadRightPanelState, surface: RightPanelSurface, @@ -259,6 +282,20 @@ export const useRightPanelStore = create()( return upsertSurface({ ...current, surfaces: withoutPlaceholder }, surface); }), })), + openThread: (ref, threadId, title) => + set((state) => ({ + byThreadKey: updateThread(state.byThreadKey, scopedThreadKey(ref), (current) => { + const surface = threadSurface(threadId, title); + const existing = current.surfaces.some((entry) => entry.id === surface.id); + return { + isOpen: true, + activeSurfaceId: surface.id, + surfaces: existing + ? current.surfaces.map((entry) => (entry.id === surface.id ? surface : entry)) + : [...current.surfaces, surface], + }; + }), + })), openFile: (ref, relativePath, line) => set((state) => ({ byThreadKey: updateThread(state.byThreadKey, scopedThreadKey(ref), (current) => { diff --git a/packages/client-runtime/src/operations/commands.ts b/packages/client-runtime/src/operations/commands.ts index a0c3cbe771f..1eb32753139 100644 --- a/packages/client-runtime/src/operations/commands.ts +++ b/packages/client-runtime/src/operations/commands.ts @@ -32,6 +32,7 @@ export type CreateProjectInput = CommandInput<"project.create">; export type UpdateProjectInput = CommandInput<"project.meta.update">; export type DeleteProjectInput = CommandInput<"project.delete">; export type CreateThreadInput = CommandInput<"thread.create">; +export type ForkThreadInput = CommandInput<"thread.fork">; export type DeleteThreadInput = CommandInput<"thread.delete">; export type ArchiveThreadInput = CommandInput<"thread.archive">; export type UnarchiveThreadInput = CommandInput<"thread.unarchive">; @@ -123,6 +124,18 @@ export const createThread: (input: CreateThreadInput) => CommandEffect = Effect. }); }); +export const forkThread: (input: ForkThreadInput) => CommandEffect = Effect.fn( + "EnvironmentCommands.forkThread", +)(function* (input) { + const metadata = yield* timestampedCommandMetadata(input); + return yield* dispatch({ + ...input, + type: "thread.fork", + commandId: metadata.commandId, + createdAt: metadata.createdAt, + }); +}); + export const deleteThread: (input: DeleteThreadInput) => CommandEffect = Effect.fn( "EnvironmentCommands.deleteThread", )(function* (input) { diff --git a/packages/client-runtime/src/state/threadCommands.ts b/packages/client-runtime/src/state/threadCommands.ts index aab5110e9cf..c73541ff101 100644 --- a/packages/client-runtime/src/state/threadCommands.ts +++ b/packages/client-runtime/src/state/threadCommands.ts @@ -5,6 +5,7 @@ import { createAtomCommandScheduler, createEnvironmentCommand } from "./runtime. import { type ArchiveThreadInput, type CreateThreadInput, + type ForkThreadInput, type DeleteThreadInput, type InterruptThreadTurnInput, type RespondToThreadApprovalInput, @@ -18,6 +19,7 @@ import { type UpdateThreadMetadataInput, archiveThread, createThread, + forkThread, deleteThread, interruptThreadTurn, respondToThreadApproval, @@ -35,6 +37,7 @@ import type { EnvironmentRegistry } from "../connection/registry.ts"; export type { ArchiveThreadInput, CreateThreadInput, + ForkThreadInput, DeleteThreadInput, InterruptThreadTurnInput, RespondToThreadApprovalInput, @@ -64,6 +67,12 @@ export function createThreadEnvironmentAtoms( scheduler, concurrency, }), + fork: createEnvironmentCommand(runtime, { + label: "environment-data:commands:thread:fork", + execute: (input: ForkThreadInput) => forkThread(input), + scheduler, + concurrency, + }), delete: createEnvironmentCommand(runtime, { label: "environment-data:commands:thread:delete", execute: (input: DeleteThreadInput) => deleteThread(input), diff --git a/packages/contracts/src/orchestration.ts b/packages/contracts/src/orchestration.ts index c0c9c62d080..ee9cc7dbc9d 100644 --- a/packages/contracts/src/orchestration.ts +++ b/packages/contracts/src/orchestration.ts @@ -341,6 +341,12 @@ export const OrchestrationLatestTurn = Schema.Struct({ }); export type OrchestrationLatestTurn = typeof OrchestrationLatestTurn.Type; +export const ThreadForkReference = Schema.Struct({ + threadId: ThreadId, + turnId: Schema.NullOr(TurnId), +}); +export type ThreadForkReference = typeof ThreadForkReference.Type; + export const OrchestrationThread = Schema.Struct({ id: ThreadId, projectId: ProjectId, @@ -352,6 +358,7 @@ export const OrchestrationThread = Schema.Struct({ ), branch: Schema.NullOr(TrimmedNonEmptyString), worktreePath: Schema.NullOr(TrimmedNonEmptyString), + forkedFrom: Schema.optional(Schema.NullOr(ThreadForkReference)), latestTurn: Schema.NullOr(OrchestrationLatestTurn), createdAt: IsoDateTime, updatedAt: IsoDateTime, @@ -398,6 +405,7 @@ export const OrchestrationThreadShell = Schema.Struct({ ), branch: Schema.NullOr(TrimmedNonEmptyString), worktreePath: Schema.NullOr(TrimmedNonEmptyString), + forkedFrom: Schema.optional(Schema.NullOr(ThreadForkReference)), latestTurn: Schema.NullOr(OrchestrationLatestTurn), createdAt: IsoDateTime, updatedAt: IsoDateTime, @@ -540,6 +548,16 @@ const ThreadCreateCommand = Schema.Struct({ createdAt: IsoDateTime, }); +const ThreadForkCommand = Schema.Struct({ + type: Schema.Literal("thread.fork"), + commandId: CommandId, + threadId: ThreadId, + sourceThreadId: ThreadId, + sourceTurnId: TurnId, + title: TrimmedNonEmptyString, + createdAt: IsoDateTime, +}); + const ThreadDeleteCommand = Schema.Struct({ type: Schema.Literal("thread.delete"), commandId: CommandId, @@ -697,6 +715,7 @@ const DispatchableClientOrchestrationCommand = Schema.Union([ ProjectMetaUpdateCommand, ProjectDeleteCommand, ThreadCreateCommand, + ThreadForkCommand, ThreadDeleteCommand, ThreadArchiveCommand, ThreadUnarchiveCommand, @@ -718,6 +737,7 @@ export const ClientOrchestrationCommand = Schema.Union([ ProjectMetaUpdateCommand, ProjectDeleteCommand, ThreadCreateCommand, + ThreadForkCommand, ThreadDeleteCommand, ThreadArchiveCommand, ThreadUnarchiveCommand, @@ -820,6 +840,7 @@ export const OrchestrationEventType = Schema.Literals([ "project.meta-updated", "project.deleted", "thread.created", + "thread.forked", "thread.deleted", "thread.archived", "thread.unarchived", @@ -886,6 +907,21 @@ export const ThreadCreatedPayload = Schema.Struct({ updatedAt: IsoDateTime, }); +export const ThreadForkedPayload = Schema.Struct({ + threadId: ThreadId, + projectId: ProjectId, + title: TrimmedNonEmptyString, + modelSelection: ModelSelection, + runtimeMode: RuntimeMode, + interactionMode: ProviderInteractionMode, + branch: Schema.NullOr(TrimmedNonEmptyString), + worktreePath: Schema.NullOr(TrimmedNonEmptyString), + forkedFrom: ThreadForkReference, + inheritedMessages: Schema.Array(OrchestrationMessage), + createdAt: IsoDateTime, + updatedAt: IsoDateTime, +}); + export const ThreadDeletedPayload = Schema.Struct({ threadId: ThreadId, deletedAt: IsoDateTime, @@ -1054,6 +1090,11 @@ export const OrchestrationEvent = Schema.Union([ type: Schema.Literal("thread.created"), payload: ThreadCreatedPayload, }), + Schema.Struct({ + ...EventBaseFields, + type: Schema.Literal("thread.forked"), + payload: ThreadForkedPayload, + }), Schema.Struct({ ...EventBaseFields, type: Schema.Literal("thread.deleted"), diff --git a/packages/contracts/src/provider.ts b/packages/contracts/src/provider.ts index 94fb007a7bc..c53e3ca4d9c 100644 --- a/packages/contracts/src/provider.ts +++ b/packages/contracts/src/provider.ts @@ -1,5 +1,5 @@ import * as Schema from "effect/Schema"; -import { TrimmedNonEmptyString } from "./baseSchemas.ts"; +import { NonNegativeInt, TrimmedNonEmptyString } from "./baseSchemas.ts"; import { ApprovalRequestId, EventId, @@ -50,6 +50,15 @@ export const ProviderSession = Schema.Struct({ }); export type ProviderSession = typeof ProviderSession.Type; +export const ProviderSessionForkSource = Schema.Struct({ + threadId: ThreadId, + sourceTurnId: Schema.optional(TurnId), + /** Zero-based provider response position for adapters whose native fork API uses a message id. */ + sourceTurnIndex: Schema.optional(NonNegativeInt), + resumeCursor: Schema.optional(Schema.Unknown), +}); +export type ProviderSessionForkSource = typeof ProviderSessionForkSource.Type; + export const ProviderSessionStartInput = Schema.Struct({ threadId: ThreadId, provider: Schema.optional(ProviderDriverKind), @@ -58,6 +67,7 @@ export const ProviderSessionStartInput = Schema.Struct({ cwd: Schema.optional(TrimmedNonEmptyString), modelSelection: Schema.optional(ModelSelection), resumeCursor: Schema.optional(Schema.Unknown), + forkFrom: Schema.optional(ProviderSessionForkSource), approvalPolicy: Schema.optional(ProviderApprovalPolicy), sandboxMode: Schema.optional(ProviderSandboxMode), runtimeMode: RuntimeMode,