diff --git a/mem0-ts/src/oss/src/memory/index.ts b/mem0-ts/src/oss/src/memory/index.ts index 5cef1c256..8ad92df5c 100644 --- a/mem0-ts/src/oss/src/memory/index.ts +++ b/mem0-ts/src/oss/src/memory/index.ts @@ -137,6 +137,27 @@ function rejectTopLevelEntityParams( } } +const ENTITY_ID_ALIASES: Record = { + userId: "user_id", + agentId: "agent_id", + runId: "run_id", +}; + +/** Rewrites camelCase entity ids inside filters to the snake_case keys vector stores index on. */ +function normalizeEntityFilterKeys( + filters: Record, +): Record { + const normalized = { ...filters }; + for (const [camelKey, snakeKey] of Object.entries(ENTITY_ID_ALIASES)) { + if (!(camelKey in normalized)) continue; + if (normalized[snakeKey] === undefined) { + normalized[snakeKey] = normalized[camelKey]; + } + delete normalized[camelKey]; + } + return normalized; +} + /** * Validates and normalizes an entity ID. * - Coerces non-string ids (e.g. numeric database keys) to string @@ -1356,16 +1377,20 @@ export class Memory { // receive `agent_id: undefined` / `run_id: undefined` and fail // (Qdrant rejects the malformed match, pgvector binds NULL, Redis // emits a literal "undefined" string in TAG filters). + const requestedFilters = normalizeEntityFilterKeys(config.filters ?? {}); const normalizedFilters: Record = config.filters ? Object.fromEntries( Object.entries({ - ...config.filters, - user_id: validateAndTrimEntityId(config.filters.user_id, "user_id"), + ...requestedFilters, + user_id: validateAndTrimEntityId( + requestedFilters.user_id, + "user_id", + ), agent_id: validateAndTrimEntityId( - config.filters.agent_id, + requestedFilters.agent_id, "agent_id", ), - run_id: validateAndTrimEntityId(config.filters.run_id, "run_id"), + run_id: validateAndTrimEntityId(requestedFilters.run_id, "run_id"), }).filter(([, v]) => v !== undefined), ) : {}; @@ -1864,12 +1889,16 @@ export class Memory { // Validate and trim entity IDs in filters. Drop keys that resolve to // undefined so downstream vector stores don't receive // `agent_id: undefined` / `run_id: undefined` and fail. + const requestedFilters = normalizeEntityFilterKeys(config.filters || {}); const filters: Record = Object.fromEntries( Object.entries({ - ...(config.filters || {}), - user_id: validateAndTrimEntityId(config.filters?.user_id, "user_id"), - agent_id: validateAndTrimEntityId(config.filters?.agent_id, "agent_id"), - run_id: validateAndTrimEntityId(config.filters?.run_id, "run_id"), + ...requestedFilters, + user_id: validateAndTrimEntityId(requestedFilters.user_id, "user_id"), + agent_id: validateAndTrimEntityId( + requestedFilters.agent_id, + "agent_id", + ), + run_id: validateAndTrimEntityId(requestedFilters.run_id, "run_id"), }).filter(([, v]) => v !== undefined), ); diff --git a/mem0-ts/src/oss/tests/memory.validation.test.ts b/mem0-ts/src/oss/tests/memory.validation.test.ts index 02aad751f..a80147c2a 100644 --- a/mem0-ts/src/oss/tests/memory.validation.test.ts +++ b/mem0-ts/src/oss/tests/memory.validation.test.ts @@ -281,6 +281,66 @@ describe("Memory Input Validation", () => { }); }); + describe("camelCase entity IDs in filters", () => { + it("maps camelCase filter keys to snake_case in search", async () => { + const searchSpy = jest + .spyOn((memory as any).vectorStore, "search") + .mockResolvedValue([]); + + await memory.search("q", { + filters: { userId: "alice", agentId: "bot", runId: "run-1" }, + }); + + const passedFilters = searchSpy.mock.calls[0][2] as Record; + expect(passedFilters).toMatchObject({ + user_id: "alice", + agent_id: "bot", + run_id: "run-1", + }); + expect(passedFilters.userId).toBeUndefined(); + expect(passedFilters.agentId).toBeUndefined(); + expect(passedFilters.runId).toBeUndefined(); + + searchSpy.mockRestore(); + }); + + it("maps camelCase filter keys to snake_case in getAll", async () => { + const listSpy = jest + .spyOn((memory as any).vectorStore, "list") + .mockResolvedValue([[], 0]); + + await memory.getAll({ filters: { userId: "alice" } }); + + const passedFilters = listSpy.mock.calls[0][0] as Record; + expect(passedFilters.user_id).toBe("alice"); + expect(passedFilters.userId).toBeUndefined(); + + listSpy.mockRestore(); + }); + + it("prefers an explicit snake_case value over its camelCase alias", async () => { + const listSpy = jest + .spyOn((memory as any).vectorStore, "list") + .mockResolvedValue([[], 0]); + + await memory.getAll({ + filters: { user_id: "snake", userId: "camel" }, + }); + + const passedFilters = listSpy.mock.calls[0][0] as Record; + expect(passedFilters.user_id).toBe("snake"); + expect(passedFilters.userId).toBeUndefined(); + + listSpy.mockRestore(); + }); + + it("validates camelCase filter values like their snake_case counterparts", async () => { + await expect( + memory.search("q", { filters: { userId: "user 123" } }), + ).rejects.toThrow("Invalid user_id: cannot contain whitespace"); + }); + }); + describe("search() filter entity ID validation", () => { it("should throw error when user_id in filters is whitespace-only", async () => { await expect(