Files
veridian/server/api/chat/generate.post.ts
T

886 lines
32 KiB
TypeScript

import type { Entity } from '@triplit/client';
import { type ModelMessage, modelMessageSchema, streamText, type StreamTextTransform, tool } from 'ai';
import { promises as fs } from 'fs';
import { glob } from 'glob';
import { nanoid } from 'nanoid';
import path from 'path';
import * as z from 'zod';
import { httpClient } from '~~/server/lib/triplit';
import { addPendingGeneration, completeGeneration } from '~~/server/utils/generations';
import type { schema } from '~~/triplit/schema';
import { spawn } from 'child_process';
import { getGateway, type ModelGateway } from '~~/server/utils/ai-provider';
export const messagesSchema = z.array(modelMessageSchema);
export default defineEventHandler(async (event) => {
await protectRoute(event);
const result = await readValidatedBody(event, (body) =>
z
.object({
messages: messagesSchema.min(1),
topicId: z.string(),
parentMessageId: z.string().nullable(),
model: z.object({
providerId: z.string(),
modelId: z.string(),
args: z.record(z.string(), z.any()),
}),
providerApiKey: z.string().optional(),
})
.safeParse(body),
);
if (!result.success) {
throw createError({
statusCode: 400,
message: result.error.issues[0]!.message,
});
}
const userId = event.context.user!.id;
const { messages, topicId, parentMessageId, model: { modelId, providerId, args }, providerApiKey } = result.data;
const fetchPromises = [];
fetchPromises.push(httpClient.fetchOne(httpClient.query('providers').Where('id', '=', providerId)));
fetchPromises.push(httpClient.fetchOne(httpClient.query('models').Where('id', '=', modelId)));
fetchPromises.push(httpClient.fetchOne(
httpClient.query('generations').Where('topicId', '=', topicId).Where('status', '=', 'pending'),
));
const [provider, model, existingPendingGenerations] = await Promise.all(fetchPromises) as [Entity<typeof schema, 'providers'> | null, Entity<typeof schema, 'models'> | null, Entity<typeof schema, 'generations'> | null];
if (provider === null || provider.userId !== userId) {
throw createError({
statusCode: 400,
message: 'Invalid provider',
});
}
if (model === null || model.providerId !== model.providerId || model.userId !== userId) {
throw createError({
statusCode: 400,
message: 'Invalid model',
});
}
if (existingPendingGenerations !== null) {
throw createError({
statusCode: 400,
message: 'There cannot be more than one active generation per topic',
});
}
const { gateway, streamTransformer: transformer } = await getGateway(provider, model, providerApiKey);
const generationId = nanoid();
const message = await httpClient.insert('messages', {
topicId,
userId,
focusedIndex: parentMessageId ? null : 0,
generationId,
parentMessageId,
content: '',
role: 'assistant',
});
await httpClient.insert('generations', {
id: generationId,
userId,
topicId,
modelId: model.externalId,
status: 'pending',
messageId: message.id,
});
let logFile: fs.FileHandle | undefined;
let logMessage: ((message: string) => void) | undefined;
if (process.env.GENERATION_DEBUG) {
logFile = await fs.open(path.join(process.env.LOG_DIR!, `${Date.now()}-${generationId}.log`), 'w');
logMessage = (message: string) => {
logFile!.write(message + '\n');
};
}
event.waitUntil(
generateResponse(message, { gateway, model, parameters: args }, generationId, userId, topicId, messages, transformer, logMessage, logFile),
);
return {
generationId,
messageId: message.id,
};
});
const INTERNAL_ERROR = 'An internal error occurred';
// todo message takes in variadics like console.log
const todo = (...args: any[]) => {
console.error('TODO', ...args);
throw new Error('TODO');
};
const evalPython = async (code: string) => {
// This is the wrapper logic from above, minified or stored as a string
// Or you can save the wrapper script to a file and call that.
const wrapper = `
import ast
import sys
code = sys.stdin.read()
tree = ast.parse(code)
last_node = tree.body[-1] if tree.body else None
namespace = {}
if len(tree.body) > 1:
exec(compile(ast.Module(tree.body[:-1], []), "<ast>", "exec"), namespace)
if isinstance(last_node, ast.Expr):
res = eval(compile(ast.Expression(last_node.value), "<ast>", "eval"), namespace)
if res is not None: print(res)
elif last_node:
exec(compile(ast.Module([last_node], []), "<ast>", "exec"), namespace)
`.trim();
return new Promise<string>((resolve, reject) => {
const child = spawn('python3', ['-c', wrapper]);
let output = '';
let errorOutput = '';
child.stdout.on('data', (data) => {
output += data.toString();
});
child.stderr.on('data', (data) => {
errorOutput += data.toString();
});
child.on('close', (exitCode) => {
if (exitCode !== 0) {
reject(errorOutput || `Exit code ${exitCode}`);
} else {
resolve(output.trim());
}
});
// Send the agent's code to the wrapper via stdin
child.stdin.write(code);
child.stdin.end();
});
};
// TODO: obviously come up with a better way for the user to define their own tools
const { listDirectoryTool, globTool, readFileTool, readFilesTool, fetchUrlTool, pythonTool, bashTool } = {
listDirectoryTool: tool({
description: 'Lists the contents of a directory',
inputSchema: z.object({
path: z.string(),
}),
outputSchema: z.object({
files: z.array(z.object({ name: z.string(), type: z.string() })),
}),
execute: async ({ path }) => {
const rawFiles = await fs.readdir(path, { withFileTypes: true });
const files = rawFiles.map((file) => ({
name: file.name,
type: file.isFile() ? 'file' : 'directory',
}));
return {
files,
};
},
}),
globTool: tool({
description: 'Lists files matching a glob pattern',
inputSchema: z.object({
pattern: z.string(),
}),
outputSchema: z.object({
files: z.array(z.object({ path: z.string(), type: z.string() })),
}),
execute: async ({ pattern }) => {
const rawFiles = await glob(pattern, { withFileTypes: true });
const files = rawFiles.map((file) => ({
path: file.parentPath + '/' + file.name,
type: file.isFile() ? 'file' : 'directory',
}));
return {
files,
};
},
}),
readFileTool: tool({
description: 'Reads the contents of a file',
inputSchema: z.object({
path: z.string(),
}),
outputSchema: z.object({
path: z.string(),
content: z.string(),
}),
execute: async ({ path }) => {
const file = await fs.readFile(path);
return {
path,
content: file.toString(),
};
},
}),
readFilesTool: tool({
description: 'Reads the contents of multiple files',
inputSchema: z.object({
paths: z.array(z.string()).describe('The file paths to read'),
}),
outputSchema: z.object({
files: z.array(
z.object({
path: z.string(),
content: z.string(),
}),
),
}),
execute: async ({ paths }) => {
const files = await Promise.all(
paths.map(async (path) => {
const file = await fs.readFile(path);
return {
path: path,
content: file.toString(),
};
}),
);
return {
files,
};
},
}),
fetchUrlTool: tool({
description: 'Fetches the content of a URL',
inputSchema: z.object({
url: z.string(),
}),
outputSchema: z.object({
content: z.string(),
}),
execute: async ({ url }) => {
const response = await fetch(url);
const content = await response.text();
return {
content,
};
},
}),
pythonTool: tool({
description: 'Executes a Python code snippet',
inputSchema: z.object({
code: z.string(),
}),
outputSchema: z.object({
output: z.string(),
}),
execute: async ({ code }) => {
const output = await evalPython(code);
return {
output,
};
},
}),
bashTool: tool({
description: 'Executes a Bash command',
inputSchema: z.object({
code: z.string(),
}),
outputSchema: z.object({
output: z.string(),
}),
execute: async ({ code }) => {
const { exec } = await import('child_process');
const { promisify } = await import('util');
const execAsync = promisify(exec);
async function runCommand(command: string) {
const { stdout, stderr } = await execAsync(command);
if (stderr) {
console.error(`Error: ${stderr}`);
return stderr;
} else {
return stdout;
}
}
const output = await runCommand(code);
return {
output,
};
},
}),
}
async function generateResponse(
message: Entity<typeof schema, 'messages'>,
model: {
gateway: ModelGateway,
model: Entity<typeof schema, 'models'>,
parameters: Record<string, any>,
},
generationId: string,
userId: string,
topicId: string,
messages: ModelMessage[],
streamTransoforms: StreamTextTransform<{}> | StreamTextTransform<{}>[] | undefined,
log?: (message: string) => void,
logFile?: fs.FileHandle,
) {
const controller = new AbortController();
addPendingGeneration(generationId, controller);
let requestStart = undefined;
let ttft = undefined;
const activeParts = new Map<string, { id: string; accumulatedContent: string; providerOptions?: any }>();
const activeToolCalls = new Map<string, void>();
const tools = {
// writeFile: tool({
// inputSchema: z.object({
// path: z.string(),
// content: z.string(),
// }),
// outputSchema: z.object({
// success: z.boolean(),
// }),
// execute: async ({ path, content }) => {
// await fs.writeFile(path, content);
// return {
// success: true,
// };
// }
// }),
listDirectory: listDirectoryTool,
glob: globTool,
readFile: readFileTool,
readFiles: readFilesTool,
// fetchUrl: fetchUrlTool,
python: pythonTool,
bash: bashTool,
};
const response = streamText({
model: model.gateway(model.model.externalId),
messages,
providerOptions: {
openrouter: {
debug: {
echo_upstream_body: true,
},
}
},
experimental_transform: streamTransoforms,
// a little trick that makes it so that the stream doesnt stop because of tool calls, and will continue an unbounded amount of time and steps
stopWhen: [],
tools: [...model.model.attributes.capabilities].includes('tools') ? tools : undefined,
onError: async (error: any) => {
console.error('generation error', error);
log?.(error);
// TODO: the docs say "The stream processing will pause until the callback promise is resolved." Suggesting that this error might not be fatal?
for (const activePart of activeParts.values()) {
await httpClient.update('message_parts', activePart.id, {
finished: true,
lastUpdatedAt: new Date(),
});
}
for (const toolCallId of activeToolCalls.keys()) {
await httpClient.update('tool_calls', toolCallId, {
status: 'failed',
error: {
type: 'text',
value: 'An unknown error occurred',
},
});
activeToolCalls.delete(toolCallId);
}
let errValue = error.message || error;
if (typeof errValue === 'object') {
errValue = JSON.stringify(errValue, null, 2);
}
await httpClient.update('generations', generationId, {
status: 'failed',
error: errValue,
});
},
abortSignal: controller.signal,
});
let curStepIdx = -1;
let key, part, type;
const pendingUpdates = new Map<string, NodeJS.Timeout>();
const TARGET_UPDATES_PER_SECOND = 24;
const scheduleUpdate = (key: string) => {
const part = activeParts.get(key);
if (!part || pendingUpdates.has(part.id)) return;
pendingUpdates.set(part.id, setTimeout(async () => {
const currentPart = activeParts.get(key);
// Only update if the part is still active and we haven't deleted it at 'text-end'
if (currentPart) {
try {
await httpClient.update('message_parts', currentPart.id, {
content: currentPart.accumulatedContent,
providerOptions: currentPart.providerOptions,
lastUpdatedAt: new Date(),
});
} catch (error) {
// the update failed, but it doesnt matter because the full message will be updated on step finish
console.warn('Failed to update message part', error);
}
}
pendingUpdates.delete(part.id);
}, 1000 / TARGET_UPDATES_PER_SECOND));
};
try {
for await (const token of response.fullStream) {
log?.(JSON.stringify(token, null, 2));
switch (token.type) {
case 'start': {
requestStart = Date.now();
} break;
case 'start-step': {
curStepIdx++;
} break;
case 'tool-input-start': {
key = `tool-call-${curStepIdx}`;
const toolCallId = token.id;
part = await httpClient.insert('message_parts', {
topicId,
messageId: message.id,
userId,
toolCallId,
type: 'tool-call',
content: '',
finished: false,
createdAt: new Date(),
lastUpdatedAt: new Date(),
});
await httpClient.insert('tool_calls', {
id: toolCallId,
userId: userId,
toolName: token.toolName,
status: 'pending',
input: null,
output: null,
error: null,
createdAt: new Date(),
});
activeToolCalls.set(toolCallId);
activeParts.set(key, { id: part.id, accumulatedContent: '' });
} break;
case 'text-start':
case 'reasoning-start': {
type = token.type.split('-')[0];
key = `${type}-${curStepIdx}`;
part = await httpClient.insert('message_parts', {
topicId,
messageId: message.id,
userId,
type: type as 'text' | 'reasoning',
content: '',
finished: false,
createdAt: new Date(),
lastUpdatedAt: new Date(),
});
activeParts.set(key, { id: part.id, accumulatedContent: '' });
} break;
case 'text-delta':
case 'reasoning-delta': {
if (ttft === undefined) {
ttft = Date.now() - requestStart!;
}
type = token.type.split('-')[0] as 'text' | 'reasoning';
key = `${type}-${curStepIdx}`;
part = activeParts.get(key);
if (part === undefined) {
console.error('Received delta without a start');
break;
}
let shouldUpdate = false;
// TODO: we should potentially merge providerOptions, but for now, just overwrite them
if (token.providerMetadata !== undefined) {
shouldUpdate = true;
part.providerOptions = token.providerMetadata;
}
// OpenRouter sometimes puts [REDACTED] in thinking if reasoning is encrypted, so we need to remove it and hide it;
// do not trim or else we lose intentional whitespace and newlines potentially breaking the UI and having words comebined e.g. "the" "\n\n" "assistant" would become "theassistant"
const text = token.text.replaceAll('[REDACTED]', '');
if (text !== '') {
shouldUpdate = true;
part.accumulatedContent += token.text;
}
if (shouldUpdate) {
scheduleUpdate(key);
}
}
break;
case 'text-end':
case 'reasoning-end': {
type = token.type.split('-')[0];
key = `${type}-${curStepIdx}`;
part = activeParts.get(key);
if (part === undefined) {
console.error('Received end without a start');
break;
}
activeParts.delete(key);
if (part.accumulatedContent === '' && part.providerOptions === undefined) {
// completely empty, delete it
await httpClient.delete('message_parts', part.id);
break;
}
await httpClient.update('message_parts', part.id, {
content: part.accumulatedContent,
providerOptions: part.providerOptions,
finished: true,
lastUpdatedAt: new Date(),
});
} break;
case 'tool-call': {
let inputType: 'text' | 'json' = 'text';
let inputValue: string = '';
switch (typeof token.input) {
case 'string':
inputType = 'text';
inputValue = token.input;
break;
case 'object':
inputType = 'json';
inputValue = JSON.stringify(token.input);
break;
default:
console.error('Unknown input type', token.input);
break;
}
if (activeToolCalls.has(token.toolCallId)) {
await httpClient.update('tool_calls', token.toolCallId, {
status: 'pending',
input: {
type: inputType,
value: inputValue,
},
});
} else {
key = `tool-call-${curStepIdx}`;
const toolCallId = token.toolCallId;
part = await httpClient.insert('message_parts', {
topicId,
messageId: message.id,
userId,
toolCallId,
type: 'tool-call',
content: '',
finished: false,
createdAt: new Date(),
lastUpdatedAt: new Date(),
});
await httpClient.insert('tool_calls', {
id: toolCallId,
userId: userId,
toolName: token.toolName,
status: 'pending',
input: {
type: inputType,
value: inputValue,
},
output: null,
error: null,
createdAt: new Date(),
});
activeToolCalls.set(toolCallId);
activeParts.set(key, { id: part.id, accumulatedContent: '' });
}
} break;
case 'tool-result': {
let outputType: 'text' | 'json' = 'text';
let outputValue: string = '';
switch (typeof token.output) {
case 'string':
outputType = 'text';
outputValue = token.output;
break;
case 'object':
outputType = 'json';
outputValue = JSON.stringify(token.output);
break;
default:
console.error('Unknown output type', token.output);
await httpClient.update('tool_calls', token.toolCallId, {
status: 'failed',
error: {
type: 'text',
value: 'Tool returned invalid output',
},
});
activeToolCalls.delete(token.toolCallId);
break;
}
await httpClient.update('tool_calls', token.toolCallId, {
status: 'completed',
output: {
type: outputType,
value: outputValue,
},
});
activeToolCalls.delete(token.toolCallId);
} break;
case 'tool-error': {
console.error('Tool error:', token);
let outputType: 'text' | 'json';
let outputValue: string;
switch (typeof token.error) {
case 'string':
outputType = 'text';
outputValue = token.error;
break;
case 'object':
outputType = 'json';
outputValue = JSON.stringify(token.error);
break;
default:
console.error('Unknown error type', token.error);
outputType = 'text';
outputValue = 'Tool returned invalid output';
break;
}
if (activeToolCalls.has(token.toolCallId)) {
await httpClient.update('tool_calls', token.toolCallId, {
status: 'failed',
error: {
type: outputType,
value: outputValue,
},
});
} else {
key = `tool-call-${curStepIdx}`;
const toolCallId = token.toolCallId;
part = await httpClient.insert('message_parts', {
topicId,
messageId: message.id,
userId,
toolCallId: token.toolCallId,
type: 'tool-call',
content: '',
finished: false,
createdAt: new Date(),
lastUpdatedAt: new Date(),
});
await httpClient.insert('tool_calls', {
id: toolCallId,
userId: userId,
toolName: token.toolName,
status: 'failed',
input: null,
output: null,
error: {
type: outputType,
value: outputValue,
},
createdAt: new Date(),
});
activeParts.set(key, { id: part.id, accumulatedContent: '' });
}
activeToolCalls.delete(token.toolCallId);
} break;
case 'error': {
console.error('Error:', token);
for (const activePart of activeParts.values()) {
await httpClient.update('message_parts', activePart.id, {
finished: true,
lastUpdatedAt: new Date(),
});
}
let error = INTERNAL_ERROR;
if (typeof token.error === 'string') {
error = token.error;
} else if (typeof token.error === 'object') {
error = JSON.stringify(token.error);
}
await httpClient.update('generations', generationId, {
status: 'failed',
error,
});
throw new Error(error);
} break;
case 'finish': {
let tps;
if (ttft !== undefined && token.totalUsage.outputTokens !== undefined) {
const tokenStreamStart = requestStart! + ttft;
// this is the *real* request duration, excluding the
// TTFT
const requestDuration = Date.now() - tokenStreamStart;
tps = token.totalUsage.outputTokens / (requestDuration / 1000);
}
for (const activePart of activeParts.values()) {
await httpClient.update('message_parts', activePart.id, {
finished: true,
lastUpdatedAt: new Date(),
});
}
switch (token.finishReason) {
case 'error':
await httpClient.update('generations', generationId, {
status: 'failed',
error: INTERNAL_ERROR,
});
break;
case 'content-filter':
await httpClient.update('generations', generationId, {
status: 'failed',
error: 'Content was filtered',
});
break;
}
const generation = await httpClient.fetchOne(httpClient.query('generations').Where('id', '=', generationId));
if (generation === null || generation.status === 'failed') return;
await httpClient.update('generations', generationId, {
status: 'completed',
tokens: {
input: token.totalUsage.inputTokens,
cache: {
read: token.totalUsage.inputTokenDetails.cacheReadTokens,
write: token.totalUsage.inputTokenDetails.cacheWriteTokens,
},
output: token.totalUsage.outputTokens,
thinking: token.totalUsage.outputTokenDetails.reasoningTokens,
ttft,
tps,
},
});
// I hate you switch fallthroughs
} break;
case 'abort': {
for (const activePart of activeParts.values()) {
await httpClient.update('message_parts', activePart.id, {
finished: true,
lastUpdatedAt: new Date(),
});
}
for (const toolCallId of activeToolCalls.keys()) {
await httpClient.update('tool_calls', toolCallId, {
status: 'cancelled',
});
}
await httpClient.update('generations', generationId, {
status: 'cancelled',
});
} break;
case 'file':
todo('file token type', token);
break;
case 'raw':
todo('raw token type', token);
break;
case 'source':
todo('source token type', token);
break;
case 'tool-approval-request':
todo('tool-approval-request token type', token);
break;
// typescript thinks this is not a real token type?
// case 'tool-output-denied':
// todo('tool-output-denied token type', token);
// break;
case 'finish-step':
case 'tool-input-delta':
case 'tool-input-end':
// handled or irrelevant
break;
}
}
} catch (error: any) {
console.error(error);
for (const activePart of activeParts.values()) {
await httpClient.update('message_parts', activePart.id, {
finished: true,
lastUpdatedAt: new Date(),
});
}
for (const toolCallId of activeToolCalls.keys()) {
await httpClient.update('tool_calls', toolCallId, {
status: 'failed',
error: {
type: 'text',
value: 'An unknown error occurred',
},
});
activeToolCalls.delete(toolCallId);
}
await httpClient.update('generations', generationId, {
status: 'failed',
error: error.message,
});
} finally {
completeGeneration(generationId);
if (logFile !== undefined) logFile.close();
}
}