chore: import upstream snapshot with attribution
Build and push multi-arch DocsGPT Docker image / build (linux/amd64, ubuntu-latest, amd64) (push) Has been cancelled
Backend release / release (push) Has been cancelled
Bandit Security Scan / bandit_scan (push) Has been cancelled
Build and push multi-arch DocsGPT Docker image / build (linux/arm64, ubuntu-24.04-arm, arm64) (push) Has been cancelled
Build and push multi-arch DocsGPT Docker image / manifest (push) Has been cancelled
Build and push DocsGPT FE Docker image for development / build (linux/amd64, ubuntu-latest, amd64) (push) Has been cancelled
Build and push DocsGPT FE Docker image for development / build (linux/arm64, ubuntu-24.04-arm, arm64) (push) Has been cancelled
Build and push DocsGPT FE Docker image for development / manifest (push) Has been cancelled
Python linting / ruff (push) Has been cancelled
Run python tests with pytest / Run tests and count coverage (3.12) (push) Has been cancelled
React Widget Build / build (push) Has been cancelled
Build and push multi-arch DocsGPT Docker image / build (linux/amd64, ubuntu-latest, amd64) (push) Has been cancelled
Backend release / release (push) Has been cancelled
Bandit Security Scan / bandit_scan (push) Has been cancelled
Build and push multi-arch DocsGPT Docker image / build (linux/arm64, ubuntu-24.04-arm, arm64) (push) Has been cancelled
Build and push multi-arch DocsGPT Docker image / manifest (push) Has been cancelled
Build and push DocsGPT FE Docker image for development / build (linux/amd64, ubuntu-latest, amd64) (push) Has been cancelled
Build and push DocsGPT FE Docker image for development / build (linux/arm64, ubuntu-24.04-arm, arm64) (push) Has been cancelled
Build and push DocsGPT FE Docker image for development / manifest (push) Has been cancelled
Python linting / ruff (push) Has been cancelled
Run python tests with pytest / Run tests and count coverage (3.12) (push) Has been cancelled
React Widget Build / build (push) Has been cancelled
This commit is contained in:
@@ -0,0 +1,386 @@
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from application.vectorstore.base import (
|
||||
BaseVectorStore,
|
||||
EmbeddingsSingleton,
|
||||
RemoteEmbeddings,
|
||||
)
|
||||
|
||||
|
||||
# --- RemoteEmbeddings ---
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestRemoteEmbeddings:
|
||||
def test_init_sets_url_and_headers(self):
|
||||
emb = RemoteEmbeddings(
|
||||
api_url="http://localhost:8080/", model_name="model-v1", api_key="sk-key"
|
||||
)
|
||||
assert emb.api_url == "http://localhost:8080"
|
||||
assert emb.model_name == "model-v1"
|
||||
assert emb.headers["Authorization"] == "Bearer sk-key"
|
||||
|
||||
def test_init_no_api_key(self):
|
||||
emb = RemoteEmbeddings(api_url="http://host", model_name="m")
|
||||
assert "Authorization" not in emb.headers
|
||||
|
||||
@patch("application.vectorstore.base.requests.post")
|
||||
def test_embed_sends_correct_payload(self, mock_post):
|
||||
mock_resp = Mock()
|
||||
mock_resp.json.return_value = {
|
||||
"data": [{"index": 0, "embedding": [0.1, 0.2]}]
|
||||
}
|
||||
mock_resp.raise_for_status = Mock()
|
||||
mock_post.return_value = mock_resp
|
||||
|
||||
emb = RemoteEmbeddings("http://host", "model-v1")
|
||||
result = emb._embed("test input")
|
||||
|
||||
mock_post.assert_called_once()
|
||||
call_kwargs = mock_post.call_args
|
||||
assert call_kwargs[1]["json"]["input"] == "test input"
|
||||
assert call_kwargs[1]["json"]["model"] == "model-v1"
|
||||
assert result == [[0.1, 0.2]]
|
||||
|
||||
@patch("application.vectorstore.base.requests.post")
|
||||
def test_embed_sorts_by_index(self, mock_post):
|
||||
mock_resp = Mock()
|
||||
mock_resp.json.return_value = {
|
||||
"data": [
|
||||
{"index": 1, "embedding": [0.3, 0.4]},
|
||||
{"index": 0, "embedding": [0.1, 0.2]},
|
||||
]
|
||||
}
|
||||
mock_resp.raise_for_status = Mock()
|
||||
mock_post.return_value = mock_resp
|
||||
|
||||
emb = RemoteEmbeddings("http://host", "m")
|
||||
result = emb._embed(["a", "b"])
|
||||
assert result == [[0.1, 0.2], [0.3, 0.4]]
|
||||
|
||||
@patch("application.vectorstore.base.requests.post")
|
||||
def test_embed_raises_on_error_response(self, mock_post):
|
||||
mock_resp = Mock()
|
||||
mock_resp.json.return_value = {"error": "rate limit exceeded"}
|
||||
mock_resp.raise_for_status = Mock()
|
||||
mock_post.return_value = mock_resp
|
||||
|
||||
emb = RemoteEmbeddings("http://host", "m")
|
||||
with pytest.raises(ValueError, match="rate limit exceeded"):
|
||||
emb._embed("test")
|
||||
|
||||
@patch("application.vectorstore.base.requests.post")
|
||||
def test_embed_raises_on_unexpected_format(self, mock_post):
|
||||
mock_resp = Mock()
|
||||
mock_resp.json.return_value = {"unexpected": True}
|
||||
mock_resp.raise_for_status = Mock()
|
||||
mock_post.return_value = mock_resp
|
||||
|
||||
emb = RemoteEmbeddings("http://host", "m")
|
||||
with pytest.raises(ValueError, match="Unexpected response format"):
|
||||
emb._embed("test")
|
||||
|
||||
@patch("application.vectorstore.base.requests.post")
|
||||
def test_embed_raises_on_non_dict_response(self, mock_post):
|
||||
mock_resp = Mock()
|
||||
mock_resp.json.return_value = [1, 2, 3]
|
||||
mock_resp.raise_for_status = Mock()
|
||||
mock_post.return_value = mock_resp
|
||||
|
||||
emb = RemoteEmbeddings("http://host", "m")
|
||||
with pytest.raises(ValueError, match="Unexpected response format"):
|
||||
emb._embed("test")
|
||||
|
||||
@patch("application.vectorstore.base.requests.post")
|
||||
def test_embed_query(self, mock_post):
|
||||
mock_resp = Mock()
|
||||
mock_resp.json.return_value = {
|
||||
"data": [{"index": 0, "embedding": [0.1, 0.2, 0.3]}]
|
||||
}
|
||||
mock_resp.raise_for_status = Mock()
|
||||
mock_post.return_value = mock_resp
|
||||
|
||||
emb = RemoteEmbeddings("http://host", "m")
|
||||
emb.dimension = None # Reset so it gets set from response
|
||||
result = emb.embed_query("hello")
|
||||
assert result == [0.1, 0.2, 0.3]
|
||||
assert emb.dimension == 3
|
||||
|
||||
@patch("application.vectorstore.base.requests.post")
|
||||
def test_embed_query_raises_on_bad_structure(self, mock_post):
|
||||
mock_resp = Mock()
|
||||
# Return multiple embeddings for a single query
|
||||
mock_resp.json.return_value = {
|
||||
"data": [
|
||||
{"index": 0, "embedding": [0.1]},
|
||||
{"index": 1, "embedding": [0.2]},
|
||||
]
|
||||
}
|
||||
mock_resp.raise_for_status = Mock()
|
||||
mock_post.return_value = mock_resp
|
||||
|
||||
emb = RemoteEmbeddings("http://host", "m")
|
||||
with pytest.raises(ValueError, match="Unexpected result structure"):
|
||||
emb.embed_query("hello")
|
||||
|
||||
@patch("application.vectorstore.base.requests.post")
|
||||
def test_embed_documents(self, mock_post):
|
||||
mock_resp = Mock()
|
||||
mock_resp.json.return_value = {
|
||||
"data": [
|
||||
{"index": 0, "embedding": [0.1, 0.2]},
|
||||
{"index": 1, "embedding": [0.3, 0.4]},
|
||||
]
|
||||
}
|
||||
mock_resp.raise_for_status = Mock()
|
||||
mock_post.return_value = mock_resp
|
||||
|
||||
emb = RemoteEmbeddings("http://host", "m")
|
||||
emb.dimension = None # Reset so it gets set from response
|
||||
result = emb.embed_documents(["doc1", "doc2"])
|
||||
assert result == [[0.1, 0.2], [0.3, 0.4]]
|
||||
assert emb.dimension == 2
|
||||
|
||||
def test_embed_documents_empty(self):
|
||||
emb = RemoteEmbeddings("http://host", "m")
|
||||
assert emb.embed_documents([]) == []
|
||||
|
||||
@patch("application.vectorstore.base.requests.post")
|
||||
def test_call_with_string(self, mock_post):
|
||||
mock_resp = Mock()
|
||||
mock_resp.json.return_value = {
|
||||
"data": [{"index": 0, "embedding": [0.5]}]
|
||||
}
|
||||
mock_resp.raise_for_status = Mock()
|
||||
mock_post.return_value = mock_resp
|
||||
|
||||
emb = RemoteEmbeddings("http://host", "m")
|
||||
result = emb("hello")
|
||||
assert result == [0.5]
|
||||
|
||||
@patch("application.vectorstore.base.requests.post")
|
||||
def test_call_with_list(self, mock_post):
|
||||
mock_resp = Mock()
|
||||
mock_resp.json.return_value = {
|
||||
"data": [{"index": 0, "embedding": [0.5]}]
|
||||
}
|
||||
mock_resp.raise_for_status = Mock()
|
||||
mock_post.return_value = mock_resp
|
||||
|
||||
emb = RemoteEmbeddings("http://host", "m")
|
||||
result = emb(["hello"])
|
||||
assert result == [[0.5]]
|
||||
|
||||
def test_call_with_invalid_type(self):
|
||||
emb = RemoteEmbeddings("http://host", "m")
|
||||
with pytest.raises(ValueError, match="Input must be a string or a list"):
|
||||
emb(123)
|
||||
|
||||
|
||||
# --- EmbeddingsSingleton ---
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestEmbeddingsSingleton:
|
||||
def setup_method(self):
|
||||
EmbeddingsSingleton._instances = {}
|
||||
|
||||
@patch("application.vectorstore.base.OpenAIEmbeddings")
|
||||
def test_get_instance_openai(self, mock_openai_cls):
|
||||
mock_instance = Mock()
|
||||
mock_openai_cls.return_value = mock_instance
|
||||
|
||||
result = EmbeddingsSingleton.get_instance("openai_text-embedding-ada-002")
|
||||
assert result is mock_instance
|
||||
|
||||
@patch("application.vectorstore.base.OpenAIEmbeddings")
|
||||
def test_singleton_returns_same_instance(self, mock_openai_cls):
|
||||
mock_instance = Mock()
|
||||
mock_openai_cls.return_value = mock_instance
|
||||
|
||||
r1 = EmbeddingsSingleton.get_instance("openai_text-embedding-ada-002")
|
||||
r2 = EmbeddingsSingleton.get_instance("openai_text-embedding-ada-002")
|
||||
assert r1 is r2
|
||||
mock_openai_cls.assert_called_once()
|
||||
|
||||
@patch("application.vectorstore.base._get_embeddings_wrapper")
|
||||
def test_get_instance_huggingface(self, mock_get_wrapper):
|
||||
mock_wrapper_cls = Mock()
|
||||
mock_instance = Mock()
|
||||
mock_wrapper_cls.return_value = mock_instance
|
||||
mock_get_wrapper.return_value = mock_wrapper_cls
|
||||
|
||||
result = EmbeddingsSingleton.get_instance(
|
||||
"huggingface_sentence-transformers/all-mpnet-base-v2"
|
||||
)
|
||||
assert result is mock_instance
|
||||
|
||||
@patch("application.vectorstore.base._get_embeddings_wrapper")
|
||||
def test_get_instance_unknown_falls_back_to_wrapper(self, mock_get_wrapper):
|
||||
mock_wrapper_cls = Mock()
|
||||
mock_instance = Mock()
|
||||
mock_wrapper_cls.return_value = mock_instance
|
||||
mock_get_wrapper.return_value = mock_wrapper_cls
|
||||
|
||||
result = EmbeddingsSingleton.get_instance("custom_model_name")
|
||||
mock_wrapper_cls.assert_called_once_with("custom_model_name")
|
||||
assert result is mock_instance
|
||||
|
||||
@patch("application.vectorstore.base.settings")
|
||||
def test_get_instance_uses_remote_when_base_url_set(self, mock_settings):
|
||||
"""Direct callers (GraphRAG, semantic chunking) must route to the
|
||||
remote embeddings API instead of loading a local model."""
|
||||
mock_settings.EMBEDDINGS_BASE_URL = "http://remote:8080"
|
||||
mock_settings.EMBEDDINGS_KEY = "sk-remote"
|
||||
|
||||
result = EmbeddingsSingleton.get_instance("embeddinggemma", "sk-remote")
|
||||
|
||||
assert isinstance(result, RemoteEmbeddings)
|
||||
assert result.api_url == "http://remote:8080"
|
||||
assert result.model_name == "embeddinggemma"
|
||||
assert result.headers["Authorization"] == "Bearer sk-remote"
|
||||
|
||||
@patch("application.vectorstore.base.settings")
|
||||
def test_get_instance_remote_falls_back_to_settings_key(self, mock_settings):
|
||||
"""When no key is passed, the remote dispatch uses EMBEDDINGS_KEY."""
|
||||
mock_settings.EMBEDDINGS_BASE_URL = "http://remote:8080"
|
||||
mock_settings.EMBEDDINGS_KEY = "sk-from-settings"
|
||||
|
||||
result = EmbeddingsSingleton.get_instance("embeddinggemma")
|
||||
|
||||
assert isinstance(result, RemoteEmbeddings)
|
||||
assert result.headers["Authorization"] == "Bearer sk-from-settings"
|
||||
|
||||
|
||||
# --- BaseVectorStore ---
|
||||
|
||||
|
||||
class ConcreteVectorStore(BaseVectorStore):
|
||||
"""Concrete implementation for testing base class methods."""
|
||||
|
||||
def search(self, *args, **kwargs):
|
||||
return []
|
||||
|
||||
def add_texts(self, texts, metadatas=None, *args, **kwargs):
|
||||
return []
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestBaseVectorStore:
|
||||
def setup_method(self):
|
||||
EmbeddingsSingleton._instances = {}
|
||||
|
||||
def test_default_methods_are_noop(self):
|
||||
store = ConcreteVectorStore()
|
||||
assert store.delete_index() is None
|
||||
assert store.save_local() is None
|
||||
assert store.get_chunks() is None
|
||||
assert store.add_chunk("text") is None
|
||||
assert store.delete_chunk("id") is None
|
||||
|
||||
@patch("application.vectorstore.base.settings")
|
||||
def test_is_azure_configured_true(self, mock_settings):
|
||||
mock_settings.OPENAI_API_BASE = "https://azure.openai.com"
|
||||
mock_settings.OPENAI_API_VERSION = "2023-05-15"
|
||||
mock_settings.AZURE_DEPLOYMENT_NAME = "my-deploy"
|
||||
|
||||
store = ConcreteVectorStore()
|
||||
assert store.is_azure_configured()
|
||||
|
||||
@patch("application.vectorstore.base.settings")
|
||||
def test_is_azure_configured_false(self, mock_settings):
|
||||
mock_settings.OPENAI_API_BASE = None
|
||||
mock_settings.OPENAI_API_VERSION = None
|
||||
mock_settings.AZURE_DEPLOYMENT_NAME = None
|
||||
|
||||
store = ConcreteVectorStore()
|
||||
assert not store.is_azure_configured()
|
||||
|
||||
@patch("application.vectorstore.base.settings")
|
||||
def test_get_embeddings_remote(self, mock_settings):
|
||||
mock_settings.EMBEDDINGS_BASE_URL = "http://remote:8080"
|
||||
|
||||
store = ConcreteVectorStore()
|
||||
result = store._get_embeddings("model-name", "api-key")
|
||||
|
||||
assert isinstance(result, RemoteEmbeddings)
|
||||
assert result.api_url == "http://remote:8080"
|
||||
|
||||
@patch("application.vectorstore.base.settings")
|
||||
@patch("application.vectorstore.base.EmbeddingsSingleton.get_instance")
|
||||
def test_get_embeddings_openai(self, mock_get_instance, mock_settings):
|
||||
mock_settings.EMBEDDINGS_BASE_URL = None
|
||||
mock_settings.OPENAI_API_BASE = None
|
||||
mock_settings.OPENAI_API_VERSION = None
|
||||
mock_settings.AZURE_DEPLOYMENT_NAME = None
|
||||
|
||||
mock_emb = Mock()
|
||||
mock_get_instance.return_value = mock_emb
|
||||
|
||||
store = ConcreteVectorStore()
|
||||
result = store._get_embeddings("openai_text-embedding-ada-002", "sk-key")
|
||||
assert result is mock_emb
|
||||
|
||||
@patch("application.vectorstore.base.settings")
|
||||
@patch("application.vectorstore.base.EmbeddingsSingleton.get_instance")
|
||||
def test_get_embeddings_openai_azure(self, mock_get_instance, mock_settings):
|
||||
mock_settings.EMBEDDINGS_BASE_URL = None
|
||||
mock_settings.OPENAI_API_BASE = "https://azure.openai.com"
|
||||
mock_settings.OPENAI_API_VERSION = "2023-05-15"
|
||||
mock_settings.AZURE_DEPLOYMENT_NAME = "deploy"
|
||||
mock_settings.AZURE_EMBEDDINGS_DEPLOYMENT_NAME = "embed-deploy"
|
||||
|
||||
mock_emb = Mock()
|
||||
mock_get_instance.return_value = mock_emb
|
||||
|
||||
store = ConcreteVectorStore()
|
||||
result = store._get_embeddings("openai_text-embedding-ada-002", "sk-key")
|
||||
assert result is mock_emb
|
||||
|
||||
@patch("application.vectorstore.base.settings")
|
||||
@patch("application.vectorstore.base.EmbeddingsSingleton.get_instance")
|
||||
@patch("os.path.exists", return_value=False)
|
||||
def test_get_embeddings_huggingface_no_local_model(
|
||||
self, mock_exists, mock_get_instance, mock_settings
|
||||
):
|
||||
mock_settings.EMBEDDINGS_BASE_URL = None
|
||||
mock_emb = Mock()
|
||||
mock_get_instance.return_value = mock_emb
|
||||
|
||||
store = ConcreteVectorStore()
|
||||
result = store._get_embeddings(
|
||||
"huggingface_sentence-transformers/all-mpnet-base-v2"
|
||||
)
|
||||
assert result is mock_emb
|
||||
|
||||
@patch("application.vectorstore.base.settings")
|
||||
@patch("application.vectorstore.base.EmbeddingsSingleton.get_instance")
|
||||
@patch("os.path.exists")
|
||||
def test_get_embeddings_huggingface_local_model(
|
||||
self, mock_exists, mock_get_instance, mock_settings
|
||||
):
|
||||
mock_settings.EMBEDDINGS_BASE_URL = None
|
||||
mock_exists.side_effect = lambda p: p == "/app/models/all-mpnet-base-v2"
|
||||
mock_emb = Mock()
|
||||
mock_get_instance.return_value = mock_emb
|
||||
|
||||
store = ConcreteVectorStore()
|
||||
result = store._get_embeddings(
|
||||
"huggingface_sentence-transformers/all-mpnet-base-v2"
|
||||
)
|
||||
assert result is mock_emb
|
||||
mock_get_instance.assert_called_with("/app/models/all-mpnet-base-v2")
|
||||
|
||||
@patch("application.vectorstore.base.settings")
|
||||
@patch("application.vectorstore.base.EmbeddingsSingleton.get_instance")
|
||||
def test_get_embeddings_generic(self, mock_get_instance, mock_settings):
|
||||
mock_settings.EMBEDDINGS_BASE_URL = None
|
||||
mock_emb = Mock()
|
||||
mock_get_instance.return_value = mock_emb
|
||||
|
||||
store = ConcreteVectorStore()
|
||||
result = store._get_embeddings("some_custom_embedding")
|
||||
assert result is mock_emb
|
||||
mock_get_instance.assert_called_with("some_custom_embedding")
|
||||
@@ -0,0 +1,39 @@
|
||||
import pytest
|
||||
from application.vectorstore.document_class import Document
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestDocument:
|
||||
def test_create_document(self):
|
||||
doc = Document(page_content="hello world", metadata={"source": "test"})
|
||||
assert doc.page_content == "hello world"
|
||||
assert doc.metadata == {"source": "test"}
|
||||
|
||||
def test_document_is_string(self):
|
||||
doc = Document(page_content="hello world", metadata={})
|
||||
assert isinstance(doc, str)
|
||||
assert str(doc) == "hello world"
|
||||
|
||||
def test_document_string_equality(self):
|
||||
doc = Document(page_content="hello", metadata={"k": "v"})
|
||||
assert doc == "hello"
|
||||
|
||||
def test_document_empty_metadata(self):
|
||||
doc = Document(page_content="text", metadata={})
|
||||
assert doc.metadata == {}
|
||||
|
||||
def test_document_empty_content(self):
|
||||
doc = Document(page_content="", metadata={"a": 1})
|
||||
assert doc.page_content == ""
|
||||
assert doc == ""
|
||||
|
||||
def test_document_preserves_complex_metadata(self):
|
||||
meta = {"source": "file.txt", "page": 3, "nested": {"key": "val"}}
|
||||
doc = Document(page_content="content", metadata=meta)
|
||||
assert doc.metadata["nested"]["key"] == "val"
|
||||
|
||||
def test_document_string_operations(self):
|
||||
doc = Document(page_content="hello world", metadata={})
|
||||
assert doc.upper() == "HELLO WORLD"
|
||||
assert doc.split() == ["hello", "world"]
|
||||
assert "world" in doc
|
||||
@@ -0,0 +1,292 @@
|
||||
from unittest.mock import MagicMock, Mock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
def _make_es_store(source_id="test-source"):
|
||||
"""Helper to create an ElasticsearchStore with mocked deps."""
|
||||
# Reset class-level connection
|
||||
from application.vectorstore.elasticsearch import ElasticsearchStore
|
||||
|
||||
ElasticsearchStore._es_connection = None
|
||||
|
||||
with patch(
|
||||
"application.vectorstore.elasticsearch.settings"
|
||||
) as mock_settings, patch.dict(
|
||||
"sys.modules", {"elasticsearch": MagicMock(), "elasticsearch.helpers": MagicMock()}
|
||||
):
|
||||
mock_settings.ELASTIC_URL = "http://localhost:9200"
|
||||
mock_settings.ELASTIC_USERNAME = "elastic"
|
||||
mock_settings.ELASTIC_PASSWORD = "password"
|
||||
mock_settings.ELASTIC_CLOUD_ID = None
|
||||
mock_settings.ELASTIC_INDEX = "test_index"
|
||||
mock_settings.EMBEDDINGS_NAME = "test_model"
|
||||
|
||||
import elasticsearch
|
||||
|
||||
mock_es = MagicMock()
|
||||
elasticsearch.Elasticsearch.return_value = mock_es
|
||||
|
||||
store = ElasticsearchStore(
|
||||
source_id=source_id,
|
||||
embeddings_key="key",
|
||||
index_name="test_index",
|
||||
)
|
||||
|
||||
return store, mock_es, mock_settings
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestElasticsearchStoreInit:
|
||||
def test_source_id_cleaned(self):
|
||||
store, _, _ = _make_es_store(source_id="application/indexes/abc123/")
|
||||
assert store.source_id == "abc123"
|
||||
|
||||
def test_init_with_url(self):
|
||||
store, mock_es, _ = _make_es_store()
|
||||
assert store.docsearch is mock_es
|
||||
assert store.index_name == "test_index"
|
||||
|
||||
def test_init_with_cloud_id(self):
|
||||
from application.vectorstore.elasticsearch import ElasticsearchStore
|
||||
|
||||
ElasticsearchStore._es_connection = None
|
||||
|
||||
with patch(
|
||||
"application.vectorstore.elasticsearch.settings"
|
||||
) as mock_settings, patch.dict(
|
||||
"sys.modules", {"elasticsearch": MagicMock()}
|
||||
):
|
||||
mock_settings.ELASTIC_URL = None
|
||||
mock_settings.ELASTIC_CLOUD_ID = "my-cloud-id"
|
||||
mock_settings.ELASTIC_USERNAME = "user"
|
||||
mock_settings.ELASTIC_PASSWORD = "pass"
|
||||
mock_settings.ELASTIC_INDEX = "idx"
|
||||
mock_settings.EMBEDDINGS_NAME = "model"
|
||||
|
||||
store = ElasticsearchStore(
|
||||
source_id="src", embeddings_key="k", index_name="idx"
|
||||
)
|
||||
assert store.docsearch is not None
|
||||
|
||||
def test_init_no_url_no_cloud_id_raises(self):
|
||||
from application.vectorstore.elasticsearch import ElasticsearchStore
|
||||
|
||||
ElasticsearchStore._es_connection = None
|
||||
|
||||
with patch(
|
||||
"application.vectorstore.elasticsearch.settings"
|
||||
) as mock_settings, patch.dict(
|
||||
"sys.modules", {"elasticsearch": MagicMock()}
|
||||
):
|
||||
mock_settings.ELASTIC_URL = None
|
||||
mock_settings.ELASTIC_CLOUD_ID = None
|
||||
mock_settings.ELASTIC_INDEX = "idx"
|
||||
mock_settings.EMBEDDINGS_NAME = "model"
|
||||
|
||||
with pytest.raises(ValueError, match="provide either"):
|
||||
ElasticsearchStore(source_id="src", embeddings_key="k")
|
||||
|
||||
def test_reuses_class_connection(self):
|
||||
from application.vectorstore.elasticsearch import ElasticsearchStore
|
||||
|
||||
ElasticsearchStore._es_connection = None
|
||||
|
||||
with patch(
|
||||
"application.vectorstore.elasticsearch.settings"
|
||||
) as mock_settings, patch.dict(
|
||||
"sys.modules", {"elasticsearch": MagicMock()}
|
||||
):
|
||||
mock_settings.ELASTIC_URL = "http://localhost:9200"
|
||||
mock_settings.ELASTIC_USERNAME = "user"
|
||||
mock_settings.ELASTIC_PASSWORD = "pass"
|
||||
mock_settings.ELASTIC_CLOUD_ID = None
|
||||
mock_settings.ELASTIC_INDEX = "idx"
|
||||
mock_settings.EMBEDDINGS_NAME = "model"
|
||||
|
||||
import elasticsearch
|
||||
|
||||
mock_es = MagicMock()
|
||||
elasticsearch.Elasticsearch.return_value = mock_es
|
||||
|
||||
store1 = ElasticsearchStore(source_id="src1", embeddings_key="k")
|
||||
store2 = ElasticsearchStore(source_id="src2", embeddings_key="k")
|
||||
|
||||
assert store1.docsearch is store2.docsearch
|
||||
elasticsearch.Elasticsearch.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestElasticsearchStoreSearch:
|
||||
def test_search_builds_query(self):
|
||||
store, mock_es, mock_settings = _make_es_store()
|
||||
|
||||
mock_emb = Mock()
|
||||
mock_emb.embed_query = Mock(return_value=[0.1, 0.2, 0.3])
|
||||
|
||||
with patch.object(store, "_get_embeddings", return_value=mock_emb):
|
||||
mock_es.search.return_value = {
|
||||
"hits": {
|
||||
"hits": [
|
||||
{
|
||||
"_source": {
|
||||
"text": "doc1",
|
||||
"metadata": {"source": "file.txt"},
|
||||
}
|
||||
},
|
||||
{
|
||||
"_source": {
|
||||
"text": "doc2",
|
||||
"metadata": {"source": "file2.txt"},
|
||||
}
|
||||
},
|
||||
]
|
||||
}
|
||||
}
|
||||
|
||||
results = store.search("query", k=2)
|
||||
|
||||
assert len(results) == 2
|
||||
assert results[0].page_content == "doc1"
|
||||
assert results[1].metadata == {"source": "file2.txt"}
|
||||
|
||||
def test_search_empty_results(self):
|
||||
store, mock_es, _ = _make_es_store()
|
||||
|
||||
mock_emb = Mock()
|
||||
mock_emb.embed_query = Mock(return_value=[0.1])
|
||||
|
||||
with patch.object(store, "_get_embeddings", return_value=mock_emb):
|
||||
mock_es.search.return_value = {"hits": {"hits": []}}
|
||||
results = store.search("query")
|
||||
|
||||
assert results == []
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestElasticsearchStoreAddTexts:
|
||||
def test_add_texts(self):
|
||||
store, mock_es, mock_settings = _make_es_store()
|
||||
|
||||
mock_emb = Mock()
|
||||
mock_emb.embed_documents = Mock(return_value=[[0.1, 0.2], [0.3, 0.4]])
|
||||
|
||||
mock_bulk = Mock(return_value=(2, 0))
|
||||
mock_helpers = MagicMock()
|
||||
mock_helpers.bulk = mock_bulk
|
||||
|
||||
with patch.object(
|
||||
store, "_get_embeddings", return_value=mock_emb
|
||||
), patch.object(
|
||||
store, "_create_index_if_not_exists"
|
||||
), patch.dict(
|
||||
"sys.modules", {"elasticsearch.helpers": mock_helpers}
|
||||
):
|
||||
ids = store.add_texts(
|
||||
["text1", "text2"],
|
||||
metadatas=[{"a": 1}, {"b": 2}],
|
||||
)
|
||||
|
||||
assert len(ids) == 2
|
||||
|
||||
def test_add_texts_empty_raises(self):
|
||||
"""Empty texts causes IndexError because code accesses vectors[0] unconditionally."""
|
||||
store, _, _ = _make_es_store()
|
||||
|
||||
mock_emb = Mock()
|
||||
mock_emb.embed_documents = Mock(return_value=[])
|
||||
|
||||
with patch.object(store, "_get_embeddings", return_value=mock_emb):
|
||||
with pytest.raises(IndexError):
|
||||
store.add_texts([], metadatas=[])
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestElasticsearchStoreDeleteIndex:
|
||||
def test_delete_index_calls_delete_by_query(self):
|
||||
store, mock_es, _ = _make_es_store(source_id="src1")
|
||||
|
||||
store.delete_index()
|
||||
|
||||
mock_es.delete_by_query.assert_called_once_with(
|
||||
index="test_index",
|
||||
query={"match": {"metadata.source_id.keyword": "src1"}},
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestElasticsearchStoreIndex:
|
||||
def test_index_returns_mapping(self):
|
||||
store, _, _ = _make_es_store()
|
||||
|
||||
mapping = store.index(dims_length=768)
|
||||
|
||||
assert mapping["mappings"]["properties"]["vector"]["type"] == "dense_vector"
|
||||
assert mapping["mappings"]["properties"]["vector"]["dims"] == 768
|
||||
assert mapping["mappings"]["properties"]["vector"]["similarity"] == "cosine"
|
||||
|
||||
def test_create_index_if_not_exists_existing(self):
|
||||
store, mock_es, _ = _make_es_store()
|
||||
mock_es.indices.exists.return_value = True
|
||||
|
||||
store._create_index_if_not_exists("test_index", 768)
|
||||
|
||||
mock_es.indices.create.assert_not_called()
|
||||
|
||||
def test_create_index_if_not_exists_new(self):
|
||||
store, mock_es, _ = _make_es_store()
|
||||
mock_es.indices.exists.return_value = False
|
||||
|
||||
store._create_index_if_not_exists("test_index", 768)
|
||||
|
||||
mock_es.indices.create.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestElasticsearchStoreConnectToElasticsearch:
|
||||
def test_connect_with_url(self):
|
||||
from application.vectorstore.elasticsearch import ElasticsearchStore
|
||||
|
||||
with patch.dict("sys.modules", {"elasticsearch": MagicMock()}):
|
||||
import elasticsearch
|
||||
|
||||
mock_es = MagicMock()
|
||||
elasticsearch.Elasticsearch.return_value = mock_es
|
||||
|
||||
result = ElasticsearchStore.connect_to_elasticsearch(
|
||||
es_url="http://localhost:9200",
|
||||
username="user",
|
||||
password="pass",
|
||||
)
|
||||
assert result is mock_es
|
||||
|
||||
def test_connect_with_both_raises(self):
|
||||
from application.vectorstore.elasticsearch import ElasticsearchStore
|
||||
|
||||
with patch.dict("sys.modules", {"elasticsearch": MagicMock()}):
|
||||
with pytest.raises(ValueError, match="Both es_url and cloud_id"):
|
||||
ElasticsearchStore.connect_to_elasticsearch(
|
||||
es_url="http://localhost", cloud_id="cloud-123"
|
||||
)
|
||||
|
||||
def test_connect_with_neither_raises(self):
|
||||
from application.vectorstore.elasticsearch import ElasticsearchStore
|
||||
|
||||
with patch.dict("sys.modules", {"elasticsearch": MagicMock()}):
|
||||
with pytest.raises(ValueError, match="provide either"):
|
||||
ElasticsearchStore.connect_to_elasticsearch()
|
||||
|
||||
def test_connect_with_api_key(self):
|
||||
from application.vectorstore.elasticsearch import ElasticsearchStore
|
||||
|
||||
with patch.dict("sys.modules", {"elasticsearch": MagicMock()}):
|
||||
import elasticsearch
|
||||
|
||||
mock_es = MagicMock()
|
||||
elasticsearch.Elasticsearch.return_value = mock_es
|
||||
|
||||
result = ElasticsearchStore.connect_to_elasticsearch(
|
||||
es_url="http://localhost:9200",
|
||||
api_key="my-api-key",
|
||||
)
|
||||
assert result is mock_es
|
||||
@@ -0,0 +1,140 @@
|
||||
from unittest.mock import MagicMock, Mock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestEmbeddingsWrapper:
|
||||
@patch("application.vectorstore.embeddings_local.SentenceTransformer")
|
||||
def test_init_success(self, mock_st_cls):
|
||||
mock_model = MagicMock()
|
||||
mock_model._first_module.return_value = MagicMock()
|
||||
mock_model.get_sentence_embedding_dimension.return_value = 768
|
||||
mock_st_cls.return_value = mock_model
|
||||
|
||||
from application.vectorstore.embeddings_local import EmbeddingsWrapper
|
||||
|
||||
wrapper = EmbeddingsWrapper("test-model")
|
||||
|
||||
mock_st_cls.assert_called_once()
|
||||
assert wrapper.dimension == 768
|
||||
|
||||
@patch("application.vectorstore.embeddings_local.SentenceTransformer")
|
||||
def test_init_failure(self, mock_st_cls):
|
||||
mock_st_cls.side_effect = Exception("model not found")
|
||||
|
||||
from application.vectorstore.embeddings_local import EmbeddingsWrapper
|
||||
|
||||
with pytest.raises(Exception, match="model not found"):
|
||||
EmbeddingsWrapper("bad-model")
|
||||
|
||||
@patch("application.vectorstore.embeddings_local.SentenceTransformer")
|
||||
def test_init_none_model(self, mock_st_cls):
|
||||
mock_st_cls.return_value = None
|
||||
|
||||
from application.vectorstore.embeddings_local import EmbeddingsWrapper
|
||||
|
||||
with pytest.raises((ValueError, AttributeError)):
|
||||
EmbeddingsWrapper("bad-model")
|
||||
|
||||
@patch("application.vectorstore.embeddings_local.SentenceTransformer")
|
||||
def test_init_null_first_module(self, mock_st_cls):
|
||||
mock_model = MagicMock()
|
||||
mock_model._first_module.return_value = None
|
||||
mock_st_cls.return_value = mock_model
|
||||
|
||||
from application.vectorstore.embeddings_local import EmbeddingsWrapper
|
||||
|
||||
with pytest.raises(ValueError, match="failed to load properly"):
|
||||
EmbeddingsWrapper("bad-model")
|
||||
|
||||
@patch("application.vectorstore.embeddings_local.SentenceTransformer")
|
||||
def test_embed_query(self, mock_st_cls):
|
||||
mock_model = MagicMock()
|
||||
mock_model._first_module.return_value = MagicMock()
|
||||
mock_model.get_sentence_embedding_dimension.return_value = 3
|
||||
mock_model.encode.return_value = MagicMock(tolist=Mock(return_value=[0.1, 0.2, 0.3]))
|
||||
mock_st_cls.return_value = mock_model
|
||||
|
||||
from application.vectorstore.embeddings_local import EmbeddingsWrapper
|
||||
|
||||
wrapper = EmbeddingsWrapper("model")
|
||||
result = wrapper.embed_query("hello world")
|
||||
|
||||
mock_model.encode.assert_called_once_with("hello world")
|
||||
assert result == [0.1, 0.2, 0.3]
|
||||
|
||||
@patch("application.vectorstore.embeddings_local.SentenceTransformer")
|
||||
def test_embed_documents(self, mock_st_cls):
|
||||
mock_model = MagicMock()
|
||||
mock_model._first_module.return_value = MagicMock()
|
||||
mock_model.get_sentence_embedding_dimension.return_value = 3
|
||||
mock_model.encode.return_value = MagicMock(
|
||||
tolist=Mock(return_value=[[0.1, 0.2], [0.3, 0.4]])
|
||||
)
|
||||
mock_st_cls.return_value = mock_model
|
||||
|
||||
from application.vectorstore.embeddings_local import EmbeddingsWrapper
|
||||
|
||||
wrapper = EmbeddingsWrapper("model")
|
||||
result = wrapper.embed_documents(["doc1", "doc2"])
|
||||
|
||||
mock_model.encode.assert_called_with(["doc1", "doc2"])
|
||||
assert result == [[0.1, 0.2], [0.3, 0.4]]
|
||||
|
||||
@patch("application.vectorstore.embeddings_local.SentenceTransformer")
|
||||
def test_call_with_string(self, mock_st_cls):
|
||||
mock_model = MagicMock()
|
||||
mock_model._first_module.return_value = MagicMock()
|
||||
mock_model.get_sentence_embedding_dimension.return_value = 3
|
||||
mock_model.encode.return_value = MagicMock(tolist=Mock(return_value=[0.1]))
|
||||
mock_st_cls.return_value = mock_model
|
||||
|
||||
from application.vectorstore.embeddings_local import EmbeddingsWrapper
|
||||
|
||||
wrapper = EmbeddingsWrapper("model")
|
||||
result = wrapper("hello")
|
||||
assert result == [0.1]
|
||||
|
||||
@patch("application.vectorstore.embeddings_local.SentenceTransformer")
|
||||
def test_call_with_list(self, mock_st_cls):
|
||||
mock_model = MagicMock()
|
||||
mock_model._first_module.return_value = MagicMock()
|
||||
mock_model.get_sentence_embedding_dimension.return_value = 3
|
||||
mock_model.encode.return_value = MagicMock(
|
||||
tolist=Mock(return_value=[[0.1], [0.2]])
|
||||
)
|
||||
mock_st_cls.return_value = mock_model
|
||||
|
||||
from application.vectorstore.embeddings_local import EmbeddingsWrapper
|
||||
|
||||
wrapper = EmbeddingsWrapper("model")
|
||||
result = wrapper(["a", "b"])
|
||||
assert result == [[0.1], [0.2]]
|
||||
|
||||
@patch("application.vectorstore.embeddings_local.SentenceTransformer")
|
||||
def test_call_with_invalid_type(self, mock_st_cls):
|
||||
mock_model = MagicMock()
|
||||
mock_model._first_module.return_value = MagicMock()
|
||||
mock_model.get_sentence_embedding_dimension.return_value = 3
|
||||
mock_st_cls.return_value = mock_model
|
||||
|
||||
from application.vectorstore.embeddings_local import EmbeddingsWrapper
|
||||
|
||||
wrapper = EmbeddingsWrapper("model")
|
||||
with pytest.raises(ValueError, match="Input must be a string or a list"):
|
||||
wrapper(123)
|
||||
|
||||
@patch("application.vectorstore.embeddings_local.SentenceTransformer")
|
||||
def test_trust_remote_code_default(self, mock_st_cls):
|
||||
mock_model = MagicMock()
|
||||
mock_model._first_module.return_value = MagicMock()
|
||||
mock_model.get_sentence_embedding_dimension.return_value = 768
|
||||
mock_st_cls.return_value = mock_model
|
||||
|
||||
from application.vectorstore.embeddings_local import EmbeddingsWrapper
|
||||
|
||||
EmbeddingsWrapper("model")
|
||||
|
||||
call_kwargs = mock_st_cls.call_args[1]
|
||||
assert call_kwargs["trust_remote_code"] is True
|
||||
@@ -0,0 +1,564 @@
|
||||
import io
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_embeddings():
|
||||
emb = Mock()
|
||||
emb.embed_query = Mock(return_value=[0.1, 0.2, 0.3])
|
||||
emb.embed_documents = Mock(return_value=[[0.1, 0.2, 0.3]])
|
||||
emb.dimension = 3
|
||||
return emb
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_storage():
|
||||
storage = Mock()
|
||||
storage.file_exists = Mock(return_value=True)
|
||||
storage.get_file = Mock(return_value=io.BytesIO(b"fake data"))
|
||||
storage.save_file = Mock()
|
||||
return storage
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_docsearch():
|
||||
ds = Mock()
|
||||
ds.similarity_search = Mock(return_value=[])
|
||||
ds.add_texts = Mock(return_value=["id1"])
|
||||
ds.add_documents = Mock(return_value=["id1"])
|
||||
ds.save_local = Mock()
|
||||
ds.delete = Mock()
|
||||
ds.index = Mock()
|
||||
ds.index.d = 3
|
||||
ds.docstore = Mock()
|
||||
ds.docstore._dict = {
|
||||
"doc1": Mock(page_content="text1", metadata={"source": "a"}),
|
||||
"doc2": Mock(page_content="text2", metadata={"source": "b"}),
|
||||
}
|
||||
return ds
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestFaissStoreInit:
|
||||
@patch("application.vectorstore.faiss.StorageCreator")
|
||||
@patch("application.vectorstore.faiss.FAISS")
|
||||
@patch.object(
|
||||
__import__("application.vectorstore.base", fromlist=["BaseVectorStore"]).BaseVectorStore,
|
||||
"_get_embeddings",
|
||||
)
|
||||
@patch("application.vectorstore.faiss.settings")
|
||||
def test_init_with_docs(self, mock_settings, mock_get_emb, mock_faiss, mock_storage_creator):
|
||||
mock_settings.EMBEDDINGS_NAME = "test_model"
|
||||
mock_emb = Mock(dimension=3)
|
||||
mock_get_emb.return_value = mock_emb
|
||||
mock_ds = Mock()
|
||||
mock_ds.index = Mock(d=3)
|
||||
mock_faiss.from_documents.return_value = mock_ds
|
||||
mock_storage_creator.get_storage.return_value = Mock()
|
||||
|
||||
from application.vectorstore.faiss import FaissStore
|
||||
|
||||
store = FaissStore(source_id="test", embeddings_key="key", docs_init=[Mock()])
|
||||
mock_faiss.from_documents.assert_called_once()
|
||||
assert store.docsearch is mock_ds
|
||||
|
||||
@patch("application.vectorstore.faiss.StorageCreator")
|
||||
@patch("application.vectorstore.faiss.FAISS")
|
||||
@patch.object(
|
||||
__import__("application.vectorstore.base", fromlist=["BaseVectorStore"]).BaseVectorStore,
|
||||
"_get_embeddings",
|
||||
)
|
||||
@patch("application.vectorstore.faiss.settings")
|
||||
def test_init_missing_index_files(
|
||||
self, mock_settings, mock_get_emb, mock_faiss, mock_storage_creator
|
||||
):
|
||||
mock_settings.EMBEDDINGS_NAME = "test_model"
|
||||
mock_emb = Mock(dimension=3)
|
||||
mock_get_emb.return_value = mock_emb
|
||||
mock_storage = Mock()
|
||||
mock_storage.file_exists.return_value = False
|
||||
mock_storage_creator.get_storage.return_value = mock_storage
|
||||
|
||||
from application.vectorstore.faiss import FaissStore
|
||||
|
||||
with pytest.raises(Exception, match="Error loading FAISS index"):
|
||||
FaissStore(source_id="test", embeddings_key="key")
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestFaissStoreSearch:
|
||||
@patch("application.vectorstore.faiss.StorageCreator")
|
||||
@patch("application.vectorstore.faiss.FAISS")
|
||||
@patch.object(
|
||||
__import__("application.vectorstore.base", fromlist=["BaseVectorStore"]).BaseVectorStore,
|
||||
"_get_embeddings",
|
||||
)
|
||||
@patch("application.vectorstore.faiss.settings")
|
||||
def test_search_delegates_to_docsearch(
|
||||
self, mock_settings, mock_get_emb, mock_faiss, mock_storage_creator
|
||||
):
|
||||
mock_settings.EMBEDDINGS_NAME = "test_model"
|
||||
mock_emb = Mock(dimension=3)
|
||||
mock_get_emb.return_value = mock_emb
|
||||
mock_ds = Mock()
|
||||
mock_ds.index = Mock(d=3)
|
||||
mock_ds.similarity_search.return_value = ["doc1"]
|
||||
mock_faiss.from_documents.return_value = mock_ds
|
||||
mock_storage_creator.get_storage.return_value = Mock()
|
||||
|
||||
from application.vectorstore.faiss import FaissStore
|
||||
|
||||
store = FaissStore(source_id="t", embeddings_key="k", docs_init=[Mock()])
|
||||
result = store.search("query", k=5)
|
||||
mock_ds.similarity_search.assert_called_once_with("query", k=5)
|
||||
assert result == ["doc1"]
|
||||
|
||||
@patch("application.vectorstore.faiss.StorageCreator")
|
||||
@patch("application.vectorstore.faiss.FAISS")
|
||||
@patch.object(
|
||||
__import__("application.vectorstore.base", fromlist=["BaseVectorStore"]).BaseVectorStore,
|
||||
"_get_embeddings",
|
||||
)
|
||||
@patch("application.vectorstore.faiss.settings")
|
||||
def test_search_ignores_score_threshold(
|
||||
self, mock_settings, mock_get_emb, mock_faiss, mock_storage_creator
|
||||
):
|
||||
# FAISS has no relevance-threshold knob; the per-source score_threshold
|
||||
# must be safely dropped, not forwarded (which would crash langchain).
|
||||
mock_settings.EMBEDDINGS_NAME = "test_model"
|
||||
mock_get_emb.return_value = Mock(dimension=3)
|
||||
mock_ds = Mock()
|
||||
mock_ds.index = Mock(d=3)
|
||||
mock_ds.similarity_search.return_value = ["doc1"]
|
||||
mock_faiss.from_documents.return_value = mock_ds
|
||||
mock_storage_creator.get_storage.return_value = Mock()
|
||||
|
||||
from application.vectorstore.faiss import FaissStore
|
||||
|
||||
store = FaissStore(source_id="t", embeddings_key="k", docs_init=[Mock()])
|
||||
result = store.search("query", k=5, score_threshold=0.9)
|
||||
# score_threshold is stripped before the forward.
|
||||
mock_ds.similarity_search.assert_called_once_with("query", k=5)
|
||||
assert result == ["doc1"]
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestFaissStoreAddTexts:
|
||||
@patch("application.vectorstore.faiss.StorageCreator")
|
||||
@patch("application.vectorstore.faiss.FAISS")
|
||||
@patch.object(
|
||||
__import__("application.vectorstore.base", fromlist=["BaseVectorStore"]).BaseVectorStore,
|
||||
"_get_embeddings",
|
||||
)
|
||||
@patch("application.vectorstore.faiss.settings")
|
||||
def test_add_texts_delegates(
|
||||
self, mock_settings, mock_get_emb, mock_faiss, mock_storage_creator
|
||||
):
|
||||
mock_settings.EMBEDDINGS_NAME = "test_model"
|
||||
mock_emb = Mock(dimension=3)
|
||||
mock_get_emb.return_value = mock_emb
|
||||
mock_ds = Mock()
|
||||
mock_ds.index = Mock(d=3)
|
||||
mock_ds.add_texts.return_value = ["id1", "id2"]
|
||||
mock_faiss.from_documents.return_value = mock_ds
|
||||
mock_storage_creator.get_storage.return_value = Mock()
|
||||
|
||||
from application.vectorstore.faiss import FaissStore
|
||||
|
||||
store = FaissStore(source_id="t", embeddings_key="k", docs_init=[Mock()])
|
||||
result = store.add_texts(["text1", "text2"])
|
||||
assert result == ["id1", "id2"]
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestFaissStoreGetChunks:
|
||||
@patch("application.vectorstore.faiss.StorageCreator")
|
||||
@patch("application.vectorstore.faiss.FAISS")
|
||||
@patch.object(
|
||||
__import__("application.vectorstore.base", fromlist=["BaseVectorStore"]).BaseVectorStore,
|
||||
"_get_embeddings",
|
||||
)
|
||||
@patch("application.vectorstore.faiss.settings")
|
||||
def test_get_chunks(self, mock_settings, mock_get_emb, mock_faiss, mock_storage_creator):
|
||||
mock_settings.EMBEDDINGS_NAME = "test_model"
|
||||
mock_emb = Mock(dimension=3)
|
||||
mock_get_emb.return_value = mock_emb
|
||||
|
||||
doc1 = Mock(page_content="text1", metadata={"source": "a"})
|
||||
doc2 = Mock(page_content="text2", metadata={"source": "b"})
|
||||
|
||||
mock_ds = Mock()
|
||||
mock_ds.index = Mock(d=3)
|
||||
mock_ds.docstore._dict = {"id1": doc1, "id2": doc2}
|
||||
mock_faiss.from_documents.return_value = mock_ds
|
||||
mock_storage_creator.get_storage.return_value = Mock()
|
||||
|
||||
from application.vectorstore.faiss import FaissStore
|
||||
|
||||
store = FaissStore(source_id="t", embeddings_key="k", docs_init=[Mock()])
|
||||
chunks = store.get_chunks()
|
||||
|
||||
assert len(chunks) == 2
|
||||
texts = {c["text"] for c in chunks}
|
||||
assert texts == {"text1", "text2"}
|
||||
|
||||
@patch("application.vectorstore.faiss.StorageCreator")
|
||||
@patch("application.vectorstore.faiss.FAISS")
|
||||
@patch.object(
|
||||
__import__("application.vectorstore.base", fromlist=["BaseVectorStore"]).BaseVectorStore,
|
||||
"_get_embeddings",
|
||||
)
|
||||
@patch("application.vectorstore.faiss.settings")
|
||||
def test_get_chunks_empty(self, mock_settings, mock_get_emb, mock_faiss, mock_storage_creator):
|
||||
mock_settings.EMBEDDINGS_NAME = "test_model"
|
||||
mock_emb = Mock(dimension=3)
|
||||
mock_get_emb.return_value = mock_emb
|
||||
mock_ds = Mock()
|
||||
mock_ds.index = Mock(d=3)
|
||||
mock_ds.docstore._dict = {}
|
||||
mock_faiss.from_documents.return_value = mock_ds
|
||||
mock_storage_creator.get_storage.return_value = Mock()
|
||||
|
||||
from application.vectorstore.faiss import FaissStore
|
||||
|
||||
store = FaissStore(source_id="t", embeddings_key="k", docs_init=[Mock()])
|
||||
assert store.get_chunks() == []
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestFaissStoreSaveLocal:
|
||||
@patch("application.vectorstore.faiss.StorageCreator")
|
||||
@patch("application.vectorstore.faiss.FAISS")
|
||||
@patch.object(
|
||||
__import__("application.vectorstore.base", fromlist=["BaseVectorStore"]).BaseVectorStore,
|
||||
"_get_embeddings",
|
||||
)
|
||||
@patch("application.vectorstore.faiss.settings")
|
||||
def test_save_local_with_path(
|
||||
self, mock_settings, mock_get_emb, mock_faiss, mock_storage_creator
|
||||
):
|
||||
mock_settings.EMBEDDINGS_NAME = "test_model"
|
||||
mock_emb = Mock(dimension=3)
|
||||
mock_get_emb.return_value = mock_emb
|
||||
mock_ds = Mock()
|
||||
mock_ds.index = Mock(d=3)
|
||||
mock_faiss.from_documents.return_value = mock_ds
|
||||
mock_storage = Mock()
|
||||
mock_storage_creator.get_storage.return_value = mock_storage
|
||||
|
||||
from application.vectorstore.faiss import FaissStore
|
||||
|
||||
store = FaissStore(source_id="t", embeddings_key="k", docs_init=[Mock()])
|
||||
|
||||
# Mock _save_to_storage to avoid file I/O
|
||||
store._save_to_storage = Mock(return_value=True)
|
||||
|
||||
with patch("os.makedirs"):
|
||||
result = store.save_local(path="/tmp/test_save")
|
||||
|
||||
mock_ds.save_local.assert_called_once_with("/tmp/test_save")
|
||||
store._save_to_storage.assert_called_once()
|
||||
assert result is True
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestFaissStoreDeleteIndex:
|
||||
@patch("application.vectorstore.faiss.StorageCreator")
|
||||
@patch("application.vectorstore.faiss.FAISS")
|
||||
@patch.object(
|
||||
__import__("application.vectorstore.base", fromlist=["BaseVectorStore"]).BaseVectorStore,
|
||||
"_get_embeddings",
|
||||
)
|
||||
@patch("application.vectorstore.faiss.settings")
|
||||
def test_delete_index_delegates(
|
||||
self, mock_settings, mock_get_emb, mock_faiss, mock_storage_creator
|
||||
):
|
||||
mock_settings.EMBEDDINGS_NAME = "test_model"
|
||||
mock_emb = Mock(dimension=3)
|
||||
mock_get_emb.return_value = mock_emb
|
||||
mock_ds = Mock()
|
||||
mock_ds.index = Mock(d=3)
|
||||
mock_faiss.from_documents.return_value = mock_ds
|
||||
mock_storage_creator.get_storage.return_value = Mock()
|
||||
|
||||
from application.vectorstore.faiss import FaissStore
|
||||
|
||||
store = FaissStore(source_id="t", embeddings_key="k", docs_init=[Mock()])
|
||||
store.delete_index(["id1"])
|
||||
mock_ds.delete.assert_called_once_with(["id1"])
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestFaissStoreAssertEmbeddingDimensions:
|
||||
@patch("application.vectorstore.faiss.StorageCreator")
|
||||
@patch("application.vectorstore.faiss.FAISS")
|
||||
@patch.object(
|
||||
__import__("application.vectorstore.base", fromlist=["BaseVectorStore"]).BaseVectorStore,
|
||||
"_get_embeddings",
|
||||
)
|
||||
@patch("application.vectorstore.faiss.settings")
|
||||
def test_dimension_mismatch_raises(
|
||||
self, mock_settings, mock_get_emb, mock_faiss, mock_storage_creator
|
||||
):
|
||||
mock_settings.EMBEDDINGS_NAME = (
|
||||
"huggingface_sentence-transformers/all-mpnet-base-v2"
|
||||
)
|
||||
mock_emb = Mock(dimension=768)
|
||||
mock_get_emb.return_value = mock_emb
|
||||
mock_ds = Mock()
|
||||
mock_ds.index = Mock(d=512) # Mismatched dimension
|
||||
mock_faiss.from_documents.return_value = mock_ds
|
||||
mock_storage_creator.get_storage.return_value = Mock()
|
||||
|
||||
from application.vectorstore.faiss import FaissStore
|
||||
|
||||
with pytest.raises(ValueError, match="Embedding dimension mismatch"):
|
||||
FaissStore(source_id="t", embeddings_key="k", docs_init=[Mock()])
|
||||
|
||||
@patch("application.vectorstore.faiss.StorageCreator")
|
||||
@patch("application.vectorstore.faiss.FAISS")
|
||||
@patch.object(
|
||||
__import__("application.vectorstore.base", fromlist=["BaseVectorStore"]).BaseVectorStore,
|
||||
"_get_embeddings",
|
||||
)
|
||||
@patch("application.vectorstore.faiss.settings")
|
||||
def test_missing_dimension_attr_raises(
|
||||
self, mock_settings, mock_get_emb, mock_faiss, mock_storage_creator
|
||||
):
|
||||
mock_settings.EMBEDDINGS_NAME = (
|
||||
"huggingface_sentence-transformers/all-mpnet-base-v2"
|
||||
)
|
||||
mock_emb = Mock(spec=[]) # No dimension attribute
|
||||
mock_get_emb.return_value = mock_emb
|
||||
mock_ds = Mock()
|
||||
mock_ds.index = Mock(d=768)
|
||||
mock_faiss.from_documents.return_value = mock_ds
|
||||
mock_storage_creator.get_storage.return_value = Mock()
|
||||
|
||||
from application.vectorstore.faiss import FaissStore
|
||||
|
||||
with pytest.raises(AttributeError, match="dimension"):
|
||||
FaissStore(source_id="t", embeddings_key="k", docs_init=[Mock()])
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestFaissStoreDeleteChunk:
|
||||
@patch("application.vectorstore.faiss.StorageCreator")
|
||||
@patch("application.vectorstore.faiss.FAISS")
|
||||
@patch.object(
|
||||
__import__("application.vectorstore.base", fromlist=["BaseVectorStore"]).BaseVectorStore,
|
||||
"_get_embeddings",
|
||||
)
|
||||
@patch("application.vectorstore.faiss.settings")
|
||||
def test_delete_chunk(self, mock_settings, mock_get_emb, mock_faiss, mock_storage_creator):
|
||||
mock_settings.EMBEDDINGS_NAME = "test_model"
|
||||
mock_emb = Mock(dimension=3)
|
||||
mock_get_emb.return_value = mock_emb
|
||||
mock_ds = Mock()
|
||||
mock_ds.index = Mock(d=3)
|
||||
mock_faiss.from_documents.return_value = mock_ds
|
||||
mock_storage = Mock()
|
||||
mock_storage_creator.get_storage.return_value = mock_storage
|
||||
|
||||
from application.vectorstore.faiss import FaissStore
|
||||
|
||||
store = FaissStore(source_id="t", embeddings_key="k", docs_init=[Mock()])
|
||||
store._save_to_storage = Mock(return_value=True)
|
||||
|
||||
result = store.delete_chunk("chunk_id")
|
||||
mock_ds.delete.assert_called_once_with(["chunk_id"])
|
||||
store._save_to_storage.assert_called_once()
|
||||
assert result is True
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestGetVectorstore:
|
||||
def test_with_path(self):
|
||||
from application.vectorstore.faiss import get_vectorstore
|
||||
|
||||
assert get_vectorstore("abc123") == "indexes/abc123"
|
||||
|
||||
def test_without_path(self):
|
||||
from application.vectorstore.faiss import get_vectorstore
|
||||
|
||||
assert get_vectorstore("") == "indexes"
|
||||
assert get_vectorstore(None) == "indexes"
|
||||
|
||||
def test_with_nested_path(self):
|
||||
from application.vectorstore.faiss import get_vectorstore
|
||||
|
||||
assert get_vectorstore("user/source123") == "indexes/user/source123"
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"malicious_path",
|
||||
[
|
||||
"../outside",
|
||||
"../../etc/passwd",
|
||||
"nested/../../../outside",
|
||||
"/tmp/evil",
|
||||
"..\\outside",
|
||||
"valid/../../escape",
|
||||
],
|
||||
)
|
||||
def test_rejects_path_traversal(self, malicious_path):
|
||||
from application.vectorstore.faiss import get_vectorstore
|
||||
|
||||
with pytest.raises(ValueError, match="Invalid source_id path"):
|
||||
get_vectorstore(malicious_path)
|
||||
|
||||
def test_allows_mongodb_style_ids(self):
|
||||
from application.vectorstore.faiss import get_vectorstore
|
||||
|
||||
assert get_vectorstore("65e8f6a8a7a96b1bdad4154f") == "indexes/65e8f6a8a7a96b1bdad4154f"
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestFaissStoreAddChunk:
|
||||
@patch("application.vectorstore.faiss.StorageCreator")
|
||||
@patch("application.vectorstore.faiss.FAISS")
|
||||
@patch.object(
|
||||
__import__("application.vectorstore.base", fromlist=["BaseVectorStore"]).BaseVectorStore,
|
||||
"_get_embeddings",
|
||||
)
|
||||
@patch("application.vectorstore.faiss.settings")
|
||||
def test_add_chunk_with_metadata(
|
||||
self, mock_settings, mock_get_emb, mock_faiss, mock_storage_creator
|
||||
):
|
||||
mock_settings.EMBEDDINGS_NAME = "test_model"
|
||||
mock_emb = Mock(dimension=3)
|
||||
mock_get_emb.return_value = mock_emb
|
||||
mock_ds = Mock()
|
||||
mock_ds.index = Mock(d=3)
|
||||
mock_ds.add_documents.return_value = ["new_id"]
|
||||
mock_faiss.from_documents.return_value = mock_ds
|
||||
mock_storage = Mock()
|
||||
mock_storage_creator.get_storage.return_value = mock_storage
|
||||
|
||||
from application.vectorstore.faiss import FaissStore
|
||||
|
||||
store = FaissStore(source_id="t", embeddings_key="k", docs_init=[Mock()])
|
||||
store._save_to_storage = Mock(return_value=True)
|
||||
|
||||
doc_id = store.add_chunk("new text", metadata={"source": "test"})
|
||||
|
||||
assert doc_id == ["new_id"]
|
||||
mock_ds.add_documents.assert_called_once()
|
||||
store._save_to_storage.assert_called_once()
|
||||
|
||||
@patch("application.vectorstore.faiss.StorageCreator")
|
||||
@patch("application.vectorstore.faiss.FAISS")
|
||||
@patch.object(
|
||||
__import__("application.vectorstore.base", fromlist=["BaseVectorStore"]).BaseVectorStore,
|
||||
"_get_embeddings",
|
||||
)
|
||||
@patch("application.vectorstore.faiss.settings")
|
||||
def test_add_chunk_default_metadata(
|
||||
self, mock_settings, mock_get_emb, mock_faiss, mock_storage_creator
|
||||
):
|
||||
mock_settings.EMBEDDINGS_NAME = "test_model"
|
||||
mock_emb = Mock(dimension=3)
|
||||
mock_get_emb.return_value = mock_emb
|
||||
mock_ds = Mock()
|
||||
mock_ds.index = Mock(d=3)
|
||||
mock_ds.add_documents.return_value = ["new_id"]
|
||||
mock_faiss.from_documents.return_value = mock_ds
|
||||
mock_storage = Mock()
|
||||
mock_storage_creator.get_storage.return_value = mock_storage
|
||||
|
||||
from application.vectorstore.faiss import FaissStore
|
||||
|
||||
store = FaissStore(source_id="t", embeddings_key="k", docs_init=[Mock()])
|
||||
store._save_to_storage = Mock(return_value=True)
|
||||
|
||||
doc_id = store.add_chunk("new text")
|
||||
|
||||
assert doc_id == ["new_id"]
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestFaissStoreSaveLocalNoPath:
|
||||
@patch("application.vectorstore.faiss.StorageCreator")
|
||||
@patch("application.vectorstore.faiss.FAISS")
|
||||
@patch.object(
|
||||
__import__("application.vectorstore.base", fromlist=["BaseVectorStore"]).BaseVectorStore,
|
||||
"_get_embeddings",
|
||||
)
|
||||
@patch("application.vectorstore.faiss.settings")
|
||||
def test_save_local_without_path(
|
||||
self, mock_settings, mock_get_emb, mock_faiss, mock_storage_creator
|
||||
):
|
||||
mock_settings.EMBEDDINGS_NAME = "test_model"
|
||||
mock_emb = Mock(dimension=3)
|
||||
mock_get_emb.return_value = mock_emb
|
||||
mock_ds = Mock()
|
||||
mock_ds.index = Mock(d=3)
|
||||
mock_faiss.from_documents.return_value = mock_ds
|
||||
mock_storage = Mock()
|
||||
mock_storage_creator.get_storage.return_value = mock_storage
|
||||
|
||||
from application.vectorstore.faiss import FaissStore
|
||||
|
||||
store = FaissStore(source_id="t", embeddings_key="k", docs_init=[Mock()])
|
||||
store._save_to_storage = Mock(return_value=True)
|
||||
|
||||
result = store.save_local()
|
||||
|
||||
# Should NOT call docsearch.save_local with a path
|
||||
mock_ds.save_local.assert_not_called()
|
||||
store._save_to_storage.assert_called_once()
|
||||
assert result is True
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestFaissStoreAssertEmbeddingDimensionsMatch:
|
||||
@patch("application.vectorstore.faiss.StorageCreator")
|
||||
@patch("application.vectorstore.faiss.FAISS")
|
||||
@patch.object(
|
||||
__import__("application.vectorstore.base", fromlist=["BaseVectorStore"]).BaseVectorStore,
|
||||
"_get_embeddings",
|
||||
)
|
||||
@patch("application.vectorstore.faiss.settings")
|
||||
def test_dimension_match_passes(
|
||||
self, mock_settings, mock_get_emb, mock_faiss, mock_storage_creator
|
||||
):
|
||||
mock_settings.EMBEDDINGS_NAME = (
|
||||
"huggingface_sentence-transformers/all-mpnet-base-v2"
|
||||
)
|
||||
mock_emb = Mock(dimension=768)
|
||||
mock_get_emb.return_value = mock_emb
|
||||
mock_ds = Mock()
|
||||
mock_ds.index = Mock(d=768) # Matching dimension
|
||||
mock_faiss.from_documents.return_value = mock_ds
|
||||
mock_storage_creator.get_storage.return_value = Mock()
|
||||
|
||||
from application.vectorstore.faiss import FaissStore
|
||||
|
||||
# Should not raise
|
||||
store = FaissStore(source_id="t", embeddings_key="k", docs_init=[Mock()])
|
||||
assert store is not None
|
||||
|
||||
@patch("application.vectorstore.faiss.StorageCreator")
|
||||
@patch("application.vectorstore.faiss.FAISS")
|
||||
@patch.object(
|
||||
__import__("application.vectorstore.base", fromlist=["BaseVectorStore"]).BaseVectorStore,
|
||||
"_get_embeddings",
|
||||
)
|
||||
@patch("application.vectorstore.faiss.settings")
|
||||
def test_non_huggingface_skips_dimension_check(
|
||||
self, mock_settings, mock_get_emb, mock_faiss, mock_storage_creator
|
||||
):
|
||||
mock_settings.EMBEDDINGS_NAME = "openai_text-embedding-ada-002"
|
||||
mock_emb = Mock(dimension=1536)
|
||||
mock_get_emb.return_value = mock_emb
|
||||
mock_ds = Mock()
|
||||
mock_ds.index = Mock(d=999) # Mismatched but doesn't matter
|
||||
mock_faiss.from_documents.return_value = mock_ds
|
||||
mock_storage_creator.get_storage.return_value = Mock()
|
||||
|
||||
from application.vectorstore.faiss import FaissStore
|
||||
|
||||
# Should not raise since embedding name is not the huggingface one
|
||||
store = FaissStore(source_id="t", embeddings_key="k", docs_init=[Mock()])
|
||||
assert store is not None
|
||||
@@ -0,0 +1,312 @@
|
||||
from unittest.mock import MagicMock, Mock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
def _make_lancedb_store(source_id="test-source"):
|
||||
"""Helper to create a LanceDBVectorStore with mocked deps."""
|
||||
with patch(
|
||||
"application.vectorstore.lancedb.settings"
|
||||
) as mock_settings:
|
||||
mock_settings.LANCEDB_PATH = "/tmp/lancedb"
|
||||
mock_settings.LANCEDB_TABLE_NAME = "docs"
|
||||
mock_settings.EMBEDDINGS_NAME = "test_model"
|
||||
|
||||
from application.vectorstore.lancedb import LanceDBVectorStore
|
||||
|
||||
store = LanceDBVectorStore(
|
||||
path="/tmp/lancedb",
|
||||
table_name_prefix="docs",
|
||||
source_id=source_id,
|
||||
embeddings_key="key",
|
||||
)
|
||||
|
||||
return store, mock_settings
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestLanceDBVectorStoreInit:
|
||||
def test_table_name_with_source_id(self):
|
||||
store, _ = _make_lancedb_store(source_id="src1")
|
||||
assert store.table_name == "docs_src1"
|
||||
|
||||
def test_table_name_without_source_id(self):
|
||||
with patch("application.vectorstore.lancedb.settings") as mock_settings:
|
||||
mock_settings.LANCEDB_PATH = "/tmp"
|
||||
mock_settings.LANCEDB_TABLE_NAME = "docs"
|
||||
|
||||
from application.vectorstore.lancedb import LanceDBVectorStore
|
||||
|
||||
store = LanceDBVectorStore(
|
||||
path="/tmp", table_name_prefix="docs", source_id=None
|
||||
)
|
||||
assert store.table_name == "docs"
|
||||
|
||||
def test_init_defaults(self):
|
||||
store, _ = _make_lancedb_store()
|
||||
assert store.path == "/tmp/lancedb"
|
||||
assert store._lance_db is None
|
||||
assert store.docsearch is None
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestLanceDBVectorStoreLazyLoading:
|
||||
def test_pa_lazy_load(self):
|
||||
store, _ = _make_lancedb_store()
|
||||
mock_pa = MagicMock()
|
||||
|
||||
with patch("importlib.import_module", return_value=mock_pa) as mock_import:
|
||||
result = store.pa
|
||||
mock_import.assert_called_with("pyarrow")
|
||||
assert result is mock_pa
|
||||
|
||||
def test_pa_cached(self):
|
||||
store, _ = _make_lancedb_store()
|
||||
mock_pa = MagicMock()
|
||||
store._pa = mock_pa
|
||||
|
||||
assert store.pa is mock_pa
|
||||
|
||||
def test_lancedb_lazy_load(self):
|
||||
store, _ = _make_lancedb_store()
|
||||
mock_ldb = MagicMock()
|
||||
|
||||
with patch("importlib.import_module", return_value=mock_ldb) as mock_import:
|
||||
result = store.lancedb
|
||||
mock_import.assert_called_with("lancedb")
|
||||
assert result is mock_ldb
|
||||
|
||||
def test_lance_db_connection(self):
|
||||
store, _ = _make_lancedb_store()
|
||||
mock_ldb_module = MagicMock()
|
||||
mock_conn = MagicMock()
|
||||
mock_ldb_module.connect.return_value = mock_conn
|
||||
store._lancedb_module = mock_ldb_module
|
||||
|
||||
result = store.lance_db
|
||||
mock_ldb_module.connect.assert_called_once_with("/tmp/lancedb")
|
||||
assert result is mock_conn
|
||||
|
||||
def test_lance_db_cached(self):
|
||||
store, _ = _make_lancedb_store()
|
||||
mock_conn = MagicMock()
|
||||
store._lance_db = mock_conn
|
||||
|
||||
assert store.lance_db is mock_conn
|
||||
|
||||
def test_table_opens_existing(self):
|
||||
store, _ = _make_lancedb_store()
|
||||
mock_conn = MagicMock()
|
||||
mock_table = MagicMock()
|
||||
mock_conn.table_names.return_value = [store.table_name]
|
||||
mock_conn.open_table.return_value = mock_table
|
||||
store._lance_db = mock_conn
|
||||
|
||||
result = store.table
|
||||
mock_conn.open_table.assert_called_once_with(store.table_name)
|
||||
assert result is mock_table
|
||||
|
||||
def test_table_returns_none_for_missing(self):
|
||||
store, _ = _make_lancedb_store()
|
||||
mock_conn = MagicMock()
|
||||
mock_conn.table_names.return_value = []
|
||||
store._lance_db = mock_conn
|
||||
|
||||
result = store.table
|
||||
assert result is None
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestLanceDBVectorStoreEnsureTableExists:
|
||||
def test_creates_table_when_missing(self):
|
||||
store, _ = _make_lancedb_store()
|
||||
mock_conn = MagicMock()
|
||||
mock_conn.table_names.return_value = []
|
||||
store._lance_db = mock_conn
|
||||
|
||||
mock_emb = MagicMock()
|
||||
mock_emb.dimension = 768
|
||||
mock_pa = MagicMock()
|
||||
store._pa = mock_pa
|
||||
|
||||
with patch.object(store, "_get_embeddings", return_value=mock_emb):
|
||||
store.ensure_table_exists()
|
||||
|
||||
mock_conn.create_table.assert_called_once()
|
||||
|
||||
def test_noop_when_table_exists(self):
|
||||
store, _ = _make_lancedb_store()
|
||||
mock_table = MagicMock()
|
||||
store.docsearch = mock_table
|
||||
mock_conn = MagicMock()
|
||||
mock_conn.table_names.return_value = [store.table_name]
|
||||
mock_conn.open_table.return_value = mock_table
|
||||
store._lance_db = mock_conn
|
||||
|
||||
store.ensure_table_exists()
|
||||
mock_conn.create_table.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestLanceDBVectorStoreAddTexts:
|
||||
def test_add_texts(self):
|
||||
store, _ = _make_lancedb_store()
|
||||
mock_table = MagicMock()
|
||||
store.docsearch = mock_table
|
||||
|
||||
mock_emb = MagicMock()
|
||||
mock_emb.embed_documents.return_value = [[0.1, 0.2], [0.3, 0.4]]
|
||||
|
||||
with patch.object(store, "_get_embeddings", return_value=mock_emb), patch.object(
|
||||
store, "ensure_table_exists"
|
||||
):
|
||||
store.add_texts(
|
||||
["text1", "text2"],
|
||||
metadatas=[{"a": "1"}, {"b": "2"}],
|
||||
source_id="src1",
|
||||
)
|
||||
|
||||
mock_table.add.assert_called_once()
|
||||
vectors = mock_table.add.call_args[0][0]
|
||||
assert len(vectors) == 2
|
||||
assert vectors[0]["text"] == "text1"
|
||||
assert vectors[0]["vector"] == [0.1, 0.2]
|
||||
|
||||
def test_add_texts_with_source_id_in_metadata(self):
|
||||
store, _ = _make_lancedb_store()
|
||||
mock_table = MagicMock()
|
||||
store.docsearch = mock_table
|
||||
|
||||
mock_emb = MagicMock()
|
||||
mock_emb.embed_documents.return_value = [[0.1]]
|
||||
|
||||
with patch.object(store, "_get_embeddings", return_value=mock_emb), patch.object(
|
||||
store, "ensure_table_exists"
|
||||
):
|
||||
store.add_texts(["text1"], metadatas=[{"k": "v"}], source_id="src1")
|
||||
|
||||
vectors = mock_table.add.call_args[0][0]
|
||||
metadata_keys = [m["key"] for m in vectors[0]["metadata"]]
|
||||
assert "source_id" in metadata_keys
|
||||
|
||||
def test_add_texts_default_metadata(self):
|
||||
store, _ = _make_lancedb_store()
|
||||
mock_table = MagicMock()
|
||||
store.docsearch = mock_table
|
||||
|
||||
mock_emb = MagicMock()
|
||||
mock_emb.embed_documents.return_value = [[0.1]]
|
||||
|
||||
with patch.object(store, "_get_embeddings", return_value=mock_emb), patch.object(
|
||||
store, "ensure_table_exists"
|
||||
):
|
||||
store.add_texts(["text1"])
|
||||
|
||||
mock_table.add.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestLanceDBVectorStoreSearch:
|
||||
def test_search(self):
|
||||
store, _ = _make_lancedb_store()
|
||||
mock_table = MagicMock()
|
||||
store.docsearch = mock_table
|
||||
|
||||
mock_emb = MagicMock()
|
||||
mock_emb.embed_query.return_value = [0.1, 0.2]
|
||||
|
||||
mock_result = MagicMock()
|
||||
mock_result.limit.return_value.to_list.return_value = [
|
||||
{"_distance": 0.1, "text": "result1", "metadata": {"k": "v"}},
|
||||
]
|
||||
mock_table.search.return_value = mock_result
|
||||
|
||||
with patch.object(store, "_get_embeddings", return_value=mock_emb), patch.object(
|
||||
store, "ensure_table_exists"
|
||||
):
|
||||
results = store.search("query", k=3)
|
||||
|
||||
assert len(results) == 1
|
||||
assert results[0][1] == "result1"
|
||||
mock_result.limit.assert_called_once_with(3)
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestLanceDBVectorStoreDeleteIndex:
|
||||
def test_delete_index_drops_table(self):
|
||||
store, _ = _make_lancedb_store()
|
||||
mock_table = MagicMock()
|
||||
store.docsearch = mock_table
|
||||
mock_conn = MagicMock()
|
||||
mock_conn.table_names.return_value = [store.table_name]
|
||||
mock_conn.open_table.return_value = mock_table
|
||||
store._lance_db = mock_conn
|
||||
|
||||
store.delete_index()
|
||||
|
||||
mock_conn.drop_table.assert_called_once_with(store.table_name)
|
||||
|
||||
def test_delete_index_noop_when_no_table(self):
|
||||
store, _ = _make_lancedb_store()
|
||||
store.docsearch = None
|
||||
mock_conn = MagicMock()
|
||||
mock_conn.table_names.return_value = []
|
||||
store._lance_db = mock_conn
|
||||
|
||||
store.delete_index()
|
||||
mock_conn.drop_table.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestLanceDBVectorStoreAssertEmbeddingDimensions:
|
||||
def test_matching_dimensions(self):
|
||||
store, _ = _make_lancedb_store()
|
||||
mock_table = MagicMock()
|
||||
mock_table.schema = {"vector": MagicMock()}
|
||||
mock_table.schema["vector"].type.value_type.__len__ = Mock(return_value=768)
|
||||
store.docsearch = mock_table
|
||||
|
||||
mock_emb = MagicMock()
|
||||
mock_emb.dimension = 768
|
||||
|
||||
# Should not raise
|
||||
store.assert_embedding_dimensions(mock_emb)
|
||||
|
||||
def test_mismatched_dimensions_raises(self):
|
||||
store, _ = _make_lancedb_store()
|
||||
mock_table = MagicMock()
|
||||
store.docsearch = mock_table
|
||||
|
||||
type_mock = MagicMock()
|
||||
type_mock.__len__ = Mock(return_value=512)
|
||||
mock_table.schema.__getitem__ = Mock(return_value=MagicMock())
|
||||
mock_table.schema["vector"].type.value_type = type_mock
|
||||
|
||||
mock_emb = MagicMock()
|
||||
mock_emb.dimension = 768
|
||||
|
||||
with pytest.raises(ValueError, match="Embedding dimension mismatch"):
|
||||
store.assert_embedding_dimensions(mock_emb)
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestLanceDBVectorStoreFilterDocuments:
|
||||
def test_filter_documents(self):
|
||||
store, _ = _make_lancedb_store()
|
||||
mock_table = MagicMock()
|
||||
mock_table.filter.return_value.to_list.return_value = [{"text": "filtered"}]
|
||||
store.docsearch = mock_table
|
||||
|
||||
with patch.object(store, "ensure_table_exists"):
|
||||
results = store.filter_documents({"source_id": "src1"})
|
||||
|
||||
assert len(results) == 1
|
||||
|
||||
def test_filter_documents_requires_source_id(self):
|
||||
store, _ = _make_lancedb_store()
|
||||
mock_table = MagicMock()
|
||||
store.docsearch = mock_table
|
||||
|
||||
with patch.object(store, "ensure_table_exists"):
|
||||
with pytest.raises(ValueError, match="must contain 'source_id'"):
|
||||
store.filter_documents({"other_key": "value"})
|
||||
@@ -0,0 +1,95 @@
|
||||
from unittest.mock import MagicMock, Mock, patch
|
||||
from uuid import UUID
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
def _make_milvus_store(source_id="test-source"):
|
||||
"""Helper to create a MilvusStore with mocked deps."""
|
||||
with patch(
|
||||
"application.vectorstore.base.BaseVectorStore._get_embeddings"
|
||||
) as mock_get_emb, patch(
|
||||
"application.vectorstore.milvus.settings"
|
||||
) as mock_settings, patch.dict(
|
||||
"sys.modules",
|
||||
{
|
||||
"langchain_milvus": MagicMock(),
|
||||
},
|
||||
):
|
||||
mock_emb = Mock()
|
||||
mock_get_emb.return_value = mock_emb
|
||||
mock_settings.EMBEDDINGS_NAME = "test_model"
|
||||
mock_settings.MILVUS_URI = "http://localhost:19530"
|
||||
mock_settings.MILVUS_TOKEN = "token"
|
||||
mock_settings.MILVUS_COLLECTION_NAME = "test_collection"
|
||||
|
||||
from langchain_milvus import Milvus
|
||||
|
||||
mock_docsearch = MagicMock()
|
||||
Milvus.return_value = mock_docsearch
|
||||
|
||||
from application.vectorstore.milvus import MilvusStore
|
||||
|
||||
store = MilvusStore(source_id=source_id, embeddings_key="key")
|
||||
store._docsearch = mock_docsearch
|
||||
|
||||
return store, mock_docsearch
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestMilvusStoreInit:
|
||||
def test_source_id_stored(self):
|
||||
store, _ = _make_milvus_store(source_id="src1")
|
||||
assert store._source_id == "src1"
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestMilvusStoreSearch:
|
||||
def test_search(self):
|
||||
store, mock_ds = _make_milvus_store(source_id="src1")
|
||||
mock_ds.similarity_search.return_value = ["doc1", "doc2"]
|
||||
|
||||
results = store.search("query", k=3)
|
||||
|
||||
mock_ds.similarity_search.assert_called_once()
|
||||
call_kwargs = mock_ds.similarity_search.call_args
|
||||
assert call_kwargs[1]["query"] == "query"
|
||||
assert call_kwargs[1]["k"] == 3
|
||||
assert call_kwargs[1]["expr"] == "source_id == 'src1'"
|
||||
assert results == ["doc1", "doc2"]
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestMilvusStoreAddTexts:
|
||||
def test_add_texts(self):
|
||||
store, mock_ds = _make_milvus_store()
|
||||
mock_ds.add_texts.return_value = ["id1", "id2"]
|
||||
|
||||
result = store.add_texts(
|
||||
["text1", "text2"], metadatas=[{"a": 1}, {"b": 2}]
|
||||
)
|
||||
|
||||
mock_ds.add_texts.assert_called_once()
|
||||
call_kwargs = mock_ds.add_texts.call_args
|
||||
assert call_kwargs[1]["texts"] == ["text1", "text2"]
|
||||
# ids should be UUIDs
|
||||
ids = call_kwargs[1]["ids"]
|
||||
assert len(ids) == 2
|
||||
for uid in ids:
|
||||
UUID(uid) # Validates it's a valid UUID
|
||||
|
||||
assert result == ["id1", "id2"]
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestMilvusStoreSaveLocal:
|
||||
def test_save_local_is_noop(self):
|
||||
store, _ = _make_milvus_store()
|
||||
assert store.save_local() is None
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestMilvusStoreDeleteIndex:
|
||||
def test_delete_index_is_noop(self):
|
||||
store, _ = _make_milvus_store()
|
||||
assert store.delete_index() is None
|
||||
@@ -0,0 +1,289 @@
|
||||
from unittest.mock import MagicMock, Mock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
def _make_mongodb_store(source_id="test-source"):
|
||||
"""Helper to create a MongoDBVectorStore with all external deps mocked."""
|
||||
with patch(
|
||||
"application.vectorstore.base.BaseVectorStore._get_embeddings"
|
||||
) as mock_get_emb, patch(
|
||||
"application.vectorstore.mongodb.settings"
|
||||
) as mock_settings, patch.dict(
|
||||
"sys.modules", {"pymongo": MagicMock()}
|
||||
):
|
||||
mock_emb = Mock()
|
||||
mock_emb.embed_query = Mock(return_value=[0.1, 0.2, 0.3])
|
||||
mock_emb.embed_documents = Mock(return_value=[[0.1, 0.2, 0.3]])
|
||||
mock_get_emb.return_value = mock_emb
|
||||
mock_settings.EMBEDDINGS_NAME = "test_model"
|
||||
mock_settings.MONGO_URI = "mongodb://localhost:27017"
|
||||
|
||||
from application.vectorstore.mongodb import MongoDBVectorStore
|
||||
|
||||
store = MongoDBVectorStore(
|
||||
source_id=source_id,
|
||||
embeddings_key="key",
|
||||
collection="test_docs",
|
||||
database="test_db",
|
||||
)
|
||||
|
||||
mock_collection = MagicMock()
|
||||
store._collection = mock_collection
|
||||
|
||||
return store, mock_collection, mock_emb
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestMongoDBVectorStoreInit:
|
||||
def test_source_id_cleaned(self):
|
||||
store, _, _ = _make_mongodb_store(source_id="application/indexes/abc123/")
|
||||
assert store._source_id == "abc123"
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestMongoDBVectorStoreSearch:
|
||||
def test_search_builds_pipeline(self):
|
||||
store, mock_collection, mock_emb = _make_mongodb_store()
|
||||
|
||||
doc1 = {
|
||||
"_id": "id1",
|
||||
"text": "hello world",
|
||||
"embedding": [0.1, 0.2],
|
||||
"source": "test",
|
||||
}
|
||||
mock_collection.aggregate.return_value = iter([doc1])
|
||||
|
||||
results = store.search("query", k=3)
|
||||
|
||||
mock_emb.embed_query.assert_called_once_with("query")
|
||||
mock_collection.aggregate.assert_called_once()
|
||||
pipeline = mock_collection.aggregate.call_args[0][0]
|
||||
assert pipeline[0]["$vectorSearch"]["limit"] == 3
|
||||
assert pipeline[0]["$vectorSearch"]["numCandidates"] == 30
|
||||
|
||||
assert len(results) == 1
|
||||
assert str(results[0]) == "hello world"
|
||||
|
||||
def test_score_threshold_adds_match_stage(self):
|
||||
store, mock_collection, _ = _make_mongodb_store()
|
||||
mock_collection.aggregate.return_value = iter([])
|
||||
|
||||
store.search("query", k=3, score_threshold=0.6)
|
||||
|
||||
pipeline = mock_collection.aggregate.call_args[0][0]
|
||||
# $addFields surfaces the vectorSearchScore, $match enforces the floor.
|
||||
assert any("$addFields" in stage for stage in pipeline)
|
||||
match_stages = [s for s in pipeline if "$match" in s]
|
||||
assert match_stages[0]["$match"]["_score"]["$gte"] == 0.6
|
||||
|
||||
def test_no_score_threshold_omits_match_stage(self):
|
||||
store, mock_collection, _ = _make_mongodb_store()
|
||||
mock_collection.aggregate.return_value = iter([])
|
||||
|
||||
store.search("query", k=3)
|
||||
|
||||
pipeline = mock_collection.aggregate.call_args[0][0]
|
||||
assert not any("$match" in stage for stage in pipeline)
|
||||
|
||||
def test_score_field_stripped_from_metadata(self):
|
||||
store, mock_collection, _ = _make_mongodb_store()
|
||||
doc = {"_id": "i", "text": "t", "embedding": [0.1], "_score": 0.9, "k": "v"}
|
||||
mock_collection.aggregate.return_value = iter([doc])
|
||||
|
||||
results = store.search("q", k=1, score_threshold=0.5)
|
||||
assert "_score" not in results[0].metadata
|
||||
assert results[0].metadata["k"] == "v"
|
||||
|
||||
def test_search_removes_id_text_embedding_from_metadata(self):
|
||||
store, mock_collection, _ = _make_mongodb_store()
|
||||
|
||||
doc = {
|
||||
"_id": "id1",
|
||||
"text": "content",
|
||||
"embedding": [0.1],
|
||||
"custom_key": "custom_val",
|
||||
}
|
||||
mock_collection.aggregate.return_value = iter([doc])
|
||||
|
||||
results = store.search("q", k=1)
|
||||
metadata = results[0].metadata
|
||||
assert "_id" not in metadata
|
||||
assert "text" not in metadata
|
||||
assert "embedding" not in metadata
|
||||
assert metadata["custom_key"] == "custom_val"
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestMongoDBVectorStoreAddTexts:
|
||||
def test_add_texts_batches(self):
|
||||
store, mock_collection, mock_emb = _make_mongodb_store()
|
||||
# Generate 150 texts to trigger batching at 100
|
||||
texts = [f"text_{i}" for i in range(150)]
|
||||
metadatas = [{"i": i} for i in range(150)]
|
||||
mock_emb.embed_documents.return_value = [[0.1]] * 100 # per batch
|
||||
|
||||
mock_collection.insert_many.return_value = Mock(
|
||||
inserted_ids=list(range(100))
|
||||
)
|
||||
|
||||
store.add_texts(texts, metadatas)
|
||||
|
||||
# Should have been called twice: batch of 100, then batch of 50
|
||||
assert mock_collection.insert_many.call_count == 2
|
||||
|
||||
def test_add_texts_default_metadata(self):
|
||||
store, mock_collection, mock_emb = _make_mongodb_store()
|
||||
mock_emb.embed_documents.return_value = [[0.1]]
|
||||
mock_collection.insert_many.return_value = Mock(inserted_ids=["id1"])
|
||||
|
||||
store.add_texts(["text1"])
|
||||
mock_collection.insert_many.assert_called_once()
|
||||
|
||||
def test_add_texts_empty(self):
|
||||
store, mock_collection, mock_emb = _make_mongodb_store()
|
||||
mock_emb.embed_documents.return_value = []
|
||||
mock_collection.insert_many.return_value = Mock(inserted_ids=[])
|
||||
|
||||
result = store.add_texts([], [])
|
||||
# _insert_texts returns [] for empty input
|
||||
assert result == []
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestMongoDBVectorStoreInsertTexts:
|
||||
def test_insert_texts_empty_returns_empty(self):
|
||||
store, _, _ = _make_mongodb_store()
|
||||
result = store._insert_texts([], [])
|
||||
assert result == []
|
||||
|
||||
def test_insert_texts_builds_correct_documents(self):
|
||||
store, mock_collection, mock_emb = _make_mongodb_store()
|
||||
mock_emb.embed_documents.return_value = [[0.1, 0.2]]
|
||||
mock_collection.insert_many.return_value = Mock(inserted_ids=["id1"])
|
||||
|
||||
store._insert_texts(["hello"], [{"source": "test"}])
|
||||
|
||||
inserted_docs = mock_collection.insert_many.call_args[0][0]
|
||||
assert len(inserted_docs) == 1
|
||||
assert inserted_docs[0]["text"] == "hello"
|
||||
assert inserted_docs[0]["embedding"] == [0.1, 0.2]
|
||||
assert inserted_docs[0]["source"] == "test"
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestMongoDBVectorStoreDeleteIndex:
|
||||
def test_delete_index_calls_delete_many(self):
|
||||
store, mock_collection, _ = _make_mongodb_store(source_id="src1")
|
||||
|
||||
store.delete_index()
|
||||
|
||||
mock_collection.delete_many.assert_called_once_with({"source_id": "src1"})
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestMongoDBVectorStoreGetChunks:
|
||||
def test_get_chunks(self):
|
||||
store, mock_collection, _ = _make_mongodb_store()
|
||||
|
||||
docs = [
|
||||
{
|
||||
"_id": "id1",
|
||||
"text": "chunk1",
|
||||
"embedding": [0.1],
|
||||
"source_id": "src",
|
||||
"extra": "val",
|
||||
},
|
||||
{
|
||||
"_id": "id2",
|
||||
"text": "chunk2",
|
||||
"embedding": [0.2],
|
||||
"source_id": "src",
|
||||
},
|
||||
]
|
||||
mock_collection.find.return_value = iter(docs)
|
||||
|
||||
chunks = store.get_chunks()
|
||||
|
||||
assert len(chunks) == 2
|
||||
assert chunks[0]["doc_id"] == "id1"
|
||||
assert chunks[0]["text"] == "chunk1"
|
||||
assert chunks[0]["metadata"] == {"extra": "val"}
|
||||
assert "embedding" not in chunks[0]["metadata"]
|
||||
assert "source_id" not in chunks[0]["metadata"]
|
||||
|
||||
def test_get_chunks_skips_empty_text(self):
|
||||
store, mock_collection, _ = _make_mongodb_store()
|
||||
|
||||
docs = [
|
||||
{"_id": "id1", "text": None, "embedding": [0.1], "source_id": "src"},
|
||||
]
|
||||
mock_collection.find.return_value = iter(docs)
|
||||
|
||||
chunks = store.get_chunks()
|
||||
assert len(chunks) == 0
|
||||
|
||||
def test_get_chunks_returns_empty_on_error(self):
|
||||
store, mock_collection, _ = _make_mongodb_store()
|
||||
mock_collection.find.side_effect = Exception("connection error")
|
||||
|
||||
assert store.get_chunks() == []
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestMongoDBVectorStoreAddChunk:
|
||||
def test_add_chunk(self):
|
||||
store, mock_collection, mock_emb = _make_mongodb_store(source_id="src1")
|
||||
mock_emb.embed_documents.return_value = [[0.1, 0.2]]
|
||||
mock_collection.insert_one.return_value = Mock(inserted_id="new_id")
|
||||
|
||||
result = store.add_chunk("hello chunk", metadata={"key": "val"})
|
||||
|
||||
assert result == "new_id"
|
||||
inserted = mock_collection.insert_one.call_args[0][0]
|
||||
assert inserted["text"] == "hello chunk"
|
||||
assert inserted["source_id"] == "src1"
|
||||
assert inserted["key"] == "val"
|
||||
|
||||
def test_add_chunk_default_metadata(self):
|
||||
store, mock_collection, mock_emb = _make_mongodb_store()
|
||||
mock_emb.embed_documents.return_value = [[0.1]]
|
||||
mock_collection.insert_one.return_value = Mock(inserted_id="id")
|
||||
|
||||
store.add_chunk("text")
|
||||
|
||||
inserted = mock_collection.insert_one.call_args[0][0]
|
||||
assert "source_id" in inserted
|
||||
|
||||
def test_add_chunk_raises_on_empty_embedding(self):
|
||||
store, _, mock_emb = _make_mongodb_store()
|
||||
mock_emb.embed_documents.return_value = []
|
||||
|
||||
with pytest.raises(ValueError, match="Could not generate embedding"):
|
||||
store.add_chunk("text")
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestMongoDBVectorStoreDeleteChunk:
|
||||
def test_delete_chunk_success(self):
|
||||
store, mock_collection, _ = _make_mongodb_store()
|
||||
mock_collection.delete_one.return_value = Mock(deleted_count=1)
|
||||
|
||||
with patch("application.vectorstore.mongodb.ObjectId", create=True):
|
||||
# We need to mock bson.objectid.ObjectId
|
||||
with patch.dict("sys.modules", {"bson": MagicMock(), "bson.objectid": MagicMock()}):
|
||||
from unittest.mock import MagicMock as MM
|
||||
mock_oid = MM()
|
||||
with patch(
|
||||
"bson.objectid.ObjectId", return_value=mock_oid
|
||||
):
|
||||
result = store.delete_chunk("507f1f77bcf86cd799439011")
|
||||
|
||||
assert result is True
|
||||
|
||||
def test_delete_chunk_returns_false_on_error(self):
|
||||
store, mock_collection, _ = _make_mongodb_store()
|
||||
mock_collection.delete_one.side_effect = Exception("fail")
|
||||
|
||||
result = store.delete_chunk("bad_id")
|
||||
assert result is False
|
||||
@@ -0,0 +1,405 @@
|
||||
from unittest.mock import MagicMock, Mock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
def _make_store(
|
||||
source_id="test-source",
|
||||
embeddings_key="key",
|
||||
connection_string="postgresql://user:pass@localhost/db",
|
||||
):
|
||||
"""Helper to create a PGVectorStore with all external deps mocked."""
|
||||
with patch(
|
||||
"application.vectorstore.base.BaseVectorStore._get_embeddings"
|
||||
) as mock_get_emb, patch(
|
||||
"application.vectorstore.pgvector.settings"
|
||||
) as mock_settings, patch.dict(
|
||||
"sys.modules",
|
||||
{
|
||||
"psycopg": MagicMock(),
|
||||
"pgvector": MagicMock(),
|
||||
"pgvector.psycopg": MagicMock(),
|
||||
},
|
||||
):
|
||||
mock_emb = Mock()
|
||||
mock_emb.embed_query = Mock(return_value=[0.1, 0.2, 0.3])
|
||||
mock_emb.embed_documents = Mock(return_value=[[0.1, 0.2, 0.3]])
|
||||
mock_emb.dimension = 768
|
||||
mock_get_emb.return_value = mock_emb
|
||||
mock_settings.EMBEDDINGS_NAME = "test_model"
|
||||
mock_settings.PGVECTOR_CONNECTION_STRING = connection_string
|
||||
|
||||
from application.vectorstore.pgvector import PGVectorStore
|
||||
|
||||
# Patch _ensure_table_exists to avoid DB calls during init
|
||||
with patch.object(PGVectorStore, "_ensure_table_exists"):
|
||||
store = PGVectorStore(
|
||||
source_id=source_id,
|
||||
embeddings_key=embeddings_key,
|
||||
connection_string=connection_string,
|
||||
)
|
||||
# Provide a mock connection
|
||||
mock_conn = MagicMock()
|
||||
mock_cursor = MagicMock()
|
||||
mock_conn.cursor.return_value = mock_cursor
|
||||
mock_conn.closed = False
|
||||
store._connection = mock_conn
|
||||
|
||||
return store, mock_conn, mock_cursor, mock_emb
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestPGVectorStoreInit:
|
||||
def test_source_id_cleaned(self):
|
||||
store, _, _, _ = _make_store(source_id="application/indexes/abc123/")
|
||||
assert store._source_id == "abc123"
|
||||
|
||||
def test_missing_connection_string_raises(self):
|
||||
with patch(
|
||||
"application.vectorstore.base.BaseVectorStore._get_embeddings"
|
||||
) as mock_get_emb, patch(
|
||||
"application.vectorstore.pgvector.settings"
|
||||
) as mock_settings, patch.dict(
|
||||
"sys.modules",
|
||||
{
|
||||
"psycopg": MagicMock(),
|
||||
"pgvector": MagicMock(),
|
||||
"pgvector.psycopg": MagicMock(),
|
||||
},
|
||||
):
|
||||
mock_get_emb.return_value = Mock(dimension=768)
|
||||
mock_settings.EMBEDDINGS_NAME = "test_model"
|
||||
mock_settings.PGVECTOR_CONNECTION_STRING = None
|
||||
mock_settings.POSTGRES_URI = None
|
||||
|
||||
from application.vectorstore.pgvector import PGVectorStore
|
||||
|
||||
with pytest.raises(ValueError, match="connection string is required"):
|
||||
PGVectorStore(
|
||||
source_id="test", embeddings_key="key", connection_string=None
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestPGVectorStoreSearch:
|
||||
def test_search_returns_documents(self):
|
||||
store, mock_conn, mock_cursor, mock_emb = _make_store()
|
||||
mock_cursor.fetchall.return_value = [
|
||||
("hello world", {"source": "test.txt"}, 0.1),
|
||||
("foo bar", {"source": "test2.txt"}, 0.2),
|
||||
]
|
||||
|
||||
results = store.search("query", k=2)
|
||||
|
||||
mock_emb.embed_query.assert_called_once_with("query")
|
||||
assert len(results) == 2
|
||||
assert results[0].page_content == "hello world"
|
||||
assert results[0].metadata == {"source": "test.txt"}
|
||||
|
||||
def test_search_returns_empty_on_error(self):
|
||||
store, mock_conn, mock_cursor, _ = _make_store()
|
||||
mock_cursor.execute.side_effect = Exception("connection lost")
|
||||
|
||||
results = store.search("query")
|
||||
assert results == []
|
||||
|
||||
def test_search_handles_null_metadata(self):
|
||||
store, _, mock_cursor, _ = _make_store()
|
||||
mock_cursor.fetchall.return_value = [("text", None, 0.5)]
|
||||
|
||||
results = store.search("query")
|
||||
assert len(results) == 1
|
||||
assert results[0].metadata == {}
|
||||
|
||||
def test_score_threshold_filters_by_distance(self):
|
||||
# similarity = 1 - distance; threshold 0.85 → keep distance <= 0.15.
|
||||
store, _, mock_cursor, _ = _make_store()
|
||||
mock_cursor.fetchall.return_value = [
|
||||
("close", {}, 0.10), # sim 0.90 → kept
|
||||
("far", {}, 0.40), # sim 0.60 → dropped
|
||||
]
|
||||
results = store.search("query", k=5, score_threshold=0.85)
|
||||
assert [r.page_content for r in results] == ["close"]
|
||||
|
||||
def test_no_score_threshold_keeps_all(self):
|
||||
store, _, mock_cursor, _ = _make_store()
|
||||
mock_cursor.fetchall.return_value = [
|
||||
("a", {}, 0.10),
|
||||
("b", {}, 0.90),
|
||||
]
|
||||
results = store.search("query", k=5)
|
||||
assert len(results) == 2
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestPGVectorStoreKeywordSearch:
|
||||
def test_keyword_search_returns_documents(self):
|
||||
store, _, mock_cursor, mock_emb = _make_store(source_id="src1")
|
||||
mock_cursor.fetchall.return_value = [
|
||||
("hello world", {"source": "a.txt"}, 0.9),
|
||||
("foo bar", {"source": "b.txt"}, 0.3),
|
||||
]
|
||||
|
||||
results = store.keyword_search("hello", k=5)
|
||||
|
||||
# Keyword search must not embed the query.
|
||||
mock_emb.embed_query.assert_not_called()
|
||||
assert len(results) == 2
|
||||
assert results[0].page_content == "hello world"
|
||||
assert results[0].metadata == {"source": "a.txt"}
|
||||
|
||||
def test_keyword_search_is_parameterized(self):
|
||||
store, _, mock_cursor, _ = _make_store(source_id="src1")
|
||||
mock_cursor.fetchall.return_value = []
|
||||
|
||||
store.keyword_search("DROP TABLE documents; --", k=7)
|
||||
|
||||
sql, params = mock_cursor.execute.call_args[0]
|
||||
# The raw question must never be interpolated into the SQL text.
|
||||
assert "DROP TABLE documents" not in sql
|
||||
assert "websearch_to_tsquery('english', %s)" in sql
|
||||
# Question is bound twice (rank + WHERE), then source_id and k.
|
||||
assert params == ("DROP TABLE documents; --", "src1", "DROP TABLE documents; --", 7)
|
||||
|
||||
def test_keyword_search_handles_null_metadata(self):
|
||||
store, _, mock_cursor, _ = _make_store()
|
||||
mock_cursor.fetchall.return_value = [("text", None, 0.5)]
|
||||
|
||||
results = store.keyword_search("query")
|
||||
assert len(results) == 1
|
||||
assert results[0].metadata == {}
|
||||
|
||||
def test_keyword_search_returns_empty_on_error(self):
|
||||
store, _, mock_cursor, _ = _make_store()
|
||||
mock_cursor.execute.side_effect = Exception("fts failed")
|
||||
|
||||
assert store.keyword_search("query") == []
|
||||
|
||||
def test_ensure_table_exists_creates_fts_index(self):
|
||||
store, mock_conn, mock_cursor, _ = _make_store()
|
||||
store._ensure_table_exists()
|
||||
|
||||
executed = " ".join(
|
||||
str(call.args[0]) for call in mock_cursor.execute.call_args_list
|
||||
)
|
||||
assert "documents_text_fts_idx" in executed
|
||||
assert "gin(to_tsvector('english'" in executed
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestPGVectorStoreAddTexts:
|
||||
def test_add_texts_inserts_and_returns_ids(self):
|
||||
store, mock_conn, mock_cursor, mock_emb = _make_store()
|
||||
mock_emb.embed_documents.return_value = [[0.1, 0.2], [0.3, 0.4]]
|
||||
mock_cursor.fetchone.side_effect = [(1,), (2,)]
|
||||
|
||||
ids = store.add_texts(["text1", "text2"], [{"a": 1}, {"b": 2}])
|
||||
|
||||
assert ids == ["1", "2"]
|
||||
assert mock_cursor.execute.call_count == 2
|
||||
mock_conn.commit.assert_called_once()
|
||||
|
||||
def test_add_texts_empty_returns_empty(self):
|
||||
store, _, _, _ = _make_store()
|
||||
assert store.add_texts([]) == []
|
||||
|
||||
def test_add_texts_default_metadatas(self):
|
||||
store, mock_conn, mock_cursor, mock_emb = _make_store()
|
||||
mock_emb.embed_documents.return_value = [[0.1, 0.2]]
|
||||
mock_cursor.fetchone.return_value = (1,)
|
||||
|
||||
ids = store.add_texts(["text1"])
|
||||
assert ids == ["1"]
|
||||
|
||||
def test_add_texts_rolls_back_on_error(self):
|
||||
store, mock_conn, mock_cursor, mock_emb = _make_store()
|
||||
mock_emb.embed_documents.return_value = [[0.1]]
|
||||
mock_cursor.execute.side_effect = Exception("insert failed")
|
||||
|
||||
with pytest.raises(Exception, match="insert failed"):
|
||||
store.add_texts(["text1"])
|
||||
|
||||
mock_conn.rollback.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestPGVectorStoreDeleteIndex:
|
||||
def test_delete_index_deletes_by_source_id(self):
|
||||
store, mock_conn, mock_cursor, _ = _make_store(source_id="src123")
|
||||
|
||||
store.delete_index()
|
||||
|
||||
mock_cursor.execute.assert_called_once()
|
||||
sql = mock_cursor.execute.call_args[0][0]
|
||||
assert "DELETE FROM" in sql
|
||||
assert mock_cursor.execute.call_args[0][1] == ("src123",)
|
||||
mock_conn.commit.assert_called_once()
|
||||
|
||||
def test_delete_index_rolls_back_on_error(self):
|
||||
store, mock_conn, mock_cursor, _ = _make_store()
|
||||
mock_cursor.execute.side_effect = Exception("fail")
|
||||
|
||||
with pytest.raises(Exception):
|
||||
store.delete_index()
|
||||
|
||||
mock_conn.rollback.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestPGVectorStoreSaveLocal:
|
||||
def test_save_local_is_noop(self):
|
||||
store, _, _, _ = _make_store()
|
||||
assert store.save_local() is None
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestPGVectorStoreGetChunks:
|
||||
def test_get_chunks(self):
|
||||
store, _, mock_cursor, _ = _make_store()
|
||||
mock_cursor.fetchall.return_value = [
|
||||
(1, "text1", {"key": "val"}),
|
||||
(2, "text2", None),
|
||||
]
|
||||
|
||||
chunks = store.get_chunks()
|
||||
assert len(chunks) == 2
|
||||
assert chunks[0] == {"doc_id": "1", "text": "text1", "metadata": {"key": "val"}}
|
||||
assert chunks[1] == {"doc_id": "2", "text": "text2", "metadata": {}}
|
||||
|
||||
def test_get_chunks_returns_empty_on_error(self):
|
||||
store, _, mock_cursor, _ = _make_store()
|
||||
mock_cursor.execute.side_effect = Exception("fail")
|
||||
|
||||
assert store.get_chunks() == []
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestPGVectorStoreAddChunk:
|
||||
def test_add_chunk(self):
|
||||
store, mock_conn, mock_cursor, mock_emb = _make_store(source_id="src1")
|
||||
mock_emb.embed_documents.return_value = [[0.1, 0.2]]
|
||||
mock_cursor.fetchone.return_value = (42,)
|
||||
|
||||
chunk_id = store.add_chunk("hello", metadata={"key": "val"})
|
||||
|
||||
assert chunk_id == "42"
|
||||
mock_conn.commit.assert_called_once()
|
||||
|
||||
def test_add_chunk_raises_on_empty_embedding(self):
|
||||
store, _, _, mock_emb = _make_store()
|
||||
mock_emb.embed_documents.return_value = []
|
||||
|
||||
with pytest.raises(ValueError, match="Could not generate embedding"):
|
||||
store.add_chunk("text")
|
||||
|
||||
def test_add_chunk_includes_source_id_in_metadata(self):
|
||||
store, mock_conn, mock_cursor, mock_emb = _make_store(source_id="src1")
|
||||
mock_emb.embed_documents.return_value = [[0.1, 0.2]]
|
||||
mock_cursor.fetchone.return_value = (1,)
|
||||
|
||||
store.add_chunk("hello", metadata={"key": "val"})
|
||||
|
||||
# Verify source_id is passed as a parameter to the INSERT
|
||||
insert_call = mock_cursor.execute.call_args
|
||||
params = insert_call[0][1]
|
||||
# source_id is the 4th param in the insert
|
||||
assert params[3] == "src1"
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestPGVectorStoreDeleteChunk:
|
||||
def test_delete_chunk_success(self):
|
||||
store, mock_conn, mock_cursor, _ = _make_store()
|
||||
mock_cursor.rowcount = 1
|
||||
|
||||
result = store.delete_chunk("42")
|
||||
assert result is True
|
||||
mock_conn.commit.assert_called_once()
|
||||
|
||||
def test_delete_chunk_not_found(self):
|
||||
store, mock_conn, mock_cursor, _ = _make_store()
|
||||
mock_cursor.rowcount = 0
|
||||
|
||||
result = store.delete_chunk("999")
|
||||
assert result is False
|
||||
|
||||
def test_delete_chunk_returns_false_on_error(self):
|
||||
store, _, mock_cursor, _ = _make_store()
|
||||
mock_cursor.execute.side_effect = Exception("fail")
|
||||
|
||||
result = store.delete_chunk("42")
|
||||
assert result is False
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestPGVectorStoreDeleteChunksBySourcePath:
|
||||
def test_targeted_delete_is_parameterized(self):
|
||||
store, mock_conn, mock_cursor, _ = _make_store(source_id="src1")
|
||||
mock_cursor.rowcount = 3
|
||||
|
||||
deleted = store.delete_chunks_by_source_path("/docs/page.md")
|
||||
|
||||
assert deleted == 3
|
||||
sql, params = mock_cursor.execute.call_args[0]
|
||||
# Single targeted DELETE; the path is a bound param, never interpolated.
|
||||
assert "/docs/page.md" not in sql
|
||||
assert "DELETE FROM" in sql
|
||||
assert "metadata->>'source' = %s" in sql
|
||||
assert "source_id = %s" in sql
|
||||
assert params == ("src1", "/docs/page.md")
|
||||
mock_conn.commit.assert_called_once()
|
||||
|
||||
def test_returns_zero_when_no_match(self):
|
||||
store, mock_conn, mock_cursor, _ = _make_store(source_id="src1")
|
||||
mock_cursor.rowcount = 0
|
||||
|
||||
assert store.delete_chunks_by_source_path("/missing.md") == 0
|
||||
mock_conn.commit.assert_called_once()
|
||||
|
||||
def test_rolls_back_and_raises_on_error(self):
|
||||
store, mock_conn, mock_cursor, _ = _make_store()
|
||||
mock_cursor.execute.side_effect = Exception("delete failed")
|
||||
|
||||
with pytest.raises(Exception, match="delete failed"):
|
||||
store.delete_chunks_by_source_path("/x.md")
|
||||
|
||||
mock_conn.rollback.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestPGVectorStoreConnection:
|
||||
def test_get_connection_creates_new_when_closed(self):
|
||||
store, mock_conn, _, _ = _make_store()
|
||||
mock_conn.closed = True
|
||||
|
||||
mock_psycopg = MagicMock()
|
||||
new_conn = MagicMock()
|
||||
mock_psycopg.connect.return_value = new_conn
|
||||
store._psycopg = mock_psycopg
|
||||
|
||||
conn = store._get_connection()
|
||||
mock_psycopg.connect.assert_called_once()
|
||||
assert conn is new_conn
|
||||
|
||||
def test_get_connection_reuses_open(self):
|
||||
store, mock_conn, _, _ = _make_store()
|
||||
mock_conn.closed = False
|
||||
|
||||
conn = store._get_connection()
|
||||
assert conn is mock_conn
|
||||
|
||||
def test_ensure_table_exists(self):
|
||||
store, mock_conn, mock_cursor, _ = _make_store()
|
||||
# Call _ensure_table_exists directly
|
||||
store._ensure_table_exists()
|
||||
|
||||
# Should execute CREATE EXTENSION, CREATE TABLE, and CREATE INDEX statements
|
||||
assert mock_cursor.execute.call_count >= 3
|
||||
mock_conn.commit.assert_called()
|
||||
|
||||
def test_del_closes_connection(self):
|
||||
store, mock_conn, _, _ = _make_store()
|
||||
mock_conn.closed = False
|
||||
|
||||
store.__del__()
|
||||
mock_conn.close.assert_called_once()
|
||||
@@ -0,0 +1,230 @@
|
||||
from unittest.mock import MagicMock, Mock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
def _make_qdrant_store(source_id="test-source"):
|
||||
"""Helper to create a QdrantStore with all external deps mocked."""
|
||||
mock_models = MagicMock()
|
||||
mock_qdrant_langchain = MagicMock()
|
||||
|
||||
with patch(
|
||||
"application.vectorstore.base.BaseVectorStore._get_embeddings"
|
||||
) as mock_get_emb, patch(
|
||||
"application.vectorstore.qdrant.settings"
|
||||
) as mock_settings, patch.dict(
|
||||
"sys.modules",
|
||||
{
|
||||
"qdrant_client": MagicMock(),
|
||||
"qdrant_client.models": mock_models,
|
||||
"langchain_community": MagicMock(),
|
||||
"langchain_community.vectorstores": MagicMock(),
|
||||
"langchain_community.vectorstores.qdrant": mock_qdrant_langchain,
|
||||
},
|
||||
):
|
||||
mock_emb = Mock()
|
||||
mock_emb.embed_query = Mock(return_value=[0.1, 0.2, 0.3])
|
||||
mock_emb.embed_documents = Mock(return_value=[[0.1, 0.2, 0.3]])
|
||||
mock_emb.client = [None, Mock(word_embedding_dimension=768)]
|
||||
mock_get_emb.return_value = mock_emb
|
||||
|
||||
mock_settings.EMBEDDINGS_NAME = "test_model"
|
||||
mock_settings.QDRANT_COLLECTION_NAME = "test_collection"
|
||||
mock_settings.QDRANT_LOCATION = ":memory:"
|
||||
mock_settings.QDRANT_URL = None
|
||||
mock_settings.QDRANT_PORT = 6333
|
||||
mock_settings.QDRANT_GRPC_PORT = 6334
|
||||
mock_settings.QDRANT_HTTPS = False
|
||||
mock_settings.QDRANT_PREFER_GRPC = False
|
||||
mock_settings.QDRANT_API_KEY = None
|
||||
mock_settings.QDRANT_PREFIX = None
|
||||
mock_settings.QDRANT_TIMEOUT = None
|
||||
mock_settings.QDRANT_PATH = None
|
||||
mock_settings.QDRANT_DISTANCE_FUNC = "Cosine"
|
||||
|
||||
mock_docsearch = MagicMock()
|
||||
mock_collections = MagicMock()
|
||||
mock_collections.collections = [MagicMock(name="test_collection")]
|
||||
mock_docsearch.client.get_collections.return_value = mock_collections
|
||||
mock_qdrant_langchain.Qdrant.construct_instance.return_value = mock_docsearch
|
||||
|
||||
from application.vectorstore.qdrant import QdrantStore
|
||||
|
||||
store = QdrantStore(source_id=source_id, embeddings_key="key")
|
||||
|
||||
return store, mock_docsearch, mock_settings
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestQdrantStoreInit:
|
||||
def test_source_id_cleaned(self):
|
||||
store, _, _ = _make_qdrant_store(source_id="application/indexes/abc123/")
|
||||
assert store._source_id == "abc123"
|
||||
|
||||
def test_filter_constructed(self):
|
||||
store, _, _ = _make_qdrant_store(source_id="src1")
|
||||
assert store._filter is not None
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestQdrantStoreSearch:
|
||||
def test_search_delegates(self):
|
||||
store, mock_ds, _ = _make_qdrant_store()
|
||||
mock_ds.similarity_search.return_value = ["result1"]
|
||||
|
||||
results = store.search("query", k=5)
|
||||
|
||||
mock_ds.similarity_search.assert_called_once()
|
||||
assert results == ["result1"]
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestQdrantStoreAddTexts:
|
||||
def test_add_texts_delegates(self):
|
||||
store, mock_ds, _ = _make_qdrant_store()
|
||||
mock_ds.add_texts.return_value = ["id1"]
|
||||
|
||||
result = store.add_texts(["text1"], metadatas=[{"a": 1}])
|
||||
mock_ds.add_texts.assert_called_once_with(["text1"], metadatas=[{"a": 1}])
|
||||
assert result == ["id1"]
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestQdrantStoreSaveLocal:
|
||||
def test_save_local_is_noop(self):
|
||||
store, _, _ = _make_qdrant_store()
|
||||
assert store.save_local() is None
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestQdrantStoreDeleteIndex:
|
||||
def test_delete_index(self):
|
||||
store, mock_ds, _ = _make_qdrant_store()
|
||||
|
||||
with patch("application.vectorstore.qdrant.settings") as ms:
|
||||
ms.QDRANT_COLLECTION_NAME = "test_collection"
|
||||
store.delete_index()
|
||||
|
||||
mock_ds.client.delete.assert_called_once_with(
|
||||
collection_name="test_collection",
|
||||
points_selector=store._filter,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestQdrantStoreGetChunks:
|
||||
def test_get_chunks(self):
|
||||
store, mock_ds, _ = _make_qdrant_store()
|
||||
|
||||
record1 = MagicMock()
|
||||
record1.id = "id1"
|
||||
record1.payload = {
|
||||
"page_content": "text1",
|
||||
"metadata": {"source": "test"},
|
||||
}
|
||||
record2 = MagicMock()
|
||||
record2.id = "id2"
|
||||
record2.payload = {
|
||||
"page_content": "text2",
|
||||
"metadata": {"source": "test2"},
|
||||
}
|
||||
|
||||
# First call returns records with offset, second returns empty with None offset
|
||||
mock_ds.client.scroll.side_effect = [
|
||||
([record1, record2], None),
|
||||
]
|
||||
|
||||
chunks = store.get_chunks()
|
||||
|
||||
assert len(chunks) == 2
|
||||
assert chunks[0] == {
|
||||
"doc_id": "id1",
|
||||
"text": "text1",
|
||||
"metadata": {"source": "test"},
|
||||
}
|
||||
|
||||
def test_get_chunks_pagination(self):
|
||||
store, mock_ds, _ = _make_qdrant_store()
|
||||
|
||||
record1 = MagicMock()
|
||||
record1.id = "id1"
|
||||
record1.payload = {"page_content": "text1", "metadata": {}}
|
||||
|
||||
record2 = MagicMock()
|
||||
record2.id = "id2"
|
||||
record2.payload = {"page_content": "text2", "metadata": {}}
|
||||
|
||||
mock_ds.client.scroll.side_effect = [
|
||||
([record1], "offset_token"),
|
||||
([record2], None),
|
||||
]
|
||||
|
||||
chunks = store.get_chunks()
|
||||
assert len(chunks) == 2
|
||||
assert mock_ds.client.scroll.call_count == 2
|
||||
|
||||
def test_get_chunks_returns_empty_on_error(self):
|
||||
store, mock_ds, _ = _make_qdrant_store()
|
||||
mock_ds.client.scroll.side_effect = Exception("fail")
|
||||
|
||||
assert store.get_chunks() == []
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestQdrantStoreAddChunk:
|
||||
def test_add_chunk(self):
|
||||
store, mock_ds, _ = _make_qdrant_store(source_id="src1")
|
||||
mock_ds.add_documents.return_value = ["new-id"]
|
||||
|
||||
result = store.add_chunk("hello", metadata={"key": "val"})
|
||||
|
||||
assert result == "new-id"
|
||||
mock_ds.add_documents.assert_called_once()
|
||||
doc = mock_ds.add_documents.call_args[0][0][0]
|
||||
assert doc.page_content == "hello"
|
||||
assert doc.metadata["source_id"] == "src1"
|
||||
assert doc.metadata["key"] == "val"
|
||||
|
||||
def test_add_chunk_default_metadata(self):
|
||||
store, mock_ds, _ = _make_qdrant_store(source_id="src1")
|
||||
mock_ds.add_documents.return_value = ["id"]
|
||||
|
||||
store.add_chunk("text")
|
||||
|
||||
doc = mock_ds.add_documents.call_args[0][0][0]
|
||||
assert doc.metadata["source_id"] == "src1"
|
||||
|
||||
def test_add_chunk_fallback_id(self):
|
||||
store, mock_ds, _ = _make_qdrant_store()
|
||||
mock_ds.add_documents.return_value = []
|
||||
|
||||
result = store.add_chunk("text")
|
||||
# Should return the uuid that was generated
|
||||
assert result is not None
|
||||
assert isinstance(result, str)
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestQdrantStoreDeleteChunk:
|
||||
def test_delete_chunk_success(self):
|
||||
store, mock_ds, _ = _make_qdrant_store()
|
||||
|
||||
with patch("application.vectorstore.qdrant.settings") as ms:
|
||||
ms.QDRANT_COLLECTION_NAME = "test_collection"
|
||||
result = store.delete_chunk("chunk-id")
|
||||
|
||||
mock_ds.client.delete.assert_called_once_with(
|
||||
collection_name="test_collection",
|
||||
points_selector=["chunk-id"],
|
||||
)
|
||||
assert result is True
|
||||
|
||||
def test_delete_chunk_returns_false_on_error(self):
|
||||
store, mock_ds, _ = _make_qdrant_store()
|
||||
|
||||
with patch("application.vectorstore.qdrant.settings") as ms:
|
||||
ms.QDRANT_COLLECTION_NAME = "test_collection"
|
||||
mock_ds.client.delete.side_effect = Exception("fail")
|
||||
result = store.delete_chunk("bad-id")
|
||||
|
||||
assert result is False
|
||||
@@ -0,0 +1,85 @@
|
||||
"""Tests for the ``EMBEDDINGS_MAX_INPUT_TOKENS`` truncation net.
|
||||
|
||||
The remote embeddings server (e.g. llama.cpp) hard-rejects any single input
|
||||
larger than its physical batch size with a 500. When the setting is
|
||||
configured, ``RemoteEmbeddings`` clips each input to that many tokens before
|
||||
the request; the overflow is dropped (lossy by design).
|
||||
"""
|
||||
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from application.core.settings import settings
|
||||
from application.utils import get_encoding
|
||||
from application.vectorstore import base
|
||||
from application.vectorstore.base import RemoteEmbeddings
|
||||
|
||||
|
||||
def _capture_post(monkeypatch):
|
||||
"""Patch ``requests.post`` and return a dict recording the sent payload."""
|
||||
captured = {}
|
||||
|
||||
def fake_post(url, headers=None, json=None, timeout=None):
|
||||
captured["payload"] = json
|
||||
n_inputs = len(json["input"]) if isinstance(json["input"], list) else 1
|
||||
resp = MagicMock()
|
||||
resp.raise_for_status.return_value = None
|
||||
resp.json.return_value = {
|
||||
"data": [{"index": i, "embedding": [0.0]} for i in range(n_inputs)]
|
||||
}
|
||||
return resp
|
||||
|
||||
monkeypatch.setattr(base.requests, "post", fake_post)
|
||||
return captured
|
||||
|
||||
|
||||
def test_truncates_oversized_input_to_limit(monkeypatch):
|
||||
monkeypatch.setattr(settings, "EMBEDDINGS_MAX_INPUT_TOKENS", 10)
|
||||
captured = _capture_post(monkeypatch)
|
||||
enc = get_encoding()
|
||||
|
||||
long_text = " ".join(["word"] * 100) # ~100 tokens, far over the limit of 10
|
||||
emb = RemoteEmbeddings(api_url="https://example.test", model_name="m")
|
||||
emb.embed_documents([long_text])
|
||||
|
||||
sent = captured["payload"]["input"][0]
|
||||
assert sent == enc.decode(enc.encode(long_text)[:10])
|
||||
assert len(enc.encode(sent)) <= 10
|
||||
|
||||
|
||||
def test_short_input_is_unchanged(monkeypatch):
|
||||
monkeypatch.setattr(settings, "EMBEDDINGS_MAX_INPUT_TOKENS", 10)
|
||||
captured = _capture_post(monkeypatch)
|
||||
|
||||
short_text = "hello world"
|
||||
emb = RemoteEmbeddings(api_url="https://example.test", model_name="m")
|
||||
emb.embed_documents([short_text])
|
||||
|
||||
assert captured["payload"]["input"][0] == short_text
|
||||
|
||||
|
||||
def test_no_truncation_when_setting_unset(monkeypatch):
|
||||
monkeypatch.setattr(settings, "EMBEDDINGS_MAX_INPUT_TOKENS", None)
|
||||
captured = _capture_post(monkeypatch)
|
||||
enc = get_encoding()
|
||||
|
||||
long_text = " ".join(["word"] * 100)
|
||||
emb = RemoteEmbeddings(api_url="https://example.test", model_name="m")
|
||||
emb.embed_documents([long_text])
|
||||
|
||||
sent = captured["payload"]["input"][0]
|
||||
assert sent == long_text
|
||||
assert len(enc.encode(sent)) > 10
|
||||
|
||||
|
||||
def test_query_path_is_truncated(monkeypatch):
|
||||
"""``embed_query`` passes a bare string through the same net."""
|
||||
monkeypatch.setattr(settings, "EMBEDDINGS_MAX_INPUT_TOKENS", 10)
|
||||
captured = _capture_post(monkeypatch)
|
||||
enc = get_encoding()
|
||||
|
||||
long_text = " ".join(["word"] * 100)
|
||||
emb = RemoteEmbeddings(api_url="https://example.test", model_name="m")
|
||||
emb.embed_query(long_text)
|
||||
|
||||
sent = captured["payload"]["input"]
|
||||
assert sent == enc.decode(enc.encode(long_text)[:10])
|
||||
@@ -0,0 +1,39 @@
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
|
||||
from application.vectorstore.vector_creator import VectorCreator
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestVectorCreator:
|
||||
def test_registered_vectorstores(self):
|
||||
assert "faiss" in VectorCreator.vectorstores
|
||||
assert "elasticsearch" in VectorCreator.vectorstores
|
||||
assert "mongodb" in VectorCreator.vectorstores
|
||||
assert "qdrant" in VectorCreator.vectorstores
|
||||
assert "milvus" in VectorCreator.vectorstores
|
||||
assert "pgvector" in VectorCreator.vectorstores
|
||||
|
||||
def test_create_vectorstore_invalid_type(self):
|
||||
with pytest.raises(ValueError, match="No vectorstore class found for type"):
|
||||
VectorCreator.create_vectorstore("nonexistent")
|
||||
|
||||
def test_create_vectorstore_case_insensitive(self):
|
||||
with patch.object(
|
||||
VectorCreator.vectorstores["faiss"], "__init__", return_value=None
|
||||
) as mock_init:
|
||||
mock_init.return_value = None
|
||||
VectorCreator.create_vectorstore("FAISS", source_id="test", embeddings_key="key")
|
||||
mock_init.assert_called_once_with(source_id="test", embeddings_key="key")
|
||||
|
||||
def test_create_vectorstore_passes_args(self):
|
||||
with patch.object(
|
||||
VectorCreator.vectorstores["mongodb"], "__init__", return_value=None
|
||||
) as mock_init:
|
||||
VectorCreator.create_vectorstore(
|
||||
"mongodb", source_id="src1", embeddings_key="ek", database="mydb"
|
||||
)
|
||||
mock_init.assert_called_once_with(
|
||||
source_id="src1", embeddings_key="ek", database="mydb"
|
||||
)
|
||||
Reference in New Issue
Block a user