diff --git a/mem0-ts/src/oss/src/memory/index.ts b/mem0-ts/src/oss/src/memory/index.ts index 8a6c2cf24..a84543e58 100644 --- a/mem0-ts/src/oss/src/memory/index.ts +++ b/mem0-ts/src/oss/src/memory/index.ts @@ -100,11 +100,20 @@ const ENTITY_PARAMS = [ "agentId", "runId", ]; - -// Identity keys stripped from update() metadata: ENTITY_PARAMS covers user_id/agent_id/run_id -// in both casings (the default store promotes camelCase on read); actor_id has no camelCase alias. +// Identity keys stripped from caller metadata in add() and update(): ENTITY_PARAMS covers +// user_id/agent_id/run_id in both casings (the default store promotes camelCase on read); +// actor_id has no camelCase alias. const IDENTITY_KEYS = [...ENTITY_PARAMS, "actor_id"]; +// Caller metadata must not overwrite or inject an identity scope (#6342 / #6367 / #6371). +function stripIdentityKeys( + metadata: Record = {}, +): Record { + return Object.fromEntries( + Object.entries(metadata).filter(([key]) => !IDENTITY_KEYS.includes(key)), + ); +} + // Batch size for deleteAll pagination. Larger than most vector store default // page limits (~100) to minimize roundtrips while bounded to avoid memory pressure. const DELETE_ALL_BATCH_SIZE = 1000; @@ -740,7 +749,8 @@ export class Memory { has_filters: !!config.filters, infer: config.infer, }); - const { metadata = {}, filters = {}, infer = true } = config; + const { filters = {}, infer = true } = config; + const metadata = stripIdentityKeys(config.metadata); // Validate and trim entity IDs const userId = validateAndTrimEntityId(config.userId, "userId"); @@ -751,6 +761,9 @@ export class Memory { if (userId) filters.user_id = metadata.user_id = userId; if (agentId) filters.agent_id = metadata.agent_id = agentId; if (runId) filters.run_id = metadata.run_id = runId; + if (filters.user_id) metadata.user_id = filters.user_id; + if (filters.agent_id) metadata.agent_id = filters.agent_id; + if (filters.run_id) metadata.run_id = filters.run_id; // Normalize expiration date into the stored metadata (round-trips via get()). if (config.expirationDate != null) { @@ -1971,10 +1984,7 @@ export class Memory { existingEmbeddings[newData] || (await this.embedder.embed(newData, "update")); - // Caller metadata must not overwrite or inject an identity scope (#6342 / #6367). - const sanitizedMetadata = Object.fromEntries( - Object.entries(metadata).filter(([k]) => !IDENTITY_KEYS.includes(k)), - ); + const sanitizedMetadata = stripIdentityKeys(metadata); const newMetadata = { ...existingMemory.payload, diff --git a/mem0-ts/src/oss/tests/memory.add.test.ts b/mem0-ts/src/oss/tests/memory.add.test.ts index 4e7ebe773..059a0ef12 100644 --- a/mem0-ts/src/oss/tests/memory.add.test.ts +++ b/mem0-ts/src/oss/tests/memory.add.test.ts @@ -8,55 +8,78 @@ import type { MemoryConfig, MemoryItem, SearchResult } from "../src/types"; jest.setTimeout(15000); -// Mock Google modules to prevent @google/genai crash in CI -jest.mock("../src/embeddings/google", () => ({ - GoogleEmbedder: jest.fn(), -})); -jest.mock("../src/llms/google", () => ({ - GoogleLLM: jest.fn(), -})); +jest.mock("../src/utils/factory", () => { + const { MemoryVectorStore } = jest.requireActual( + "../src/vector_stores/memory", + ); + const { MemoryHistoryManager } = jest.requireActual( + "../src/storage/MemoryHistoryManager", + ); + const testEmbedding = new Array(1536).fill(0.1); -jest.mock("../src/llms/openai", () => ({ - OpenAILLM: jest.fn().mockImplementation(() => ({ - generateResponse: jest - .fn() - .mockImplementation( - (messages: Array<{ role: string; content: string }>) => { - // V3 pipeline: single LLM call with additive extraction prompt. - const userMsg = messages.find((m) => m.role === "user"); - const content = userMsg?.content ?? ""; - const newMsgMatch = content.match( - /## New Messages\n([\s\S]*?)(?=\n##|$)/, + class MockEmbedder { + embeddingDims = 1536; + + async embed(): Promise { + return testEmbedding; + } + + async embedBatch(texts: string[]): Promise { + return texts.map(() => testEmbedding); + } + } + + class MockLLM { + async generateResponse(messages: Array<{ role: string; content: string }>) { + const userMsg = messages.find((m) => m.role === "user"); + const content = userMsg?.content ?? ""; + const newMsgMatch = content.match( + /## New Messages\n([\s\S]*?)(?=\n##|$)/, + ); + const extracted = newMsgMatch + ? newMsgMatch[1].trim() + : "extracted fact from input"; + + return JSON.stringify({ + memory: [ + { + id: "0", + text: extracted, + attributed_to: "user", + }, + ], + }); + } + } + + return { + __esModule: true, + EmbedderFactory: { + create: jest.fn(() => new MockEmbedder()), + }, + LLMFactory: { + create: jest.fn(() => new MockLLM()), + }, + VectorStoreFactory: { + create: jest.fn((provider: string, config: any) => { + if (provider.toLowerCase() !== "memory") { + throw new Error( + `Unsupported vector store provider in test: ${provider}`, ); - const extracted = newMsgMatch - ? newMsgMatch[1].trim() - : "extracted fact from input"; - return JSON.stringify({ - memory: [ - { - id: "0", - text: extracted, - attributed_to: "user", - }, - ], - }); - }, - ), - })), -})); - -const mockEmbedding = new Array(1536).fill(0.1); -jest.mock("../src/embeddings/openai", () => ({ - OpenAIEmbedder: jest.fn().mockImplementation(() => ({ - embed: jest.fn().mockResolvedValue(mockEmbedding), - embedBatch: jest - .fn() - .mockImplementation((texts: string[]) => - Promise.resolve(texts.map(() => mockEmbedding)), - ), - embeddingDims: 1536, - })), -})); + } + return new MemoryVectorStore(config); + }), + }, + HistoryManagerFactory: { + create: jest.fn(() => new MemoryHistoryManager()), + }, + RerankerFactory: { + create: jest.fn(() => { + throw new Error("RerankerFactory is not used in memory.add.test.ts"); + }), + }, + }; +}); function createMemory(overrides: Partial = {}): Memory { return new Memory({ @@ -172,6 +195,138 @@ describe("Memory - add()", () => { ); }); + test("does not allow metadata to set identity scope", async () => { + const result: SearchResult = await memory.add("I am a software engineer", { + userId: "u1", + metadata: { + agent_id: "other", + agentId: "other-camel", + run_id: "other-run", + runId: "other-run-camel", + actor_id: "x", + source: "issue-6371", + nested: { preserved: true }, + }, + }); + const stored: MemoryItem | null = await memory.get(result.results[0].id); + + expect(stored).toEqual( + expect.objectContaining({ + user_id: "u1", + metadata: expect.objectContaining({ + source: "issue-6371", + nested: { preserved: true }, + }), + }), + ); + expect(stored).not.toHaveProperty("agent_id"); + expect(stored).not.toHaveProperty("run_id"); + expect(stored!.metadata).not.toHaveProperty("agent_id"); + expect(stored!.metadata).not.toHaveProperty("agentId"); + expect(stored!.metadata).not.toHaveProperty("run_id"); + expect(stored!.metadata).not.toHaveProperty("runId"); + expect(stored!.metadata).not.toHaveProperty("actor_id"); + }); + + test.each([ + ["userId", { userId: "u1" }, true, "user_id", "u1"], + ["agentId", { agentId: "a1" }, false, "agent_id", "a1"], + ["runId", { runId: "r1" }, true, "run_id", "r1"], + ] as const)( + "preserves typed %s scope while stripping conflicting metadata identities", + async (_mode, scope, infer, canonicalKey, canonicalValue) => { + const result: SearchResult = await memory.add("scoped content", { + ...scope, + infer, + metadata: { + user_id: "metadata-user", + userId: "metadata-user-camel", + agent_id: "metadata-agent", + agentId: "metadata-agent-camel", + run_id: "metadata-run", + runId: "metadata-run-camel", + actor_id: "metadata-actor", + ordinary: "preserved", + }, + }); + const stored: MemoryItem | null = await memory.get(result.results[0].id); + + expect(stored).toHaveProperty(canonicalKey, canonicalValue); + for (const key of ["user_id", "agent_id", "run_id"]) { + if (key !== canonicalKey) { + expect(stored).not.toHaveProperty(key); + } + } + expect(stored!.metadata).toEqual( + expect.objectContaining({ ordinary: "preserved" }), + ); + for (const key of [ + "user_id", + "userId", + "agent_id", + "agentId", + "run_id", + "runId", + "actor_id", + ]) { + if (key !== canonicalKey) { + expect(stored!.metadata).not.toHaveProperty(key); + } + } + }, + ); + + test.each([ + [true, "user_id", "filter-user"], + [false, "user_id", "filter-user"], + [false, "agent_id", "filter-agent"], + [false, "run_id", "filter-run"], + ] as const)( + "preserves infer=%s %s filters scope after sanitization", + async (infer, filterKey, filterValue) => { + const result: SearchResult = await memory.add("filter-scoped content", { + filters: { [filterKey]: filterValue }, + infer, + metadata: { + user_id: "metadata-user", + userId: "metadata-user-camel", + agent_id: "metadata-agent", + agentId: "metadata-agent-camel", + run_id: "metadata-run", + runId: "metadata-run-camel", + actor_id: "metadata-actor", + ordinary: "preserved", + }, + }); + const stored: MemoryItem | null = await memory.get(result.results[0].id); + + expect(stored).toHaveProperty(filterKey, filterValue); + for (const key of ["user_id", "agent_id", "run_id"]) { + if (key !== filterKey) { + expect(stored).not.toHaveProperty(key); + } + } + expect(stored!.metadata).toEqual( + expect.objectContaining({ + [filterKey]: filterValue, + ordinary: "preserved", + }), + ); + for (const key of [ + "user_id", + "userId", + "agent_id", + "agentId", + "run_id", + "runId", + "actor_id", + ]) { + if (key === filterKey) continue; + expect(stored!.metadata).not.toHaveProperty(key); + } + }, + ); + test("with infer=false skips LLM and stores messages directly", async () => { const result: SearchResult = await memory.add("Direct storage content", { userId,