diff --git a/src/cli/daemon/agent/__tests__/claude.test.ts b/src/cli/daemon/agent/__tests__/claude.test.ts index b0f68328a..c35f67bbb 100644 --- a/src/cli/daemon/agent/__tests__/claude.test.ts +++ b/src/cli/daemon/agent/__tests__/claude.test.ts @@ -72,6 +72,23 @@ describe("ClaudeBackend", () => { return currentMockProc!; } + it("passes --fork-session when resuming a branch", async () => { + const session = backend.execute("hello", { + cwd: "/tmp", + resumeSessionId: "parent_session", + forkSession: true, + }); + const mock = getMock(); + + const spawnCall = (spawn as any).mock.calls[0]; + expect(spawnCall[1]).toContain("--resume"); + expect(spawnCall[1]).toContain("parent_session"); + expect(spawnCall[1]).toContain("--fork-session"); + + mock.proc.emit("close", 0); + await session.result; + }); + it("emits MessageText for assistant text blocks", async () => { const session = backend.execute("hello", { cwd: "/tmp" }); const mock = getMock(); diff --git a/src/cli/daemon/agent/__tests__/codex.test.ts b/src/cli/daemon/agent/__tests__/codex.test.ts index 6dcc68d11..91da229cb 100644 --- a/src/cli/daemon/agent/__tests__/codex.test.ts +++ b/src/cli/daemon/agent/__tests__/codex.test.ts @@ -214,6 +214,40 @@ describe("CodexBackend", () => { await session.result; }); + it("uses thread/fork instead of thread/resume when resuming a branch", async () => { + const session = backend.execute("branch prompt", { + cwd: "/tmp", + model: "gpt-4", + resumeSessionId: "parent_thread", + forkSession: true, + }); + const mock = getMock(); + + await tick(); + sendResponse(1, {}); + await tick(); + + const forkWrite = mock.stdinWrites.find((w) => w.includes('"thread/fork"')); + expect(forkWrite).toBeDefined(); + const parsedFork = JSON.parse(forkWrite!); + expect(parsedFork.params.threadId).toBe("parent_thread"); + expect(parsedFork.params.model).toBe("gpt-4"); + expect(mock.stdinWrites.some((w) => w.includes('"thread/resume"'))).toBe(false); + + sendResponse(2, { thread: { id: "branch_thread" } }); + await tick(); + + const turnWrite = mock.stdinWrites.find((w) => w.includes('"turn/start"')); + expect(turnWrite).toBeDefined(); + const parsedTurn = JSON.parse(turnWrite!); + expect(parsedTurn.params.threadId).toBe("branch_thread"); + expect(parsedTurn.params.input).toEqual([{ type: "text", text: "branch prompt" }]); + + sendResponse(3, {}); + mock.proc.emit("close", 0); + await session.result; + }); + it("extracts session ID from thread/start response", async () => { const session = backend.execute("hello", { cwd: "/tmp" }); const mock = getMock(); diff --git a/src/cli/daemon/agent/claude.ts b/src/cli/daemon/agent/claude.ts index 214e252d9..f391edf9d 100644 --- a/src/cli/daemon/agent/claude.ts +++ b/src/cli/daemon/agent/claude.ts @@ -28,6 +28,9 @@ export class ClaudeBackend implements AgentBackend { } if (options.resumeSessionId) { args.push("--resume", options.resumeSessionId); + if (options.forkSession) { + args.push("--fork-session"); + } } const proc = spawn(this.cliPath, args, { diff --git a/src/cli/daemon/agent/codex.ts b/src/cli/daemon/agent/codex.ts index bc4334712..d9aee5077 100644 --- a/src/cli/daemon/agent/codex.ts +++ b/src/cli/daemon/agent/codex.ts @@ -474,7 +474,13 @@ export class CodexBackend implements AgentBackend { // 3. Start or resume thread let threadResponse: unknown; - if (options.resumeSessionId) { + if (options.resumeSessionId && options.forkSession) { + threadResponse = await sendRpc("thread/fork", { + threadId: options.resumeSessionId, + ...(options.model ? { model: options.model } : {}), + }); + sessionId = extractThreadID(threadResponse); + } else if (options.resumeSessionId) { // thread/resume reopens an existing thread by ID threadResponse = await sendRpc("thread/resume", { threadId: options.resumeSessionId, diff --git a/src/cli/daemon/session-runner.test.ts b/src/cli/daemon/session-runner.test.ts index 727428227..305c0be18 100644 --- a/src/cli/daemon/session-runner.test.ts +++ b/src/cli/daemon/session-runner.test.ts @@ -854,6 +854,101 @@ describe("session-runner runSession", () => { ); }); + it("forks from the pinned root-message parent session for branch tasks", async () => { + setupBackend([], { + status: "completed", + output: "Forked", + error: "", + durationMs: 100, + sessionId: "branch-session", + }); + + await runSession( + makeInput({ + task: { + ...makeInput().task, + conversationId: "branch_c", + contextKey: "branch_c", + context: { + runtime_branch: { + parent_context_key: "parent_c", + parent_task_id: "root_task", + parent_session_id: "root-message-session", + root_message_id: "root_m", + provider: "claude", + }, + }, + }, + }), + ); + + expect(mockFindResumableSessionByContextKey).not.toHaveBeenCalled(); + expect(mockBackendExecute).toHaveBeenCalledWith( + expect.any(String), + expect.objectContaining({ + resumeSessionId: "root-message-session", + forkSession: true, + }), + ); + }); + + it("fails branch tasks before spawning when runtime provider does not match branch provider", async () => { + await runSession( + makeInput({ + provider: "claude", + task: { + ...makeInput().task, + conversationId: "branch_c", + contextKey: "branch_c", + context: { + runtime_branch: { + parent_context_key: "parent_c", + parent_task_id: "root_task", + parent_session_id: "root-session", + root_message_id: "root_m", + provider: "codex", + }, + }, + }, + }), + ); + + expect(mockFindResumableSessionByContextKey).not.toHaveBeenCalled(); + expect(mockBackendExecute).not.toHaveBeenCalled(); + expect(mockClientInstance.failTask).toHaveBeenCalledWith( + "test_token", + "t1", + expect.stringContaining("does not match runtime provider claude"), + ); + }); + + it("fails branch tasks before spawning when no pinned parent session is provided", async () => { + await runSession( + makeInput({ + task: { + ...makeInput().task, + conversationId: "branch_c", + contextKey: "branch_c", + context: { + runtime_branch: { + parent_context_key: "parent_c", + root_message_id: "root_m", + provider: "claude", + }, + }, + }, + }), + ); + + expect(mockFindResumableSessionByContextKey).not.toHaveBeenCalled(); + expect(mockBackendExecute).not.toHaveBeenCalled(); + expect(mockClientInstance.failTask).toHaveBeenCalledWith( + "test_token", + "t1", + expect.stringContaining("pinned parent session is required"), + ); + }); + it("session starts fresh when provider has no matching prior entry", async () => { mockFindResumableSessionByContextKey.mockReturnValueOnce(null); diff --git a/src/cli/daemon/session-runner.ts b/src/cli/daemon/session-runner.ts index cf524bb25..85fcb8f75 100644 --- a/src/cli/daemon/session-runner.ts +++ b/src/cli/daemon/session-runner.ts @@ -318,11 +318,83 @@ export async function runSession(input: SessionRunnerInput): Promise { const prompt = input.promptOverride ?? buildPrompt(task, attachments); - const resumeSessionId = task.contextKey - ? findResumableSessionByContextKey(timelineDir, task.contextKey, provider) ?? undefined - : undefined; + async function failBeforeSpawn(errMsg: string) { + log.error(errMsg); + updateEntry(timelineDir, task.id, (entry) => { + entry.pid = null; + entry.status = "failed"; + entry.errmsg = errMsg; + }); + await reportToServer( + () => client.failTask(token, task.id, errMsg), + { taskId: task.id, type: "fail", payload: { error: errMsg }, token, serverURL, createdAt: new Date().toISOString() }, + workspacesRoot, + ); + process.removeListener("SIGTERM", onKill); + process.removeListener("SIGINT", onKill); + } + + const runtimeBranch = task.context?.runtime_branch as + | { + parent_context_key?: unknown; + parent_task_id?: unknown; + parent_session_id?: unknown; + root_message_id?: unknown; + provider?: unknown; + } + | undefined; + const branchTask = Boolean(runtimeBranch); + const branchParentContextKey = + typeof runtimeBranch?.parent_context_key === "string" + ? runtimeBranch.parent_context_key + : null; + const branchParentTaskId = + typeof runtimeBranch?.parent_task_id === "string" + ? runtimeBranch.parent_task_id + : null; + const branchParentSessionId = + typeof runtimeBranch?.parent_session_id === "string" + ? runtimeBranch.parent_session_id + : null; + const branchRootMessageId = + typeof runtimeBranch?.root_message_id === "string" + ? runtimeBranch.root_message_id + : null; + const branchProvider = + typeof runtimeBranch?.provider === "string" ? runtimeBranch.provider : null; + const forkSession = branchTask; + if (forkSession && !branchParentContextKey) { + await failBeforeSpawn("cannot branch: parent context key is required"); + return; + } + if (forkSession && !branchParentSessionId) { + await failBeforeSpawn( + `cannot branch: pinned parent session is required for root_message_id ${branchRootMessageId ?? "unknown"}`, + ); + return; + } + if (forkSession && !branchProvider) { + await failBeforeSpawn("cannot branch: branch provider is required"); + return; + } + if (forkSession && branchProvider !== provider) { + await failBeforeSpawn( + `cannot branch: branch provider ${branchProvider} does not match runtime provider ${provider}`, + ); + return; + } + const resumeContextKey = forkSession ? branchParentContextKey : task.contextKey ?? null; + const resumeSessionId = forkSession + ? branchParentSessionId! + : resumeContextKey + ? findResumableSessionByContextKey(timelineDir, resumeContextKey, provider) ?? undefined + : undefined; if (resumeSessionId) { - log.info(`resuming session ${resumeSessionId} (context_key: ${task.contextKey})`); + log.info( + `${forkSession ? "forking" : "resuming"} session ${resumeSessionId} (context_key: ${resumeContextKey}${ + branchParentTaskId ? `, parent_task_id: ${branchParentTaskId}` : "" + })`, + ); } const session = backend.execute(prompt, { @@ -331,6 +403,7 @@ export async function runSession(input: SessionRunnerInput): Promise { env, timeout: agentTimeout, resumeSessionId, + forkSession, }); // Capture agent PID so the handler can reap it. backend.execute() spawns diff --git a/src/cli/daemon/types.ts b/src/cli/daemon/types.ts index 04bb157cb..7949ce4d7 100644 --- a/src/cli/daemon/types.ts +++ b/src/cli/daemon/types.ts @@ -88,6 +88,7 @@ export interface ExecOptions { maxTurns?: number; timeout?: number; resumeSessionId?: string; + forkSession?: boolean; } /** Serialized input passed from daemon to the detached session-runner process. */ diff --git a/src/shared/src/api-types.ts b/src/shared/src/api-types.ts index 40840e2ee..29349d376 100644 --- a/src/shared/src/api-types.ts +++ b/src/shared/src/api-types.ts @@ -129,6 +129,7 @@ export interface UpdateAgentRequest { export interface SendMessageRequest { content: string; + metadata?: Record; } export interface CreateMachineTokenRequest { diff --git a/src/shared/src/constants.ts b/src/shared/src/constants.ts index 65933c8f0..bb1d2d8e6 100644 --- a/src/shared/src/constants.ts +++ b/src/shared/src/constants.ts @@ -48,6 +48,17 @@ export const TASK_TYPES = { export type TaskType = (typeof TASK_TYPES)[keyof typeof TASK_TYPES]; +export const CONVERSATION_TYPES = { + USER_DM_MESSAGE: TASK_TYPES.USER_DM_MESSAGE, + EMAIL_NOTIFICATION: TASK_TYPES.EMAIL_NOTIFICATION, + CALENDAR_EVENT: TASK_TYPES.CALENDAR_EVENT, + ISSUE_EVENT: TASK_TYPES.ISSUE_EVENT, + MESSAGE_BRANCH: "message_branch", +} as const; + +export type ConversationType = + (typeof CONVERSATION_TYPES)[keyof typeof CONVERSATION_TYPES]; + export const IssueStatus = { TODO: "todo", IN_PROGRESS: "in_progress", diff --git a/src/shared/src/db/queries-index.ts b/src/shared/src/db/queries-index.ts index 8915995dd..b7b8c1450 100644 --- a/src/shared/src/db/queries-index.ts +++ b/src/shared/src/db/queries-index.ts @@ -4,6 +4,7 @@ export * as member from "./queries/member"; export * as agent from "./queries/agent"; export * as runtime from "./queries/runtime"; export * as conversation from "./queries/conversation"; +export * as conversationBranch from "./queries/conversation-branch"; export * as message from "./queries/message"; export * as task from "./queries/task"; export * as taskMessage from "./queries/task-message"; diff --git a/src/shared/src/db/queries/conversation-branch.ts b/src/shared/src/db/queries/conversation-branch.ts new file mode 100644 index 000000000..c50e6ffe9 --- /dev/null +++ b/src/shared/src/db/queries/conversation-branch.ts @@ -0,0 +1,99 @@ +import { and, asc, eq } from "drizzle-orm"; +import { conversationBranch, message } from "../schema"; +import type { Database } from "../index"; + +export async function createBranch( + db: Database, + data: { + workspaceId: string; + parentConversationId: string; + branchConversationId: string; + rootMessageId: string; + provider: string; + forkSourceTaskId?: string | null; + forkSourceSessionId?: string | null; + createdBy: string; + }, +) { + const rows = await db + .insert(conversationBranch) + .values(data) + .returning(); + return rows[0]!; +} + +export async function listBranchesByParent( + db: Database, + data: { workspaceId: string; parentConversationId: string }, +) { + return db + .select() + .from(conversationBranch) + .where( + and( + eq(conversationBranch.workspaceId, data.workspaceId), + eq(conversationBranch.parentConversationId, data.parentConversationId), + ), + ) + .orderBy(asc(conversationBranch.createdAt)); +} + +export async function getBranchForRoot( + db: Database, + data: { + workspaceId: string; + parentConversationId: string; + rootMessageId: string; + }, +) { + const rows = await db + .select() + .from(conversationBranch) + .where( + and( + eq(conversationBranch.workspaceId, data.workspaceId), + eq(conversationBranch.parentConversationId, data.parentConversationId), + eq(conversationBranch.rootMessageId, data.rootMessageId), + ), + ) + .limit(1); + return rows[0] ?? null; +} + +export async function getBranchByConversation( + db: Database, + data: { workspaceId: string; branchConversationId: string }, +) { + const rows = await db + .select() + .from(conversationBranch) + .where( + and( + eq(conversationBranch.workspaceId, data.workspaceId), + eq(conversationBranch.branchConversationId, data.branchConversationId), + ), + ) + .limit(1); + return rows[0] ?? null; +} + +export async function getBranchOrigin( + db: Database, + data: { workspaceId: string; branchConversationId: string }, +) { + const rows = await db + .select({ + branch: conversationBranch, + rootMessage: message, + }) + .from(conversationBranch) + .innerJoin(message, eq(message.id, conversationBranch.rootMessageId)) + .where( + and( + eq(conversationBranch.workspaceId, data.workspaceId), + eq(conversationBranch.branchConversationId, data.branchConversationId), + ), + ) + .limit(1); + return rows[0] ?? null; +} diff --git a/src/shared/src/db/queries/conversation.ts b/src/shared/src/db/queries/conversation.ts index 21466062f..d4a4c504d 100644 --- a/src/shared/src/db/queries/conversation.ts +++ b/src/shared/src/db/queries/conversation.ts @@ -1,7 +1,11 @@ import { eq, and, desc, ne, lt, sql, count as drizzleCount, inArray } from "drizzle-orm"; import { conversation, message } from "../schema"; import type { Database } from "../index"; -import { TASK_TYPES, type TaskType } from "../../constants"; +import { + CONVERSATION_TYPES, + TASK_TYPES, + type ConversationType, +} from "../../constants"; export async function createConversation( @@ -11,7 +15,7 @@ export async function createConversation( agentId: string; userId: string; title: string; - type?: TaskType; + type?: ConversationType; channel?: string; } ) { @@ -22,7 +26,7 @@ export async function createConversation( agentId: data.agentId, userId: data.userId, title: data.title, - type: data.type ?? TASK_TYPES.USER_DM_MESSAGE, + type: data.type ?? CONVERSATION_TYPES.USER_DM_MESSAGE, channel: data.channel ?? "default", }) .returning(); @@ -54,6 +58,7 @@ export async function listConversations( const conditions = [ eq(conversation.workspaceId, workspaceId), eq(conversation.userId, userId), + eq(conversation.type, CONVERSATION_TYPES.USER_DM_MESSAGE), ]; if (channel) { conditions.push(eq(conversation.channel, channel)); @@ -76,6 +81,7 @@ export async function listConversationsByAgent( eq(conversation.workspaceId, workspaceId), eq(conversation.userId, userId), eq(conversation.agentId, agentId), + eq(conversation.type, CONVERSATION_TYPES.USER_DM_MESSAGE), ]; if (channel) { conditions.push(eq(conversation.channel, channel)); @@ -87,6 +93,7 @@ export async function listConversationsByAgent( agentId: conversation.agentId, userId: conversation.userId, title: conversation.title, + type: conversation.type, channel: conversation.channel, createdAt: conversation.createdAt, messageCount: diff --git a/src/shared/src/db/queries/message.ts b/src/shared/src/db/queries/message.ts index d4d7be173..9d54ff489 100644 --- a/src/shared/src/db/queries/message.ts +++ b/src/shared/src/db/queries/message.ts @@ -1,5 +1,5 @@ -import { eq, asc, desc, and, lt, gte, or, count } from "drizzle-orm"; -import { message } from "../schema"; +import { eq, asc, desc, and, lt, gte, or, count, ne, isNotNull } from "drizzle-orm"; +import { agentTaskQueue, message } from "../schema"; import type { Database } from "../index"; export async function createMessage( @@ -28,6 +28,55 @@ export async function createMessage( } const DEFAULT_MESSAGE_LIMIT = 20; +const NON_BRANCHABLE_MESSAGE_KINDS = new Set([ + "event", + "lifecycle", + "process", + "progress", + "status", + "transient", + "typing", +]); + +type BranchableMessageCandidate = { + role: string; + status?: string | null; + metadata?: string | Record | null; +}; + +function parseMetadata( + metadata: BranchableMessageCandidate["metadata"], +): Record | null { + if (!metadata) return null; + if (typeof metadata === "object") return metadata; + try { + const parsed = JSON.parse(metadata); + return parsed && typeof parsed === "object" && !Array.isArray(parsed) + ? parsed + : null; + } catch { + return null; + } +} + +export function isBranchableMessageRoot( + candidate: BranchableMessageCandidate, +): boolean { + const status = candidate.status as string | null | undefined; + if (status && status !== "active") return false; + if (candidate.role !== "user" && candidate.role !== "assistant") return false; + + const metadata = parseMetadata(candidate.metadata); + const kind = + typeof metadata?.kind === "string" + ? metadata.kind.toLowerCase() + : null; + if (kind && NON_BRANCHABLE_MESSAGE_KINDS.has(kind)) return false; + if (metadata?.transient === true) return false; + if (metadata?.error_source) return false; + if (candidate.role === "assistant") return kind === null || kind === "dm"; + return true; +} export async function getNewestMessageId( db: Database, @@ -110,6 +159,54 @@ export async function getMessage(db: Database, id: string) { return rows[0] ?? null; } +export async function getMessageForConversation( + db: Database, + conversationId: string, + id: string, +) { + const rows = await db + .select() + .from(message) + .where(and(eq(message.conversationId, conversationId), eq(message.id, id))) + .limit(1); + return rows[0] ?? null; +} + +export async function getLatestNonEventMessage(db: Database, conversationId: string) { + const rows = await db + .select() + .from(message) + .where( + and( + eq(message.conversationId, conversationId), + eq(message.status, "active"), + ne(message.role, "event"), + ), + ) + .orderBy(desc(message.createdAt), desc(message.id)) + .limit(1); + return rows[0] ?? null; +} + +export async function getLatestBranchableMessage(db: Database, conversationId: string) { + const rows = await db + .select({ msg: message }) + .from(message) + .innerJoin(agentTaskQueue, eq(message.taskId, agentTaskQueue.id)) + .where( + and( + eq(message.conversationId, conversationId), + eq(message.status, "active"), + ne(message.role, "event"), + eq(agentTaskQueue.status, "completed"), + isNotNull(agentTaskQueue.sessionId), + ), + ) + .orderBy(desc(message.createdAt), desc(message.id)) + .limit(100); + return rows.find((row) => isBranchableMessageRoot(row.msg))?.msg ?? null; +} + export async function updateMessageTaskId(db: Database, messageId: string, taskId: string) { await db.update(message).set({ taskId }).where(eq(message.id, messageId)); } diff --git a/src/shared/src/db/queries/task.ts b/src/shared/src/db/queries/task.ts index 86aa841b1..d282eddad 100644 --- a/src/shared/src/db/queries/task.ts +++ b/src/shared/src/db/queries/task.ts @@ -1,4 +1,4 @@ -import { eq, and, desc, asc, inArray, notInArray, ne, count, lt, or, sql, exists } from "drizzle-orm"; +import { eq, and, desc, asc, inArray, notInArray, ne, count, lt, or, sql, exists, isNotNull } from "drizzle-orm"; import { alias } from "drizzle-orm/sqlite-core"; import { agentTaskQueue, taskMessage, conversation } from "../schema"; import type { Database } from "../index"; @@ -61,6 +61,31 @@ export async function getLatestTaskForConversation(db: Database, conversationId: return rows[0] ?? null; } +export async function getLatestCompletedTaskWithSessionForConversation( + db: Database, + data: { workspaceId: string; conversationId: string }, +) { + const rows = await db + .select({ + id: agentTaskQueue.id, + runtimeId: agentTaskQueue.runtimeId, + sessionId: agentTaskQueue.sessionId, + traceId: agentTaskQueue.traceId, + }) + .from(agentTaskQueue) + .where( + and( + eq(agentTaskQueue.workspaceId, data.workspaceId), + eq(agentTaskQueue.conversationId, data.conversationId), + eq(agentTaskQueue.status, "completed"), + isNotNull(agentTaskQueue.sessionId), + ), + ) + .orderBy(desc(agentTaskQueue.completedAt), desc(agentTaskQueue.createdAt)) + .limit(1); + return rows[0] ?? null; +} + export async function getTask(db: Database, id: string, workspaceId?: string) { const conditions = [eq(agentTaskQueue.id, id)]; if (workspaceId) conditions.push(eq(agentTaskQueue.workspaceId, workspaceId)); diff --git a/src/shared/src/db/schema.ts b/src/shared/src/db/schema.ts index 31313ec87..671ac5fd6 100644 --- a/src/shared/src/db/schema.ts +++ b/src/shared/src/db/schema.ts @@ -9,7 +9,7 @@ import { } from "drizzle-orm/sqlite-core"; import { sql } from "drizzle-orm"; import { nanoid } from "nanoid"; -import { TASK_TYPES } from "../constants"; +import { CONVERSATION_TYPES, TASK_TYPES } from "../constants"; // --------------------------------------------------------------------------- // Better Auth tables @@ -302,7 +302,7 @@ export const conversation = sqliteTable( .notNull() .references(() => user.id, { onDelete: "cascade" }), title: text("title").notNull().default(""), - type: text("type").notNull().default(TASK_TYPES.USER_DM_MESSAGE), + type: text("type").notNull().default(CONVERSATION_TYPES.USER_DM_MESSAGE), channel: text("channel").notNull().default("default"), createdAt: text("created_at").notNull().$defaultFn(() => new Date().toISOString()), }, @@ -337,6 +337,44 @@ export const message = sqliteTable( ] ); +export const conversationBranch = sqliteTable( + "conversation_branch", + { + id: text("id").primaryKey().$defaultFn(() => "br_" + nanoid()), + workspaceId: text("workspace_id") + .notNull() + .references(() => workspace.id, { onDelete: "cascade" }), + parentConversationId: text("parent_conversation_id") + .notNull() + .references(() => conversation.id, { onDelete: "cascade" }), + branchConversationId: text("branch_conversation_id") + .notNull() + .references(() => conversation.id, { onDelete: "cascade" }), + rootMessageId: text("root_message_id") + .notNull() + .references(() => message.id, { onDelete: "cascade" }), + provider: text("provider").notNull(), + forkSourceTaskId: text("fork_source_task_id"), + forkSourceSessionId: text("fork_source_session_id"), + createdBy: text("created_by") + .notNull() + .references(() => user.id, { onDelete: "cascade" }), + createdAt: text("created_at").notNull().$defaultFn(() => new Date().toISOString()), + }, + (t) => [ + unique("conversation_branch_root_unique").on( + t.workspaceId, + t.parentConversationId, + t.rootMessageId, + ), + unique("conversation_branch_conversation_unique").on(t.branchConversationId), + index("idx_conversation_branch_parent").on( + t.workspaceId, + t.parentConversationId, + ), + ], +); + export const agentTaskQueue = sqliteTable( "agent_task_queue", { diff --git a/src/shared/src/index.ts b/src/shared/src/index.ts index 0c181e7c5..26a88da74 100644 --- a/src/shared/src/index.ts +++ b/src/shared/src/index.ts @@ -7,6 +7,7 @@ export type { RuntimeMetadata, Machine, Conversation, + ConversationBranch, Message, TaskMessage, TaskMessageResponse, @@ -75,6 +76,7 @@ export { TERMINAL_TASK_STATUSES, isTerminalTaskStatus, TASK_TYPES, + CONVERSATION_TYPES, IssueStatus, ACTIVE_ISSUE_STATUSES, TERMINAL_ISSUE_STATUSES, @@ -98,6 +100,7 @@ export type { RuntimeStatusType, TaskStatusType, TaskType, + ConversationType, IssueStatusType, MessageRoleType, MeetingStatusType, @@ -145,6 +148,7 @@ export { CreateAgentRequestSchema, UpdateAgentRequestSchema, CreateConversationRequestSchema, + CreateBranchRequestSchema, CreateMessageRequestSchema, AgentDmRequestSchema, EmailAttachmentSchema, @@ -205,6 +209,7 @@ export type { UpdateAgentLinkRequestInput, UpsertAgentLinkRequestInput, AddWhitelistRequest, + CreateBranchRequest, CreateEmailAccountRequest, UpdateMemberRequest, UpdateEmailAccountRequest, diff --git a/src/shared/src/schemas.ts b/src/shared/src/schemas.ts index 63e632961..ebbfd52c7 100644 --- a/src/shared/src/schemas.ts +++ b/src/shared/src/schemas.ts @@ -539,6 +539,11 @@ export type CreateConversationRequest = z.infer< typeof CreateConversationRequestSchema >; +export const CreateBranchRequestSchema = z.object({ + root_message_id: z.string().min(1, "root_message_id is required"), +}); +export type CreateBranchRequest = z.infer; + // --------------------------------------------------------------------------- // Message request schema (JSON body only — FormData path is separate) // --------------------------------------------------------------------------- diff --git a/src/shared/src/types.ts b/src/shared/src/types.ts index 60fce14fd..3da148366 100644 --- a/src/shared/src/types.ts +++ b/src/shared/src/types.ts @@ -71,6 +71,19 @@ export interface Conversation { message_count?: number; } +export interface ConversationBranch { + id: string; + workspace_id: string; + parent_conversation_id: string; + branch_conversation_id: string; + root_message_id: string; + provider: string; + fork_source_task_id: string | null; + fork_source_session_id: string | null; + created_by: string; + created_at: string; +} + export interface Channel { id: string; workspace_id: string; diff --git a/src/shared/test/queries/conversation-branch.test.ts b/src/shared/test/queries/conversation-branch.test.ts new file mode 100644 index 000000000..e1b192815 --- /dev/null +++ b/src/shared/test/queries/conversation-branch.test.ts @@ -0,0 +1,108 @@ +import { describe, it, expect, vi } from "vitest"; +import * as branchQueries from "../../src/db/queries/conversation-branch"; + +function createMockDb(rows: any[]) { + const chain: any = {}; + chain.select = vi.fn(() => chain); + chain.from = vi.fn(() => chain); + chain.where = vi.fn(() => chain); + chain.orderBy = vi.fn(() => Promise.resolve(rows)); + chain.limit = vi.fn(() => Promise.resolve(rows)); + chain.insert = vi.fn(() => chain); + chain.values = vi.fn(() => chain); + chain.returning = vi.fn(() => Promise.resolve(rows)); + chain.innerJoin = vi.fn(() => chain); + return chain; +} + +describe("conversation branch query module exports", () => { + it("exports createBranch", () => { + expect(typeof branchQueries.createBranch).toBe("function"); + }); + + it("exports listBranchesByParent", () => { + expect(typeof branchQueries.listBranchesByParent).toBe("function"); + }); + + it("exports getBranchForRoot", () => { + expect(typeof branchQueries.getBranchForRoot).toBe("function"); + }); + + it("exports getBranchByConversation", () => { + expect(typeof branchQueries.getBranchByConversation).toBe("function"); + }); + + it("exports getBranchOrigin", () => { + expect(typeof branchQueries.getBranchOrigin).toBe("function"); + }); +}); + +describe("conversation branch queries", () => { + it("creates and returns a branch row", async () => { + const branch = { id: "br_1", rootMessageId: "m1" }; + const mockDb = createMockDb([branch]); + + const result = await branchQueries.createBranch(mockDb, { + workspaceId: "w1", + parentConversationId: "parent_c", + branchConversationId: "branch_c", + rootMessageId: "m1", + provider: "claude", + forkSourceTaskId: "task_1", + forkSourceSessionId: "session_1", + createdBy: "u1", + }); + + expect(result).toEqual(branch); + expect(mockDb.values).toHaveBeenCalledWith({ + workspaceId: "w1", + parentConversationId: "parent_c", + branchConversationId: "branch_c", + rootMessageId: "m1", + provider: "claude", + forkSourceTaskId: "task_1", + forkSourceSessionId: "session_1", + createdBy: "u1", + }); + }); + + it("returns null when no root branch exists", async () => { + const mockDb = createMockDb([]); + + const result = await branchQueries.getBranchForRoot(mockDb, { + workspaceId: "w1", + parentConversationId: "parent_c", + rootMessageId: "m_missing", + }); + + expect(result).toBeNull(); + }); + + it("returns null when no branch conversation mapping exists", async () => { + const mockDb = createMockDb([]); + + const result = await branchQueries.getBranchByConversation(mockDb, { + workspaceId: "w1", + branchConversationId: "branch_missing", + }); + + expect(result).toBeNull(); + }); + + it("loads branch origin with the root message join", async () => { + const origin = { + branch: { id: "br_1", rootMessageId: "m1" }, + rootMessage: { id: "m1", content: "last message" }, + }; + const mockDb = createMockDb([origin]); + + const result = await branchQueries.getBranchOrigin(mockDb, { + workspaceId: "w1", + branchConversationId: "branch_c", + }); + + expect(result).toEqual(origin); + expect(mockDb.innerJoin).toHaveBeenCalled(); + expect(mockDb.limit).toHaveBeenCalledWith(1); + }); +}); diff --git a/src/shared/test/queries/conversation.test.ts b/src/shared/test/queries/conversation.test.ts index a14718050..d01c18473 100644 --- a/src/shared/test/queries/conversation.test.ts +++ b/src/shared/test/queries/conversation.test.ts @@ -1,6 +1,12 @@ import { describe, it, expect, vi } from "vitest"; +import { drizzle } from "drizzle-orm/d1"; +import { desc } from "drizzle-orm"; +import { CONVERSATION_TYPES } from "../../src/constants"; +import { conversation } from "../../src/db/schema"; import * as conversationQueries from "../../src/db/queries/conversation"; +const fakeDb = drizzle({} as never); + function createMockDb(rows: any[]) { const chain: any = {}; chain.select = vi.fn(() => chain); @@ -19,6 +25,24 @@ function createMockDb(rows: any[]) { return chain; } +function createOrderedMockDb(rows: any[]) { + const calls: { select?: unknown; where?: unknown } = {}; + const chain: any = {}; + chain.select = vi.fn((selection?: unknown) => { + calls.select = selection; + return chain; + }); + chain.from = vi.fn(() => chain); + chain.where = vi.fn((where: unknown) => { + calls.where = where; + return chain; + }); + chain.orderBy = vi.fn(() => Promise.resolve(rows)); + chain.leftJoin = vi.fn(() => chain); + chain.groupBy = vi.fn(() => chain); + return { chain, calls }; +} + describe("conversation query module exports", () => { it("exports createConversation", () => { expect(typeof conversationQueries.createConversation).toBe("function"); @@ -101,6 +125,67 @@ describe("getConversation", () => { }); }); +describe("listConversations", () => { + it("filters normal conversation history to user DM conversations", async () => { + const rows = [{ id: "conv_1", type: "user_dm_message" }]; + const { chain, calls } = createOrderedMockDb(rows); + + const result = await conversationQueries.listConversations( + chain, + "ws_1", + "usr_1", + ); + + expect(result).toEqual(rows); + expect(calls.where).toBeDefined(); + const { sql, params } = fakeDb + .select() + .from(conversation) + .where(calls.where as any) + .orderBy(desc(conversation.createdAt)) + .toSQL(); + + expect(sql).toContain('"conversation"."type" = ?'); + expect(params).toEqual([ + "ws_1", + "usr_1", + CONVERSATION_TYPES.USER_DM_MESSAGE, + ]); + }); +}); + +describe("listConversationsByAgent", () => { + it("filters agent history to user DM conversations and selects type", async () => { + const rows = [{ id: "conv_1", type: "user_dm_message" }]; + const { chain, calls } = createOrderedMockDb(rows); + + const result = await conversationQueries.listConversationsByAgent( + chain, + "ws_1", + "usr_1", + "ag_1", + ); + + expect(result).toEqual(rows); + expect((calls.select as { type?: unknown }).type).toBe(conversation.type); + expect(calls.where).toBeDefined(); + const { sql, params } = fakeDb + .select() + .from(conversation) + .where(calls.where as any) + .orderBy(desc(conversation.createdAt)) + .toSQL(); + + expect(sql).toContain('"conversation"."type" = ?'); + expect(params).toEqual([ + "ws_1", + "usr_1", + "ag_1", + CONVERSATION_TYPES.USER_DM_MESSAGE, + ]); + }); +}); + describe("updateConversationTitle", () => { it("returns null when no row updated", async () => { const chain: any = {}; diff --git a/src/shared/test/queries/message.test.ts b/src/shared/test/queries/message.test.ts index 62cb3d259..4bcacbec0 100644 --- a/src/shared/test/queries/message.test.ts +++ b/src/shared/test/queries/message.test.ts @@ -6,6 +6,7 @@ function createMockDb(rows: any[]) { const chain: any = {}; chain.select = vi.fn(() => chain); chain.from = vi.fn(() => chain); + chain.innerJoin = vi.fn(() => chain); chain.where = vi.fn(() => chain); chain.orderBy = vi.fn(() => chain); chain.limit = vi.fn(() => Promise.resolve(rows)); @@ -39,6 +40,10 @@ describe("message query module exports", () => { expect(typeof messageQueries.getMessage).toBe("function"); }); + it("exports getMessageForConversation", () => { + expect(typeof messageQueries.getMessageForConversation).toBe("function"); + }); + it("exports updateMessageTaskId", () => { expect(typeof messageQueries.updateMessageTaskId).toBe("function"); }); @@ -46,6 +51,18 @@ describe("message query module exports", () => { it("exports listMessagesAroundTask", () => { expect(typeof messageQueries.listMessagesAroundTask).toBe("function"); }); + + it("exports getLatestNonEventMessage", () => { + expect(typeof messageQueries.getLatestNonEventMessage).toBe("function"); + }); + + it("exports getLatestBranchableMessage", () => { + expect(typeof messageQueries.getLatestBranchableMessage).toBe("function"); + }); + + it("exports isBranchableMessageRoot", () => { + expect(typeof messageQueries.isBranchableMessageRoot).toBe("function"); + }); }); describe("getNewestMessageId", () => { @@ -62,6 +79,142 @@ describe("getNewestMessageId", () => { }); }); +describe("getLatestNonEventMessage", () => { + it("returns the newest active non-event message", async () => { + const latest = { id: "msg_latest", role: "assistant", status: "active" }; + const mockDb = createMockDb([latest]); + + const result = await messageQueries.getLatestNonEventMessage( + mockDb, + "conv_1", + ); + + expect(result).toEqual(latest); + expect(mockDb.orderBy).toHaveBeenCalled(); + expect(mockDb.limit).toHaveBeenCalledWith(1); + }); + + it("returns null when no active non-event message exists", async () => { + const mockDb = createMockDb([]); + + const result = await messageQueries.getLatestNonEventMessage( + mockDb, + "conv_empty", + ); + + expect(result).toBeNull(); + }); +}); + +describe("getLatestBranchableMessage", () => { + it("returns the newest active user/assistant message with a completed session task", async () => { + const latest = { + id: "msg_latest", + role: "assistant", + status: "active", + metadata: JSON.stringify({ kind: "dm" }), + }; + const mockDb = createMockDb([{ msg: latest }]); + + const result = await messageQueries.getLatestBranchableMessage( + mockDb, + "conv_1", + ); + + expect(result).toEqual(latest); + expect(mockDb.innerJoin).toHaveBeenCalled(); + expect(mockDb.orderBy).toHaveBeenCalled(); + expect(mockDb.limit).toHaveBeenCalledWith(100); + }); + + it("skips transient/process rows and returns the previous complete reply", async () => { + const previousComplete = { + id: "msg_done", + role: "assistant", + status: "active", + metadata: JSON.stringify({ kind: "dm" }), + }; + const mockDb = createMockDb([ + { + msg: { + id: "msg_process", + role: "assistant", + status: "active", + metadata: JSON.stringify({ kind: "process", transient: true }), + }, + }, + { msg: previousComplete }, + ]); + + const result = await messageQueries.getLatestBranchableMessage( + mockDb, + "conv_1", + ); + + expect(result).toEqual(previousComplete); + }); + + it("returns null when no completed session-backed branchable message exists", async () => { + const mockDb = createMockDb([]); + + const result = await messageQueries.getLatestBranchableMessage( + mockDb, + "conv_empty", + ); + + expect(result).toBeNull(); + }); +}); + +describe("isBranchableMessageRoot", () => { + it("accepts user roots, current assistant DM roots, and legacy assistant roots", () => { + expect( + messageQueries.isBranchableMessageRoot({ + role: "user", + status: "active", + metadata: null, + }), + ).toBe(true); + expect( + messageQueries.isBranchableMessageRoot({ + role: "assistant", + status: "active", + metadata: JSON.stringify({ kind: "dm" }), + }), + ).toBe(true); + }); + + it("accepts legacy assistant roots without a DM marker", () => { + expect( + messageQueries.isBranchableMessageRoot({ + role: "assistant", + status: "active", + metadata: null, + }), + ).toBe(true); + }); + + it("rejects transient/process roots", () => { + expect( + messageQueries.isBranchableMessageRoot({ + role: "assistant", + status: "active", + metadata: JSON.stringify({ kind: "progress", transient: true }), + }), + ).toBe(false); + }); + + it("rejects runtime error assistant rows", () => { + expect( + messageQueries.isBranchableMessageRoot({ + role: "assistant", + status: "active", + metadata: JSON.stringify({ error_source: "runtime" }), + }), + ).toBe(false); + }); +}); + describe("getActiveMessageCount", () => { it("returns count from query result", async () => { const chain: any = {}; @@ -99,6 +252,35 @@ describe("getMessage", () => { }); }); +describe("getMessageForConversation", () => { + it("scopes the lookup by conversation id and message id", async () => { + const msg = { id: "msg_1", conversationId: "conv_1", content: "hello" }; + const mockDb = createMockDb([msg]); + + const result = await messageQueries.getMessageForConversation( + mockDb, + "conv_1", + "msg_1", + ); + + expect(result).toEqual(msg); + expect(mockDb.where).toHaveBeenCalled(); + expect(mockDb.limit).toHaveBeenCalledWith(1); + }); + + it("returns null when no scoped message matches", async () => { + const mockDb = createMockDb([]); + + const result = await messageQueries.getMessageForConversation( + mockDb, + "conv_1", + "msg_missing", + ); + + expect(result).toBeNull(); + }); +}); + // TC9 — the message.status column (default "active") and the // idx_message_conversation_status index survive the buffer teardown. diff --git a/src/shared/test/queries/task.test.ts b/src/shared/test/queries/task.test.ts index c4526ee13..9f4cd1671 100644 --- a/src/shared/test/queries/task.test.ts +++ b/src/shared/test/queries/task.test.ts @@ -39,6 +39,12 @@ describe("task query module exports", () => { it("exports failStaleRunningTasks", () => { expect(typeof taskQueries.failStaleRunningTasks).toBe("function"); }); + + it("exports getLatestCompletedTaskWithSessionForConversation", () => { + expect( + typeof taskQueries.getLatestCompletedTaskWithSessionForConversation, + ).toBe("function"); + }); }); describe("task query function signatures", () => { @@ -105,6 +111,36 @@ describe("getLatestTaskForConversation", () => { }); }); +describe("getLatestCompletedTaskWithSessionForConversation", () => { + it("returns null when no completed session-backed task exists", async () => { + const mockDb = createMockDb([]); + const result = + await taskQueries.getLatestCompletedTaskWithSessionForConversation( + mockDb, + { workspaceId: "w1", conversationId: "conv_empty" }, + ); + expect(result).toBeNull(); + }); + + it("returns the latest completed task with a session", async () => { + const task = { + id: "task_done", + runtimeId: "rt1", + sessionId: "session_1", + traceId: "trace_1", + }; + const mockDb = createMockDb([task]); + const result = + await taskQueries.getLatestCompletedTaskWithSessionForConversation( + mockDb, + { workspaceId: "w1", conversationId: "conv_1" }, + ); + expect(result).toEqual(task); + expect(mockDb.orderBy).toHaveBeenCalled(); + expect(mockDb.limit).toHaveBeenCalledWith(1); + }); +}); + describe("getTask", () => { it("returns null when task not found", async () => { const chain: any = {}; diff --git a/src/web/migrations/0039_conversation_branch.sql b/src/web/migrations/0039_conversation_branch.sql new file mode 100644 index 000000000..b53799e35 --- /dev/null +++ b/src/web/migrations/0039_conversation_branch.sql @@ -0,0 +1,17 @@ +CREATE TABLE IF NOT EXISTS conversation_branch ( + id TEXT PRIMARY KEY, + workspace_id TEXT NOT NULL REFERENCES workspace(id) ON DELETE CASCADE, + parent_conversation_id TEXT NOT NULL REFERENCES conversation(id) ON DELETE CASCADE, + branch_conversation_id TEXT NOT NULL REFERENCES conversation(id) ON DELETE CASCADE, + root_message_id TEXT NOT NULL REFERENCES message(id) ON DELETE CASCADE, + provider TEXT NOT NULL, + fork_source_task_id TEXT, + fork_source_session_id TEXT, + created_by TEXT NOT NULL REFERENCES "user"(id) ON DELETE CASCADE, + created_at TEXT NOT NULL DEFAULT (datetime('now')), + UNIQUE(workspace_id, parent_conversation_id, root_message_id), + UNIQUE(branch_conversation_id) +); + +CREATE INDEX IF NOT EXISTS idx_conversation_branch_parent + ON conversation_branch(workspace_id, parent_conversation_id); diff --git a/src/web/src/app/api/conversations/[id]/branch-origin/route.test.ts b/src/web/src/app/api/conversations/[id]/branch-origin/route.test.ts new file mode 100644 index 000000000..14a13c21f --- /dev/null +++ b/src/web/src/app/api/conversations/[id]/branch-origin/route.test.ts @@ -0,0 +1,104 @@ +import { describe, it, expect, vi, beforeEach } from "vitest"; +import { NextRequest } from "next/server"; + +const mockGetConversation = vi.fn(); +const mockGetBranchOrigin = vi.fn(); +const mockConversationBranchToResponse = vi.fn((b: any) => ({ + id: b.id, + root_message_id: b.rootMessageId, +})); +const mockMessageToResponse = vi.fn((m: any) => ({ + id: m.id, + content: m.content, +})); + +vi.mock("@opennextjs/cloudflare", () => ({ + getCloudflareContext: vi.fn(() => ({ env: { DB: {} } })), +})); +vi.mock("@/lib/db", () => ({ getDb: vi.fn(() => ({})) })); +vi.mock("@/lib/middleware/auth", () => ({ + withAuth: vi.fn((handler: any) => async (req: any, ctx?: any) => { + const params = ctx?.params instanceof Promise ? await ctx.params : ctx?.params; + return handler(req, { userId: "u1", email: "u@t.com", params }); + }), +})); +vi.mock("@/lib/middleware/workspace", () => ({ + withWorkspaceMember: vi.fn(async () => ({ workspaceId: "w1" })), +})); +vi.mock("@/lib/api/responses", () => ({ + conversationBranchToResponse: (...args: any[]) => + mockConversationBranchToResponse(...args), + messageToResponse: (...args: any[]) => mockMessageToResponse(...args), +})); +vi.mock("@alook/shared", async () => { + const actual = await vi.importActual("@alook/shared"); + return { + ...actual, + CONVERSATION_TYPES: { + USER_DM_MESSAGE: "user_dm_message", + MESSAGE_BRANCH: "message_branch", + }, + queries: { + conversation: { + getConversation: (...args: any[]) => mockGetConversation(...args), + }, + conversationBranch: { + getBranchOrigin: (...args: any[]) => mockGetBranchOrigin(...args), + }, + }, + }; +}); + +import { GET } from "./route"; + +const withParams = (id: string) => ({ params: Promise.resolve({ id }) }); + +describe("GET /api/conversations/[id]/branch-origin", () => { + beforeEach(() => vi.clearAllMocks()); + + it("returns the origin branch and root message for an owned branch conversation", async () => { + mockGetConversation.mockResolvedValue({ + id: "branch_c", + userId: "u1", + type: "message_branch", + }); + mockGetBranchOrigin.mockResolvedValue({ + branch: { id: "br_1", rootMessageId: "m_root" }, + rootMessage: { id: "m_root", content: "original last message" }, + }); + + const res = await GET( + new NextRequest("http://localhost/api/conversations/branch_c/branch-origin"), + withParams("branch_c"), + ); + const body = await res.json(); + + expect(res.status).toBe(200); + expect(body).toEqual({ + branch: { id: "br_1", root_message_id: "m_root" }, + root_message: { id: "m_root", content: "original last message" }, + }); + expect(mockGetBranchOrigin).toHaveBeenCalledWith({}, { + workspaceId: "w1", + branchConversationId: "branch_c", + }); + }); + + it("rejects normal conversations", async () => { + mockGetConversation.mockResolvedValue({ + id: "parent_c", + userId: "u1", + type: "user_dm_message", + }); + + const res = await GET( + new NextRequest("http://localhost/api/conversations/parent_c/branch-origin"), + withParams("parent_c"), + ); + const body = await res.json(); + + expect(res.status).toBe(404); + expect(body.error).toBe("conversation is not a message branch"); + expect(mockGetBranchOrigin).not.toHaveBeenCalled(); + }); +}); diff --git a/src/web/src/app/api/conversations/[id]/branch-origin/route.ts b/src/web/src/app/api/conversations/[id]/branch-origin/route.ts new file mode 100644 index 000000000..4fd8af0e2 --- /dev/null +++ b/src/web/src/app/api/conversations/[id]/branch-origin/route.ts @@ -0,0 +1,46 @@ +import { NextRequest } from "next/server"; +import { getCloudflareContext } from "@opennextjs/cloudflare"; +import { CONVERSATION_TYPES, queries } from "@alook/shared"; +import { getDb } from "@/lib/db"; +import { withAuth } from "@/lib/middleware/auth"; +import { withWorkspaceMember } from "@/lib/middleware/workspace"; +import { writeError, writeJSON } from "@/lib/middleware/helpers"; +import { + conversationBranchToResponse, + messageToResponse, +} from "@/lib/api/responses"; + +export const GET = withAuth(async (req: NextRequest, ctx) => { + const ws = await withWorkspaceMember(req, ctx); + if (ws instanceof Response) return ws; + + const { env } = getCloudflareContext(); + const db = getDb((env as Env).DB); + + const id = ctx.params?.id; + if (!id) return writeError("conversation id is required", 400); + + const conversation = await queries.conversation.getConversation( + db, + id, + ws.workspaceId, + ); + if (!conversation) return writeError("conversation not found", 404); + if (conversation.userId !== ctx.userId) { + return writeError("conversation access denied", 403); + } + if (conversation.type !== CONVERSATION_TYPES.MESSAGE_BRANCH) { + return writeError("conversation is not a message branch", 404); + } + + const origin = await queries.conversationBranch.getBranchOrigin(db, { + workspaceId: ws.workspaceId, + branchConversationId: id, + }); + if (!origin) return writeError("branch origin not found", 404); + + return writeJSON({ + branch: conversationBranchToResponse(origin.branch), + root_message: messageToResponse(origin.rootMessage), + }); +}); diff --git a/src/web/src/app/api/conversations/[id]/branches/route.test.ts b/src/web/src/app/api/conversations/[id]/branches/route.test.ts new file mode 100644 index 000000000..caa374911 --- /dev/null +++ b/src/web/src/app/api/conversations/[id]/branches/route.test.ts @@ -0,0 +1,454 @@ +import { describe, it, expect, vi, beforeEach } from "vitest"; +import { NextRequest } from "next/server"; + +const mockGetConversation = vi.fn(); +const mockCreateConversation = vi.fn(); +const mockListBranchesByParent = vi.fn(); +const mockGetBranchForRoot = vi.fn(); +const mockCreateBranch = vi.fn(); +const mockGetMessageForConversation = vi.fn(); +const mockGetLatestBranchableMessage = vi.fn(); +const mockGetLatestCompletedTaskWithSessionForConversation = vi.fn(); +const mockGetAgent = vi.fn(); +const mockGetRuntime = vi.fn(); +const mockGetTask = vi.fn(); + +const mockConversationBranchToResponse = vi.fn((b: any) => ({ + id: b.id, + root_message_id: b.rootMessageId, + branch_conversation_id: b.branchConversationId, + provider: b.provider, + fork_source_task_id: b.forkSourceTaskId ?? null, + fork_source_session_id: b.forkSourceSessionId ?? null, +})); +const mockConversationToResponse = vi.fn((c: any) => ({ + id: c.id, + type: c.type, + title: c.title, +})); + +vi.mock("@opennextjs/cloudflare", () => ({ + getCloudflareContext: vi.fn(() => ({ env: { DB: {} } })), +})); +vi.mock("@/lib/db", () => ({ getDb: vi.fn(() => ({})) })); +vi.mock("@/lib/middleware/auth", () => ({ + withAuth: vi.fn((handler: any) => async (req: any, ctx?: any) => { + const params = ctx?.params instanceof Promise ? await ctx.params : ctx?.params; + return handler(req, { userId: "u1", email: "u@t.com", params }); + }), +})); +vi.mock("@/lib/middleware/workspace", () => ({ + withWorkspaceMember: vi.fn(async () => ({ workspaceId: "w1" })), +})); +vi.mock("@/lib/api/responses", () => ({ + conversationBranchToResponse: (...args: any[]) => + mockConversationBranchToResponse(...args), + conversationToResponse: (...args: any[]) => mockConversationToResponse(...args), +})); +vi.mock("@alook/shared", async () => { + const actual = await vi.importActual("@alook/shared"); + return { + ...actual, + CONVERSATION_TYPES: { + USER_DM_MESSAGE: "user_dm_message", + MESSAGE_BRANCH: "message_branch", + }, + queries: { + conversation: { + getConversation: (...args: any[]) => mockGetConversation(...args), + createConversation: (...args: any[]) => + mockCreateConversation(...args), + }, + conversationBranch: { + listBranchesByParent: (...args: any[]) => + mockListBranchesByParent(...args), + getBranchForRoot: (...args: any[]) => mockGetBranchForRoot(...args), + createBranch: (...args: any[]) => mockCreateBranch(...args), + }, + message: { + isBranchableMessageRoot: (...args: any[]) => + (actual as any).queries.message.isBranchableMessageRoot(...args), + getMessageForConversation: (...args: any[]) => + mockGetMessageForConversation(...args), + getLatestBranchableMessage: (...args: any[]) => + mockGetLatestBranchableMessage(...args), + }, + agent: { + getAgent: (...args: any[]) => mockGetAgent(...args), + }, + runtime: { + getAgentRuntimeForWorkspace: (...args: any[]) => mockGetRuntime(...args), + }, + task: { + getLatestCompletedTaskWithSessionForConversation: (...args: any[]) => + mockGetLatestCompletedTaskWithSessionForConversation(...args), + getTask: (...args: any[]) => mockGetTask(...args), + }, + }, + }; +}); + +import { GET, POST } from "./route"; + +const withParams = (id: string) => ({ params: Promise.resolve({ id }) }); +const parentConversation = { + id: "parent_c", + workspaceId: "w1", + userId: "u1", + agentId: "a1", + title: "Parent", + type: "user_dm_message", + channel: "default", +}; + +describe("/api/conversations/[id]/branches", () => { + beforeEach(() => { + vi.clearAllMocks(); + }); + + it("lists branches for an owned parent conversation", async () => { + mockGetConversation.mockResolvedValue(parentConversation); + mockListBranchesByParent.mockResolvedValue([ + { + id: "br_1", + rootMessageId: "m1", + branchConversationId: "branch_c", + provider: "claude", + }, + ]); + + const res = await GET( + new NextRequest("http://localhost/api/conversations/parent_c/branches"), + withParams("parent_c"), + ); + const body = await res.json(); + + expect(res.status).toBe(200); + expect(body.branches).toEqual([ + { + id: "br_1", + root_message_id: "m1", + branch_conversation_id: "branch_c", + provider: "claude", + fork_source_task_id: null, + fork_source_session_id: null, + }, + ]); + expect(mockListBranchesByParent).toHaveBeenCalledWith({}, { + workspaceId: "w1", + parentConversationId: "parent_c", + }); + }); + + it("reuses an existing branch conversation for the same root message", async () => { + const branch = { + id: "br_1", + rootMessageId: "m1", + branchConversationId: "branch_c", + provider: "claude", + }; + const branchConversation = { + id: "branch_c", + title: "Parent", + type: "message_branch", + }; + mockGetConversation + .mockResolvedValueOnce(parentConversation) + .mockResolvedValueOnce(branchConversation); + mockGetMessageForConversation.mockResolvedValue({ + id: "m1", + conversationId: "parent_c", + role: "assistant", + status: "active", + metadata: { kind: "dm" }, + }); + mockGetBranchForRoot.mockResolvedValue(branch); + + const res = await POST( + new NextRequest("http://localhost/api/conversations/parent_c/branches", { + method: "POST", + body: JSON.stringify({ root_message_id: "m1" }), + headers: { "Content-Type": "application/json" }, + }), + withParams("parent_c"), + ); + const body = await res.json(); + + expect(res.status).toBe(200); + expect(body.conversation.id).toBe("branch_c"); + expect(mockCreateConversation).not.toHaveBeenCalled(); + expect(mockCreateBranch).not.toHaveBeenCalled(); + }); + + it("rejects a transient root before reusing an existing branch row", async () => { + mockGetConversation.mockResolvedValue(parentConversation); + mockGetMessageForConversation.mockResolvedValue({ + id: "m_process", + conversationId: "parent_c", + role: "assistant", + status: "active", + metadata: { kind: "process", transient: true }, + }); + mockGetBranchForRoot.mockResolvedValue({ + id: "br_bad", + rootMessageId: "m_process", + branchConversationId: "branch_c", + provider: "codex", + }); + + const res = await POST( + new NextRequest("http://localhost/api/conversations/parent_c/branches", { + method: "POST", + body: JSON.stringify({ root_message_id: "m_process" }), + headers: { "Content-Type": "application/json" }, + }), + withParams("parent_c"), + ); + const body = await res.json(); + + expect(res.status).toBe(400); + expect(body.error).toBe("root message cannot be branched"); + expect(mockGetBranchForRoot).not.toHaveBeenCalled(); + expect(mockCreateConversation).not.toHaveBeenCalled(); + expect(mockCreateBranch).not.toHaveBeenCalled(); + }); + + it("creates a branch for the latest completed branchable message on a supported runtime", async () => { + const rootMessage = { + id: "m_latest", + conversationId: "parent_c", + role: "assistant", + status: "active", + taskId: "root_task", + metadata: { kind: "dm" }, + }; + const branchConversation = { + id: "branch_c", + title: "Parent", + type: "message_branch", + }; + const branch = { + id: "br_1", + rootMessageId: "m_latest", + branchConversationId: "branch_c", + provider: "codex", + forkSourceTaskId: "root_task", + forkSourceSessionId: "root_session", + }; + mockGetConversation.mockResolvedValue(parentConversation); + mockGetBranchForRoot.mockResolvedValue(null); + mockGetMessageForConversation.mockResolvedValue(rootMessage); + mockGetAgent.mockResolvedValue({ id: "a1", runtimeId: "rt1" }); + mockGetLatestCompletedTaskWithSessionForConversation.mockResolvedValue({ + id: "root_task", + runtimeId: "rt1", + sessionId: "root_session", + }); + mockGetRuntime.mockResolvedValue({ id: "rt1", provider: "codex" }); + mockCreateConversation.mockResolvedValue(branchConversation); + mockCreateBranch.mockResolvedValue(branch); + + const res = await POST( + new NextRequest("http://localhost/api/conversations/parent_c/branches", { + method: "POST", + body: JSON.stringify({ root_message_id: "m_latest" }), + headers: { "Content-Type": "application/json" }, + }), + withParams("parent_c"), + ); + const body = await res.json(); + + expect(res.status).toBe(201); + expect(body.branch.provider).toBe("codex"); + expect(mockCreateConversation).toHaveBeenCalledWith( + {}, + expect.objectContaining({ + type: "message_branch", + agentId: "a1", + userId: "u1", + }), + ); + expect(mockCreateBranch).toHaveBeenCalledWith( + {}, + expect.objectContaining({ + parentConversationId: "parent_c", + branchConversationId: "branch_c", + rootMessageId: "m_latest", + provider: "codex", + forkSourceTaskId: "root_task", + forkSourceSessionId: "root_session", + }), + ); + }); + + it("rejects branch creation when no completed fork source session exists", async () => { + const rootMessage = { + id: "m_latest", + conversationId: "parent_c", + role: "assistant", + status: "active", + metadata: { kind: "dm" }, + }; + mockGetConversation.mockResolvedValue(parentConversation); + mockGetBranchForRoot.mockResolvedValue(null); + mockGetMessageForConversation.mockResolvedValue(rootMessage); + mockGetAgent.mockResolvedValue({ id: "a1", runtimeId: "rt1" }); + mockGetLatestCompletedTaskWithSessionForConversation.mockResolvedValue(null); + + const res = await POST( + new NextRequest("http://localhost/api/conversations/parent_c/branches", { + method: "POST", + body: JSON.stringify({ root_message_id: "m_latest" }), + headers: { "Content-Type": "application/json" }, + }), + withParams("parent_c"), + ); + const body = await res.json(); + + expect(res.status).toBe(409); + expect(body.error).toBe("branch fork source session is not available"); + expect(mockCreateConversation).not.toHaveBeenCalled(); + expect(mockCreateBranch).not.toHaveBeenCalled(); + }); + + it("creates a branch from the previous completed message while the parent conversation has an active task", async () => { + const rootMessage = { + id: "m_previous_completed", + conversationId: "parent_c", + role: "assistant", + status: "active", + taskId: "root_task", + metadata: { kind: "dm" }, + }; + const branchConversation = { + id: "branch_c", + title: "Parent", + type: "message_branch", + }; + const branch = { + id: "br_1", + rootMessageId: "m_previous_completed", + branchConversationId: "branch_c", + provider: "claude", + forkSourceTaskId: "latest_done_task", + forkSourceSessionId: "latest_done_session", + }; + mockGetConversation.mockResolvedValue(parentConversation); + mockGetBranchForRoot.mockResolvedValue(null); + mockGetMessageForConversation.mockResolvedValue(rootMessage); + mockGetAgent.mockResolvedValue({ id: "a1", runtimeId: "rt1" }); + mockGetLatestCompletedTaskWithSessionForConversation.mockResolvedValue({ + id: "latest_done_task", + runtimeId: "rt1", + sessionId: "latest_done_session", + }); + mockGetRuntime.mockResolvedValue({ id: "rt1", provider: "claude" }); + mockCreateConversation.mockResolvedValue(branchConversation); + mockCreateBranch.mockResolvedValue(branch); + + const res = await POST( + new NextRequest("http://localhost/api/conversations/parent_c/branches", { + method: "POST", + body: JSON.stringify({ root_message_id: "m_previous_completed" }), + headers: { "Content-Type": "application/json" }, + }), + withParams("parent_c"), + ); + const body = await res.json(); + + expect(res.status).toBe(201); + expect(body.branch.root_message_id).toBe("m_previous_completed"); + expect(mockCreateBranch).toHaveBeenCalledWith( + {}, + expect.objectContaining({ + rootMessageId: "m_previous_completed", + provider: "claude", + forkSourceTaskId: "latest_done_task", + forkSourceSessionId: "latest_done_session", + }), + ); + }); + + it("rejects transient process roots even if stale UI sends one", async () => { + mockGetConversation.mockResolvedValue(parentConversation); + mockGetBranchForRoot.mockResolvedValue(null); + mockGetMessageForConversation.mockResolvedValue({ + id: "m_process", + conversationId: "parent_c", + role: "assistant", + status: "active", + metadata: JSON.stringify({ kind: "process", transient: true }), + }); + + const res = await POST( + new NextRequest("http://localhost/api/conversations/parent_c/branches", { + method: "POST", + body: JSON.stringify({ root_message_id: "m_process" }), + headers: { "Content-Type": "application/json" }, + }), + withParams("parent_c"), + ); + const body = await res.json(); + + expect(res.status).toBe(400); + expect(body.error).toBe("root message cannot be branched"); + expect(mockGetLatestBranchableMessage).not.toHaveBeenCalled(); + expect(mockCreateConversation).not.toHaveBeenCalled(); + expect(mockCreateBranch).not.toHaveBeenCalled(); + }); + + it("creates a branch for an older user message using the latest completed fork source", async () => { + const branchConversation = { + id: "branch_c", + title: "Parent", + type: "message_branch", + }; + const branch = { + id: "br_1", + rootMessageId: "m_old", + branchConversationId: "branch_c", + provider: "codex", + forkSourceTaskId: "latest_done_task", + forkSourceSessionId: "latest_done_session", + }; + mockGetConversation.mockResolvedValue(parentConversation); + mockGetBranchForRoot.mockResolvedValue(null); + mockGetMessageForConversation.mockResolvedValue({ + id: "m_old", + conversationId: "parent_c", + role: "user", + status: "active", + }); + mockGetAgent.mockResolvedValue({ id: "a1", runtimeId: "rt1" }); + mockGetLatestCompletedTaskWithSessionForConversation.mockResolvedValue({ + id: "latest_done_task", + runtimeId: "rt1", + sessionId: "latest_done_session", + }); + mockGetRuntime.mockResolvedValue({ id: "rt1", provider: "codex" }); + mockCreateConversation.mockResolvedValue(branchConversation); + mockCreateBranch.mockResolvedValue(branch); + + const res = await POST( + new NextRequest("http://localhost/api/conversations/parent_c/branches", { + method: "POST", + body: JSON.stringify({ root_message_id: "m_old" }), + headers: { "Content-Type": "application/json" }, + }), + withParams("parent_c"), + ); + const body = await res.json(); + + expect(res.status).toBe(201); + expect(body.branch.root_message_id).toBe("m_old"); + expect(body.branch.fork_source_task_id).toBe("latest_done_task"); + expect(mockCreateBranch).toHaveBeenCalledWith( + {}, + expect.objectContaining({ + rootMessageId: "m_old", + forkSourceTaskId: "latest_done_task", + forkSourceSessionId: "latest_done_session", + }), + ); + }); +}); diff --git a/src/web/src/app/api/conversations/[id]/branches/route.ts b/src/web/src/app/api/conversations/[id]/branches/route.ts new file mode 100644 index 000000000..a1e467805 --- /dev/null +++ b/src/web/src/app/api/conversations/[id]/branches/route.ts @@ -0,0 +1,173 @@ +import { NextRequest } from "next/server"; +import { getCloudflareContext } from "@opennextjs/cloudflare"; +import { CONVERSATION_TYPES, CreateBranchRequestSchema, queries } from "@alook/shared"; +import { getDb } from "@/lib/db"; +import { withAuth } from "@/lib/middleware/auth"; +import { withWorkspaceMember } from "@/lib/middleware/workspace"; +import { parseBody, writeError, writeJSON } from "@/lib/middleware/helpers"; +import { + conversationBranchToResponse, + conversationToResponse, +} from "@/lib/api/responses"; + +const BRANCHABLE_PROVIDERS = new Set(["claude", "codex"]); +const FORK_SOURCE_UNAVAILABLE = "branch fork source session is not available"; + +export const GET = withAuth(async (_req: NextRequest, ctx) => { + const ws = await withWorkspaceMember(_req, ctx); + if (ws instanceof Response) return ws; + + const { env } = getCloudflareContext(); + const db = getDb((env as Env).DB); + + const id = ctx.params?.id; + if (!id) return writeError("conversation id is required", 400); + + const conversation = await queries.conversation.getConversation( + db, + id, + ws.workspaceId, + ); + if (!conversation) return writeError("conversation not found", 404); + if (conversation.userId !== ctx.userId) { + return writeError("conversation access denied", 403); + } + + const branches = await queries.conversationBranch.listBranchesByParent(db, { + workspaceId: ws.workspaceId, + parentConversationId: id, + }); + return writeJSON({ branches: branches.map(conversationBranchToResponse) }); +}); + +export const POST = withAuth(async (req: NextRequest, ctx) => { + const ws = await withWorkspaceMember(req, ctx); + if (ws instanceof Response) return ws; + + const { env } = getCloudflareContext(); + const db = getDb((env as Env).DB); + + const id = ctx.params?.id; + if (!id) return writeError("conversation id is required", 400); + + const [body, err] = await parseBody(req, CreateBranchRequestSchema); + if (err) return err; + + const parentConversation = await queries.conversation.getConversation( + db, + id, + ws.workspaceId, + ); + if (!parentConversation) return writeError("conversation not found", 404); + if (parentConversation.userId !== ctx.userId) { + return writeError("conversation access denied", 403); + } + if (parentConversation.type !== CONVERSATION_TYPES.USER_DM_MESSAGE) { + return writeError("only normal conversations can be branched", 400); + } + + const rootMessage = await queries.message.getMessageForConversation( + db, + id, + body.root_message_id, + ); + if ( + !rootMessage || + rootMessage.status !== "active" + ) { + return writeError("root message not found", 404); + } + if (!queries.message.isBranchableMessageRoot(rootMessage)) { + return writeError("root message cannot be branched", 400); + } + + const existingBranch = await queries.conversationBranch.getBranchForRoot(db, { + workspaceId: ws.workspaceId, + parentConversationId: id, + rootMessageId: body.root_message_id, + }); + if (existingBranch) { + const existingConversation = await queries.conversation.getConversation( + db, + existingBranch.branchConversationId, + ws.workspaceId, + ); + if (!existingConversation) { + return writeError("branch conversation not found", 500); + } + return writeJSON({ + branch: conversationBranchToResponse(existingBranch), + conversation: conversationToResponse(existingConversation), + }); + } + + const agent = await queries.agent.getAgent( + db, + parentConversation.agentId, + ws.workspaceId, + ctx.userId, + ); + if (!agent) return writeError("agent not found", 404); + if (!agent.runtimeId) return writeError("agent has no runtime", 400); + + const forkSource = await queries.task.getLatestCompletedTaskWithSessionForConversation( + db, + { workspaceId: ws.workspaceId, conversationId: id }, + ); + if (!forkSource?.sessionId) { + return writeError(FORK_SOURCE_UNAVAILABLE, 409); + } + const forkRuntime = await queries.runtime.getAgentRuntimeForWorkspace( + db, + forkSource.runtimeId, + ws.workspaceId, + ); + const forkProvider = forkRuntime?.provider ?? null; + if (!forkProvider || !BRANCHABLE_PROVIDERS.has(forkProvider)) { + return writeError("fork source runtime does not support message branching", 400); + } + + const runtime = await queries.runtime.getAgentRuntimeForWorkspace( + db, + agent.runtimeId, + ws.workspaceId, + ); + const provider = runtime?.provider ?? null; + if (!provider || !BRANCHABLE_PROVIDERS.has(provider)) { + return writeError("agent runtime does not support message branching", 400); + } + if (provider !== forkProvider) { + return writeError( + "agent runtime provider does not match fork source provider", + 409, + ); + } + + const branchConversation = await queries.conversation.createConversation(db, { + workspaceId: ws.workspaceId, + agentId: parentConversation.agentId, + userId: ctx.userId, + title: parentConversation.title, + type: CONVERSATION_TYPES.MESSAGE_BRANCH, + channel: parentConversation.channel, + }); + + const branch = await queries.conversationBranch.createBranch(db, { + workspaceId: ws.workspaceId, + parentConversationId: parentConversation.id, + branchConversationId: branchConversation.id, + rootMessageId: rootMessage.id, + provider: forkProvider, + forkSourceTaskId: forkSource.id, + forkSourceSessionId: forkSource.sessionId, + createdBy: ctx.userId, + }); + + return writeJSON( + { + branch: conversationBranchToResponse(branch), + conversation: conversationToResponse(branchConversation), + }, + 201, + ); +}); diff --git a/src/web/src/app/api/conversations/[id]/messages/route.test.ts b/src/web/src/app/api/conversations/[id]/messages/route.test.ts index d13c371bb..c3f936eac 100644 --- a/src/web/src/app/api/conversations/[id]/messages/route.test.ts +++ b/src/web/src/app/api/conversations/[id]/messages/route.test.ts @@ -6,6 +6,12 @@ const mockUpdateConversationTitle = vi.fn(); const mockListMessages = vi.fn(); const mockCreateMessage = vi.fn(); const mockEnqueueTask = vi.fn(); +const mockGetBranchByConversation = vi.fn(); +const mockGetLatestTaskForConversation = vi.fn(); +const mockGetLatestCompletedTaskWithSessionForConversation = vi.fn(); +const mockGetMessageForConversation = vi.fn(); +const mockGetTask = vi.fn(); +const mockGetRuntime = vi.fn(); const mockMessageToResponse = vi.fn((m: any) => ({ id: m.id, content: m.content })); const mockTaskToResponse = vi.fn((t: any) => ({ id: t.id, status: t.status })); @@ -36,8 +42,24 @@ vi.mock("@alook/shared", async () => { message: { listMessages: (...args: any[]) => mockListMessages(...args), createMessage: (...args: any[]) => mockCreateMessage(...args), + getMessageForConversation: (...args: any[]) => + mockGetMessageForConversation(...args), updateMessageTaskId: vi.fn().mockResolvedValue(undefined), }, + task: { + getTask: (...args: any[]) => mockGetTask(...args), + getLatestTaskForConversation: (...args: any[]) => + mockGetLatestTaskForConversation(...args), + getLatestCompletedTaskWithSessionForConversation: (...args: any[]) => + mockGetLatestCompletedTaskWithSessionForConversation(...args), + }, + runtime: { + getAgentRuntimeForWorkspace: (...args: any[]) => mockGetRuntime(...args), + }, + conversationBranch: { + getBranchByConversation: (...args: any[]) => + mockGetBranchByConversation(...args), + }, }, }; }); @@ -179,7 +201,12 @@ describe("GET /api/conversations/[id]/messages", () => { }); describe("POST /api/conversations/[id]/messages", () => { - beforeEach(() => vi.clearAllMocks()); + beforeEach(() => { + vi.clearAllMocks(); + mockGetBranchByConversation.mockResolvedValue(null); + mockGetLatestTaskForConversation.mockResolvedValue(null); + mockGetLatestCompletedTaskWithSessionForConversation.mockResolvedValue(null); + }); it("sends message and enqueues task, returns 201", async () => { const conv = { id: "c1", workspaceId: "w1", agentId: "a1" }; @@ -206,6 +233,304 @@ describe("POST /api/conversations/[id]/messages", () => { expect(mockEnqueueTask).toHaveBeenCalledWith("a1", "c1", "w1", "Hi there", "user_dm_message", expect.objectContaining({ contextKey: "c1", traceId: expect.stringMatching(/^tr_/), parentTaskId: null })); }); + it("adds runtime branch context to the first message in a branch conversation", async () => { + const conv = { id: "branch_c", workspaceId: "w1", agentId: "a1" }; + const msg = { id: "m1", content: "fork from here" }; + const task = { id: "t1", status: "pending" }; + mockGetConversation.mockResolvedValue(conv); + mockCreateMessage.mockResolvedValue(msg); + mockUpdateConversationTitle.mockResolvedValue(undefined); + mockGetBranchByConversation.mockResolvedValue({ + parentConversationId: "parent_c", + rootMessageId: "root_m", + provider: "codex", + forkSourceTaskId: "fork_task", + forkSourceSessionId: "fork-session-before-parent-advanced", + }); + mockGetLatestTaskForConversation.mockResolvedValue(null); + mockGetMessageForConversation.mockResolvedValue({ + id: "root_m", + conversationId: "parent_c", + content: "Historical root text", + }); + mockEnqueueTask.mockResolvedValue(task); + + const res = await POST( + new NextRequest("http://localhost/api/conversations/branch_c/messages", { + method: "POST", + body: JSON.stringify({ content: "fork from here" }), + headers: { "Content-Type": "application/json" }, + }), + withParams("branch_c") + ); + + expect(res.status).toBe(201); + expect(mockEnqueueTask).toHaveBeenCalledWith( + "a1", + "branch_c", + "w1", + "fork from here", + "user_dm_message", + expect.objectContaining({ + contextKey: "branch_c", + context: expect.objectContaining({ + message_id: "m1", + quoted_message: { + message_id: "root_m", + excerpt: "Historical root text", + }, + runtime_branch: { + parent_context_key: "parent_c", + parent_task_id: "fork_task", + parent_session_id: "fork-session-before-parent-advanced", + root_message_id: "root_m", + provider: "codex", + }, + }), + traceId: expect.stringMatching(/^tr_/), + parentTaskId: null, + }), + ); + expect(mockCreateMessage).toHaveBeenCalledWith( + {}, + expect.objectContaining({ + metadata: JSON.stringify({ + quote: { messageId: "root_m", excerpt: "Historical root text" }, + }), + }), + ); + }); + + it("uses an explicit quote instead of the branch root quote on first branch send", async () => { + const conv = { id: "branch_c", workspaceId: "w1", agentId: "a1" }; + const msg = { id: "m1", content: "fork from here" }; + const task = { id: "t1", status: "pending" }; + mockGetConversation.mockResolvedValue(conv); + mockCreateMessage.mockResolvedValue(msg); + mockUpdateConversationTitle.mockResolvedValue(undefined); + mockGetBranchByConversation.mockResolvedValue({ + parentConversationId: "parent_c", + rootMessageId: "root_m", + provider: "codex", + forkSourceTaskId: "fork_task", + forkSourceSessionId: "fork_session", + }); + mockGetLatestTaskForConversation.mockResolvedValue(null); + mockGetMessageForConversation.mockResolvedValue({ + id: "root_m", + conversationId: "parent_c", + content: "Historical root text", + }); + mockEnqueueTask.mockResolvedValue(task); + + const res = await POST( + new NextRequest("http://localhost/api/conversations/branch_c/messages", { + method: "POST", + body: JSON.stringify({ + content: "fork from here", + metadata: { + quote: { messageId: "explicit_m", excerpt: "Explicit quote" }, + }, + }), + headers: { "Content-Type": "application/json" }, + }), + withParams("branch_c") + ); + + expect(res.status).toBe(201); + expect(mockCreateMessage).toHaveBeenCalledWith( + {}, + expect.objectContaining({ + metadata: JSON.stringify({ + quote: { messageId: "explicit_m", excerpt: "Explicit quote" }, + }), + }), + ); + expect(mockEnqueueTask).toHaveBeenCalledWith( + "a1", + "branch_c", + "w1", + "fork from here", + "user_dm_message", + expect.objectContaining({ + context: expect.objectContaining({ + quoted_message: { + message_id: "explicit_m", + excerpt: "Explicit quote", + }, + runtime_branch: expect.objectContaining({ + parent_task_id: "fork_task", + parent_session_id: "fork_session", + }), + }), + }), + ); + }); + + it("keeps runtime branch context after a failed branch task without a session", async () => { + const conv = { id: "branch_c", workspaceId: "w1", agentId: "a1" }; + const msg = { id: "m1", content: "retry fork" }; + const task = { id: "t_retry", status: "pending" }; + mockGetConversation.mockResolvedValue(conv); + mockCreateMessage.mockResolvedValue(msg); + mockUpdateConversationTitle.mockResolvedValue(undefined); + mockGetBranchByConversation.mockResolvedValue({ + parentConversationId: "parent_c", + rootMessageId: "root_m", + provider: "codex", + forkSourceTaskId: "fork_task", + forkSourceSessionId: "fork_session", + }); + mockGetLatestTaskForConversation.mockResolvedValue({ + id: "failed_before_session", + traceId: "trace_failed", + status: "failed", + }); + mockGetLatestCompletedTaskWithSessionForConversation.mockResolvedValue(null); + mockGetMessageForConversation.mockResolvedValue({ + id: "root_m", + conversationId: "parent_c", + content: "Historical root text", + }); + mockEnqueueTask.mockResolvedValue(task); + + const res = await POST( + new NextRequest("http://localhost/api/conversations/branch_c/messages", { + method: "POST", + body: JSON.stringify({ content: "retry fork" }), + headers: { "Content-Type": "application/json" }, + }), + withParams("branch_c") + ); + + expect(res.status).toBe(201); + const options = mockEnqueueTask.mock.calls[0][5]; + expect(options.parentTaskId).toBeNull(); + expect(options.traceId).toMatch(/^tr_/); + expect(options.context.runtime_branch).toEqual({ + parent_context_key: "parent_c", + parent_task_id: "fork_task", + parent_session_id: "fork_session", + root_message_id: "root_m", + provider: "codex", + }); + }); + + it("keeps runtime branch context after a queued branch task without a session", async () => { + const conv = { id: "branch_c", workspaceId: "w1", agentId: "a1" }; + const msg = { id: "m1", content: "retry fork" }; + const task = { id: "t_retry", status: "pending" }; + mockGetConversation.mockResolvedValue(conv); + mockCreateMessage.mockResolvedValue(msg); + mockUpdateConversationTitle.mockResolvedValue(undefined); + mockGetBranchByConversation.mockResolvedValue({ + parentConversationId: "parent_c", + rootMessageId: "root_m", + provider: "codex", + forkSourceTaskId: "fork_task", + forkSourceSessionId: "fork_session", + }); + mockGetLatestTaskForConversation.mockResolvedValue({ + id: "queued_before_session", + traceId: "trace_queued", + status: "queued", + }); + mockGetLatestCompletedTaskWithSessionForConversation.mockResolvedValue(null); + mockGetMessageForConversation.mockResolvedValue({ + id: "root_m", + conversationId: "parent_c", + content: "Historical root text", + }); + mockEnqueueTask.mockResolvedValue(task); + + const res = await POST( + new NextRequest("http://localhost/api/conversations/branch_c/messages", { + method: "POST", + body: JSON.stringify({ content: "retry fork" }), + headers: { "Content-Type": "application/json" }, + }), + withParams("branch_c") + ); + + expect(res.status).toBe(201); + const options = mockEnqueueTask.mock.calls[0][5]; + expect(options.parentTaskId).toBeNull(); + expect(options.traceId).toMatch(/^tr_/); + expect(options.context.runtime_branch).toEqual({ + parent_context_key: "parent_c", + parent_task_id: "fork_task", + parent_session_id: "fork_session", + root_message_id: "root_m", + provider: "codex", + }); + }); + + it("resumes branch conversation after a completed session-backed branch task exists", async () => { + const conv = { id: "branch_c", workspaceId: "w1", agentId: "a1" }; + const msg = { id: "m1", content: "continue branch" }; + const task = { id: "t_next", status: "pending" }; + mockGetConversation.mockResolvedValue(conv); + mockCreateMessage.mockResolvedValue(msg); + mockUpdateConversationTitle.mockResolvedValue(undefined); + mockGetBranchByConversation.mockResolvedValue({ + parentConversationId: "parent_c", + rootMessageId: "root_m", + provider: "codex", + forkSourceTaskId: "fork_task", + forkSourceSessionId: "fork_session", + }); + mockGetLatestCompletedTaskWithSessionForConversation.mockResolvedValue({ + id: "branch_done", + traceId: "trace_branch", + sessionId: "branch_session", + }); + mockEnqueueTask.mockResolvedValue(task); + + const res = await POST( + new NextRequest("http://localhost/api/conversations/branch_c/messages", { + method: "POST", + body: JSON.stringify({ content: "continue branch" }), + headers: { "Content-Type": "application/json" }, + }), + withParams("branch_c") + ); + + expect(res.status).toBe(201); + const options = mockEnqueueTask.mock.calls[0][5]; + expect(options.parentTaskId).toBe("branch_done"); + expect(options.traceId).toBe("trace_branch"); + expect(options.context).toEqual({ message_id: "m1" }); + expect(mockGetMessageForConversation).not.toHaveBeenCalled(); + }); + + it("rejects the first branch message when the persisted fork source is missing", async () => { + const conv = { id: "branch_c", workspaceId: "w1", agentId: "a1" }; + mockGetConversation.mockResolvedValue(conv); + mockGetBranchByConversation.mockResolvedValue({ + parentConversationId: "parent_c", + rootMessageId: "root_m", + provider: "codex", + forkSourceTaskId: null, + forkSourceSessionId: null, + }); + mockGetLatestTaskForConversation.mockResolvedValue(null); + + const res = await POST( + new NextRequest("http://localhost/api/conversations/branch_c/messages", { + method: "POST", + body: JSON.stringify({ content: "fork from here" }), + headers: { "Content-Type": "application/json" }, + }), + withParams("branch_c") + ); + const body = await res.json(); + + expect(res.status).toBe(409); + expect(body.error).toBe("branch fork source session is not available"); + expect(mockCreateMessage).not.toHaveBeenCalled(); + expect(mockEnqueueTask).not.toHaveBeenCalled(); + }); + it("auto-titles conversation with truncated first message", async () => { const longContent = "A ".repeat(40).trim(); const conv = { id: "c1", workspaceId: "w1", agentId: "a1" }; diff --git a/src/web/src/app/api/conversations/[id]/messages/route.ts b/src/web/src/app/api/conversations/[id]/messages/route.ts index dc9df6a9b..93c8d7a47 100644 --- a/src/web/src/app/api/conversations/[id]/messages/route.ts +++ b/src/web/src/app/api/conversations/[id]/messages/route.ts @@ -14,11 +14,28 @@ import { invalidate, cacheKeys } from "@/lib/cache"; const MAX_FILE_SIZE = 10 * 1024 * 1024; // 10 MB const MAX_FILES = 10; +const FORK_SOURCE_UNAVAILABLE = "branch fork source session is not available"; + +type QuoteMetadata = { messageId: string; excerpt: string }; function sanitizeFilename(name: string): string { return name.replace(/[/\\]/g, "_").replace(/\.\./g, "_").slice(0, 255) || "file"; } +function normalizeQuote(value: unknown): QuoteMetadata | null { + if (!value || typeof value !== "object") return null; + const quote = value as { messageId?: unknown; excerpt?: unknown }; + if (typeof quote.messageId !== "string" || !quote.messageId) return null; + return { + messageId: quote.messageId, + excerpt: typeof quote.excerpt === "string" ? quote.excerpt : "", + }; +} + +function quoteExcerpt(content: string): string { + return content.trim().replace(/\s+/g, " ").slice(0, 100); +} + export const GET = withAuth(async (req, ctx) => { const ws = await withWorkspaceMember(req, ctx); if (ws instanceof Response) return ws; @@ -143,12 +160,55 @@ export const POST = withAuth(async (req: NextRequest, ctx) => { artifactIds.push(artifactId); } + const branch = await queries.conversationBranch.getBranchByConversation(db, { + workspaceId: ws.workspaceId, + branchConversationId: id, + }); + const resumableBranchTask = branch + ? await queries.task.getLatestCompletedTaskWithSessionForConversation(db, { + workspaceId: ws.workspaceId, + conversationId: id, + }) + : null; + let runtimeBranchContext: Record | undefined; + let effectiveMessageMetadata = messageMetadata; + if (branch && !resumableBranchTask) { + if (!branch.forkSourceTaskId || !branch.forkSourceSessionId) { + return writeError(FORK_SOURCE_UNAVAILABLE, 409); + } + const rootMessage = await queries.message.getMessageForConversation( + db, + branch.parentConversationId, + branch.rootMessageId, + ); + if (!rootMessage) { + return writeError("branch root message not found", 409); + } + const explicitQuote = normalizeQuote(messageMetadata?.quote); + if (!explicitQuote) { + effectiveMessageMetadata = { + ...(messageMetadata ?? {}), + quote: { + messageId: branch.rootMessageId, + excerpt: quoteExcerpt(rootMessage.content), + }, + }; + } + runtimeBranchContext = { + parent_context_key: branch.parentConversationId, + parent_task_id: branch.forkSourceTaskId, + parent_session_id: branch.forkSourceSessionId, + root_message_id: branch.rootMessageId, + provider: branch.provider, + }; + } + const message = await queries.message.createMessage(db, { conversationId: id, role: "user", content, attachmentIds: artifactIds.length > 0 ? JSON.stringify(artifactIds) : null, - metadata: messageMetadata ? JSON.stringify(messageMetadata) : null, + metadata: effectiveMessageMetadata ? JSON.stringify(effectiveMessageMetadata) : null, }); broadcastToUser(ctx.userId, { @@ -188,14 +248,15 @@ export const POST = withAuth(async (req: NextRequest, ctx) => { } const contextKey = id; - const quote = messageMetadata?.quote as { messageId?: string; excerpt?: string } | undefined; + const quote = normalizeQuote(effectiveMessageMetadata?.quote); const taskContext: Record = { message_id: message.id, ...(artifactIds.length > 0 ? { attachment_ids: artifactIds } : {}), ...mentionContext, ...(quote ? { quoted_message: { message_id: quote.messageId, excerpt: quote.excerpt } } : {}), + ...(runtimeBranchContext ? { runtime_branch: runtimeBranchContext } : {}), }; - const traceId = "tr_" + nanoid(); + const traceId = resumableBranchTask?.traceId ?? "tr_" + nanoid(); const taskService = new TaskService(db); try { const task = await taskService.enqueueTask( @@ -208,7 +269,7 @@ export const POST = withAuth(async (req: NextRequest, ctx) => { contextKey, context: Object.keys(taskContext).length > 0 ? taskContext : undefined, traceId, - parentTaskId: null, + parentTaskId: resumableBranchTask?.id ?? null, }, ); queries.message.updateMessageTaskId(db, message.id, task.id).catch(() => {}); diff --git a/src/web/src/components/agent-chat/agent-chat-view.test.ts b/src/web/src/components/agent-chat/agent-chat-view.test.ts index 9483cd4b4..01e2aef82 100644 --- a/src/web/src/components/agent-chat/agent-chat-view.test.ts +++ b/src/web/src/components/agent-chat/agent-chat-view.test.ts @@ -1,6 +1,20 @@ import { describe, it, expect } from "vitest"; import type { Message, Artifact } from "@alook/shared"; -import { sortMessages, mergeMessages, buildTimeline, computeGroupPositions, getEventIconType, eventTypeFromMessage, shouldPersistPointerForLoad, pointerRefreshTargetForTaskCreated } from "./chat-message-utils"; +import { + canShowBranchAction, + computeGroupPositions, + getBranchReturnTarget, + getEventIconType, + getLatestBranchableMessageId, + isBranchableMessage, + isMessageBranchActionCandidate, + eventTypeFromMessage, + mergeMessages, + pointerRefreshTargetForTaskCreated, + shouldPersistPointerForLoad, + sortMessages, + buildTimeline, +} from "./chat-message-utils"; import type { NapMarker } from "./chat-message-utils"; function msg(id: string, created_at: string, role: "user" | "assistant" | "event" = "user", content = ""): Message { @@ -586,6 +600,251 @@ describe("pointerRefreshTargetForTaskCreated (TODO-2: WS-driven refresh scope)", }); }); +describe("last-message branch action safety", () => { + it("ignores event rows when choosing the latest branchable message", () => { + const messages = [ + { + ...msg("assistant_done", "2024-01-01T00:00:00Z", "assistant", "done"), + metadata: { kind: "dm" }, + }, + msg("event_status", "2024-01-01T00:00:10Z", "event", "Agent is typing"), + ]; + + expect(getLatestBranchableMessageId(messages)).toBe("assistant_done"); + }); + + it("uses the previous complete visible reply while the current task is active", () => { + const messages: Message[] = [ + { + ...msg("assistant_done", "2024-01-01T00:00:00Z", "assistant", "done"), + task_id: "task_done", + metadata: { kind: "dm" }, + }, + { + ...msg("active_user", "2024-01-01T00:00:05Z", "user", "next"), + task_id: null, + }, + { + ...msg("process_row", "2024-01-01T00:00:10Z", "assistant", "working"), + task_id: "task_running", + metadata: { kind: "process", transient: true }, + }, + ]; + + expect(getLatestBranchableMessageId(messages, "task_running")).toBe( + "assistant_done", + ); + }); + + it("hides Branch on a freshly sent active user prompt while keeping the completed active root", () => { + const messages: Message[] = [ + { + ...msg("assistant_done", "2024-01-01T00:00:00Z", "assistant", "done"), + task_id: "task_done", + metadata: { kind: "dm" }, + }, + { + ...msg("active_user", "2024-01-01T00:00:05Z", "user", "next"), + task_id: null, + }, + { + ...msg("process_row", "2024-01-01T00:00:10Z", "assistant", "working"), + task_id: "task_running", + metadata: { kind: "process", transient: true }, + }, + ]; + const activeBranchRootMessageId = getLatestBranchableMessageId( + messages, + "task_running", + ); + + expect(activeBranchRootMessageId).toBe("assistant_done"); + expect( + isMessageBranchActionCandidate({ + message: messages[1], + isTaskActive: true, + activeBranchRootMessageId, + hasExistingBranch: false, + }), + ).toBe(false); + expect( + canShowBranchAction({ + conversationType: "user_dm_message", + supportsBranch: true, + branchingMessageId: null, + hasExistingBranch: false, + messageIsBranchable: isMessageBranchActionCandidate({ + message: messages[1], + isTaskActive: true, + activeBranchRootMessageId, + hasExistingBranch: false, + }), + }), + ).toBe(false); + expect( + isMessageBranchActionCandidate({ + message: messages[0], + isTaskActive: true, + activeBranchRootMessageId, + hasExistingBranch: false, + }), + ).toBe(true); + }); + + it("skips runtime-error assistant rows when choosing a branch root", () => { + const messages: Message[] = [ + { + ...msg("assistant_done", "2024-01-01T00:00:00Z", "assistant", "done"), + metadata: { kind: "dm" }, + }, + { + ...msg("runtime_error", "2024-01-01T00:00:10Z", "assistant", "boom"), + metadata: { error_source: "runtime" }, + }, + ]; + + expect(getLatestBranchableMessageId(messages)).toBe("assistant_done"); + }); + + it("shows branch creation on the previous completed reply while typing is active", () => { + expect( + canShowBranchAction({ + conversationType: "user_dm_message", + supportsBranch: true, + branchingMessageId: null, + hasExistingBranch: false, + messageIsBranchable: true, + }), + ).toBe(true); + }); + + it("shows branch creation on historical user messages", () => { + const historicalUser = msg( + "historical_user", + "2024-01-01T00:00:00Z", + "user", + "previous prompt", + ); + + expect(isBranchableMessage(historicalUser)).toBe(true); + expect( + canShowBranchAction({ + conversationType: "user_dm_message", + supportsBranch: true, + branchingMessageId: null, + hasExistingBranch: false, + messageIsBranchable: isBranchableMessage(historicalUser), + }), + ).toBe(true); + }); + + it("hides branch creation on non-branchable process rows", () => { + const processMessage: Message = { + ...msg("process_row", "2024-01-01T00:00:00Z", "assistant", "working"), + metadata: { kind: "process", transient: true }, + }; + + expect(isBranchableMessage(processMessage)).toBe(false); + expect( + canShowBranchAction({ + conversationType: "user_dm_message", + supportsBranch: true, + branchingMessageId: null, + hasExistingBranch: false, + messageIsBranchable: isBranchableMessage(processMessage), + }), + ).toBe(false); + }); + + it("still allows opening an existing branch while the parent task is active", () => { + const historicalUser = msg( + "historical_user", + "2024-01-01T00:00:00Z", + "user", + "previous prompt", + ); + + expect( + isMessageBranchActionCandidate({ + message: historicalUser, + isTaskActive: true, + activeBranchRootMessageId: "assistant_done", + hasExistingBranch: true, + }), + ).toBe(true); + expect( + canShowBranchAction({ + conversationType: "user_dm_message", + supportsBranch: true, + branchingMessageId: null, + hasExistingBranch: true, + messageIsBranchable: true, + }), + ).toBe(true); + }); + + it("never shows another branch action inside a branch conversation", () => { + expect( + canShowBranchAction({ + conversationType: "message_branch", + supportsBranch: true, + branchingMessageId: null, + hasExistingBranch: false, + messageIsBranchable: true, + }), + ).toBe(false); + }); + + it("returns to the clicked historical message for new branch paths", () => { + const historicalUser = { + ...msg( + "historical_user", + "2024-01-01T00:00:00Z", + "user", + "previous prompt", + ), + task_id: "task_historical", + }; + + expect( + getBranchReturnTarget({ + agentId: "agent_1", + conversationId: "conv_parent", + message: historicalUser, + fallbackTaskId: "task_visible", + }), + ).toEqual({ + agentId: "agent_1", + conversationId: "conv_parent", + taskId: "task_historical", + messageId: "historical_user", + }); + }); + + it("returns to the clicked message for existing branch paths when the message has no task id", () => { + const historicalUser = msg( + "historical_user_no_task", + "2024-01-01T00:00:00Z", + "user", + "previous prompt", + ); + + expect( + getBranchReturnTarget({ + agentId: "agent_1", + conversationId: "conv_parent", + message: historicalUser, + fallbackTaskId: "task_visible", + }), + ).toEqual({ + agentId: "agent_1", + conversationId: "conv_parent", + taskId: "task_visible", + messageId: "historical_user_no_task", + }); + }); +}); + // TC1 — grouping: same-role + <60s clusters into first/middle/last/solo; // event messages never group (stay null). describe("computeGroupPositions (TC1)", () => { diff --git a/src/web/src/components/agent-chat/agent-chat-view.tsx b/src/web/src/components/agent-chat/agent-chat-view.tsx index 6462555a6..b86af4d2f 100644 --- a/src/web/src/components/agent-chat/agent-chat-view.tsx +++ b/src/web/src/components/agent-chat/agent-chat-view.tsx @@ -13,15 +13,25 @@ import { useWorkspace } from "@/contexts/workspace-context"; import { Button } from "@/components/ui/button"; import { TaskStream } from "@/components/task-stream"; import { + createBranch, + getBranchOrigin, getTask, + listBranches, updateIssue, getAgentSkills, cancelActiveTask, } from "@/lib/api"; -import { useLatest } from "@/components/agent-chat/chat-message-utils"; +import { + canShowBranchAction, + getBranchReturnTarget, + getLatestBranchableMessageId, + isMessageBranchActionCandidate, + useLatest, +} from "@/components/agent-chat/chat-message-utils"; import type { Artifact, Issue, + Message, SkillEntry, WsMessage, } from "@alook/shared"; @@ -86,6 +96,11 @@ import { PopoverTrigger, PopoverContent, } from "@/components/ui/popover"; +import { useAgentChatSheet } from "@/contexts/agent-chat-sheet-context"; +import type { + BranchOriginResponse, + CreateBranchResponse, +} from "@/lib/api"; export function AgentChatView({ agentId: propAgentId, @@ -110,6 +125,7 @@ export function AgentChatView({ subscribeReconnect, } = useAgentContext(); const { refresh: refreshInboxCount } = useInboxCount(); + const { openAgentChat } = useAgentChatSheet(); const { activeChannel, loading: channelLoading, @@ -123,6 +139,7 @@ export function AgentChatView({ ? runtimes.find((r) => r.id === activeAgent.runtime_id) : null; const runtimeProvider = activeRuntime?.provider ?? null; + const supportsBranch = runtimeProvider === "claude" || runtimeProvider === "codex"; const scrollToTaskId = propScrollToTaskId !== undefined ? propScrollToTaskId @@ -291,6 +308,139 @@ export function AgentChatView({ handleNap, } = chat; + const isTaskActive = + !!activeTask && + !["completed", "failed", "cancelled", "superseded"].includes( + activeTask.status, + ); + + const [branchingMessageId, setBranchingMessageId] = useState(null); + const [branchByRootMessageId, setBranchByRootMessageId] = useState< + Map + >(() => new Map()); + const [branchOrigin, setBranchOrigin] = useState(null); + + const activeBranchRootMessageId = useMemo( + () => + isTaskActive + ? getLatestBranchableMessageId(messages, activeTask?.id ?? null) + : null, + [messages, isTaskActive, activeTask?.id], + ); + + useEffect(() => { + if (!conversation || !supportsBranch || conversation.type === "message_branch") { + setBranchByRootMessageId(new Map()); + return; + } + + let ignore = false; + listBranches(conversation.id, workspaceId) + .then((res) => { + if (ignore) return; + setBranchByRootMessageId( + new Map(res.branches.map((branch) => [branch.root_message_id, branch])), + ); + }) + .catch(() => { + if (!ignore) setBranchByRootMessageId(new Map()); + }); + + return () => { + ignore = true; + }; + }, [conversation, supportsBranch, workspaceId]); + + useEffect(() => { + if (!conversation || conversation.type !== "message_branch") { + setBranchOrigin(null); + return; + } + + let ignore = false; + getBranchOrigin(conversation.id, workspaceId) + .then((origin) => { + if (!ignore) setBranchOrigin(origin); + }) + .catch(() => { + if (!ignore) setBranchOrigin(null); + }); + + return () => { + ignore = true; + }; + }, [conversation, workspaceId]); + + const handleBranch = useCallback( + async (message: Message) => { + if (!conversation || !supportsBranch || conversation.type === "message_branch") return; + + const returnTo = getBranchReturnTarget({ + agentId, + conversationId: conversation.id, + message, + fallbackTaskId: scrollToTaskId, + }); + const existingBranch = branchByRootMessageId.get(message.id); + if (existingBranch) { + openAgentChat(agentId, { + conversationId: existingBranch.branch_conversation_id, + mode: "branch", + returnTo, + }); + return; + } + + setBranchingMessageId(message.id); + try { + const res = await createBranch(conversation.id, message.id, workspaceId); + setBranchByRootMessageId((prev) => { + const next = new Map(prev); + next.set(message.id, res.branch); + return next; + }); + if (typeof window !== "undefined") { + const key = `chat-draft-meta:${agentId}:${res.conversation.id}`; + const meta = (() => { + try { + return JSON.parse(localStorage.getItem(key) ?? "{}"); + } catch { + return {}; + } + })(); + localStorage.setItem( + key, + JSON.stringify({ + ...meta, + quote: { + id: message.id, + excerpt: message.content.trim().replace(/\s+/g, " ").slice(0, 100), + }, + }), + ); + } + openAgentChat(agentId, { + conversationId: res.conversation.id, + mode: "branch", + returnTo, + }); + } catch (err) { + toast.error(err instanceof Error ? err.message : "Failed to create branch"); + } finally { + setBranchingMessageId(null); + } + }, + [ + agentId, + branchByRootMessageId, + conversation, + openAgentChat, + supportsBranch, + scrollToTaskId, + workspaceId, + ], + ); + const { versionMap, duplicateFilenames } = useMemo( () => computeArtifactVersions(agentArtifacts), [agentArtifacts], @@ -532,12 +682,6 @@ export function AgentChatView({ const [menuOpen, setMenuOpen] = useState(false); const [stopping, setStopping] = useState(false); - const isTaskActive = - !!activeTask && - !["completed", "failed", "cancelled", "superseded"].includes( - activeTask.status, - ); - const handleStop = useCallback(async () => { if (!conversation?.id || stopping) return; setStopping(true); @@ -665,6 +809,18 @@ export function AgentChatView({ )} + {branchOrigin && ( +
+
+ + Branched from {branchOrigin.root_message.role} message +
+
+ {branchOrigin.root_message.content || "(empty message)"} +
+
+ )} + {messages.length === 0 && !activeTask && (() => { @@ -747,6 +903,13 @@ export function AgentChatView({ } const msg = item.data; + const hasExistingBranch = branchByRootMessageId.has(msg.id); + const messageIsBranchable = isMessageBranchActionCandidate({ + message: msg, + isTaskActive, + activeBranchRootMessageId, + hasExistingBranch, + }); return (
openIssue(issueId)} onRetry={handleRetryTask} + canBranch={canShowBranchAction({ + conversationType: conversation?.type, + supportsBranch, + branchingMessageId, + hasExistingBranch, + messageIsBranchable, + })} + isBranched={hasExistingBranch} + onBranchClick={handleBranch} mentionComponents={MENTION_COMPONENTS} isFlagged={flaggedIds.has(msg.id)} onToggleFlag={ diff --git a/src/web/src/components/agent-chat/chat-message-utils.ts b/src/web/src/components/agent-chat/chat-message-utils.ts index b7f98bb64..a207f2413 100644 --- a/src/web/src/components/agent-chat/chat-message-utils.ts +++ b/src/web/src/components/agent-chat/chat-message-utils.ts @@ -53,6 +53,82 @@ export function mergeMessages( return sortMessages([...merged.values()]); } +const NON_BRANCHABLE_MESSAGE_KINDS = new Set([ + "event", + "lifecycle", + "process", + "progress", + "status", + "transient", + "typing", +]); + +export function isBranchableMessage( + message: Message, +): boolean { + const status = message.status as string | undefined; + if (status && status !== "active") return false; + if (message.role !== "user" && message.role !== "assistant") return false; + + const kind = + typeof message.metadata?.kind === "string" + ? message.metadata.kind.toLowerCase() + : null; + if (kind && NON_BRANCHABLE_MESSAGE_KINDS.has(kind)) return false; + if (message.metadata?.transient === true) return false; + if (message.metadata?.error_source) return false; + if (message.role === "assistant") return kind === null || kind === "dm"; + + return true; +} + +export function getLatestBranchableMessageId( + messages: Message[], + activeTaskId?: string | null, +): string | null { + for (let i = messages.length - 1; i >= 0; i--) { + if (activeTaskId && messages[i].role === "user") continue; + if (activeTaskId && messages[i].task_id === activeTaskId) continue; + if (isBranchableMessage(messages[i])) return messages[i].id; + } + return null; +} + +export function isMessageBranchActionCandidate({ + message, + isTaskActive, + activeBranchRootMessageId, + hasExistingBranch, +}: { + message: Message; + isTaskActive: boolean; + activeBranchRootMessageId: string | null; + hasExistingBranch: boolean; +}): boolean { + if (!isBranchableMessage(message)) return false; + if (!isTaskActive) return true; + return hasExistingBranch || message.id === activeBranchRootMessageId; +} + +export function getBranchReturnTarget({ + agentId, + conversationId, + message, + fallbackTaskId, +}: { + agentId: string; + conversationId: string; + message: Message; + fallbackTaskId?: string | null; +}) { + return { + agentId, + conversationId, + taskId: message.task_id ?? fallbackTaskId ?? null, + messageId: message.id, + }; +} + export type NapMarker = { agentName: string; created_at: string; id: string }; type TimelineItem = @@ -264,6 +340,19 @@ export function pointerRefreshTargetForTaskCreated(args: { return task.conversation_id; } +export function canShowBranchAction(args: { + conversationType?: string | null; + supportsBranch: boolean; + branchingMessageId: string | null; + hasExistingBranch: boolean; + messageIsBranchable: boolean; +}): boolean { + if (args.conversationType === "message_branch") return false; + if (!args.supportsBranch || args.branchingMessageId !== null) return false; + if (!args.messageIsBranchable) return false; + return true; +} + export function useLatest(value: T) { const ref = useRef(value); useEffect(() => { diff --git a/src/web/src/components/agent-chat/message-list.tsx b/src/web/src/components/agent-chat/message-list.tsx index 8537a5994..e3a90737c 100644 --- a/src/web/src/components/agent-chat/message-list.tsx +++ b/src/web/src/components/agent-chat/message-list.tsx @@ -11,7 +11,7 @@ import { highlightMentions } from "@/lib/highlight-mentions"; import { TaskStream } from "@/components/task-stream"; import { RuntimeErrorBlock } from "@/components/agent-chat/runtime-error-block"; import { AnimatedAvatar, type AvatarConfig } from "@/components/avatar"; -import { FileText, Flag, Copy, Check, MessageSquareQuote } from "lucide-react"; +import { FileText, Flag, Copy, Check, GitBranch, MessageSquareQuote } from "lucide-react"; import { EmailCard } from "@/components/agent-chat/event-cards/email-card"; import { CalendarCard } from "@/components/agent-chat/event-cards/calendar-card"; import { IssueCard } from "@/components/agent-chat/event-cards/issue-card"; @@ -70,6 +70,9 @@ export interface MessageItemProps { onIssueClick: (issueId: string) => void; onCalendarEventClick: (calendarEventId: string) => void; onRetry?: () => void; + canBranch?: boolean; + isBranched?: boolean; + onBranchClick?: (message: Message) => void; mentionComponents: Record & { children?: React.ReactNode }>>; isFlagged?: boolean; onToggleFlag?: (messageId: string) => void; @@ -388,6 +391,9 @@ export const MessageItem = memo(function MessageItem({ onIssueClick, onCalendarEventClick, onRetry, + canBranch = false, + isBranched = false, + onBranchClick, mentionComponents, isFlagged, onToggleFlag, @@ -453,6 +459,15 @@ export const MessageItem = memo(function MessageItem({ const canCopy = msg.role === "assistant" || msg.role === "user"; const messageActions: MessageAction[] = []; + if (canBranch && onBranchClick) { + messageActions.push({ + key: "branch", + label: isBranched ? "Open branch" : "Branch", + icon: , + onClick: () => onBranchClick(msg), + active: isBranched, + }); + } if (canCopy) { messageActions.push({ key: "copy", diff --git a/src/web/src/components/canvas/agent-chat-sheet.tsx b/src/web/src/components/canvas/agent-chat-sheet.tsx index 3385c2f3a..adf15787b 100644 --- a/src/web/src/components/canvas/agent-chat-sheet.tsx +++ b/src/web/src/components/canvas/agent-chat-sheet.tsx @@ -17,9 +17,10 @@ import { useAgentContext } from "@/contexts/agent-context"; import { useWorkspace } from "@/contexts/workspace-context"; import { ChannelBar } from "@/components/channel-bar"; import { AgentChatView } from "@/components/agent-chat/agent-chat-view"; -import { ArrowUpRight, XIcon } from "lucide-react"; +import { ArrowUpRight, CornerUpLeft, XIcon } from "lucide-react"; import { Button, buttonVariants } from "@/components/ui/button"; import { Tooltip, TooltipTrigger, TooltipContent } from "@/components/ui/tooltip"; +import type { AgentChatSheetMode } from "@/contexts/agent-chat-sheet-state"; interface AgentChatSheetProps { open: boolean; @@ -34,13 +35,27 @@ interface AgentChatSheetProps { targetConvId?: string | null; scrollToTaskId?: string | null; scrollToMessageId?: string | null; + mode?: AgentChatSheetMode; + canReturnToPreviousMain?: boolean; + onReturnToPreviousMain?: () => void; } const MIN_WIDTH = 320; const MAX_WIDTH_RATIO = 0.8; const DEFAULT_WIDTH = 480; -export function AgentChatSheet({ open, onOpenChange, agentId, agent, targetConvId, scrollToTaskId, scrollToMessageId }: AgentChatSheetProps) { +export function AgentChatSheet({ + open, + onOpenChange, + agentId, + agent, + targetConvId, + scrollToTaskId, + scrollToMessageId, + mode = "main", + canReturnToPreviousMain = false, + onReturnToPreviousMain, +}: AgentChatSheetProps) { const { runtimes, activeTaskCounts } = useAgentContext(); const { slug } = useWorkspace(); const router = useRouter(); @@ -62,7 +77,24 @@ export function AgentChatSheet({ open, onOpenChange, agentId, agent, targetConvI {/* Top-right action buttons */}
- {agentId && (() => { + {canReturnToPreviousMain && onReturnToPreviousMain && ( + + + } + > + + + Return to main conversation + + )} + {agentId && mode !== "branch" && (() => { const params = new URLSearchParams(); if (scrollToTaskId) params.set("task", scrollToTaskId); if (scrollToMessageId) params.set("msg", scrollToMessageId); @@ -114,6 +146,11 @@ export function AgentChatSheet({ open, onOpenChange, agentId, agent, targetConvI {agent?.name ?? "Chat"} + {mode === "branch" && ( + + Branch + + )} {agent?.email_handle && ( {toAlookAddress(agent.email_handle)} )} diff --git a/src/web/src/contexts/agent-chat-sheet-context.tsx b/src/web/src/contexts/agent-chat-sheet-context.tsx index 3469ed229..a97f7a5a9 100644 --- a/src/web/src/contexts/agent-chat-sheet-context.tsx +++ b/src/web/src/contexts/agent-chat-sheet-context.tsx @@ -13,12 +13,16 @@ import { useRouter } from "next/navigation"; import { useAgentContext } from "@/contexts/agent-context"; import { useWorkspace } from "@/contexts/workspace-context"; import { AgentChatSheet } from "@/components/canvas/agent-chat-sheet"; +import { + capturePreviousMainTargetForBranch, + normalizeSheetTarget, + type AgentChatSheetMode, + type AgentChatSheetOpenOptions, + type AgentChatSheetTarget, +} from "@/contexts/agent-chat-sheet-state"; interface AgentChatSheetContextValue { - openAgentChat: ( - agentId: string, - opts?: { conversationId?: string; taskId?: string; messageId?: string }, - ) => void; + openAgentChat: (agentId: string, opts?: AgentChatSheetOpenOptions) => void; } const AgentChatSheetContext = createContext( @@ -44,6 +48,9 @@ export function AgentChatSheetProvider({ children }: { children: ReactNode }) { const [targetConvId, setTargetConvId] = useState(null); const [scrollToTaskId, setScrollToTaskId] = useState(null); const [scrollToMessageId, setScrollToMessageId] = useState(null); + const [sheetMode, setSheetMode] = useState("main"); + const [previousMainTarget, setPreviousMainTarget] = + useState(null); const agent = agentId ? agents.find((a) => a.id === agentId) ?? null : null; @@ -51,19 +58,48 @@ export function AgentChatSheetProvider({ children }: { children: ReactNode }) { useEffect(() => { agentsRef.current = agents; }); const openAgentChat = useCallback( - (id: string, opts?: { conversationId?: string; taskId?: string; messageId?: string }) => { + (id: string, opts?: AgentChatSheetOpenOptions) => { const found = agentsRef.current.find((a) => a.id === id); if (!found) { router.push(`/w/${slug}/agents/${id}`); return; } + const nextMode = opts?.mode ?? "main"; + const currentTarget = agentId + ? normalizeSheetTarget(agentId, { + conversationId: targetConvId, + taskId: scrollToTaskId, + messageId: scrollToMessageId, + }) + : null; + const nextPreviousMainTarget = + nextMode === "branch" + ? capturePreviousMainTargetForBranch({ + isOpen: open, + currentMode: sheetMode, + currentTarget, + explicitReturnTarget: opts?.returnTo, + }) + : null; + setAgentId(id); setTargetConvId(opts?.conversationId ?? null); setScrollToTaskId(opts?.taskId ?? null); setScrollToMessageId(opts?.messageId ?? null); + setSheetMode(nextMode); + setPreviousMainTarget(nextPreviousMainTarget); setOpen(true); }, - [router, slug], + [ + agentId, + open, + router, + scrollToMessageId, + scrollToTaskId, + sheetMode, + slug, + targetConvId, + ], ); const handleOpenChange = useCallback((nextOpen: boolean) => { @@ -72,9 +108,23 @@ export function AgentChatSheetProvider({ children }: { children: ReactNode }) { setTargetConvId(null); setScrollToTaskId(null); setScrollToMessageId(null); + setSheetMode("main"); + setPreviousMainTarget(null); } }, []); + const handleReturnToPreviousMain = useCallback(() => { + if (!previousMainTarget) return; + + setAgentId(previousMainTarget.agentId); + setTargetConvId(previousMainTarget.conversationId); + setScrollToTaskId(previousMainTarget.taskId); + setScrollToMessageId(previousMainTarget.messageId); + setSheetMode("main"); + setPreviousMainTarget(null); + setOpen(true); + }, [previousMainTarget]); + return ( {children} @@ -86,6 +136,11 @@ export function AgentChatSheetProvider({ children }: { children: ReactNode }) { targetConvId={targetConvId} scrollToTaskId={scrollToTaskId} scrollToMessageId={scrollToMessageId} + mode={sheetMode} + canReturnToPreviousMain={ + sheetMode === "branch" && previousMainTarget !== null + } + onReturnToPreviousMain={handleReturnToPreviousMain} /> ); diff --git a/src/web/src/contexts/agent-chat-sheet-state.test.ts b/src/web/src/contexts/agent-chat-sheet-state.test.ts new file mode 100644 index 000000000..e42719678 --- /dev/null +++ b/src/web/src/contexts/agent-chat-sheet-state.test.ts @@ -0,0 +1,95 @@ +import { describe, expect, it } from "vitest"; +import { + capturePreviousMainTargetForBranch, + normalizeSheetTarget, + type AgentChatSheetTarget, +} from "./agent-chat-sheet-state"; + +const mainTarget: AgentChatSheetTarget = { + agentId: "agent_main", + conversationId: "conv_main", + taskId: "task_1", + messageId: "msg_1", +}; + +describe("normalizeSheetTarget", () => { + it("normalizes omitted optional fields to null", () => { + expect(normalizeSheetTarget("agent_main")).toEqual({ + agentId: "agent_main", + conversationId: null, + taskId: null, + messageId: null, + }); + }); +}); + +describe("capturePreviousMainTargetForBranch", () => { + it("captures the current main sheet target for branch return", () => { + expect( + capturePreviousMainTargetForBranch({ + isOpen: true, + currentMode: "main", + currentTarget: mainTarget, + }), + ).toEqual(mainTarget); + }); + + it("uses the explicit return target when the current sheet target has no conversation id", () => { + const explicitReturnTarget = { + agentId: "agent_main", + conversationId: "conv_resolved", + taskId: null, + messageId: null, + }; + + expect( + capturePreviousMainTargetForBranch({ + isOpen: true, + currentMode: "main", + currentTarget: { + agentId: "agent_main", + conversationId: null, + taskId: null, + messageId: null, + }, + explicitReturnTarget, + }), + ).toEqual(explicitReturnTarget); + }); + + it("does not capture when no sheet is open", () => { + expect( + capturePreviousMainTargetForBranch({ + isOpen: false, + currentMode: "main", + currentTarget: mainTarget, + explicitReturnTarget: mainTarget, + }), + ).toBeNull(); + }); + + it("does not capture when the current sheet is already a branch", () => { + expect( + capturePreviousMainTargetForBranch({ + isOpen: true, + currentMode: "branch", + currentTarget: mainTarget, + }), + ).toBeNull(); + }); + + it("does not expose a return target when the main sheet cannot be restored exactly", () => { + expect( + capturePreviousMainTargetForBranch({ + isOpen: true, + currentMode: "main", + currentTarget: { + agentId: "agent_main", + conversationId: null, + taskId: null, + messageId: null, + }, + }), + ).toBeNull(); + }); +}); diff --git a/src/web/src/contexts/agent-chat-sheet-state.ts b/src/web/src/contexts/agent-chat-sheet-state.ts new file mode 100644 index 000000000..c642ec75c --- /dev/null +++ b/src/web/src/contexts/agent-chat-sheet-state.ts @@ -0,0 +1,54 @@ +export type AgentChatSheetMode = "main" | "branch"; + +export interface AgentChatSheetTarget { + agentId: string; + conversationId: string | null; + taskId: string | null; + messageId: string | null; +} + +export interface AgentChatSheetOpenOptions { + conversationId?: string | null; + taskId?: string | null; + messageId?: string | null; + mode?: AgentChatSheetMode; + returnTo?: AgentChatSheetTarget | null; +} + +export function normalizeSheetTarget( + agentId: string, + opts?: AgentChatSheetOpenOptions, +): AgentChatSheetTarget { + return { + agentId, + conversationId: opts?.conversationId ?? null, + taskId: opts?.taskId ?? null, + messageId: opts?.messageId ?? null, + }; +} + +export function capturePreviousMainTargetForBranch({ + isOpen, + currentMode, + currentTarget, + explicitReturnTarget, +}: { + isOpen: boolean; + currentMode: AgentChatSheetMode; + currentTarget: AgentChatSheetTarget | null; + explicitReturnTarget?: AgentChatSheetTarget | null; +}): AgentChatSheetTarget | null { + if (!isOpen || currentMode !== "main" || !currentTarget) return null; + + if (explicitReturnTarget) return explicitReturnTarget; + + if ( + !currentTarget.conversationId && + !currentTarget.taskId && + !currentTarget.messageId + ) { + return null; + } + + return currentTarget; +} diff --git a/src/web/src/lib/api.ts b/src/web/src/lib/api.ts index 1fb91f722..1b8689a42 100644 --- a/src/web/src/lib/api.ts +++ b/src/web/src/lib/api.ts @@ -7,6 +7,7 @@ import type { CalendarEvent, Channel, Conversation, + ConversationBranch, CreateIssueRequest, CreateAgentLinkRequest, CreateAgentRequest, @@ -251,6 +252,43 @@ export const getOrCreateAgentConversation = (agentId: string, workspaceId: strin body: JSON.stringify({ ...(channel ? { channel } : {}) }), }); +export interface CreateBranchResponse { + branch: ConversationBranch; + conversation: Conversation; +} + +export interface ListBranchesResponse { + branches: ConversationBranch[]; +} + +export interface BranchOriginResponse { + branch: ConversationBranch; + root_message: Message; +} + +export const listBranches = (conversationId: string, workspaceId: string) => + apiFetch( + `/api/conversations/${conversationId}/branches${wsQuery(workspaceId)}`, + ); + +export const createBranch = ( + conversationId: string, + rootMessageId: string, + workspaceId: string, +) => + apiFetch( + `/api/conversations/${conversationId}/branches${wsQuery(workspaceId)}`, + { + method: "POST", + body: JSON.stringify({ root_message_id: rootMessageId }), + }, + ); + +export const getBranchOrigin = (conversationId: string, workspaceId: string) => + apiFetch( + `/api/conversations/${conversationId}/branch-origin${wsQuery(workspaceId)}`, + ); + export interface PreviousConversation { id: string; created_at: string; diff --git a/src/web/src/lib/api/responses.ts b/src/web/src/lib/api/responses.ts index 028b7e5b8..1f3446fe4 100644 --- a/src/web/src/lib/api/responses.ts +++ b/src/web/src/lib/api/responses.ts @@ -130,6 +130,21 @@ export function conversationToResponse(c: any) { return resp; } +export function conversationBranchToResponse(b: any) { + return { + id: b.id, + workspace_id: b.workspaceId, + parent_conversation_id: b.parentConversationId, + branch_conversation_id: b.branchConversationId, + root_message_id: b.rootMessageId, + provider: b.provider, + fork_source_task_id: b.forkSourceTaskId ?? null, + fork_source_session_id: b.forkSourceSessionId ?? null, + created_by: b.createdBy, + created_at: formatTimestamp(b.createdAt), + }; +} + export function channelToResponse(c: any) { return { id: c.id,