Luigit
repositories / pi-ext

pi-ext

bugabingas pi extensions

owned by admin

extensions/decay/core.ts

Raw
import type { Usage } from "@earendil-works/pi-ai";
import type {
	SessionBeforeCompactEvent,
	SessionEntry,
} from "@earendil-works/pi-coding-agent";
import {
	convertToLlm,
	estimateTokens,
	serializeConversation,
	sessionEntryToContextMessages,
} from "@earendil-works/pi-coding-agent";
import { type Static, Type } from "typebox";
import { Value } from "typebox/value";

export const DECAY_DETAILS_VERSION = 1;
export const CHUNK_TARGET_TOKENS = 6_000;
export const MAX_ATOM_TEXT_CHARS = 600;
export const MAX_ATOM_KEY_CHARS = 120;
export const RECURRENCE_LOG_CAP = 2;
export const MIN_PROTECTED_ALLOCATION = 24;
export const MAX_ATOM_ALLOCATION = 128;
export const SUMMARY_MEMORY_SHARE = 0.75;

export const ATOM_KINDS = [
	"constraint",
	"decision",
	"goal",
	"state",
	"blocker",
	"exact",
	"context",
] as const;
export type AtomKind = (typeof ATOM_KINDS)[number];
export type AtomStatus = "active" | "superseded";
export type ChunkSource = "legacy-summary" | "history" | "split-prefix";
export type DegradedStage = "summarizer";

type Preparation = SessionBeforeCompactEvent["preparation"];
type SourceMessage = Preparation["messagesToSummarize"][number];

export interface SourceChunk {
	chunk: number;
	of: number;
	source: ChunkSource;
	firstPosition: number;
	lastPosition: number;
	text: string;
	estimatedTokens: number;
}

export const ClassifiedAtomSchema = Type.Object(
	{
		key: Type.String({ minLength: 1, maxLength: MAX_ATOM_KEY_CHARS }),
		text: Type.String({ minLength: 1, maxLength: MAX_ATOM_TEXT_CHARS }),
		kind: Type.Union(ATOM_KINDS.map((kind) => Type.Literal(kind))),
		status: Type.Union([Type.Literal("active"), Type.Literal("superseded")]),
	},
	{ additionalProperties: false },
);

export const RecordChunkSchema = Type.Object(
	{
		chunk: Type.Integer({ minimum: 1 }),
		of: Type.Integer({ minimum: 1 }),
		atoms: Type.Array(ClassifiedAtomSchema),
	},
	{ additionalProperties: false },
);

export type ClassifiedAtom = Static<typeof ClassifiedAtomSchema>;
export type RecordChunk = Static<typeof RecordChunkSchema>;

export interface MemoryAtom {
	key: string;
	text: string;
	kind: AtomKind;
	firstSeen: number;
	lastSeen: number;
	seenCount: number;
	status: AtomStatus;
	sourcePosition: number;
}

export interface DecayDetails {
	decay: {
		version: typeof DECAY_DETAILS_VERSION;
		model: string;
		elapsedMs: number;
		chunkCount: number;
		keepRecentTokens: number;
		estimatedKeptTokens?: number;
		degraded?: DegradedStage;
		atoms: MemoryAtom[];
		readFiles: string[];
		modifiedFiles: string[];
	};
}

export interface AllocatedAtom extends MemoryAtom {
	allocation: number;
	weight: number;
	protected: boolean;
	recurring: boolean;
}

export function isRecord(value: unknown): value is Record<string, unknown> {
	return value !== null && typeof value === "object" && !Array.isArray(value);
}

export function parseModelReference(
	reference: string,
): { provider: string; modelId: string } | undefined {
	const separator = reference.indexOf("/");
	if (separator <= 0 || separator === reference.length - 1) return undefined;
	return {
		provider: reference.slice(0, separator),
		modelId: reference.slice(separator + 1),
	};
}

export function createSourceChunks(options: {
	messages: SourceMessage[];
	turnPrefixMessages: SourceMessage[];
	legacySummary?: string;
	targetTokens?: number;
}): SourceChunk[] {
	const targetTokens = options.targetTokens ?? CHUNK_TARGET_TOKENS;
	let position = 0;

	const pending: Omit<SourceChunk, "chunk" | "of">[] = [];
	let current:
		| {
				source: ChunkSource;
				firstPosition: number;
				lastPosition: number;
				texts: string[];
				estimatedTokens: number;
		  }
		| undefined;
	let activitySinceUser = false;
	let pendingToolResults = 0;
	let lastRole: SourceMessage["role"] | undefined;

	const flush = () => {
		if (!current) return;
		pending.push({
			source: current.source,
			firstPosition: current.firstPosition,
			lastPosition: current.lastPosition,
			text: current.texts.join("\n\n"),
			estimatedTokens: current.estimatedTokens,
		});
		current = undefined;
	};

	const appendMessages = (messages: SourceMessage[], source: ChunkSource) => {
		flush();
		activitySinceUser = false;
		pendingToolResults = 0;
		lastRole = undefined;
		for (const message of messages) {
			if (message.role === "user") {
				if (activitySinceUser) flush();
				activitySinceUser = false;
				pendingToolResults = 0;
			} else {
				activitySinceUser = true;
			}
			const serialized = serializeConversation(convertToLlm([message])).trim();
			if (!serialized) continue;
			const text = `[MESSAGE ${position} SOURCE ${source}]\n${serialized}`;
			const tokens = Math.max(1, Math.ceil(text.length / 4));
			const hasToolCalls =
				message.role === "assistant" &&
				message.content.some((item) => item.type === "toolCall");
			if (
				current &&
				current.estimatedTokens + tokens > targetTokens &&
				!(message.role === "toolResult" && pendingToolResults > 0) &&
				!(
					message.role === "assistant" &&
					lastRole === "toolResult" &&
					!hasToolCalls
				)
			)
				flush();
			current ??= {
				source,
				firstPosition: position,
				lastPosition: position,
				texts: [],
				estimatedTokens: 0,
			};
			current.texts.push(text);
			current.lastPosition = position++;
			current.estimatedTokens += tokens;
			if (message.role === "assistant")
				pendingToolResults = message.content.filter(
					(item) => item.type === "toolCall",
				).length;
			else if (message.role === "toolResult" && pendingToolResults > 0)
				pendingToolResults--;
			lastRole = message.role;
		}
	};
	if (options.legacySummary?.trim()) {
		const text = `[PREVIOUS SUMMARY]\n${options.legacySummary.trim()}`;
		pending.push({
			source: "legacy-summary",
			firstPosition: position,
			lastPosition: position++,
			text,
			estimatedTokens: Math.ceil(text.length / 4),
		});
	}
	appendMessages(options.messages, "history");
	appendMessages(options.turnPrefixMessages, "split-prefix");
	flush();

	return pending.map((chunk, index) => ({
		...chunk,
		chunk: index + 1,
		of: pending.length,
	}));
}

export function validateRecordChunk(value: unknown): RecordChunk | undefined {
	return Value.Check(RecordChunkSchema, value)
		? (value as RecordChunk)
		: undefined;
}

export function indexRecords(
	records: readonly RecordChunk[],
	totalChunks: number,
): Map<number, RecordChunk> {
	if (totalChunks < 1) throw new Error("Decay has no source chunks");
	const byChunk = new Map<number, RecordChunk>();
	for (const record of records) {
		if (record.of !== totalChunks)
			throw new Error(
				`Decay classifier reported ${record.of} chunks, expected ${totalChunks}`,
			);
		if (record.chunk < 1 || record.chunk > totalChunks)
			throw new Error(
				`Decay classifier reported invalid chunk ${record.chunk}`,
			);
		const existing = byChunk.get(record.chunk);
		if (existing) {
			if (canonicalRecord(existing) !== canonicalRecord(record))
				throw new Error(
					`Decay classifier reclassified chunk ${record.chunk} with conflicting atoms`,
				);
			continue;
		}
		byChunk.set(record.chunk, record);
	}
	return byChunk;
}

export function missingChunks(
	records: readonly RecordChunk[],
	totalChunks: number,
): number[] {
	const byChunk = indexRecords(records, totalChunks);
	return Array.from({ length: totalChunks }, (_, index) => index + 1).filter(
		(chunk) => !byChunk.has(chunk),
	);
}

export function reconcileRecords(
	records: readonly RecordChunk[],
	totalChunks: number,
): RecordChunk[] {
	const byChunk = indexRecords(records, totalChunks);
	if (byChunk.size !== totalChunks) {
		const missing = Array.from({ length: totalChunks }, (_, index) => index + 1)
			.filter((chunk) => !byChunk.has(chunk))
			.join(", ");
		throw new Error(`Decay classifier omitted chunks: ${missing}`);
	}
	return [...byChunk.values()].sort((left, right) => left.chunk - right.chunk);
}

function canonicalRecord(record: RecordChunk): string {
	return JSON.stringify({
		chunk: record.chunk,
		of: record.of,
		atoms: record.atoms.map((atom) => [
			atom.key,
			atom.text,
			atom.kind,
			atom.status,
		]),
	});
}

export function estimateKeptTokens(
	event: SessionBeforeCompactEvent,
): number | undefined {
	const boundaryIndex = event.branchEntries.findIndex(
		(entry) => entry.id === event.preparation.firstKeptEntryId,
	);
	if (boundaryIndex < 0) return undefined;
	return (
		event.branchEntries
			.slice(boundaryIndex)
			// Pi's projection maps older compaction entries to no context messages;
			// only the newest compaction (not yet appended) contributes a summary.
			.flatMap((entry) =>
				entry.type === "compaction" ? [] : sessionEntryToContextMessages(entry),
			)
			.reduce((total, message) => total + estimateTokens(message), 0)
	);
}

export function findPreviousDecayDetails(
	entries: readonly SessionEntry[],
): DecayDetails["decay"] | undefined {
	for (let index = entries.length - 1; index >= 0; index--) {
		const entry = entries[index];
		if (entry.type !== "compaction") continue;
		return parseDecayDetails(entry.details)?.decay;
	}
	return undefined;
}

export function parseDecayDetails(value: unknown): DecayDetails | undefined {
	if (!isRecord(value) || !isRecord(value.decay)) return undefined;
	const decay = value.decay;
	if (
		decay.version !== DECAY_DETAILS_VERSION ||
		typeof decay.model !== "string" ||
		!Number.isFinite(decay.elapsedMs) ||
		!Number.isInteger(decay.chunkCount) ||
		!Number.isInteger(decay.keepRecentTokens) ||
		("estimatedKeptTokens" in decay &&
			(!Number.isInteger(decay.estimatedKeptTokens) ||
				(decay.estimatedKeptTokens as number) < 0)) ||
		(decay.degraded !== undefined && decay.degraded !== "summarizer") ||
		!Array.isArray(decay.atoms) ||
		!Array.isArray(decay.readFiles) ||
		!Array.isArray(decay.modifiedFiles)
	)
		return undefined;
	const atoms = decay.atoms.filter(isMemoryAtom);
	if (atoms.length !== decay.atoms.length) return undefined;
	if (
		!decay.readFiles.every((path) => typeof path === "string") ||
		!decay.modifiedFiles.every((path) => typeof path === "string")
	)
		return undefined;
	return {
		decay: {
			version: DECAY_DETAILS_VERSION,
			model: decay.model,
			elapsedMs: decay.elapsedMs as number,
			chunkCount: decay.chunkCount as number,
			keepRecentTokens: decay.keepRecentTokens as number,
			...(decay.estimatedKeptTokens === undefined
				? {}
				: { estimatedKeptTokens: decay.estimatedKeptTokens as number }),
			...(decay.degraded === undefined
				? {}
				: { degraded: "summarizer" as const }),
			atoms,
			readFiles: [...(decay.readFiles as string[])],
			modifiedFiles: [...(decay.modifiedFiles as string[])],
		},
	};
}

function isMemoryAtom(value: unknown): value is MemoryAtom {
	if (!isRecord(value)) return false;
	return (
		typeof value.key === "string" &&
		typeof value.text === "string" &&
		ATOM_KINDS.includes(value.kind as AtomKind) &&
		Number.isInteger(value.firstSeen) &&
		Number.isInteger(value.lastSeen) &&
		Number.isInteger(value.seenCount) &&
		(value.status === "active" || value.status === "superseded") &&
		Number.isInteger(value.sourcePosition)
	);
}

export function mergeMemory(
	previous: readonly MemoryAtom[],
	records: readonly RecordChunk[],
	chunks: readonly SourceChunk[],
): MemoryAtom[] {
	const previousMaxEpoch = previous.reduce(
		(maximum, atom) => Math.max(maximum, atom.lastSeen),
		0,
	);
	const currentBase = previousMaxEpoch + 1;
	const currentPosition = new Map<number, number>();
	let position = 0;
	for (const chunk of chunks) {
		if (chunk.source !== "legacy-summary")
			currentPosition.set(chunk.chunk, currentBase + position++);
	}
	const merged = new Map(
		previous.map((atom) => [normalizeAtomKey(atom.key), { ...atom }]),
	);
	const occurrences = new Map<
		string,
		Array<{
			atom: ClassifiedAtom;
			epoch: number;
			chunk: number;
			sourcePosition: number;
		}>
	>();

	for (const record of records) {
		const sourceChunk = chunks[record.chunk - 1];
		const epoch =
			sourceChunk?.source === "legacy-summary"
				? 0
				: (currentPosition.get(record.chunk) ?? currentBase);
		const sourcePosition = sourceChunk?.firstPosition ?? record.chunk - 1;
		const latestInChunk = new Map<string, ClassifiedAtom>();
		for (const atom of record.atoms)
			latestInChunk.set(normalizeAtomKey(atom.key), atom);
		for (const [key, atom] of latestInChunk) {
			const list = occurrences.get(key) ?? [];
			list.push({ atom, epoch, chunk: record.chunk, sourcePosition });
			occurrences.set(key, list);
		}
	}

	for (const [key, found] of occurrences) {
		found.sort(
			(left, right) =>
				left.epoch - right.epoch ||
				left.chunk - right.chunk ||
				left.sourcePosition - right.sourcePosition,
		);
		const latest = found.at(-1);
		if (!latest) continue;
		const prior = merged.get(key);
		const firstSeen = prior
			? Math.min(prior.firstSeen, ...found.map((item) => item.epoch))
			: Math.min(...found.map((item) => item.epoch));
		merged.set(key, {
			key,
			text: latest.atom.text.trim(),
			kind: latest.atom.kind,
			firstSeen,
			lastSeen: Math.max(prior?.lastSeen ?? 0, latest.epoch),
			seenCount: (prior?.seenCount ?? 0) + found.length,
			status: latest.atom.status,
			sourcePosition: latest.sourcePosition,
		});
	}

	return [...merged.values()].sort(
		(left, right) =>
			left.firstSeen - right.firstSeen ||
			left.sourcePosition - right.sourcePosition ||
			left.key.localeCompare(right.key),
	);
}

function normalizeAtomKey(key: string): string {
	return key.trim().toLowerCase().replace(/\s+/g, " ");
}

export function isProtectedKind(kind: AtomKind): boolean {
	return kind !== "context";
}

export function atomWeight(atom: MemoryAtom, newestEpoch: number): number {
	const age = Math.max(0, newestEpoch - atom.lastSeen);
	const recency = 1 / (1 + age);
	const recurrence =
		1 + Math.min(RECURRENCE_LOG_CAP, Math.log2(Math.max(1, atom.seenCount)));
	return recency * recurrence;
}

export function allocateAtoms(
	atoms: readonly MemoryAtom[],
	outputTokenBudget: number,
): AllocatedAtom[] {
	const active = atoms.filter((atom) => atom.status === "active");
	const newestEpoch = active.reduce(
		(maximum, atom) => Math.max(maximum, atom.lastSeen),
		0,
	);
	const memoryBudget = Math.max(
		1,
		Math.floor(outputTokenBudget * SUMMARY_MEMORY_SHARE),
	);
	const result = active.map((atom) => ({
		...atom,
		allocation: 0,
		weight: atomWeight(atom, newestEpoch),
		protected: isProtectedKind(atom.kind),
		recurring: atom.seenCount > 1,
	}));
	const protectedAtoms = result.filter((atom) => atom.protected);
	for (const atom of protectedAtoms) {
		atom.allocation = Math.min(
			MAX_ATOM_ALLOCATION,
			Math.max(MIN_PROTECTED_ALLOCATION, Math.ceil(atom.text.length / 4)),
		);
	}
	const used = protectedAtoms.reduce(
		(total, atom) => total + atom.allocation,
		0,
	);
	const ordinary = result.filter((atom) => !atom.protected);
	const remaining = Math.max(0, memoryBudget - used);
	const totalWeight = ordinary.reduce((total, atom) => total + atom.weight, 0);
	for (const atom of ordinary) {
		atom.allocation = Math.min(
			MAX_ATOM_ALLOCATION,
			Math.floor(remaining * (atom.weight / Math.max(totalWeight, 1))),
		);
	}
	return result
		.filter((atom) => atom.protected || atom.allocation > 0)
		.sort(
			(left, right) =>
				Number(right.protected) - Number(left.protected) ||
				right.weight - left.weight ||
				right.lastSeen - left.lastSeen ||
				left.key.localeCompare(right.key),
		);
}

export function mergeFileLists(
	previous:
		| Pick<DecayDetails["decay"], "readFiles" | "modifiedFiles">
		| undefined,
	fileOps: Preparation["fileOps"],
): { readFiles: string[]; modifiedFiles: string[] } {
	const modified = new Set([
		...(previous?.modifiedFiles ?? []),
		...fileOps.written,
		...fileOps.edited,
	]);
	const read = new Set([...(previous?.readFiles ?? []), ...fileOps.read]);
	return {
		readFiles: [...read].filter((path) => !modified.has(path)).sort(),
		modifiedFiles: [...modified].sort(),
	};
}

const FALLBACK_NOTE =
	"> Decay summarizer unavailable. Memory below is rendered deterministically from classified atoms without prose synthesis.";

export function renderFallbackSummary(atoms: readonly AllocatedAtom[]): string {
	const byKind = (kinds: readonly AtomKind[]): string[] =>
		atoms
			.filter((atom) => kinds.includes(atom.kind))
			.map((atom) => `- ${atom.text.replace(/\s+/g, " ").trim()}`);
	const section = (heading: string, kinds: readonly AtomKind[]): string => {
		const lines = byKind(kinds);
		return lines.length > 0 ? `${heading}\n\n${lines.join("\n")}` : heading;
	};
	return [
		FALLBACK_NOTE,
		section("## Goal", ["goal"]),
		section("## Constraints & Preferences", ["constraint"]),
		"## Progress",
		"### Done",
		section("### In Progress", ["state"]),
		section("### Blocked", ["blocker"]),
		section("## Key Decisions", ["decision"]),
		"## Next Steps",
		section("## Critical Context", ["exact", "context"]),
	].join("\n\n");
}

const TRANSIENT_PATTERN =
	/\b(429|500|502|503|504)\b|rate.?limit|too many requests|overloaded|temporarily unavailable|service unavailable|timeout|timed out|socket hang up|stream (?:ended without completion|closed|disconnected)|network|fetch failed|ECONNRESET|ECONNREFUSED|EPIPE|ETIMEDOUT|EAI_AGAIN|ENOTFOUND/i;

export function isTransientFailure(error: Error): boolean {
	if (error.name === "AbortError") return false;
	return TRANSIENT_PATTERN.test(`${error.name}: ${error.message}`);
}

export function formatFileAppendix(
	readFiles: readonly string[],
	modifiedFiles: readonly string[],
): string {
	const sections: string[] = [];
	if (readFiles.length > 0)
		sections.push(`<read-files>\n${readFiles.join("\n")}\n</read-files>`);
	if (modifiedFiles.length > 0)
		sections.push(
			`<modified-files>\n${modifiedFiles.join("\n")}\n</modified-files>`,
		);
	return sections.length > 0 ? `\n\n${sections.join("\n\n")}` : "";
}

export function combineUsage(...items: readonly Usage[]): Usage {
	const usage: Usage = {
		input: 0,
		output: 0,
		cacheRead: 0,
		cacheWrite: 0,
		totalTokens: 0,
		cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0, total: 0 },
	};
	let hasReasoning = false;
	let hasCacheWrite1h = false;
	for (const item of items) {
		usage.input += item.input;
		usage.output += item.output;
		usage.cacheRead += item.cacheRead;
		usage.cacheWrite += item.cacheWrite;
		usage.totalTokens += item.totalTokens;
		usage.cost.input += item.cost.input;
		usage.cost.output += item.cost.output;
		usage.cost.cacheRead += item.cost.cacheRead;
		usage.cost.cacheWrite += item.cost.cacheWrite;
		usage.cost.total += item.cost.total;
		if (item.reasoning !== undefined) {
			hasReasoning = true;
			usage.reasoning = (usage.reasoning ?? 0) + item.reasoning;
		}
		if (item.cacheWrite1h !== undefined) {
			hasCacheWrite1h = true;
			usage.cacheWrite1h = (usage.cacheWrite1h ?? 0) + item.cacheWrite1h;
		}
	}
	if (!hasReasoning) delete usage.reasoning;
	if (!hasCacheWrite1h) delete usage.cacheWrite1h;
	return usage;
}