Luigit
repositories / pi-ext

pi-ext

bugabingas pi extensions

owned by admin

extensions/cost/index.ts

Raw
/**
 * Cost and token footer segments.
 *
 * Emits into the shared footer "llm" zone:
 * - tokens: cumulative session token and cache usage
 * - cost: turn/session/monthly spend when non-zero
 */

import type { Dirent } from "node:fs";
import { readdir, readFile, stat } from "node:fs/promises";
import { dirname, join } from "node:path";
import type {
	ExtensionAPI,
	ExtensionContext,
	ThemeColor,
} from "@earendil-works/pi-coding-agent";
import { parseSessionEntries } from "@earendil-works/pi-coding-agent";
import { closeDebug, dbg, span } from "./src/debug.ts";
import {
	offerFooterSegment,
	removeFooterSegment,
} from "./src/pi-ext-footer-segment.ts";

// ── Constants ─────────────────────────────────────────────────────────────

const STATUS_ID = "cost";
const CACHE_ICON = "";
const PRICE_UNIT_TOKENS = 1_000_000;
const COUNT_DECIMAL_THRESHOLD = 1_000;
const COUNT_ROUND_THRESHOLD = 10_000;
const MILLION_THRESHOLD = 1_000_000;
const MILLION_DECIMAL_THRESHOLD = 10_000_000;
const LOW_COST_THRESHOLD = 0.01;
const MEDIUM_COST_THRESHOLD = 10;
const HIGH_COST_THRESHOLD = 1_000;
const SESSION_EXT = ".jsonl";
const SESSION_MONTH_RE = /^(\d{4})-(\d{2})/;

const FOOTER = {
	tokens: { id: "tokens", color: "muted", order: 2 },
	cost: { id: "cost", color: "success", order: 3 },
	zone: "llm",
} as const satisfies Record<string, unknown>;

const UPDATE_EVENTS = [
	"agent_start",
	"agent_end",
	"tool_execution_end",
] as const;
const RESET_SIGNATURE_EVENTS = [
	"model_select",
	"session_compact",
	"session_tree",
] as const;

// ── Types ─────────────────────────────────────────────────────────────────

type PriceRow = {
	input: number;
	cacheRead: number;
	output: number;
};

type TokenUsage = {
	input: number;
	output: number;
	cacheRead: number;
	cacheWrite: number;
};

type UsageWithCost = Partial<TokenUsage> & {
	cost?: { total?: number };
};

type AssistantLikeMessage = {
	role?: unknown;
	usage?: UsageWithCost;
	provider?: string;
	model?: string;
};

type Snapshot = {
	usage: TokenUsage;
	cacheHitRate?: number;
	turnCost: number;
	sessionCost: number;
	monthlyCost: number;
};

type SegmentId = "tokens" | "cost";

type MonthlyCostCache = Map<
	string,
	{ mtimeMs: number; size: number; cost: number }
>;

// ── Pricing ───────────────────────────────────────────────────────────────

// Prices per 1M tokens. Used only when providers report $0 cost.
const PRICING: Record<string, Record<string, PriceRow>> = {
	zai: {
		"glm-5.1": { input: 1.4, cacheRead: 0.26, output: 4.4 },
		"glm-5": { input: 1.0, cacheRead: 0.2, output: 4.0 },
		"glm-5-turbo": { input: 0.5, cacheRead: 0.1, output: 2.0 },
		"glm-4.7": { input: 1.0, cacheRead: 0.2, output: 4.0 },
		"glm-4.7-flash": { input: 0.1, cacheRead: 0.02, output: 0.4 },
		"glm-4.7-flashx": { input: 0.1, cacheRead: 0.02, output: 0.4 },
		"glm-4.6": { input: 1.0, cacheRead: 0.2, output: 4.0 },
		"glm-4.5": { input: 1.0, cacheRead: 0.2, output: 4.0 },
		"glm-4.5-air": { input: 0.1, cacheRead: 0.02, output: 0.4 },
		"glm-4.5-flash": { input: 0.1, cacheRead: 0.02, output: 0.4 },
	},
	minimax: {
		"MiniMax-M2.7": { input: 1.0, cacheRead: 0.06, output: 4.0 },
		"MiniMax-M2.7-highspeed": { input: 0.6, cacheRead: 0.06, output: 2.4 },
		"MiniMax-M1": { input: 1.0, cacheRead: 0.06, output: 4.0 },
	},
};

function priceFor(provider: string, model: string): PriceRow | undefined {
	const providerPrices = PRICING[provider];
	if (!providerPrices) return undefined;

	return (
		providerPrices[model] ??
		Object.entries(providerPrices).find(([prefix]) =>
			model.startsWith(prefix),
		)?.[1]
	);
}

function manualTokenCost(
	provider: string,
	model: string,
	usage: TokenUsage,
): number {
	const price = priceFor(provider, model);
	if (!price) return 0;

	return (
		(usage.input * price.input +
			usage.cacheRead * price.cacheRead +
			usage.output * price.output) /
		PRICE_UNIT_TOKENS
	);
}

// ── Message extraction ────────────────────────────────────────────────────

function isAssistantMessageEntry(
	entry: unknown,
): entry is { type: "message"; message: AssistantLikeMessage } {
	const candidate = entry as
		| { type?: unknown; message?: AssistantLikeMessage }
		| undefined;
	return (
		candidate?.type === "message" && candidate.message?.role === "assistant"
	);
}

function normalizeUsage(usage: UsageWithCost | undefined): TokenUsage {
	return {
		input: usage?.input ?? 0,
		output: usage?.output ?? 0,
		cacheRead: usage?.cacheRead ?? 0,
		cacheWrite: usage?.cacheWrite ?? 0,
	};
}

function messageCost(message: AssistantLikeMessage): number {
	const usage = normalizeUsage(message.usage);
	const reported = message.usage?.cost?.total ?? 0;
	if (reported > 0) return reported;

	return manualTokenCost(message.provider ?? "", message.model ?? "", usage);
}

function sumMessageCosts(entries: readonly unknown[]): number {
	let total = 0;
	for (const entry of entries) {
		if (isAssistantMessageEntry(entry)) total += messageCost(entry.message);
	}
	return total;
}

function usageFromEntry(entry: unknown): UsageWithCost | undefined {
	const candidate = entry as
		| {
				type?: unknown;
				message?: { role?: unknown; usage?: UsageWithCost };
				usage?: UsageWithCost;
		  }
		| undefined;
	if (candidate?.type === "message") {
		if (
			candidate.message?.role === "assistant" ||
			candidate.message?.role === "toolResult"
		)
			return candidate.message.usage;
		return undefined;
	}
	if (candidate?.type === "compaction" || candidate?.type === "branch_summary")
		return candidate.usage;
	return undefined;
}

function sessionUsage(entries: readonly unknown[]): TokenUsage {
	const total: TokenUsage = {
		input: 0,
		output: 0,
		cacheRead: 0,
		cacheWrite: 0,
	};
	for (const entry of entries) {
		const usage = normalizeUsage(usageFromEntry(entry));
		total.input += usage.input;
		total.output += usage.output;
		total.cacheRead += usage.cacheRead;
		total.cacheWrite += usage.cacheWrite;
	}
	return total;
}

function latestAssistantCacheHitRate(
	entries: readonly unknown[],
): number | undefined {
	for (let i = entries.length - 1; i >= 0; i--) {
		const entry = entries[i];
		if (!isAssistantMessageEntry(entry) || !entry.message.usage) continue;
		const usage = normalizeUsage(entry.message.usage);
		const prompt = usage.input + usage.cacheRead + usage.cacheWrite;
		return prompt > 0 ? (usage.cacheRead / prompt) * 100 : undefined;
	}
	return undefined;
}

// ── Session file scan ─────────────────────────────────────────────────────

async function safeReadDir(path: string): Promise<Dirent[]> {
	try {
		return await readdir(path, { withFileTypes: true });
	} catch {
		dbg?.("monthly.file", { outcome: "directory_failed" });
		return [];
	}
}

function currentMonthKey(date = new Date()): string {
	return `${date.getFullYear()}-${String(date.getMonth() + 1).padStart(2, "0")}`;
}

function sessionFileMonth(name: string): string | undefined {
	const match = name.match(SESSION_MONTH_RE);
	return match ? `${match[1]}-${match[2]}` : undefined;
}

function isSessionFileForMonth(file: Dirent, month: string): boolean {
	return (
		file.isFile() &&
		file.name.endsWith(SESSION_EXT) &&
		sessionFileMonth(file.name) === month
	);
}

async function sessionFilesForMonth(
	sessionDir: string,
	month: string,
	isCurrent: () => boolean,
): Promise<string[] | undefined> {
	const root = dirname(sessionDir);
	const files: string[] = [];

	for (const entry of await safeReadDir(root)) {
		if (!isCurrent()) return undefined;
		if (entry.isFile() && isSessionFileForMonth(entry, month)) {
			files.push(join(root, entry.name));
			continue;
		}

		if (!entry.isDirectory()) continue;
		const childDir = join(root, entry.name);
		for (const file of await safeReadDir(childDir)) {
			if (!isCurrent()) return undefined;
			if (isSessionFileForMonth(file, month))
				files.push(join(childDir, file.name));
		}
	}

	return files;
}

async function sumFileCosts(filePath: string): Promise<number | undefined> {
	try {
		return sumMessageCosts(
			parseSessionEntries(await readFile(filePath, "utf-8")),
		);
	} catch {
		dbg?.("monthly.file", { outcome: "parse_failed" });
		return undefined;
	}
}

async function calculateMonthlyCost(
	sessionDir: string,
	isCurrent: () => boolean = () => true,
	cache: MonthlyCostCache = new Map(),
): Promise<number | undefined> {
	if (!isCurrent()) return undefined;
	const month = currentMonthKey();
	const files = await sessionFilesForMonth(sessionDir, month, isCurrent);
	if (!files) return undefined;
	const next: MonthlyCostCache = new Map();
	let total = 0;
	for (const file of files) {
		if (!isCurrent()) return undefined;
		try {
			const { mtimeMs, size } = await stat(file);
			if (!isCurrent()) return undefined;
			const cached = cache.get(file);
			const isCached = cached?.mtimeMs === mtimeMs && cached.size === size;
			const cost = isCached ? cached.cost : await sumFileCosts(file);
			if (isCached) dbg?.("monthly.file", { outcome: "cached" });
			else if (cost !== undefined) dbg?.("monthly.file", { outcome: "read" });
			if (cost === undefined) continue;
			next.set(file, { mtimeMs, size, cost });
			total += cost;
		} catch {
			dbg?.("monthly.file", { outcome: "stat_failed" });
			// Deleted or inaccessible files are retried on the next refresh.
		}
	}
	if (!isCurrent() || month !== currentMonthKey()) return undefined;
	cache.clear();
	for (const [file, subtotal] of next) cache.set(file, subtotal);
	return total;
}

// ── Formatting ────────────────────────────────────────────────────────────

/** Human-readable token count: 978, 1.1k, 448k, 1.6M, 122M. */
function formatCount(value: number): string {
	if (value < COUNT_DECIMAL_THRESHOLD) return `${value}`;
	if (value < COUNT_ROUND_THRESHOLD)
		return `${(value / COUNT_DECIMAL_THRESHOLD).toFixed(1)}k`;
	if (value < MILLION_THRESHOLD)
		return `${Math.round(value / COUNT_DECIMAL_THRESHOLD)}k`;
	if (value < MILLION_DECIMAL_THRESHOLD)
		return `${(value / MILLION_THRESHOLD).toFixed(1)}M`;
	return `${Math.round(value / MILLION_THRESHOLD)}M`;
}

function formatCost(value: number): string {
	if (value < LOW_COST_THRESHOLD) return `$${value.toFixed(4)}`;
	if (value < MEDIUM_COST_THRESHOLD) return `$${value.toFixed(2)}`;
	if (value < HIGH_COST_THRESHOLD) return `$${value.toFixed(1)}`;
	return `$${(value / HIGH_COST_THRESHOLD).toFixed(1)}k`;
}

/**
 * Cumulative prompt and completion volume. Cache reads and writes are omitted
 * on purpose: as raw totals they are unreadable and unactionable, while the
 * cache hit rate already carries the only decision-relevant part.
 */
function tokenLabel(usage: TokenUsage, cacheHitRate?: number): string {
	if (
		usage.input === 0 &&
		usage.output === 0 &&
		usage.cacheRead === 0 &&
		usage.cacheWrite === 0
	)
		return "↑0";

	const parts: string[] = [];
	if (usage.input > 0) parts.push(`↑${formatCount(usage.input)}`);
	if (usage.output > 0) parts.push(`↓${formatCount(usage.output)}`);
	if (cacheHitRate !== undefined)
		parts.push(`${CACHE_ICON} ${Math.round(cacheHitRate)}%`);
	return parts.join(" ");
}

function costLabel(snapshot: Snapshot): string | undefined {
	const parts: string[] = [];
	if (snapshot.turnCost > 0) parts.push(formatCost(snapshot.turnCost));
	if (snapshot.sessionCost > 0)
		parts.push(`Σ${formatCost(snapshot.sessionCost)}`);
	if (snapshot.monthlyCost > snapshot.sessionCost)
		parts.push(` ${formatCost(snapshot.monthlyCost)}`);
	return parts.length > 0 ? parts.join(" · ") : undefined;
}

// ── Footer emission ───────────────────────────────────────────────────────

function emitSegment(
	pi: ExtensionAPI,
	id: SegmentId,
	text: string | undefined,
): boolean {
	if (text === undefined) return removeFooterSegment(pi, FOOTER[id].id);
	return offerFooterSegment(pi, {
		id: FOOTER[id].id,
		text,
		color: FOOTER[id].color as ThemeColor,
		zone: FOOTER.zone,
		order: FOOTER[id].order,
	});
}

/** TUI fallback when the footer rejects the segments; glyphs stay readable. */
function statusLabel(snapshot: Snapshot): string {
	return [
		tokenLabel(snapshot.usage, snapshot.cacheHitRate),
		costLabel(snapshot),
	]
		.filter((part): part is string => !!part)
		.join(" · ");
}

/**
 * Non-TUI label: words instead of arrows and private-use icons, because RPC
 * clients render the status line as plain text in an unknown font.
 */
function plainLabel(snapshot: Snapshot): string {
	const { input, output } = snapshot.usage;
	const parts = [`in ${formatCount(input)}`, `out ${formatCount(output)}`];
	if (snapshot.cacheHitRate !== undefined)
		parts.push(`cache hit ${Math.round(snapshot.cacheHitRate)}%`);
	if (snapshot.turnCost > 0)
		parts.push(`turn ${formatCost(snapshot.turnCost)}`);
	if (snapshot.sessionCost > 0)
		parts.push(`session ${formatCost(snapshot.sessionCost)}`);
	if (snapshot.monthlyCost > snapshot.sessionCost)
		parts.push(`month ${formatCost(snapshot.monthlyCost)}`);
	return parts.join(", ");
}

function publishSnapshot(
	pi: ExtensionAPI,
	ctx: ExtensionContext,
	snapshot: Snapshot,
) {
	if (ctx.mode !== "tui") {
		clearSegments(pi);
		ctx.ui.setStatus(STATUS_ID, plainLabel(snapshot));
		return;
	}

	const tokensAccepted = emitSegment(
		pi,
		"tokens",
		tokenLabel(snapshot.usage, snapshot.cacheHitRate),
	);
	const costAccepted = emitSegment(pi, "cost", costLabel(snapshot));
	const accepted = tokensAccepted && costAccepted;

	if (!accepted) clearSegments(pi);
	ctx.ui.setStatus(STATUS_ID, accepted ? undefined : statusLabel(snapshot));
}

function clearSegments(pi: ExtensionAPI) {
	emitSegment(pi, "tokens", undefined);
	emitSegment(pi, "cost", undefined);
}

function clearPublished(pi: ExtensionAPI, ctx?: ExtensionContext) {
	clearSegments(pi);
	if (ctx?.hasUI) ctx.ui.setStatus(STATUS_ID, undefined);
}

// ── Snapshot ──────────────────────────────────────────────────────────────

function buildSnapshot(
	ctx: ExtensionContext,
	previousSessionCost: number,
	monthlyCost: number,
): Snapshot {
	const branch = ctx.sessionManager.getBranch();
	const entries = ctx.sessionManager.getEntries();
	const sessionCost = sumMessageCosts(branch);

	return {
		usage: sessionUsage(entries),
		cacheHitRate: latestAssistantCacheHitRate(entries),
		turnCost: Math.max(0, sessionCost - previousSessionCost),
		sessionCost,
		monthlyCost,
	};
}

function signature(snapshot: Snapshot): string {
	const usage = snapshot.usage;
	return [
		usage.input,
		usage.output,
		usage.cacheRead,
		usage.cacheWrite,
		snapshot.cacheHitRate,
		snapshot.sessionCost,
		snapshot.monthlyCost,
	].join("|");
}

// ── Extension ─────────────────────────────────────────────────────────────

export default function costExtension(pi: ExtensionAPI) {
	let previousSessionCost = 0;
	let monthlyCost = 0;
	let lastSignature = "";
	let hasUI = false;
	let generation = 0;
	const monthlyCache: MonthlyCostCache = new Map();
	let refreshRunning = false;
	let pendingRefresh: { ctx: ExtensionContext; ticket: number } | undefined;

	function requestMonthlyRefresh(ctx: ExtensionContext) {
		if (!hasUI) return;
		pendingRefresh = { ctx, ticket: generation };
		if (refreshRunning) return;
		refreshRunning = true;
		setTimeout(() => {
			void (async () => {
				while (pendingRefresh) {
					const { ctx: refreshCtx, ticket } = pendingRefresh;
					pendingRefresh = undefined;
					const finish = span?.("monthly.refresh");
					try {
						const total = await calculateMonthlyCost(
							refreshCtx.sessionManager.getSessionDir(),
							() => ticket === generation,
							monthlyCache,
						);
						if (total === undefined || ticket !== generation) {
							finish?.("finish", { outcome: "cancelled" });
							continue;
						}
						monthlyCost = total;
						await syncBestEffort(refreshCtx);
						finish?.();
					} catch {
						finish?.("error");
						// Background telemetry must never disturb the agent loop.
					}
				}
				refreshRunning = false;
			})();
		}, 0);
	}

	async function sync(ctx: ExtensionContext) {
		if (!hasUI) return;

		const snapshot = buildSnapshot(ctx, previousSessionCost, monthlyCost);
		previousSessionCost = snapshot.sessionCost;

		const nextSignature = signature(snapshot);
		if (nextSignature === lastSignature) return;
		lastSignature = nextSignature;

		publishSnapshot(pi, ctx, snapshot);
	}

	async function syncBestEffort(ctx: ExtensionContext) {
		try {
			await sync(ctx);
		} catch {
			// Footer telemetry must never disturb the agent loop.
		}
	}

	pi.on("session_start", (_event, ctx) => {
		const ticket = ++generation;
		dbg?.("session.start");
		hasUI = ctx.hasUI;
		if (!hasUI) return;

		previousSessionCost = 0;
		lastSignature = "";
		monthlyCost = 0;
		monthlyCache.clear();

		setTimeout(() => {
			if (ticket !== generation) return;
			void syncBestEffort(ctx);
		}, 0);

		requestMonthlyRefresh(ctx);
	});

	for (const eventName of UPDATE_EVENTS) {
		pi.on(eventName as any, async (_event: any, ctx: any) => {
			await syncBestEffort(ctx);
			if (eventName === "agent_end") requestMonthlyRefresh(ctx);
		});
	}

	for (const eventName of RESET_SIGNATURE_EVENTS) {
		pi.on(eventName as any, async (_event: any, ctx: any) => {
			lastSignature = "";
			await syncBestEffort(ctx);
		});
	}

	pi.on("session_shutdown", async (_event, ctx) => {
		generation++;
		hasUI = false;
		pendingRefresh = undefined;
		monthlyCache.clear();
		clearPublished(pi, ctx);
		dbg?.("session.shutdown");
		closeDebug();
	});
}

export const __test = {
	calculateMonthlyCost,
	costLabel,
	latestAssistantCacheHitRate,
	plainLabel,
	sessionUsage,
	tokenLabel,
};