diff --git a/.env.example b/.env.example index d33a6400..0d3c7340 100644 --- a/.env.example +++ b/.env.example @@ -1,5 +1,6 @@ # Application NEXT_PUBLIC_APP_URL=http://localhost:3000 +NEXT_PUBLIC_USE_NEW_CHAT=false # Better Auth (Authentication) BETTER_AUTH_SECRET=your-better-auth-secret diff --git a/package.json b/package.json index 92ea88fa..c4afca71 100644 --- a/package.json +++ b/package.json @@ -8,6 +8,7 @@ "start": "next start", "lint": "oxlint", "lint:fix": "oxlint --fix", + "typecheck": "pnpm tc", "tc": "tsgo --noEmit", "tc:tsc": "tsc --noEmit", "clean": "rm -rf .next out dist", @@ -32,6 +33,7 @@ "dependencies": { "@ai-sdk/devtools": "^0.0.15", "@ai-sdk/google": "^3.0.62", + "@ai-sdk/react": "^3.0.170", "@assistant-ui/react": "^0.12.24", "@assistant-ui/react-ai-sdk": "^1.3.18", "@assistant-ui/react-devtools": "^1.0.5", @@ -163,6 +165,7 @@ "tailwind-merge": "^3.5.0", "tw-shimmer": "^0.4.9", "unicodeit": "^0.7.5", + "virtua": "^0.49.1", "workflow": "4.2.1", "zod": "^4.3.6", "zustand": "^5.0.11" diff --git a/src/app/api/chat-v2/edit/route.ts b/src/app/api/chat-v2/edit/route.ts new file mode 100644 index 00000000..34614acb --- /dev/null +++ b/src/app/api/chat-v2/edit/route.ts @@ -0,0 +1,64 @@ +import { NextRequest, NextResponse } from "next/server"; +import { and, eq } from "drizzle-orm"; +import { withServerObservability } from "@/lib/with-server-observability"; +import { db } from "@/lib/db/client"; +import { chatMessages, chatThreads } from "@/lib/db/schema"; +import { requireAuth, verifyThreadOwnership, verifyWorkspaceAccess } from "@/lib/api/workspace-helpers"; +import type { UIMessage } from "ai"; + +export const POST = withServerObservability(async function POST(req: NextRequest) { + try { + const userId = await requireAuth(); + const body = await req.json().catch(() => ({})) as { threadId?: string; messageId?: string; text?: string }; + if (!body.threadId || !body.messageId || typeof body.text !== "string") { + return NextResponse.json({ error: "threadId, messageId, and text are required" }, { status: 400 }); + } + + const [thread] = await db.select().from(chatThreads).where(eq(chatThreads.id, body.threadId)).limit(1); + if (!thread) return NextResponse.json({ error: "Thread not found" }, { status: 404 }); + await verifyWorkspaceAccess(thread.workspaceId, userId); + verifyThreadOwnership(thread, userId); + + const [stored] = await db.select().from(chatMessages).where(and(eq(chatMessages.threadId, body.threadId), eq(chatMessages.messageId, body.messageId))).limit(1); + if (!stored) return NextResponse.json({ error: "Message not found" }, { status: 404 }); + + const content = stored.content as UIMessage; + if (content.role !== "user") { + return NextResponse.json({ error: "Only user messages can be edited" }, { status: 400 }); + } + + const newMessageId = crypto.randomUUID(); + const nextMessage: UIMessage = { + ...content, + id: newMessageId, + parts: ( + content.parts.some((part) => part.type === "text") + ? content.parts.map((part) => + part.type === "text" + ? { ...part, text: body.text } + : part, + ) + : [...content.parts, { type: "text", text: body.text }] + ) as UIMessage["parts"], + }; + + await db.insert(chatMessages).values({ + threadId: body.threadId, + messageId: newMessageId, + parentId: stored.parentId, + format: "ai-sdk/v6", + content: nextMessage, + }); + + await db.update(chatThreads).set({ + headMessageId: newMessageId, + updatedAt: new Date().toISOString(), + lastMessageAt: new Date().toISOString(), + }).where(eq(chatThreads.id, body.threadId)); + + return NextResponse.json({ newMessageId }); + } catch (error) { + if (error instanceof Response) return error; + return NextResponse.json({ error: "Internal server error" }, { status: 500 }); + } +}, { routeName: "POST /api/chat-v2/edit" }); diff --git a/src/app/api/chat-v2/route.ts b/src/app/api/chat-v2/route.ts new file mode 100644 index 00000000..eb5ef659 --- /dev/null +++ b/src/app/api/chat-v2/route.ts @@ -0,0 +1,317 @@ +import { + convertToModelMessages, + createUIMessageStream, + createUIMessageStreamResponse, + pruneMessages, + safeValidateUIMessages, + smoothStream, + stepCountIs, + streamText, + wrapLanguageModel, + type UIMessage, +} from "ai"; +import { devToolsMiddleware } from "@ai-sdk/devtools"; +import { withTracing } from "@posthog/ai"; +import { headers } from "next/headers"; +import { and, eq } from "drizzle-orm"; +import { auth } from "@/lib/auth"; +import { createChatTools } from "@/lib/ai/tools"; +import { getDefaultChatModelId, resolveGatewayModelId } from "@/lib/ai/models"; +import { normalizeLegacyToolMessages } from "@/lib/ai/legacy-tool-message-compat"; +import { maybeWithSupermemory } from "@/lib/ai/supermemory"; +import { buildGatewayProviderOptions, createGatewayLanguageModel, getGatewayAttributionHeaders } from "@/lib/ai/gateway-provider-options"; +import { capturePostHogServerException, getPostHogServerClient } from "@/lib/posthog-server"; +import { withServerObservability } from "@/lib/with-server-observability"; +import { db } from "@/lib/db/client"; +import { chatMessages, chatThreads } from "@/lib/db/schema"; +import { logger } from "@/lib/utils/logger"; +import type { ReplySelection } from "@/lib/stores/ui-store"; +import { verifyThreadOwnership, verifyWorkspaceAccess } from "@/lib/api/workspace-helpers"; +import { getNewFinishedMessages, resolveInitialParentId } from "@/lib/chat-v2/stream-persistence"; + +function getSelectedCardsContext(body: Record): string { + return typeof body.selectedCardsContext === "string" ? body.selectedCardsContext : ""; +} + +function injectSelectionContext( + messages: Awaited>, + metadata?: { replySelections?: ReplySelection[] }, + selectedCardsContext?: string, +): void { + const parts: string[] = []; + if (selectedCardsContext?.trim()) { + parts.push(`[Selected cards context:\n${selectedCardsContext.trim()}]`); + } + if (metadata?.replySelections?.length) { + const quoted = metadata.replySelections + .map((selection) => selection.title ? `> From: ${selection.title}\n> ${selection.text}` : `> ${selection.text}`) + .join("\n\n"); + parts.push(`[Referring to:\n${quoted}]`); + } + if (parts.length === 0) return; + const prefix = `${parts.join("\n")}\n\n`; + + for (let index = messages.length - 1; index >= 0; index -= 1) { + const message = messages[index]; + if (message.role !== "user") continue; + for (const part of message.content) { + if (typeof part !== "string" && part.type === "text") { + part.text = prefix + part.text; + return; + } + } + } +} + +async function saveMessage({ + threadId, + message, + parentId, +}: { + threadId: string; + message: UIMessage; + parentId: string | null; +}) { + await db.insert(chatMessages).values({ + threadId, + messageId: message.id, + parentId, + format: "ai-sdk/v6", + content: message, + }).onConflictDoUpdate({ + target: [chatMessages.threadId, chatMessages.messageId], + set: { + parentId, + format: "ai-sdk/v6", + content: message, + }, + }); +} + +async function updateThreadHead(threadId: string, headMessageId: string | null) { + await db.update(chatThreads).set({ + headMessageId, + lastMessageAt: new Date().toISOString(), + updatedAt: new Date().toISOString(), + }).where(eq(chatThreads.id, threadId)); +} + +async function handlePOST(req: Request) { + let workspaceId: string | null = null; + let userId: string | null = null; + + try { + const [headersObj, body] = await Promise.all([headers(), req.json() as Promise>]); + const session = await auth.api.getSession({ headers: headersObj }); + userId = session?.user?.id ?? null; + + const threadId = typeof body.id === "string" ? body.id : null; + const trigger = + body.trigger === "regenerate-message" + ? "regenerate-message" + : "submit-message"; + const regenerateMessageId = + typeof body.messageId === "string" ? body.messageId : null; + const activeFolderId = typeof body.activeFolderId === "string" ? body.activeFolderId : undefined; + const memoryEnabled = body.memoryEnabled === true; + const system = typeof body.system === "string" ? body.system : ""; + const selectedCardsContext = getSelectedCardsContext(body); + const messages = Array.isArray(body.messages) ? (body.messages as UIMessage[]) : []; + + if (!threadId) { + return new Response(JSON.stringify({ error: "Thread id is required" }), { status: 400, headers: { "Content-Type": "application/json" } }); + } + + const [thread] = await db.select().from(chatThreads).where(eq(chatThreads.id, threadId)).limit(1); + if (!thread) { + return new Response(JSON.stringify({ error: "Thread not found" }), { status: 404, headers: { "Content-Type": "application/json" } }); + } + if (!userId) { + return new Response(JSON.stringify({ error: "Unauthorized" }), { status: 401, headers: { "Content-Type": "application/json" } }); + } + await verifyWorkspaceAccess(thread.workspaceId, userId); + verifyThreadOwnership(thread, userId); + workspaceId = thread.workspaceId; + + const tools = createChatTools({ + workspaceId, + userId, + activeFolderId, + threadId, + }); + + const compatibleMessages = normalizeLegacyToolMessages(messages, { + availableToolNames: Object.keys(tools), + }); + const validation = await safeValidateUIMessages({ messages: compatibleMessages, tools }); + if (!validation.success) { + throw validation.error; + } + + const validatedMessages = validation.data; + const lastMessage = validatedMessages.at(-1); + const previousMessageId = validatedMessages.at(-2)?.id ?? null; + if (trigger === "submit-message" && lastMessage?.role === "user") { + await saveMessage({ threadId, message: lastMessage, parentId: previousMessageId }); + await updateThreadHead(threadId, lastMessage.id); + } + + const initialParentId = await resolveInitialParentId({ + trigger, + regenerateMessageId, + lastMessage, + getStoredMessageParentId: async (messageId) => { + const [siblingRow] = await db + .select({ parentId: chatMessages.parentId }) + .from(chatMessages) + .where(and(eq(chatMessages.threadId, threadId), eq(chatMessages.messageId, messageId))) + .limit(1); + return siblingRow?.parentId ?? null; + }, + }); + + let convertedMessages = await convertToModelMessages(validatedMessages, { tools }); + convertedMessages = pruneMessages({ + messages: convertedMessages, + reasoning: "before-last-message", + toolCalls: "before-last-5-messages", + emptyMessages: "remove", + }); + const lastMetadata = + lastMessage?.metadata && + typeof lastMessage.metadata === "object" && + "replySelections" in lastMessage.metadata + ? (lastMessage.metadata as { replySelections?: ReplySelection[] }) + : undefined; + injectSelectionContext(convertedMessages, lastMetadata, selectedCardsContext); + + const modelId = resolveGatewayModelId(typeof body.modelId === "string" ? body.modelId : getDefaultChatModelId()); + const posthogClient = getPostHogServerClient(); + const baseGatewayModel = createGatewayLanguageModel(modelId); + const tracedModel = posthogClient + ? withTracing(baseGatewayModel, posthogClient, { + posthogDistinctId: userId || "anonymous", + posthogProperties: { + workspaceId, + activeFolderId, + modelId, + memoryEnabled, + }, + }) + : baseGatewayModel; + + const memoryWrappedModel = maybeWithSupermemory(tracedModel, { + userId: userId ?? "", + threadId, + memoryEnabled, + }); + + const model = wrapLanguageModel({ + model: memoryWrappedModel, + middleware: process.env.NODE_ENV === "development" ? devToolsMiddleware() : [], + }); + + const providerOptions = buildGatewayProviderOptions(modelId, { userId }); + + const result = streamText({ + model, + temperature: 1, + system, + messages: convertedMessages, + stopWhen: stepCountIs(25), + tools, + providerOptions: providerOptions as Parameters[0]["providerOptions"], + headers: getGatewayAttributionHeaders(), + experimental_telemetry: { + isEnabled: true, + metadata: { + "tcc.conversational": "true", + ...(threadId ? { "tcc.sessionId": String(threadId) } : {}), + ...(userId ? { userId } : {}), + }, + }, + experimental_transform: smoothStream({ chunking: "word", delayInMs: 15 }), + onFinish: ({ usage, finishReason }) => { + logger.info("📊 [CHAT-V2] Final Token Usage:", { + inputTokens: usage?.inputTokens, + outputTokens: usage?.outputTokens, + totalTokens: usage?.totalTokens, + cachedInputTokens: usage?.cachedInputTokens, + reasoningTokens: usage?.reasoningTokens, + finishReason, + }); + }, + }); + + const stream = createUIMessageStream({ + originalMessages: validatedMessages, + execute: ({ writer }) => { + // TODO: emit data-chat-title here if chat-v2 gets a shared title-generation utility. + writer.merge(result.toUIMessageStream({ + sendReasoning: true, + sendSources: true, + })); + }, + onFinish: async ({ messages: finishedMessages }) => { + let parentId = initialParentId; + const newMessages = getNewFinishedMessages({ + finishedMessages, + validatedMessages, + }); + for (const message of newMessages) { + await saveMessage({ threadId, message, parentId }); + parentId = message.id; + } + await updateThreadHead(threadId, parentId); + }, + }); + + return createUIMessageStreamResponse({ stream }); + } catch (error) { + if (error instanceof Response) { + return error; + } + + const errorMessage = error instanceof Error ? error.message : String(error); + const isTimeout = + errorMessage.includes("timeout") || + errorMessage.includes("TIMEOUT") || + errorMessage.includes("Function execution exceeded") || + errorMessage.includes("Execution timeout") || + (error && typeof error === "object" && "code" in error && error.code === "TIMEOUT"); + + if (isTimeout) { + capturePostHogServerException(error, { + distinctId: userId ?? undefined, + properties: { + route_name: "POST /api/chat-v2", + workspaceId: workspaceId ?? undefined, + chat_error_kind: "timeout", + }, + }); + return new Response(JSON.stringify({ + error: "Request timeout", + message: "The request took too long to process.", + code: "TIMEOUT", + }), { status: 504, headers: { "Content-Type": "application/json" } }); + } + + capturePostHogServerException(error, { + distinctId: userId ?? undefined, + properties: { + route_name: "POST /api/chat-v2", + workspaceId: workspaceId ?? undefined, + chat_error_kind: "internal", + }, + }); + + return new Response(JSON.stringify({ + error: "Internal server error", + message: "An unexpected error occurred while processing your request.", + details: process.env.NODE_ENV === "development" ? errorMessage : undefined, + code: "INTERNAL_ERROR", + }), { status: 500, headers: { "Content-Type": "application/json" } }); + } +} + +export const POST = withServerObservability(handlePOST, { routeName: "POST /api/chat-v2" }); diff --git a/src/app/api/threads/[id]/messages/[messageId]/branches/route.ts b/src/app/api/threads/[id]/messages/[messageId]/branches/route.ts new file mode 100644 index 00000000..84888408 --- /dev/null +++ b/src/app/api/threads/[id]/messages/[messageId]/branches/route.ts @@ -0,0 +1,88 @@ +import { NextRequest, NextResponse } from "next/server"; +import { and, asc, eq, isNull } from "drizzle-orm"; +import { db } from "@/lib/db/client"; +import { chatMessages, chatThreads } from "@/lib/db/schema"; +import { requireAuth, verifyThreadOwnership, verifyWorkspaceAccess } from "@/lib/api/workspace-helpers"; +import { withServerObservability } from "@/lib/with-server-observability"; + +async function getThreadAndMessage(threadId: string, messageId: string, userId: string) { + const [thread] = await db.select().from(chatThreads).where(eq(chatThreads.id, threadId)).limit(1); + if (!thread) throw NextResponse.json({ error: "Thread not found" }, { status: 404 }); + await verifyWorkspaceAccess(thread.workspaceId, userId); + verifyThreadOwnership(thread, userId); + + const [message] = await db.select().from(chatMessages).where(and(eq(chatMessages.threadId, threadId), eq(chatMessages.messageId, messageId))).limit(1); + if (!message) throw NextResponse.json({ error: "Message not found" }, { status: 404 }); + + return { thread, message }; +} + +async function findLeaf(threadId: string, messageId: string): Promise { + let currentId = messageId; + while (true) { + const [child] = await db.select({ messageId: chatMessages.messageId }).from(chatMessages).where(and(eq(chatMessages.threadId, threadId), eq(chatMessages.parentId, currentId))).orderBy(asc(chatMessages.createdAt)).limit(1); + if (!child) return currentId; + currentId = child.messageId; + } +} + +export const GET = withServerObservability(async function GET(_req: NextRequest, { params }: { params: Promise<{ id: string; messageId: string }> }) { + try { + const userId = await requireAuth(); + const { id, messageId } = await params; + const { message } = await getThreadAndMessage(id, messageId, userId); + const siblings = await db + .select({ id: chatMessages.messageId, parentId: chatMessages.parentId, createdAt: chatMessages.createdAt }) + .from(chatMessages) + .where( + and( + eq(chatMessages.threadId, id), + message.parentId == null + ? isNull(chatMessages.parentId) + : eq(chatMessages.parentId, message.parentId), + ), + ) + .orderBy(asc(chatMessages.createdAt)); + return NextResponse.json({ + siblings, + currentIndex: siblings.findIndex((candidate) => candidate.id === messageId), + }); + } catch (error) { + if (error instanceof Response) return error; + return NextResponse.json({ error: "Internal server error" }, { status: 500 }); + } +}, { routeName: "GET /api/threads/[id]/messages/[messageId]/branches" }); + +export const POST = withServerObservability(async function POST(req: NextRequest, { params }: { params: Promise<{ id: string; messageId: string }> }) { + try { + const userId = await requireAuth(); + const { id, messageId } = await params; + const body = await req.json().catch(() => ({})) as { targetBranchId?: string }; + if (!body.targetBranchId) { + return NextResponse.json({ error: "targetBranchId is required" }, { status: 400 }); + } + + const { message } = await getThreadAndMessage(id, messageId, userId); + const [target] = await db + .select() + .from(chatMessages) + .where(and(eq(chatMessages.threadId, id), eq(chatMessages.messageId, body.targetBranchId))) + .limit(1); + if (!target) { + return NextResponse.json({ error: "Target branch not found" }, { status: 404 }); + } + if (target.parentId !== message.parentId) { + return NextResponse.json( + { error: "Target is not a sibling of the anchor message" }, + { status: 400 }, + ); + } + + const leafId = await findLeaf(id, body.targetBranchId); + await db.update(chatThreads).set({ headMessageId: leafId, updatedAt: new Date().toISOString() }).where(eq(chatThreads.id, id)); + return NextResponse.json({ headMessageId: leafId }); + } catch (error) { + if (error instanceof Response) return error; + return NextResponse.json({ error: "Internal server error" }, { status: 500 }); + } +}, { routeName: "POST /api/threads/[id]/messages/[messageId]/branches" }); diff --git a/src/app/api/threads/[id]/messages/route.ts b/src/app/api/threads/[id]/messages/route.ts index cf9a2113..83cb5dad 100644 --- a/src/app/api/threads/[id]/messages/route.ts +++ b/src/app/api/threads/[id]/messages/route.ts @@ -6,7 +6,7 @@ import { verifyWorkspaceAccess, verifyThreadOwnership, } from "@/lib/api/workspace-helpers"; -import { eq, and, desc } from "drizzle-orm"; +import { eq, and } from "drizzle-orm"; import { withServerObservability } from "@/lib/with-server-observability"; async function getThreadAndVerify(id: string, userId: string) { @@ -45,16 +45,41 @@ export const GET = withServerObservability(async function GET( const rows = await db .select() .from(chatMessages) - .where(and(eq(chatMessages.threadId, id), eq(chatMessages.format, format))) - .orderBy(desc(chatMessages.createdAt)); - - const messages = rows.map((r) => ({ - id: r.messageId, - parent_id: r.parentId, - format: r.format, - content: r.content, - created_at: r.createdAt, - })); + .where(and(eq(chatMessages.threadId, id), eq(chatMessages.format, format))); + + const byId = new Map(rows.map((row) => [row.messageId, row])); + const ordered: typeof rows = []; + + if (thread.headMessageId) { + let cursor: string | null = thread.headMessageId; + const seen = new Set(); + + while (cursor && !seen.has(cursor)) { + seen.add(cursor); + const row = byId.get(cursor); + if (!row) break; + ordered.push(row); + cursor = row.parentId; + } + + ordered.reverse(); + } else { + ordered.push( + ...rows.sort((a, b) => { + const at = a.createdAt ? new Date(a.createdAt).getTime() : 0; + const bt = b.createdAt ? new Date(b.createdAt).getTime() : 0; + return at - bt; + }), + ); + } + + const messages = ordered.map((r) => ({ + id: r.messageId, + parent_id: r.parentId, + format: r.format, + content: r.content, + created_at: r.createdAt, + })); return NextResponse.json({ messages, diff --git a/src/components/assistant-ui/AssistantPanel.tsx b/src/components/assistant-ui/AssistantPanel.tsx index deacaf6d..434bed2e 100644 --- a/src/components/assistant-ui/AssistantPanel.tsx +++ b/src/components/assistant-ui/AssistantPanel.tsx @@ -9,6 +9,8 @@ import AppChatHeader from "@/components/chat/AppChatHeader"; import { cn } from "@/lib/utils"; import AssistantTextSelectionManager from "@/components/assistant-ui/AssistantTextSelectionManager"; import { useEffect } from "react"; +import { USE_NEW_CHAT } from "@/lib/chat-v2/feature-flag"; +import { ChatV2Panel } from "@/components/chat-v2/ChatV2Panel"; interface AssistantPanelProps { workspaceId?: string | null; @@ -108,9 +110,22 @@ function WorkspaceContextWrapperContent({ // Workspace name comes from canonical workspace metadata. const { currentWorkspace } = useWorkspaceContext(); - // Inject minimal workspace context (metadata and system instructions only) - // Cards register their own context individually - useWorkspaceContextProvider(workspaceId || null, state, currentWorkspace?.name); + if (!USE_NEW_CHAT) { + useWorkspaceContextProvider(workspaceId || null, state, currentWorkspace?.name); + } + + if (USE_NEW_CHAT) { + return ( + + ); + } diff --git a/src/components/assistant-ui/WorkspaceRuntimeProvider.tsx b/src/components/assistant-ui/WorkspaceRuntimeProvider.tsx index dbaffa5b..45ea007e 100644 --- a/src/components/assistant-ui/WorkspaceRuntimeProvider.tsx +++ b/src/components/assistant-ui/WorkspaceRuntimeProvider.tsx @@ -24,43 +24,17 @@ import { chatToolToolkit } from "@/components/assistant-ui/chat-toolkit"; interface WorkspaceRuntimeProviderProps { workspaceId: string; + disableChatRuntime?: boolean; children: React.ReactNode; } -function createWorkspaceChatRuntimeHook( - transport: AssistantChatTransport, - onError: (error: Error) => void, -) { - return function useWorkspaceChatRuntimeHook() { - return useChatRuntime({ - transport, - onError, - toCreateMessage: toCreateMessageWithContext, - }); - }; -} - -/** - * Bridges the outer RemoteThreadListRuntime's initialize() into a ref - * so the transport's prepareSendMessagesRequest can eagerly resolve the - * real thread remoteId before each chat request. - * - * Workaround for https://github.com/assistant-ui/assistant-ui/issues/3578 - */ -function ThreadInitBridge({ - initRef, -}: { - initRef: React.MutableRefObject<(() => Promise<{ remoteId: string }>) | null>; -}) { - const aui = useAui(); - initRef.current = () => aui.threadListItem().initialize(); - return null; -} - -export function WorkspaceRuntimeProvider({ +function WorkspaceRuntimeWithChatRuntime({ workspaceId, children, -}: WorkspaceRuntimeProviderProps) { +}: { + workspaceId: string; + children: React.ReactNode; +}) { const selectedModelId = useUIStore((state) => state.selectedModelId); const memoryEnabled = useUIStore((state) => state.memoryEnabled); const activeFolderId = useUIStore((state) => state.activeFolderId); @@ -71,7 +45,6 @@ export function WorkspaceRuntimeProvider({ const { state: workspaceState } = useWorkspaceState(workspaceId); const viewingItemIds = useViewingItemIds(); - /** Union of selected cards and items open in the workspace viewer (primary and/or secondary). */ const contextCardIds = useMemo(() => { const ids = new Set(selectedCardIdsSet); viewingItemIds.forEach((id) => ids.add(id)); @@ -99,13 +72,6 @@ export function WorkspaceRuntimeProvider({ ); }, [workspaceState, contextCardIds, activePdfPageByItemId, viewingItemIds]); - // Per AI SDK, transport `body` is `Resolvable` — if it is a function, `resolve()` - // calls it on every sendMessages (see @ai-sdk/provider-utils resolve()). That gives - // fresh metadata without recreating AssistantChatTransport (which would change the - // transport passed into useChatRuntime and stress useRemoteThreadListRuntime). - // @assistant-ui/react-ai-sdk also wraps the transport in useDynamicChatTransport (Proxy - // + ref) so the chat layer can follow transport updates; keeping one transport instance - // is still the least surprising option for our custom runtimeHook wrapper. const chatApiPayloadRef = useRef({ workspaceId, modelId: selectedModelId, @@ -125,7 +91,6 @@ export function WorkspaceRuntimeProvider({ const handleChatError = useCallback((error: Error) => { console.error("[Chat Error]", error); - // Extract error message from various sources (error.message, responseBody, data, etc.) const errorMessage = error.message?.toLowerCase() || ""; const responseBody = (error as any).responseBody?.toLowerCase() || ""; const errorData = (error as any).data?.error?.message?.toLowerCase() || ""; @@ -184,7 +149,6 @@ export function WorkspaceRuntimeProvider({ "Please check your GOOGLE_GENERATIVE_AI_API_KEY in your environment variables.", }); } else { - // Generic error fallback toast.error("Something went wrong", { description: error.message || "An unexpected error occurred. Please try again.", @@ -231,8 +195,6 @@ export function WorkspaceRuntimeProvider({ }; }, }), - // Body snapshot comes from the ref via Resolvable function above. - // eslint-disable-next-line react-hooks/exhaustive-deps -- stable transport instance [], ); @@ -257,3 +219,50 @@ export function WorkspaceRuntimeProvider({ ); } + +function createWorkspaceChatRuntimeHook( + transport: AssistantChatTransport, + onError: (error: Error) => void, +) { + return function useWorkspaceChatRuntimeHook() { + return useChatRuntime({ + transport, + onError, + toCreateMessage: toCreateMessageWithContext, + }); + }; +} + +/** + * Bridges the outer RemoteThreadListRuntime's initialize() into a ref + * so the transport's prepareSendMessagesRequest can eagerly resolve the + * real thread remoteId before each chat request. + * + * Workaround for https://github.com/assistant-ui/assistant-ui/issues/3578 + */ +function ThreadInitBridge({ + initRef, +}: { + initRef: React.MutableRefObject<(() => Promise<{ remoteId: string }>) | null>; +}) { + const aui = useAui(); + initRef.current = () => aui.threadListItem().initialize(); + return null; +} + +export function WorkspaceRuntimeProvider({ + workspaceId, + disableChatRuntime = false, + children, +}: WorkspaceRuntimeProviderProps) { + if (disableChatRuntime) { + // Keep assistant availability context for the dashboard, but skip assistant-ui's thread/runtime wiring when chat-v2 is enabled. + return {children}; + } + + return ( + + {children} + + ); +} diff --git a/src/components/chat-v2/AssistantActionBar.tsx b/src/components/chat-v2/AssistantActionBar.tsx new file mode 100644 index 00000000..48d70da6 --- /dev/null +++ b/src/components/chat-v2/AssistantActionBar.tsx @@ -0,0 +1,35 @@ +"use client"; + +import { CheckIcon, CopyIcon, RefreshCwIcon } from "lucide-react"; +import { useEffect, useState } from "react"; +import { Button } from "@/components/ui/button"; + +interface AssistantActionBarProps { + textContent: string; + onRefresh: () => void; +} + +export function AssistantActionBar({ textContent, onRefresh }: AssistantActionBarProps) { + const [copied, setCopied] = useState(false); + + useEffect(() => { + if (!copied) return; + const timeout = window.setTimeout(() => setCopied(false), 2000); + return () => window.clearTimeout(timeout); + }, [copied]); + + return ( +
+ + +
+ ); +} diff --git a/src/components/chat-v2/AssistantMessage.tsx b/src/components/chat-v2/AssistantMessage.tsx new file mode 100644 index 00000000..79d9d036 --- /dev/null +++ b/src/components/chat-v2/AssistantMessage.tsx @@ -0,0 +1,70 @@ +"use client"; + +import { Loader2 } from "lucide-react"; +import { useMemo } from "react"; +import { AssistantActionBar } from "./AssistantActionBar"; +import { BranchNav } from "./BranchNav"; +import { TextPart } from "./parts/TextPart"; +import { ToolGroup } from "./parts/ToolGroup"; +import { ReasoningPart } from "./parts/ReasoningPart"; +import { groupParts } from "./parts/group-parts"; +import type { ChatMessage } from "@/lib/chat-v2/types"; + +interface AssistantMessageProps { + threadId?: string | null; + message: ChatMessage; + isStreaming: boolean; + blankSize?: number; + onRefresh: (messageId: string) => void; + onReloadThread: () => Promise; +} + +export function AssistantMessage({ threadId, message, isStreaming, blankSize, onRefresh, onReloadThread }: AssistantMessageProps) { + const grouped = useMemo(() => groupParts(message.parts), [message.parts]); + const textContent = useMemo( + () => message.parts.filter((part) => part.type === "text").map((part) => part.text).join("\n\n"), + [message.parts], + ); + + return ( +
+
+ {isStreaming && message.parts.length === 0 ? : null} + {grouped.map((segment, index) => { + if (segment.kind === "reasoning") { + return ( + + ); + } + if (segment.kind === "tools") { + return ; + } + if (segment.part.type === "text") { + return ; + } + if (segment.part.type === "file") { + return {segment.part.filename ?? segment.part.url}; + } + if (segment.part.type === "source-url") { + return {segment.part.title ?? segment.part.url}; + } + if (segment.part.type === "source-document") { + return
{segment.part.title}
; + } + return null; + })} +
+ +
+ + onRefresh(message.id)} /> +
+
+ ); +} diff --git a/src/components/chat-v2/BranchNav.tsx b/src/components/chat-v2/BranchNav.tsx new file mode 100644 index 00000000..b87d5550 --- /dev/null +++ b/src/components/chat-v2/BranchNav.tsx @@ -0,0 +1,77 @@ +"use client"; + +import { ChevronLeftIcon, ChevronRightIcon } from "lucide-react"; +import { useEffect, useState } from "react"; +import { Button } from "@/components/ui/button"; +import { fetchBranches, switchBranch } from "@/lib/chat-v2/branches"; +import { useChatRuntime } from "@/lib/chat-v2/use-chat-runtime"; + +interface BranchNavProps { + threadId?: string | null; + messageId: string; + className?: string; +} + +export function BranchNav({ threadId, messageId, className }: BranchNavProps) { + const [count, setCount] = useState(0); + const [currentIndex, setCurrentIndex] = useState(0); + const [siblings, setSiblings] = useState>([]); + const [loading, setLoading] = useState(false); + const { status, stop, refreshMessagesIfSafe } = useChatRuntime(); + const isBusy = status === "streaming" || status === "submitted"; + + useEffect(() => { + if (!threadId) return; + let cancelled = false; + void (async () => { + try { + const data = await fetchBranches(threadId, messageId); + if (cancelled) return; + setSiblings(data.siblings); + setCount(data.siblings.length); + setCurrentIndex(data.currentIndex); + } catch { + if (!cancelled) { + setSiblings([]); + setCount(0); + setCurrentIndex(0); + } + } + })(); + + return () => { + cancelled = true; + }; + }, [messageId, threadId]); + + if (!threadId || count <= 1) return null; + + const handleSwitch = async (delta: -1 | 1) => { + const nextIndex = (currentIndex + delta + siblings.length) % siblings.length; + const target = siblings[nextIndex]; + if (!target) return; + setLoading(true); + try { + if (isBusy) { + await stop(); + } + await switchBranch(threadId, messageId, target.id); + await refreshMessagesIfSafe(); + setCurrentIndex(nextIndex); + } finally { + setLoading(false); + } + }; + + return ( +
+ + {currentIndex + 1} / {count} + +
+ ); +} diff --git a/src/components/chat-v2/ChatHeader.tsx b/src/components/chat-v2/ChatHeader.tsx new file mode 100644 index 00000000..d9b9a790 --- /dev/null +++ b/src/components/chat-v2/ChatHeader.tsx @@ -0,0 +1,99 @@ +"use client"; + +import { ChevronDown, Check, Edit2, X } from "lucide-react"; +import { LuMaximize2, LuMinimize2, LuPanelRightClose } from "react-icons/lu"; +import { useEffect, useMemo, useState } from "react"; +import { toast } from "sonner"; +import { Button } from "@/components/ui/button"; +import { Textarea } from "@/components/ui/textarea"; +import { ThreadListDropdown } from "./ThreadListDropdown"; + +interface ChatHeaderProps { + workspaceId: string; + threadId?: string | null; + onSelectThread: (threadId: string | null) => void; + onCollapse?: () => void; + isMaximized?: boolean; + onToggleMaximize?: () => void; +} + +export function ChatHeader({ workspaceId, threadId, onSelectThread, onCollapse, isMaximized, onToggleMaximize }: ChatHeaderProps) { + const [threadTitle, setThreadTitle] = useState("New Chat"); + const [editing, setEditing] = useState(false); + const [draft, setDraft] = useState("New Chat"); + + useEffect(() => { + if (!threadId) { + setThreadTitle("New Chat"); + setDraft("New Chat"); + return; + } + + let cancelled = false; + + void (async () => { + try { + const response = await fetch(`/api/threads/${threadId}`, { cache: "no-store" }); + if (!response.ok || cancelled) return; + const data = (await response.json()) as { title?: string }; + if (cancelled) return; + setThreadTitle(data.title ?? "New Chat"); + setDraft(data.title ?? "New Chat"); + } catch { + return; + } + })(); + + return () => { + cancelled = true; + }; + }, [threadId]); + + const title = useMemo(() => threadTitle || "New Chat", [threadTitle]); + + const save = async () => { + if (!threadId) return; + const nextTitle = draft.trim() || "New Chat"; + const response = await fetch(`/api/threads/${threadId}`, { + method: "PATCH", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify({ title: nextTitle }), + }); + if (!response.ok) { + toast.error("Failed to rename chat"); + return; + } + setThreadTitle(nextTitle); + setEditing(false); + }; + + return ( +
+
+ {title}} + /> + {editing ? ( +
+