feat: add es headers config (#3088)

This commit is contained in:
AkisAya
2025-08-19 03:45:23 +08:00
committed by GitHub
parent 49ad64708b
commit 346b913ace
4 changed files with 69 additions and 1 deletions
@@ -54,7 +54,8 @@ Let's see the available parameters for the `elasticsearch` config:
| `password` | Password for basic authentication | `None` |
| `verify_certs` | Whether to verify SSL certificates | `True` |
| `auto_create_index` | Whether to automatically create the index | `True` |
| `custom_search_query` | Function returning a custom search query | `None` |
| `custom_search_query` | Function returning a custom search query | `None` |
| `headers` | Custom headers to include in requests | `None` |
### Features
@@ -19,6 +19,7 @@ class ElasticsearchConfig(BaseModel):
custom_search_query: Optional[Callable[[List[float], int, Optional[Dict]], Dict]] = Field(
None, description="Custom search query function. Parameters: (query, limit, filters) -> Dict"
)
headers: Optional[Dict[str, str]] = Field(None, description="Custom headers to include in requests")
@model_validator(mode="before")
@classmethod
@@ -33,6 +34,23 @@ class ElasticsearchConfig(BaseModel):
return values
@model_validator(mode="before")
@classmethod
def validate_headers(cls, values: Dict[str, Any]) -> Dict[str, Any]:
"""Validate headers format and content"""
headers = values.get("headers")
if headers is not None:
# Check if headers is a dictionary
if not isinstance(headers, dict):
raise ValueError("headers must be a dictionary")
# Check if all keys and values are strings
for key, value in headers.items():
if not isinstance(key, str) or not isinstance(value, str):
raise ValueError("All header keys and values must be strings")
return values
@model_validator(mode="before")
@classmethod
def validate_extra_fields(cls, values: Dict[str, Any]) -> Dict[str, Any]:
+2
View File
@@ -31,12 +31,14 @@ class ElasticsearchDB(VectorStoreBase):
cloud_id=config.cloud_id,
api_key=config.api_key,
verify_certs=config.verify_certs,
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,
headers= config.headers or {},
)
self.collection_name = config.collection_name
+47
View File
@@ -10,6 +10,7 @@ except ImportError:
raise ImportError("Elasticsearch requires extra dependencies. Install with `pip install elasticsearch`") from None
from mem0.vector_stores.elasticsearch import ElasticsearchDB, OutputData
from mem0.configs.vector_stores.elasticsearch import ElasticsearchConfig
class TestElasticsearchDB(unittest.TestCase):
@@ -309,3 +310,49 @@ class TestElasticsearchDB(unittest.TestCase):
# Verify delete call
self.client_mock.indices.delete.assert_called_once_with(index="test_collection")
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)
self.assertEqual(es_config.user, "elastic")
self.assertEqual(es_config.password, "password")
def test_es_valid_headers(self):
config = {
"host": "localhost",
"port": 9200,
"user": "elastic",
"password": "password",
"headers": {"x-extra-info": "my-mem0-instance"},
}
es_config = ElasticsearchConfig(**config)
self.assertIsNotNone(es_config.headers)
self.assertEqual(len(es_config.headers), 1)
self.assertEqual(es_config.headers["x-extra-info"], "my-mem0-instance")
def test_es_invalid_headers(self):
base_config = {
"host": "localhost",
"port": 9200,
"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}
ElasticsearchConfig(**config)