From d26c56c2f0a5bf1f2dd47af9d37d886f00dbc99b Mon Sep 17 00:00:00 2001 From: Joey Yakimowich-Payne Date: Fri, 13 Mar 2026 11:40:56 -0600 Subject: [PATCH] feat: add multi-fidelity compression engine 5-level fidelity manager (L0-Full to L4-Evicted) with helper LLM (Haiku 4.5) for intelligent summarization during degradation. Ultraworked with [Sisyphus](https://github.com/code-yeongyu/oh-my-opencode) Co-authored-by: Sisyphus --- src/mnemosyne/fidelity.py | 615 +++++++++++++++++++++++ src/mnemosyne/helper_llm.py | 442 +++++++++++++++++ tests/test_fidelity.py | 861 +++++++++++++++++++++++++++++++++ tests/test_gateway_fidelity.py | 513 ++++++++++++++++++++ tests/test_helper_llm.py | 572 ++++++++++++++++++++++ 5 files changed, 3003 insertions(+) create mode 100644 src/mnemosyne/fidelity.py create mode 100644 src/mnemosyne/helper_llm.py create mode 100644 tests/test_fidelity.py create mode 100644 tests/test_gateway_fidelity.py create mode 100644 tests/test_helper_llm.py diff --git a/src/mnemosyne/fidelity.py b/src/mnemosyne/fidelity.py new file mode 100644 index 0000000..d68fced --- /dev/null +++ b/src/mnemosyne/fidelity.py @@ -0,0 +1,615 @@ +"""Multi-fidelity state machine for semantic objects. + +Manages the fidelity ladder (L0-L4) for each semantic object in a session. +Pressure zones drive degradation; access patterns drive upgrades; faults +drive pinning. This is the core scheduling logic — it decides what stays +in context at what resolution. + +Fidelity levels: + L0: Full content, no compression + L1: Detailed summary, ~30% of original size, with declared losses + L2: Compact summary, ~5% of original size, with accumulated losses + L3: Metadata stub, ~50-100 tokens, one-line description + L4: Evicted, not in context at all + +Pressure zones (configurable thresholds): + Normal < 50% — no action + Caution 50-70% — degrade oldest L0 → L1 + Warning 70-85% — degrade L0→L1, L1→L2, oldest L2→L3 + Critical 85-95% — aggressive L2+→L3, evict L3→L4 + Emergency >95% — force-evict everything except last 2 user turns + system prompt +""" + +from __future__ import annotations + +import uuid +from dataclasses import dataclass, field +from enum import IntEnum + + +class FidelityLevel(IntEnum): + """Fidelity levels for semantic objects, ordered from highest to lowest.""" + + L0 = 0 # Full content, no compression + L1 = 1 # Detailed summary, ~30% of original size + L2 = 2 # Compact summary, ~5% of original size + L3 = 3 # Metadata stub, ~50-100 tokens + L4 = 4 # Evicted, not in context + + +class PressureZone(IntEnum): + """Context window pressure zones, ordered by severity.""" + + NORMAL = 0 # < 50% + CAUTION = 1 # 50-70% + WARNING = 2 # 70-85% + CRITICAL = 3 # 85-95% + EMERGENCY = 4 # > 95% + + +# Valid object types per ARCHITECTURE.md §5.1 +VALID_OBJECT_TYPES = frozenset( + { + "conversation_phase", + "design_decision", + "debugging_session", + "file_context", + "tool_result", + "plan", + "error_context", + "external_reference", + } +) + + +@dataclass +class SemanticObject: + """A tracked semantic object with multi-fidelity representations. + + Each object carries its full content (L0) and optional compressed + representations (L1-L3). The fidelity manager decides which level + is currently active in the context window. + """ + + id: str + object_type: str + + # Content at each fidelity level + content_full: str # L0 + summary_detailed: str | None = None # L1 + summary_compact: str | None = None # L2 + stub: str | None = None # L3 + + # Declared losses — what information was dropped at each level + losses_l1: list[str] = field(default_factory=list) + losses_l2: list[str] = field(default_factory=list) + + # Queryability metadata + can_answer: list[str] = field(default_factory=list) + fault_when: list[str] = field(default_factory=list) + key_entities: list[str] = field(default_factory=list) + + # State + current_fidelity: FidelityLevel = FidelityLevel.L0 + pinned: bool = False + pin_until_turn: int | None = None + created_at_turn: int = 0 + last_accessed_turn: int = 0 + fault_count: int = 0 + + # Token estimates (len(content) / 4 heuristic) + token_count_l0: int = 0 + token_count_l1: int | None = None + token_count_l2: int | None = None + token_count_l3: int = 25 # Stubs are typically ~25 tokens + + def tokens_at(self, level: FidelityLevel) -> int: + """Return estimated token count at the given fidelity level.""" + if level == FidelityLevel.L0: + return self.token_count_l0 + if level == FidelityLevel.L1: + return self.token_count_l1 if self.token_count_l1 is not None else 0 + if level == FidelityLevel.L2: + return self.token_count_l2 if self.token_count_l2 is not None else 0 + if level == FidelityLevel.L3: + return self.token_count_l3 + # L4: evicted, zero tokens in context + return 0 + + @property + def current_tokens(self) -> int: + """Token count at the object's current fidelity level.""" + return self.tokens_at(self.current_fidelity) + + +def _estimate_tokens(text: str | None) -> int: + """Estimate token count using ~4 chars per token heuristic.""" + if text is None: + return 0 + return max(1, len(text) // 4) + + +def make_object( + *, + object_type: str, + content_full: str, + created_at_turn: int = 0, + summary_detailed: str | None = None, + summary_compact: str | None = None, + stub: str | None = None, + losses_l1: list[str] | None = None, + losses_l2: list[str] | None = None, + can_answer: list[str] | None = None, + fault_when: list[str] | None = None, + key_entities: list[str] | None = None, +) -> SemanticObject: + """Factory for creating a SemanticObject with auto-computed token estimates. + + Preferred over constructing SemanticObject directly — handles ID + generation and token estimation. + """ + obj = SemanticObject( + id=uuid.uuid4().hex[:16], + object_type=object_type, + content_full=content_full, + summary_detailed=summary_detailed, + summary_compact=summary_compact, + stub=stub, + losses_l1=losses_l1 or [], + losses_l2=losses_l2 or [], + can_answer=can_answer or [], + fault_when=fault_when or [], + key_entities=key_entities or [], + current_fidelity=FidelityLevel.L0, + created_at_turn=created_at_turn, + last_accessed_turn=created_at_turn, + token_count_l0=_estimate_tokens(content_full), + token_count_l1=_estimate_tokens(summary_detailed) if summary_detailed else None, + token_count_l2=_estimate_tokens(summary_compact) if summary_compact else None, + token_count_l3=_estimate_tokens(stub) if stub else 25, + ) + return obj + + +class FidelityManager: + """Manages fidelity levels for all semantic objects in a session. + + Core responsibilities: + - Register new objects (start at L0) + - Degrade objects under pressure (oldest-first) + - Upgrade objects on access + - Pin objects after faults + - Track token budget + """ + + def __init__(self, window_size: int = 200_000) -> None: + """Initialize the fidelity manager. + + Args: + window_size: Total context window size in tokens. Pressure + zones are computed as percentages of this value. + """ + self.window_size = window_size + self._objects: dict[str, SemanticObject] = {} + + # Pressure zone thresholds (fraction of window_size) + self.threshold_caution = 0.50 + self.threshold_warning = 0.70 + self.threshold_critical = 0.85 + self.threshold_emergency = 0.95 + + def register_object(self, obj: SemanticObject) -> str: + """Register a new semantic object. Returns its ID. + + New objects always start at L0 (full content). Token estimates + are recomputed if they look unset. + """ + # Ensure token estimates are populated + if obj.token_count_l0 == 0 and obj.content_full: + obj.token_count_l0 = _estimate_tokens(obj.content_full) + if obj.token_count_l1 is None and obj.summary_detailed: + obj.token_count_l1 = _estimate_tokens(obj.summary_detailed) + if obj.token_count_l2 is None and obj.summary_compact: + obj.token_count_l2 = _estimate_tokens(obj.summary_compact) + if obj.stub and obj.token_count_l3 == 25: + obj.token_count_l3 = _estimate_tokens(obj.stub) + + self._objects[obj.id] = obj + return obj.id + + def get_object(self, object_id: str) -> SemanticObject | None: + """Retrieve a semantic object by ID, or None if not found.""" + return self._objects.get(object_id) + + def total_tokens(self) -> int: + """Sum of tokens across all objects at their current fidelity level.""" + return sum(obj.current_tokens for obj in self._objects.values()) + + def current_pressure(self) -> PressureZone: + """Calculate the current pressure zone from total token usage. + + Returns the pressure zone enum value based on the ratio of + total tokens to window size. + """ + if self.window_size <= 0: + return PressureZone.EMERGENCY + + ratio = self.total_tokens() / self.window_size + + if ratio >= self.threshold_emergency: + return PressureZone.EMERGENCY + if ratio >= self.threshold_critical: + return PressureZone.CRITICAL + if ratio >= self.threshold_warning: + return PressureZone.WARNING + if ratio >= self.threshold_caution: + return PressureZone.CAUTION + return PressureZone.NORMAL + + def _sorted_by_age(self, objects: list[SemanticObject]) -> list[SemanticObject]: + """Sort objects oldest-first by last_accessed_turn, then created_at_turn.""" + return sorted(objects, key=lambda o: (o.last_accessed_turn, o.created_at_turn)) + + def _degrade_one(self, obj: SemanticObject) -> FidelityLevel | None: + """Degrade an object by one fidelity level. Returns new level, or None if already L4.""" + old = obj.current_fidelity + if old >= FidelityLevel.L4: + return None + + # Check if the next level has content available + new = FidelityLevel(old + 1) + + # For L0→L1: need summary_detailed (or degrade anyway — content just won't be shown) + # For L1→L2: need summary_compact + # For L2→L3: need stub + # For L3→L4: always possible (eviction) + obj.current_fidelity = new + return new + + def degrade(self, current_turn: int) -> list[tuple[str, FidelityLevel, FidelityLevel]]: + """Degrade objects based on current pressure zone. + + Walks objects oldest-first and degrades based on the pressure zone: + - Normal: no action + - Caution: degrade oldest L0 → L1 + - Warning: degrade L0→L1, L1→L2, oldest L2→L3 + - Critical: aggressive L2+→L3, evict L3→L4 + - Emergency: force-evict everything except pinned objects + + Args: + current_turn: Current conversation turn number. + + Returns: + List of (object_id, old_level, new_level) for each degradation. + """ + transitions: list[tuple[str, FidelityLevel, FidelityLevel]] = [] + zone = self.current_pressure() + + if zone == PressureZone.NORMAL: + return transitions + + # Expire stale pins + self._expire_pins(current_turn) + + if zone == PressureZone.CAUTION: + transitions.extend(self._degrade_caution(current_turn)) + elif zone == PressureZone.WARNING: + transitions.extend(self._degrade_warning(current_turn)) + elif zone == PressureZone.CRITICAL: + transitions.extend(self._degrade_critical(current_turn)) + elif zone == PressureZone.EMERGENCY: + transitions.extend(self._degrade_emergency(current_turn)) + + return transitions + + def _expire_pins(self, current_turn: int) -> None: + """Unpin objects whose pin duration has expired.""" + for obj in self._objects.values(): + if obj.pinned and obj.pin_until_turn is not None: + if current_turn >= obj.pin_until_turn: + obj.pinned = False + obj.pin_until_turn = None + + def _degradable(self, obj: SemanticObject) -> bool: + """Check if an object can be degraded (not pinned, not already L4).""" + return not obj.pinned and obj.current_fidelity < FidelityLevel.L4 + + def _degrade_caution( + self, + current_turn: int, + ) -> list[tuple[str, FidelityLevel, FidelityLevel]]: + """Caution zone: degrade oldest L0 objects to L1.""" + transitions: list[tuple[str, FidelityLevel, FidelityLevel]] = [] + candidates = [ + o + for o in self._objects.values() + if o.current_fidelity == FidelityLevel.L0 and self._degradable(o) + ] + candidates = self._sorted_by_age(candidates) + + for obj in candidates: + old = obj.current_fidelity + new = self._degrade_one(obj) + if new is not None: + transitions.append((obj.id, old, new)) + # Re-check pressure after each degradation + if self.current_pressure() <= PressureZone.NORMAL: + break + + return transitions + + def _degrade_warning( + self, + current_turn: int, + ) -> list[tuple[str, FidelityLevel, FidelityLevel]]: + """Warning zone: degrade L0→L1, L1→L2, oldest L2→L3.""" + transitions: list[tuple[str, FidelityLevel, FidelityLevel]] = [] + + # First pass: L0 → L1 + l0_objs = self._sorted_by_age( + [ + o + for o in self._objects.values() + if o.current_fidelity == FidelityLevel.L0 and self._degradable(o) + ] + ) + for obj in l0_objs: + old = obj.current_fidelity + new = self._degrade_one(obj) + if new is not None: + transitions.append((obj.id, old, new)) + if self.current_pressure() <= PressureZone.NORMAL: + return transitions + + # Second pass: L1 → L2 + l1_objs = self._sorted_by_age( + [ + o + for o in self._objects.values() + if o.current_fidelity == FidelityLevel.L1 and self._degradable(o) + ] + ) + for obj in l1_objs: + old = obj.current_fidelity + new = self._degrade_one(obj) + if new is not None: + transitions.append((obj.id, old, new)) + if self.current_pressure() <= PressureZone.NORMAL: + return transitions + + # Third pass: oldest L2 → L3 + l2_objs = self._sorted_by_age( + [ + o + for o in self._objects.values() + if o.current_fidelity == FidelityLevel.L2 and self._degradable(o) + ] + ) + for obj in l2_objs: + old = obj.current_fidelity + new = self._degrade_one(obj) + if new is not None: + transitions.append((obj.id, old, new)) + if self.current_pressure() <= PressureZone.NORMAL: + return transitions + + return transitions + + def _degrade_critical( + self, + current_turn: int, + ) -> list[tuple[str, FidelityLevel, FidelityLevel]]: + """Critical zone: aggressive L2+→L3, evict L3→L4.""" + transitions: list[tuple[str, FidelityLevel, FidelityLevel]] = [] + + # First: degrade everything L0/L1/L2 down aggressively + for target_level in (FidelityLevel.L0, FidelityLevel.L1, FidelityLevel.L2): + candidates = self._sorted_by_age( + [ + o + for o in self._objects.values() + if o.current_fidelity == target_level and self._degradable(o) + ] + ) + for obj in candidates: + # Degrade all the way to L3 + while obj.current_fidelity < FidelityLevel.L3 and self._degradable(obj): + old = obj.current_fidelity + new = self._degrade_one(obj) + if new is not None: + transitions.append((obj.id, old, new)) + if self.current_pressure() <= PressureZone.NORMAL: + return transitions + + # Then: evict L3 → L4 + l3_objs = self._sorted_by_age( + [ + o + for o in self._objects.values() + if o.current_fidelity == FidelityLevel.L3 and self._degradable(o) + ] + ) + for obj in l3_objs: + old = obj.current_fidelity + new = self._degrade_one(obj) + if new is not None: + transitions.append((obj.id, old, new)) + if self.current_pressure() <= PressureZone.NORMAL: + return transitions + + return transitions + + def _degrade_emergency( + self, + current_turn: int, + ) -> list[tuple[str, FidelityLevel, FidelityLevel]]: + """Emergency zone: force-evict everything except pinned objects. + + In a real system, "last 2 user turns + system prompt" would be + handled by the context assembler. Here we evict all non-pinned + objects to L4. + """ + transitions: list[tuple[str, FidelityLevel, FidelityLevel]] = [] + + all_objs = self._sorted_by_age( + [ + o + for o in self._objects.values() + if o.current_fidelity < FidelityLevel.L4 and self._degradable(o) + ] + ) + for obj in all_objs: + while obj.current_fidelity < FidelityLevel.L4: + old = obj.current_fidelity + new = self._degrade_one(obj) + if new is not None: + transitions.append((obj.id, old, new)) + + return transitions + + def upgrade( + self, + object_id: str, + target: FidelityLevel, + current_turn: int, + ) -> bool: + """Upgrade an object's fidelity level (e.g., on model access). + + Can only upgrade to a higher fidelity (lower numeric level). + Updates last_accessed_turn as a side effect. + + Args: + object_id: The object to upgrade. + target: The target fidelity level (must be < current level). + current_turn: Current conversation turn. + + Returns: + True if the upgrade was performed, False otherwise. + """ + obj = self._objects.get(object_id) + if obj is None: + return False + + if target >= obj.current_fidelity: + return False # Not an upgrade + + # Check that the target level has content available + if target == FidelityLevel.L0 and not obj.content_full: + return False + if target == FidelityLevel.L1 and not obj.summary_detailed: + return False + if target == FidelityLevel.L2 and not obj.summary_compact: + return False + + obj.current_fidelity = target + obj.last_accessed_turn = current_turn + return True + + def record_fault( + self, + object_id: str, + current_turn: int, + pin_duration: int = 5, + ) -> None: + """Record a fault for an object and pin it at current fidelity. + + After a fault (model needed content that was degraded), pin the + object so it won't be degraded again for `pin_duration` turns. + + Args: + object_id: The faulted object. + current_turn: Current conversation turn. + pin_duration: Number of turns to pin the object. + """ + obj = self._objects.get(object_id) + if obj is None: + return + + obj.fault_count += 1 + obj.pinned = True + obj.pin_until_turn = current_turn + pin_duration + obj.last_accessed_turn = current_turn + + def mark_accessed(self, object_id: str, current_turn: int) -> None: + """Update an object's last_accessed_turn. + + Called when the model references an object, even without + upgrading its fidelity level. + + Args: + object_id: The accessed object. + current_turn: Current conversation turn. + """ + obj = self._objects.get(object_id) + if obj is None: + return + obj.last_accessed_turn = current_turn + + def eviction_candidates(self, current_turn: int) -> list[SemanticObject]: + """Return objects sorted by eviction priority (most evictable first). + + Priority order: + 1. Not pinned before pinned + 2. Lower fidelity (closer to eviction) before higher + 3. Older last_accessed_turn before newer + 4. Older created_at_turn before newer + + Only includes objects not already at L4. + + Args: + current_turn: Current conversation turn (used for pin expiry check). + + Returns: + Objects sorted most-evictable-first. + """ + self._expire_pins(current_turn) + + candidates = [o for o in self._objects.values() if o.current_fidelity < FidelityLevel.L4] + + # Sort: not-pinned first, then by fidelity (higher numeric = closer to eviction), + # then oldest-accessed first + return sorted( + candidates, + key=lambda o: ( + o.pinned, # False (0) before True (1) — not-pinned first + -o.current_fidelity, # Higher fidelity number = closer to eviction + o.last_accessed_turn, # Older access = more evictable + o.created_at_turn, # Older creation = more evictable + ), + ) + + def objects_at_fidelity(self, level: FidelityLevel) -> list[SemanticObject]: + """Return all objects currently at the given fidelity level. + + Args: + level: The fidelity level to filter by. + + Returns: + List of objects at that level. + """ + return [o for o in self._objects.values() if o.current_fidelity == level] + + def summary(self) -> dict: + """Current state summary for telemetry and debugging. + + Returns: + Dict with object counts per fidelity level, total tokens, + pressure zone, pinned count, and fault count. + """ + by_level = {} + for level in FidelityLevel: + objs = self.objects_at_fidelity(level) + by_level[level.name] = len(objs) + + total_faults = sum(o.fault_count for o in self._objects.values()) + pinned_count = sum(1 for o in self._objects.values() if o.pinned) + + return { + "total_objects": len(self._objects), + "total_tokens": self.total_tokens(), + "window_size": self.window_size, + "pressure_zone": self.current_pressure().name, + "objects_by_level": by_level, + "pinned_count": pinned_count, + "total_faults": total_faults, + } diff --git a/src/mnemosyne/helper_llm.py b/src/mnemosyne/helper_llm.py new file mode 100644 index 0000000..21c7d6c --- /dev/null +++ b/src/mnemosyne/helper_llm.py @@ -0,0 +1,442 @@ +"""Helper LLM client for multi-fidelity compression and micro-fault QA. + +Wraps the Anthropic API (Haiku) for cheap, fast operations: + - Summarization with declared losses (L0→L1, L1→L2, L2→L3) + - Micro-fault question-answering (query evicted content without page-in) + - Goal classification for context-aware retrieval + +Cost per session (~200 turns): ~$0.006 total helper spend. +See ARCHITECTURE.md §6 and §11.3 for full specification. +""" + +from __future__ import annotations + +import json +import logging +import os +from dataclasses import dataclass, field + +import anthropic + +logger = logging.getLogger(__name__) + +# --------------------------------------------------------------------------- +# Data classes +# --------------------------------------------------------------------------- + + +@dataclass +class SummaryResult: + """Result of a summarization call at any fidelity level.""" + + summary: str + losses: list[str] = field(default_factory=list) + can_answer: list[str] = field(default_factory=list) + key_entities: list[str] = field(default_factory=list) + + +@dataclass +class GoalClassification: + """Result of classifying the user's current goal.""" + + goal: str + relevant_types: list[str] = field(default_factory=list) + relevant_tags: list[str] = field(default_factory=list) + predicted_needs: list[str] = field(default_factory=list) + + +# --------------------------------------------------------------------------- +# Prompt templates +# --------------------------------------------------------------------------- + +_L0_TO_L1_PROMPT = """\ +You are a context compression engine for a coding agent. Summarize the following \ +content while preserving maximum utility for future reference. + +CONTENT TYPE: {object_type} +CONTENT: +{content} + +INSTRUCTIONS: +1. Write a detailed summary (~30% of original length) +2. MUST preserve: file paths, function names, variable names, library names, \ +error messages, decision rationale, specific values that may be referenced later +3. List DECLARED LOSSES: specific information you omitted that someone might need. \ +Be precise -- "specific error codes" not "some details" +4. List CAN_ANSWER: categories of questions this summary can answer without \ +needing the original content + +OUTPUT FORMAT (JSON): +{{ + "summary": "...", + "losses": ["exact error code for token expiry", ...], + "can_answer": ["auth approach used", ...], + "key_entities": ["src/auth/middleware.ts", ...] +}}""" + +_L1_TO_L2_PROMPT = """\ +You are a context compression engine. Compress the following L1 summary into a \ +compact L2 summary (~5% of the original content length). + +CONTENT TYPE: {object_type} +L1 SUMMARY: +{l1_summary} + +KNOWN LOSSES FROM L1: +{l1_losses} + +INSTRUCTIONS: +1. Write a compact summary preserving only: what was done, what was decided, key files +2. List ADDITIONAL DECLARED LOSSES beyond those already declared in L1 +3. List CAN_ANSWER: what questions this compact summary can still answer + +OUTPUT FORMAT (JSON): +{{ + "summary": "...", + "losses": ["additional loss 1", ...], + "can_answer": ["what was decided", ...], + "key_entities": ["key/file.ts", ...] +}}""" + +_L2_TO_L3_PROMPT = """\ +Generate a one-line metadata stub for this content. + +CONTENT TYPE: {object_type} +TIMESTAMP: {timestamp} +L2 SUMMARY: +{l2_summary} + +Return ONLY a single line in this exact format (no JSON, no markdown): +[{object_type} | {timestamp} | | related objects]""" + +_MICRO_FAULT_PROMPT = """\ +Answer the following question using ONLY the provided context. \ +Be concise and precise (50-200 tokens). If the context doesn't contain \ +the answer, say so explicitly. + +QUESTION: {question} + +CONTEXT: +{context} + +ANSWER:""" + +_GOAL_CLASSIFICATION_PROMPT = """\ +Given the user's message and recent conversation context, classify the user's \ +current goal and predict what context would be most relevant. + +USER MESSAGE: {user_message} + +RECENT CONTEXT: +{recent_context} + +Return JSON: +{{ + "goal": "one-sentence description of current goal", + "relevant_types": ["conversation_phase", "file_context", ...], + "relevant_tags": ["auth", "testing", ...], + "predicted_needs": ["auth middleware implementation", "test patterns", ...] +}}""" + + +# --------------------------------------------------------------------------- +# Helper LLM client +# --------------------------------------------------------------------------- + + +class HelperLLM: + """Async client for cheap helper model (Haiku) operations. + + Handles summarization, declared loss generation, micro-fault QA, + and goal classification. All methods return structured results + with graceful fallbacks on error. + """ + + def __init__( + self, + api_key: str | None = None, + model: str = "claude-haiku-4-5-20251001", + base_url: str = "https://api.anthropic.com", + ) -> None: + """Initialize the helper LLM client. + + Args: + api_key: Anthropic API key. Falls back to ANTHROPIC_API_KEY env var. + model: Model identifier for helper calls. + base_url: Anthropic API base URL. + """ + resolved_key = api_key or os.environ.get("ANTHROPIC_API_KEY") or "" + auth_token = os.environ.get("ANTHROPIC_AUTH_TOKEN") or "" + + # OAuth: use auth_token (Bearer) when available. The SDK does NOT + # add the required beta header automatically — we must pass it. + # Falls back to api_key when no OAuth token is available. + if auth_token and not resolved_key: + self._client = anthropic.AsyncAnthropic( + auth_token=auth_token, + base_url=base_url, + timeout=10.0, + max_retries=2, + default_headers={ + "anthropic-beta": "oauth-2025-04-20", + "user-agent": "claude-cli/2.1.2 (external, cli)", + }, + ) + else: + self._client = anthropic.AsyncAnthropic( + api_key=resolved_key, + base_url=base_url, + timeout=10.0, + max_retries=2, + ) + self._model = model + + async def summarize_l0_to_l1( + self, + content: str, + object_type: str, + max_summary_tokens: int = 1024, + ) -> SummaryResult: + """Summarize full content (L0) into a detailed summary (L1). + + Produces a ~30% compression with declared losses and entity + extraction. See ARCHITECTURE.md §6.1 for the prompt template. + + Args: + content: Full L0 content to summarize. + object_type: Semantic object type (e.g. 'file_context'). + max_summary_tokens: Max tokens for the response. + + Returns: + SummaryResult with summary, losses, can_answer, key_entities. + """ + prompt = _L0_TO_L1_PROMPT.format( + object_type=object_type, + content=content, + ) + raw = await self._call(prompt, max_tokens=max_summary_tokens) + result = self._parse_summary_json(raw) + if result is not None: + return result + # Fallback: first ~30% of content as summary + logger.warning("L0→L1 JSON parse failed, using fallback summary") + cutoff = max(1, len(content) * 30 // 100) + return SummaryResult(summary=content[:cutoff]) + + async def compress_l1_to_l2( + self, + l1_summary: str, + l1_losses: list[str], + object_type: str, + max_tokens: int = 256, + ) -> SummaryResult: + """Compress an L1 summary into a compact L2 summary (~5% of original). + + Accumulates losses from L1 into the L2 result. + + Args: + l1_summary: The L1 detailed summary. + l1_losses: Declared losses from the L1 level. + object_type: Semantic object type. + max_tokens: Max tokens for the response. + + Returns: + SummaryResult with accumulated losses from both L1 and L2. + """ + losses_text = "\n".join(f"- {loss}" for loss in l1_losses) if l1_losses else "(none)" + prompt = _L1_TO_L2_PROMPT.format( + object_type=object_type, + l1_summary=l1_summary, + l1_losses=losses_text, + ) + raw = await self._call(prompt, max_tokens=max_tokens) + result = self._parse_summary_json(raw) + if result is not None: + # Accumulate L1 losses into L2 + result.losses = l1_losses + result.losses + return result + logger.warning("L1→L2 JSON parse failed, using fallback summary") + cutoff = max(1, len(l1_summary) * 30 // 100) + return SummaryResult(summary=l1_summary[:cutoff], losses=list(l1_losses)) + + async def generate_stub( + self, + l2_summary: str, + object_type: str, + timestamp: str, + ) -> str: + """Generate a one-line L3 metadata stub from an L2 summary. + + Format: [{type} | {timestamp} | {description} | {n} related objects] + + Args: + l2_summary: The L2 compact summary. + object_type: Semantic object type. + timestamp: ISO-ish timestamp string. + + Returns: + A single-line stub string. + """ + prompt = _L2_TO_L3_PROMPT.format( + object_type=object_type, + timestamp=timestamp, + l2_summary=l2_summary, + ) + raw = await self._call(prompt, max_tokens=100) + # Take only the first line, strip whitespace + stub = raw.strip().split("\n")[0].strip() + if not stub: + stub = f"[{object_type} | {timestamp} | (summary unavailable) | 0 related objects]" + return stub + + async def answer_micro_fault( + self, + question: str, + relevant_contents: list[str], + max_tokens: int = 200, + ) -> str: + """Answer a question using retrieved full-content from the backing store. + + Used by the memory_query phantom tool to avoid full page-in. + Returns a targeted 50-200 token answer. + + Args: + question: The question to answer. + relevant_contents: List of full-content strings from backing store. + max_tokens: Max tokens for the answer. + + Returns: + Concise answer string. + """ + context = "\n\n---\n\n".join(relevant_contents) + prompt = _MICRO_FAULT_PROMPT.format( + question=question, + context=context, + ) + raw = await self._call(prompt, max_tokens=max_tokens) + return raw.strip() if raw.strip() else "Unable to answer from available context." + + async def classify_goal( + self, + user_message: str, + recent_context: str, + ) -> GoalClassification: + """Classify the user's current goal for context-aware retrieval. + + Takes the current user message and last ~2 turns of context. + Returns structured goal classification for the Context Assembler. + + Args: + user_message: The current user message. + recent_context: Last 2 turns of conversation context. + + Returns: + GoalClassification with goal, relevant_types, tags, predicted_needs. + """ + prompt = _GOAL_CLASSIFICATION_PROMPT.format( + user_message=user_message, + recent_context=recent_context, + ) + raw = await self._call(prompt, max_tokens=200) + result = self._parse_goal_json(raw) + if result is not None: + return result + logger.warning("Goal classification JSON parse failed, using fallback") + return GoalClassification(goal=user_message[:200]) + + # ------------------------------------------------------------------ + # Internal helpers + # ------------------------------------------------------------------ + + async def _call(self, prompt: str, max_tokens: int) -> str: + """Make a single Haiku API call. Returns raw text or empty on error.""" + try: + response = await self._client.messages.create( + model=self._model, + max_tokens=max_tokens, + messages=[{"role": "user", "content": prompt}], + ) + # Extract text from content blocks + parts: list[str] = [] + for block in response.content: + if block.type == "text": + parts.append(block.text) + return "\n".join(parts) + except anthropic.APITimeoutError: + logger.warning("Helper LLM call timed out") + return "" + except anthropic.APIError as exc: + logger.warning("Helper LLM API error: %s", exc) + return "" + + @staticmethod + def _parse_summary_json(raw: str) -> SummaryResult | None: + """Try to parse a SummaryResult from raw LLM output. + + Handles JSON embedded in markdown code fences or bare JSON. + Returns None on parse failure. + """ + text = raw.strip() + if not text: + return None + # Strip markdown code fences if present + if text.startswith("```"): + lines = text.split("\n") + # Remove first and last fence lines + lines = [ln for ln in lines if not ln.strip().startswith("```")] + text = "\n".join(lines).strip() + try: + data = json.loads(text) + except json.JSONDecodeError: + # Try to find JSON object in the text + start = text.find("{") + end = text.rfind("}") + if start >= 0 and end > start: + try: + data = json.loads(text[start : end + 1]) + except json.JSONDecodeError: + return None + else: + return None + if not isinstance(data, dict): + return None + return SummaryResult( + summary=data.get("summary", ""), + losses=data.get("losses", []), + can_answer=data.get("can_answer", []), + key_entities=data.get("key_entities", []), + ) + + @staticmethod + def _parse_goal_json(raw: str) -> GoalClassification | None: + """Try to parse a GoalClassification from raw LLM output. + + Returns None on parse failure. + """ + text = raw.strip() + if not text: + return None + if text.startswith("```"): + lines = text.split("\n") + lines = [ln for ln in lines if not ln.strip().startswith("```")] + text = "\n".join(lines).strip() + try: + data = json.loads(text) + except json.JSONDecodeError: + start = text.find("{") + end = text.rfind("}") + if start >= 0 and end > start: + try: + data = json.loads(text[start : end + 1]) + except json.JSONDecodeError: + return None + else: + return None + if not isinstance(data, dict): + return None + return GoalClassification( + goal=data.get("goal", ""), + relevant_types=data.get("relevant_types", []), + relevant_tags=data.get("relevant_tags", []), + predicted_needs=data.get("predicted_needs", []), + ) diff --git a/tests/test_fidelity.py b/tests/test_fidelity.py new file mode 100644 index 0000000..f9e8316 --- /dev/null +++ b/tests/test_fidelity.py @@ -0,0 +1,861 @@ +"""Tests for the multi-fidelity state machine.""" + +from __future__ import annotations + +from mnemosyne.fidelity import ( + VALID_OBJECT_TYPES, + FidelityLevel, + FidelityManager, + PressureZone, + SemanticObject, + _estimate_tokens, + make_object, +) + +# ── Helpers ────────────────────────────────────────────────── + + +def _make_obj( + content: str = "x" * 400, + *, + object_type: str = "file_context", + turn: int = 0, + summary_detailed: str | None = None, + summary_compact: str | None = None, + stub: str | None = None, +) -> SemanticObject: + """Create a SemanticObject with sensible defaults for testing.""" + return make_object( + object_type=object_type, + content_full=content, + created_at_turn=turn, + summary_detailed=summary_detailed or ("summary " * 20), # ~140 chars + summary_compact=summary_compact or ("compact " * 5), # ~40 chars + stub=stub or "file_context: test object", + ) + + +def _fill_manager( + manager: FidelityManager, + count: int, + *, + content_size: int = 400, + turn: int = 0, +) -> list[str]: + """Register `count` objects and return their IDs.""" + ids = [] + for i in range(count): + obj = _make_obj("x" * content_size, turn=turn + i) + oid = manager.register_object(obj) + ids.append(oid) + return ids + + +# ── FidelityLevel enum ─────────────────────────────────────── + + +class TestFidelityLevel: + def test_five_levels(self): + assert len(FidelityLevel) == 5 + + def test_ordering(self): + assert FidelityLevel.L0 < FidelityLevel.L1 < FidelityLevel.L2 + assert FidelityLevel.L2 < FidelityLevel.L3 < FidelityLevel.L4 + + def test_values(self): + assert FidelityLevel.L0 == 0 + assert FidelityLevel.L4 == 4 + + def test_names(self): + assert FidelityLevel.L0.name == "L0" + assert FidelityLevel.L4.name == "L4" + + +# ── PressureZone enum ──────────────────────────────────────── + + +class TestPressureZone: + def test_five_zones(self): + assert len(PressureZone) == 5 + + def test_ordering(self): + assert PressureZone.NORMAL < PressureZone.CAUTION + assert PressureZone.CAUTION < PressureZone.WARNING + assert PressureZone.WARNING < PressureZone.CRITICAL + assert PressureZone.CRITICAL < PressureZone.EMERGENCY + + def test_names(self): + expected = {"NORMAL", "CAUTION", "WARNING", "CRITICAL", "EMERGENCY"} + assert {z.name for z in PressureZone} == expected + + +# ── SemanticObject ──────────────────────────────────────────── + + +class TestSemanticObject: + def test_required_fields(self): + obj = _make_obj() + assert obj.id + assert obj.object_type == "file_context" + assert obj.content_full == "x" * 400 + assert obj.current_fidelity == FidelityLevel.L0 + assert obj.pinned is False + assert obj.pin_until_turn is None + assert obj.fault_count == 0 + + def test_token_estimation(self): + obj = _make_obj("a" * 1000) + assert obj.token_count_l0 == 250 # 1000 / 4 + + def test_tokens_at_each_level(self): + obj = _make_obj( + "a" * 400, + summary_detailed="b" * 120, + summary_compact="c" * 20, + stub="d" * 100, + ) + assert obj.tokens_at(FidelityLevel.L0) == 100 # 400/4 + assert obj.tokens_at(FidelityLevel.L1) == 30 # 120/4 + assert obj.tokens_at(FidelityLevel.L2) == 5 # 20/4 + assert obj.tokens_at(FidelityLevel.L3) == 25 # 100/4 + assert obj.tokens_at(FidelityLevel.L4) == 0 # evicted + + def test_current_tokens_tracks_fidelity(self): + obj = _make_obj("a" * 400, summary_detailed="b" * 120) + assert obj.current_tokens == 100 # L0: 400/4 + obj.current_fidelity = FidelityLevel.L1 + assert obj.current_tokens == 30 # L1: 120/4 + + def test_losses_default_empty(self): + obj = _make_obj() + assert obj.losses_l1 == [] + assert obj.losses_l2 == [] + + def test_losses_populated(self): + obj = make_object( + object_type="design_decision", + content_full="decision content", + losses_l1=["exact error codes"], + losses_l2=["exact error codes", "function signatures"], + ) + assert "exact error codes" in obj.losses_l1 + assert len(obj.losses_l2) == 2 + + def test_valid_object_types(self): + expected = { + "conversation_phase", + "design_decision", + "debugging_session", + "file_context", + "tool_result", + "plan", + "error_context", + "external_reference", + } + assert VALID_OBJECT_TYPES == expected + + def test_queryability_fields(self): + obj = make_object( + object_type="file_context", + content_full="content", + can_answer=["what functions are defined"], + fault_when=["need exact line numbers"], + key_entities=["auth.py", "middleware"], + ) + assert obj.can_answer == ["what functions are defined"] + assert obj.fault_when == ["need exact line numbers"] + assert obj.key_entities == ["auth.py", "middleware"] + + +# ── Token estimation ────────────────────────────────────────── + + +class TestTokenEstimation: + def test_basic(self): + assert _estimate_tokens("a" * 100) == 25 + + def test_none(self): + assert _estimate_tokens(None) == 0 + + def test_empty(self): + # len("") = 0, 0 // 4 = 0, max(1, 0) = 1 + assert _estimate_tokens("") == 1 + + def test_empty_returns_minimum(self): + # Empty string: len=0, 0//4=0, max(1,0)=1 + assert _estimate_tokens("") == 1 + + def test_short(self): + assert _estimate_tokens("hi") == 1 # 2//4=0, max(1,0)=1 + + def test_exact_multiple(self): + assert _estimate_tokens("a" * 400) == 100 + + +# ── make_object factory ─────────────────────────────────────── + + +class TestMakeObject: + def test_generates_id(self): + obj = make_object(object_type="plan", content_full="plan content") + assert len(obj.id) == 16 + assert obj.id.isalnum() + + def test_unique_ids(self): + ids = {make_object(object_type="plan", content_full="x").id for _ in range(100)} + assert len(ids) == 100 + + def test_auto_token_estimates(self): + obj = make_object( + object_type="file_context", + content_full="a" * 800, + summary_detailed="b" * 240, + summary_compact="c" * 40, + stub="d" * 100, + ) + assert obj.token_count_l0 == 200 + assert obj.token_count_l1 == 60 + assert obj.token_count_l2 == 10 + assert obj.token_count_l3 == 25 # 100/4 + + def test_default_stub_tokens(self): + obj = make_object(object_type="plan", content_full="content") + assert obj.token_count_l3 == 25 # default when no stub provided + + def test_starts_at_l0(self): + obj = make_object(object_type="plan", content_full="content") + assert obj.current_fidelity == FidelityLevel.L0 + + +# ── FidelityManager basics ──────────────────────────────────── + + +class TestFidelityManagerBasics: + def test_default_window_size(self): + fm = FidelityManager() + assert fm.window_size == 200_000 + + def test_custom_window_size(self): + fm = FidelityManager(window_size=100_000) + assert fm.window_size == 100_000 + + def test_register_and_get(self): + fm = FidelityManager() + obj = _make_obj() + oid = fm.register_object(obj) + assert fm.get_object(oid) is obj + + def test_get_nonexistent(self): + fm = FidelityManager() + assert fm.get_object("nonexistent") is None + + def test_register_returns_id(self): + fm = FidelityManager() + obj = _make_obj() + oid = fm.register_object(obj) + assert oid == obj.id + + def test_total_tokens_empty(self): + fm = FidelityManager() + assert fm.total_tokens() == 0 + + def test_total_tokens_sums_current_fidelity(self): + fm = FidelityManager() + obj1 = _make_obj("a" * 400) # 100 tokens at L0 + obj2 = _make_obj("b" * 800) # 200 tokens at L0 + fm.register_object(obj1) + fm.register_object(obj2) + assert fm.total_tokens() == 300 + + def test_total_tokens_respects_fidelity_change(self): + fm = FidelityManager() + obj = _make_obj("a" * 400, summary_detailed="b" * 120) + fm.register_object(obj) + assert fm.total_tokens() == 100 # L0: 400/4 + obj.current_fidelity = FidelityLevel.L1 + assert fm.total_tokens() == 30 # L1: 120/4 + + +# ── Pressure zones ──────────────────────────────────────────── + + +class TestPressureZones: + def test_normal_zone(self): + fm = FidelityManager(window_size=1000) + # 400 chars = 100 tokens, 10% of 1000 → NORMAL + _fill_manager(fm, 1, content_size=400) + assert fm.current_pressure() == PressureZone.NORMAL + + def test_caution_zone(self): + fm = FidelityManager(window_size=1000) + # 2400 chars = 600 tokens, 60% of 1000 → CAUTION + _fill_manager(fm, 1, content_size=2400) + assert fm.current_pressure() == PressureZone.CAUTION + + def test_warning_zone(self): + fm = FidelityManager(window_size=1000) + # 3200 chars = 800 tokens, 80% of 1000 → WARNING + _fill_manager(fm, 1, content_size=3200) + assert fm.current_pressure() == PressureZone.WARNING + + def test_critical_zone(self): + fm = FidelityManager(window_size=1000) + # 3600 chars = 900 tokens, 90% of 1000 → CRITICAL + _fill_manager(fm, 1, content_size=3600) + assert fm.current_pressure() == PressureZone.CRITICAL + + def test_emergency_zone(self): + fm = FidelityManager(window_size=1000) + # 3840 chars = 960 tokens, 96% of 1000 → EMERGENCY + _fill_manager(fm, 1, content_size=3840) + assert fm.current_pressure() == PressureZone.EMERGENCY + + def test_exact_boundary_50pct(self): + fm = FidelityManager(window_size=1000) + # 2000 chars = 500 tokens, exactly 50% → CAUTION (>= threshold) + _fill_manager(fm, 1, content_size=2000) + assert fm.current_pressure() == PressureZone.CAUTION + + def test_just_below_50pct(self): + fm = FidelityManager(window_size=1000) + # 1996 chars = 499 tokens, 49.9% → NORMAL + _fill_manager(fm, 1, content_size=1996) + assert fm.current_pressure() == PressureZone.NORMAL + + def test_zero_window_is_emergency(self): + fm = FidelityManager(window_size=0) + assert fm.current_pressure() == PressureZone.EMERGENCY + + +# ── Degradation ─────────────────────────────────────────────── + + +class TestDegradation: + def test_no_degradation_in_normal(self): + fm = FidelityManager(window_size=10_000) + _fill_manager(fm, 1, content_size=400) # 100 tokens, 1% → NORMAL + transitions = fm.degrade(current_turn=1) + assert transitions == [] + + def test_caution_degrades_l0_to_l1(self): + fm = FidelityManager(window_size=1000) + # 2400 chars = 600 tokens → CAUTION + _fill_manager(fm, 3, content_size=800) + assert fm.current_pressure() == PressureZone.CAUTION + + transitions = fm.degrade(current_turn=5) + # Should have degraded some L0 objects to L1 + assert len(transitions) > 0 + for _oid, old, new in transitions: + assert old == FidelityLevel.L0 + assert new == FidelityLevel.L1 + + def test_caution_degrades_oldest_first(self): + fm = FidelityManager(window_size=1000) + # Create objects at different turns + # 1200 chars each = 300 tokens each, 600 total = 60% → CAUTION + obj_old = _make_obj("a" * 1200, turn=0) + obj_new = _make_obj("b" * 1200, turn=5) + id_old = fm.register_object(obj_old) + id_new = fm.register_object(obj_new) + + # Mark new one as recently accessed + fm.mark_accessed(id_new, current_turn=10) + + assert fm.current_pressure() == PressureZone.CAUTION + + transitions = fm.degrade(current_turn=10) + # Oldest should be degraded first + if transitions: + assert transitions[0][0] == id_old + + def test_warning_degrades_multiple_levels(self): + fm = FidelityManager(window_size=1000) + # 3200 chars = 800 tokens → WARNING + _fill_manager(fm, 4, content_size=800) + assert fm.current_pressure() == PressureZone.WARNING + + transitions = fm.degrade(current_turn=10) + assert len(transitions) > 0 + # Should see L0→L1 transitions at minimum + levels_seen = {(old, new) for _, old, new in transitions} + assert (FidelityLevel.L0, FidelityLevel.L1) in levels_seen + + def test_critical_degrades_aggressively(self): + fm = FidelityManager(window_size=1000) + # 3600 chars = 900 tokens → CRITICAL + _fill_manager(fm, 3, content_size=1200) + assert fm.current_pressure() == PressureZone.CRITICAL + + transitions = fm.degrade(current_turn=10) + assert len(transitions) > 0 + # Should see objects pushed to L3 or L4 + final_levels = {new for _, _, new in transitions} + assert FidelityLevel.L3 in final_levels or FidelityLevel.L4 in final_levels + + def test_emergency_evicts_all_unpinned(self): + fm = FidelityManager(window_size=1000) + # 3840 chars = 960 tokens → EMERGENCY + ids = _fill_manager(fm, 4, content_size=960) + assert fm.current_pressure() == PressureZone.EMERGENCY + + fm.degrade(current_turn=10) + # All objects should end at L4 + for obj_id in ids: + obj = fm.get_object(obj_id) + assert obj is not None + assert obj.current_fidelity == FidelityLevel.L4 + + def test_emergency_respects_pins(self): + fm = FidelityManager(window_size=1000) + ids = _fill_manager(fm, 4, content_size=960) + + # Pin one object + fm.record_fault(ids[0], current_turn=5, pin_duration=20) + + fm.degrade(current_turn=10) + # Pinned object should NOT be at L4 + pinned_obj = fm.get_object(ids[0]) + assert pinned_obj is not None + assert pinned_obj.current_fidelity < FidelityLevel.L4 + + # Others should be at L4 + for oid in ids[1:]: + obj = fm.get_object(oid) + assert obj is not None + assert obj.current_fidelity == FidelityLevel.L4 + + def test_degrade_returns_transitions(self): + fm = FidelityManager(window_size=1000) + _fill_manager(fm, 3, content_size=800) + transitions = fm.degrade(current_turn=5) + for oid, old, new in transitions: + assert isinstance(oid, str) + assert isinstance(old, FidelityLevel) + assert isinstance(new, FidelityLevel) + assert new > old # Degradation means higher numeric level + + def test_degrade_stops_when_pressure_relieved(self): + fm = FidelityManager(window_size=1000) + # Create objects with large L0 but small L1 + for i in range(3): + obj = _make_obj( + "a" * 800, # 200 tokens at L0 + summary_detailed="b" * 40, # 10 tokens at L1 + turn=i, + ) + fm.register_object(obj) + + # 600 tokens → CAUTION + assert fm.current_pressure() == PressureZone.CAUTION + + fm.degrade(current_turn=5) + # After degrading enough objects, pressure should drop + # Not all objects need to be degraded + assert fm.current_pressure() <= PressureZone.CAUTION + + +# ── Upgrade ─────────────────────────────────────────────────── + + +class TestUpgrade: + def test_upgrade_l1_to_l0(self): + fm = FidelityManager() + obj = _make_obj("content" * 50) + oid = fm.register_object(obj) + obj.current_fidelity = FidelityLevel.L1 + + result = fm.upgrade(oid, FidelityLevel.L0, current_turn=5) + assert result is True + assert obj.current_fidelity == FidelityLevel.L0 + + def test_upgrade_updates_last_accessed(self): + fm = FidelityManager() + obj = _make_obj("content" * 50) + oid = fm.register_object(obj) + obj.current_fidelity = FidelityLevel.L2 + + fm.upgrade(oid, FidelityLevel.L0, current_turn=42) + assert obj.last_accessed_turn == 42 + + def test_upgrade_rejects_same_level(self): + fm = FidelityManager() + obj = _make_obj() + oid = fm.register_object(obj) + + result = fm.upgrade(oid, FidelityLevel.L0, current_turn=5) + assert result is False # Already at L0 + + def test_upgrade_rejects_downgrade(self): + fm = FidelityManager() + obj = _make_obj() + oid = fm.register_object(obj) + + result = fm.upgrade(oid, FidelityLevel.L1, current_turn=5) + assert result is False # L1 > L0, not an upgrade + + def test_upgrade_nonexistent_object(self): + fm = FidelityManager() + result = fm.upgrade("nonexistent", FidelityLevel.L0, current_turn=5) + assert result is False + + def test_upgrade_l3_to_l1(self): + fm = FidelityManager() + obj = _make_obj(summary_detailed="detailed summary content") + oid = fm.register_object(obj) + obj.current_fidelity = FidelityLevel.L3 + + result = fm.upgrade(oid, FidelityLevel.L1, current_turn=10) + assert result is True + assert obj.current_fidelity == FidelityLevel.L1 + + def test_upgrade_requires_content_at_target(self): + fm = FidelityManager() + obj = make_object( + object_type="file_context", + content_full="full content", + # No summary_detailed provided + ) + oid = fm.register_object(obj) + obj.current_fidelity = FidelityLevel.L3 + + # Can upgrade to L0 (has content_full) + result = fm.upgrade(oid, FidelityLevel.L0, current_turn=5) + assert result is True + + def test_upgrade_rejects_missing_l1_content(self): + fm = FidelityManager() + obj = SemanticObject( + id="test123", + object_type="file_context", + content_full="full", + summary_detailed=None, # No L1 content + token_count_l0=1, + ) + fm.register_object(obj) + obj.current_fidelity = FidelityLevel.L3 + + result = fm.upgrade("test123", FidelityLevel.L1, current_turn=5) + assert result is False + + +# ── Fault-driven pinning ────────────────────────────────────── + + +class TestFaultPinning: + def test_record_fault_pins_object(self): + fm = FidelityManager() + obj = _make_obj() + oid = fm.register_object(obj) + + fm.record_fault(oid, current_turn=10, pin_duration=5) + assert obj.pinned is True + assert obj.pin_until_turn == 15 + assert obj.fault_count == 1 + + def test_record_fault_updates_access(self): + fm = FidelityManager() + obj = _make_obj() + oid = fm.register_object(obj) + + fm.record_fault(oid, current_turn=10) + assert obj.last_accessed_turn == 10 + + def test_multiple_faults_increment_count(self): + fm = FidelityManager() + obj = _make_obj() + oid = fm.register_object(obj) + + fm.record_fault(oid, current_turn=10) + fm.record_fault(oid, current_turn=12) + assert obj.fault_count == 2 + + def test_pin_prevents_degradation(self): + fm = FidelityManager(window_size=1000) + ids = _fill_manager(fm, 3, content_size=800) + + # Pin the oldest object + fm.record_fault(ids[0], current_turn=5, pin_duration=20) + + transitions = fm.degrade(current_turn=10) + # Pinned object should not appear in transitions + degraded_ids = {oid for oid, _, _ in transitions} + assert ids[0] not in degraded_ids + + def test_pin_expires(self): + fm = FidelityManager(window_size=1000) + ids = _fill_manager(fm, 3, content_size=800) + + fm.record_fault(ids[0], current_turn=5, pin_duration=3) + # Pin expires at turn 8 + + fm.degrade(current_turn=9) + # Now the object can be degraded + obj = fm.get_object(ids[0]) + assert obj is not None + assert obj.pinned is False + + def test_default_pin_duration(self): + fm = FidelityManager() + obj = _make_obj() + oid = fm.register_object(obj) + + fm.record_fault(oid, current_turn=10) + assert obj.pin_until_turn == 15 # default duration = 5 + + def test_fault_nonexistent_object(self): + fm = FidelityManager() + # Should not raise + fm.record_fault("nonexistent", current_turn=10) + + +# ── mark_accessed ───────────────────────────────────────────── + + +class TestMarkAccessed: + def test_updates_last_accessed(self): + fm = FidelityManager() + obj = _make_obj(turn=0) + oid = fm.register_object(obj) + + fm.mark_accessed(oid, current_turn=42) + assert obj.last_accessed_turn == 42 + + def test_nonexistent_object(self): + fm = FidelityManager() + # Should not raise + fm.mark_accessed("nonexistent", current_turn=10) + + +# ── eviction_candidates ────────────────────────────────────── + + +class TestEvictionCandidates: + def test_empty_manager(self): + fm = FidelityManager() + assert fm.eviction_candidates(current_turn=0) == [] + + def test_excludes_l4_objects(self): + fm = FidelityManager() + obj = _make_obj() + fm.register_object(obj) + obj.current_fidelity = FidelityLevel.L4 + + candidates = fm.eviction_candidates(current_turn=0) + assert len(candidates) == 0 + + def test_unpinned_before_pinned(self): + fm = FidelityManager() + obj_pinned = _make_obj(turn=0) + obj_free = _make_obj(turn=1) + id_pinned = fm.register_object(obj_pinned) + id_free = fm.register_object(obj_free) + + fm.record_fault(id_pinned, current_turn=5, pin_duration=20) + + candidates = fm.eviction_candidates(current_turn=6) + assert len(candidates) == 2 + assert candidates[0].id == id_free # Unpinned first + + def test_older_access_more_evictable(self): + fm = FidelityManager() + obj_old = _make_obj(turn=0) + obj_new = _make_obj(turn=0) + id_old = fm.register_object(obj_old) + id_new = fm.register_object(obj_new) + + fm.mark_accessed(id_new, current_turn=10) + + candidates = fm.eviction_candidates(current_turn=10) + assert candidates[0].id == id_old + + def test_lower_fidelity_more_evictable(self): + fm = FidelityManager() + obj_l0 = _make_obj(turn=0) + obj_l3 = _make_obj(turn=0) + fm.register_object(obj_l0) + id_l3 = fm.register_object(obj_l3) + obj_l3.current_fidelity = FidelityLevel.L3 + + candidates = fm.eviction_candidates(current_turn=0) + # L3 (closer to eviction) should come first + assert candidates[0].id == id_l3 + + def test_expired_pins_are_evictable(self): + fm = FidelityManager() + obj = _make_obj(turn=0) + oid = fm.register_object(obj) + + fm.record_fault(oid, current_turn=5, pin_duration=3) + # Pin expires at turn 8 + + candidates = fm.eviction_candidates(current_turn=9) + assert len(candidates) == 1 + assert candidates[0].pinned is False + + +# ── objects_at_fidelity ─────────────────────────────────────── + + +class TestObjectsAtFidelity: + def test_all_at_l0(self): + fm = FidelityManager() + _fill_manager(fm, 3) + assert len(fm.objects_at_fidelity(FidelityLevel.L0)) == 3 + assert len(fm.objects_at_fidelity(FidelityLevel.L1)) == 0 + + def test_mixed_levels(self): + fm = FidelityManager() + ids = _fill_manager(fm, 3) + obj0 = fm.get_object(ids[0]) + obj1 = fm.get_object(ids[1]) + assert obj0 is not None + assert obj1 is not None + obj0.current_fidelity = FidelityLevel.L1 + obj1.current_fidelity = FidelityLevel.L3 + + assert len(fm.objects_at_fidelity(FidelityLevel.L0)) == 1 + assert len(fm.objects_at_fidelity(FidelityLevel.L1)) == 1 + assert len(fm.objects_at_fidelity(FidelityLevel.L3)) == 1 + + +# ── summary ─────────────────────────────────────────────────── + + +class TestSummary: + def test_summary_structure(self): + fm = FidelityManager(window_size=10_000) + _fill_manager(fm, 2, content_size=400) + + s = fm.summary() + assert "total_objects" in s + assert "total_tokens" in s + assert "window_size" in s + assert "pressure_zone" in s + assert "objects_by_level" in s + assert "pinned_count" in s + assert "total_faults" in s + + def test_summary_values(self): + fm = FidelityManager(window_size=10_000) + ids = _fill_manager(fm, 3, content_size=400) + fm.record_fault(ids[0], current_turn=5) + + s = fm.summary() + assert s["total_objects"] == 3 + assert s["total_tokens"] == 300 # 3 * 100 + assert s["window_size"] == 10_000 + assert s["pressure_zone"] == "NORMAL" + assert s["objects_by_level"]["L0"] == 3 + assert s["pinned_count"] == 1 + assert s["total_faults"] == 1 + + +# ── Integration: full lifecycle ─────────────────────────────── + + +class TestIntegration: + def test_register_degrade_upgrade_cycle(self): + """Full lifecycle: register → pressure → degrade → access → upgrade.""" + fm = FidelityManager(window_size=500) + + # Register objects that push into CAUTION + ids = [] + for i in range(5): + obj = _make_obj( + "x" * 400, # 100 tokens each + summary_detailed="y" * 120, # 30 tokens + summary_compact="z" * 20, # 5 tokens + stub="stub text", + turn=i, + ) + ids.append(fm.register_object(obj)) + + # 500 tokens in 500 window → 100% → EMERGENCY + assert fm.current_pressure() == PressureZone.EMERGENCY + + # Degrade + transitions = fm.degrade(current_turn=10) + assert len(transitions) > 0 + + # All should be evicted (emergency) + for oid in ids: + obj = fm.get_object(oid) + assert obj is not None + assert obj.current_fidelity == FidelityLevel.L4 + + # Upgrade one back to L0 + result = fm.upgrade(ids[0], FidelityLevel.L0, current_turn=11) + assert result is True + upgraded = fm.get_object(ids[0]) + assert upgraded is not None + assert upgraded.current_fidelity == FidelityLevel.L0 + + def test_fault_pin_degrade_cycle(self): + """Fault → pin → degrade respects pin → pin expires → degrade works.""" + fm = FidelityManager(window_size=1000) + ids = _fill_manager(fm, 4, content_size=800) + + # Record fault on first object + fm.record_fault(ids[0], current_turn=5, pin_duration=3) + + # Degrade at turn 6 — pinned object survives + transitions = fm.degrade(current_turn=6) + degraded_ids = {oid for oid, _, _ in transitions} + assert ids[0] not in degraded_ids + + # At turn 9, pin expired — now it can be degraded + obj = fm.get_object(ids[0]) + assert obj is not None + # Reset to L0 for clean test + obj.current_fidelity = FidelityLevel.L0 + + # Re-fill to get pressure back up + _fill_manager(fm, 2, content_size=800, turn=9) + + fm.degrade(current_turn=9) + # Now the previously-pinned object should be degradable + obj_after = fm.get_object(ids[0]) + assert obj_after is not None + assert obj_after.pinned is False + + def test_token_accounting_through_degradation(self): + """Token count decreases as objects are degraded.""" + fm = FidelityManager(window_size=1000) + + for i in range(3): + obj = _make_obj( + "a" * 800, # 200 tokens at L0 + summary_detailed="b" * 120, # 30 tokens at L1 + turn=i, + ) + fm.register_object(obj) + + initial_tokens = fm.total_tokens() + assert initial_tokens == 600 # 3 * 200 + + fm.degrade(current_turn=10) + + # Tokens should have decreased + assert fm.total_tokens() < initial_tokens + + def test_multiple_degrade_passes(self): + """Multiple degrade calls progressively reduce fidelity.""" + fm = FidelityManager(window_size=200) + + for i in range(3): + obj = _make_obj( + "a" * 400, # 100 tokens at L0 + summary_detailed="b" * 120, # 30 tokens at L1 + summary_compact="c" * 20, # 5 tokens at L2 + stub="stub", + turn=i, + ) + fm.register_object(obj) + + # 300 tokens in 200 window → EMERGENCY + fm.degrade(current_turn=10) + + # After emergency, all should be L4 + for obj in fm._objects.values(): + assert obj.current_fidelity == FidelityLevel.L4 diff --git a/tests/test_gateway_fidelity.py b/tests/test_gateway_fidelity.py new file mode 100644 index 0000000..f34e570 --- /dev/null +++ b/tests/test_gateway_fidelity.py @@ -0,0 +1,513 @@ +"""Integration tests for FidelityManager integration in the gateway. + +Tests the Phase 2.4 fidelity pipeline: object registration, pressure +calculation, degradation, and content replacement in ephemeral payloads. +Does NOT test HelperLLM integration (requires mocking). +""" + +from __future__ import annotations + +import copy +from pathlib import Path +from tempfile import TemporaryDirectory + +import pytest + +from mnemosyne.fidelity import FidelityLevel, FidelityManager, PressureZone, make_object +from mnemosyne.gateway import ( + Session, + _apply_fidelity, + _auto_stub, + _block_text, + _content_key, +) + + +# ── Fixtures ───────────────────────────────────────────────────────────── + + +@pytest.fixture +def tmp_log_dir(): + with TemporaryDirectory() as d: + yield Path(d) + + +@pytest.fixture +def session(tmp_log_dir): + return Session("test01", tmp_log_dir) + + +def _make_large_text(size: int = 600) -> str: + """Generate a text string of approximately `size` bytes.""" + return "x" * size + + +def _tool_result_block(tool_use_id: str, content: str) -> dict: + return { + "type": "tool_result", + "tool_use_id": tool_use_id, + "content": content, + } + + +def _text_block(text: str) -> dict: + return {"type": "text", "text": text} + + +def _msg(role: str, blocks: list[dict]) -> dict: + return {"role": role, "content": blocks} + + +# ── Test: FidelityManager created per session ──────────────────────────── + + +class TestSessionFidelityManager: + def test_session_has_fidelity_manager(self, session): + assert hasattr(session, "fidelity_manager") + assert isinstance(session.fidelity_manager, FidelityManager) + + def test_fidelity_manager_default_window(self, session): + assert session.fidelity_manager.window_size == 200_000 + + def test_fidelity_manager_per_session(self, tmp_log_dir): + s1 = Session("sess_a", tmp_log_dir) + s2 = Session("sess_b", tmp_log_dir) + assert s1.fidelity_manager is not s2.fidelity_manager + + def test_session_has_content_map(self, session): + assert hasattr(session, "_fidelity_content_map") + assert isinstance(session._fidelity_content_map, dict) + assert len(session._fidelity_content_map) == 0 + + +# ── Test: Content key derivation ───────────────────────────────────────── + + +class TestContentKey: + def test_tool_result_key(self): + block = _tool_result_block("toolu_abc123", "some content") + key = _content_key(block, {}) + assert key == "tool:toolu_abc123" + + def test_large_text_key(self): + text = _make_large_text(600) + block = _text_block(text) + key = _content_key(block, {}) + assert key is not None + assert key.startswith("text:") + + def test_small_text_returns_none(self): + block = _text_block("short") + key = _content_key(block, {}) + assert key is None + + def test_stable_key_for_same_content(self): + text = _make_large_text(600) + block1 = _text_block(text) + block2 = _text_block(text) + assert _content_key(block1, {}) == _content_key(block2, {}) + + def test_different_key_for_different_content(self): + block1 = _text_block("a" * 600) + block2 = _text_block("b" * 600) + assert _content_key(block1, {}) != _content_key(block2, {}) + + +# ── Test: Block text extraction ────────────────────────────────────────── + + +class TestBlockText: + def test_text_block(self): + assert _block_text(_text_block("hello")) == "hello" + + def test_tool_result_string_content(self): + block = _tool_result_block("id1", "result text") + assert _block_text(block) == "result text" + + def test_tool_result_list_content(self): + block = { + "type": "tool_result", + "tool_use_id": "id2", + "content": [ + {"type": "text", "text": "line 1"}, + {"type": "text", "text": "line 2"}, + ], + } + assert _block_text(block) == "line 1\nline 2" + + +# ── Test: Auto stub generation ─────────────────────────────────────────── + + +class TestAutoStub: + def test_short_content(self): + stub = _auto_stub("Hello world") + assert stub == "[evicted content: Hello world]" + + def test_multiline_uses_first_line(self): + stub = _auto_stub("First line\nSecond line\nThird line") + assert "First line" in stub + assert "Second line" not in stub + + def test_long_first_line_truncated(self): + long_line = "x" * 200 + stub = _auto_stub(long_line) + assert len(stub) < 200 + assert "..." in stub + + +# ── Test: Object registration via _apply_fidelity ─────────────────────── + + +class TestApplyFidelityRegistration: + def test_registers_tool_result(self, session): + session.token_state["turn"] = 1 + payload = { + "messages": [ + _msg("user", [_tool_result_block("tool_1", _make_large_text(600))]), + ] + } + _apply_fidelity(payload, session) + + fm = session.fidelity_manager + assert fm.total_tokens() > 0 + assert len(session._fidelity_content_map) == 1 + assert "tool:tool_1" in session._fidelity_content_map + + def test_registers_large_text_block(self, session): + session.token_state["turn"] = 1 + large_text = _make_large_text(800) + payload = { + "messages": [ + _msg("assistant", [_text_block(large_text)]), + ] + } + _apply_fidelity(payload, session) + + fm = session.fidelity_manager + assert fm.total_tokens() > 0 + assert len(session._fidelity_content_map) == 1 + + def test_ignores_small_text_block(self, session): + session.token_state["turn"] = 1 + payload = { + "messages": [ + _msg("user", [_text_block("small content")]), + ] + } + _apply_fidelity(payload, session) + + assert len(session._fidelity_content_map) == 0 + + def test_multiple_blocks_registered(self, session): + session.token_state["turn"] = 1 + payload = { + "messages": [ + _msg( + "user", + [ + _tool_result_block("tool_a", _make_large_text(600)), + _tool_result_block("tool_b", _make_large_text(700)), + _text_block(_make_large_text(800)), + ], + ), + ] + } + _apply_fidelity(payload, session) + + assert len(session._fidelity_content_map) == 3 + + def test_idempotent_on_second_call(self, session): + session.token_state["turn"] = 1 + payload = { + "messages": [ + _msg("user", [_tool_result_block("tool_1", _make_large_text(600))]), + ] + } + _apply_fidelity(payload, session) + count_after_first = len(session._fidelity_content_map) + + # Second call with same content — should not re-register + payload2 = { + "messages": [ + _msg("user", [_tool_result_block("tool_1", _make_large_text(600))]), + ] + } + session.token_state["turn"] = 2 + _apply_fidelity(payload2, session) + assert len(session._fidelity_content_map) == count_after_first + + def test_new_objects_start_at_l0(self, session): + session.token_state["turn"] = 1 + payload = { + "messages": [ + _msg("user", [_tool_result_block("tool_1", _make_large_text(600))]), + ] + } + _apply_fidelity(payload, session) + + obj_id = session._fidelity_content_map["tool:tool_1"] + obj = session.fidelity_manager.get_object(obj_id) + assert obj is not None + assert obj.current_fidelity == FidelityLevel.L0 + + +# ── Test: Pressure calculation with real payloads ──────────────────────── + + +class TestPressureCalculation: + def test_normal_pressure_small_payload(self, session): + session.token_state["turn"] = 1 + payload = { + "messages": [ + _msg("user", [_tool_result_block("t1", _make_large_text(600))]), + ] + } + _apply_fidelity(payload, session) + + # 600 bytes ≈ 150 tokens, window=200k → well under 50% + assert session.fidelity_manager.current_pressure() == PressureZone.NORMAL + + def test_high_pressure_large_payload(self, session): + # Fill the fidelity manager with enough objects to exceed caution threshold + fm = session.fidelity_manager + # 200k window, caution at 50% = 100k tokens + # Each object: ~25k tokens (100k chars / 4) + for i in range(5): + obj = make_object( + object_type="tool_result", + content_full="x" * 100_000, + created_at_turn=1, + stub=f"[stub {i}]", + ) + fm.register_object(obj) + + # 5 * 25k = 125k tokens > 100k caution threshold + assert fm.current_pressure() >= PressureZone.CAUTION + + def test_pressure_zones_ordered(self, session): + fm = session.fidelity_manager + # Verify zone thresholds are ordered correctly + assert fm.threshold_caution < fm.threshold_warning + assert fm.threshold_warning < fm.threshold_critical + assert fm.threshold_critical < fm.threshold_emergency + + +# ── Test: Degradation replaces content in ephemeral payload ────────────── + + +class TestDegradationReplacement: + def test_degraded_tool_result_replaced_with_stub(self, session): + session.token_state["turn"] = 1 + original_content = _make_large_text(600) + tool_id = "tool_degrade" + + # Register the object + payload1 = { + "messages": [ + _msg("user", [_tool_result_block(tool_id, original_content)]), + ] + } + _apply_fidelity(payload1, session) + + # Manually degrade the object to L3 (stub) + obj_id = session._fidelity_content_map[f"tool:{tool_id}"] + obj = session.fidelity_manager.get_object(obj_id) + obj.current_fidelity = FidelityLevel.L3 + + # Apply fidelity again — should replace content with stub + session.token_state["turn"] = 2 + payload2 = { + "messages": [ + _msg("user", [_tool_result_block(tool_id, original_content)]), + ] + } + _apply_fidelity(payload2, session) + + replaced_content = payload2["messages"][0]["content"][0]["content"] + assert replaced_content != original_content + assert "[evicted content:" in replaced_content + + def test_degraded_text_block_replaced_with_stub(self, session): + session.token_state["turn"] = 1 + original_text = "Important data: " + _make_large_text(600) + + payload1 = { + "messages": [ + _msg("assistant", [_text_block(original_text)]), + ] + } + _apply_fidelity(payload1, session) + + # Find the content key and degrade + assert len(session._fidelity_content_map) == 1 + obj_id = list(session._fidelity_content_map.values())[0] + obj = session.fidelity_manager.get_object(obj_id) + obj.current_fidelity = FidelityLevel.L3 + + # Apply again + session.token_state["turn"] = 2 + payload2 = { + "messages": [ + _msg("assistant", [_text_block(original_text)]), + ] + } + _apply_fidelity(payload2, session) + + replaced_text = payload2["messages"][0]["content"][0]["text"] + assert replaced_text != original_text + assert "[evicted content:" in replaced_text + + def test_l0_content_not_replaced(self, session): + session.token_state["turn"] = 1 + original_content = _make_large_text(600) + tool_id = "tool_keep" + + payload = { + "messages": [ + _msg("user", [_tool_result_block(tool_id, original_content)]), + ] + } + _apply_fidelity(payload, session) + + # Object stays at L0 — content should be unchanged + result_content = payload["messages"][0]["content"][0]["content"] + assert result_content == original_content + + def test_l4_evicted_also_replaced(self, session): + session.token_state["turn"] = 1 + original_content = _make_large_text(600) + tool_id = "tool_evict" + + payload1 = { + "messages": [ + _msg("user", [_tool_result_block(tool_id, original_content)]), + ] + } + _apply_fidelity(payload1, session) + + # Degrade to L4 (evicted) + obj_id = session._fidelity_content_map[f"tool:{tool_id}"] + obj = session.fidelity_manager.get_object(obj_id) + obj.current_fidelity = FidelityLevel.L4 + + # Apply again — L4 is >= L3, so should still replace + session.token_state["turn"] = 2 + payload2 = { + "messages": [ + _msg("user", [_tool_result_block(tool_id, original_content)]), + ] + } + _apply_fidelity(payload2, session) + + replaced_content = payload2["messages"][0]["content"][0]["content"] + assert "[evicted content:" in replaced_content + + +# ── Test: Degradation triggered by pressure ────────────────────────────── + + +class TestPressureDegradation: + def test_degrade_under_pressure(self, session): + """When pressure exceeds NORMAL, _apply_fidelity triggers degradation.""" + fm = session.fidelity_manager + # Use a tiny window to force pressure + fm.window_size = 1000 # 1000 tokens + + session.token_state["turn"] = 1 + # Register objects that exceed the window + # 2000 bytes ≈ 500 tokens per object, 3 objects = 1500 tokens > 1000 + payload = { + "messages": [ + _msg( + "user", + [ + _tool_result_block("t1", _make_large_text(2000)), + _tool_result_block("t2", _make_large_text(2000)), + _tool_result_block("t3", _make_large_text(2000)), + ], + ), + ] + } + _apply_fidelity(payload, session) + + # After apply, some objects should have been degraded + degraded = [obj for obj in fm._objects.values() if obj.current_fidelity > FidelityLevel.L0] + # At least some degradation should have occurred + assert len(degraded) > 0 or fm.current_pressure() == PressureZone.NORMAL + + def test_mark_accessed_updates_turn(self, session): + """Objects seen again get their last_accessed_turn updated.""" + session.token_state["turn"] = 1 + payload = { + "messages": [ + _msg("user", [_tool_result_block("t1", _make_large_text(600))]), + ] + } + _apply_fidelity(payload, session) + + obj_id = session._fidelity_content_map["tool:t1"] + obj = session.fidelity_manager.get_object(obj_id) + assert obj.last_accessed_turn == 1 + + # See it again at turn 5 + session.token_state["turn"] = 5 + payload2 = { + "messages": [ + _msg("user", [_tool_result_block("t1", _make_large_text(600))]), + ] + } + _apply_fidelity(payload2, session) + assert obj.last_accessed_turn == 5 + + +# ── Test: Mixed content payloads ───────────────────────────────────────── + + +class TestMixedPayloads: + def test_mixed_small_and_large_blocks(self, session): + """Only large blocks get tracked; small ones pass through.""" + session.token_state["turn"] = 1 + payload = { + "messages": [ + _msg( + "user", + [ + _text_block("small question"), + _tool_result_block("t1", _make_large_text(600)), + _text_block("another small bit"), + ], + ), + ] + } + _apply_fidelity(payload, session) + + # Only the tool_result should be tracked + assert len(session._fidelity_content_map) == 1 + assert "tool:t1" in session._fidelity_content_map + + def test_string_content_messages_ignored(self, session): + """Messages with string content (not list) are skipped.""" + session.token_state["turn"] = 1 + payload = { + "messages": [ + {"role": "user", "content": "Just a plain string message"}, + ] + } + _apply_fidelity(payload, session) + assert len(session._fidelity_content_map) == 0 + + def test_multiple_messages_all_tracked(self, session): + """Objects across multiple messages are all registered.""" + session.token_state["turn"] = 1 + payload = { + "messages": [ + _msg("user", [_tool_result_block("t1", _make_large_text(600))]), + _msg("assistant", [_text_block(_make_large_text(800))]), + _msg("user", [_tool_result_block("t2", _make_large_text(700))]), + ] + } + _apply_fidelity(payload, session) + + # t1, large text, t2 = 3 tracked objects + assert len(session._fidelity_content_map) == 3 diff --git a/tests/test_helper_llm.py b/tests/test_helper_llm.py new file mode 100644 index 0000000..c3c4f7e --- /dev/null +++ b/tests/test_helper_llm.py @@ -0,0 +1,572 @@ +"""Tests for the Helper LLM client. + +All tests use mocked Anthropic API responses — no real API calls. +""" + +from __future__ import annotations + +import json +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +from mnemosyne.helper_llm import ( + GoalClassification, + HelperLLM, + SummaryResult, + _L0_TO_L1_PROMPT, + _L1_TO_L2_PROMPT, + _GOAL_CLASSIFICATION_PROMPT, + _MICRO_FAULT_PROMPT, +) + + +# --------------------------------------------------------------------------- +# Fixtures +# --------------------------------------------------------------------------- + + +def _make_api_response(text: str) -> MagicMock: + """Build a mock Anthropic Messages response with a single text block.""" + block = MagicMock() + block.type = "text" + block.text = text + response = MagicMock() + response.content = [block] + return response + + +@pytest.fixture() +def helper() -> HelperLLM: + """Return a HelperLLM with a mocked async client.""" + with patch.dict("os.environ", {"ANTHROPIC_API_KEY": "test-key-000"}): + h = HelperLLM() + h._client = MagicMock() + h._client.messages = MagicMock() + h._client.messages.create = AsyncMock() + return h + + +# --------------------------------------------------------------------------- +# Dataclass tests +# --------------------------------------------------------------------------- + + +class TestDataclasses: + def test_summary_result_defaults(self) -> None: + r = SummaryResult(summary="hello") + assert r.summary == "hello" + assert r.losses == [] + assert r.can_answer == [] + assert r.key_entities == [] + + def test_summary_result_with_fields(self) -> None: + r = SummaryResult( + summary="s", + losses=["a"], + can_answer=["b"], + key_entities=["c"], + ) + assert r.losses == ["a"] + assert r.can_answer == ["b"] + assert r.key_entities == ["c"] + + def test_goal_classification_defaults(self) -> None: + g = GoalClassification(goal="write tests") + assert g.goal == "write tests" + assert g.relevant_types == [] + assert g.relevant_tags == [] + assert g.predicted_needs == [] + + +# --------------------------------------------------------------------------- +# JSON parsing tests (static methods, no API calls) +# --------------------------------------------------------------------------- + + +class TestParseSummaryJson: + def test_valid_json(self) -> None: + raw = json.dumps( + { + "summary": "Auth middleware uses JWT.", + "losses": ["exact error codes"], + "can_answer": ["auth approach"], + "key_entities": ["src/auth/middleware.ts"], + } + ) + result = HelperLLM._parse_summary_json(raw) + assert result is not None + assert result.summary == "Auth middleware uses JWT." + assert result.losses == ["exact error codes"] + assert result.can_answer == ["auth approach"] + assert result.key_entities == ["src/auth/middleware.ts"] + + def test_json_in_code_fence(self) -> None: + raw = ( + '```json\n{"summary": "test", "losses": [], "can_answer": [], "key_entities": []}\n```' + ) + result = HelperLLM._parse_summary_json(raw) + assert result is not None + assert result.summary == "test" + + def test_json_with_surrounding_text(self) -> None: + raw = 'Here is the result:\n{"summary": "ok", "losses": ["x"]}\nDone.' + result = HelperLLM._parse_summary_json(raw) + assert result is not None + assert result.summary == "ok" + assert result.losses == ["x"] + + def test_invalid_json_returns_none(self) -> None: + assert HelperLLM._parse_summary_json("not json at all") is None + + def test_empty_string_returns_none(self) -> None: + assert HelperLLM._parse_summary_json("") is None + + def test_non_dict_json_returns_none(self) -> None: + assert HelperLLM._parse_summary_json("[1, 2, 3]") is None + + def test_missing_fields_default_to_empty(self) -> None: + raw = '{"summary": "minimal"}' + result = HelperLLM._parse_summary_json(raw) + assert result is not None + assert result.summary == "minimal" + assert result.losses == [] + assert result.can_answer == [] + assert result.key_entities == [] + + +class TestParseGoalJson: + def test_valid_json(self) -> None: + raw = json.dumps( + { + "goal": "write auth tests", + "relevant_types": ["file_context", "design_decision"], + "relevant_tags": ["auth", "testing"], + "predicted_needs": ["auth middleware impl"], + } + ) + result = HelperLLM._parse_goal_json(raw) + assert result is not None + assert result.goal == "write auth tests" + assert "file_context" in result.relevant_types + assert "auth" in result.relevant_tags + + def test_invalid_json_returns_none(self) -> None: + assert HelperLLM._parse_goal_json("garbage") is None + + def test_json_in_code_fence(self) -> None: + raw = '```\n{"goal": "deploy", "relevant_types": [], "relevant_tags": [], "predicted_needs": []}\n```' + result = HelperLLM._parse_goal_json(raw) + assert result is not None + assert result.goal == "deploy" + + +# --------------------------------------------------------------------------- +# Prompt construction tests +# --------------------------------------------------------------------------- + + +class TestPromptConstruction: + def test_l0_to_l1_prompt_contains_content_and_type(self) -> None: + prompt = _L0_TO_L1_PROMPT.format( + object_type="file_context", + content="def hello(): pass", + ) + assert "file_context" in prompt + assert "def hello(): pass" in prompt + assert "DECLARED LOSSES" in prompt + assert "CAN_ANSWER" in prompt + assert "OUTPUT FORMAT (JSON)" in prompt + + def test_l1_to_l2_prompt_includes_losses(self) -> None: + prompt = _L1_TO_L2_PROMPT.format( + object_type="debugging_session", + l1_summary="Fixed race condition", + l1_losses="- exact error codes\n- line numbers", + ) + assert "debugging_session" in prompt + assert "Fixed race condition" in prompt + assert "exact error codes" in prompt + + def test_goal_prompt_includes_message_and_context(self) -> None: + prompt = _GOAL_CLASSIFICATION_PROMPT.format( + user_message="now write tests", + recent_context="We just implemented auth.", + ) + assert "now write tests" in prompt + assert "We just implemented auth." in prompt + + def test_micro_fault_prompt_includes_question_and_context(self) -> None: + prompt = _MICRO_FAULT_PROMPT.format( + question="What error code?", + context="Error 401 unauthorized", + ) + assert "What error code?" in prompt + assert "Error 401 unauthorized" in prompt + + +# --------------------------------------------------------------------------- +# API call tests (mocked) +# --------------------------------------------------------------------------- + + +class TestSummarizeL0ToL1: + async def test_successful_summarization(self, helper: HelperLLM) -> None: + api_response = _make_api_response( + json.dumps( + { + "summary": "Auth middleware validates JWT tokens.", + "losses": ["exact error codes for token expiry"], + "can_answer": ["auth approach used"], + "key_entities": ["src/auth/middleware.ts", "jsonwebtoken"], + } + ) + ) + helper._client.messages.create = AsyncMock(return_value=api_response) + + result = await helper.summarize_l0_to_l1( + content="Full auth middleware source code here...", + object_type="file_context", + ) + + assert isinstance(result, SummaryResult) + assert result.summary == "Auth middleware validates JWT tokens." + assert "exact error codes for token expiry" in result.losses + assert "src/auth/middleware.ts" in result.key_entities + + # Verify the API was called with correct model + call_kwargs = helper._client.messages.create.call_args.kwargs + assert call_kwargs["model"] == "claude-haiku-4-5-20251001" + # Verify prompt contains the content + user_msg = call_kwargs["messages"][0]["content"] + assert "file_context" in user_msg + assert "Full auth middleware source code here..." in user_msg + + async def test_json_parse_failure_returns_fallback(self, helper: HelperLLM) -> None: + api_response = _make_api_response("This is not valid JSON at all.") + helper._client.messages.create = AsyncMock(return_value=api_response) + + content = "A" * 1000 + result = await helper.summarize_l0_to_l1( + content=content, + object_type="file_context", + ) + + assert isinstance(result, SummaryResult) + # Fallback: first 30% of content + assert len(result.summary) == 300 + assert result.summary == "A" * 300 + assert result.losses == [] + + async def test_respects_max_summary_tokens(self, helper: HelperLLM) -> None: + api_response = _make_api_response('{"summary": "ok"}') + helper._client.messages.create = AsyncMock(return_value=api_response) + + await helper.summarize_l0_to_l1( + content="test", + object_type="tool_result", + max_summary_tokens=512, + ) + + call_kwargs = helper._client.messages.create.call_args.kwargs + assert call_kwargs["max_tokens"] == 512 + + +class TestCompressL1ToL2: + async def test_successful_compression(self, helper: HelperLLM) -> None: + api_response = _make_api_response( + json.dumps( + { + "summary": "Auth uses JWT with refresh tokens.", + "losses": ["function signatures"], + "can_answer": ["what was decided"], + "key_entities": ["middleware.ts"], + } + ) + ) + helper._client.messages.create = AsyncMock(return_value=api_response) + + l1_losses = ["exact error codes", "line-by-line implementation"] + result = await helper.compress_l1_to_l2( + l1_summary="Detailed auth summary...", + l1_losses=l1_losses, + object_type="file_context", + ) + + assert isinstance(result, SummaryResult) + assert result.summary == "Auth uses JWT with refresh tokens." + # L1 losses should be accumulated (prepended) + assert result.losses[0] == "exact error codes" + assert result.losses[1] == "line-by-line implementation" + assert "function signatures" in result.losses + + async def test_loss_accumulation(self, helper: HelperLLM) -> None: + api_response = _make_api_response( + json.dumps( + { + "summary": "compact", + "losses": ["new_loss"], + "can_answer": [], + "key_entities": [], + } + ) + ) + helper._client.messages.create = AsyncMock(return_value=api_response) + + result = await helper.compress_l1_to_l2( + l1_summary="summary", + l1_losses=["old_loss_1", "old_loss_2"], + object_type="design_decision", + ) + + assert result.losses == ["old_loss_1", "old_loss_2", "new_loss"] + + async def test_fallback_on_parse_failure(self, helper: HelperLLM) -> None: + api_response = _make_api_response("broken response") + helper._client.messages.create = AsyncMock(return_value=api_response) + + l1_losses = ["loss_a"] + result = await helper.compress_l1_to_l2( + l1_summary="A" * 100, + l1_losses=l1_losses, + object_type="file_context", + ) + + assert isinstance(result, SummaryResult) + # Fallback: 30% of l1_summary + assert len(result.summary) == 30 + # L1 losses preserved in fallback + assert result.losses == ["loss_a"] + + +class TestGenerateStub: + async def test_successful_stub(self, helper: HelperLLM) -> None: + stub_text = "[debugging_session | 2026-03-13 14:30 | Fixed race condition in auth | 12 related objects]" + api_response = _make_api_response(stub_text) + helper._client.messages.create = AsyncMock(return_value=api_response) + + result = await helper.generate_stub( + l2_summary="Fixed race condition in auth token refresh.", + object_type="debugging_session", + timestamp="2026-03-13 14:30", + ) + + assert result == stub_text + assert "debugging_session" in result + assert "2026-03-13 14:30" in result + + async def test_multiline_response_takes_first_line(self, helper: HelperLLM) -> None: + api_response = _make_api_response( + "[type | ts | desc | 0 related objects]\nExtra line\nAnother" + ) + helper._client.messages.create = AsyncMock(return_value=api_response) + + result = await helper.generate_stub("summary", "type", "ts") + assert "\n" not in result + assert result == "[type | ts | desc | 0 related objects]" + + async def test_empty_response_fallback(self, helper: HelperLLM) -> None: + api_response = _make_api_response("") + helper._client.messages.create = AsyncMock(return_value=api_response) + + result = await helper.generate_stub("summary", "file_context", "2026-01-01") + assert "file_context" in result + assert "2026-01-01" in result + assert "summary unavailable" in result + + +class TestAnswerMicroFault: + async def test_successful_answer(self, helper: HelperLLM) -> None: + api_response = _make_api_response("The error code is 401 UNAUTHORIZED.") + helper._client.messages.create = AsyncMock(return_value=api_response) + + result = await helper.answer_micro_fault( + question="What error code does auth return for expired tokens?", + relevant_contents=[ + "Auth middleware returns 401 for expired tokens.", + "Token refresh logic in refresh.ts.", + ], + ) + + assert result == "The error code is 401 UNAUTHORIZED." + + # Verify context was joined with separator + call_kwargs = helper._client.messages.create.call_args.kwargs + prompt = call_kwargs["messages"][0]["content"] + assert "What error code" in prompt + assert "Auth middleware returns 401" in prompt + assert "---" in prompt # separator between contents + + async def test_empty_response_fallback(self, helper: HelperLLM) -> None: + api_response = _make_api_response(" ") + helper._client.messages.create = AsyncMock(return_value=api_response) + + result = await helper.answer_micro_fault( + question="anything", + relevant_contents=["content"], + ) + + assert result == "Unable to answer from available context." + + async def test_respects_max_tokens(self, helper: HelperLLM) -> None: + api_response = _make_api_response("answer") + helper._client.messages.create = AsyncMock(return_value=api_response) + + await helper.answer_micro_fault("q", ["c"], max_tokens=150) + + call_kwargs = helper._client.messages.create.call_args.kwargs + assert call_kwargs["max_tokens"] == 150 + + +class TestClassifyGoal: + async def test_successful_classification(self, helper: HelperLLM) -> None: + api_response = _make_api_response( + json.dumps( + { + "goal": "write unit tests for auth module", + "relevant_types": ["file_context", "design_decision"], + "relevant_tags": ["auth", "testing"], + "predicted_needs": ["auth middleware implementation", "test patterns"], + } + ) + ) + helper._client.messages.create = AsyncMock(return_value=api_response) + + result = await helper.classify_goal( + user_message="Now write tests for the auth middleware", + recent_context="We just finished implementing JWT auth.", + ) + + assert isinstance(result, GoalClassification) + assert result.goal == "write unit tests for auth module" + assert "file_context" in result.relevant_types + assert "auth" in result.relevant_tags + assert "auth middleware implementation" in result.predicted_needs + + async def test_fallback_on_parse_failure(self, helper: HelperLLM) -> None: + api_response = _make_api_response("I don't understand the format") + helper._client.messages.create = AsyncMock(return_value=api_response) + + result = await helper.classify_goal( + user_message="deploy to production", + recent_context="context", + ) + + assert isinstance(result, GoalClassification) + assert result.goal == "deploy to production" + assert result.relevant_types == [] + + +# --------------------------------------------------------------------------- +# Error handling tests +# --------------------------------------------------------------------------- + + +class TestErrorHandling: + async def test_timeout_returns_fallback(self, helper: HelperLLM) -> None: + import anthropic as anth + + helper._client.messages.create = AsyncMock( + side_effect=anth.APITimeoutError(request=MagicMock()) + ) + + result = await helper.summarize_l0_to_l1( + content="some content here", + object_type="file_context", + ) + + # Should get fallback (30% of content) + assert isinstance(result, SummaryResult) + assert len(result.summary) == 5 # 30% of 18 chars ≈ 5 + + async def test_api_error_returns_fallback(self, helper: HelperLLM) -> None: + import anthropic as anth + + helper._client.messages.create = AsyncMock( + side_effect=anth.APIError( + message="Internal server error", + request=MagicMock(), + body=None, + ) + ) + + result = await helper.summarize_l0_to_l1( + content="test content", + object_type="tool_result", + ) + + assert isinstance(result, SummaryResult) + # Fallback summary + assert result.summary == "tes" # 30% of 12 chars = 3 + + async def test_timeout_on_micro_fault(self, helper: HelperLLM) -> None: + import anthropic as anth + + helper._client.messages.create = AsyncMock( + side_effect=anth.APITimeoutError(request=MagicMock()) + ) + + result = await helper.answer_micro_fault("question", ["content"]) + assert result == "Unable to answer from available context." + + async def test_timeout_on_goal_classification(self, helper: HelperLLM) -> None: + import anthropic as anth + + helper._client.messages.create = AsyncMock( + side_effect=anth.APITimeoutError(request=MagicMock()) + ) + + result = await helper.classify_goal("do something", "context") + assert isinstance(result, GoalClassification) + assert result.goal == "do something" + + async def test_timeout_on_generate_stub(self, helper: HelperLLM) -> None: + import anthropic as anth + + helper._client.messages.create = AsyncMock( + side_effect=anth.APITimeoutError(request=MagicMock()) + ) + + result = await helper.generate_stub("summary", "file_context", "2026-01-01") + assert "file_context" in result + assert "summary unavailable" in result + + +# --------------------------------------------------------------------------- +# Constructor tests +# --------------------------------------------------------------------------- + + +class TestConstructor: + def test_uses_provided_api_key(self) -> None: + with patch("mnemosyne.helper_llm.anthropic.AsyncAnthropic") as mock_cls: + HelperLLM(api_key="sk-test-123") + call_kwargs = mock_cls.call_args.kwargs + assert call_kwargs["api_key"] == "sk-test-123" + + def test_falls_back_to_env_var(self) -> None: + with ( + patch.dict("os.environ", {"ANTHROPIC_API_KEY": "sk-env-456"}), + patch("mnemosyne.helper_llm.anthropic.AsyncAnthropic") as mock_cls, + ): + HelperLLM() + call_kwargs = mock_cls.call_args.kwargs + assert call_kwargs["api_key"] == "sk-env-456" + + def test_custom_model_and_base_url(self) -> None: + with patch("mnemosyne.helper_llm.anthropic.AsyncAnthropic") as mock_cls: + h = HelperLLM( + api_key="key", + model="claude-3-haiku-20240307", + base_url="http://localhost:8080", + ) + call_kwargs = mock_cls.call_args.kwargs + assert call_kwargs["base_url"] == "http://localhost:8080" + assert h._model == "claude-3-haiku-20240307" + + def test_default_timeout_and_retries(self) -> None: + with patch("mnemosyne.helper_llm.anthropic.AsyncAnthropic") as mock_cls: + HelperLLM(api_key="key") + call_kwargs = mock_cls.call_args.kwargs + assert call_kwargs["timeout"] == 10.0 + assert call_kwargs["max_retries"] == 2