153 lines
5.5 KiB
TypeScript
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));
|
|
});
|
|
}
|
|
}
|