feat: add memory management pipeline
Admission control, entropy-based micro-faulting, phantom tool injection for backing store queries, and xMemory session hierarchy for long conversations (50+ turns). 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
a13719f754
commit
681c1454cb
9 changed files with 3540 additions and 0 deletions
156
tests/test_admission.py
Normal file
156
tests/test_admission.py
Normal file
|
|
@ -0,0 +1,156 @@
|
|||
"""Tests for the admission control module.
|
||||
|
||||
Tests the AdmissionController scoring logic across all four axes:
|
||||
type_score, novelty_score, utility_score, size_score, and the
|
||||
threshold-based admit/reject decision.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from mnemosyne.admission import AdmissionController, AdmissionScore, DEFAULT_THRESHOLD
|
||||
|
||||
|
||||
# ── Type score ───────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_admission_score_design_decision():
|
||||
"""design_decision has the highest type weight (0.9)."""
|
||||
ctrl = AdmissionController()
|
||||
score = ctrl.score("x" * 300, "design_decision", has_duplicate=False)
|
||||
assert score.type_score == 0.9
|
||||
|
||||
|
||||
def test_admission_score_conversation_phase():
|
||||
"""conversation_phase has the lowest type weight (0.3)."""
|
||||
ctrl = AdmissionController()
|
||||
score = ctrl.score("x" * 300, "conversation_phase", has_duplicate=False)
|
||||
assert score.type_score == 0.3
|
||||
|
||||
|
||||
def test_admission_unknown_type_default():
|
||||
"""Unknown object types default to 0.3."""
|
||||
ctrl = AdmissionController()
|
||||
score = ctrl.score("x" * 300, "totally_unknown_type", has_duplicate=False)
|
||||
assert score.type_score == 0.3
|
||||
|
||||
|
||||
# ── Novelty score ────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_admission_duplicate_penalized():
|
||||
"""Duplicates get novelty_score=0.2 instead of 1.0."""
|
||||
ctrl = AdmissionController()
|
||||
score_unique = ctrl.score("x" * 300, "file_context", has_duplicate=False)
|
||||
score_dup = ctrl.score("x" * 300, "file_context", has_duplicate=True)
|
||||
assert score_unique.novelty_score == 1.0
|
||||
assert score_dup.novelty_score == 0.2
|
||||
assert score_dup.total < score_unique.total
|
||||
|
||||
|
||||
# ── Utility score ────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_admission_utility_code_blocks():
|
||||
"""Content with code blocks gets utility boost."""
|
||||
ctrl = AdmissionController()
|
||||
content_with_code = "Here is code:\n```python\ndef foo(): pass\n```\n" + "x" * 200
|
||||
content_plain = "Just some plain text without code blocks. " + "x" * 200
|
||||
score_code = ctrl.score(content_with_code, "file_context", has_duplicate=False)
|
||||
score_plain = ctrl.score(content_plain, "file_context", has_duplicate=False)
|
||||
assert score_code.utility_score > score_plain.utility_score
|
||||
|
||||
|
||||
def test_admission_utility_entities():
|
||||
"""Content with many key_entities gets utility boost."""
|
||||
ctrl = AdmissionController()
|
||||
entities = ["foo.py", "bar.py", "baz.py"]
|
||||
score_with = ctrl.score("x" * 300, "file_context", has_duplicate=False, key_entities=entities)
|
||||
score_without = ctrl.score("x" * 300, "file_context", has_duplicate=False, key_entities=[])
|
||||
assert score_with.utility_score > score_without.utility_score
|
||||
|
||||
|
||||
# ── Size score ───────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_admission_size_small_penalized():
|
||||
"""Very small content (< 100 chars) gets penalized size_score."""
|
||||
ctrl = AdmissionController()
|
||||
score = ctrl.score("tiny", "file_context", has_duplicate=False)
|
||||
assert score.size_score < 1.0
|
||||
assert score.size_score == len("tiny") / 100
|
||||
|
||||
|
||||
def test_admission_size_large_penalized():
|
||||
"""Very large content (> 50k chars) gets penalized size_score."""
|
||||
ctrl = AdmissionController()
|
||||
score = ctrl.score("x" * 80000, "file_context", has_duplicate=False)
|
||||
assert score.size_score < 1.0
|
||||
assert score.size_score >= 0.3
|
||||
|
||||
|
||||
def test_admission_size_normal():
|
||||
"""Normal-sized content (100-50k chars) gets size_score=1.0."""
|
||||
ctrl = AdmissionController()
|
||||
score = ctrl.score("x" * 500, "file_context", has_duplicate=False)
|
||||
assert score.size_score == 1.0
|
||||
|
||||
|
||||
# ── Threshold decisions ──────────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_admission_threshold_admit():
|
||||
"""High-value content is admitted."""
|
||||
ctrl = AdmissionController()
|
||||
# design_decision (0.9 type) + unique (1.0 novelty) + code + entities + length
|
||||
content = "```python\ndef important(): pass\n```\n" + "x" * 300
|
||||
admitted, score = ctrl.should_admit(
|
||||
content,
|
||||
"design_decision",
|
||||
has_duplicate=False,
|
||||
key_entities=["foo.py", "bar.py", "baz.py"],
|
||||
)
|
||||
assert admitted is True
|
||||
assert score.total >= DEFAULT_THRESHOLD
|
||||
|
||||
|
||||
def test_admission_threshold_reject():
|
||||
"""Low-value duplicate conversation_phase is rejected."""
|
||||
ctrl = AdmissionController()
|
||||
# conversation_phase (0.3 type) + duplicate (0.2 novelty) + tiny content
|
||||
admitted, score = ctrl.should_admit("hi", "conversation_phase", has_duplicate=True)
|
||||
assert admitted is False
|
||||
assert score.total < DEFAULT_THRESHOLD
|
||||
|
||||
|
||||
# ── Stats tracking ───────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_admission_stats_tracking():
|
||||
"""Stats correctly count admitted and rejected objects."""
|
||||
ctrl = AdmissionController()
|
||||
assert ctrl.stats == {"admitted": 0, "rejected": 0}
|
||||
|
||||
# Admit a high-value object
|
||||
ctrl.should_admit(
|
||||
"```python\ndef foo(): pass\n```\n" + "x" * 300,
|
||||
"design_decision",
|
||||
has_duplicate=False,
|
||||
key_entities=["a.py", "b.py", "c.py"],
|
||||
)
|
||||
assert ctrl.stats["admitted"] == 1
|
||||
|
||||
# Reject a low-value object
|
||||
ctrl.should_admit("hi", "conversation_phase", has_duplicate=True)
|
||||
assert ctrl.stats["rejected"] == 1
|
||||
assert ctrl.stats["admitted"] == 1
|
||||
|
||||
|
||||
# ── Edge cases ───────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_admission_empty_content():
|
||||
"""Empty content gets size_score=0 and low utility."""
|
||||
ctrl = AdmissionController()
|
||||
score = ctrl.score("", "file_context", has_duplicate=False)
|
||||
assert score.size_score == 0.0
|
||||
assert score.utility_score == 0.0
|
||||
182
tests/test_entropy.py
Normal file
182
tests/test_entropy.py
Normal file
|
|
@ -0,0 +1,182 @@
|
|||
"""Tests for the entropy-gated faulting module.
|
||||
|
||||
Tests the EntropyDetector's ability to detect hedging, uncertainty,
|
||||
evicted entity references, and the should_fault decision logic.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from mnemosyne.entropy import (
|
||||
DEFAULT_FAULT_THRESHOLD,
|
||||
EntropyDetector,
|
||||
EntropySignal,
|
||||
)
|
||||
|
||||
|
||||
# ── Hedging detection ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_entropy_hedging_detected():
|
||||
"""Hedging language like 'I think' and 'probably' is detected."""
|
||||
detector = EntropyDetector()
|
||||
signal = detector.analyze_response(
|
||||
"I think the function is probably in utils.py",
|
||||
evicted_entities=[],
|
||||
)
|
||||
assert signal.has_hedging is True
|
||||
assert signal.score > 0
|
||||
|
||||
|
||||
def test_entropy_multiple_patterns():
|
||||
"""Multiple hedging patterns increase the score."""
|
||||
detector = EntropyDetector()
|
||||
signal_one = detector.analyze_response(
|
||||
"I think it might work.",
|
||||
evicted_entities=[],
|
||||
)
|
||||
signal_many = detector.analyze_response(
|
||||
"I think it probably might be in the file, if I recall correctly. I believe so.",
|
||||
evicted_entities=[],
|
||||
)
|
||||
assert signal_many.score >= signal_one.score
|
||||
|
||||
|
||||
# ── Uncertainty detection ────────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_entropy_uncertainty_detected():
|
||||
"""Explicit uncertainty like 'I'm not sure' is detected."""
|
||||
detector = EntropyDetector()
|
||||
signal = detector.analyze_response(
|
||||
"I'm not sure about the exact implementation. It's unclear to me.",
|
||||
evicted_entities=[],
|
||||
)
|
||||
assert signal.has_uncertainty is True
|
||||
assert signal.score > 0
|
||||
|
||||
|
||||
# ── No signals ───────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_entropy_no_signals():
|
||||
"""Clean, confident response has no entropy signals."""
|
||||
detector = EntropyDetector()
|
||||
signal = detector.analyze_response(
|
||||
"The function `calculate_total` is defined in src/utils.py at line 42. "
|
||||
"It takes two arguments: price and quantity.",
|
||||
evicted_entities=[],
|
||||
)
|
||||
assert signal.has_hedging is False
|
||||
assert signal.has_uncertainty is False
|
||||
assert signal.references_evicted is False
|
||||
assert signal.has_hallucination_risk is False
|
||||
assert signal.score == 0.0
|
||||
|
||||
|
||||
# ── Evicted entity references ───────────────────────────────────────────
|
||||
|
||||
|
||||
def test_entropy_evicted_references():
|
||||
"""References to evicted entities are detected."""
|
||||
detector = EntropyDetector()
|
||||
signal = detector.analyze_response(
|
||||
"The config is in settings.py and the handler is in views.py",
|
||||
evicted_entities=["settings.py", "views.py", "models.py"],
|
||||
)
|
||||
assert signal.references_evicted is True
|
||||
assert "settings.py" in signal.referenced_entities
|
||||
assert "views.py" in signal.referenced_entities
|
||||
assert "models.py" not in signal.referenced_entities
|
||||
|
||||
|
||||
def test_entropy_hallucination_risk():
|
||||
"""Hedging + evicted references = hallucination risk."""
|
||||
detector = EntropyDetector()
|
||||
signal = detector.analyze_response(
|
||||
"I think the config is probably in settings.py somewhere",
|
||||
evicted_entities=["settings.py"],
|
||||
)
|
||||
assert signal.has_hedging is True
|
||||
assert signal.references_evicted is True
|
||||
assert signal.has_hallucination_risk is True
|
||||
# Hallucination risk adds a bonus to the score
|
||||
assert signal.score > 0.3
|
||||
|
||||
|
||||
# ── Composite score ──────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_entropy_composite_score():
|
||||
"""Score combines hedging, uncertainty, and evicted references."""
|
||||
detector = EntropyDetector()
|
||||
# All signals present
|
||||
signal = detector.analyze_response(
|
||||
"I'm not sure, but I think the code is probably in utils.py",
|
||||
evicted_entities=["utils.py"],
|
||||
)
|
||||
assert signal.has_hedging is True
|
||||
assert signal.has_uncertainty is True
|
||||
assert signal.references_evicted is True
|
||||
assert signal.score > 0.5 # Should be high with all signals
|
||||
|
||||
|
||||
# ── should_fault decisions ───────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_entropy_should_fault_above_threshold():
|
||||
"""Faulting triggers when score >= threshold AND references evicted."""
|
||||
detector = EntropyDetector(fault_threshold=0.4)
|
||||
signal = detector.analyze_response(
|
||||
"I think the code is probably in utils.py, if I recall correctly. I believe so.",
|
||||
evicted_entities=["utils.py"],
|
||||
)
|
||||
# Score should be above 0.4 with hedging + evicted ref + hallucination bonus
|
||||
entities = detector.should_fault(signal)
|
||||
assert len(entities) > 0
|
||||
assert "utils.py" in entities
|
||||
|
||||
|
||||
def test_entropy_should_fault_below_threshold():
|
||||
"""No faulting when score is below threshold."""
|
||||
detector = EntropyDetector(fault_threshold=0.99) # Very high threshold
|
||||
signal = detector.analyze_response(
|
||||
"I think it might be there.",
|
||||
evicted_entities=[],
|
||||
)
|
||||
entities = detector.should_fault(signal)
|
||||
assert entities == []
|
||||
|
||||
|
||||
def test_entropy_should_fault_no_evicted_refs():
|
||||
"""No faulting when there are no evicted entity references, even with high hedging."""
|
||||
detector = EntropyDetector(fault_threshold=0.1) # Very low threshold
|
||||
signal = detector.analyze_response(
|
||||
"I think it probably might be somewhere, if I recall. I believe so.",
|
||||
evicted_entities=[],
|
||||
)
|
||||
# Even though hedging score is high, no evicted refs → no fault
|
||||
entities = detector.should_fault(signal)
|
||||
assert entities == []
|
||||
|
||||
|
||||
# ── Stats tracking ───────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_entropy_stats_tracking():
|
||||
"""Stats correctly count analyzed responses and triggered faults."""
|
||||
detector = EntropyDetector(fault_threshold=0.3)
|
||||
assert detector.stats == {"analyzed": 0, "faults_triggered": 0}
|
||||
|
||||
# Analyze a clean response
|
||||
detector.analyze_response("The function is in utils.py.", evicted_entities=[])
|
||||
assert detector.stats["analyzed"] == 1
|
||||
assert detector.stats["faults_triggered"] == 0
|
||||
|
||||
# Analyze and fault a hedging response with evicted refs
|
||||
signal = detector.analyze_response(
|
||||
"I think the code is probably in utils.py, if I recall correctly. I believe so.",
|
||||
evicted_entities=["utils.py"],
|
||||
)
|
||||
assert detector.stats["analyzed"] == 2
|
||||
detector.should_fault(signal)
|
||||
assert detector.stats["faults_triggered"] == 1
|
||||
130
tests/test_goal_aware.py
Normal file
130
tests/test_goal_aware.py
Normal file
|
|
@ -0,0 +1,130 @@
|
|||
"""Tests for goal-aware retrieval integration in the gateway.
|
||||
|
||||
Tests the Session attributes for goal tracking, cosine similarity-based
|
||||
topic shift detection, and graceful degradation when helper_llm is absent.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
from tempfile import TemporaryDirectory
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from mnemosyne.gateway import Session
|
||||
from mnemosyne.helper_llm import GoalClassification
|
||||
from mnemosyne.object_store import DummyEmbedder, _cosine_similarity
|
||||
|
||||
|
||||
# ── Fixtures ─────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def tmp_log_dir():
|
||||
with TemporaryDirectory() as d:
|
||||
yield Path(d)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def session(tmp_log_dir):
|
||||
return Session("goal01", tmp_log_dir)
|
||||
|
||||
|
||||
# ── Session attributes ───────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestGoalSessionAttributes:
|
||||
"""Verify Session.__init__ creates goal-tracking attributes."""
|
||||
|
||||
def test_goal_session_attributes_exist(self, session):
|
||||
"""Session has _last_user_embedding and _current_goal attributes."""
|
||||
assert hasattr(session, "_last_user_embedding")
|
||||
assert session._last_user_embedding is None
|
||||
assert hasattr(session, "_current_goal")
|
||||
assert session._current_goal is None
|
||||
assert hasattr(session, "entropy_detector")
|
||||
assert session.entropy_detector is not None
|
||||
|
||||
|
||||
# ── Cosine similarity topic shift ────────────────────────────────────────
|
||||
|
||||
|
||||
class TestCosineTopicShift:
|
||||
"""Test topic shift detection via cosine similarity."""
|
||||
|
||||
def test_goal_cosine_similarity_topic_shift(self):
|
||||
"""Different topics produce low cosine similarity (< 0.5)."""
|
||||
embedder = DummyEmbedder()
|
||||
# Two very different texts should produce different embeddings
|
||||
emb_a = embedder.embed("How do I configure the database connection pool?")
|
||||
emb_b = embedder.embed("What color should the login button be?")
|
||||
sim = _cosine_similarity(emb_a, emb_b)
|
||||
# DummyEmbedder uses hash-based vectors, so different texts
|
||||
# produce essentially random vectors with low expected similarity
|
||||
# For 384-dim random unit vectors, expected |cos| ≈ 0.05
|
||||
assert sim < 0.5
|
||||
|
||||
def test_goal_cosine_similarity_same_topic(self):
|
||||
"""Identical text produces cosine similarity of 1.0."""
|
||||
embedder = DummyEmbedder()
|
||||
text = "How do I configure the database connection pool?"
|
||||
emb_a = embedder.embed(text)
|
||||
emb_b = embedder.embed(text)
|
||||
sim = _cosine_similarity(emb_a, emb_b)
|
||||
assert sim > 0.99 # Same text → same hash → same embedding
|
||||
|
||||
def test_goal_first_message_always_classifies(self, session):
|
||||
"""First message (no prior embedding) should always trigger classification."""
|
||||
# _last_user_embedding is None → goal_changed should be True
|
||||
assert session._last_user_embedding is None
|
||||
# This is the logic from gateway._preprocess step 1c:
|
||||
# if session._last_user_embedding is not None: check sim
|
||||
# else: goal_changed = True
|
||||
goal_changed = session._last_user_embedding is None
|
||||
assert goal_changed is True
|
||||
|
||||
def test_goal_classification_fallback_no_helper(self, session):
|
||||
"""Goal detection is skipped gracefully when helper_llm is None.
|
||||
|
||||
The gateway code wraps goal detection in try/except and checks
|
||||
`if goal_changed and helper_llm is not None`. When helper_llm
|
||||
is None, no classification occurs and _current_goal stays None.
|
||||
"""
|
||||
# Simulate the gateway logic: helper_llm is None
|
||||
helper_llm = None
|
||||
goal_changed = True # First message
|
||||
|
||||
# This mirrors the gateway code path
|
||||
if goal_changed and helper_llm is not None:
|
||||
# Would call helper_llm.classify_goal(...)
|
||||
session._current_goal = GoalClassification(goal="test")
|
||||
|
||||
# Goal should remain None since helper_llm is None
|
||||
assert session._current_goal is None
|
||||
|
||||
def test_goal_embedding_updated_after_message(self, session):
|
||||
"""_last_user_embedding is updated after processing a message."""
|
||||
embedder = DummyEmbedder()
|
||||
user_text = "Tell me about the authentication system"
|
||||
current_embedding = embedder.embed(user_text)
|
||||
|
||||
# Simulate the gateway update
|
||||
session._last_user_embedding = current_embedding
|
||||
|
||||
assert session._last_user_embedding is not None
|
||||
assert len(session._last_user_embedding) == 384
|
||||
|
||||
def test_goal_classification_stored_on_session(self, session):
|
||||
"""GoalClassification is stored on session._current_goal."""
|
||||
goal = GoalClassification(
|
||||
goal="Implement authentication",
|
||||
relevant_types=["file_context", "design_decision"],
|
||||
relevant_tags=["auth", "security"],
|
||||
predicted_needs=["auth.py", "middleware.py"],
|
||||
)
|
||||
session._current_goal = goal
|
||||
|
||||
assert session._current_goal.goal == "Implement authentication"
|
||||
assert "file_context" in session._current_goal.relevant_types
|
||||
assert "auth" in session._current_goal.relevant_tags
|
||||
893
tests/test_hierarchy.py
Normal file
893
tests/test_hierarchy.py
Normal file
|
|
@ -0,0 +1,893 @@
|
|||
"""Tests for the hierarchical segmentation system (Strategy B).
|
||||
|
||||
Covers:
|
||||
- Episode creation and clustering
|
||||
- SemanticTheme creation
|
||||
- Theme creation
|
||||
- SessionHierarchy: add_object, rebuild, retrieve, maintenance
|
||||
- Edge cases: empty, single object, identical embeddings, no embeddings
|
||||
- Vector math utilities
|
||||
- Incremental vs full rebuild consistency
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
|
||||
from mnemosyne.hierarchy import (
|
||||
EPISODE_RETRIEVAL_THRESHOLD,
|
||||
EPISODE_SIMILARITY_THRESHOLD,
|
||||
Episode,
|
||||
RECLUSTER_INTERVAL,
|
||||
SEMANTIC_SIMILARITY_THRESHOLD,
|
||||
SemanticTheme,
|
||||
SessionHierarchy,
|
||||
THEME_RETRIEVAL_THRESHOLD,
|
||||
THEME_SIMILARITY_THRESHOLD,
|
||||
Theme,
|
||||
_agglomerative_cluster,
|
||||
_centroid,
|
||||
_cosine_similarity,
|
||||
_cosine_similarity_matrix,
|
||||
_generate_episode_summary,
|
||||
_generate_theme_label,
|
||||
_make_episode,
|
||||
_make_semantic_theme,
|
||||
_make_theme,
|
||||
)
|
||||
from mnemosyne.object_store import DummyEmbedder, StoredObject, _estimate_tokens
|
||||
|
||||
|
||||
# ── Helpers ──────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _make_stored_object(
|
||||
content: str = "test content",
|
||||
*,
|
||||
session_id: str = "sess-1",
|
||||
object_type: str = "file_context",
|
||||
source_tool: str | None = "Read",
|
||||
source_key: str | None = None,
|
||||
stub: str | None = None,
|
||||
embedding: list[float] | None = None,
|
||||
object_id: str | None = None,
|
||||
turn: int | None = None,
|
||||
) -> StoredObject:
|
||||
"""Create a StoredObject with sensible defaults for testing."""
|
||||
return StoredObject(
|
||||
id=object_id or f"obj-{hash(content) % 100000:05d}",
|
||||
session_id=session_id,
|
||||
object_type=object_type,
|
||||
source_tool=source_tool,
|
||||
source_key=source_key,
|
||||
content_full=content,
|
||||
summary_detailed=None,
|
||||
summary_compact=None,
|
||||
stub=stub or f"{object_type}: {content[:40]}",
|
||||
tokens_l0=_estimate_tokens(content),
|
||||
tokens_l3=_estimate_tokens(stub or f"{object_type}: test"),
|
||||
embedding=embedding or [],
|
||||
created_at="2025-01-01T00:00:00+00:00",
|
||||
last_accessed="2025-01-01T00:00:00+00:00",
|
||||
source_turn_start=turn,
|
||||
source_turn_end=turn,
|
||||
)
|
||||
|
||||
|
||||
def _make_unit_vector(dim: int = 384, seed: int = 42) -> list[float]:
|
||||
"""Create a deterministic unit vector."""
|
||||
rng = np.random.default_rng(seed)
|
||||
vec = rng.standard_normal(dim)
|
||||
vec = vec / np.linalg.norm(vec)
|
||||
return vec.tolist()
|
||||
|
||||
|
||||
def _make_similar_vector(base: list[float], noise: float = 0.05, seed: int = 99) -> list[float]:
|
||||
"""Create a vector similar to `base` by adding small noise."""
|
||||
rng = np.random.default_rng(seed)
|
||||
arr = np.asarray(base, dtype=np.float64)
|
||||
perturbation = rng.standard_normal(len(base)) * noise
|
||||
result = arr + perturbation
|
||||
result = result / np.linalg.norm(result)
|
||||
return result.tolist()
|
||||
|
||||
|
||||
def _make_orthogonal_vector(base: list[float], seed: int = 77) -> list[float]:
|
||||
"""Create a vector roughly orthogonal to `base`."""
|
||||
rng = np.random.default_rng(seed)
|
||||
random_vec = rng.standard_normal(len(base))
|
||||
arr = np.asarray(base, dtype=np.float64)
|
||||
# Gram-Schmidt: subtract projection onto base
|
||||
proj = np.dot(random_vec, arr) / np.dot(arr, arr) * arr
|
||||
ortho = random_vec - proj
|
||||
norm = np.linalg.norm(ortho)
|
||||
if norm > 0:
|
||||
ortho = ortho / norm
|
||||
return ortho.tolist()
|
||||
|
||||
|
||||
# ── Vector math tests ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestCosineSimlarity:
|
||||
def test_identical_vectors(self):
|
||||
v = _make_unit_vector(seed=1)
|
||||
assert _cosine_similarity(v, v) == pytest.approx(1.0, abs=1e-6)
|
||||
|
||||
def test_orthogonal_vectors(self):
|
||||
v1 = _make_unit_vector(seed=1)
|
||||
v2 = _make_orthogonal_vector(v1, seed=2)
|
||||
assert abs(_cosine_similarity(v1, v2)) < 0.05
|
||||
|
||||
def test_opposite_vectors(self):
|
||||
v = _make_unit_vector(seed=1)
|
||||
neg_v = [-x for x in v]
|
||||
assert _cosine_similarity(v, neg_v) == pytest.approx(-1.0, abs=1e-6)
|
||||
|
||||
def test_empty_vectors(self):
|
||||
assert _cosine_similarity([], []) == 0.0
|
||||
assert _cosine_similarity([1.0, 0.0], []) == 0.0
|
||||
assert _cosine_similarity([], [1.0, 0.0]) == 0.0
|
||||
|
||||
def test_zero_vector(self):
|
||||
v = _make_unit_vector(seed=1)
|
||||
zero = [0.0] * len(v)
|
||||
assert _cosine_similarity(v, zero) == 0.0
|
||||
|
||||
def test_similar_vectors_high_similarity(self):
|
||||
v1 = _make_unit_vector(seed=1)
|
||||
v2 = _make_similar_vector(v1, noise=0.01, seed=2)
|
||||
sim = _cosine_similarity(v1, v2)
|
||||
assert sim > 0.9
|
||||
|
||||
|
||||
class TestCosineSimlarityMatrix:
|
||||
def test_single_vector(self):
|
||||
v = _make_unit_vector(seed=1)
|
||||
mat = _cosine_similarity_matrix(np.array([v]))
|
||||
assert mat.shape == (1, 1)
|
||||
assert mat[0, 0] == pytest.approx(1.0, abs=1e-6)
|
||||
|
||||
def test_identity_diagonal(self):
|
||||
vecs = [_make_unit_vector(seed=i) for i in range(5)]
|
||||
mat = _cosine_similarity_matrix(np.array(vecs))
|
||||
for i in range(5):
|
||||
assert mat[i, i] == pytest.approx(1.0, abs=1e-6)
|
||||
|
||||
def test_symmetry(self):
|
||||
vecs = [_make_unit_vector(seed=i) for i in range(3)]
|
||||
mat = _cosine_similarity_matrix(np.array(vecs))
|
||||
for i in range(3):
|
||||
for j in range(3):
|
||||
assert mat[i, j] == pytest.approx(mat[j, i], abs=1e-10)
|
||||
|
||||
def test_empty(self):
|
||||
mat = _cosine_similarity_matrix(np.empty((0, 384)))
|
||||
assert mat.shape == (0, 0)
|
||||
|
||||
|
||||
class TestCentroid:
|
||||
def test_single_vector(self):
|
||||
v = _make_unit_vector(seed=1)
|
||||
c = _centroid([v])
|
||||
# Centroid of a single unit vector is itself
|
||||
assert _cosine_similarity(v, c) == pytest.approx(1.0, abs=1e-6)
|
||||
|
||||
def test_identical_vectors(self):
|
||||
v = _make_unit_vector(seed=1)
|
||||
c = _centroid([v, v, v])
|
||||
assert _cosine_similarity(v, c) == pytest.approx(1.0, abs=1e-6)
|
||||
|
||||
def test_empty(self):
|
||||
assert _centroid([]) == []
|
||||
|
||||
def test_centroid_is_normalized(self):
|
||||
vecs = [_make_unit_vector(seed=i) for i in range(5)]
|
||||
c = _centroid(vecs)
|
||||
norm = float(np.linalg.norm(c))
|
||||
assert norm == pytest.approx(1.0, abs=1e-6)
|
||||
|
||||
|
||||
# ── Agglomerative clustering tests ──────────────────────────────────────
|
||||
|
||||
|
||||
class TestAgglomerativeClustering:
|
||||
def test_empty(self):
|
||||
assert _agglomerative_cluster([], 0.5) == []
|
||||
|
||||
def test_single_item(self):
|
||||
v = _make_unit_vector(seed=1)
|
||||
clusters = _agglomerative_cluster([v], 0.5)
|
||||
assert len(clusters) == 1
|
||||
assert clusters[0] == [0]
|
||||
|
||||
def test_identical_items_merge(self):
|
||||
v = _make_unit_vector(seed=1)
|
||||
clusters = _agglomerative_cluster([v, v, v], 0.5)
|
||||
assert len(clusters) == 1
|
||||
assert sorted(clusters[0]) == [0, 1, 2]
|
||||
|
||||
def test_dissimilar_items_separate(self):
|
||||
v1 = _make_unit_vector(seed=1)
|
||||
v2 = _make_orthogonal_vector(v1, seed=2)
|
||||
clusters = _agglomerative_cluster([v1, v2], 0.5)
|
||||
assert len(clusters) == 2
|
||||
|
||||
def test_similar_items_merge(self):
|
||||
v1 = _make_unit_vector(seed=1)
|
||||
v2 = _make_similar_vector(v1, noise=0.02, seed=2)
|
||||
clusters = _agglomerative_cluster([v1, v2], 0.5)
|
||||
assert len(clusters) == 1
|
||||
|
||||
def test_mixed_clusters(self):
|
||||
"""Two groups of similar vectors should form two clusters."""
|
||||
v1 = _make_unit_vector(seed=1)
|
||||
v1b = _make_similar_vector(v1, noise=0.02, seed=10)
|
||||
v2 = _make_orthogonal_vector(v1, seed=2)
|
||||
v2b = _make_similar_vector(v2, noise=0.02, seed=20)
|
||||
clusters = _agglomerative_cluster([v1, v1b, v2, v2b], 0.5)
|
||||
assert len(clusters) == 2
|
||||
|
||||
def test_threshold_boundary(self):
|
||||
"""Items with high similarity should merge at a reasonable threshold."""
|
||||
v = _make_unit_vector(seed=1)
|
||||
# Identical vectors have similarity ~1.0 (floating point)
|
||||
clusters = _agglomerative_cluster([v, v], 0.99)
|
||||
assert len(clusters) == 1
|
||||
|
||||
|
||||
# ── Episode dataclass tests ─────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestEpisode:
|
||||
def test_creation(self):
|
||||
ep = Episode(id="ep-1")
|
||||
assert ep.id == "ep-1"
|
||||
assert ep.objects == []
|
||||
assert ep.embedding == []
|
||||
assert ep.summary == ""
|
||||
assert ep.turn_range == (0, 0)
|
||||
assert ep.object_types == set()
|
||||
|
||||
def test_hash_and_equality(self):
|
||||
ep1 = Episode(id="ep-1")
|
||||
ep2 = Episode(id="ep-1")
|
||||
ep3 = Episode(id="ep-2")
|
||||
assert ep1 == ep2
|
||||
assert ep1 != ep3
|
||||
assert hash(ep1) == hash(ep2)
|
||||
|
||||
def test_make_episode(self):
|
||||
embedder = DummyEmbedder()
|
||||
obj1 = _make_stored_object(
|
||||
"content A", object_id="a", turn=5, embedding=embedder.embed("content A")
|
||||
)
|
||||
obj2 = _make_stored_object(
|
||||
"content B", object_id="b", turn=8, embedding=embedder.embed("content B")
|
||||
)
|
||||
ep = _make_episode([obj1, obj2])
|
||||
assert len(ep.objects) == 2
|
||||
assert ep.turn_range == (5, 8)
|
||||
assert "file_context" in ep.object_types
|
||||
assert ep.embedding # Should have a centroid
|
||||
assert ep.summary # Should have auto-generated summary
|
||||
|
||||
|
||||
class TestSemanticTheme:
|
||||
def test_creation(self):
|
||||
st = SemanticTheme(id="st-1")
|
||||
assert st.id == "st-1"
|
||||
assert st.episodes == []
|
||||
assert st.embedding == []
|
||||
assert st.label == ""
|
||||
|
||||
def test_hash_and_equality(self):
|
||||
st1 = SemanticTheme(id="st-1")
|
||||
st2 = SemanticTheme(id="st-1")
|
||||
st3 = SemanticTheme(id="st-2")
|
||||
assert st1 == st2
|
||||
assert st1 != st3
|
||||
|
||||
|
||||
class TestTheme:
|
||||
def test_creation(self):
|
||||
t = Theme(id="t-1")
|
||||
assert t.id == "t-1"
|
||||
assert t.semantic_themes == []
|
||||
|
||||
def test_hash_and_equality(self):
|
||||
t1 = Theme(id="t-1")
|
||||
t2 = Theme(id="t-1")
|
||||
t3 = Theme(id="t-2")
|
||||
assert t1 == t2
|
||||
assert t1 != t3
|
||||
|
||||
|
||||
# ── Summary generation tests ────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestSummaryGeneration:
|
||||
def test_episode_summary_from_stubs(self):
|
||||
obj = _make_stored_object("hello world", stub="file_context: hello")
|
||||
summary = _generate_episode_summary([obj])
|
||||
assert "file_context: hello" in summary
|
||||
|
||||
def test_episode_summary_empty(self):
|
||||
summary = _generate_episode_summary([])
|
||||
assert "empty" in summary.lower()
|
||||
|
||||
def test_episode_summary_many_objects(self):
|
||||
objs = [
|
||||
_make_stored_object(f"content {i}", object_id=f"o{i}", stub=f"stub {i}")
|
||||
for i in range(10)
|
||||
]
|
||||
summary = _generate_episode_summary(objs)
|
||||
assert "+5 more" in summary
|
||||
|
||||
def test_theme_label_from_episodes(self):
|
||||
ep = Episode(id="ep-1", object_types={"file_context", "tool_result"})
|
||||
label = _generate_theme_label([ep])
|
||||
assert "file_context" in label or "tool_result" in label
|
||||
assert "1 episodes" in label
|
||||
|
||||
def test_theme_label_empty(self):
|
||||
label = _generate_theme_label([])
|
||||
assert "empty" in label.lower()
|
||||
|
||||
|
||||
# ── SessionHierarchy tests ──────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestSessionHierarchyInit:
|
||||
def test_init(self):
|
||||
embedder = DummyEmbedder()
|
||||
h = SessionHierarchy(embedder)
|
||||
assert h.enabled is True
|
||||
assert h.object_count == 0
|
||||
assert h.episode_count == 0
|
||||
assert h.theme_count == 0
|
||||
|
||||
def test_disable(self):
|
||||
embedder = DummyEmbedder()
|
||||
h = SessionHierarchy(embedder)
|
||||
h.enabled = False
|
||||
assert h.enabled is False
|
||||
|
||||
def test_summary_empty(self):
|
||||
embedder = DummyEmbedder()
|
||||
h = SessionHierarchy(embedder)
|
||||
s = h.summary()
|
||||
assert s["objects"] == 0
|
||||
assert s["episodes"] == 0
|
||||
assert s["themes"] == 0
|
||||
|
||||
|
||||
class TestSessionHierarchyAddObject:
|
||||
def test_add_single_object(self):
|
||||
embedder = DummyEmbedder()
|
||||
h = SessionHierarchy(embedder)
|
||||
obj = _make_stored_object(
|
||||
"test content", object_id="o1", embedding=embedder.embed("test content")
|
||||
)
|
||||
ep = h.add_object(obj)
|
||||
assert ep is not None
|
||||
assert h.object_count == 1
|
||||
assert h.episode_count == 1
|
||||
|
||||
def test_add_duplicate_ignored(self):
|
||||
embedder = DummyEmbedder()
|
||||
h = SessionHierarchy(embedder)
|
||||
obj = _make_stored_object("test", object_id="o1", embedding=embedder.embed("test"))
|
||||
h.add_object(obj)
|
||||
result = h.add_object(obj)
|
||||
assert result is None
|
||||
assert h.object_count == 1
|
||||
|
||||
def test_add_similar_objects_same_episode(self):
|
||||
"""Objects with very similar embeddings should join the same episode."""
|
||||
embedder = DummyEmbedder()
|
||||
h = SessionHierarchy(embedder)
|
||||
|
||||
# Use the same text to get identical embeddings
|
||||
emb = embedder.embed("shared topic content")
|
||||
obj1 = _make_stored_object("shared topic content A", object_id="o1", embedding=emb)
|
||||
obj2 = _make_stored_object(
|
||||
"shared topic content B", object_id="o2", embedding=emb
|
||||
) # Same embedding
|
||||
|
||||
ep1 = h.add_object(obj1)
|
||||
ep2 = h.add_object(obj2)
|
||||
assert ep1 is not None
|
||||
assert ep2 is not None
|
||||
assert ep1.id == ep2.id # Same episode
|
||||
assert h.episode_count == 1
|
||||
|
||||
def test_add_dissimilar_objects_different_episodes(self):
|
||||
"""Objects with very different embeddings should go to different episodes."""
|
||||
embedder = DummyEmbedder()
|
||||
h = SessionHierarchy(embedder)
|
||||
|
||||
# DummyEmbedder produces different vectors for different text
|
||||
obj1 = _make_stored_object(
|
||||
"python programming language",
|
||||
object_id="o1",
|
||||
embedding=embedder.embed("python programming language"),
|
||||
)
|
||||
obj2 = _make_stored_object(
|
||||
"cooking recipes for dinner",
|
||||
object_id="o2",
|
||||
embedding=embedder.embed("cooking recipes for dinner"),
|
||||
)
|
||||
|
||||
h.add_object(obj1)
|
||||
h.add_object(obj2)
|
||||
# DummyEmbedder is hash-based, so different texts → different vectors
|
||||
# They should be in different episodes (similarity < 0.6)
|
||||
assert h.episode_count == 2
|
||||
|
||||
def test_add_object_without_embedding(self):
|
||||
embedder = DummyEmbedder()
|
||||
h = SessionHierarchy(embedder)
|
||||
obj = _make_stored_object("no embedding", object_id="o1", embedding=[])
|
||||
ep = h.add_object(obj)
|
||||
assert ep is not None
|
||||
assert h.episode_count == 1
|
||||
|
||||
def test_add_object_disabled(self):
|
||||
embedder = DummyEmbedder()
|
||||
h = SessionHierarchy(embedder)
|
||||
h.enabled = False
|
||||
obj = _make_stored_object("test", object_id="o1", embedding=embedder.embed("test"))
|
||||
result = h.add_object(obj)
|
||||
assert result is None
|
||||
assert h.object_count == 0
|
||||
|
||||
def test_add_updates_episode_metadata(self):
|
||||
embedder = DummyEmbedder()
|
||||
h = SessionHierarchy(embedder)
|
||||
emb = embedder.embed("shared content")
|
||||
obj1 = _make_stored_object(
|
||||
"shared content A", object_id="o1", object_type="file_context", embedding=emb, turn=5
|
||||
)
|
||||
obj2 = _make_stored_object(
|
||||
"shared content B", object_id="o2", object_type="tool_result", embedding=emb, turn=10
|
||||
)
|
||||
h.add_object(obj1)
|
||||
ep = h.add_object(obj2)
|
||||
assert ep is not None
|
||||
assert "file_context" in ep.object_types
|
||||
assert "tool_result" in ep.object_types
|
||||
assert ep.turn_range == (5, 10)
|
||||
|
||||
|
||||
class TestSessionHierarchyRebuild:
|
||||
def test_rebuild_empty(self):
|
||||
embedder = DummyEmbedder()
|
||||
h = SessionHierarchy(embedder)
|
||||
h.rebuild([])
|
||||
assert h.object_count == 0
|
||||
assert h.episode_count == 0
|
||||
|
||||
def test_rebuild_single_object(self):
|
||||
embedder = DummyEmbedder()
|
||||
h = SessionHierarchy(embedder)
|
||||
obj = _make_stored_object("test", object_id="o1", embedding=embedder.embed("test"))
|
||||
h.rebuild([obj])
|
||||
assert h.object_count == 1
|
||||
assert h.episode_count == 1
|
||||
assert len(h.get_themes()) >= 1
|
||||
assert len(h.get_top_themes()) >= 1
|
||||
|
||||
def test_rebuild_clusters_similar_objects(self):
|
||||
embedder = DummyEmbedder()
|
||||
h = SessionHierarchy(embedder)
|
||||
emb = embedder.embed("shared topic")
|
||||
objects = [
|
||||
_make_stored_object(f"shared topic {i}", object_id=f"o{i}", embedding=emb)
|
||||
for i in range(5)
|
||||
]
|
||||
h.rebuild(objects)
|
||||
assert h.object_count == 5
|
||||
# All identical embeddings → 1 episode
|
||||
assert h.episode_count == 1
|
||||
|
||||
def test_rebuild_separates_dissimilar_objects(self):
|
||||
embedder = DummyEmbedder()
|
||||
h = SessionHierarchy(embedder)
|
||||
objects = [
|
||||
_make_stored_object(
|
||||
f"unique topic number {i} with distinct content",
|
||||
object_id=f"o{i}",
|
||||
embedding=embedder.embed(f"unique topic number {i} with distinct content"),
|
||||
)
|
||||
for i in range(10)
|
||||
]
|
||||
h.rebuild(objects)
|
||||
assert h.object_count == 10
|
||||
# DummyEmbedder with different texts → different vectors → multiple episodes
|
||||
assert h.episode_count > 1
|
||||
|
||||
def test_rebuild_handles_mixed_embeddings(self):
|
||||
"""Objects with and without embeddings should both be included."""
|
||||
embedder = DummyEmbedder()
|
||||
h = SessionHierarchy(embedder)
|
||||
obj_with = _make_stored_object(
|
||||
"with embedding", object_id="o1", embedding=embedder.embed("with embedding")
|
||||
)
|
||||
obj_without = _make_stored_object("no embedding", object_id="o2", embedding=[])
|
||||
h.rebuild([obj_with, obj_without])
|
||||
assert h.object_count == 2
|
||||
assert h.episode_count == 2 # Each in its own episode
|
||||
|
||||
def test_rebuild_clears_previous_state(self):
|
||||
embedder = DummyEmbedder()
|
||||
h = SessionHierarchy(embedder)
|
||||
obj1 = _make_stored_object("first", object_id="o1", embedding=embedder.embed("first"))
|
||||
h.rebuild([obj1])
|
||||
assert h.object_count == 1
|
||||
|
||||
obj2 = _make_stored_object("second", object_id="o2", embedding=embedder.embed("second"))
|
||||
h.rebuild([obj2])
|
||||
assert h.object_count == 1 # Only the new object
|
||||
assert h.episode_count == 1
|
||||
|
||||
def test_rebuild_creates_all_levels(self):
|
||||
embedder = DummyEmbedder()
|
||||
h = SessionHierarchy(embedder)
|
||||
objects = [
|
||||
_make_stored_object(
|
||||
f"content {i}", object_id=f"o{i}", embedding=embedder.embed(f"content {i}")
|
||||
)
|
||||
for i in range(20)
|
||||
]
|
||||
h.rebuild(objects)
|
||||
assert h.object_count == 20
|
||||
assert h.episode_count > 0
|
||||
assert len(h.get_themes()) > 0
|
||||
assert len(h.get_top_themes()) > 0
|
||||
|
||||
|
||||
class TestSessionHierarchyRetrieve:
|
||||
def test_retrieve_empty_hierarchy(self):
|
||||
embedder = DummyEmbedder()
|
||||
h = SessionHierarchy(embedder)
|
||||
query = embedder.embed("test query")
|
||||
results = h.retrieve(query, limit=5)
|
||||
assert results == []
|
||||
|
||||
def test_retrieve_returns_relevant_objects(self):
|
||||
embedder = DummyEmbedder()
|
||||
h = SessionHierarchy(embedder)
|
||||
emb = embedder.embed("python programming")
|
||||
obj = _make_stored_object("python programming guide", object_id="o1", embedding=emb)
|
||||
h.rebuild([obj])
|
||||
results = h.retrieve(emb, limit=5)
|
||||
assert len(results) == 1
|
||||
assert results[0].id == "o1"
|
||||
|
||||
def test_retrieve_respects_limit(self):
|
||||
embedder = DummyEmbedder()
|
||||
h = SessionHierarchy(embedder)
|
||||
emb = embedder.embed("shared")
|
||||
objects = [
|
||||
_make_stored_object(f"shared content {i}", object_id=f"o{i}", embedding=emb)
|
||||
for i in range(20)
|
||||
]
|
||||
h.rebuild(objects)
|
||||
results = h.retrieve(emb, limit=5)
|
||||
assert len(results) <= 5
|
||||
|
||||
def test_retrieve_empty_query(self):
|
||||
embedder = DummyEmbedder()
|
||||
h = SessionHierarchy(embedder)
|
||||
obj = _make_stored_object("test", object_id="o1", embedding=embedder.embed("test"))
|
||||
h.rebuild([obj])
|
||||
results = h.retrieve([], limit=5)
|
||||
assert results == []
|
||||
|
||||
def test_retrieve_disabled_falls_back_to_flat(self):
|
||||
embedder = DummyEmbedder()
|
||||
h = SessionHierarchy(embedder)
|
||||
emb = embedder.embed("test")
|
||||
obj = _make_stored_object("test content", object_id="o1", embedding=emb)
|
||||
h.rebuild([obj])
|
||||
h.enabled = False
|
||||
results = h.retrieve(emb, limit=5)
|
||||
# Falls back to flat search
|
||||
assert len(results) == 1
|
||||
|
||||
def test_retrieve_sorted_by_similarity(self):
|
||||
"""Results should be sorted by similarity to query (highest first)."""
|
||||
embedder = DummyEmbedder()
|
||||
h = SessionHierarchy(embedder)
|
||||
|
||||
query_emb = embedder.embed("target query")
|
||||
# Create objects with varying similarity to query
|
||||
obj_close = _make_stored_object(
|
||||
"target query exact", object_id="close", embedding=query_emb
|
||||
)
|
||||
obj_far = _make_stored_object(
|
||||
"completely unrelated xyz",
|
||||
object_id="far",
|
||||
embedding=embedder.embed("completely unrelated xyz"),
|
||||
)
|
||||
|
||||
h.rebuild([obj_close, obj_far])
|
||||
results = h.retrieve(query_emb, limit=10)
|
||||
assert len(results) >= 1
|
||||
# The close object should be first
|
||||
assert results[0].id == "close"
|
||||
|
||||
def test_retrieve_no_duplicates(self):
|
||||
"""Retrieve should not return duplicate objects."""
|
||||
embedder = DummyEmbedder()
|
||||
h = SessionHierarchy(embedder)
|
||||
emb = embedder.embed("shared")
|
||||
objects = [
|
||||
_make_stored_object(f"shared {i}", object_id=f"o{i}", embedding=emb) for i in range(5)
|
||||
]
|
||||
h.rebuild(objects)
|
||||
results = h.retrieve(emb, limit=10)
|
||||
ids = [r.id for r in results]
|
||||
assert len(ids) == len(set(ids))
|
||||
|
||||
def test_flat_search_fallback(self):
|
||||
"""When hierarchy has no themes, should fall back to flat search."""
|
||||
embedder = DummyEmbedder()
|
||||
h = SessionHierarchy(embedder)
|
||||
emb = embedder.embed("test")
|
||||
obj = _make_stored_object("test", object_id="o1", embedding=emb)
|
||||
# Add object but clear themes to force fallback
|
||||
h._all_objects = [obj]
|
||||
h._object_ids = {obj.id}
|
||||
h._themes = []
|
||||
results = h.retrieve(emb, limit=5)
|
||||
assert len(results) == 1
|
||||
|
||||
|
||||
class TestSessionHierarchyMaintenance:
|
||||
def test_maintenance_at_interval(self):
|
||||
embedder = DummyEmbedder()
|
||||
h = SessionHierarchy(embedder)
|
||||
emb = embedder.embed("test")
|
||||
obj = _make_stored_object("test", object_id="o1", embedding=emb)
|
||||
h.add_object(obj)
|
||||
|
||||
# Turn 20 should trigger rebuild
|
||||
rebuilt = h.maintenance(turn=RECLUSTER_INTERVAL)
|
||||
assert rebuilt is True
|
||||
|
||||
def test_maintenance_not_at_interval(self):
|
||||
embedder = DummyEmbedder()
|
||||
h = SessionHierarchy(embedder)
|
||||
emb = embedder.embed("test")
|
||||
obj = _make_stored_object("test", object_id="o1", embedding=emb)
|
||||
h.add_object(obj)
|
||||
|
||||
rebuilt = h.maintenance(turn=15)
|
||||
assert rebuilt is False
|
||||
|
||||
def test_maintenance_on_goal_change(self):
|
||||
embedder = DummyEmbedder()
|
||||
h = SessionHierarchy(embedder)
|
||||
emb = embedder.embed("test")
|
||||
obj = _make_stored_object("test", object_id="o1", embedding=emb)
|
||||
h.add_object(obj)
|
||||
|
||||
rebuilt = h.maintenance(turn=5, goal_hash="goal-1")
|
||||
assert rebuilt is True # First goal → always rebuild
|
||||
|
||||
rebuilt = h.maintenance(turn=6, goal_hash="goal-1")
|
||||
assert rebuilt is False # Same goal, not at interval
|
||||
|
||||
rebuilt = h.maintenance(turn=7, goal_hash="goal-2")
|
||||
assert rebuilt is True # Goal changed
|
||||
|
||||
def test_maintenance_disabled(self):
|
||||
embedder = DummyEmbedder()
|
||||
h = SessionHierarchy(embedder)
|
||||
h.enabled = False
|
||||
rebuilt = h.maintenance(turn=RECLUSTER_INTERVAL)
|
||||
assert rebuilt is False
|
||||
|
||||
def test_maintenance_no_double_rebuild(self):
|
||||
"""Same turn should not trigger rebuild twice."""
|
||||
embedder = DummyEmbedder()
|
||||
h = SessionHierarchy(embedder)
|
||||
emb = embedder.embed("test")
|
||||
obj = _make_stored_object("test", object_id="o1", embedding=emb)
|
||||
h.add_object(obj)
|
||||
|
||||
rebuilt1 = h.maintenance(turn=RECLUSTER_INTERVAL)
|
||||
assert rebuilt1 is True
|
||||
rebuilt2 = h.maintenance(turn=RECLUSTER_INTERVAL)
|
||||
assert rebuilt2 is False
|
||||
|
||||
def test_maintenance_at_multiple_intervals(self):
|
||||
embedder = DummyEmbedder()
|
||||
h = SessionHierarchy(embedder)
|
||||
emb = embedder.embed("test")
|
||||
obj = _make_stored_object("test", object_id="o1", embedding=emb)
|
||||
h.add_object(obj)
|
||||
|
||||
assert h.maintenance(turn=20) is True
|
||||
assert h.maintenance(turn=40) is True
|
||||
assert h.maintenance(turn=60) is True
|
||||
|
||||
def test_maintenance_turn_zero_no_rebuild(self):
|
||||
embedder = DummyEmbedder()
|
||||
h = SessionHierarchy(embedder)
|
||||
rebuilt = h.maintenance(turn=0)
|
||||
assert rebuilt is False
|
||||
|
||||
|
||||
class TestSessionHierarchyGetters:
|
||||
def test_get_episodes(self):
|
||||
embedder = DummyEmbedder()
|
||||
h = SessionHierarchy(embedder)
|
||||
emb = embedder.embed("test")
|
||||
obj = _make_stored_object("test", object_id="o1", embedding=emb)
|
||||
h.add_object(obj)
|
||||
episodes = h.get_episodes()
|
||||
assert len(episodes) == 1
|
||||
assert episodes[0].objects[0].id == "o1"
|
||||
|
||||
def test_get_themes(self):
|
||||
embedder = DummyEmbedder()
|
||||
h = SessionHierarchy(embedder)
|
||||
emb = embedder.embed("test")
|
||||
obj = _make_stored_object("test", object_id="o1", embedding=emb)
|
||||
h.add_object(obj)
|
||||
themes = h.get_themes()
|
||||
assert len(themes) >= 1
|
||||
|
||||
def test_get_top_themes(self):
|
||||
embedder = DummyEmbedder()
|
||||
h = SessionHierarchy(embedder)
|
||||
emb = embedder.embed("test")
|
||||
obj = _make_stored_object("test", object_id="o1", embedding=emb)
|
||||
h.add_object(obj)
|
||||
top_themes = h.get_top_themes()
|
||||
assert len(top_themes) >= 1
|
||||
|
||||
|
||||
class TestSessionHierarchySummary:
|
||||
def test_summary_after_rebuild(self):
|
||||
embedder = DummyEmbedder()
|
||||
h = SessionHierarchy(embedder)
|
||||
objects = [
|
||||
_make_stored_object(
|
||||
f"content {i}", object_id=f"o{i}", embedding=embedder.embed(f"content {i}")
|
||||
)
|
||||
for i in range(10)
|
||||
]
|
||||
h.rebuild(objects)
|
||||
s = h.summary()
|
||||
assert s["objects"] == 10
|
||||
assert s["episodes"] > 0
|
||||
assert s["themes"] > 0
|
||||
assert s["enabled"] is True
|
||||
assert isinstance(s["avg_episode_size"], float)
|
||||
|
||||
|
||||
# ── Edge case tests ─────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestEdgeCases:
|
||||
def test_all_identical_embeddings(self):
|
||||
"""All objects with identical embeddings → single episode."""
|
||||
embedder = DummyEmbedder()
|
||||
h = SessionHierarchy(embedder)
|
||||
emb = embedder.embed("identical")
|
||||
objects = [
|
||||
_make_stored_object(f"identical {i}", object_id=f"o{i}", embedding=emb)
|
||||
for i in range(10)
|
||||
]
|
||||
h.rebuild(objects)
|
||||
assert h.episode_count == 1
|
||||
assert h.object_count == 10
|
||||
|
||||
def test_all_objects_no_embeddings(self):
|
||||
"""Objects without embeddings each get their own episode."""
|
||||
embedder = DummyEmbedder()
|
||||
h = SessionHierarchy(embedder)
|
||||
objects = [
|
||||
_make_stored_object(f"no emb {i}", object_id=f"o{i}", embedding=[]) for i in range(5)
|
||||
]
|
||||
h.rebuild(objects)
|
||||
assert h.episode_count == 5
|
||||
|
||||
def test_large_number_of_objects(self):
|
||||
"""Hierarchy should handle 100+ objects without error."""
|
||||
embedder = DummyEmbedder()
|
||||
h = SessionHierarchy(embedder)
|
||||
objects = [
|
||||
_make_stored_object(
|
||||
f"object number {i} with unique content",
|
||||
object_id=f"o{i}",
|
||||
embedding=embedder.embed(f"object number {i} with unique content"),
|
||||
)
|
||||
for i in range(100)
|
||||
]
|
||||
h.rebuild(objects)
|
||||
assert h.object_count == 100
|
||||
assert h.episode_count > 0
|
||||
assert len(h.get_top_themes()) > 0
|
||||
|
||||
def test_retrieve_with_many_objects(self):
|
||||
"""Retrieval should work efficiently with many objects."""
|
||||
embedder = DummyEmbedder()
|
||||
h = SessionHierarchy(embedder)
|
||||
objects = [
|
||||
_make_stored_object(
|
||||
f"content {i}", object_id=f"o{i}", embedding=embedder.embed(f"content {i}")
|
||||
)
|
||||
for i in range(50)
|
||||
]
|
||||
h.rebuild(objects)
|
||||
query = embedder.embed("content 0")
|
||||
results = h.retrieve(query, limit=5)
|
||||
assert len(results) <= 5
|
||||
# The exact match should be in results
|
||||
result_ids = [r.id for r in results]
|
||||
assert "o0" in result_ids
|
||||
|
||||
def test_incremental_then_rebuild_consistency(self):
|
||||
"""Incremental adds followed by rebuild should produce valid hierarchy."""
|
||||
embedder = DummyEmbedder()
|
||||
h = SessionHierarchy(embedder)
|
||||
|
||||
objects = []
|
||||
for i in range(10):
|
||||
obj = _make_stored_object(
|
||||
f"content {i}", object_id=f"o{i}", embedding=embedder.embed(f"content {i}")
|
||||
)
|
||||
objects.append(obj)
|
||||
h.add_object(obj)
|
||||
|
||||
# Now rebuild from scratch
|
||||
h.rebuild(objects)
|
||||
assert h.object_count == 10
|
||||
assert h.episode_count > 0
|
||||
|
||||
def test_rebuild_with_single_object_no_crash(self):
|
||||
embedder = DummyEmbedder()
|
||||
h = SessionHierarchy(embedder)
|
||||
obj = _make_stored_object("solo", object_id="o1", embedding=embedder.embed("solo"))
|
||||
h.rebuild([obj])
|
||||
results = h.retrieve(embedder.embed("solo"), limit=5)
|
||||
assert len(results) == 1
|
||||
|
||||
def test_episode_not_equal_to_non_episode(self):
|
||||
ep = Episode(id="ep-1")
|
||||
assert ep != "not an episode"
|
||||
|
||||
def test_semantic_theme_not_equal_to_non_theme(self):
|
||||
st = SemanticTheme(id="st-1")
|
||||
assert st != 42
|
||||
|
||||
def test_theme_not_equal_to_non_theme(self):
|
||||
t = Theme(id="t-1")
|
||||
assert t != []
|
||||
|
||||
|
||||
# ── Constants tests ─────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestConstants:
|
||||
def test_episode_threshold(self):
|
||||
assert EPISODE_SIMILARITY_THRESHOLD == 0.6
|
||||
|
||||
def test_semantic_threshold(self):
|
||||
assert SEMANTIC_SIMILARITY_THRESHOLD == 0.4
|
||||
|
||||
def test_theme_threshold(self):
|
||||
assert THEME_SIMILARITY_THRESHOLD == 0.25
|
||||
|
||||
def test_retrieval_thresholds(self):
|
||||
assert THEME_RETRIEVAL_THRESHOLD == 0.3
|
||||
assert EPISODE_RETRIEVAL_THRESHOLD == 0.4
|
||||
|
||||
def test_recluster_interval(self):
|
||||
assert RECLUSTER_INTERVAL == 20
|
||||
180
tests/test_phantom_memory_query.py
Normal file
180
tests/test_phantom_memory_query.py
Normal file
|
|
@ -0,0 +1,180 @@
|
|||
"""Tests for the memory_query phantom tool addition to phantom.py."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import copy
|
||||
|
||||
from mnemosyne.phantom import (
|
||||
PHANTOM_TOOL_DEFINITIONS,
|
||||
PHANTOM_TOOL_NAMES,
|
||||
PhantomCall,
|
||||
_handle_phantom_call,
|
||||
inject_tools,
|
||||
)
|
||||
|
||||
|
||||
# ── memory_query in PHANTOM_TOOL_NAMES ───────────────────────
|
||||
|
||||
|
||||
def test_memory_query_in_phantom_tool_names():
|
||||
"""memory_query must be a recognized phantom tool name."""
|
||||
assert "memory_query" in PHANTOM_TOOL_NAMES
|
||||
|
||||
|
||||
def test_phantom_tool_names_is_frozenset():
|
||||
"""PHANTOM_TOOL_NAMES must remain a frozenset (immutable)."""
|
||||
assert isinstance(PHANTOM_TOOL_NAMES, frozenset)
|
||||
|
||||
|
||||
# ── memory_query tool definition schema ──────────────────────
|
||||
|
||||
|
||||
def _get_memory_query_def() -> dict:
|
||||
"""Helper: find the memory_query definition from PHANTOM_TOOL_DEFINITIONS."""
|
||||
for defn in PHANTOM_TOOL_DEFINITIONS:
|
||||
if defn["name"] == "memory_query":
|
||||
return defn
|
||||
raise AssertionError("memory_query not found in PHANTOM_TOOL_DEFINITIONS")
|
||||
|
||||
|
||||
def test_memory_query_definition_exists():
|
||||
"""memory_query must have a definition in PHANTOM_TOOL_DEFINITIONS."""
|
||||
defn = _get_memory_query_def()
|
||||
assert defn["name"] == "memory_query"
|
||||
|
||||
|
||||
def test_memory_query_has_description():
|
||||
"""memory_query definition must have a non-empty description."""
|
||||
defn = _get_memory_query_def()
|
||||
assert isinstance(defn["description"], str)
|
||||
assert len(defn["description"]) > 20
|
||||
|
||||
|
||||
def test_memory_query_schema_properties():
|
||||
"""memory_query input_schema must have question, scope, max_tokens."""
|
||||
defn = _get_memory_query_def()
|
||||
schema = defn["input_schema"]
|
||||
assert schema["type"] == "object"
|
||||
props = schema["properties"]
|
||||
assert "question" in props
|
||||
assert props["question"]["type"] == "string"
|
||||
assert "scope" in props
|
||||
assert props["scope"]["type"] == "string"
|
||||
assert "max_tokens" in props
|
||||
assert props["max_tokens"]["type"] == "integer"
|
||||
|
||||
|
||||
def test_memory_query_required_fields():
|
||||
"""Only 'question' should be required for memory_query."""
|
||||
defn = _get_memory_query_def()
|
||||
required = defn["input_schema"]["required"]
|
||||
assert required == ["question"]
|
||||
|
||||
|
||||
# ── inject_tools adds memory_query ───────────────────────────
|
||||
|
||||
|
||||
def test_inject_tools_adds_memory_query():
|
||||
"""inject_tools should add memory_query to the tools array."""
|
||||
body = {"tools": [], "messages": []}
|
||||
inject_tools(body)
|
||||
tool_names = {t["name"] for t in body["tools"]}
|
||||
assert "memory_query" in tool_names
|
||||
|
||||
|
||||
def test_inject_tools_does_not_duplicate_memory_query():
|
||||
"""inject_tools should not duplicate memory_query if already present."""
|
||||
existing_def = {
|
||||
"name": "memory_query",
|
||||
"description": "already here",
|
||||
"input_schema": {"type": "object", "properties": {}},
|
||||
}
|
||||
body = {"tools": [existing_def], "messages": []}
|
||||
inject_tools(body)
|
||||
mq_count = sum(1 for t in body["tools"] if t["name"] == "memory_query")
|
||||
assert mq_count == 1
|
||||
|
||||
|
||||
def test_inject_tools_returns_memory_query_as_observe_only_when_framework_provides():
|
||||
"""If framework already has memory_query, it should be in observe_only set."""
|
||||
existing_def = {
|
||||
"name": "memory_query",
|
||||
"description": "framework provided",
|
||||
"input_schema": {"type": "object", "properties": {}},
|
||||
}
|
||||
body = {"tools": [existing_def], "messages": []}
|
||||
observe_only = inject_tools(body)
|
||||
assert "memory_query" in observe_only
|
||||
|
||||
|
||||
# ── _handle_phantom_call returns pending placeholder ─────────
|
||||
|
||||
|
||||
def test_handle_phantom_call_memory_query_returns_pending():
|
||||
"""memory_query handler should return a pending placeholder string."""
|
||||
call = PhantomCall(
|
||||
name="memory_query",
|
||||
tool_use_id="toolu_test123",
|
||||
input={"question": "What auth library is used?"},
|
||||
)
|
||||
result = _handle_phantom_call(call, page_store=None)
|
||||
assert "[memory_query:pending]" in result
|
||||
assert "What auth library is used?" in result
|
||||
|
||||
|
||||
def test_handle_phantom_call_memory_query_includes_scope():
|
||||
"""memory_query handler should include scope in the placeholder."""
|
||||
call = PhantomCall(
|
||||
name="memory_query",
|
||||
tool_use_id="toolu_test456",
|
||||
input={
|
||||
"question": "What port does the server run on?",
|
||||
"scope": "config files",
|
||||
},
|
||||
)
|
||||
result = _handle_phantom_call(call, page_store=None)
|
||||
assert "[memory_query:pending]" in result
|
||||
assert "config files" in result
|
||||
|
||||
|
||||
def test_handle_phantom_call_memory_query_includes_max_tokens():
|
||||
"""memory_query handler should include max_tokens in the placeholder."""
|
||||
call = PhantomCall(
|
||||
name="memory_query",
|
||||
tool_use_id="toolu_test789",
|
||||
input={
|
||||
"question": "What is the DB schema?",
|
||||
"max_tokens": 100,
|
||||
},
|
||||
)
|
||||
result = _handle_phantom_call(call, page_store=None)
|
||||
assert "[memory_query:pending]" in result
|
||||
assert "100" in result
|
||||
|
||||
|
||||
def test_handle_phantom_call_memory_query_defaults():
|
||||
"""memory_query handler should use defaults for optional fields."""
|
||||
call = PhantomCall(
|
||||
name="memory_query",
|
||||
tool_use_id="toolu_defaults",
|
||||
input={"question": "test question"},
|
||||
)
|
||||
result = _handle_phantom_call(call, page_store=None)
|
||||
assert "[memory_query:pending]" in result
|
||||
# Default scope is None, default max_tokens is 200
|
||||
assert "None" in result
|
||||
assert "200" in result
|
||||
|
||||
|
||||
# ── Existing tools still work ────────────────────────────────
|
||||
|
||||
|
||||
def test_existing_phantom_tools_still_present():
|
||||
"""All original phantom tool names must still be in PHANTOM_TOOL_NAMES."""
|
||||
for name in ("yuyay", "recall", "memory_fault", "qunqay", "tiqsiy"):
|
||||
assert name in PHANTOM_TOOL_NAMES, f"{name} missing from PHANTOM_TOOL_NAMES"
|
||||
|
||||
|
||||
def test_existing_tool_definitions_count():
|
||||
"""PHANTOM_TOOL_DEFINITIONS should now have 4 tools (3 original + memory_query)."""
|
||||
assert len(PHANTOM_TOOL_DEFINITIONS) == 4
|
||||
Loading…
Add table
Add a link
Reference in a new issue