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