feat: add AI catalog description generation
This commit is contained in:
@@ -0,0 +1,150 @@
|
||||
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() }).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;
|
||||
}
|
||||
|
||||
/** The provider boundary used by Description Generation. */
|
||||
export interface ModelCompleter {
|
||||
complete(request: ModelCompletionRequest): Promise<string>;
|
||||
}
|
||||
|
||||
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<string> {
|
||||
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<string>((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(output.content);
|
||||
} catch {
|
||||
fail();
|
||||
}
|
||||
});
|
||||
child.stdin.once("error", () => {
|
||||
if (terminatingWith) return;
|
||||
terminate(new ModelCompletionProviderError());
|
||||
});
|
||||
child.stdin.end(JSON.stringify(payload));
|
||||
});
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user