feat: add es headers config (#3088)
This commit is contained in:
@@ -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]:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user