feat(memory): expose expiration controls in client docs (#5874)
Co-authored-by: Soumil Rathi <soumilrathi@gmail.com>
This commit is contained in:
@@ -167,6 +167,33 @@ class TestAsyncUpdate:
|
||||
"test_id", "Updated memory", {"Updated memory": [0.1, 0.2, 0.3]}, {}
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_update_can_change_expiration_date_without_changing_text(self, mock_async_memory, mocker):
|
||||
mock_async_memory.embedding_model.embed = Mock(return_value=[0.1, 0.2, 0.3])
|
||||
mock_async_memory.vector_store.get = Mock(
|
||||
return_value=Mock(
|
||||
payload={
|
||||
"data": "Existing memory",
|
||||
"user_id": "test_user",
|
||||
"created_at": "2026-01-01T00:00:00+00:00",
|
||||
"expiration_date": "2026-12-31",
|
||||
}
|
||||
)
|
||||
)
|
||||
mock_async_memory.vector_store.update = Mock()
|
||||
mock_async_memory.db.add_history = Mock()
|
||||
mock_async_memory._remove_memory_from_entity_store = mocker.AsyncMock()
|
||||
mock_async_memory._link_entities_for_memory = mocker.AsyncMock()
|
||||
|
||||
result = await mock_async_memory.update("test_id", expiration_date="2999-01-01")
|
||||
|
||||
assert result["message"] == "Memory updated successfully!"
|
||||
payload = mock_async_memory.vector_store.update.call_args.kwargs["payload"]
|
||||
assert payload["data"] == "Existing memory"
|
||||
assert payload["expiration_date"] == "2999-01-01"
|
||||
mock_async_memory._remove_memory_from_entity_store.assert_not_called()
|
||||
mock_async_memory._link_entities_for_memory.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
class TestAsyncAddToVectorStoreErrors:
|
||||
|
||||
@@ -51,6 +51,20 @@ class TestSearchEntityParamRejection:
|
||||
json={"query": "test query", "filters": {"user_id": "u1"}},
|
||||
)
|
||||
|
||||
def test_search_passes_show_expired(self, mock_memory_client):
|
||||
"""search() should pass show_expired to the API."""
|
||||
mock_response = MagicMock()
|
||||
mock_response.json.return_value = {"results": []}
|
||||
mock_response.raise_for_status.return_value = None
|
||||
mock_memory_client.client.post.return_value = mock_response
|
||||
|
||||
mock_memory_client.search("test query", filters={"user_id": "u1"}, show_expired=True)
|
||||
|
||||
mock_memory_client.client.post.assert_called_once_with(
|
||||
"/v3/memories/search/",
|
||||
json={"query": "test query", "filters": {"user_id": "u1"}, "show_expired": True},
|
||||
)
|
||||
|
||||
def test_search_rejects_user_id_kwarg(self, mock_memory_client):
|
||||
"""search() should reject user_id as top-level kwarg."""
|
||||
with pytest.raises(ValueError, match=r"user_id"):
|
||||
@@ -100,6 +114,39 @@ class TestGetAllEntityParamRejection:
|
||||
with pytest.raises(ValueError, match=r"run_id"):
|
||||
mock_memory_client.get_all(run_id="r1")
|
||||
|
||||
def test_get_all_passes_show_expired(self, mock_memory_client):
|
||||
"""get_all() should pass show_expired to the API."""
|
||||
mock_response = MagicMock()
|
||||
mock_response.json.return_value = {"results": []}
|
||||
mock_response.raise_for_status.return_value = None
|
||||
mock_memory_client.client.post.return_value = mock_response
|
||||
|
||||
mock_memory_client.get_all(filters={"user_id": "u1"}, show_expired=True)
|
||||
|
||||
mock_memory_client.client.post.assert_called_once_with(
|
||||
"/v3/memories/",
|
||||
json={"filters": {"user_id": "u1"}, "show_expired": True},
|
||||
)
|
||||
|
||||
|
||||
class TestUpdateExpirationDate:
|
||||
"""Tests for update expiration_date payload handling."""
|
||||
|
||||
def test_update_preserves_null_expiration_date(self, mock_memory_client):
|
||||
"""update() should send expiration_date=None so the API can clear it."""
|
||||
mock_response = MagicMock()
|
||||
mock_response.json.return_value = {"id": "mem_1", "expiration_date": None}
|
||||
mock_response.raise_for_status.return_value = None
|
||||
mock_memory_client.client.put.return_value = mock_response
|
||||
|
||||
mock_memory_client.update("mem_1", expiration_date=None)
|
||||
|
||||
mock_memory_client.client.put.assert_called_once_with(
|
||||
"/v1/memories/mem_1/",
|
||||
json={"expiration_date": None},
|
||||
params={},
|
||||
)
|
||||
|
||||
|
||||
class TestFilterOperatorPassthrough:
|
||||
"""Tests that AND/OR/NOT filter operators are passed through to the API."""
|
||||
|
||||
+111
-1
@@ -65,6 +65,24 @@ def test_add(memory_instance):
|
||||
)
|
||||
|
||||
|
||||
def test_add_stores_expiration_date(memory_instance):
|
||||
memory_instance._add_to_vector_store = Mock(return_value=[{"memory": "Test memory", "event": "ADD"}])
|
||||
|
||||
memory_instance.add(
|
||||
messages=[{"role": "user", "content": "Test message"}],
|
||||
user_id="test_user",
|
||||
expiration_date="2999-01-01",
|
||||
)
|
||||
|
||||
memory_instance._add_to_vector_store.assert_called_once_with(
|
||||
[{"role": "user", "content": "Test message"}],
|
||||
{"user_id": "test_user", "expiration_date": "2999-01-01"},
|
||||
{"user_id": "test_user"},
|
||||
True,
|
||||
prompt=None,
|
||||
)
|
||||
|
||||
|
||||
def test_get(memory_instance):
|
||||
mock_memory = Mock(
|
||||
id="test_id",
|
||||
@@ -117,6 +135,39 @@ def test_search(memory_instance):
|
||||
)
|
||||
|
||||
|
||||
def test_search_hides_expired_memories_by_default(memory_instance):
|
||||
mock_memories = [
|
||||
Mock(id="1", payload={"data": "Expired memory", "user_id": "test_user", "expiration_date": "2000-01-01"}, score=0.9),
|
||||
Mock(id="2", payload={"data": "Active memory", "user_id": "test_user", "expiration_date": "2999-01-01"}, score=0.8),
|
||||
]
|
||||
memory_instance.vector_store.search = Mock(return_value=mock_memories)
|
||||
memory_instance.vector_store.keyword_search = Mock(return_value=None)
|
||||
memory_instance.embedding_model.embed = Mock(return_value=[0.1, 0.2, 0.3])
|
||||
|
||||
with patch("mem0.memory.main.lemmatize_for_bm25", return_value="test query"), \
|
||||
patch("mem0.memory.main.extract_entities", return_value=[]):
|
||||
result = memory_instance.search("test query", filters={"user_id": "test_user"})
|
||||
|
||||
assert [memory["memory"] for memory in result["results"]] == ["Active memory"]
|
||||
assert result["results"][0]["expiration_date"] == "2999-01-01"
|
||||
|
||||
|
||||
def test_search_can_show_expired_memories(memory_instance):
|
||||
mock_memories = [
|
||||
Mock(id="1", payload={"data": "Expired memory", "user_id": "test_user", "expiration_date": "2000-01-01"}, score=0.9),
|
||||
Mock(id="2", payload={"data": "Active memory", "user_id": "test_user"}, score=0.8),
|
||||
]
|
||||
memory_instance.vector_store.search = Mock(return_value=mock_memories)
|
||||
memory_instance.vector_store.keyword_search = Mock(return_value=None)
|
||||
memory_instance.embedding_model.embed = Mock(return_value=[0.1, 0.2, 0.3])
|
||||
|
||||
with patch("mem0.memory.main.lemmatize_for_bm25", return_value="test query"), \
|
||||
patch("mem0.memory.main.extract_entities", return_value=[]):
|
||||
result = memory_instance.search("test query", filters={"user_id": "test_user"}, show_expired=True)
|
||||
|
||||
assert [memory["memory"] for memory in result["results"]] == ["Expired memory", "Active memory"]
|
||||
|
||||
|
||||
def test_update(memory_instance):
|
||||
memory_instance.embedding_model = Mock()
|
||||
memory_instance.embedding_model.embed = Mock(return_value=[0.1, 0.2, 0.3])
|
||||
@@ -161,6 +212,42 @@ def test_update_with_empty_metadata(memory_instance):
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("expiration_date", "expected_expiration_date"),
|
||||
[
|
||||
("2999-01-01", "2999-01-01"),
|
||||
(None, None),
|
||||
],
|
||||
)
|
||||
def test_update_can_change_expiration_date_without_changing_text(
|
||||
memory_instance, expiration_date, expected_expiration_date
|
||||
):
|
||||
memory_instance.embedding_model.embed = Mock(return_value=[0.1, 0.2, 0.3])
|
||||
memory_instance.vector_store.get = Mock(
|
||||
return_value=Mock(
|
||||
payload={
|
||||
"data": "Existing memory",
|
||||
"user_id": "test_user",
|
||||
"created_at": "2026-01-01T00:00:00+00:00",
|
||||
"expiration_date": "2026-12-31",
|
||||
}
|
||||
)
|
||||
)
|
||||
memory_instance.vector_store.update = Mock()
|
||||
memory_instance.db.add_history = Mock()
|
||||
memory_instance._remove_memory_from_entity_store = Mock()
|
||||
memory_instance._link_entities_for_memory = Mock()
|
||||
|
||||
result = memory_instance.update("test_id", expiration_date=expiration_date)
|
||||
|
||||
assert result["message"] == "Memory updated successfully!"
|
||||
payload = memory_instance.vector_store.update.call_args.kwargs["payload"]
|
||||
assert payload["data"] == "Existing memory"
|
||||
assert payload["expiration_date"] == expected_expiration_date
|
||||
memory_instance._remove_memory_from_entity_store.assert_not_called()
|
||||
memory_instance._link_entities_for_memory.assert_not_called()
|
||||
|
||||
|
||||
def test_delete(memory_instance):
|
||||
memory_instance._delete_memory = Mock()
|
||||
|
||||
@@ -200,7 +287,30 @@ def test_get_all(memory_instance):
|
||||
assert result["results"][0]["memory"] == "Memory 1"
|
||||
assert result["results"][0]["user_id"] == "test_user"
|
||||
|
||||
memory_instance.vector_store.list.assert_called_once_with(filters={"user_id": "test_user"}, top_k=20)
|
||||
|
||||
def test_get_all_hides_expired_memories_by_default(memory_instance):
|
||||
mock_memories = [
|
||||
Mock(id="1", payload={"data": "Expired memory", "user_id": "test_user", "expiration_date": "2000-01-01"}),
|
||||
Mock(id="2", payload={"data": "Active memory", "user_id": "test_user", "expiration_date": "2999-01-01"}),
|
||||
]
|
||||
memory_instance.vector_store.list = Mock(return_value=(mock_memories, None))
|
||||
|
||||
result = memory_instance.get_all(filters={"user_id": "test_user"})
|
||||
|
||||
assert [memory["memory"] for memory in result["results"]] == ["Active memory"]
|
||||
assert result["results"][0]["expiration_date"] == "2999-01-01"
|
||||
|
||||
|
||||
def test_get_all_can_show_expired_memories(memory_instance):
|
||||
mock_memories = [
|
||||
Mock(id="1", payload={"data": "Expired memory", "user_id": "test_user", "expiration_date": "2000-01-01"}),
|
||||
Mock(id="2", payload={"data": "Active memory", "user_id": "test_user"}),
|
||||
]
|
||||
memory_instance.vector_store.list = Mock(return_value=(mock_memories, None))
|
||||
|
||||
result = memory_instance.get_all(filters={"user_id": "test_user"}, show_expired=True)
|
||||
|
||||
assert [memory["memory"] for memory in result["results"]] == ["Expired memory", "Active memory"]
|
||||
|
||||
|
||||
def test_no_telemetry_vector_store_when_disabled():
|
||||
|
||||
@@ -564,12 +564,20 @@ class TestUpdateMemory:
|
||||
resp = client.put("/memories/mem-1", json={"text": "Likes tennis"})
|
||||
assert resp.status_code == 200
|
||||
_, kwargs = mock_memory.update.call_args
|
||||
assert kwargs["metadata"] is None
|
||||
assert "metadata" not in kwargs
|
||||
|
||||
def test_missing_text_returns_422(self, client):
|
||||
"""text is required — omitting it should fail validation."""
|
||||
resp = client.put("/memories/mem-1", json={"metadata": {"k": "v"}})
|
||||
assert resp.status_code == 422
|
||||
def test_expiration_date_forwarded_without_text(self, client, mock_memory):
|
||||
resp = client.put("/memories/mem-1", json={"expiration_date": "2999-01-01"})
|
||||
assert resp.status_code == 200
|
||||
_, kwargs = mock_memory.update.call_args
|
||||
assert kwargs["expiration_date"] == "2999-01-01"
|
||||
assert "data" not in kwargs
|
||||
|
||||
def test_null_expiration_date_forwarded_for_clear(self, client, mock_memory):
|
||||
resp = client.put("/memories/mem-1", json={"expiration_date": None})
|
||||
assert resp.status_code == 200
|
||||
_, kwargs = mock_memory.update.call_args
|
||||
assert kwargs["expiration_date"] is None
|
||||
|
||||
def test_dict_not_passed_as_data(self, client, mock_memory):
|
||||
"""Regression test for #3933: the entire dict must NOT be passed as data."""
|
||||
|
||||
Reference in New Issue
Block a user