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 <clio-agent@sisyphuslabs.ai>
This commit is contained in:
Joey Yakimowich-Payne 2026-03-13 11:40:56 -06:00
commit d26c56c2f0
5 changed files with 3003 additions and 0 deletions

861
tests/test_fidelity.py Normal file
View file

@ -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

View file

@ -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

572
tests/test_helper_llm.py Normal file
View file

@ -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