Files
ThothII/backend/src/catalog/model-completer.ts
T

153 lines
5.5 KiB
TypeScript

import { spawn } from "node:child_process";
import { z } from "zod";
import type { ResolvedMetadataGenerationModel } from "./metadata-generation-models.js";
const MAX_HELPER_OUTPUT_BYTES = 64 * 1024;
const helperOutputSchema = z.discriminatedUnion("ok", [
z.object({ ok: z.literal(true), content: z.string(), usage: z.object({ input: z.number().int().nonnegative(), cacheRead: z.number().int().nonnegative(), output: z.number().int().nonnegative() }).strict().optional() }).strict(),
z.object({ ok: z.literal(false), error: z.literal("provider_failure") }).strict(),
]);
export interface ModelCompletionMessage {
role: "system" | "user";
content: string;
}
export interface ModelCompletionRequest {
model: ResolvedMetadataGenerationModel;
messages: readonly ModelCompletionMessage[];
signal: AbortSignal;
}
export interface ModelCompletionUsage { input: number; cacheRead: number; output: number; }
export interface ModelCompletionResult { content: string; usage: ModelCompletionUsage; }
/** The provider boundary used by Description Generation. */
export interface ModelCompleter {
complete(request: ModelCompletionRequest): Promise<string | ModelCompletionResult>;
}
export class ModelCompletionProviderError extends Error {
constructor() {
super("model completion failed");
this.name = "ModelCompletionProviderError";
}
}
export class ModelCompletionCancelledError extends Error {
constructor() {
super("model completion cancelled");
this.name = "ModelCompletionCancelledError";
}
}
export class PythonModelCompleter implements ModelCompleter {
constructor(private readonly options: {
pythonExecutable: string;
cwd: string;
helperModule?: string;
timeoutMs?: number;
terminationGraceMs?: number;
}) {}
async complete(request: ModelCompletionRequest): Promise<ModelCompletionResult> {
if (request.signal.aborted) throw new ModelCompletionCancelledError();
const payload = {
model: `${request.model.provider}/${request.model.model}`,
...(request.model.apiKey === undefined ? {} : { api_key: request.model.apiKey }),
messages: request.messages.map((message) => ({ ...message })),
...(request.model.endpoint?.baseUrl === undefined
? {}
: { api_base: request.model.endpoint.baseUrl }),
...(request.model.endpoint?.apiVersion === undefined
? {}
: { api_version: request.model.endpoint.apiVersion }),
...(request.model.disableThinking === true ? { disable_thinking: true } : {}),
};
return await new Promise<ModelCompletionResult>((resolve, reject) => {
const child = spawn(
this.options.pythonExecutable,
["-m", this.options.helperModule ?? "tht.internal.litellm_completion"],
{
cwd: this.options.cwd,
stdio: ["pipe", "pipe", "pipe"],
},
);
let stdout = "";
let settled = false;
let timeout: ReturnType<typeof setTimeout> | undefined;
let killFallback: ReturnType<typeof setTimeout> | undefined;
let terminatingWith: Error | undefined;
const cleanup = () => {
if (timeout) clearTimeout(timeout);
if (killFallback) clearTimeout(killFallback);
request.signal.removeEventListener("abort", cancel);
};
const fail = (error: Error = new ModelCompletionProviderError()) => {
if (settled) return;
settled = true;
cleanup();
reject(error);
};
const terminate = (error: Error) => {
if (settled || terminatingWith) return;
terminatingWith = error;
if (timeout) clearTimeout(timeout);
try {
child.kill("SIGTERM");
} catch {
fail(error);
return;
}
killFallback = setTimeout(() => {
if (settled || child.exitCode !== null || child.signalCode !== null) return;
try {
child.kill("SIGKILL");
} catch {
fail(error);
}
}, this.options.terminationGraceMs ?? 250);
};
const cancel = () => {
terminate(new ModelCompletionCancelledError());
};
timeout = setTimeout(() => {
terminate(new ModelCompletionProviderError());
}, this.options.timeoutMs ?? 120_000);
request.signal.addEventListener("abort", cancel, { once: true });
if (request.signal.aborted) cancel();
child.stdout.setEncoding("utf8");
child.stdout.on("data", (chunk: string) => {
if (terminatingWith) return;
stdout += chunk;
if (Buffer.byteLength(stdout, "utf8") > MAX_HELPER_OUTPUT_BYTES) {
terminate(new ModelCompletionProviderError());
}
});
// Helper and provider diagnostics are deliberately not copied into application logs.
child.stderr.resume();
child.once("error", () => fail(terminatingWith ?? new ModelCompletionProviderError()));
child.once("close", (code) => {
if (settled) return;
if (terminatingWith) return fail(terminatingWith);
try {
if (code !== 0) return fail();
const output = helperOutputSchema.parse(JSON.parse(stdout));
if (!output.ok) return fail();
settled = true;
cleanup();
resolve({ content: output.content, usage: output.usage ?? { input: 0, cacheRead: 0, output: 0 } });
} catch {
fail();
}
});
child.stdin.once("error", () => {
if (terminatingWith) return;
terminate(new ModelCompletionProviderError());
});
child.stdin.end(JSON.stringify(payload));
});
}
}