diff --git a/mem0/vector_stores/elasticsearch.py b/mem0/vector_stores/elasticsearch.py index 6578e9863d..09f3668fbc 100644 --- a/mem0/vector_stores/elasticsearch.py +++ b/mem0/vector_stores/elasticsearch.py @@ -3,7 +3,7 @@ from typing import Any, Dict, List, Optional try: - from elasticsearch import Elasticsearch + from elasticsearch import Elasticsearch, NotFoundError from elasticsearch.helpers import bulk except ImportError: raise ImportError("Elasticsearch requires extra dependencies. Install with `pip install elasticsearch`") from None @@ -29,10 +29,7 @@ def _validate_filter(key: str, value: Any) -> None: if not isinstance(key, str) or not _SAFE_FILTER_KEY.match(key): raise ValueError(f"Invalid filter key: {key!r}") if not isinstance(value, (str, int, float, bool)): - raise ValueError( - f"Filter value for {key!r} must be str, int, float, or bool, " - f"got {type(value).__name__}" - ) + raise ValueError(f"Filter value for {key!r} must be str, int, float, or bool, got {type(value).__name__}") class ElasticsearchDB(VectorStoreBase): @@ -46,7 +43,7 @@ def __init__(self, **kwargs): api_key=config.api_key, verify_certs=config.verify_certs, ca_certs=config.ca_certs, - headers= config.headers or {}, + headers=config.headers or {}, ) else: self.client = Elasticsearch( @@ -54,7 +51,7 @@ def __init__(self, **kwargs): basic_auth=(config.user, config.password) if (config.user and config.password) else None, verify_certs=config.verify_certs, ca_certs=config.ca_certs, - headers= config.headers or {}, + headers=config.headers or {}, ) self.collection_name = config.collection_name @@ -244,22 +241,28 @@ def update(self, vector_id: str, vector: Optional[List[float]] = None, payload: self.client.update(index=self.collection_name, id=vector_id, body={"doc": doc}) def get(self, vector_id: str) -> Optional[OutputData]: - """Retrieve a vector by ID.""" + """Retrieve a vector by ID. + + Returns None only when the vector does not exist or the response is + malformed; transport, auth and server failures raise, so callers can + tell a backend outage apart from a missing vector. + """ try: response = self.client.get(index=self.collection_name, id=vector_id) + except NotFoundError: + return None + except Exception as e: + logger.error(f"Failed to fetch vector {vector_id} from Elasticsearch: {e}") + raise + + try: return OutputData( id=response["_id"], score=1.0, # Default score for direct get payload=response["_source"].get("metadata", {}), ) - except KeyError as e: - logger.warning(f"Missing key in Elasticsearch response: {e}") - return None - except TypeError as e: - logger.warning(f"Invalid response type from Elasticsearch: {e}") - return None - except Exception as e: - logger.error(f"Unexpected error while parsing Elasticsearch response: {e}") + except (KeyError, TypeError) as e: + logger.warning(f"Malformed Elasticsearch response for vector {vector_id}: {e}") return None def list_cols(self) -> List[str]: diff --git a/tests/vector_stores/test_elasticsearch.py b/tests/vector_stores/test_elasticsearch.py index 57bc9d81fe..d96fa7f24b 100644 --- a/tests/vector_stores/test_elasticsearch.py +++ b/tests/vector_stores/test_elasticsearch.py @@ -5,7 +5,7 @@ import dotenv try: - from elasticsearch import Elasticsearch + from elasticsearch import ConnectionError, Elasticsearch, NotFoundError except ImportError: raise ImportError("Elasticsearch requires extra dependencies. Install with `pip install elasticsearch`") from None @@ -295,13 +295,31 @@ def test_get(self): self.assertEqual(result.payload, {"key": "value"}) def test_get_not_found(self): - # Mock get raising exception - self.client_mock.get.side_effect = Exception("Not found") + # elasticsearch signals a missing document with NotFoundError + self.client_mock.get.side_effect = NotFoundError( + "document_missing_exception", meta=Mock(), body={"found": False} + ) # Verify get returns None when document not found result = self.es_db.get(vector_id="nonexistent") self.assertIsNone(result) + def test_get_transport_error_propagates(self): + # Transport/auth/server failures must not read as "no such vector": + # callers treat None as missing and would insert duplicates (issue #7516). + self.client_mock.get.side_effect = ConnectionError("connection refused") + + with self.assertRaises(ConnectionError): + self.es_db.get(vector_id="id1") + + def test_get_malformed_response_returns_none(self): + # Parse-shape failures stay non-fatal + self.client_mock.get.return_value = {"_id": "id1"} # missing "_source" + self.assertIsNone(self.es_db.get(vector_id="id1")) + + self.client_mock.get.return_value = None # wrong response type + self.assertIsNone(self.es_db.get(vector_id="id1")) + def test_list(self): # Mock search response with scores mock_response = { @@ -357,11 +375,11 @@ def test_delete_col(self): def test_es_config(self): config = {"host": "localhost", "port": 9200, "user": "elastic", "password": "password"} es_config = ElasticsearchConfig(**config) - + # Assert that the config object was created successfully self.assertIsNotNone(es_config) self.assertIsInstance(es_config, ElasticsearchConfig) - + # Assert that the configuration values are correctly set self.assertEqual(es_config.host, "localhost") self.assertEqual(es_config.port, 9200) @@ -388,13 +406,13 @@ def test_es_invalid_headers(self): "user": "elastic", "password": "password", } - + invalid_headers = [ "not-a-dict", # Non-dict headers {"x-extra-info": 123}, # Non-string values {123: "456"}, # Non-string keys ] - + for headers in invalid_headers: with self.assertRaises(ValueError): config = {**base_config, "headers": headers}