From 974863e7b3779143b249ac518dc3730ef22e7d12 Mon Sep 17 00:00:00 2001 From: Joey Yakimowich-Payne Date: Fri, 13 Mar 2026 11:40:49 -0600 Subject: [PATCH] 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 --- src/mnemosyne/__init__.py | 6 + src/mnemosyne/__main__.py | 307 ++++ src/mnemosyne/blocks.py | 377 +++++ src/mnemosyne/config.py | 176 ++ src/mnemosyne/core/__init__.py | 1 + src/mnemosyne/core/models.py | 40 + src/mnemosyne/core/pipeline.py | 64 + src/mnemosyne/core/policy.py | 150 ++ src/mnemosyne/core/utils.py | 24 + src/mnemosyne/gateway.py | 2222 ++++++++++++++++++++++++++ src/mnemosyne/launcher.py | 45 + src/mnemosyne/message_ops.py | 572 +++++++ src/mnemosyne/message_store.py | 283 ++++ src/mnemosyne/mnemosyne_config.py | 249 +++ src/mnemosyne/pager.py | 979 ++++++++++++ src/mnemosyne/providers/__init__.py | 9 + src/mnemosyne/providers/anthropic.py | 63 + src/mnemosyne/providers/base.py | 15 + src/mnemosyne/providers/openai.py | 76 + src/mnemosyne/tags.py | 199 +++ src/mnemosyne/trimmer.py | 734 +++++++++ 21 files changed, 6591 insertions(+) create mode 100644 src/mnemosyne/__init__.py create mode 100644 src/mnemosyne/__main__.py create mode 100644 src/mnemosyne/blocks.py create mode 100644 src/mnemosyne/config.py create mode 100644 src/mnemosyne/core/__init__.py create mode 100644 src/mnemosyne/core/models.py create mode 100644 src/mnemosyne/core/pipeline.py create mode 100644 src/mnemosyne/core/policy.py create mode 100644 src/mnemosyne/core/utils.py create mode 100644 src/mnemosyne/gateway.py create mode 100644 src/mnemosyne/launcher.py create mode 100644 src/mnemosyne/message_ops.py create mode 100644 src/mnemosyne/message_store.py create mode 100644 src/mnemosyne/mnemosyne_config.py create mode 100644 src/mnemosyne/pager.py create mode 100644 src/mnemosyne/providers/__init__.py create mode 100644 src/mnemosyne/providers/anthropic.py create mode 100644 src/mnemosyne/providers/base.py create mode 100644 src/mnemosyne/providers/openai.py create mode 100644 src/mnemosyne/tags.py create mode 100644 src/mnemosyne/trimmer.py diff --git a/src/mnemosyne/__init__.py b/src/mnemosyne/__init__.py new file mode 100644 index 0000000..6c6e2b8 --- /dev/null +++ b/src/mnemosyne/__init__.py @@ -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. +""" diff --git a/src/mnemosyne/__main__.py b/src/mnemosyne/__main__.py new file mode 100644 index 0000000..9ff5313 --- /dev/null +++ b/src/mnemosyne/__main__.py @@ -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() diff --git a/src/mnemosyne/blocks.py b/src/mnemosyne/blocks.py new file mode 100644 index 0000000..05ec772 --- /dev/null +++ b/src/mnemosyne/blocks.py @@ -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 diff --git a/src/mnemosyne/config.py b/src/mnemosyne/config.py new file mode 100644 index 0000000..f3cb3fc --- /dev/null +++ b/src/mnemosyne/config.py @@ -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), + ) diff --git a/src/mnemosyne/core/__init__.py b/src/mnemosyne/core/__init__.py new file mode 100644 index 0000000..33c64bb --- /dev/null +++ b/src/mnemosyne/core/__init__.py @@ -0,0 +1 @@ +"""Core gateway abstractions: canonical model + policy pipeline.""" diff --git a/src/mnemosyne/core/models.py b/src/mnemosyne/core/models.py new file mode 100644 index 0000000..1963a5a --- /dev/null +++ b/src/mnemosyne/core/models.py @@ -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) diff --git a/src/mnemosyne/core/pipeline.py b/src/mnemosyne/core/pipeline.py new file mode 100644 index 0000000..c832cb6 --- /dev/null +++ b/src/mnemosyne/core/pipeline.py @@ -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 diff --git a/src/mnemosyne/core/policy.py b/src/mnemosyne/core/policy.py new file mode 100644 index 0000000..55c1a0a --- /dev/null +++ b/src/mnemosyne/core/policy.py @@ -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 diff --git a/src/mnemosyne/core/utils.py b/src/mnemosyne/core/utils.py new file mode 100644 index 0000000..48c971b --- /dev/null +++ b/src/mnemosyne/core/utils.py @@ -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")) diff --git a/src/mnemosyne/gateway.py b/src/mnemosyne/gateway.py new file mode 100644 index 0000000..0428e08 --- /dev/null +++ b/src/mnemosyne/gateway.py @@ -0,0 +1,2222 @@ +"""Pichay gateway — cache-aware multi-provider proxy for LLM context management. + +Primary entry point. Supersedes deprecated/proxy.py (Flask). + +Architecture: + Claude Code → Gateway (FastAPI) → Provider (Anthropic, OpenAI) + +The gateway intercepts, transforms, and forwards API requests. Features: +- Per-conversation session management with stable fingerprinting +- Static system prompt injection (cache-friendly, no prefix mutation) +- Dynamic status anchor (end-of-messages, after cache breakpoints) +- Token cap enforcement with cache hit rate tracking +- Policy pipeline: phantom protection → paging → trim +- Multi-provider support via adapter pattern +- Prometheus metrics + live HTML dashboard +""" + +from __future__ import annotations + +import argparse +import asyncio +import concurrent.futures +import copy +import hashlib +from contextlib import asynccontextmanager +import json +import os +import socket +import sys +import threading +import time +import uuid +from datetime import datetime, timezone +from pathlib import Path +from typing import Any + +import httpx +import uvicorn +from fastapi import FastAPI, HTTPException, Request, Response +from fastapi.responses import HTMLResponse, StreamingResponse + +from mnemosyne.blocks import BlockStore +from mnemosyne.core.models import CanonicalRequest +from mnemosyne.core.pipeline import Pipeline +from mnemosyne.core.policy import PolicyConfig, collect_blocks +from mnemosyne.core.utils import parse_duration +from mnemosyne.launcher import LaunchSpec, launch +from mnemosyne.message_ops import ( + check_inbound_for_injected_tags, + get_system_prompt, + inject_system_status, + process_cleanup_tags, + PICHAY_STATUS_MARKER, +) +from mnemosyne.fidelity import FidelityLevel, FidelityManager, PressureZone, make_object +from mnemosyne.object_store import ObjectStoreBackend +from mnemosyne.pager import PageStore, compact_messages +from mnemosyne.message_store import MessageStore +from mnemosyne.providers import adapters +from mnemosyne.telemetry import Telemetry + + +# ANSI for stderr status lines +_DIM = "\033[2m" +_YELLOW = "\033[33m" +_RED = "\033[31m" +_RESET = "\033[0m" + + +def _run_async(coro): + """Run an async coroutine from sync context within uvicorn. + + When called from a sync function that's running inside an async + event loop (e.g. FastAPI route → sync _preprocess), we can't use + asyncio.run() (loop already running). Instead, dispatch to a + thread pool and run a fresh event loop there. + """ + try: + asyncio.get_running_loop() + # We're inside an event loop — run in a thread with its own loop + with concurrent.futures.ThreadPoolExecutor(max_workers=1) as pool: + future = pool.submit(asyncio.run, coro) + return future.result(timeout=10) + except RuntimeError: + # No event loop running — safe to use asyncio.run directly + return asyncio.run(coro) + + +def find_free_port() -> int: + with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s: + s.bind(("127.0.0.1", 0)) + return s.getsockname()[1] + + +def _session_id(body: dict) -> str: + """Derive a stable session fingerprint from the conversation. + + Uses the system prompt (stable across turns within a session) combined + with the first user message text (unique per conversation). Falls back + to first message hash if no system prompt is present. + """ + parts: list[str] = [] + + # System prompt is the most stable identifier — same across all turns + system = body.get("system") + if isinstance(system, list): + # Extract text blocks only (ignore cache_control, etc.) + sys_text = "".join( + b.get("text", "") for b in system if isinstance(b, dict) and b.get("type") == "text" + ) + if sys_text: + # Use first 500 chars — enough to distinguish sessions, + # avoids hashing megabytes of system prompt + parts.append(sys_text[:500]) + elif isinstance(system, str) and system: + parts.append(system[:500]) + + # Add first user message text for uniqueness across concurrent sessions + # with the same system prompt + messages = body.get("messages", []) + for msg in messages: + if msg.get("role") == "user": + content = msg.get("content", "") + if isinstance(content, str): + parts.append(content[:200]) + elif isinstance(content, list): + for block in content: + if isinstance(block, dict) and block.get("type") == "text": + parts.append(block.get("text", "")[:200]) + break + break + + fingerprint = ( + "|".join(parts) if parts else json.dumps(messages[0], sort_keys=True) if messages else "" + ) + return hashlib.sha256(fingerprint.encode()).hexdigest()[:8] + + +def _duplication_score(req: CanonicalRequest) -> float: + texts = [b.text for b in collect_blocks(req) if b.text] + if not texts: + return 0.0 + unique = len(set(texts)) + return max(0.0, float(len(texts) - unique) / float(len(texts))) + + +def _add_cache_control(msg: dict) -> None: + """Add cache_control marker to the last content block of a message.""" + content = msg.get("content") + if isinstance(content, str): + msg["content"] = [{"type": "text", "text": content, "cache_control": {"type": "ephemeral"}}] + elif isinstance(content, list) and content: + last = content[-1] + if isinstance(last, dict): + last["cache_control"] = {"type": "ephemeral"} + + +def _strip_all_cache_controls(payload: dict) -> None: + """Remove all existing cache_control markers from system and messages. + + Must run before _place_cache_controls to prevent exceeding the 4-block + API limit. Claude Code places its own markers on the system prompt; + without stripping first we can end up with 5+. + """ + system = payload.get("system") + if isinstance(system, list): + for block in system: + if isinstance(block, dict): + block.pop("cache_control", None) + for msg in payload.get("messages", []): + content = msg.get("content") + if isinstance(content, list): + for block in content: + if isinstance(block, dict): + block.pop("cache_control", None) + + +def _place_cache_controls(payload: dict) -> None: + """Place up to 4 cache_control markers in the outbound payload. + + Mutates payload in place (the ephemeral copy — never the physical store). + Strips all existing markers first, then places fresh ones at: + 1. System prompt (always stable) + 2. 25% through messages + 3. 75% through messages + 4. Second-to-last message (previous assistant response) + """ + _strip_all_cache_controls(payload) + + # 1. System prompt + system = payload.get("system") + if system is not None: + if isinstance(system, str): + payload["system"] = [ + {"type": "text", "text": system, "cache_control": {"type": "ephemeral"}} + ] + elif isinstance(system, list) and system: + last = system[-1] + if isinstance(last, dict): + last["cache_control"] = {"type": "ephemeral"} + + messages = payload.get("messages", []) + n = len(messages) + if n < 2: + return + + placed_indices: set[int] = set() + + def _try_place(idx: int) -> None: + """Place cache_control at messages[idx] if valid (not last, not tool role).""" + if idx < 0 or idx >= n - 1: + return + if idx in placed_indices: + return + msg = messages[idx] + if msg.get("role") == "tool": + return + _add_cache_control(msg) + placed_indices.add(idx) + + # 2. 25% of the way through + _try_place(n // 4) + + # 3. 75% of the way through + _try_place(3 * n // 4) + + # 4. Second-to-last message + _try_place(n - 2) + + +def _copy_headers(headers: httpx.Headers) -> dict[str, str]: + dropped = { + "content-length", + "transfer-encoding", + "connection", + "content-encoding", + } + out: dict[str, str] = {} + for k, v in headers.items(): + if k.lower() in dropped: + continue + out[k] = v + return out + + +def _inspect_sse_chunk( + chunk: bytes, + *, + buffer: bytearray, + emit_event, + request_id: str, + session_id: str, + provider: str, + usage_accumulator: dict[str, Any] | None = None, +) -> None: + """Best-effort SSE validation for telemetry; never mutates payload.""" + buffer.extend(chunk) + while b"\n\n" in buffer: + raw_event, rest = buffer.split(b"\n\n", 1) + buffer[:] = rest + if not raw_event: + continue + try: + text = raw_event.decode("utf-8") + except UnicodeDecodeError as e: + emit_event( + "anomaly", + kind="malformed_stream_chunk", + request_id=request_id, + session_id=session_id, + provider=provider, + error=f"invalid_utf8:{e}", + ) + continue + + for line in text.splitlines(): + if not line.startswith("data: "): + continue + payload = line[6:] + if payload == "[DONE]": + continue + try: + evt = json.loads(payload) + if usage_accumulator is not None and isinstance(evt, dict): + etype = evt.get("type") + if etype == "message_start": + msg = evt.get("message", {}) + if isinstance(msg, dict): + usage = msg.get("usage", {}) + if isinstance(usage, dict): + usage_accumulator.update(usage) + elif etype == "message_delta": + usage = evt.get("usage", {}) + if isinstance(usage, dict): + usage_accumulator.update(usage) + except json.JSONDecodeError as e: + emit_event( + "anomaly", + kind="malformed_stream_chunk", + request_id=request_id, + session_id=session_id, + provider=provider, + error=f"invalid_json:{e}", + ) + + +def _dashboard_html() -> str: + return """ + +Mnemosyne + + +
+ +
Mnemosyne
Context Memory Proxy
+
+
connecting…
+
+
+

Context Reduction

+
—
token savings
+

Active Sessions

+
0
connected clients
+

Objects Stored

+
0
in memory store
+

Fidelity Degradations

+
0
total downgrades
+

Shrink Ratio Trend

+
+

Fidelity Distribution

+
L0
0
+
L1
0
+
L2
0
+
L3
0
+
L4
0
+
+

Admission Control

+ + +— +
Admitted
0
+
Rejected
0
+
+

Micro-Faults

+
Attempts0
+
Successes0
+
Tokens Saved0
+
Success Rate—
+

Latency (ms)

+
Average—
+
Minimum—
+
Maximum—
+
P95—
+

Sessions

+ + +
SessionRequestsObjectsFidelityCtx ReductionEvictionsFaultsFault RateIncomingOutgoing
+
Event Log +
+
+ +""" + + +# ── Session state ──────────────────────────────────────────────────────── + + +class Session: + """Per-conversation state. Isolated by first-message fingerprint.""" + + def __init__( + self, + sid: str, + log_dir: Path, + object_store_config: Any | None = None, + ): + self.id = sid + self.token_state = { + "last_effective": 0, + "blocked": False, + "turn": 0, + } + self.page_store = PageStore( + log_path=log_dir / f"pages_{sid}.jsonl", + ) + self._block_checkpoint = log_dir / f"blocks_{sid}.json" + self._page_checkpoint = log_dir / f"pages_{sid}_checkpoint.json" + if self._page_checkpoint.is_file(): + import json as _json + + try: + with open(self._page_checkpoint) as f: + data = _json.load(f) + self.page_store.restore(data) + print( + f" {_DIM}[{sid}] restored page_store: " + f"{len(self.page_store._released_handles)} released handles, " + f"{len(self.page_store._pinned)} pinned{_RESET}", + file=sys.stderr, + ) + except (json.JSONDecodeError, ValueError, KeyError) as exc: + print( + f" {_YELLOW}[{sid}] corrupt page checkpoint, ignoring: {exc}{_RESET}", + file=sys.stderr, + ) + if self._block_checkpoint.is_file(): + self.block_store = BlockStore.from_checkpoint(self._block_checkpoint) + print( + f" {_DIM}[{sid}] restored {self.block_store.block_count} blocks{_RESET}", + file=sys.stderr, + ) + else: + self.block_store = BlockStore() + self.message_store = MessageStore( + sid, + self.page_store, + log_path=log_dir / f"violations_{sid}.jsonl" if log_dir else None, + ) + self.fidelity_manager = FidelityManager(window_size=200_000) + # Maps content keys → fidelity object IDs for lookup during apply + self._fidelity_content_map: dict[str, str] = {} + self.last_cleanup_stats: str | None = None + + # Semantic object management + from mnemosyne.segmenter import Segmenter, SegmentedObject + from mnemosyne.object_store import ObjectStore, InMemoryBackend, DummyEmbedder + from mnemosyne.context_assembler import ContextAssembler + from mnemosyne.embedder import try_get_embedder + + embedder = try_get_embedder() or DummyEmbedder() + self.segmenter = Segmenter() + + # Choose backend based on config + backend: ObjectStoreBackend + if ( + object_store_config is not None + and getattr(object_store_config, "backend", "memory") == "postgresql" + ): + from mnemosyne.pgvector_backend import PgVectorBackend + + pg_backend = PgVectorBackend( + host=object_store_config.pg_host, + port=object_store_config.pg_port, + database=object_store_config.pg_database, + user=object_store_config.pg_user, + password=object_store_config.pg_password, + ) + _run_async(pg_backend.connect()) + backend = pg_backend + print( + f" {_DIM}[{sid}] using PostgreSQL backend " + f"({object_store_config.pg_host}:{object_store_config.pg_port}){_RESET}", + file=sys.stderr, + ) + else: + backend = InMemoryBackend() + + self.object_store = ObjectStore(backend, embedder=embedder) + self._segmented_objects: list[SegmentedObject] = [] + # ContextAssembler will be set by create_app after helper_llm is resolved + self.context_assembler: ContextAssembler | None = None + + # Phase 4c: Goal-aware retrieval + from mnemosyne.helper_llm import GoalClassification + + self._last_user_embedding: list[float] | None = None + self._current_goal: GoalClassification | None = None + + # Phase 4d: Admission control + from mnemosyne.admission import AdmissionController + + self.admission = AdmissionController() + + # Phase 4e: Entropy-gated faulting + from mnemosyne.entropy import EntropyDetector + + self.entropy_detector = EntropyDetector() + + # Phase 6: Hierarchical segmentation (Strategy B, turns 50+) + from mnemosyne.hierarchy import SessionHierarchy + + self.hierarchy = SessionHierarchy(embedder) + + # Benchmark metrics + from mnemosyne.benchmark import SessionBenchmark + + self.benchmark = SessionBenchmark(sid) + + # Summary cache: (object_id, FidelityLevel) → summary text + # Avoids re-summarizing the same content on repeated _apply_fidelity calls + self._summary_cache: dict[tuple[str, int], str] = {} + + def track_usage(self, usage: dict) -> None: + """Update token state from API response usage.""" + effective = ( + usage.get("input_tokens", 0) + + usage.get("cache_creation_input_tokens", 0) + + usage.get("cache_read_input_tokens", 0) + ) + if effective > 0: + self.token_state["last_effective"] = effective + + def increment_turn(self) -> None: + self.token_state["turn"] += 1 + + +class SessionStore: + """Maps conversation fingerprints to session state.""" + + def __init__(self, log_dir: Path, object_store_config: Any | None = None): + self._sessions: dict[str, Session] = {} + self._log_dir = log_dir + self._object_store_config = object_store_config + + def get(self, body: dict) -> Session: + sid = _session_id(body) + if sid not in self._sessions: + self._sessions[sid] = Session( + sid, self._log_dir, object_store_config=self._object_store_config + ) + print( + f" {_DIM}[{sid}] new session{_RESET}", + file=sys.stderr, + ) + return self._sessions[sid] + + def get_by_id(self, session_id: str) -> Session | None: + """Look up a session by its ID string. Returns None if not found.""" + return self._sessions.get(session_id) + + def all(self) -> dict[str, Session]: + return dict(self._sessions) + + +# ── Fidelity integration helpers ───────────────────────────────────────── + +_FIDELITY_MIN_BLOCK_SIZE = 500 # Only track content blocks >= 500 bytes + + +def _content_key(block: dict, msg: dict) -> str | None: + """Derive a stable tracking key for a content block. + + Tool results use tool_use_id; large text blocks use a content hash. + Returns None for blocks that shouldn't be tracked. + """ + # Tool result blocks carry a tool_use_id + if block.get("type") == "tool_result" and block.get("tool_use_id"): + return f"tool:{block['tool_use_id']}" + + # Text blocks — only track if large enough + text = block.get("text", "") + if ( + isinstance(text, str) + and len(text.encode("utf-8", errors="replace")) >= _FIDELITY_MIN_BLOCK_SIZE + ): + return ( + f"text:{hashlib.sha256(text[:200].encode('utf-8', errors='replace')).hexdigest()[:12]}" + ) + + return None + + +def _block_text(block: dict) -> str: + """Extract the text content from a content block.""" + if block.get("type") == "tool_result": + # tool_result content can be a string or list of blocks + content = block.get("content", "") + if isinstance(content, str): + return content + if isinstance(content, list): + parts = [] + for sub in content: + if isinstance(sub, dict) and sub.get("type") == "text": + parts.append(sub.get("text", "")) + return "\n".join(parts) + return "" + return block.get("text", "") + + +def _auto_stub(text: str) -> str: + """Generate a simple stub from the first line of content.""" + first_line = text.split("\n", 1)[0].strip() + if len(first_line) > 120: + first_line = first_line[:117] + "..." + return f"[evicted content: {first_line}]" + + +import logging as _logging + +_summarization_logger = _logging.getLogger("mnemosyne.summarization") + + +async def _generate_degradation_summaries( + transitions: list[tuple[str, "FidelityLevel", "FidelityLevel"]], + session: "Session", + helper_llm: Any, +) -> None: + """Generate LLM-powered summaries for degraded objects (async, non-blocking). + + For each transition: + L0→L1: Call helper_llm.summarize_l0_to_l1() → store as summary_detailed + L1→L2: Call helper_llm.compress_l1_to_l2() → store as summary_compact + L2→L3: Simple stub (no LLM call) + L3→L4: No content needed (evicted) + + Results are cached in session._summary_cache to avoid re-summarization. + On any HelperLLM failure, falls back to existing stub text silently. + """ + if helper_llm is None: + return + + fm = session.fidelity_manager + cache = session._summary_cache + + for obj_id, old_level, new_level in transitions: + obj = fm.get_object(obj_id) + if obj is None: + continue + + cache_key = (obj_id, int(new_level)) + + # Skip if already cached + if cache_key in cache: + continue + + try: + if old_level == FidelityLevel.L0 and new_level == FidelityLevel.L1: + # L0→L1: Generate detailed summary from full content + result = await helper_llm.summarize_l0_to_l1( + content=obj.content_full, + object_type=obj.object_type, + max_summary_tokens=1024, + ) + if result.summary: + obj.summary_detailed = result.summary + obj.losses_l1 = result.losses + obj.can_answer = result.can_answer + obj.key_entities = result.key_entities + obj.token_count_l1 = max(1, len(result.summary) // 4) + cache[cache_key] = result.summary + _summarization_logger.debug( + "L0→L1 summary generated for %s (%d chars)", + obj_id[:8], + len(result.summary), + ) + else: + # Empty result — use fallback + _summarization_logger.warning( + "L0→L1 empty summary for %s, using fallback", obj_id[:8] + ) + + elif old_level == FidelityLevel.L1 and new_level == FidelityLevel.L2: + # L1→L2: Compress L1 summary further + l1_text = obj.summary_detailed or cache.get((obj_id, int(FidelityLevel.L1)), "") + if not l1_text: + # No L1 text available — skip + _summarization_logger.warning("L1→L2 no L1 text for %s, skipping", obj_id[:8]) + continue + result = await helper_llm.compress_l1_to_l2( + l1_summary=l1_text, + l1_losses=obj.losses_l1, + object_type=obj.object_type, + max_tokens=256, + ) + if result.summary: + obj.summary_compact = result.summary + obj.losses_l2 = result.losses + obj.token_count_l2 = max(1, len(result.summary) // 4) + cache[cache_key] = result.summary + _summarization_logger.debug( + "L1→L2 summary generated for %s (%d chars)", + obj_id[:8], + len(result.summary), + ) + else: + _summarization_logger.warning( + "L1→L2 empty summary for %s, using fallback", obj_id[:8] + ) + + elif old_level == FidelityLevel.L2 and new_level == FidelityLevel.L3: + # L2→L3: Simple stub, no LLM call + stub = f"[evicted: {obj.object_type} from turn {obj.created_at_turn}]" + obj.stub = stub + obj.token_count_l3 = max(1, len(stub) // 4) + cache[cache_key] = stub + + # L3→L4: Full eviction — no content needed, nothing to generate + + except Exception: + # Graceful fallback: log and continue with existing stub text + _summarization_logger.warning( + "Summarization failed for %s (%s→%s), using fallback", + obj_id[:8], + old_level.name, + new_level.name, + exc_info=True, + ) + + +def _apply_fidelity(payload: dict, session: "Session") -> None: + """Register and apply fidelity-based content replacement on the ephemeral payload. + + For each message, scans content blocks for tool_results and large text blocks. + Registers new objects in the FidelityManager, and replaces degraded objects + with their stub/tombstone representation. + + This runs on the ephemeral copy — never mutates the physical message store. + """ + fm = session.fidelity_manager + content_map = session._fidelity_content_map + turn = session.token_state.get("turn", 0) + transitions_logged: list[str] = [] + + messages = payload.get("messages", []) + for msg in messages: + content = msg.get("content") + if not isinstance(content, list): + continue + + for i, block in enumerate(content): + if not isinstance(block, dict): + continue + + key = _content_key(block, msg) + if key is None: + continue + + text = _block_text(block) + if not text: + continue + + if key in content_map: + # Already tracked — check if degraded + obj_id = content_map[key] + obj = fm.get_object(obj_id) + if obj is None: + continue + + # Mark as accessed this turn + fm.mark_accessed(obj_id, turn) + + if obj.current_fidelity >= FidelityLevel.L3: + # Replace with stub + stub = obj.stub or _auto_stub(obj.content_full) + if block.get("type") == "tool_result": + block["content"] = stub + else: + block["text"] = stub + elif obj.current_fidelity == FidelityLevel.L2: + # L2: use compact summary if available (from LLM or pre-set) + summary = ( + session._summary_cache.get((obj_id, int(FidelityLevel.L2))) + or obj.summary_compact + ) + if summary: + if block.get("type") == "tool_result": + block["content"] = summary + else: + block["text"] = summary + # else: keep full content as fallback + elif obj.current_fidelity == FidelityLevel.L1: + # L1: use detailed summary if available (from LLM or pre-set) + summary = ( + session._summary_cache.get((obj_id, int(FidelityLevel.L1))) + or obj.summary_detailed + ) + if summary: + if block.get("type") == "tool_result": + block["content"] = summary + else: + block["text"] = summary + # else: keep full content as fallback + else: + # New object — register it + obj_type = "tool_result" if block.get("type") == "tool_result" else "file_context" + stub = _auto_stub(text) + obj = make_object( + object_type=obj_type, + content_full=text, + created_at_turn=turn, + stub=stub, + ) + obj_id = fm.register_object(obj) + content_map[key] = obj_id + + # Check pressure and log + pressure = fm.current_pressure() + if pressure > PressureZone.NORMAL: + transitions = fm.degrade(turn) + for obj_id, old_level, new_level in transitions: + transitions_logged.append(f"{obj_id[:8]}:{old_level.name}→{new_level.name}") + + if transitions_logged: + print( + f" {_DIM}[{session.id}] fidelity: pressure={pressure.name}, " + f"degraded {len(transitions_logged)} objects: " + f"{', '.join(transitions_logged[:5])}" + f"{'...' if len(transitions_logged) > 5 else ''}{_RESET}", + file=sys.stderr, + ) + + +# ── Gateway factory ────────────────────────────────────────────────────── + + +def create_app( + *, + log_dir: Path, + anthropic_upstream: str = "https://api.anthropic.com", + openai_upstream: str = "https://api.openai.com", + hydration_window_seconds: int = 86400, + enable_paging: bool = True, + enable_trim: bool = True, + min_evict_size: int = 500, + token_cap: int = 0, + anthropic_model_override: str | None = None, + openai_model_override: str | None = None, + process_session_id: str | None = None, +) -> FastAPI: + clients: dict[str, httpx.Client] = {} + + if process_session_id is None: + process_session_id = datetime.now(timezone.utc).strftime("proc_%Y%m%d_%H%M%S") + + @asynccontextmanager + async def lifespan(app: FastAPI): + try: + yield + finally: + for c in clients.values(): + c.close() + + app = FastAPI(title="Pichay Gateway", version="0.3.0", lifespan=lifespan) + + log_dir.mkdir(parents=True, exist_ok=True) + log_path = log_dir / f"gateway_{datetime.now(timezone.utc).strftime('%Y%m%d_%H%M%S')}.jsonl" + telemetry = Telemetry(log_path=log_path, hydration_window_seconds=hydration_window_seconds) + app.state.telemetry = telemetry + app.state.process_session_id = process_session_id + + # Load Mnemosyne config for backend selection + from mnemosyne.mnemosyne_config import load_config + + mnemosyne_config = load_config() + object_store_config = mnemosyne_config.object_store + + sessions = SessionStore(log_dir, object_store_config=object_store_config) + app.state.sessions = sessions + + # Configure OAuth / API key authentication + from mnemosyne.oauth import configure_environment + + auth_method = configure_environment() + print(f"Auth: {auth_method}", file=sys.stderr) + + # Shared HelperLLM — one per gateway process (graceful degradation without auth) + helper_llm = None + if os.environ.get("ANTHROPIC_AUTH_TOKEN") or os.environ.get("ANTHROPIC_API_KEY"): + from mnemosyne.helper_llm import HelperLLM + + helper_llm = HelperLLM() + app.state.helper_llm = helper_llm + + _token_warning_threshold = int(token_cap * 0.80) if token_cap > 0 else 0 + + cfg = PolicyConfig( + enable_paging=enable_paging, + enable_trim=enable_trim, + min_evict_size=min_evict_size, + ) + + def emit_event(event_type: str, **fields: Any) -> None: + telemetry.emit( + event_type, + process_session_id=process_session_id, + **fields, + ) + + pipeline = Pipeline(cfg=cfg, emit_event=emit_event) + app.state.pipeline = pipeline + + app.state.adapters = adapters() + clients.update( + { + "anthropic": httpx.Client( + base_url=anthropic_upstream, timeout=httpx.Timeout(300.0, connect=30.0) + ), + "openai": httpx.Client( + base_url=openai_upstream, timeout=httpx.Timeout(300.0, connect=30.0) + ), + } + ) + app.state.clients = clients + + # ── Endpoints ──────────────────────────────────────────────────── + + @app.get("/health") + def health() -> dict[str, Any]: + all_sessions = sessions.all() + session_summaries = {} + for sid, s in all_sessions.items(): + ps = s.page_store + session_summaries[sid] = { + "turn": s.token_state["turn"], + "last_effective_tokens": s.token_state["last_effective"], + "evictions": ps.unique_evictions, + "gc": ps.gc_count, + "faults": len(ps.faults), + } + return { + "status": "ok", + "process_session_id": process_session_id, + "providers": ["anthropic", "openai"], + "log_path": str(log_path), + "token_cap": token_cap, + "model_overrides": { + "anthropic": anthropic_model_override, + "openai": openai_model_override, + }, + "sessions": session_summaries, + } + + @app.get("/metrics") + def metrics() -> Response: + return Response(content=telemetry.get_metrics(), media_type="text/plain; version=0.0.4") + + @app.get("/api/sessions") + def api_sessions() -> dict[str, Any]: + return { + "process_session_id": process_session_id, + "sessions": telemetry.session_summary(), + } + + @app.get("/api/events") + def api_events(window: str | None = None) -> dict[str, Any]: + window_seconds = None + if window: + try: + window_seconds = parse_duration(window) + except ValueError as e: + raise HTTPException(status_code=400, detail=str(e)) + return {"events": telemetry.recent_events(window_seconds)} + + @app.get("/api/cost") + def api_cost(window: str | None = None) -> dict[str, Any]: + window_seconds = None + if window: + try: + window_seconds = parse_duration(window) + except ValueError as e: + raise HTTPException(status_code=400, detail=str(e)) + return telemetry.cost_summary(window_seconds) + + # ── Plugin API endpoints ──────────────────────────────────────── + # Called by the opencode-mnemosyne plugin for context retrieval. + + @app.get("/api/memory") + async def api_memory( + session_id: str | None = None, + query: str | None = None, + limit: int = 5, + ) -> dict[str, Any]: + """Search the object store for a session, or return aggregate summary. + + When called without params (from dashboard), returns per-session object counts. + When called with session_id + query (from plugin), performs semantic search. + """ + # Aggregate summary mode — used by the dashboard + if session_id is None or query is None: + all_sessions = sessions.all() + session_summaries: dict[str, Any] = {} + for sid, sess in all_sessions.items(): + try: + objs = await sess.object_store.get_session_objects(sid) + total_tokens = sum(obj.current_tokens for obj in objs) + fidelity_dist: dict[int, int] = {} + for obj in objs: + fidelity_dist[obj.current_fidelity] = ( + fidelity_dist.get(obj.current_fidelity, 0) + 1 + ) + session_summaries[sid] = { + "total_objects": len(objs), + "total_tokens": total_tokens, + "fidelity_distribution": fidelity_dist, + } + except Exception: + session_summaries[sid] = { + "total_objects": 0, + "total_tokens": 0, + "fidelity_distribution": {}, + } + return {"sessions": session_summaries} + + # Semantic search mode — used by the plugin + session = sessions.get_by_id(session_id) + if session is None: + raise HTTPException(status_code=404, detail=f"Session '{session_id}' not found") + + results = await session.object_store.semantic_search(session_id, query, limit=limit) + all_objects = await session.object_store.get_session_objects(session_id) + + total_tokens = sum(obj.current_tokens for obj in all_objects) + + objects_out = [] + for obj, score in results: + objects_out.append( + { + "id": obj.id, + "object_type": obj.object_type, + "stub": obj.stub, + "current_fidelity": obj.current_fidelity, + "tokens": obj.current_tokens, + "similarity": round(score, 4), + } + ) + + return { + "objects": objects_out, + "total_objects": len(all_objects), + "session_tokens": total_tokens, + } + + @app.get("/api/compaction-context") + async def api_compaction_context( + session_id: str, + max_objects: int = 20, + ) -> dict[str, Any]: + """Return a summary of the object store for compaction prompt injection.""" + session = sessions.get_by_id(session_id) + if session is None: + raise HTTPException(status_code=404, detail=f"Session '{session_id}' not found") + + all_objects = await session.object_store.get_session_objects(session_id) + # Sort by fidelity (lower = higher fidelity = more important), then by access count + all_objects.sort( + key=lambda o: (o.current_fidelity, -o.access_count), + ) + + # Take the top N objects + selected = all_objects[:max_objects] + + if not selected: + return {"context": "No objects stored in memory."} + + lines = [ + f"Memory contains {len(all_objects)} semantic objects " + f"({sum(o.current_tokens for o in all_objects):,} tokens total):", + "", + ] + fidelity_labels = {0: "Full", 1: "Summary", 2: "Compact", 3: "Stub", 4: "Evicted"} + for obj in selected: + fl = fidelity_labels.get(obj.current_fidelity, "?") + lines.append( + f"- [{fl}] {obj.object_type}: {obj.stub[:120]}" + f" ({obj.current_tokens} tok, accessed {obj.access_count}x)" + ) + + return {"context": "\n".join(lines)} + + @app.get("/api/benchmark") + def api_benchmark(session_id: str | None = None) -> dict[str, Any]: + """Return benchmark metrics for a session or all sessions.""" + from mnemosyne.benchmark import BenchmarkCollector + + # Collect all session benchmarks into a temporary collector + collector = BenchmarkCollector() + all_sess = sessions.all() + for sid, sess in all_sess.items(): + collector._sessions[sid] = sess.benchmark + + if session_id is not None: + report = collector.session_report(session_id) + if report is None: + raise HTTPException(status_code=404, detail=f"Session '{session_id}' not found") + return {"type": "session", **report} + + result = collector.aggregate() + result["type"] = "aggregate" + result["per_session"] = {sid: sess.benchmark.to_dict() for sid, sess in all_sess.items()} + return result + + @app.get("/dashboard") + def dashboard() -> HTMLResponse: + return HTMLResponse(_dashboard_html()) + + # ── Pre/post processing ────────────────────────────────────────── + + def _preprocess(payload: dict, session: Session) -> dict: + """Apply gateway transformations before pipeline and forwarding. + + Operates on the raw Anthropic-format payload (before normalization) + because system prompt injection needs access to the system field + directly. + """ + from mnemosyne.benchmark import Timer + + incoming_messages = payload.get("messages", []) + + # Guard: don't forward turns with empty user content (injection-only turns) + if incoming_messages and incoming_messages[-1].get("role") == "user": + last_content = incoming_messages[-1].get("content", "") + if isinstance(last_content, str): + _is_empty = not last_content.strip() + elif isinstance(last_content, list): + _is_empty = len(last_content) == 0 + else: + _is_empty = True + if _is_empty: + raise HTTPException( + status_code=422, detail="Empty user turn — not forwarded to provider" + ) + + session.increment_turn() + ts = session.token_state + request_time = datetime.now(timezone.utc) + ps = session.page_store + ms = session.message_store + + # 1. Ingest into MessageStore (asserts append-only, compacts) + ingest = ms.ingest( + incoming_messages, + age_threshold=4, + min_evict_size=min_evict_size, + ) + if ingest.new_count > 0 or ingest.compacted_count > 0: + parts = [] + if ingest.new_count: + parts.append(f"+{ingest.new_count} msgs") + if ingest.compacted_count: + parts.append(f"{ingest.compacted_count} evicted") + if ingest.bytes_saved: + parts.append(f"{ingest.bytes_saved:,} bytes saved") + if ingest.mutations_detected: + parts.append(f"{ingest.mutations_detected} MUTATIONS") + if ingest.deletions_detected: + parts.append(f"{ingest.deletions_detected} DELETIONS") + print( + f" {_DIM}[{session.id}] ingest: {', '.join(parts)}{_RESET}", + file=sys.stderr, + ) + + # 1b. Segment new messages into semantic objects and store in ObjectStore + # Phase 4d: admission control gates each object before storage + try: + with Timer(session.benchmark.latency["segmentation"]): + segmented = session.segmenter.segment_incremental( + ms.messages, + session._segmented_objects, + start_turn=0, + ) + new_count = len(segmented) - len(session._segmented_objects) + if new_count > 0: + new_objects = segmented[len(session._segmented_objects) :] + admitted_count = 0 + rejected_count = 0 + for seg_obj in new_objects: + # Phase 4d: admission control — score and gate + has_dup = False + if seg_obj.source_key: + dup = _run_async( + session.object_store.find_duplicate(session.id, seg_obj.source_key) + ) + has_dup = dup is not None + with Timer(session.benchmark.latency["admission"]): + admitted, _score = session.admission.should_admit( + seg_obj.content, + seg_obj.object_type, + has_dup, + seg_obj.key_entities, + ) + session.benchmark.admission.record_decision( + admitted, _score.total, seg_obj.object_type + ) + if not admitted: + rejected_count += 1 + continue + admitted_count += 1 + with Timer(session.benchmark.latency["embedding"]): + _run_async( + session.object_store.store_object( + session_id=session.id, + content=seg_obj.content, + object_type=seg_obj.object_type, + source_tool=seg_obj.source_tool, + source_key=seg_obj.source_key, + stub=seg_obj.stub, + tags=seg_obj.tags, + key_entities=seg_obj.key_entities, + turn=seg_obj.turn_start, + ) + ) + session.benchmark.segmentation.record_object( + seg_obj.object_type, + max(1, len(seg_obj.content) // 4), + ) + session._segmented_objects = segmented + parts = [ + f"{admitted_count} new objects ({len(segmented)} total)", + ] + if rejected_count: + parts.append(f"{rejected_count} rejected by admission") + print( + f" {_DIM}[{session.id}] segmented: {', '.join(parts)}{_RESET}", + file=sys.stderr, + ) + except Exception as exc: + print( + f" {_YELLOW}[{session.id}] segmenter error (non-fatal): {exc}{_RESET}", + file=sys.stderr, + ) + + # 1b-ii. Hierarchical segmentation (Strategy B, turns 50+) + # Feed stored objects into the hierarchy for clustering. + # After turn 50, run periodic maintenance (re-cluster every 20 turns). + try: + turn = ts.get("turn", 0) + if turn >= 50: + # Enable hierarchy after turn 50 + session.hierarchy.enabled = True + + # Feed all session objects into hierarchy if not yet populated + if session.hierarchy.object_count == 0: + all_stored = _run_async(session.object_store.get_session_objects(session.id)) + if all_stored: + session.hierarchy.rebuild(all_stored) + print( + f" {_DIM}[{session.id}] hierarchy: initial build " + f"({session.hierarchy.object_count} objects, " + f"{session.hierarchy.episode_count} episodes, " + f"{session.hierarchy.theme_count} themes){_RESET}", + file=sys.stderr, + ) + else: + # Incremental: add any new objects from this turn + all_stored = _run_async(session.object_store.get_session_objects(session.id)) + for obj in all_stored: + session.hierarchy.add_object(obj) + + # Periodic maintenance (re-cluster every 20 turns or on goal change) + goal_hash = None + if session._current_goal is not None: + goal_hash = session._current_goal.goal + rebuilt = session.hierarchy.maintenance(turn, goal_hash=goal_hash) + if rebuilt: + print( + f" {_DIM}[{session.id}] hierarchy: rebuilt at turn {turn} " + f"({session.hierarchy.episode_count} episodes, " + f"{session.hierarchy.theme_count} themes){_RESET}", + file=sys.stderr, + ) + else: + # Before turn 50, hierarchy is disabled (Strategy A only) + session.hierarchy.enabled = False + except Exception as exc: + print( + f" {_YELLOW}[{session.id}] hierarchy error (non-fatal): {exc}{_RESET}", + file=sys.stderr, + ) + + # 1c. Goal-aware retrieval — detect topic shifts and reclassify + try: + # Extract last user message text + user_text = "" + for msg in reversed(incoming_messages): + if msg.get("role") == "user": + content = msg.get("content", "") + if isinstance(content, str): + user_text = content + elif isinstance(content, list): + user_text = " ".join( + b.get("text", "") + for b in content + if isinstance(b, dict) and b.get("type") == "text" + ) + break + + if user_text and session.object_store._embedder is not None: + from mnemosyne.object_store import _cosine_similarity + + embedder = session.object_store._embedder + current_embedding = embedder.embed(user_text) + + # Check for topic shift + goal_changed = False + if session._last_user_embedding is not None: + sim = _cosine_similarity(current_embedding, session._last_user_embedding) + goal_changed = sim < 0.5 + else: + goal_changed = True # First message — classify + + session._last_user_embedding = current_embedding + + if goal_changed and helper_llm is not None: + session.benchmark.goals.record_topic_shift() + # Build recent context (last 2 messages) + recent = ( + incoming_messages[-4:] if len(incoming_messages) > 4 else incoming_messages + ) + recent_text = "\n".join(str(m.get("content", ""))[:500] for m in recent) + with Timer(session.benchmark.latency["goal_classification"]): + goal = _run_async(helper_llm.classify_goal(user_text, recent_text)) + session._current_goal = goal + session.benchmark.goals.record_reclassification() + print( + f" {_DIM}[{session.id}] goal transition: {goal.goal[:80]}, " + f"types={goal.relevant_types}, tags={goal.relevant_tags}{_RESET}", + file=sys.stderr, + ) + + # Adjust fidelity based on goal relevance + turn = ts.get("turn", 0) + fm = session.fidelity_manager + promotions = 0 + for obj_id in list(fm._objects.keys()): + obj = fm.get_object(obj_id) + if obj is None or obj.pinned: + continue + # Check if object matches goal via stored object tags + stored = _run_async(session.object_store.get(obj_id)) + obj_tags = stored.tags if stored else [] + is_relevant = obj.object_type in goal.relevant_types or any( + t in obj_tags for t in goal.relevant_tags + ) + if is_relevant and obj.current_fidelity > FidelityLevel.L1: + # Promote relevant objects + fm.upgrade(obj_id, FidelityLevel.L1, turn) + promotions += 1 + elif not is_relevant and obj.current_fidelity < FidelityLevel.L2: + # Don't aggressively degrade — let pressure handle it + pass + if promotions: + session.benchmark.goals.record_promotion(promotions) + except Exception as exc: + print( + f" {_YELLOW}[{session.id}] goal detection error (non-fatal): {exc}{_RESET}", + file=sys.stderr, + ) + + # 2. Process model releases on OUR messages + cleanup_stats = process_cleanup_tags( + ms.messages, + session.block_store, + ps, + ) + session.last_cleanup_stats = cleanup_stats + if cleanup_stats: + print( + f" {_DIM}[{session.id}] cleanup: {cleanup_stats}{_RESET}", + file=sys.stderr, + ) + + # 3. Detect page faults + faults = ps.detect_faults(ms.messages) + if faults: + for fault in faults: + ago = time.monotonic() - fault.original_eviction.evicted_at + print( + f" {_YELLOW}[{session.id}] PAGE FAULT: " + f"{fault.tool_name} (evicted {ago:.0f}s ago){_RESET}", + file=sys.stderr, + ) + + # 3b. Entropy-gated faulting — check last assistant response for uncertainty + try: + last_assistant_text = "" + for msg in reversed(ms.messages): + if msg.get("role") == "assistant": + content = msg.get("content", "") + if isinstance(content, str): + last_assistant_text = content + elif isinstance(content, list): + last_assistant_text = " ".join( + b.get("text", "") + for b in content + if isinstance(b, dict) and b.get("type") == "text" + ) + break + + if last_assistant_text: + # Collect evicted entities + evicted_entities: list[str] = [] + fm = session.fidelity_manager + for obj_id, obj in fm._objects.items(): + if obj.current_fidelity >= FidelityLevel.L4: + # Get key_entities from object store + stored = _run_async(session.object_store.get(obj_id)) + if stored and stored.key_entities: + evicted_entities.extend(stored.key_entities) + + if evicted_entities: + signal = session.entropy_detector.analyze_response( + last_assistant_text, evicted_entities + ) + fault_entities = session.entropy_detector.should_fault(signal) + if fault_entities: + print( + f" {_YELLOW}[{session.id}] entropy fault: score={signal.score:.2f}, " + f"faulting {len(fault_entities)} entities{_RESET}", + file=sys.stderr, + ) + # Proactively restore via micro-fault for next turn + if session.context_assembler: + for entity in fault_entities[:3]: # Limit to 3 to avoid latency + try: + _run_async( + session.context_assembler.handle_micro_fault( + session_id=session.id, + question=f"What was the content related to {entity}?", + scope=entity, + ) + ) + except Exception: + pass # Best effort + except Exception as exc: + print( + f" {_YELLOW}[{session.id}] entropy check error (non-fatal): {exc}{_RESET}", + file=sys.stderr, + ) + + # 4. Build ephemeral outbound view — never mutate the physical store + payload["messages"] = copy.deepcopy(ms.messages) + + # 4a. Inject phantom tool definitions and resolve pending calls + from mnemosyne.phantom import inject_tools, inject_phantom_results, extract_phantom_calls + + observe_only = inject_tools(payload) + # Check for pending phantom calls from the previous assistant turn + pending_calls = extract_phantom_calls(payload.get("messages", [])) + if pending_calls: + inject_phantom_results( + payload["messages"], + pending_calls, + ps, + observe_only, + context_assembler=session.context_assembler, + session_id=session.id, + ) + + # 4b. Apply fidelity-based content replacement + _apply_fidelity(payload, session) + + # Place cache_control markers at optimal positions + _place_cache_controls(payload) + + # System status: static system prompt + dynamic anchor + inject_system_status( + payload, + ts, + token_cap, + request_time, + block_store=session.block_store, + page_store=ps, + last_cleanup_stats=session.last_cleanup_stats, + ) + + # Block labeling (with injection-safe validation) + session.block_store.label_messages(ms.messages, ts["turn"]) + + # Checkpoint page_store (releases, pins survive restart) + import json as _json + + with open(session._page_checkpoint, "w") as f: + _json.dump(session.page_store.checkpoint(), f) + + return payload + + def _check_token_cap(usage: dict, session: Session) -> None: + """Track usage and enforce token cap.""" + session.track_usage(usage) + if token_cap <= 0: + return + + effective = session.token_state["last_effective"] + pct = effective / token_cap * 100 + sid = session.id + + if effective > token_cap: + session.token_state["blocked"] = True + print( + f"{_RED} [{sid}] TOKEN CAP EXCEEDED: {effective:,} / {token_cap:,} " + f"({pct:.0f}%) — next request will be blocked{_RESET}", + file=sys.stderr, + ) + + def _update_fidelity_pressure(usage: dict, session: Session) -> None: + """Update FidelityManager window understanding from actual API usage. + + After receiving the response, we know the real input token count. + Update the fidelity manager's window_size understanding and schedule + degradation if pressure is above NORMAL for the next turn. + """ + input_tokens = usage.get("input_tokens", 0) + if input_tokens <= 0: + return + + fm = session.fidelity_manager + turn = session.token_state.get("turn", 0) + + # The FidelityManager tracks its own token budget via registered objects. + # Here we use the real API token count to check if we need proactive degradation. + # If real usage exceeds the fidelity window threshold, trigger degradation now + # so the NEXT turn benefits from reduced content. + pressure_ratio = input_tokens / fm.window_size if fm.window_size > 0 else 1.0 + + if pressure_ratio >= fm.threshold_caution: + transitions = fm.degrade(turn) + if transitions: + zone = fm.current_pressure() + for _obj_id, old_level, new_level in transitions: + session.benchmark.fidelity.record_transition(old_level.value, new_level.value) + + # Schedule async background summarization for degraded objects. + # This runs without blocking the response stream. + if helper_llm is not None: + try: + loop = asyncio.get_running_loop() + loop.create_task( + _generate_degradation_summaries(transitions, session, helper_llm) + ) + except RuntimeError: + # No running event loop (e.g. in sync test context) — + # run in a thread pool as a best-effort fallback + try: + with concurrent.futures.ThreadPoolExecutor(max_workers=1) as pool: + pool.submit( + asyncio.run, + _generate_degradation_summaries( + transitions, session, helper_llm + ), + ) + except Exception: + _summarization_logger.warning( + "Failed to schedule background summarization", + exc_info=True, + ) + + def _display_turn_status(usage: dict, session: Session) -> None: + """Post-response status line with cache hit rate.""" + sid = session.id + ts = session.token_state + input_tok = usage.get("input_tokens", 0) + cache_read = usage.get("cache_read_input_tokens", 0) + cache_create = usage.get("cache_creation_input_tokens", 0) + effective = input_tok + cache_read + cache_create + if effective == 0: + return + + # Record token metrics for benchmarking + session.benchmark.tokens.record_turn( + input_tokens=input_tok, + effective_tokens=effective, + cache_read=cache_read, + cache_create=cache_create, + incoming_bytes=0, # filled by telemetry record_request + outgoing_bytes=0, + ) + + cap_str = "" + if token_cap > 0: + pct = effective / token_cap * 100 + cap_str = f"/{token_cap // 1000}k ({pct:.0f}%)" + + # Cache hit rate + cache_read = usage.get("cache_read_input_tokens", 0) + cache_create = usage.get("cache_creation_input_tokens", 0) + cache_total = cache_read + cache_create + cache_str = "" + if cache_total > 0: + cache_pct = cache_read / cache_total * 100 + cache_str = f" | cache {cache_pct:.0f}%" + + ps = session.page_store + ev_str = "" + if ps.unique_evictions > 0 or ps.gc_count > 0: + pin_str = f" pin {len(ps._pinned)}" if ps._pinned else "" + ev_str = ( + f" | ev {ps.unique_evictions} gc {ps.gc_count}{pin_str}" + f" | faults {len(ps.faults)}/{ps.unique_evictions}" + ) + + print( + f" [{sid}] [Turn {ts['turn']}] {effective:,} tok{cap_str}{cache_str}{ev_str}", + file=sys.stderr, + ) + + # ── Request handling ───────────────────────────────────────────── + + def _handle_provider_request( + provider: str, + endpoint: str, + payload: dict[str, Any], + query_string: str = "", + ): + payload = dict(payload) + forwarded_headers = payload.pop("_headers", {}) + + # Session management + session = sessions.get(payload) + + # Lazy-inject ContextAssembler once helper_llm is known + if session.context_assembler is None: + from mnemosyne.context_assembler import ContextAssembler + + session.context_assembler = ContextAssembler(session.object_store, helper_llm) + + # Token cap enforcement + if token_cap > 0 and session.token_state.get("blocked"): + raise HTTPException( + status_code=429, + detail=f"Token cap exceeded ({session.token_state['last_effective']:,}/{token_cap:,})", + ) + + # Reject inbound cleanup tag injection + # tag_error = check_inbound_for_injected_tags(payload) + # if tag_error: + # print(f" {_RED}[{session.id}] {tag_error}{_RESET}", file=sys.stderr) + # raise HTTPException(status_code=400, detail=tag_error) + + # Cleanup now runs inside _preprocess (before manifest injection) + + # Pre-process: system status, block labeling + if endpoint == "messages": + payload = _preprocess(payload, session) + + # Optional provider-level model override for cost-controlled runs. + if provider == "anthropic" and anthropic_model_override: + payload["model"] = anthropic_model_override + elif provider == "openai" and openai_model_override: + payload["model"] = openai_model_override + + adapter = app.state.adapters[provider] + client: httpx.Client = app.state.clients[provider] + + request_id = str(uuid.uuid4()) + started = time.perf_counter() + incoming_bytes = len(json.dumps(payload, default=str).encode("utf-8")) + session_id = session.id + + req = adapter.normalize_request(payload) + req = pipeline.run(req) + duplication_score = _duplication_score(req) + outgoing_body = adapter.denormalize_request(req) + if endpoint == "count_tokens": + outgoing_body.pop("stream", None) + outgoing_body.pop("max_tokens", None) + outgoing_bytes = len(json.dumps(outgoing_body, default=str).encode("utf-8")) + + upstream_path = adapter.upstream_path(req, endpoint=endpoint) + if query_string: + upstream_path = f"{upstream_path}?{query_string}" + headers = { + k: v + for k, v in forwarded_headers.items() + if k.lower() not in {"host", "content-length", "transfer-encoding"} + } + + # OAuth: replace dummy x-api-key with Bearer token when available. + # Anthropic requires the "oauth-2025-04-20" beta header and a + # Claude-CLI-style user-agent for OAuth Bearer auth to be accepted. + # See: treplay/server.mjs proxyAnthropicContinuation() + oauth_token = os.environ.get("ANTHROPIC_AUTH_TOKEN") + if oauth_token and provider == "anthropic": + headers.pop("x-api-key", None) + headers.pop("X-Api-Key", None) + headers["authorization"] = f"Bearer {oauth_token}" + # Beta flag is REQUIRED — without it, Anthropic returns + # "OAuth authentication is currently not supported" + existing_beta = headers.get("anthropic-beta", "") + oauth_beta = "oauth-2025-04-20" + if oauth_beta not in existing_beta: + if existing_beta: + headers["anthropic-beta"] = f"{existing_beta},{oauth_beta}" + else: + headers["anthropic-beta"] = oauth_beta + headers["user-agent"] = "claude-cli/2.1.2 (external, cli)" + + if req.stream: + + def generate(): + status_code = 599 + bytes_out = 0 + chunk_count = 0 + stream_error = False + sse_buffer = bytearray() + usage: dict[str, Any] = {} + try: + with client.stream( + "POST", upstream_path, json=outgoing_body, headers=headers + ) as resp: + status_code = resp.status_code + if status_code >= 400: + body = resp.read() + bytes_out = len(body) + yield body + else: + for chunk in resp.iter_bytes(): + bytes_out += len(chunk) + chunk_count += 1 + _inspect_sse_chunk( + chunk, + buffer=sse_buffer, + emit_event=emit_event, + request_id=request_id, + session_id=session_id, + provider=provider, + usage_accumulator=usage, + ) + yield chunk + except Exception as e: + stream_error = True + emit_event( + "stream_error", + request_id=request_id, + session_id=session_id, + provider=provider, + error=str(e), + ) + raise + finally: + if stream_error: + status_code = 599 + # Post-response: track usage, display status + if usage: + _check_token_cap(usage, session) + _display_turn_status(usage, session) + _update_fidelity_pressure(usage, session) + emit_event( + "response_observed", + request_id=request_id, + session_id=session_id, + provider=provider, + endpoint=endpoint, + status=status_code, + response_bytes=bytes_out, + chunk_count=chunk_count, + usage=usage, + ) + latency_ms = (time.perf_counter() - started) * 1000 + telemetry.record_request( + request_id=request_id, + session_id=session_id, + provider=provider, + model=req.model, + status=status_code, + incoming_bytes=incoming_bytes, + outgoing_bytes=outgoing_bytes, + latency_ms=latency_ms, + streaming=True, + duplication_score=duplication_score, + usage=usage, + messages_full=payload.get("messages", []), + ) + + return StreamingResponse(generate(), media_type="text/event-stream") + + # Non-streaming + try: + resp = client.post(upstream_path, json=outgoing_body, headers=headers) + except httpx.HTTPError as e: + emit_event( + "provider_error", + request_id=request_id, + session_id=session_id, + provider=provider, + error=str(e), + ) + raise HTTPException(status_code=502, detail=str(e)) + + usage: dict[str, Any] = {} + try: + parsed = resp.json() + if isinstance(parsed, dict): + u = parsed.get("usage", {}) + if isinstance(u, dict): + usage = u + except Exception: + usage = {} + + if usage: + _check_token_cap(usage, session) + _display_turn_status(usage, session) + _update_fidelity_pressure(usage, session) + + emit_event( + "response_observed", + request_id=request_id, + session_id=session_id, + provider=provider, + endpoint=endpoint, + status=resp.status_code, + response_bytes=len(resp.content), + chunk_count=0, + usage=usage, + ) + + latency_ms = (time.perf_counter() - started) * 1000 + telemetry.record_request( + request_id=request_id, + session_id=session_id, + provider=provider, + model=req.model, + status=resp.status_code, + incoming_bytes=incoming_bytes, + outgoing_bytes=outgoing_bytes, + latency_ms=latency_ms, + streaming=False, + duplication_score=duplication_score, + usage=usage, + messages_full=payload.get("messages", []), + ) + return Response( + content=resp.content, status_code=resp.status_code, headers=_copy_headers(resp.headers) + ) + + def _payload_with_headers(req: Request, payload: dict[str, Any]) -> dict[str, Any]: + payload = dict(payload) + payload["_headers"] = dict(req.headers) + return payload + + @app.post("/v1/messages") + @app.post("/messages") + async def anthropic_messages(request: Request): + payload = await request.json() + return _handle_provider_request( + "anthropic", + "messages", + _payload_with_headers(request, payload), + query_string=request.url.query, + ) + + @app.post("/v1/messages/count_tokens") + @app.post("/messages/count_tokens") + async def anthropic_count_tokens(request: Request): + payload = await request.json() + payload = _payload_with_headers(request, payload) + payload["stream"] = False + return _handle_provider_request( + "anthropic", + "count_tokens", + payload, + query_string=request.url.query, + ) + + @app.post("/v1/chat/completions") + async def openai_chat_completions(request: Request): + payload = await request.json() + return _handle_provider_request( + "openai", + "chat_completions", + _payload_with_headers(request, payload), + query_string=request.url.query, + ) + + # ── Catch-all proxy for unhandled /v1/ routes ────────────────── + # Claude Code may hit /v1/models or other endpoints for validation. + # Forward anything we don't explicitly handle to the real API. + @app.api_route("/v1/{path:path}", methods=["GET", "POST", "PUT", "DELETE"]) + async def proxy_passthrough(request: Request, path: str): + import httpx + + api_key = request.headers.get("x-api-key", "") + headers = { + "anthropic-version": request.headers.get("anthropic-version", "2023-06-01"), + "content-type": request.headers.get("content-type", "application/json"), + } + # OAuth: use Bearer token if available, otherwise x-api-key + oauth_token = os.environ.get("ANTHROPIC_AUTH_TOKEN") + if oauth_token: + headers["authorization"] = f"Bearer {oauth_token}" + existing_beta = headers.get("anthropic-beta", "") + oauth_beta = "oauth-2025-04-20" + if oauth_beta not in existing_beta: + headers["anthropic-beta"] = ( + f"{existing_beta},{oauth_beta}" if existing_beta else oauth_beta + ) + headers["user-agent"] = "claude-cli/2.1.2 (external, cli)" + else: + headers["x-api-key"] = api_key + url = f"https://api.anthropic.com/v1/{path}" + if request.url.query: + url += f"?{request.url.query}" + async with httpx.AsyncClient(timeout=30) as client: + if request.method == "GET": + resp = await client.get(url, headers=headers) + else: + body = await request.body() + resp = await client.request(request.method, url, headers=headers, content=body) + return Response( + content=resp.content, + status_code=resp.status_code, + headers=dict(resp.headers), + ) + + return app + + +def _run_server_in_thread(app: FastAPI, port: int) -> tuple[uvicorn.Server, threading.Thread]: + config = uvicorn.Config(app, host="127.0.0.1", port=port, log_level="warning") + server = uvicorn.Server(config=config) + + t = threading.Thread(target=server.run, daemon=True) + t.start() + + # Wait up to 10s for bind. + deadline = time.time() + 10.0 + while time.time() < deadline: + if server.started: + return server, t + time.sleep(0.05) + raise RuntimeError("gateway server did not start") + + +def main() -> None: + parser = argparse.ArgumentParser(description="Mnemosyne context memory proxy") + parser.add_argument("--login", action="store_true", help="OAuth login for Anthropic API") + parser.add_argument("--claude", action="store_true", help="Launch Claude CLI against gateway") + parser.add_argument("--codex", action="store_true", help="Launch Codex CLI against gateway") + parser.add_argument("--gemini", action="store_true", help="Gemini mode (stub in v1)") + parser.add_argument( + "--no-launch", action="store_true", help="Run gateway as persistent service" + ) + parser.add_argument("--port", type=int, default=0, help="Gateway port (0=random)") + parser.add_argument( + "--log-dir", type=Path, default=Path("logs"), help="Telemetry log directory" + ) + parser.add_argument( + "--hydration-window", default="24h", help="Dashboard hydration window (e.g. 24h, 7d)" + ) + parser.add_argument("--anthropic-upstream", default="https://api.anthropic.com") + parser.add_argument("--openai-upstream", default="https://api.openai.com") + parser.add_argument("--token-cap", type=int, default=0, help="Token cap (0=unlimited)") + parser.add_argument( + "--anthropic-model-override", + default=None, + help="Force Anthropic model ID for all Anthropic requests", + ) + parser.add_argument( + "--openai-model-override", + default=None, + help="Force OpenAI model ID for all OpenAI requests", + ) + parser.add_argument("--disable-paging", action="store_true") + parser.add_argument("--disable-trim", action="store_true") + parser.add_argument("--min-evict-size", type=int, default=500) + parser.add_argument( + "--benchmark", action="store_true", help="Show benchmark metrics from running gateway" + ) + parser.add_argument( + "--session", type=str, default=None, help="Session ID for benchmark (with --benchmark)" + ) + parser.add_argument("--json-output", action="store_true", help="Output benchmark as raw JSON") + parser.add_argument( + "cli_args", nargs=argparse.REMAINDER, help="Args passed to launched CLI after '--'" + ) + args = parser.parse_args() + + if args.login: + from mnemosyne.oauth import login_interactive + + login_interactive() + return + + if args.benchmark: + from mnemosyne.benchmark_cli import run_benchmark_cli + + # When used with --benchmark, treat --port 0 as 8080 + benchmark_port = args.port if args.port != 0 else 8080 + run_benchmark_cli( + port=benchmark_port, + session_id=args.session, + json_output=args.json_output, + ) + return + + selected_modes = [ + m + for m, enabled in (("claude", args.claude), ("codex", args.codex), ("gemini", args.gemini)) + if enabled + ] + if not args.no_launch and len(selected_modes) != 1: + parser.error( + "choose exactly one launch target: --claude | --codex | --gemini, or pass --no-launch" + ) + + try: + hydration_seconds = parse_duration(args.hydration_window) + except ValueError as e: + parser.error(str(e)) + + port = args.port if args.port != 0 else find_free_port() + + app = create_app( + log_dir=args.log_dir, + anthropic_upstream=args.anthropic_upstream, + openai_upstream=args.openai_upstream, + hydration_window_seconds=hydration_seconds, + enable_paging=not args.disable_paging, + enable_trim=not args.disable_trim, + min_evict_size=args.min_evict_size, + token_cap=args.token_cap, + anthropic_model_override=args.anthropic_model_override, + openai_model_override=args.openai_model_override, + ) + + if args.no_launch: + print(f"Gateway listening on http://127.0.0.1:{port}", file=sys.stderr) + uvicorn.run(app, host="127.0.0.1", port=port) + return + + mode = selected_modes[0] + if mode == "gemini": + print("Gemini adapter is not enabled in v1. Use --claude or --codex.", file=sys.stderr) + raise SystemExit(2) + + server, thread = _run_server_in_thread(app, port) + print(f"Gateway listening on http://127.0.0.1:{port}", file=sys.stderr) + + extra = args.cli_args + if extra and extra[0] == "--": + extra = extra[1:] + + rc = 1 + try: + rc = launch(LaunchSpec(mode=mode, port=port, extra_args=extra)) + finally: + server.should_exit = True + thread.join(timeout=5) + raise SystemExit(rc) + + +if __name__ == "__main__": + main() diff --git a/src/mnemosyne/launcher.py b/src/mnemosyne/launcher.py new file mode 100644 index 0000000..633c55a --- /dev/null +++ b/src/mnemosyne/launcher.py @@ -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 diff --git a/src/mnemosyne/message_ops.py b/src/mnemosyne/message_ops.py new file mode 100644 index 0000000..2123d67 --- /dev/null +++ b/src/mnemosyne/message_ops.py @@ -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"\s*.*?\s*", 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" str | None: + """Scan inbound messages for injected 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 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 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 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 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 block, Pichay is asking you to advise on " + "memory management. Respond with a block using this format:\n\n" + "\n" + '\n' + '\n' + "\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' ' + ) + if tensor_lines or pruned_count > 0: + manifest_parts = ["\n\n"] + # Feedback: what happened last turn (closed-loop) + if last_cleanup_stats: + manifest_parts.append( + f" {_escape_xml_attr(last_cleanup_stats)}" + f"\n" + ) + if pruned_count > 0: + manifest_parts.append( + f' \n' + ) + if tensor_lines: + manifest_parts.append( + f" \n" + + "\n".join(tensor_lines[:15]) # cap at 15 to avoid bloat + + "\n \n" + ) + manifest_parts.append("") + 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 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( + "\nContext pressure is high. " + "Review the manifest above. Which tensors can be " + "released? Respond in a block with " + "release decisions before your normal response." + "" + ) + + 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} diff --git a/src/mnemosyne/message_store.py b/src/mnemosyne/message_store.py new file mode 100644 index 0000000..43bc5a3 --- /dev/null +++ b/src/mnemosyne/message_store.py @@ -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 diff --git a/src/mnemosyne/mnemosyne_config.py b/src/mnemosyne/mnemosyne_config.py new file mode 100644 index 0000000..ed207c7 --- /dev/null +++ b/src/mnemosyne/mnemosyne_config.py @@ -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")), + ) diff --git a/src/mnemosyne/pager.py b/src/mnemosyne/pager.py new file mode 100644 index 0000000..bc77c18 --- /dev/null +++ b/src/mnemosyne/pager.py @@ -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 ( 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) diff --git a/src/mnemosyne/providers/__init__.py b/src/mnemosyne/providers/__init__.py new file mode 100644 index 0000000..3d610d6 --- /dev/null +++ b/src/mnemosyne/providers/__init__.py @@ -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(), + } diff --git a/src/mnemosyne/providers/anthropic.py b/src/mnemosyne/providers/anthropic.py new file mode 100644 index 0000000..613ef46 --- /dev/null +++ b/src/mnemosyne/providers/anthropic.py @@ -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" diff --git a/src/mnemosyne/providers/base.py b/src/mnemosyne/providers/base.py new file mode 100644 index 0000000..8a02fec --- /dev/null +++ b/src/mnemosyne/providers/base.py @@ -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: ... diff --git a/src/mnemosyne/providers/openai.py b/src/mnemosyne/providers/openai.py new file mode 100644 index 0000000..c588b79 --- /dev/null +++ b/src/mnemosyne/providers/openai.py @@ -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" diff --git a/src/mnemosyne/tags.py b/src/mnemosyne/tags.py new file mode 100644 index 0000000..a200911 --- /dev/null +++ b/src/mnemosyne/tags.py @@ -0,0 +1,199 @@ +"""Cleanup tag parser — extracts cooperative memory operations from model text. + +The model emits 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: + + 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 + +""" + +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 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 ... blocks (non-greedy) +_TAG_PATTERN = re.compile( + r"\s*(.*?)\s*", + 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 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 ... blocks +_YUYAY_RESPONSE_PATTERN = re.compile( + r"\s*(.*?)\s*", + re.DOTALL, +) + +# Match structured eviction decisions: +_YUYAY_RELEASE = re.compile(r' CleanupOps: + """Extract memory operations from blocks. + + The model responds to with structured eviction + decisions. These are converted to CleanupOps for execution + through the same pipeline as tags. + + Supports two formats: + - Structured XML: + - 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 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 ... 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 diff --git a/src/mnemosyne/trimmer.py b/src/mnemosyne/trimmer.py new file mode 100644 index 0000000..ce6defa --- /dev/null +++ b/src/mnemosyne/trimmer.py @@ -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 "" 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 "" 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"(.*?)", + 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 "" + "\n".join(output_lines) + "" + + new_text = re.sub( + r"(.*?)", + _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"(.*?)", + 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 "" 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 "" 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()