acb / tests /test_embeddings.py
ktek's picture
Unit tests finalized
8853fb6
Raw
History Blame Contribute Delete
5.16 kB
import sys
from pathlib import Path
sys.path.insert(0, str(Path(__file__).parent.parent / "src"))
import pytest
import numpy as np
from unittest.mock import Mock, patch, MagicMock
class TestTurkishEmbedder:
"""Tests for TurkishEmbedder class."""
@patch("embeddings.SentenceTransformer")
def test_lazy_loading(self, mock_st):
"""Model should not load until first use."""
from embeddings import TurkishEmbedder, reset_embedder
reset_embedder()
embedder = TurkishEmbedder()
mock_st.assert_not_called()
_ = embedder.model
mock_st.assert_called_once()
@patch("embeddings.SentenceTransformer")
def test_embed_documents_adds_prefix(self, mock_st):
"""Documents should be prefixed with 'passage:'."""
from embeddings import TurkishEmbedder, reset_embedder
reset_embedder()
mock_model = MagicMock()
mock_model.encode.return_value = np.array([[0.1, 0.2], [0.3, 0.4]])
mock_st.return_value = mock_model
embedder = TurkishEmbedder()
embedder.embed_documents(["test1", "test2"])
call_args = mock_model.encode.call_args[0][0]
assert call_args == ["passage: test1", "passage: test2"]
@patch("embeddings.SentenceTransformer")
def test_embed_query_adds_prefix(self, mock_st):
"""Query should be prefixed with 'query:'."""
from embeddings import TurkishEmbedder, reset_embedder
reset_embedder()
mock_model = MagicMock()
mock_model.encode.return_value = np.array([0.1, 0.2])
mock_st.return_value = mock_model
embedder = TurkishEmbedder()
embedder.embed_query("test query")
call_args = mock_model.encode.call_args[0][0]
assert call_args == "query: test query"
@patch("embeddings.SentenceTransformer")
def test_query_and_document_prefix_differ(self, mock_st):
"""Query and document prefixes should be different."""
from embeddings import TurkishEmbedder, reset_embedder
reset_embedder()
mock_model = MagicMock()
mock_model.encode.return_value = np.array([0.1, 0.2])
mock_st.return_value = mock_model
embedder = TurkishEmbedder()
embedder.embed_query("test")
query_prefix = mock_model.encode.call_args[0][0]
mock_model.encode.return_value = np.array([[0.1, 0.2]])
embedder.embed_documents(["test"])
doc_prefix = mock_model.encode.call_args[0][0][0]
assert query_prefix != doc_prefix
assert query_prefix.startswith("query:")
assert doc_prefix.startswith("passage:")
@patch("embeddings.SentenceTransformer")
def test_embed_documents_empty_list(self, mock_st):
"""Empty input should return empty list."""
from embeddings import TurkishEmbedder, reset_embedder
reset_embedder()
embedder = TurkishEmbedder()
result = embedder.embed_documents([])
assert result == []
mock_st.assert_not_called()
@patch("embeddings.SentenceTransformer")
def test_embed_documents_returns_list(self, mock_st):
"""Should return list of lists."""
from embeddings import TurkishEmbedder, reset_embedder
reset_embedder()
mock_model = MagicMock()
mock_model.encode.return_value = np.array([[0.1, 0.2], [0.3, 0.4]])
mock_st.return_value = mock_model
embedder = TurkishEmbedder()
result = embedder.embed_documents(["a", "b"])
assert isinstance(result, list)
assert isinstance(result[0], list)
@patch("embeddings.SentenceTransformer")
def test_embed_passages_returns_numpy(self, mock_st):
"""embed_passages should return numpy array."""
from embeddings import TurkishEmbedder, reset_embedder
reset_embedder()
mock_model = MagicMock()
mock_model.encode.return_value = np.array([[0.1, 0.2]])
mock_st.return_value = mock_model
embedder = TurkishEmbedder()
result = embedder.embed_passages(["test"])
assert isinstance(result, np.ndarray)
class TestSingleton:
"""Tests for singleton pattern."""
@patch("embeddings.SentenceTransformer")
def test_get_embedder_returns_same_instance(self, mock_st):
"""get_embedder should return the same instance."""
from embeddings import get_embedder, reset_embedder
reset_embedder()
instance1 = get_embedder()
instance2 = get_embedder()
assert instance1 is instance2
@patch("embeddings.SentenceTransformer")
def test_reset_embedder_clears_instance(self, mock_st):
"""reset_embedder should clear the singleton."""
from embeddings import get_embedder, reset_embedder
reset_embedder()
instance1 = get_embedder()
reset_embedder()
instance2 = get_embedder()
assert instance1 is not instance2