244 lines
11 KiB
Python
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
|