From 346b913ace0a4b3794ab9c86e13339e35fb6db31 Mon Sep 17 00:00:00 2001 From: AkisAya Date: Tue, 19 Aug 2025 03:45:23 +0800 Subject: [PATCH] feat: add es headers config (#3088) --- .../vectordbs/dbs/elasticsearch.mdx | 3 +- mem0/configs/vector_stores/elasticsearch.py | 18 +++++++ mem0/vector_stores/elasticsearch.py | 2 + tests/vector_stores/test_elasticsearch.py | 47 +++++++++++++++++++ 4 files changed, 69 insertions(+), 1 deletion(-) diff --git a/docs/components/vectordbs/dbs/elasticsearch.mdx b/docs/components/vectordbs/dbs/elasticsearch.mdx index d8918112b..5e735d232 100644 --- a/docs/components/vectordbs/dbs/elasticsearch.mdx +++ b/docs/components/vectordbs/dbs/elasticsearch.mdx @@ -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 diff --git a/mem0/configs/vector_stores/elasticsearch.py b/mem0/configs/vector_stores/elasticsearch.py index 7f76c238d..ed12d8625 100644 --- a/mem0/configs/vector_stores/elasticsearch.py +++ b/mem0/configs/vector_stores/elasticsearch.py @@ -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]: diff --git a/mem0/vector_stores/elasticsearch.py b/mem0/vector_stores/elasticsearch.py index 016f3783b..b73eedcdd 100644 --- a/mem0/vector_stores/elasticsearch.py +++ b/mem0/vector_stores/elasticsearch.py @@ -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 diff --git a/tests/vector_stores/test_elasticsearch.py b/tests/vector_stores/test_elasticsearch.py index 8fd7b5500..db7a82e14 100644 --- a/tests/vector_stores/test_elasticsearch.py +++ b/tests/vector_stores/test_elasticsearch.py @@ -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)