132 lines
4.1 KiB
TypeScript
132 lines
4.1 KiB
TypeScript
import { existsSync, mkdtempSync, readFileSync, rmSync, writeFileSync } from "node:fs";
|
|
import { tmpdir } from "node:os";
|
|
import { join } from "node:path";
|
|
import { afterEach, expect, test, vi } from "vitest";
|
|
import { PythonLocalNerDetector } from "../src/catalog/local-ner-detector.js";
|
|
|
|
const roots: string[] = [];
|
|
|
|
afterEach(() => {
|
|
vi.unstubAllEnvs();
|
|
for (const root of roots.splice(0)) rmSync(root, { recursive: true, force: true });
|
|
});
|
|
|
|
test("keeps a CPU-only local worker warm and returns sanitized evidence", async () => {
|
|
vi.stubEnv("THT_MODEL_API_KEY", "must-not-reach-worker");
|
|
const root = mkdtempSync(join(tmpdir(), "thothii-local-ner-"));
|
|
roots.push(root);
|
|
const helper = join(root, "fake_ner_worker.py");
|
|
writeFileSync(helper, `
|
|
import json
|
|
import os
|
|
import pathlib
|
|
import sys
|
|
|
|
root = pathlib.Path.cwd()
|
|
root.joinpath("runtime.json").write_text(json.dumps({
|
|
"argv": sys.argv,
|
|
"cuda": os.environ.get("CUDA_VISIBLE_DEVICES"),
|
|
"hip": os.environ.get("HIP_VISIBLE_DEVICES"),
|
|
"offline": os.environ.get("HF_HUB_OFFLINE"),
|
|
"inherited_secret": os.environ.get("THT_MODEL_API_KEY"),
|
|
"pid": os.getpid(),
|
|
}), encoding="utf-8")
|
|
print(json.dumps({"ready": True}), flush=True)
|
|
for line in sys.stdin:
|
|
request = json.loads(line)
|
|
root.joinpath("request.json").write_text(json.dumps(request), encoding="utf-8")
|
|
print(json.dumps({
|
|
"id": request["id"],
|
|
"ok": True,
|
|
"evidence": [{
|
|
"columnId": request["candidates"][0]["columnId"],
|
|
"label": "person",
|
|
"confidence": 0.93,
|
|
}],
|
|
}), flush=True)
|
|
`, "utf8");
|
|
const detector = new PythonLocalNerDetector({
|
|
pythonExecutable: "python3",
|
|
workerScript: helper,
|
|
modelPath: join(root, "pinned-model"),
|
|
cwd: root,
|
|
threads: 2,
|
|
startupTimeoutMs: 5_000,
|
|
});
|
|
const candidate = {
|
|
columnId: "33333333-3333-4333-8333-333333333333",
|
|
text: "Dimesso Mario Rossi",
|
|
};
|
|
|
|
try {
|
|
expect(detector.isReady()).toBe(false);
|
|
await detector.warmup();
|
|
expect(detector.isReady()).toBe(true);
|
|
expect(existsSync(join(root, "request.json"))).toBe(false);
|
|
|
|
await expect(detector.detect(
|
|
[candidate],
|
|
new AbortController().signal,
|
|
Date.now() + 5_000,
|
|
)).resolves.toEqual([{
|
|
columnId: candidate.columnId,
|
|
label: "person",
|
|
confidence: 0.93,
|
|
}]);
|
|
const firstRuntime = JSON.parse(readFileSync(join(root, "runtime.json"), "utf8"));
|
|
expect(firstRuntime).toMatchObject({
|
|
cuda: "",
|
|
hip: "",
|
|
offline: "1",
|
|
inherited_secret: null,
|
|
});
|
|
expect(JSON.stringify(firstRuntime.argv)).not.toContain(candidate.text);
|
|
expect(JSON.parse(readFileSync(join(root, "request.json"), "utf8")).candidates).toEqual([candidate]);
|
|
|
|
await detector.detect([candidate], new AbortController().signal, Date.now() + 5_000);
|
|
const secondRuntime = JSON.parse(readFileSync(join(root, "runtime.json"), "utf8"));
|
|
expect(secondRuntime.pid).toBe(firstRuntime.pid);
|
|
} finally {
|
|
await detector.close();
|
|
}
|
|
});
|
|
|
|
test("bounds worker startup by the caller deadline", async () => {
|
|
const root = mkdtempSync(join(tmpdir(), "thothii-local-ner-deadline-"));
|
|
roots.push(root);
|
|
const helper = join(root, "slow_ner_worker.py");
|
|
writeFileSync(helper, `
|
|
import json
|
|
import sys
|
|
import time
|
|
|
|
time.sleep(2)
|
|
print(json.dumps({"ready": True}), flush=True)
|
|
for line in sys.stdin:
|
|
request = json.loads(line)
|
|
print(json.dumps({"id": request["id"], "ok": True, "evidence": []}), flush=True)
|
|
`, "utf8");
|
|
const detector = new PythonLocalNerDetector({
|
|
pythonExecutable: "python3",
|
|
workerScript: helper,
|
|
modelPath: join(root, "pinned-model"),
|
|
cwd: root,
|
|
startupTimeoutMs: 5_000,
|
|
});
|
|
const startedAt = Date.now();
|
|
|
|
try {
|
|
await expect(detector.detect(
|
|
[{
|
|
columnId: "33333333-3333-4333-8333-333333333333",
|
|
text: "Dimesso Mario Rossi",
|
|
}],
|
|
new AbortController().signal,
|
|
startedAt + 50,
|
|
)).rejects.toThrow("local NER is unavailable");
|
|
expect(Date.now() - startedAt).toBeLessThan(1_000);
|
|
} finally {
|
|
await detector.close();
|
|
}
|
|
});
|