From eac554014a6f86271d0ce544ec148225aefe3274 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Samy=20Pess=C3=A9?= Date: Tue, 10 Jun 2025 12:08:19 +0200 Subject: [PATCH] Continue --- packages/gitbook-v2/src/lib/data/api.ts | 2 + packages/gitbook-v2/src/lib/data/types.ts | 2 + .../gitbook/src/components/Ask/AskDialog.tsx | 5 +- .../gitbook/src/components/Ask/AskInput.tsx | 1 + .../src/components/Ask/AskMessages.tsx | 15 ++ .../Ask/server-actions/AIMessageView.tsx | 31 +++ .../src/components/Ask/server-actions/api.tsx | 214 ++++++++++++++++++ .../src/components/Ask/server-actions/ask.ts | 44 ++++ .../components/Ask/server-actions/index.ts | 2 + packages/gitbook/src/components/Ask/state.tsx | 71 +++++- 10 files changed, 383 insertions(+), 4 deletions(-) create mode 100644 packages/gitbook/src/components/Ask/AskMessages.tsx create mode 100644 packages/gitbook/src/components/Ask/server-actions/AIMessageView.tsx create mode 100644 packages/gitbook/src/components/Ask/server-actions/api.tsx create mode 100644 packages/gitbook/src/components/Ask/server-actions/ask.ts create mode 100644 packages/gitbook/src/components/Ask/server-actions/index.ts diff --git a/packages/gitbook-v2/src/lib/data/api.ts b/packages/gitbook-v2/src/lib/data/api.ts index 44313bfc5..932973484 100644 --- a/packages/gitbook-v2/src/lib/data/api.ts +++ b/packages/gitbook-v2/src/lib/data/api.ts @@ -719,6 +719,8 @@ async function* streamAIResponse( input: params.input, output: params.output, model: params.model, + instructions: params.instructions, + previousResponseId: params.previousResponseId, }, { ...noCacheFetchOptions, diff --git a/packages/gitbook-v2/src/lib/data/types.ts b/packages/gitbook-v2/src/lib/data/types.ts index 178a0ba77..86d6b4130 100644 --- a/packages/gitbook-v2/src/lib/data/types.ts +++ b/packages/gitbook-v2/src/lib/data/types.ts @@ -189,5 +189,7 @@ export interface GitBookDataFetcher { input: api.AIMessageInput[]; output: api.AIOutputFormat; model: api.AIModel; + instructions?: string; + previousResponseId?: string; }): AsyncGenerator; } diff --git a/packages/gitbook/src/components/Ask/AskDialog.tsx b/packages/gitbook/src/components/Ask/AskDialog.tsx index e31256529..631902434 100644 --- a/packages/gitbook/src/components/Ask/AskDialog.tsx +++ b/packages/gitbook/src/components/Ask/AskDialog.tsx @@ -3,6 +3,7 @@ import { tcls } from '@/lib/tailwind'; import { Button } from '../primitives'; import { AskInput } from './AskInput'; +import { AskMessages } from './AskMessages'; import { useAskController, useAskState } from './state'; export function AskDialog() { @@ -51,7 +52,9 @@ export function AskDialog() { /> -
+
+ +
diff --git a/packages/gitbook/src/components/Ask/AskInput.tsx b/packages/gitbook/src/components/Ask/AskInput.tsx index e29949247..bf7b478f8 100644 --- a/packages/gitbook/src/components/Ask/AskInput.tsx +++ b/packages/gitbook/src/components/Ask/AskInput.tsx @@ -20,6 +20,7 @@ export function AskInput() { controller.postMessage({ message: value, }); + setValue(''); } }} /> diff --git a/packages/gitbook/src/components/Ask/AskMessages.tsx b/packages/gitbook/src/components/Ask/AskMessages.tsx new file mode 100644 index 000000000..296b8dfb5 --- /dev/null +++ b/packages/gitbook/src/components/Ask/AskMessages.tsx @@ -0,0 +1,15 @@ +import type { AskSession } from './state'; + +export function AskMessages(props: { + session: AskSession; +}) { + const { session } = props; + + return ( +
+ {session.messages.map((message, index) => { + return
{message.content}
; + })} +
+ ); +} diff --git a/packages/gitbook/src/components/Ask/server-actions/AIMessageView.tsx b/packages/gitbook/src/components/Ask/server-actions/AIMessageView.tsx new file mode 100644 index 000000000..42479fbc5 --- /dev/null +++ b/packages/gitbook/src/components/Ask/server-actions/AIMessageView.tsx @@ -0,0 +1,31 @@ +import type { AIMessage } from '@gitbook/api'; +import { DocumentView } from '../../DocumentView'; + +/** + * Render a message from the API backend. + */ +export function AIMessageView(props: { + message: AIMessage; +}) { + const { message } = props; + + return ( +
+ {message.steps.map((step, index) => { + return ( +
+ +
+ ); + })} +
+ ); +} diff --git a/packages/gitbook/src/components/Ask/server-actions/api.tsx b/packages/gitbook/src/components/Ask/server-actions/api.tsx new file mode 100644 index 000000000..d500d4c49 --- /dev/null +++ b/packages/gitbook/src/components/Ask/server-actions/api.tsx @@ -0,0 +1,214 @@ +'use server'; +import { + type AIMessage, + AIMessageRole, + type AIMessageStep, + type AIStreamResponse, +} from '@gitbook/api'; +import type { GitBookBaseContext } from '@v2/lib/context'; +import type { GitBookDataFetcher } from '@v2/lib/data'; +import { EventIterator } from 'event-iterator'; +import type { MaybePromise } from 'p-map'; +import * as partialJson from 'partial-json'; +import type { DeepPartial } from 'ts-essentials'; +import type { z } from 'zod'; +import { zodToJsonSchema } from 'zod-to-json-schema'; +import { AIMessageView } from './AIMessageView'; + +/** + * Get the latest value from a stream and the response id. + */ +export async function generate( + promise: MaybePromise<{ + stream: EventIterator; + response: Promise<{ responseId: string }>; + }> +) { + const input = await promise; + let value: T | undefined; + + for await (const event of input.stream) { + value = event; + } + + const { responseId } = await input.response; + return { + responseId, + value, + }; +} + +/** + * Stream the generation of an object using the AI. + */ +export async function streamGenerateObject( + context: GitBookBaseContext, + { + schema, + ...input + }: Omit[0], 'output'> & { + schema: z.ZodSchema; + } +) { + const rawStream = context.dataFetcher.streamAIResponse({ + ...input, + output: { + type: 'object', + schema: zodToJsonSchema(schema), + }, + }); + + let json = ''; + return parseResponse>(rawStream, (event) => { + if (event.type === 'response_object') { + json += event.jsonChunk; + + const parsed = partialJson.parse(json, partialJson.ALL); + return parsed; + } + }); +} + +/** + * Stream the generation of a document. + */ +export async function streamGenerateDocument( + context: GitBookBaseContext, + input: Omit[0], 'output'> +) { + const rawStream = context.dataFetcher.streamAIResponse({ + ...input, + output: { + type: 'document', + }, + }); + + const message: AIMessage = { + id: '', + role: AIMessageRole.Assistant, + steps: [], + }; + + const updateProcessingMessageStep = ( + stepIndex: number, + callback: (step: AIMessageStep) => void + ) => { + if (stepIndex > message.steps.length) { + throw new Error( + `Step index out of bounds ${stepIndex} (${message.steps.length} steps)` + ); + } + + if (message.steps[stepIndex]) { + message.steps = [...message.steps]; + message.steps[stepIndex] = { ...message.steps[stepIndex] }; + callback(message.steps[stepIndex]); + } else { + message.steps = [ + ...message.steps, + { + content: { + object: 'document', + data: {}, + nodes: [], + }, + }, + ]; + callback(message.steps[stepIndex]); + } + }; + + return parseResponse(rawStream, (event) => { + switch (event.type) { + /** + * The agent is processing a tool call in a new message. + */ + case 'response_tool_call': { + updateProcessingMessageStep(event.stepIndex, (step) => { + step.toolCalls ??= []; + step.toolCalls.push(event.toolCall); + }); + break; + } + + /** + * The agent is writing the content of a new message. + */ + case 'response_reasoning': + case 'response_document': { + updateProcessingMessageStep(event.stepIndex, (step) => { + const container = event.type === 'response_reasoning' ? 'reasoning' : 'content'; + + step[container] ??= { + object: 'document', + data: {}, + nodes: [], + }; + step[container] = { + ...step[container], + nodes: [...step[container].nodes], + }; + if (event.operation === 'insert') { + step[container].nodes.push(...event.blocks); + } else { + step[container].nodes.splice( + -event.blocks.length, + event.blocks.length, + ...event.blocks + ); + } + }); + break; + } + } + + return ; + }); +} + +/** + * Parse a stream from the API to extract the responseId. + */ +function parseResponse( + responseStream: EventIterator, + parse: (response: AIStreamResponse) => T | undefined +): { + stream: EventIterator; + response: Promise<{ responseId: string }>; +} { + let resolveResponse: (value: { responseId: string }) => void; + const response = new Promise<{ responseId: string }>((resolve) => { + resolveResponse = resolve; + }); + + const stream = new EventIterator((queue) => { + (async () => { + let foundResponse = false; + + for await (const event of responseStream) { + if (event.type === 'response_finish') { + foundResponse = true; + resolveResponse({ responseId: event.responseId }); + } else { + const parsed = parse(event); + if (parsed !== undefined) { + queue.push(parsed); + } + } + } + + if (!foundResponse) { + throw new Error('No response found'); + } + })().then( + () => { + queue.stop(); + }, + (error) => { + queue.fail(error); + } + ); + }); + + return { stream, response }; +} diff --git a/packages/gitbook/src/components/Ask/server-actions/ask.ts b/packages/gitbook/src/components/Ask/server-actions/ask.ts new file mode 100644 index 000000000..02ad705eb --- /dev/null +++ b/packages/gitbook/src/components/Ask/server-actions/ask.ts @@ -0,0 +1,44 @@ +'use server'; +import { getV1BaseContext } from '@/lib/v1'; +import { isV2 } from '@/lib/v2'; +import { AIMessageRole, AIModel } from '@gitbook/api'; +import { getSiteURLDataFromMiddleware } from '@v2/lib/middleware'; +import { getServerActionBaseContext } from '@v2/lib/server-actions'; +import { streamGenerateDocument } from './api'; + +const PROMPT = ` +You are a helpful assistant that can answer questions about the content of the site. +`; + +export async function* streamAsk({ + query, + previousResponseId, +}: { + query: string; + previousResponseId?: string; +}) { + const baseContext = isV2() ? await getServerActionBaseContext() : await getV1BaseContext(); + const siteURLData = await getSiteURLDataFromMiddleware(); + + const { response, stream } = await streamGenerateDocument(baseContext, { + organizationId: siteURLData.organization, + siteId: siteURLData.site, + model: AIModel.Fast, + instructions: PROMPT, + previousResponseId, + input: [ + { + role: AIMessageRole.User, + content: query, + }, + ], + }); + + for await (const output of stream) { + yield { output }; + } + + // Wait for the responseId to be available and yield one final time + const { responseId } = await response; + yield { responseId }; +} diff --git a/packages/gitbook/src/components/Ask/server-actions/index.ts b/packages/gitbook/src/components/Ask/server-actions/index.ts new file mode 100644 index 000000000..ff3414980 --- /dev/null +++ b/packages/gitbook/src/components/Ask/server-actions/index.ts @@ -0,0 +1,2 @@ +export * from './api'; +export * from './ask'; diff --git a/packages/gitbook/src/components/Ask/state.tsx b/packages/gitbook/src/components/Ask/state.tsx index a7302b01e..d693d549d 100644 --- a/packages/gitbook/src/components/Ask/state.tsx +++ b/packages/gitbook/src/components/Ask/state.tsx @@ -1,7 +1,8 @@ 'use client'; -import type { AIMessageRole } from '@gitbook/api'; +import { AIMessageRole } from '@gitbook/api'; import * as React from 'react'; +import { streamAsk } from './server-actions'; export type AskMessage = { role: AIMessageRole; @@ -68,6 +69,9 @@ export function AskStateProvider(props: React.PropsWithChildren) { }, }); + const stateRef = React.useRef(state); + stateRef.current = state; + const controller = React.useMemo(() => { return { open: () => { @@ -86,8 +90,69 @@ export function AskStateProvider(props: React.PropsWithChildren) { }; }); }, - postMessage: (input: { title?: string; message: string }) => { - // TODO + postMessage: async (input: { title?: string; message: string }) => { + try { + const stream = await streamAsk({ + query: input.message, + previousResponseId: stateRef.current.session.responseId ?? undefined, + }); + + setState((previous) => { + return { + ...previous, + session: { + ...previous.session, + messages: [ + ...previous.session.messages, + { + role: AIMessageRole.User, + content: input.message, + }, + { + role: AIMessageRole.Assistant, + content: null, + }, + ], + }, + }; + }); + + for await (const data of stream) { + if (!data) continue; + + if (data.responseId) { + setState((previous) => { + return { + ...previous, + session: { + ...previous.session, + responseId: data.responseId, + }, + }; + }); + } + + if (data.output) { + setState((previous) => { + return { + ...previous, + session: { + ...previous.session, + messages: [ + ...previous.session.messages.slice(0, -1), + { + role: AIMessageRole.Assistant, + content: data.output, + }, + ], + }, + }; + }); + } + } + } catch (error) { + console.error('Error in summary stream:', error); + } }, }; }, []);