Files
ThothII/backend/src/catalog/local-ner-detector.ts
T

255 lines
8.1 KiB
TypeScript

import { randomUUID } from "node:crypto";
import { spawn, type ChildProcessWithoutNullStreams } from "node:child_process";
import { tmpdir } from "node:os";
import { z } from "zod";
import type {
LocalNerCandidate,
LocalNerDetector,
LocalNerEvidence,
} from "./sensitivity-classifier.js";
const MAX_LINE_BYTES = 64 * 1024;
const candidateSchema = z.object({
columnId: z.uuid(),
text: z.string().min(1).max(500),
}).strict();
const workerMessageSchema = z.union([
z.object({ ready: z.literal(true) }).strict(),
z.object({
id: z.uuid(),
ok: z.literal(true),
evidence: z.array(z.object({
columnId: z.uuid(),
label: z.string().min(1).max(80),
confidence: z.number().min(0).max(1),
}).strict()).max(1_000),
}).strict(),
z.object({ id: z.uuid(), ok: z.literal(false), error: z.string().min(1).max(80) }).strict(),
]);
export class LocalNerUnavailableError extends Error {
constructor() {
super("local NER is unavailable");
this.name = "LocalNerUnavailableError";
}
}
interface PendingRequest {
resolve: (value: readonly LocalNerEvidence[]) => void;
reject: (error: Error) => void;
timer: ReturnType<typeof setTimeout>;
signal: AbortSignal;
cancel: () => void;
}
/** Persistent JSONL adapter for the optional, CPU-only Python NER worker. */
export class PythonLocalNerDetector implements LocalNerDetector {
private child?: ChildProcessWithoutNullStreams;
private ready?: Promise<void>;
private readyResolve?: () => void;
private readyReject?: (error: Error) => void;
private workerReady = false;
private stdout = "";
private readonly pending = new Map<string, PendingRequest>();
constructor(private readonly options: {
pythonExecutable: string;
workerScript: string;
modelPath: string;
cwd: string;
threads?: number;
startupTimeoutMs?: number;
}) {}
async warmup(): Promise<void> {
await this.ensureStarted();
}
isReady(): boolean {
return this.workerReady
&& this.child !== undefined
&& this.child.exitCode === null
&& this.child.signalCode === null;
}
async detect(
candidates: readonly LocalNerCandidate[],
signal: AbortSignal,
deadline: number,
): Promise<readonly LocalNerEvidence[]> {
const parsed = z.array(candidateSchema).min(1).max(128).parse(candidates);
if (signal.aborted || deadline <= Date.now()) throw new LocalNerUnavailableError();
await this.ensureStartedWithin(signal, deadline);
if (!this.child || this.child.exitCode !== null || this.child.signalCode !== null) {
throw new LocalNerUnavailableError();
}
const id = randomUUID();
return await new Promise<readonly LocalNerEvidence[]>((resolve, reject) => {
const fail = () => {
this.finishPending(id);
reject(new LocalNerUnavailableError());
this.stopWorker();
};
const timer = setTimeout(fail, Math.max(1, Math.floor(deadline - Date.now())));
const cancel = fail;
const pending: PendingRequest = { resolve, reject, timer, signal, cancel };
this.pending.set(id, pending);
signal.addEventListener("abort", cancel, { once: true });
this.child!.stdin.write(`${JSON.stringify({ id, candidates: parsed })}\n`, (error) => {
if (error) fail();
});
});
}
async close(): Promise<void> {
const child = this.child;
if (!child || child.exitCode !== null || child.signalCode !== null) return;
await new Promise<void>((resolve) => {
child.once("close", () => resolve());
child.kill("SIGTERM");
setTimeout(() => {
if (child.exitCode === null && child.signalCode === null) child.kill("SIGKILL");
}, 250).unref();
});
}
private async ensureStarted(): Promise<void> {
if (this.ready) return await this.ready;
this.ready = new Promise<void>((resolve, reject) => {
this.readyResolve = resolve;
this.readyReject = reject;
});
const threads = String(this.options.threads ?? 2);
const inheritedRuntimeEnvironment = Object.fromEntries([
"PATH", "SystemRoot", "WINDIR", "PATHEXT", "TMPDIR", "TEMP", "TMP", "LANG", "LC_ALL",
].flatMap((name) => process.env[name] === undefined ? [] : [[name, process.env[name]!]]));
const child = spawn(this.options.pythonExecutable, [
"-I",
"-B",
this.options.workerScript,
"--model",
this.options.modelPath,
"--threads",
threads,
], {
cwd: this.options.cwd,
stdio: ["pipe", "pipe", "pipe"],
env: {
...inheritedRuntimeEnvironment,
HOME: process.env.HOME ?? tmpdir(),
CUDA_VISIBLE_DEVICES: "",
HIP_VISIBLE_DEVICES: "",
HF_HUB_OFFLINE: "1",
HF_HUB_DISABLE_TELEMETRY: "1",
TRANSFORMERS_OFFLINE: "1",
TOKENIZERS_PARALLELISM: "false",
PYTHONNOUSERSITE: "1",
OMP_NUM_THREADS: threads,
MKL_NUM_THREADS: threads,
OPENBLAS_NUM_THREADS: threads,
HTTP_PROXY: "",
HTTPS_PROXY: "",
ALL_PROXY: "",
NO_PROXY: "*",
},
});
this.child = child;
child.stdout.setEncoding("utf8");
child.stdout.on("data", (chunk: string) => this.receive(chunk));
child.stderr.resume();
child.once("error", () => this.failWorker());
child.once("close", () => this.failWorker());
const startupTimer = setTimeout(() => this.failWorker(), this.options.startupTimeoutMs ?? 120_000);
startupTimer.unref();
try {
await this.ready;
} finally {
clearTimeout(startupTimer);
}
}
private async ensureStartedWithin(signal: AbortSignal, deadline: number): Promise<void> {
const started = this.ensureStarted();
await new Promise<void>((resolve, reject) => {
let settled = false;
const finish = (error?: Error, stopWorker = false) => {
if (settled) return;
settled = true;
clearTimeout(timer);
signal.removeEventListener("abort", cancel);
if (stopWorker) this.failWorker();
if (error) reject(error);
else resolve();
};
const cancel = () => finish(new LocalNerUnavailableError(), true);
const timer = setTimeout(cancel, Math.max(1, Math.floor(deadline - Date.now())));
signal.addEventListener("abort", cancel, { once: true });
void started.then(
() => finish(),
() => finish(new LocalNerUnavailableError()),
);
});
}
private receive(chunk: string): void {
this.stdout += chunk;
if (Buffer.byteLength(this.stdout, "utf8") > MAX_LINE_BYTES) {
this.failWorker();
return;
}
let newline: number;
while ((newline = this.stdout.indexOf("\n")) >= 0) {
const line = this.stdout.slice(0, newline);
this.stdout = this.stdout.slice(newline + 1);
if (!line) continue;
try {
const message = workerMessageSchema.parse(JSON.parse(line));
if ("ready" in message) {
this.workerReady = true;
this.readyResolve?.();
this.readyResolve = undefined;
this.readyReject = undefined;
continue;
}
const pending = this.pending.get(message.id);
if (!pending) continue;
this.finishPending(message.id);
if (message.ok) pending.resolve(message.evidence);
else pending.reject(new LocalNerUnavailableError());
} catch {
this.failWorker();
return;
}
}
}
private finishPending(id: string): void {
const pending = this.pending.get(id);
if (!pending) return;
clearTimeout(pending.timer);
pending.signal.removeEventListener("abort", pending.cancel);
this.pending.delete(id);
}
private stopWorker(): void {
const child = this.child;
if (child && child.exitCode === null && child.signalCode === null) child.kill("SIGTERM");
}
private failWorker(): void {
const error = new LocalNerUnavailableError();
this.readyReject?.(error);
this.readyResolve = undefined;
this.readyReject = undefined;
for (const [id, pending] of this.pending) {
this.finishPending(id);
pending.reject(error);
}
this.stopWorker();
this.child = undefined;
this.ready = undefined;
this.workerReady = false;
this.stdout = "";
}
}