50 lines
1.9 KiB
TypeScript
50 lines
1.9 KiB
TypeScript
import { describe, expect, it } from "vitest";
|
|
import { InferenceWorkerHost, type InferenceWorker } from "../src/background/inference-worker";
|
|
import type { InferenceRequest } from "../src/shared/types";
|
|
|
|
const request: InferenceRequest = { id: "post", text: "text", site: "threads", navigationId: "nav", priority: 0, requestId: "request" };
|
|
class FakeWorker implements InferenceWorker {
|
|
readonly listeners = new Map<string, EventListener[]>();
|
|
readonly requests: Array<{ requests: InferenceRequest[]; modelBaseUrl: string }> = [];
|
|
terminated = false;
|
|
|
|
addEventListener(type: "message" | "error", listener: EventListener): void {
|
|
this.listeners.set(type, [...(this.listeners.get(type) ?? []), listener]);
|
|
}
|
|
|
|
removeEventListener(type: "message" | "error", listener: EventListener): void {
|
|
this.listeners.set(type, (this.listeners.get(type) ?? []).filter((candidate) => candidate !== listener));
|
|
}
|
|
|
|
postMessage(message: { requests: InferenceRequest[]; modelBaseUrl: string }): void {
|
|
this.requests.push(message);
|
|
}
|
|
|
|
terminate(): void { this.terminated = true; }
|
|
|
|
reply(): void {
|
|
const event = new MessageEvent("message", { data: { results: [{ requestId: request.requestId, id: request.id, textHash: "hash", label: "not_toxic", probability: .1, navigationId: request.navigationId }] } });
|
|
this.listeners.get("message")?.forEach((listener) => listener(event));
|
|
}
|
|
}
|
|
|
|
describe("InferenceWorkerHost", () => {
|
|
it("reuses its worker", async () => {
|
|
const workers: FakeWorker[] = [];
|
|
const host = new InferenceWorkerHost(() => {
|
|
const worker = new FakeWorker();
|
|
workers.push(worker);
|
|
return worker;
|
|
}, "chrome-extension://id/models/toxicity/");
|
|
|
|
const first = host.run([request]);
|
|
workers[0]!.reply();
|
|
await first;
|
|
const second = host.run([request]);
|
|
workers[0]!.reply();
|
|
await second;
|
|
expect(workers).toHaveLength(1);
|
|
|
|
expect(workers[0]!.requests).toHaveLength(2);
|
|
});
|
|
});
|