diff --git a/mem0-ts/src/oss/src/memory/index.ts b/mem0-ts/src/oss/src/memory/index.ts index 363e31f01..b571e3912 100644 --- a/mem0-ts/src/oss/src/memory/index.ts +++ b/mem0-ts/src/oss/src/memory/index.ts @@ -1601,7 +1601,9 @@ export class Memory { has_agent_id: !!config.agentId, has_run_id: !!config.runId, }); - const { userId, agentId, runId } = config; + const userId = validateAndTrimEntityId(config.userId, "userId"); + const agentId = validateAndTrimEntityId(config.agentId, "agentId"); + const runId = validateAndTrimEntityId(config.runId, "runId"); // Convert camelCase entity params to snake_case for filters (matches storage and search/getAll) const filters: SearchFilters = {}; diff --git a/mem0-ts/src/oss/tests/memory.validation.test.ts b/mem0-ts/src/oss/tests/memory.validation.test.ts index 73808ba95..5503a8902 100644 --- a/mem0-ts/src/oss/tests/memory.validation.test.ts +++ b/mem0-ts/src/oss/tests/memory.validation.test.ts @@ -314,4 +314,27 @@ describe("Memory Input Validation", () => { expect(result.results).toBeDefined(); }); }); + + describe("deleteAll() entity ID validation", () => { + it("should throw error when userId is whitespace-only", async () => { + await expect(memory.deleteAll({ userId: " " })).rejects.toThrow( + "Invalid userId", + ); + }); + + it("should throw error when userId contains internal whitespace", async () => { + await expect(memory.deleteAll({ userId: "user 123" })).rejects.toThrow( + "Invalid userId: cannot contain whitespace", + ); + }); + + it("should trim userId before listing memories", async () => { + const listSpy = jest.spyOn(memory["vectorStore"], "list"); + listSpy.mockResolvedValue([[], null]); + + await memory.deleteAll({ userId: " alice " }); + + expect(listSpy).toHaveBeenCalledWith({ user_id: "alice" }); + }); + }); }); diff --git a/mem0/memory/main.py b/mem0/memory/main.py index 8aa2b59b7..947e63530 100644 --- a/mem0/memory/main.py +++ b/mem0/memory/main.py @@ -1762,6 +1762,10 @@ class Memory(MemoryBase): agent_id (str, optional): ID of the agent to delete memories for. Defaults to None. run_id (str, optional): ID of the run to delete memories for. Defaults to None. """ + user_id = _validate_and_trim_entity_id(user_id, "user_id") + agent_id = _validate_and_trim_entity_id(agent_id, "agent_id") + run_id = _validate_and_trim_entity_id(run_id, "run_id") + filters: Dict[str, Any] = {} if user_id: filters["user_id"] = user_id @@ -3331,6 +3335,10 @@ class AsyncMemory(MemoryBase): agent_id (str, optional): ID of the agent to delete memories for. Defaults to None. run_id (str, optional): ID of the run to delete memories for. Defaults to None. """ + user_id = _validate_and_trim_entity_id(user_id, "user_id") + agent_id = _validate_and_trim_entity_id(agent_id, "agent_id") + run_id = _validate_and_trim_entity_id(run_id, "run_id") + filters = {} if user_id: filters["user_id"] = user_id diff --git a/tests/test_main.py b/tests/test_main.py index e4396ebee..8c8fc516a 100644 --- a/tests/test_main.py +++ b/tests/test_main.py @@ -286,6 +286,24 @@ class TestEntityIdValidation: with pytest.raises(ValueError, match="Invalid user_id.*cannot contain whitespace"): memory_instance.add("test message", user_id="user 123") + def test_delete_all_rejects_whitespace_only_user_id(self, memory_instance): + """delete_all should reject whitespace-only user_id.""" + with pytest.raises(ValueError, match="Invalid user_id.*cannot be empty"): + memory_instance.delete_all(user_id=" ") + + def test_delete_all_rejects_internal_whitespace_user_id(self, memory_instance): + """delete_all should reject user_id with internal whitespace.""" + with pytest.raises(ValueError, match="Invalid user_id.*cannot contain whitespace"): + memory_instance.delete_all(user_id="user 123") + + def test_delete_all_trims_user_id_before_list(self, memory_instance): + """delete_all should trim leading/trailing whitespace on entity IDs.""" + memory_instance.vector_store.list = Mock(return_value=([], None)) + + memory_instance.delete_all(user_id=" alice ") + + memory_instance.vector_store.list.assert_called_once_with(filters={"user_id": "alice"}) + class TestSearchParamValidation: """Tests for search parameter validation (threshold and top_k)."""