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 <clio-agent@sisyphuslabs.ai>
This commit is contained in:
Joey Yakimowich-Payne 2026-03-13 11:41:22 -06:00
commit d660414ad7
12 changed files with 6276 additions and 0 deletions

665
src/mnemosyne/analyzer.py Normal file
View file

@ -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 <proxy-log.jsonl> [--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"^<fast_mode_info>"),
]
# 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 "<system-reminder>" 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 "<system-reminder>" 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"<system-reminder>(.*?)</system-reminder>", 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()

724
src/mnemosyne/benchmark.py Normal file
View file

@ -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)

View file

@ -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 <id> # 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))

576
src/mnemosyne/cost.py Normal file
View file

@ -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()

View file

@ -0,0 +1,2 @@
# Deprecated modules — kept for transitional import compatibility.
# New development targets gateway.py.

File diff suppressed because it is too large Load diff

496
src/mnemosyne/eval.py Normal file
View file

@ -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()

301
src/mnemosyne/oauth.py Normal file
View file

@ -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")

326
src/mnemosyne/replay.py Normal file
View file

@ -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()

292
src/mnemosyne/telemetry.py Normal file
View file

@ -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()
}

812
tests/test_benchmark.py Normal file
View file

@ -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

View file

@ -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