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:
parent
681c1454cb
commit
d660414ad7
12 changed files with 6276 additions and 0 deletions
665
src/mnemosyne/analyzer.py
Normal file
665
src/mnemosyne/analyzer.py
Normal 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
724
src/mnemosyne/benchmark.py
Normal 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)
|
||||
316
src/mnemosyne/benchmark_cli.py
Normal file
316
src/mnemosyne/benchmark_cli.py
Normal 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
576
src/mnemosyne/cost.py
Normal 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()
|
||||
2
src/mnemosyne/deprecated/__init__.py
Normal file
2
src/mnemosyne/deprecated/__init__.py
Normal file
|
|
@ -0,0 +1,2 @@
|
|||
# Deprecated modules — kept for transitional import compatibility.
|
||||
# New development targets gateway.py.
|
||||
1208
src/mnemosyne/deprecated/proxy.py
Normal file
1208
src/mnemosyne/deprecated/proxy.py
Normal file
File diff suppressed because it is too large
Load diff
496
src/mnemosyne/eval.py
Normal file
496
src/mnemosyne/eval.py
Normal 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
301
src/mnemosyne/oauth.py
Normal 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
326
src/mnemosyne/replay.py
Normal 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
292
src/mnemosyne/telemetry.py
Normal 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
812
tests/test_benchmark.py
Normal 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
|
||||
558
tests/test_summarization_pipeline.py
Normal file
558
tests/test_summarization_pipeline.py
Normal 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
|
||||
Loading…
Add table
Add a link
Reference in a new issue