Files
ObsiGate/tests/test_semantic_search.py
T

244 lines
11 KiB
Python

# tests/test_semantic_search.py — Tests for semantic search (embeddings + RRF)
import math
from backend.search import advanced_search
from backend.semantic_search import (
EMBEDDING_DIM,
HashEmbeddingProvider,
SemanticIndex,
VectorStore,
chunk_text,
get_semantic_index,
on_index_change,
reset_semantic_index,
rrf_fuse,
semantic_status,
)
# ═══════════════════════════════════════════════════════════════════
# Chunking
# ═══════════════════════════════════════════════════════════════════
class TestChunkText:
def test_empty(self):
assert chunk_text("") == []
assert chunk_text(" ") == []
def test_short_text_single_chunk(self):
chunks = chunk_text("un deux trois", chunk_tokens=512)
assert chunks == ["un deux trois"]
def test_long_text_multiple_chunks(self):
words = " ".join(f"mot{i}" for i in range(1200))
chunks = chunk_text(words, chunk_tokens=512, overlap=64)
assert len(chunks) >= 3
# First chunk has 512 words
assert len(chunks[0].split()) == 512
def test_overlap_shared_tokens(self):
words = [f"w{i}" for i in range(600)]
chunks = chunk_text(" ".join(words), chunk_tokens=100, overlap=20)
first_tail = chunks[0].split()[-20:]
second_head = chunks[1].split()[:20]
assert first_tail == second_head
# ═══════════════════════════════════════════════════════════════════
# Hash embedding provider (dependency-free fallback)
# ═══════════════════════════════════════════════════════════════════
class TestHashEmbeddingProvider:
def setup_method(self):
self.provider = HashEmbeddingProvider()
def test_dimension(self):
vec = self.provider.encode_one("hello world")
assert len(vec) == EMBEDDING_DIM
def test_l2_normalized(self):
vec = self.provider.encode_one("un texte de test assez long")
norm = math.sqrt(sum(v * v for v in vec))
assert abs(norm - 1.0) < 1e-6
def test_empty_text_zero_vector(self):
vec = self.provider.encode_one("")
assert all(v == 0.0 for v in vec)
def test_deterministic(self):
a = self.provider.encode_one("sauvegarde des données")
b = self.provider.encode_one("sauvegarde des données")
assert a == b
def test_similar_more_similar_than_dissimilar(self):
base = self.provider.encode_one("stratégie de sauvegarde automatique des données")
close = self.provider.encode_one("sauvegarde automatique des données")
far = self.provider.encode_one("recette de cuisine au chocolat")
sim_close = sum(x * y for x, y in zip(base, close))
sim_far = sum(x * y for x, y in zip(base, far))
assert sim_close > sim_far
def test_batch_matches_single(self):
texts = ["premier document", "deuxième document"]
batch = self.provider.encode(texts)
assert batch[0] == self.provider.encode_one(texts[0])
assert batch[1] == self.provider.encode_one(texts[1])
# ═══════════════════════════════════════════════════════════════════
# Vector store
# ═══════════════════════════════════════════════════════════════════
class TestVectorStore:
def test_add_and_search(self):
provider = HashEmbeddingProvider()
store = VectorStore(provider.dimension)
store.add("v::a.md", "a", provider.encode_one("python programmation"))
store.add("v::b.md", "b", provider.encode_one("recette cuisine chocolat"))
hits = store.search(provider.encode_one("python"), top_k=2)
assert hits[0][0] == "v::a.md"
assert hits[0][1] > hits[1][1]
def test_remove_document(self):
provider = HashEmbeddingProvider()
store = VectorStore(provider.dimension)
store.add("v::a.md", "a", provider.encode_one("python"))
store.add("v::a.md", "a2", provider.encode_one("code"))
store.add("v::b.md", "b", provider.encode_one("cuisine"))
assert len(store) == 3
store.remove_document("v::a.md")
assert len(store) == 1
hits = store.search(provider.encode_one("python"), top_k=5)
assert all(key != "v::a.md" for key, _ in hits)
def test_empty_search(self):
store = VectorStore(EMBEDDING_DIM)
assert store.search([0.0] * EMBEDDING_DIM) == []
# ═══════════════════════════════════════════════════════════════════
# Reciprocal Rank Fusion
# ═══════════════════════════════════════════════════════════════════
class TestRRF:
def test_single_ranking(self):
scores = rrf_fuse([["a", "b"]])
assert scores["a"] > scores["b"]
def test_fusion_promotes_consensus(self):
scores = rrf_fuse([["a", "b", "c"], ["b", "c", "d"]])
# "b" is well ranked by both methods -> best fused score
assert max(scores, key=scores.get) == "b"
def test_dedup_within_ranking(self):
scores = rrf_fuse([["a", "a", "b"]])
single = rrf_fuse([["a", "b"]])
assert abs(scores["a"] - single["a"]) < 1e-9
# ═══════════════════════════════════════════════════════════════════
# SemanticIndex — unit (injected provider, manual documents)
# ═══════════════════════════════════════════════════════════════════
class TestSemanticIndexUnit:
def test_add_search_remove(self):
index = SemanticIndex(provider=HashEmbeddingProvider())
index._ready = True
index.add_document("V", "backup.md", {
"path": "backup.md",
"title": "Stratégie de backup",
"content": "Protection et sauvegarde automatique des données",
"tags": [],
})
index.add_document("V", "cuisine.md", {
"path": "cuisine.md",
"title": "Recette au chocolat",
"content": "Faire fondre le chocolat puis ajouter la farine",
"tags": [],
})
hits = index.search("sauvegarde des données", top_k=5)
assert hits
assert hits[0][0] == "V::backup.md"
index.remove_document("V", "backup.md")
hits = index.search("sauvegarde des données", top_k=5)
assert all(key != "V::backup.md" for key, _ in hits)
def test_vault_filter(self):
index = SemanticIndex(provider=HashEmbeddingProvider())
index._ready = True
index.add_document("V1", "a.md", {"path": "a.md", "title": "python", "content": "python", "tags": []})
index.add_document("V2", "b.md", {"path": "b.md", "title": "python", "content": "python", "tags": []})
hits = index.search("python", vault_filter="V1", top_k=5)
assert hits
assert all(key.startswith("V1::") for key, _ in hits)
def test_not_ready_is_noop(self):
index = SemanticIndex(provider=HashEmbeddingProvider())
index.add_document("V", "a.md", {"path": "a.md", "title": "x", "content": "x", "tags": []})
assert len(index.store) == 0
assert index.search("x") == []
# ═══════════════════════════════════════════════════════════════════
# Integration with the global index / advanced_search
# ═══════════════════════════════════════════════════════════════════
class TestSemanticIntegration:
def test_rebuild_from_global_index(self, client):
index = get_semantic_index()
assert index.is_ready()
assert len(index.doc_keys) >= 3
def test_on_index_change_hook(self, client):
index = get_semantic_index()
file_info = {
"path": "semantic_hook.md",
"title": "Sauvegarde",
"content": "Sauvegarde automatique des données",
"tags": [],
}
on_index_change("add", "TestVault", "semantic_hook.md", file_info)
assert "TestVault::semantic_hook.md" in index.doc_keys
on_index_change("remove", "TestVault", "semantic_hook.md", file_info)
assert "TestVault::semantic_hook.md" not in index.doc_keys
def test_advanced_search_semantic(self, client):
result = advanced_search("python", vault_filter="all", semantic=True)
assert result["semantic_available"] is True
assert len(result["results"]) >= 1
assert any(r.get("semantic_score", 0) > 0 for r in result["results"])
def test_advanced_search_semantic_field_default(self, client):
result = advanced_search("python", vault_filter="all")
for r in result["results"]:
assert "semantic_score" in r
def test_semantic_status(self, client):
status = semantic_status()
assert status["available"] is True
assert status["documents"] >= 3
assert status["dimension"] > 0
def test_reset_semantic_index(self, client):
reset_semantic_index()
assert get_semantic_index().is_ready() is False
# Restore for subsequent tests
from backend.semantic_search import init_semantic_index
init_semantic_index()
class TestSemanticAPI:
def test_api_semantic_flag(self, client):
resp = client.get("/api/search/advanced?q=python&vault=all&semantic=true")
assert resp.status_code == 200
data = resp.json()
assert data["semantic_available"] is True
assert len(data["results"]) >= 1
assert "semantic_score" in data["results"][0]
def test_api_without_semantic(self, client):
resp = client.get("/api/search/advanced?q=python&vault=all")
assert resp.status_code == 200
data = resp.json()
assert "semantic_available" in data