import { Emitter } from "@noya-app/emitter"; import { encodeSchema } from "@noya-app/noya-schemas"; import { Base64, isDeepEqual } from "@noya-app/noya-utils"; import { Observable } from "@noya-app/observable"; import type { Static, TObject, TSchema } from "@sinclair/typebox"; import { Value } from "@sinclair/typebox/value"; import type { AIGenerateImage, AIGenerateImageRequest, AIGenerateImageResponse, AIGenerateRequest, AIGenerateResponse, AIGenerateTextRequest, AIGenerateTextResponse, AIGenerationStreamEvent, AISamplingOptions, } from "./rpc/routes"; import { RPCManager } from "./rpcManager"; export type AIToolDefinition = { functionName: string; description: string; parameters: Record; onCall: (parameters: Record) => Promise; }; export type CallableAIToolsMap = Record< string, { description: string; parameters: Record } >; export type AIToolInvocation = { id: string; name: string; parameters: Record; }; export type AIImageInput = | Blob | { data: Uint8Array; mediaType: string } | { url: string | URL }; export type AIGenerateObjectOptions = AISamplingOptions & { prompt: string; schema: Schema; system?: string; images?: AIImageInput[]; model?: string; }; export type AIGenerateTextOptions = AISamplingOptions & { prompt: string; system?: string; images?: AIImageInput[]; model?: string; }; export type AIDeepPartial = T extends readonly (infer Item)[] ? AIDeepPartial[] : T extends object ? { [Key in keyof T]?: AIDeepPartial } : T; export type AIStreamObjectResult = { partialObjectStream: AsyncIterable>>; object: Promise>; }; export type AIStreamTextResult = { textStream: AsyncIterable; text: Promise; }; export type AIGenerateImageOptions = { prompt: string; model?: string; }; export type AIGeneratedImage = { data: Uint8Array; mediaType: string; }; export class AIManager { invocationEmitter = new Emitter<[AIToolInvocation]>(); responseEmitter = new Emitter<[AIToolInvocation, string | undefined]>(); tools$ = new Observable([]); systemMessage$ = new Observable(undefined); documentState$ = new Observable(undefined); documentSchema$ = new Observable(undefined); callableTools$ = this.tools$.map( (tools): CallableAIToolsMap => Object.fromEntries( tools.map(({ functionName, description, parameters }) => [ functionName, { description, parameters }, ]) ) ); configuration$ = Observable.combine( [this.callableTools$, this.systemMessage$], ([callableTools, systemMessage]) => ({ tools: callableTools, systemMessage, }), { isEqual: isDeepEqual } ); get callableTools() { return this.callableTools$.get(); } constructor(public rpcManager: RPCManager) { this.fetch = this.createFetch(); } private createFetch(): typeof globalThis.fetch { return (async (_input: RequestInfo | URL, init?: RequestInit) => { // Create a TransformStream to handle the response body const { readable, writable } = new TransformStream(); const writer = writable.getWriter(); // Make the RPC request with streaming this.rpcManager .requestStreamingRoute("POST /api/ai", { body: init?.body?.toString(), headers: init?.headers ? Object.fromEntries(new Headers(init.headers).entries()) : undefined, onStreamChunk: async (chunk) => { // Write each chunk to the stream await writer.write(new TextEncoder().encode(chunk)); }, onStreamEnd: async () => { await writer.close(); }, }) .then( async () => { // Nothing to do here - writer is closed in onStreamEnd }, async (error) => { await writer.abort(error); } ); // Create response object immediately with default values const response = new Response(readable, { status: 200, statusText: "OK", headers: new Headers({ "content-type": "text/plain; charset=utf-8", }), }); return response; }) as typeof globalThis.fetch; } fetch: typeof globalThis.fetch; registerTool(tool: AIToolDefinition) { this.tools$.set([...this.tools$.get(), tool]); return () => this.unregisterTool(tool.functionName); } unregisterTool(name: string) { this.tools$.set( this.tools$.get().filter((tool) => tool.functionName !== name) ); } callTool(toolInvocation: AIToolInvocation) { const tool = this.tools$ .get() .find((tool) => tool.functionName === toolInvocation.name); if (!tool) { console.error(`Tool ${toolInvocation.name} not found`); return; } return tool.onCall(toolInvocation.parameters); } private async requestGenerationStream( route: "POST /api/ai/generate" | "POST /api/ai/generate/text", payload: AIGenerateRequest | AIGenerateTextRequest, onEvent: (event: AIGenerationStreamEvent) => void ): Promise { let buffer = ""; let streamError: string | undefined; let protocolError: Error | undefined; const handleLine = (line: string) => { if (!line.trim()) return; const event = JSON.parse(line) as AIGenerationStreamEvent; if (event.type === "error") { streamError = event.error; } else if ( event.type === "text-delta" || event.type === "partial-object" || event.type === "object" || event.type === "text" ) { onEvent(event); } else { throw new Error("AI returned an unexpected stream event"); } }; const flushLines = (flushRemainder = false) => { const lines = buffer.split("\n"); buffer = lines.pop() ?? ""; if (flushRemainder && buffer.trim()) { lines.push(buffer); buffer = ""; } for (const line of lines) { handleLine(line); } }; await this.rpcManager.requestStreamingRoute(route, { body: JSON.stringify(payload), headers: { "Content-Type": "application/json" }, onStreamChunk: (chunk) => { if (protocolError) return; try { buffer += chunk; flushLines(); } catch (error) { protocolError = error instanceof Error ? error : new Error("AI returned invalid stream data"); } }, onStreamEnd: async () => { if (protocolError) return; try { flushLines(true); } catch (error) { protocolError = error instanceof Error ? error : new Error("AI returned invalid stream data"); } }, }); if (protocolError) throw protocolError; if (streamError) throw new Error(streamError); } async generateObject( options: AIGenerateObjectOptions ): Promise> { const { prompt, schema, system, images = [], model } = options; validatePromptAndModel(prompt, model); const samplingOptions = validateSamplingOptions(options); const payload: AIGenerateRequest = { prompt, schema, encodedSchema: encodeSchema(schema), ...samplingOptions, ...(system ? { system } : {}), ...(model ? { model } : {}), ...(images.length ? { images: await Promise.all(images.map(normalizeImageInput)) } : {}), }; const response = await this.rpcManager.requestRoute( "POST /api/ai/generate", { body: JSON.stringify(payload), headers: { "Content-Type": "application/json" }, } ); const result = this.rpcManager.getResponseBody(response) as | AIGenerateResponse | undefined; if (!result || !("object" in result)) { throw new Error("AI returned an unexpected response"); } return validateGeneratedObject(schema, result.object); } streamObject( options: AIGenerateObjectOptions ): AIStreamObjectResult { const { prompt, schema, system, images = [], model } = options; validatePromptAndModel(prompt, model); const samplingOptions = validateSamplingOptions(options); const partials = createAsyncQueue>>(); const objectResult = Promise.withResolvers>(); void (async () => { try { const payload: AIGenerateRequest = { prompt, schema, encodedSchema: encodeSchema(schema), stream: true, ...samplingOptions, ...(system ? { system } : {}), ...(model ? { model } : {}), ...(images.length ? { images: await Promise.all(images.map(normalizeImageInput)) } : {}), }; let finalObject: unknown; let hasFinalObject = false; await this.requestGenerationStream( "POST /api/ai/generate", payload, (event) => { if (event.type === "partial-object") { partials.push(event.object as AIDeepPartial>); } else if (event.type === "object") { finalObject = event.object; hasFinalObject = true; } else { throw new Error("AI returned an unexpected object stream event"); } } ); if (!hasFinalObject) { throw new Error("AI returned no result in its object stream"); } const object = validateGeneratedObject(schema, finalObject); partials.end(); objectResult.resolve(object); } catch (error) { const streamError = toError(error); partials.fail(streamError); objectResult.reject(streamError); } })(); return { partialObjectStream: partials.iterable, object: objectResult.promise, }; } async generateText(options: AIGenerateTextOptions): Promise { const { prompt, system, images = [], model } = options; validatePromptAndModel(prompt, model); const samplingOptions = validateSamplingOptions(options); const payload: AIGenerateTextRequest = { prompt, ...samplingOptions, ...(system ? { system } : {}), ...(model ? { model } : {}), ...(images.length ? { images: await Promise.all(images.map(normalizeImageInput)) } : {}), }; const response = await this.rpcManager.requestRoute( "POST /api/ai/generate/text", { body: JSON.stringify(payload), headers: { "Content-Type": "application/json" }, } ); const result = this.rpcManager.getResponseBody(response) as | AIGenerateTextResponse | undefined; if (!result || typeof result.text !== "string") { throw new Error("AI returned an unexpected text response"); } return result.text; } streamText(options: AIGenerateTextOptions): AIStreamTextResult { const { prompt, system, images = [], model } = options; validatePromptAndModel(prompt, model); const samplingOptions = validateSamplingOptions(options); const chunks = createAsyncQueue(); const textResult = Promise.withResolvers(); void (async () => { try { const payload: AIGenerateTextRequest = { prompt, stream: true, ...samplingOptions, ...(system ? { system } : {}), ...(model ? { model } : {}), ...(images.length ? { images: await Promise.all(images.map(normalizeImageInput)) } : {}), }; let finalText: string | undefined; await this.requestGenerationStream( "POST /api/ai/generate/text", payload, (event) => { if (event.type === "text-delta") { chunks.push(event.textDelta); } else if (event.type === "text") { finalText = event.text; } else { throw new Error("AI returned an unexpected text stream event"); } } ); if (finalText === undefined) { throw new Error("AI returned no result in its text stream"); } chunks.end(); textResult.resolve(finalText); } catch (error) { const streamError = toError(error); chunks.fail(streamError); textResult.reject(streamError); } })(); return { textStream: chunks.iterable, text: textResult.promise }; } async generateImage({ prompt, model, }: AIGenerateImageOptions): Promise { validatePromptAndModel(prompt, model); const payload: AIGenerateImageRequest = { prompt, ...(model ? { model } : {}), }; const response = await this.rpcManager.requestRoute( "POST /api/ai/generate/image", { body: JSON.stringify(payload), headers: { "Content-Type": "application/json" }, } ); const result = this.rpcManager.getResponseBody(response) as | AIGenerateImageResponse | undefined; if ( !result || typeof result.base64 !== "string" || typeof result.mediaType !== "string" ) { throw new Error("AI returned an unexpected image response"); } return { data: Base64.decode(result.base64), mediaType: result.mediaType, }; } } function validateGeneratedObject( schema: Schema, object: unknown ): Static { if (!Value.Check(schema, object)) { const errors = [...Value.Errors(schema, object)] .map(({ path, message }) => `${path || "/"}: ${message}`) .join(", "); throw new Error( `AI response does not match the requested schema: ${errors}` ); } return object as Static; } function createAsyncQueue() { const values: T[] = []; let isDone = false; let error: Error | undefined; let wake: (() => void) | undefined; const notify = () => { wake?.(); wake = undefined; }; const next = async (): Promise> => { if (values.length > 0) { return { done: false, value: values.shift()! }; } if (error) throw error; if (isDone) return { done: true, value: undefined }; await new Promise((resolve) => { wake = resolve; }); return next(); }; const iterable: AsyncIterable = { [Symbol.asyncIterator]() { return { next }; }, }; return { iterable, push(value: T) { if (isDone || error) return; values.push(value); notify(); }, end() { isDone = true; notify(); }, fail(reason: Error) { error = reason; notify(); }, }; } function toError(error: unknown): Error { return error instanceof Error ? error : new Error(String(error)); } function validateSamplingOptions( options: AISamplingOptions ): AISamplingOptions { validateSamplingRange(options.temperature, "temperature", 0, 1); validateSamplingRange(options.topP, "topP", 0, 1); validateSamplingRange(options.presencePenalty, "presencePenalty", -1, 1); validateSamplingRange(options.frequencyPenalty, "frequencyPenalty", -1, 1); if ( options.topK !== undefined && (!Number.isSafeInteger(options.topK) || options.topK <= 0) ) { throw new Error("AI topK must be a positive integer"); } if (options.seed !== undefined && !Number.isSafeInteger(options.seed)) { throw new Error("AI seed must be an integer"); } return { temperature: options.temperature, topP: options.topP, topK: options.topK, presencePenalty: options.presencePenalty, frequencyPenalty: options.frequencyPenalty, seed: options.seed, }; } function validateSamplingRange( value: number | undefined, name: string, min: number, max: number ) { if ( value !== undefined && (!Number.isFinite(value) || value < min || value > max) ) { throw new Error(`AI ${name} must be between ${min} and ${max}`); } } function validatePromptAndModel(prompt: string, model?: string) { if (!prompt.trim()) { throw new Error("AI prompt must not be empty"); } if (model !== undefined && !model.trim()) { throw new Error("AI model must not be empty"); } } async function normalizeImageInput( image: AIImageInput ): Promise { if (image instanceof Blob) { assertImageMediaType(image.type); return { type: "data", data: Base64.encode(await image.arrayBuffer()), mediaType: image.type, }; } if ("url" in image) { const url = new URL(image.url).toString(); if (!["http:", "https:"].includes(new URL(url).protocol)) { throw new Error("AI image URLs must use http or https"); } return { type: "url", url }; } assertImageMediaType(image.mediaType); return { type: "data", data: Base64.encode(image.data), mediaType: image.mediaType, }; } function assertImageMediaType(mediaType: string) { if (!mediaType.toLowerCase().startsWith("image/")) { throw new Error(`Invalid AI image media type: ${mediaType || "(missing)"}`); } }