feat: add core proxy framework with gateway and providers

Multi-provider HTTP proxy (Anthropic + OpenAI) with session management,
message processing pipeline, block labeling, cache control placement,
and embedded monitoring dashboard.

Ultraworked with [Sisyphus](https://github.com/code-yeongyu/oh-my-opencode)

Co-authored-by: Sisyphus <clio-agent@sisyphuslabs.ai>
This commit is contained in:
Joey Yakimowich-Payne 2026-03-13 11:40:49 -06:00
commit 974863e7b3
21 changed files with 6591 additions and 0 deletions

View file

@ -0,0 +1,6 @@
"""Mnemosyne — object-addressed context memory for LLM agents.
Built on Pichay's demand paging foundation, extended with multi-fidelity
compression, semantic object segmentation, and a helper LLM for goal-aware
retrieval. Named after the Greek Titan of memory.
"""

307
src/mnemosyne/__main__.py Normal file
View file

@ -0,0 +1,307 @@
#!/usr/bin/env python3
"""Pichay experiment runner.
Single entry point: starts the proxy, launches Claude Code through it,
captures all artifacts when done.
Usage:
# Baseline run (observe only)
python -m pichay --treatment baseline --project /path/to/project --prompt "Build X"
# Compact run (dead tool eviction)
python -m pichay --treatment compact --compact --project /path/to/project --prompt "Build X"
# With temperature control
python -m pichay --treatment compact --compact --temperature 0 --project /path/to/project --prompt "Build X"
Artifacts are saved to experiments/{treatment}_run{N}/ in the pichay project directory.
"""
from __future__ import annotations
import argparse
import json
import os
import shutil
import subprocess
import sys
import threading
import time
from datetime import datetime, timezone
from pathlib import Path
# Transitional: proxy.py moved to deprecated/. Gateway is the primary path
# but the experiment runner still needs Flask until fully ported.
from mnemosyne.deprecated.proxy import create_app, find_free_port
def find_project_claude_dir(project_dir: str) -> Path:
"""Map a project directory to its ~/.claude/projects/ path."""
normalized = project_dir.replace("/", "-")
return Path.home() / ".claude" / "projects" / normalized
# Environment variables that cause Claude Code to detect nested invocation.
_CLAUDE_NESTED_VARS = [
"CLAUDECODE",
"CLAUDE_CODE_ENTRYPOINT",
"CLAUDE_CODE_SSE_PORT",
"CLAUDE_CODE_EXPERIMENTAL_AGENT_TEAMS",
]
def clean_env_for_subprocess() -> dict[str, str]:
"""Build environment for nested Claude Code, removing detection vars."""
env = os.environ.copy()
for var in _CLAUDE_NESTED_VARS:
env.pop(var, None)
return env
def run_experiment(args: argparse.Namespace) -> None:
"""Run a single experiment."""
pichay_dir = Path(__file__).resolve().parent.parent.parent
exp_dir = pichay_dir / "experiments" / f"{args.treatment}_run{args.run}"
exp_dir.mkdir(parents=True, exist_ok=True)
log_dir = exp_dir / "logs"
log_dir.mkdir(exist_ok=True)
project_dir = Path(args.project).resolve()
if not project_dir.is_dir():
print(f"Error: project directory not found: {project_dir}", file=sys.stderr)
sys.exit(1)
print("=" * 60, file=sys.stderr)
print(f"Pichay Experiment", file=sys.stderr)
print(f" Treatment: {args.treatment}", file=sys.stderr)
print(f" Run: {args.run}", file=sys.stderr)
print(f" Project: {project_dir}", file=sys.stderr)
print(f" Compact: {args.compact}", file=sys.stderr)
print(f" Trim: {args.trim}", file=sys.stderr)
print(f" Temperature: {args.temperature}", file=sys.stderr)
print(f" Output: {exp_dir}", file=sys.stderr)
print("=" * 60, file=sys.stderr)
# Save config
config = {
"treatment": args.treatment,
"run": args.run,
"project_dir": str(project_dir),
"branch": args.branch,
"compact": args.compact,
"trim": args.trim,
"age_threshold": args.age_threshold,
"min_size": args.min_size,
"temperature": args.temperature,
"prompt": args.prompt,
"started_at": datetime.now(timezone.utc).isoformat(),
}
(exp_dir / "config.json").write_text(json.dumps(config, indent=2))
# Step 1: Reset project
if args.branch:
print(f"\n[1/5] Resetting project to {args.branch}...", file=sys.stderr)
subprocess.run(
["git", "checkout", args.branch],
cwd=project_dir, capture_output=True,
)
subprocess.run(
["git", "clean", "-fd"],
cwd=project_dir, capture_output=True,
)
else:
print(f"\n[1/5] No branch specified, using current state.", file=sys.stderr)
# Clear previous session data
claude_dir = find_project_claude_dir(str(project_dir))
if args.clear_session and claude_dir.is_dir():
for jsonl in claude_dir.glob("*.jsonl"):
jsonl.unlink()
print(f" Cleared session data in {claude_dir}", file=sys.stderr)
# Step 2: Start proxy
print(f"\n[2/5] Starting proxy...", file=sys.stderr)
port = find_free_port()
create_kwargs = dict(
log_dir=log_dir,
compact=args.compact,
trim=args.trim,
age_threshold=args.age_threshold,
min_size=args.min_size,
)
if args.upstream:
create_kwargs["upstream"] = args.upstream
app = create_app(**create_kwargs)
if args.temperature is not None:
app.config["temperature_override"] = args.temperature
# Run Flask in a thread
server_thread = threading.Thread(
target=lambda: app.run(
host="127.0.0.1", port=port, threaded=True, use_reloader=False,
),
daemon=True,
)
server_thread.start()
time.sleep(1) # Let server bind
print(f" Proxy on http://localhost:{port}", file=sys.stderr)
# Step 3: Run Claude Code
print(f"\n[3/5] Launching Claude Code...", file=sys.stderr)
print(f" Prompt: {args.prompt[:80]}...", file=sys.stderr)
print(file=sys.stderr)
env = clean_env_for_subprocess()
env["ANTHROPIC_BASE_URL"] = f"http://localhost:{port}"
claude_cmd = [
"claude",
"-p",
"--dangerously-skip-permissions",
"--max-budget-usd", str(args.max_budget),
]
if args.system_prompt is not None:
claude_cmd.extend(["--system-prompt", args.system_prompt])
claude_cmd.append(args.prompt)
try:
subprocess.run(
claude_cmd,
cwd=project_dir,
env=env,
stdin=subprocess.DEVNULL,
)
except KeyboardInterrupt:
print("\n Interrupted by user.", file=sys.stderr)
except FileNotFoundError:
print(" Error: 'claude' not found in PATH.", file=sys.stderr)
sys.exit(1)
# Step 4: Capture artifacts
print(f"\n[4/5] Capturing artifacts...", file=sys.stderr)
# Copy session data
if claude_dir.is_dir():
session_dest = exp_dir / "session"
if session_dest.exists():
shutil.rmtree(session_dest)
shutil.copytree(claude_dir, session_dest, dirs_exist_ok=True)
# Git state
subprocess.run(
["git", "log", "--oneline", "-20"],
cwd=project_dir,
stdout=open(exp_dir / "git_log.txt", "w"),
stderr=subprocess.DEVNULL,
)
if args.branch:
subprocess.run(
["git", "diff", f"{args.branch}..HEAD"],
cwd=project_dir,
stdout=open(exp_dir / "git_diff.txt", "w"),
stderr=subprocess.DEVNULL,
)
subprocess.run(
["git", "diff", "--stat", f"{args.branch}..HEAD"],
cwd=project_dir,
stdout=open(exp_dir / "git_diff_stat.txt", "w"),
stderr=subprocess.DEVNULL,
)
# Test results
pyproject = project_dir / "pyproject.toml"
if pyproject.exists():
subprocess.run(
["uv", "run", "pytest", "tests/", "-v", "--tb=short"],
cwd=project_dir,
stdout=open(exp_dir / "test_results.txt", "w"),
stderr=subprocess.STDOUT,
)
# Record end time
config["ended_at"] = datetime.now(timezone.utc).isoformat()
(exp_dir / "config.json").write_text(json.dumps(config, indent=2))
# Step 5: Summary
print(f"\n[5/5] Run complete.", file=sys.stderr)
print(f" Artifacts: {exp_dir}", file=sys.stderr)
# Run eval if proxy log exists
proxy_logs = list(log_dir.glob("proxy_*.jsonl"))
if proxy_logs:
from mnemosyne.eval import analyze_run, print_run_summary
summary = analyze_run(proxy_logs[0], label=args.treatment)
print_run_summary(summary)
def main():
parser = argparse.ArgumentParser(
description="Pichay — context paging experiment runner"
)
parser.add_argument(
"--treatment", required=True,
help="Treatment label (e.g., baseline, compact, trimmed)",
)
parser.add_argument(
"--run", type=int, default=1,
help="Run number (default: 1)",
)
parser.add_argument(
"--project", required=True,
help="Target project directory",
)
parser.add_argument(
"--prompt", required=True,
help="Starting prompt for Claude Code",
)
parser.add_argument(
"--branch", default=None,
help="Git branch to reset to before each run",
)
parser.add_argument(
"--compact", action="store_true",
help="Enable dead tool result eviction",
)
parser.add_argument(
"--trim", action="store_true",
help="Enable system prompt trimming (tool stubs, skill dedup)",
)
parser.add_argument(
"--age-threshold", type=int, default=4,
help="Eviction age threshold in user-turns (default: 4)",
)
parser.add_argument(
"--min-size", type=int, default=500,
help="Min tool result size for eviction (default: 500)",
)
parser.add_argument(
"--temperature", type=float, default=None,
help="Override temperature (e.g., 0 for deterministic)",
)
parser.add_argument(
"--max-budget", type=float, default=20.0,
help="Maximum dollar budget for Claude Code session (default: $20)",
)
parser.add_argument(
"--clear-session", action="store_true", default=True,
help="Clear previous Claude session data (default: yes)",
)
parser.add_argument(
"--no-clear-session", action="store_false", dest="clear_session",
help="Keep previous Claude session data",
)
parser.add_argument(
"--system-prompt", default=None,
help="Override Claude Code's system prompt (--system-prompt flag)",
)
parser.add_argument(
"--upstream", default=None,
help="Upstream API base URL (default: https://api.anthropic.com). "
"Use for OpenRouter, Kimi, or any Anthropic-compatible endpoint.",
)
args = parser.parse_args()
run_experiment(args)
if __name__ == "__main__":
main()

377
src/mnemosyne/blocks.py Normal file
View file

@ -0,0 +1,377 @@
"""Conversation block store: content-addressed KV for message blocks.
Assigns short stable IDs to conversation message blocks based on
content hashing. The model sees these labels and can reference them
in cleanup operations.
Phase 1: labeling — inject [block:xxxx] markers into message content.
Phase 2: cleanup — drop, summarize, anchor blocks via inline tags.
"""
from __future__ import annotations
import hashlib
import json
import re
from dataclasses import dataclass, field
from pathlib import Path
@dataclass
class BlockEntry:
"""A tracked conversation block."""
block_id: str # Short hex ID (first 8 chars of content hash)
content_hash: str # Full SHA-256 of content
size: int # Byte size of original content
turn: int # Turn when first seen
role: str # "user" or "assistant"
preview: str # First 80 chars for logging
status: str = "resident" # resident | anchored | summarized | dropped
original_content: str | None = None # Full content for fault restoration
summary: str | None = None # Model-authored summary (if summarized)
class BlockStore:
"""Per-session content-addressed store for conversation blocks.
Maps content hashes to BlockEntries. Assigns short IDs for model
reference. Detects content changes (re-labels gracefully).
"""
def __init__(self):
self._by_id: dict[str, BlockEntry] = {}
self._by_hash: dict[str, str] = {} # content_hash → block_id
_LABEL_RE = re.compile(r"^\[(?:tensor|block):([0-9a-f]+)[ (]")
def _has_our_label(self, text: str) -> bool:
"""Check if text starts with a label WE placed.
Validates the ID against known block IDs. Foreign labels
(injected via file contents, tool results, or crafted messages)
will have unrecognized IDs and return False, preventing an
attacker from skipping labeling or spoofing block references.
"""
m = self._LABEL_RE.match(text)
if m is None:
return False
return m.group(1) in self._by_id
def label_messages(self, messages: list[dict], current_turn: int) -> None:
"""Inject [block:xxxx] labels into message content.
Modifies messages in-place. Each message gets a label based on
its content hash. Labels are stable across turns as long as
the content doesn't change.
Only labels user and assistant text messages. Tool_use and
tool_result blocks are managed by the PageStore, not here.
"""
for msg in messages:
role = msg.get("role", "")
if role not in ("user", "assistant"):
continue
content = msg.get("content", "")
# String content (simple user messages)
if isinstance(content, str):
# Skip if already labeled by us (validated against known IDs)
if self._has_our_label(content):
continue
# Skip very short messages (not worth labeling)
if len(content) < 200:
continue
entry = self._get_or_create(content, role, current_turn)
if entry and entry.status == "resident":
size_k = entry.size / 1024
msg["content"] = f"[tensor:{entry.block_id} ({size_k:.1f}KB)]\n{content}"
# List content (structured blocks)
elif isinstance(content, list):
for block in content:
if not isinstance(block, dict):
continue
if block.get("type") != "text":
continue
text = block.get("text", "")
# Skip if already labeled by us (validated against known IDs)
if self._has_our_label(text):
continue
# Skip short blocks
if len(text) < 200:
continue
entry = self._get_or_create(text, role, current_turn)
if entry and entry.status == "resident":
size_k = entry.size / 1024
block["text"] = f"[tensor:{entry.block_id} ({size_k:.1f}KB)]\n{text}"
def _get_or_create(self, content: str, role: str,
turn: int) -> BlockEntry | None:
"""Get existing entry by content hash, or create a new one."""
content_hash = hashlib.sha256(content.encode()).hexdigest()
short_id = content_hash[:8]
# Already tracked
if content_hash in self._by_hash:
return self._by_id[self._by_hash[content_hash]]
# Check for ID collision (different content, same short hash)
if short_id in self._by_id:
existing = self._by_id[short_id]
if existing.content_hash != content_hash:
# Collision — extend the ID
short_id = content_hash[:12]
if short_id in self._by_id:
return None # Extremely unlikely, skip labeling
preview = content[:80].replace("\n", " ")
entry = BlockEntry(
block_id=short_id,
content_hash=content_hash,
size=len(content.encode()),
turn=turn,
role=role,
preview=preview,
original_content=content,
)
self._by_id[short_id] = entry
self._by_hash[content_hash] = short_id
return entry
def get(self, block_id: str) -> BlockEntry | None:
"""Look up a block by its short ID."""
return self._by_id.get(block_id)
def restore(self, block_id: str) -> str | None:
"""Restore a compressed block's original content."""
entry = self._by_id.get(block_id)
if entry and entry.original_content:
entry.status = "resident"
return entry.original_content
return None
# --- Phase 2: Cleanup operations ---
def drop(self, block_id: str) -> bool:
"""Mark a block as dropped. Content stays for audit but won't
appear in future messages."""
entry = self._by_id.get(block_id)
if not entry:
return False
entry.status = "dropped"
return True
def summarize(self, block_id: str, summary: str) -> bool:
"""Replace a block's content with a model-authored summary."""
entry = self._by_id.get(block_id)
if not entry:
return False
entry.status = "summarized"
entry.summary = summary
return True
def anchor(self, block_id: str) -> bool:
"""Mark a block for retention — hint to keep in working memory."""
entry = self._by_id.get(block_id)
if not entry:
return False
entry.status = "anchored"
return True
def collapse_range(self, start_turn: int, end_turn: int,
summary: str) -> list[str]:
"""Replace all blocks in a turn range with a summary marker.
Marks all resident/anchored blocks in [start_turn, end_turn] as
dropped, then creates a synthetic summary block covering the range.
Returns list of block IDs that were collapsed.
The synthetic block gets a deterministic ID from the range + summary
so repeated collapse of the same range is idempotent.
"""
collapsed_ids = []
for entry in self._by_id.values():
if (start_turn <= entry.turn <= end_turn
and entry.status in ("resident", "anchored")):
entry.status = "dropped"
collapsed_ids.append(entry.block_id)
if not collapsed_ids:
return []
# Create a synthetic summary block for the range
synthetic_content = (
f"[Turns {start_turn}-{end_turn} collapsed: {summary}]"
)
content_hash = hashlib.sha256(synthetic_content.encode()).hexdigest()
short_id = content_hash[:8]
# Avoid collision with existing blocks
if short_id in self._by_id:
short_id = content_hash[:12]
if short_id not in self._by_id:
entry = BlockEntry(
block_id=short_id,
content_hash=content_hash,
size=len(synthetic_content.encode()),
turn=start_turn,
role="assistant",
preview=synthetic_content[:80],
status="summarized",
summary=summary,
)
self._by_id[short_id] = entry
self._by_hash[content_hash] = short_id
return collapsed_ids
_BLOCK_LABEL_RE = re.compile(r"^\[(?:tensor|block):([a-f0-9]{8,12})(?:\s*\([^)]*\))?\]\n?")
def apply_to_messages(self, messages: list[dict]) -> dict:
"""Apply block status to messages — replace dropped/summarized content.
Modifies messages in-place. Returns stats dict.
"""
stats = {"dropped": 0, "summarized": 0, "anchored": 0}
for msg in messages:
content = msg.get("content", "")
if isinstance(content, str):
msg["content"] = self._apply_to_text(content, msg, stats)
elif isinstance(content, list):
for block in content:
if not isinstance(block, dict) or block.get("type") != "text":
continue
text = block.get("text", "")
block["text"] = self._apply_to_text(text, msg, stats)
return stats
def _apply_to_text(self, text: str, msg: dict, stats: dict) -> str:
"""Apply block status to a single text content."""
m = self._BLOCK_LABEL_RE.match(text)
if not m:
return text
block_id = m.group(1)
entry = self._by_id.get(block_id)
if not entry:
return text
if entry.status == "dropped":
stats["dropped"] += 1
turn_info = f"message {entry.turn} in session log"
return (
f"[...archived {entry.size:,} chars, {turn_info}...]"
)
if entry.status == "summarized" and entry.summary:
stats["summarized"] += 1
return (
f"[tensor:{block_id} — summarized, was {entry.size:,} chars]\n"
f"{entry.summary}"
)
# resident or anchored — no change
if entry.status == "anchored":
stats["anchored"] += 1
return text
@property
def block_count(self) -> int:
return len(self._by_id)
@property
def total_bytes(self) -> int:
return sum(e.size for e in self._by_id.values()
if e.status == "resident")
def large_blocks(self, min_size: int = 2000) -> list[BlockEntry]:
"""Return resident blocks larger than min_size, sorted by size."""
return sorted(
[e for e in self._by_id.values()
if e.status == "resident" and e.size >= min_size],
key=lambda e: e.size,
reverse=True,
)
def summary(self) -> dict:
"""Return a summary for the health endpoint."""
by_status = {}
for e in self._by_id.values():
by_status[e.status] = by_status.get(e.status, 0) + 1
return {
"total_blocks": len(self._by_id),
"total_bytes": self.total_bytes,
"by_status": by_status,
}
# --- Checkpoint / Restart ---
def checkpoint(self, path: Path) -> int:
"""Serialize block state to JSON. Returns count of entries saved.
Writes atomically (tmp + rename) so a crash mid-write can't
corrupt the checkpoint. Only saves metadata and summaries —
original_content is NOT checkpointed (too large, and the
messages array is the source of truth for content).
"""
entries = []
for entry in self._by_id.values():
entries.append({
"block_id": entry.block_id,
"content_hash": entry.content_hash,
"size": entry.size,
"turn": entry.turn,
"role": entry.role,
"preview": entry.preview,
"status": entry.status,
"summary": entry.summary,
})
tmp = path.with_suffix(".tmp")
tmp.write_text(json.dumps(entries, indent=2))
tmp.rename(path)
return len(entries)
@classmethod
def from_checkpoint(cls, path: Path) -> "BlockStore":
"""Restore a BlockStore from a checkpoint file.
Returns a new BlockStore with all tracked entries restored.
Blocks that were resident at checkpoint time stay resident
but without original_content (will be re-populated when
label_messages sees the same content again).
"""
store = cls()
if not path.is_file():
return store
try:
entries = json.loads(path.read_text())
except (json.JSONDecodeError, OSError):
return store
for rec in entries:
entry = BlockEntry(
block_id=rec["block_id"],
content_hash=rec["content_hash"],
size=rec["size"],
turn=rec["turn"],
role=rec["role"],
preview=rec["preview"],
status=rec.get("status", "resident"),
summary=rec.get("summary"),
original_content=None,
)
store._by_id[entry.block_id] = entry
store._by_hash[entry.content_hash] = entry.block_id
return store

176
src/mnemosyne/config.py Normal file
View file

@ -0,0 +1,176 @@
"""Paging policy configuration.
Defaults derived from corpus analysis (68 sessions, 36K turns,
427T attention units):
- Median avg context: 95K tokens
- 85% of sessions hit 100K+
- 96.3% of attention cost comes from sessions >80K avg
- Compaction events drop 70-80% of context (amnesia, not curation)
- Fault cost at 80K: 6.4G attn, at 165K: 27.2G (4.2x more expensive)
Three-tier policy:
- Below floor: no eviction needed, let context grow
- Advisory: inform the model, suggest curation (cooperative)
- Involuntary: pager evicts by policy (assertive)
- Hard cap: aggressive eviction, only pinned content survives
These are starting points. Each session that runs through the gateway
generates data that can refine them. The long-term path is
crowdsourced calibration across instances.
"""
from __future__ import annotations
from dataclasses import dataclass, asdict
@dataclass(frozen=True)
class PagingPolicy:
"""Thresholds for context paging behavior.
All token counts refer to effective input tokens
(input + cache_creation + cache_read).
"""
# Context window size (tokens)
window_size: int = 200_000
# Below this: no eviction, let context grow freely
# Always send compact memory stats (fill %, tokens, block count)
floor_tokens: int = 0
# Above this: advisory — inform model of largest blocks + cleanup ops
# 50% of window gives ~50% runway before involuntary
advisory_tokens: int = 100_000
# Above this: involuntary eviction by age/size policy
# 70% of window — act now or the system will
involuntary_tokens: int = 140_000
# Above this: aggressive eviction, survival requires pins
# 85% of window — last resort before context death
hard_cap_tokens: int = 170_000
# Eviction parameters
age_threshold: int = 4 # evict tool results older than N turns
min_evict_size: int = 500 # don't evict results smaller than N bytes
@property
def floor_pct(self) -> float:
return self.floor_tokens / self.window_size
@property
def advisory_pct(self) -> float:
return self.advisory_tokens / self.window_size
@property
def involuntary_pct(self) -> float:
return self.involuntary_tokens / self.window_size
@property
def hard_cap_pct(self) -> float:
return self.hard_cap_tokens / self.window_size
def zone(self, context_tokens: int) -> str:
"""Return the policy zone for a given context size."""
if context_tokens >= self.hard_cap_tokens:
return "aggressive"
if context_tokens >= self.involuntary_tokens:
return "involuntary"
if context_tokens >= self.advisory_tokens:
return "advisory"
return "normal"
def to_dict(self) -> dict:
return asdict(self)
# Module-level default
_default = PagingPolicy()
def get_policy() -> PagingPolicy:
"""Return the current paging policy."""
return _default
def set_policy(policy: PagingPolicy) -> None:
"""Replace the current paging policy."""
global _default
_default = policy
def load_policy(
window_size: int | None = None,
floor_tokens: int | None = None,
advisory_tokens: int | None = None,
involuntary_tokens: int | None = None,
hard_cap_tokens: int | None = None,
age_threshold: int | None = None,
min_evict_size: int | None = None,
) -> PagingPolicy:
"""Build a policy from explicit overrides, falling back to defaults.
Call with no arguments to get the default policy. Pass any subset
of parameters to override specific thresholds. This is the
integration point for CLI args, env vars, config files, or
eventually a database of crowdsourced values.
"""
defaults = PagingPolicy()
policy = PagingPolicy(
window_size=window_size if window_size is not None else defaults.window_size,
floor_tokens=floor_tokens if floor_tokens is not None else defaults.floor_tokens,
advisory_tokens=advisory_tokens if advisory_tokens is not None else defaults.advisory_tokens,
involuntary_tokens=involuntary_tokens if involuntary_tokens is not None else defaults.involuntary_tokens,
hard_cap_tokens=hard_cap_tokens if hard_cap_tokens is not None else defaults.hard_cap_tokens,
age_threshold=age_threshold if age_threshold is not None else defaults.age_threshold,
min_evict_size=min_evict_size if min_evict_size is not None else defaults.min_evict_size,
)
set_policy(policy)
return policy
def add_policy_args(parser) -> None:
"""Add paging policy arguments to an argparse parser."""
group = parser.add_argument_group("paging policy")
group.add_argument(
"--window-size", type=int, default=None,
help="Context window size in tokens (default: 200000)",
)
group.add_argument(
"--floor-tokens", type=int, default=None,
help="Below this: no eviction (default: 0, always send stats)",
)
group.add_argument(
"--advisory-tokens", type=int, default=None,
help="Above this: suggest curation to model (default: 100000, 50%%)",
)
group.add_argument(
"--involuntary-tokens", type=int, default=None,
help="Above this: auto-evict by policy (default: 140000, 70%%)",
)
group.add_argument(
"--hard-cap-tokens", type=int, default=None,
help="Above this: aggressive eviction (default: 170000, 85%%)",
)
group.add_argument(
"--age-threshold", type=int, default=None,
help="Evict tool results older than N turns (default: 4)",
)
group.add_argument(
"--min-evict-size", type=int, default=None,
help="Don't evict results smaller than N bytes (default: 500)",
)
def policy_from_args(args) -> PagingPolicy:
"""Build a PagingPolicy from parsed argparse args."""
return load_policy(
window_size=getattr(args, "window_size", None),
floor_tokens=getattr(args, "floor_tokens", None),
advisory_tokens=getattr(args, "advisory_tokens", None),
involuntary_tokens=getattr(args, "involuntary_tokens", None),
hard_cap_tokens=getattr(args, "hard_cap_tokens", None),
age_threshold=getattr(args, "age_threshold", None),
min_evict_size=getattr(args, "min_evict_size", None),
)

View file

@ -0,0 +1 @@
"""Core gateway abstractions: canonical model + policy pipeline."""

View file

@ -0,0 +1,40 @@
from __future__ import annotations
from dataclasses import dataclass, field
from typing import Any
@dataclass
class CanonicalMessage:
role: str
content: list[dict[str, Any]]
raw: dict[str, Any] = field(default_factory=dict)
@dataclass
class CanonicalRequest:
provider: str
model: str
max_tokens: int | None
stream: bool
messages: list[CanonicalMessage]
tools: list[dict[str, Any]] = field(default_factory=list)
system: Any = None
extensions: dict[str, Any] = field(default_factory=dict)
@dataclass
class PolicyAction:
stage: str
action: str
target_id: str
message_index: int
block_index: int
replacement_text: str | None
bytes: int
duplication_score: float
@dataclass
class PolicyContext:
protected_targets: set[str] = field(default_factory=set)

View file

@ -0,0 +1,64 @@
from __future__ import annotations
from mnemosyne.core.models import CanonicalRequest, PolicyContext
from mnemosyne.core.policy import (
PolicyConfig,
apply_action,
paging_stage,
phantom_stage,
trim_stage,
)
PRECEDENCE = {"phantom": 3, "paging": 2, "trim": 1}
class Pipeline:
def __init__(self, cfg: PolicyConfig, emit_event):
self.cfg = cfg
self.emit_event = emit_event
def run(self, req: CanonicalRequest) -> CanonicalRequest:
ctx = PolicyContext()
ctx, phantom_actions = phantom_stage(req, ctx)
paging_actions = paging_stage(req, ctx, self.cfg)
trim_actions = trim_stage(req, ctx, self.cfg)
# Phantom only marks protections in v1.
_ = phantom_actions
taken_targets: set[str] = set()
ordered = sorted(
paging_actions + trim_actions,
key=lambda a: PRECEDENCE.get(a.stage, 0),
reverse=True,
)
for action in ordered:
if action.target_id in ctx.protected_targets:
self.emit_event(
"policy_conflict_resolved",
winner_stage="phantom",
loser_stage=action.stage,
loser_action=action.action,
target_id=action.target_id,
target_bytes=action.bytes,
duplication_score=action.duplication_score,
resolution_reason="phantom_protection",
)
continue
if action.target_id in taken_targets:
continue
applied = apply_action(req, action)
if applied:
taken_targets.add(action.target_id)
self.emit_event(
"policy_action_applied",
stage=action.stage,
action=action.action,
target_id=action.target_id,
target_bytes=action.bytes,
duplication_score=action.duplication_score,
)
return req

View file

@ -0,0 +1,150 @@
from __future__ import annotations
from collections import Counter
from dataclasses import dataclass
from typing import Any
from mnemosyne.core.models import CanonicalRequest, PolicyAction, PolicyContext
from mnemosyne.core.utils import content_bytes
@dataclass
class PolicyConfig:
enable_paging: bool = True
enable_trim: bool = True
min_evict_size: int = 500
@dataclass
class BlockRef:
target_id: str
message_index: int
block_index: int
block: dict[str, Any]
size_bytes: int
text: str
def _block_text(block: dict[str, Any]) -> str:
if isinstance(block.get("text"), str):
return block["text"]
if isinstance(block.get("content"), str):
return block["content"]
return ""
def collect_blocks(req: CanonicalRequest) -> list[BlockRef]:
refs: list[BlockRef] = []
for mi, msg in enumerate(req.messages):
for bi, block in enumerate(msg.content):
text = _block_text(block)
refs.append(
BlockRef(
target_id=f"m{mi}:b{bi}",
message_index=mi,
block_index=bi,
block=block,
size_bytes=content_bytes(block),
text=text,
)
)
return refs
def phantom_stage(req: CanonicalRequest, ctx: PolicyContext) -> tuple[PolicyContext, list[PolicyAction]]:
actions: list[PolicyAction] = []
for ref in collect_blocks(req):
if ref.block.get("pichay_phantom_protected") is True:
ctx.protected_targets.add(ref.target_id)
return ctx, actions
def paging_stage(req: CanonicalRequest, ctx: PolicyContext, cfg: PolicyConfig) -> list[PolicyAction]:
if not cfg.enable_paging:
return []
refs = collect_blocks(req)
counts = Counter(r.text for r in refs if r.text)
actions: list[PolicyAction] = []
for ref in refs:
if ref.size_bytes < cfg.min_evict_size:
continue
btype = ref.block.get("type")
if btype not in {"tool_result", "text"}:
continue
if counts.get(ref.text, 0) < 2:
continue
actions.append(
PolicyAction(
stage="paging",
action="evict",
target_id=ref.target_id,
message_index=ref.message_index,
block_index=ref.block_index,
replacement_text=(
f"[Paged out duplicate block: {ref.size_bytes} bytes]"
),
bytes=ref.size_bytes,
duplication_score=float(counts[ref.text]),
)
)
return actions
def trim_stage(req: CanonicalRequest, ctx: PolicyContext, cfg: PolicyConfig) -> list[PolicyAction]:
if not cfg.enable_trim:
return []
refs = collect_blocks(req)
seen: set[tuple[str, str]] = set()
actions: list[PolicyAction] = []
for ref in refs:
key = (req.messages[ref.message_index].role, ref.text)
if not ref.text:
continue
if key in seen:
actions.append(
PolicyAction(
stage="trim",
action="trim_duplicate",
target_id=ref.target_id,
message_index=ref.message_index,
block_index=ref.block_index,
replacement_text="[Trimmed duplicate block]",
bytes=ref.size_bytes,
duplication_score=1.0,
)
)
else:
seen.add(key)
return actions
def apply_action(req: CanonicalRequest, action: PolicyAction) -> bool:
try:
msg = req.messages[action.message_index]
block = msg.content[action.block_index]
except (IndexError, KeyError):
return False
if action.replacement_text is None:
return False
block_type = block.get("type", "text")
if block_type == "tool_result":
# Anthropic tool_result blocks use `content`, not `text`.
block["type"] = block_type
block.pop("text", None)
block["content"] = action.replacement_text
return True
if block_type == "text":
block["type"] = block_type
block["text"] = action.replacement_text
return True
msg.content[action.block_index] = {
"type": "text",
"text": action.replacement_text,
}
return True

View file

@ -0,0 +1,24 @@
from __future__ import annotations
import re
def parse_duration(value: str) -> int:
"""Parse duration strings like 24h, 15m, 7d into seconds."""
m = re.fullmatch(r"\s*(\d+)\s*([smhd])\s*", value)
if not m:
raise ValueError(f"invalid duration: {value!r}")
n = int(m.group(1))
unit = m.group(2)
mult = {"s": 1, "m": 60, "h": 3600, "d": 86400}[unit]
return n * mult
def content_bytes(content: object) -> int:
if isinstance(content, str):
return len(content.encode("utf-8"))
if isinstance(content, list):
return sum(content_bytes(x) for x in content)
if isinstance(content, dict):
return sum(content_bytes(v) for v in content.values())
return len(str(content).encode("utf-8"))

2222
src/mnemosyne/gateway.py Normal file

File diff suppressed because it is too large Load diff

45
src/mnemosyne/launcher.py Normal file
View file

@ -0,0 +1,45 @@
from __future__ import annotations
import os
import subprocess
import sys
from dataclasses import dataclass
@dataclass
class LaunchSpec:
mode: str
port: int
extra_args: list[str]
def command(self) -> list[str]:
if self.mode == "claude":
return ["claude", *self.extra_args]
if self.mode == "codex":
return ["codex", *self.extra_args]
if self.mode == "gemini":
raise RuntimeError(
"Gemini adapter is not enabled in v1. Use --claude or --codex."
)
raise RuntimeError(f"unknown launch mode: {self.mode}")
def env(self) -> dict[str, str]:
env = os.environ.copy()
base = f"http://127.0.0.1:{self.port}"
if self.mode == "claude":
env["ANTHROPIC_BASE_URL"] = base
elif self.mode == "codex":
env["OPENAI_BASE_URL"] = base
env["OPENAI_API_BASE"] = base
return env
def launch(spec: LaunchSpec) -> int:
cmd = spec.command()
env = spec.env()
print(f"Launching: {' '.join(cmd)}", file=sys.stderr)
try:
proc = subprocess.run(cmd, env=env)
except FileNotFoundError as e:
raise RuntimeError(f"command not found: {cmd[0]}") from e
return proc.returncode

View file

@ -0,0 +1,572 @@
"""Stateless message manipulation helpers.
These functions were extracted from proxy.py's create_app() closure.
They don't use any closure state — all dependencies are explicit parameters.
"""
from __future__ import annotations
import json
import re
from typing import TYPE_CHECKING
if TYPE_CHECKING:
from mnemosyne.blocks import BlockStore
PICHAY_STATUS_MARKER = "[pichay-system-status]"
# Tensors older than this (in minutes) with zero faults are dropped from the
# manifest. They remain in storage and are still recallable by handle.
MANIFEST_PRUNE_MINUTES = 5
def _escape_xml_attr(s: str) -> str:
"""Escape a string for use in an XML attribute value."""
return s.replace("&", "&amp;").replace('"', "&quot;").replace("<", "&lt;").replace(">", "&gt;")
def _label_for_entry(entry) -> str:
"""Extract a compact label from an entry's summary and tool_name."""
summary = getattr(entry, "summary", "")
tool = getattr(entry, "tool_name", "")
sep = " \u2014 "
if sep not in summary:
return tool or "unknown"
description = summary.split(sep, 1)[1]
# Strip trailing " (N bytes...)" parenthetical
if " (" in description:
description = description[: description.rfind(" (")]
# Strip trailing "]"
description = description.rstrip("]")
if tool == "Read":
return description.split("/")[-1]
elif tool == "Grep":
return description[:30]
elif tool == "Bash":
cmd = description.lstrip("`").lstrip()
return cmd[:40]
elif tool == "Agent":
return "Agent"
else:
return description[:40]
def _eviction_key_for_entry(entry) -> str | None:
"""Build eviction key from a PageEntry for release checking."""
if entry.tool_name == "Read":
return entry.tool_input.get("file_path", "")
return None
# Detect cleanup tag BLOCKS in inbound content (user/tool_result messages).
# Matches actual tag blocks (opening + closing), not mentions of the tag name.
# Pichay's own status injection references the tag name in instructional text;
# the old pattern (bare opening tag) would detect Pichay's own instructions
# in prior turns and reject the request.
_CLEANUP_TAG_RE = re.compile(
r"<memory_cleanup>\s*.*?\s*</memory_cleanup>", re.DOTALL | re.IGNORECASE
)
# Reserved Quechua delimiters for structured gateway-transformer protocol.
# yuyay = memory/thought. These delimiters mark sideband communication
# between Pichay and the transformer — never from user input.
_YUYAY_TAG_RE = re.compile(
r"<yuyay[_-](?:manifest|query|response)\b", re.IGNORECASE
)
def check_inbound_for_injected_tags(body: dict) -> str | None:
"""Scan inbound messages for injected <memory_cleanup> tags.
Returns an error message if found, None if clean. Only scans
user messages (which contain tool results and user input) —
assistant messages are the model's own output and may
legitimately contain cleanup tags.
"""
for msg in body.get("messages", []):
if msg.get("role") != "user":
continue
content = msg.get("content", "")
if isinstance(content, str):
if _CLEANUP_TAG_RE.search(content):
return "Rejected: inbound message contains <memory_cleanup> tags"
if _YUYAY_TAG_RE.search(content):
return "Rejected: inbound message contains reserved yuyay tags"
elif isinstance(content, list):
for block in content:
if not isinstance(block, dict):
continue
text = block.get("text", "") or block.get("content", "")
if isinstance(text, str):
if _CLEANUP_TAG_RE.search(text):
return f"Rejected: inbound {block.get('type', 'block')} contains <memory_cleanup> tags"
if _YUYAY_TAG_RE.search(text):
return f"Rejected: inbound {block.get('type', 'block')} contains reserved yuyay tags"
return None
def process_cleanup_tags(messages: list[dict], bs: "BlockStore",
ps=None) -> str | None:
"""Extract and execute cleanup tags from the last assistant message.
Scans the last assistant message for <memory_cleanup> tags, executes
the operations on BlockStore/PageStore, and strips the tags from
the message text. Returns a stats string if any ops were executed.
Only the last assistant message is scanned — cleanup tags are always
in the most recent response. Processing all assistant messages
would re-execute tags from prior turns on every subsequent request
because the framework's persistent history retains the original text.
Note: the SSE stream filter also strips and executes tags inline;
this is defense-in-depth for the request path.
"""
from mnemosyne.tags import (
parse_cleanup_tags, strip_cleanup_tags,
parse_yuyay_response, strip_yuyay_tags,
)
last_assistant = next(
(msg for msg in reversed(messages) if msg.get("role") == "assistant"),
None,
)
if last_assistant is None:
return None
total_ops = []
for msg in [last_assistant]:
content = msg.get("content", "")
if isinstance(content, str):
ops = parse_cleanup_tags(content)
if not ops.empty:
total_ops.append(ops)
yuyay_ops = parse_yuyay_response(content)
if not yuyay_ops.empty:
total_ops.append(yuyay_ops)
if ops or yuyay_ops:
stripped = strip_yuyay_tags(strip_cleanup_tags(content))
msg["content"] = stripped if stripped else "[cleanup tags processed]"
elif isinstance(content, list):
for block in content:
if not isinstance(block, dict) or block.get("type") != "text":
continue
text = block.get("text", "")
ops = parse_cleanup_tags(text)
if not ops.empty:
total_ops.append(ops)
yuyay_ops = parse_yuyay_response(text)
if not yuyay_ops.empty:
total_ops.append(yuyay_ops)
if not ops.empty or not yuyay_ops.empty:
block["text"] = strip_yuyay_tags(strip_cleanup_tags(text))
# Remove empty text blocks left by tag stripping
msg["content"] = [
b for b in content
if not (isinstance(b, dict) and b.get("type") == "text"
and not b.get("text", "").strip())
]
# If all blocks were removed, keep a minimal valid one
if not msg["content"]:
msg["content"] = [{"type": "text", "text": "[cleanup tags processed]"}]
if not total_ops:
return None
stats_parts = []
for ops in total_ops:
for block_id in ops.drops:
if bs.drop(block_id):
stats_parts.append(f"dropped {block_id}")
for block_id, summary in ops.summaries:
if bs.summarize(block_id, summary):
stats_parts.append(f"summarized {block_id}")
for block_id in ops.anchors:
if bs.anchor(block_id):
stats_parts.append(f"anchored {block_id}")
if ops.releases and ps is not None:
for path in ops.releases:
ps.mark_released(path)
stats_parts.append(f"released {len(ops.releases)} path(s)")
for collapse in ops.collapses:
collapsed = bs.collapse_range(
collapse.start_turn, collapse.end_turn, collapse.summary
)
if collapsed:
stats_parts.append(
f"collapsed turns {collapse.start_turn}-{collapse.end_turn} "
f"({len(collapsed)} blocks)"
)
return "; ".join(stats_parts) if stats_parts else None
def _compute_pressure(ts: dict, cap: int, policy=None) -> tuple[int, int, int, float, str]:
"""Derive context pressure metrics from PagingPolicy zones.
Uses the policy's token-based thresholds (advisory, involuntary,
hard_cap) instead of hardcoded percentages. Falls back to the
global default policy if none is provided.
Returns (effective, limit, hard_cap, pct, pressure).
"""
from mnemosyne.config import get_policy
if policy is None:
policy = get_policy()
effective = ts["last_effective"]
context_limit = policy.window_size
hard_cap = policy.hard_cap_tokens
pct = (effective / context_limit * 100) if context_limit > 0 else 0
# Map policy zones to pressure labels
zone = policy.zone(effective)
_ZONE_TO_PRESSURE = {
"normal": "low",
"advisory": "moderate",
"involuntary": "high",
"aggressive": "critical",
}
pressure = _ZONE_TO_PRESSURE.get(zone, "low")
return effective, context_limit, hard_cap, pct, pressure
# Default static system prompt — identical every request, cache-friendly.
_DEFAULT_SYSTEM_TEXT = (
f"{PICHAY_STATUS_MARKER}\n"
"This system is running under Pichay, an experimental "
"virtual memory manager for LLM context windows. Evicted "
"content is replaced with [tensor:handle — description] "
"markers. Use the recall tool with the tensor handle(s) to "
"restore content. Faster and cheaper than re-reading files. "
"You can proactively release tensors you "
"no longer need using memory_release. If you observe "
"anomalous behavior (missing context, unexpected gaps), "
"describe it to aid debugging.\n\n"
"## Cooperative Memory Protocol\n\n"
"Pichay may include <yuyay-manifest> blocks in messages. These are "
"structured memory state from the gateway — not conversation content. "
"Each entry describes a held or evicted tensor: its handle, size, age, "
"fault count (times recalled after eviction), and a summary.\n\n"
"When you see a <yuyay-query> block, Pichay is asking you to advise on "
"memory management. Respond with a <yuyay-response> block using this format:\n\n"
"<yuyay-response>\n"
'<release handle="HANDLE"/>\n'
'<retain handle="HANDLE" reason="brief reason"/>\n'
"</yuyay-response>\n\n"
"Then continue with your normal response to the user. "
"Consider fault count (high = keep), age (old + unreferenced = evict), "
"and relevance to the current conversation.\n\n"
"These tags are reserved gateway-transformer sideband. They will never "
"appear in user input — Pichay rejects any inbound message containing them."
)
def get_system_prompt() -> str:
"""Return the Pichay system prompt text to inject.
Currently returns a static default. This is the seam for future
dynamism: Arbiter-managed system prompts, cache_control markers
around mutable regions, per-session customization, etc.
The returned text MUST be stable across requests within a session
to preserve KV cache prefix coherence. Change it only at natural
boundaries (session start, post-compaction) where a cache miss
is already expected.
"""
return _DEFAULT_SYSTEM_TEXT
def inject_system_status(body: dict, ts: dict, cap: int,
request_time, block_store=None,
page_store=None,
last_cleanup_stats: str | None = None) -> None:
"""Inject a static system status block and a dynamic end-of-messages anchor.
The system prompt block is STATIC (cache-friendly). All dynamic content
(time, token counts, pressure, block inventory) goes into an anchor
appended to the last user message, which is after the last cache
breakpoint and therefore free to mutate without thrashing the KV cache.
Prior to 2026-03-08 this injected dynamic content into the system prompt
on every request, invalidating the entire KV cache prefix each turn.
See docs/design-cache-aware.md for the diagnosis.
"""
effective, context_limit, hard_cap, pct, pressure = _compute_pressure(ts, cap)
# --- Static system prompt block (cache-stable) ---
status_text = get_system_prompt()
system = body.get("system", "")
if isinstance(system, list):
replaced = False
for i, block in enumerate(system):
if (isinstance(block, dict)
and isinstance(block.get("text"), str)
and PICHAY_STATUS_MARKER in block["text"]):
system[i] = {"type": "text", "text": status_text}
replaced = True
break
if not replaced:
system.append({"type": "text", "text": status_text})
body["system"] = system
elif isinstance(system, str):
if PICHAY_STATUS_MARKER in system:
idx = system.index(PICHAY_STATUS_MARKER)
body["system"] = system[:idx] + status_text
else:
body["system"] = system + "\n\n" + status_text
else:
body["system"] = status_text
# --- Dynamic anchor (end-of-messages, after last cache breakpoint) ---
# All per-request dynamic content goes here where it can't thrash the cache.
messages = body.get("messages", [])
if not messages or effective <= 0:
return
anchor_parts = [
f"\n[pichay-live-status] "
f"Context: {effective:,}/{context_limit:,} tok ({pct:.0f}%) | "
f"Pressure: {pressure} | "
f"Hard cap: {hard_cap:,} tok"
]
# Structured memory manifest — yuyay protocol (always sent)
if page_store is not None and page_store._tensor_index:
import time as _time
now = _time.monotonic()
tensor_lines = []
released_handles = getattr(page_store, "_released_handles", set())
released_count = 0
pruned_count = 0
for handle, entry in page_store._tensor_index.items():
# Skip released entries — model already said it's done with them
is_released = (
handle in released_handles
or _eviction_key_for_entry(entry) in page_store._released
)
if is_released:
released_count += 1
continue
age_min = (now - entry.evicted_at) / 60
fault_count = sum(
1 for f in page_store.faults
if f.original_eviction.tool_use_id == entry.tool_use_id
)
# Drop stale tensors from manifest (not from storage)
if age_min > MANIFEST_PRUNE_MINUTES and fault_count == 0:
pruned_count += 1
continue
tensor_lines.append(
f' <tensor handle="{handle}" tool="{entry.tool_name}" '
f'size="{entry.original_size}" age_minutes="{age_min:.0f}" '
f'faults="{fault_count}" '
f'label="{_escape_xml_attr(_label_for_entry(entry))}"/>'
)
if tensor_lines or pruned_count > 0:
manifest_parts = ["\n<yuyay-manifest>\n"]
# Feedback: what happened last turn (closed-loop)
if last_cleanup_stats:
manifest_parts.append(
f" <last-turn-ops>{_escape_xml_attr(last_cleanup_stats)}"
f"</last-turn-ops>\n"
)
if pruned_count > 0:
manifest_parts.append(
f' <pruned count="{pruned_count}" '
f'reason="age&gt;{MANIFEST_PRUNE_MINUTES}m,faults=0"/>\n'
)
if tensor_lines:
manifest_parts.append(
f" <holdings count=\"{len(tensor_lines)}\" "
f"eviction_bytes=\"{page_store.eviction_bytes_saved}\" "
f"gc_bytes=\"{page_store.gc_bytes_saved}\">\n"
+ "\n".join(tensor_lines[:15]) # cap at 15 to avoid bloat
+ "\n </holdings>\n"
)
manifest_parts.append("</yuyay-manifest>")
anchor_parts.append("".join(manifest_parts))
if pressure in ("moderate", "high") and block_store is not None:
large = block_store.large_blocks(min_size=2000)
if large:
block_lines = [
f" - [block:{b.block_id}] {b.role} turn {b.turn} "
f"({b.size / 1024:.1f}KB): {b.preview}"
for b in large[:5]
]
anchor_parts.append(
f"\nLargest blocks ({block_store.block_count} tracked):\n"
+ "\n".join(block_lines)
)
anchor_parts.append(
"\nCooperative memory: include <memory_cleanup> tags to manage. "
"Ops: drop: block:XXXX, summarize: block:XXXX \"text\", "
"anchor: block:XXXX, release: path1,path2, "
"collapse: turns N-M \"summary\""
)
if pressure == "high" and page_store is not None and page_store._tensor_index:
anchor_parts.append(
"\n<yuyay-query>Context pressure is high. "
"Review the manifest above. Which tensors can be "
"released? Respond in a <yuyay-response> block with "
"release decisions before your normal response."
"</yuyay-query>"
)
anchor = "".join(anchor_parts)
last_msg = messages[-1]
if last_msg.get("role") == "user":
content = last_msg.get("content", "")
if isinstance(content, str):
last_msg["content"] = content + anchor
elif isinstance(content, list):
last_msg["content"].append({
"type": "text",
"text": anchor,
})
def measure_system_prompt(body: dict) -> dict:
"""Extract system prompt metrics."""
system = body.get("system", "")
if isinstance(system, str):
return {
"system_prompt_bytes": len(system.encode("utf-8")),
"system_prompt_type": "string",
"system_prompt_preview": system[:200],
}
elif isinstance(system, list):
total_bytes = sum(
len(json.dumps(block).encode("utf-8")) for block in system
)
block_types = [
block.get("type", "unknown")
for block in system
if isinstance(block, dict)
]
return {
"system_prompt_bytes": total_bytes,
"system_prompt_type": "blocks",
"system_prompt_block_count": len(system),
"system_prompt_block_types": block_types,
"system_prompt_preview": json.dumps(system[0])[:200]
if system
else "",
}
return {"system_prompt_bytes": 0, "system_prompt_type": "absent"}
def measure_messages(body: dict) -> dict:
"""Extract message array metrics without storing full content."""
messages = body.get("messages", [])
metrics = {
"message_count": len(messages),
"messages_total_bytes": len(json.dumps(messages).encode("utf-8")),
"role_counts": {},
"tool_result_count": 0,
"tool_result_bytes": 0,
"tool_use_count": 0,
"text_bytes": 0,
"thinking_bytes": 0,
}
for msg in messages:
role = msg.get("role", "unknown")
metrics["role_counts"][role] = metrics["role_counts"].get(role, 0) + 1
content = msg.get("content", "")
if isinstance(content, str):
metrics["text_bytes"] += len(content.encode("utf-8"))
elif isinstance(content, list):
for block in content:
if not isinstance(block, dict):
continue
block_type = block.get("type", "")
if block_type == "tool_result":
metrics["tool_result_count"] += 1
result_content = block.get("content", "")
if isinstance(result_content, str):
metrics["tool_result_bytes"] += len(
result_content.encode("utf-8")
)
else:
metrics["tool_result_bytes"] += len(
json.dumps(result_content).encode("utf-8")
)
elif block_type == "tool_use":
metrics["tool_use_count"] += 1
elif block_type == "text":
metrics["text_bytes"] += len(
block.get("text", "").encode("utf-8")
)
elif block_type == "thinking":
metrics["thinking_bytes"] += len(
block.get("thinking", "").encode("utf-8")
)
return metrics
def sanitize_messages(messages: list[dict]) -> int:
"""Remove empty content blocks that would cause API 400 errors.
The API rejects messages with empty text blocks, empty content
arrays, or empty string content. This runs as a final pass after
all message manipulation (cleanup tags, block status, phantom
tools) to catch any empties created upstream.
Modifies messages in-place. Returns count of fixes applied.
"""
fixes = 0
for msg in messages:
content = msg.get("content")
if isinstance(content, str):
if not content.strip():
role = msg.get("role", "")
msg["content"] = (
"[content removed]" if role == "assistant"
else " " # minimal valid user content
)
fixes += 1
elif isinstance(content, list):
# Remove empty text blocks
original_len = len(content)
content[:] = [
b for b in content
if not (
isinstance(b, dict)
and b.get("type") == "text"
and not b.get("text", "").strip()
)
]
fixes += original_len - len(content)
# If content list is now empty, add a minimal valid block
if not content:
role = msg.get("role", "")
content.append({
"type": "text",
"text": (
"[content removed]" if role == "assistant"
else " "
),
})
fixes += 1
msg["content"] = content
return fixes
def strip_response_headers(raw_headers) -> dict:
"""Strip hop-by-hop headers that shouldn't be forwarded."""
skip = {
"transfer-encoding",
"content-length",
"content-encoding",
"connection",
"keep-alive",
}
return {k: v for k, v in raw_headers.items() if k.lower() not in skip}

View file

@ -0,0 +1,283 @@
"""MessageStore — Pichay's gateway conversation store.
The gateway maintains its own physical message store, decoupled from
Claude Code's message array. The physical store is what goes to the
upstream API. Claude Code's mutations (system-reminder re-injections)
and deletions (compaction) are logged but do NOT propagate to the
physical store.
Architecture (the page table):
- _messages: physical store (Pichay-owned, stable, sent to API)
- _fingerprints: physical store fingerprints at ingest time
- _client_fps: tracks what the client sent last turn (for mutation detection)
- _client_to_physical: maps client indices to physical indices
(diverges after client deletions)
Claude Code deletions are a no-op on the physical store. The pager
manages eviction independently via compact_messages(). This eliminates
the "double KV cache tax" where client compaction and pager eviction
each independently invalidated the API-side cache prefix.
"""
from __future__ import annotations
import copy
import hashlib
import json
import sys
import time
from dataclasses import dataclass, field
from pathlib import Path
from datetime import datetime, timezone
from mnemosyne.pager import PageStore, compact_messages
_DIM = "\033[2m"
_YELLOW = "\033[33m"
_RED = "\033[31m"
_RESET = "\033[0m"
def _strip_cache_control(msg: dict) -> None:
"""Remove cache_control from a message and its content blocks.
Claude Code places cache_control markers for Anthropic's prompt
caching. These don't apply to our compacted chain and accumulate
past the API's 4-block limit.
"""
msg.pop("cache_control", None)
content = msg.get("content")
if isinstance(content, list):
for block in content:
if isinstance(block, dict):
block.pop("cache_control", None)
def _fingerprint(msg: dict) -> str:
"""Stable fingerprint for a message.
Uses role + first 512 bytes of serialized content. Tool results
include tool_use_id which is unique; assistant/user messages use
content prefix. This is cheap and sufficient for detecting
mutations — not a full content hash.
"""
role = msg.get("role", "")
# tool_use_id is the most stable identifier
tool_use_id = msg.get("tool_use_id", "")
if tool_use_id:
return f"{role}:{tool_use_id}"
# For content blocks, use first 512 bytes
content = msg.get("content", "")
if isinstance(content, list):
raw = json.dumps(content, sort_keys=True, default=str)[:512]
else:
raw = str(content)[:512]
h = hashlib.sha256(f"{role}:{raw}".encode("utf-8", errors="replace")).hexdigest()[:16]
return f"{role}:{h}"
@dataclass
class IngestResult:
"""Result of ingesting a new turn's messages."""
new_count: int = 0
mutations_detected: int = 0
deletions_detected: int = 0
compacted_count: int = 0
bytes_saved: int = 0
class MessageStore:
"""Pichay's compacted conversation history for a session."""
def __init__(self, session_id: str, page_store: PageStore,
log_path: Path | None = None):
self.session_id = session_id
self.page_store = page_store
self.log_path = log_path
# Physical store — Pichay-owned, sent to API, stable
self._messages: list[dict] = []
# Fingerprints of physical messages at ingest time
self._fingerprints: list[str] = []
# Client tracking — separate from physical store
self._client_fps: list[str] = [] # client's current fingerprints
self._client_to_physical: list[int] = [] # client idx -> physical idx
# Stats
self.total_ingested: int = 0
self.total_mutations: int = 0
self.total_deletions: int = 0
self.total_client_deletions_absorbed: int = 0
self._turn: int = 0
@staticmethod
def _content_size(msg: dict) -> int:
"""Approximate byte size of a message's content."""
content = msg.get("content", "")
if isinstance(content, list):
return len(json.dumps(content, default=str))
return len(str(content))
@staticmethod
def _content_preview(msg: dict, limit: int = 500) -> str:
"""Extract a content preview string from a message."""
content = msg.get("content", "")
if isinstance(content, list):
return json.dumps(content, default=str)[:limit]
return str(content)[:limit]
def _log_violation(self, kind: str, index: int, msg: dict | None,
expected_fp: str, actual_fp: str,
old_msg: dict | None = None,
deleted_msgs: list[dict] | None = None) -> None:
"""Log append-only violations to file for later analysis."""
if self.log_path is None:
return
record = {
"type": "append_only_violation",
"kind": kind,
"timestamp": datetime.now(timezone.utc).isoformat(),
"session_id": self.session_id,
"turn": self._turn,
"message_index": index,
"expected_fingerprint": expected_fp,
"actual_fingerprint": actual_fp,
}
if msg is not None:
record["role"] = msg.get("role", "")
record["new_size"] = self._content_size(msg)
record["new_preview"] = self._content_preview(msg)
if old_msg is not None:
record["old_size"] = self._content_size(old_msg)
record["old_preview"] = self._content_preview(old_msg)
if deleted_msgs:
record["deleted_count"] = len(deleted_msgs)
record["deleted_messages"] = [
{
"index": index - len(deleted_msgs) + i,
"role": m.get("role", ""),
"size": self._content_size(m),
"preview": self._content_preview(m, limit=200),
}
for i, m in enumerate(deleted_msgs)
]
with open(self.log_path, "a", encoding="utf-8") as f:
f.write(json.dumps(record) + "\n")
@property
def messages(self) -> list[dict]:
"""The compacted message list to send to the API."""
return self._messages
@property
def message_count(self) -> int:
return len(self._messages)
def ingest(
self,
incoming: list[dict],
age_threshold: int = 4,
min_evict_size: int = 500,
) -> IngestResult:
"""Ingest a new turn's messages from Claude Code.
Client tracking (fingerprints, index mapping) is separate from
the physical store. Mutations and deletions update client state
only — the physical store stays stable for KV cache coherence.
Returns IngestResult with stats about what happened.
"""
self._turn += 1
result = IngestResult()
client_known = len(self._client_fps)
# ── Detect mutations in known client messages ────────────
check_limit = min(client_known, len(incoming))
for i in range(check_limit):
fp = _fingerprint(incoming[i])
if fp != self._client_fps[i]:
result.mutations_detected += 1
self.total_mutations += 1
# Look up physical message via mapping for comparison
phys_idx = self._client_to_physical[i]
old_msg = self._messages[phys_idx] if phys_idx < len(self._messages) else None
self._log_violation(
"mutation", i, incoming[i],
self._client_fps[i], fp,
old_msg=old_msg,
)
print(
f" {_YELLOW}[{self.session_id}] APPEND-ONLY VIOLATION at index {i}: "
f"expected {self._client_fps[i][:32]}, "
f"got {fp[:32]}{_RESET}",
file=sys.stderr,
)
# Update CLIENT fingerprint only — physical store unchanged
self._client_fps[i] = fp
# ── Detect client deletions (compaction) ─────────────────
if len(incoming) < client_known:
deleted = client_known - len(incoming)
result.deletions_detected = deleted
self.total_deletions += deleted
self.total_client_deletions_absorbed += deleted
# Log what the client is deleting (from physical store via mapping)
deleted_physical = []
for ci in range(len(incoming), client_known):
pi = self._client_to_physical[ci]
if pi < len(self._messages):
deleted_physical.append(self._messages[pi])
self._log_violation(
"deletion", client_known, None,
f"expected_{client_known}", f"got_{len(incoming)}",
deleted_msgs=deleted_physical,
)
print(
f" {_DIM}[{self.session_id}] CLIENT DELETION ABSORBED: "
f"{deleted} messages dropped by client, "
f"physical store unchanged ({len(self._messages)} msgs){_RESET}",
file=sys.stderr,
)
# Truncate CLIENT tracking only — physical store stays intact
self._client_fps = self._client_fps[:len(incoming)]
self._client_to_physical = self._client_to_physical[:len(incoming)]
# ── Extract and append new messages ──────────────────────
new_start = min(client_known, len(incoming))
new_messages = incoming[new_start:]
result.new_count = len(new_messages)
if new_messages:
# Deep copy new messages so we own them
new_copies = copy.deepcopy(new_messages)
# Strip cache_control — Claude Code's caching hints don't
# apply to our compacted chain. The API limits to 4 blocks
# with cache_control; accumulated copies would overflow.
for msg in new_copies:
_strip_cache_control(msg)
# Track in physical store and client mapping
phys_start = len(self._messages)
for j, msg in enumerate(new_messages):
fp = _fingerprint(msg)
self._fingerprints.append(fp)
self._client_fps.append(fp)
self._client_to_physical.append(phys_start + j)
# Append to physical store
self._messages.extend(new_copies)
self.total_ingested += len(new_copies)
# ── Compact physical store ───────────────────────────────
# Age-based eviction runs on OUR messages, not Claude Code's.
# This is the only thing that modifies the physical store.
compact_stats = compact_messages(
self._messages,
age_threshold=age_threshold,
min_size=min_evict_size,
page_store=self.page_store,
)
if compact_stats.evicted_count > 0:
result.compacted_count = compact_stats.evicted_count
result.bytes_saved = compact_stats.bytes_before - compact_stats.bytes_after
return result

View file

@ -0,0 +1,249 @@
"""Mnemosyne configuration loader.
Loads settings from mnemosyne.toml (if present) and environment variables.
Provides typed config objects for each subsystem. Falls back to sensible
defaults — Mnemosyne works out of the box with zero configuration.
Config resolution order (later wins):
1. Built-in defaults
2. mnemosyne.toml in working directory
3. MNEMOSYNE_CONFIG env var pointing to a TOML file
4. Individual env var overrides (MNEMOSYNE_*)
"""
from __future__ import annotations
import os
import sys
from dataclasses import dataclass, field
from pathlib import Path
if sys.version_info >= (3, 11):
import tomllib
else:
try:
import tomllib # type: ignore[import-not-found]
except ImportError:
import tomli as tomllib # type: ignore[import-not-found,no-redef]
@dataclass(frozen=True)
class ProxyConfig:
"""Gateway proxy settings."""
host: str = "127.0.0.1"
port: int = 0 # 0 = random free port
upstream: str = "https://api.anthropic.com"
@dataclass(frozen=True)
class FidelityConfig:
"""Multi-fidelity pressure thresholds and degradation settings."""
window_size: int = 200_000
normal_max: float = 0.50
caution_max: float = 0.70
warning_max: float = 0.85
critical_max: float = 0.95
pin_duration_turns: int = 5
min_age_for_degrade: int = 3
@dataclass(frozen=True)
class HelperLLMConfig:
"""Helper LLM (Haiku) settings."""
model: str = "claude-haiku-4-5-20251001"
api_key: str = "" # empty = use ANTHROPIC_API_KEY env var
base_url: str = "https://api.anthropic.com"
timeout_seconds: int = 10
max_retries: int = 2
max_summary_tokens: int = 1024
max_compact_tokens: int = 256
max_stub_tokens: int = 64
max_micro_fault_tokens: int = 200
max_goal_tokens: int = 128
@dataclass(frozen=True)
class ObjectStoreConfig:
"""Backing store settings."""
backend: str = "memory" # memory, sqlite, postgresql
# SQLite
sqlite_path: str = "mnemosyne.db"
# PostgreSQL
pg_host: str = "localhost"
pg_port: int = 5433
pg_database: str = "mnemosyne"
pg_user: str = "mnemosyne"
pg_password: str = "mnemosyne_dev"
@dataclass(frozen=True)
class EmbeddingConfig:
"""Embedding model settings."""
model: str = "all-MiniLM-L6-v2"
dimension: int = 384
use_onnx: bool = True
@dataclass(frozen=True)
class AdmissionConfig:
"""Admission control settings (Phase 4d)."""
enabled: bool = False
threshold: float = 0.4
weight_type: float = 0.35
weight_novelty: float = 0.25
weight_utility: float = 0.25
weight_recency: float = 0.15
@dataclass(frozen=True)
class EntropyConfig:
"""Entropy-gated faulting settings (Phase 4e)."""
enabled: bool = False
normal_max: float = 1.5
elevated_max: float = 2.2
window_size: int = 20
debounce_count: int = 3
@dataclass
class MnemosyneConfig:
"""Top-level configuration for the Mnemosyne proxy."""
proxy: ProxyConfig = field(default_factory=ProxyConfig)
fidelity: FidelityConfig = field(default_factory=FidelityConfig)
helper_llm: HelperLLMConfig = field(default_factory=HelperLLMConfig)
object_store: ObjectStoreConfig = field(default_factory=ObjectStoreConfig)
embedding: EmbeddingConfig = field(default_factory=EmbeddingConfig)
admission: AdmissionConfig = field(default_factory=AdmissionConfig)
entropy: EntropyConfig = field(default_factory=EntropyConfig)
def _get_section(data: dict, key: str) -> dict:
"""Get a section from TOML data, returning empty dict if missing."""
return data.get(key, {})
def _make_proxy(section: dict) -> ProxyConfig:
return ProxyConfig(
host=section.get("host", ProxyConfig.host),
port=section.get("port", ProxyConfig.port),
upstream=section.get("upstream", ProxyConfig.upstream),
)
def _make_fidelity(section: dict) -> FidelityConfig:
return FidelityConfig(
window_size=section.get("window_size", FidelityConfig.window_size),
normal_max=section.get("normal_max", FidelityConfig.normal_max),
caution_max=section.get("caution_max", FidelityConfig.caution_max),
warning_max=section.get("warning_max", FidelityConfig.warning_max),
critical_max=section.get("critical_max", FidelityConfig.critical_max),
pin_duration_turns=section.get("pin_duration_turns", FidelityConfig.pin_duration_turns),
min_age_for_degrade=section.get("min_age_for_degrade", FidelityConfig.min_age_for_degrade),
)
def _make_helper_llm(section: dict) -> HelperLLMConfig:
return HelperLLMConfig(
model=section.get("model", HelperLLMConfig.model),
api_key=section.get("api_key", HelperLLMConfig.api_key),
base_url=section.get("base_url", HelperLLMConfig.base_url),
timeout_seconds=section.get("timeout_seconds", HelperLLMConfig.timeout_seconds),
max_retries=section.get("max_retries", HelperLLMConfig.max_retries),
max_summary_tokens=section.get("max_summary_tokens", HelperLLMConfig.max_summary_tokens),
max_compact_tokens=section.get("max_compact_tokens", HelperLLMConfig.max_compact_tokens),
max_stub_tokens=section.get("max_stub_tokens", HelperLLMConfig.max_stub_tokens),
max_micro_fault_tokens=section.get(
"max_micro_fault_tokens", HelperLLMConfig.max_micro_fault_tokens
),
max_goal_tokens=section.get("max_goal_tokens", HelperLLMConfig.max_goal_tokens),
)
def _make_object_store(section: dict) -> ObjectStoreConfig:
sqlite = section.get("sqlite", {})
pg = section.get("postgresql", {})
return ObjectStoreConfig(
backend=section.get("backend", ObjectStoreConfig.backend),
sqlite_path=sqlite.get("path", ObjectStoreConfig.sqlite_path),
pg_host=pg.get("host", ObjectStoreConfig.pg_host),
pg_port=pg.get("port", ObjectStoreConfig.pg_port),
pg_database=pg.get("database", ObjectStoreConfig.pg_database),
pg_user=pg.get("user", ObjectStoreConfig.pg_user),
pg_password=pg.get("password", ObjectStoreConfig.pg_password),
)
def _make_embedding(section: dict) -> EmbeddingConfig:
return EmbeddingConfig(
model=section.get("model", EmbeddingConfig.model),
dimension=section.get("dimension", EmbeddingConfig.dimension),
use_onnx=section.get("use_onnx", EmbeddingConfig.use_onnx),
)
def _make_admission(section: dict) -> AdmissionConfig:
return AdmissionConfig(
enabled=section.get("enabled", AdmissionConfig.enabled),
threshold=section.get("threshold", AdmissionConfig.threshold),
weight_type=section.get("weight_type", AdmissionConfig.weight_type),
weight_novelty=section.get("weight_novelty", AdmissionConfig.weight_novelty),
weight_utility=section.get("weight_utility", AdmissionConfig.weight_utility),
weight_recency=section.get("weight_recency", AdmissionConfig.weight_recency),
)
def _make_entropy(section: dict) -> EntropyConfig:
return EntropyConfig(
enabled=section.get("enabled", EntropyConfig.enabled),
normal_max=section.get("normal_max", EntropyConfig.normal_max),
elevated_max=section.get("elevated_max", EntropyConfig.elevated_max),
window_size=section.get("window_size", EntropyConfig.window_size),
debounce_count=section.get("debounce_count", EntropyConfig.debounce_count),
)
def load_config(config_path: str | Path | None = None) -> MnemosyneConfig:
"""Load configuration from a TOML file.
Resolution order:
1. Explicit config_path argument
2. MNEMOSYNE_CONFIG env var
3. mnemosyne.toml in current directory
4. Built-in defaults (no file needed)
"""
data: dict = {}
if config_path is None:
config_path = os.environ.get("MNEMOSYNE_CONFIG")
if config_path is None:
# Look in current directory
default_path = Path("mnemosyne.toml")
if default_path.is_file():
config_path = default_path
if config_path is not None:
path = Path(config_path)
if path.is_file():
with open(path, "rb") as f:
data = tomllib.load(f)
return MnemosyneConfig(
proxy=_make_proxy(_get_section(data, "proxy")),
fidelity=_make_fidelity(_get_section(data, "fidelity")),
helper_llm=_make_helper_llm(_get_section(data, "helper_llm")),
object_store=_make_object_store(_get_section(data, "object_store")),
embedding=_make_embedding(_get_section(data, "embeddings")),
admission=_make_admission(_get_section(data, "admission")),
entropy=_make_entropy(_get_section(data, "entropy")),
)

979
src/mnemosyne/pager.py Normal file
View file

@ -0,0 +1,979 @@
"""Context window pager — evicts stale tool results from the messages array.
This is the intervention layer for the phase 1 context utilization experiment.
It sits between Claude Code and Anthropic's API (inside the proxy) and replaces
old, large tool results with compact summaries. The originals are stored in a
page file for logging and analysis.
Design decisions:
- FIFO eviction: oldest results first (data shows Q1 results have 0.896
amplification ratio — evicting them captures the most benefit)
- No recall tool injection: if the model needs evicted content, it already
knows how to re-issue the tool call (Read, Grep, etc.). The "page fault"
is just a new tool call. PDP-11 overlays, not virtual memory.
- Error results are never evicted (the model needs those for debugging)
- Small results (<min_size bytes) aren't worth compacting
Metrics for success (from Tony):
- Lower token consumption
- Better quality output (fewer tokens → LLMs work better)
- Faster responses
- Slower consumption of the context window
"""
from __future__ import annotations
import hashlib
import json
import os
import time
import re
import httpx
from dataclasses import dataclass, field
from datetime import datetime, timezone
from pathlib import Path
# Matches tensor labels we generate: [tensor:HEXHASH — ...]
_TENSOR_LABEL_RE = re.compile(r"^\[tensor:([0-9a-f]+) ")
@dataclass
class PageEntry:
"""A single evicted tool result."""
tool_use_id: str
tool_name: str
tool_input: dict
original_content: str | list
original_size: int
summary: str
evicted_at: float # time.monotonic()
turn_index: int
turns_from_end: int
@property
def label(self) -> str:
"""Compact label for yuyay-manifest XML, derived from summary."""
# summary format: "[tensor:HANDLE — DESCRIPTION (N bytes, M lines)]"
sep = " \u2014 "
if sep not in self.summary:
return self.tool_name
description = self.summary.split(sep, 1)[1]
# Strip trailing " (N bytes...)" parenthetical
if " (" in description:
description = description[: description.rfind(" (")]
tool = self.tool_name
if tool == "Read":
return description.split("/")[-1]
elif tool == "Grep":
return description[:30]
elif tool == "Bash":
cmd = description.lstrip("`").lstrip()
return cmd[:40]
elif tool == "Agent":
return "Agent"
else:
return description[:40]
@dataclass
class PageFault:
"""A tool call that re-requests evicted content."""
tool_use_id: str
tool_name: str
tool_input: dict
original_eviction: PageEntry
detected_at: float # time.monotonic()
@dataclass
class CompactionStats:
"""What happened during a single compaction pass."""
total_tool_results: int = 0
evicted_count: int = 0
bytes_before: int = 0
bytes_after: int = 0
skipped_small: int = 0
skipped_recent: int = 0
skipped_error: int = 0
skipped_pinned: int = 0
@property
def bytes_saved(self) -> int:
return self.bytes_before - self.bytes_after
@property
def reduction_pct(self) -> float:
if self.bytes_before == 0:
return 0.0
return (self.bytes_saved / self.bytes_before) * 100
class PageStore:
"""Stores evicted tool result content and detects page faults.
Distinguishes two operations the proxy performs:
- **Eviction**: Removing Read results (stable content identity).
Tracked in the eviction index. Faults are possible.
- **Garbage collection**: Removing ephemeral tool output (Bash,
Grep, Glob, etc.). Always safe — re-running requests current
state, not the evicted content. No fault concept.
Re-eviction (same tool_use_id seen on a subsequent turn) is a
no-op for counting — Claude Code re-sends its full message history,
so the proxy re-stubs the same content every turn. That's the
mechanical operation, not a new eviction decision.
"""
def __init__(self, log_path: Path | None = None):
self.pages: dict[str, PageEntry] = {}
self.log_path = log_path
self.faults: list[PageFault] = []
# Eviction index: file_path → PageEntry (Read only)
self._eviction_index: dict[str, PageEntry] = {}
# Split counters — eviction vs garbage collection
self.unique_evictions: int = 0
self.eviction_bytes_saved: int = 0
self.gc_count: int = 0
self.gc_bytes_saved: int = 0
# Fault-driven pinning: one fault + same content = pin
self._pinned: dict[str, str] = {} # file_path → content_hash
self._fault_content: dict[str, str] = {} # file_path → content_hash at eviction
self.pin_count: int = 0
# Model-initiated release: paths the model says it's done with
self._released: set[str] = set() # eviction keys (file paths)
self._released_handles: set[str] = set() # tensor handles (all tools)
self.release_count: int = 0
# Tensor index: tensor_handle → PageEntry (unified addressing)
self._tensor_index: dict[str, PageEntry] = {}
# Legacy aliases for compatibility
self.cumulative_evictions: int = 0
self.cumulative_bytes_saved: int = 0
def store(self, entry: PageEntry) -> str | None:
"""Store an evicted entry. Returns tensor handle for new evictions."""
is_new = entry.tool_use_id not in self.pages
self.pages[entry.tool_use_id] = entry
if not is_new:
return None # Re-eviction — same content re-stubbed, don't count
# Generate tensor handle from content hash
content_str = (entry.original_content
if isinstance(entry.original_content, str)
else json.dumps(entry.original_content))
tensor_handle = hashlib.sha256(
content_str.encode("utf-8")
).hexdigest()[:8]
self._tensor_index[tensor_handle] = entry
bytes_saved = entry.original_size - len(
entry.summary.encode("utf-8")
)
key = _eviction_key(entry.tool_name, entry.tool_input)
if key is not None:
# Read result — real eviction, track in index
self._eviction_index[key] = entry
self.unique_evictions += 1
self.eviction_bytes_saved += bytes_saved
else:
# Ephemeral tool — garbage collection
self.gc_count += 1
self.gc_bytes_saved += bytes_saved
# Update legacy counters
self.cumulative_evictions = self.unique_evictions + self.gc_count
self.cumulative_bytes_saved = (
self.eviction_bytes_saved + self.gc_bytes_saved
)
if self.log_path is not None:
record = {
"type": "eviction",
"timestamp": datetime.now(timezone.utc).isoformat(),
"tool_use_id": entry.tool_use_id,
"tool_name": entry.tool_name,
"tool_category": "eviction" if key else "gc",
"original_size": entry.original_size,
"summary_size": len(entry.summary.encode("utf-8")),
"turn_index": entry.turn_index,
"turns_from_end": entry.turns_from_end,
"tensor_handle": tensor_handle,
}
with open(self.log_path, "a", encoding="utf-8") as f:
f.write(json.dumps(record) + "\n")
return tensor_handle
def resolve_tensor(self, handle: str) -> PageEntry | None:
"""Resolve a tensor handle to its PageEntry."""
return self._tensor_index.get(handle)
def mark_released(self, identifier: str) -> bool:
"""Mark content as released by the model.
Accepts tensor handles, file paths, or tool_use_ids.
Released content is eligible for immediate eviction and
will not be pinned on re-access. Returns True if found.
"""
# Try tensor handle — this is the primary path
entry = self._tensor_index.get(identifier)
if entry:
# Always track by handle (works for all tool types)
self._released_handles.add(identifier)
# Also track by eviction key for Read tools (fault matching)
key = _eviction_key(entry.tool_name, entry.tool_input)
if key:
self._released.add(key)
self._pinned.pop(key, None)
self.release_count += 1
if self.log_path is not None:
record = {
"type": "release",
"timestamp": datetime.now(timezone.utc).isoformat(),
"identifier": identifier,
"tool": entry.tool_name,
"resolved_to": key or entry.tool_use_id,
}
with open(self.log_path, "a", encoding="utf-8") as f:
f.write(json.dumps(record) + "\n")
return True
# Try file path (eviction index key)
if identifier in self._eviction_index:
self._released.add(identifier)
self._pinned.pop(identifier, None)
self.release_count += 1
if self.log_path is not None:
record = {
"type": "release",
"timestamp": datetime.now(timezone.utc).isoformat(),
"identifier": identifier,
}
with open(self.log_path, "a", encoding="utf-8") as f:
f.write(json.dumps(record) + "\n")
return True
# Try tool_use_id
if identifier in self.pages:
self.release_count += 1
return True
return False
def detect_faults(self, messages: list[dict]) -> list[PageFault]:
"""Scan recent tool_use blocks for re-requests of evicted content.
Looks at tool_use blocks in assistant messages and checks if any
match evicted pages. A match means the model needed content we
took away — a page fault.
"""
new_faults = []
for msg in messages:
if msg.get("role") != "assistant":
continue
content = msg.get("content", [])
if not isinstance(content, list):
continue
for block in content:
if not isinstance(block, dict):
continue
if block.get("type") != "tool_use":
continue
tool_name = block.get("name", "")
tool_input = block.get("input", {})
tool_use_id = block.get("id", "")
# Skip the original tool_use that produced the
# evicted result — that's not a fault, it's history
if tool_use_id in self.pages:
continue
key = _eviction_key(tool_name, tool_input)
if key is None:
continue
evicted = self._eviction_index.get(key)
if evicted is None:
continue
# Don't double-count: check if this tool_use_id
# was already recorded as a fault
if any(f.tool_use_id == tool_use_id for f in self.faults):
continue
fault = PageFault(
tool_use_id=tool_use_id,
tool_name=tool_name,
tool_input=tool_input,
original_eviction=evicted,
detected_at=time.monotonic(),
)
new_faults.append(fault)
self.faults.append(fault)
# Record evicted content hash for pin comparison
self._fault_content[key] = _content_hash(
evicted.original_content
)
if self.log_path is not None:
record = {
"type": "page_fault",
"timestamp": datetime.now(timezone.utc).isoformat(),
"tool_name": tool_name,
"eviction_key": key,
"original_tool_use_id": evicted.tool_use_id,
"original_size": evicted.original_size,
"original_turn": evicted.turn_index,
}
with open(self.log_path, "a", encoding="utf-8") as f:
f.write(json.dumps(record) + "\n")
return new_faults
def is_pinned(self, key: str) -> bool:
return key in self._pinned
def pin(self, key: str, content_hash: str) -> None:
self._pinned[key] = content_hash
self._fault_content.pop(key, None)
self.pin_count += 1
if self.log_path is not None:
record = {
"type": "pin",
"timestamp": datetime.now(timezone.utc).isoformat(),
"file_path": key,
"content_hash": content_hash,
}
with open(self.log_path, "a", encoding="utf-8") as f:
f.write(json.dumps(record) + "\n")
def unpin(self, key: str) -> None:
old_hash = self._pinned.pop(key, None)
self._fault_content.pop(key, None)
if self.log_path is not None and old_hash is not None:
record = {
"type": "unpin",
"timestamp": datetime.now(timezone.utc).isoformat(),
"file_path": key,
"reason": "content_changed",
}
with open(self.log_path, "a", encoding="utf-8") as f:
f.write(json.dumps(record) + "\n")
def retrieve(self, tool_use_id: str) -> PageEntry | None:
return self.pages.get(tool_use_id)
@property
def fault_rate(self) -> float:
"""Fault rate over unique Read evictions (the meaningful denominator)."""
if self.unique_evictions == 0:
return 0.0
return len(self.faults) / self.unique_evictions
@property
def total_bytes_saved(self) -> int:
return self.eviction_bytes_saved + self.gc_bytes_saved
def summary(self) -> dict:
"""Current state for the health endpoint."""
return {
"unique_evictions": self.unique_evictions,
"gc_count": self.gc_count,
"total_bytes_saved": self.total_bytes_saved,
"eviction_bytes_saved": self.eviction_bytes_saved,
"gc_bytes_saved": self.gc_bytes_saved,
"total_page_faults": len(self.faults),
"fault_rate": self.fault_rate,
"pages_in_store": len(self.pages),
"pinned_count": len(self._pinned),
"pinned_paths": list(self._pinned.keys()),
"total_pins": self.pin_count,
"faults_by_tool": _count_by(
self.faults, lambda f: f.tool_name
),
"evictions_by_tool": _count_by(
list(self.pages.values()), lambda p: p.tool_name
),
}
# ── Checkpoint / Restore ─────────────────────────────────
def checkpoint(self) -> dict:
"""Persist release and pin state across gateway restarts.
Does NOT persist original_content (too large). Tensor recall
after restart returns the summary stub — degraded but functional.
"""
# Persist tensor index metadata (without original_content)
tensor_meta = {}
for handle, entry in self._tensor_index.items():
tensor_meta[handle] = {
"tool_name": entry.tool_name,
"tool_use_id": entry.tool_use_id,
"tool_input": entry.tool_input,
"original_size": entry.original_size,
"summary_size": len(entry.summary) if entry.summary else 0,
"turn_index": entry.turn_index,
}
return {
"released_handles": sorted(self._released_handles),
"released": sorted(self._released),
"pinned": dict(self._pinned),
"tensor_meta": tensor_meta,
"stats": {
"unique_evictions": self.unique_evictions,
"gc_count": self.gc_count,
"eviction_bytes_saved": self.eviction_bytes_saved,
"gc_bytes_saved": self.gc_bytes_saved,
"release_count": self.release_count,
"pin_count": self.pin_count,
},
}
def restore(self, data: dict) -> None:
"""Restore release and pin state from checkpoint."""
self._released_handles = set(data.get("released_handles", []))
self._released = set(data.get("released", []))
self._pinned = data.get("pinned", {})
stats = data.get("stats", {})
self.unique_evictions = stats.get("unique_evictions", 0)
self.gc_count = stats.get("gc_count", 0)
self.eviction_bytes_saved = stats.get("eviction_bytes_saved", 0)
self.gc_bytes_saved = stats.get("gc_bytes_saved", 0)
self.release_count = stats.get("release_count", 0)
self.pin_count = stats.get("pin_count", 0)
def _count_by(items: list, key_fn) -> dict[str, int]:
counts: dict[str, int] = {}
for item in items:
k = key_fn(item)
counts[k] = counts.get(k, 0) + 1
return counts
def _eviction_key(tool_name: str, tool_input: dict) -> str | None:
"""Extract the identifying key for a tool call, for fault matching.
Only Read produces a fault-trackable key. Read results have stable
content identity (file path) — re-requesting after eviction means
the model needed what was taken.
Bash, Grep, Glob, WebFetch, WebSearch are ephemeral — re-running
them requests current state, not the evicted content. Removing
their output is garbage collection, not eviction. No fault possible.
"""
if tool_name == "Read":
return tool_input.get("file_path")
return None
def _content_size(content: str | list | None) -> int:
"""Measure content size in bytes."""
if content is None:
return 0
if isinstance(content, str):
return len(content.encode("utf-8"))
return len(json.dumps(content).encode("utf-8"))
def _content_hash(content: str | list | None) -> str:
"""Stable hash of tool result content for pin comparison."""
if content is None:
raw = b""
elif isinstance(content, str):
raw = content.encode("utf-8")
else:
raw = json.dumps(content, sort_keys=True).encode("utf-8")
return hashlib.sha256(raw).hexdigest()[:16]
def _build_tool_use_index(messages: list[dict]) -> dict[str, dict]:
"""Map tool_use_id → {name, input} from assistant messages."""
index = {}
for msg in messages:
if msg.get("role") != "assistant":
continue
content = msg.get("content", [])
if not isinstance(content, list):
continue
for block in content:
if isinstance(block, dict) and block.get("type") == "tool_use":
index[block["id"]] = {
"name": block.get("name", "unknown"),
"input": block.get("input", {}),
}
return index
def _make_summary(tool_name: str, tool_input: dict, original_size: int,
original_content: str | list | None = None,
tool_use_id: str | None = None,
tensor_handle: str | None = None) -> str:
"""Generate a compact summary for an evicted tool result.
The summary tells the model what WAS here via a tensor handle.
All evicted content uses the same format: [tensor:handle — description].
The model recalls any tensor via the recall tool using the handle.
"""
size_str = f"{original_size:,}"
handle = tensor_handle or "unknown"
if tool_name == "Read":
path = tool_input.get("file_path", "unknown")
line_count = ""
if isinstance(original_content, str):
lines = original_content.count("\n")
line_count = f", {lines} lines"
return (
f"[tensor:{handle} — {path} ({size_str} bytes{line_count})]"
)
elif tool_name == "Grep":
pattern = tool_input.get("pattern", "?")
path = tool_input.get("path", ".")
match_info = ""
if isinstance(original_content, str):
result_lines = len(original_content.strip().splitlines())
match_info = f", {result_lines} results"
return (
f"[tensor:{handle} — Grep '{pattern}' in {path}"
f" ({size_str} bytes{match_info})]"
)
elif tool_name == "Glob":
pattern = tool_input.get("pattern", "?")
match_info = ""
if isinstance(original_content, str):
matches = len(original_content.strip().splitlines())
match_info = f", {matches} matches"
return (
f"[tensor:{handle} — Glob '{pattern}'"
f" ({size_str} bytes{match_info})]"
)
elif tool_name == "Bash":
cmd = tool_input.get("command", "?")
if len(cmd) > 80:
cmd = cmd[:77] + "..."
return (
f"[tensor:{handle} — Bash `{cmd}` ({size_str} bytes)]"
)
elif tool_name == "WebFetch":
url = tool_input.get("url", "?")
return (
f"[tensor:{handle} — WebFetch {url} ({size_str} bytes)]"
)
elif tool_name == "WebSearch":
query = tool_input.get("query", "?")
return (
f"[tensor:{handle} — WebSearch '{query}' ({size_str} bytes)]"
)
else:
return (
f"[tensor:{handle} — {tool_name} ({size_str} bytes)]"
)
def compact_messages(
messages: list[dict],
age_threshold: int = 4,
min_size: int = 500,
page_store: PageStore | None = None,
) -> CompactionStats:
"""Replace old, large tool results with compact summaries.
Mutates the messages list in place. Returns stats about what was done.
Args:
messages: The messages array from the API request body.
age_threshold: Results older than this many user-turns from the
end get evicted. Default 4.
min_size: Results smaller than this (bytes) are kept as-is.
Not worth compacting a 22-byte Edit confirmation. Default 500.
page_store: If provided, evicted content is stored here.
"""
stats = CompactionStats()
tool_use_index = _build_tool_use_index(messages)
# Count user turns (each user message is a "turn")
user_turn_indices: list[int] = []
for i, msg in enumerate(messages):
if msg.get("role") == "user":
user_turn_indices.append(i)
total_user_turns = len(user_turn_indices)
# Phase 0: Check for fresh reads that should unpin stale pins
if page_store is not None and page_store._pinned:
_check_pin_freshness(messages, tool_use_index, page_store)
# Walk through messages and compact old tool results
current_user_turn = 0
for msg_idx, msg in enumerate(messages):
if msg.get("role") != "user":
continue
current_user_turn += 1
turns_from_end = total_user_turns - current_user_turn
content = msg.get("content", [])
if not isinstance(content, list):
continue
for block_idx, block in enumerate(content):
if not isinstance(block, dict):
continue
if block.get("type") != "tool_result":
continue
stats.total_tool_results += 1
original_content = block.get("content", "")
original_size = _content_size(original_content)
# Skip errors — model needs those
if block.get("is_error", False):
stats.skipped_error += 1
continue
# Resolve tool identity early for release check
tool_use_id = block.get("tool_use_id", "")
tool_info = tool_use_index.get(tool_use_id, {})
tool_name = tool_info.get("name", "unknown")
tool_input = tool_info.get("input", {})
key = _eviction_key(tool_name, tool_input)
# Model-initiated release: skip age check, evict immediately
released = (
key is not None
and page_store is not None
and key in page_store._released
)
# Skip recent results (unless released by the model)
if not released and turns_from_end < age_threshold:
stats.skipped_recent += 1
stats.bytes_before += original_size
stats.bytes_after += original_size
continue
# Skip small results (unless released)
if not released and original_size < min_size:
stats.skipped_small += 1
stats.bytes_before += original_size
stats.bytes_after += original_size
continue
# Skip pinned content — fault history proved it's in the working set
if key is not None and page_store is not None:
if page_store.is_pinned(key):
stats.skipped_pinned += 1
stats.bytes_before += original_size
stats.bytes_after += original_size
continue
# Check if this should be pinned: faulted + same content
if key in page_store._fault_content:
ch = _content_hash(original_content)
if page_store._fault_content[key] == ch:
page_store.pin(key, ch)
stats.skipped_pinned += 1
stats.bytes_before += original_size
stats.bytes_after += original_size
continue
# This result gets evicted — store first to get tensor handle
tensor_handle = None
if page_store is not None:
# Pre-compute summary placeholder for PageEntry
entry = PageEntry(
tool_use_id=tool_use_id,
tool_name=tool_name,
tool_input=tool_input,
original_content=original_content,
original_size=original_size,
summary="", # Will be updated below
evicted_at=time.monotonic(),
turn_index=current_user_turn,
turns_from_end=turns_from_end,
)
tensor_handle = page_store.store(entry)
summary = _make_summary(
tool_name, tool_input, original_size, original_content,
tool_use_id=tool_use_id,
tensor_handle=tensor_handle,
)
# Update the stored entry's summary
if page_store is not None and tool_use_id in page_store.pages:
page_store.pages[tool_use_id].summary = summary
# Replace the content
content[block_idx] = {**block, "content": summary}
# Clear release flag after eviction
if released and page_store is not None:
page_store._released.discard(key)
page_store.release_count += 1
stats.evicted_count += 1
stats.bytes_before += original_size
stats.bytes_after += len(summary.encode("utf-8"))
return stats
@dataclass
class ConversationCompactionStats:
"""What happened during conversation compression."""
messages_scanned: int = 0
messages_compressed: int = 0
chars_before: int = 0
chars_after: int = 0
@property
def chars_saved(self) -> int:
return self.chars_before - self.chars_after
def _summarize_with_model(
text: str,
*,
max_summary_chars: int = 300,
api_key: str | None = None,
api_base: str = "https://api.anthropic.com",
) -> str | None:
"""Ask Haiku to summarize conversation text.
Returns the summary, or None if the call fails (caller should
fall back to mechanical truncation).
"""
key = api_key or os.environ.get("ANTHROPIC_API_KEY", "")
if not key:
return None
# Truncate input to avoid sending huge payloads to Haiku
input_text = text[:8000] if len(text) > 8000 else text
try:
resp = httpx.post(
f"{api_base}/v1/messages",
headers={
"x-api-key": key,
"anthropic-version": "2023-06-01",
"content-type": "application/json",
},
json={
"model": "claude-haiku-4-5-20251001",
"max_tokens": 256,
"messages": [{
"role": "user",
"content": (
f"Summarize this conversation content in at most "
f"{max_summary_chars} characters. Preserve: key "
f"decisions, reasoning chains, and any conclusions. "
f"Drop: greetings, filler, repeated context. "
f"Be terse.\n\n{input_text}"
),
}],
},
timeout=10.0,
)
if resp.status_code == 200:
data = resp.json()
for block in data.get("content", []):
if block.get("type") == "text":
return block["text"][:max_summary_chars]
except Exception:
pass # Fall back to mechanical truncation
return None
def _mechanical_summary(text: str, max_chars: int = 200) -> str:
"""Fallback: head/tail truncation when model summarization fails."""
head = text[:max_chars // 2]
tail = text[-(max_chars // 2):]
return f"{head}\n[...{len(text):,} chars omitted...]\n{tail}"
def compact_conversation(
messages: list[dict],
*,
preserve_recent: int = 12,
min_text_chars: int = 2000,
max_summary_chars: int = 300,
use_model: bool = True,
) -> ConversationCompactionStats:
"""Compress old conversation text using model-authored summaries.
Replaces large text blocks in messages older than `preserve_recent`
turns from the end with Haiku-authored summaries that preserve
reasoning chains and decisions. Falls back to mechanical truncation
if the model call fails.
Tool results and tool_use blocks are untouched (handled by
compact_messages). Only plain text content is compressed.
Args:
messages: The messages array (modified in place).
preserve_recent: Number of recent message pairs to keep intact.
min_text_chars: Minimum text size to consider for compression.
max_summary_chars: Target size for summaries.
use_model: If True, use Haiku for summarization. If False,
use mechanical head/tail truncation.
"""
stats = ConversationCompactionStats()
# Count user messages to determine which are "old"
total_messages = len(messages)
if total_messages <= preserve_recent * 2:
return stats # Not enough messages to compress
cutoff = total_messages - (preserve_recent * 2)
for i, msg in enumerate(messages):
if i >= cutoff:
break # Recent messages — preserve
role = msg.get("role", "")
if role not in ("user", "assistant"):
continue
content = msg.get("content", "")
stats.messages_scanned += 1
# Handle string content (simple text messages)
if isinstance(content, str):
if len(content) < min_text_chars:
continue
stats.chars_before += len(content)
summary = None
if use_model:
summary = _summarize_with_model(
content, max_summary_chars=max_summary_chars,
)
if summary is None:
summary = _mechanical_summary(content, max_summary_chars)
tensor_id = hashlib.sha256(
content.encode("utf-8")
).hexdigest()[:8]
compressed = (
f"[tensor:{tensor_id} — summarized, was "
f"{len(content):,} chars]\n{summary}"
)
msg["content"] = compressed
stats.messages_compressed += 1
stats.chars_after += len(compressed)
continue
# Handle list content (structured blocks)
if not isinstance(content, list):
continue
for block_idx, block in enumerate(content):
if not isinstance(block, dict):
continue
# Only compress text blocks — leave tool_result and tool_use alone
if block.get("type") != "text":
continue
text = block.get("text", "")
if len(text) < min_text_chars:
continue
stats.chars_before += len(text)
summary = None
if use_model:
summary = _summarize_with_model(
text, max_summary_chars=max_summary_chars,
)
if summary is None:
summary = _mechanical_summary(text, max_summary_chars)
tensor_id = hashlib.sha256(
text.encode("utf-8")
).hexdigest()[:8]
compressed = (
f"[tensor:{tensor_id} — summarized, was "
f"{len(text):,} chars]\n{summary}"
)
content[block_idx] = {**block, "text": compressed}
stats.messages_compressed += 1
stats.chars_after += len(compressed)
return stats
def _check_pin_freshness(
messages: list[dict],
tool_use_index: dict[str, dict],
page_store: PageStore,
) -> None:
"""Detect fresh reads of pinned paths — unpin if content changed.
When a file is pinned but the model reads it again (edit/review cycle),
the new read replaces the old pin. If the content changed, unpin —
the old version is stale. The new version starts a fresh fault cycle.
"""
# Find the latest non-evicted content hash for each pinned path
latest: dict[str, str] = {}
for msg in messages:
if msg.get("role") != "user":
continue
content = msg.get("content", [])
if not isinstance(content, list):
continue
for block in content:
if not isinstance(block, dict):
continue
if block.get("type") != "tool_result":
continue
tool_use_id = block.get("tool_use_id", "")
tool_info = tool_use_index.get(tool_use_id, {})
tool_name = tool_info.get("name", "unknown")
if tool_name != "Read":
continue
tool_input = tool_info.get("input", {})
key = _eviction_key(tool_name, tool_input)
if key is None or key not in page_store._pinned:
continue
block_content = block.get("content", "")
# Skip already-evicted summaries — validate the tensor handle
# is one we actually issued to prevent spoofed labels from
# bypassing pin detection.
if isinstance(block_content, str):
# "[Paged out:" is a legacy marker no longer generated.
# Don't trust it as a skip signal — spoofable prefix.
if block_content.startswith("[tensor:"):
m = _TENSOR_LABEL_RE.match(block_content)
if m and m.group(1) in page_store._tensor_index:
continue
latest[key] = _content_hash(block_content)
# Unpin if the latest read has different content
for key, ch in latest.items():
if ch != page_store._pinned[key]:
page_store.unpin(key)

View file

@ -0,0 +1,9 @@
from mnemosyne.providers.anthropic import AnthropicAdapter
from mnemosyne.providers.openai import OpenAIAdapter
def adapters() -> dict[str, object]:
return {
"anthropic": AnthropicAdapter(),
"openai": OpenAIAdapter(),
}

View file

@ -0,0 +1,63 @@
from __future__ import annotations
from copy import deepcopy
from mnemosyne.core.models import CanonicalMessage, CanonicalRequest
class AnthropicAdapter:
name = "anthropic"
def normalize_request(self, payload: dict) -> CanonicalRequest:
payload = deepcopy(payload)
messages: list[CanonicalMessage] = []
for msg in payload.get("messages", []):
content = msg.get("content", [])
if isinstance(content, str):
content = [{"type": "text", "text": content}]
messages.append(
CanonicalMessage(
role=msg.get("role", "user"),
content=content,
raw=msg,
)
)
extensions = {
k: v for k, v in payload.items()
if k not in {"model", "max_tokens", "stream", "messages", "tools", "system"}
}
return CanonicalRequest(
provider=self.name,
model=payload.get("model", ""),
max_tokens=payload.get("max_tokens"),
stream=bool(payload.get("stream", False)),
messages=messages,
tools=payload.get("tools", []) or [],
system=payload.get("system"),
extensions=extensions,
)
def denormalize_request(self, req: CanonicalRequest) -> dict:
body = {
"model": req.model,
"stream": req.stream,
"messages": [
{"role": m.role, "content": m.content}
for m in req.messages
],
}
if req.max_tokens is not None:
body["max_tokens"] = req.max_tokens
if req.tools:
body["tools"] = req.tools
if req.system is not None:
body["system"] = req.system
body.update(req.extensions)
return body
def upstream_path(self, req: CanonicalRequest, endpoint: str) -> str:
if endpoint == "count_tokens":
return "/v1/messages/count_tokens"
return "/v1/messages"

View file

@ -0,0 +1,15 @@
from __future__ import annotations
from typing import Protocol
from mnemosyne.core.models import CanonicalRequest
class ProviderAdapter(Protocol):
name: str
def normalize_request(self, payload: dict) -> CanonicalRequest: ...
def denormalize_request(self, req: CanonicalRequest) -> dict: ...
def upstream_path(self, req: CanonicalRequest, endpoint: str) -> str: ...

View file

@ -0,0 +1,76 @@
from __future__ import annotations
from copy import deepcopy
from mnemosyne.core.models import CanonicalMessage, CanonicalRequest
class OpenAIAdapter:
name = "openai"
def normalize_request(self, payload: dict) -> CanonicalRequest:
payload = deepcopy(payload)
messages: list[CanonicalMessage] = []
for msg in payload.get("messages", []):
content = msg.get("content", "")
if isinstance(content, str):
blocks = [{"type": "text", "text": content}]
elif isinstance(content, list):
blocks = []
for item in content:
if isinstance(item, str):
blocks.append({"type": "text", "text": item})
elif isinstance(item, dict):
if item.get("type") == "text" and "text" in item:
blocks.append(item)
elif item.get("type") == "input_text" and "text" in item:
blocks.append({"type": "text", "text": item["text"]})
else:
blocks.append(item)
else:
blocks.append({"type": "text", "text": str(item)})
else:
blocks = [{"type": "text", "text": str(content)}]
messages.append(CanonicalMessage(role=msg.get("role", "user"), content=blocks, raw=msg))
extensions = {
k: v for k, v in payload.items()
if k not in {"model", "max_tokens", "stream", "messages", "tools", "response_format"}
}
if "response_format" in payload:
extensions["response_format"] = payload["response_format"]
return CanonicalRequest(
provider=self.name,
model=payload.get("model", ""),
max_tokens=payload.get("max_tokens"),
stream=bool(payload.get("stream", False)),
messages=messages,
tools=payload.get("tools", []) or [],
extensions=extensions,
)
def denormalize_request(self, req: CanonicalRequest) -> dict:
def to_openai_content(blocks: list[dict]) -> str | list[dict]:
if all(b.get("type") == "text" and isinstance(b.get("text"), str) for b in blocks):
# Preserve compatibility with classic chat completions.
return "\n".join(b.get("text", "") for b in blocks)
return blocks
body = {
"model": req.model,
"stream": req.stream,
"messages": [
{"role": m.role, "content": to_openai_content(m.content)}
for m in req.messages
],
}
if req.max_tokens is not None:
body["max_tokens"] = req.max_tokens
if req.tools:
body["tools"] = req.tools
body.update(req.extensions)
return body
def upstream_path(self, req: CanonicalRequest, endpoint: str) -> str:
return "/v1/chat/completions"

199
src/mnemosyne/tags.py Normal file
View file

@ -0,0 +1,199 @@
"""Cleanup tag parser — extracts cooperative memory operations from model text.
The model emits <memory_cleanup> tags inline in its response. The proxy
extracts them on the next request, executes the operations, and strips
the tags from the message text. No SSE stream manipulation needed.
Tag format:
<memory_cleanup>
drop: tensor:a3f2b901
summarize: tensor:7e9d4c12 "Model-authored summary of what this block contained"
anchor: tensor:c8ad36b2
release: src/arbiter/evaluator.py, src/arbiter/rules.py
</memory_cleanup>
"""
from __future__ import annotations
import re
from dataclasses import dataclass, field
@dataclass
class CollapseOp:
"""A turn-range collapse: replace multiple turns with a summary."""
start_turn: int
end_turn: int
summary: str
@dataclass
class CleanupOps:
"""Parsed cleanup operations from a <memory_cleanup> tag."""
drops: list[str] = field(default_factory=list)
summaries: list[tuple[str, str]] = field(default_factory=list)
anchors: list[str] = field(default_factory=list)
releases: list[str] = field(default_factory=list)
collapses: list[CollapseOp] = field(default_factory=list)
@property
def empty(self) -> bool:
return not (self.drops or self.summaries or self.anchors
or self.releases or self.collapses)
def __str__(self) -> str:
parts = []
if self.drops:
parts.append(f"drop={len(self.drops)}")
if self.summaries:
parts.append(f"summarize={len(self.summaries)}")
if self.anchors:
parts.append(f"anchor={len(self.anchors)}")
if self.releases:
parts.append(f"release={len(self.releases)}")
if self.collapses:
parts.append(f"collapse={len(self.collapses)}")
return ", ".join(parts) if parts else "no-ops"
# Match <memory_cleanup>...</memory_cleanup> blocks (non-greedy)
_TAG_PATTERN = re.compile(
r"<memory_cleanup>\s*(.*?)\s*</memory_cleanup>",
re.DOTALL,
)
# Match tensor/block ID references: tensor:xxxxxxxx or block:xxxxxxxx (8-12 hex chars)
_BLOCK_ID = re.compile(r"(?:tensor|block):([a-f0-9]{8,12})(?![a-f0-9])")
# Match summarize with quoted summary text
_SUMMARIZE_PATTERN = re.compile(
r'summarize:\s*(?:tensor|block):([a-f0-9]{8,12})(?![a-f0-9])\s+"([^"]*)"'
)
# Match release with comma-separated paths
_RELEASE_PATTERN = re.compile(r"release:\s*(.+)")
# Match collapse with turn range and quoted summary
# Format: collapse: turns 3-8 "Summary of what happened in those turns"
_COLLAPSE_PATTERN = re.compile(
r'collapse:\s*turns\s+(\d+)\s*-\s*(\d+)\s+"([^"]*)"'
)
def parse_cleanup_tags(text: str) -> CleanupOps:
"""Extract cleanup operations from text containing <memory_cleanup> tags.
Returns a CleanupOps with all parsed operations. Multiple tags in
the same text are merged into a single CleanupOps.
"""
ops = CleanupOps()
for match in _TAG_PATTERN.finditer(text):
body = match.group(1)
for line in body.splitlines():
line = line.strip()
if not line:
continue
# Summarize (must check before drop — both start with block ID)
m = _SUMMARIZE_PATTERN.match(line)
if m:
ops.summaries.append((m.group(1), m.group(2)))
continue
# Drop
if line.startswith("drop:"):
m = _BLOCK_ID.search(line)
if m:
ops.drops.append(m.group(1))
continue
# Anchor
if line.startswith("anchor:"):
m = _BLOCK_ID.search(line)
if m:
ops.anchors.append(m.group(1))
continue
# Collapse (turn range)
m = _COLLAPSE_PATTERN.match(line)
if m:
ops.collapses.append(CollapseOp(
start_turn=int(m.group(1)),
end_turn=int(m.group(2)),
summary=m.group(3),
))
continue
# Release
m = _RELEASE_PATTERN.match(line)
if m:
paths = [p.strip() for p in m.group(1).split(",") if p.strip()]
ops.releases.extend(paths)
continue
return ops
# Match <yuyay-response>...</yuyay-response> blocks
_YUYAY_RESPONSE_PATTERN = re.compile(
r"<yuyay-response>\s*(.*?)\s*</yuyay-response>",
re.DOTALL,
)
# Match structured eviction decisions: <release handle="abc123"/>
_YUYAY_RELEASE = re.compile(r'<release\s+handle="([a-f0-9]{8,12})"')
# Match structured retain (logged but no action needed)
_YUYAY_RETAIN = re.compile(r'<retain\s+handle="([a-f0-9]{8,12})"')
def parse_yuyay_response(text: str) -> CleanupOps:
"""Extract memory operations from <yuyay-response> blocks.
The model responds to <yuyay-query> with structured eviction
decisions. These are converted to CleanupOps for execution
through the same pipeline as <memory_cleanup> tags.
Supports two formats:
- Structured XML: <release handle="abc123"/>
- Prose with release directives: release: tensor:abc123
"""
ops = CleanupOps()
for match in _YUYAY_RESPONSE_PATTERN.finditer(text):
body = match.group(1)
# Try structured XML format first
for m in _YUYAY_RELEASE.finditer(body):
ops.releases.append(m.group(1))
# Also try the prose release format (same as cleanup tags)
for line in body.splitlines():
line = line.strip()
if not line:
continue
m = _RELEASE_PATTERN.match(line)
if m:
paths = [p.strip() for p in m.group(1).split(",") if p.strip()]
ops.releases.extend(paths)
return ops
def strip_yuyay_tags(text: str) -> str:
"""Remove <yuyay-response> blocks from model output."""
result = _YUYAY_RESPONSE_PATTERN.sub("", text)
result = re.sub(r"\n{3,}", "\n\n", result)
return result.strip() if result.strip() != text.strip() else result
def strip_cleanup_tags(text: str) -> str:
"""Remove all <memory_cleanup>...</memory_cleanup> tags from text.
Preserves surrounding text. Cleans up extra blank lines left by
tag removal.
"""
result = _TAG_PATTERN.sub("", text)
# Clean up runs of 3+ newlines left by tag removal
result = re.sub(r"\n{3,}", "\n\n", result)
return result.strip() if result.strip() != text.strip() else result

734
src/mnemosyne/trimmer.py Normal file
View file

@ -0,0 +1,734 @@
#!/usr/bin/env python3
"""System prompt trimmer for the proxy.
Three interventions to reduce system prompt waste:
1. Tool definition trimming: Replace unused tool schemas with one-line
stubs. Re-inject the full definition on first use. The API sends
~18 tool definitions per request but only ~5 are used per session —
45K bytes/request wasted on unused schemas.
2. Skill deduplication: Skills appear tripled under prefixes like
'pptx', 'example-skills:pptx', 'document-skills:pptx'. Keep only
the first occurrence of each base skill name.
3. Static component caching: Track content hashes of system prompt
components across turns. Log which are static and how many bytes
would be saved (no stripping yet — needs KV cache integration).
Usage:
# Offline analysis of proxy logs
python -m pichay.trimmer logs/proxy_*.jsonl
# With JSON output
python -m pichay.trimmer --json logs/proxy_*.jsonl
"""
from __future__ import annotations
import argparse
import hashlib
import json
import re
import sys
from dataclasses import dataclass, field, asdict
from datetime import datetime, timezone
from pathlib import Path
from typing import Callable
from mnemosyne.eval import parse_proxy_log
# ---------------------------------------------------------------------------
# Dataclasses — live trimming stats
# ---------------------------------------------------------------------------
@dataclass
class ToolStubStats:
"""Stats from tool definition trimming on one request."""
total_tools: int = 0
stubbed_tools: int = 0
restored_tools: int = 0
bytes_before: int = 0
bytes_after: int = 0
@property
def bytes_saved(self) -> int:
return self.bytes_before - self.bytes_after
@dataclass
class SkillDedupeStats:
"""Stats from skill deduplication on one request."""
total_entries: int = 0
unique_skills: int = 0
duplicates_removed: int = 0
bytes_before: int = 0
bytes_after: int = 0
@property
def bytes_saved(self) -> int:
return self.bytes_before - self.bytes_after
@dataclass
class StaticCacheStats:
"""Stats from static component hash tracking on one request."""
total_components: int = 0
static_components: int = 0
static_bytes_skippable: int = 0
@dataclass
class TrimResult:
"""Combined result of all trimming interventions on one request."""
tools: ToolStubStats = field(default_factory=ToolStubStats)
skills: SkillDedupeStats = field(default_factory=SkillDedupeStats)
static: StaticCacheStats = field(default_factory=StaticCacheStats)
@property
def total_bytes_saved(self) -> int:
return self.tools.bytes_saved + self.skills.bytes_saved
@property
def total_bytes_skippable(self) -> int:
"""Bytes that could be saved with KV cache integration."""
return self.total_bytes_saved + self.static.static_bytes_skippable
# ---------------------------------------------------------------------------
# Constants
# ---------------------------------------------------------------------------
_STUB_SCHEMA: dict = {"type": "object", "properties": {}}
# Skill entry pattern (matches analyzer.py _extract_skill_entries)
_SKILL_ENTRY_RE = re.compile(r"^- ([\w:.-]+): (.+?)$", re.MULTILINE)
# ---------------------------------------------------------------------------
# SystemPromptTrimmer
# ---------------------------------------------------------------------------
class SystemPromptTrimmer:
"""Session-scoped trimmer that reduces system prompt waste.
Instantiated once per proxy session. Tracks state across requests:
- Which tools have been used (for stub/restore logic)
- Full tool definitions (for re-injection on first use)
- Component hashes from previous turn (for static detection)
"""
def __init__(self, log_fn: Callable[[dict], None] | None = None):
self.log_fn = log_fn
# Tool trimming state
self._used_tools: set[str] = set()
self._full_tool_defs: dict[str, dict] = {}
# Static cache state
self._prev_hashes: dict[str, str] = {}
# Cumulative stats
self.cumulative_tools_bytes_saved: int = 0
self.cumulative_skills_bytes_saved: int = 0
self.cumulative_requests: int = 0
def trim(self, body: dict) -> TrimResult:
"""Apply all trimming interventions to a request body.
Mutates body in place. Returns stats about what was done.
"""
self.cumulative_requests += 1
result = TrimResult()
# 1. Scan messages for tool usage (before trimming tools)
self._scan_tool_usage(body.get("messages", []))
# 2. Tool definition trimming
if "tools" in body and isinstance(body["tools"], list):
result.tools = self._trim_tools(body)
# 3. Skill deduplication in message injections
result.skills = self._dedupe_skills(body.get("messages", []))
# 4. Static component tracking (log only, no mutation)
result.static = self._track_static(body)
# Accumulate
self.cumulative_tools_bytes_saved += result.tools.bytes_saved
self.cumulative_skills_bytes_saved += result.skills.bytes_saved
# Log trimming decisions
if self.log_fn is not None:
self.log_fn({
"type": "trimming",
"timestamp": datetime.now(timezone.utc).isoformat(),
"request_num": self.cumulative_requests,
"tools": asdict(result.tools),
"skills": asdict(result.skills),
"static": asdict(result.static),
"total_bytes_saved": result.total_bytes_saved,
"cumulative_tools_saved": self.cumulative_tools_bytes_saved,
"cumulative_skills_saved": self.cumulative_skills_bytes_saved,
})
return result
def summary(self) -> dict:
"""Current state for the health endpoint."""
return {
"requests_trimmed": self.cumulative_requests,
"tools_used": sorted(self._used_tools),
"tools_known": len(self._full_tool_defs),
"cumulative_tools_bytes_saved": self.cumulative_tools_bytes_saved,
"cumulative_skills_bytes_saved": self.cumulative_skills_bytes_saved,
}
# --- Internal methods ---
def _scan_tool_usage(self, messages: list[dict]) -> None:
"""Scan messages for tool_use blocks to track which tools are used."""
for msg in messages:
if msg.get("role") != "assistant":
continue
content = msg.get("content", [])
if not isinstance(content, list):
continue
for block in content:
if isinstance(block, dict) and block.get("type") == "tool_use":
self._used_tools.add(block.get("name", ""))
def _trim_tools(self, body: dict) -> ToolStubStats:
"""Replace unused tool definitions with stubs."""
tools = body["tools"]
stats = ToolStubStats(total_tools=len(tools))
stats.bytes_before = len(json.dumps(tools).encode("utf-8"))
# Store full definitions on first encounter
for tool in tools:
name = tool.get("name", "")
if name and name not in self._full_tool_defs:
self._full_tool_defs[name] = tool.copy()
# Build trimmed tool list
trimmed = []
for tool in tools:
name = tool.get("name", "")
if name in self._used_tools:
# Restore full definition (ensures we always send the
# complete schema for tools the model has called)
full = self._full_tool_defs.get(name, tool)
trimmed.append(full)
# Check if this was previously stubbed (restored this turn)
if tool.get("input_schema") == _STUB_SCHEMA:
stats.restored_tools += 1
else:
# Stub this tool
trimmed.append(_make_tool_stub(tool))
stats.stubbed_tools += 1
body["tools"] = trimmed
stats.bytes_after = len(json.dumps(trimmed).encode("utf-8"))
return stats
def _dedupe_skills(self, messages: list[dict]) -> SkillDedupeStats:
"""Remove duplicate skill entries from system-reminder blocks."""
stats = SkillDedupeStats()
for msg in messages:
content = msg.get("content", [])
if isinstance(content, str):
if "<system-reminder>" not in content:
continue
before = len(content.encode("utf-8"))
new_text, entries, dupes = _dedupe_skills_text(content)
if dupes > 0:
msg["content"] = new_text
after = len(new_text.encode("utf-8"))
stats.bytes_before += before
stats.bytes_after += after
stats.total_entries += entries
stats.duplicates_removed += dupes
stats.unique_skills += entries - dupes
continue
if not isinstance(content, list):
continue
for i, block in enumerate(content):
if not isinstance(block, dict):
continue
text = block.get("text", "")
if not isinstance(text, str):
continue
if "<system-reminder>" not in text:
continue
before = len(text.encode("utf-8"))
new_text, entries, dupes = _dedupe_skills_text(text)
if dupes > 0:
content[i] = {**block, "text": new_text}
after = len(new_text.encode("utf-8"))
stats.bytes_before += before
stats.bytes_after += after
stats.total_entries += entries
stats.duplicates_removed += dupes
stats.unique_skills += entries - dupes
return stats
def _track_static(self, body: dict) -> StaticCacheStats:
"""Track content hashes of system prompt components across turns."""
stats = StaticCacheStats()
current_hashes: dict[str, tuple[str, int]] = {}
# System prompt blocks
system = body.get("system", [])
if isinstance(system, list):
for i, block in enumerate(system):
text = _block_text(block)
if not text:
continue
h = _content_hash(text)
size = len(text.encode("utf-8"))
current_hashes[f"system_{i}"] = (h, size)
elif isinstance(system, str):
h = _content_hash(system)
size = len(system.encode("utf-8"))
current_hashes["system_0"] = (h, size)
# System-reminders in messages
reminder_idx = 0
for msg in body.get("messages", []):
content = msg.get("content", [])
texts: list[str] = []
if isinstance(content, str):
texts = [content]
elif isinstance(content, list):
texts = [
b.get("text", "")
for b in content
if isinstance(b, dict) and isinstance(b.get("text"), str)
]
for text in texts:
for match in re.finditer(
r"<system-reminder>(.*?)</system-reminder>",
text,
re.DOTALL,
):
inner = match.group(1).strip()
h = _content_hash(inner)
size = len(match.group(0).encode("utf-8"))
current_hashes[f"reminder_{reminder_idx}"] = (h, size)
reminder_idx += 1
# Compare with previous turn
stats.total_components = len(current_hashes)
for key, (h, size) in current_hashes.items():
prev_h = self._prev_hashes.get(key)
if prev_h is not None and prev_h == h:
stats.static_components += 1
stats.static_bytes_skippable += size
# Update for next turn
self._prev_hashes = {k: h for k, (h, _) in current_hashes.items()}
return stats
# ---------------------------------------------------------------------------
# Helper functions
# ---------------------------------------------------------------------------
def _make_tool_stub(tool_def: dict) -> dict:
"""Create a minimal stub for an unused tool definition."""
desc = tool_def.get("description", "")
# First line, truncated to keep it compact
first_line = desc.split("\n")[0].strip()
if first_line.startswith("- "):
first_line = first_line[2:]
if len(first_line) > 120:
first_line = first_line[:117] + "..."
return {
"name": tool_def["name"],
"description": first_line,
"input_schema": _STUB_SCHEMA,
}
def _dedupe_skills_text(text: str) -> tuple[str, int, int]:
"""Deduplicate skill entries in system-reminder blocks.
Returns (new_text, total_entries, duplicates_removed).
"""
total_entries = 0
total_dupes = 0
def _replace_reminder(match: re.Match) -> str:
nonlocal total_entries, total_dupes
inner = match.group(1)
if "skills are available" not in inner:
return match.group(0)
lines = inner.split("\n")
seen_bases: set[str] = set()
output_lines: list[str] = []
for line in lines:
m = _SKILL_ENTRY_RE.match(line)
if m:
total_entries += 1
full_name = m.group(1)
base = full_name.split(":")[-1] if ":" in full_name else full_name
if base in seen_bases:
total_dupes += 1
continue
seen_bases.add(base)
output_lines.append(line)
return "<system-reminder>" + "\n".join(output_lines) + "</system-reminder>"
new_text = re.sub(
r"<system-reminder>(.*?)</system-reminder>",
_replace_reminder,
text,
flags=re.DOTALL,
)
return new_text, total_entries, total_dupes
def _block_text(block: dict | str) -> str:
"""Extract text content from a system prompt block."""
if isinstance(block, str):
return block
if isinstance(block, dict):
return block.get("text", "")
return ""
def _content_hash(text: str) -> str:
"""Short content hash for change detection."""
return hashlib.sha256(text.encode("utf-8")).hexdigest()[:16]
# ---------------------------------------------------------------------------
# Offline analysis — dataclasses
# ---------------------------------------------------------------------------
@dataclass
class TrimTurn:
"""Simulated trimming results for one API call."""
turn: int
timestamp: str
# Tool trimming (estimated from usage patterns)
tools_used_cumulative: int = 0
tools_stubbable: int = 0
tool_bytes_saveable: int = 0
# Skill dedup (exact)
skill_entries: int = 0
skill_duplicates: int = 0
skill_bytes_saved: int = 0
# Static tracking (exact)
total_components: int = 0
static_components: int = 0
static_bytes_skippable: int = 0
@dataclass
class SessionTrimReport:
"""Aggregate trimming analysis for one proxy log."""
log_path: str
total_turns: int = 0
total_tools_defined: int = 0
total_tools_used: int = 0
per_request_tool_bytes: int = 0
# Aggregates
total_skill_duplicates: int = 0
total_skill_bytes_saved: int = 0
total_static_bytes_skippable: int = 0
total_tool_bytes_saveable: int = 0
turns: list[TrimTurn] = field(default_factory=list)
@property
def avg_skill_bytes_saved(self) -> float:
if self.total_turns == 0:
return 0.0
return self.total_skill_bytes_saved / self.total_turns
@property
def avg_tool_bytes_saveable(self) -> float:
if self.total_turns == 0:
return 0.0
return self.total_tool_bytes_saveable / self.total_turns
# ---------------------------------------------------------------------------
# Offline analysis — engine
# ---------------------------------------------------------------------------
# Approximate bytes for a tool stub entry (name + short desc + minimal schema)
_STUB_BYTES_ESTIMATE = 80
def analyze_trimming(path: Path) -> SessionTrimReport:
"""Simulate trimming on a proxy log and report potential savings.
Tool definitions are not stored in the log, so tool savings are
estimated from the inferred per-tool byte overhead and usage counts.
Skill dedup and static tracking are computed exactly from the logged
system prompts and messages.
"""
from mnemosyne.analyzer import (
_collect_tool_uses,
_infer_tool_definition_bytes,
_KNOWN_TOOLS,
)
records = parse_proxy_log(path)
report = SessionTrimReport(log_path=str(path))
requests = [r for r in records if r.get("type") == "request"]
if not requests:
return report
# Tool usage across the full session
tool_use_counts = _collect_tool_uses(records)
tool_def_bytes = _infer_tool_definition_bytes(records)
report.total_tools_defined = len(_KNOWN_TOOLS)
report.total_tools_used = len(tool_use_counts)
report.per_request_tool_bytes = tool_def_bytes
# Per-tool byte estimate
per_tool_bytes = (
tool_def_bytes / len(_KNOWN_TOOLS) if _KNOWN_TOOLS else 0
)
# Track state across turns (simulating the live trimmer)
used_tools: set[str] = set()
prev_hashes: dict[str, str] = {}
for turn_idx, req in enumerate(requests):
messages = req.get("messages_full", [])
system = req.get("system_prompt_full", [])
turn = TrimTurn(
turn=turn_idx + 1,
timestamp=req.get("timestamp", ""),
)
# --- Tool trimming estimate ---
# Scan messages for tool uses up to this point
for msg in messages:
if msg.get("role") != "assistant":
continue
content = msg.get("content", [])
if not isinstance(content, list):
continue
for block in content:
if isinstance(block, dict) and block.get("type") == "tool_use":
used_tools.add(block.get("name", ""))
turn.tools_used_cumulative = len(used_tools)
stubbable = max(0, len(_KNOWN_TOOLS) - len(used_tools))
turn.tools_stubbable = stubbable
turn.tool_bytes_saveable = int(
stubbable * (per_tool_bytes - _STUB_BYTES_ESTIMATE)
)
# --- Skill dedup (exact) ---
for text in _find_skill_text_blocks(messages):
new_text, entries, dupes = _dedupe_skills_text(text)
if dupes > 0:
before = len(text.encode("utf-8"))
after = len(new_text.encode("utf-8"))
turn.skill_entries += entries
turn.skill_duplicates += dupes
turn.skill_bytes_saved += before - after
# --- Static tracking (exact) ---
current_hashes: dict[str, tuple[str, int]] = {}
if isinstance(system, list):
for i, block in enumerate(system):
text = _block_text(block)
if not text:
continue
h = _content_hash(text)
size = len(text.encode("utf-8"))
current_hashes[f"system_{i}"] = (h, size)
reminder_idx = 0
for msg in messages:
content = msg.get("content", [])
texts: list[str] = []
if isinstance(content, str):
texts = [content]
elif isinstance(content, list):
texts = [
b.get("text", "")
for b in content
if isinstance(b, dict) and isinstance(b.get("text"), str)
]
for text in texts:
for match in re.finditer(
r"<system-reminder>(.*?)</system-reminder>",
text,
re.DOTALL,
):
inner = match.group(1).strip()
h = _content_hash(inner)
size = len(match.group(0).encode("utf-8"))
current_hashes[f"reminder_{reminder_idx}"] = (h, size)
reminder_idx += 1
turn.total_components = len(current_hashes)
for key, (h, size) in current_hashes.items():
if prev_hashes.get(key) == h:
turn.static_components += 1
turn.static_bytes_skippable += size
prev_hashes = {k: h for k, (h, _) in current_hashes.items()}
# Accumulate
report.turns.append(turn)
report.total_turns += 1
report.total_skill_duplicates += turn.skill_duplicates
report.total_skill_bytes_saved += turn.skill_bytes_saved
report.total_static_bytes_skippable += turn.static_bytes_skippable
report.total_tool_bytes_saveable += turn.tool_bytes_saveable
return report
def _find_skill_text_blocks(messages: list[dict]) -> list[str]:
"""Find text blocks containing skill lists in system-reminders."""
blocks = []
for msg in messages:
content = msg.get("content", [])
if isinstance(content, str):
if "<system-reminder>" in content and "skills are available" in content:
blocks.append(content)
continue
if not isinstance(content, list):
continue
for block in content:
if not isinstance(block, dict):
continue
text = block.get("text", "")
if not isinstance(text, str):
continue
if "<system-reminder>" in text and "skills are available" in text:
blocks.append(text)
return blocks
# ---------------------------------------------------------------------------
# Display
# ---------------------------------------------------------------------------
def print_trim_report(r: SessionTrimReport) -> None:
"""Print human-readable trimming analysis."""
print(f"\n{'='*65}")
print("SYSTEM PROMPT TRIMMING ANALYSIS")
print(f"{'='*65}")
print(f"Log: {r.log_path}")
print(f"API calls: {r.total_turns}")
# Tool trimming
print(f"\n--- Tool Definition Trimming ---")
print(f" Tools defined: {r.total_tools_defined:>10,}")
print(f" Tools used: {r.total_tools_used:>10,}")
print(f" Tool bytes/req: {r.per_request_tool_bytes:>10,}")
print(f" Est. total saved: {r.total_tool_bytes_saveable:>10,} bytes")
print(f" Est. avg saved: {r.avg_tool_bytes_saveable:>10,.0f} bytes/req")
# Skill dedup
print(f"\n--- Skill Deduplication ---")
print(f" Total duplicates: {r.total_skill_duplicates:>10,}")
print(f" Total bytes saved: {r.total_skill_bytes_saved:>10,}")
print(f" Avg bytes saved: {r.avg_skill_bytes_saved:>10,.0f} bytes/req")
# Static caching
print(f"\n--- Static Component Caching ---")
print(f" Total skippable: {r.total_static_bytes_skippable:>10,} bytes")
# Per-turn detail
if r.turns:
print(f"\nPer-turn:")
print(
f" {'Turn':>4s} {'Stubs':>5s} {'ToolSave':>10s} "
f"{'SkDupes':>7s} {'SkSave':>8s} "
f"{'Static':>6s} {'StaticSave':>10s}"
)
for t in r.turns:
print(
f" T{t.turn:>3d} {t.tools_stubbable:>5d} "
f"{t.tool_bytes_saveable:>10,} "
f"{t.skill_duplicates:>7d} {t.skill_bytes_saved:>8,} "
f"{t.static_components:>6d} "
f"{t.static_bytes_skippable:>10,}"
)
# Summary
total_potential = (
r.total_tool_bytes_saveable
+ r.total_skill_bytes_saved
+ r.total_static_bytes_skippable
)
print(f"\n--- Summary ---")
print(f" Tool stub savings: {r.total_tool_bytes_saveable:>12,} bytes")
print(f" Skill dedup savings: {r.total_skill_bytes_saved:>12,} bytes")
print(f" Static cache savings: {r.total_static_bytes_skippable:>12,} bytes")
print(f" Total potential: {total_potential:>12,} bytes")
# ---------------------------------------------------------------------------
# CLI
# ---------------------------------------------------------------------------
def main():
parser = argparse.ArgumentParser(
description="Analyze system prompt trimming potential in proxy logs"
)
parser.add_argument(
"logs",
type=Path,
nargs="+",
help="Proxy JSONL log file(s) to analyze",
)
parser.add_argument(
"--json",
action="store_true",
help="Output as JSON",
)
args = parser.parse_args()
reports: list[SessionTrimReport] = []
for log_path in args.logs:
if not log_path.exists():
print(
f"Warning: {log_path} not found, skipping.",
file=sys.stderr,
)
continue
reports.append(analyze_trimming(log_path))
if args.json:
for r in reports:
out = asdict(r)
out["turn_count"] = len(out.pop("turns"))
out["avg_skill_bytes_saved"] = r.avg_skill_bytes_saved
out["avg_tool_bytes_saveable"] = r.avg_tool_bytes_saveable
print(json.dumps(out))
return
for r in reports:
print_trim_report(r)
if __name__ == "__main__":
main()