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:
parent
974863e7b3
commit
d26c56c2f0
5 changed files with 3003 additions and 0 deletions
861
tests/test_fidelity.py
Normal file
861
tests/test_fidelity.py
Normal 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
|
||||
513
tests/test_gateway_fidelity.py
Normal file
513
tests/test_gateway_fidelity.py
Normal 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
572
tests/test_helper_llm.py
Normal 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
|
||||
Loading…
Add table
Add a link
Reference in a new issue