"""test_search.py — verify the search tool: retrieval, return format, error handling. Since search() imports get_embedder and get_conn lazily (inside the function body), we patch them at their source modules: codex.embed and codex.db. """ from __future__ import annotations from contextlib import contextmanager from typing import Any from unittest.mock import MagicMock, patch import numpy as np # --------------------------------------------------------------------------- # Helpers # --------------------------------------------------------------------------- def _make_conn_cm(mock_conn: MagicMock) -> Any: """Return a context-manager factory that yields *mock_conn*.""" @contextmanager # type: ignore[arg-type] def _cm() -> Any: yield mock_conn return _cm def _fake_row( bibkey: str, paper_id: str, ord_val: int, content: str, dist: float, ) -> dict[str, Any]: return { "bibkey": bibkey, "paper_id": paper_id, "ord": ord_val, "content": content, "dist": dist, } # --------------------------------------------------------------------------- # Tests # --------------------------------------------------------------------------- def test_search_returns_correct_keys() -> None: """search() must return dicts with the required five keys.""" from codex.mcp_server import search mock_conn = MagicMock() mock_conn.execute.return_value.fetchall.return_value = [ _fake_row("Springborn2008", "springborn-2008", 16, "The volume formula is V = ...", 0.2), ] fake_dense = np.zeros((1, 1024), dtype=np.float32) mock_embedder = MagicMock() mock_embedder.encode_dense.return_value = fake_dense with ( patch("codex.embed.get_embedder", return_value=mock_embedder), patch("codex.db.get_conn", side_effect=_make_conn_cm(mock_conn)), ): results = search("volume formula", limit=5) assert len(results) == 1 hit = results[0] assert set(hit.keys()) == {"bibkey", "paper_id", "locator", "score", "snippet"} def test_search_score_is_similarity_not_distance() -> None: """score = 1 - distance, so a dist of 0.0 should yield score 1.0.""" from codex.mcp_server import search mock_conn = MagicMock() mock_conn.execute.return_value.fetchall.return_value = [ _fake_row("Author2020", "paper-x", 3, "Some chunk content here.", 0.0), ] fake_dense = np.zeros((1, 1024), dtype=np.float32) mock_embedder = MagicMock() mock_embedder.encode_dense.return_value = fake_dense with ( patch("codex.embed.get_embedder", return_value=mock_embedder), patch("codex.db.get_conn", side_effect=_make_conn_cm(mock_conn)), ): results = search("some query") assert results[0]["score"] == 1.0 def test_search_locator_format() -> None: """locator must be 'chunk '.""" from codex.mcp_server import search mock_conn = MagicMock() mock_conn.execute.return_value.fetchall.return_value = [ _fake_row("Auth2021", "paper-y", 42, "Content.", 0.3), ] fake_dense = np.zeros((1, 1024), dtype=np.float32) mock_embedder = MagicMock() mock_embedder.encode_dense.return_value = fake_dense with ( patch("codex.embed.get_embedder", return_value=mock_embedder), patch("codex.db.get_conn", side_effect=_make_conn_cm(mock_conn)), ): results = search("query") assert results[0]["locator"] == "chunk 42" def test_search_snippet_truncated_to_300() -> None: """snippet must be truncated to 300 characters.""" from codex.mcp_server import search long_content = "x" * 500 mock_conn = MagicMock() mock_conn.execute.return_value.fetchall.return_value = [ _fake_row("Auth2022", "paper-z", 1, long_content, 0.1), ] fake_dense = np.zeros((1, 1024), dtype=np.float32) mock_embedder = MagicMock() mock_embedder.encode_dense.return_value = fake_dense with ( patch("codex.embed.get_embedder", return_value=mock_embedder), patch("codex.db.get_conn", side_effect=_make_conn_cm(mock_conn)), ): results = search("query") assert len(results[0]["snippet"]) == 300 # type: ignore[arg-type] def test_search_returns_error_dict_on_exception() -> None: """When the embedder raises, search must return [{"error": ...}], not crash.""" from codex.mcp_server import search mock_embedder = MagicMock() mock_embedder.encode_dense.side_effect = RuntimeError("DB down") with patch("codex.embed.get_embedder", return_value=mock_embedder): results = search("query") assert len(results) == 1 assert "error" in results[0] def test_search_provenance_keys() -> None: """Each hit must carry bibkey and paper_id for provenance traceability.""" from codex.mcp_server import search mock_conn = MagicMock() mock_conn.execute.return_value.fetchall.return_value = [ _fake_row("Crane1999", "crane-1999", 7, "Discrete exterior calculus ...", 0.15), ] fake_dense = np.zeros((1, 1024), dtype=np.float32) mock_embedder = MagicMock() mock_embedder.encode_dense.return_value = fake_dense with ( patch("codex.embed.get_embedder", return_value=mock_embedder), patch("codex.db.get_conn", side_effect=_make_conn_cm(mock_conn)), ): results = search("exterior calculus") assert results[0]["bibkey"] == "Crane1999" assert results[0]["paper_id"] == "crane-1999"