100 lines
4.1 KiB
Python
100 lines
4.1 KiB
Python
"""Tests for embedding records and provider-independent similarity math."""
|
|
|
|
import hashlib
|
|
import tempfile
|
|
import unittest
|
|
from pathlib import Path
|
|
|
|
from processors.rag_ingestion.chunking import ChunkConfiguration, chunk_document
|
|
from processors.rag_ingestion.embeddings import embed_chunk
|
|
from processors.rag_ingestion.loader import load_markdown_document
|
|
from processors.rag_ingestion.similarity import cosine_similarity
|
|
|
|
|
|
class FixedEmbeddingProvider:
|
|
provider_name = "Test Provider"
|
|
model_name = "test-embedding-model:v1"
|
|
expected_dimension = 3
|
|
|
|
def __init__(self, vector: list[float] | None = None) -> None:
|
|
self.vector = vector or [0.25, -0.5, 0.75]
|
|
self.received_text: str | None = None
|
|
|
|
def embed(self, text: str) -> list[float]:
|
|
self.received_text = text
|
|
return list(self.vector)
|
|
|
|
|
|
class CosineSimilarityTests(unittest.TestCase):
|
|
def test_identical_vectors_have_similarity_one(self) -> None:
|
|
self.assertAlmostEqual(cosine_similarity([1.0, 2.0], [1.0, 2.0]), 1.0)
|
|
|
|
def test_orthogonal_vectors_have_similarity_zero(self) -> None:
|
|
self.assertAlmostEqual(cosine_similarity([1.0, 0.0], [0.0, 1.0]), 0.0)
|
|
|
|
def test_opposite_vectors_have_similarity_negative_one(self) -> None:
|
|
self.assertAlmostEqual(cosine_similarity([1.0, 2.0], [-1.0, -2.0]), -1.0)
|
|
|
|
def test_dimension_mismatch_is_rejected(self) -> None:
|
|
with self.assertRaisesRegex(ValueError, "equal dimensions"):
|
|
cosine_similarity([1.0], [1.0, 2.0])
|
|
|
|
def test_zero_vector_is_rejected(self) -> None:
|
|
with self.assertRaisesRegex(ValueError, "zero vector"):
|
|
cosine_similarity([0.0, 0.0], [1.0, 1.0])
|
|
|
|
def test_empty_vectors_are_rejected(self) -> None:
|
|
with self.assertRaisesRegex(ValueError, "non-empty"):
|
|
cosine_similarity([], [])
|
|
|
|
|
|
class EmbeddingRecordTests(unittest.TestCase):
|
|
def test_embedding_uses_exact_chunk_text_and_preserves_linkage(self) -> None:
|
|
document = self._load_document("# Conversation\n\nExact chunk text.")
|
|
chunk = chunk_document(document, ChunkConfiguration(100, 0))[0]
|
|
provider = FixedEmbeddingProvider()
|
|
|
|
embedding = embed_chunk(chunk, provider)
|
|
|
|
self.assertEqual(provider.received_text, chunk.text)
|
|
self.assertEqual(embedding.chunk_id, chunk.chunk_id)
|
|
self.assertEqual(embedding.document_id, document.document_id)
|
|
self.assertEqual(embedding.embedding_model, provider.model_name)
|
|
self.assertEqual(embedding.embedding_provider, provider.provider_name)
|
|
self.assertEqual(embedding.embedding_dimension, 3)
|
|
self.assertEqual(embedding.vector, (0.25, -0.5, 0.75))
|
|
|
|
def test_unexpected_embedding_dimension_is_rejected(self) -> None:
|
|
document = self._load_document("source")
|
|
chunk = chunk_document(document, ChunkConfiguration(100, 0))[0]
|
|
|
|
with self.assertRaisesRegex(ValueError, "returned 2 dimensions"):
|
|
embed_chunk(chunk, FixedEmbeddingProvider([1.0, 2.0]))
|
|
|
|
def test_embedding_preserves_source_document_and_chunk_text(self) -> None:
|
|
with tempfile.TemporaryDirectory() as directory:
|
|
source = Path(directory) / "conversation.md"
|
|
source.write_text("Primary Source\n" * 10, encoding="utf-8")
|
|
source_hash = hashlib.sha256(source.read_bytes()).hexdigest()
|
|
document = load_markdown_document(source)
|
|
chunk = chunk_document(document, ChunkConfiguration(50, 10))[0]
|
|
document_text = document.raw_text
|
|
chunk_text = chunk.text
|
|
|
|
embed_chunk(chunk, FixedEmbeddingProvider())
|
|
|
|
self.assertEqual(hashlib.sha256(source.read_bytes()).hexdigest(), source_hash)
|
|
self.assertEqual(document.raw_text, document_text)
|
|
self.assertEqual(chunk.text, chunk_text)
|
|
|
|
def _load_document(self, text: str):
|
|
temporary_directory = tempfile.TemporaryDirectory()
|
|
self.addCleanup(temporary_directory.cleanup)
|
|
source = Path(temporary_directory.name) / "conversation.md"
|
|
source.write_text(text, encoding="utf-8")
|
|
return load_markdown_document(source)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|