feat: add object store with semantic segmentation

Object-addressed memory: segment messages into semantic objects,
embed with sentence-transformers, store in pgvector-backed store,
and reassemble context via goal-aware retrieval.

Ultraworked with [Sisyphus](https://github.com/code-yeongyu/oh-my-opencode)

Co-authored-by: Sisyphus <clio-agent@sisyphuslabs.ai>
This commit is contained in:
Joey Yakimowich-Payne 2026-03-13 11:41:04 -06:00
commit a13719f754
9 changed files with 5644 additions and 0 deletions

View file

@ -0,0 +1,408 @@
"""Tests for the ContextAssembler micro-fault and context assembly module."""
from __future__ import annotations
import asyncio
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from mnemosyne.context_assembler import (
ContextAssembler,
ContextWindow,
MicroFaultResult,
_estimate_tokens,
)
from mnemosyne.fidelity import FidelityLevel, FidelityManager, make_object
from mnemosyne.object_store import (
DummyEmbedder,
InMemoryBackend,
ObjectStore,
StoredObject,
)
# ── Fixtures ─────────────────────────────────────────────────
@pytest.fixture
def object_store():
"""Create an ObjectStore with InMemoryBackend and DummyEmbedder."""
backend = InMemoryBackend()
embedder = DummyEmbedder()
return ObjectStore(backend=backend, embedder=embedder)
@pytest.fixture
def mock_helper_llm():
"""Create a mock HelperLLM with answer_micro_fault returning a canned answer."""
helper = AsyncMock()
helper.answer_micro_fault = AsyncMock(
return_value="The auth library used is passport.js with JWT strategy."
)
return helper
@pytest.fixture
def assembler(object_store, mock_helper_llm):
"""Create a ContextAssembler with real ObjectStore and mock HelperLLM."""
return ContextAssembler(object_store=object_store, helper_llm=mock_helper_llm)
@pytest.fixture
def assembler_no_helper(object_store):
"""Create a ContextAssembler without HelperLLM (fallback mode)."""
return ContextAssembler(object_store=object_store, helper_llm=None)
async def _seed_objects(store: ObjectStore, session_id: str = "sess_test") -> list[StoredObject]:
"""Seed the store with a few test objects and return them."""
obj1 = await store.store_object(
session_id=session_id,
content="Authentication uses passport.js with JWT strategy. "
"The secret key is loaded from AUTH_SECRET env var. "
"Token expiry is set to 24 hours.",
object_type="file_context",
source_tool="Read",
source_key="src/auth/middleware.ts",
stub="[file_context | auth middleware | passport.js JWT]",
key_entities=["passport.js", "JWT", "AUTH_SECRET"],
)
obj2 = await store.store_object(
session_id=session_id,
content="Database schema has users, posts, and comments tables. "
"Users table has id, email, password_hash, created_at columns. "
"Posts reference users via author_id foreign key.",
object_type="file_context",
source_tool="Read",
source_key="schema.sql",
stub="[file_context | DB schema | users, posts, comments]",
key_entities=["users", "posts", "comments", "schema.sql"],
)
obj3 = await store.store_object(
session_id=session_id,
content="Server runs on port 3000 by default. "
"Configuration loaded from .env file. "
"CORS enabled for localhost:5173 in development.",
object_type="file_context",
source_tool="Read",
source_key="src/config.ts",
stub="[file_context | server config | port 3000, CORS]",
key_entities=["port 3000", "CORS", ".env"],
)
return [obj1, obj2, obj3]
# ── MicroFaultResult dataclass ───────────────────────────────
def test_micro_fault_result_fields():
"""MicroFaultResult should store all expected fields."""
result = MicroFaultResult(
answer="The answer is 42.",
sources=["obj_abc", "obj_def"],
answer_tokens=5,
avoided_tokens=500,
latency_ms=12.5,
)
assert result.answer == "The answer is 42."
assert result.sources == ["obj_abc", "obj_def"]
assert result.answer_tokens == 5
assert result.avoided_tokens == 500
assert result.latency_ms == 12.5
def test_micro_fault_result_token_savings():
"""avoided_tokens should represent tokens saved vs full restore."""
result = MicroFaultResult(
answer="short answer",
sources=["obj1"],
answer_tokens=3,
avoided_tokens=997,
latency_ms=10.0,
)
# The savings ratio: avoided / (avoided + answer_tokens)
savings_ratio = result.avoided_tokens / (result.avoided_tokens + result.answer_tokens)
assert savings_ratio > 0.99 # 99%+ savings
# ── ContextWindow dataclass ──────────────────────────────────
def test_context_window_fields():
"""ContextWindow should store objects, total_tokens, and pressure_zone."""
window = ContextWindow(
objects=[("obj1", 0, "full content"), ("obj2", 2, "compact summary")],
total_tokens=150,
pressure_zone="NORMAL",
)
assert len(window.objects) == 2
assert window.total_tokens == 150
assert window.pressure_zone == "NORMAL"
# ── handle_micro_fault with HelperLLM ────────────────────────
@pytest.mark.asyncio
async def test_micro_fault_with_helper(assembler, object_store, mock_helper_llm):
"""Micro-fault with HelperLLM should search, call helper, return answer."""
await _seed_objects(object_store, "sess_test")
result = await assembler.handle_micro_fault(
session_id="sess_test",
question="What auth library is used?",
)
assert isinstance(result, MicroFaultResult)
assert result.answer == "The auth library used is passport.js with JWT strategy."
assert len(result.sources) > 0
assert result.answer_tokens > 0
assert result.avoided_tokens > 0
assert result.latency_ms >= 0
# Verify HelperLLM was called
mock_helper_llm.answer_micro_fault.assert_called_once()
call_kwargs = mock_helper_llm.answer_micro_fault.call_args
assert call_kwargs.kwargs["question"] == "What auth library is used?"
assert len(call_kwargs.kwargs["relevant_contents"]) > 0
@pytest.mark.asyncio
async def test_micro_fault_with_scope(assembler, object_store, mock_helper_llm):
"""Micro-fault with scope hint should incorporate it into the search query."""
await _seed_objects(object_store, "sess_test")
result = await assembler.handle_micro_fault(
session_id="sess_test",
question="What port?",
scope="config files",
)
assert isinstance(result, MicroFaultResult)
assert len(result.sources) > 0
mock_helper_llm.answer_micro_fault.assert_called_once()
@pytest.mark.asyncio
async def test_micro_fault_with_max_tokens(assembler, object_store, mock_helper_llm):
"""Micro-fault should pass max_tokens to the HelperLLM."""
await _seed_objects(object_store, "sess_test")
await assembler.handle_micro_fault(
session_id="sess_test",
question="What is the DB schema?",
max_tokens=100,
)
call_kwargs = mock_helper_llm.answer_micro_fault.call_args
assert call_kwargs.kwargs["max_tokens"] == 100
# ── handle_micro_fault fallback (no HelperLLM) ──────────────
@pytest.mark.asyncio
async def test_micro_fault_fallback_no_helper(assembler_no_helper, object_store):
"""Without HelperLLM, micro-fault should return summaries as fallback."""
objects = await _seed_objects(object_store, "sess_test")
# Set a summary on one object so fallback has something to show
await object_store.update_fidelity(objects[0].id, 2, summary="Auth uses passport.js JWT")
result = await assembler_no_helper.handle_micro_fault(
session_id="sess_test",
question="What auth library is used?",
)
assert isinstance(result, MicroFaultResult)
assert "HelperLLM unavailable" in result.answer
assert len(result.sources) > 0
assert result.answer_tokens > 0
@pytest.mark.asyncio
async def test_micro_fault_fallback_uses_stub_when_no_summary(assembler_no_helper, object_store):
"""Fallback should use stub when no summaries are available."""
await _seed_objects(object_store, "sess_test")
result = await assembler_no_helper.handle_micro_fault(
session_id="sess_test",
question="What auth library is used?",
)
assert isinstance(result, MicroFaultResult)
assert "HelperLLM unavailable" in result.answer
# Should contain stub content since no summaries exist
assert len(result.sources) > 0
# ── handle_micro_fault with no search results ────────────────
@pytest.mark.asyncio
async def test_micro_fault_no_results(assembler):
"""Micro-fault with no matching objects should return a 'not found' message."""
# Don't seed any objects — empty store
result = await assembler.handle_micro_fault(
session_id="sess_empty",
question="What is the meaning of life?",
)
assert isinstance(result, MicroFaultResult)
assert "No relevant content found" in result.answer
assert result.sources == []
assert result.avoided_tokens == 0
# ── Micro-fault records access on consulted objects ──────────
@pytest.mark.asyncio
async def test_micro_fault_records_access(assembler, object_store):
"""Micro-fault should call record_fault(is_micro=True) on consulted objects."""
objects = await _seed_objects(object_store, "sess_test")
result = await assembler.handle_micro_fault(
session_id="sess_test",
question="What auth library is used?",
)
# At least one source should have been consulted
assert len(result.sources) > 0
# Check that micro_fault_count was incremented on consulted objects
for source_id in result.sources:
obj = await object_store.get(source_id)
assert obj is not None
assert obj.micro_fault_count > 0
# ── assemble_context ─────────────────────────────────────────
@pytest.mark.asyncio
async def test_assemble_context_all_l0(assembler, object_store):
"""assemble_context with all L0 objects should include full content."""
objects = await _seed_objects(object_store, "sess_test")
fm = FidelityManager(window_size=200_000)
for obj in objects:
fm_obj = make_object(
object_type=obj.object_type,
content_full=obj.content_full,
stub=obj.stub,
)
fm_obj.id = obj.id # match IDs
fm.register_object(fm_obj)
blocks = await assembler.assemble_context(
session_id="sess_test",
fidelity_manager=fm,
current_turn=1,
)
assert len(blocks) == 3
for block in blocks:
assert block["fidelity"] == 0 # L0
assert len(block["content"]) > 50 # full content
assert block["tokens"] > 0
@pytest.mark.asyncio
async def test_assemble_context_mixed_fidelity(assembler, object_store):
"""assemble_context should respect per-object fidelity levels."""
objects = await _seed_objects(object_store, "sess_test")
# Set different fidelity levels in the store
await object_store.update_fidelity(
objects[0].id,
0, # L0: full
)
await object_store.update_fidelity(
objects[1].id,
2,
summary="DB has users, posts, comments tables", # L2: compact
)
await object_store.update_fidelity(
objects[2].id,
3, # L3: stub
)
fm = FidelityManager(window_size=200_000)
for obj in objects:
fm_obj = make_object(
object_type=obj.object_type,
content_full=obj.content_full,
stub=obj.stub,
)
fm_obj.id = obj.id
# Set fidelity to match what we set in the store
stored = await object_store.get(obj.id)
fm_obj.current_fidelity = FidelityLevel(stored.current_fidelity)
if stored.summary_compact:
fm_obj.summary_compact = stored.summary_compact
fm.register_object(fm_obj)
blocks = await assembler.assemble_context(
session_id="sess_test",
fidelity_manager=fm,
current_turn=5,
)
assert len(blocks) == 3
fidelities = {b["object_id"]: b["fidelity"] for b in blocks}
assert fidelities[objects[0].id] == 0 # L0
assert fidelities[objects[1].id] == 2 # L2
assert fidelities[objects[2].id] == 3 # L3
@pytest.mark.asyncio
async def test_assemble_context_excludes_evicted(assembler, object_store):
"""assemble_context should exclude L4 (evicted) objects."""
objects = await _seed_objects(object_store, "sess_test")
# Evict one object
await object_store.update_fidelity(objects[1].id, 4)
fm = FidelityManager(window_size=200_000)
for obj in objects:
fm_obj = make_object(
object_type=obj.object_type,
content_full=obj.content_full,
stub=obj.stub,
)
fm_obj.id = obj.id
stored = await object_store.get(obj.id)
fm_obj.current_fidelity = FidelityLevel(min(stored.current_fidelity, 4))
fm.register_object(fm_obj)
blocks = await assembler.assemble_context(
session_id="sess_test",
fidelity_manager=fm,
current_turn=5,
)
# Only 2 objects should be in context (one was evicted)
block_ids = {b["object_id"] for b in blocks}
assert objects[1].id not in block_ids
assert len(blocks) == 2
# ── _estimate_tokens helper ──────────────────────────────────
def test_estimate_tokens_none():
"""_estimate_tokens(None) should return 0."""
assert _estimate_tokens(None) == 0
def test_estimate_tokens_empty():
"""_estimate_tokens('') should return 1 (minimum)."""
assert _estimate_tokens("") == 1
def test_estimate_tokens_normal():
"""_estimate_tokens should use ~4 chars per token heuristic."""
text = "a" * 400
assert _estimate_tokens(text) == 100

996
tests/test_object_store.py Normal file
View file

@ -0,0 +1,996 @@
"""Tests for the semantic object backing store."""
from __future__ import annotations
import asyncio
import pytest
from mnemosyne.object_store import (
DummyEmbedder,
InMemoryBackend,
ObjectStore,
ObjectStoreBackend,
StoredObject,
_cosine_similarity,
_estimate_tokens,
)
# ── Helpers ──────────────────────────────────────────────────
def _make_stored_object(
session_id: str = "sess-1",
content: str = "x" * 400,
*,
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,
) -> StoredObject:
"""Create a StoredObject with sensible defaults for testing."""
return StoredObject(
id=object_id or f"obj-{id(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}: test object",
tokens_l0=_estimate_tokens(content),
tokens_l3=_estimate_tokens(stub or f"{object_type}: test object"),
embedding=embedding or [],
created_at="2025-01-01T00:00:00+00:00",
last_accessed="2025-01-01T00:00:00+00:00",
)
# ── StoredObject dataclass ───────────────────────────────────
class TestStoredObject:
def test_required_fields(self):
obj = _make_stored_object()
assert obj.id
assert obj.session_id == "sess-1"
assert obj.object_type == "file_context"
assert obj.content_full == "x" * 400
assert obj.current_fidelity == 0
assert obj.pinned is False
assert obj.fault_count == 0
assert obj.micro_fault_count == 0
def test_default_lists_are_empty(self):
obj = _make_stored_object()
assert obj.losses_l1 == []
assert obj.losses_l2 == []
assert obj.can_answer_l1 == []
assert obj.can_answer_l2 == []
assert obj.fault_when == []
assert obj.key_entities == []
assert obj.tags == []
def test_content_at_levels(self):
obj = _make_stored_object(content="full content")
obj.summary_detailed = "detailed"
obj.summary_compact = "compact"
obj.stub = "stub line"
assert obj.content_at(0) == "full content"
assert obj.content_at(1) == "detailed"
assert obj.content_at(2) == "compact"
assert obj.content_at(3) == "stub line"
assert obj.content_at(4) is None
def test_tokens_at_levels(self):
obj = _make_stored_object(content="a" * 400)
obj.tokens_l0 = 100
obj.tokens_l1 = 30
obj.tokens_l2 = 5
obj.tokens_l3 = 3
assert obj.tokens_at(0) == 100
assert obj.tokens_at(1) == 30
assert obj.tokens_at(2) == 5
assert obj.tokens_at(3) == 3
assert obj.tokens_at(4) == 0
def test_tokens_at_none_levels(self):
obj = _make_stored_object()
obj.tokens_l1 = None
obj.tokens_l2 = None
assert obj.tokens_at(1) == 0
assert obj.tokens_at(2) == 0
def test_current_tokens_tracks_fidelity(self):
obj = _make_stored_object(content="a" * 400)
obj.tokens_l0 = 100
obj.tokens_l1 = 30
assert obj.current_tokens == 100 # L0
obj.current_fidelity = 1
assert obj.current_tokens == 30 # L1
def test_embedding_default_empty(self):
obj = _make_stored_object()
assert obj.embedding == []
# ── DummyEmbedder ────────────────────────────────────────────
class TestDummyEmbedder:
def test_embed_returns_384_dim(self):
emb = DummyEmbedder()
vec = emb.embed("hello world")
assert len(vec) == 384
def test_embed_deterministic(self):
emb = DummyEmbedder()
v1 = emb.embed("test input")
v2 = emb.embed("test input")
assert v1 == v2
def test_embed_different_inputs_differ(self):
emb = DummyEmbedder()
v1 = emb.embed("input A")
v2 = emb.embed("input B")
assert v1 != v2
def test_embed_is_normalized(self):
import numpy as np
emb = DummyEmbedder()
vec = emb.embed("normalize me")
norm = float(np.linalg.norm(vec))
assert abs(norm - 1.0) < 1e-6
def test_embed_batch(self):
emb = DummyEmbedder()
texts = ["alpha", "beta", "gamma"]
vecs = emb.embed_batch(texts)
assert len(vecs) == 3
assert all(len(v) == 384 for v in vecs)
def test_embed_batch_matches_single(self):
emb = DummyEmbedder()
texts = ["one", "two"]
batch = emb.embed_batch(texts)
singles = [emb.embed(t) for t in texts]
assert batch == singles
# ── Cosine similarity ────────────────────────────────────────
class TestCosineSimilarity:
def test_identical_vectors(self):
v = [1.0, 0.0, 0.0]
assert abs(_cosine_similarity(v, v) - 1.0) < 1e-9
def test_orthogonal_vectors(self):
a = [1.0, 0.0, 0.0]
b = [0.0, 1.0, 0.0]
assert abs(_cosine_similarity(a, b)) < 1e-9
def test_opposite_vectors(self):
a = [1.0, 0.0]
b = [-1.0, 0.0]
assert abs(_cosine_similarity(a, b) - (-1.0)) < 1e-9
def test_empty_vectors(self):
assert _cosine_similarity([], []) == 0.0
assert _cosine_similarity([1.0], []) == 0.0
def test_zero_vector(self):
assert _cosine_similarity([0.0, 0.0], [1.0, 0.0]) == 0.0
# ── InMemoryBackend CRUD ─────────────────────────────────────
class TestInMemoryBackendCRUD:
async def test_store_and_get(self):
backend = InMemoryBackend()
obj = _make_stored_object(object_id="obj-001")
await backend.store(obj)
result = await backend.get("obj-001")
assert result is obj
async def test_get_nonexistent(self):
backend = InMemoryBackend()
assert await backend.get("nonexistent") is None
async def test_store_overwrites(self):
backend = InMemoryBackend()
obj1 = _make_stored_object(object_id="obj-001", content="original")
obj2 = _make_stored_object(object_id="obj-001", content="updated")
await backend.store(obj1)
await backend.store(obj2)
result = await backend.get("obj-001")
assert result is not None
assert result.content_full == "updated"
async def test_get_by_session(self):
backend = InMemoryBackend()
obj1 = _make_stored_object(session_id="s1", object_id="o1")
obj2 = _make_stored_object(session_id="s1", object_id="o2")
obj3 = _make_stored_object(session_id="s2", object_id="o3")
await backend.store(obj1)
await backend.store(obj2)
await backend.store(obj3)
s1_objs = await backend.get_by_session("s1")
assert len(s1_objs) == 2
assert {o.id for o in s1_objs} == {"o1", "o2"}
async def test_get_by_session_filters_fidelity(self):
backend = InMemoryBackend()
obj1 = _make_stored_object(session_id="s1", object_id="o1")
obj2 = _make_stored_object(session_id="s1", object_id="o2")
obj2.current_fidelity = 4 # evicted
await backend.store(obj1)
await backend.store(obj2)
# Default: include evicted
all_objs = await backend.get_by_session("s1", fidelity_max=4)
assert len(all_objs) == 2
# Exclude evicted
active_objs = await backend.get_by_session("s1", fidelity_max=3)
assert len(active_objs) == 1
assert active_objs[0].id == "o1"
async def test_get_by_session_empty(self):
backend = InMemoryBackend()
assert await backend.get_by_session("nonexistent") == []
async def test_delete_session(self):
backend = InMemoryBackend()
obj1 = _make_stored_object(session_id="s1", object_id="o1")
obj2 = _make_stored_object(session_id="s1", object_id="o2")
obj3 = _make_stored_object(session_id="s2", object_id="o3")
await backend.store(obj1)
await backend.store(obj2)
await backend.store(obj3)
count = await backend.delete_session("s1")
assert count == 2
assert await backend.get("o1") is None
assert await backend.get("o2") is None
assert await backend.get("o3") is not None
async def test_delete_session_nonexistent(self):
backend = InMemoryBackend()
count = await backend.delete_session("nonexistent")
assert count == 0
# ── InMemoryBackend fidelity updates ─────────────────────────
class TestInMemoryBackendFidelity:
async def test_update_fidelity_basic(self):
backend = InMemoryBackend()
obj = _make_stored_object(object_id="o1")
await backend.store(obj)
await backend.update_fidelity("o1", 2)
result = await backend.get("o1")
assert result is not None
assert result.current_fidelity == 2
async def test_update_fidelity_with_summary_l1(self):
backend = InMemoryBackend()
obj = _make_stored_object(object_id="o1")
await backend.store(obj)
await backend.update_fidelity(
"o1", 1, summary="detailed summary", losses=["exact line numbers"]
)
result = await backend.get("o1")
assert result is not None
assert result.current_fidelity == 1
assert result.summary_detailed == "detailed summary"
assert result.losses_l1 == ["exact line numbers"]
assert result.tokens_l1 is not None
assert result.tokens_l1 > 0
async def test_update_fidelity_with_summary_l2(self):
backend = InMemoryBackend()
obj = _make_stored_object(object_id="o1")
await backend.store(obj)
await backend.update_fidelity(
"o1", 2, summary="compact summary", losses=["function bodies"]
)
result = await backend.get("o1")
assert result is not None
assert result.summary_compact == "compact summary"
assert result.losses_l2 == ["function bodies"]
async def test_update_fidelity_nonexistent(self):
backend = InMemoryBackend()
# Should not raise
await backend.update_fidelity("nonexistent", 1)
# ── InMemoryBackend source_key dedup ─────────────────────────
class TestInMemoryBackendSourceKey:
async def test_get_by_source_key(self):
backend = InMemoryBackend()
obj = _make_stored_object(session_id="s1", object_id="o1", source_key="src/auth.py")
await backend.store(obj)
result = await backend.get_by_source_key("s1", "src/auth.py")
assert result is not None
assert result.id == "o1"
async def test_get_by_source_key_not_found(self):
backend = InMemoryBackend()
assert await backend.get_by_source_key("s1", "nonexistent.py") is None
async def test_get_by_source_key_session_isolation(self):
backend = InMemoryBackend()
obj = _make_stored_object(session_id="s1", object_id="o1", source_key="src/auth.py")
await backend.store(obj)
# Different session should not find it
assert await backend.get_by_source_key("s2", "src/auth.py") is None
async def test_get_by_source_key_returns_latest(self):
backend = InMemoryBackend()
obj_old = _make_stored_object(session_id="s1", object_id="o1", source_key="src/auth.py")
obj_old.created_at = "2025-01-01T00:00:00+00:00"
obj_new = _make_stored_object(session_id="s1", object_id="o2", source_key="src/auth.py")
obj_new.created_at = "2025-01-02T00:00:00+00:00"
await backend.store(obj_old)
await backend.store(obj_new)
result = await backend.get_by_source_key("s1", "src/auth.py")
assert result is not None
assert result.id == "o2"
# ── InMemoryBackend embedding search ─────────────────────────
class TestInMemoryBackendEmbeddingSearch:
async def test_search_by_embedding_basic(self):
emb = DummyEmbedder()
backend = InMemoryBackend()
obj = _make_stored_object(
session_id="s1",
object_id="o1",
content="authentication middleware",
embedding=emb.embed("authentication middleware"),
)
await backend.store(obj)
query_vec = emb.embed("authentication middleware")
results = await backend.search_by_embedding("s1", query_vec, limit=5)
assert len(results) == 1
assert results[0][0].id == "o1"
# Same text → same embedding → similarity ≈ 1.0
assert results[0][1] > 0.99
async def test_search_by_embedding_ranking(self):
emb = DummyEmbedder()
backend = InMemoryBackend()
# Store objects with different content
for i, content in enumerate(["auth login", "database schema", "test runner"]):
obj = _make_stored_object(
session_id="s1",
object_id=f"o{i}",
content=content,
embedding=emb.embed(content),
)
await backend.store(obj)
# Search for exact match
query_vec = emb.embed("auth login")
results = await backend.search_by_embedding("s1", query_vec, limit=3)
assert len(results) == 3
# Exact match should be first with highest similarity
assert results[0][0].id == "o0"
assert results[0][1] > results[1][1]
async def test_search_by_embedding_session_isolation(self):
emb = DummyEmbedder()
backend = InMemoryBackend()
obj_s1 = _make_stored_object(
session_id="s1",
object_id="o1",
content="hello",
embedding=emb.embed("hello"),
)
obj_s2 = _make_stored_object(
session_id="s2",
object_id="o2",
content="hello",
embedding=emb.embed("hello"),
)
await backend.store(obj_s1)
await backend.store(obj_s2)
results = await backend.search_by_embedding("s1", emb.embed("hello"), limit=10)
assert len(results) == 1
assert results[0][0].id == "o1"
async def test_search_by_embedding_respects_limit(self):
emb = DummyEmbedder()
backend = InMemoryBackend()
for i in range(10):
obj = _make_stored_object(
session_id="s1",
object_id=f"o{i}",
content=f"content {i}",
embedding=emb.embed(f"content {i}"),
)
await backend.store(obj)
results = await backend.search_by_embedding("s1", emb.embed("content 0"), limit=3)
assert len(results) == 3
async def test_search_by_embedding_skips_no_embedding(self):
backend = InMemoryBackend()
obj = _make_stored_object(session_id="s1", object_id="o1", embedding=[])
await backend.store(obj)
results = await backend.search_by_embedding("s1", [0.1] * 384, limit=5)
assert len(results) == 0
# ── InMemoryBackend text search ──────────────────────────────
class TestInMemoryBackendTextSearch:
async def test_search_by_text_content(self):
backend = InMemoryBackend()
obj = _make_stored_object(
session_id="s1",
object_id="o1",
content="authentication middleware for JWT tokens",
)
await backend.store(obj)
results = await backend.search_by_text("s1", "JWT")
assert len(results) == 1
assert results[0].id == "o1"
async def test_search_by_text_case_insensitive(self):
backend = InMemoryBackend()
obj = _make_stored_object(session_id="s1", object_id="o1", content="Authentication")
await backend.store(obj)
results = await backend.search_by_text("s1", "authentication")
assert len(results) == 1
async def test_search_by_text_in_stub(self):
backend = InMemoryBackend()
obj = _make_stored_object(
session_id="s1",
object_id="o1",
content="some content",
stub="file_context: auth middleware",
)
await backend.store(obj)
results = await backend.search_by_text("s1", "auth middleware")
assert len(results) == 1
async def test_search_by_text_in_key_entities(self):
backend = InMemoryBackend()
obj = _make_stored_object(session_id="s1", object_id="o1", content="code")
obj.key_entities = ["AuthService", "JWTValidator"]
await backend.store(obj)
results = await backend.search_by_text("s1", "JWTValidator")
assert len(results) == 1
async def test_search_by_text_no_match(self):
backend = InMemoryBackend()
obj = _make_stored_object(session_id="s1", object_id="o1", content="hello")
await backend.store(obj)
results = await backend.search_by_text("s1", "nonexistent")
assert len(results) == 0
async def test_search_by_text_session_isolation(self):
backend = InMemoryBackend()
obj = _make_stored_object(session_id="s1", object_id="o1", content="shared keyword")
await backend.store(obj)
results = await backend.search_by_text("s2", "shared keyword")
assert len(results) == 0
async def test_search_by_text_respects_limit(self):
backend = InMemoryBackend()
for i in range(10):
obj = _make_stored_object(
session_id="s1",
object_id=f"o{i}",
content=f"common keyword item {i}",
)
await backend.store(obj)
results = await backend.search_by_text("s1", "common keyword", limit=3)
assert len(results) == 3
# ── Session isolation ────────────────────────────────────────
class TestSessionIsolation:
async def test_objects_isolated_by_session(self):
backend = InMemoryBackend()
obj_s1 = _make_stored_object(session_id="s1", object_id="o1")
obj_s2 = _make_stored_object(session_id="s2", object_id="o2")
await backend.store(obj_s1)
await backend.store(obj_s2)
s1_objs = await backend.get_by_session("s1")
s2_objs = await backend.get_by_session("s2")
assert len(s1_objs) == 1
assert s1_objs[0].id == "o1"
assert len(s2_objs) == 1
assert s2_objs[0].id == "o2"
async def test_delete_session_does_not_affect_other(self):
backend = InMemoryBackend()
obj_s1 = _make_stored_object(session_id="s1", object_id="o1")
obj_s2 = _make_stored_object(session_id="s2", object_id="o2")
await backend.store(obj_s1)
await backend.store(obj_s2)
await backend.delete_session("s1")
assert await backend.get("o1") is None
assert await backend.get("o2") is not None
assert len(await backend.get_by_session("s2")) == 1
async def test_source_key_scoped_to_session(self):
backend = InMemoryBackend()
obj_s1 = _make_stored_object(session_id="s1", object_id="o1", source_key="file.py")
obj_s2 = _make_stored_object(session_id="s2", object_id="o2", source_key="file.py")
await backend.store(obj_s1)
await backend.store(obj_s2)
result = await backend.get_by_source_key("s1", "file.py")
assert result is not None
assert result.id == "o1"
# ── ObjectStore facade ───────────────────────────────────────
class TestObjectStoreFacade:
async def test_store_object_creates_with_defaults(self):
store = ObjectStore(InMemoryBackend(), DummyEmbedder())
obj = await store.store_object(
session_id="s1",
content="def authenticate(token): ...",
object_type="file_context",
source_tool="Read",
)
assert obj.id
assert len(obj.id) == 16
assert obj.session_id == "s1"
assert obj.object_type == "file_context"
assert obj.content_full == "def authenticate(token): ..."
assert obj.source_tool == "Read"
assert obj.current_fidelity == 0
assert obj.created_at
assert obj.last_accessed
assert obj.tokens_l0 > 0
assert len(obj.embedding) == 384
async def test_store_object_auto_generates_stub(self):
store = ObjectStore(InMemoryBackend(), DummyEmbedder())
obj = await store.store_object(
session_id="s1",
content="some content here",
object_type="tool_result",
)
assert "tool_result:" in obj.stub
async def test_store_object_custom_stub(self):
store = ObjectStore(InMemoryBackend(), DummyEmbedder())
obj = await store.store_object(
session_id="s1",
content="content",
object_type="file_context",
stub="Read src/auth.py (150 lines)",
)
assert obj.stub == "Read src/auth.py (150 lines)"
async def test_store_object_with_tags_and_entities(self):
store = ObjectStore(InMemoryBackend(), DummyEmbedder())
obj = await store.store_object(
session_id="s1",
content="content",
object_type="file_context",
tags=["auth", "middleware"],
key_entities=["AuthService"],
)
assert obj.tags == ["auth", "middleware"]
assert obj.key_entities == ["AuthService"]
async def test_store_object_with_turn(self):
store = ObjectStore(InMemoryBackend(), DummyEmbedder())
obj = await store.store_object(
session_id="s1",
content="content",
object_type="file_context",
turn=7,
)
assert obj.source_turn_start == 7
assert obj.source_turn_end == 7
async def test_store_object_no_embedder(self):
store = ObjectStore(InMemoryBackend(), embedder=None)
obj = await store.store_object(
session_id="s1",
content="content",
object_type="file_context",
)
assert obj.embedding == []
async def test_get_retrieves_stored(self):
store = ObjectStore(InMemoryBackend(), DummyEmbedder())
obj = await store.store_object(
session_id="s1", content="content", object_type="file_context"
)
result = await store.get(obj.id)
assert result is obj
async def test_get_nonexistent(self):
store = ObjectStore(InMemoryBackend())
assert await store.get("nonexistent") is None
async def test_get_session_objects_excludes_evicted(self):
store = ObjectStore(InMemoryBackend(), DummyEmbedder())
obj1 = await store.store_object(
session_id="s1", content="active", object_type="file_context"
)
obj2 = await store.store_object(
session_id="s1", content="evicted", object_type="file_context"
)
await store.update_fidelity(obj2.id, 4)
active = await store.get_session_objects("s1")
assert len(active) == 1
assert active[0].id == obj1.id
async def test_get_session_objects_includes_evicted(self):
store = ObjectStore(InMemoryBackend(), DummyEmbedder())
await store.store_object(session_id="s1", content="active", object_type="file_context")
obj2 = await store.store_object(
session_id="s1", content="evicted", object_type="file_context"
)
await store.update_fidelity(obj2.id, 4)
all_objs = await store.get_session_objects("s1", include_evicted=True)
assert len(all_objs) == 2
async def test_update_fidelity(self):
store = ObjectStore(InMemoryBackend(), DummyEmbedder())
obj = await store.store_object(
session_id="s1", content="content", object_type="file_context"
)
await store.update_fidelity(obj.id, 1, summary="summary", losses=["details"])
result = await store.get(obj.id)
assert result is not None
assert result.current_fidelity == 1
assert result.summary_detailed == "summary"
assert result.losses_l1 == ["details"]
# ── ObjectStore semantic search ──────────────────────────────
class TestObjectStoreSemanticSearch:
async def test_semantic_search_finds_exact_match(self):
store = ObjectStore(InMemoryBackend(), DummyEmbedder())
await store.store_object(
session_id="s1",
content="authentication middleware for JWT",
object_type="file_context",
)
await store.store_object(
session_id="s1",
content="database migration script",
object_type="file_context",
)
results = await store.semantic_search("s1", "authentication middleware for JWT")
assert len(results) >= 1
# Exact content match should rank first
assert results[0][0].content_full == "authentication middleware for JWT"
async def test_semantic_search_hybrid_boost(self):
"""Text match should boost ranking via hybrid scoring."""
store = ObjectStore(InMemoryBackend(), DummyEmbedder())
await store.store_object(
session_id="s1",
content="the quick brown fox jumps over the lazy dog",
object_type="file_context",
)
await store.store_object(
session_id="s1",
content="unrelated content about databases",
object_type="file_context",
)
results = await store.semantic_search("s1", "quick brown fox")
assert len(results) >= 1
# Text match should help the fox content rank higher
assert "fox" in results[0][0].content_full
async def test_semantic_search_session_isolation(self):
store = ObjectStore(InMemoryBackend(), DummyEmbedder())
await store.store_object(
session_id="s1", content="session one content", object_type="file_context"
)
await store.store_object(
session_id="s2", content="session two content", object_type="file_context"
)
results = await store.semantic_search("s1", "session one content")
assert all(r[0].session_id == "s1" for r in results)
async def test_semantic_search_respects_limit(self):
store = ObjectStore(InMemoryBackend(), DummyEmbedder())
for i in range(10):
await store.store_object(
session_id="s1",
content=f"content item {i}",
object_type="file_context",
)
results = await store.semantic_search("s1", "content item", limit=3)
assert len(results) <= 3
async def test_semantic_search_no_embedder(self):
"""Without embedder, search falls back to text-only."""
store = ObjectStore(InMemoryBackend(), embedder=None)
await store.store_object(
session_id="s1",
content="findable keyword here",
object_type="file_context",
)
results = await store.semantic_search("s1", "findable keyword")
assert len(results) == 1
assert results[0][1] == 0.3 # text-only score
async def test_semantic_search_empty_session(self):
store = ObjectStore(InMemoryBackend(), DummyEmbedder())
results = await store.semantic_search("empty", "anything")
assert results == []
# ── ObjectStore deduplication ────────────────────────────────
class TestObjectStoreDedup:
async def test_find_duplicate_exists(self):
store = ObjectStore(InMemoryBackend(), DummyEmbedder())
await store.store_object(
session_id="s1",
content="file content",
object_type="file_context",
source_key="src/auth.py",
)
dup = await store.find_duplicate("s1", "src/auth.py")
assert dup is not None
assert dup.source_key == "src/auth.py"
async def test_find_duplicate_not_found(self):
store = ObjectStore(InMemoryBackend(), DummyEmbedder())
dup = await store.find_duplicate("s1", "nonexistent.py")
assert dup is None
async def test_find_duplicate_session_scoped(self):
store = ObjectStore(InMemoryBackend(), DummyEmbedder())
await store.store_object(
session_id="s1",
content="content",
object_type="file_context",
source_key="src/auth.py",
)
dup = await store.find_duplicate("s2", "src/auth.py")
assert dup is None
# ── ObjectStore access/fault tracking ────────────────────────
class TestObjectStoreTracking:
async def test_record_access(self):
store = ObjectStore(InMemoryBackend(), DummyEmbedder())
obj = await store.store_object(
session_id="s1", content="content", object_type="file_context"
)
original_accessed = obj.last_accessed
await store.record_access(obj.id)
assert obj.access_count == 1
assert obj.last_accessed >= original_accessed
async def test_record_access_increments(self):
store = ObjectStore(InMemoryBackend(), DummyEmbedder())
obj = await store.store_object(
session_id="s1", content="content", object_type="file_context"
)
await store.record_access(obj.id)
await store.record_access(obj.id)
await store.record_access(obj.id)
assert obj.access_count == 3
async def test_record_access_nonexistent(self):
store = ObjectStore(InMemoryBackend())
# Should not raise
await store.record_access("nonexistent")
async def test_record_fault(self):
store = ObjectStore(InMemoryBackend(), DummyEmbedder())
obj = await store.store_object(
session_id="s1", content="content", object_type="file_context"
)
await store.record_fault(obj.id)
assert obj.fault_count == 1
assert obj.micro_fault_count == 0
async def test_record_micro_fault(self):
store = ObjectStore(InMemoryBackend(), DummyEmbedder())
obj = await store.store_object(
session_id="s1", content="content", object_type="file_context"
)
await store.record_fault(obj.id, is_micro=True)
assert obj.fault_count == 0
assert obj.micro_fault_count == 1
async def test_record_fault_nonexistent(self):
store = ObjectStore(InMemoryBackend())
# Should not raise
await store.record_fault("nonexistent")
async def test_mixed_faults(self):
store = ObjectStore(InMemoryBackend(), DummyEmbedder())
obj = await store.store_object(
session_id="s1", content="content", object_type="file_context"
)
await store.record_fault(obj.id)
await store.record_fault(obj.id, is_micro=True)
await store.record_fault(obj.id)
await store.record_fault(obj.id, is_micro=True)
await store.record_fault(obj.id, is_micro=True)
assert obj.fault_count == 2
assert obj.micro_fault_count == 3
# ── Integration: full lifecycle ──────────────────────────────
class TestIntegration:
async def test_store_search_degrade_cycle(self):
"""Full lifecycle: store → search → degrade → search again."""
store = ObjectStore(InMemoryBackend(), DummyEmbedder())
# Store several objects
obj1 = await store.store_object(
session_id="s1",
content="authentication middleware handles JWT validation",
object_type="file_context",
source_tool="Read",
source_key="src/auth/middleware.ts",
)
obj2 = await store.store_object(
session_id="s1",
content="database migration adds users table with email column",
object_type="file_context",
source_tool="Read",
source_key="migrations/001_users.sql",
)
# Search finds relevant object
results = await store.semantic_search("s1", "JWT validation")
assert len(results) >= 1
# Degrade first object
await store.update_fidelity(
obj1.id,
1,
summary="Auth middleware: validates JWT tokens",
losses=["exact error handling code"],
)
degraded = await store.get(obj1.id)
assert degraded is not None
assert degraded.current_fidelity == 1
assert degraded.summary_detailed is not None
# Search still works after degradation
results2 = await store.semantic_search("s1", "JWT validation")
assert len(results2) >= 1
async def test_dedup_workflow(self):
"""Dedup: check for existing → store if new."""
store = ObjectStore(InMemoryBackend(), DummyEmbedder())
# First read
dup = await store.find_duplicate("s1", "src/auth.py")
assert dup is None
obj = await store.store_object(
session_id="s1",
content="original content",
object_type="file_context",
source_key="src/auth.py",
)
# Second read — duplicate found
dup = await store.find_duplicate("s1", "src/auth.py")
assert dup is not None
assert dup.id == obj.id
async def test_multi_session_lifecycle(self):
"""Multiple sessions operate independently."""
backend = InMemoryBackend()
store = ObjectStore(backend, DummyEmbedder())
await store.store_object(
session_id="s1", content="session 1 auth code", object_type="file_context"
)
await store.store_object(
session_id="s2", content="session 2 db code", object_type="file_context"
)
s1_objs = await store.get_session_objects("s1")
s2_objs = await store.get_session_objects("s2")
assert len(s1_objs) == 1
assert len(s2_objs) == 1
# Delete session 1
count = await backend.delete_session("s1")
assert count == 1
# Session 2 unaffected
s2_objs = await store.get_session_objects("s2")
assert len(s2_objs) == 1
# Session 1 empty
s1_objs = await store.get_session_objects("s1")
assert len(s1_objs) == 0

View file

@ -0,0 +1,882 @@
"""Tests for the PostgreSQL + pgvector backend.
Unit tests use mocked asyncpg connections to verify SQL generation and
data conversion without requiring a database. Integration tests require
a running PostgreSQL instance (Docker on port 5433) and are skipped
automatically when unavailable.
"""
from __future__ import annotations
import asyncio
import json
import os
import socket
import uuid
from datetime import datetime, timezone
from typing import Any
from unittest.mock import AsyncMock, MagicMock, patch
import numpy as np
import pytest
from mnemosyne.object_store import StoredObject, _estimate_tokens
from mnemosyne.pgvector_backend import (
PgVectorBackend,
_parse_jsonb,
_parse_timestamp,
_row_to_stored_object,
)
class _MockPool:
"""A mock asyncpg pool that properly supports async context manager on acquire()."""
def __init__(self, conn: AsyncMock):
self._conn = conn
def acquire(self):
return _MockAcquire(self._conn)
class _MockAcquire:
"""Async context manager returned by pool.acquire()."""
def __init__(self, conn: AsyncMock):
self._conn = conn
async def __aenter__(self):
return self._conn
async def __aexit__(self, *args):
return False
def _make_mock_pool(conn: AsyncMock) -> _MockPool:
"""Create a mock pool with proper async context manager support."""
return _MockPool(conn)
# ── Helpers ──────────────────────────────────────────────────
def _make_stored_object(
session_id: str = "test-session",
content: str = "Test content for pgvector backend",
*,
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,
) -> StoredObject:
"""Create a StoredObject with sensible defaults for testing."""
oid = object_id or uuid.uuid4().hex
return StoredObject(
id=oid,
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}: test object",
tokens_l0=_estimate_tokens(content),
tokens_l3=_estimate_tokens(stub or f"{object_type}: test object"),
embedding=embedding or [0.1] * 384,
created_at="2025-01-01T00:00:00+00:00",
last_accessed="2025-01-01T00:00:00+00:00",
)
def _make_mock_row(
obj: StoredObject | None = None,
*,
session_external_id: str = "test-session",
similarity: float | None = None,
) -> MagicMock:
"""Create a mock asyncpg.Record from a StoredObject."""
if obj is None:
obj = _make_stored_object()
row = MagicMock()
row_data: dict[str, Any] = {
"id": uuid.UUID(obj.id) if len(obj.id) == 32 else uuid.uuid4(),
"session_id": uuid.uuid4(),
"session_external_id": session_external_id,
"object_type": obj.object_type,
"source_tool": obj.source_tool,
"source_key": obj.source_key,
"content_full": obj.content_full,
"summary_detailed": obj.summary_detailed,
"summary_compact": obj.summary_compact,
"stub": obj.stub,
"losses_l1": obj.losses_l1,
"losses_l2": obj.losses_l2,
"can_answer_l1": obj.can_answer_l1,
"can_answer_l2": obj.can_answer_l2,
"fault_when": obj.fault_when,
"key_entities": obj.key_entities,
"tags": obj.tags,
"current_fidelity": obj.current_fidelity,
"pinned": obj.pinned,
"tokens_l0": obj.tokens_l0,
"tokens_l1": obj.tokens_l1,
"tokens_l2": obj.tokens_l2,
"tokens_l3": obj.tokens_l3,
"source_turn_start": obj.source_turn_start,
"source_turn_end": obj.source_turn_end,
"embedding": np.array(obj.embedding, dtype=np.float32) if obj.embedding else None,
"created_at": datetime(2025, 1, 1, tzinfo=timezone.utc),
"last_accessed": datetime(2025, 1, 1, tzinfo=timezone.utc),
"access_count": obj.access_count,
"fault_count": obj.fault_count,
"micro_fault_count": obj.micro_fault_count,
}
if similarity is not None:
row_data["similarity"] = similarity
row.__getitem__ = lambda self, key: row_data[key]
return row
def _pg_available() -> bool:
"""Check if PostgreSQL is reachable on port 5433."""
try:
with socket.create_connection(("localhost", 5433), timeout=1):
return True
except (OSError, ConnectionRefusedError):
return False
# ── Unit Tests: Data Conversion ──────────────────────────────
class TestParseJsonb:
def test_none_returns_empty_list(self):
assert _parse_jsonb(None) == []
def test_list_passthrough(self):
assert _parse_jsonb(["a", "b", "c"]) == ["a", "b", "c"]
def test_list_converts_to_strings(self):
assert _parse_jsonb([1, 2, 3]) == ["1", "2", "3"]
def test_json_string(self):
assert _parse_jsonb('["x", "y"]') == ["x", "y"]
def test_invalid_json_string(self):
assert _parse_jsonb("not json") == []
def test_empty_list(self):
assert _parse_jsonb([]) == []
class TestParseTimestamp:
def test_iso_format(self):
dt = _parse_timestamp("2025-01-01T00:00:00+00:00")
assert dt.year == 2025
assert dt.tzinfo is not None
def test_empty_string_returns_now(self):
dt = _parse_timestamp("")
assert dt.tzinfo is not None
# Should be close to now
diff = abs((datetime.now(timezone.utc) - dt).total_seconds())
assert diff < 5
def test_naive_timestamp_gets_utc(self):
dt = _parse_timestamp("2025-06-15T12:00:00")
assert dt.tzinfo == timezone.utc
def test_invalid_returns_now(self):
dt = _parse_timestamp("not-a-date")
assert dt.tzinfo is not None
class TestRowToStoredObject:
def test_basic_conversion(self):
obj = _make_stored_object()
row = _make_mock_row(obj)
result = _row_to_stored_object(row)
assert result.object_type == "file_context"
assert result.content_full == obj.content_full
assert result.stub == obj.stub
assert result.session_id == "test-session"
assert result.current_fidelity == 0
assert result.pinned is False
def test_embedding_conversion(self):
obj = _make_stored_object(embedding=[0.5] * 384)
row = _make_mock_row(obj)
result = _row_to_stored_object(row)
assert len(result.embedding) == 384
assert abs(result.embedding[0] - 0.5) < 1e-6
def test_jsonb_fields_parsed(self):
obj = _make_stored_object()
obj.losses_l1 = ["detail_a", "detail_b"]
obj.key_entities = ["src/main.py", "Config"]
row = _make_mock_row(obj)
result = _row_to_stored_object(row)
assert result.losses_l1 == ["detail_a", "detail_b"]
assert result.key_entities == ["src/main.py", "Config"]
def test_tags_conversion(self):
obj = _make_stored_object()
obj.tags = ["auth", "middleware"]
row = _make_mock_row(obj)
result = _row_to_stored_object(row)
assert result.tags == ["auth", "middleware"]
def test_null_embedding(self):
obj = _make_stored_object(embedding=[])
row_data: dict[str, Any] = {
"id": uuid.uuid4(),
"session_id": uuid.uuid4(),
"session_external_id": "test-session",
"object_type": obj.object_type,
"source_tool": obj.source_tool,
"source_key": obj.source_key,
"content_full": obj.content_full,
"summary_detailed": obj.summary_detailed,
"summary_compact": obj.summary_compact,
"stub": obj.stub,
"losses_l1": [],
"losses_l2": [],
"can_answer_l1": [],
"can_answer_l2": [],
"fault_when": [],
"key_entities": [],
"tags": [],
"current_fidelity": 0,
"pinned": False,
"tokens_l0": obj.tokens_l0,
"tokens_l1": None,
"tokens_l2": None,
"tokens_l3": obj.tokens_l3,
"source_turn_start": None,
"source_turn_end": None,
"embedding": None,
"created_at": datetime(2025, 1, 1, tzinfo=timezone.utc),
"last_accessed": datetime(2025, 1, 1, tzinfo=timezone.utc),
"access_count": 0,
"fault_count": 0,
"micro_fault_count": 0,
}
row = MagicMock()
row.__getitem__ = lambda self, key: row_data[key]
result = _row_to_stored_object(row)
assert result.embedding == []
# ── Unit Tests: Backend Methods (Mocked DB) ─────────────────
class TestPgVectorBackendInit:
def test_default_config(self):
backend = PgVectorBackend()
assert backend._host == "localhost"
assert backend._port == 5433
assert backend._database == "mnemosyne"
assert backend._pool is None
def test_custom_config(self):
backend = PgVectorBackend(
host="db.example.com",
port=5432,
database="mydb",
user="myuser",
password="secret",
min_connections=5,
max_connections=20,
)
assert backend._host == "db.example.com"
assert backend._port == 5432
assert backend._min_connections == 5
assert backend._max_connections == 20
def test_get_pool_raises_when_not_connected(self):
backend = PgVectorBackend()
with pytest.raises(RuntimeError, match="not connected"):
backend._get_pool()
class TestPgVectorBackendStore:
"""Test store() with mocked pool."""
async def test_store_calls_execute(self):
backend = PgVectorBackend()
mock_conn = AsyncMock()
mock_conn.execute = AsyncMock()
mock_conn.fetchrow = AsyncMock(return_value={"id": uuid.uuid4()})
backend._pool = _make_mock_pool(mock_conn)
obj = _make_stored_object(object_id=uuid.uuid4().hex)
await backend.store(obj)
# Should have called execute for the INSERT (store) and fetchrow for session upsert
assert mock_conn.execute.called or mock_conn.fetchrow.called
async def test_store_creates_session_if_needed(self):
backend = PgVectorBackend()
session_uuid = uuid.uuid4()
mock_conn = AsyncMock()
mock_conn.fetchrow = AsyncMock(return_value={"id": session_uuid})
mock_conn.execute = AsyncMock()
backend._pool = _make_mock_pool(mock_conn)
obj = _make_stored_object(object_id=uuid.uuid4().hex)
await backend.store(obj)
# Session should be cached after creation
assert "test-session" in backend._session_cache
class TestPgVectorBackendGet:
"""Test get() with mocked pool."""
async def test_get_returns_none_when_not_found(self):
backend = PgVectorBackend()
mock_conn = AsyncMock()
mock_conn.fetchrow = AsyncMock(return_value=None)
backend._pool = _make_mock_pool(mock_conn)
result = await backend.get(uuid.uuid4().hex)
assert result is None
async def test_get_returns_stored_object(self):
backend = PgVectorBackend()
obj = _make_stored_object(object_id=uuid.uuid4().hex)
mock_row = _make_mock_row(obj)
mock_conn = AsyncMock()
mock_conn.fetchrow = AsyncMock(return_value=mock_row)
backend._pool = _make_mock_pool(mock_conn)
result = await backend.get(obj.id)
assert result is not None
assert result.content_full == obj.content_full
class TestPgVectorBackendGetBySession:
"""Test get_by_session() with mocked pool."""
async def test_returns_empty_for_unknown_session(self):
backend = PgVectorBackend()
mock_conn = AsyncMock()
mock_conn.fetchrow = AsyncMock(return_value=None)
backend._pool = _make_mock_pool(mock_conn)
result = await backend.get_by_session("nonexistent")
assert result == []
class TestPgVectorBackendUpdateFidelity:
"""Test update_fidelity() with mocked pool."""
async def test_update_fidelity_basic(self):
backend = PgVectorBackend()
mock_conn = AsyncMock()
mock_conn.execute = AsyncMock()
backend._pool = _make_mock_pool(mock_conn)
await backend.update_fidelity(uuid.uuid4().hex, 1)
assert mock_conn.execute.called
async def test_update_fidelity_with_summary_l1(self):
backend = PgVectorBackend()
mock_conn = AsyncMock()
mock_conn.execute = AsyncMock()
backend._pool = _make_mock_pool(mock_conn)
await backend.update_fidelity(
uuid.uuid4().hex, 1, summary="A summary", losses=["lost detail"]
)
# Verify the SQL includes summary_detailed and losses_l1
call_args = mock_conn.execute.call_args
query = call_args[0][0]
assert "summary_detailed" in query
assert "losses_l1" in query
async def test_update_fidelity_with_summary_l2(self):
backend = PgVectorBackend()
mock_conn = AsyncMock()
mock_conn.execute = AsyncMock()
backend._pool = _make_mock_pool(mock_conn)
await backend.update_fidelity(uuid.uuid4().hex, 2, summary="Compact", losses=["more lost"])
call_args = mock_conn.execute.call_args
query = call_args[0][0]
assert "summary_compact" in query
assert "losses_l2" in query
class TestPgVectorBackendSearchByEmbedding:
"""Test search_by_embedding() with mocked pool."""
async def test_returns_empty_for_unknown_session(self):
backend = PgVectorBackend()
mock_conn = AsyncMock()
mock_conn.fetchrow = AsyncMock(return_value=None)
backend._pool = _make_mock_pool(mock_conn)
result = await backend.search_by_embedding("nonexistent", [0.1] * 384)
assert result == []
class TestPgVectorBackendSearchByText:
"""Test search_by_text() with mocked pool."""
async def test_returns_empty_for_unknown_session(self):
backend = PgVectorBackend()
mock_conn = AsyncMock()
mock_conn.fetchrow = AsyncMock(return_value=None)
backend._pool = _make_mock_pool(mock_conn)
result = await backend.search_by_text("nonexistent", "test query")
assert result == []
class TestPgVectorBackendDeleteSession:
"""Test delete_session() with mocked pool."""
async def test_returns_zero_for_unknown_session(self):
backend = PgVectorBackend()
mock_conn = AsyncMock()
mock_conn.fetchrow = AsyncMock(return_value=None)
backend._pool = _make_mock_pool(mock_conn)
result = await backend.delete_session("nonexistent")
assert result == 0
class TestPgVectorBackendGetBySourceKey:
"""Test get_by_source_key() with mocked pool."""
async def test_returns_none_for_unknown_session(self):
backend = PgVectorBackend()
mock_conn = AsyncMock()
mock_conn.fetchrow = AsyncMock(return_value=None)
backend._pool = _make_mock_pool(mock_conn)
result = await backend.get_by_source_key("nonexistent", "src/main.py")
assert result is None
class TestPgVectorBackendHealthCheck:
"""Test health_check()."""
async def test_returns_false_when_not_connected(self):
backend = PgVectorBackend()
assert await backend.health_check() is False
async def test_returns_true_with_healthy_pool(self):
backend = PgVectorBackend()
mock_conn = AsyncMock()
mock_conn.fetchval = AsyncMock(return_value=1)
backend._pool = _make_mock_pool(mock_conn)
assert await backend.health_check() is True
async def test_returns_false_on_error(self):
backend = PgVectorBackend()
mock_conn = AsyncMock()
mock_conn.fetchval = AsyncMock(side_effect=Exception("connection lost"))
backend._pool = _make_mock_pool(mock_conn)
assert await backend.health_check() is False
class TestPgVectorBackendSessionCache:
"""Test session ID caching behavior."""
async def test_session_cache_populated_on_ensure(self):
backend = PgVectorBackend()
session_uuid = uuid.uuid4()
mock_conn = AsyncMock()
mock_conn.fetchrow = AsyncMock(return_value={"id": session_uuid})
backend._pool = _make_mock_pool(mock_conn)
result = await backend.ensure_session("my-session", "claude-opus-4")
assert result == session_uuid
assert backend._session_cache["my-session"] == session_uuid
async def test_session_cache_avoids_db_lookup(self):
backend = PgVectorBackend()
session_uuid = uuid.uuid4()
backend._session_cache["cached-session"] = session_uuid
# Pool shouldn't be needed since cache is populated
backend._pool = _make_mock_pool(AsyncMock())
result = await backend._resolve_session_id("cached-session")
assert result == session_uuid
async def test_close_clears_cache(self):
backend = PgVectorBackend()
backend._session_cache["test"] = uuid.uuid4()
mock_pool = AsyncMock()
mock_pool.close = AsyncMock()
backend._pool = mock_pool
await backend.close()
assert len(backend._session_cache) == 0
assert backend._pool is None
# ── Integration Tests (require Docker PostgreSQL) ────────────
_skip_no_pg = pytest.mark.skipif(
not _pg_available(),
reason="PostgreSQL not available on localhost:5433 (run: docker compose up -d)",
)
@_skip_no_pg
class TestPgVectorBackendIntegration:
"""Integration tests against a real PostgreSQL instance.
These tests require Docker PostgreSQL running on port 5433 with
the schema from sql/init.sql applied. Start with:
cd ~/Projects/contextmanager && docker compose up -d
"""
@pytest.fixture
async def backend(self):
"""Create a connected backend and clean up after test."""
b = PgVectorBackend(
host="localhost",
port=5433,
database="mnemosyne",
user="mnemosyne",
password="mnemosyne_dev",
min_connections=1,
max_connections=3,
)
await b.connect()
# Create a unique session for this test
test_session = f"integration-test-{uuid.uuid4().hex[:8]}"
yield b, test_session
# Cleanup: delete test session data
try:
session_uuid = await b._resolve_session_id(test_session)
if session_uuid is not None:
pool = b._get_pool()
async with pool.acquire() as conn:
await conn.execute(
"DELETE FROM semantic_objects WHERE session_id = $1",
session_uuid,
)
await conn.execute(
"DELETE FROM sessions WHERE id = $1",
session_uuid,
)
except Exception:
pass
await b.close()
async def test_health_check(self, backend):
b, _ = backend
assert await b.health_check() is True
async def test_ensure_session(self, backend):
b, test_session = backend
session_uuid = await b.ensure_session(test_session, "claude-opus-4")
assert isinstance(session_uuid, uuid.UUID)
# Second call should return same UUID
session_uuid2 = await b.ensure_session(test_session, "claude-opus-4")
assert session_uuid == session_uuid2
async def test_store_and_get(self, backend):
b, test_session = backend
await b.ensure_session(test_session)
obj_id = uuid.uuid4().hex
obj = _make_stored_object(
session_id=test_session,
content="Integration test content",
object_id=obj_id,
)
await b.store(obj)
retrieved = await b.get(obj_id)
assert retrieved is not None
assert retrieved.content_full == "Integration test content"
assert retrieved.session_id == test_session
assert retrieved.object_type == "file_context"
async def test_store_upsert(self, backend):
b, test_session = backend
await b.ensure_session(test_session)
obj_id = uuid.uuid4().hex
obj = _make_stored_object(
session_id=test_session,
content="Original content",
object_id=obj_id,
)
await b.store(obj)
# Update the same object
obj.content_full = "Updated content"
await b.store(obj)
retrieved = await b.get(obj_id)
assert retrieved is not None
assert retrieved.content_full == "Updated content"
async def test_get_by_session(self, backend):
b, test_session = backend
await b.ensure_session(test_session)
# Store 3 objects
for i in range(3):
obj = _make_stored_object(
session_id=test_session,
content=f"Content {i}",
object_id=uuid.uuid4().hex,
)
await b.store(obj)
results = await b.get_by_session(test_session)
assert len(results) == 3
async def test_get_by_session_fidelity_filter(self, backend):
b, test_session = backend
await b.ensure_session(test_session)
# Store object at fidelity 0
obj0 = _make_stored_object(
session_id=test_session,
content="Fidelity 0",
object_id=uuid.uuid4().hex,
)
await b.store(obj0)
# Store object at fidelity 4 (evicted)
obj4 = _make_stored_object(
session_id=test_session,
content="Fidelity 4",
object_id=uuid.uuid4().hex,
)
obj4.current_fidelity = 4
await b.store(obj4)
# Default: include all
all_results = await b.get_by_session(test_session, fidelity_max=4)
assert len(all_results) == 2
# Exclude evicted
active_results = await b.get_by_session(test_session, fidelity_max=3)
assert len(active_results) == 1
assert active_results[0].content_full == "Fidelity 0"
async def test_update_fidelity(self, backend):
b, test_session = backend
await b.ensure_session(test_session)
obj_id = uuid.uuid4().hex
obj = _make_stored_object(
session_id=test_session,
content="Will be degraded",
object_id=obj_id,
)
await b.store(obj)
await b.update_fidelity(
obj_id,
1,
summary="Detailed summary",
losses=["lost some detail"],
)
retrieved = await b.get(obj_id)
assert retrieved is not None
assert retrieved.current_fidelity == 1
assert retrieved.summary_detailed == "Detailed summary"
async def test_search_by_embedding(self, backend):
b, test_session = backend
await b.ensure_session(test_session)
# Store objects with different embeddings
rng = np.random.default_rng(42)
for i in range(5):
vec = rng.standard_normal(384).astype(np.float32)
vec = vec / np.linalg.norm(vec)
obj = _make_stored_object(
session_id=test_session,
content=f"Embedding test {i}",
object_id=uuid.uuid4().hex,
embedding=vec.tolist(),
)
await b.store(obj)
# Search with the first object's embedding
query_vec = rng.standard_normal(384).astype(np.float32)
query_vec = query_vec / np.linalg.norm(query_vec)
results = await b.search_by_embedding(test_session, query_vec.tolist(), limit=3)
assert len(results) <= 3
# Results should be (StoredObject, float) tuples
for obj, score in results:
assert isinstance(obj, StoredObject)
assert isinstance(score, float)
assert -1.0 <= score <= 1.0
# Scores should be in descending order
scores = [s for _, s in results]
assert scores == sorted(scores, reverse=True)
async def test_search_by_text(self, backend):
b, test_session = backend
await b.ensure_session(test_session)
obj = _make_stored_object(
session_id=test_session,
content="The authentication middleware validates JWT tokens and checks expiration dates",
object_id=uuid.uuid4().hex,
)
await b.store(obj)
obj2 = _make_stored_object(
session_id=test_session,
content="Database connection pooling configuration for PostgreSQL",
object_id=uuid.uuid4().hex,
)
await b.store(obj2)
# Search for auth-related content
results = await b.search_by_text(test_session, "authentication JWT tokens")
assert len(results) >= 1
assert any("authentication" in r.content_full for r in results)
async def test_delete_session(self, backend):
b, test_session = backend
await b.ensure_session(test_session)
# Store some objects
for i in range(3):
obj = _make_stored_object(
session_id=test_session,
content=f"Delete test {i}",
object_id=uuid.uuid4().hex,
)
await b.store(obj)
count = await b.delete_session(test_session)
assert count == 3
# Verify they're gone
results = await b.get_by_session(test_session)
assert len(results) == 0
async def test_get_by_source_key(self, backend):
b, test_session = backend
await b.ensure_session(test_session)
obj = _make_stored_object(
session_id=test_session,
content="File content",
source_key="src/main.py",
object_id=uuid.uuid4().hex,
)
await b.store(obj)
result = await b.get_by_source_key(test_session, "src/main.py")
assert result is not None
assert result.source_key == "src/main.py"
# Non-existent key
result2 = await b.get_by_source_key(test_session, "nonexistent.py")
assert result2 is None
async def test_get_by_source_key_returns_most_recent(self, backend):
b, test_session = backend
await b.ensure_session(test_session)
# Store two objects with same source_key
obj1 = _make_stored_object(
session_id=test_session,
content="Old version",
source_key="src/config.py",
object_id=uuid.uuid4().hex,
)
obj1.created_at = "2025-01-01T00:00:00+00:00"
await b.store(obj1)
obj2 = _make_stored_object(
session_id=test_session,
content="New version",
source_key="src/config.py",
object_id=uuid.uuid4().hex,
)
obj2.created_at = "2025-06-01T00:00:00+00:00"
await b.store(obj2)
result = await b.get_by_source_key(test_session, "src/config.py")
assert result is not None
assert result.content_full == "New version"
async def test_session_isolation(self, backend):
b, test_session = backend
other_session = f"other-{uuid.uuid4().hex[:8]}"
await b.ensure_session(test_session)
await b.ensure_session(other_session)
# Store in test_session
obj = _make_stored_object(
session_id=test_session,
content="Session A content",
object_id=uuid.uuid4().hex,
)
await b.store(obj)
# Store in other_session
obj2 = _make_stored_object(
session_id=other_session,
content="Session B content",
object_id=uuid.uuid4().hex,
)
await b.store(obj2)
# Each session should only see its own objects
results_a = await b.get_by_session(test_session)
results_b = await b.get_by_session(other_session)
assert len(results_a) == 1
assert results_a[0].content_full == "Session A content"
assert len(results_b) == 1
assert results_b[0].content_full == "Session B content"
# Cleanup other session
try:
session_uuid = await b._resolve_session_id(other_session)
if session_uuid:
pool = b._get_pool()
async with pool.acquire() as conn:
await conn.execute(
"DELETE FROM semantic_objects WHERE session_id = $1",
session_uuid,
)
await conn.execute(
"DELETE FROM sessions WHERE id = $1",
session_uuid,
)
except Exception:
pass

980
tests/test_segmenter.py Normal file
View file

@ -0,0 +1,980 @@
"""Tests for the semantic object segmenter.
Covers: tool result segmentation, text classification, user message handling,
entity extraction, stub/tag generation, merging, incremental segmentation,
realistic multi-turn payloads, and edge cases.
"""
from __future__ import annotations
import pytest
from mnemosyne.segmenter import (
VALID_OBJECT_TYPES,
Segmenter,
SegmentedObject,
_estimate_tokens,
)
# ---------------------------------------------------------------------------
# Fixtures
# ---------------------------------------------------------------------------
@pytest.fixture
def seg() -> Segmenter:
"""Default segmenter with standard thresholds."""
return Segmenter()
@pytest.fixture
def small_seg() -> Segmenter:
"""Segmenter with low min_object_tokens for testing merging."""
return Segmenter(min_object_tokens=10, max_object_tokens=500)
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _msg(role: str, content: str | list) -> dict:
"""Shorthand for creating a message dict."""
return {"role": role, "content": content}
def _tool_use_block(tool_id: str, name: str, input_: dict | None = None) -> dict:
return {
"type": "tool_use",
"id": tool_id,
"name": name,
"input": input_ or {},
}
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 _long_text(base: str = "x", tokens: int = 200) -> str:
"""Generate text of approximately `tokens` estimated tokens."""
return base * (tokens * 4)
# ===========================================================================
# 1. Tool result segmentation
# ===========================================================================
class TestToolResultSegmentation:
"""Test that tool results are classified by tool name."""
def test_read_tool_produces_file_context(self, seg: Segmenter):
messages = [
_msg(
"assistant",
[
_text_block("Let me read the file."),
_tool_use_block("t1", "Read", {"file_path": "src/main.py"}),
],
),
_msg(
"user",
[
_tool_result_block("t1", _long_text("def main():\n pass\n")),
],
),
]
objects = seg.segment_messages(messages)
file_objs = [o for o in objects if o.object_type == "file_context"]
assert len(file_objs) >= 1
assert file_objs[0].source_tool == "Read"
assert file_objs[0].source_key == "src/main.py"
def test_bash_tool_produces_tool_result(self, seg: Segmenter):
messages = [
_msg(
"assistant",
[
_text_block("Running tests."),
_tool_use_block("t1", "Bash", {"command": "pytest"}),
],
),
_msg(
"user",
[
_tool_result_block("t1", _long_text("All 42 tests passed.\n")),
],
),
]
objects = seg.segment_messages(messages)
tool_objs = [o for o in objects if o.object_type == "tool_result"]
assert len(tool_objs) >= 1
assert tool_objs[0].source_tool == "Bash"
def test_grep_tool_produces_tool_result(self, seg: Segmenter):
messages = [
_msg(
"assistant",
[
_text_block("Searching for usage."),
_tool_use_block("t1", "Grep", {"pattern": "handleAuth"}),
],
),
_msg(
"user",
[
_tool_result_block("t1", _long_text("src/auth.ts:15: handleAuth()\n")),
],
),
]
objects = seg.segment_messages(messages)
tool_objs = [o for o in objects if o.object_type == "tool_result"]
assert len(tool_objs) >= 1
assert tool_objs[0].source_tool == "Grep"
def test_write_tool_produces_tool_result(self, seg: Segmenter):
messages = [
_msg(
"assistant",
[
_text_block("Writing the file."),
_tool_use_block("t1", "Write", {"file_path": "out.py"}),
],
),
_msg(
"user",
[
_tool_result_block("t1", _long_text("File written successfully.\n")),
],
),
]
objects = seg.segment_messages(messages)
tool_objs = [o for o in objects if o.object_type == "tool_result"]
assert len(tool_objs) >= 1
def test_unknown_tool_produces_tool_result(self, seg: Segmenter):
messages = [
_msg(
"assistant",
[
_text_block("Using custom tool."),
_tool_use_block("t1", "CustomTool", {}),
],
),
_msg(
"user",
[
_tool_result_block("t1", _long_text("Custom output here.\n")),
],
),
]
objects = seg.segment_messages(messages)
tool_objs = [o for o in objects if o.object_type == "tool_result"]
assert len(tool_objs) >= 1
assert tool_objs[0].source_tool == "CustomTool"
def test_web_fetch_produces_external_reference(self, seg: Segmenter):
messages = [
_msg(
"assistant",
[
_text_block("Fetching docs."),
_tool_use_block("t1", "WebFetch", {"url": "https://docs.example.com"}),
],
),
_msg(
"user",
[
_tool_result_block("t1", _long_text("# API Documentation\n")),
],
),
]
objects = seg.segment_messages(messages)
ext_objs = [o for o in objects if o.object_type == "external_reference"]
assert len(ext_objs) >= 1
assert ext_objs[0].source_key == "https://docs.example.com"
def test_read_with_filePath_key(self, seg: Segmenter):
"""Test that filePath (camelCase) is also recognized."""
messages = [
_msg(
"assistant",
[
_tool_use_block("t1", "Read", {"filePath": "/etc/config.toml"}),
],
),
_msg(
"user",
[
_tool_result_block("t1", _long_text("key = 'value'\n")),
],
),
]
objects = seg.segment_messages(messages)
file_objs = [o for o in objects if o.object_type == "file_context"]
assert len(file_objs) >= 1
assert file_objs[0].source_key == "/etc/config.toml"
# ===========================================================================
# 2. Text classification
# ===========================================================================
class TestTextClassification:
"""Test assistant text classification into object types."""
def test_error_text_classified(self, seg: Segmenter):
error_text = _long_text(
"Traceback (most recent call last):\n"
' File "src/main.py", line 42, in run\n'
" raise ValueError('bad input')\n"
"ValueError: bad input\n"
)
messages = [_msg("assistant", error_text)]
objects = seg.segment_messages(messages)
assert any(o.object_type == "error_context" for o in objects)
def test_plan_text_classified(self, seg: Segmenter):
plan_text = _long_text(
"Here's the implementation plan:\n"
"1. Create the database schema\n"
"2. Implement the API endpoints\n"
"3. Write integration tests\n"
"4. Deploy to staging\n"
)
messages = [_msg("assistant", plan_text)]
objects = seg.segment_messages(messages)
assert any(o.object_type == "plan" for o in objects)
def test_decision_text_classified(self, seg: Segmenter):
decision_text = _long_text(
"I decided to use JWT tokens for authentication because "
"they're stateless and work well with our microservice architecture. "
"I chose JWT over session cookies because we need cross-domain support.\n"
)
messages = [_msg("assistant", decision_text)]
objects = seg.segment_messages(messages)
assert any(o.object_type == "design_decision" for o in objects)
def test_debug_text_classified(self, seg: Segmenter):
debug_text = _long_text(
"I'm investigating the race condition. After debugging, "
"I found the root cause: the mutex wasn't being held during "
"the token refresh. Fixed by adding a lock around the critical section.\n"
)
messages = [_msg("assistant", debug_text)]
objects = seg.segment_messages(messages)
assert any(o.object_type == "debugging_session" for o in objects)
def test_default_conversation_phase(self, seg: Segmenter):
text = _long_text(
"Sure, I'll help you with that. Let me take a look at the code "
"and see what we can improve here.\n"
)
messages = [_msg("assistant", text)]
objects = seg.segment_messages(messages)
assert any(o.object_type == "conversation_phase" for o in objects)
def test_error_needs_multiple_signals(self, seg: Segmenter):
"""A single 'error' word shouldn't trigger error_context."""
text = _long_text("There might be an error somewhere in the logic.\n")
messages = [_msg("assistant", text)]
objects = seg.segment_messages(messages)
# Should NOT be error_context with just one weak signal
assert all(o.object_type != "error_context" for o in objects)
# ===========================================================================
# 3. User message handling
# ===========================================================================
class TestUserMessageHandling:
"""Test user message boundary and merge behavior."""
def test_short_user_message_merges(self, seg: Segmenter):
"""Short messages like 'ok' should merge with next segment."""
messages = [
_msg("user", "ok"),
_msg("assistant", _long_text("I'll proceed with the implementation.\n")),
]
objects = seg.segment_messages(messages)
# Should produce a single merged object, not two
assert len(objects) <= 2
def test_normal_user_message_creates_boundary(self, seg: Segmenter):
long_user = _long_text(
"Can you refactor the authentication module to use OAuth2 "
"instead of the current basic auth approach?\n"
)
messages = [
_msg("user", long_user),
_msg("assistant", _long_text("I'll refactor the auth module.\n")),
]
objects = seg.segment_messages(messages)
assert len(objects) >= 1
def test_short_messages_recognized(self, seg: Segmenter):
"""Various short messages should be recognized."""
short_msgs = ["ok", "yes", "continue", "go ahead", "thanks", "lgtm", "done"]
for short in short_msgs:
messages = [
_msg("user", short),
_msg("assistant", _long_text("Continuing...\n")),
]
objects = seg.segment_messages(messages)
# Should merge — at most 1 object
assert len(objects) <= 2, f"'{short}' was not merged"
def test_user_message_with_tool_results(self, seg: Segmenter):
"""User messages can contain both text and tool_result blocks."""
messages = [
_msg(
"assistant",
[
_text_block("Let me read the file."),
_tool_use_block("t1", "Read", {"file_path": "src/app.py"}),
],
),
_msg(
"user",
[
_tool_result_block(
"t1", _long_text("import flask\napp = flask.Flask(__name__)\n")
),
],
),
]
objects = seg.segment_messages(messages)
assert any(o.object_type == "file_context" for o in objects)
# ===========================================================================
# 4. Entity extraction
# ===========================================================================
class TestEntityExtraction:
"""Test extraction of file paths, function names, and packages."""
def test_file_paths_extracted(self, seg: Segmenter):
text = _long_text(
"I read src/auth/middleware.ts and tests/auth.test.ts. "
"The main logic is in src/core/handler.py.\n"
)
messages = [_msg("assistant", text)]
objects = seg.segment_messages(messages)
entities = objects[0].key_entities
assert any("middleware.ts" in e for e in entities)
assert any("handler.py" in e for e in entities)
def test_function_names_extracted(self, seg: Segmenter):
text = _long_text("def handleAuth(request):\n return authenticate(request.token)\n")
messages = [_msg("assistant", text)]
objects = seg.segment_messages(messages)
entities = objects[0].key_entities
assert any("handleAuth" in e for e in entities)
def test_import_names_extracted(self, seg: Segmenter):
text = _long_text("import flask\nfrom sqlalchemy import Column\nimport os\n")
messages = [_msg("assistant", text)]
objects = seg.segment_messages(messages)
entities = objects[0].key_entities
assert any("flask" in e for e in entities)
assert any("sqlalchemy" in e for e in entities)
def test_js_require_extracted(self, seg: Segmenter):
text = _long_text(
"const express = require('express');\nconst jwt = require('jsonwebtoken');\n"
)
messages = [_msg("assistant", text)]
objects = seg.segment_messages(messages)
entities = objects[0].key_entities
assert any("express" in e for e in entities)
assert any("jsonwebtoken" in e for e in entities)
def test_entity_cap_at_20(self, seg: Segmenter):
"""Entities should be capped at 20."""
# Generate text with many unique file paths
paths = [f"src/module{i}/file{i}.py" for i in range(30)]
text = _long_text(" ".join(paths) + "\n")
messages = [_msg("assistant", text)]
objects = seg.segment_messages(messages)
assert len(objects[0].key_entities) <= 20
# ===========================================================================
# 5. Stub generation
# ===========================================================================
class TestStubGeneration:
"""Test auto-generated stubs."""
def test_stub_format(self, seg: Segmenter):
text = _long_text("This is the first line of content.\nSecond line here.\n")
messages = [_msg("assistant", text)]
objects = seg.segment_messages(messages)
stub = objects[0].stub
assert stub.startswith("[")
assert stub.endswith("]")
assert ":" in stub
def test_stub_contains_type(self, seg: Segmenter):
text = _long_text("Some conversation content here.\n")
messages = [_msg("assistant", text)]
objects = seg.segment_messages(messages)
assert objects[0].object_type in objects[0].stub
def test_stub_truncates_long_first_line(self, seg: Segmenter):
long_line = "A" * 200 + "\nSecond line."
text = _long_text(long_line)
messages = [_msg("assistant", text)]
objects = seg.segment_messages(messages)
stub = objects[0].stub
# Type prefix + ": " + content, content part should be <= 100 chars
content_part = stub.split(": ", 1)[1].rstrip("]")
assert len(content_part) <= 103 # 100 + "..."
def test_stub_uses_first_line_only(self, seg: Segmenter):
text = _long_text("First line here.\nSecond line should not appear.\n")
messages = [_msg("assistant", text)]
objects = seg.segment_messages(messages)
assert "Second line" not in objects[0].stub
# ===========================================================================
# 6. Tag generation
# ===========================================================================
class TestTagGeneration:
"""Test auto-generated tags."""
def test_tags_include_object_type(self, seg: Segmenter):
text = _long_text("Some content.\n")
messages = [_msg("assistant", text)]
objects = seg.segment_messages(messages)
assert objects[0].object_type in objects[0].tags
def test_tags_include_source_tool(self, seg: Segmenter):
messages = [
_msg(
"assistant",
[
_tool_use_block("t1", "Read", {"file_path": "x.py"}),
],
),
_msg(
"user",
[
_tool_result_block("t1", _long_text("content of x.py\n")),
],
),
]
objects = seg.segment_messages(messages)
file_objs = [o for o in objects if o.source_tool == "Read"]
assert len(file_objs) >= 1
assert "read" in file_objs[0].tags
def test_tags_include_file_extensions(self, seg: Segmenter):
text = _long_text("Modified src/auth.ts and src/handler.py to fix the issue.\n")
messages = [_msg("assistant", text)]
objects = seg.segment_messages(messages)
tags = objects[0].tags
assert ".ts" in tags
assert ".py" in tags
def test_tags_no_duplicate_extensions(self, seg: Segmenter):
text = _long_text("Read src/a.py and src/b.py and src/c.py.\n")
messages = [_msg("assistant", text)]
objects = seg.segment_messages(messages)
py_count = sum(1 for t in objects[0].tags if t == ".py")
assert py_count == 1
# ===========================================================================
# 7. Minimum object size merging
# ===========================================================================
class TestMinObjectSizeMerging:
"""Test that undersized objects get merged."""
def test_small_segments_merge(self, seg: Segmenter):
"""Very small adjacent segments should merge."""
messages = [
_msg("assistant", "Hi."),
_msg("user", "Hello."),
_msg("assistant", "How can I help?"),
]
objects = seg.segment_messages(messages)
# These are all tiny — should merge into fewer objects
assert len(objects) <= 2
def test_large_segments_stay_separate(self, seg: Segmenter):
"""Segments above min_object_tokens stay separate."""
messages = [
_msg("assistant", _long_text("First large block of content.\n")),
_msg("user", _long_text("Second large block of content.\n")),
]
objects = seg.segment_messages(messages)
assert len(objects) >= 1 # At least one object
def test_incompatible_types_dont_merge(self, small_seg: Segmenter):
"""file_context and plan shouldn't merge even if small."""
messages = [
_msg(
"assistant",
[
_tool_use_block("t1", "Read", {"file_path": "a.py"}),
],
),
_msg(
"user",
[
_tool_result_block("t1", "x = 1"),
],
),
_msg(
"assistant",
[
_tool_use_block("t2", "Read", {"file_path": "b.py"}),
],
),
_msg(
"user",
[
_tool_result_block("t2", "y = 2"),
],
),
]
objects = small_seg.segment_messages(messages)
file_objs = [o for o in objects if o.object_type == "file_context"]
# Even if small, file_context objects from different files should
# remain separate (they have different source_keys)
# But they might merge if compatible — the key thing is they exist
assert len(file_objs) >= 1
# ===========================================================================
# 8. Incremental segmentation
# ===========================================================================
class TestIncrementalSegmentation:
"""Test segment_incremental for extending existing objects."""
def test_incremental_appends_new(self, seg: Segmenter):
existing = [
SegmentedObject(
content=_long_text("Previous conversation.\n"),
object_type="conversation_phase",
turn_start=0,
turn_end=1,
token_estimate=200,
)
]
new_messages = [
_msg(
"assistant",
[
_tool_use_block("t1", "Read", {"file_path": "new.py"}),
],
),
_msg(
"user",
[
_tool_result_block("t1", _long_text("new file content\n")),
],
),
]
result = seg.segment_incremental(new_messages, existing, start_turn=2)
assert len(result) > len(existing)
assert result[0] is existing[0] # First object unchanged
def test_incremental_merges_compatible(self, seg: Segmenter):
existing = [
SegmentedObject(
content=_long_text("Starting the discussion.\n"),
object_type="conversation_phase",
turn_start=0,
turn_end=0,
token_estimate=200,
)
]
new_messages = [
_msg("assistant", _long_text("Continuing the discussion.\n")),
]
result = seg.segment_incremental(new_messages, existing, start_turn=1)
# Should merge the conversation_phase objects
assert len(result) >= 1
def test_incremental_empty_new(self, seg: Segmenter):
existing = [
SegmentedObject(
content="test",
object_type="conversation_phase",
turn_start=0,
turn_end=0,
token_estimate=1,
)
]
result = seg.segment_incremental([], existing)
assert len(result) == 1
def test_incremental_empty_existing(self, seg: Segmenter):
new_messages = [
_msg("assistant", _long_text("Hello world.\n")),
]
result = seg.segment_incremental(new_messages, [])
assert len(result) >= 1
def test_incremental_no_merge_across_types(self, seg: Segmenter):
"""file_context shouldn't merge with conversation_phase."""
existing = [
SegmentedObject(
content=_long_text("file content here\n"),
object_type="file_context",
source_tool="Read",
turn_start=0,
turn_end=0,
token_estimate=200,
)
]
new_messages = [
_msg("assistant", _long_text("Now let me explain what I found.\n")),
]
result = seg.segment_incremental(new_messages, existing, start_turn=1)
# Should NOT merge file_context with conversation_phase
assert len(result) >= 2
# ===========================================================================
# 9. Realistic multi-turn conversation payloads
# ===========================================================================
class TestRealisticPayloads:
"""Test with realistic multi-turn conversation structures."""
def test_typical_coding_session(self, seg: Segmenter):
"""Simulate: user asks → assistant reads file → assistant explains."""
messages = [
_msg(
"user",
_long_text(
"Can you look at the auth middleware and tell me "
"how it handles token refresh?\n"
),
),
_msg(
"assistant",
[
_text_block("I'll read the auth middleware file."),
_tool_use_block("t1", "Read", {"file_path": "src/auth/middleware.ts"}),
],
),
_msg(
"user",
[
_tool_result_block(
"t1",
_long_text(
"import jwt from 'jsonwebtoken';\n"
"export function handleAuth(req, res, next) {\n"
" const token = req.headers.authorization;\n"
" // ... token validation logic\n"
"}\n"
),
),
],
),
_msg(
"assistant",
_long_text(
"The auth middleware in src/auth/middleware.ts handles "
"token refresh by checking the JWT expiry and issuing "
"a new token if within the refresh window.\n"
),
),
]
objects = seg.segment_messages(messages)
types = {o.object_type for o in objects}
assert "file_context" in types
assert len(objects) >= 2
def test_debugging_flow(self, seg: Segmenter):
"""Simulate: error → investigation → fix."""
messages = [
_msg("user", _long_text("The tests are failing with this error.\n")),
_msg(
"assistant",
_long_text(
"I'm investigating the test failure. Let me look at the "
"root cause. The error seems to be a TypeError in the "
"authentication module.\n"
),
),
_msg(
"assistant",
[
_text_block("Let me run the failing test."),
_tool_use_block("t1", "Bash", {"command": "pytest tests/test_auth.py -v"}),
],
),
_msg(
"user",
[
_tool_result_block(
"t1",
_long_text(
"FAILED tests/test_auth.py::test_refresh - TypeError: "
"'NoneType' object is not subscriptable\n"
"Traceback (most recent call last):\n"
' File "tests/test_auth.py", line 42\n'
" token = response['access_token']\n"
"TypeError: 'NoneType' object is not subscriptable\n"
),
),
],
),
_msg(
"assistant",
_long_text(
"Found the problem! The root cause is that the refresh "
"endpoint returns None when the token is expired. "
"Fixed by adding a null check before accessing the response.\n"
),
),
]
objects = seg.segment_messages(messages)
types = {o.object_type for o in objects}
# Should have debugging and/or error objects
assert types & {"debugging_session", "error_context", "tool_result"}
def test_planning_session(self, seg: Segmenter):
"""Simulate: user asks for plan → assistant creates plan."""
messages = [
_msg(
"user",
_long_text(
"I need to add OAuth2 support. Can you create an implementation plan?\n"
),
),
_msg(
"assistant",
_long_text(
"Here's the implementation plan for OAuth2:\n"
"1. Install the oauth2 library\n"
"2. Create the OAuth2 provider configuration\n"
"3. Implement the authorization flow\n"
"4. Add callback handling\n"
"5. Write integration tests\n"
"Step 1 involves adding the dependency to package.json.\n"
),
),
]
objects = seg.segment_messages(messages)
assert any(o.object_type == "plan" for o in objects)
def test_multi_tool_turn(self, seg: Segmenter):
"""Assistant uses multiple tools in one turn."""
messages = [
_msg(
"assistant",
[
_text_block("Let me check both files."),
_tool_use_block("t1", "Read", {"file_path": "src/a.py"}),
_tool_use_block("t2", "Read", {"file_path": "src/b.py"}),
],
),
_msg(
"user",
[
_tool_result_block("t1", _long_text("# File A content\nclass A:\n pass\n")),
_tool_result_block("t2", _long_text("# File B content\nclass B:\n pass\n")),
],
),
]
objects = seg.segment_messages(messages)
file_objs = [o for o in objects if o.object_type == "file_context"]
assert len(file_objs) >= 2
def test_long_conversation_produces_multiple_objects(self, seg: Segmenter):
"""A long conversation should produce multiple objects."""
messages = []
for i in range(10):
messages.append(_msg("user", _long_text(f"Question {i} about the codebase.\n")))
messages.append(_msg("assistant", _long_text(f"Answer {i} with details.\n")))
objects = seg.segment_messages(messages)
assert len(objects) >= 3 # Should have multiple objects
# ===========================================================================
# 10. Empty and edge cases
# ===========================================================================
class TestEdgeCases:
"""Test empty inputs, malformed messages, and boundary conditions."""
def test_empty_messages(self, seg: Segmenter):
assert seg.segment_messages([]) == []
def test_empty_content_string(self, seg: Segmenter):
messages = [_msg("assistant", "")]
assert seg.segment_messages(messages) == []
def test_empty_content_list(self, seg: Segmenter):
messages = [_msg("assistant", [])]
assert seg.segment_messages(messages) == []
def test_whitespace_only_content(self, seg: Segmenter):
messages = [_msg("assistant", " \n\t ")]
assert seg.segment_messages(messages) == []
def test_missing_role(self, seg: Segmenter):
"""Messages without role should be handled gracefully."""
messages = [{"content": "no role here"}]
# Should not crash
result = seg.segment_messages(messages)
assert isinstance(result, list)
def test_non_dict_content_blocks(self, seg: Segmenter):
"""Content list with non-dict items should be handled."""
messages = [_msg("assistant", ["not a dict", 42, None])]
result = seg.segment_messages(messages)
assert isinstance(result, list)
def test_tool_result_with_list_content(self, seg: Segmenter):
"""tool_result content can be a list of text blocks."""
messages = [
_msg(
"assistant",
[
_tool_use_block("t1", "Read", {"file_path": "x.py"}),
],
),
_msg(
"user",
[
{
"type": "tool_result",
"tool_use_id": "t1",
"content": [{"type": "text", "text": _long_text("file content\n")}],
},
],
),
]
objects = seg.segment_messages(messages)
file_objs = [o for o in objects if o.object_type == "file_context"]
assert len(file_objs) >= 1
def test_token_estimate_accuracy(self, seg: Segmenter):
"""Token estimate should be approximately len/4."""
text = "a" * 400
assert _estimate_tokens(text) == 100
def test_token_estimate_minimum(self):
"""Token estimate should be at least 1."""
assert _estimate_tokens("") == 1
assert _estimate_tokens("a") == 1
def test_all_object_types_valid(self, seg: Segmenter):
"""All produced object types should be in VALID_OBJECT_TYPES."""
messages = [
_msg("user", _long_text("Question.\n")),
_msg(
"assistant",
_long_text(
'Traceback (most recent call last):\n File "x.py", line 1\nValueError: bad\n'
),
),
_msg("assistant", _long_text("Here's the plan:\n1. Do this\n2. Do that\n")),
_msg(
"assistant",
[
_tool_use_block("t1", "Read", {"file_path": "f.py"}),
],
),
_msg(
"user",
[
_tool_result_block("t1", _long_text("content\n")),
],
),
]
objects = seg.segment_messages(messages)
for obj in objects:
assert obj.object_type in VALID_OBJECT_TYPES, f"Invalid type: {obj.object_type}"
def test_oversized_content_gets_split(self, seg: Segmenter):
"""Content exceeding max_object_tokens should be split."""
small_seg = Segmenter(min_object_tokens=10, max_object_tokens=100)
huge_text = "x" * 2000 # ~500 tokens, well over 100
messages = [_msg("assistant", huge_text)]
objects = small_seg.segment_messages(messages)
assert len(objects) >= 2
for obj in objects:
assert obj.token_estimate <= 200 # Some slack for splitting
def test_segmented_object_defaults(self):
"""SegmentedObject should have sensible defaults."""
obj = SegmentedObject(content="test", object_type="conversation_phase")
assert obj.source_tool is None
assert obj.source_key is None
assert obj.stub == ""
assert obj.turn_start == 0
assert obj.turn_end == 0
assert obj.token_estimate == 0
assert obj.key_entities == []
assert obj.tags == []
def test_tool_result_with_error_content(self, seg: Segmenter):
"""Tool result containing errors should be classified as error_context."""
messages = [
_msg(
"assistant",
[
_tool_use_block("t1", "Bash", {"command": "npm test"}),
],
),
_msg(
"user",
[
_tool_result_block(
"t1",
_long_text(
"FAILED test_auth.py\n"
"TypeError: Cannot read property 'token' of undefined\n"
" at handleAuth (src/auth.js:42:15)\n"
" at processTicksAndRejections (internal/process/task_queues.js:95:5)\n"
),
),
],
),
]
objects = seg.segment_messages(messages)
assert any(o.object_type == "error_context" for o in objects)
def test_system_role_ignored(self, seg: Segmenter):
"""System messages should be ignored (not user or assistant)."""
messages = [
{"role": "system", "content": "You are a helpful assistant."},
_msg("assistant", _long_text("Hello!\n")),
]
objects = seg.segment_messages(messages)
# Should only have objects from the assistant message
assert all("helpful assistant" not in o.content for o in objects)