from unittest.mock import Mock, patch import pytest from docsgpt.vectorstore.base import ( BaseVectorStore, EmbeddingsSingleton, RemoteEmbeddings, get_embeddings, ) HF_MPNET = "huggingface_sentence-transformers/all-mpnet-base-v2" LOCAL_MPNET = "/app/models/all-mpnet-base-v2" # --- 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("docsgpt.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("docsgpt.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("docsgpt.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("docsgpt.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("docsgpt.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("docsgpt.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("docsgpt.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("docsgpt.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("docsgpt.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("docsgpt.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("docsgpt.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("docsgpt.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("docsgpt.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("docsgpt.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("docsgpt.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("docsgpt.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" @patch("docsgpt.vectorstore.base.settings") @patch("docsgpt.vectorstore.base._get_embeddings_wrapper") def test_get_instance_hf_ignores_positional_key( self, mock_get_wrapper, mock_settings ): """A stray key must not reach the wrapper for a registered model. Registered models take their whole configuration from the registry, so a caller that passes ``settings.EMBEDDINGS_KEY`` positionally (as the vector stores do) must have it dropped rather than forwarded. """ mock_settings.EMBEDDINGS_BASE_URL = None 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(HF_MPNET, None) assert result is mock_instance # The configured name is passed through; the registry maps it to a repo. mock_wrapper_cls.assert_called_once_with(HF_MPNET) @patch("docsgpt.vectorstore.base.settings") @patch("docsgpt.vectorstore.base._get_embeddings_wrapper") def test_get_instance_hf_ignores_keyword_args( self, mock_get_wrapper, mock_settings ): mock_settings.EMBEDDINGS_BASE_URL = None mock_wrapper_cls = Mock() mock_get_wrapper.return_value = mock_wrapper_cls EmbeddingsSingleton.get_instance(HF_MPNET, openai_api_key="sk-nope") mock_wrapper_cls.assert_called_once_with(HF_MPNET) # --- 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("docsgpt.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("docsgpt.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("docsgpt.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("docsgpt.vectorstore.base.settings") @patch("docsgpt.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("docsgpt.vectorstore.base.settings") @patch("docsgpt.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("docsgpt.vectorstore.base.settings") @patch("docsgpt.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("docsgpt.vectorstore.base.settings") @patch("docsgpt.vectorstore.base.EmbeddingsSingleton.get_instance") def test_get_embeddings_registered_model_passes_configured_name( self, mock_get_instance, mock_settings ): """No bundled-path branch any more: the name goes straight through. FastEmbed resolves artifacts through its own cache (warmed in the image), so the old ``/app/models/...`` probe has no job to do. """ 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 mock_get_instance.assert_called_with( "huggingface_sentence-transformers/all-mpnet-base-v2" ) @patch("docsgpt.vectorstore.base.settings") @patch("docsgpt.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") @pytest.mark.unit class TestSearchWithScoresDefault: def test_pairs_hits_with_none(self): """A store that reports no score still satisfies the contract, so the retriever never has to special-case it.""" from docsgpt.vectorstore.base import BaseVectorStore class _Store(BaseVectorStore): def search(self, question, k=2, *args, **kwargs): return ["a", "b"] def add_texts(self, texts, metadatas=None, *args, **kwargs): return [] store = _Store() assert store.score_kind is None assert store.search_with_scores("q", k=2) == [("a", None), ("b", None)] def test_handles_store_returning_none(self): from docsgpt.vectorstore.base import BaseVectorStore class _Store(BaseVectorStore): def search(self, question, k=2, *args, **kwargs): return None def add_texts(self, texts, metadatas=None, *args, **kwargs): return [] assert _Store().search_with_scores("q") == [] # --- get_embeddings (the single resolver) --- @pytest.mark.unit class TestGetEmbeddingsResolver: """``get_embeddings`` is the one entry point every caller must use. Calling ``EmbeddingsSingleton.get_instance`` directly reproduces neither the bundled local-model path nor the OpenAI/Azure key handling. """ def setup_method(self): EmbeddingsSingleton._instances = {} @patch("docsgpt.vectorstore.base.settings") @patch("docsgpt.vectorstore.base._get_embeddings_wrapper") @patch("os.path.exists", return_value=False) def test_defaults_from_settings_do_not_raise( self, _mock_exists, mock_get_wrapper, mock_settings ): """The default config (HF mpnet name, no key) must resolve, not crash.""" mock_settings.EMBEDDINGS_BASE_URL = None mock_settings.EMBEDDINGS_NAME = HF_MPNET mock_settings.EMBEDDINGS_KEY = None mock_wrapper_cls = Mock() mock_instance = Mock() mock_wrapper_cls.return_value = mock_instance mock_get_wrapper.return_value = mock_wrapper_cls result = get_embeddings() assert result is mock_instance assert set(EmbeddingsSingleton._instances) == {HF_MPNET} @patch("docsgpt.vectorstore.base.settings") @patch("docsgpt.vectorstore.base._get_embeddings_wrapper") @patch("os.path.exists", return_value=False) def test_shares_cache_entry_with_vectorstore_helper( self, _mock_exists, mock_get_wrapper, mock_settings ): """Same object, same cache key as the vector stores get — one model.""" mock_settings.EMBEDDINGS_BASE_URL = None mock_settings.EMBEDDINGS_NAME = HF_MPNET mock_settings.EMBEDDINGS_KEY = None mock_wrapper_cls = Mock() mock_wrapper_cls.return_value = Mock() mock_get_wrapper.return_value = mock_wrapper_cls store_result = ConcreteVectorStore()._get_embeddings(HF_MPNET, None) resolver_result = get_embeddings() assert resolver_result is store_result assert set(EmbeddingsSingleton._instances) == {HF_MPNET} mock_wrapper_cls.assert_called_once() @patch("docsgpt.vectorstore.base.settings") @patch("docsgpt.vectorstore.base._get_embeddings_wrapper") def test_repeated_resolution_loads_one_model( self, mock_get_wrapper, mock_settings ): """A second call must not load a second copy of the model. The instance is keyed by the configured name. It used to be keyed by a bundled filesystem path when one happened to exist, which meant the same model could be cached twice under two keys. """ mock_settings.EMBEDDINGS_BASE_URL = None mock_settings.EMBEDDINGS_NAME = HF_MPNET mock_settings.EMBEDDINGS_KEY = None mock_wrapper_cls = Mock() mock_wrapper_cls.return_value = Mock() mock_get_wrapper.return_value = mock_wrapper_cls first = get_embeddings() second = get_embeddings() assert first is second assert set(EmbeddingsSingleton._instances) == {HF_MPNET} mock_wrapper_cls.assert_called_once_with(HF_MPNET) @patch("docsgpt.vectorstore.base.settings") def test_remote_when_base_url_configured(self, mock_settings): mock_settings.EMBEDDINGS_BASE_URL = "http://remote:8080" mock_settings.EMBEDDINGS_NAME = HF_MPNET mock_settings.EMBEDDINGS_KEY = "sk-remote" result = get_embeddings() assert isinstance(result, RemoteEmbeddings) assert result.api_url == "http://remote:8080" assert result.model_name == HF_MPNET assert result.headers["Authorization"] == "Bearer sk-remote" @patch("docsgpt.vectorstore.base.settings") @patch("docsgpt.vectorstore.base.EmbeddingsSingleton.get_instance") def test_openai_passes_key(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_settings.EMBEDDINGS_NAME = "openai_text-embedding-ada-002" mock_settings.EMBEDDINGS_KEY = "sk-from-settings" get_embeddings() mock_get_instance.assert_called_once_with( "openai_text-embedding-ada-002", openai_api_key="sk-from-settings" ) @patch("docsgpt.vectorstore.base.settings") @patch("docsgpt.vectorstore.base.EmbeddingsSingleton.get_instance") def test_openai_azure_uses_deployment_name( 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_settings.EMBEDDINGS_NAME = "openai_text-embedding-ada-002" mock_settings.EMBEDDINGS_KEY = "sk-key" get_embeddings() mock_get_instance.assert_called_once_with( "openai_text-embedding-ada-002", model="embed-deploy" ) @patch("docsgpt.vectorstore.base.settings") @patch("docsgpt.vectorstore.base.EmbeddingsSingleton.get_instance") def test_openai_alias_also_reaches_the_azure_deployment( self, mock_get_instance, mock_settings ): """The registry accepts the bare alias, so the key handling must too. Matching on the canonical string alone sent the alias down the generic branch, where the deployment name is never passed and Azure answers every embed with DeploymentNotFound. """ 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_settings.EMBEDDINGS_NAME = "text-embedding-ada-002" mock_settings.EMBEDDINGS_KEY = "sk-key" get_embeddings() mock_get_instance.assert_called_once_with( "text-embedding-ada-002", model="embed-deploy" ) @patch("docsgpt.vectorstore.base.settings") @patch("docsgpt.vectorstore.base.EmbeddingsSingleton.get_instance") def test_openai_name_is_matched_case_insensitively( 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_settings.EMBEDDINGS_NAME = "OpenAI_Text-Embedding-Ada-002" mock_settings.EMBEDDINGS_KEY = "sk-from-settings" get_embeddings() mock_get_instance.assert_called_once_with( "OpenAI_Text-Embedding-Ada-002", openai_api_key="sk-from-settings" ) @patch("docsgpt.vectorstore.base.settings") @patch("docsgpt.vectorstore.base.EmbeddingsSingleton.get_instance") def test_explicit_arguments_win_over_settings( self, mock_get_instance, mock_settings ): mock_settings.EMBEDDINGS_BASE_URL = None mock_settings.EMBEDDINGS_NAME = HF_MPNET mock_settings.EMBEDDINGS_KEY = "sk-from-settings" get_embeddings("some_custom_embedding", "sk-explicit") mock_get_instance.assert_called_once_with("some_custom_embedding") @patch("docsgpt.vectorstore.base.settings") @patch("docsgpt.vectorstore.base.get_embeddings") def test_vectorstore_helper_delegates_to_resolver( self, mock_resolver, _mock_settings ): """``BaseVectorStore._get_embeddings`` is a thin delegate now.""" sentinel = Mock() mock_resolver.return_value = sentinel result = ConcreteVectorStore()._get_embeddings("a-name", "a-key") assert result is sentinel mock_resolver.assert_called_once_with("a-name", "a-key") class _RecordingStore(ConcreteVectorStore): """Store whose add/delete calls are recorded, for the update fallback.""" def __init__(self, delete_result=True, delete_error=None): super().__init__() self.calls = [] self._delete_result = delete_result self._delete_error = delete_error def add_chunk(self, text, metadata=None, *args, **kwargs): self.calls.append(("add", text, metadata)) return "new-id" def delete_chunk(self, chunk_id, *args, **kwargs): self.calls.append(("delete", chunk_id)) if chunk_id != "old-id": return True if self._delete_error is not None: raise self._delete_error return self._delete_result @pytest.mark.unit class TestBaseUpdateChunkFallback: def test_adds_then_deletes_and_returns_the_new_id(self): store = _RecordingStore() new_id = store.update_chunk("old-id", "new text", {"k": "v"}) assert new_id == "new-id" assert store.calls == [("add", "new text", {"k": "v"}), ("delete", "old-id")] def test_false_delete_rolls_back_the_new_chunk_and_raises(self): store = _RecordingStore(delete_result=False) with pytest.raises(RuntimeError, match="old-id"): store.update_chunk("old-id", "new text", {}) assert store.calls == [ ("add", "new text", {}), ("delete", "old-id"), ("delete", "new-id"), ] def test_raising_delete_rolls_back_the_new_chunk_and_raises(self): store = _RecordingStore(delete_error=ConnectionError("milvus down")) with pytest.raises(RuntimeError, match="old-id") as excinfo: store.update_chunk("old-id", "new text", {}) assert isinstance(excinfo.value.__cause__, ConnectionError) assert store.calls[-1] == ("delete", "new-id") def test_failed_rollback_still_raises_the_update_error(self, caplog): store = _RecordingStore(delete_result=False) real_delete = store.delete_chunk def delete(chunk_id, *args, **kwargs): if chunk_id == "new-id": store.calls.append(("delete", chunk_id)) raise ConnectionError("rollback failed") return real_delete(chunk_id) store.delete_chunk = delete with caplog.at_level("ERROR"), pytest.raises(RuntimeError, match="old-id"): store.update_chunk("old-id", "new text", {}) assert ("delete", "new-id") in store.calls assert "new-id" in caplog.text def test_failed_add_skips_the_delete(self): store = _RecordingStore() store.add_chunk = Mock(side_effect=RuntimeError("embed down")) with pytest.raises(RuntimeError): store.update_chunk("old-id", "new text", {}) assert ("delete", "old-id") not in store.calls def test_milvus_keeps_the_default(self): """Milvus has no in-place update here; it re-adds under a new id.""" from docsgpt.vectorstore.milvus import MilvusStore assert MilvusStore.update_chunk is BaseVectorStore.update_chunk