diff --git a/mem0/memory/main.py b/mem0/memory/main.py index 333d2510f..80a30cd68 100644 --- a/mem0/memory/main.py +++ b/mem0/memory/main.py @@ -341,6 +341,11 @@ def _build_session_scope(filters): return "&".join(parts) +def _entity_collection_name(provider: str, collection_name: str) -> str: + separator = "-" if provider == "s3_vectors" else "_" + return f"{collection_name}{separator}entities" + + setup_config() logger = logging.getLogger(__name__) @@ -417,7 +422,7 @@ class Memory(MemoryBase): """Lazily initialize entity store on first use.""" if self._entity_store is None: entity_config = _safe_deepcopy_config(self.config.vector_store.config) - entity_collection = f"{self.collection_name}_entities" + entity_collection = _entity_collection_name(self.config.vector_store.provider, self.collection_name) # Set collection name on the cloned config if hasattr(entity_config, 'collection_name'): entity_config.collection_name = entity_collection @@ -1889,7 +1894,7 @@ class AsyncMemory(MemoryBase): """Lazily initialize entity store on first use.""" if self._entity_store is None: entity_config = _safe_deepcopy_config(self.config.vector_store.config) - entity_collection = f"{self.collection_name}_entities" + entity_collection = _entity_collection_name(self.config.vector_store.provider, self.collection_name) if hasattr(entity_config, 'collection_name'): entity_config.collection_name = entity_collection elif isinstance(entity_config, dict): diff --git a/tests/test_memory.py b/tests/test_memory.py index ae3f3ae89..371dd71ae 100644 --- a/tests/test_memory.py +++ b/tests/test_memory.py @@ -6,6 +6,7 @@ import pytest from mem0 import Memory from mem0.configs.base import MemoryConfig +from mem0.memory.main import _entity_collection_name from mem0.memory.utils import normalize_facts @@ -37,6 +38,14 @@ def test_create_memory(memory_client): assert result["results"][0]["memory"] == data +def test_entity_collection_name_uses_dash_for_s3_vectors(): + assert _entity_collection_name("s3_vectors", "test-index") == "test-index-entities" + + +def test_entity_collection_name_keeps_underscore_for_other_stores(): + assert _entity_collection_name("qdrant", "mem0") == "mem0_entities" + + def test_get_memory(memory_client): data = "Name is John Doe." memory_client.add([{"role": "user", "content": data}], user_id="test_user") diff --git a/tests/vector_stores/test_s3_vectors.py b/tests/vector_stores/test_s3_vectors.py index aaa0e62dc..68b7fa0f6 100644 --- a/tests/vector_stores/test_s3_vectors.py +++ b/tests/vector_stores/test_s3_vectors.py @@ -110,6 +110,32 @@ def test_memory_initialization_with_config(mock_boto_client, mock_llm, mock_embe pytest.fail("Memory initialization failed") +def test_memory_entity_store_uses_s3_valid_index_name(mock_boto_client, mock_llm, mock_embedder): + mock_boto_client.get_vector_bucket.return_value = {} + mock_boto_client.get_index.return_value = {} + + config = { + "vector_store": { + "provider": "s3_vectors", + "config": { + "vector_bucket_name": BUCKET_NAME, + "collection_name": INDEX_NAME, + "embedding_model_dims": EMBEDDING_DIMS, + "distance_metric": "cosine", + "region_name": REGION, + }, + } + } + + memory = Memory.from_config(config) + assert memory.entity_store is not None + + index_names = [call.kwargs["indexName"] for call in mock_boto_client.get_index.call_args_list] + assert INDEX_NAME in index_names + assert f"{INDEX_NAME}-entities" in index_names + assert f"{INDEX_NAME}_entities" not in index_names + + def test_insert(mock_boto_client): """Test inserting vectors.""" store = S3Vectors(