e4c5582808
Co-authored-by: parshvadaftari <daftariparshva@gmail.com>
98 lines
4.0 KiB
Python
98 lines
4.0 KiB
Python
from unittest.mock import MagicMock, patch
|
|
|
|
import pytest
|
|
|
|
from mem0 import Memory
|
|
from mem0.configs.base import MemoryConfig
|
|
|
|
|
|
@pytest.fixture
|
|
def memory_client():
|
|
with patch.object(Memory, "__init__", return_value=None):
|
|
client = Memory()
|
|
client.add = MagicMock(return_value={"results": [{"id": "1", "memory": "Name is John Doe.", "event": "ADD"}]})
|
|
client.get = MagicMock(return_value={"id": "1", "memory": "Name is John Doe."})
|
|
client.update = MagicMock(return_value={"message": "Memory updated successfully!"})
|
|
client.delete = MagicMock(return_value={"message": "Memory deleted successfully!"})
|
|
client.history = MagicMock(return_value=[{"memory": "I like Indian food."}, {"memory": "I like Italian food."}])
|
|
client.get_all = MagicMock(return_value=["Name is John Doe.", "Name is John Doe. I like to code in Python."])
|
|
yield client
|
|
|
|
|
|
def test_create_memory(memory_client):
|
|
data = "Name is John Doe."
|
|
result = memory_client.add([{"role": "user", "content": data}], user_id="test_user")
|
|
assert result["results"][0]["memory"] == data
|
|
|
|
|
|
def test_get_memory(memory_client):
|
|
data = "Name is John Doe."
|
|
memory_client.add([{"role": "user", "content": data}], user_id="test_user")
|
|
result = memory_client.get("1")
|
|
assert result["memory"] == data
|
|
|
|
|
|
def test_update_memory(memory_client):
|
|
data = "Name is John Doe."
|
|
memory_client.add([{"role": "user", "content": data}], user_id="test_user")
|
|
new_data = "Name is John Kapoor."
|
|
update_result = memory_client.update("1", text=new_data)
|
|
assert update_result["message"] == "Memory updated successfully!"
|
|
|
|
|
|
def test_delete_memory(memory_client):
|
|
data = "Name is John Doe."
|
|
memory_client.add([{"role": "user", "content": data}], user_id="test_user")
|
|
delete_result = memory_client.delete("1")
|
|
assert delete_result["message"] == "Memory deleted successfully!"
|
|
|
|
|
|
def test_history(memory_client):
|
|
data = "I like Indian food."
|
|
memory_client.add([{"role": "user", "content": data}], user_id="test_user")
|
|
memory_client.update("1", text="I like Italian food.")
|
|
history = memory_client.history("1")
|
|
assert history[0]["memory"] == "I like Indian food."
|
|
assert history[1]["memory"] == "I like Italian food."
|
|
|
|
|
|
def test_list_memories(memory_client):
|
|
data1 = "Name is John Doe."
|
|
data2 = "Name is John Doe. I like to code in Python."
|
|
memory_client.add([{"role": "user", "content": data1}], user_id="test_user")
|
|
memory_client.add([{"role": "user", "content": data2}], user_id="test_user")
|
|
memories = memory_client.get_all(user_id="test_user")
|
|
assert data1 in memories
|
|
assert data2 in memories
|
|
|
|
|
|
@patch('mem0.utils.factory.EmbedderFactory.create')
|
|
@patch('mem0.utils.factory.VectorStoreFactory.create')
|
|
@patch('mem0.utils.factory.LlmFactory.create')
|
|
@patch('mem0.memory.storage.SQLiteManager')
|
|
def test_collection_name_preserved_after_reset(mock_sqlite, mock_llm_factory, mock_vector_factory, mock_embedder_factory):
|
|
mock_embedder_factory.return_value = MagicMock()
|
|
mock_vector_store = MagicMock()
|
|
mock_vector_factory.return_value = mock_vector_store
|
|
mock_llm_factory.return_value = MagicMock()
|
|
mock_sqlite.return_value = MagicMock()
|
|
|
|
test_collection_name = "mem0"
|
|
config = MemoryConfig()
|
|
config.vector_store.config.collection_name = test_collection_name
|
|
|
|
memory = Memory(config)
|
|
|
|
assert memory.collection_name == test_collection_name
|
|
assert memory.config.vector_store.config.collection_name == test_collection_name
|
|
|
|
memory.reset()
|
|
|
|
assert memory.collection_name == test_collection_name
|
|
assert memory.config.vector_store.config.collection_name == test_collection_name
|
|
|
|
reset_calls = [call for call in mock_vector_factory.call_args_list if len(mock_vector_factory.call_args_list) > 2]
|
|
if reset_calls:
|
|
reset_config = reset_calls[-1][0][1]
|
|
assert reset_config.collection_name == test_collection_name, f"Reset used wrong collection name: {reset_config.collection_name}"
|