This commit is contained in:
Samy Pessé
2025-06-10 12:08:19 +02:00
parent fef62b12b4
commit eac554014a
10 changed files with 383 additions and 4 deletions
+2
View File
@@ -719,6 +719,8 @@ async function* streamAIResponse(
input: params.input,
output: params.output,
model: params.model,
instructions: params.instructions,
previousResponseId: params.previousResponseId,
},
{
...noCacheFetchOptions,
@@ -189,5 +189,7 @@ export interface GitBookDataFetcher {
input: api.AIMessageInput[];
output: api.AIOutputFormat;
model: api.AIModel;
instructions?: string;
previousResponseId?: string;
}): AsyncGenerator<api.AIStreamResponse, void, unknown>;
}
@@ -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() {
/>
</div>
</div>
<div className="flex-1"></div>
<div className="flex-1">
<AskMessages session={state.session} />
</div>
<div className="flex flex-row">
<AskInput />
</div>
@@ -20,6 +20,7 @@ export function AskInput() {
controller.postMessage({
message: value,
});
setValue('');
}
}}
/>
@@ -0,0 +1,15 @@
import type { AskSession } from './state';
export function AskMessages(props: {
session: AskSession;
}) {
const { session } = props;
return (
<div className="flex flex-col gap-2">
{session.messages.map((message, index) => {
return <div key={index}>{message.content}</div>;
})}
</div>
);
}
@@ -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 (
<div className="flex flex-col gap-2">
{message.steps.map((step, index) => {
return (
<div key={index} className="flex flex-col gap-2">
<DocumentView
document={step.content}
context={{
mode: 'default',
contentContext: undefined,
wrapBlocksInSuspense: false,
}}
style={['space-y-5']}
/>
</div>
);
})}
</div>
);
}
@@ -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<T>(
promise: MaybePromise<{
stream: EventIterator<T>;
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<T>(
context: GitBookBaseContext,
{
schema,
...input
}: Omit<Parameters<GitBookDataFetcher['streamAIResponse']>[0], 'output'> & {
schema: z.ZodSchema<T>;
}
) {
const rawStream = context.dataFetcher.streamAIResponse({
...input,
output: {
type: 'object',
schema: zodToJsonSchema(schema),
},
});
let json = '';
return parseResponse<DeepPartial<T>>(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<Parameters<GitBookDataFetcher['streamAIResponse']>[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<React.ReactNode>(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 <AIMessageView message={message} />;
});
}
/**
* Parse a stream from the API to extract the responseId.
*/
function parseResponse<T>(
responseStream: EventIterator<AIStreamResponse>,
parse: (response: AIStreamResponse) => T | undefined
): {
stream: EventIterator<T>;
response: Promise<{ responseId: string }>;
} {
let resolveResponse: (value: { responseId: string }) => void;
const response = new Promise<{ responseId: string }>((resolve) => {
resolveResponse = resolve;
});
const stream = new EventIterator<T>((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 };
}
@@ -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 };
}
@@ -0,0 +1,2 @@
export * from './api';
export * from './ask';
+68 -3
View File
@@ -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<AskState>(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);
}
},
};
}, []);