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:
parent
ed0361f97c
commit
974863e7b3
21 changed files with 6591 additions and 0 deletions
6
src/mnemosyne/__init__.py
Normal file
6
src/mnemosyne/__init__.py
Normal 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
307
src/mnemosyne/__main__.py
Normal 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
377
src/mnemosyne/blocks.py
Normal 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
176
src/mnemosyne/config.py
Normal 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),
|
||||
)
|
||||
1
src/mnemosyne/core/__init__.py
Normal file
1
src/mnemosyne/core/__init__.py
Normal file
|
|
@ -0,0 +1 @@
|
|||
"""Core gateway abstractions: canonical model + policy pipeline."""
|
||||
40
src/mnemosyne/core/models.py
Normal file
40
src/mnemosyne/core/models.py
Normal 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)
|
||||
64
src/mnemosyne/core/pipeline.py
Normal file
64
src/mnemosyne/core/pipeline.py
Normal 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
|
||||
150
src/mnemosyne/core/policy.py
Normal file
150
src/mnemosyne/core/policy.py
Normal 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
|
||||
24
src/mnemosyne/core/utils.py
Normal file
24
src/mnemosyne/core/utils.py
Normal 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
2222
src/mnemosyne/gateway.py
Normal file
File diff suppressed because it is too large
Load diff
45
src/mnemosyne/launcher.py
Normal file
45
src/mnemosyne/launcher.py
Normal 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
|
||||
572
src/mnemosyne/message_ops.py
Normal file
572
src/mnemosyne/message_ops.py
Normal 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("&", "&").replace('"', """).replace("<", "<").replace(">", ">")
|
||||
|
||||
|
||||
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>{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}
|
||||
283
src/mnemosyne/message_store.py
Normal file
283
src/mnemosyne/message_store.py
Normal 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
|
||||
249
src/mnemosyne/mnemosyne_config.py
Normal file
249
src/mnemosyne/mnemosyne_config.py
Normal 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
979
src/mnemosyne/pager.py
Normal 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)
|
||||
9
src/mnemosyne/providers/__init__.py
Normal file
9
src/mnemosyne/providers/__init__.py
Normal 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(),
|
||||
}
|
||||
63
src/mnemosyne/providers/anthropic.py
Normal file
63
src/mnemosyne/providers/anthropic.py
Normal 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"
|
||||
15
src/mnemosyne/providers/base.py
Normal file
15
src/mnemosyne/providers/base.py
Normal 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: ...
|
||||
76
src/mnemosyne/providers/openai.py
Normal file
76
src/mnemosyne/providers/openai.py
Normal 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
199
src/mnemosyne/tags.py
Normal 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
734
src/mnemosyne/trimmer.py
Normal 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()
|
||||
Loading…
Add table
Add a link
Reference in a new issue