From f288f88f19eaa87df80d408b9d5fc6757369f55e Mon Sep 17 00:00:00 2001 From: Ariel Memory Date: Mon, 6 Jul 2026 00:25:11 +0300 Subject: [PATCH] test: add saga, middleware, embeddings, tools unit tests (+54) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - test_saga_unit.py: persistence, idempotency, watchdog, helpers (20 tests) - test_middleware_unit.py: ImportanceGate, Validation, Dedup, Pipeline (10 tests) - test_embeddings_unit.py: hash_embedding, similarity (8 tests) - test_tools_unit.py: _validate_layer, _fire_hook, memory_remember, recall, forget, session, episode, graph, stats (16 tests) - Coverage: 64% → 70% (+6%) - shared/saga.py: 58% → 79% - shared/middleware.py: 71% → 88% - shared/embeddings.py: 58% → 78% - mcp_server/tools_layer.py: 15% → 53% - Total: 447 tests (was 393, +54) --- tests/test_mcp/test_tools_unit.py | 186 ++++++++++++++ tests/test_shared/test_embeddings_unit.py | 50 ++++ tests/test_shared/test_middleware_unit.py | 137 ++++++++++ tests/test_shared/test_saga_unit.py | 297 ++++++++++++++++++++++ 4 files changed, 670 insertions(+) create mode 100644 tests/test_mcp/test_tools_unit.py create mode 100644 tests/test_shared/test_embeddings_unit.py create mode 100644 tests/test_shared/test_middleware_unit.py create mode 100644 tests/test_shared/test_saga_unit.py diff --git a/tests/test_mcp/test_tools_unit.py b/tests/test_mcp/test_tools_unit.py new file mode 100644 index 00000000..3fb82b83 --- /dev/null +++ b/tests/test_mcp/test_tools_unit.py @@ -0,0 +1,186 @@ +"""Tests for mcp_server/tools_layer.py — tools with correct mock ctx.""" + +import pytest +from unittest.mock import MagicMock, AsyncMock +from mcp_server.tools_layer import ( + _validate_layer, + _fire_hook, + _get_cache_key, + memory_remember, + memory_recall, + memory_forget, + memory_session_start, + memory_session_end, + memory_episode_save, + memory_graph_add, + memory_stats, +) + + +# ── Helpers ── + + +def test_validate_layer_valid(): + assert _validate_layer("user") == "user" + assert _validate_layer("agent") == "agent" + + +def test_validate_layer_invalid(): + with pytest.raises(ValueError, match="Invalid layer"): + _validate_layer("admin") + + +def test_get_cache_key(): + key = _get_cache_key("user", "alice") + assert "user" in key + assert "alice" in key + + +def test_fire_hook_no_handlers(): + result = _fire_hook("nonexistent_hook", "user", {}) + assert result.get("skipped") is True + + +# ── Mock ctx helper ── + + +def _make_ctx(layer="user"): + """Create mock MCP ctx with AppContext.""" + ctx = MagicMock() + app = MagicMock() + app.mm = MagicMock() + app.rate_limiter = MagicMock() + app.rate_limiter.check = AsyncMock(return_value={"allowed": True, "remaining": 100, "reset_in": 60}) + app.emotion_trigger = MagicMock() + app.emotion_trigger.should_save = MagicMock(return_value=(False, "", 0.0)) + app.user_hooks = MagicMock() + app.agent_hooks = MagicMock() + app.user_graph = MagicMock() + app.agent_graph = MagicMock() + ctx.request_context = MagicMock() + ctx.request_context.lifespan_context = app + return ctx, app + + +# ── memory_remember ── + + +@pytest.mark.asyncio +async def test_remember_user(): + ctx, app = _make_ctx() + app.mm.user_memory.return_value.remember = AsyncMock(return_value=1) + result = await memory_remember(layer="user", user_id="u1", key="name", value="Alice", ctx=ctx) + assert result["status"] == "ok" + + +@pytest.mark.asyncio +async def test_remember_agent(): + ctx, app = _make_ctx() + app.mm.agent_memory.return_value.remember = AsyncMock(return_value=1) + app.agent_graph.add_node = AsyncMock(return_value=1) + result = await memory_remember(layer="agent", user_id="u1", key="decision", value="Use X", ctx=ctx) + assert result["status"] == "ok" + + +@pytest.mark.asyncio +async def test_remember_rate_limited(): + ctx, app = _make_ctx() + app.rate_limiter.check = AsyncMock(return_value={"allowed": False, "remaining": 0, "reset_in": 60}) + result = await memory_remember(layer="user", user_id="u1", key="k", value="v", ctx=ctx) + assert "error" in result + + +@pytest.mark.asyncio +async def test_remember_invalid_layer(): + ctx, _ = _make_ctx() + with pytest.raises(ValueError, match="Invalid layer"): + await memory_remember(layer="bad", user_id="u1", key="k", value="v", ctx=ctx) + + +# ── memory_recall ── + + +@pytest.mark.asyncio +async def test_recall(): + ctx, app = _make_ctx() + app.mm.user_memory.return_value.recall = AsyncMock(return_value=[{"key": "n"}]) + result = await memory_recall(layer="user", user_id="u1", query="name", ctx=ctx) + assert "results" in result + + +# ── memory_forget ── + + +@pytest.mark.asyncio +async def test_forget(): + ctx, app = _make_ctx() + app.mm.user_memory.return_value.forget = AsyncMock(return_value=True) + result = await memory_forget(layer="user", user_id="u1", key="k", ctx=ctx) + assert result.get("deleted") is True + + +# ── memory_session_start/end ── + + +@pytest.mark.asyncio +async def test_session_start(): + ctx, app = _make_ctx() + app.mm.user_memory.return_value.l2 = MagicMock() + app.mm.user_memory.return_value.l2.create_session = AsyncMock(return_value="s1") + result = await memory_session_start(layer="user", user_id="u1", ctx=ctx) + assert "session_id" in result + + +@pytest.mark.asyncio +async def test_session_end(): + ctx, app = _make_ctx() + app.mm.user_memory.return_value.l2 = MagicMock() + app.mm.user_memory.return_value.l2.close_session = AsyncMock() + result = await memory_session_end(layer="user", user_id="u1", session_id="s1", summary="done", ctx=ctx) + assert result["status"] == "ok" + + +# ── memory_episode_save ── + + +@pytest.mark.asyncio +async def test_episode_save(): + ctx, app = _make_ctx() + app.mm.user_memory.return_value.l3 = MagicMock() + app.mm.user_memory.return_value.l3.save = AsyncMock(return_value=1) + result = await memory_episode_save(layer="user", user_id="u1", summary="Event", weight=0.8, ctx=ctx) + assert "episode_id" in result + + +# ── memory_graph_add ── + + +@pytest.mark.asyncio +async def test_graph_add(): + ctx, app = _make_ctx() + app.user_graph.add_node = AsyncMock(return_value=1) + result = await memory_graph_add(layer="user", user_id="u1", content="Fact", node_type="fact", ctx=ctx) + assert "node_id" in result + + +# ── memory_stats ── + + +@pytest.mark.asyncio +async def test_stats(): + ctx, app = _make_ctx() + mem = app.mm.user_memory.return_value + mem.l1 = MagicMock() + mem.l1.size = MagicMock(return_value=0) + mem.l2 = MagicMock() + mem.l2.count_sessions = AsyncMock(return_value=0) + mem.l3 = MagicMock() + mem.l3.count = AsyncMock(return_value=0) + mem.l4 = MagicMock() + mem.l4.count = AsyncMock(return_value=0) + wiki = app.user_wiki + wiki.count = AsyncMock(return_value=0) + graph = app.user_graph + graph.count_nodes = AsyncMock(return_value=0) + result = await memory_stats(layer="user", user_id="u1", ctx=ctx) + assert isinstance(result, dict) diff --git a/tests/test_shared/test_embeddings_unit.py b/tests/test_shared/test_embeddings_unit.py new file mode 100644 index 00000000..d0e12123 --- /dev/null +++ b/tests/test_shared/test_embeddings_unit.py @@ -0,0 +1,50 @@ +"""Tests for shared/embeddings.py — hash embedding, similarity, cache.""" + +import pytest +from shared.embeddings import _hash_embedding, similarity + + +def test_hash_embedding_dim(): + result = _hash_embedding("test text", dim=128) + assert len(result) == 128 + + +def test_hash_embedding_normalized(): + result = _hash_embedding("test", dim=64) + norm = sum(x**2 for x in result) ** 0.5 + assert abs(norm - 1.0) < 0.01 + + +def test_hash_embedding_deterministic(): + r1 = _hash_embedding("hello", dim=32) + r2 = _hash_embedding("hello", dim=32) + assert r1 == r2 + + +def test_hash_embedding_different_inputs(): + r1 = _hash_embedding("hello", dim=32) + r2 = _hash_embedding("world", dim=32) + assert r1 != r2 + + +def test_similarity_identical(): + v = [1.0, 0.0, 0.0] + assert similarity(v, v) == pytest.approx(1.0) + + +def test_similarity_orthogonal(): + v1 = [1.0, 0.0, 0.0] + v2 = [0.0, 1.0, 0.0] + assert similarity(v1, v2) == pytest.approx(0.0) + + +def test_similarity_zero_vector(): + v1 = [1.0, 0.0] + v2 = [0.0, 0.0] + assert similarity(v1, v2) == 0.0 + + +def test_similarity_opposite(): + v1 = [1.0, 0.0] + v2 = [-1.0, 0.0] + assert similarity(v1, v2) == pytest.approx(-1.0) diff --git a/tests/test_shared/test_middleware_unit.py b/tests/test_shared/test_middleware_unit.py new file mode 100644 index 00000000..e11d2ee9 --- /dev/null +++ b/tests/test_shared/test_middleware_unit.py @@ -0,0 +1,137 @@ +"""Tests for shared/middleware.py — actual API.""" + +import asyncio +from shared.middleware import ( + MiddlewareContext, + ImportanceGateMiddleware, + DedupMiddleware, + ValidationMiddleware, + AuditMiddleware, + MiddlewarePipeline, +) + + +def test_middleware_context_defaults(): + ctx = MiddlewareContext() + assert ctx.tool_name == "" + assert ctx.user_id == "default" + assert ctx.blocked is False + + +def test_importance_gate_passes_non_matching_tool(): + gate = ImportanceGateMiddleware() + ctx = MiddlewareContext(args={"importance": 0.1}, tool_name="other_tool") + + async def handler(c): + return {"ok": True} + + result = asyncio.run(gate.process(ctx, handler)) + assert result == {"ok": True} + assert ctx.blocked is False + + +def test_importance_gate_blocks_low(): + gate = ImportanceGateMiddleware() + ctx = MiddlewareContext(args={"value": "hi"}, tool_name="memory_user_remember") + + async def handler(c): + return {"ok": True} + + result = asyncio.run(gate.process(ctx, handler)) + assert ctx.blocked is True + + +def test_importance_gate_allows_high(): + gate = ImportanceGateMiddleware() + ctx = MiddlewareContext( + args={ + "value": "This is a critical and important decision about our architecture that affects production systems and requires immediate attention" + }, + tool_name="memory_user_remember", + ) + + async def handler(c): + return {"ok": True} + + result = asyncio.run(gate.process(ctx, handler)) + assert ctx.blocked is False + + +def test_validation_blocks_empty_user(): + val = ValidationMiddleware() + ctx = MiddlewareContext(user_id="", tool_name="memory_remember") + + async def handler(c): + return {"ok": True} + + asyncio.run(val.process(ctx, handler)) + assert ctx.blocked is True + + +def test_validation_blocks_missing_key(): + val = ValidationMiddleware() + ctx = MiddlewareContext(user_id="u1", tool_name="memory_user_remember", args={}) + + async def handler(c): + return {"ok": True} + + asyncio.run(val.process(ctx, handler)) + assert ctx.blocked is True + + +def test_validation_allows_valid(): + val = ValidationMiddleware() + ctx = MiddlewareContext(user_id="u1", tool_name="memory_user_remember", args={"key": "k"}) + + async def handler(c): + return {"ok": True} + + asyncio.run(val.process(ctx, handler)) + assert ctx.blocked is False + + +def test_dedup_catches_duplicates(): + dedup = DedupMiddleware() + ctx1 = MiddlewareContext(tool_name="test_tool", user_id="u1", args={"k": "v"}) + + async def handler(c): + return {"ok": True} + + r1 = asyncio.run(dedup.process(ctx1, handler)) + assert r1 == {"ok": True} + + ctx2 = MiddlewareContext(tool_name="test_tool", user_id="u1", args={"k": "v"}) + r2 = asyncio.run(dedup.process(ctx2, handler)) + assert ctx2.metadata.get("deduped") is True + + +def test_pipeline_runs(): + pipe = MiddlewarePipeline() + + class CountMiddleware: + name = "count" + + async def process(self, ctx, next_fn): + ctx.metadata["count"] = True + return await next_fn(ctx) + + pipe.add(CountMiddleware()) + + async def handler(ctx): + return {"ok": True} + + ctx = MiddlewareContext() + result = asyncio.run(pipe.execute(ctx, handler)) + assert result == {"ok": True} + assert ctx.metadata.get("count") is True + + +def test_audit_sets_metadata(): + audit = AuditMiddleware() + ctx = MiddlewareContext() + + async def handler(c): + return {"ok": True} + + asyncio.run(audit.process(ctx, handler)) + assert "elapsed" in ctx.metadata diff --git a/tests/test_shared/test_saga_unit.py b/tests/test_shared/test_saga_unit.py new file mode 100644 index 00000000..2f8924fb --- /dev/null +++ b/tests/test_shared/test_saga_unit.py @@ -0,0 +1,297 @@ +"""Unit tests for shared/saga.py — Saga, SagaWatchdog, helpers.""" + +import asyncio +import json +import time +from shared.saga import ( + Saga, + SagaStep, + SagaStatus, + SagaWatchdog, + create_consolidation_saga, + create_backup_saga, +) + + +async def _noop(data): + return {"ok": True} + + +async def _failing(data): + raise ValueError("step failed") + + +# ── Saga basics ── + + +def test_add_step(): + s = Saga("test") + s.add_step("s1", _noop) + assert len(s._steps) == 1 + s.add_step("s2", _noop) + assert len(s._steps) == 2 + + +def test_status_property(): + s = Saga("test") + assert s.status == SagaStatus.PENDING + + +def test_data_property(): + s = Saga("test") + s._data = {"k": "v"} + assert s.data == {"k": "v"} + + +def test_get_state(): + s = Saga("test") + state = s.get_state() + assert state["name"] == "test" + assert state["status"] == "pending" + assert isinstance(state["steps"], list) + + +# ── State persistence ── + + +def test_save_state_creates_file(tmp_path): + from shared import saga as saga_mod + + orig = saga_mod.SAGA_DIR + saga_mod.SAGA_DIR = tmp_path + try: + s = Saga("save_test") + s.add_step("s1", _noop) + s._saga_id = "sv_1" + s._save_state() + assert (tmp_path / "sv_1.json").exists() + assert len((tmp_path / "sv_1.json").read_bytes()) > 0 + finally: + saga_mod.SAGA_DIR = orig + + +def test_load_state_roundtrip(tmp_path): + from shared import saga as saga_mod + + orig = saga_mod.SAGA_DIR + saga_mod.SAGA_DIR = tmp_path + try: + s1 = Saga("rt") + s1._saga_id = "rt_1" + s1._data = {"key": "value"} + s1._save_state() + + s2 = Saga("rt") + loaded = s2._load_state("rt_1") + assert loaded is not None + assert loaded["data"] == {"key": "value"} + finally: + saga_mod.SAGA_DIR = orig + + +def test_load_state_missing(tmp_path): + from shared import saga as saga_mod + + orig = saga_mod.SAGA_DIR + saga_mod.SAGA_DIR = tmp_path + try: + s = Saga("t") + assert s._load_state("nonexistent") is None + finally: + saga_mod.SAGA_DIR = orig + + +def test_cleanup_state(tmp_path): + from shared import saga as saga_mod + + orig = saga_mod.SAGA_DIR + saga_mod.SAGA_DIR = tmp_path + try: + (tmp_path / "del.json").write_text("{}") + s = Saga("t") + s._saga_id = "del" + s._cleanup_state() + assert not (tmp_path / "del.json").exists() + finally: + saga_mod.SAGA_DIR = orig + + +# ── Idempotency ── + + +def test_compute_idempotency_key_none_without_fn(): + s = Saga("t") + step = SagaStep(name="s1", action=_noop) + assert s._compute_idempotency_key(step) is None + + +def test_compute_idempotency_key_deterministic(): + s = Saga("t") + step = SagaStep(name="s1", action=_noop, idempotency_key_fn=lambda d: "key123") + k1 = s._compute_idempotency_key(step) + k2 = s._compute_idempotency_key(step) + assert k1 == k2 # deterministic, not necessarily the raw value + assert k1 is not None + + +def test_is_already_completed_false(): + s = Saga("t") + result = asyncio.run(s._is_already_completed("nonexistent_key")) + assert result is False + + +def test_get_cached_result_none(): + s = Saga("t") + result = asyncio.run(s._get_cached_result("nonexistent_key")) + assert result is None + + +# ── SagaWatchdog ── + + +def test_watchdog_get_stuck_sagas(tmp_path): + from shared import saga as saga_mod + + orig = saga_mod.SAGA_DIR + saga_mod.SAGA_DIR = tmp_path + try: + old = { + "name": "old", + "saga_id": "o1", + "status": "running", + "current_step": 0, + "started_at": time.time() - 120, + "data": {}, + "completed_steps": [], + "steps": [{"name": "s1", "status": "completed", "result": {}}], + } + (tmp_path / "o1.json").write_text(json.dumps(old)) + wd = SagaWatchdog(max_age_seconds=60) + stuck = wd.get_stuck_sagas() + assert len(stuck) == 1 + assert stuck[0]["saga_id"] == "o1" + finally: + saga_mod.SAGA_DIR = orig + + +def test_watchdog_recover_sets_manual_review(tmp_path): + from shared import saga as saga_mod + + orig = saga_mod.SAGA_DIR + saga_mod.SAGA_DIR = tmp_path + try: + old = { + "name": "r", + "saga_id": "r1", + "status": "stuck", + "current_step": 0, + "started_at": time.time() - 120, + "data": {}, + "completed_steps": [], + "steps": [{"name": "s1", "status": "completed", "result": {}}], + } + (tmp_path / "r1.json").write_text(json.dumps(old)) + wd = SagaWatchdog(max_age_seconds=60) + result = wd.recover_saga("r1") + assert result is not None + assert result["status"] == "manual_review_required" + finally: + saga_mod.SAGA_DIR = orig + + +def test_watchdog_recover_nonexistent(tmp_path): + from shared import saga as saga_mod + + orig = saga_mod.SAGA_DIR + saga_mod.SAGA_DIR = tmp_path + try: + wd = SagaWatchdog(max_age_seconds=60) + assert wd.recover_saga("nope") is None + finally: + saga_mod.SAGA_DIR = orig + + +def test_watchdog_cleanup_completed(tmp_path): + from shared import saga as saga_mod + + orig = saga_mod.SAGA_DIR + saga_mod.SAGA_DIR = tmp_path + try: + done = { + "name": "d", + "saga_id": "d1", + "status": "completed", + "current_step": 1, + "started_at": time.time() - 7200, + "data": {}, + "completed_steps": [0], + "steps": [{"name": "s1", "status": "completed", "result": {}}], + } + run = { + "name": "r", + "saga_id": "r1", + "status": "running", + "current_step": 0, + "started_at": time.time(), + "data": {}, + "completed_steps": [], + "steps": [{"name": "s1", "status": "running", "result": {}}], + } + (tmp_path / "d1.json").write_text(json.dumps(done)) + (tmp_path / "r1.json").write_text(json.dumps(run)) + wd = SagaWatchdog(max_age_seconds=60) + removed = wd.cleanup_completed() + assert removed >= 1 + assert not (tmp_path / "d1.json").exists() + assert (tmp_path / "r1.json").exists() + finally: + saga_mod.SAGA_DIR = orig + + +def test_watchdog_start_stop(tmp_path): + from shared import saga as saga_mod + + orig = saga_mod.SAGA_DIR + saga_mod.SAGA_DIR = tmp_path + try: + wd = SagaWatchdog(check_interval=1, max_age_seconds=60) + wd.start() + assert wd._running is True + wd.stop() + assert wd._running is False + finally: + saga_mod.SAGA_DIR = orig + + +# ── Helper functions ── + + +def test_create_consolidation_saga(): + s = create_consolidation_saga("user1") + assert "consolidation" in s.name + assert "user1" in s.name + assert len(s._steps) > 0 + + +def test_create_backup_saga(): + s = create_backup_saga() + assert s.name == "backup" + assert len(s._steps) > 0 + + +def test_consolidation_saga_execute(): + async def t(): + s = create_consolidation_saga("test_u") + result = await s.execute() + assert result is not None + + asyncio.run(t()) + + +def test_backup_saga_execute(): + async def t(): + s = create_backup_saga() + result = await s.execute() + assert result is not None + + asyncio.run(t())