From c2bc28e589ed01c822e5d72a61d175cbb05e64fc Mon Sep 17 00:00:00 2001 From: Aditya Jethani Date: Tue, 21 Jul 2026 12:44:22 +0530 Subject: [PATCH] fix(memory): don't let update() metadata overwrite user_id/agent_id/run_id (#6278) Co-authored-by: kartik-mem0 --- mem0/memory/main.py | 30 ++++++++---- tests/memory/test_main.py | 98 +++++++++++++++++++++------------------ 2 files changed, 72 insertions(+), 56 deletions(-) diff --git a/mem0/memory/main.py b/mem0/memory/main.py index 3d74cc87a..4aac8f8e5 100644 --- a/mem0/memory/main.py +++ b/mem0/memory/main.py @@ -134,6 +134,20 @@ _SENSITIVE_SUFFIXES = ( # Entity parameters that must be passed via filters, not top-level kwargs ENTITY_PARAMS = frozenset({"user_id", "agent_id", "run_id"}) +# Tenant-scoping fields that update() must never let caller-supplied metadata overwrite (issues #4490, #6277). +_IDENTITY_KEYS = ENTITY_PARAMS | {"actor_id"} + + +def _strip_identity_keys(metadata: Dict[str, Any], existing_payload: Dict[str, Any]) -> Dict[str, Any]: + """Drop identity keys from caller metadata; they are immutable after creation (issues #4490, #6277).""" + clean = {} + for key, value in metadata.items(): + if key not in _IDENTITY_KEYS: + clean[key] = value + elif value != existing_payload.get(key): + logger.warning(f"update(): ignoring metadata['{key}'] - identity fields are immutable after creation") + return clean + def _reject_top_level_entity_params(kwargs: Dict[str, Any], method_name: str) -> None: """Reject top-level entity parameters - must use filters instead.""" @@ -1783,6 +1797,8 @@ class Memory(MemoryBase): memory_id (str): ID of the memory to update. text (str, optional): New content to update the memory with. metadata (dict, optional): Metadata to update with the memory. Defaults to None. + ``user_id``/``agent_id``/``run_id``/``actor_id`` are ignored here - they are + immutable after creation. expiration_date (Any, optional): Date in YYYY-MM-DD format, or None to clear it. data (str, optional): Deprecated alias for ``text``. Will be removed in the next major release; use ``text`` instead. @@ -1991,7 +2007,7 @@ class Memory(MemoryBase): new_metadata = deepcopy(existing_memory.payload) if metadata is not None: - new_metadata.update(metadata) + new_metadata.update(_strip_identity_keys(metadata, existing_memory.payload)) new_metadata["data"] = data new_metadata["hash"] = hashlib.md5(data.encode()).hexdigest() @@ -1999,10 +2015,6 @@ class Memory(MemoryBase): new_metadata["created_at"] = existing_memory.payload.get("created_at") new_metadata["updated_at"] = datetime.now(timezone.utc).isoformat() - # actor_id is immutable after creation (issue #4490) - if "actor_id" in existing_memory.payload: - new_metadata["actor_id"] = existing_memory.payload["actor_id"] - if data in existing_embeddings: embeddings = existing_embeddings[data] else: @@ -3414,6 +3426,8 @@ class AsyncMemory(MemoryBase): memory_id (str): ID of the memory to update. text (str, optional): New content to update the memory with. metadata (dict, optional): Metadata to update with the memory. Defaults to None. + ``user_id``/``agent_id``/``run_id``/``actor_id`` are ignored here - they are + immutable after creation. expiration_date (Any, optional): Date in YYYY-MM-DD format, or None to clear it. data (str, optional): Deprecated alias for ``text``. Will be removed in the next major release; use ``text`` instead. @@ -3657,7 +3671,7 @@ class AsyncMemory(MemoryBase): new_metadata = deepcopy(existing_memory.payload) if metadata is not None: - new_metadata.update(metadata) + new_metadata.update(_strip_identity_keys(metadata, existing_memory.payload)) new_metadata["data"] = data new_metadata["hash"] = hashlib.md5(data.encode()).hexdigest() @@ -3665,10 +3679,6 @@ class AsyncMemory(MemoryBase): new_metadata["created_at"] = existing_memory.payload.get("created_at") new_metadata["updated_at"] = datetime.now(timezone.utc).isoformat() - # actor_id is immutable after creation (issue #4490) - if "actor_id" in existing_memory.payload: - new_metadata["actor_id"] = existing_memory.payload["actor_id"] - if data in existing_embeddings: embeddings = existing_embeddings[data] else: diff --git a/tests/memory/test_main.py b/tests/memory/test_main.py index fd4d68852..d00e953b0 100644 --- a/tests/memory/test_main.py +++ b/tests/memory/test_main.py @@ -418,6 +418,58 @@ async def test_async_update_memory_uses_utc_timestamps(mocker): assert payload["updated_at"] is not None +_ATTACKER_UPDATE_METADATA = { + "user_id": "attacker_tenant", + "agent_id": "attacker_agent", + "run_id": "attacker_run", + "actor_id": "attacker_actor", + "category": "sports", +} + +# Omits agent_id on purpose, so one payload covers both overwriting and injecting an identity field. +_EXISTING_UPDATE_PAYLOAD = { + "data": "old memory", + "user_id": "tenant_a", + "run_id": "run_a", + "actor_id": "actor_a", +} + + +def test_update_memory_metadata_cannot_change_identity_fields(mocker, caplog): + """Regression (issues #4490, #6277): update() metadata must not overwrite or inject identity fields.""" + memory = _build_memory_instance(mocker, Memory) + memory.vector_store.get.return_value = MagicMock(payload=dict(_EXISTING_UPDATE_PAYLOAD)) + + with caplog.at_level(logging.WARNING, logger="mem0.memory.main"): + memory._update_memory("memory-id", "new memory", {}, metadata=dict(_ATTACKER_UPDATE_METADATA)) + + payload = memory.vector_store.update.call_args.kwargs["payload"] + assert payload["user_id"] == "tenant_a" + assert payload["run_id"] == "run_a" + assert "agent_id" not in payload + assert payload["actor_id"] == "actor_a" + assert payload["category"] == "sports" + assert payload["data"] == "new memory" + assert "ignoring metadata['user_id']" in caplog.text + + +@pytest.mark.asyncio +async def test_async_update_memory_metadata_cannot_change_identity_fields(mocker): + """Async counterpart of test_update_memory_metadata_cannot_change_identity_fields.""" + memory = _build_memory_instance(mocker, AsyncMemory) + memory.vector_store.get.return_value = MagicMock(payload=dict(_EXISTING_UPDATE_PAYLOAD)) + + await memory._update_memory("memory-id", "new memory", {}, metadata=dict(_ATTACKER_UPDATE_METADATA)) + + payload = memory.vector_store.update.call_args.kwargs["payload"] + assert payload["user_id"] == "tenant_a" + assert payload["run_id"] == "run_a" + assert "agent_id" not in payload + assert payload["actor_id"] == "actor_a" + assert payload["category"] == "sports" + assert payload["data"] == "new memory" + + def test_create_then_search_and_get_all_return_same_timestamps(mocker): """Reproduces issue #3720: created_at must be identical in search() and get_all().""" memory = _build_memory_instance(mocker, Memory) @@ -714,52 +766,6 @@ class TestMetadataNotMutated: ) -def test_update_preserves_actor_id_when_different_actor_updates(mocker): - """actor_id must be preserved from the original memory even when the - updating caller passes a different actor_id in metadata (issue #4490).""" - memory = _build_memory_instance(mocker, Memory) - memory.vector_store.get.return_value = MagicMock( - payload={ - "data": "I am player #1", - "user_id": "team", - "actor_id": "Alice", - "created_at": "2026-01-01T00:00:00+00:00", - } - ) - - memory._update_memory( - "mem-id", "Player #1 is a good person", - {"Player #1 is a good person": [0.1, 0.2, 0.3]}, - metadata={"user_id": "team", "actor_id": "Bob"}, - ) - - stored = memory.vector_store.update.call_args.kwargs["payload"] - assert stored["actor_id"] == "Alice" - - -@pytest.mark.asyncio -async def test_async_update_preserves_actor_id_when_different_actor_updates(mocker): - """Async variant: actor_id must be preserved from the original memory (issue #4490).""" - memory = _build_memory_instance(mocker, AsyncMemory) - memory.vector_store.get.return_value = MagicMock( - payload={ - "data": "I am player #1", - "user_id": "team", - "actor_id": "Alice", - "created_at": "2026-01-01T00:00:00+00:00", - } - ) - - await memory._update_memory( - "mem-id", "Player #1 is a good person", - {"Player #1 is a good person": [0.1, 0.2, 0.3]}, - metadata={"user_id": "team", "actor_id": "Bob"}, - ) - - stored = memory.vector_store.update.call_args.kwargs["payload"] - assert stored["actor_id"] == "Alice" - - def _make_match(score, linked_memory_ids): return SimpleNamespace(score=score, payload={"linked_memory_ids": linked_memory_ids})