Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
35 changes: 19 additions & 16 deletions mem0/vector_stores/elasticsearch.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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):
Expand All @@ -46,15 +43,15 @@ 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(
hosts=[f"{config.host}" if config.port is None else f"{config.host}:{config.port}"],
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
Expand Down Expand Up @@ -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]:
Expand Down
32 changes: 25 additions & 7 deletions tests/vector_stores/test_elasticsearch.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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 = {
Expand Down Expand Up @@ -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)
Expand All @@ -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}
Expand Down
Loading