Luigit
repositories / pi-ext

pi-ext

bugabingas pi extensions

owned by admin

extensions/bak/archive.ts

Raw
import { randomUUID } from "node:crypto";
import {
	link,
	mkdir,
	open,
	readdir,
	readFile,
	unlink,
	writeFile,
} from "node:fs/promises";
import { basename, dirname, join, resolve } from "node:path";
import type { SessionInfo } from "@earendil-works/pi-coding-agent";
import { S3Error, type S3Store } from "./s3.js";
import {
	atomicWriteJson,
	type CatalogCache,
	type HostCatalog,
	type HostIdentity,
	type HostIndex,
	newIdentity,
	normalizeAlias,
	readJson,
	type SessionRecord,
	type StatePaths,
	writeJsonExclusive,
} from "./state.js";

const HOST_CATALOG_KEY = "catalog/hosts.json";
const CATALOG_VERSION = 1;
const CAS_ATTEMPTS = 5;
const SESSION_HEADER_LIMIT = 64 * 1024;
const TRANSFER_CONCURRENCY = 6;
const UUID = /^[0-9a-f]{8}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{12}$/i;
const decoder = new TextDecoder();

export interface BackupResult {
	uploaded: number;
	indexed: number;
}

export interface RestoreResult {
	restored: number;
	skipped: number;
	conflicts: number;
}

interface SessionSnapshot {
	record: SessionRecord;
	body: Uint8Array;
	sourcePath: string;
}

export class PartialBackupError extends Error {
	constructor(
		message: string,
		readonly succeededPaths: string[],
		options?: ErrorOptions,
	) {
		super(message, options);
		this.name = "PartialBackupError";
	}
}

class AliasCollisionError extends Error {}

async function forEachConcurrent<T>(
	items: T[],
	worker: (item: T) => Promise<void>,
): Promise<{ failed: false } | { failed: true; error: unknown }> {
	let cursor = 0;
	let failed = false;
	let failure: unknown;
	const run = async () => {
		while (cursor < items.length) {
			const item = items[cursor++];
			if (item === undefined) continue;
			try {
				await worker(item);
			} catch (error) {
				if (!failed) failure = error;
				failed = true;
			}
		}
	};
	await Promise.all(
		Array.from({ length: Math.min(TRANSFER_CONCURRENCY, items.length) }, run),
	);
	return failed ? { failed: true, error: failure } : { failed: false };
}

function hostIndexKey(hostId: string): string {
	return `hosts/${hostId}.json`;
}

function sessionKey(hostId: string, sessionId: string): string {
	return `sessions/${hostId}/${sessionId}.jsonl`;
}

function parseJson<T>(body: Uint8Array, key: string): T {
	try {
		return JSON.parse(decoder.decode(body)) as T;
	} catch {
		throw new Error(`Invalid JSON in ${key}`);
	}
}

function emptyCatalog(): HostCatalog {
	return { version: CATALOG_VERSION, hosts: [] };
}

function emptyIndex(identity: HostIdentity): HostIndex {
	return { version: CATALOG_VERSION, ...identity, sessions: [] };
}

function validIdentity(value: unknown): value is HostIdentity {
	if (!value || typeof value !== "object") return false;
	const identity = value as Partial<HostIdentity>;
	return (
		typeof identity.hostId === "string" &&
		UUID.test(identity.hostId) &&
		typeof identity.alias === "string" &&
		normalizeAlias(identity.alias) === identity.alias
	);
}

function validTimestamp(value: unknown): value is string {
	if (typeof value !== "string") return false;
	try {
		return new Date(value).toISOString() === value;
	} catch {
		return false;
	}
}

function validRecord(value: unknown, hostId: string): value is SessionRecord {
	if (!value || typeof value !== "object") return false;
	const record = value as Partial<SessionRecord>;
	return (
		typeof record.id === "string" &&
		UUID.test(record.id) &&
		record.key === sessionKey(hostId, record.id) &&
		(record.name === undefined || typeof record.name === "string") &&
		typeof record.cwd === "string" &&
		validTimestamp(record.created)
	);
}

function parseCatalog(body: Uint8Array): HostCatalog {
	const catalog = parseJson<HostCatalog>(body, HOST_CATALOG_KEY);
	if (
		catalog.version !== CATALOG_VERSION ||
		!Array.isArray(catalog.hosts) ||
		!catalog.hosts.every(validIdentity)
	)
		throw new Error(`Unsupported ${HOST_CATALOG_KEY}`);
	const ids = new Set(catalog.hosts.map((host) => host.hostId));
	const aliases = new Set(catalog.hosts.map((host) => host.alias));
	if (
		ids.size !== catalog.hosts.length ||
		aliases.size !== catalog.hosts.length
	)
		throw new Error(`Duplicate host in ${HOST_CATALOG_KEY}`);
	return catalog;
}

function parseIndex(
	body: Uint8Array,
	key: string,
	expected: HostIdentity,
): HostIndex {
	const index = parseJson<HostIndex>(body, key);
	if (
		index.version !== CATALOG_VERSION ||
		index.hostId !== expected.hostId ||
		index.alias !== expected.alias ||
		!Array.isArray(index.sessions) ||
		!index.sessions.every((record) => validRecord(record, expected.hostId))
	)
		throw new Error(`Unsupported ${key}`);
	if (
		new Set(index.sessions.map((record) => record.id)).size !==
		index.sessions.length
	)
		throw new Error(`Duplicate session in ${key}`);
	return index;
}

function stableIndex(index: HostIndex): string {
	return JSON.stringify({
		...index,
		sessions: [...index.sessions].sort((a, b) => a.id.localeCompare(b.id)),
	});
}

function validateSessionBody(
	body: Uint8Array,
	expectedId?: string,
): { id: string; timestamp: string; cwd: string } {
	const text = decoder.decode(body);
	const lines = text.split("\n");
	if (lines.at(-1) === "") lines.pop();
	if (lines.length === 0) throw new Error("Empty Pi session");
	type SessionHeaderCandidate = {
		type?: unknown;
		id?: unknown;
		timestamp?: unknown;
		cwd?: unknown;
	};
	let header: SessionHeaderCandidate | undefined;
	for (let index = 0; index < lines.length; index++) {
		const line = lines[index];
		if (!line?.trim()) continue;
		let entry: unknown;
		try {
			entry = JSON.parse(line);
		} catch {
			throw new Error(`Invalid Pi session JSONL at line ${index + 1}`);
		}
		if (
			!entry ||
			typeof entry !== "object" ||
			typeof (entry as { type?: unknown }).type !== "string"
		)
			throw new Error(`Invalid Pi session JSONL at line ${index + 1}`);
		if (index === 0) header = entry as SessionHeaderCandidate;
	}
	if (
		header?.type !== "session" ||
		typeof header.id !== "string" ||
		!UUID.test(header.id) ||
		(expectedId !== undefined && header.id !== expectedId) ||
		!validTimestamp(header.timestamp) ||
		typeof header.cwd !== "string"
	)
		throw new Error("Invalid Pi session header");
	return { id: header.id, timestamp: header.timestamp, cwd: header.cwd };
}

export async function sessionInfoFromPath(path: string): Promise<SessionInfo> {
	const body = new Uint8Array(await readFile(path));
	const header = validateSessionBody(body);
	let name: string | undefined;
	for (const line of decoder.decode(body).split("\n")) {
		if (!line.trim()) continue;
		const entry = JSON.parse(line) as { type?: unknown; name?: unknown };
		if (entry.type === "session_info" && typeof entry.name === "string")
			name = entry.name;
	}
	return {
		name,
		path,
		id: header.id,
		cwd: header.cwd,
		created: new Date(header.timestamp),
		modified: new Date(),
		messageCount: 0,
		firstMessage: "",
		allMessagesText: "",
	};
}

async function snapshot(
	info: SessionInfo,
	hostId: string,
): Promise<SessionSnapshot> {
	const body = new Uint8Array(await readFile(info.path));
	const header = validateSessionBody(body, info.id);
	return {
		body,
		sourcePath: info.path,
		record: {
			id: header.id,
			key: sessionKey(hostId, header.id),
			...(info.name ? { name: info.name } : {}),
			cwd: header.cwd,
			created: header.timestamp,
		},
	};
}

function sameBytes(left: Uint8Array, right: Uint8Array): boolean {
	return (
		left.byteLength === right.byteLength &&
		left.every((value, index) => value === right[index])
	);
}

function startsWithBytes(value: Uint8Array, prefix: Uint8Array): boolean {
	return (
		value.byteLength >= prefix.byteLength &&
		prefix.every((byte, index) => value[index] === byte)
	);
}

async function sameFile(path: string, body: Uint8Array): Promise<boolean> {
	try {
		const local = await readFile(path);
		return local.byteLength === body.byteLength && local.equals(body);
	} catch (error) {
		if ((error as NodeJS.ErrnoException).code === "ENOENT") return false;
		throw error;
	}
}

function isLinkUnsupported(error: unknown): boolean {
	if (!(error instanceof Error && "code" in error)) return false;
	const { code } = error as NodeJS.ErrnoException;
	return (
		code === "EACCES" ||
		code === "EPERM" ||
		code === "ENOSYS" ||
		code === "EOPNOTSUPP" ||
		code === "EXDEV" ||
		code === "EMLINK"
	);
}

async function exclusiveWrite(
	path: string,
	body: Uint8Array,
): Promise<boolean> {
	try {
		await writeFile(path, body, { mode: 0o600, flag: "wx" });
		return true;
	} catch (error) {
		if ((error as NodeJS.ErrnoException).code === "EEXIST") return false;
		throw error;
	}
}

async function publishNoReplace(
	path: string,
	body: Uint8Array,
): Promise<boolean> {
	await mkdir(dirname(path), { recursive: true, mode: 0o700 });
	const temporary = join(
		dirname(path),
		`.${basename(path)}.${randomUUID()}.tmp`,
	);
	try {
		await writeFile(temporary, body, { mode: 0o600, flag: "wx" });
		try {
			await link(temporary, path);
			return true;
		} catch (error) {
			if ((error as NodeJS.ErrnoException).code === "EEXIST") return false;
			// Android SELinux and some filesystems forbid hard links.
			if (!isLinkUnsupported(error)) throw error;
			return exclusiveWrite(path, body);
		}
	} finally {
		await unlink(temporary).catch(() => undefined);
	}
}

function restoredFilename(record: SessionRecord): string {
	const timestamp = record.created.replace(/[:.]/g, "-");
	return `${timestamp}_${record.id}.jsonl`;
}

async function sessionIdFromFile(path: string): Promise<string | undefined> {
	let handle: Awaited<ReturnType<typeof open>> | undefined;
	try {
		handle = await open(path, "r");
		const buffer = Buffer.allocUnsafe(SESSION_HEADER_LIMIT);
		const { bytesRead } = await handle.read(buffer, 0, buffer.length, 0);
		const newline = buffer.subarray(0, bytesRead).indexOf(10);
		if (newline < 0) return undefined;
		const header = JSON.parse(buffer.subarray(0, newline).toString("utf8")) as {
			type?: unknown;
			id?: unknown;
		};
		return header.type === "session" && typeof header.id === "string"
			? header.id
			: undefined;
	} catch {
		return undefined;
	} finally {
		await handle?.close();
	}
}

async function destinationSessions(
	directory: string,
): Promise<Map<string, string[]>> {
	const sessions = new Map<string, string[]>();
	for (const name of await readdir(directory)) {
		if (!name.endsWith(".jsonl")) continue;
		const path = join(directory, name);
		const id = await sessionIdFromFile(path);
		if (!id) continue;
		const paths = sessions.get(id) ?? [];
		paths.push(path);
		sessions.set(id, paths);
	}
	return sessions;
}

export class BakArchive {
	private initializedHostId?: string;

	constructor(
		private readonly store: S3Store,
		private readonly paths: StatePaths,
	) {}

	async identity(): Promise<HostIdentity | undefined> {
		const value = await readJson<unknown>(this.paths.identity);
		if (value === undefined) return undefined;
		if (!validIdentity(value)) throw new Error("Invalid bak identity file");
		return value;
	}

	async cachedCatalog(): Promise<CatalogCache | undefined> {
		return readJson<CatalogCache>(this.paths.cache);
	}

	private async remoteCatalog(signal?: AbortSignal): Promise<{
		catalog: HostCatalog;
		etag?: string;
		exists: boolean;
	}> {
		const object = await this.store.get(HOST_CATALOG_KEY, { signal });
		if (!object) return { catalog: emptyCatalog(), exists: false };
		if (!object.etag)
			throw new Error(`Missing ETag for existing ${HOST_CATALOG_KEY}`);
		return {
			catalog: parseCatalog(object.body),
			etag: object.etag,
			exists: true,
		};
	}

	private async register(
		identity: HostIdentity,
		signal?: AbortSignal,
	): Promise<void> {
		for (let attempt = 0; attempt < CAS_ATTEMPTS; attempt++) {
			const remote = await this.remoteCatalog(signal);
			const collision = remote.catalog.hosts.find(
				(host) =>
					host.alias === identity.alias && host.hostId !== identity.hostId,
			);
			if (collision)
				throw new AliasCollisionError(
					`Alias already exists: ${identity.alias}`,
				);
			const registered = remote.catalog.hosts.find(
				(host) => host.hostId === identity.hostId,
			);
			if (registered?.alias === identity.alias) return;
			const hosts = remote.catalog.hosts.filter(
				(host) => host.hostId !== identity.hostId,
			);
			hosts.push(identity);
			hosts.sort((a, b) => a.alias.localeCompare(b.alias));
			try {
				await this.store.put(
					HOST_CATALOG_KEY,
					`${JSON.stringify({ version: CATALOG_VERSION, hosts })}\n`,
					{
						contentType: "application/json",
						...(remote.exists
							? { ifMatch: remote.etag }
							: { ifNoneMatch: "*" }),
						signal,
					},
				);
				return;
			} catch (error) {
				if (!(error instanceof S3Error) || error.status !== 412) throw error;
			}
		}
		throw new Error("Host catalog changed repeatedly; retry /bak init");
	}

	async ensureInitialized(
		identity: HostIdentity,
		signal?: AbortSignal,
	): Promise<void> {
		if (this.initializedHostId === identity.hostId) return;
		await this.register(identity, signal);
		const key = hostIndexKey(identity.hostId);
		const indexObject = await this.store.get(key, { signal });
		if (!indexObject) {
			try {
				await this.store.put(key, `${stableIndex(emptyIndex(identity))}\n`, {
					contentType: "application/json",
					ifNoneMatch: "*",
					signal,
				});
			} catch (error) {
				if (!(error instanceof S3Error) || error.status !== 412) throw error;
			}
		} else {
			parseIndex(indexObject.body, key, identity);
		}
		this.initializedHostId = identity.hostId;
	}

	async init(aliasInput: string, signal?: AbortSignal): Promise<HostIdentity> {
		const alias = normalizeAlias(aliasInput);
		if (!alias)
			throw new Error(
				"Alias must use 1-64 lowercase letters, digits, ., _, or -",
			);
		let identity = await this.identity();
		if (identity && identity.alias !== alias)
			throw new Error(`Already initialized as ${identity.alias}`);
		let created = false;
		if (!identity) {
			const candidate = newIdentity(alias);
			created = await writeJsonExclusive(this.paths.identity, candidate);
			identity = created ? candidate : await this.identity();
			if (!identity) throw new Error("Could not create bak identity");
			if (identity.alias !== alias)
				throw new Error(`Already initialized as ${identity.alias}`);
		}
		try {
			await this.ensureInitialized(identity, signal);
			return identity;
		} catch (error) {
			if (created && error instanceof AliasCollisionError)
				await unlink(this.paths.identity).catch(() => undefined);
			throw error;
		}
	}

	async refresh(signal?: AbortSignal): Promise<CatalogCache> {
		const { catalog } = await this.remoteCatalog(signal);
		const indexes = (
			await Promise.all(
				catalog.hosts.map(async (host) => {
					const key = hostIndexKey(host.hostId);
					const object = await this.store.get(key, { signal });
					if (!object) return emptyIndex(host);
					return parseIndex(object.body, key, host);
				}),
			)
		).sort((a, b) => a.alias.localeCompare(b.alias));
		const cache = {
			refreshedAt: new Date().toISOString(),
			catalog,
			indexes,
		};
		await atomicWriteJson(this.paths.cache, cache);
		return cache;
	}

	private async putSession(
		snapshot: SessionSnapshot,
		signal?: AbortSignal,
	): Promise<void> {
		for (let attempt = 0; attempt < CAS_ATTEMPTS; attempt++) {
			const current = await this.store.get(snapshot.record.key, { signal });
			if (current) {
				validateSessionBody(current.body, snapshot.record.id);
				if (sameBytes(current.body, snapshot.body)) return;
				if (startsWithBytes(current.body, snapshot.body)) return;
				if (!startsWithBytes(snapshot.body, current.body))
					throw new Error(`Archived session diverged: ${snapshot.record.id}`);
				if (!current.etag)
					throw new Error(`Missing ETag for existing ${snapshot.record.key}`);
			}
			try {
				await this.store.put(snapshot.record.key, snapshot.body, {
					contentType: "application/x-ndjson",
					...(current ? { ifMatch: current.etag } : { ifNoneMatch: "*" }),
					signal,
				});
				return;
			} catch (error) {
				if (!(error instanceof S3Error) || error.status !== 412) throw error;
			}
		}
		throw new Error(`Session changed repeatedly: ${snapshot.record.id}`);
	}

	async backup(
		infos: SessionInfo[],
		onProgress?: (done: number, total: number) => void,
		signal?: AbortSignal,
	): Promise<BackupResult> {
		const identity = await this.identity();
		if (!identity) throw new Error("Run /bak init <alias> first");
		await this.ensureInitialized(identity, signal);
		const snapshots: SessionSnapshot[] = [];
		let done = 0;
		const transferError = await forEachConcurrent(infos, async (info) => {
			signal?.throwIfAborted();
			const next = await snapshot(info, identity.hostId);
			await this.putSession(next, signal);
			snapshots.push(next);
			done++;
			onProgress?.(done, infos.length);
		});
		const key = hostIndexKey(identity.hostId);
		let indexed: number | undefined;
		for (let attempt = 0; attempt < CAS_ATTEMPTS; attempt++) {
			signal?.throwIfAborted();
			const currentObject = await this.store.get(key, { signal });
			if (currentObject && !currentObject.etag)
				throw new Error(`Missing ETag for existing ${key}`);
			const current = currentObject
				? parseIndex(currentObject.body, key, identity)
				: emptyIndex(identity);
			const records = new Map(
				current.sessions.map((record) => [record.id, record]),
			);
			for (const item of snapshots) {
				const previous = records.get(item.record.id);
				records.set(item.record.id, {
					...item.record,
					name: item.record.name ?? previous?.name,
				});
			}
			const index: HostIndex = {
				version: CATALOG_VERSION,
				...identity,
				sessions: [...records.values()],
			};
			if (stableIndex(index) === stableIndex(current)) {
				indexed = index.sessions.length;
				break;
			}
			try {
				await this.store.put(key, `${stableIndex(index)}\n`, {
					contentType: "application/json",
					...(currentObject
						? { ifMatch: currentObject.etag }
						: { ifNoneMatch: "*" }),
					signal,
				});
				indexed = index.sessions.length;
				break;
			} catch (error) {
				if (!(error instanceof S3Error) || error.status !== 412) throw error;
			}
		}
		if (indexed === undefined)
			throw new Error("Host index changed repeatedly; retry backup");
		if (transferError.failed) {
			const message =
				transferError.error instanceof Error
					? transferError.error.message
					: String(transferError.error);
			throw new PartialBackupError(
				message,
				snapshots.map((item) => item.sourcePath),
				{ cause: transferError.error },
			);
		}
		return { uploaded: done, indexed };
	}

	async restore(
		records: SessionRecord[],
		destination: string,
		onProgress?: (done: number, total: number) => void,
		signal?: AbortSignal,
	): Promise<RestoreResult> {
		const targetDirectory = resolve(destination);
		await mkdir(targetDirectory, { recursive: true, mode: 0o700 });
		const localSessions = await destinationSessions(targetDirectory);
		const result = { restored: 0, skipped: 0, conflicts: 0 };
		const groups = new Map<string, SessionRecord[]>();
		for (const record of records) {
			const group = groups.get(record.id) ?? [];
			group.push(record);
			groups.set(record.id, group);
		}
		let done = 0;
		const transferError = await forEachConcurrent(
			[...groups.values()],
			async (group) => {
				for (const record of group) {
					signal?.throwIfAborted();
					const hostId = record.key.split("/")[1];
					if (!hostId || !validRecord(record, hostId))
						throw new Error(`Invalid archived session record: ${record.id}`);
					const object = await this.store.get(record.key, { signal });
					if (!object) throw new Error(`Missing archive object: ${record.key}`);
					validateSessionBody(object.body, record.id);
					const existing = localSessions.get(record.id) ?? [];
					if (existing.length > 0) {
						if (
							(
								await Promise.all(
									existing.map((path) => sameFile(path, object.body)),
								)
							).some(Boolean)
						)
							result.skipped++;
						else result.conflicts++;
					} else {
						const path = resolve(targetDirectory, restoredFilename(record));
						if (dirname(path) !== targetDirectory)
							throw new Error(`Unsafe restore path for ${record.id}`);
						if (await publishNoReplace(path, object.body)) {
							result.restored++;
							localSessions.set(record.id, [path]);
						} else if (await sameFile(path, object.body)) {
							result.skipped++;
							localSessions.set(record.id, [path]);
						} else result.conflicts++;
					}
					done++;
					onProgress?.(done, records.length);
				}
			},
		);
		if (transferError.failed) throw transferError.error;
		return result;
	}
}