Luigit
repositories / pi-ext

pi-ext

bugabingas pi extensions

owned by admin

extensions/strata/server.ts

Raw
import { randomBytes } from "node:crypto";
import { readFile } from "node:fs/promises";
import {
	createServer,
	type IncomingMessage,
	type Server,
	type ServerResponse,
} from "node:http";
import { renderReviewPage } from "./render.js";
import type {
	AskRequest,
	AskThread,
	Draft,
	Feedback,
	Finding,
	Review,
	ReviewServer,
	ServerOptions,
	Snapshot,
} from "./types.js";

const BODY_LIMIT = 64 * 1024;
const MAX_NOTES = 20_000;
const MAX_FINDINGS = 1_000;
const MAX_TEXT = 4_000;
const MAX_QUESTION_BYTES = 4 * 1024;
const MAX_ANSWER_BYTES = 16 * 1024;
const MAX_ID = 128;
const MAX_ASK_EXCHANGES = 128;
const SECURITY_HEADERS = {
	"Cache-Control": "no-store",
	"Content-Security-Policy":
		"default-src 'none'; script-src 'self'; style-src 'self'; connect-src 'self'; img-src 'self'; base-uri 'none'; form-action 'self'; frame-ancestors 'none'",
	"Referrer-Policy": "no-referrer",
	"X-Content-Type-Options": "nosniff",
	"X-Frame-Options": "DENY",
} as const;

type JsonRecord = Record<string, unknown>;

export async function startReviewServer(
	options: ServerOptions,
): Promise<ReviewServer> {
	let review = options.review;
	let revision = 0;
	let threads: AskThread[] = [];
	let closed = false;
	let busy = false;
	const sessionToken = randomBytes(32).toString("base64url");
	const cookieName = `strata_${randomBytes(8).toString("hex")}`;
	const controllers = new Set<AbortController>();
	const requests = new Set<Promise<void>>();
	let closing: Promise<void> | undefined;
	const [script, stylesheet] = await Promise.all([
		readFile(new URL("./assets/app.js", import.meta.url), "utf8"),
		readFile(new URL("./assets/style.css", import.meta.url), "utf8"),
	]);
	const page = renderReviewPage();
	let expectedHost = "";
	let expectedOrigin = "";

	const close = (): Promise<void> => {
		if (closing) return closing;
		closed = true;
		for (const controller of controllers) controller.abort();
		const stopped = new Promise<void>((resolve) =>
			server.close(() => resolve()),
		);
		server.closeAllConnections();
		closing = Promise.allSettled([...requests, stopped]).then(() => {});
		options.onClose?.();
		return closing;
	};

	const runExclusive = async <T>(
		req: IncomingMessage,
		res: ServerResponse,
		work: (signal: AbortSignal) => Promise<T> | T,
	): Promise<T> => {
		if (busy) throw new HttpError(409, "review server is busy", "busy");
		if (closed) throw new HttpError(410, "review server is closed");
		busy = true;
		const controller = new AbortController();
		const abort = () => controller.abort();
		const abortOnPrematureClose = () => {
			if (!res.writableEnded) controller.abort();
		};
		controllers.add(controller);
		req.once("aborted", abort);
		res.once("close", abortOnPrematureClose);
		try {
			const result = await work(controller.signal);
			if (closed || controller.signal.aborted) {
				throw new HttpError(410, "review server is closed");
			}
			return result;
		} finally {
			req.off("aborted", abort);
			res.off("close", abortOnPrematureClose);
			controllers.delete(controller);
			busy = false;
		}
	};

	const responseVersion = (value: JsonRecord): JsonRecord => ({
		...value,
		revision,
		snapshotId: review.snapshot.id,
	});
	const reportFailure = (error: unknown): void => {
		if (error instanceof HttpError && error.code === "stale") return;
		if (
			error instanceof HttpError &&
			["busy", "confirm_required", "stale_client"].includes(error.code ?? "")
		)
			return;
		options.onLifecycle?.({
			type: "error",
			message: error instanceof Error ? error.message : "request failed",
		});
	};

	const handle = async (req: IncomingMessage, res: ServerResponse) => {
		setHeaders(res);
		if (closed) return json(res, 410, { error: "review server is closed" });
		if (req.headers.host !== expectedHost) {
			return json(res, 400, { error: "invalid host" });
		}

		const url = new URL(req.url ?? "/", expectedOrigin);
		if (req.method === "GET" && url.pathname === "/" && url.search) {
			if (
				url.searchParams.size !== 1 ||
				url.searchParams.get("token") !== sessionToken
			) {
				return json(res, 401, { error: "unauthorized" });
			}
			res.writeHead(303, {
				Location: "/",
				"Set-Cookie": `${cookieName}=${sessionToken}; HttpOnly; SameSite=Strict; Path=/`,
			});
			res.end();
			return;
		}

		if (!authenticated(req, cookieName, sessionToken)) {
			return json(res, 401, { error: "unauthorized" });
		}
		const origin = req.headers.origin;
		if (
			(req.method !== "GET" && origin !== expectedOrigin) ||
			(origin !== undefined && origin !== expectedOrigin)
		) {
			return json(res, 403, { error: "invalid origin" });
		}

		if (req.method === "GET" && url.pathname === "/") {
			return text(res, 200, "text/html; charset=utf-8", page);
		}
		if (req.method === "GET" && url.pathname === "/assets/app.js") {
			return text(res, 200, "text/javascript; charset=utf-8", script);
		}
		if (req.method === "GET" && url.pathname === "/assets/style.css") {
			return text(res, 200, "text/css; charset=utf-8", stylesheet);
		}
		if (req.method === "GET" && url.pathname === "/api/state") {
			return json(res, 200, {
				review,
				threads,
				busy,
				revision,
				snapshotId: review.snapshot.id,
			});
		}
		if (req.method !== "POST") {
			return json(res, 404, { error: "not found" });
		}

		if (url.pathname === "/api/close") {
			ensureEmptyObject(await readJson(req));
			json(res, 200, responseVersion({ closed: true }));
			setImmediate(close);
			return;
		}
		if (url.pathname === "/api/draft") {
			const body = await readJson(req);
			const draft = await runExclusive(req, res, () => {
				assertClientVersion(req, revision, review.snapshot.id);
				const next = validateDraft(body, review.snapshot);
				options.save(next);
				review = { ...review, draft: next };
				revision += 1;
				options.onLifecycle?.({ type: "draft", draft: next });
				return next;
			});
			return json(res, 200, responseVersion({ draft }));
		}
		if (url.pathname === "/api/ask") {
			const body = exactObject(await readJson(req), [
				"layerId",
				"hunkId",
				"question",
			]);
			const answer = await runExclusive(req, res, async (signal) => {
				assertClientVersion(req, revision, review.snapshot.id);
				const layerId = nullableString(body.layerId, "layerId", MAX_ID);
				const hunkId = nullableString(body.hunkId, "hunkId", MAX_ID);
				const question = boundedString(
					body.question,
					"question",
					1,
					MAX_TEXT,
				).trim();
				if (!question)
					throw new HttpError(400, "question must not be whitespace");
				assertUtf8Bytes(question, MAX_QUESTION_BYTES, "question", 400);
				const layer = layerId === null ? undefined : findLayer(review, layerId);
				if (layerId !== null && !layer)
					throw new HttpError(400, "unknown layerId");
				const hunk =
					hunkId === null
						? undefined
						: review.snapshot.hunks.find(
								(candidate) => candidate.id === hunkId,
							);
				if (hunkId !== null && !hunk)
					throw new HttpError(400, "unknown hunkId");
				if (
					hunkId !== null &&
					layer &&
					!layer.hunks.some((candidate) => candidate.id === hunkId)
				)
					throw new HttpError(400, "hunkId is outside layer scope");
				const exchangeCount = threads.reduce(
					(total, thread) => total + thread.messages.length,
					0,
				);
				if (exchangeCount >= MAX_ASK_EXCHANGES)
					throw new HttpError(
						409,
						`review already contains ${MAX_ASK_EXCHANGES} Ask exchanges`,
						"history_limit",
					);
				const thread = threads.find(
					(candidate) => candidate.layerId === layerId,
				);
				const request: AskRequest = {
					layerId,
					hunkId,
					question,
					history: (thread?.messages ?? []).map((exchange) => ({
						...exchange,
					})),
				};
				if (!(await options.isCurrent(signal))) {
					options.onLifecycle?.({ type: "stale" });
					throw new HttpError(
						409,
						"snapshot is stale; refresh before asking",
						"stale",
					);
				}
				options.onLifecycle?.({ type: "verified-current" });
				signal.throwIfAborted();
				options.onLifecycle?.({ type: "answering" });
				const next = await options.ask(request, signal);
				signal.throwIfAborted();
				if (typeof next !== "string" || !next.trim())
					throw new HttpError(500, "ask returned an invalid answer");
				assertUtf8Bytes(next, MAX_ANSWER_BYTES, "ask answer", 500);
				if (!(await options.isCurrent(signal))) {
					options.onLifecycle?.({ type: "stale" });
					throw new HttpError(
						409,
						"snapshot became stale while answering; response was discarded",
						"stale",
					);
				}
				options.onLifecycle?.({ type: "verified-current" });
				signal.throwIfAborted();
				const completed = {
					question,
					answer: next,
					hunkId,
				};
				if (thread) thread.messages.push(completed);
				else threads.push({ layerId, messages: [completed] });
				revision += 1;
				options.onLifecycle?.({ type: "ready", review });
				return next;
			}).catch((error: unknown) => {
				reportFailure(error);
				throw error;
			});
			return json(res, 200, responseVersion({ answer, threads }));
		}
		if (url.pathname === "/api/refresh") {
			const body = exactObject(await readJson(req), ["confirm"]);
			if (typeof body.confirm !== "boolean") {
				throw new HttpError(400, "confirm must be a boolean");
			}
			const refreshed = await runExclusive(req, res, async (signal) => {
				assertClientVersion(req, revision, review.snapshot.id);
				const wasCurrent = await options.isCurrent(signal);
				if (wasCurrent) options.onLifecycle?.({ type: "verified-current" });
				else options.onLifecycle?.({ type: "stale" });
				signal.throwIfAborted();
				if (hasDraft(review.draft) && !body.confirm) {
					throw new HttpError(
						409,
						"refresh requires confirmation before discarding the draft",
						"confirm_required",
					);
				}
				if (wasCurrent) options.onLifecycle?.({ type: "refreshing" });
				const next = await options.refresh(signal);
				signal.throwIfAborted();
				const draft = validateDraft(next.draft, next.snapshot);
				review = { ...next, draft };
				threads = [];
				revision += 1;
				options.onLifecycle?.({ type: "ready", review });
				return { wasCurrent };
			}).catch((error: unknown) => {
				reportFailure(error);
				throw error;
			});
			return json(res, 200, responseVersion({ review, threads, ...refreshed }));
		}
		if (url.pathname === "/api/submit") {
			const body = exactObject(await readJson(req), ["draft"]);
			const draft = await runExclusive(req, res, async (signal) => {
				assertClientVersion(req, revision, review.snapshot.id);
				const next = validateDraft(body.draft, review.snapshot);
				options.save(next);
				review = { ...review, draft: next };
				options.onLifecycle?.({ type: "draft", draft: next });
				const current = await options.isCurrent(signal);
				if (current) options.onLifecycle?.({ type: "verified-current" });
				else options.onLifecycle?.({ type: "stale" });
				signal.throwIfAborted();
				if (!current) {
					revision += 1;
					throw new HttpError(
						409,
						"snapshot is stale; refresh before sending feedback",
						"stale",
					);
				}
				const feedback: Feedback = {
					...next,
					snapshotId: review.snapshot.id,
				};
				await options.submit(feedback);
				signal.throwIfAborted();
				options.onLifecycle?.({ type: "feedback-sent" });
				const empty: Draft = { reviewed: [], findings: [], notes: "" };
				options.save(empty);
				review = { ...review, draft: empty };
				revision += 1;
				options.onLifecycle?.({ type: "draft", draft: empty });
				return empty;
			}).catch((error: unknown) => {
				reportFailure(error);
				throw error;
			});
			return json(res, 200, responseVersion({ draft, submitted: true }));
		}
		return json(res, 404, { error: "not found" });
	};

	const server = createServer((req, res) => {
		const request = handle(req, res).catch((error: unknown) => {
			if (res.headersSent || res.destroyed) return;
			const status = error instanceof HttpError ? error.status : 500;
			const body: JsonRecord = {
				error: error instanceof Error ? error.message : "request failed",
			};
			if (error instanceof HttpError && error.code) body.code = error.code;
			if (error instanceof HttpError && error.status === 409) {
				body.revision = revision;
				body.snapshotId = review.snapshot.id;
			}
			json(res, status, body);
		});
		requests.add(request);
		void request.finally(() => requests.delete(request));
	});
	const port = await listen(server);
	expectedHost = `127.0.0.1:${port}`;
	expectedOrigin = `http://${expectedHost}`;

	return {
		url: `${expectedOrigin}/?token=${encodeURIComponent(sessionToken)}`,
		close,
	};
}

export function validateDraft(value: unknown, snapshot: Snapshot): Draft {
	const input = exactObject(value, ["reviewed", "findings", "notes"]);
	if (!Array.isArray(input.reviewed)) {
		throw new HttpError(400, "reviewed must be an array");
	}
	if (!Array.isArray(input.findings)) {
		throw new HttpError(400, "findings must be an array");
	}
	if (input.reviewed.length > snapshot.hunks.length) {
		throw new HttpError(400, "reviewed contains too many entries");
	}
	if (input.findings.length > MAX_FINDINGS) {
		throw new HttpError(400, "findings contains too many entries");
	}
	const hunkIds = new Set(snapshot.hunks.map((hunk) => hunk.id));
	const reviewed = input.reviewed.map((value, index) =>
		boundedString(value, `reviewed[${index}]`, 1, MAX_ID),
	);
	if (new Set(reviewed).size !== reviewed.length) {
		throw new HttpError(400, "reviewed IDs must be unique");
	}
	for (const id of reviewed) {
		if (!hunkIds.has(id))
			throw new HttpError(400, `unknown reviewed hunk: ${id}`);
	}

	const findingIds = new Set<string>();
	const findings = input.findings.map((value, index): Finding => {
		const finding = exactObject(value, [
			"id",
			"hunkId",
			"side",
			"line",
			"severity",
			"text",
		]);
		const id = boundedString(finding.id, `findings[${index}].id`, 1, MAX_ID);
		if (findingIds.has(id))
			throw new HttpError(400, "finding IDs must be unique");
		findingIds.add(id);
		const hunkId = boundedString(
			finding.hunkId,
			`findings[${index}].hunkId`,
			1,
			MAX_ID,
		);
		const hunk = snapshot.hunks.find((candidate) => candidate.id === hunkId);
		if (!hunk) throw new HttpError(400, `unknown finding hunk: ${hunkId}`);
		const side = finding.side;
		if (side !== "old" && side !== "new") {
			throw new HttpError(400, `findings[${index}].side is invalid`);
		}
		if (!Number.isSafeInteger(finding.line) || Number(finding.line) < 1) {
			throw new HttpError(400, `findings[${index}].line is invalid`);
		}
		const line = Number(finding.line);
		const lineExists = hunk.lines.some((candidate) =>
			side === "old" ? candidate.oldLine === line : candidate.newLine === line,
		);
		if (!lineExists) {
			throw new HttpError(400, `finding line does not exist: ${hunkId}`);
		}
		const severity = finding.severity;
		if (
			severity !== "question" &&
			severity !== "minor" &&
			severity !== "major"
		) {
			throw new HttpError(400, `findings[${index}].severity is invalid`);
		}
		const text = boundedString(
			finding.text,
			`findings[${index}].text`,
			1,
			MAX_TEXT,
		);
		return { id, hunkId, side, line, severity, text };
	});
	const notes = boundedString(input.notes, "notes", 0, MAX_NOTES);
	return { reviewed, findings, notes };
}

function findLayer(review: Review, layerId: string) {
	for (const cohort of review.plan.cohorts) {
		const layer = cohort.layers.find((candidate) => candidate.id === layerId);
		if (layer) return layer;
	}
	return undefined;
}

function nullableString(
	value: unknown,
	name: string,
	maximum: number,
): string | null {
	return value === null ? null : boundedString(value, name, 1, maximum);
}

function assertUtf8Bytes(
	value: string,
	maximum: number,
	name: string,
	status: number,
): void {
	if (Buffer.byteLength(value, "utf8") > maximum)
		throw new HttpError(status, `${name} exceeds ${maximum} byte limit`);
}

function hasDraft(draft: Draft): boolean {
	return (
		draft.reviewed.length > 0 ||
		draft.findings.length > 0 ||
		draft.notes.length > 0
	);
}

function assertClientVersion(
	req: IncomingMessage,
	revision: number,
	snapshotId: string,
): void {
	if (
		req.headers["if-match"] !== String(revision) ||
		req.headers["x-strata-snapshot"] !== snapshotId
	) {
		throw new HttpError(
			409,
			"review changed in another tab; reload before continuing",
			"stale_client",
		);
	}
}

function authenticated(
	req: IncomingMessage,
	name: string,
	token: string,
): boolean {
	const cookies = req.headers.cookie?.split(";") ?? [];
	return cookies.some((cookie) => cookie.trim() === `${name}=${token}`);
}

async function readJson(req: IncomingMessage): Promise<unknown> {
	const type = req.headers["content-type"]?.split(";", 1)[0]?.trim();
	if (type !== "application/json") throw new HttpError(415, "expected JSON");
	const declared = Number(req.headers["content-length"] ?? 0);
	if (Number.isFinite(declared) && declared > BODY_LIMIT) {
		throw new HttpError(413, "request body is too large");
	}
	const chunks: Buffer[] = [];
	let size = 0;
	for await (const chunk of req) {
		const buffer = Buffer.isBuffer(chunk) ? chunk : Buffer.from(chunk);
		size += buffer.length;
		if (size > BODY_LIMIT)
			throw new HttpError(413, "request body is too large");
		chunks.push(buffer);
	}
	try {
		return JSON.parse(Buffer.concat(chunks).toString("utf8"));
	} catch {
		throw new HttpError(400, "malformed JSON");
	}
}

function exactObject(value: unknown, keys: string[]): JsonRecord {
	if (!value || typeof value !== "object" || Array.isArray(value)) {
		throw new HttpError(400, "request body must be an object");
	}
	const input = value as JsonRecord;
	const actual = Object.keys(input).sort();
	const expected = [...keys].sort();
	if (
		actual.length !== expected.length ||
		actual.some((key, index) => key !== expected[index])
	) {
		throw new HttpError(400, "request body has unexpected fields");
	}
	return input;
}

function ensureEmptyObject(value: unknown): void {
	exactObject(value, []);
}

function boundedString(
	value: unknown,
	name: string,
	minimum: number,
	maximum: number,
): string {
	if (
		typeof value !== "string" ||
		value.length < minimum ||
		value.length > maximum
	) {
		throw new HttpError(
			400,
			`${name} must be a string between ${minimum} and ${maximum} characters`,
		);
	}
	return value;
}

function setHeaders(res: ServerResponse): void {
	for (const [name, value] of Object.entries(SECURITY_HEADERS)) {
		res.setHeader(name, value);
	}
}

function json(res: ServerResponse, status: number, value: unknown): void {
	res.writeHead(status, { "Content-Type": "application/json; charset=utf-8" });
	res.end(JSON.stringify(value));
}

function text(
	res: ServerResponse,
	status: number,
	contentType: string,
	value: string,
): void {
	res.writeHead(status, { "Content-Type": contentType });
	res.end(value);
}

function listen(server: Server): Promise<number> {
	return new Promise((resolve, reject) => {
		const onError = (error: Error) => reject(error);
		server.once("error", onError);
		server.listen(0, "127.0.0.1", () => {
			server.off("error", onError);
			const address = server.address();
			if (!address || typeof address === "string") {
				reject(new Error("Failed to bind Strata review server"));
				return;
			}
			resolve(address.port);
		});
	});
}

class HttpError extends Error {
	constructor(
		readonly status: number,
		message: string,
		readonly code?: string,
	) {
		super(message);
	}
}