From d660414ad721d0f4ae625a1852bf0399274b4f71 Mon Sep 17 00:00:00 2001 From: Joey Yakimowich-Payne Date: Fri, 13 Mar 2026 11:41:22 -0600 Subject: [PATCH] feat: add benchmarking, auth, and utility modules CLI benchmark command, threshold auto-tuning, OAuth PKCE auth (same flow as Claude Code), cost tracking, telemetry, and replay. Ultraworked with [Sisyphus](https://github.com/code-yeongyu/oh-my-opencode) Co-authored-by: Sisyphus --- src/mnemosyne/analyzer.py | 665 ++++++++++++++ src/mnemosyne/benchmark.py | 724 +++++++++++++++ src/mnemosyne/benchmark_cli.py | 316 +++++++ src/mnemosyne/cost.py | 576 ++++++++++++ src/mnemosyne/deprecated/__init__.py | 2 + src/mnemosyne/deprecated/proxy.py | 1208 ++++++++++++++++++++++++++ src/mnemosyne/eval.py | 496 +++++++++++ src/mnemosyne/oauth.py | 301 +++++++ src/mnemosyne/replay.py | 326 +++++++ src/mnemosyne/telemetry.py | 292 +++++++ tests/test_benchmark.py | 812 +++++++++++++++++ tests/test_summarization_pipeline.py | 558 ++++++++++++ 12 files changed, 6276 insertions(+) create mode 100644 src/mnemosyne/analyzer.py create mode 100644 src/mnemosyne/benchmark.py create mode 100644 src/mnemosyne/benchmark_cli.py create mode 100644 src/mnemosyne/cost.py create mode 100644 src/mnemosyne/deprecated/__init__.py create mode 100644 src/mnemosyne/deprecated/proxy.py create mode 100644 src/mnemosyne/eval.py create mode 100644 src/mnemosyne/oauth.py create mode 100644 src/mnemosyne/replay.py create mode 100644 src/mnemosyne/telemetry.py create mode 100644 tests/test_benchmark.py create mode 100644 tests/test_summarization_pipeline.py diff --git a/src/mnemosyne/analyzer.py b/src/mnemosyne/analyzer.py new file mode 100644 index 0000000..a04eaf7 --- /dev/null +++ b/src/mnemosyne/analyzer.py @@ -0,0 +1,665 @@ +#!/usr/bin/env python3 +"""System prompt waste analyzer for proxy JSONL logs. + +Decomposes the system prompt into semantic components, tracks static vs +dynamic content across turns, cross-references tool definitions against +actual usage, and detects duplicates. + +The system prompt in Claude Code sessions is a composite of: + - Agent identity (short, fixed preamble) + - Conversation instructions (tool usage, tone, git, etc.) + - Auto memory configuration + - Environment info (platform, shell, model) + - Git status snapshot + - CLAUDE.md project instructions (in system-reminder blocks in messages) + - Skills list (in system-reminder blocks in messages) + - Budget reminders (in system-reminder blocks in messages) + +Tool JSON schemas are sent via the API's `tools` parameter, which the +proxy captures in `total_request_bytes` but doesn't log separately. We +infer the tool definition overhead from the gap between total_request_bytes +and (system_prompt_bytes + messages_bytes). + +Usage: + python -m pichay.analyzer [--json] +""" + +from __future__ import annotations + +import argparse +import hashlib +import json +import re +import sys +from dataclasses import dataclass, field, asdict +from pathlib import Path + +from mnemosyne.eval import parse_proxy_log + + +# --------------------------------------------------------------------------- +# Dataclasses +# --------------------------------------------------------------------------- + +@dataclass +class Component: + """A semantic component of the system prompt or message injection.""" + + name: str + source: str # "system_prompt" or "message_injection" + text: str + bytes: int + block_index: int = 0 # which block in the system prompt list + content_hash: str = "" + + def __post_init__(self): + if not self.content_hash: + self.content_hash = hashlib.sha256( + self.text.encode("utf-8") + ).hexdigest()[:16] + + +@dataclass +class ComponentTrack: + """Tracks a component across turns for static/dynamic analysis.""" + + name: str + hashes: list[str] = field(default_factory=list) + sizes: list[int] = field(default_factory=list) + + @property + def is_static(self) -> bool: + return len(set(self.hashes)) <= 1 + + @property + def total_bytes_sent(self) -> int: + return sum(self.sizes) + + @property + def unique_bytes(self) -> int: + """Bytes of the first occurrence — the only send that mattered.""" + return self.sizes[0] if self.sizes else 0 + + @property + def wasted_bytes(self) -> int: + """Bytes re-sent that were identical to the first send.""" + if not self.is_static or len(self.sizes) <= 1: + return 0 + return sum(self.sizes[1:]) + + @property + def turns_present(self) -> int: + return len(self.hashes) + + @property + def change_count(self) -> int: + """How many times the content changed between consecutive turns.""" + changes = 0 + for i in range(1, len(self.hashes)): + if self.hashes[i] != self.hashes[i - 1]: + changes += 1 + return changes + + +@dataclass +class DuplicateGroup: + """A group of components with identical or near-identical content.""" + + canonical_name: str + members: list[str] = field(default_factory=list) + instance_bytes: int = 0 + total_bytes: int = 0 + + @property + def duplicate_bytes(self) -> int: + """Bytes wasted by sending copies beyond the first.""" + return max(0, self.total_bytes - self.instance_bytes) + + +@dataclass +class ToolUsageReport: + """Cross-reference of defined tools vs actual usage.""" + + defined_tools: list[str] = field(default_factory=list) + used_tools: list[str] = field(default_factory=list) + unused_tools: list[str] = field(default_factory=list) + tool_use_counts: dict[str, int] = field(default_factory=dict) + tool_definition_bytes: int = 0 # inferred from request gap + + +@dataclass +class SessionAnalysis: + """Full analysis of system prompt waste for one proxy log.""" + + proxy_log: str + api_calls: int = 0 + # Total bytes + total_system_prompt_bytes: int = 0 + total_message_injection_bytes: int = 0 + total_tool_definition_bytes: int = 0 + total_request_bytes: int = 0 + # Static vs dynamic + static_bytes: int = 0 + dynamic_bytes: int = 0 + static_wasted_bytes: int = 0 # re-sent identical content + # Duplicates (within a single turn) + duplicate_bytes: int = 0 + duplicate_groups: list[DuplicateGroup] = field(default_factory=list) + # Tools + tool_usage: ToolUsageReport = field(default_factory=ToolUsageReport) + # Per-component tracking + component_tracks: list[ComponentTrack] = field(default_factory=list) + # Per-turn decomposition (first turn only, for display) + sample_decomposition: list[dict] = field(default_factory=list) + + @property + def total_overhead_bytes(self) -> int: + """Total system prompt + injection + tool definition bytes.""" + return ( + self.total_system_prompt_bytes + + self.total_message_injection_bytes + + self.total_tool_definition_bytes + ) + + @property + def static_pct(self) -> float: + total = self.static_bytes + self.dynamic_bytes + if total == 0: + return 0.0 + return self.static_bytes / total + + @property + def unused_tool_bytes(self) -> int: + """Estimated bytes wasted on unused tool definitions.""" + if not self.tool_usage.defined_tools: + return 0 + n_defined = len(self.tool_usage.defined_tools) + n_unused = len(self.tool_usage.unused_tools) + if n_defined == 0: + return 0 + per_tool = self.tool_usage.tool_definition_bytes / n_defined + return int(per_tool * n_unused) + + +# --------------------------------------------------------------------------- +# Component extraction +# --------------------------------------------------------------------------- + +# Section boundaries in the main system prompt text. +_SP_SECTIONS = [ + ("agent_identity", r"^You are a Claude agent"), + ("conversation_instructions", r"^You are an interactive agent"), + ("system_section", r"^# System\b"), + ("doing_tasks", r"^# Doing tasks\b"), + ("executing_actions", r"^# Executing actions with care\b"), + ("using_tools", r"^# Using your tools\b"), + ("tone_and_style", r"^# Tone and style\b"), + ("auto_memory", r"^# auto memory\b"), + ("environment", r"^# Environment\b"), + ("git_status", r"^gitStatus:"), + ("fast_mode", r"^"), +] + +# Patterns for system-reminder injections in messages. +_REMINDER_PATTERNS = [ + ("skills_list", r"skills are available for use"), + ("budget_reminder", r"USD budget:"), + ("claude_md", r"# claudeMd\b"), + ("memory_md", r"# memoryMd\b|MEMORY\.md"), + ("current_date", r"# currentDate\b"), + ("todo_reminder", r"TodoWrite tool hasn't been used"), +] + + +def _extract_sp_components(system_blocks: list[dict]) -> list[Component]: + """Decompose system prompt blocks into semantic components.""" + components = [] + + for block_idx, block in enumerate(system_blocks): + if not isinstance(block, dict): + continue + text = block.get("text", "") + if not text: + continue + + # Try to split the text into sections by heading patterns + sections = _split_into_sections(text) + if sections: + for name, section_text in sections: + components.append(Component( + name=name, + source="system_prompt", + text=section_text, + bytes=len(section_text.encode("utf-8")), + block_index=block_idx, + )) + else: + # Single undivided block + components.append(Component( + name=f"system_block_{block_idx}", + source="system_prompt", + text=text, + bytes=len(text.encode("utf-8")), + block_index=block_idx, + )) + + return components + + +def _split_into_sections(text: str) -> list[tuple[str, str]]: + """Split system prompt text into named sections by heading patterns.""" + # Find all section starts + hits: list[tuple[int, str]] = [] + for name, pattern in _SP_SECTIONS: + m = re.search(pattern, text, re.MULTILINE) + if m: + hits.append((m.start(), name)) + + if not hits: + return [] + + hits.sort(key=lambda x: x[0]) + + # If there's content before the first hit, capture it as preamble + sections = [] + if hits[0][0] > 0: + preamble = text[: hits[0][0]].strip() + if preamble: + sections.append(("preamble", preamble)) + + for i, (start, name) in enumerate(hits): + end = hits[i + 1][0] if i + 1 < len(hits) else len(text) + section_text = text[start:end].strip() + if section_text: + sections.append((name, section_text)) + + return sections + + +def _extract_message_injections( + messages: list[dict], +) -> list[Component]: + """Extract system-reminder injections from message content blocks.""" + components = [] + + for msg in messages: + content = msg.get("content", []) + if not isinstance(content, list): + if isinstance(content, str) and "" in content: + for comp in _parse_reminder_text(content): + components.append(comp) + continue + + for block in content: + if not isinstance(block, dict): + continue + text = block.get("text", block.get("content", "")) + if not isinstance(text, str): + continue + if "" not in text: + continue + for comp in _parse_reminder_text(text): + components.append(comp) + + return components + + +def _parse_reminder_text(text: str) -> list[Component]: + """Parse system-reminder tags out of a text block.""" + components = [] + for match in re.finditer( + r"(.*?)", text, re.DOTALL + ): + inner = match.group(1).strip() + if not inner: + continue + name = _classify_reminder(inner) + components.append(Component( + name=name, + source="message_injection", + text=inner, + bytes=len(match.group(0).encode("utf-8")), + )) + return components + + +def _classify_reminder(text: str) -> str: + """Classify a system-reminder block by its content.""" + for name, pattern in _REMINDER_PATTERNS: + if re.search(pattern, text): + return name + return "unknown_reminder" + + +# --------------------------------------------------------------------------- +# Skill duplicate detection +# --------------------------------------------------------------------------- + +def _extract_skill_entries(text: str) -> list[tuple[str, str]]: + """Parse skill entries from a skills list block. + + Returns (full_name, base_name) pairs. + E.g. ("example-skills:pptx", "pptx"), ("pptx", "pptx") + """ + entries = [] + for match in re.finditer( + r"^- ([\w:.-]+): (.+?)$", text, re.MULTILINE + ): + full_name = match.group(1) + base = full_name.split(":")[-1] if ":" in full_name else full_name + entries.append((full_name, base)) + return entries + + +def _find_duplicate_skills(components: list[Component]) -> list[DuplicateGroup]: + """Find skills that appear multiple times under different prefixes.""" + groups: list[DuplicateGroup] = [] + + for comp in components: + if comp.name != "skills_list": + continue + + entries = _extract_skill_entries(comp.text) + if not entries: + continue + + # Group by base name + by_base: dict[str, list[str]] = {} + for full_name, base in entries: + by_base.setdefault(base, []).append(full_name) + + for base, names in by_base.items(): + if len(names) <= 1: + continue + # Estimate bytes per skill entry (average across the block) + avg_bytes = comp.bytes // max(1, len(entries)) + groups.append(DuplicateGroup( + canonical_name=base, + members=names, + instance_bytes=avg_bytes, + total_bytes=avg_bytes * len(names), + )) + + return groups + + +# --------------------------------------------------------------------------- +# Tool usage tracking +# --------------------------------------------------------------------------- + +# Known Claude Code built-in tools. +_KNOWN_TOOLS = [ + "Agent", "Bash", "Edit", "Glob", "Grep", "Read", "Write", + "WebFetch", "WebSearch", "NotebookEdit", "TodoWrite", "Skill", + "AskUserQuestion", "EnterPlanMode", "ExitPlanMode", "TaskOutput", + "TaskStop", "EnterWorktree", +] + + +def _collect_tool_uses(records: list[dict]) -> dict[str, int]: + """Count tool_use occurrences across all request messages.""" + counts: dict[str, int] = {} + for rec in records: + if rec.get("type") != "request": + continue + for msg in rec.get("messages_full", []): + content = msg.get("content", []) + if not isinstance(content, list): + continue + for block in content: + if isinstance(block, dict) and block.get("type") == "tool_use": + name = block.get("name", "") + counts[name] = counts.get(name, 0) + 1 + return counts + + +def _infer_tool_definition_bytes(records: list[dict]) -> int: + """Infer tool definition overhead from the request size gap. + + total_request_bytes = system_prompt + messages + tools + metadata. + The gap between total and (system + messages) is mostly tool schemas. + """ + gaps = [] + for rec in records: + if rec.get("type") != "request": + continue + total = rec.get("total_request_bytes", 0) + sys_bytes = rec.get("system", {}).get("system_prompt_bytes", 0) + msg_bytes = rec.get("messages", {}).get("messages_total_bytes", 0) + gap = total - sys_bytes - msg_bytes + if gap > 0: + gaps.append(gap) + + if not gaps: + return 0 + + # The gap should be roughly constant (tool schemas don't change). + # Use the median to be robust against outliers. + gaps.sort() + return gaps[len(gaps) // 2] + + +# --------------------------------------------------------------------------- +# Main analysis +# --------------------------------------------------------------------------- + +def analyze_system_prompts(proxy_path: Path) -> SessionAnalysis: + """Analyze system prompt waste across a proxy JSONL log.""" + records = parse_proxy_log(proxy_path) + analysis = SessionAnalysis(proxy_log=str(proxy_path)) + + # Separate request records + requests = [r for r in records if r.get("type") == "request"] + analysis.api_calls = len(requests) + if not requests: + return analysis + + # Tool usage + tool_use_counts = _collect_tool_uses(records) + tool_def_bytes = _infer_tool_definition_bytes(records) + analysis.tool_usage = ToolUsageReport( + defined_tools=list(_KNOWN_TOOLS), + used_tools=sorted(tool_use_counts.keys()), + unused_tools=sorted( + set(_KNOWN_TOOLS) - set(tool_use_counts.keys()) + ), + tool_use_counts=tool_use_counts, + tool_definition_bytes=tool_def_bytes, + ) + analysis.total_tool_definition_bytes = tool_def_bytes * len(requests) + + # Component tracking across turns + tracks: dict[str, ComponentTrack] = {} + + for turn_idx, rec in enumerate(requests): + sp = rec.get("system_prompt_full", []) + msgs = rec.get("messages_full", []) + + # Extract components + sp_components = _extract_sp_components( + sp if isinstance(sp, list) else [] + ) + msg_components = _extract_message_injections(msgs) + + # System prompt bytes + sp_bytes = sum(c.bytes for c in sp_components) + analysis.total_system_prompt_bytes += sp_bytes + + # Message injection bytes + inj_bytes = sum(c.bytes for c in msg_components) + analysis.total_message_injection_bytes += inj_bytes + + analysis.total_request_bytes += rec.get("total_request_bytes", 0) + + # Track each component + all_components = sp_components + msg_components + for comp in all_components: + if comp.name not in tracks: + tracks[comp.name] = ComponentTrack(name=comp.name) + tracks[comp.name].hashes.append(comp.content_hash) + tracks[comp.name].sizes.append(comp.bytes) + + # Duplicate detection (skills within this turn) + if turn_idx == 0: + dupe_groups = _find_duplicate_skills(msg_components) + analysis.duplicate_groups = dupe_groups + analysis.duplicate_bytes = sum( + g.duplicate_bytes for g in dupe_groups + ) + + # Sample decomposition (first turn) + if turn_idx == 0: + analysis.sample_decomposition = [ + {"name": c.name, "source": c.source, "bytes": c.bytes} + for c in all_components + ] + + # Aggregate static/dynamic + for track in tracks.values(): + if track.is_static: + analysis.static_bytes += track.total_bytes_sent + analysis.static_wasted_bytes += track.wasted_bytes + else: + analysis.dynamic_bytes += track.total_bytes_sent + + analysis.component_tracks = sorted( + tracks.values(), key=lambda t: t.total_bytes_sent, reverse=True + ) + + return analysis + + +# --------------------------------------------------------------------------- +# Display +# --------------------------------------------------------------------------- + +def print_analysis(a: SessionAnalysis) -> None: + """Print human-readable analysis.""" + print(f"\n{'=' * 65}") + print("SYSTEM PROMPT WASTE ANALYSIS") + print(f"{'=' * 65}") + print(f"Log: {a.proxy_log}") + print(f"API calls: {a.api_calls}") + + # Component decomposition (first turn) + print(f"\n--- Component Decomposition (Turn 1) ---") + for comp in a.sample_decomposition: + print(f" {comp['name']:<35s} {comp['bytes']:>8,} bytes ({comp['source']})") + + # Static vs dynamic + total_content = a.static_bytes + a.dynamic_bytes + print(f"\n--- Static vs Dynamic ---") + print(f" Static content: {a.static_bytes:>12,} bytes ({_pct(a.static_bytes, total_content)})") + print(f" Dynamic content: {a.dynamic_bytes:>12,} bytes ({_pct(a.dynamic_bytes, total_content)})") + print(f" Static re-send waste: {a.static_wasted_bytes:>12,} bytes") + + # Per-component tracking + print(f"\n--- Per-Component Tracking ---") + print(f" {'Component':<35s} {'Total':>10s} {'Turns':>6s} {'Changes':>8s} {'Static':>7s}") + print(f" {'-'*35} {'-'*10} {'-'*6} {'-'*8} {'-'*7}") + for t in a.component_tracks: + static_label = "yes" if t.is_static else "no" + print( + f" {t.name:<35s} {t.total_bytes_sent:>10,} " + f"{t.turns_present:>6d} {t.change_count:>8d} {static_label:>7s}" + ) + + # Tool usage + tu = a.tool_usage + print(f"\n--- Tool Usage ---") + print(f" Tool definition overhead: {tu.tool_definition_bytes:>10,} bytes/request") + print(f" Total across session: {a.total_tool_definition_bytes:>10,} bytes") + print(f" Defined tools: {len(tu.defined_tools)}") + print(f" Used tools: {len(tu.used_tools)} {tu.used_tools}") + print(f" Unused tools: {len(tu.unused_tools)} {tu.unused_tools}") + if tu.defined_tools: + print(f" Est. unused tool bytes: {a.unused_tool_bytes:>10,} bytes/request") + if tu.tool_use_counts: + print(f"\n Tool call counts:") + for name, count in sorted( + tu.tool_use_counts.items(), key=lambda x: -x[1] + ): + print(f" {name:<25s} {count:>6d}") + + # Duplicates + if a.duplicate_groups: + print(f"\n--- Duplicate Content ---") + print(f" Total duplicate bytes (per turn): {a.duplicate_bytes:>10,}") + for g in sorted(a.duplicate_groups, key=lambda x: -x.duplicate_bytes): + print( + f" {g.canonical_name:<30s} " + f"x{len(g.members)} copies, " + f"{g.duplicate_bytes:>6,} bytes wasted" + ) + for m in g.members: + print(f" - {m}") + + # Summary + print(f"\n--- Session Summary ---") + print(f" Total request bytes: {a.total_request_bytes:>12,}") + print(f" System prompt bytes: {a.total_system_prompt_bytes:>12,} ({_pct(a.total_system_prompt_bytes, a.total_request_bytes)})") + print(f" Message injection bytes: {a.total_message_injection_bytes:>12,} ({_pct(a.total_message_injection_bytes, a.total_request_bytes)})") + print(f" Tool definition bytes: {a.total_tool_definition_bytes:>12,} ({_pct(a.total_tool_definition_bytes, a.total_request_bytes)})") + print(f" Total overhead: {a.total_overhead_bytes:>12,} ({_pct(a.total_overhead_bytes, a.total_request_bytes)})") + print(f" Static re-send waste: {a.static_wasted_bytes:>12,}") + print(f" Duplicate waste (per turn): {a.duplicate_bytes:>12,}") + print(f" Unused tool waste (per req): {a.unused_tool_bytes:>12,}") + + +def _pct(part: int, whole: int) -> str: + if whole == 0: + return "0.0%" + return f"{part / whole:.1%}" + + +# --------------------------------------------------------------------------- +# CLI +# --------------------------------------------------------------------------- + +def main(): + parser = argparse.ArgumentParser( + description="Analyze system prompt waste in proxy JSONL logs" + ) + parser.add_argument( + "proxy_log", + type=Path, + help="Path to proxy JSONL log file", + ) + parser.add_argument( + "--json", action="store_true", + help="Output as JSON", + ) + args = parser.parse_args() + + if not args.proxy_log.exists(): + print(f"Error: {args.proxy_log} not found", file=sys.stderr) + sys.exit(1) + + analysis = analyze_system_prompts(args.proxy_log) + + if args.json: + out = asdict(analysis) + # Trim large fields + out.pop("sample_decomposition", None) + ct = out.pop("component_tracks", None) + if ct: + out["component_tracks"] = [ + { + "name": t["name"], + "total_bytes_sent": sum(t["sizes"]), + "turns_present": len(t["hashes"]), + "is_static": len(set(t["hashes"])) <= 1, + "change_count": sum( + 1 for i in range(1, len(t["hashes"])) + if t["hashes"][i] != t["hashes"][i - 1] + ), + } + for t in ct + ] + print(json.dumps(out, indent=2)) + else: + print_analysis(analysis) + + +if __name__ == "__main__": + main() diff --git a/src/mnemosyne/benchmark.py b/src/mnemosyne/benchmark.py new file mode 100644 index 0000000..789474a --- /dev/null +++ b/src/mnemosyne/benchmark.py @@ -0,0 +1,724 @@ +"""Mnemosyne benchmarking and metrics collection. + +Tracks per-session and aggregate metrics for: + - Token savings (context reduction ratio) + - Cache hit rates + - Fidelity transitions (degradation/upgrade counts by level) + - Admission control (accepted/rejected, scores) + - Micro-fault effectiveness (attempts, successes, tokens saved) + - Segmentation quality (object counts, sizes, types) + - Goal-aware retrieval (topic shifts, reclassifications) + - Entropy-gated faulting (triggers, false positives) + - Latency overhead per subsystem + +All metrics are collected in-memory and can be dumped to JSON or +exposed via the /api/benchmark endpoint. +""" + +from __future__ import annotations + +import time +from dataclasses import dataclass, field +from typing import Any + + +@dataclass +class LatencyTracker: + """Tracks min/max/avg/count for a named operation.""" + + count: int = 0 + total_ms: float = 0.0 + min_ms: float = float("inf") + max_ms: float = 0.0 + + def record(self, ms: float) -> None: + self.count += 1 + self.total_ms += ms + if ms < self.min_ms: + self.min_ms = ms + if ms > self.max_ms: + self.max_ms = ms + + @property + def avg_ms(self) -> float: + return self.total_ms / self.count if self.count > 0 else 0.0 + + def to_dict(self) -> dict[str, Any]: + if self.count == 0: + return {"count": 0, "avg_ms": 0, "min_ms": 0, "max_ms": 0} + return { + "count": self.count, + "avg_ms": round(self.avg_ms, 2), + "min_ms": round(self.min_ms, 2), + "max_ms": round(self.max_ms, 2), + "total_ms": round(self.total_ms, 2), + } + + +@dataclass +class FidelityMetrics: + """Fidelity transition tracking.""" + + degradations: int = 0 + upgrades: int = 0 + # Count by level: key = "L0->L1", "L1->L2", etc. + transitions_by_level: dict[str, int] = field(default_factory=dict) + # Objects at each level at last snapshot + objects_by_level: dict[str, int] = field(default_factory=dict) + # Tokens at each level at last snapshot + tokens_by_level: dict[str, int] = field(default_factory=dict) + + def record_transition(self, from_level: int, to_level: int) -> None: + key = f"L{from_level}->L{to_level}" + self.transitions_by_level[key] = self.transitions_by_level.get(key, 0) + 1 + if to_level > from_level: + self.degradations += 1 + else: + self.upgrades += 1 + + def snapshot_levels( + self, + objects_by_level: dict[int, int], + tokens_by_level: dict[int, int], + ) -> None: + self.objects_by_level = {f"L{k}": v for k, v in objects_by_level.items()} + self.tokens_by_level = {f"L{k}": v for k, v in tokens_by_level.items()} + + def to_dict(self) -> dict[str, Any]: + return { + "degradations": self.degradations, + "upgrades": self.upgrades, + "transitions_by_level": self.transitions_by_level, + "objects_by_level": self.objects_by_level, + "tokens_by_level": self.tokens_by_level, + } + + +@dataclass +class AdmissionMetrics: + """Admission control tracking.""" + + total_evaluated: int = 0 + admitted: int = 0 + rejected: int = 0 + score_sum: float = 0.0 + admitted_score_sum: float = 0.0 + rejected_score_sum: float = 0.0 + # Rejected by type + rejected_by_type: dict[str, int] = field(default_factory=dict) + + def record_decision(self, admitted: bool, score: float, object_type: str) -> None: + self.total_evaluated += 1 + self.score_sum += score + if admitted: + self.admitted += 1 + self.admitted_score_sum += score + else: + self.rejected += 1 + self.rejected_score_sum += score + self.rejected_by_type[object_type] = self.rejected_by_type.get(object_type, 0) + 1 + + @property + def rejection_rate(self) -> float: + return self.rejected / self.total_evaluated if self.total_evaluated > 0 else 0.0 + + @property + def avg_score(self) -> float: + return self.score_sum / self.total_evaluated if self.total_evaluated > 0 else 0.0 + + def to_dict(self) -> dict[str, Any]: + return { + "total_evaluated": self.total_evaluated, + "admitted": self.admitted, + "rejected": self.rejected, + "rejection_rate": round(self.rejection_rate, 4), + "avg_score": round(self.avg_score, 4), + "avg_admitted_score": round( + self.admitted_score_sum / self.admitted if self.admitted > 0 else 0.0, 4 + ), + "avg_rejected_score": round( + self.rejected_score_sum / self.rejected if self.rejected > 0 else 0.0, 4 + ), + "rejected_by_type": self.rejected_by_type, + } + + +@dataclass +class MicroFaultMetrics: + """Micro-fault (memory_query phantom tool) tracking.""" + + attempts: int = 0 + successes: int = 0 # model didn't need full restore after + tokens_saved: int = 0 # tokens avoided vs full restore + tokens_used: int = 0 # tokens of micro-fault answers + + def record_fault(self, tokens_answer: int, tokens_avoided: int, success: bool = True) -> None: + self.attempts += 1 + self.tokens_used += tokens_answer + self.tokens_saved += tokens_avoided + if success: + self.successes += 1 + + @property + def success_rate(self) -> float: + return self.successes / self.attempts if self.attempts > 0 else 0.0 + + @property + def avg_tokens_saved(self) -> float: + return self.tokens_saved / self.attempts if self.attempts > 0 else 0.0 + + def to_dict(self) -> dict[str, Any]: + return { + "attempts": self.attempts, + "successes": self.successes, + "success_rate": round(self.success_rate, 4), + "tokens_saved": self.tokens_saved, + "tokens_used": self.tokens_used, + "avg_tokens_saved": round(self.avg_tokens_saved, 1), + "savings_ratio": round( + self.tokens_saved / (self.tokens_saved + self.tokens_used) + if (self.tokens_saved + self.tokens_used) > 0 + else 0.0, + 4, + ), + } + + +@dataclass +class SegmentationMetrics: + """Object segmentation tracking.""" + + total_objects_created: int = 0 + objects_by_type: dict[str, int] = field(default_factory=dict) + total_tokens_stored: int = 0 + # Size distribution + sizes: list[int] = field(default_factory=list) + + def record_object(self, object_type: str, tokens: int) -> None: + self.total_objects_created += 1 + self.objects_by_type[object_type] = self.objects_by_type.get(object_type, 0) + 1 + self.total_tokens_stored += tokens + self.sizes.append(tokens) + + @property + def avg_object_size(self) -> float: + return sum(self.sizes) / len(self.sizes) if self.sizes else 0.0 + + def to_dict(self) -> dict[str, Any]: + sorted_sizes = sorted(self.sizes) + p50 = sorted_sizes[len(sorted_sizes) // 2] if sorted_sizes else 0 + p90 = sorted_sizes[int(len(sorted_sizes) * 0.9)] if sorted_sizes else 0 + return { + "total_objects_created": self.total_objects_created, + "objects_by_type": self.objects_by_type, + "total_tokens_stored": self.total_tokens_stored, + "avg_object_size_tokens": round(self.avg_object_size, 1), + "p50_object_size_tokens": p50, + "p90_object_size_tokens": p90, + "min_object_size_tokens": sorted_sizes[0] if sorted_sizes else 0, + "max_object_size_tokens": sorted_sizes[-1] if sorted_sizes else 0, + } + + +@dataclass +class GoalMetrics: + """Goal-aware retrieval tracking.""" + + topic_shifts_detected: int = 0 + goal_reclassifications: int = 0 + promotions_triggered: int = 0 + + def record_topic_shift(self) -> None: + self.topic_shifts_detected += 1 + + def record_reclassification(self) -> None: + self.goal_reclassifications += 1 + + def record_promotion(self, count: int = 1) -> None: + self.promotions_triggered += count + + def to_dict(self) -> dict[str, Any]: + return { + "topic_shifts_detected": self.topic_shifts_detected, + "goal_reclassifications": self.goal_reclassifications, + "promotions_triggered": self.promotions_triggered, + } + + +@dataclass +class EntropyMetrics: + """Entropy-gated faulting tracking.""" + + checks: int = 0 + triggers: int = 0 + entities_faulted: int = 0 + + def record_check(self, triggered: bool, entities_count: int = 0) -> None: + self.checks += 1 + if triggered: + self.triggers += 1 + self.entities_faulted += entities_count + + @property + def trigger_rate(self) -> float: + return self.triggers / self.checks if self.checks > 0 else 0.0 + + def to_dict(self) -> dict[str, Any]: + return { + "checks": self.checks, + "triggers": self.triggers, + "trigger_rate": round(self.trigger_rate, 4), + "entities_faulted": self.entities_faulted, + } + + +@dataclass +class TokenMetrics: + """Per-turn token tracking for context reduction measurement.""" + + turns: int = 0 + # Raw token counts per turn + input_tokens_per_turn: list[int] = field(default_factory=list) + # Effective (including cache) per turn + effective_tokens_per_turn: list[int] = field(default_factory=list) + # Cache stats per turn + cache_read_per_turn: list[int] = field(default_factory=list) + cache_create_per_turn: list[int] = field(default_factory=list) + # Payload sizes + incoming_bytes_per_turn: list[int] = field(default_factory=list) + outgoing_bytes_per_turn: list[int] = field(default_factory=list) + + def record_turn( + self, + input_tokens: int, + effective_tokens: int, + cache_read: int, + cache_create: int, + incoming_bytes: int, + outgoing_bytes: int, + ) -> None: + self.turns += 1 + self.input_tokens_per_turn.append(input_tokens) + self.effective_tokens_per_turn.append(effective_tokens) + self.cache_read_per_turn.append(cache_read) + self.cache_create_per_turn.append(cache_create) + self.incoming_bytes_per_turn.append(incoming_bytes) + self.outgoing_bytes_per_turn.append(outgoing_bytes) + + @property + def total_input_tokens(self) -> int: + return sum(self.input_tokens_per_turn) + + @property + def total_effective_tokens(self) -> int: + return sum(self.effective_tokens_per_turn) + + @property + def total_cache_read(self) -> int: + return sum(self.cache_read_per_turn) + + @property + def avg_cache_hit_rate(self) -> float: + total_cache = self.total_cache_read + sum(self.cache_create_per_turn) + return self.total_cache_read / total_cache if total_cache > 0 else 0.0 + + @property + def context_reduction_ratio(self) -> float: + """How much smaller outgoing payloads are vs incoming. + 1.0 = no reduction, 0.5 = halved, 0.2 = 80% reduction. + """ + total_in = sum(self.incoming_bytes_per_turn) + total_out = sum(self.outgoing_bytes_per_turn) + return total_out / total_in if total_in > 0 else 1.0 + + def to_dict(self) -> dict[str, Any]: + return { + "turns": self.turns, + "total_input_tokens": self.total_input_tokens, + "total_effective_tokens": self.total_effective_tokens, + "total_cache_read_tokens": self.total_cache_read, + "avg_cache_hit_rate": round(self.avg_cache_hit_rate, 4), + "context_reduction_ratio": round(self.context_reduction_ratio, 4), + "context_reduction_pct": round((1.0 - self.context_reduction_ratio) * 100, 1), + "avg_input_tokens_per_turn": round( + self.total_input_tokens / self.turns if self.turns > 0 else 0, 1 + ), + "peak_input_tokens": max(self.input_tokens_per_turn) + if self.input_tokens_per_turn + else 0, + "total_incoming_bytes": sum(self.incoming_bytes_per_turn), + "total_outgoing_bytes": sum(self.outgoing_bytes_per_turn), + } + + +class SessionBenchmark: + """Collects all metrics for a single session.""" + + def __init__(self, session_id: str) -> None: + self.session_id = session_id + self.start_time = time.monotonic() + self.tokens = TokenMetrics() + self.fidelity = FidelityMetrics() + self.admission = AdmissionMetrics() + self.micro_faults = MicroFaultMetrics() + self.segmentation = SegmentationMetrics() + self.goals = GoalMetrics() + self.entropy = EntropyMetrics() + self.latency: dict[str, LatencyTracker] = { + "preprocess": LatencyTracker(), + "postprocess": LatencyTracker(), + "segmentation": LatencyTracker(), + "embedding": LatencyTracker(), + "admission": LatencyTracker(), + "fidelity_apply": LatencyTracker(), + "goal_classification": LatencyTracker(), + "entropy_check": LatencyTracker(), + "memory_query": LatencyTracker(), + "helper_llm": LatencyTracker(), + } + + @property + def elapsed_seconds(self) -> float: + return time.monotonic() - self.start_time + + def to_dict(self) -> dict[str, Any]: + return { + "session_id": self.session_id, + "elapsed_seconds": round(self.elapsed_seconds, 1), + "tokens": self.tokens.to_dict(), + "fidelity": self.fidelity.to_dict(), + "admission": self.admission.to_dict(), + "micro_faults": self.micro_faults.to_dict(), + "segmentation": self.segmentation.to_dict(), + "goals": self.goals.to_dict(), + "entropy": self.entropy.to_dict(), + "latency": {k: v.to_dict() for k, v in self.latency.items()}, + } + + +class BenchmarkCollector: + """Global benchmark collector across all sessions.""" + + def __init__(self) -> None: + self._sessions: dict[str, SessionBenchmark] = {} + + def get_session(self, session_id: str) -> SessionBenchmark: + if session_id not in self._sessions: + self._sessions[session_id] = SessionBenchmark(session_id) + return self._sessions[session_id] + + def aggregate(self) -> dict[str, Any]: + """Return aggregate metrics across all sessions.""" + if not self._sessions: + return {"sessions": 0, "message": "No sessions recorded"} + + sessions = list(self._sessions.values()) + n = len(sessions) + + # Aggregate tokens + total_turns = sum(s.tokens.turns for s in sessions) + total_input = sum(s.tokens.total_input_tokens for s in sessions) + total_effective = sum(s.tokens.total_effective_tokens for s in sessions) + total_cache_read = sum(s.tokens.total_cache_read for s in sessions) + total_incoming = sum(sum(s.tokens.incoming_bytes_per_turn) for s in sessions) + total_outgoing = sum(sum(s.tokens.outgoing_bytes_per_turn) for s in sessions) + + # Aggregate fidelity + total_degradations = sum(s.fidelity.degradations for s in sessions) + total_upgrades = sum(s.fidelity.upgrades for s in sessions) + + # Aggregate admission + total_admitted = sum(s.admission.admitted for s in sessions) + total_rejected = sum(s.admission.rejected for s in sessions) + + # Aggregate micro-faults + total_fault_attempts = sum(s.micro_faults.attempts for s in sessions) + total_fault_successes = sum(s.micro_faults.successes for s in sessions) + total_tokens_saved = sum(s.micro_faults.tokens_saved for s in sessions) + + # Aggregate segmentation + total_objects = sum(s.segmentation.total_objects_created for s in sessions) + + # Aggregate latency + latency_agg: dict[str, dict[str, Any]] = {} + for key in sessions[0].latency: + counts = [s.latency[key].count for s in sessions if key in s.latency] + total_count = sum(counts) + if total_count > 0: + total_time = sum(s.latency[key].total_ms for s in sessions if key in s.latency) + min_time = min( + ( + s.latency[key].min_ms + for s in sessions + if key in s.latency and s.latency[key].count > 0 + ), + default=0, + ) + max_time = max( + ( + s.latency[key].max_ms + for s in sessions + if key in s.latency and s.latency[key].count > 0 + ), + default=0, + ) + latency_agg[key] = { + "total_calls": total_count, + "avg_ms": round(total_time / total_count, 2), + "min_ms": round(min_time, 2), + "max_ms": round(max_time, 2), + } + + context_reduction = total_outgoing / total_incoming if total_incoming > 0 else 1.0 + + return { + "sessions": n, + "total_turns": total_turns, + "tokens": { + "total_input_tokens": total_input, + "total_effective_tokens": total_effective, + "total_cache_read_tokens": total_cache_read, + "context_reduction_ratio": round(context_reduction, 4), + "context_reduction_pct": round((1.0 - context_reduction) * 100, 1), + "avg_input_tokens_per_turn": round( + total_input / total_turns if total_turns > 0 else 0, 1 + ), + }, + "fidelity": { + "total_degradations": total_degradations, + "total_upgrades": total_upgrades, + }, + "admission": { + "total_admitted": total_admitted, + "total_rejected": total_rejected, + "rejection_rate": round( + total_rejected / (total_admitted + total_rejected) + if (total_admitted + total_rejected) > 0 + else 0.0, + 4, + ), + }, + "micro_faults": { + "total_attempts": total_fault_attempts, + "total_successes": total_fault_successes, + "success_rate": round( + total_fault_successes / total_fault_attempts + if total_fault_attempts > 0 + else 0.0, + 4, + ), + "total_tokens_saved": total_tokens_saved, + }, + "segmentation": { + "total_objects_created": total_objects, + "avg_objects_per_session": round(total_objects / n, 1), + }, + "latency": latency_agg, + } + + def session_report(self, session_id: str) -> dict[str, Any] | None: + """Return detailed metrics for a specific session.""" + bench = self._sessions.get(session_id) + if bench is None: + return None + return bench.to_dict() + + def all_sessions(self) -> dict[str, dict[str, Any]]: + """Return metrics for all sessions.""" + return {sid: b.to_dict() for sid, b in self._sessions.items()} + + def reset(self) -> None: + """Clear all collected metrics.""" + self._sessions.clear() + + +def suggest_thresholds(metrics: dict[str, Any]) -> dict[str, Any]: + """Analyze aggregate metrics and return suggested threshold adjustments. + + Takes the dict returned by ``BenchmarkCollector.aggregate()`` (or the + ``/api/benchmark`` endpoint for aggregate data) and produces actionable + tuning suggestions. + + Returns a dict keyed by parameter name, each containing: + current_value – inferred current operating point + suggested_value – recommended new value (same as current when OK) + reason – human-readable explanation + status – "ok" | "adjust" + """ + suggestions: dict[str, Any] = {} + + # ── Admission threshold ───────────────────────────────────────── + admission = metrics.get("admission", {}) + rejection_rate = admission.get("rejection_rate", 0.0) + total_evaluated = admission.get("total_admitted", admission.get("admitted", 0)) + admission.get( + "total_rejected", admission.get("rejected", 0) + ) + if total_evaluated == 0: + suggestions["admission_threshold"] = { + "current_value": 0.0, + "suggested_value": 0.0, + "reason": "No admission decisions recorded", + "status": "ok", + } + elif rejection_rate > 0.30: + suggestions["admission_threshold"] = { + "current_value": round(rejection_rate, 4), + "suggested_value": round(rejection_rate - 0.05, 4), + "reason": ( + f"Rejection rate {rejection_rate:.1%} exceeds 30% — " + "consider lowering admission threshold to admit more objects" + ), + "status": "adjust", + } + elif rejection_rate < 0.10: + suggestions["admission_threshold"] = { + "current_value": round(rejection_rate, 4), + "suggested_value": round(rejection_rate + 0.05, 4), + "reason": ( + f"Rejection rate {rejection_rate:.1%} is below 10% — " + "consider raising admission threshold to be more selective" + ), + "status": "adjust", + } + else: + suggestions["admission_threshold"] = { + "current_value": round(rejection_rate, 4), + "suggested_value": round(rejection_rate, 4), + "reason": f"Rejection rate {rejection_rate:.1%} is within healthy range (10-30%)", + "status": "ok", + } + + # ── Entropy sensitivity ───────────────────────────────────────── + # Entropy data may come from per-session detail or be absent in aggregate + entropy = metrics.get("entropy", {}) + trigger_rate = entropy.get("trigger_rate", 0.0) + checks = entropy.get("checks", 0) + if checks > 0: + if trigger_rate > 0.25: + suggestions["entropy_sensitivity"] = { + "current_value": round(trigger_rate, 4), + "suggested_value": round(trigger_rate - 0.05, 4), + "reason": ( + f"Entropy trigger rate {trigger_rate:.1%} exceeds 25% — " + "too many false positives, consider lowering sensitivity" + ), + "status": "adjust", + } + elif trigger_rate < 0.05: + suggestions["entropy_sensitivity"] = { + "current_value": round(trigger_rate, 4), + "suggested_value": round(trigger_rate + 0.05, 4), + "reason": ( + f"Entropy trigger rate {trigger_rate:.1%} is below 5% — " + "consider raising sensitivity to catch more entropy spikes" + ), + "status": "adjust", + } + else: + suggestions["entropy_sensitivity"] = { + "current_value": round(trigger_rate, 4), + "suggested_value": round(trigger_rate, 4), + "reason": f"Entropy trigger rate {trigger_rate:.1%} is within healthy range (5-25%)", + "status": "ok", + } + else: + suggestions["entropy_sensitivity"] = { + "current_value": 0.0, + "suggested_value": 0.0, + "reason": "No entropy checks recorded", + "status": "ok", + } + + # ── Fidelity pressure ─────────────────────────────────────────── + fidelity = metrics.get("fidelity", {}) + degradations = fidelity.get("total_degradations", fidelity.get("degradations", 0)) + upgrades = fidelity.get("total_upgrades", fidelity.get("upgrades", 0)) + if degradations > 0 and upgrades > 0: + ratio = degradations / upgrades + elif degradations > 0: + ratio = float(degradations) # effectively infinite + else: + ratio = 0.0 + + if ratio > 5.0: + suggestions["fidelity_pressure"] = { + "current_value": round(ratio, 2), + "suggested_value": 3.0, + "reason": ( + f"Degradation/upgrade ratio {ratio:.1f}:1 — context is over-pressured, " + "consider raising token cap or lowering admission rate" + ), + "status": "adjust", + } + elif upgrades > 0 and degradations > 0 and (upgrades / degradations) > 5.0: + suggestions["fidelity_pressure"] = { + "current_value": round(ratio, 2), + "suggested_value": 1.0, + "reason": ( + f"Upgrade/degradation ratio {upgrades / degradations:.1f}:1 — " + "system is too conservative, context budget is underutilized" + ), + "status": "adjust", + } + else: + suggestions["fidelity_pressure"] = { + "current_value": round(ratio, 2), + "suggested_value": round(ratio, 2), + "reason": "Fidelity pressure is balanced", + "status": "ok", + } + + # ── Micro-fault quality ───────────────────────────────────────── + micro = metrics.get("micro_faults", {}) + attempts = micro.get("total_attempts", micro.get("attempts", 0)) + success_rate = micro.get("success_rate", 0.0) + if attempts > 0: + if success_rate < 0.50: + suggestions["micro_fault_quality"] = { + "current_value": round(success_rate, 4), + "suggested_value": 0.50, + "reason": ( + f"Micro-fault success rate {success_rate:.1%} is below 50% — " + "faults are failing too often, consider tuning query strategies" + ), + "status": "adjust", + } + else: + suggestions["micro_fault_quality"] = { + "current_value": round(success_rate, 4), + "suggested_value": round(success_rate, 4), + "reason": f"Micro-fault success rate {success_rate:.1%} is healthy", + "status": "ok", + } + else: + suggestions["micro_fault_quality"] = { + "current_value": 0.0, + "suggested_value": 0.0, + "reason": "No micro-fault attempts recorded", + "status": "ok", + } + + return suggestions + + +class Timer: + """Context manager for measuring operation latency. + + Usage: + with Timer(benchmark.latency["preprocess"]) as t: + do_work() + # t.elapsed_ms available after exit + """ + + def __init__(self, tracker: LatencyTracker) -> None: + self._tracker = tracker + self._start: float = 0.0 + self.elapsed_ms: float = 0.0 + + def __enter__(self) -> "Timer": + self._start = time.monotonic() + return self + + def __exit__(self, *_: Any) -> None: + self.elapsed_ms = (time.monotonic() - self._start) * 1000 + self._tracker.record(self.elapsed_ms) diff --git a/src/mnemosyne/benchmark_cli.py b/src/mnemosyne/benchmark_cli.py new file mode 100644 index 0000000..edd88f1 --- /dev/null +++ b/src/mnemosyne/benchmark_cli.py @@ -0,0 +1,316 @@ +"""Benchmark CLI — connects to a running Mnemosyne gateway and displays metrics. + +Usage (via the ``mnemosyne`` entry-point): + mnemosyne --benchmark # aggregate report + mnemosyne --benchmark --session # per-session detail + mnemosyne --benchmark --json-output # raw JSON + mnemosyne --benchmark --port 9090 # custom gateway port +""" + +from __future__ import annotations + +import json +import sys +from typing import Any + +import httpx + +from mnemosyne.benchmark import suggest_thresholds + +# ── Box-drawing constants ─────────────────────────────────────────────── + +_TL = "\u2554" # ╔ +_TR = "\u2557" # ╗ +_BL = "\u255a" # ╚ +_BR = "\u255d" # ╝ +_H = "\u2550" # ═ +_V = "\u2551" # ║ +_ML = "\u2560" # ╠ +_MR = "\u2563" # ╣ + +_WIDTH = 58 # inner width (between ║ … ║) + + +def _box_top() -> str: + return f"{_TL}{_H * _WIDTH}{_TR}" + + +def _box_bottom() -> str: + return f"{_BL}{_H * _WIDTH}{_BR}" + + +def _box_sep() -> str: + return f"{_ML}{_H * _WIDTH}{_MR}" + + +def _box_line(text: str) -> str: + """Pad *text* to _WIDTH and wrap with ║.""" + return f"{_V} {text:<{_WIDTH - 2}} {_V}" + + +def _box_title(text: str) -> str: + """Centre *text* inside the box.""" + return f"{_V}{text:^{_WIDTH}}{_V}" + + +def _fmt_num(n: int | float) -> str: + """Format a number with thousands separators.""" + if isinstance(n, float): + return f"{n:,.1f}" + return f"{n:,}" + + +def _fmt_pct(v: float) -> str: + """Format a 0-1 ratio as a percentage string.""" + return f"{v * 100:.1f}%" + + +# ── Report formatting ────────────────────────────────────────────────── + + +def format_aggregate_report(data: dict[str, Any]) -> str: + """Build a Unicode box-drawing report from aggregate benchmark data.""" + lines: list[str] = [] + + lines.append(_box_top()) + lines.append(_box_title("Mnemosyne Benchmark Report")) + lines.append(_box_sep()) + + sessions = data.get("sessions", 0) + turns = data.get("total_turns", 0) + lines.append(_box_line(f"Sessions: {sessions} Total Turns: {turns}")) + + # ── Token savings ─────────────────────────────────────────────── + tokens = data.get("tokens", {}) + if tokens: + lines.append(_box_sep()) + lines.append(_box_line("TOKEN SAVINGS")) + reduction_pct = tokens.get("context_reduction_pct", 0.0) + avg_input = tokens.get("avg_input_tokens_per_turn", 0) + cache_read = tokens.get("total_cache_read_tokens", 0) + lines.append(_box_line(f" Context Reduction: {reduction_pct:.1f}%")) + lines.append(_box_line(f" Avg Input Tokens/Turn: {_fmt_num(avg_input)}")) + lines.append(_box_line(f" Cache Read Tokens: {_fmt_num(cache_read)}")) + + # ── Fidelity ──────────────────────────────────────────────────── + fidelity = data.get("fidelity", {}) + if fidelity: + lines.append(_box_sep()) + lines.append(_box_line("FIDELITY")) + deg = fidelity.get("total_degradations", fidelity.get("degradations", 0)) + upg = fidelity.get("total_upgrades", fidelity.get("upgrades", 0)) + lines.append(_box_line(f" Degradations: {deg} Upgrades: {upg}")) + + # ── Admission control ─────────────────────────────────────────── + admission = data.get("admission", {}) + if admission: + lines.append(_box_sep()) + lines.append(_box_line("ADMISSION CONTROL")) + admitted = admission.get("total_admitted", admission.get("admitted", 0)) + rejected = admission.get("total_rejected", admission.get("rejected", 0)) + rate = admission.get("rejection_rate", 0.0) + lines.append( + _box_line(f" Admitted: {admitted} Rejected: {rejected} Rate: {_fmt_pct(rate)}") + ) + + # ── Micro-faults ──────────────────────────────────────────────── + micro = data.get("micro_faults", {}) + if micro: + lines.append(_box_sep()) + lines.append(_box_line("MICRO-FAULTS")) + attempts = micro.get("total_attempts", micro.get("attempts", 0)) + successes = micro.get("total_successes", micro.get("successes", 0)) + rate = micro.get("success_rate", 0.0) + saved = micro.get("total_tokens_saved", micro.get("tokens_saved", 0)) + lines.append( + _box_line(f" Attempts: {attempts} Successes: {successes} Rate: {_fmt_pct(rate)}") + ) + lines.append(_box_line(f" Tokens Saved: {_fmt_num(saved)}")) + + # ── Latency ───────────────────────────────────────────────────── + latency = data.get("latency", {}) + if latency: + lines.append(_box_sep()) + lines.append(_box_line("LATENCY")) + # Find longest key for alignment + max_key = max(len(k) for k in latency) if latency else 0 + for key, vals in latency.items(): + avg = vals.get("avg_ms", 0.0) + mn = vals.get("min_ms", 0.0) + mx = vals.get("max_ms", 0.0) + lines.append( + _box_line( + f" {key:<{max_key}}: avg {avg:>6.1f}ms min {mn:>5.1f}ms max {mx:>6.1f}ms" + ) + ) + + # ── Threshold suggestions ─────────────────────────────────────── + suggestions = suggest_thresholds(data) + if suggestions: + lines.append(_box_sep()) + lines.append(_box_line("THRESHOLD SUGGESTIONS")) + for param, info in suggestions.items(): + if info["status"] == "ok": + lines.append(_box_line(f" {param}: OK ({info['reason']})")) + else: + cur = info["current_value"] + sug = info["suggested_value"] + lines.append(_box_line(f" {param}: {cur} -> {sug} ({info['reason']})")) + + lines.append(_box_bottom()) + return "\n".join(lines) + + +def format_session_report(data: dict[str, Any]) -> str: + """Build a Unicode box-drawing report from per-session benchmark data.""" + lines: list[str] = [] + + sid = data.get("session_id", "unknown") + elapsed = data.get("elapsed_seconds", 0) + + lines.append(_box_top()) + lines.append(_box_title(f"Session: {sid}")) + lines.append(_box_sep()) + lines.append(_box_line(f"Elapsed: {elapsed:.1f}s")) + + # Tokens + tokens = data.get("tokens", {}) + if tokens: + lines.append(_box_sep()) + lines.append(_box_line("TOKEN SAVINGS")) + lines.append(_box_line(f" Turns: {tokens.get('turns', 0)}")) + lines.append( + _box_line(f" Context Reduction: {tokens.get('context_reduction_pct', 0.0):.1f}%") + ) + lines.append( + _box_line( + f" Avg Input Tokens/Turn: {_fmt_num(tokens.get('avg_input_tokens_per_turn', 0))}" + ) + ) + lines.append( + _box_line(f" Cache Hit Rate: {_fmt_pct(tokens.get('avg_cache_hit_rate', 0.0))}") + ) + + # Fidelity + fidelity = data.get("fidelity", {}) + if fidelity: + lines.append(_box_sep()) + lines.append(_box_line("FIDELITY")) + lines.append( + _box_line( + f" Degradations: {fidelity.get('degradations', 0)}" + f" Upgrades: {fidelity.get('upgrades', 0)}" + ) + ) + + # Admission + admission = data.get("admission", {}) + if admission: + lines.append(_box_sep()) + lines.append(_box_line("ADMISSION CONTROL")) + lines.append( + _box_line( + f" Admitted: {admission.get('admitted', 0)}" + f" Rejected: {admission.get('rejected', 0)}" + f" Rate: {_fmt_pct(admission.get('rejection_rate', 0.0))}" + ) + ) + + # Micro-faults + micro = data.get("micro_faults", {}) + if micro: + lines.append(_box_sep()) + lines.append(_box_line("MICRO-FAULTS")) + lines.append( + _box_line( + f" Attempts: {micro.get('attempts', 0)}" + f" Successes: {micro.get('successes', 0)}" + f" Rate: {_fmt_pct(micro.get('success_rate', 0.0))}" + ) + ) + lines.append(_box_line(f" Tokens Saved: {_fmt_num(micro.get('tokens_saved', 0))}")) + + # Entropy + entropy = data.get("entropy", {}) + if entropy and entropy.get("checks", 0) > 0: + lines.append(_box_sep()) + lines.append(_box_line("ENTROPY")) + lines.append( + _box_line( + f" Checks: {entropy.get('checks', 0)}" + f" Triggers: {entropy.get('triggers', 0)}" + f" Rate: {_fmt_pct(entropy.get('trigger_rate', 0.0))}" + ) + ) + + # Latency + latency = data.get("latency", {}) + if latency: + lines.append(_box_sep()) + lines.append(_box_line("LATENCY")) + max_key = max(len(k) for k in latency) if latency else 0 + for key, vals in latency.items(): + if isinstance(vals, dict) and vals.get("count", vals.get("total_calls", 0)) > 0: + avg = vals.get("avg_ms", 0.0) + mn = vals.get("min_ms", 0.0) + mx = vals.get("max_ms", 0.0) + lines.append( + _box_line( + f" {key:<{max_key}}: avg {avg:>6.1f}ms min {mn:>5.1f}ms max {mx:>6.1f}ms" + ) + ) + + # Threshold suggestions + suggestions = suggest_thresholds(data) + if suggestions: + lines.append(_box_sep()) + lines.append(_box_line("THRESHOLD SUGGESTIONS")) + for param, info in suggestions.items(): + if info["status"] == "ok": + lines.append(_box_line(f" {param}: OK ({info['reason']})")) + else: + cur = info["current_value"] + sug = info["suggested_value"] + lines.append(_box_line(f" {param}: {cur} -> {sug} ({info['reason']})")) + + lines.append(_box_bottom()) + return "\n".join(lines) + + +# ── CLI entry point ──────────────────────────────────────────────────── + + +def run_benchmark_cli( + port: int, + session_id: str | None = None, + json_output: bool = False, +) -> None: + """Fetch benchmark data from a running gateway and display it.""" + url = f"http://127.0.0.1:{port}/api/benchmark" + if session_id is not None: + url += f"?session_id={session_id}" + + try: + resp = httpx.get(url, timeout=10.0) + resp.raise_for_status() + except httpx.ConnectError: + print( + f"Error: Could not connect to gateway at 127.0.0.1:{port}. Is the gateway running?", + file=sys.stderr, + ) + raise SystemExit(1) + except httpx.HTTPStatusError as exc: + print(f"Error: {exc.response.status_code} — {exc.response.text}", file=sys.stderr) + raise SystemExit(1) + + data: dict[str, Any] = resp.json() + + if json_output: + print(json.dumps(data, indent=2)) + return + + report_type = data.get("type", "aggregate") + if report_type == "session": + print(format_session_report(data)) + else: + print(format_aggregate_report(data)) diff --git a/src/mnemosyne/cost.py b/src/mnemosyne/cost.py new file mode 100644 index 0000000..8054bb2 --- /dev/null +++ b/src/mnemosyne/cost.py @@ -0,0 +1,576 @@ +"""Cost simulation for context paging under the inverted cost model. + +Computes three cost metrics from proxy JSONL logs: + +1. Cumulative token cost: Σ n_t (what the API charges for) +2. Cumulative attention cost: Σ n_t² (proportional to actual compute) +3. Fault-adjusted cost: attention cost including extra inference + passes from page faults at (n_t + |p|)² + +The inverted cost model (Section 6.4 of the paper): keeping is +expensive, faulting is cheap. A page sitting in context for T turns +costs |p| · T tokens of processing. Faulting it back costs one extra +inference pass at the current context size — O(n²), not O(|p|). + +This produces a counter-intuitive policy gradient: + - Low fill: faults cheap → evict aggressively + - High fill: faults expensive → evict conservatively + +Usage: + # Simulate cost on a proxy log (uses actual token counts) + uv run python -m pichay.cost experiments/baseline_run2/logs/proxy_*.jsonl + + # Replay with eviction and compare + uv run python -m pichay.cost --replay --age-threshold 4 \\ + experiments/baseline_run2/logs/proxy_*.jsonl + + # JSON output + uv run python -m pichay.cost --json experiments/*/logs/proxy_*.jsonl +""" + +from __future__ import annotations + +import argparse +import copy +import json +import sys +from dataclasses import dataclass, field, asdict +from pathlib import Path + +from mnemosyne.eval import parse_proxy_log +from mnemosyne.pager import PageStore, compact_messages +from mnemosyne.replay import _apply_evictions + + +@dataclass +class TurnCost: + """Cost metrics for a single API call.""" + + turn: int + timestamp: str = "" + # Context size (effective input tokens) + context_tokens: int = 0 + # Linear cost (what you pay for) + token_cost: int = 0 + cumulative_token_cost: int = 0 + # Quadratic attention cost (proportional to compute) + attention_cost: int = 0 + cumulative_attention_cost: int = 0 + # Eviction/fault info + evictions: int = 0 + faults: int = 0 + fault_tokens: int = 0 # extra tokens from fault inference passes + fault_attention_cost: int = 0 # extra attention cost from faults + + +@dataclass +class CostSummary: + """Aggregate cost metrics for one simulation.""" + + label: str + log_path: str = "" + total_turns: int = 0 + # Linear costs + cumulative_token_cost: int = 0 + # Quadratic costs + cumulative_attention_cost: int = 0 + # Fault overhead + total_fault_attention_cost: int = 0 + total_faults: int = 0 + total_evictions: int = 0 + # Context size stats + max_context_tokens: int = 0 + avg_context_tokens: float = 0.0 + # Per-turn data + turns: list[TurnCost] = field(default_factory=list) + + +@dataclass +class CostComparison: + """Side-by-side comparison of baseline vs managed costs.""" + + baseline: CostSummary + managed: CostSummary + + @property + def token_savings_pct(self) -> float: + if self.baseline.cumulative_token_cost == 0: + return 0.0 + saved = self.baseline.cumulative_token_cost - self.managed.cumulative_token_cost + return saved / self.baseline.cumulative_token_cost + + @property + def attention_savings_pct(self) -> float: + if self.baseline.cumulative_attention_cost == 0: + return 0.0 + saved = self.baseline.cumulative_attention_cost - self.managed.cumulative_attention_cost + return saved / self.baseline.cumulative_attention_cost + + @property + def net_attention_savings_pct(self) -> float: + """Attention savings after accounting for fault costs.""" + if self.baseline.cumulative_attention_cost == 0: + return 0.0 + managed_total = ( + self.managed.cumulative_attention_cost + + self.managed.total_fault_attention_cost + ) + saved = self.baseline.cumulative_attention_cost - managed_total + return saved / self.baseline.cumulative_attention_cost + + +def _effective_input(usage: dict) -> int: + """Extract effective input tokens from a usage dict.""" + return ( + usage.get("input_tokens", 0) + + usage.get("cache_creation_input_tokens", 0) + + usage.get("cache_read_input_tokens", 0) + ) + + +def _detect_log_format(records: list[dict]) -> str: + """Detect whether records are proxy format or native Claude Code format. + + Proxy format: {"type": "request", ...} / {"type": "response_stream", ...} + Native format: {"type": "assistant", "message": {"usage": ...}, ...} + """ + for rec in records[:10]: + if rec.get("type") in ("request", "response_stream", "response"): + return "proxy" + if rec.get("type") == "assistant" and "message" in rec: + return "native" + return "unknown" + + +def _extract_native_turns(records: list[dict]) -> list[tuple[str, int]]: + """Extract (timestamp, effective_input_tokens) from native Claude Code logs. + + Native logs have one record per assistant message, each with usage + data embedded in message.usage. We extract the effective input + tokens (input + cache_creation + cache_read) from each. + """ + turns = [] + for rec in records: + if rec.get("type") != "assistant": + continue + msg = rec.get("message", {}) + if not isinstance(msg, dict): + continue + usage = msg.get("usage", {}) + if not usage: + continue + n = _effective_input(usage) + if n > 0: + ts = rec.get("timestamp", "") + turns.append((ts, n)) + return turns + + +def compute_baseline_cost(path: Path, label: str = "") -> CostSummary: + """Compute cost metrics from a proxy or native Claude Code log. + + Accepts both formats: + - Proxy JSONL: request/response pairs with usage in response + - Native Claude Code JSONL: assistant messages with usage in message + """ + records = parse_proxy_log(path) + if not label: + label = path.stem + + fmt = _detect_log_format(records) + + summary = CostSummary(label=label, log_path=str(path)) + cum_token = 0 + cum_attention = 0 + turn_num = 0 + context_sizes = [] + + if fmt == "native": + turn_data = _extract_native_turns(records) + for ts, n in turn_data: + turn_num += 1 + cum_token += n + attention = n * n + cum_attention += attention + context_sizes.append(n) + + turn = TurnCost( + turn=turn_num, + timestamp=ts, + context_tokens=n, + token_cost=n, + cumulative_token_cost=cum_token, + attention_cost=attention, + cumulative_attention_cost=cum_attention, + ) + summary.turns.append(turn) + else: + # Proxy format: pair requests with responses + pending_request = None + for rec in records: + if rec.get("type") == "request": + pending_request = rec + continue + + if rec.get("type") in ("response_stream", "response") and pending_request: + turn_num += 1 + usage = rec.get("usage", {}) + n = _effective_input(usage) + + cum_token += n + attention = n * n + cum_attention += attention + context_sizes.append(n) + + turn = TurnCost( + turn=turn_num, + timestamp=rec.get("timestamp", ""), + context_tokens=n, + token_cost=n, + cumulative_token_cost=cum_token, + attention_cost=attention, + cumulative_attention_cost=cum_attention, + ) + summary.turns.append(turn) + pending_request = None + + summary.total_turns = turn_num + summary.cumulative_token_cost = cum_token + summary.cumulative_attention_cost = cum_attention + if context_sizes: + summary.max_context_tokens = max(context_sizes) + summary.avg_context_tokens = sum(context_sizes) / len(context_sizes) + + return summary + + +def simulate_managed_cost( + path: Path, + age_threshold: int = 4, + min_size: int = 500, + label: str = "", +) -> CostSummary: + """Simulate cost with eviction, including fault overhead. + + Replays the proxy log with compaction, then estimates: + - Reduced context size per turn (from eviction) + - Extra inference passes from faults (at quadratic cost) + + Fault cost model: each fault triggers one additional inference + pass over the full context (n + |p|)² ≈ n² for large n. + """ + records = parse_proxy_log(path) + if not label: + label = f"{path.stem}_managed" + + summary = CostSummary(label=label, log_path=str(path)) + + requests = [ + r for r in records + if r.get("type") == "request" and "messages_full" in r + ] + # Pair with responses for actual token counts + responses = [] + pending = None + for rec in records: + if rec.get("type") == "request": + pending = rec + elif rec.get("type") in ("response_stream", "response") and pending: + responses.append(rec) + pending = None + + if not requests: + return summary + + page_store = PageStore() + cum_token = 0 + cum_attention = 0 + total_fault_attention = 0 + context_sizes = [] + + for turn_idx, req in enumerate(requests): + messages = copy.deepcopy(req["messages_full"]) + + # Apply previous evictions + _apply_evictions(messages, page_store) + + # Detect faults before compaction + faults = page_store.detect_faults(messages) + + # Run compaction + stats = compact_messages( + messages, + age_threshold=age_threshold, + min_size=min_size, + page_store=page_store, + ) + + # Estimate managed context size: + # Use the ratio of compacted/original bytes to scale the + # actual token count from the API response. + bytes_compacted = len(json.dumps(messages).encode("utf-8")) + bytes_original = len( + json.dumps(req["messages_full"]).encode("utf-8") + ) + + if turn_idx < len(responses): + usage = responses[turn_idx].get("usage", {}) + n_baseline = _effective_input(usage) + else: + n_baseline = 0 + + if bytes_original > 0 and n_baseline > 0: + ratio = bytes_compacted / bytes_original + n_managed = max(1, int(n_baseline * ratio)) + else: + n_managed = n_baseline + + # Token cost (linear) + cum_token += n_managed + + # Attention cost (quadratic) for the main inference + attention = n_managed * n_managed + cum_attention += attention + + # Fault cost: each fault is an extra inference pass + # The fault restores |p| tokens into a context of size n_managed + fault_attention = 0 + fault_tokens = 0 + for fault in faults: + # Estimate restored page size from the fault + p_size = getattr(fault, 'original_size', 0) or 500 + p_tokens = p_size // 4 # rough bytes-to-tokens + n_with_fault = n_managed + p_tokens + fault_attention += n_with_fault * n_with_fault + fault_tokens += n_with_fault + + total_fault_attention += fault_attention + context_sizes.append(n_managed) + + turn = TurnCost( + turn=turn_idx + 1, + timestamp=req.get("timestamp", ""), + context_tokens=n_managed, + token_cost=n_managed, + cumulative_token_cost=cum_token, + attention_cost=attention, + cumulative_attention_cost=cum_attention, + evictions=stats.evicted_count, + faults=len(faults), + fault_tokens=fault_tokens, + fault_attention_cost=fault_attention, + ) + summary.turns.append(turn) + + summary.total_turns = len(requests) + summary.cumulative_token_cost = cum_token + summary.cumulative_attention_cost = cum_attention + summary.total_fault_attention_cost = total_fault_attention + summary.total_faults = sum(t.faults for t in summary.turns) + summary.total_evictions = sum(t.evictions for t in summary.turns) + if context_sizes: + summary.max_context_tokens = max(context_sizes) + summary.avg_context_tokens = sum(context_sizes) / len(context_sizes) + + return summary + + +def compare(baseline: CostSummary, managed: CostSummary) -> CostComparison: + """Build a comparison between baseline and managed runs.""" + return CostComparison(baseline=baseline, managed=managed) + + +# --------------------------------------------------------------------------- +# Display +# --------------------------------------------------------------------------- + +def _fmt_attention(n: int) -> str: + """Format attention cost in human-readable units.""" + if n >= 1e12: + return f"{n / 1e12:.2f}T" + if n >= 1e9: + return f"{n / 1e9:.2f}G" + if n >= 1e6: + return f"{n / 1e6:.2f}M" + if n >= 1e3: + return f"{n / 1e3:.1f}K" + return str(n) + + +def print_cost_summary(s: CostSummary) -> None: + """Print cost summary for one simulation.""" + print(f"\n{'=' * 60}") + print(f"Cost: {s.label}") + print(f"{'=' * 60}") + print(f"Turns: {s.total_turns:>10,}") + print(f"Cumulative tokens: {s.cumulative_token_cost:>10,}") + print(f"Cumulative attention: {_fmt_attention(s.cumulative_attention_cost):>10s}") + print(f"Max context (tokens): {s.max_context_tokens:>10,}") + print(f"Avg context (tokens): {s.avg_context_tokens:>10,.0f}") + + if s.total_evictions > 0: + print(f"Evictions: {s.total_evictions:>10,}") + print(f"Faults: {s.total_faults:>10,}") + print(f"Fault attention cost: {_fmt_attention(s.total_fault_attention_cost):>10s}") + + # Per-turn curve + if len(s.turns) > 3: + print(f"\nPer-turn context and attention:") + print(f" {'Turn':>4s} {'Context':>10s} {'Attention':>12s} {'Cum Attn':>12s} {'Evict':>5s} {'Fault':>5s}") + step = max(1, len(s.turns) // 20) + for i, t in enumerate(s.turns): + if i % step == 0 or i == len(s.turns) - 1: + fault_str = "" + if t.fault_attention_cost > 0: + fault_str = f" +{_fmt_attention(t.fault_attention_cost)}" + print( + f" T{t.turn:>3d} {t.context_tokens:>10,} " + f"{_fmt_attention(t.attention_cost):>12s} " + f"{_fmt_attention(t.cumulative_attention_cost):>12s} " + f"{t.evictions:>5d} {t.faults:>5d}{fault_str}" + ) + + +def print_comparison(comp: CostComparison) -> None: + """Print side-by-side cost comparison.""" + b, m = comp.baseline, comp.managed + + print(f"\n{'=' * 65}") + print("COST COMPARISON") + print(f"{'=' * 65}") + + def row(label: str, bval: str, mval: str, delta: str = ""): + print(f" {label:<28s} {bval:>14s} {mval:>14s} {delta:>10s}") + + row("", b.label, m.label, "delta") + print(f" {'-' * 28} {'-' * 14} {'-' * 14} {'-' * 10}") + + row("Turns", str(b.total_turns), str(m.total_turns)) + row( + "Cumulative tokens", + f"{b.cumulative_token_cost:,}", + f"{m.cumulative_token_cost:,}", + f"{comp.token_savings_pct:+.1%}", + ) + row( + "Cumulative attention", + _fmt_attention(b.cumulative_attention_cost), + _fmt_attention(m.cumulative_attention_cost), + f"{comp.attention_savings_pct:+.1%}", + ) + row( + "Fault attention overhead", + "0", + _fmt_attention(m.total_fault_attention_cost), + ) + row( + "Net attention (incl faults)", + _fmt_attention(b.cumulative_attention_cost), + _fmt_attention(m.cumulative_attention_cost + m.total_fault_attention_cost), + f"{comp.net_attention_savings_pct:+.1%}", + ) + row( + "Max context", + f"{b.max_context_tokens:,}", + f"{m.max_context_tokens:,}", + ) + row( + "Avg context", + f"{b.avg_context_tokens:,.0f}", + f"{m.avg_context_tokens:,.0f}", + ) + row("Evictions", "0", f"{m.total_evictions:,}") + row("Faults", "0", f"{m.total_faults:,}") + + # Dollar cost estimate (Opus pricing) + def _dollar(s: CostSummary) -> float: + # Simplified: all effective input at $15/M + return (s.cumulative_token_cost / 1e6) * 15.0 + + b_cost = _dollar(b) + m_cost = _dollar(m) + print(f"\n Est. input cost (Opus $15/M):") + print(f" Baseline: ${b_cost:,.2f}") + print(f" Managed: ${m_cost:,.2f}") + if b_cost > 0: + print(f" Savings: ${b_cost - m_cost:,.2f} ({(b_cost - m_cost) / b_cost:.1%})") + + +# --------------------------------------------------------------------------- +# CLI +# --------------------------------------------------------------------------- + +def main(): + parser = argparse.ArgumentParser( + description="Cost simulation under the inverted cost model" + ) + parser.add_argument( + "logs", + type=Path, + nargs="+", + help="Proxy JSONL log file(s)", + ) + parser.add_argument( + "--replay", + action="store_true", + help="Also simulate managed cost with eviction and compare", + ) + parser.add_argument( + "--age-threshold", + type=int, + default=4, + help="Evict tool results older than N turns (default: 4)", + ) + parser.add_argument( + "--min-size", + type=int, + default=500, + help="Don't evict results smaller than N bytes (default: 500)", + ) + parser.add_argument( + "--json", + action="store_true", + help="Output as JSON", + ) + args = parser.parse_args() + + for log_path in args.logs: + if not log_path.exists(): + print(f"Warning: {log_path} not found, skipping.", file=sys.stderr) + continue + + baseline = compute_baseline_cost(log_path) + + if args.replay: + managed = simulate_managed_cost( + log_path, + age_threshold=args.age_threshold, + min_size=args.min_size, + ) + comp = compare(baseline, managed) + + if args.json: + out = { + "baseline": asdict(baseline), + "managed": asdict(managed), + "token_savings_pct": comp.token_savings_pct, + "attention_savings_pct": comp.attention_savings_pct, + "net_attention_savings_pct": comp.net_attention_savings_pct, + } + # Trim per-turn data for JSON + for key in ("baseline", "managed"): + out[key]["turn_count"] = len(out[key].pop("turns")) + print(json.dumps(out, indent=2)) + else: + print_comparison(comp) + else: + if args.json: + out = asdict(baseline) + out["turn_count"] = len(out.pop("turns")) + print(json.dumps(out, indent=2)) + else: + print_cost_summary(baseline) + + +if __name__ == "__main__": + main() diff --git a/src/mnemosyne/deprecated/__init__.py b/src/mnemosyne/deprecated/__init__.py new file mode 100644 index 0000000..7abfddb --- /dev/null +++ b/src/mnemosyne/deprecated/__init__.py @@ -0,0 +1,2 @@ +# Deprecated modules — kept for transitional import compatibility. +# New development targets gateway.py. diff --git a/src/mnemosyne/deprecated/proxy.py b/src/mnemosyne/deprecated/proxy.py new file mode 100644 index 0000000..fd31627 --- /dev/null +++ b/src/mnemosyne/deprecated/proxy.py @@ -0,0 +1,1208 @@ +#!/usr/bin/env python3 +"""Logging proxy for Claude API calls, with optional context paging. + +Two modes: + Observe (default): Logs request/response metrics. Pure observation. + Compact (--compact): Also evicts stale tool results from the messages + array before forwarding, replacing them with compact summaries. + +Usage: + # Observation only + python -m pichay.proxy [--port 0] [--log-dir logs] + + # With context paging + python -m pichay.proxy --compact [--age-threshold 4] [--min-size 500] + + # Point Claude Code at it + ANTHROPIC_BASE_URL=http://localhost: claude + +Port 0 (default) picks a random free port to avoid collisions. + +Session isolation: Each distinct conversation (identified by a fingerprint +of the first user message) gets its own token state, page store, and +phantom call queue. This prevents cross-contamination when Claude Code +spawns concurrent subagents through the same proxy. +""" + +from __future__ import annotations + +import argparse +import hashlib +import json +import os +import socket +import sys +import time +from datetime import datetime, timezone +from pathlib import Path + +import httpx +from flask import Flask, Response, request + +from mnemosyne.blocks import BlockStore +from mnemosyne.pager import PageStore, compact_messages, compact_conversation +from mnemosyne.tags import parse_cleanup_tags, strip_cleanup_tags +from mnemosyne.message_ops import ( + PICHAY_STATUS_MARKER, + process_cleanup_tags, + inject_system_status, + measure_system_prompt, + measure_messages, + sanitize_messages, + strip_response_headers, +) +from mnemosyne.phantom import ( + PhantomCall, + _handle_phantom_call, + apply_compaction, + filtered_stream, + inject_phantom_results, + inject_tools, +) + +DEFAULT_API_BASE = "https://api.anthropic.com" + + +def _phantom_continuation( + body: dict, + phantom_calls: list[PhantomCall], + page_store, + block_store, + headers: dict, + upstream_path: str, + client: httpx.Client, + chunks_collected: list, + sid: str, +): + """Gateway auto-continue: execute phantom tools and stream continuation. + + When the model calls phantom tools (yuyay, qunqay), the proxy handles + them and immediately sends a continuation request with the tool results. + The model continues generating without a turn boundary. The framework + never sees the phantom tools — just one seamless response. + + Yields SSE bytes that extend the original response stream. + """ + # Build the assistant message content with phantom tool_use blocks + # (We need to reconstruct what the model actually said) + assistant_content = [] + for pc in phantom_calls: + assistant_content.append({ + "type": "tool_use", + "id": pc.tool_use_id, + "name": pc.name, + "input": pc.input, + }) + + # Build tool_result messages + tool_results = [] + for pc in phantom_calls: + result_text = _handle_phantom_call(pc, page_store, + block_store=block_store) + tool_results.append({ + "type": "tool_result", + "tool_use_id": pc.tool_use_id, + "content": result_text, + }) + + # Continuation messages: original messages + assistant + user with results + messages = list(body.get("messages", [])) + messages.append({"role": "assistant", "content": assistant_content}) + messages.append({"role": "user", "content": tool_results}) + + # Build continuation request (same params, updated messages) + cont_body = {**body, "messages": messages, "stream": True} + + cont_size = len(json.dumps(cont_body)) + print( + f" [{sid}] CONTINUATION: {len(phantom_calls)} phantom call(s), " + f"sending {len(messages)} messages, {cont_size // 1024}KB", + file=sys.stderr, + ) + + try: + cont_resp = client.send( + client.build_request( + "POST", upstream_path, + json=cont_body, headers=headers, + ), + stream=True, + ) + + # Stream continuation, suppressing message_start + # (framework already received one from the first response) + buffer = b"" + for chunk in cont_resp.iter_bytes(): + chunks_collected.append(chunk) + buffer += chunk + + while b"\n\n" in buffer: + event_bytes, buffer = buffer.split(b"\n\n", 1) + event_text = event_bytes.decode("utf-8", errors="replace") + + data_str = None + for line in event_text.split("\n"): + if line.startswith("data: "): + data_str = line[6:] + + suppress = False + if data_str and data_str != "[DONE]": + try: + evt = json.loads(data_str) + # Suppress envelope events — the framework already + # has message_start/message_delta/message_stop from + # the original response. The continuation is an + # internal gateway round-trip; only content_block + # events should leak through to the client. + if evt.get("type") in ( + "message_start", + "message_delta", + "message_stop", + ): + suppress = True + except json.JSONDecodeError: + pass + + if not suppress: + yield event_bytes + b"\n\n" + + if buffer: + yield buffer + + cont_resp.close() + except Exception as e: + print(f" [{sid}] CONTINUATION STREAM ERROR: {e}", file=sys.stderr) + raise + + +def _session_id(body: dict) -> str: + """Derive a stable session fingerprint from the first user message. + + Different conversations have different first messages, so this + is unique per conversation and stable across turns (the first + message stays in the array for the conversation's lifetime). + + The system prompt is excluded because it contains dynamic content + (timestamps, git status, cache blocks) that changes between requests. + """ + messages = body.get("messages", []) + first = json.dumps(messages[0], sort_keys=True) if messages else "" + return hashlib.sha256(first.encode()).hexdigest()[:8] + + +def find_free_port() -> int: + """Find a free port by binding to port 0 and reading the assignment.""" + with socket.socket(socket.AF_INET6, socket.SOCK_STREAM) as s: + s.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1) + try: + # Try IPv6 loopback first (reduces collision space) + s.bind(("::1", 0)) + except OSError: + # Fall back to IPv4 + s.close() + with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s4: + s4.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1) + s4.bind(("127.0.0.1", 0)) + return s4.getsockname()[1] + return s.getsockname()[1] + + +def create_app( + log_dir: Path, + compact: bool = False, + trim: bool = False, + age_threshold: int = 4, + min_size: int = 500, + upstream: str = DEFAULT_API_BASE, + token_cap: int = 0, +) -> Flask: + """Create the proxy Flask app.""" + app = Flask(__name__) + log_dir.mkdir(parents=True, exist_ok=True) + + # One log file per proxy session + session_start = datetime.now(timezone.utc).strftime("%Y%m%d_%H%M%S") + log_file = log_dir / f"proxy_{session_start}.jsonl" + + _token_warning_threshold = int(token_cap * 0.8) if token_cap > 0 else 0 + + # ANSI colors for terminal output + _YELLOW = "\033[33m" + _RED = "\033[31m" + _DIM = "\033[2m" + _RESET = "\033[0m" + + if token_cap > 0: + print( + f"Token cap: {token_cap:,} " + f"(warning at {_token_warning_threshold:,})", + file=sys.stderr, + ) + + # --- Session-keyed state --- + # Each conversation gets its own token state, page store, and phantom + # call queue. This prevents cross-contamination when multiple Claude + # Code conversations (including subagents) hit the proxy concurrently. + _sessions: dict[str, dict] = {} + + def _get_session(body: dict) -> dict: + """Get or create per-conversation session state.""" + sid = _session_id(body) + if sid not in _sessions: + page_log = log_dir / f"pages_{session_start}_{sid}.jsonl" + ps = PageStore(log_path=page_log) + _sessions[sid] = { + "id": sid, + "token_state": { + "last_effective": 0, + "blocked": False, + "turn": 0, + "calibrated": False, + }, + "page_store": ps, + "block_store": BlockStore(), + "phantom_pending": [], + "observe_only": set(), + "pending_compaction": None, + } + print( + f" {_DIM}[{sid}] new session{_RESET}", + file=sys.stderr, + ) + return _sessions[sid] + + # Trimmer state (shared — it's stateless per-request, just caches patterns) + trimmer = None + if trim: + from mnemosyne.trimmer import SystemPromptTrimmer + trimmer = SystemPromptTrimmer() + print( + "Trim mode: tool stubs + skill dedup + static tracking", + file=sys.stderr, + ) + + if compact: + print( + f"Compact mode: age_threshold={age_threshold}, " + f"min_size={min_size}", + file=sys.stderr, + ) + + # Persistent HTTP client for forwarding + client = httpx.Client( + base_url=upstream, + timeout=httpx.Timeout(300.0, connect=30.0), + ) + if upstream != DEFAULT_API_BASE: + print(f"Upstream: {upstream}", file=sys.stderr) + + def log_record(record: dict) -> None: + """Append a record to the log file.""" + with open(log_file, "a", encoding="utf-8") as f: + f.write(json.dumps(record, default=str) + "\n") + + def _check_token_cap(usage: dict, session: dict) -> None: + """Check effective tokens against cap, update state, emit warnings.""" + ts = session["token_state"] + sid = session["id"] + effective = ( + usage.get("input_tokens", 0) + + usage.get("cache_creation_input_tokens", 0) + + usage.get("cache_read_input_tokens", 0) + ) + # Always track usage (needed for system status injection) + ts["last_effective"] = effective + if token_cap <= 0: + return + pct = (effective / token_cap * 100) if token_cap > 0 else 0 + + if effective > token_cap: + ts["blocked"] = True + print( + f"{_RED} [{sid}] TOKEN CAP EXCEEDED: {effective:,} / {token_cap:,} " + f"({pct:.0f}%) — next request will be blocked{_RESET}", + file=sys.stderr, + ) + log_record({ + "type": "token_cap", + "timestamp": datetime.now(timezone.utc).isoformat(), + "session": sid, + "action": "exceeded", + "effective_tokens": effective, + "token_cap": token_cap, + "pct": round(pct, 1), + "turn": ts["turn"], + }) + elif effective > _token_warning_threshold: + print( + f"{_YELLOW} [{sid}] TOKEN WARNING: {effective:,} / {token_cap:,} " + f"({pct:.0f}%) — approaching cap{_RESET}", + file=sys.stderr, + ) + log_record({ + "type": "token_cap", + "timestamp": datetime.now(timezone.utc).isoformat(), + "session": sid, + "action": "warning", + "effective_tokens": effective, + "token_cap": token_cap, + "pct": round(pct, 1), + "turn": ts["turn"], + }) + + def _display_turn_status(usage: dict, session: dict) -> None: + """Post-response status line with real token count.""" + ts = session["token_state"] + sid = session["id"] + ps = session["page_store"] + effective = ( + usage.get("input_tokens", 0) + + usage.get("cache_creation_input_tokens", 0) + + usage.get("cache_read_input_tokens", 0) + ) + if effective == 0: + return + + cap_str = "" + if token_cap > 0: + pct = effective / token_cap * 100 + cap_str = f"/{token_cap // 1000}k ({pct:.0f}%)" + + # Cache hit rate: fraction of input tokens served from KV cache + cache_read = usage.get("cache_read_input_tokens", 0) + cache_create = usage.get("cache_creation_input_tokens", 0) + cache_total = cache_read + cache_create + cache_str = "" + if cache_total > 0: + cache_pct = cache_read / cache_total * 100 + cache_str = f" | cache {cache_pct:.0f}%" + + ev_str = "" + if ps is not None: + pin_str = f" pin {len(ps._pinned)}" if ps._pinned else "" + ev_str = ( + f" | ev {ps.unique_evictions} gc {ps.gc_count}{pin_str}" + f" | faults {len(ps.faults)}/{ps.unique_evictions}" + ) + + print( + f" [{sid}] [Turn {ts['turn']}] " + f"{effective:,} tok{cap_str}{cache_str}{ev_str}", + file=sys.stderr, + ) + + # Wire up trimmer logging now that log_record exists + if trimmer is not None: + trimmer.log_fn = log_record + + @app.route("/v1/messages/count_tokens", methods=["POST"]) + def proxy_count_tokens(): + """Forward count_tokens to upstream, applying compaction first.""" + body = request.get_json(force=True) + session = _get_session(body) + sid = session["id"] + ps = session["page_store"] + + # Apply the same compaction so the count reflects reality + if compact and ps is not None: + messages = body.get("messages", []) + stats = compact_messages( + messages, + age_threshold=age_threshold, + min_size=min_size, + page_store=ps, + ) + if stats.evicted_count > 0: + print( + f" [{sid}] [count_tokens] compacted {stats.evicted_count} results " + f"before counting", + file=sys.stderr, + ) + + headers = dict(request.headers) + for h in ["Host", "Content-Length", "Transfer-Encoding"]: + headers.pop(h, None) + + query = "/v1/messages/count_tokens" + if request.query_string: + query += "?" + request.query_string.decode("utf-8") + + count_body_bytes = len(json.dumps(body).encode("utf-8")) + count_tools = len(body.get("tools", [])) + count_messages = len(body.get("messages", [])) + + try: + resp = client.post(query, json=body, headers=headers) + # Parse and log the token count Anthropic returns + input_tokens = None + try: + count_data = json.loads(resp.content) + input_tokens = count_data.get("input_tokens") + except Exception: + pass + # Stash for comparison with outgoing /v1/messages + session["last_count"] = { + "input_tokens": input_tokens, + "body_bytes": count_body_bytes, + "tools": count_tools, + "messages": count_messages, + } + log_record({ + "type": "count_tokens", + "timestamp": datetime.now(timezone.utc).isoformat(), + "session": sid, + "input_tokens": input_tokens, + "body_bytes": count_body_bytes, + "tools": count_tools, + "messages": count_messages, + "status_code": resp.status_code, + }) + print( + f" [{sid}] COUNT: {input_tokens} tokens" + f" ({count_body_bytes / 1024:.0f}KB," + f" {count_messages} msgs, {count_tools} tools)", + file=sys.stderr, + ) + return ( + resp.content, + resp.status_code, + strip_response_headers(resp.headers), + ) + except Exception as e: + print(f" [{sid}] [count_tokens] error: {e}", file=sys.stderr) + return Response( + json.dumps({"error": str(e)}), + status=502, + content_type="application/json", + ) + + @app.route("/v1/messages", methods=["POST"]) + def proxy_messages(): + """Proxy the messages endpoint with full logging.""" + request_time = datetime.now(timezone.utc) + body = request.get_json(force=True) + session = _get_session(body) + sid = session["id"] + ts = session["token_state"] + ps = session["page_store"] + + # Measure without storing full content + system_metrics = measure_system_prompt(body) + message_metrics = measure_messages(body) + + request_record = { + "type": "request", + "timestamp": request_time.isoformat(), + "session": sid, + "model": body.get("model", "unknown"), + "max_tokens": body.get("max_tokens"), + "stream": body.get("stream", False), + "system": system_metrics, + "messages": message_metrics, + "total_request_bytes": len( + json.dumps(body).encode("utf-8") + ), + } + + if body.get("system"): + request_record["system_prompt_full"] = body["system"] + + # Store full messages for offline replay + request_record["messages_full"] = body.get("messages", []) + + log_record(request_record) + + # On first request for this session, calibrate turn counter + if not ts["calibrated"]: + messages = body.get("messages", []) + prior_turns = sum(1 for m in messages if m.get("role") == "assistant") + ts["turn"] = prior_turns + ts["calibrated"] = True + + ts["turn"] += 1 + + # --- Token cap gate --- + if token_cap > 0 and ts["blocked"]: + last = ts["last_effective"] + msg = ( + f"Token cap exceeded. Last effective context: " + f"{last:,} tokens (cap: {token_cap:,}). " + f"Session has been blocked to prevent billing. " + f"Restart with a fresh context or raise --token-cap." + ) + print( + f"{_RED} [{sid}] BLOCKED: {msg}{_RESET}", + file=sys.stderr, + ) + log_record({ + "type": "token_cap", + "timestamp": datetime.now(timezone.utc).isoformat(), + "session": sid, + "action": "blocked", + "last_effective_tokens": last, + "token_cap": token_cap, + "turn": ts["turn"], + }) + return Response( + json.dumps({ + "type": "error", + "error": { + "type": "rate_limit_error", + "message": msg, + }, + }), + status=429, + content_type="application/json", + ) + + # --- System prompt trimming (if enabled) --- + if trim and trimmer is not None: + trim_result = trimmer.trim(body) + if trim_result.total_bytes_saved > 0: + print( + f" [{sid}] Trimmed: {trim_result.tools.stubbed_tools} stubs, " + f"{trim_result.skills.duplicates_removed} skill dupes, " + f"saved {trim_result.total_bytes_saved:,} bytes", + file=sys.stderr, + ) + + # --- Temperature override (if configured) --- + # Extended thinking requires temperature=1; skip override when thinking is active. + thinking_enabled = bool(body.get("thinking")) + if app.config.get("temperature_override") is not None and not thinking_enabled: + body["temperature"] = app.config["temperature_override"] + elif app.config.get("temperature_override") is not None and thinking_enabled: + print( + f" [{sid}] [proxy] Skipping temperature override " + f"(thinking enabled: {body.get('thinking')})", + file=sys.stderr, + ) + + # --- System status injection (always active) --- + # Give the model context pressure awareness and system identification. + # Replaced each turn (not appended) so the model sees current state. + inject_system_status(body, ts, token_cap, request_time, + block_store=session.get("block_store")) + + # --- Block labeling (conversation memory management) --- + bs = session["block_store"] + if bs is not None: + messages = body.get("messages", []) + bs.label_messages(messages, ts["turn"]) + + # --- Cleanup tag processing (Phase 2) --- + if bs is not None: + messages = body.get("messages", []) + cleanup_stats = process_cleanup_tags(messages, bs, + session.get("page_store")) + if cleanup_stats: + print( + f" [{sid}] CLEANUP: {cleanup_stats}", + file=sys.stderr, + ) + # Apply block status to messages (replace dropped/summarized) + apply_stats = bs.apply_to_messages(messages) + applied = sum(apply_stats.values()) + if applied: + print( + f" [{sid}] BLOCKS: {apply_stats}", + file=sys.stderr, + ) + + # --- Phantom tool handling (always active) --- + observe_only = session["observe_only"] + if ps is not None: + # Inject results for phantom calls from the previous turn + pending = session["phantom_pending"] + if pending: + messages = body.get("messages", []) + inject_phantom_results(messages, pending, ps, observe_only) + session["phantom_pending"] = [] + released = [ + c for c in pending + if c.name in ("qunqay", "memory_release") + ] + faulted = [ + c for c in pending + if c.name in ("yuyay", "recall", "memory_fault") + ] + if released: + paths = [] + for c in released: + paths.extend(c.input.get("paths", [])) + print( + f" [{sid}] RELEASE: model released {len(paths)} path(s)", + file=sys.stderr, + ) + if faulted: + handles = [] + for c in faulted: + handles.extend(c.input.get("handles", + c.input.get("paths", []))) + print( + f" [{sid}] RECALL: restored {len(handles)} tensor(s) " + f"from cache", + file=sys.stderr, + ) + + # Inject phantom tools into the tools array + observe_only = inject_tools(body) + session["observe_only"] = observe_only + + # --- Deferred structural compaction (tiqsiy) --- + pending_compact = session.get("pending_compaction") + if pending_compact is not None: + messages = body.get("messages", []) + cr = apply_compaction( + messages, + older_than=pending_compact["older_than"], + summary=pending_compact["summary"], + ) + session["pending_compaction"] = None + if cr.messages_removed > 0: + print( + f" [{sid}] COMPACT: {cr.messages_removed} messages → " + f"summary pair, {cr.chars_removed:,} chars archived, " + f"{cr.messages_after} messages remain", + file=sys.stderr, + ) + log_record({ + "type": "structural_compaction", + "timestamp": datetime.now(timezone.utc).isoformat(), + "session": sid, + "messages_removed": cr.messages_removed, + "messages_after": cr.messages_after, + "chars_removed": cr.chars_removed, + "older_than": pending_compact["older_than"], + "summary_length": len(pending_compact["summary"]), + }) + + # --- Context paging (if enabled) --- + if compact and ps is not None: + messages = body.get("messages", []) + + # Detect page faults BEFORE compaction + faults = ps.detect_faults(messages) + if faults: + for fault in faults: + ago = time.monotonic() - fault.original_eviction.evicted_at + if ago < 60: + age_str = f"{ago:.0f}s ago" + else: + age_str = f"{ago / 60:.1f}min ago" + print( + f" [{sid}] PAGE FAULT: {fault.tool_name} " + f"re-requested (evicted {age_str}, " + f"{fault.original_eviction.original_size:,} bytes)", + file=sys.stderr, + ) + log_record({ + "type": "page_faults", + "timestamp": datetime.now(timezone.utc).isoformat(), + "session": sid, + "count": len(faults), + "faults": [ + { + "tool_name": f.tool_name, + "original_turn": f.original_eviction.turn_index, + "original_size": f.original_eviction.original_size, + } + for f in faults + ], + "cumulative_faults": len(ps.faults), + "unique_evictions": ps.unique_evictions, + "gc_count": ps.gc_count, + "fault_rate": ps.fault_rate, + }) + + # Compact + stats = compact_messages( + messages, + age_threshold=age_threshold, + min_size=min_size, + page_store=ps, + ) + if stats.evicted_count > 0: + post_metrics = measure_messages(body) + compaction_record = { + "type": "compaction", + "timestamp": datetime.now(timezone.utc).isoformat(), + "session": sid, + "evicted": stats.evicted_count, + "total_tool_results": stats.total_tool_results, + "bytes_before": stats.bytes_before, + "bytes_after": stats.bytes_after, + "bytes_saved": stats.bytes_saved, + "reduction_pct": round(stats.reduction_pct, 1), + "skipped_small": stats.skipped_small, + "skipped_recent": stats.skipped_recent, + "skipped_error": stats.skipped_error, + "unique_evictions": ps.unique_evictions, + "gc_count": ps.gc_count, + "eviction_bytes_saved": ps.eviction_bytes_saved, + "gc_bytes_saved": ps.gc_bytes_saved, + "cumulative_faults": len(ps.faults), + "fault_rate": round(ps.fault_rate * 100, 1), + "messages_bytes_before": message_metrics[ + "messages_total_bytes" + ], + "messages_bytes_after": post_metrics[ + "messages_total_bytes" + ], + } + log_record(compaction_record) + + # Dashboard status line (pre-forward, bytes only) + before_kb = message_metrics["messages_total_bytes"] / 1024 + after_kb = post_metrics["messages_total_bytes"] / 1024 + saved_pct = ( + (1 - after_kb / before_kb) * 100 + if before_kb > 0 + else 0 + ) + pin_str = f" pin {len(ps._pinned)}" if ps._pinned else "" + print( + f" [{sid}] [{before_kb:.0f}KB \u2192 {after_kb:.0f}KB]" + f" ({saved_pct:.0f}% saved) " + f"ev {ps.unique_evictions} gc {ps.gc_count}" + f"{pin_str} | " + f"faults {len(ps.faults)}/{ps.unique_evictions}", + file=sys.stderr, + ) + else: + # No eviction this turn — still show working set size + ws_kb = message_metrics["messages_total_bytes"] / 1024 + print( + f" [{sid}] [{ws_kb:.0f}KB]", + file=sys.stderr, + ) + + # Conversation compression — DISABLED by default. + # Third-party summarization loses model-relevant context; + # the model should manage its own memory via qunqay/compact. + # To re-enable, set PICHAY_CONV_COMPACT=1. + conv_stats = None + if os.environ.get("PICHAY_CONV_COMPACT") == "1": + conv_stats = compact_conversation(messages, preserve_recent=12) + if conv_stats and conv_stats.messages_compressed > 0: + log_record({ + "type": "conversation_compaction", + "timestamp": datetime.now(timezone.utc).isoformat(), + "session": sid, + "messages_compressed": conv_stats.messages_compressed, + "chars_saved": conv_stats.chars_saved, + }) + print( + f" [{sid}] CONV: {conv_stats.messages_compressed} msgs compressed, " + f"{conv_stats.chars_saved:,} chars saved", + file=sys.stderr, + ) + + # --- Final sanitization (catch empty blocks from any upstream step) --- + messages = body.get("messages", []) + sanitize_fixes = sanitize_messages(messages) + if sanitize_fixes: + print( + f" [{sid}] SANITIZE: fixed {sanitize_fixes} empty block(s)", + file=sys.stderr, + ) + + # Forward to Anthropic — measure the final outgoing payload + outgoing_bytes = len(json.dumps(body).encode("utf-8")) + outgoing_tools = len(body.get("tools", [])) + outgoing_messages = len(body.get("messages", [])) + log_record({ + "type": "outgoing", + "timestamp": datetime.now(timezone.utc).isoformat(), + "session": sid, + "turn": ts["turn"], + "outgoing_bytes": outgoing_bytes, + "outgoing_tools": outgoing_tools, + "outgoing_messages": outgoing_messages, + "max_tokens": body.get("max_tokens"), + "model": body.get("model", "unknown"), + }) + # Compare with last count_tokens to detect divergence + last_count = session.get("last_count") + divergence_pct = None + if last_count and last_count["body_bytes"] > 0: + divergence_pct = ( + (outgoing_bytes - last_count["body_bytes"]) + / last_count["body_bytes"] + * 100 + ) + + if divergence_pct is not None and divergence_pct > 10: + print( + f"{_RED} [{sid}] DIVERGENCE WARNING: outgoing " + f"{outgoing_bytes / 1024:.0f}KB vs count_tokens " + f"{last_count['body_bytes'] / 1024:.0f}KB " + f"(+{divergence_pct:.0f}%) — gateway is injecting " + f"{(outgoing_bytes - last_count['body_bytes']) / 1024:.0f}KB" + f"{_RESET}", + file=sys.stderr, + ) + print( + f"{_RED} [{sid}] count: {last_count['messages']} msgs, " + f"{last_count['tools']} tools → outgoing: " + f"{outgoing_messages} msgs, {outgoing_tools} tools" + f"{_RESET}", + file=sys.stderr, + ) + else: + print( + f" [{sid}] OUTGOING: {outgoing_bytes / 1024:.0f}KB" + f" ({outgoing_messages} msgs, {outgoing_tools} tools)", + file=sys.stderr, + ) + + headers = dict(request.headers) + for h in ["Host", "Content-Length", "Transfer-Encoding"]: + headers.pop(h, None) + + upstream_path = "/v1/messages" + if request.query_string: + upstream_path += "?" + request.query_string.decode("utf-8") + + if body.get("stream", False): + return _proxy_streaming( + body, headers, request_time, upstream_path, session, + ) + else: + return _proxy_direct( + body, headers, request_time, upstream_path, session, + ) + + def _proxy_direct(body, headers, request_time, upstream_path, session): + sid = session["id"] + try: + resp = client.post( + upstream_path, json=body, headers=headers, + ) + response_time = datetime.now(timezone.utc) + try: + resp_body = resp.json() + usage = resp_body.get("usage", {}) + log_record({ + "type": "response", + "timestamp": response_time.isoformat(), + "session": sid, + "duration_ms": int( + (response_time - request_time).total_seconds() * 1000 + ), + "status_code": resp.status_code, + "usage": usage, + "stop_reason": resp_body.get("stop_reason"), + }) + _check_token_cap(usage, session) + _display_turn_status(usage, session) + except Exception: + log_record({ + "type": "response_error", + "timestamp": response_time.isoformat(), + "session": sid, + "status_code": resp.status_code, + }) + return Response( + resp.content, + status=resp.status_code, + headers=strip_response_headers(resp.headers), + ) + except Exception as e: + log_record({ + "type": "proxy_error", + "timestamp": datetime.now(timezone.utc).isoformat(), + "session": sid, + "error": str(e), + }) + return Response( + json.dumps({"error": str(e)}), + status=502, + content_type="application/json", + ) + + def _proxy_streaming(body, headers, request_time, upstream_path, session): + sid = session["id"] + ps = session["page_store"] + observe_only = session["observe_only"] + try: + resp = client.send( + client.build_request( + "POST", upstream_path, json=body, headers=headers, + ), + stream=True, + ) + # Log non-200 responses immediately (error bodies are not SSE) + if resp.status_code != 200: + error_body = resp.read().decode("utf-8", errors="replace") + log_record({ + "type": "api_error", + "timestamp": datetime.now(timezone.utc).isoformat(), + "session": sid, + "status_code": resp.status_code, + "error_body": error_body[:2000], + }) + print( + f" [{sid}] API ERROR {resp.status_code}: " + f"{error_body[:200]}", + file=sys.stderr, + ) + resp.close() + return Response( + error_body, + status=resp.status_code, + headers=strip_response_headers(resp.headers), + content_type=resp.headers.get( + "content-type", "application/json" + ), + ) + response_headers = strip_response_headers(resp.headers) + + def generate(): + first_byte_time = None + chunks_collected = [] + phantom_calls: list[PhantomCall] = [] + continuation_needed: list[bool] = [] + try: + # Always filter stream for phantom tool calls + for chunk in filtered_stream( + resp.iter_bytes(), + chunks_collected, + phantom_calls, + observe_only=observe_only, + block_store=session["block_store"], + page_store=session.get("page_store"), + session_id=sid, + continuation_needed=continuation_needed, + ): + if first_byte_time is None: + first_byte_time = datetime.now(timezone.utc) + yield chunk + finally: + resp.close() + response_time = datetime.now(timezone.utc) + full_response = b"".join(chunks_collected) + + # Handle phantom calls — gateway mode. + # The proxy executes phantom tools and auto-continues + # so the model gets results without a turn boundary. + if phantom_calls: + for pc in phantom_calls: + _handle_phantom_call(pc, ps, + block_store=session.get("block_store")) + is_observed = ( + isinstance(observe_only, set) + and pc.name in observe_only + ) + print( + f" [{sid}] PHANTOM" + f"{'(observe)' if is_observed else ''}: " + f"{pc.name}({pc.input})", + file=sys.stderr, + ) + # Store tiqsiy compaction for deferred execution + if pc.name == "tiqsiy": + session["pending_compaction"] = { + "older_than": pc.input.get("older_than", 20), + "summary": pc.input.get("summary", ""), + } + + # Auto-continue: if filtered_stream suppressed the stop + # events, we need to execute phantom tools and send the + # results back so the model can continue generating. + # Only continue for intercepted calls (not observe-only). + intercepted = [ + pc for pc in phantom_calls + if not (isinstance(observe_only, set) + and pc.name in observe_only) + ] + if continuation_needed and intercepted: + try: + yield from _phantom_continuation( + body, intercepted, ps, + session.get("block_store"), + headers, upstream_path, client, + chunks_collected, sid, + ) + except Exception as e: + print( + f" [{sid}] CONTINUATION ERROR: {e}", + file=sys.stderr, + ) + # Fallback: yield stop events so stream is valid + stop_delta = json.dumps({ + "type": "message_delta", + "delta": {"stop_reason": "end_turn"}, + "usage": {}, + }) + stop_msg = json.dumps({"type": "message_stop"}) + yield f"data: {stop_delta}\n\n".encode() + yield f"data: {stop_msg}\n\n".encode() + yield b"data: [DONE]\n\n" + + usage = {} + try: + text = full_response.decode("utf-8", errors="replace") + # Merge usage from both message_start (input tokens) + # and message_delta (output tokens). Don't break early. + for line in text.split("\n"): + if not line.startswith("data: "): + continue + if "message_start" in line and "usage" in line: + event_data = json.loads(line[6:]) + msg = event_data.get("message", {}) + usage.update(msg.get("usage", {})) + elif "message_delta" in line and "usage" in line: + event_data = json.loads(line[6:]) + usage.update(event_data.get("usage", {})) + except Exception: + pass + + log_record({ + "type": "response_stream", + "timestamp": response_time.isoformat(), + "session": sid, + "duration_ms": int( + (response_time - request_time).total_seconds() + * 1000 + ), + "first_byte_ms": int( + (first_byte_time - request_time).total_seconds() + * 1000 + ) + if first_byte_time + else None, + "status_code": resp.status_code, + "response_bytes": len(full_response), + "usage": usage, + }) + _check_token_cap(usage, session) + _display_turn_status(usage, session) + + return Response( + generate(), + status=resp.status_code, + headers=response_headers, + content_type=response_headers.get( + "content-type", "text/event-stream" + ), + ) + except Exception as e: + log_record({ + "type": "proxy_error", + "timestamp": datetime.now(timezone.utc).isoformat(), + "session": sid, + "error": str(e), + }) + return Response( + json.dumps({"error": str(e)}), + status=502, + content_type="application/json", + ) + + @app.route("/health") + def health(): + parts = [] + if compact: + parts.append("compact") + if trim: + parts.append("trim") + mode = "+".join(parts) if parts else "observe" + + result = { + "status": "ok", + "mode": mode, + "log_file": str(log_file), + "active_sessions": len(_sessions), + } + result["sessions"] = {} + for sid, s in _sessions.items(): + sess_info = s["page_store"].summary() if s["page_store"] else {} + if s.get("block_store"): + sess_info["blocks"] = s["block_store"].summary() + result["sessions"][sid] = sess_info + if trim and trimmer is not None: + result["trimmer"] = trimmer.summary() + return result + + print(f"Logging to: {log_file}", file=sys.stderr) + return app + + +def main(): + parser = argparse.ArgumentParser( + description="Logging proxy for Claude API (with optional context paging)" + ) + parser.add_argument( + "--port", type=int, default=0, + help="Port to listen on (0 = random free port, default: 0)", + ) + parser.add_argument( + "--log-dir", type=Path, default=Path("logs"), + help="Directory for proxy logs (default: logs/)", + ) + parser.add_argument( + "--compact", action="store_true", + help="Enable context paging: evict stale tool results", + ) + parser.add_argument( + "--trim", action="store_true", + help="Enable system prompt trimming: tool stubs, skill dedup, static tracking", + ) + parser.add_argument( + "--age-threshold", type=int, default=4, + help="Evict tool results older than N user-turns (default: 4)", + ) + parser.add_argument( + "--min-size", type=int, default=500, + help="Don't evict results smaller than N bytes (default: 500)", + ) + parser.add_argument( + "--temperature", type=float, default=None, + help="Override temperature on all requests (e.g., 0 for deterministic)", + ) + parser.add_argument( + "--upstream", type=str, default=DEFAULT_API_BASE, + help=f"Upstream API base URL (default: {DEFAULT_API_BASE}). " + "Any Anthropic-compatible endpoint: OpenRouter, Kimi, etc.", + ) + parser.add_argument( + "--token-cap", type=int, default=0, + help="Hard cap on effective input tokens. Blocks requests after " + "exceeding this. Warning at 80%%. 0 = no cap (default: 0). " + "Set to 200000 for subscription billing threshold.", + ) + args = parser.parse_args() + + app = create_app( + args.log_dir, + compact=args.compact, + trim=args.trim, + age_threshold=args.age_threshold, + min_size=args.min_size, + upstream=args.upstream, + token_cap=args.token_cap, + ) + if args.temperature is not None: + app.config["temperature_override"] = args.temperature + + port = args.port if args.port != 0 else find_free_port() + parts = [] + if args.compact: + parts.append("COMPACT") + if args.trim: + parts.append("TRIM") + mode = "+".join(parts) if parts else "OBSERVE" + print( + f"Proxy [{mode}] listening on http://localhost:{port}", + file=sys.stderr, + ) + print( + f"Use: ANTHROPIC_BASE_URL=http://localhost:{port} claude", + file=sys.stderr, + ) + app.run(host="127.0.0.1", port=port, threaded=True) + + +if __name__ == "__main__": + main() diff --git a/src/mnemosyne/eval.py b/src/mnemosyne/eval.py new file mode 100644 index 0000000..ae3351f --- /dev/null +++ b/src/mnemosyne/eval.py @@ -0,0 +1,496 @@ +#!/usr/bin/env python3 +"""Evaluation framework for context paging experiments. + +Parses proxy JSONL logs from experimental runs, computes per-turn and +cumulative metrics, and compares across treatment conditions. + +Data sources: + - proxy_*.jsonl: request/response records with usage data + - pages_*.jsonl: eviction and page fault records + +Metrics computed: + - Token consumption (input, output, cache hits/misses) + - API call count + - Context size growth curve + - Compaction events and savings + - Fault rate + - Wall-clock time + - System prompt overhead + +Usage: + # Analyze a single run + uv run python tools/phase1/experiment_eval.py --run tmp/api_logs/proxy_*.jsonl + + # Compare treatments + uv run python tools/phase1/experiment_eval.py \\ + --compare baseline=tmp/exp/baseline/proxy.jsonl \\ + t1_compact=tmp/exp/t1/proxy.jsonl \\ + t2_trimmed=tmp/exp/t2/proxy.jsonl + + # JSON output for downstream analysis + uv run python tools/phase1/experiment_eval.py --run proxy.jsonl --json +""" + +from __future__ import annotations + +import argparse +import json +import sys +from dataclasses import dataclass, field, asdict +from datetime import datetime +from pathlib import Path + + +@dataclass +class TurnMetrics: + """Metrics for a single API call (request + response pair).""" + + turn: int + timestamp: str + model: str + # Token counts from Anthropic's response + # input_tokens = non-cached only. Effective = input + cache_creation + cache_read + input_tokens: int = 0 + output_tokens: int = 0 + cache_creation_tokens: int = 0 + cache_read_tokens: int = 0 + + @property + def effective_input_tokens(self) -> int: + """Total tokens in context = non-cached + cache_creation + cache_read.""" + return self.input_tokens + self.cache_creation_tokens + self.cache_read_tokens + # Sizes from proxy measurement + total_request_bytes: int = 0 + messages_bytes: int = 0 + system_prompt_bytes: int = 0 + tool_result_count: int = 0 + tool_result_bytes: int = 0 + tool_use_count: int = 0 + # Compaction (if active) + evictions: int = 0 + compaction_bytes_saved: int = 0 + faults: int = 0 + # Timing + duration_ms: int = 0 + first_byte_ms: int = 0 + + +@dataclass +class RunSummary: + """Aggregate metrics for one experimental run.""" + + label: str + proxy_log: str + # Totals + api_calls: int = 0 + total_input_tokens: int = 0 + total_output_tokens: int = 0 + total_cache_creation: int = 0 + total_cache_read: int = 0 + total_tokens: int = 0 # input + output + total_effective_input: int = 0 # input + cache_creation + cache_read (context size) + # Bytes + total_request_bytes: int = 0 + total_messages_bytes: int = 0 + total_system_prompt_bytes: int = 0 + total_tool_result_bytes: int = 0 + # Compaction + total_evictions: int = 0 + total_compaction_bytes_saved: int = 0 + total_faults: int = 0 + fault_rate: float = 0.0 + # Timing + total_duration_ms: int = 0 + wall_clock_seconds: float = 0.0 + avg_first_byte_ms: float = 0.0 + # Context growth + max_messages_bytes: int = 0 + max_input_tokens: int = 0 + # Per-turn data for curves + turns: list[TurnMetrics] = field(default_factory=list) + + @property + def avg_input_tokens(self) -> float: + if self.api_calls == 0: + return 0.0 + return self.total_input_tokens / self.api_calls + + @property + def avg_effective_input(self) -> float: + if self.api_calls == 0: + return 0.0 + return self.total_effective_input / self.api_calls + + @property + def n_squared_cost(self) -> int: + """Cumulative effective input tokens — the n² metric.""" + return self.total_effective_input + + @property + def system_prompt_fraction(self) -> float: + if self.total_request_bytes == 0: + return 0.0 + return self.total_system_prompt_bytes / self.total_request_bytes + + +def parse_proxy_log(path: Path) -> list[dict]: + """Read all records from a proxy JSONL log.""" + records = [] + with open(path, "r", encoding="utf-8", errors="replace") as f: + for line in f: + line = line.strip() + if not line: + continue + try: + records.append(json.loads(line)) + except json.JSONDecodeError: + continue + return records + + +def parse_page_log(path: Path) -> list[dict]: + """Read eviction/fault records from a page log.""" + if not path.exists(): + return [] + return parse_proxy_log(path) + + +def analyze_run(proxy_path: Path, label: str = "") -> RunSummary: + """Analyze a single experimental run from its proxy log.""" + records = parse_proxy_log(proxy_path) + if not label: + label = proxy_path.stem + + # Find matching page log + page_path = proxy_path.parent / proxy_path.name.replace("proxy_", "pages_") + page_records = parse_page_log(page_path) + + summary = RunSummary(label=label, proxy_log=str(proxy_path)) + + # Pair requests with responses + pending_request: dict | None = None + turn_num = 0 + first_timestamp: str | None = None + last_timestamp: str | None = None + + # Index compaction records by timestamp for matching + compaction_by_ts: dict[str, dict] = {} + fault_by_ts: dict[str, dict] = {} + for rec in records: + if rec["type"] == "compaction": + compaction_by_ts[rec["timestamp"]] = rec + elif rec["type"] == "page_faults": + fault_by_ts[rec["timestamp"]] = rec + + for rec in records: + rtype = rec.get("type", "") + + if rtype == "request": + pending_request = rec + if first_timestamp is None: + first_timestamp = rec["timestamp"] + continue + + if rtype in ("response_stream", "response") and pending_request is not None: + turn_num += 1 + last_timestamp = rec["timestamp"] + + usage = rec.get("usage", {}) + messages = pending_request.get("messages", {}) + system = pending_request.get("system", {}) + + input_tokens = usage.get("input_tokens", 0) + output_tokens = usage.get("output_tokens", 0) + cache_creation = usage.get("cache_creation_input_tokens", 0) + cache_read = usage.get("cache_read_input_tokens", 0) + + # System prompt bytes + if isinstance(system, dict): + sp_bytes = system.get("system_prompt_bytes", 0) + if isinstance(sp_bytes, str): + sp_bytes = 0 + else: + sp_bytes = 0 + + # Messages metrics + if isinstance(messages, dict): + msg_bytes = messages.get("messages_total_bytes", 0) + tr_count = messages.get("tool_result_count", 0) + tr_bytes = messages.get("tool_result_bytes", 0) + tu_count = messages.get("tool_use_count", 0) + else: + msg_bytes = tr_count = tr_bytes = tu_count = 0 + + duration = rec.get("duration_ms", 0) + first_byte = rec.get("first_byte_ms") or 0 + + # Find compaction for this turn (closest timestamp before response) + turn_evictions = 0 + turn_comp_saved = 0 + turn_faults = 0 + + # Simple: find compaction/fault records between request and response + req_ts = pending_request.get("timestamp", "") + resp_ts = rec.get("timestamp", "") + for cts, crec in compaction_by_ts.items(): + if req_ts <= cts <= resp_ts: + turn_evictions += crec.get("evicted", 0) + turn_comp_saved += crec.get("bytes_saved", 0) + for fts, frec in fault_by_ts.items(): + if req_ts <= fts <= resp_ts: + turn_faults += frec.get("count", 0) + + turn = TurnMetrics( + turn=turn_num, + timestamp=rec["timestamp"], + model=pending_request.get("model", "unknown"), + input_tokens=input_tokens, + output_tokens=output_tokens, + cache_creation_tokens=cache_creation, + cache_read_tokens=cache_read, + total_request_bytes=pending_request.get("total_request_bytes", 0), + messages_bytes=msg_bytes, + system_prompt_bytes=sp_bytes, + tool_result_count=tr_count, + tool_result_bytes=tr_bytes, + tool_use_count=tu_count, + evictions=turn_evictions, + compaction_bytes_saved=turn_comp_saved, + faults=turn_faults, + duration_ms=duration, + first_byte_ms=first_byte, + ) + summary.turns.append(turn) + + # Accumulate + summary.api_calls += 1 + summary.total_input_tokens += input_tokens + summary.total_output_tokens += output_tokens + summary.total_cache_creation += cache_creation + summary.total_cache_read += cache_read + summary.total_request_bytes += pending_request.get( + "total_request_bytes", 0 + ) + summary.total_messages_bytes += msg_bytes + summary.total_system_prompt_bytes += sp_bytes + summary.total_tool_result_bytes += tr_bytes + summary.total_evictions += turn_evictions + summary.total_compaction_bytes_saved += turn_comp_saved + summary.total_faults += turn_faults + summary.total_duration_ms += duration + if msg_bytes > summary.max_messages_bytes: + summary.max_messages_bytes = msg_bytes + if input_tokens > summary.max_input_tokens: + summary.max_input_tokens = input_tokens + + pending_request = None + + summary.total_tokens = summary.total_input_tokens + summary.total_output_tokens + summary.total_effective_input = ( + summary.total_input_tokens + + summary.total_cache_creation + + summary.total_cache_read + ) + + if summary.total_evictions > 0: + summary.fault_rate = summary.total_faults / summary.total_evictions + + # Wall clock + if first_timestamp and last_timestamp: + try: + t0 = datetime.fromisoformat(first_timestamp) + t1 = datetime.fromisoformat(last_timestamp) + summary.wall_clock_seconds = (t1 - t0).total_seconds() + except (ValueError, TypeError): + pass + + # Average first byte latency + fb_times = [t.first_byte_ms for t in summary.turns if t.first_byte_ms > 0] + if fb_times: + summary.avg_first_byte_ms = sum(fb_times) / len(fb_times) + + return summary + + +def print_run_summary(s: RunSummary) -> None: + """Print a human-readable summary of one run.""" + print(f"\n{'='*60}") + print(f"Run: {s.label}") + print(f"Log: {s.proxy_log}") + print(f"{'='*60}") + + print(f"\nAPI calls: {s.api_calls:>10,}") + print(f"Wall clock: {s.wall_clock_seconds:>10.0f}s") + print(f"Avg first byte: {s.avg_first_byte_ms:>10.0f}ms") + + print(f"\nTokens:") + print(f" Input (non-cached):{s.total_input_tokens:>12,}") + print(f" Cache creation: {s.total_cache_creation:>12,}") + print(f" Cache read: {s.total_cache_read:>12,}") + print(f" Effective input: {s.total_effective_input:>12,} (context size)") + print(f" Output: {s.total_output_tokens:>12,}") + print(f" Avg effective/call:{s.avg_effective_input:>12,.0f}") + print(f" Max input (1 call):{s.max_input_tokens:>12,}") + + print(f"\nBytes:") + print(f" Total request: {s.total_request_bytes:>12,}") + print(f" Messages: {s.total_messages_bytes:>12,}") + print(f" System prompt: {s.total_system_prompt_bytes:>12,}") + print(f" Tool results: {s.total_tool_result_bytes:>12,}") + print(f" Sys prompt %: {s.system_prompt_fraction:>11.1%}") + print(f" Max messages: {s.max_messages_bytes:>12,}") + + if s.total_evictions > 0: + print(f"\nCompaction:") + print(f" Evictions: {s.total_evictions:>12,}") + print(f" Bytes saved: {s.total_compaction_bytes_saved:>12,}") + print(f" Faults: {s.total_faults:>12,}") + print(f" Fault rate: {s.fault_rate:>11.2%}") + + # Context growth curve (show every Nth turn for readability) + if len(s.turns) > 5: + print(f"\nContext growth (effective input tokens per call):") + step = max(1, len(s.turns) // 15) + cum_input = 0 + max_eff = max((t.effective_input_tokens for t in s.turns), default=1) + for i, t in enumerate(s.turns): + eff = t.effective_input_tokens + cum_input += eff + if i % step == 0 or i == len(s.turns) - 1: + bar_len = min(50, eff * 50 // max(1, max_eff)) + bar = "#" * bar_len + print(f" T{t.turn:>4d}: {eff:>8,} {bar}") + print(f" Cumulative (n²): {cum_input:>12,}") + + +def print_comparison(summaries: list[RunSummary]) -> None: + """Print side-by-side comparison of treatment runs.""" + if not summaries: + return + + baseline = summaries[0] + print(f"\n{'='*70}") + print("TREATMENT COMPARISON") + print(f"{'='*70}") + + # Header + labels = [s.label for s in summaries] + header = f"{'Metric':<30s}" + for label in labels: + header += f" {label:>14s}" + if len(summaries) > 1: + header += f" {'vs baseline':>14s}" + print(f"\n{header}") + print("-" * len(header)) + + def row(metric: str, values: list, fmt: str = ",d", pct: bool = True): + line = f"{metric:<30s}" + for v in values: + line += f" {v:>14{fmt}}" + if pct and len(values) > 1 and values[0] != 0: + delta = (values[-1] - values[0]) / values[0] + line += f" {delta:>+13.1%}" + print(line) + + row("API calls", [s.api_calls for s in summaries]) + row("Effective input tokens", [s.total_effective_input for s in summaries]) + row(" Non-cached", [s.total_input_tokens for s in summaries]) + row(" Cache creation", [s.total_cache_creation for s in summaries]) + row(" Cache read", [s.total_cache_read for s in summaries]) + row("Output tokens", [s.total_output_tokens for s in summaries]) + row("Avg eff input/call", [int(s.avg_effective_input) for s in summaries]) + row("Max input (1 call)", [s.max_input_tokens for s in summaries]) + row("Wall clock (s)", [int(s.wall_clock_seconds) for s in summaries]) + row("Avg first byte (ms)", + [int(s.avg_first_byte_ms) for s in summaries]) + row("Total request bytes", [s.total_request_bytes for s in summaries]) + row("System prompt bytes", [s.total_system_prompt_bytes for s in summaries]) + row("Tool result bytes", [s.total_tool_result_bytes for s in summaries]) + row("Evictions", [s.total_evictions for s in summaries]) + row("Faults", [s.total_faults for s in summaries]) + + # Cost estimate (Opus pricing as of early 2026) + # Non-cached input: $15/M, Cache write: $18.75/M, Cache read: $1.50/M, Output: $75/M + def _estimate_cost(s: RunSummary) -> tuple[float, float, float, float]: + nc = (s.total_input_tokens / 1e6) * 15.0 + cw = (s.total_cache_creation / 1e6) * 18.75 + cr = (s.total_cache_read / 1e6) * 1.50 + out = (s.total_output_tokens / 1e6) * 75.0 + return nc, cw, cr, out + + print() + costs = [] + for s in summaries: + nc, cw, cr, out = _estimate_cost(s) + total = nc + cw + cr + out + costs.append(total) + print(f" {s.label} est. cost: ${total:,.2f} " + f"(input ${nc:.2f} + cache_write ${cw:.2f} + " + f"cache_read ${cr:.2f} + output ${out:.2f})") + + if len(summaries) > 1 and costs[0] > 0: + savings = costs[0] - costs[-1] + print(f"\n Savings ({summaries[-1].label} vs {baseline.label}): " + f"${savings:,.2f} ({savings/costs[0]:.1%})") + + +def main(): + parser = argparse.ArgumentParser( + description="Evaluate context paging experiment runs" + ) + parser.add_argument( + "--run", + type=Path, + nargs="*", + help="Proxy JSONL log(s) to analyze individually", + ) + parser.add_argument( + "--compare", + nargs="*", + metavar="LABEL=PATH", + help="Compare treatments: baseline=path/proxy.jsonl t1=path/proxy.jsonl", + ) + parser.add_argument( + "--json", action="store_true", + help="Output as JSON", + ) + args = parser.parse_args() + + if not args.run and not args.compare: + parser.print_help() + return + + summaries: list[RunSummary] = [] + + if args.run: + for path in args.run: + s = analyze_run(path) + summaries.append(s) + + if args.compare: + for spec in args.compare: + if "=" in spec: + label, path_str = spec.split("=", 1) + else: + label = Path(spec).stem + path_str = spec + s = analyze_run(Path(path_str), label=label) + summaries.append(s) + + if args.json: + for s in summaries: + out = asdict(s) + # Remove per-turn data from JSON summary (too large) + out["turn_count"] = len(out.pop("turns")) + print(json.dumps(out)) + return + + if args.compare and len(summaries) > 1: + print_comparison(summaries) + else: + for s in summaries: + print_run_summary(s) + + +if __name__ == "__main__": + main() diff --git a/src/mnemosyne/oauth.py b/src/mnemosyne/oauth.py new file mode 100644 index 0000000..7e9a545 --- /dev/null +++ b/src/mnemosyne/oauth.py @@ -0,0 +1,301 @@ +"""OAuth PKCE authentication for Anthropic API via Claude Pro/Max subscription. + +Implements the same OAuth flow as Claude Code / opencode, using Anthropic's +public OAuth client. Users authenticate via browser, and the proxy uses +Bearer tokens instead of API keys. + +Usage:: + + mnemosyne login # opens browser, saves tokens + mnemosyne --no-launch # automatically uses OAuth if tokens exist + +Token storage: ~/.config/mnemosyne/auth.json (XDG_CONFIG_HOME respected) +Adapted from trajectory-labs oauth.py (same flow, different storage path). +""" + +from __future__ import annotations + +import base64 +import hashlib +import json +import logging +import os +import secrets +import tempfile +import time +import webbrowser +from dataclasses import asdict, dataclass +from pathlib import Path +from urllib.parse import urlencode + +import httpx + +logger = logging.getLogger(__name__) + +# --------------------------------------------------------------------------- +# Constants — same as Claude Code's OAuth client +# --------------------------------------------------------------------------- + +CLIENT_ID = "9d1c250a-e61b-44d9-88ed-5944d1962f5e" +TOKEN_URL = "https://console.anthropic.com/v1/oauth/token" +AUTHORIZE_URL = "https://claude.ai/oauth/authorize" +REDIRECT_URI = "https://console.anthropic.com/oauth/code/callback" +SCOPE = "org:create_api_key user:profile user:inference" + +_CONFIG_DIR_NAME = "mnemosyne" +TOKEN_FILE = ( + Path(os.environ.get("XDG_CONFIG_HOME") or (Path.home() / ".config")) + / _CONFIG_DIR_NAME + / "auth.json" +) + +# Refresh this many seconds before the token actually expires. +_EXPIRY_BUFFER_SECS = 60 + + +# --------------------------------------------------------------------------- +# Token dataclass +# --------------------------------------------------------------------------- + + +@dataclass +class OAuthTokens: + """Container for OAuth token data.""" + + access_token: str + refresh_token: str + expires_at: float # Unix timestamp + + @property + def expired(self) -> bool: + """True when the access token is expired (or about to expire).""" + return time.time() >= self.expires_at - _EXPIRY_BUFFER_SECS + + +# --------------------------------------------------------------------------- +# PKCE helpers +# --------------------------------------------------------------------------- + + +def _generate_pkce() -> tuple[str, str]: + """Generate a PKCE code verifier and its S256 challenge.""" + verifier = secrets.token_urlsafe(32) + digest = hashlib.sha256(verifier.encode("ascii")).digest() + challenge = base64.urlsafe_b64encode(digest).rstrip(b"=").decode("ascii") + return verifier, challenge + + +# --------------------------------------------------------------------------- +# OAuth flow +# --------------------------------------------------------------------------- + + +def get_authorize_url() -> tuple[str, str, str]: + """Build the OAuth authorization URL. + + Returns: + (authorize_url, verifier, state) + """ + verifier, challenge = _generate_pkce() + params = { + "code": "true", + "client_id": CLIENT_ID, + "response_type": "code", + "redirect_uri": REDIRECT_URI, + "scope": SCOPE, + "state": verifier, + "code_challenge": challenge, + "code_challenge_method": "S256", + } + url = f"{AUTHORIZE_URL}?{urlencode(params)}" + return url, verifier, verifier + + +def exchange_code(code: str, verifier: str) -> OAuthTokens: + """Exchange an authorization code for access + refresh tokens.""" + parts = code.strip().split("#", 1) + auth_code = parts[0] + state = parts[1] if len(parts) > 1 else None + + with httpx.Client(timeout=30) as client: + payload: dict[str, str] = { + "grant_type": "authorization_code", + "client_id": CLIENT_ID, + "code": auth_code, + "redirect_uri": REDIRECT_URI, + "code_verifier": verifier, + } + if state: + payload["state"] = state + resp = client.post(TOKEN_URL, json=payload) + resp.raise_for_status() + body = resp.json() + + return OAuthTokens( + access_token=body["access_token"], + refresh_token=body["refresh_token"], + expires_at=time.time() + body["expires_in"], + ) + + +def refresh_tokens(refresh_token: str) -> OAuthTokens: + """Use a refresh token to obtain a new access token.""" + with httpx.Client(timeout=30) as client: + resp = client.post( + TOKEN_URL, + json={ + "grant_type": "refresh_token", + "client_id": CLIENT_ID, + "refresh_token": refresh_token, + }, + ) + resp.raise_for_status() + body = resp.json() + + return OAuthTokens( + access_token=body["access_token"], + refresh_token=body.get("refresh_token", refresh_token), + expires_at=time.time() + body["expires_in"], + ) + + +# --------------------------------------------------------------------------- +# Token persistence (atomic writes, 0600 permissions) +# --------------------------------------------------------------------------- + + +def save_tokens(tokens: OAuthTokens) -> None: + """Persist tokens to disk atomically with restrictive permissions (0600).""" + TOKEN_FILE.parent.mkdir(parents=True, exist_ok=True) + content = json.dumps(asdict(tokens), indent=2) + + tmp_fd, tmp_path = tempfile.mkstemp(dir=TOKEN_FILE.parent, prefix=".tokens-", suffix=".tmp") + try: + os.write(tmp_fd, content.encode()) + if hasattr(os, "fchmod"): + os.fchmod(tmp_fd, 0o600) + os.close(tmp_fd) + tmp_fd = -1 + os.replace(tmp_path, TOKEN_FILE) + except BaseException: + if tmp_fd >= 0: + os.close(tmp_fd) + Path(tmp_path).unlink(missing_ok=True) + raise + + +def load_tokens() -> OAuthTokens | None: + """Load tokens from disk, or return None if missing/corrupt.""" + if not TOKEN_FILE.exists(): + return None + try: + data = json.loads(TOKEN_FILE.read_text()) + return OAuthTokens( + access_token=data["access_token"], + refresh_token=data["refresh_token"], + expires_at=data["expires_at"], + ) + except (json.JSONDecodeError, KeyError, TypeError): + logger.warning("Corrupt token file at %s — ignoring", TOKEN_FILE) + return None + + +# --------------------------------------------------------------------------- +# Auth resolution — OAuth first, then API key fallback +# --------------------------------------------------------------------------- + + +def get_auth_token() -> str | None: + """Get a valid Bearer token, refreshing if needed. + + Returns the access_token string, or None if no OAuth tokens are saved. + Does NOT fall back to ANTHROPIC_API_KEY — the caller decides. + """ + tokens = load_tokens() + if tokens is None: + return None + + if tokens.expired: + logger.info("OAuth access token expired — refreshing") + try: + tokens = refresh_tokens(tokens.refresh_token) + save_tokens(tokens) + except Exception: + logger.exception("Failed to refresh OAuth token") + return None + + return tokens.access_token + + +def configure_environment() -> str: + """Set up environment for Anthropic SDK to use OAuth. + + Resolution order: + 1. Existing ANTHROPIC_AUTH_TOKEN env var (already set) + 2. OAuth tokens from ~/.config/mnemosyne/auth.json + 3. ANTHROPIC_API_KEY env var (no changes needed) + + Returns a description of the auth method being used. + """ + # Already configured via env + if os.environ.get("ANTHROPIC_AUTH_TOKEN"): + return "bearer (ANTHROPIC_AUTH_TOKEN env var)" + + # Try OAuth + token = get_auth_token() + if token is not None: + # The Anthropic SDK reads ANTHROPIC_AUTH_TOKEN and sends + # Authorization: Bearer instead of X-Api-Key. + os.environ["ANTHROPIC_AUTH_TOKEN"] = token + return f"oauth (token from {TOKEN_FILE})" + + # Fall back to API key + if os.environ.get("ANTHROPIC_API_KEY"): + return "api-key (ANTHROPIC_API_KEY env var)" + + return "none (no authentication configured)" + + +def has_auth() -> bool: + """Check if any Anthropic authentication is available.""" + if os.environ.get("ANTHROPIC_AUTH_TOKEN"): + return True + if os.environ.get("ANTHROPIC_API_KEY"): + return True + return load_tokens() is not None + + +# --------------------------------------------------------------------------- +# Interactive login +# --------------------------------------------------------------------------- + + +def login_interactive() -> None: + """Run the browser-based OAuth PKCE login flow. + + Opens the authorization URL in the default browser, waits for the user + to paste the authorization code, exchanges it for tokens, and saves them. + """ + url, verifier, expected_state = get_authorize_url() + print(f"\n Opening browser for OAuth login...\n {url}\n") + webbrowser.open(url) + code = input(" Paste the authorization code here: ").strip() + if not code: + print(" No code provided — aborting.") + return + + # Validate CSRF state if callback included it (code#state format) + if "#" in code: + _, returned_state = code.split("#", 1) + if returned_state != expected_state: + print(" State mismatch — possible CSRF. Aborting.") + return + + try: + tokens = exchange_code(code, verifier) + except httpx.HTTPStatusError as exc: + print(f" Token exchange failed: {exc.response.status_code} {exc.response.text}") + return + + save_tokens(tokens) + print(f" Login successful! Token saved to {TOKEN_FILE}\n") diff --git a/src/mnemosyne/replay.py b/src/mnemosyne/replay.py new file mode 100644 index 0000000..534998b --- /dev/null +++ b/src/mnemosyne/replay.py @@ -0,0 +1,326 @@ +"""Offline replay of context paging on proxy logs. + +Reads proxy JSONL logs (from observe-mode runs), reconstructs the +messages array at each API call turn, applies pager compaction, and +reports what WOULD have happened: how much content would be evicted, +and whether the model's actual next actions would have triggered page +faults. + +This enables "what-if" analysis: run once without compaction, then +replay offline with different thresholds to find optimal parameters +without burning API credits. + +Usage: + # Replay one session + python -m pichay.replay logs/proxy_20260302_024551.jsonl + + # Replay with custom thresholds + python -m pichay.replay --age-threshold 6 --min-size 1000 logs/*.jsonl + + # JSON output for downstream analysis + python -m pichay.replay --json logs/proxy_*.jsonl +""" + +from __future__ import annotations + +import argparse +import copy +import json +import sys +from dataclasses import dataclass, field, asdict +from pathlib import Path + +from mnemosyne.eval import parse_proxy_log +from mnemosyne.pager import PageStore, compact_messages + + +@dataclass +class ReplayTurn: + """Results of simulated compaction at one API call.""" + + turn: int + timestamp: str + message_count: int + bytes_original: int + bytes_compacted: int + bytes_saved: int + reduction_pct: float + evictions: int + faults: int + cumulative_evictions: int + cumulative_faults: int + + +@dataclass +class SessionReplay: + """Aggregate replay results for one proxy log.""" + + log_path: str + total_turns: int = 0 + total_evictions: int = 0 + total_bytes_saved: int = 0 + total_bytes_original: int = 0 + total_faults: int = 0 + turns: list[ReplayTurn] = field(default_factory=list) + + @property + def fault_rate(self) -> float: + if self.total_evictions == 0: + return 0.0 + return self.total_faults / self.total_evictions + + @property + def reduction_pct(self) -> float: + if self.total_bytes_original == 0: + return 0.0 + return (self.total_bytes_saved / self.total_bytes_original) * 100 + + +def replay_session( + path: Path, + age_threshold: int = 4, + min_size: int = 500, +) -> SessionReplay: + """Replay compaction on a proxy log session. + + For each API call in the log, reconstructs the messages array, + applies cumulative compaction (simulating what would have happened + if the pager had been active from the start), and checks for + page faults against the model's actual next actions. + """ + records = parse_proxy_log(path) + result = SessionReplay(log_path=str(path)) + + requests = [ + r for r in records + if r.get("type") == "request" and "messages_full" in r + ] + + if not requests: + return result + + page_store = PageStore() + + for turn_idx, req in enumerate(requests): + messages = copy.deepcopy(req["messages_full"]) + bytes_original = len(json.dumps(messages).encode("utf-8")) + + # Apply previous evictions (simulate cumulative compaction) + _apply_evictions(messages, page_store) + + # Detect faults: did the model re-request evicted content? + faults = page_store.detect_faults(messages) + + # Run compaction on this turn + stats = compact_messages( + messages, + age_threshold=age_threshold, + min_size=min_size, + page_store=page_store, + ) + + bytes_compacted = len(json.dumps(messages).encode("utf-8")) + bytes_saved = bytes_original - bytes_compacted + + turn = ReplayTurn( + turn=turn_idx + 1, + timestamp=req.get("timestamp", ""), + message_count=len(req["messages_full"]), + bytes_original=bytes_original, + bytes_compacted=bytes_compacted, + bytes_saved=bytes_saved, + reduction_pct=( + bytes_saved / bytes_original * 100 + if bytes_original > 0 + else 0.0 + ), + evictions=stats.evicted_count, + faults=len(faults), + cumulative_evictions=page_store.cumulative_evictions, + cumulative_faults=len(page_store.faults), + ) + result.turns.append(turn) + result.total_turns += 1 + result.total_evictions += stats.evicted_count + result.total_bytes_saved += bytes_saved + result.total_bytes_original += bytes_original + result.total_faults += len(faults) + + return result + + +def _apply_evictions(messages: list[dict], page_store: PageStore) -> None: + """Replace tool result content for previously evicted results.""" + for msg in messages: + if msg.get("role") != "user": + continue + content = msg.get("content", []) + if not isinstance(content, list): + continue + for i, block in enumerate(content): + if not isinstance(block, dict): + continue + if block.get("type") != "tool_result": + continue + tool_use_id = block.get("tool_use_id", "") + entry = page_store.retrieve(tool_use_id) + if entry is not None: + content[i] = {**block, "content": entry.summary} + + +def print_session_report(s: SessionReplay) -> None: + """Print a human-readable replay report.""" + print(f"\n{'='*60}") + print(f"Replay: {s.log_path}") + print(f"{'='*60}") + + print(f"\nTurns: {s.total_turns:>10,}") + print(f"Evictions: {s.total_evictions:>10,}") + print(f"Bytes saved: {s.total_bytes_saved:>10,}") + print(f"Bytes original: {s.total_bytes_original:>10,}") + print(f"Reduction: {s.reduction_pct:>9.1f}%") + print(f"Faults: {s.total_faults:>10,}") + print(f"Fault rate: {s.fault_rate:>9.2%}") + + if not s.turns: + return + + # Per-turn detail + print(f"\nPer-turn:") + print( + f" {'Turn':>4s} {'Original':>10s} {'Compacted':>10s} " + f"{'Saved':>10s} {'Reduc%':>6s} {'Evict':>5s} {'Fault':>5s}" + ) + for t in s.turns: + print( + f" T{t.turn:>3d} {t.bytes_original:>10,} " + f"{t.bytes_compacted:>10,} {t.bytes_saved:>10,} " + f"{t.reduction_pct:>5.1f}% " + f"{t.evictions:>5d} {t.faults:>5d}" + ) + + # Cumulative curve (final state) + last = s.turns[-1] + print(f"\nFinal turn context: {last.bytes_original:,} → " + f"{last.bytes_compacted:,} bytes " + f"({last.reduction_pct:.1f}% reduction)") + + +def print_aggregate(sessions: list[SessionReplay]) -> None: + """Print aggregate stats across multiple sessions.""" + if not sessions: + return + + total_turns = sum(s.total_turns for s in sessions) + total_evictions = sum(s.total_evictions for s in sessions) + total_bytes_saved = sum(s.total_bytes_saved for s in sessions) + total_bytes_original = sum(s.total_bytes_original for s in sessions) + total_faults = sum(s.total_faults for s in sessions) + + fault_rate = ( + total_faults / total_evictions if total_evictions > 0 else 0.0 + ) + reduction_pct = ( + total_bytes_saved / total_bytes_original * 100 + if total_bytes_original > 0 + else 0.0 + ) + + print(f"\n{'='*60}") + print(f"AGGREGATE ({len(sessions)} sessions)") + print(f"{'='*60}") + print(f"Total turns: {total_turns:>10,}") + print(f"Total evictions: {total_evictions:>10,}") + print(f"Total bytes saved: {total_bytes_saved:>10,}") + print(f"Total bytes original:{total_bytes_original:>10,}") + print(f"Reduction: {reduction_pct:>9.1f}%") + print(f"Faults: {total_faults:>10,}") + print(f"Fault rate: {fault_rate:>9.2%}") + + +def main(): + parser = argparse.ArgumentParser( + description="Offline replay of context paging on proxy logs" + ) + parser.add_argument( + "logs", + type=Path, + nargs="+", + help="Proxy JSONL log file(s) to replay", + ) + parser.add_argument( + "--age-threshold", + type=int, + default=4, + help="Evict tool results older than N user-turns (default: 4)", + ) + parser.add_argument( + "--min-size", + type=int, + default=500, + help="Don't evict results smaller than N bytes (default: 500)", + ) + parser.add_argument( + "--json", + action="store_true", + help="Output as JSON", + ) + args = parser.parse_args() + + sessions: list[SessionReplay] = [] + for log_path in args.logs: + if not log_path.exists(): + print( + f"Warning: {log_path} not found, skipping.", + file=sys.stderr, + ) + continue + sessions.append(replay_session( + log_path, + age_threshold=args.age_threshold, + min_size=args.min_size, + )) + + if args.json: + for s in sessions: + out = asdict(s) + out["fault_rate"] = s.fault_rate + out["reduction_pct"] = s.reduction_pct + out["turn_count"] = len(out.pop("turns")) + print(json.dumps(out)) + if len(sessions) > 1: + total_e = sum(s.total_evictions for s in sessions) + total_o = sum(s.total_bytes_original for s in sessions) + agg = { + "type": "aggregate", + "sessions": len(sessions), + "total_turns": sum(s.total_turns for s in sessions), + "total_evictions": total_e, + "total_bytes_saved": sum( + s.total_bytes_saved for s in sessions + ), + "total_bytes_original": total_o, + "total_faults": sum(s.total_faults for s in sessions), + "fault_rate": ( + sum(s.total_faults for s in sessions) / total_e + if total_e > 0 + else 0.0 + ), + "reduction_pct": ( + sum(s.total_bytes_saved for s in sessions) / total_o * 100 + if total_o > 0 + else 0.0 + ), + } + print(json.dumps(agg)) + return + + for s in sessions: + print_session_report(s) + + if len(sessions) > 1: + print_aggregate(sessions) + + +if __name__ == "__main__": + main() diff --git a/src/mnemosyne/telemetry.py b/src/mnemosyne/telemetry.py new file mode 100644 index 0000000..b5717b0 --- /dev/null +++ b/src/mnemosyne/telemetry.py @@ -0,0 +1,292 @@ +from __future__ import annotations + +import json +import threading +from collections import defaultdict, deque +from dataclasses import dataclass +from datetime import datetime, timezone +from pathlib import Path +from typing import Any + +from prometheus_client import Counter, Histogram, generate_latest + + +REQ_TOTAL = Counter( + "pichay_requests_total", + "Total gateway requests", + ["provider", "status"], +) + +REQ_LATENCY_MS = Histogram( + "pichay_request_latency_ms", + "Gateway request latency milliseconds", + ["provider"], +) + +SHRINK_RATIO = Histogram( + "pichay_shrink_ratio", + "Outgoing/incoming payload ratio", + ["provider"], +) + +POLICY_CONFLICTS = Counter( + "pichay_policy_conflicts_total", + "Policy conflict resolution count", + ["winner_stage", "loser_stage"], +) + +ANOMALIES = Counter( + "pichay_anomalies_total", + "Data anomalies detected", + ["kind"], +) + +CACHE_READ = Counter( + "pichay_cache_read_tokens_total", + "Tokens read from provider cache", + ["provider"], +) + +CACHE_CREATE = Counter( + "pichay_cache_create_tokens_total", + "Tokens written to provider cache", + ["provider"], +) + +CACHE_MISS_EVENTS = Counter( + "pichay_cache_miss_events_total", + "Requests with zero cache read (unexpected misses)", + ["provider"], +) + + +@dataclass +class SessionSummary: + request_count: int = 0 + incoming_bytes: int = 0 + outgoing_bytes: int = 0 + + +class Telemetry: + def __init__(self, log_path: Path, hydration_window_seconds: int, max_events: int = 5000): + self.log_path = log_path + self.hydration_window_seconds = hydration_window_seconds + self._lock = threading.Lock() + self.events: deque[dict[str, Any]] = deque(maxlen=max_events) + self.sessions: dict[str, SessionSummary] = defaultdict(SessionSummary) + + self.log_path.parent.mkdir(parents=True, exist_ok=True) + self._hydrate() + + def _now(self) -> str: + return datetime.now(timezone.utc).isoformat() + + def _hydrate(self) -> None: + if not self.log_path.exists(): + return + cutoff = datetime.now(timezone.utc).timestamp() - self.hydration_window_seconds + try: + with open(self.log_path, "r", encoding="utf-8") as f: + for line in f: + try: + event = json.loads(line) + except json.JSONDecodeError: + continue + ts = event.get("timestamp") + if not isinstance(ts, str): + continue + try: + epoch = datetime.fromisoformat(ts).timestamp() + except ValueError: + continue + if epoch < cutoff: + continue + self.events.append(event) + except OSError: + return + + def emit(self, event_type: str, **fields: Any) -> None: + record = { + "type": event_type, + "timestamp": self._now(), + **fields, + } + with self._lock: + self.events.append(record) + with open(self.log_path, "a", encoding="utf-8") as f: + f.write(json.dumps(record, default=str) + "\n") + + if event_type == "policy_conflict_resolved": + POLICY_CONFLICTS.labels( + winner_stage=fields.get("winner_stage", "unknown"), + loser_stage=fields.get("loser_stage", "unknown"), + ).inc() + if event_type == "anomaly": + ANOMALIES.labels(kind=fields.get("kind", "unknown")).inc() + + def record_request( + self, + *, + session_id: str, + provider: str, + status: int, + incoming_bytes: int, + outgoing_bytes: int, + latency_ms: float, + streaming: bool, + model: str, + request_id: str, + duplication_score: float, + usage: dict[str, Any] | None = None, + messages_full: list[dict] | None = None, + ) -> None: + shrink_ratio = (outgoing_bytes / incoming_bytes) if incoming_bytes > 0 else 1.0 + + with self._lock: + s = self.sessions[session_id] + s.request_count += 1 + s.incoming_bytes += incoming_bytes + s.outgoing_bytes += outgoing_bytes + + REQ_TOTAL.labels(provider=provider, status=str(status)).inc() + REQ_LATENCY_MS.labels(provider=provider).observe(latency_ms) + SHRINK_RATIO.labels(provider=provider).observe(shrink_ratio) + + # Cache analysis + input_tokens = 0 + cache_read = 0 + cache_create = 0 + cache_read_pct = 0.0 + effective_tokens = 0 + if usage: + input_tokens = usage.get("input_tokens", 0) + cache_read = usage.get("cache_read_input_tokens", 0) + cache_create = usage.get("cache_creation_input_tokens", 0) + effective_tokens = input_tokens + cache_read + cache_create + cache_read_pct = (cache_read / effective_tokens * 100) if effective_tokens > 0 else 0.0 + if cache_read: + CACHE_READ.labels(provider=provider).inc(cache_read) + if cache_create: + CACHE_CREATE.labels(provider=provider).inc(cache_create) + + # Token-equivalent economics for subscription-style pricing: + # cached read ~= 0.1x cost; miss ~= 1.0x cost. + # A full miss therefore carries roughly a 0.9x premium on effective tokens. + miss_penalty_tokens_est = 0.0 + if effective_tokens > 0 and cache_read == 0: + miss_penalty_tokens_est = 0.9 * effective_tokens + + # Rough local savings signal from payload shrink (bytes to tokens @ ~4 bytes/token). + size_saved_tokens_est = max(0.0, (incoming_bytes - outgoing_bytes) / 4.0) + net_token_value_est = size_saved_tokens_est - miss_penalty_tokens_est + + self.emit( + "request_metrics", + request_id=request_id, + session_id=session_id, + provider=provider, + model=model, + status=status, + incoming_bytes=incoming_bytes, + outgoing_bytes=outgoing_bytes, + shrink_ratio=shrink_ratio, + duplication_score=duplication_score, + latency_ms=latency_ms, + streaming=streaming, + input_tokens=input_tokens, + cache_read_tokens=cache_read, + cache_create_tokens=cache_create, + effective_tokens=effective_tokens, + cache_read_pct=round(cache_read_pct, 1), + miss_penalty_tokens_est=round(miss_penalty_tokens_est, 1), + size_saved_tokens_est=round(size_saved_tokens_est, 1), + net_token_value_est=round(net_token_value_est, 1), + messages_full=messages_full or [], + ) + + # Small increases are expected — Pichay injects tensor handles, + # yuyay manifests, system reminders. Only flag when growth exceeds + # 5% of incoming (suggesting duplicated message blocks, not injection). + if incoming_bytes > 0 and outgoing_bytes > incoming_bytes: + growth_pct = (outgoing_bytes - incoming_bytes) / incoming_bytes + if growth_pct > 0.05: + self.emit( + "anomaly", + kind="outgoing_growth_suspicious", + request_id=request_id, + session_id=session_id, + provider=provider, + incoming_bytes=incoming_bytes, + outgoing_bytes=outgoing_bytes, + growth_pct=round(growth_pct * 100, 1), + duplication_score=duplication_score, + ) + + # Flag unexpected cache misses: request had enough tokens to cache + # but got zero cache reads. Skip first request per session (cold start). + if usage and status == 200 and s.request_count > 1: + if effective_tokens > 4096 and cache_read == 0 and cache_create > 0: + CACHE_MISS_EVENTS.labels(provider=provider).inc() + self.emit( + "cache_miss_unexpected", + request_id=request_id, + session_id=session_id, + provider=provider, + model=model, + effective_tokens=effective_tokens, + cache_create_tokens=cache_create, + ) + + def cost_summary(self, window_seconds: int | None = None) -> dict[str, Any]: + events = self.recent_events(window_seconds) + req = [e for e in events if e.get("type") == "request_metrics"] + if not req: + return { + "requests": 0, + "avg_cache_read_pct": 0.0, + "zero_cache_read_requests": 0, + "avg_effective_tokens": 0.0, + "avg_miss_penalty_tokens_est": 0.0, + "avg_size_saved_tokens_est": 0.0, + "avg_net_token_value_est": 0.0, + } + n = len(req) + return { + "requests": n, + "avg_cache_read_pct": round(sum(float(e.get("cache_read_pct", 0.0)) for e in req) / n, 2), + "zero_cache_read_requests": sum(1 for e in req if float(e.get("cache_read_tokens", 0.0)) == 0.0), + "avg_effective_tokens": round(sum(float(e.get("effective_tokens", 0.0)) for e in req) / n, 1), + "avg_miss_penalty_tokens_est": round(sum(float(e.get("miss_penalty_tokens_est", 0.0)) for e in req) / n, 1), + "avg_size_saved_tokens_est": round(sum(float(e.get("size_saved_tokens_est", 0.0)) for e in req) / n, 1), + "avg_net_token_value_est": round(sum(float(e.get("net_token_value_est", 0.0)) for e in req) / n, 1), + } + + def get_metrics(self) -> bytes: + return generate_latest() + + def recent_events(self, window_seconds: int | None = None) -> list[dict[str, Any]]: + if window_seconds is None: + return list(self.events) + cutoff = datetime.now(timezone.utc).timestamp() - window_seconds + out: list[dict[str, Any]] = [] + for e in self.events: + ts = e.get("timestamp") + if not isinstance(ts, str): + continue + try: + epoch = datetime.fromisoformat(ts).timestamp() + except ValueError: + continue + if epoch >= cutoff: + out.append(e) + return out + + def session_summary(self) -> dict[str, dict[str, int]]: + return { + sid: { + "request_count": s.request_count, + "incoming_bytes": s.incoming_bytes, + "outgoing_bytes": s.outgoing_bytes, + } + for sid, s in self.sessions.items() + } diff --git a/tests/test_benchmark.py b/tests/test_benchmark.py new file mode 100644 index 0000000..339c76a --- /dev/null +++ b/tests/test_benchmark.py @@ -0,0 +1,812 @@ +"""Tests for the Mnemosyne benchmark module.""" + +from __future__ import annotations + +import time + +from mnemosyne.benchmark import ( + AdmissionMetrics, + BenchmarkCollector, + EntropyMetrics, + FidelityMetrics, + GoalMetrics, + LatencyTracker, + MicroFaultMetrics, + SegmentationMetrics, + SessionBenchmark, + Timer, + TokenMetrics, +) + + +# ── LatencyTracker ────────────────────────────────────────────────────── + + +class TestLatencyTracker: + def test_empty(self): + t = LatencyTracker() + assert t.count == 0 + assert t.avg_ms == 0.0 + d = t.to_dict() + assert d["count"] == 0 + assert d["avg_ms"] == 0 + + def test_record_single(self): + t = LatencyTracker() + t.record(5.0) + assert t.count == 1 + assert t.avg_ms == 5.0 + assert t.min_ms == 5.0 + assert t.max_ms == 5.0 + + def test_record_multiple(self): + t = LatencyTracker() + t.record(2.0) + t.record(8.0) + t.record(5.0) + assert t.count == 3 + assert t.avg_ms == 5.0 + assert t.min_ms == 2.0 + assert t.max_ms == 8.0 + + def test_to_dict(self): + t = LatencyTracker() + t.record(10.0) + t.record(20.0) + d = t.to_dict() + assert d["count"] == 2 + assert d["avg_ms"] == 15.0 + assert d["min_ms"] == 10.0 + assert d["max_ms"] == 20.0 + assert d["total_ms"] == 30.0 + + +# ── Timer ─────────────────────────────────────────────────────────────── + + +class TestTimer: + def test_timer_records(self): + tracker = LatencyTracker() + with Timer(tracker) as t: + time.sleep(0.01) # 10ms + assert t.elapsed_ms > 5 # at least 5ms + assert tracker.count == 1 + assert tracker.total_ms > 5 + + def test_timer_records_on_exception(self): + tracker = LatencyTracker() + try: + with Timer(tracker): + raise ValueError("test") + except ValueError: + pass + assert tracker.count == 1 + + +# ── FidelityMetrics ───────────────────────────────────────────────────── + + +class TestFidelityMetrics: + def test_empty(self): + m = FidelityMetrics() + assert m.degradations == 0 + assert m.upgrades == 0 + + def test_record_degradation(self): + m = FidelityMetrics() + m.record_transition(0, 1) # L0 -> L1 + assert m.degradations == 1 + assert m.upgrades == 0 + assert m.transitions_by_level["L0->L1"] == 1 + + def test_record_upgrade(self): + m = FidelityMetrics() + m.record_transition(2, 0) # L2 -> L0 + assert m.degradations == 0 + assert m.upgrades == 1 + assert m.transitions_by_level["L2->L0"] == 1 + + def test_multiple_transitions(self): + m = FidelityMetrics() + m.record_transition(0, 1) + m.record_transition(0, 1) + m.record_transition(1, 2) + m.record_transition(2, 0) + assert m.degradations == 3 + assert m.upgrades == 1 + assert m.transitions_by_level["L0->L1"] == 2 + + def test_snapshot_levels(self): + m = FidelityMetrics() + m.snapshot_levels({0: 5, 1: 3, 2: 2}, {0: 5000, 1: 1500, 2: 200}) + assert m.objects_by_level["L0"] == 5 + assert m.tokens_by_level["L1"] == 1500 + + def test_to_dict(self): + m = FidelityMetrics() + m.record_transition(0, 1) + d = m.to_dict() + assert "degradations" in d + assert "transitions_by_level" in d + + +# ── AdmissionMetrics ──────────────────────────────────────────────────── + + +class TestAdmissionMetrics: + def test_empty(self): + m = AdmissionMetrics() + assert m.rejection_rate == 0.0 + assert m.avg_score == 0.0 + + def test_record_admitted(self): + m = AdmissionMetrics() + m.record_decision(True, 0.8, "file_context") + assert m.admitted == 1 + assert m.rejected == 0 + assert m.rejection_rate == 0.0 + + def test_record_rejected(self): + m = AdmissionMetrics() + m.record_decision(False, 0.2, "tool_result") + assert m.admitted == 0 + assert m.rejected == 1 + assert m.rejection_rate == 1.0 + assert m.rejected_by_type["tool_result"] == 1 + + def test_mixed_decisions(self): + m = AdmissionMetrics() + m.record_decision(True, 0.8, "file_context") + m.record_decision(True, 0.7, "plan") + m.record_decision(False, 0.2, "tool_result") + assert m.total_evaluated == 3 + assert m.admitted == 2 + assert m.rejected == 1 + assert abs(m.rejection_rate - 1 / 3) < 0.01 + + def test_to_dict(self): + m = AdmissionMetrics() + m.record_decision(True, 0.8, "file_context") + d = m.to_dict() + assert d["admitted"] == 1 + assert d["rejection_rate"] == 0.0 + + +# ── MicroFaultMetrics ─────────────────────────────────────────────────── + + +class TestMicroFaultMetrics: + def test_empty(self): + m = MicroFaultMetrics() + assert m.success_rate == 0.0 + assert m.avg_tokens_saved == 0.0 + + def test_record_fault(self): + m = MicroFaultMetrics() + m.record_fault(tokens_answer=50, tokens_avoided=3000, success=True) + assert m.attempts == 1 + assert m.successes == 1 + assert m.tokens_saved == 3000 + assert m.tokens_used == 50 + assert m.success_rate == 1.0 + + def test_mixed_faults(self): + m = MicroFaultMetrics() + m.record_fault(50, 3000, success=True) + m.record_fault(60, 2000, success=False) + assert m.attempts == 2 + assert m.successes == 1 + assert m.success_rate == 0.5 + assert m.tokens_saved == 5000 + + def test_to_dict_savings_ratio(self): + m = MicroFaultMetrics() + m.record_fault(100, 900, success=True) + d = m.to_dict() + assert d["savings_ratio"] == 0.9 # 900 / (900 + 100) + + +# ── SegmentationMetrics ───────────────────────────────────────────────── + + +class TestSegmentationMetrics: + def test_empty(self): + m = SegmentationMetrics() + assert m.total_objects_created == 0 + assert m.avg_object_size == 0.0 + + def test_record_objects(self): + m = SegmentationMetrics() + m.record_object("file_context", 500) + m.record_object("file_context", 1000) + m.record_object("tool_result", 200) + assert m.total_objects_created == 3 + assert m.objects_by_type["file_context"] == 2 + assert m.objects_by_type["tool_result"] == 1 + assert m.total_tokens_stored == 1700 + + def test_to_dict_percentiles(self): + m = SegmentationMetrics() + for i in range(10): + m.record_object("test", (i + 1) * 100) + d = m.to_dict() + assert d["min_object_size_tokens"] == 100 + assert d["max_object_size_tokens"] == 1000 + assert d["p50_object_size_tokens"] == 600 # nearest-rank: index 5 of [100..1000] + + +# ── GoalMetrics ───────────────────────────────────────────────────────── + + +class TestGoalMetrics: + def test_empty(self): + m = GoalMetrics() + assert m.topic_shifts_detected == 0 + + def test_record_events(self): + m = GoalMetrics() + m.record_topic_shift() + m.record_reclassification() + m.record_promotion(3) + d = m.to_dict() + assert d["topic_shifts_detected"] == 1 + assert d["goal_reclassifications"] == 1 + assert d["promotions_triggered"] == 3 + + +# ── EntropyMetrics ────────────────────────────────────────────────────── + + +class TestEntropyMetrics: + def test_empty(self): + m = EntropyMetrics() + assert m.trigger_rate == 0.0 + + def test_record_checks(self): + m = EntropyMetrics() + m.record_check(triggered=False) + m.record_check(triggered=True, entities_count=2) + m.record_check(triggered=False) + assert m.checks == 3 + assert m.triggers == 1 + assert m.entities_faulted == 2 + assert abs(m.trigger_rate - 1 / 3) < 0.01 + + +# ── TokenMetrics ──────────────────────────────────────────────────────── + + +class TestTokenMetrics: + def test_empty(self): + m = TokenMetrics() + assert m.turns == 0 + assert m.context_reduction_ratio == 1.0 + + def test_record_turn(self): + m = TokenMetrics() + m.record_turn( + input_tokens=10000, + effective_tokens=15000, + cache_read=3000, + cache_create=2000, + incoming_bytes=50000, + outgoing_bytes=30000, + ) + assert m.turns == 1 + assert m.total_input_tokens == 10000 + assert m.total_effective_tokens == 15000 + assert m.total_cache_read == 3000 + + def test_context_reduction(self): + m = TokenMetrics() + m.record_turn(0, 0, 0, 0, incoming_bytes=10000, outgoing_bytes=4000) + assert m.context_reduction_ratio == 0.4 # 60% reduction + d = m.to_dict() + assert d["context_reduction_pct"] == 60.0 + + def test_cache_hit_rate(self): + m = TokenMetrics() + m.record_turn(0, 0, cache_read=8000, cache_create=2000, incoming_bytes=0, outgoing_bytes=0) + assert m.avg_cache_hit_rate == 0.8 + + def test_multiple_turns(self): + m = TokenMetrics() + m.record_turn(5000, 7000, 1500, 500, 20000, 15000) + m.record_turn(8000, 12000, 3000, 1000, 40000, 20000) + assert m.turns == 2 + assert m.total_input_tokens == 13000 + d = m.to_dict() + assert d["peak_input_tokens"] == 8000 + + +# ── SessionBenchmark ──────────────────────────────────────────────────── + + +class TestSessionBenchmark: + def test_creation(self): + b = SessionBenchmark("test-session") + assert b.session_id == "test-session" + assert b.elapsed_seconds >= 0 + + def test_to_dict_has_all_sections(self): + b = SessionBenchmark("test-session") + d = b.to_dict() + assert "session_id" in d + assert "tokens" in d + assert "fidelity" in d + assert "admission" in d + assert "micro_faults" in d + assert "segmentation" in d + assert "goals" in d + assert "entropy" in d + assert "latency" in d + + def test_latency_trackers_initialized(self): + b = SessionBenchmark("test") + assert "preprocess" in b.latency + assert "segmentation" in b.latency + assert "embedding" in b.latency + assert "admission" in b.latency + assert "goal_classification" in b.latency + assert "entropy_check" in b.latency + assert "memory_query" in b.latency + assert "helper_llm" in b.latency + + +# ── BenchmarkCollector ────────────────────────────────────────────────── + + +class TestBenchmarkCollector: + def test_empty_aggregate(self): + c = BenchmarkCollector() + result = c.aggregate() + assert result["sessions"] == 0 + + def test_get_session_creates(self): + c = BenchmarkCollector() + s = c.get_session("s1") + assert s.session_id == "s1" + # Getting again returns same instance + s2 = c.get_session("s1") + assert s is s2 + + def test_aggregate_single_session(self): + c = BenchmarkCollector() + s = c.get_session("s1") + s.tokens.record_turn(5000, 7000, 1500, 500, 20000, 12000) + s.tokens.record_turn(8000, 12000, 3000, 1000, 40000, 18000) + s.fidelity.record_transition(0, 1) + s.admission.record_decision(True, 0.8, "file_context") + s.admission.record_decision(False, 0.2, "tool_result") + + result = c.aggregate() + assert result["sessions"] == 1 + assert result["total_turns"] == 2 + assert result["tokens"]["total_input_tokens"] == 13000 + assert result["fidelity"]["total_degradations"] == 1 + assert result["admission"]["total_admitted"] == 1 + assert result["admission"]["total_rejected"] == 1 + + def test_aggregate_multiple_sessions(self): + c = BenchmarkCollector() + s1 = c.get_session("s1") + s1.tokens.record_turn(5000, 7000, 1500, 500, 20000, 12000) + s1.fidelity.record_transition(0, 1) + + s2 = c.get_session("s2") + s2.tokens.record_turn(3000, 4000, 800, 200, 10000, 7000) + s2.fidelity.record_transition(1, 2) + + result = c.aggregate() + assert result["sessions"] == 2 + assert result["total_turns"] == 2 + assert result["tokens"]["total_input_tokens"] == 8000 + assert result["fidelity"]["total_degradations"] == 2 + + def test_session_report(self): + c = BenchmarkCollector() + s = c.get_session("s1") + s.segmentation.record_object("file_context", 500) + report = c.session_report("s1") + assert report is not None + assert report["segmentation"]["total_objects_created"] == 1 + + def test_session_report_not_found(self): + c = BenchmarkCollector() + assert c.session_report("nonexistent") is None + + def test_all_sessions(self): + c = BenchmarkCollector() + c.get_session("s1") + c.get_session("s2") + all_s = c.all_sessions() + assert len(all_s) == 2 + assert "s1" in all_s + assert "s2" in all_s + + def test_reset(self): + c = BenchmarkCollector() + c.get_session("s1") + c.reset() + assert c.aggregate()["sessions"] == 0 + + def test_context_reduction_calculation(self): + c = BenchmarkCollector() + s = c.get_session("s1") + # 100k incoming, 20k outgoing = 80% reduction + s.tokens.record_turn(0, 0, 0, 0, 100000, 20000) + result = c.aggregate() + assert result["tokens"]["context_reduction_pct"] == 80.0 + + +# ── suggest_thresholds ────────────────────────────────────────────────── + +from mnemosyne.benchmark import suggest_thresholds + + +class TestSuggestThresholds: + def test_high_rejection_rate_suggests_lower_threshold(self): + metrics = { + "admission": {"rejection_rate": 0.45, "total_admitted": 55, "total_rejected": 45}, + "fidelity": {"total_degradations": 10, "total_upgrades": 10}, + "micro_faults": {"total_attempts": 10, "success_rate": 0.8, "total_tokens_saved": 500}, + } + result = suggest_thresholds(metrics) + assert result["admission_threshold"]["status"] == "adjust" + assert "lowering" in result["admission_threshold"]["reason"] + assert ( + result["admission_threshold"]["suggested_value"] + < result["admission_threshold"]["current_value"] + ) + + def test_low_rejection_rate_suggests_higher_threshold(self): + metrics = { + "admission": {"rejection_rate": 0.05, "total_admitted": 95, "total_rejected": 5}, + "fidelity": {"total_degradations": 5, "total_upgrades": 5}, + "micro_faults": {"total_attempts": 10, "success_rate": 0.8, "total_tokens_saved": 500}, + } + result = suggest_thresholds(metrics) + assert result["admission_threshold"]["status"] == "adjust" + assert "raising" in result["admission_threshold"]["reason"] + + def test_balanced_metrics_all_ok(self): + metrics = { + "admission": {"rejection_rate": 0.20, "total_admitted": 80, "total_rejected": 20}, + "entropy": {"trigger_rate": 0.12, "checks": 100, "triggers": 12}, + "fidelity": {"total_degradations": 10, "total_upgrades": 8}, + "micro_faults": { + "total_attempts": 20, + "success_rate": 0.75, + "total_tokens_saved": 1000, + }, + } + result = suggest_thresholds(metrics) + assert result["admission_threshold"]["status"] == "ok" + assert result["entropy_sensitivity"]["status"] == "ok" + assert result["fidelity_pressure"]["status"] == "ok" + assert result["micro_fault_quality"]["status"] == "ok" + + def test_high_entropy_trigger_rate_suggests_adjustment(self): + metrics = { + "admission": {"rejection_rate": 0.20}, + "entropy": {"trigger_rate": 0.35, "checks": 100, "triggers": 35}, + "fidelity": {"total_degradations": 5, "total_upgrades": 5}, + "micro_faults": {"total_attempts": 0, "success_rate": 0.0}, + } + result = suggest_thresholds(metrics) + assert result["entropy_sensitivity"]["status"] == "adjust" + assert "lowering" in result["entropy_sensitivity"]["reason"] + + def test_low_entropy_trigger_rate_suggests_raising(self): + metrics = { + "admission": {"rejection_rate": 0.20}, + "entropy": {"trigger_rate": 0.02, "checks": 100, "triggers": 2}, + "fidelity": {"total_degradations": 5, "total_upgrades": 5}, + "micro_faults": {"total_attempts": 0, "success_rate": 0.0}, + } + result = suggest_thresholds(metrics) + assert result["entropy_sensitivity"]["status"] == "adjust" + assert "raising" in result["entropy_sensitivity"]["reason"] + + def test_high_fidelity_pressure_suggests_adjustment(self): + metrics = { + "admission": {"rejection_rate": 0.20}, + "fidelity": {"total_degradations": 60, "total_upgrades": 5}, + "micro_faults": {"total_attempts": 10, "success_rate": 0.8, "total_tokens_saved": 500}, + } + result = suggest_thresholds(metrics) + assert result["fidelity_pressure"]["status"] == "adjust" + assert "over-pressured" in result["fidelity_pressure"]["reason"] + + def test_too_conservative_fidelity(self): + metrics = { + "admission": {"rejection_rate": 0.20}, + "fidelity": {"total_degradations": 2, "total_upgrades": 20}, + "micro_faults": {"total_attempts": 10, "success_rate": 0.8, "total_tokens_saved": 500}, + } + result = suggest_thresholds(metrics) + assert result["fidelity_pressure"]["status"] == "adjust" + assert "conservative" in result["fidelity_pressure"]["reason"] + + def test_low_micro_fault_success_rate(self): + metrics = { + "admission": {"rejection_rate": 0.20}, + "fidelity": {"total_degradations": 5, "total_upgrades": 5}, + "micro_faults": {"total_attempts": 20, "success_rate": 0.30, "total_tokens_saved": 100}, + } + result = suggest_thresholds(metrics) + assert result["micro_fault_quality"]["status"] == "adjust" + assert "below 50%" in result["micro_fault_quality"]["reason"] + + def test_no_entropy_checks_is_ok(self): + metrics = { + "admission": {"rejection_rate": 0.20}, + "fidelity": {"total_degradations": 5, "total_upgrades": 5}, + "micro_faults": {"total_attempts": 0, "success_rate": 0.0}, + } + result = suggest_thresholds(metrics) + assert result["entropy_sensitivity"]["status"] == "ok" + assert "No entropy checks" in result["entropy_sensitivity"]["reason"] + + def test_no_micro_fault_attempts_is_ok(self): + metrics = { + "admission": {"rejection_rate": 0.20}, + "fidelity": {"total_degradations": 5, "total_upgrades": 5}, + "micro_faults": {"total_attempts": 0, "success_rate": 0.0}, + } + result = suggest_thresholds(metrics) + assert result["micro_fault_quality"]["status"] == "ok" + + def test_empty_metrics(self): + result = suggest_thresholds({}) + # Should not crash, all should be "ok" with defaults + assert result["admission_threshold"]["status"] == "ok" + assert result["entropy_sensitivity"]["status"] == "ok" + assert result["fidelity_pressure"]["status"] == "ok" + assert result["micro_fault_quality"]["status"] == "ok" + + +# ── Benchmark CLI ─────────────────────────────────────────────────────── + +from unittest.mock import patch, MagicMock +import json as json_mod + +import httpx + +from mnemosyne.benchmark_cli import ( + format_aggregate_report, + format_session_report, + run_benchmark_cli, + _box_top, + _box_bottom, + _box_sep, + _box_line, +) + + +class TestBenchmarkCLIFormatting: + def test_box_drawing_characters(self): + top = _box_top() + assert top.startswith("\u2554") + assert top.endswith("\u2557") + bottom = _box_bottom() + assert bottom.startswith("\u255a") + assert bottom.endswith("\u255d") + sep = _box_sep() + assert sep.startswith("\u2560") + assert sep.endswith("\u2563") + + def test_box_line_pads_correctly(self): + line = _box_line("hello") + assert line.startswith("\u2551") + assert line.endswith("\u2551") + assert "hello" in line + + def test_format_aggregate_report_structure(self): + data = { + "type": "aggregate", + "sessions": 3, + "total_turns": 147, + "tokens": { + "total_input_tokens": 50000, + "total_effective_tokens": 70000, + "total_cache_read_tokens": 892100, + "context_reduction_ratio": 0.268, + "context_reduction_pct": 73.2, + "avg_input_tokens_per_turn": 12450.0, + }, + "fidelity": {"total_degradations": 45, "total_upgrades": 12}, + "admission": { + "total_admitted": 312, + "total_rejected": 89, + "rejection_rate": 0.222, + }, + "micro_faults": { + "total_attempts": 23, + "total_successes": 18, + "success_rate": 0.783, + "total_tokens_saved": 45200, + }, + "latency": { + "segmentation": {"total_calls": 50, "avg_ms": 12.3, "min_ms": 5.1, "max_ms": 45.2}, + "admission": {"total_calls": 40, "avg_ms": 2.1, "min_ms": 0.8, "max_ms": 8.3}, + }, + } + report = format_aggregate_report(data) + assert "Mnemosyne Benchmark Report" in report + assert "Sessions: 3" in report + assert "Total Turns: 147" in report + assert "TOKEN SAVINGS" in report + assert "73.2%" in report + assert "FIDELITY" in report + assert "Degradations: 45" in report + assert "ADMISSION CONTROL" in report + assert "MICRO-FAULTS" in report + assert "LATENCY" in report + assert "segmentation" in report + assert "THRESHOLD SUGGESTIONS" in report + # Box drawing chars present + assert "\u2554" in report + assert "\u255d" in report + + def test_format_session_report_structure(self): + data = { + "type": "session", + "session_id": "test-abc", + "elapsed_seconds": 120.5, + "tokens": { + "turns": 10, + "total_input_tokens": 5000, + "context_reduction_pct": 60.0, + "avg_input_tokens_per_turn": 500.0, + "avg_cache_hit_rate": 0.75, + }, + "fidelity": {"degradations": 3, "upgrades": 1}, + "admission": {"admitted": 20, "rejected": 5, "rejection_rate": 0.2}, + "micro_faults": { + "attempts": 5, + "successes": 4, + "success_rate": 0.8, + "tokens_saved": 1000, + }, + "entropy": {"checks": 10, "triggers": 2, "trigger_rate": 0.2}, + "latency": {}, + } + report = format_session_report(data) + assert "Session: test-abc" in report + assert "Elapsed: 120.5s" in report + assert "TOKEN SAVINGS" in report + assert "ENTROPY" in report + + def test_format_aggregate_report_empty_sessions(self): + data = {"sessions": 0, "message": "No sessions recorded"} + report = format_aggregate_report(data) + assert "Sessions: 0" in report + + +class TestBenchmarkCLIExecution: + def test_run_benchmark_cli_json_output(self, capsys): + mock_data = {"type": "aggregate", "sessions": 2, "total_turns": 50} + mock_response = MagicMock() + mock_response.json.return_value = mock_data + mock_response.raise_for_status = MagicMock() + + with patch("mnemosyne.benchmark_cli.httpx.get", return_value=mock_response) as mock_get: + run_benchmark_cli(port=8080, json_output=True) + mock_get.assert_called_once_with("http://127.0.0.1:8080/api/benchmark", timeout=10.0) + + captured = capsys.readouterr() + parsed = json_mod.loads(captured.out) + assert parsed["sessions"] == 2 + + def test_run_benchmark_cli_with_session_id(self, capsys): + mock_data = { + "type": "session", + "session_id": "s1", + "elapsed_seconds": 10.0, + "tokens": {}, + "fidelity": {}, + "admission": {}, + "micro_faults": {}, + "entropy": {}, + "latency": {}, + } + mock_response = MagicMock() + mock_response.json.return_value = mock_data + mock_response.raise_for_status = MagicMock() + + with patch("mnemosyne.benchmark_cli.httpx.get", return_value=mock_response) as mock_get: + run_benchmark_cli(port=9090, session_id="s1") + mock_get.assert_called_once_with( + "http://127.0.0.1:9090/api/benchmark?session_id=s1", timeout=10.0 + ) + + captured = capsys.readouterr() + assert "Session: s1" in captured.out + + def test_run_benchmark_cli_pretty_output(self, capsys): + mock_data = { + "type": "aggregate", + "sessions": 1, + "total_turns": 10, + "tokens": { + "context_reduction_pct": 50.0, + "avg_input_tokens_per_turn": 1000, + "total_cache_read_tokens": 5000, + }, + "fidelity": {"total_degradations": 2, "total_upgrades": 1}, + "admission": {"total_admitted": 10, "total_rejected": 3, "rejection_rate": 0.23}, + "micro_faults": { + "total_attempts": 5, + "total_successes": 4, + "success_rate": 0.8, + "total_tokens_saved": 2000, + }, + "latency": {}, + } + mock_response = MagicMock() + mock_response.json.return_value = mock_data + mock_response.raise_for_status = MagicMock() + + with patch("mnemosyne.benchmark_cli.httpx.get", return_value=mock_response): + run_benchmark_cli(port=8080) + + captured = capsys.readouterr() + assert "Mnemosyne Benchmark Report" in captured.out + assert "\u2554" in captured.out # box top + + def test_run_benchmark_cli_connection_error(self): + with patch("mnemosyne.benchmark_cli.httpx.get", side_effect=httpx.ConnectError("refused")): + import pytest + + with pytest.raises(SystemExit) as exc_info: + run_benchmark_cli(port=8080) + assert exc_info.value.code == 1 + + +class TestBenchmarkCLIArgParsing: + def test_benchmark_flag_parsed(self): + """Verify --benchmark is recognized by the argparse parser.""" + import argparse + + # Reconstruct the parser as main() does + parser = argparse.ArgumentParser() + parser.add_argument("--benchmark", action="store_true") + parser.add_argument("--session", type=str, default=None) + parser.add_argument("--json-output", action="store_true") + parser.add_argument("--port", type=int, default=0) + + args = parser.parse_args(["--benchmark"]) + assert args.benchmark is True + assert args.session is None + assert args.json_output is False + + def test_benchmark_with_session_and_json(self): + import argparse + + parser = argparse.ArgumentParser() + parser.add_argument("--benchmark", action="store_true") + parser.add_argument("--session", type=str, default=None) + parser.add_argument("--json-output", action="store_true") + parser.add_argument("--port", type=int, default=0) + + args = parser.parse_args( + ["--benchmark", "--session", "abc-123", "--json-output", "--port", "9090"] + ) + assert args.benchmark is True + assert args.session == "abc-123" + assert args.json_output is True + assert args.port == 9090 + + def test_benchmark_port_default(self): + """When --benchmark is used without --port, port defaults to 0 (gateway treats as 8080).""" + import argparse + + parser = argparse.ArgumentParser() + parser.add_argument("--benchmark", action="store_true") + parser.add_argument("--port", type=int, default=0) + + args = parser.parse_args(["--benchmark"]) + # The gateway main() maps port=0 to 8080 for benchmark mode + benchmark_port = args.port if args.port != 0 else 8080 + assert benchmark_port == 8080 diff --git a/tests/test_summarization_pipeline.py b/tests/test_summarization_pipeline.py new file mode 100644 index 0000000..c027c9f --- /dev/null +++ b/tests/test_summarization_pipeline.py @@ -0,0 +1,558 @@ +"""Tests for the HelperLLM summarization → fidelity degradation pipeline. + +Verifies that when FidelityManager degrades objects, the gateway calls +HelperLLM to generate compressed summaries instead of using stub text. + +All tests use mocked HelperLLM — no real API calls. +""" + +from __future__ import annotations + +import asyncio +import json +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +from mnemosyne.fidelity import ( + FidelityLevel, + FidelityManager, + SemanticObject, + make_object, +) +from mnemosyne.gateway import ( + Session, + _apply_fidelity, + _auto_stub, + _generate_degradation_summaries, +) +from mnemosyne.helper_llm import HelperLLM, SummaryResult + + +# --------------------------------------------------------------------------- +# Fixtures +# --------------------------------------------------------------------------- + + +def _make_obj( + content: str = "x" * 800, + *, + object_type: str = "file_context", + turn: int = 0, + summary_detailed: str | None = None, + summary_compact: str | None = None, + stub: str | None = None, +) -> SemanticObject: + """Create a SemanticObject with sensible defaults for testing.""" + return make_object( + object_type=object_type, + content_full=content, + created_at_turn=turn, + summary_detailed=summary_detailed, + summary_compact=summary_compact, + stub=stub or f"[evicted content: {content[:40]}]", + ) + + +def _make_mock_helper() -> MagicMock: + """Create a mock HelperLLM with async methods.""" + helper = MagicMock(spec=HelperLLM) + helper.summarize_l0_to_l1 = AsyncMock( + return_value=SummaryResult( + summary="LLM-generated L1 summary of the content.", + losses=["exact error codes"], + can_answer=["what was implemented"], + key_entities=["src/main.py"], + ) + ) + helper.compress_l1_to_l2 = AsyncMock( + return_value=SummaryResult( + summary="LLM-generated compact L2 summary.", + losses=["exact error codes", "function signatures"], + can_answer=["what was decided"], + key_entities=["main.py"], + ) + ) + helper.generate_stub = AsyncMock( + return_value="[file_context | turn 0 | test stub | 0 related objects]" + ) + return helper + + +def _make_session() -> Session: + """Create a minimal Session for testing (mocked dependencies).""" + import tempfile + from pathlib import Path + + with tempfile.TemporaryDirectory() as tmpdir: + log_dir = Path(tmpdir) + # Session.__init__ imports many things; patch what we need + session = Session.__new__(Session) + session.id = "test-sess" + session.token_state = {"last_effective": 0, "blocked": False, "turn": 5} + session.fidelity_manager = FidelityManager(window_size=200_000) + session._fidelity_content_map = {} + session._summary_cache = {} + return session + + +# --------------------------------------------------------------------------- +# _generate_degradation_summaries tests +# --------------------------------------------------------------------------- + + +class TestGenerateDegradationSummaries: + """Test the async summarization function directly.""" + + async def test_l0_to_l1_calls_summarize(self) -> None: + """L0→L1 transition calls helper_llm.summarize_l0_to_l1().""" + session = _make_session() + helper = _make_mock_helper() + fm = session.fidelity_manager + + obj = _make_obj(content="Full content for summarization " * 20) + obj_id = fm.register_object(obj) + + transitions = [(obj_id, FidelityLevel.L0, FidelityLevel.L1)] + await _generate_degradation_summaries(transitions, session, helper) + + # Verify LLM was called + helper.summarize_l0_to_l1.assert_called_once() + call_kwargs = helper.summarize_l0_to_l1.call_args.kwargs + assert call_kwargs["object_type"] == "file_context" + assert call_kwargs["max_summary_tokens"] == 1024 + + # Verify summary was stored on the object + assert obj.summary_detailed == "LLM-generated L1 summary of the content." + assert obj.losses_l1 == ["exact error codes"] + assert obj.can_answer == ["what was implemented"] + assert obj.key_entities == ["src/main.py"] + + # Verify cache was populated + assert (obj_id, int(FidelityLevel.L1)) in session._summary_cache + assert session._summary_cache[(obj_id, int(FidelityLevel.L1))] == ( + "LLM-generated L1 summary of the content." + ) + + async def test_l1_to_l2_calls_compress(self) -> None: + """L1→L2 transition calls helper_llm.compress_l1_to_l2().""" + session = _make_session() + helper = _make_mock_helper() + fm = session.fidelity_manager + + obj = _make_obj( + summary_detailed="Existing L1 detailed summary of auth middleware.", + ) + obj.losses_l1 = ["line numbers"] + obj_id = fm.register_object(obj) + + transitions = [(obj_id, FidelityLevel.L1, FidelityLevel.L2)] + await _generate_degradation_summaries(transitions, session, helper) + + # Verify LLM was called with L1 text + helper.compress_l1_to_l2.assert_called_once() + call_kwargs = helper.compress_l1_to_l2.call_args.kwargs + assert call_kwargs["l1_summary"] == "Existing L1 detailed summary of auth middleware." + assert call_kwargs["l1_losses"] == ["line numbers"] + assert call_kwargs["object_type"] == "file_context" + + # Verify compact summary was stored + assert obj.summary_compact == "LLM-generated compact L2 summary." + assert (obj_id, int(FidelityLevel.L2)) in session._summary_cache + + async def test_l1_to_l2_uses_cache_for_l1_text(self) -> None: + """L1→L2 uses cached L1 summary when obj.summary_detailed is None.""" + session = _make_session() + helper = _make_mock_helper() + fm = session.fidelity_manager + + obj = _make_obj() + obj.summary_detailed = None # No pre-set L1 summary + obj_id = fm.register_object(obj) + + # Pre-populate cache with L1 summary + session._summary_cache[(obj_id, int(FidelityLevel.L1))] = "Cached L1 summary text." + + transitions = [(obj_id, FidelityLevel.L1, FidelityLevel.L2)] + await _generate_degradation_summaries(transitions, session, helper) + + # Should use cached L1 text + call_kwargs = helper.compress_l1_to_l2.call_args.kwargs + assert call_kwargs["l1_summary"] == "Cached L1 summary text." + + async def test_l2_to_l3_no_llm_call(self) -> None: + """L2→L3 generates a simple stub without calling the LLM.""" + session = _make_session() + helper = _make_mock_helper() + fm = session.fidelity_manager + + obj = _make_obj(object_type="tool_result") + obj.created_at_turn = 7 + obj_id = fm.register_object(obj) + + transitions = [(obj_id, FidelityLevel.L2, FidelityLevel.L3)] + await _generate_degradation_summaries(transitions, session, helper) + + # No LLM calls should have been made + helper.summarize_l0_to_l1.assert_not_called() + helper.compress_l1_to_l2.assert_not_called() + helper.generate_stub.assert_not_called() + + # Stub should be set + assert obj.stub == "[evicted: tool_result from turn 7]" + assert (obj_id, int(FidelityLevel.L3)) in session._summary_cache + + async def test_l3_to_l4_no_action(self) -> None: + """L3→L4 (full eviction) doesn't generate any content.""" + session = _make_session() + helper = _make_mock_helper() + fm = session.fidelity_manager + + obj = _make_obj() + obj_id = fm.register_object(obj) + + transitions = [(obj_id, FidelityLevel.L3, FidelityLevel.L4)] + await _generate_degradation_summaries(transitions, session, helper) + + # No LLM calls + helper.summarize_l0_to_l1.assert_not_called() + helper.compress_l1_to_l2.assert_not_called() + + # No cache entry for L4 + assert (obj_id, int(FidelityLevel.L4)) not in session._summary_cache + + async def test_fallback_on_llm_exception(self) -> None: + """If HelperLLM raises an exception, fall back silently.""" + session = _make_session() + helper = _make_mock_helper() + helper.summarize_l0_to_l1 = AsyncMock(side_effect=RuntimeError("API down")) + fm = session.fidelity_manager + + obj = _make_obj() + original_stub = obj.stub + obj_id = fm.register_object(obj) + + transitions = [(obj_id, FidelityLevel.L0, FidelityLevel.L1)] + # Should NOT raise + await _generate_degradation_summaries(transitions, session, helper) + + # Cache should NOT have an entry (fallback means no summary generated) + assert (obj_id, int(FidelityLevel.L1)) not in session._summary_cache + # Original stub should be preserved + assert obj.stub == original_stub + + async def test_fallback_on_empty_summary(self) -> None: + """If HelperLLM returns empty summary, don't cache it.""" + session = _make_session() + helper = _make_mock_helper() + helper.summarize_l0_to_l1 = AsyncMock( + return_value=SummaryResult(summary="", losses=[], can_answer=[], key_entities=[]) + ) + fm = session.fidelity_manager + + obj = _make_obj() + obj_id = fm.register_object(obj) + + transitions = [(obj_id, FidelityLevel.L0, FidelityLevel.L1)] + await _generate_degradation_summaries(transitions, session, helper) + + # Empty summary should NOT be cached + assert (obj_id, int(FidelityLevel.L1)) not in session._summary_cache + # summary_detailed should NOT be set to empty + assert obj.summary_detailed is None or obj.summary_detailed != "" + + async def test_summary_cache_hit_skips_llm(self) -> None: + """If summary is already cached, don't call the LLM again.""" + session = _make_session() + helper = _make_mock_helper() + fm = session.fidelity_manager + + obj = _make_obj() + obj_id = fm.register_object(obj) + + # Pre-populate cache + session._summary_cache[(obj_id, int(FidelityLevel.L1))] = "Already cached summary." + + transitions = [(obj_id, FidelityLevel.L0, FidelityLevel.L1)] + await _generate_degradation_summaries(transitions, session, helper) + + # LLM should NOT have been called + helper.summarize_l0_to_l1.assert_not_called() + + async def test_none_helper_llm_is_noop(self) -> None: + """If helper_llm is None, the function is a no-op.""" + session = _make_session() + fm = session.fidelity_manager + + obj = _make_obj() + obj_id = fm.register_object(obj) + + transitions = [(obj_id, FidelityLevel.L0, FidelityLevel.L1)] + # Should not raise + await _generate_degradation_summaries(transitions, session, None) + + # No cache entries + assert len(session._summary_cache) == 0 + + async def test_multiple_transitions_processed(self) -> None: + """Multiple transitions in one batch are all processed.""" + session = _make_session() + helper = _make_mock_helper() + fm = session.fidelity_manager + + obj1 = _make_obj(content="Content A " * 50) + obj2 = _make_obj(content="Content B " * 50, object_type="tool_result") + obj2.summary_detailed = "L1 summary for obj2" + obj2.created_at_turn = 3 + + id1 = fm.register_object(obj1) + id2 = fm.register_object(obj2) + + transitions = [ + (id1, FidelityLevel.L0, FidelityLevel.L1), + (id2, FidelityLevel.L2, FidelityLevel.L3), + ] + await _generate_degradation_summaries(transitions, session, helper) + + # obj1: L0→L1 should have called summarize + helper.summarize_l0_to_l1.assert_called_once() + assert (id1, int(FidelityLevel.L1)) in session._summary_cache + + # obj2: L2→L3 should have set stub without LLM + assert obj2.stub == "[evicted: tool_result from turn 3]" + assert (id2, int(FidelityLevel.L3)) in session._summary_cache + + +# --------------------------------------------------------------------------- +# _apply_fidelity integration tests +# --------------------------------------------------------------------------- + + +class TestApplyFidelityWithSummaries: + """Test that _apply_fidelity uses LLM-generated summaries.""" + + def test_l1_object_uses_summary_from_cache(self) -> None: + """An L1 object should have its content replaced with the cached summary.""" + session = _make_session() + fm = session.fidelity_manager + + # Create and register an object + content = "x" * 800 + obj = _make_obj(content=content) + obj_id = fm.register_object(obj) + + # Set up content map (simulate previous registration) + content_key = ( + f"text:{__import__('hashlib').sha256(content[:200].encode()).hexdigest()[:12]}" + ) + session._fidelity_content_map[content_key] = obj_id + + # Degrade to L1 and cache a summary + obj.current_fidelity = FidelityLevel.L1 + session._summary_cache[(obj_id, int(FidelityLevel.L1))] = "Cached L1 summary." + + # Build payload with the original content + payload = { + "messages": [ + { + "role": "assistant", + "content": [{"type": "text", "text": content}], + } + ] + } + + _apply_fidelity(payload, session) + + # Content should be replaced with the cached summary + assert payload["messages"][0]["content"][0]["text"] == "Cached L1 summary." + + def test_l2_object_uses_compact_summary(self) -> None: + """An L2 object should have its content replaced with the compact summary.""" + session = _make_session() + fm = session.fidelity_manager + + content = "y" * 800 + obj = _make_obj(content=content) + obj.summary_compact = "Pre-set compact summary." + obj_id = fm.register_object(obj) + + content_key = ( + f"text:{__import__('hashlib').sha256(content[:200].encode()).hexdigest()[:12]}" + ) + session._fidelity_content_map[content_key] = obj_id + + obj.current_fidelity = FidelityLevel.L2 + + payload = { + "messages": [ + { + "role": "user", + "content": [{"type": "text", "text": content}], + } + ] + } + + _apply_fidelity(payload, session) + + assert payload["messages"][0]["content"][0]["text"] == "Pre-set compact summary." + + def test_l1_object_falls_back_to_full_content(self) -> None: + """If no L1 summary is available, keep full content.""" + session = _make_session() + fm = session.fidelity_manager + + content = "z" * 800 + obj = _make_obj(content=content) + obj.summary_detailed = None # No summary available + obj_id = fm.register_object(obj) + + content_key = ( + f"text:{__import__('hashlib').sha256(content[:200].encode()).hexdigest()[:12]}" + ) + session._fidelity_content_map[content_key] = obj_id + + obj.current_fidelity = FidelityLevel.L1 + + payload = { + "messages": [ + { + "role": "assistant", + "content": [{"type": "text", "text": content}], + } + ] + } + + _apply_fidelity(payload, session) + + # Should keep original content as fallback + assert payload["messages"][0]["content"][0]["text"] == content + + def test_l3_object_uses_stub(self) -> None: + """L3 objects should use stub text (unchanged behavior).""" + session = _make_session() + fm = session.fidelity_manager + + content = "w" * 800 + obj = _make_obj(content=content, stub="[evicted content: test stub]") + obj_id = fm.register_object(obj) + + content_key = ( + f"text:{__import__('hashlib').sha256(content[:200].encode()).hexdigest()[:12]}" + ) + session._fidelity_content_map[content_key] = obj_id + + obj.current_fidelity = FidelityLevel.L3 + + payload = { + "messages": [ + { + "role": "assistant", + "content": [{"type": "text", "text": content}], + } + ] + } + + _apply_fidelity(payload, session) + + assert payload["messages"][0]["content"][0]["text"] == "[evicted content: test stub]" + + def test_l2_cache_takes_priority_over_object_field(self) -> None: + """Cache entry should take priority over obj.summary_compact.""" + session = _make_session() + fm = session.fidelity_manager + + content = "a" * 800 + obj = _make_obj(content=content) + obj.summary_compact = "Old compact summary." + obj_id = fm.register_object(obj) + + content_key = ( + f"text:{__import__('hashlib').sha256(content[:200].encode()).hexdigest()[:12]}" + ) + session._fidelity_content_map[content_key] = obj_id + + obj.current_fidelity = FidelityLevel.L2 + session._summary_cache[(obj_id, int(FidelityLevel.L2))] = "Fresh cached L2 summary." + + payload = { + "messages": [ + { + "role": "assistant", + "content": [{"type": "text", "text": content}], + } + ] + } + + _apply_fidelity(payload, session) + + assert payload["messages"][0]["content"][0]["text"] == "Fresh cached L2 summary." + + def test_tool_result_block_uses_summary(self) -> None: + """Tool result blocks should also get their content replaced.""" + session = _make_session() + fm = session.fidelity_manager + + content = "tool output " * 100 + obj = _make_obj(content=content, object_type="tool_result") + obj_id = fm.register_object(obj) + + content_key = f"tool:use_123" + session._fidelity_content_map[content_key] = obj_id + + obj.current_fidelity = FidelityLevel.L1 + session._summary_cache[(obj_id, int(FidelityLevel.L1))] = "Tool result L1 summary." + + payload = { + "messages": [ + { + "role": "user", + "content": [ + { + "type": "tool_result", + "tool_use_id": "use_123", + "content": content, + } + ], + } + ] + } + + _apply_fidelity(payload, session) + + assert payload["messages"][0]["content"][0]["content"] == "Tool result L1 summary." + + +# --------------------------------------------------------------------------- +# Async non-blocking behavior test +# --------------------------------------------------------------------------- + + +class TestAsyncNonBlocking: + """Verify summarization runs asynchronously without blocking.""" + + async def test_summarization_completes_independently(self) -> None: + """The summarization coroutine completes without blocking the caller.""" + session = _make_session() + helper = _make_mock_helper() + + # Add a small delay to simulate LLM latency + original_summarize = helper.summarize_l0_to_l1 + + async def slow_summarize(**kwargs): + await asyncio.sleep(0.01) # 10ms simulated latency + return await original_summarize(**kwargs) + + helper.summarize_l0_to_l1 = slow_summarize + + fm = session.fidelity_manager + obj = _make_obj() + obj_id = fm.register_object(obj) + + transitions = [(obj_id, FidelityLevel.L0, FidelityLevel.L1)] + + # Run as a task — should complete without issues + task = asyncio.create_task(_generate_degradation_summaries(transitions, session, helper)) + await task + + # Summary should be cached after completion + assert (obj_id, int(FidelityLevel.L1)) in session._summary_cache