From e44b46ef2ea559953190d0bfb5a58bf10ad799bb Mon Sep 17 00:00:00 2001 From: Kartik Date: Sun, 12 Apr 2026 00:34:58 +0530 Subject: [PATCH] fix(sdk): removing deprecating param from our sdk and docs changes with it (#4740) --- LLM.md | 4 +- docs/api-reference/memory/add-memories.mdx | 3 +- docs/api-reference/memory/get-memories.mdx | 3 +- docs/api-reference/organizations-projects.mdx | 9 +- docs/components/vectordbs/dbs/langchain.mdx | 2 +- .../companions/voice-companion-openai.mdx | 12 +- docs/core-concepts/memory-operations/add.mdx | 1 - docs/docs.json | 53 +- docs/integrations/google-ai-adk.mdx | 2 +- docs/integrations/keywords.mdx | 3 +- docs/integrations/langchain-tools.mdx | 5 +- docs/integrations/openai-agents-sdk.mdx | 4 +- docs/integrations/openclaw.mdx | 2 +- docs/integrations/vercel-ai-sdk.mdx | 2 +- docs/llms.txt | 2 +- docs/migration/api-changes.mdx | 4 +- docs/migration/breaking-changes.mdx | 383 -------------- docs/migration/oss-to-platform.mdx | 16 +- docs/migration/v0-to-v1.mdx | 481 ------------------ docs/open-source/features/async-memory.mdx | 4 +- ...ion-prompt.mdx => custom-instructions.mdx} | 30 +- .../features/custom-update-memory-prompt.mdx | 6 +- docs/open-source/features/graph-memory.mdx | 8 +- .../features/openai_compatibility.mdx | 2 +- docs/open-source/features/overview.mdx | 2 +- docs/open-source/features/reranker-search.mdx | 4 +- docs/open-source/node-quickstart.mdx | 4 +- docs/openapi.json | 52 +- docs/platform/features/advanced-retrieval.mdx | 118 +---- .../features/async-mode-default-change.mdx | 197 ------- docs/platform/features/criteria-retrieval.mdx | 6 +- docs/platform/features/custom-categories.mdx | 4 +- docs/platform/features/expiration-date.mdx | 109 ---- docs/platform/features/graph-memory.mdx | 24 +- docs/platform/features/group-chat.mdx | 7 +- mem0-ts/src/client/index.ts | 9 +- mem0-ts/src/client/mem0.ts | 262 +++------- mem0-ts/src/client/mem0.types.ts | 139 +++-- .../src/client/tests/integration/crud.test.ts | 8 +- .../src/client/tests/integration/helpers.ts | 4 +- .../tests/integration/initialization.test.ts | 6 +- .../client/tests/integration/search.test.ts | 21 +- .../client/tests/memoryClient.crud.test.ts | 101 +--- .../client/tests/memoryClient.init.test.ts | 62 +-- .../client/tests/memoryClient.project.test.ts | 79 +-- .../client/tests/memoryClient.search.test.ts | 78 ++- .../client/tests/memoryClient.users.test.ts | 115 +---- .../tests/memoryClient.webhooks.test.ts | 28 +- .../src/integrations/langchain/mem0.ts | 15 +- mem0-ts/src/oss/src/config/manager.ts | 2 +- mem0-ts/src/oss/src/graphs/configs.ts | 2 +- mem0-ts/src/oss/src/memory/graph_memory.ts | 14 +- mem0-ts/src/oss/src/memory/index.ts | 24 +- mem0-ts/src/oss/src/memory/memory.types.ts | 4 +- .../src/tests/sqlite-backward-compat.test.ts | 6 +- mem0-ts/src/oss/src/types/index.ts | 8 +- .../oss/src/vector_stores/azure_ai_search.ts | 12 +- mem0-ts/src/oss/src/vector_stores/base.ts | 4 +- .../src/oss/src/vector_stores/langchain.ts | 6 +- mem0-ts/src/oss/src/vector_stores/memory.ts | 8 +- mem0-ts/src/oss/src/vector_stores/pgvector.ts | 8 +- mem0-ts/src/oss/src/vector_stores/qdrant.ts | 8 +- mem0-ts/src/oss/src/vector_stores/redis.ts | 10 +- mem0-ts/src/oss/src/vector_stores/supabase.ts | 8 +- .../src/oss/src/vector_stores/vectorize.ts | 8 +- .../oss/tests/graph-memory-parsing.test.ts | 6 +- mem0/client/main.py | 431 ++++++++-------- mem0/client/types.py | 103 ++++ mem0/configs/base.py | 6 +- mem0/configs/vector_stores/elasticsearch.py | 2 +- mem0/graphs/neptune/base.py | 18 +- mem0/graphs/neptune/neptunedb.py | 18 +- mem0/graphs/neptune/neptunegraph.py | 14 +- mem0/memory/apache_age_memory.py | 14 +- mem0/memory/graph_memory.py | 14 +- mem0/memory/kuzu_memory.py | 18 +- mem0/memory/main.py | 78 +-- mem0/memory/memgraph_memory.py | 18 +- mem0/memory/utils.py | 2 +- mem0/proxy/main.py | 8 +- mem0/vector_stores/azure_ai_search.py | 16 +- mem0/vector_stores/azure_mysql.py | 16 +- mem0/vector_stores/baidu.py | 12 +- mem0/vector_stores/base.py | 4 +- mem0/vector_stores/cassandra.py | 16 +- mem0/vector_stores/chroma.py | 12 +- mem0/vector_stores/databricks.py | 30 +- mem0/vector_stores/elasticsearch.py | 12 +- mem0/vector_stores/faiss.py | 31 +- mem0/vector_stores/langchain.py | 10 +- mem0/vector_stores/milvus.py | 12 +- mem0/vector_stores/mongodb.py | 14 +- mem0/vector_stores/neptune_analytics.py | 12 +- mem0/vector_stores/opensearch.py | 14 +- mem0/vector_stores/pgvector.py | 12 +- mem0/vector_stores/pinecone.py | 12 +- mem0/vector_stores/qdrant.py | 12 +- mem0/vector_stores/redis.py | 10 +- mem0/vector_stores/s3_vectors.py | 10 +- mem0/vector_stores/supabase.py | 12 +- mem0/vector_stores/turbopuffer.py | 12 +- mem0/vector_stores/upstash_vector.py | 14 +- mem0/vector_stores/valkey.py | 16 +- mem0/vector_stores/vertex_ai_vector_search.py | 26 +- mem0/vector_stores/weaviate.py | 8 +- server/main.py | 2 +- skills/mem0/client/differences.md | 2 +- skills/mem0/client/python.md | 2 +- tests/memory/test_apache_age_e2e.py | 4 +- tests/memory/test_apache_age_memory.py | 4 +- tests/memory/test_graph_memory_soft_delete.py | 4 +- tests/memory/test_main.py | 14 +- tests/memory/test_neptune_analytics_memory.py | 10 +- tests/memory/test_neptune_memory.py | 6 +- tests/test_main.py | 8 +- tests/vector_stores/test_azure_ai_search.py | 2 +- tests/vector_stores/test_azure_mysql.py | 7 +- tests/vector_stores/test_baidu.py | 6 +- tests/vector_stores/test_cassandra.py | 9 +- tests/vector_stores/test_chroma.py | 16 +- tests/vector_stores/test_databricks.py | 40 +- tests/vector_stores/test_elasticsearch.py | 8 +- tests/vector_stores/test_faiss.py | 6 +- .../test_langchain_vector_store.py | 18 +- tests/vector_stores/test_milvus.py | 10 +- tests/vector_stores/test_mongodb.py | 16 +- tests/vector_stores/test_neptune_analytics.py | 2 +- tests/vector_stores/test_opensearch.py | 4 +- tests/vector_stores/test_pgvector.py | 32 +- tests/vector_stores/test_pinecone.py | 2 +- tests/vector_stores/test_qdrant.py | 14 +- tests/vector_stores/test_s3_vectors.py | 4 +- tests/vector_stores/test_supabase.py | 4 +- tests/vector_stores/test_turbopuffer.py | 6 +- tests/vector_stores/test_upstash_vector.py | 8 +- tests/vector_stores/test_valkey.py | 14 +- .../test_vertex_ai_vector_search.py | 2 +- tests/vector_stores/test_weaviate.py | 8 +- 138 files changed, 1186 insertions(+), 2840 deletions(-) delete mode 100644 docs/migration/breaking-changes.mdx delete mode 100644 docs/migration/v0-to-v1.mdx rename docs/open-source/features/{custom-fact-extraction-prompt.mdx => custom-instructions.mdx} (85%) delete mode 100644 docs/platform/features/async-mode-default-change.mdx delete mode 100644 docs/platform/features/expiration-date.mdx create mode 100644 mem0/client/types.py diff --git a/LLM.md b/LLM.md index c97564070..693f69ae2 100644 --- a/LLM.md +++ b/LLM.md @@ -266,7 +266,7 @@ config = MemoryConfig( graph_store=GraphStoreConfig(provider="neo4j", config={...}), # optional history_db_path="~/.mem0/history.db", version="v1.1", - custom_fact_extraction_prompt="Custom prompt...", + custom_instructions="Custom prompt...", custom_update_memory_prompt="Custom prompt..." ) ``` @@ -684,7 +684,7 @@ Conversation: {messages} """ config = MemoryConfig( - custom_fact_extraction_prompt=custom_extraction_prompt + custom_instructions=custom_extraction_prompt ) memory = Memory(config) ``` diff --git a/docs/api-reference/memory/add-memories.mdx b/docs/api-reference/memory/add-memories.mdx index 144ce1217..a44c7d384 100644 --- a/docs/api-reference/memory/add-memories.mdx +++ b/docs/api-reference/memory/add-memories.mdx @@ -47,8 +47,7 @@ Provide at least one message or direct memory string. Most callers supply `messa | `messages` | array | No* | Conversation turns for Mem0 to infer memories from. Each object should include `role` and `content`. | | `metadata` | object | Optional | Custom key/value metadata (e.g., `{"topic": "preferences"}`). | | `infer` | boolean (default `true`) | Optional | Set to `false` to skip inference and store the provided text as-is. | -| `async_mode` | boolean (default `true`) | Optional | Controls asynchronous processing. Most clients leave this enabled. | -| `output_format` | string (default `v1.1`) | Optional | Response format. `v1.1` wraps results in a `results` array. | +| `enable_graph` | boolean (default `false`) | Optional | Set to `true` to create entity nodes and relationships for knowledge graph. | > \* Provide at least one `messages` entry to describe what you are storing. For scoped memories, include `user_id`. You can also attach `agent_id`, `app_id`, `run_id`, `project_id`, or `org_id` to refine ownership. diff --git a/docs/api-reference/memory/get-memories.mdx b/docs/api-reference/memory/get-memories.mdx index de1035213..be63b949e 100644 --- a/docs/api-reference/memory/get-memories.mdx +++ b/docs/api-reference/memory/get-memories.mdx @@ -62,8 +62,7 @@ To retrieve graph memory relationships between entities, pass `output_format="v1 memories = client.get_all( filters={ "user_id": "alex" - }, - output_format="v1.1" + } ) ``` diff --git a/docs/api-reference/organizations-projects.mdx b/docs/api-reference/organizations-projects.mdx index 69d6c9140..22bd9350c 100644 --- a/docs/api-reference/organizations-projects.mdx +++ b/docs/api-reference/organizations-projects.mdx @@ -32,7 +32,7 @@ Example with the mem0 Python package: ```python from mem0 import MemoryClient -client = MemoryClient(org_id='YOUR_ORG_ID', project_id='YOUR_PROJECT_ID') +client = MemoryClient(api_key="your-api-key") ``` @@ -41,10 +41,7 @@ client = MemoryClient(org_id='YOUR_ORG_ID', project_id='YOUR_PROJECT_ID') ```javascript import { MemoryClient } from "mem0ai"; -const client = new MemoryClient({ - organizationId: "YOUR_ORG_ID", - projectId: "YOUR_PROJECT_ID" -}); +const client = new MemoryClient({ apiKey: "your-api-key" }); ``` @@ -172,7 +169,7 @@ All project methods are available in async mode: from mem0 import AsyncMemoryClient async def manage_project(): - client = AsyncMemoryClient(org_id='YOUR_ORG_ID', project_id='YOUR_PROJECT_ID') + client = AsyncMemoryClient(api_key="your-api-key") # All methods support async/await project_info = await client.project.get() diff --git a/docs/components/vectordbs/dbs/langchain.mdx b/docs/components/vectordbs/dbs/langchain.mdx index a6f5a0e60..edaedab6c 100644 --- a/docs/components/vectordbs/dbs/langchain.mdx +++ b/docs/components/vectordbs/dbs/langchain.mdx @@ -15,7 +15,7 @@ Mem0 supports LangChain as a provider for vector store integration. LangChain pr ```python Python import os from mem0 import Memory -from langchain_community.vectorstores import Chroma +from langchain_chroma import Chroma from langchain_openai import OpenAIEmbeddings # Initialize a LangChain vector store diff --git a/docs/cookbooks/companions/voice-companion-openai.mdx b/docs/cookbooks/companions/voice-companion-openai.mdx index b74bd7930..d54dd570d 100644 --- a/docs/cookbooks/companions/voice-companion-openai.mdx +++ b/docs/cookbooks/companions/voice-companion-openai.mdx @@ -129,13 +129,13 @@ async def search_memories( user_id=USER_ID, limit=5, threshold=0.7, # Higher threshold for more relevant results - + ) - + # Format and return the results if not results.get('results', []): return "I don't have any relevant memories about this topic." - + memories = [f"• {result['memory']}" for result in results.get('results', [])] return "Here's what I remember that might be relevant:\n" + "\n".join(memories) ``` @@ -345,13 +345,13 @@ async def search_memories( user_id=USER_ID, limit=5, threshold=0.7, # Higher threshold for more relevant results - + ) - + # Format and return the results if not results.get('results', []): return "I don't have any relevant memories about this topic." - + memories = [f"• {result['memory']}" for result in results.get('results', [])] return "Here's what I remember that might be relevant:\n" + "\n".join(memories) diff --git a/docs/core-concepts/memory-operations/add.mdx b/docs/core-concepts/memory-operations/add.mdx index 6b2c9deeb..28aa3fe59 100644 --- a/docs/core-concepts/memory-operations/add.mdx +++ b/docs/core-concepts/memory-operations/add.mdx @@ -85,7 +85,6 @@ const messages = [ await client.add(messages, { user_id: "alice", - version: "v2", }); ``` diff --git a/docs/docs.json b/docs/docs.json index 118bb696d..18bb24613 100644 --- a/docs/docs.json +++ b/docs/docs.json @@ -70,7 +70,6 @@ "platform/features/v2-memory-filters", "platform/features/entity-scoped-memory", "platform/features/async-client", - "platform/features/async-mode-default-change", "platform/features/multimodal-support", "platform/features/custom-categories" ] @@ -94,8 +93,7 @@ "pages": [ "platform/features/direct-import", "platform/features/memory-export", - "platform/features/timestamp", - "platform/features/expiration-date" + "platform/features/timestamp" ] }, { @@ -122,8 +120,6 @@ "icon": "arrow-right", "pages": [ "migration/oss-to-platform", - "migration/v0-to-v1", - "migration/breaking-changes", "migration/api-changes" ] }, @@ -133,16 +129,6 @@ "pages": [ "platform/contribute" ] - }, - { - "group": "Release Notes", - "icon": "rocket", - "pages": [ - "changelog/highlights", - "changelog/sdk", - "changelog/platform", - "changelog/openclaw" - ] } ] }, @@ -182,7 +168,7 @@ "open-source/features/reranker-search", "open-source/features/async-memory", "open-source/features/multimodal-support", - "open-source/features/custom-fact-extraction-prompt", + "open-source/features/custom-instructions", "open-source/features/custom-update-memory-prompt", "open-source/features/rest-api", "open-source/features/openai_compatibility" @@ -562,6 +548,21 @@ ] } ] + }, + { + "tab": "Release Notes", + "groups": [ + { + "group": "Release Notes", + "icon": "rocket", + "pages": [ + "changelog/highlights", + "changelog/sdk", + "changelog/platform", + "changelog/openclaw" + ] + } + ] } ] } @@ -612,6 +613,26 @@ ] }, "redirects": [ + { + "source": "/migration/breaking-changes", + "destination": "/" + }, + { + "source": "/migration/v0-to-v1", + "destination": "/" + }, + { + "source": "/platform/features/expiration-date", + "destination": "/" + }, + { + "source": "/platform/features/async-mode-default-change", + "destination": "/" + }, + { + "source": "/open-source/features/custom-fact-extraction-prompt", + "destination": "/open-source/features/custom-instructions" + }, { "source": "/changelog", "destination": "/changelog/highlights" diff --git a/docs/integrations/google-ai-adk.mdx b/docs/integrations/google-ai-adk.mdx index a5274be00..43f6019a5 100644 --- a/docs/integrations/google-ai-adk.mdx +++ b/docs/integrations/google-ai-adk.mdx @@ -268,7 +268,7 @@ memories = mem0.search( {"categories": {"contains": "travel"}} ] }, - limit=5 + top_k=5 ) # Configure agent with custom model settings diff --git a/docs/integrations/keywords.mdx b/docs/integrations/keywords.mdx index 35f515d3d..c537cf805 100644 --- a/docs/integrations/keywords.mdx +++ b/docs/integrations/keywords.mdx @@ -16,7 +16,7 @@ Combining Mem0 with Keywords AI allows you to: 4. Optimize token usage and reduce costs -You can get your Mem0 API key, user_id, and org_id from the Mem0 dashboard. These are required for proper integration. +You can get your Mem0 API key from the Mem0 dashboard. ## Setup and Configuration @@ -107,7 +107,6 @@ response = client.chat.completions.create( extra_body={ "mem0_params": { "user_id": "test_user", - "org_id": "org_1", "api_key": os.environ.get("MEM0_API_KEY"), "add_memories": { "messages": messages, diff --git a/docs/integrations/langchain-tools.mdx b/docs/integrations/langchain-tools.mdx index 3710107b6..833e2bc0e 100644 --- a/docs/integrations/langchain-tools.mdx +++ b/docs/integrations/langchain-tools.mdx @@ -29,10 +29,7 @@ import os os.environ["MEM0_API_KEY"] = "your-api-key" -client = MemoryClient( - org_id=your_org_id, - project_id=your_project_id -) +client = MemoryClient() ``` ## Available Tools diff --git a/docs/integrations/openai-agents-sdk.mdx b/docs/integrations/openai-agents-sdk.mdx index b030a7ff7..e05c133fb 100644 --- a/docs/integrations/openai-agents-sdk.mdx +++ b/docs/integrations/openai-agents-sdk.mdx @@ -45,7 +45,7 @@ mem0 = MemoryClient() @function_tool def search_memory(query: str, user_id: str) -> str: """Search through past conversations and memories""" - memories = mem0.search(query, user_id=user_id, limit=3) + memories = mem0.search(query, user_id=user_id, top_k=3) if memories and memories.get('results'): return "\n".join([f"- {mem['memory']}" for mem in memories['results']]) return "No relevant memories found." @@ -215,7 +215,7 @@ Customize memory behavior: memories = mem0.search( query="travel preferences", user_id="alex", - limit=5 # Number of memories to retrieve + top_k=5 # Number of memories to retrieve ) # Add metadata to memories diff --git a/docs/integrations/openclaw.mdx b/docs/integrations/openclaw.mdx index 56b4fb33a..762041751 100644 --- a/docs/integrations/openclaw.mdx +++ b/docs/integrations/openclaw.mdx @@ -155,7 +155,7 @@ openclaw mem0 stats | Key | Type | Default | Description | |-----|------|---------|-------------| -| `customPrompt` | `string` | *(built-in)* | Extraction prompt for memory processing | +| `customInstructions` | `string` | *(built-in)* | Extraction prompt for memory processing | | `oss.embedder.provider` | `string` | `"openai"` | Embedding provider (`"openai"`, `"ollama"`, etc.) | | `oss.embedder.config` | `object` | — | Provider config: `apiKey`, `model`, `baseURL` | | `oss.vectorStore.provider` | `string` | `"memory"` | Vector store (`"memory"`, `"qdrant"`, `"chroma"`, etc.) | diff --git a/docs/integrations/vercel-ai-sdk.mdx b/docs/integrations/vercel-ai-sdk.mdx index e2d6da497..30df8aa95 100644 --- a/docs/integrations/vercel-ai-sdk.mdx +++ b/docs/integrations/vercel-ai-sdk.mdx @@ -52,7 +52,7 @@ npm install @mem0/vercel-ai-provider > **Note**: The `openai` provider is set as default. Consider using `MEM0_API_KEY` and `OPENAI_API_KEY` as environment variables for security. - > **Note**: The `mem0Config` is optional. It is used to set the global config for the Mem0 Client (eg. `user_id`, `agent_id`, `app_id`, `run_id`, `org_id`, `project_id` etc). + > **Note**: The `mem0Config` is optional. It is used to set the global config for the Mem0 Client (eg. `user_id`, `agent_id`, `app_id`, `run_id` etc). 3. Add Memories to Enhance Context: diff --git a/docs/llms.txt b/docs/llms.txt index dbf690120..0bde0ad34 100644 --- a/docs/llms.txt +++ b/docs/llms.txt @@ -86,7 +86,7 @@ Key differentiators: - [Reranker Search](https://docs.mem0.ai/open-source/features/reranker-search): Enhanced search results with reranking models - [Async Memory](https://docs.mem0.ai/open-source/features/async-memory): Asynchronous memory operations for better performance - [Multimodal Support](https://docs.mem0.ai/open-source/features/multimodal-support): Handle text, images, and documents in self-hosted setup -- [Custom Fact Extraction](https://docs.mem0.ai/open-source/features/custom-fact-extraction-prompt): Tailor information extraction for specific use cases +- [Custom Instructions](https://docs.mem0.ai/open-source/features/custom-instructions): Tailor information extraction for specific use cases - [Custom Memory Update Prompt](https://docs.mem0.ai/open-source/features/custom-update-memory-prompt): Customize how memories are updated and merged - [REST API Server](https://docs.mem0.ai/open-source/features/rest-api): FastAPI-based server with core operations and OpenAPI documentation - [OpenAI Compatibility](https://docs.mem0.ai/open-source/features/openai_compatibility): Seamless integration with OpenAI-compatible APIs diff --git a/docs/migration/api-changes.mdx b/docs/migration/api-changes.mdx index 15ebdbe3c..65e4b3a35 100644 --- a/docs/migration/api-changes.mdx +++ b/docs/migration/api-changes.mdx @@ -319,7 +319,7 @@ config = { "graph_store": {...}, "version": "v1.0", # ❌ v1.0 no longer supported "history_db_path": "...", - "custom_fact_extraction_prompt": "..." + "custom_instructions": "..." } ``` @@ -336,7 +336,7 @@ config = { }, "version": "v1.1", # ✅ v1.1+ only "history_db_path": "...", - "custom_fact_extraction_prompt": "...", + "custom_instructions": "...", "custom_update_memory_prompt": "..." # ✅ NEW: Custom update prompt } ``` diff --git a/docs/migration/breaking-changes.mdx b/docs/migration/breaking-changes.mdx deleted file mode 100644 index e84ec7fd5..000000000 --- a/docs/migration/breaking-changes.mdx +++ /dev/null @@ -1,383 +0,0 @@ ---- -title: Breaking Changes in v1.0.0 -description: 'Complete list of breaking changes when upgrading from v0.x to v1.0.0 ' -icon: "triangle-exclamation" -iconType: "solid" ---- - - -**Important:** This page lists all breaking changes. Please review carefully before upgrading. - - -## API Version Changes - -### Removed v1.0 API Support - -**Breaking Change:** The v1.0 API format is completely removed and no longer supported. - -#### Before (v0.x) -```python -# This was supported in v0.x -config = { - "version": "v1.0" # ❌ No longer supported -} - -result = m.add( - "memory content", - user_id="alice" -) -``` - -#### After (v1.0.0 ) -```python -# v1.1 is the minimum supported version -config = { - "version": "v1.1" # ✅ Required minimum -} - -result = m.add( - "memory content", - user_id="alice" -) -``` - -**Error Message:** -``` -ValueError: The v1.0 API format is no longer supported in mem0ai 1.0.0+. -Please use v1.1 format which returns a dict with 'results' key. -``` - -## Parameter Removals - -### 1. version Parameter in Method Calls - -**Breaking Change:** Version parameter removed from method calls. - -#### Before (v0.x) -```python -result = m.add("content", user_id="alice", version="v1.0") -``` - -#### After (v1.0.0 ) -```python -result = m.add("content", user_id="alice") -``` - -### 2. async_mode Parameter (Platform Client) - -**Change:** For `MemoryClient` (Platform API), `async_mode` now defaults to `True` but can still be configured. - -#### Before (v0.x) -```python -from mem0 import MemoryClient - -client = MemoryClient(api_key="your-key") -result = client.add("content", user_id="alice", async_mode=True) -result = client.add("content", user_id="alice", async_mode=False) -``` - -#### After (v1.0.0 ) -```python -from mem0 import MemoryClient - -client = MemoryClient(api_key="your-key") - -# async_mode now defaults to True, but you can still override it -result = client.add("content", user_id="alice") # Uses async_mode=True by default - -# You can still explicitly set it to False if needed -result = client.add("content", user_id="alice", async_mode=False) -``` - -## Response Format Changes - -### Standardized Response Structure - -**Breaking Change:** All responses now return a standardized dictionary format. - -#### Before (v0.x) -```python -# Could return different formats based on version configuration -result = m.add("content", user_id="alice") -# With v1.0: Returns [{"id": "...", "memory": "...", "event": "ADD"}] -# With v1.1: Returns {"results": [{"id": "...", "memory": "...", "event": "ADD"}]} -``` - -#### After (v1.0.0 ) -```python -# Always returns standardized format -result = m.add("content", user_id="alice") -# Always returns: {"results": [{"id": "...", "memory": "...", "event": "ADD"}]} - -# Access results consistently -for memory in result["results"]: - print(memory["memory"]) -``` - -## Configuration Changes - -### Version Configuration - -**Breaking Change:** Default API version changed. - -#### Before (v0.x) -```python -# v1.0 was supported -config = { - "version": "v1.0" # ❌ No longer supported -} -``` - -#### After (v1.0.0 ) -```python -# v1.1 is minimum, v1.1 is default -config = { - "version": "v1.1" # ✅ Minimum supported -} - -# Or omit for default -config = { - # version defaults to v1.1 -} -``` - -### Memory Configuration - -**Breaking Change:** Some configuration options have changed defaults. - -#### Before (v0.x) -```python -from mem0 import Memory - -# Default configuration in v0.x -m = Memory() # Used default settings suitable for v0.x -``` - -#### After (v1.0.0 ) -```python -from mem0 import Memory - -# Default configuration optimized for v1.0.0 -m = Memory() # Uses v1.1+ optimized defaults - -# Explicit configuration recommended -config = { - "version": "v1.1", - "vector_store": { - "provider": "qdrant", - "config": { - "host": "localhost", - "port": 6333 - } - } -} -m = Memory.from_config(config) -``` - -## Method Signature Changes - -### Search Method - -**Enhanced but backward compatible:** - -#### Before (v0.x) -```python -results = m.search( - "query", - user_id="alice", - filters={"key": "value"} # Simple key-value only -) -``` - -#### After (v1.0.0 ) -```python -# Basic usage remains the same -results = m.search("query", user_id="alice") - -# Enhanced filtering available (optional) -results = m.search( - "query", - user_id="alice", - filters={ - "AND": [ - {"key": "value"}, - {"score": {"gte": 0.8}} - ] - }, - rerank=True # New parameter -) -``` - -## Error Handling Changes - -### New Error Types - -**Breaking Change:** More specific error types and messages. - -#### Before (v0.x) -```python -try: - result = m.add("content", user_id="alice", version="v1.0") -except Exception as e: - print(f"Generic error: {e}") -``` - -#### After (v1.0.0 ) -```python -try: - result = m.add("content", user_id="alice") -except ValueError as e: - if "v1.0 API format is no longer supported" in str(e): - # Handle version error specifically - print("Please upgrade your code to use v1.1+ format") - else: - print(f"Value error: {e}") -except Exception as e: - print(f"Unexpected error: {e}") -``` - -### Validation Changes - -**Breaking Change:** Stricter parameter validation. - -#### Before (v0.x) -```python -# Some invalid parameters might have been ignored -result = m.add( - "content", - user_id="alice", - invalid_param="ignored" # Might have been silently ignored -) -``` - -#### After (v1.0.0 ) -```python -# Strict validation - unknown parameters cause errors -try: - result = m.add( - "content", - user_id="alice", - invalid_param="value" # ❌ Will raise TypeError - ) -except TypeError as e: - print(f"Invalid parameter: {e}") -``` - -## Import Changes - -### No Breaking Changes in Imports - -**Good News:** Import statements remain the same. - -```python -# These imports work in both v0.x and v1.0.0 -from mem0 import Memory, AsyncMemory -from mem0 import MemoryConfig -``` - -## Dependency Changes - -### Minimum Python Version - -**Potential Breaking Change:** Check Python version requirements. - -#### Before (v0.x) -- Python 3.8+ supported - -#### After (v1.0.0 ) -- Python 3.9+ required (check current requirements) - -### Package Dependencies - -**Breaking Change:** Some dependencies updated with potential breaking changes. - -```bash -# Check for conflicts after upgrade -pip install --upgrade mem0ai -pip check # Verify no dependency conflicts -``` - -## Data Migration - -### Database Schema - -**Good News:** No database schema changes required. - -- Existing memories remain compatible -- No data migration required -- Vector store data unchanged - -### Memory Format - -**Good News:** Memory storage format unchanged. - -- Existing memories work with v1.0.0 -- Search continues to work with old memories -- No re-indexing required - -## Testing Changes - -### Test Updates Required - -**Breaking Change:** Update tests for new response format. - -#### Before (v0.x) -```python -def test_add_memory(): - result = m.add("content", user_id="alice") - assert isinstance(result, list) # ❌ No longer true - assert len(result) > 0 -``` - -#### After (v1.0.0 ) -```python -def test_add_memory(): - result = m.add("content", user_id="alice") - assert isinstance(result, dict) # ✅ Always dict - assert "results" in result # ✅ Always has results key - assert len(result["results"]) > 0 -``` - -## Rollback Considerations - -### Safe Rollback Process - -If you need to rollback: - -```bash -# 1. Rollback package -pip install mem0ai==0.1.20 # Last stable v0.x - -# 2. Revert code changes -git checkout previous_commit - -# 3. Test functionality -python test_mem0_functionality.py -``` - -### Data Safety - -- **Safe:** Memories stored in v0.x format work with v1.0.0 -- **Safe:** Rollback doesn't lose data -- **Safe:** Vector store data remains intact - -## Next Steps - -1. **Review all breaking changes** in your codebase -2. **Update method calls** to remove deprecated parameters -3. **Update response handling** to use standardized format -4. **Test thoroughly** with your existing data -5. **Update error handling** for new error types - - - - Step-by-step migration instructions - - - Complete API reference changes - - - - -**Need Help?** If you encounter issues during migration, check our [GitHub Discussions](https://github.com/mem0ai/mem0/discussions) or community support channels. - \ No newline at end of file diff --git a/docs/migration/oss-to-platform.mdx b/docs/migration/oss-to-platform.mdx index 71fbd2668..24f5c1494 100644 --- a/docs/migration/oss-to-platform.mdx +++ b/docs/migration/oss-to-platform.mdx @@ -81,6 +81,10 @@ client = MemoryClient(api_key="m0-...") **Critical Change**: Platform uses v2 endpoints that require filtering parameters to be nested inside a `filters` dictionary. + + The `limit` parameter has been removed in favor of `top_k` across all SDKs. Update any code using `limit=` to use `top_k=` instead. + + | Method | Open Source | Platform | | ------ | ----------- | -------- | | `search()` | `m.search(query, user_id="alex")` | `client.search(query, filters={"user_id": "alex"})` | @@ -121,18 +125,18 @@ Note: `add()` and `delete()` methods remain unchanged. The `update()` method is ```python Open Source (Old) # Get all memories for a user - memories = m.get_all(user_id="alex", limit=10) + memories = m.get_all(user_id="alex", top_k=10) # Get memories with pagination - memories = m.get_all(user_id="alex", limit=5, offset=10) + memories = m.get_all(user_id="alex", top_k=5, offset=10) ``` ```python Platform (New) # Get all memories for a user - memories = client.get_all(filters={"user_id": "alex"}, limit=10) + memories = client.get_all(filters={"user_id": "alex"}, top_k=10) # Get memories with pagination - memories = client.get_all(filters={"user_id": "alex"}, limit=5, offset=10) + memories = client.get_all(filters={"user_id": "alex"}, top_k=5, offset=10) ``` @@ -283,7 +287,7 @@ The Platform introduces powerful capabilities not available in OSS: "user preferences", filters={"user_id": "alex"}, rerank=True, # Platform exclusive - limit=5 + top_k=5 ) # Search with keyword expansion @@ -335,7 +339,7 @@ The Platform introduces powerful capabilities not available in OSS: {"timestamp": {"gte": "2024-01-01"}} ] }, - limit=100 + top_k=100 ) # Monitor usage patterns diff --git a/docs/migration/v0-to-v1.mdx b/docs/migration/v0-to-v1.mdx deleted file mode 100644 index 6d611eaa6..000000000 --- a/docs/migration/v0-to-v1.mdx +++ /dev/null @@ -1,481 +0,0 @@ ---- -title: Migrating from v0.x to v1.0.0 -description: 'Complete guide to upgrade your Mem0 implementation to version 1.0.0 ' -icon: "arrow-right" -iconType: "solid" ---- - - -**Breaking Changes Ahead!** Mem0 1.0.0 introduces several breaking changes. Please read this guide carefully before upgrading. - - -## Overview - -Mem0 1.0.0 is a major release that modernizes the API, improves performance, and adds powerful new features. This guide will help you migrate your existing v0.x implementation to the new version. - -## Key Changes Summary - -| Feature | v0.x | v1.0.0 | Migration Required | -|---------|------|-------------|-------------------| -| API Version | v1.0 supported | v1.0 **removed**, v1.1+ only | ✅ Yes | -| Async Mode (Platform Client) | Optional/manual | Defaults to `True`, configurable | ⚠️ Partial | -| Metadata Filtering | Basic | Enhanced with operators | ⚠️ Optional | -| Reranking | Not available | Full support | ⚠️ Optional | - -## Step-by-Step Migration - -### 1. Update Installation - -```bash -# Update to the latest version -pip install --upgrade mem0ai -``` - -### 2. Remove Deprecated Parameters - -#### Before (v0.x) -```python -from mem0 import Memory - -# These parameters are no longer supported -m = Memory() -result = m.add( - "I love pizza", - user_id="alice", - version="v1.0" # ❌ REMOVED -) -``` - -#### After (v1.0.0 ) -```python -from mem0 import Memory - -# Clean, simplified API -m = Memory() -result = m.add( - "I love pizza", - user_id="alice" - # version parameter removed -) -``` - -### 3. Update Configuration - -#### Before (v0.x) -```python -config = { - "vector_store": { - "provider": "qdrant", - "config": { - "host": "localhost", - "port": 6333 - } - }, - "version": "v1.0" # ❌ No longer supported -} - -m = Memory.from_config(config) -``` - -#### After (v1.0.0 ) -```python -config = { - "vector_store": { - "provider": "qdrant", - "config": { - "host": "localhost", - "port": 6333 - } - }, - "version": "v1.1" # ✅ v1.1 is the minimum supported version -} - -m = Memory.from_config(config) -``` - -### 4. Handle Response Format Changes - -#### Before (v0.x) -```python -# Response could be a list or dict depending on version -result = m.add("I love coffee", user_id="alice") - -if isinstance(result, list): - # Handle list format - for item in result: - print(item["memory"]) -else: - # Handle dict format - print(result["results"]) -``` - -#### After (v1.0.0 ) -```python -# Response is always a standardized dict with "results" key -result = m.add("I love coffee", user_id="alice") - -# Always access via "results" key -for item in result["results"]: - print(item["memory"]) -``` - -### 5. Update Search Operations - -#### Before (v0.x) -```python -# Basic search -results = m.search("What do I like?", user_id="alice") - -# With filters -results = m.search( - "What do I like?", - user_id="alice", - filters={"category": "food"} -) -``` - -#### After (v1.0.0 ) -```python -# Same basic search API -results = m.search("What do I like?", user_id="alice") - -# Enhanced filtering with operators (optional upgrade) -results = m.search( - "What do I like?", - user_id="alice", - filters={ - "AND": [ - {"category": "food"}, - {"rating": {"gte": 8}} - ] - } -) - -# New: Reranking support (optional) -results = m.search( - "What do I like?", - user_id="alice", - rerank=True # Requires reranker configuration -) -``` - -### 6. Platform Client async_mode Default Changed - -**Change:** For `MemoryClient`, the `async_mode` parameter now defaults to `True` for better performance. - -#### Before (v0.x) -```python -from mem0 import MemoryClient - -client = MemoryClient(api_key="your-key") - -# Had to explicitly set async_mode -result = client.add("I enjoy hiking", user_id="alice", async_mode=True) -``` - -#### After (v1.0.0 ) -```python -from mem0 import MemoryClient - -client = MemoryClient(api_key="your-key") - -# async_mode now defaults to True (best performance) -result = client.add("I enjoy hiking", user_id="alice") - -# You can still override if needed for synchronous processing -result = client.add("I enjoy hiking", user_id="alice", async_mode=False) -``` - -## Configuration Migration - -### Basic Configuration - -#### Before (v0.x) -```python -config = { - "vector_store": { - "provider": "qdrant", - "config": { - "host": "localhost", - "port": 6333 - } - }, - "llm": { - "provider": "openai", - "config": { - "model": "gpt-3.5-turbo", - "api_key": "your-key" - } - }, - "version": "v1.0" -} -``` - -#### After (v1.0.0 ) -```python -config = { - "vector_store": { - "provider": "qdrant", - "config": { - "host": "localhost", - "port": 6333 - } - }, - "llm": { - "provider": "openai", - "config": { - "model": "gpt-3.5-turbo", - "api_key": "your-key" - } - }, - "version": "v1.1", # Minimum supported version - - # New optional features - "reranker": { - "provider": "cohere", - "config": { - "model": "rerank-english-v3.0", - "api_key": "your-cohere-key" - } - } -} -``` - -### Enhanced Features (Optional) - -```python -# Take advantage of new features -config = { - "vector_store": { - "provider": "qdrant", - "config": { - "host": "localhost", - "port": 6333 - } - }, - "llm": { - "provider": "openai", - "config": { - "model": "gpt-4", - "api_key": "your-key" - } - }, - "embedder": { - "provider": "openai", - "config": { - "model": "text-embedding-3-small", - "api_key": "your-key" - } - }, - "reranker": { - "provider": "sentence_transformer", - "config": { - "model": "cross-encoder/ms-marco-MiniLM-L-6-v2" - } - }, - "version": "v1.1" -} -``` - -## Error Handling Migration - -### Before (v0.x) -```python -try: - result = m.add("memory", user_id="alice", version="v1.0") -except Exception as e: - print(f"Error: {e}") -``` - -### After (v1.0.0 ) -```python -try: - result = m.add("memory", user_id="alice") -except ValueError as e: - if "v1.0 API format is no longer supported" in str(e): - print("Please upgrade your code to use v1.1+ format") - else: - print(f"Error: {e}") -except Exception as e: - print(f"Unexpected error: {e}") -``` - -## Testing Your Migration - -### 1. Basic Functionality Test - -```python -def test_basic_functionality(): - m = Memory() - - # Test add - result = m.add("I love testing", user_id="test_user") - assert "results" in result - assert len(result["results"]) > 0 - - # Test search - search_results = m.search("testing", user_id="test_user") - assert "results" in search_results - - # Test get_all - all_memories = m.get_all(user_id="test_user") - assert "results" in all_memories - - print("✅ Basic functionality test passed") - -test_basic_functionality() -``` - -### 2. Enhanced Features Test - -```python -def test_enhanced_features(): - config = { - "reranker": { - "provider": "sentence_transformer", - "config": { - "model": "cross-encoder/ms-marco-MiniLM-L-6-v2" - } - } - } - - m = Memory.from_config(config) - - # Test reranking - m.add("I love advanced features", user_id="test_user") - results = m.search("features", user_id="test_user", rerank=True) - assert "results" in results - - # Test enhanced filtering - results = m.search( - "features", - user_id="test_user", - filters={"user_id": {"eq": "test_user"}} - ) - assert "results" in results - - print("✅ Enhanced features test passed") - -test_enhanced_features() -``` - -## Common Migration Issues - -### Issue 1: Version Error - -**Error:** -``` -ValueError: The v1.0 API format is no longer supported in mem0ai 1.0.0+ -``` - -**Solution:** -```python -# Remove version parameters or set to v1.1+ -config = { - # ... other config - "version": "v1.1" # or remove entirely for default -} -``` - -### Issue 2: Response Format Error - -**Error:** -``` -KeyError: 'results' -``` - -**Solution:** -```python -# Always access response via "results" key -result = m.add("memory", user_id="alice") -memories = result["results"] # Not result directly -``` - -### Issue 3: Parameter Error - -**Error:** -``` -TypeError: add() got an unexpected keyword argument 'output_format' -``` - -**Solution:** -```python -# Remove deprecated parameters -result = m.add( - "memory", - user_id="alice" - # Remove: version -) -``` - -## Rollback Plan - -If you encounter issues during migration: - -### 1. Immediate Rollback - -```bash -# Downgrade to last v0.x version -pip install mem0ai==0.1.20 # Replace with your last working version -``` - -### 2. Gradual Migration - -```python -# Test both versions side by side -import mem0_v0 # Your old version -import mem0 # New version - -def compare_results(query, user_id): - old_results = mem0_v0.search(query, user_id=user_id) - new_results = mem0.search(query, user_id=user_id) - - print("Old format:", old_results) - print("New format:", new_results["results"]) -``` - -## Performance Improvements - -### Before (v0.x) -```python -# Sequential operations -result1 = m.add("memory 1", user_id="alice") -result2 = m.add("memory 2", user_id="alice") -result3 = m.search("query", user_id="alice") -``` - -### After (v1.0.0 ) -```python -# Better async performance -async def batch_operations(): - async_memory = AsyncMemory() - - # Concurrent operations - results = await asyncio.gather( - async_memory.add("memory 1", user_id="alice"), - async_memory.add("memory 2", user_id="alice"), - async_memory.search("query", user_id="alice") - ) - return results -``` - -## Next Steps - -1. **Complete the migration** using this guide -2. **Test thoroughly** with your existing data -3. **Explore new features** like enhanced filtering and reranking -4. **Update your documentation** to reflect the new API -5. **Monitor performance** and optimize as needed - - - - Detailed list of all breaking changes - - - Complete API reference changes - - - - -Need help with migration? Check our [GitHub Discussions](https://github.com/mem0ai/mem0/discussions) or reach out to our community for support. - \ No newline at end of file diff --git a/docs/open-source/features/async-memory.mdx b/docs/open-source/features/async-memory.mdx index f19425bba..1f3e96f59 100644 --- a/docs/open-source/features/async-memory.mdx +++ b/docs/open-source/features/async-memory.mdx @@ -240,7 +240,7 @@ async_openai_client = AsyncOpenAI() async_memory = AsyncMemory() async def chat_with_memories(message: str, user_id: str = "default_user") -> str: - search_result = await async_memory.search(query=message, user_id=user_id, limit=3) + search_result = await async_memory.search(query=message, user_id=user_id, top_k=3) relevant_memories = search_result["results"] memories_str = "\n".join(f"- {entry['memory']}" for entry in relevant_memories) @@ -326,7 +326,7 @@ async def add_memory(messages: list, user_id: str): @app.get("/memories/search") async def search_memories(query: str, user_id: str, limit: int = 10): try: - result = await memory.search(query=query, user_id=user_id, limit=limit) + result = await memory.search(query=query, user_id=user_id, top_k=limit) return {"status": "success", "data": result} except Exception as exc: raise HTTPException(status_code=500, detail=str(exc)) diff --git a/docs/open-source/features/custom-fact-extraction-prompt.mdx b/docs/open-source/features/custom-instructions.mdx similarity index 85% rename from docs/open-source/features/custom-fact-extraction-prompt.mdx rename to docs/open-source/features/custom-instructions.mdx index 3738becbd..ea0155644 100644 --- a/docs/open-source/features/custom-fact-extraction-prompt.mdx +++ b/docs/open-source/features/custom-instructions.mdx @@ -1,13 +1,13 @@ --- -title: Custom Fact Extraction Prompt +title: Custom Instructions description: Tailor fact extraction so Mem0 stores only the details you care about. icon: "wand-magic-sparkles" --- -Custom fact extraction prompts let you decide exactly which facts Mem0 records from a conversation. Define a focused prompt, give a few examples, and Mem0 will add only the memories that match your use case. +Custom instructions let you decide exactly which facts Mem0 records from a conversation. Define a focused prompt, give a few examples, and Mem0 will add only the memories that match your use case. - **You’ll use this when…** + **You'll use this when...** - A project needs domain-specific facts (order numbers, customer info) without storing casual chatter. - You already have a clear schema for memories and want the LLM to follow it. - You must prevent irrelevant details from entering long-term storage. @@ -17,6 +17,10 @@ Custom fact extraction prompts let you decide exactly which facts Mem0 records f Prompts that are too broad cause unrelated facts to slip through. Keep instructions tight and test them with real transcripts. + + The `custom_fact_extraction_prompt` parameter has been renamed to `custom_instructions`. If you are upgrading from an older version, update your configuration accordingly. + + --- ## Feature anatomy @@ -24,13 +28,13 @@ Custom fact extraction prompts let you decide exactly which facts Mem0 records f - **Prompt instructions:** Describe which entities or phrases to keep. Specific guidance keeps the extractor focused. - **Few-shot examples:** Show positive and negative cases so the model copies the right format. - **Structured output:** Responses return JSON with a `facts` array that Mem0 converts into individual memories. -- **LLM configuration:** `custom_fact_extraction_prompt` (Python) or `customPrompt` (TypeScript) lives alongside your model settings. +- **LLM configuration:** `custom_instructions` (Python) or `customInstructions` (TypeScript) lives alongside your model settings. - 1. State the allowed fact types. - 2. Include short examples that mirror production messages. - 3. Show both empty (`[]`) and populated outputs. + 1. State the allowed fact types. + 2. Include short examples that mirror production messages. + 3. Show both empty (`[]`) and populated outputs. 4. Remind the model to return JSON with a `facts` key only. @@ -43,8 +47,8 @@ Custom fact extraction prompts let you decide exactly which facts Mem0 records f ```python Python -custom_fact_extraction_prompt = """ -Please only extract entities containing customer support information, order details, and user information. +custom_instructions = """ +Please only extract entities containing customer support information, order details, and user information. Here are some few shot examples: Input: Hi. @@ -67,8 +71,8 @@ Return the facts and customer information in a json format as shown above. ``` ```ts TypeScript -const customPrompt = ` -Please only extract entities containing customer support information, order details, and user information. +const customInstructions = ` +Please only extract entities containing customer support information, order details, and user information. Here are some few shot examples: Input: Hi. @@ -110,7 +114,7 @@ config = { "max_tokens": 2000, } }, - "custom_fact_extraction_prompt": custom_fact_extraction_prompt, + "custom_instructions": custom_instructions, "version": "v1.1" } @@ -131,7 +135,7 @@ const config = { maxTokens: 1500, }, }, - customPrompt: customPrompt, + customInstructions: customInstructions, }; const memory = new Memory(config); diff --git a/docs/open-source/features/custom-update-memory-prompt.mdx b/docs/open-source/features/custom-update-memory-prompt.mdx index 06ce6566a..ba6d1ca66 100644 --- a/docs/open-source/features/custom-update-memory-prompt.mdx +++ b/docs/open-source/features/custom-update-memory-prompt.mdx @@ -263,7 +263,7 @@ Please note to return the IDs in the output from the input IDs only and do not g - Log each decision so product teams can review why a change happened. - The prompt works alongside `custom_fact_extraction_prompt`—fact extraction identifies candidate facts, and the update prompt decides how to merge them into long-term storage. + The prompt works alongside `custom_instructions`—fact extraction identifies candidate facts, and the update prompt decides how to merge them into long-term storage. --- @@ -288,7 +288,7 @@ Please note to return the IDs in the output from the input IDs only and do not g ## Compare prompts -| Feature | `custom_update_memory_prompt` | `custom_fact_extraction_prompt` | +| Feature | `custom_update_memory_prompt` | `custom_instructions` | | --- | --- | --- | | Primary job | Decide memory actions (ADD/UPDATE/DELETE/NONE) | Pull facts from user and assistant messages | | Inputs | Retrieved facts + existing memory entries | Raw conversation turns | @@ -297,7 +297,7 @@ Please note to return the IDs in the output from the input IDs only and do not g --- - + Coordinate both prompts so fact extraction feeds clean inputs into the update flow. diff --git a/docs/open-source/features/graph-memory.mdx b/docs/open-source/features/graph-memory.mdx index c7499d4de..a9ef43768 100644 --- a/docs/open-source/features/graph-memory.mdx +++ b/docs/open-source/features/graph-memory.mdx @@ -94,7 +94,7 @@ memory.add(conversation, user_id="demo-user") results = memory.search( "Who did Alice meet at GraphConf?", user_id="demo-user", - limit=3, + top_k=3, rerank=True, ) @@ -146,7 +146,7 @@ await memory.add(conversation, { userId: "demo-user" }); const results = await memory.search( "Who did Alice meet at GraphConf?", - { userId: "demo-user", limit: 3, rerank: true } + { userId: "demo-user", topK: 3, rerank: true } ); results.results.forEach((hit) => { @@ -204,7 +204,7 @@ const config = { username: process.env.NEO4J_USERNAME!, password: process.env.NEO4J_PASSWORD!, }, - customPrompt: "Please only capture people, organisations, and project links.", + customInstructions: "Please only capture people, organisations, and project links.", } }; @@ -265,7 +265,7 @@ Monitor graph growth, especially on free tiers, by periodically cleaning dormant ## Decision Points - Select the graph store that fits your deployment (managed Aura vs. self-hosted Neo4j vs. AWS Neptune vs. local Kuzu vs. Apache AGE on PostgreSQL). -- Decide when to enable graph writes per request; routine conversations may stay vector-only to save latency. +- Decide whether to include a graph store in your config; routine conversations may stay vector-only to save latency. - Set a policy for pruning stale relationships so your graph stays fast and affordable. ## Provider setup diff --git a/docs/open-source/features/openai_compatibility.mdx b/docs/open-source/features/openai_compatibility.mdx index c430e8b7f..b027daaf7 100644 --- a/docs/open-source/features/openai_compatibility.mdx +++ b/docs/open-source/features/openai_compatibility.mdx @@ -131,7 +131,7 @@ print(response.choices[0].message.content) | `run_id` | `str` | Optional session/run identifier for short-lived flows. | | `metadata` | `dict` | Store extra fields alongside each memory entry. | | `filters` | `dict` | Restrict retrieval to specific memories while responding. | -| `limit` | `int` | Cap how many memories Mem0 pulls into the context (default 10). | +| `top_k` | `int` | Cap how many memories Mem0 pulls into the context (default 10). | Other request fields mirror OpenAI’s chat completion API. diff --git a/docs/open-source/features/overview.mdx b/docs/open-source/features/overview.mdx index dfef5cdbf..520b7fb8c 100644 --- a/docs/open-source/features/overview.mdx +++ b/docs/open-source/features/overview.mdx @@ -30,7 +30,7 @@ Mem0 Open Source ships with capabilities that adapt memory behavior for producti Process images, audio, and video memories. - + Tailor how facts are extracted from text. diff --git a/docs/open-source/features/reranker-search.mdx b/docs/open-source/features/reranker-search.mdx index 893d9a7af..937a63098 100644 --- a/docs/open-source/features/reranker-search.mdx +++ b/docs/open-source/features/reranker-search.mdx @@ -321,7 +321,7 @@ results = m.search( ] }, rerank=True, - limit=20 + top_k=20 ) ``` @@ -366,7 +366,7 @@ results = m.search( user_id="reader123", filters={"content_type": "book_recommendation"}, rerank=True, - limit=10 + top_k=10 ) for result in results["results"]: diff --git a/docs/open-source/node-quickstart.mdx b/docs/open-source/node-quickstart.mdx index 2128f79b5..e42fcf20b 100644 --- a/docs/open-source/node-quickstart.mdx +++ b/docs/open-source/node-quickstart.mdx @@ -230,7 +230,7 @@ Mem0 offers granular configuration across vector stores, LLMs, embedders, and hi | --- | --- | --- | | `historyDbPath` | Path to history database | `"{mem0_dir}/history.db"` | | `version` | API version | `"v1.0"` | -| `customPrompt` | Custom processing prompt | `undefined` | +| `customInstructions` | Custom processing prompt | `undefined` | | Parameter | Description | Default | @@ -273,7 +273,7 @@ const config = { } }, disableHistory: false, - customPrompt: "I'm a virtual assistant. I'm here to help you with your queries." + customInstructions: "I'm a virtual assistant. I'm here to help you with your queries." }; ``` diff --git a/docs/openapi.json b/docs/openapi.json index 5dd03e6eb..53c04ccf6 100644 --- a/docs/openapi.json +++ b/docs/openapi.json @@ -183,7 +183,7 @@ "x-code-samples": [ { "lang": "Python", - "source": "# To use the Python SDK, install the package:\n# pip install mem0ai\n\nfrom mem0 import MemoryClient\nclient = MemoryClient(api_key=\"your_api_key\", org_id=\"your_org_id\", project_id=\"your_project_id\")\nusers = client.users()\nprint(users)" + "source": "# To use the Python SDK, install the package:\n# pip install mem0ai\n\nfrom mem0 import MemoryClient\nclient = MemoryClient(api_key=\"your_api_key\")\nusers = client.users()\nprint(users)" }, { "lang": "JavaScript", @@ -717,7 +717,7 @@ "x-code-samples": [ { "lang": "Python", - "source": "# To use the Python SDK, install the package:\n# pip install mem0ai\n\nfrom mem0 import MemoryClient\n\nclient = MemoryClient(api_key=\"your_api_key\", org_id=\"your_org_id\", project_id=\"your_project_id\")\n\njson_schema = {pydantic_json_schema}\nfilters = {\n \"AND\": [\n {\"user_id\": \"alex\"}\n ]\n}\n\nresponse = client.create_memory_export(\n schema=json_schema,\n filters=filters\n)\nprint(response)" + "source": "# To use the Python SDK, install the package:\n# pip install mem0ai\n\nfrom mem0 import MemoryClient\n\nclient = MemoryClient(api_key=\"your_api_key\")\n\njson_schema = {pydantic_json_schema}\nfilters = {\n \"AND\": [\n {\"user_id\": \"alex\"}\n ]\n}\n\nresponse = client.create_memory_export(\n schema=json_schema,\n filters=filters\n)\nprint(response)" }, { "lang": "JavaScript", @@ -845,7 +845,7 @@ "x-code-samples": [ { "lang": "Python", - "source": "# To use the Python SDK, install the package:\n# pip install mem0ai\n\nfrom mem0 import MemoryClient\n\nclient = MemoryClient(api_key=\"your_api_key\", org_id=\"your_org_id\", project_id=\"project_id\")\n\nmemory_export_id = \"\"\n\nresponse = client.get_memory_export(memory_export_id=memory_export_id)\nprint(response)" + "source": "# To use the Python SDK, install the package:\n# pip install mem0ai\n\nfrom mem0 import MemoryClient\n\nclient = MemoryClient(api_key=\"your_api_key\")\n\nmemory_export_id = \"\"\n\nresponse = client.get_memory_export(memory_export_id=memory_export_id)\nprint(response)" }, { "lang": "JavaScript", @@ -1096,7 +1096,7 @@ "x-code-samples": [ { "lang": "Python", - "source": "# To use the Python SDK, install the package:\n# pip install mem0ai\n\nfrom mem0 import MemoryClient\nclient = MemoryClient(api_key=\"your_api_key\", org_id=\"your_org_id\", project_id=\"your_project_id\")\n\n# Retrieve memories for a specific user\nuser_memories = client.get_all(user_id=\"\")\n\nprint(user_memories)" + "source": "# To use the Python SDK, install the package:\n# pip install mem0ai\n\nfrom mem0 import MemoryClient\nclient = MemoryClient(api_key=\"your_api_key\")\n\n# Retrieve memories for a specific user\nuser_memories = client.get_all(user_id=\"\")\n\nprint(user_memories)" }, { "lang": "JavaScript", @@ -1203,11 +1203,11 @@ "x-code-samples": [ { "lang": "Python", - "source": "# To use the Python SDK, install the package:\n# pip install mem0ai\n\nfrom mem0 import MemoryClient\n\nclient = MemoryClient(api_key=\"your_api_key\", org_id=\"your_org_id\", project_id=\"your_project_id\")\n\nmessages = [\n {\"role\": \"user\", \"content\": \"\"},\n {\"role\": \"assistant\", \"content\": \"\"}\n]\n\nclient.add(messages, user_id=\"\", version=\"v2\")" + "source": "# To use the Python SDK, install the package:\n# pip install mem0ai\n\nfrom mem0 import MemoryClient\n\nclient = MemoryClient(api_key=\"your_api_key\")\n\nmessages = [\n {\"role\": \"user\", \"content\": \"\"},\n {\"role\": \"assistant\", \"content\": \"\"}\n]\n\nclient.add(messages, user_id=\"\")" }, { "lang": "JavaScript", - "source": "// To use the JavaScript SDK, install the package:\n// npm i mem0ai\n\nimport MemoryClient from 'mem0ai';\nconst client = new MemoryClient({ apiKey: \"your-api-key\" });\n\nconst messages = [\n { role: \"user\", content: \"Hi, I'm Alex. I'm a vegetarian and I'm allergic to nuts.\" },\n { role: \"assistant\", content: \"Hello Alex! I've noted that you're a vegetarian and have a nut allergy. I'll keep this in mind for any food-related recommendations or discussions.\" }\n];\n\nclient.add(messages, { user_id: \"\", version: \"v2\" })\n .then(result => console.log(result))\n .catch(error => console.error(error));" + "source": "// To use the JavaScript SDK, install the package:\n// npm i mem0ai\n\nimport MemoryClient from 'mem0ai';\nconst client = new MemoryClient({ apiKey: \"your-api-key\" });\n\nconst messages = [\n { role: \"user\", content: \"Hi, I'm Alex. I'm a vegetarian and I'm allergic to nuts.\" },\n { role: \"assistant\", content: \"Hello Alex! I've noted that you're a vegetarian and have a nut allergy. I'll keep this in mind for any food-related recommendations or discussions.\" }\n];\n\nclient.add(messages, { user_id: \"\" })\n .then(result => console.log(result))\n .catch(error => console.error(error));" }, { "lang": "cURL", @@ -1232,7 +1232,7 @@ "tags": [ "memories" ], - "description": "Delete memories by filter. At least one filter is required — previously omitting all filters silently deleted everything; now it returns a validation error.", + "description": "Delete memories by filter. At least one filter is required \u2014 previously omitting all filters silently deleted everything; now it returns a validation error.", "operationId": "memories_delete", "parameters": [ { @@ -1315,15 +1315,15 @@ "x-code-samples": [ { "lang": "Python", - "source": "# To use the Python SDK, install the package:\n# pip install mem0ai\n\nfrom mem0 import MemoryClient\nclient = MemoryClient(api_key=\"your_api_key\", org_id=\"your_org_id\", project_id=\"your_project_id\")\n\n# Delete all memories for a specific user\nclient.delete_all(user_id=\"\")\n\n# Delete all memories for every user in the project (wildcard)\nclient.delete_all(user_id=\"*\")\n\n# Full project wipe — all four filters must be explicitly set to \"*\"\nclient.delete_all(user_id=\"*\", agent_id=\"*\", app_id=\"*\", run_id=\"*\")\n\n# NOTE: Calling delete_all() with no filters raises a validation error.\n# At least one filter is required to prevent accidental data loss." + "source": "# To use the Python SDK, install the package:\n# pip install mem0ai\n\nfrom mem0 import MemoryClient\nclient = MemoryClient(api_key=\"your_api_key\")\n\n# Delete all memories for a specific user\nclient.delete_all(user_id=\"\")\n\n# Delete all memories for every user in the project (wildcard)\nclient.delete_all(user_id=\"*\")\n\n# Full project wipe \u2014 all four filters must be explicitly set to \"*\"\nclient.delete_all(user_id=\"*\", agent_id=\"*\", app_id=\"*\", run_id=\"*\")\n\n# NOTE: Calling delete_all() with no filters raises a validation error.\n# At least one filter is required to prevent accidental data loss." }, { "lang": "JavaScript", - "source": "// To use the JavaScript SDK, install the package:\n// npm i mem0ai\n\nimport MemoryClient from 'mem0ai';\nconst client = new MemoryClient({ apiKey: \"your-api-key\" });\n\n// Delete all memories for a specific user\nclient.deleteAll({ user_id: \"\" })\n .then(result => console.log(result))\n .catch(error => console.error(error));\n\n// Delete all memories for every user in the project (wildcard)\nclient.deleteAll({ user_id: \"*\" })\n .then(result => console.log(result))\n .catch(error => console.error(error));\n\n// Full project wipe — all four filters must be explicitly set to \"*\"\nclient.deleteAll({ user_id: \"*\", agent_id: \"*\", app_id: \"*\", run_id: \"*\" })\n .then(result => console.log(result))\n .catch(error => console.error(error));" + "source": "// To use the JavaScript SDK, install the package:\n// npm i mem0ai\n\nimport MemoryClient from 'mem0ai';\nconst client = new MemoryClient({ apiKey: \"your-api-key\" });\n\n// Delete all memories for a specific user\nclient.deleteAll({ user_id: \"\" })\n .then(result => console.log(result))\n .catch(error => console.error(error));\n\n// Delete all memories for every user in the project (wildcard)\nclient.deleteAll({ user_id: \"*\" })\n .then(result => console.log(result))\n .catch(error => console.error(error));\n\n// Full project wipe \u2014 all four filters must be explicitly set to \"*\"\nclient.deleteAll({ user_id: \"*\", agent_id: \"*\", app_id: \"*\", run_id: \"*\" })\n .then(result => console.log(result))\n .catch(error => console.error(error));" }, { "lang": "cURL", - "source": "# Delete memories for a specific user\ncurl --request DELETE \\\n --url 'https://api.mem0.ai/v1/memories/?user_id=' \\\n --header 'Authorization: Token '\n\n# Delete memories for all users (wildcard)\ncurl --request DELETE \\\n --url 'https://api.mem0.ai/v1/memories/?user_id=*' \\\n --header 'Authorization: Token '\n\n# Full project wipe — all four filters must be set to *\ncurl --request DELETE \\\n --url 'https://api.mem0.ai/v1/memories/?user_id=*&agent_id=*&app_id=*&run_id=*' \\\n --header 'Authorization: Token '" + "source": "# Delete memories for a specific user\ncurl --request DELETE \\\n --url 'https://api.mem0.ai/v1/memories/?user_id=' \\\n --header 'Authorization: Token '\n\n# Delete memories for all users (wildcard)\ncurl --request DELETE \\\n --url 'https://api.mem0.ai/v1/memories/?user_id=*' \\\n --header 'Authorization: Token '\n\n# Full project wipe \u2014 all four filters must be set to *\ncurl --request DELETE \\\n --url 'https://api.mem0.ai/v1/memories/?user_id=*&agent_id=*&app_id=*&run_id=*' \\\n --header 'Authorization: Token '" }, { "lang": "Go", @@ -1441,11 +1441,11 @@ "x-code-samples": [ { "lang": "Python", - "source": "# To use the Python SDK, install the package:\n# pip install mem0ai\n\nfrom mem0 import MemoryClient\nclient = MemoryClient(api_key=\"your_api_key\", org_id=\"your_org_id\", project_id=\"your_project_id\")\n\n# Retrieve memories with filters\nmemories = client.get_all(\n filters={\n \"AND\": [\n {\n \"user_id\": \"alex\"\n },\n {\n \"created_at\": {\n \"gte\": \"2024-07-01\",\n \"lte\": \"2024-07-31\"\n }\n }\n ]\n },\n version=\"v2\"\n)\n\nprint(memories)" + "source": "# To use the Python SDK, install the package:\n# pip install mem0ai\n\nfrom mem0 import MemoryClient\nclient = MemoryClient(api_key=\"your_api_key\")\n\n# Retrieve memories with filters\nmemories = client.get_all(\n filters={\n \"AND\": [\n {\n \"user_id\": \"alex\"\n },\n {\n \"created_at\": {\n \"gte\": \"2024-07-01\",\n \"lte\": \"2024-07-31\"\n }\n }\n ]\n }\n)\n\nprint(memories)" }, { "lang": "JavaScript", - "source": "// To use the JavaScript SDK, install the package:\n// npm i mem0ai\n\nimport MemoryClient from 'mem0ai';\nconst client = new MemoryClient({ apiKey: \"your-api-key\" });\n\nconst filters = {\n AND: [\n { user_id: 'alex' },\n { created_at: { gte: '2024-07-01', lte: '2024-07-31' } }\n ]\n};\n\nclient.getAll({ filters, api_version: 'v2' })\n .then(result => console.log(result))\n .catch(error => console.error(error));" + "source": "// To use the JavaScript SDK, install the package:\n// npm i mem0ai\n\nimport MemoryClient from 'mem0ai';\nconst client = new MemoryClient({ apiKey: \"your-api-key\" });\n\nconst filters = {\n AND: [\n { user_id: 'alex' },\n { created_at: { gte: '2024-07-01', lte: '2024-07-31' } }\n ]\n};\n\nclient.getAll({ filters })\n .then(result => console.log(result))\n .catch(error => console.error(error));" }, { "lang": "cURL", @@ -1589,11 +1589,11 @@ "x-code-samples": [ { "lang": "Python", - "source": "# To use the Python SDK, install the package:\n# pip install mem0ai\n\nfrom mem0 import MemoryClient\nclient = MemoryClient(api_key=\"your_api_key\", org_id=\"your_org_id\", project_id=\"your_project_id\")\n\nquery = \"Your search query here\"\n\nresults = client.search(query, user_id=\"\", output_format=\"v1.1\")\nprint(results)" + "source": "# To use the Python SDK, install the package:\n# pip install mem0ai\n\nfrom mem0 import MemoryClient\nclient = MemoryClient(api_key=\"your_api_key\")\n\nquery = \"Your search query here\"\n\nresults = client.search(query, user_id=\"\")\nprint(results)" }, { "lang": "JavaScript", - "source": "// To use the JavaScript SDK, install the package:\n// npm i mem0ai\n\nimport MemoryClient from 'mem0ai';\nconst client = new MemoryClient({ apiKey: \"your-api-key\" });\n\nconst query = \"Your search query here\";\n\nclient.search(query, { user_id: \"\", output_format: \"v1.1\" })\n .then(result => console.log(result))\n .catch(error => console.error(error));" + "source": "// To use the JavaScript SDK, install the package:\n// npm i mem0ai\n\nimport MemoryClient from 'mem0ai';\nconst client = new MemoryClient({ apiKey: \"your-api-key\" });\n\nconst query = \"Your search query here\";\n\nclient.search(query, { user_id: \"\" })\n .then(result => console.log(result))\n .catch(error => console.error(error));" }, { "lang": "cURL", @@ -1708,11 +1708,11 @@ "x-code-samples": [ { "lang": "Python", - "source": "# To use the Python SDK, install the package:\n# pip install mem0ai\n\nfrom mem0 import MemoryClient\nclient = MemoryClient(api_key=\"your_api_key\", org_id=\"your_org_id\", project_id=\"your_project_id\")\n\nquery = \"What do you know about me?\"\nfilters = {\n \"OR\":[\n {\n \"user_id\":\"alex\"\n },\n {\n \"agent_id\":{\n \"in\":[\n \"travel-assistant\",\n \"customer-support\"\n ]\n }\n }\n ]\n}\nclient.search(query, version=\"v2\", filters=filters)" + "source": "# To use the Python SDK, install the package:\n# pip install mem0ai\n\nfrom mem0 import MemoryClient\nclient = MemoryClient(api_key=\"your_api_key\")\n\nquery = \"What do you know about me?\"\nfilters = {\n \"OR\":[\n {\n \"user_id\":\"alex\"\n },\n {\n \"agent_id\":{\n \"in\":[\n \"travel-assistant\",\n \"customer-support\"\n ]\n }\n }\n ]\n}\nclient.search(query, filters=filters)" }, { "lang": "JavaScript", - "source": "// To use the JavaScript SDK, install the package:\n// npm i mem0ai\n\nimport MemoryClient from 'mem0ai';\nconst client = new MemoryClient({ apiKey: \"your-api-key\" });\n\nconst query = \"What do you know about me?\";\nconst filters = {\n OR: [\n { user_id: \"alex\" },\n { agent_id: { in: [\"travel-assistant\", \"customer-support\"] } }\n ]\n};\n\nclient.search(query, { api_version: \"v2\", filters })\n .then(result => console.log(result))\n .catch(error => console.error(error));" + "source": "// To use the JavaScript SDK, install the package:\n// npm i mem0ai\n\nimport MemoryClient from 'mem0ai';\nconst client = new MemoryClient({ apiKey: \"your-api-key\" });\n\nconst query = \"What do you know about me?\";\nconst filters = {\n OR: [\n { user_id: \"alex\" },\n { agent_id: { in: [\"travel-assistant\", \"customer-support\"] } }\n ]\n};\n\nclient.search(query, { filters })\n .then(result => console.log(result))\n .catch(error => console.error(error));" }, { "lang": "cURL", @@ -1864,7 +1864,7 @@ "x-code-samples": [ { "lang": "Python", - "source": "# To use the Python SDK, install the package:\n# pip install mem0ai\n\nfrom mem0 import MemoryClient\nclient = MemoryClient(api_key=\"your_api_key\", org_id=\"your_org_id\", project_id=\"your_project_id\")\n\nmemory = client.get(memory_id=\"\")" + "source": "# To use the Python SDK, install the package:\n# pip install mem0ai\n\nfrom mem0 import MemoryClient\nclient = MemoryClient(api_key=\"your_api_key\")\n\nmemory = client.get(memory_id=\"\")" }, { "lang": "JavaScript", @@ -1989,7 +1989,7 @@ "x-code-samples": [ { "lang": "Python", - "source": "# To use the Python SDK, install the package:\n# pip install mem0ai\n\nfrom mem0 import MemoryClient\nclient = MemoryClient(api_key=\"your_api_key\", org_id=\"your_org_id\", project_id=\"your_project_id\")\n\n# Update a memory\nmemory_id = \"\"\nclient.update(\n memory_id=memory_id,\n text=\"Your updated memory message here\",\n metadata={\"category\": \"example\"}\n)" + "source": "# To use the Python SDK, install the package:\n# pip install mem0ai\n\nfrom mem0 import MemoryClient\nclient = MemoryClient(api_key=\"your_api_key\")\n\n# Update a memory\nmemory_id = \"\"\nclient.update(\n memory_id=memory_id,\n text=\"Your updated memory message here\",\n metadata={\"category\": \"example\"}\n)" }, { "lang": "JavaScript", @@ -2053,7 +2053,7 @@ "x-code-samples": [ { "lang": "Python", - "source": "# To use the Python SDK, install the package:\n# pip install mem0ai\n\nfrom mem0 import MemoryClient\nclient = MemoryClient(api_key=\"your_api_key\", org_id=\"your_org_id\", project_id=\"your_project_id\")\n\nmemory_id = \"\"\nclient.delete(memory_id=memory_id)" + "source": "# To use the Python SDK, install the package:\n# pip install mem0ai\n\nfrom mem0 import MemoryClient\nclient = MemoryClient(api_key=\"your_api_key\")\n\nmemory_id = \"\"\nclient.delete(memory_id=memory_id)" }, { "lang": "JavaScript", @@ -2199,7 +2199,7 @@ "x-code-samples": [ { "lang": "Python", - "source": "# To use the Python SDK, install the package:\n# pip install mem0ai\n\nfrom mem0 import MemoryClient\nclient = MemoryClient(api_key=\"your_api_key\", org_id=\"your_org_id\", project_id=\"your_project_id\")\n\n# Add some message to create history\nmessages = [{\"role\": \"user\", \"content\": \"\"}]\nclient.add(messages, user_id=\"\")\n\n# Add second message to update history\nmessages.append({\"role\": \"user\", \"content\": \"\"})\nclient.add(messages, user_id=\"\")\n\n# Get history of how memory changed over time\nmemory_id = \"\"\nhistory = client.history(memory_id)" + "source": "# To use the Python SDK, install the package:\n# pip install mem0ai\n\nfrom mem0 import MemoryClient\nclient = MemoryClient(api_key=\"your_api_key\")\n\n# Add some message to create history\nmessages = [{\"role\": \"user\", \"content\": \"\"}]\nclient.add(messages, user_id=\"\")\n\n# Add second message to update history\nmessages.append({\"role\": \"user\", \"content\": \"\"})\nclient.add(messages, user_id=\"\")\n\n# Get history of how memory changed over time\nmemory_id = \"\"\nhistory = client.history(memory_id)" }, { "lang": "JavaScript", @@ -3608,7 +3608,7 @@ "x-code-samples": [ { "lang": "Python", - "source": "# To use the Python SDK, install the package:\n# pip install mem0ai\n\nfrom mem0 import MemoryClient\n\nclient = MemoryClient(api_key=\"your_api_key\", org_id=\"your_org_id\", project_id=\"your_project_id\")\n\nresponse = client.get_project()\nprint(response)" + "source": "# To use the Python SDK, install the package:\n# pip install mem0ai\n\nfrom mem0 import MemoryClient\n\nclient = MemoryClient(api_key=\"your_api_key\")\n\nresponse = client.get_project()\nprint(response)" }, { "lang": "JavaScript", @@ -4411,7 +4411,7 @@ "x-code-samples": [ { "lang": "Python", - "source": "# To use the Python SDK, install the package:\n# pip install mem0ai\n\nfrom mem0 import MemoryClient\nclient = MemoryClient(api_key=\"your_api_key\", org_id=\"your_org_id\", project_id=\"your_project_id\")\n\nupdate_memories = [\n {\n \"memory_id\": \"285ed74b-6e05-4043-b16b-3abd5b533496\",\n \"text\": \"Watches football\"\n },\n {\n \"memory_id\": \"2c9bd859-d1b7-4d33-a6b8-94e0147c4f07\",\n \"text\": \"Likes to travel\"\n }\n]\n\nresponse = client.batch_update(update_memories)\nprint(response)" + "source": "# To use the Python SDK, install the package:\n# pip install mem0ai\n\nfrom mem0 import MemoryClient\nclient = MemoryClient(api_key=\"your_api_key\")\n\nupdate_memories = [\n {\n \"memory_id\": \"285ed74b-6e05-4043-b16b-3abd5b533496\",\n \"text\": \"Watches football\"\n },\n {\n \"memory_id\": \"2c9bd859-d1b7-4d33-a6b8-94e0147c4f07\",\n \"text\": \"Likes to travel\"\n }\n]\n\nresponse = client.batch_update(update_memories)\nprint(response)" }, { "lang": "JavaScript", @@ -4490,7 +4490,7 @@ "x-code-samples": [ { "lang": "Python", - "source": "# To use the Python SDK, install the package:\n# pip install mem0ai\n\nfrom mem0 import MemoryClient\nclient = MemoryClient(api_key=\"your_api_key\", org_id=\"your_org_id\", project_id=\"your_project_id\")\n\ndelete_memories = [\n {\"memory_id\": \"285ed74b-6e05-4043-b16b-3abd5b533496\"},\n {\"memory_id\": \"2c9bd859-d1b7-4d33-a6b8-94e0147c4f07\"}\n]\n\nresponse = client.batch_delete(delete_memories)\nprint(response)" + "source": "# To use the Python SDK, install the package:\n# pip install mem0ai\n\nfrom mem0 import MemoryClient\nclient = MemoryClient(api_key=\"your_api_key\")\n\ndelete_memories = [\n {\"memory_id\": \"285ed74b-6e05-4043-b16b-3abd5b533496\"},\n {\"memory_id\": \"2c9bd859-d1b7-4d33-a6b8-94e0147c4f07\"}\n]\n\nresponse = client.batch_delete(delete_memories)\nprint(response)" }, { "lang": "JavaScript", @@ -4758,7 +4758,7 @@ "x-code-samples": [ { "lang": "Python", - "source": "# To use the Python SDK, install the package:\n# pip install mem0ai\n\nfrom mem0 import MemoryClient\nclient = MemoryClient(api_key=\"your_api_key\", org_id=\"your_org_id\", project_id=\"your_project_id\")\n\n# Create a webhook\nwebhook = client.create_webhook(\n url=\"https://your-webhook-url.com\",\n name=\"My Webhook\",\n project_id=\"your_project_id\",\n event_types=[\"memory:add\", \"memory:categorize\"]\n)\nprint(webhook)" + "source": "# To use the Python SDK, install the package:\n# pip install mem0ai\n\nfrom mem0 import MemoryClient\nclient = MemoryClient(api_key=\"your_api_key\")\n\n# Create a webhook\nwebhook = client.create_webhook(\n url=\"https://your-webhook-url.com\",\n name=\"My Webhook\",\n project_id=\"your_project_id\",\n event_types=[\"memory:add\", \"memory:categorize\"]\n)\nprint(webhook)" }, { "lang": "JavaScript", @@ -4998,7 +4998,7 @@ "x-code-samples": [ { "lang": "Python", - "source": "# To use the Python SDK, install the package:\n# pip install mem0ai\n\nfrom mem0 import MemoryClient\nclient = MemoryClient(api_key=\"your_api_key\", org_id=\"your_org_id\", project_id=\"your_project_id\")\n\n# Delete a webhook\nresponse = client.delete_webhook(webhook_id=\"your_webhook_id\")\nprint(response)" + "source": "# To use the Python SDK, install the package:\n# pip install mem0ai\n\nfrom mem0 import MemoryClient\nclient = MemoryClient(api_key=\"your_api_key\")\n\n# Delete a webhook\nresponse = client.delete_webhook(webhook_id=\"your_webhook_id\")\nprint(response)" }, { "lang": "JavaScript", @@ -5742,4 +5742,4 @@ } }, "x-original-swagger-version": "2.0" -} \ No newline at end of file +} diff --git a/docs/platform/features/advanced-retrieval.mdx b/docs/platform/features/advanced-retrieval.mdx index 4b0c21490..0dce8efbf 100644 --- a/docs/platform/features/advanced-retrieval.mdx +++ b/docs/platform/features/advanced-retrieval.mdx @@ -1,6 +1,6 @@ --- title: Advanced Retrieval -description: "Advanced memory search with keyword expansion, intelligent reranking, and precision filtering" +description: "Advanced memory search with intelligent reranking for precise results" --- ## What is Advanced Retrieval? @@ -9,40 +9,6 @@ Advanced Retrieval gives you precise control over how memories are found and ran ## Search Enhancement Options -### Keyword Search - -Expands results to include memories with specific terms, names, and technical keywords. - - - -- Searching for specific entities, names, or technical terms -- Need comprehensive coverage of a topic -- Want broader recall even if some results are less relevant -- Working with domain-specific terminology - - -```python Python -# Find memories containing specific food-related terms -results = client.search( - query="What foods should I avoid?", - keyword_search=True, - user_id="user123" -) - -# Results might include: -# ✓ "Allergic to peanuts and shellfish" -# ✓ "Lactose intolerant - avoid dairy" -# ✓ "Mentioned avoiding gluten last week" -``` - - -- **Latency**: ~10ms additional -- **Recall**: Significantly increased -- **Precision**: Slightly decreased -- **Best for**: Entity search, comprehensive coverage - - - ### Reranking Reorders results using deep semantic understanding to put the most relevant memories first. @@ -77,41 +43,6 @@ results = client.search( -### Memory Filtering - -Filters results to keep only the most precisely relevant memories. - - - -- Need highly specific, focused results -- Working with large datasets where noise is problematic -- Quality over quantity is essential -- Building production or safety-critical applications - - -```python Python -# Get only the most relevant dietary restrictions -results = client.search( - query="What are my dietary restrictions?", - filter_memories=True, - user_id="user123" -) - -# Before filtering: After filtering: -# • "Allergic to nuts" → • "Allergic to nuts" -# • "Likes Italian food" → • "Vegetarian diet" -# • "Vegetarian diet" → -# • "Eats dinner at 7pm" → -``` - - -- **Latency**: 200-300ms additional -- **Precision**: Maximized -- **Recall**: May be reduced -- **Best for**: Focused queries, production systems - - - ## Real-World Use Cases @@ -120,7 +51,6 @@ results = client.search( # Smart home assistant finding device preferences results = client.search( query="How do I like my bedroom temperature?", - keyword_search=True, # Find specific temperature mentions rerank=True, # Get most recent preferences first user_id="user123" ) @@ -133,8 +63,6 @@ results = client.search( # Find specific product issues with high precision results = client.search( query="Problems with premium subscription billing", - keyword_search=True, # Find "premium", "billing", "subscription" - filter_memories=True, # Only billing-related issues user_id="customer456" ) @@ -147,11 +75,10 @@ results = client.search( results = client.search( query="Patient allergies and contraindications", rerank=True, # Most important info first - filter_memories=True, # Only medical restrictions user_id="patient789" ) -# Ensures critical allergy info appears first and filters out non-medical data +# Ensures critical allergy info appears first ``` @@ -159,7 +86,6 @@ results = client.search( # Find learning progress for specific topics results = client.search( query="Python programming progress and difficulties", - keyword_search=True, # Find "Python", "programming", specific concepts rerank=True, # Recent progress first user_id="student123" ) @@ -169,63 +95,57 @@ results = client.search( -## Choosing the Right Combination +## Choosing the Right Configuration ### Recommended Configurations ```python Python -# Fast and broad - good for exploration +# Basic search - good for exploration def quick_search(query, user_id): return client.search( query=query, - keyword_search=True, user_id=user_id ) -# Balanced - good for most applications +# Reranked search - good for most applications def standard_search(query, user_id): return client.search( query=query, - keyword_search=True, rerank=True, user_id=user_id ) -# High precision - good for critical applications +# Reranked search - good for critical applications def precise_search(query, user_id): return client.search( query=query, rerank=True, - filter_memories=True, user_id=user_id ) ``` ```javascript JavaScript -// Fast and broad - good for exploration +// Basic search - good for exploration function quickSearch(query, userId) { return client.search(query, { - user_id: userId, - keyword_search: true + user_id: userId }); } -// Balanced - good for most applications +// Reranked search - good for most applications function standardSearch(query, userId) { return client.search(query, { user_id: userId, - keyword_search: true, rerank: true }); } -// High precision - good for critical applications +// Reranked search - good for critical applications function preciseSearch(query, userId) { return client.search(query, { user_id: userId, - rerank: true, - filter_memories: true + rerank: true }); } ``` @@ -235,19 +155,15 @@ function preciseSearch(query, userId) { ### Do -- Start simple with just one enhancement and measure impact -- Use keyword search for entity-heavy queries (names, places, technical terms) +- Start simple with basic search and measure impact before enabling reranking - Use reranking when the top result quality matters most -- Use filtering for production systems where precision is critical -- Handle empty results gracefully when filtering is too aggressive - Monitor latency and adjust based on your application's needs +- Handle empty results gracefully ### Don't -- Enable all options by default without measuring necessity -- Use filtering for broad exploratory queries +- Enable reranking by default without measuring necessity - Ignore latency impact in real-time applications -- Forget to handle cases where filtering returns no results - Use advanced retrieval for simple, fast lookup scenarios ## Performance Guidelines @@ -261,20 +177,18 @@ import time start_time = time.time() results = client.search( query="user preferences", - keyword_search=True, # +10ms rerank=True, # +150ms - filter_memories=True, # +250ms user_id="user123" ) latency = time.time() - start_time -print(f"Search completed in {latency:.2f}s") # ~0.41s expected +print(f"Search completed in {latency:.2f}s") ``` ### Optimization Tips 1. **Cache frequent queries** to avoid repeated advanced processing 2. **Use session-specific search** with `run_id` to reduce search space -3. **Implement fallback logic** when filtering returns empty results +3. **Implement fallback logic** when search returns empty results 4. **Monitor and alert** on search latency patterns diff --git a/docs/platform/features/async-mode-default-change.mdx b/docs/platform/features/async-mode-default-change.mdx deleted file mode 100644 index 7fb39017d..000000000 --- a/docs/platform/features/async-mode-default-change.mdx +++ /dev/null @@ -1,197 +0,0 @@ ---- -title: Async Mode Default Change -description: "The async_mode parameter now defaults to true for all memory additions, changing from synchronous processing." ---- - - - **Important Change** - - The `async_mode` parameter defaults to `true` for all memory additions, changing the default API behavior to asynchronous processing. - - -## Overview - -The Memory Addition API processes all memory additions asynchronously by default. This change improves performance and scalability by queuing memory operations in the background, allowing your application to continue without waiting for memory processing to complete. - -## What's Changing - -The parameter `async_mode` will default to `true` instead of `false`. - -This means memory additions will be **processed asynchronously** by default - queued for background execution instead of waiting for processing to complete. - -## Behavior Comparison - -### Old Default Behavior (async_mode = false) - -When `async_mode` was set to `false`, the API returned fully processed memory objects immediately: - -```json -{ - "results": [ - { - "id": "de0ee948-af6a-436c-835c-efb6705207de", - "event": "ADD", - "memory": "User Order #1234 was for a 'Nova 2000'", - "structured_attributes": { - "day": 13, - "hour": 16, - "year": 2025, - "month": 10, - "minute": 59, - "quarter": 4, - "is_weekend": false, - "day_of_week": "monday", - "day_of_year": 286, - "week_of_year": 42 - } - } - ] -} -``` - -### New Default Behavior (async_mode = true) - -With `async_mode` defaulting to `true`, memory processing is queued in the background and the API returns immediately: - -```json -{ - "results": [ - { - "message": "Memory processing has been queued for background execution", - "status": "PENDING", - "event_id": "d7b5282a-0031-4cc2-98ba-5a02d8531e17" - } - ] -} -``` - -## Migration Guide - -### If You Need Synchronous Processing - -If your integration relies on receiving the processed memory object immediately, you can explicitly set `async_mode` to `false` in your requests: - - - -```python Python -from mem0 import MemoryClient - -client = MemoryClient(api_key="your-api-key") - -# Explicitly set async_mode=False to preserve synchronous behavior -messages = [ - {"role": "user", "content": "I ordered a Nova 2000"} -] - -result = client.add( - messages, - user_id="user-123", - async_mode=False # This ensures synchronous processing -) -``` - -```javascript JavaScript -const { MemoryClient } = require('mem0ai'); - -const client = new MemoryClient({ apiKey: 'your-api-key' }); - -// Explicitly set async_mode: false to preserve synchronous behavior -const messages = [ - { role: "user", content: "I ordered a Nova 2000" } -]; - -const result = await client.add(messages, { - user_id: "user-123", - async_mode: false // This ensures synchronous processing -}); -``` - -```bash cURL -curl -X POST https://api.mem0.ai/v1/memories/ \ - -H "Authorization: Token your-api-key" \ - -H "Content-Type: application/json" \ - -d '{ - "messages": [ - {"role": "user", "content": "I ordered a Nova 2000"} - ], - "user_id": "user-123", - "async_mode": false - }' -``` - - - -### If You Want to Adopt Asynchronous Processing - -If you want to benefit from the improved performance of asynchronous processing: - -1. **Remove** any explicit `async_mode=False` parameters from your code -2. **Use webhooks** to receive notifications when memory processing completes - - -Learn more about [Webhooks](/platform/features/webhooks) for real-time notifications about memory events. - - -## Benefits of Asynchronous Processing - -Switching to asynchronous processing provides several advantages: - -- **Faster API Response Times**: Your application doesn't wait for memory processing -- **Better Scalability**: Handle more memory additions concurrently -- **Improved User Experience**: Reduced latency in your application -- **Resource Efficiency**: Background processing optimizes server resources - -## Important Notes - -- The default behavior is now `async_mode=true` for asynchronous processing -- Explicitly set `async_mode=false` if you need synchronous behavior -- Use webhooks to receive notifications when memories are processed - -## Monitoring Memory Processing - -When using asynchronous mode, use webhooks to receive notifications about memory events: - - - Learn how to set up webhooks for memory processing events - - -You can also retrieve all processed memories at any time: - - - -```python Python -# Retrieve all memories for a user -# Note: get_all now requires filters -memories = client.get_all(filters={"AND": [{"user_id": "user-123"}]}) -``` - -```javascript JavaScript -// Retrieve all memories for a user -// Note: getAll now requires filters -const memories = await client.getAll({ filters: {"AND": [{"user_id": "user-123"}]} }); -``` - - - -## Need Help? - -If you have questions about this change or need assistance updating your integration: - - - -## Related Documentation - - - - Learn about the asynchronous client for Mem0 - - - View the complete API reference for adding memories - - - Configure webhooks for memory processing events - - - Understand memory addition operations - - diff --git a/docs/platform/features/criteria-retrieval.mdx b/docs/platform/features/criteria-retrieval.mdx index aabc27c76..d0ea0a3fb 100644 --- a/docs/platform/features/criteria-retrieval.mdx +++ b/docs/platform/features/criteria-retrieval.mdx @@ -38,11 +38,7 @@ Before defining any criteria, make sure to initialize the `MemoryClient` with yo ```python from mem0 import MemoryClient -client = MemoryClient( - api_key="your_mem0_api_key", - org_id="your_organization_id", - project_id="your_project_id" -) +client = MemoryClient(api_key="your_mem0_api_key") ``` ### Define Your Criteria diff --git a/docs/platform/features/custom-categories.mdx b/docs/platform/features/custom-categories.mdx index 1b6a9310b..1c7446fd9 100644 --- a/docs/platform/features/custom-categories.mdx +++ b/docs/platform/features/custom-categories.mdx @@ -98,7 +98,7 @@ messages = [ ] # Add memories with project-level custom categories -client.add(messages, user_id="alice", async_mode=False) +client.add(messages, user_id="alice") ``` @@ -187,7 +187,7 @@ messages = [ ] # Add memories with default categories -client.add(messages, user_id='alice', async_mode=False) +client.add(messages, user_id='alice') ``` ```python Memories with categories diff --git a/docs/platform/features/expiration-date.mdx b/docs/platform/features/expiration-date.mdx deleted file mode 100644 index 86a1b3ac7..000000000 --- a/docs/platform/features/expiration-date.mdx +++ /dev/null @@ -1,109 +0,0 @@ ---- -title: Expiration Date -description: 'Set time-bound memories in Mem0 with automatic expiration dates to manage temporal information effectively.' ---- - -## Benefits of Memory Expiration - -Setting expiration dates for memories offers several advantages: - -- **Time-Sensitive Information Management**: Handle information that is only relevant for a specific time period. -- **Event-Based Memory**: Manage information related to upcoming events that becomes irrelevant after the event passes. - -These benefits enable more sophisticated memory management for applications where temporal context matters. - -## Setting Memory Expiration Date - -You can set an expiration date for memories, after which they will no longer be retrieved in searches. This is useful for creating temporary memories or memories that are relevant only for a specific time period. - - - -```python Python -import datetime -from mem0 import MemoryClient - -client = MemoryClient(api_key="your-api-key") - -messages = [ - { - "role": "user", - "content": "I'll be in San Francisco until the end of this month." - } -] - -# Set an expiration date for this memory -client.add(messages=messages, user_id="alex", expiration_date=str(datetime.datetime.now().date() + datetime.timedelta(days=30))) - -# You can also use an explicit date string -client.add(messages=messages, user_id="alex", expiration_date="2023-08-31") -``` - -```javascript JavaScript -import MemoryClient from 'mem0ai'; -const client = new MemoryClient({ apiKey: 'your-api-key' }); - -const messages = [ - { - "role": "user", - "content": "I'll be in San Francisco until the end of this month." - } -]; - -// Set an expiration date 30 days from now -const expirationDate = new Date(); -expirationDate.setDate(expirationDate.getDate() + 30); -client.add(messages, { - user_id: "alex", - expiration_date: expirationDate.toISOString().split('T')[0] -}) - .then(response => console.log(response)) - .catch(error => console.error(error)); - -// You can also use an explicit date string -client.add(messages, { - user_id: "alex", - expiration_date: "2023-08-31" -}) - .then(response => console.log(response)) - .catch(error => console.error(error)); -``` - -```bash cURL -curl -X POST "https://api.mem0.ai/v1/memories/" \ - -H "Authorization: Token your-api-key" \ - -H "Content-Type: application/json" \ - -d '{ - "messages": [ - { - "role": "user", - "content": "I'll be in San Francisco until the end of this month." - } - ], - "user_id": "alex", - "expiration_date": "2023-08-31" - }' -``` - -```json Output -{ - "results": [ - { - "id": "a1b2c3d4-e5f6-4g7h-8i9j-k0l1m2n3o4p5", - "data": { - "memory": "In San Francisco until the end of this month" - }, - "event": "ADD" - } - ] -} -``` - - - - -Once a memory reaches its expiration date, it will not be included in search or get results, though the data remains stored in the system. - - -If you have any questions, please feel free to reach out to us using one of the following methods: - - diff --git a/docs/platform/features/graph-memory.mdx b/docs/platform/features/graph-memory.mdx index 4a7e4dcc6..8bb947667 100644 --- a/docs/platform/features/graph-memory.mdx +++ b/docs/platform/features/graph-memory.mdx @@ -30,11 +30,7 @@ When adding new memories, enable Graph Memory to automatically build relationshi ```python Python from mem0 import MemoryClient -client = MemoryClient( - api_key="your-api-key", - org_id="your-org-id", - project_id="your-project-id" -) +client = MemoryClient(api_key="your-api-key") messages = [ {"role": "user", "content": "My name is Joseph"}, @@ -53,11 +49,7 @@ client.add( ```javascript JavaScript import { MemoryClient } from "mem0"; -const client = new MemoryClient({ - apiKey: "your-api-key", - org_id: "your-org-id", - project_id: "your-project-id" -}); +const client = new MemoryClient({ apiKey: "your-api-key" }); const messages = [ { role: "user", content: "My name is Joseph" }, @@ -284,11 +276,7 @@ Instead of passing `enable_graph=True` to every add call, you can enable it once ```python Python from mem0 import MemoryClient -client = MemoryClient( - api_key="your-api-key", - org_id="your-org-id", - project_id="your-project-id" -) +client = MemoryClient(api_key="your-api-key") # Enable graph memory for all operations in this project client.project.update(enable_graph=True) @@ -309,11 +297,7 @@ client.add( ```javascript JavaScript import { MemoryClient } from "mem0"; -const client = new MemoryClient({ - apiKey: "your-api-key", - org_id: "your-org-id", - project_id: "your-project-id" -}); +const client = new MemoryClient({ apiKey: "your-api-key" }); // Enable graph memory for all operations in this project await client.project.update({ enable_graph: true }); diff --git a/docs/platform/features/group-chat.mdx b/docs/platform/features/group-chat.mdx index 9b2c051cd..f0f964336 100644 --- a/docs/platform/features/group-chat.mdx +++ b/docs/platform/features/group-chat.mdx @@ -217,17 +217,16 @@ print(search_response) ## Async Mode Support -Group chat also supports async processing for improved performance: +Group chat supports async processing for improved performance. Memory additions are processed asynchronously by default. ```python Python -# Group chat with async mode +# Group chat — async processing is the default response = client.add( messages, run_id="groupchat_async", infer=True, - async_mode=True ) print(response) ``` @@ -268,7 +267,7 @@ Each message in a group chat must include: 4. **Memory Filtering**: Use filters to retrieve memories from specific participants or sessions when needed. -5. **Async Processing**: Use `async_mode=True` for large group conversations to improve performance. +5. **Async Processing**: Memory additions are processed asynchronously by default, which is ideal for large group conversations. 6. **Search Context**: Leverage the search functionality to find specific information within group chat contexts. diff --git a/mem0-ts/src/client/index.ts b/mem0-ts/src/client/index.ts index b6d9be766..b925909e2 100644 --- a/mem0-ts/src/client/index.ts +++ b/mem0-ts/src/client/index.ts @@ -3,14 +3,17 @@ import type * as MemoryTypes from "./mem0.types"; // Re-export all types from mem0.types export type { - MemoryOptions, + EntityOptions, + AddMemoryOptions, + SearchMemoryOptions, + GetAllMemoryOptions, + DeleteAllMemoryOptions, ProjectOptions, Memory, MemoryHistory, MemoryUpdateBody, ProjectResponse, PromptUpdatePayload, - SearchOptions, Webhook, WebhookCreatePayload, WebhookUpdatePayload, @@ -19,6 +22,8 @@ export type { AllUsers, User, FeedbackPayload, + CreateMemoryExportPayload, + GetMemoryExportPayload, } from "./mem0.types"; // Re-export enums as values (not type-only) diff --git a/mem0-ts/src/client/mem0.ts b/mem0-ts/src/client/mem0.ts index 1c48f6e23..51c658611 100644 --- a/mem0-ts/src/client/mem0.ts +++ b/mem0-ts/src/client/mem0.ts @@ -4,11 +4,13 @@ import { ProjectOptions, Memory, MemoryHistory, - MemoryOptions, + AddMemoryOptions, + SearchMemoryOptions, + GetAllMemoryOptions, + DeleteAllMemoryOptions, MemoryUpdateBody, ProjectResponse, PromptUpdatePayload, - SearchOptions, Webhook, WebhookCreatePayload, WebhookUpdatePayload, @@ -30,19 +32,13 @@ class APIError extends Error { interface ClientOptions { apiKey: string; host?: string; - organizationName?: string; - projectName?: string; - organizationId?: string; - projectId?: string; } export default class MemoryClient { apiKey: string; host: string; - organizationName: string | null; - projectName: string | null; - organizationId: string | number | null; - projectId: string | number | null; + private organizationId: string | number | null; + private projectId: string | number | null; headers: Record; client: any; telemetryId: string; @@ -59,35 +55,11 @@ export default class MemoryClient { } } - _validateOrgProject(): void { - // Check for organizationName/projectName pair - if ( - (this.organizationName === null && this.projectName !== null) || - (this.organizationName !== null && this.projectName === null) - ) { - console.warn( - "Warning: Both organizationName and projectName must be provided together when using either. This will be removed from version 1.0.40. Note that organizationName/projectName are being deprecated in favor of organizationId/projectId.", - ); - } - - // Check for organizationId/projectId pair - if ( - (this.organizationId === null && this.projectId !== null) || - (this.organizationId !== null && this.projectId === null) - ) { - console.warn( - "Warning: Both organizationId and projectId must be provided together when using either. This will be removed from version 1.0.40.", - ); - } - } - constructor(options: ClientOptions) { this.apiKey = options.apiKey; this.host = options.host || "https://api.mem0.ai"; - this.organizationName = options.organizationName || null; - this.projectName = options.projectName || null; - this.organizationId = options.organizationId || null; - this.projectId = options.projectId || null; + this.organizationId = null; + this.projectId = null; this.headers = { Authorization: `Token ${this.apiKey}`, @@ -101,28 +73,19 @@ export default class MemoryClient { }); this._validateApiKey(); - - // Initialize with a temporary ID that will be updated this.telemetryId = ""; - - // Initialize the client this._initializeClient(); } private async _initializeClient() { try { - // Generate telemetry ID await this.ping(); if (!this.telemetryId) { this.telemetryId = generateHash(this.apiKey); } - this._validateOrgProject(); - - // Capture initialization event captureClientEvent("init", this, { - api_version: "v1", client_type: "MemoryClient", }).catch((error: any) => { console.error("Failed to capture event:", error); @@ -163,13 +126,16 @@ export default class MemoryClient { return jsonResponse; } - _preparePayload(messages: Array, options: MemoryOptions): object { + _preparePayload( + messages: Array, + options: Record, + ): object { const payload: any = {}; payload.messages = messages; return { ...payload, ...options }; } - _prepareParams(options: MemoryOptions): object { + _prepareParams(options: Record): object { return Object.fromEntries( Object.entries(options).filter(([_, v]) => v != null), ); @@ -197,9 +163,8 @@ export default class MemoryClient { const { org_id, project_id, user_email } = response; - // Only update if values are actually present - if (org_id && !this.organizationId) this.organizationId = org_id; - if (project_id && !this.projectId) this.projectId = project_id; + if (org_id) this.organizationId = org_id; + if (project_id) this.projectId = project_id; if (user_email) this.telemetryId = user_email; } catch (error: any) { // Pass through structured exceptions and APIError @@ -215,30 +180,11 @@ export default class MemoryClient { async add( messages: Array, - options: MemoryOptions & Record = {}, + options: AddMemoryOptions & Record = {}, ): Promise> { if (this.telemetryId === "") await this.ping(); - this._validateOrgProject(); - if (this.organizationName != null && this.projectName != null) { - options.org_name = this.organizationName; - options.project_name = this.projectName; - } - - if (this.organizationId != null && this.projectId != null) { - options.org_id = this.organizationId; - options.project_id = this.projectId; - - if (options.org_name) delete options.org_name; - if (options.project_name) delete options.project_name; - } - - if (options.api_version) { - options.version = options.api_version.toString() || "v2"; - } const payload = this._preparePayload(messages, options); - - // get payload keys whose value is not null or undefined const payloadKeys = Object.keys(payload); this._captureEvent("add", [payloadKeys]); @@ -276,7 +222,6 @@ export default class MemoryClient { } if (this.telemetryId === "") await this.ping(); - this._validateOrgProject(); const payload: Record = {}; if (text !== undefined) payload.text = text; if (metadata !== undefined) payload.metadata = metadata; @@ -307,87 +252,53 @@ export default class MemoryClient { ); } - async getAll(options?: SearchOptions): Promise> { + async getAll(options?: GetAllMemoryOptions): Promise> { if (this.telemetryId === "") await this.ping(); - this._validateOrgProject(); const payloadKeys = Object.keys(options || {}); this._captureEvent("get_all", [payloadKeys]); - const { api_version, page, page_size, ...otherOptions } = options ?? {}; - if (this.organizationName != null && this.projectName != null) { - otherOptions.org_name = this.organizationName; - otherOptions.project_name = this.projectName; - } - - let appendedParams = ""; - let paginated_response = false; + const { page, page_size, ...rest } = options ?? {}; + const body: Record = { + output_format: "v1.1", + ...rest, + }; + let url = `${this.host}/v2/memories/`; if (page && page_size) { - appendedParams += `page=${page}&page_size=${page_size}`; - paginated_response = true; + url += `?page=${page}&page_size=${page_size}`; } - if (this.organizationId != null && this.projectId != null) { - otherOptions.org_id = this.organizationId; - otherOptions.project_id = this.projectId; - - if (otherOptions.org_name) delete otherOptions.org_name; - if (otherOptions.project_name) delete otherOptions.project_name; - } - - if (api_version === "v2") { - let url = paginated_response - ? `${this.host}/v2/memories/?${appendedParams}` - : `${this.host}/v2/memories/`; - return this._fetchWithErrorHandling(url, { - method: "POST", - headers: this.headers, - body: JSON.stringify(otherOptions), - }); - } else { - // @ts-ignore - const params = new URLSearchParams(this._prepareParams(otherOptions)); - const url = paginated_response - ? `${this.host}/v1/memories/?${params}&${appendedParams}` - : `${this.host}/v1/memories/?${params}`; - return this._fetchWithErrorHandling(url, { - headers: this.headers, - }); - } + const response = await this._fetchWithErrorHandling(url, { + method: "POST", + headers: this.headers, + body: JSON.stringify(body), + }); + // Unwrap v1.1 format: { results: [...] } → [...] + return Array.isArray(response) ? response : (response?.results ?? response); } async search( query: string, - options?: SearchOptions & Record, + options?: SearchMemoryOptions, ): Promise> { if (this.telemetryId === "") await this.ping(); - this._validateOrgProject(); const payloadKeys = Object.keys(options || {}); this._captureEvent("search", [payloadKeys]); - const { api_version, ...otherOptions } = options ?? {}; - const payload = { query, ...otherOptions }; - if (this.organizationName != null && this.projectName != null) { - payload.org_name = this.organizationName; - payload.project_name = this.projectName; - } + const payload: Record = { + query, + output_format: "v1.1", + ...(options ?? {}), + }; - if (this.organizationId != null && this.projectId != null) { - payload.org_id = this.organizationId; - payload.project_id = this.projectId; - - if (payload.org_name) delete payload.org_name; - if (payload.project_name) delete payload.project_name; - } - const endpoint = - api_version === "v2" ? "/v2/memories/search/" : "/v1/memories/search/"; const response = await this._fetchWithErrorHandling( - `${this.host}${endpoint}`, + `${this.host}/v2/memories/search/`, { method: "POST", headers: this.headers, body: JSON.stringify(payload), }, ); - return response; + // Unwrap v1.1 format: { results: [...] } → [...] + return Array.isArray(response) ? response : (response?.results ?? response); } async delete(memoryId: string): Promise<{ message: string }> { @@ -402,23 +313,12 @@ export default class MemoryClient { ); } - async deleteAll(options: MemoryOptions = {}): Promise<{ message: string }> { + async deleteAll( + options: DeleteAllMemoryOptions = {}, + ): Promise<{ message: string }> { if (this.telemetryId === "") await this.ping(); - this._validateOrgProject(); const payloadKeys = Object.keys(options || {}); this._captureEvent("delete_all", [payloadKeys]); - if (this.organizationName != null && this.projectName != null) { - options.org_name = this.organizationName; - options.project_name = this.projectName; - } - - if (this.organizationId != null && this.projectId != null) { - options.org_id = this.organizationId; - options.project_id = this.projectId; - - if (options.org_name) delete options.org_name; - if (options.project_name) delete options.project_name; - } // @ts-ignore const params = new URLSearchParams(this._prepareParams(options)); const response = await this._fetchWithErrorHandling( @@ -443,31 +343,20 @@ export default class MemoryClient { return response; } - async users(): Promise { + async users(options?: { + page?: number; + page_size?: number; + }): Promise { if (this.telemetryId === "") await this.ping(); - this._validateOrgProject(); this._captureEvent("users", []); - const options: MemoryOptions = {}; - if (this.organizationName != null && this.projectName != null) { - options.org_name = this.organizationName; - options.project_name = this.projectName; - } - - if (this.organizationId != null && this.projectId != null) { - options.org_id = this.organizationId; - options.project_id = this.projectId; - - if (options.org_name) delete options.org_name; - if (options.project_name) delete options.project_name; - } - // @ts-ignore - const params = new URLSearchParams(options); - const response = await this._fetchWithErrorHandling( - `${this.host}/v1/entities/?${params}`, - { - headers: this.headers, - }, - ); + let url = `${this.host}/v1/entities/`; + const params: string[] = []; + if (options?.page) params.push(`page=${options.page}`); + if (options?.page_size) params.push(`page_size=${options.page_size}`); + if (params.length) url += `?${params.join("&")}`; + const response = await this._fetchWithErrorHandling(url, { + headers: this.headers, + }); return response; } @@ -502,7 +391,6 @@ export default class MemoryClient { } = {}, ): Promise<{ message: string }> { if (this.telemetryId === "") await this.ping(); - this._validateOrgProject(); let to_delete: Array<{ type: string; name: string }> = []; const { user_id, agent_id, app_id, run_id } = params; @@ -527,29 +415,9 @@ export default class MemoryClient { throw new Error("No entities to delete"); } - const requestOptions: MemoryOptions = {}; - if (this.organizationName != null && this.projectName != null) { - requestOptions.org_name = this.organizationName; - requestOptions.project_name = this.projectName; - } - - if (this.organizationId != null && this.projectId != null) { - requestOptions.org_id = this.organizationId; - requestOptions.project_id = this.projectId; - - if (requestOptions.org_name) delete requestOptions.org_name; - if (requestOptions.project_name) delete requestOptions.project_name; - } - - // Delete each entity and handle errors for (const entity of to_delete) { try { - await this.client.delete( - `/v2/entities/${entity.type}/${entity.name}/`, - { - params: requestOptions, - }, - ); + await this.client.delete(`/v2/entities/${entity.type}/${entity.name}/`); } catch (error: any) { throw new APIError( `Failed to delete ${entity.type} ${entity.name}: ${error.message}`, @@ -558,13 +426,7 @@ export default class MemoryClient { } this._captureEvent("delete_users", [ - { - user_id: user_id, - agent_id: agent_id, - app_id: app_id, - run_id: run_id, - sync_type: "sync", - }, + { user_id, agent_id, app_id, run_id, sync_type: "sync" }, ]); return { @@ -612,7 +474,6 @@ export default class MemoryClient { async getProject(options: ProjectOptions): Promise { if (this.telemetryId === "") await this.ping(); - this._validateOrgProject(); const payloadKeys = Object.keys(options || {}); this._captureEvent("get_project", [payloadKeys]); const { fields } = options; @@ -639,7 +500,6 @@ export default class MemoryClient { prompts: PromptUpdatePayload, ): Promise> { if (this.telemetryId === "") await this.ping(); - this._validateOrgProject(); this._captureEvent("update_project", []); if (!(this.organizationId && this.projectId)) { throw new Error( @@ -748,15 +608,10 @@ export default class MemoryClient { if (this.telemetryId === "") await this.ping(); this._captureEvent("create_memory_export", []); - // Return if missing filters or schema if (!data.filters || !data.schema) { throw new Error("Missing filters or schema"); } - // Add Org and Project ID - data.org_id = this.organizationId?.toString() || null; - data.project_id = this.projectId?.toString() || null; - const response = await this._fetchWithErrorHandling( `${this.host}/v1/exports/`, { @@ -779,9 +634,6 @@ export default class MemoryClient { throw new Error("Missing memory_export_id or filters"); } - data.org_id = this.organizationId?.toString() || ""; - data.project_id = this.projectId?.toString() || ""; - const response = await this._fetchWithErrorHandling( `${this.host}/v1/exports/get/`, { diff --git a/mem0-ts/src/client/mem0.types.ts b/mem0-ts/src/client/mem0.types.ts index 6db195eac..342bab674 100644 --- a/mem0-ts/src/client/mem0.types.ts +++ b/mem0-ts/src/client/mem0.types.ts @@ -1,59 +1,70 @@ -interface Common { - project_id?: string | null; - org_id?: string | null; -} - -export interface MemoryOptions { - api_version?: API_VERSION | string; - version?: API_VERSION | string; +// ─── Entity Options (for add/delete — top-level identity) ─── +export interface EntityOptions { user_id?: string; agent_id?: string; app_id?: string; run_id?: string; +} + +// ─── Per-Method Options ───────────────────────────────────── +export interface AddMemoryOptions extends EntityOptions { metadata?: Record; - filters?: Record; - org_name?: string | null; // Deprecated - project_name?: string | null; // Deprecated - org_id?: string | number | null; - project_id?: string | number | null; infer?: boolean; - page?: number; - page_size?: number; - includes?: string; - excludes?: string; - enable_graph?: boolean; - start_date?: string; - end_date?: string; custom_categories?: custom_categories[]; custom_instructions?: string; timestamp?: number; - output_format?: string | OutputFormat; - async_mode?: boolean; - filter_memories?: boolean; - immutable?: boolean; structured_data_schema?: Record; + enable_graph?: boolean; } +export interface SearchMemoryOptions { + filters?: Record; + metadata?: Record; + top_k?: number; + threshold?: number; + rerank?: boolean; + fields?: string[]; + categories?: string[]; + enable_graph?: boolean; +} + +export interface GetAllMemoryOptions { + filters?: Record; + page?: number; + page_size?: number; + start_date?: string; + end_date?: string; + categories?: string[]; + enable_graph?: boolean; +} + +export interface DeleteAllMemoryOptions extends EntityOptions {} + +// ─── Project Options ──────────────────────────────────────── export interface ProjectOptions { fields?: string[]; } -export enum OutputFormat { - V1 = "v1.0", - V1_1 = "v1.1", -} - -export enum API_VERSION { - V1 = "v1", - V2 = "v2", +export interface PromptUpdatePayload { + custom_instructions?: string; + custom_categories?: custom_categories[]; + retrieval_criteria?: any[]; + enable_graph?: boolean; + version?: string; + memory_depth?: string | null; + usecase_setting?: string | number; + multilingual?: boolean; + [key: string]: any; } +// ─── Enums ────────────────────────────────────────────────── export enum Feedback { POSITIVE = "POSITIVE", NEGATIVE = "NEGATIVE", VERY_NEGATIVE = "VERY_NEGATIVE", } +// ─── Message Types ────────────────────────────────────────── export interface MultiModalMessages { type: "image_url"; image_url: { @@ -68,30 +79,9 @@ export interface Messages { export interface Message extends Messages {} -export interface MemoryHistory { - id: string; - memory_id: string; - input: Array; - old_memory: string | null; - new_memory: string | null; - user_id: string; - categories: Array; - event: Event | string; - created_at: Date; - updated_at: Date; -} - -export interface SearchOptions extends MemoryOptions { - api_version?: API_VERSION | string; - limit?: number; - enable_graph?: boolean; - threshold?: number; - top_k?: number; - only_metadata_based_search?: boolean; - keyword_search?: boolean; - fields?: string[]; - categories?: string[]; - rerank?: boolean; +// ─── Response Types (reflect API shapes, unchanged) ───────── +export interface MemoryData { + memory: string; } enum Event { @@ -101,10 +91,6 @@ enum Event { NOOP = "NOOP", } -export interface MemoryData { - memory: string; -} - export interface Memory { id: string; messages?: Array; @@ -125,6 +111,19 @@ export interface Memory { run_id?: string | null; } +export interface MemoryHistory { + id: string; + memory_id: string; + input: Array; + old_memory: string | null; + new_memory: string | null; + user_id: string; + categories: Array; + event: Event | string; + created_at: Date; + updated_at: Date; +} + export interface MemoryUpdateBody { memoryId: string; text: string; @@ -157,20 +156,7 @@ interface custom_categories { [key: string]: any; } -export interface PromptUpdatePayload { - custom_instructions?: string; - custom_categories?: custom_categories[]; - retrieval_criteria?: any[]; - enable_graph?: boolean; - version?: string; - inclusion_prompt?: string; - exclusion_prompt?: string; - memory_depth?: string | null; - usecase_setting?: string | number; - multilingual?: boolean; - [key: string]: any; -} - +// ─── Webhook Types ────────────────────────────────────────── export enum WebhookEvent { MEMORY_ADDED = "memory_add", MEMORY_UPDATED = "memory_update", @@ -202,19 +188,20 @@ export interface WebhookUpdatePayload { eventTypes?: WebhookEvent[]; } +// ─── Feedback & Export Types ──────────────────────────────── export interface FeedbackPayload { memory_id: string; feedback?: Feedback | null; feedback_reason?: string | null; } -export interface CreateMemoryExportPayload extends Common { +export interface CreateMemoryExportPayload { schema: Record; filters: Record; export_instructions?: string; } -export interface GetMemoryExportPayload extends Common { +export interface GetMemoryExportPayload { filters?: Record; memory_export_id?: string; } diff --git a/mem0-ts/src/client/tests/integration/crud.test.ts b/mem0-ts/src/client/tests/integration/crud.test.ts index 5cd5421cf..2b6e2e771 100644 --- a/mem0-ts/src/client/tests/integration/crud.test.ts +++ b/mem0-ts/src/client/tests/integration/crud.test.ts @@ -122,7 +122,9 @@ describeIntegration("MemoryClient Integration — CRUD", () => { // ─── Get all ────────────────────────────────────────────── describe("get all memories", () => { test("returns all memories for test user", async () => { - const memories = await client.getAll({ user_id: TEST_USER_ID }); + const memories = await client.getAll({ + filters: { user_id: TEST_USER_ID }, + }); expect(Array.isArray(memories)).toBe(true); expect(memories.length).toBeGreaterThanOrEqual(memoryIds.length); @@ -135,7 +137,7 @@ describeIntegration("MemoryClient Integration — CRUD", () => { test("returns paginated results with page and page_size", async () => { const page1 = await client.getAll({ - user_id: TEST_USER_ID, + filters: { user_id: TEST_USER_ID }, page: 1, page_size: 1, }); @@ -200,7 +202,7 @@ describeIntegration("MemoryClient Integration — CRUD", () => { test("getAll for non-existent user returns empty array", async () => { const memories = await client.getAll({ - user_id: `nonexistent-user-${randomUUID()}`, + filters: { user_id: `nonexistent-user-${randomUUID()}` }, }); expect(Array.isArray(memories)).toBe(true); diff --git a/mem0-ts/src/client/tests/integration/helpers.ts b/mem0-ts/src/client/tests/integration/helpers.ts index df4378b33..3f0b8bfc1 100644 --- a/mem0-ts/src/client/tests/integration/helpers.ts +++ b/mem0-ts/src/client/tests/integration/helpers.ts @@ -63,7 +63,9 @@ export async function waitForMemories( maxRetries = 4, ): Promise { for (let attempt = 1; attempt <= maxRetries; attempt++) { - const memories = await withRetry(() => client.getAll({ user_id: userId })); + const memories = await withRetry(() => + client.getAll({ filters: { user_id: userId } }), + ); if (Array.isArray(memories) && memories.length >= minCount) { return memories; } diff --git a/mem0-ts/src/client/tests/integration/initialization.test.ts b/mem0-ts/src/client/tests/integration/initialization.test.ts index 6336d0c11..6472a14d6 100644 --- a/mem0-ts/src/client/tests/integration/initialization.test.ts +++ b/mem0-ts/src/client/tests/integration/initialization.test.ts @@ -31,10 +31,10 @@ describeIntegration("MemoryClient Integration — Initialization", () => { afterAll(() => cleanup()); - test("client pings successfully and resolves org/project", async () => { + test("client pings successfully", async () => { await client.ping(); - expect(client.organizationId).toBeTruthy(); - expect(client.projectId).toBeTruthy(); + // org/project are now resolved internally from the API key + expect(client.telemetryId).toBeTruthy(); }); test("get with invalid ID throws ValidationError", async () => { diff --git a/mem0-ts/src/client/tests/integration/search.test.ts b/mem0-ts/src/client/tests/integration/search.test.ts index ac572f63d..0591f2ffa 100644 --- a/mem0-ts/src/client/tests/integration/search.test.ts +++ b/mem0-ts/src/client/tests/integration/search.test.ts @@ -1,7 +1,7 @@ /** * Integration tests: Search and history operations. * - * Tests search v1, search v2, and memory history against the real API. + * Tests search, filtered search, and memory history against the real API. * * Run: MEM0_API_KEY=your-key npx jest search.test.ts --forceExit */ @@ -36,14 +36,14 @@ describeIntegration("MemoryClient Integration — Search & History", () => { cleanup(); }); - // ─── Search v1 ──────────────────────────────────────────── - describe("search v1", () => { + // ─── Search ───────────────────────────────────────────── + describe("search", () => { test("searches memories by user_id and returns results with scores", async () => { // Search index may lag behind listing index — poll until ready const results = await waitForSearchResults( client, "What is my favorite color?", - { user_id: TEST_USER_ID }, + { filters: { user_id: TEST_USER_ID } }, ); expect(Array.isArray(results)).toBe(true); @@ -57,15 +57,14 @@ describeIntegration("MemoryClient Integration — Search & History", () => { }); }); - // ─── Search v2 ──────────────────────────────────────────── - describe("search v2", () => { + // ─── Search with filters ───────────────────────────────── + describe("search with filters", () => { test("searches with OR filters and returns results", async () => { const results = await waitForSearchResults( client, "What do you know about me?", { filters: { OR: [{ user_id: TEST_USER_ID }] }, - api_version: "v2", }, ); @@ -110,19 +109,19 @@ describeIntegration("MemoryClient Integration — Search & History", () => { describe("edge cases", () => { test("search for non-existent user returns empty results", async () => { const results = await client.search("anything", { - user_id: `nonexistent-user-${randomUUID()}`, + filters: { user_id: `nonexistent-user-${randomUUID()}` }, }); expect(Array.isArray(results)).toBe(true); expect(results.length).toBe(0); }); - test("search with limit param does not throw", async () => { + test("search with top_k param does not throw", async () => { const results = await client.search( "Tell me about integration test user", { - user_id: TEST_USER_ID, - limit: 1, + filters: { user_id: TEST_USER_ID }, + top_k: 1, }, ); diff --git a/mem0-ts/src/client/tests/memoryClient.crud.test.ts b/mem0-ts/src/client/tests/memoryClient.crud.test.ts index 58ada7bbf..7391f8994 100644 --- a/mem0-ts/src/client/tests/memoryClient.crud.test.ts +++ b/mem0-ts/src/client/tests/memoryClient.crud.test.ts @@ -1,15 +1,13 @@ /** - * MemoryClient unit tests — add, get, getAll, update, delete, deleteAll, history. + * MemoryClient unit tests — add, get, update, delete, deleteAll, history. * Tests verify request construction, not mock response echo. */ import { MemoryClient } from "../mem0"; -import type { Memory, MemoryHistory } from "../mem0.types"; +import type { MemoryHistory } from "../mem0.types"; import { createMockMemory, createMockMemoryHistory, TEST_API_KEY, - TEST_ORG_ID, - TEST_PROJECT_ID, } from "./helpers"; import { setupMockFetch, @@ -61,40 +59,6 @@ describe("MemoryClient - add()", () => { expect(getFetchBody(call!).user_id).toBe("user_1"); }); - test("attaches org_id from constructor to payload", async () => { - const extra = new Map(); - extra.set("/v1/memories/", { status: 200, body: [createMockMemory()] }); - const mock = setupMockFetch(extra); - - const client = new MemoryClient({ - apiKey: TEST_API_KEY, - organizationId: TEST_ORG_ID, - projectId: TEST_PROJECT_ID, - }); - await client.add([{ role: "user", content: "test" }], { user_id: "u1" }); - - const call = findFetchCall(mock, "/v1/memories/", "POST"); - const body = getFetchBody(call!); - expect(body.org_id).toBe(TEST_ORG_ID); - }); - - test("attaches project_id from constructor to payload", async () => { - const extra = new Map(); - extra.set("/v1/memories/", { status: 200, body: [createMockMemory()] }); - const mock = setupMockFetch(extra); - - const client = new MemoryClient({ - apiKey: TEST_API_KEY, - organizationId: TEST_ORG_ID, - projectId: TEST_PROJECT_ID, - }); - await client.add([{ role: "user", content: "test" }], { user_id: "u1" }); - - const call = findFetchCall(mock, "/v1/memories/", "POST"); - const body = getFetchBody(call!); - expect(body.project_id).toBe(TEST_PROJECT_ID); - }); - test("sends empty messages array without crashing", async () => { const extra = new Map(); extra.set("/v1/memories/", { status: 200, body: [] }); @@ -142,67 +106,6 @@ describe("MemoryClient - get()", () => { }); }); -// ─── getAll() ──────────────────────────────────────────── - -describe("MemoryClient - getAll()", () => { - test("uses v2 POST endpoint when api_version=v2", async () => { - const extra = new Map(); - extra.set("/v2/memories/", { status: 200, body: [] }); - const mock = setupMockFetch(extra); - - const client = new MemoryClient({ apiKey: TEST_API_KEY }); - await client.getAll({ user_id: "u1", api_version: "v2" }); - - expect(findFetchCall(mock, "/v2/memories/", "POST")).toBeDefined(); - }); - - test("uses v1 GET endpoint by default with user_id as query param", async () => { - const extra = new Map(); - extra.set("/v1/memories/", { status: 200, body: [] }); - const mock = setupMockFetch(extra); - - const client = new MemoryClient({ apiKey: TEST_API_KEY }); - await client.getAll({ user_id: "u1" }); - - const call = mock.mock.calls.find( - (c: [string, RequestInit]) => - c[0].includes("/v1/memories/?") && !c[1]?.method, - ); - expect(call).toBeDefined(); - expect(call![0]).toContain("user_id=u1"); - }); - - test("appends page and page_size to URL as query params", async () => { - const extra = new Map(); - extra.set("/v2/memories/", { status: 200, body: [] }); - const mock = setupMockFetch(extra); - - const client = new MemoryClient({ apiKey: TEST_API_KEY }); - await client.getAll({ - user_id: "u1", - api_version: "v2", - page: 2, - page_size: 25, - }); - - const call = mock.mock.calls.find((c: [string, RequestInit]) => - c[0].includes("page="), - ); - expect(call![0]).toContain("page=2"); - expect(call![0]).toContain("page_size=25"); - }); - - test("does not crash when called without options", async () => { - const extra = new Map(); - extra.set("/v1/memories/", { status: 200, body: [] }); - setupMockFetch(extra); - - const client = new MemoryClient({ apiKey: TEST_API_KEY }); - const result: Memory[] = await client.getAll(); - expect(Array.isArray(result)).toBe(true); - }); -}); - // ─── update() ──────────────────────────────────────────── describe("MemoryClient - update()", () => { diff --git a/mem0-ts/src/client/tests/memoryClient.init.test.ts b/mem0-ts/src/client/tests/memoryClient.init.test.ts index f31e89cdb..17de53b55 100644 --- a/mem0-ts/src/client/tests/memoryClient.init.test.ts +++ b/mem0-ts/src/client/tests/memoryClient.init.test.ts @@ -7,13 +7,7 @@ import { ValidationError, MemoryError, } from "../../common/exceptions"; -import { - createMockFetch, - TEST_API_KEY, - TEST_HOST, - TEST_ORG_ID, - TEST_PROJECT_ID, -} from "./helpers"; +import { createMockFetch, TEST_API_KEY, TEST_HOST } from "./helpers"; import { setupMockFetch, installConsoleSuppression, @@ -55,24 +49,6 @@ describe("MemoryClient - Initialization", () => { expect(client.host).toBe(TEST_HOST); }); - test("sets organizationId from constructor", () => { - const client = new MemoryClient({ - apiKey: TEST_API_KEY, - organizationId: TEST_ORG_ID, - projectId: TEST_PROJECT_ID, - }); - expect(client.organizationId).toBe(TEST_ORG_ID); - }); - - test("sets projectId from constructor", () => { - const client = new MemoryClient({ - apiKey: TEST_API_KEY, - organizationId: TEST_ORG_ID, - projectId: TEST_PROJECT_ID, - }); - expect(client.projectId).toBe(TEST_PROJECT_ID); - }); - test("sets Authorization header with Token prefix", () => { const client = new MemoryClient({ apiKey: TEST_API_KEY }); expect(client.headers["Authorization"]).toBe(`Token ${TEST_API_KEY}`); @@ -87,20 +63,6 @@ describe("MemoryClient - Initialization", () => { // ─── Ping ──────────────────────────────────────────────── describe("MemoryClient - ping()", () => { - test("sets organizationId from ping response", async () => { - setupMockFetch(); - const client = new MemoryClient({ apiKey: TEST_API_KEY }); - await client.ping(); - expect(client.organizationId).toBe(TEST_ORG_ID); - }); - - test("sets projectId from ping response", async () => { - setupMockFetch(); - const client = new MemoryClient({ apiKey: TEST_API_KEY }); - await client.ping(); - expect(client.projectId).toBe(TEST_PROJECT_ID); - }); - test("sets telemetryId from user_email in response", async () => { setupMockFetch(); const client = new MemoryClient({ apiKey: TEST_API_KEY }); @@ -108,28 +70,6 @@ describe("MemoryClient - ping()", () => { expect(client.telemetryId).toBe("test@example.com"); }); - test("preserves constructor organizationId over ping response", async () => { - setupMockFetch(); - const client = new MemoryClient({ - apiKey: TEST_API_KEY, - organizationId: "my_org", - projectId: "my_proj", - }); - await client.ping(); - expect(client.organizationId).toBe("my_org"); - }); - - test("preserves constructor projectId over ping response", async () => { - setupMockFetch(); - const client = new MemoryClient({ - apiKey: TEST_API_KEY, - organizationId: "my_org", - projectId: "my_proj", - }); - await client.ping(); - expect(client.projectId).toBe("my_proj"); - }); - test("throws AuthenticationError on 401 response", async () => { const { AuthenticationError } = await import("../../common/exceptions"); const responses = new Map(); diff --git a/mem0-ts/src/client/tests/memoryClient.project.test.ts b/mem0-ts/src/client/tests/memoryClient.project.test.ts index 01f02c760..4e5c9014e 100644 --- a/mem0-ts/src/client/tests/memoryClient.project.test.ts +++ b/mem0-ts/src/client/tests/memoryClient.project.test.ts @@ -4,12 +4,7 @@ */ import { MemoryClient } from "../mem0"; import { Feedback } from "../mem0.types"; -import { - createMockFetch, - TEST_API_KEY, - TEST_ORG_ID, - TEST_PROJECT_ID, -} from "./helpers"; +import { createMockFetch, TEST_API_KEY } from "./helpers"; import { setupMockFetch, findFetchCall, @@ -22,7 +17,7 @@ installConsoleSuppression(); // ─── getProject() ─────────────────────────────────────── describe("MemoryClient - getProject()", () => { - test("throws when organizationId and projectId not set", async () => { + test("throws when organizationId and projectId not set (ping returns no org)", async () => { const responses = new Map(); responses.set("/v1/ping/", { status: 200, body: { status: "ok" } }); global.fetch = createMockFetch(responses); @@ -47,11 +42,9 @@ describe("MemoryClient - getProject()", () => { }); const mock = setupMockFetch(extra); - const client = new MemoryClient({ - apiKey: TEST_API_KEY, - organizationId: TEST_ORG_ID, - projectId: TEST_PROJECT_ID, - }); + // org/project come from ping mock response + const client = new MemoryClient({ apiKey: TEST_API_KEY }); + await client.ping(); await client.getProject({ fields: ["custom_instructions"] }); const call = mock.mock.calls.find( @@ -74,11 +67,8 @@ describe("MemoryClient - updateProject()", () => { }); const mock = setupMockFetch(extra); - const client = new MemoryClient({ - apiKey: TEST_API_KEY, - organizationId: TEST_ORG_ID, - projectId: TEST_PROJECT_ID, - }); + const client = new MemoryClient({ apiKey: TEST_API_KEY }); + await client.ping(); await client.updateProject({ custom_instructions: "Updated instructions", }); @@ -95,11 +85,8 @@ describe("MemoryClient - updateProject()", () => { }); const mock = setupMockFetch(extra); - const client = new MemoryClient({ - apiKey: TEST_API_KEY, - organizationId: TEST_ORG_ID, - projectId: TEST_PROJECT_ID, - }); + const client = new MemoryClient({ apiKey: TEST_API_KEY }); + await client.ping(); await client.updateProject({ custom_instructions: "Updated instructions", }); @@ -161,11 +148,7 @@ describe("MemoryClient - feedback()", () => { describe("MemoryClient - Memory Exports", () => { test("createMemoryExport throws when missing filters or schema", async () => { setupMockFetch(); - const client = new MemoryClient({ - apiKey: TEST_API_KEY, - organizationId: TEST_ORG_ID, - projectId: TEST_PROJECT_ID, - }); + const client = new MemoryClient({ apiKey: TEST_API_KEY }); await expect( client.createMemoryExport({ filters: null as never, @@ -182,11 +165,7 @@ describe("MemoryClient - Memory Exports", () => { }); const mock = setupMockFetch(extra); - const client = new MemoryClient({ - apiKey: TEST_API_KEY, - organizationId: TEST_ORG_ID, - projectId: TEST_PROJECT_ID, - }); + const client = new MemoryClient({ apiKey: TEST_API_KEY }); await client.createMemoryExport({ schema: { fields: ["memory", "user_id"] }, filters: { user_id: "u1" }, @@ -195,37 +174,9 @@ describe("MemoryClient - Memory Exports", () => { expect(findFetchCall(mock, "/v1/exports/", "POST")).toBeDefined(); }); - test("createMemoryExport attaches org_id and project_id to body", async () => { - const extra = new Map(); - extra.set("/v1/exports/", { - status: 200, - body: { message: "Created", id: "exp_1" }, - }); - const mock = setupMockFetch(extra); - - const client = new MemoryClient({ - apiKey: TEST_API_KEY, - organizationId: TEST_ORG_ID, - projectId: TEST_PROJECT_ID, - }); - await client.createMemoryExport({ - schema: { fields: ["memory"] }, - filters: { user_id: "u1" }, - }); - - const call = findFetchCall(mock, "/v1/exports/", "POST"); - const body = getFetchBody(call!); - expect(body.org_id).toBe(TEST_ORG_ID); - expect(body.project_id).toBe(TEST_PROJECT_ID); - }); - test("getMemoryExport throws when missing both id and filters", async () => { setupMockFetch(); - const client = new MemoryClient({ - apiKey: TEST_API_KEY, - organizationId: TEST_ORG_ID, - projectId: TEST_PROJECT_ID, - }); + const client = new MemoryClient({ apiKey: TEST_API_KEY }); await expect(client.getMemoryExport({} as never)).rejects.toThrow( "Missing memory_export_id or filters", ); @@ -239,11 +190,7 @@ describe("MemoryClient - Memory Exports", () => { }); const mock = setupMockFetch(extra); - const client = new MemoryClient({ - apiKey: TEST_API_KEY, - organizationId: TEST_ORG_ID, - projectId: TEST_PROJECT_ID, - }); + const client = new MemoryClient({ apiKey: TEST_API_KEY }); await client.getMemoryExport({ memory_export_id: "exp_123" }); expect(findFetchCall(mock, "/v1/exports/get/", "POST")).toBeDefined(); diff --git a/mem0-ts/src/client/tests/memoryClient.search.test.ts b/mem0-ts/src/client/tests/memoryClient.search.test.ts index b7c97595f..6efbaf6b5 100644 --- a/mem0-ts/src/client/tests/memoryClient.search.test.ts +++ b/mem0-ts/src/client/tests/memoryClient.search.test.ts @@ -1,5 +1,5 @@ /** - * MemoryClient unit tests — search (v1/v2 routing, filters). + * MemoryClient unit tests — search (v2 default, filters). * Tests verify request construction, not mock response echo. */ import { MemoryClient } from "../mem0"; @@ -15,60 +15,52 @@ import { installConsoleSuppression(); describe("MemoryClient - search()", () => { - test("sends POST to /v1/memories/search/ by default", async () => { - const extra = new Map(); - extra.set("/v1/memories/search/", { status: 200, body: [] }); - const mock = setupMockFetch(extra); - - const client = new MemoryClient({ apiKey: TEST_API_KEY }); - await client.search("What is my name?", { user_id: "u1" }); - - expect(findFetchCall(mock, "/v1/memories/search/", "POST")).toBeDefined(); - }); - - test("includes query in request body", async () => { - const extra = new Map(); - extra.set("/v1/memories/search/", { status: 200, body: [] }); - const mock = setupMockFetch(extra); - - const client = new MemoryClient({ apiKey: TEST_API_KEY }); - await client.search("What is my name?", { user_id: "u1" }); - - const call = findFetchCall(mock, "/v1/memories/search/", "POST"); - expect(getFetchBody(call!).query).toBe("What is my name?"); - }); - - test("includes user_id in request body", async () => { - const extra = new Map(); - extra.set("/v1/memories/search/", { status: 200, body: [] }); - const mock = setupMockFetch(extra); - - const client = new MemoryClient({ apiKey: TEST_API_KEY }); - await client.search("test", { user_id: "u1" }); - - const call = findFetchCall(mock, "/v1/memories/search/", "POST"); - expect(getFetchBody(call!).user_id).toBe("u1"); - }); - - test("uses /v2/memories/search/ when api_version=v2", async () => { + test("sends POST to /v2/memories/search/ by default", async () => { const extra = new Map(); extra.set("/v2/memories/search/", { status: 200, body: [] }); const mock = setupMockFetch(extra); const client = new MemoryClient({ apiKey: TEST_API_KEY }); - await client.search("test", { user_id: "u1", api_version: "v2" }); + await client.search("What is my name?", { + filters: { user_id: "u1" }, + }); expect(findFetchCall(mock, "/v2/memories/search/", "POST")).toBeDefined(); }); - test("passes filters through to the v2 API body", async () => { + test("includes query in request body", async () => { + const extra = new Map(); + extra.set("/v2/memories/search/", { status: 200, body: [] }); + const mock = setupMockFetch(extra); + + const client = new MemoryClient({ apiKey: TEST_API_KEY }); + await client.search("What is my name?", { + filters: { user_id: "u1" }, + }); + + const call = findFetchCall(mock, "/v2/memories/search/", "POST"); + expect(getFetchBody(call!).query).toBe("What is my name?"); + }); + + test("passes filters through to the API body", async () => { + const extra = new Map(); + extra.set("/v2/memories/search/", { status: 200, body: [] }); + const mock = setupMockFetch(extra); + + const client = new MemoryClient({ apiKey: TEST_API_KEY }); + await client.search("test", { filters: { user_id: "u1" } }); + + const call = findFetchCall(mock, "/v2/memories/search/", "POST"); + expect(getFetchBody(call!).filters).toEqual({ user_id: "u1" }); + }); + + test("passes complex OR filters through to the API body", async () => { const extra = new Map(); extra.set("/v2/memories/search/", { status: 200, body: [] }); const mock = setupMockFetch(extra); const client = new MemoryClient({ apiKey: TEST_API_KEY }); await client.search("query", { - api_version: "v2", filters: { OR: [{ user_id: "u1" }, { agent_id: "a1" }] }, }); @@ -81,7 +73,7 @@ describe("MemoryClient - search()", () => { test("does not crash when called without options", async () => { const extra = new Map(); - extra.set("/v1/memories/search/", { status: 200, body: [] }); + extra.set("/v2/memories/search/", { status: 200, body: [] }); setupMockFetch(extra); const client = new MemoryClient({ apiKey: TEST_API_KEY }); @@ -91,12 +83,12 @@ describe("MemoryClient - search()", () => { test("handles empty results array", async () => { const extra = new Map(); - extra.set("/v1/memories/search/", { status: 200, body: [] }); + extra.set("/v2/memories/search/", { status: 200, body: [] }); setupMockFetch(extra); const client = new MemoryClient({ apiKey: TEST_API_KEY }); const result: Memory[] = await client.search("nonexistent query", { - user_id: "u1", + filters: { user_id: "u1" }, }); expect(result).toHaveLength(0); }); diff --git a/mem0-ts/src/client/tests/memoryClient.users.test.ts b/mem0-ts/src/client/tests/memoryClient.users.test.ts index f4f6d5688..ba232f21c 100644 --- a/mem0-ts/src/client/tests/memoryClient.users.test.ts +++ b/mem0-ts/src/client/tests/memoryClient.users.test.ts @@ -1,15 +1,9 @@ /** - * MemoryClient unit tests — users, deleteUser, deleteUsers. + * MemoryClient unit tests — users, deleteUser. * Tests verify entity type routing and request construction. */ import { MemoryClient } from "../mem0"; -import { - createMockUser, - createMockAllUsers, - TEST_API_KEY, - TEST_ORG_ID, - TEST_PROJECT_ID, -} from "./helpers"; +import { createMockUser, createMockAllUsers, TEST_API_KEY } from "./helpers"; import { setupMockFetch, findFetchCall, @@ -40,111 +34,6 @@ describe("MemoryClient - users()", () => { }); }); -// ─── deleteUsers() ────────────────────────────────────── - -describe("MemoryClient - deleteUsers()", () => { - function createClientWithMockedAxios() { - setupMockFetch(); - const client = new MemoryClient({ - apiKey: TEST_API_KEY, - organizationId: TEST_ORG_ID, - projectId: TEST_PROJECT_ID, - }); - const axiosDeleteMock = jest - .fn() - .mockResolvedValue({ data: { message: "Deleted" } }); - client.client.delete = axiosDeleteMock; - return { client, axiosDeleteMock }; - } - - test("routes user_id to DELETE /v2/entities/user/:name/", async () => { - const { client, axiosDeleteMock } = createClientWithMockedAxios(); - await client.deleteUsers({ user_id: "u1" }); - - expect(axiosDeleteMock).toHaveBeenCalledWith("/v2/entities/user/u1/", { - params: expect.objectContaining({ - org_id: TEST_ORG_ID, - project_id: TEST_PROJECT_ID, - }), - }); - }); - - test("routes agent_id to DELETE /v2/entities/agent/:name/", async () => { - const { client, axiosDeleteMock } = createClientWithMockedAxios(); - await client.deleteUsers({ agent_id: "agent_1" }); - - expect(axiosDeleteMock).toHaveBeenCalledWith( - "/v2/entities/agent/agent_1/", - expect.any(Object), - ); - }); - - test("routes app_id to DELETE /v2/entities/app/:name/", async () => { - const { client, axiosDeleteMock } = createClientWithMockedAxios(); - await client.deleteUsers({ app_id: "app_1" }); - - expect(axiosDeleteMock).toHaveBeenCalledWith( - "/v2/entities/app/app_1/", - expect.any(Object), - ); - }); - - test("routes run_id to DELETE /v2/entities/run/:name/", async () => { - const { client, axiosDeleteMock } = createClientWithMockedAxios(); - await client.deleteUsers({ run_id: "run_1" }); - - expect(axiosDeleteMock).toHaveBeenCalledWith( - "/v2/entities/run/run_1/", - expect.any(Object), - ); - }); - - test("returns 'Entity deleted successfully.' for single entity", async () => { - const { client } = createClientWithMockedAxios(); - const result = await client.deleteUsers({ user_id: "u1" }); - expect(result.message).toBe("Entity deleted successfully."); - }); - - test("returns 'All users, agents, apps and runs deleted.' when no params given", async () => { - const extra = new Map(); - extra.set("/v1/entities/", { - status: 200, - body: createMockAllUsers([createMockUser({ name: "u1", type: "user" })]), - }); - setupMockFetch(extra); - - const client = new MemoryClient({ - apiKey: TEST_API_KEY, - organizationId: TEST_ORG_ID, - projectId: TEST_PROJECT_ID, - }); - client.client.delete = jest - .fn() - .mockResolvedValue({ data: { message: "Deleted" } }); - - const result = await client.deleteUsers(); - expect(result.message).toBe("All users, agents, apps and runs deleted."); - }); - - test("throws when no entities exist to delete", async () => { - const extra = new Map(); - extra.set("/v1/entities/", { - status: 200, - body: createMockAllUsers([]), - }); - setupMockFetch(extra); - - const client = new MemoryClient({ - apiKey: TEST_API_KEY, - organizationId: TEST_ORG_ID, - projectId: TEST_PROJECT_ID, - }); - client.client.delete = jest.fn(); - - await expect(client.deleteUsers()).rejects.toThrow("No entities to delete"); - }); -}); - // ─── deleteUser() (deprecated) ────────────────────────── describe("MemoryClient - deleteUser() (deprecated)", () => { diff --git a/mem0-ts/src/client/tests/memoryClient.webhooks.test.ts b/mem0-ts/src/client/tests/memoryClient.webhooks.test.ts index 6e4de484f..c3d1f54e7 100644 --- a/mem0-ts/src/client/tests/memoryClient.webhooks.test.ts +++ b/mem0-ts/src/client/tests/memoryClient.webhooks.test.ts @@ -5,7 +5,7 @@ */ import { MemoryClient } from "../mem0"; import { WebhookEvent } from "../mem0.types"; -import { TEST_API_KEY, TEST_ORG_ID, TEST_PROJECT_ID } from "./helpers"; +import { TEST_API_KEY } from "./helpers"; import { setupMockFetch, findFetchCall, @@ -23,8 +23,6 @@ function webhookMock(extra?: Map) { function createClient() { return new MemoryClient({ apiKey: TEST_API_KEY, - organizationId: TEST_ORG_ID, - projectId: TEST_PROJECT_ID, }); } @@ -100,14 +98,6 @@ describe("MemoryClient - createWebhook", () => { expect(body.eventTypes).toBeUndefined(); }); - test("body does not contain projectId", async () => { - const mock = await callCreate(); - const body = getFetchBody( - findFetchCall(mock, "/api/v1/webhooks/", "POST")!, - ); - expect(body.projectId).toBeUndefined(); - }); - test("body does not contain webhookId", async () => { const mock = await callCreate(); const body = getFetchBody( @@ -173,22 +163,6 @@ describe("MemoryClient - updateWebhook", () => { expect(body.eventTypes).toBeUndefined(); }); - test("body does not contain project_id", async () => { - const mock = await callUpdate(); - const body = getFetchBody( - findFetchCall(mock, "/api/v1/webhooks/wh_1/", "PUT")!, - ); - expect(body.project_id).toBeUndefined(); - }); - - test("body does not contain projectId", async () => { - const mock = await callUpdate(); - const body = getFetchBody( - findFetchCall(mock, "/api/v1/webhooks/wh_1/", "PUT")!, - ); - expect(body.projectId).toBeUndefined(); - }); - test("body does not contain webhookId", async () => { const mock = await callUpdate(); const body = getFetchBody( diff --git a/mem0-ts/src/community/src/integrations/langchain/mem0.ts b/mem0-ts/src/community/src/integrations/langchain/mem0.ts index 315cdd32c..1aac51358 100644 --- a/mem0-ts/src/community/src/integrations/langchain/mem0.ts +++ b/mem0-ts/src/community/src/integrations/langchain/mem0.ts @@ -1,5 +1,10 @@ import { MemoryClient } from "mem0ai"; -import type { Memory, MemoryOptions, SearchOptions } from "mem0ai"; +import type { + Memory, + AddMemoryOptions, + SearchMemoryOptions, + GetAllMemoryOptions, +} from "mem0ai"; import { InputValues, @@ -102,10 +107,6 @@ export const mem0MemoryToMessages = (memories: Memory[]): BaseMessage[] => { export interface ClientOptions { apiKey: string; host?: string; - organizationName?: string; - projectName?: string; - organizationId?: string; - projectId?: string; } /** @@ -117,7 +118,7 @@ export interface Mem0MemoryInput extends BaseChatMemoryInput { apiKey: string; humanPrefix?: string; aiPrefix?: string; - memoryOptions?: MemoryOptions | SearchOptions; + memoryOptions?: AddMemoryOptions | SearchMemoryOptions | GetAllMemoryOptions; mem0Options?: ClientOptions; separateMessages?: boolean; } @@ -160,7 +161,7 @@ export class Mem0Memory extends BaseChatMemory implements Mem0MemoryInput { mem0Client: InstanceType; - memoryOptions: MemoryOptions | SearchOptions; + memoryOptions: AddMemoryOptions | SearchMemoryOptions | GetAllMemoryOptions; mem0Options: ClientOptions; diff --git a/mem0-ts/src/oss/src/config/manager.ts b/mem0-ts/src/oss/src/config/manager.ts index 6a4a85376..6235157ce 100644 --- a/mem0-ts/src/oss/src/config/manager.ts +++ b/mem0-ts/src/oss/src/config/manager.ts @@ -131,7 +131,7 @@ export class ConfigManager { userConfig.historyDbPath || userConfig.historyStore?.config?.historyDbPath || DEFAULT_MEMORY_CONFIG.historyStore?.config?.historyDbPath, - customPrompt: userConfig.customPrompt, + customInstructions: userConfig.customInstructions, graphStore: { ...DEFAULT_MEMORY_CONFIG.graphStore, ...userConfig.graphStore, diff --git a/mem0-ts/src/oss/src/graphs/configs.ts b/mem0-ts/src/oss/src/graphs/configs.ts index bb7491dfa..a3a7579e2 100644 --- a/mem0-ts/src/oss/src/graphs/configs.ts +++ b/mem0-ts/src/oss/src/graphs/configs.ts @@ -10,7 +10,7 @@ export interface GraphStoreConfig { provider: string; config: Neo4jConfig; llm?: LLMConfig; - customPrompt?: string; + customInstructions?: string; } export function validateNeo4jConfig(config: Neo4jConfig): void { diff --git a/mem0-ts/src/oss/src/memory/graph_memory.ts b/mem0-ts/src/oss/src/memory/graph_memory.ts index 3996cffb0..da865da65 100644 --- a/mem0-ts/src/oss/src/memory/graph_memory.ts +++ b/mem0-ts/src/oss/src/memory/graph_memory.ts @@ -136,7 +136,7 @@ export class MemoryGraph { }; } - async search(query: string, filters: Record, limit = 100) { + async search(query: string, filters: Record, topK = 100) { const entityTypeMap = await this._retrieveNodesFromData(query, filters); const searchOutput = await this._searchGraphDb( Object.keys(entityTypeMap), @@ -178,7 +178,7 @@ export class MemoryGraph { } } - async getAll(filters: Record, limit = 100) { + async getAll(filters: Record, topK = 100) { const session = this.graph.session(); try { const result = await session.run( @@ -187,7 +187,7 @@ export class MemoryGraph { RETURN n.name AS source, type(r) AS relationship, m.name AS target LIMIT toInteger($limit) `, - { user_id: filters["userId"], limit: Math.floor(Number(limit)) }, + { user_id: filters["userId"], limit: Math.floor(Number(topK)) }, ); const finalResults = result.records.map((record) => ({ @@ -253,7 +253,7 @@ export class MemoryGraph { entityTypeMap: Record, ) { let messages; - if (this.config.graphStore?.customPrompt) { + if (this.config.graphStore?.customInstructions) { messages = [ { role: "system", @@ -263,7 +263,7 @@ export class MemoryGraph { filters["userId"], ).replace( "CUSTOM_PROMPT", - `4. ${this.config.graphStore.customPrompt}`, + `4. ${this.config.graphStore.customInstructions}`, ) + "\nPlease provide your response in JSON format.", }, { role: "user", content: data }, @@ -307,7 +307,7 @@ export class MemoryGraph { private async _searchGraphDb( nodeList: string[], filters: Record, - limit = 100, + topK = 100, ): Promise { const resultRelations: SearchOutput[] = []; const session = this.graph.session(); @@ -344,7 +344,7 @@ export class MemoryGraph { n_embedding: nEmbedding, threshold: this.threshold, user_id: filters["userId"], - limit: Math.floor(Number(limit)), + limit: Math.floor(Number(topK)), }); resultRelations.push( diff --git a/mem0-ts/src/oss/src/memory/index.ts b/mem0-ts/src/oss/src/memory/index.ts index 01828d04d..d4c4e25d0 100644 --- a/mem0-ts/src/oss/src/memory/index.ts +++ b/mem0-ts/src/oss/src/memory/index.ts @@ -39,7 +39,7 @@ import { captureClientEvent } from "../utils/telemetry"; export class Memory { private config: MemoryConfig; - private customPrompt: string | undefined; + private customInstructions: string | undefined; private embedder: Embedder; private vectorStore!: VectorStore; private llm: LLM; @@ -56,7 +56,7 @@ export class Memory { // Merge and validate config this.config = ConfigManager.mergeConfig(config); - this.customPrompt = this.config.customPrompt; + this.customInstructions = this.config.customInstructions; this.embedder = EmbedderFactory.create( this.config.embedder.provider, this.config.embedder.config, @@ -293,11 +293,11 @@ export class Memory { } const parsedMessages = messages.map((m) => m.content).join("\n"); - const [systemPrompt, userPrompt] = this.customPrompt + const [systemPrompt, userPrompt] = this.customInstructions ? [ - this.customPrompt.toLowerCase().includes("json") - ? this.customPrompt - : `${this.customPrompt}\n\nYou MUST return a valid JSON object with a 'facts' key containing an array of strings.`, + this.customInstructions.toLowerCase().includes("json") + ? this.customInstructions + : `${this.customInstructions}\n\nYou MUST return a valid JSON object with a 'facts' key containing an array of strings.`, `Input:\n${parsedMessages}`, ] : getFactRetrievalMessages(parsedMessages); @@ -478,10 +478,10 @@ export class Memory { await this._ensureInitialized(); await this._captureEvent("search", { query_length: query.length, - limit: config.limit, + topK: config.topK, has_filters: !!config.filters, }); - const { userId, agentId, runId, limit = 100, filters = {} } = config; + const { userId, agentId, runId, topK = 100, filters = {} } = config; if (userId) filters.userId = userId; if (agentId) filters.agentId = agentId; @@ -497,7 +497,7 @@ export class Memory { const queryEmbedding = await this.embedder.embed(query); const memories = await this.vectorStore.search( queryEmbedding, - limit, + topK, filters, ); @@ -642,19 +642,19 @@ export class Memory { async getAll(config: GetAllMemoryOptions): Promise { await this._ensureInitialized(); await this._captureEvent("get_all", { - limit: config.limit, + topK: config.topK, has_user_id: !!config.userId, has_agent_id: !!config.agentId, has_run_id: !!config.runId, }); - const { userId, agentId, runId, limit = 100 } = config; + const { userId, agentId, runId, topK = 100 } = config; const filters: SearchFilters = {}; if (userId) filters.userId = userId; if (agentId) filters.agentId = agentId; if (runId) filters.runId = runId; - const [memories] = await this.vectorStore.list(filters, limit); + const [memories] = await this.vectorStore.list(filters, topK); const excludedKeys = new Set([ "userId", diff --git a/mem0-ts/src/oss/src/memory/memory.types.ts b/mem0-ts/src/oss/src/memory/memory.types.ts index 82bb59cd1..26d5cc352 100644 --- a/mem0-ts/src/oss/src/memory/memory.types.ts +++ b/mem0-ts/src/oss/src/memory/memory.types.ts @@ -14,12 +14,12 @@ export interface AddMemoryOptions extends Entity { } export interface SearchMemoryOptions extends Entity { - limit?: number; + topK?: number; filters?: SearchFilters; } export interface GetAllMemoryOptions extends Entity { - limit?: number; + topK?: number; } export interface DeleteAllMemoryOptions extends Entity {} diff --git a/mem0-ts/src/oss/src/tests/sqlite-backward-compat.test.ts b/mem0-ts/src/oss/src/tests/sqlite-backward-compat.test.ts index 2a9d29253..8e2eb7750 100644 --- a/mem0-ts/src/oss/src/tests/sqlite-backward-compat.test.ts +++ b/mem0-ts/src/oss/src/tests/sqlite-backward-compat.test.ts @@ -120,11 +120,11 @@ describe("backward compat: ConfigManager.mergeConfig", () => { expect(cfg.graphStore!.config.url).toBe("neo4j://custom:7687"); }); - it("customPrompt passes through unchanged", () => { + it("customInstructions passes through unchanged", () => { const cfg = ConfigManager.mergeConfig({ - customPrompt: "You are a helpful assistant", + customInstructions: "You are a helpful assistant", }); - expect(cfg.customPrompt).toBe("You are a helpful assistant"); + expect(cfg.customInstructions).toBe("You are a helpful assistant"); }); it("version override passes through unchanged", () => { diff --git a/mem0-ts/src/oss/src/types/index.ts b/mem0-ts/src/oss/src/types/index.ts index 58bfb544d..0476ef515 100644 --- a/mem0-ts/src/oss/src/types/index.ts +++ b/mem0-ts/src/oss/src/types/index.ts @@ -60,7 +60,7 @@ export interface GraphStoreConfig { provider: string; config: Neo4jConfig; llm?: LLMConfig; - customPrompt?: string; + customInstructions?: string; } export interface MemoryConfig { @@ -80,7 +80,7 @@ export interface MemoryConfig { historyStore?: HistoryStoreConfig; disableHistory?: boolean; historyDbPath?: string; - customPrompt?: string; + customInstructions?: string; graphStore?: GraphStoreConfig; enableGraph?: boolean; } @@ -148,7 +148,7 @@ export const MemoryConfigSchema = z.object({ }), }), historyDbPath: z.string().optional(), - customPrompt: z.string().optional(), + customInstructions: z.string().optional(), enableGraph: z.boolean().optional(), graphStore: z .object({ @@ -164,7 +164,7 @@ export const MemoryConfigSchema = z.object({ config: z.record(z.string(), z.any()), }) .optional(), - customPrompt: z.string().optional(), + customInstructions: z.string().optional(), }) .optional(), historyStore: z diff --git a/mem0-ts/src/oss/src/vector_stores/azure_ai_search.ts b/mem0-ts/src/oss/src/vector_stores/azure_ai_search.ts index 3e3b586fe..3f765dda5 100644 --- a/mem0-ts/src/oss/src/vector_stores/azure_ai_search.ts +++ b/mem0-ts/src/oss/src/vector_stores/azure_ai_search.ts @@ -331,7 +331,7 @@ export class AzureAISearch implements VectorStore { */ async search( query: number[], - limit: number = 5, + topK: number = 5, filters?: SearchFilters, ): Promise { const filterExpression = filters @@ -341,7 +341,7 @@ export class AzureAISearch implements VectorStore { const vectorQuery: VectorizedQuery = { kind: "vector", vector: query, - kNearestNeighborsCount: limit, + kNearestNeighborsCount: topK, fields: ["vector"], }; @@ -355,7 +355,7 @@ export class AzureAISearch implements VectorStore { filterMode: this.vectorFilterMode as any, }, filter: filterExpression, - top: limit, + top: topK, searchFields: ["payload"], }); } else { @@ -366,7 +366,7 @@ export class AzureAISearch implements VectorStore { filterMode: this.vectorFilterMode as any, }, filter: filterExpression, - top: limit, + top: topK, }); } @@ -501,7 +501,7 @@ export class AzureAISearch implements VectorStore { */ async list( filters?: SearchFilters, - limit: number = 100, + topK: number = 100, ): Promise<[VectorStoreResult[], number]> { const filterExpression = filters ? this.buildFilterExpression(filters) @@ -509,7 +509,7 @@ export class AzureAISearch implements VectorStore { const searchResults = await this.searchClient.search("*", { filter: filterExpression, - top: limit, + top: topK, }); const results: VectorStoreResult[] = []; diff --git a/mem0-ts/src/oss/src/vector_stores/base.ts b/mem0-ts/src/oss/src/vector_stores/base.ts index cb6f24aa4..ef2d00b2d 100644 --- a/mem0-ts/src/oss/src/vector_stores/base.ts +++ b/mem0-ts/src/oss/src/vector_stores/base.ts @@ -8,7 +8,7 @@ export interface VectorStore { ): Promise; search( query: number[], - limit?: number, + topK?: number, filters?: SearchFilters, ): Promise; get(vectorId: string): Promise; @@ -21,7 +21,7 @@ export interface VectorStore { deleteCol(): Promise; list( filters?: SearchFilters, - limit?: number, + topK?: number, ): Promise<[VectorStoreResult[], number]>; getUserId(): Promise; setUserId(userId: string): Promise; diff --git a/mem0-ts/src/oss/src/vector_stores/langchain.ts b/mem0-ts/src/oss/src/vector_stores/langchain.ts index 852ecaa44..161e024cc 100644 --- a/mem0-ts/src/oss/src/vector_stores/langchain.ts +++ b/mem0-ts/src/oss/src/vector_stores/langchain.ts @@ -101,7 +101,7 @@ export class LangchainVectorStore implements VectorStore { async search( query: number[], - limit: number = 5, + topK: number = 5, filters?: SearchFilters, // filters parameter is received but will be ignored ): Promise { if (this.dimension && query.length !== this.dimension) { @@ -119,7 +119,7 @@ export class LangchainVectorStore implements VectorStore { // Call similaritySearchVectorWithScore WITHOUT the filter argument const results = await this.lcStore.similaritySearchVectorWithScore( query, - limit, + topK, // Do not pass lcFilter here ); @@ -192,7 +192,7 @@ export class LangchainVectorStore implements VectorStore { async list( filters?: SearchFilters, - limit: number = 100, + topK: number = 100, ): Promise<[VectorStoreResult[], number]> { // No standard list method in Langchain core interface. console.error( diff --git a/mem0-ts/src/oss/src/vector_stores/memory.ts b/mem0-ts/src/oss/src/vector_stores/memory.ts index 62a825507..2891bf70a 100644 --- a/mem0-ts/src/oss/src/vector_stores/memory.ts +++ b/mem0-ts/src/oss/src/vector_stores/memory.ts @@ -100,7 +100,7 @@ export class MemoryVectorStore implements VectorStore { async search( query: number[], - limit: number = 10, + topK: number = 10, filters?: SearchFilters, ): Promise { if (query.length !== this.dimension) { @@ -136,7 +136,7 @@ export class MemoryVectorStore implements VectorStore { } results.sort((a, b) => (b.score || 0) - (a.score || 0)); - return results.slice(0, limit); + return results.slice(0, topK); } async get(vectorId: string): Promise { @@ -179,7 +179,7 @@ export class MemoryVectorStore implements VectorStore { async list( filters?: SearchFilters, - limit: number = 100, + topK: number = 100, ): Promise<[VectorStoreResult[], number]> { const rows = this.db.prepare(`SELECT * FROM vectors`).all() as any[]; const results: VectorStoreResult[] = []; @@ -206,7 +206,7 @@ export class MemoryVectorStore implements VectorStore { } } - return [results.slice(0, limit), results.length]; + return [results.slice(0, topK), results.length]; } async getUserId(): Promise { diff --git a/mem0-ts/src/oss/src/vector_stores/pgvector.ts b/mem0-ts/src/oss/src/vector_stores/pgvector.ts index 78b8a7ef8..7f7a58c6c 100644 --- a/mem0-ts/src/oss/src/vector_stores/pgvector.ts +++ b/mem0-ts/src/oss/src/vector_stores/pgvector.ts @@ -164,12 +164,12 @@ export class PGVector implements VectorStore { async search( query: number[], - limit: number = 5, + topK: number = 5, filters?: SearchFilters, ): Promise { const filterConditions: string[] = []; const queryVector = `[${query.join(",")}]`; // Format query vector as string with square brackets - const filterValues: any[] = [queryVector, limit]; + const filterValues: any[] = [queryVector, topK]; let filterIndex = 3; if (filters) { @@ -254,7 +254,7 @@ export class PGVector implements VectorStore { async list( filters?: SearchFilters, - limit: number = 100, + topK: number = 100, ): Promise<[VectorStoreResult[], number]> { const filterConditions: string[] = []; const filterValues: any[] = []; @@ -286,7 +286,7 @@ export class PGVector implements VectorStore { ${filterClause} `; - filterValues.push(limit); // Add limit as the last parameter + filterValues.push(topK); // Add limit as the last parameter const [listResult, countResult] = await Promise.all([ this.client.query(listQuery, filterValues), diff --git a/mem0-ts/src/oss/src/vector_stores/qdrant.ts b/mem0-ts/src/oss/src/vector_stores/qdrant.ts index 81c6faea8..83eb05325 100644 --- a/mem0-ts/src/oss/src/vector_stores/qdrant.ts +++ b/mem0-ts/src/oss/src/vector_stores/qdrant.ts @@ -139,14 +139,14 @@ export class Qdrant implements VectorStore { async search( query: number[], - limit: number = 5, + topK: number = 5, filters?: SearchFilters, ): Promise { const queryFilter = this.createFilter(filters); const results = await this.client.search(this.collectionName, { vector: query, filter: queryFilter, - limit, + limit: topK, }); return results.map((hit) => ({ @@ -198,10 +198,10 @@ export class Qdrant implements VectorStore { async list( filters?: SearchFilters, - limit: number = 100, + topK: number = 100, ): Promise<[VectorStoreResult[], number]> { const scrollRequest = { - limit, + limit: topK, filter: this.createFilter(filters), with_payload: true, with_vectors: false, diff --git a/mem0-ts/src/oss/src/vector_stores/redis.ts b/mem0-ts/src/oss/src/vector_stores/redis.ts index 5d97b3a9c..41767df36 100644 --- a/mem0-ts/src/oss/src/vector_stores/redis.ts +++ b/mem0-ts/src/oss/src/vector_stores/redis.ts @@ -359,7 +359,7 @@ export class RedisDB implements VectorStore { async search( query: number[], - limit: number = 5, + topK: number = 5, filters?: SearchFilters, ): Promise { const snakeFilters = filters ? toSnakeCase(filters) : undefined; @@ -391,14 +391,14 @@ export class RedisDB implements VectorStore { DIALECT: 2, LIMIT: { from: 0, - size: limit, + size: topK, }, }; try { const results = (await this.client.ft.search( this.indexName, - `${filterExpr} =>[KNN ${limit} @embedding $vec AS __vector_score]`, + `${filterExpr} =>[KNN ${topK} @embedding $vec AS __vector_score]`, searchOptions, )) as unknown as RedisSearchResult; @@ -598,7 +598,7 @@ export class RedisDB implements VectorStore { async list( filters?: SearchFilters, - limit: number = 100, + topK: number = 100, ): Promise<[VectorStoreResult[], number]> { const snakeFilters = filters ? toSnakeCase(filters) : undefined; const filterExpr = snakeFilters @@ -613,7 +613,7 @@ export class RedisDB implements VectorStore { SORTDIR: "DESC", LIMIT: { from: 0, - size: limit, + size: topK, }, }; diff --git a/mem0-ts/src/oss/src/vector_stores/supabase.ts b/mem0-ts/src/oss/src/vector_stores/supabase.ts index 43db08486..e3d754ef1 100644 --- a/mem0-ts/src/oss/src/vector_stores/supabase.ts +++ b/mem0-ts/src/oss/src/vector_stores/supabase.ts @@ -231,13 +231,13 @@ See the SQL migration instructions in the code comments.`, async search( query: number[], - limit: number = 5, + topK: number = 5, filters?: SearchFilters, ): Promise { try { const rpcQuery: VectorQueryParams = { query_embedding: query, - match_count: limit, + match_count: topK, }; if (filters) { @@ -336,13 +336,13 @@ See the SQL migration instructions in the code comments.`, async list( filters?: SearchFilters, - limit: number = 100, + topK: number = 100, ): Promise<[VectorStoreResult[], number]> { try { let query = this.client .from(this.tableName) .select("*", { count: "exact" }) - .limit(limit); + .limit(topK); if (filters) { Object.entries(filters).forEach(([key, value]) => { diff --git a/mem0-ts/src/oss/src/vector_stores/vectorize.ts b/mem0-ts/src/oss/src/vector_stores/vectorize.ts index ca19a8675..dec5c3f4d 100644 --- a/mem0-ts/src/oss/src/vector_stores/vectorize.ts +++ b/mem0-ts/src/oss/src/vector_stores/vectorize.ts @@ -76,7 +76,7 @@ export class VectorizeDB implements VectorStore { async search( query: number[], - limit: number = 5, + topK: number = 5, filters?: SearchFilters, ): Promise { try { @@ -87,7 +87,7 @@ export class VectorizeDB implements VectorStore { vector: query, filter: filters, returnMetadata: "all", - topK: limit, + topK: topK, }, ); @@ -197,7 +197,7 @@ export class VectorizeDB implements VectorStore { async list( filters?: SearchFilters, - limit: number = 20, + topK: number = 20, ): Promise<[VectorStoreResult[], number]> { try { const result = await this.client?.vectorize.indexes.query( @@ -206,7 +206,7 @@ export class VectorizeDB implements VectorStore { account_id: this.accountId, vector: Array(this.dimensions).fill(0), // Dummy vector for listing filter: filters, - topK: limit, + topK: topK, returnMetadata: "all", }, ); diff --git a/mem0-ts/src/oss/tests/graph-memory-parsing.test.ts b/mem0-ts/src/oss/tests/graph-memory-parsing.test.ts index 0a3689d22..58c58e5f6 100644 --- a/mem0-ts/src/oss/tests/graph-memory-parsing.test.ts +++ b/mem0-ts/src/oss/tests/graph-memory-parsing.test.ts @@ -305,13 +305,15 @@ describe("_establishNodesRelationsFromData", () => { expect(systemContent).toContain("test-user"); expect(systemContent).not.toContain("USER_ID"); // CUSTOM_PROMPT placeholder stays when no custom prompt is configured - // (only replaced when config.graphStore.customPrompt is set) + // (only replaced when config.graphStore.customInstructions is set) }); it("appends JSON format suffix and custom prompt when configured", async () => { mockGenerateResponse.mockResolvedValueOnce({ toolCalls: [] }); - const mg = graph({ customPrompt: "Focus on food relationships only." }); + const mg = graph({ + customInstructions: "Focus on food relationships only.", + }); await mg._establishNodesRelationsFromData("data", FILTERS, {}); const [messages] = mockGenerateResponse.mock.calls[0]; diff --git a/mem0/client/main.py b/mem0/client/main.py index aef020862..6511e73ca 100644 --- a/mem0/client/main.py +++ b/mem0/client/main.py @@ -2,13 +2,22 @@ import hashlib import logging import os import warnings -from typing import Any, Dict, List, Optional, Union +from typing import Any, Dict, List, Optional import httpx import requests from mem0.client.project import AsyncProject, Project +from mem0.client.types import ( + AddMemoryOptions, + DeleteAllMemoryOptions, + GetAllMemoryOptions, + ProjectUpdateOptions, + SearchMemoryOptions, + UpdateMemoryOptions, +) from mem0.client.utils import api_error_handler + # Exception classes are referenced in docstrings only from mem0.memory.setup import get_user_id, setup_config from mem0.memory.telemetry import capture_client_event @@ -31,8 +40,6 @@ class MemoryClient: api_key (str): The API key for authenticating with the Mem0 API. host (str): The base URL for the Mem0 API. client (httpx.Client): The HTTP client used for making API requests. - org_id (str, optional): Organization ID. - project_id (str, optional): Project ID. user_id (str): Unique identifier for the user. """ @@ -40,8 +47,6 @@ class MemoryClient: self, api_key: Optional[str] = None, host: Optional[str] = None, - org_id: Optional[str] = None, - project_id: Optional[str] = None, client: Optional[httpx.Client] = None, ): """Initialize the MemoryClient. @@ -52,8 +57,6 @@ class MemoryClient: environment variable. host: The base URL for the Mem0 API. Defaults to "https://api.mem0.ai". - org_id: The ID of the organization. - project_id: The ID of the project. client: A custom httpx.Client instance. If provided, it will be used instead of creating a new one. Note that base_url and headers will be set/overridden as needed. @@ -63,8 +66,8 @@ class MemoryClient: """ self.api_key = api_key or os.getenv("MEM0_API_KEY") self.host = host or "https://api.mem0.ai" - self.org_id = org_id - self.project_id = project_id + self.org_id = None + self.project_id = None self.user_id = get_user_id() if not self.api_key: @@ -128,15 +131,16 @@ class MemoryClient: raise ValueError(f"Error: {error_message}") @api_error_handler - def add(self, messages, **kwargs) -> Dict[str, Any]: + def add(self, messages, options: Optional[AddMemoryOptions] = None, **kwargs) -> Dict[str, Any]: """Add a new memory. Args: messages: A list of message dictionaries, a single message dictionary, or a string. If a string is provided, it will be converted to a user message. + options: Typed options for the add operation (AddMemoryOptions). **kwargs: Additional parameters such as user_id, agent_id, app_id, - metadata, filters, async_mode. + metadata, filters. Returns: A dictionary containing the API response in v1.1 format. @@ -149,24 +153,19 @@ class MemoryClient: NetworkError: If network connectivity issues occur. MemoryNotFoundError: If the memory doesn't exist (for updates/deletes). """ + kwargs = {**(options.model_dump(exclude_unset=True) if options else {}), **kwargs} # Handle different message input formats (align with OSS behavior) if isinstance(messages, str): messages = [{"role": "user", "content": messages}] elif isinstance(messages, dict): messages = [messages] elif not isinstance(messages, list): - raise ValueError( - f"messages must be str, dict, or list[dict], got {type(messages).__name__}" - ) - - kwargs = self._prepare_params(kwargs) - - # Set async_mode to True by default, but allow user override - if "async_mode" not in kwargs: - kwargs["async_mode"] = True + raise ValueError(f"messages must be str, dict, or list[dict], got {type(messages).__name__}") # Force v1.1 format for all add operations kwargs["output_format"] = "v1.1" + + kwargs = self._prepare_params(kwargs) payload = self._prepare_payload(messages, kwargs) response = self.client.post("/v1/memories/", json=payload) response.raise_for_status() @@ -200,10 +199,11 @@ class MemoryClient: return response.json() @api_error_handler - def get_all(self, **kwargs) -> Dict[str, Any]: + def get_all(self, options: Optional[GetAllMemoryOptions] = None, **kwargs) -> Dict[str, Any]: """Retrieve all memories, with optional filtering. Args: + options: Typed options for the get_all operation (GetAllMemoryOptions). **kwargs: Optional parameters for filtering (user_id, agent_id, app_id, top_k, page, page_size). @@ -218,8 +218,8 @@ class MemoryClient: NetworkError: If network connectivity issues occur. MemoryNotFoundError: If the memory doesn't exist (for updates/deletes). """ + kwargs = {**(options.model_dump(exclude_unset=True) if options else {}), **kwargs} params = self._prepare_params(kwargs) - params.pop("async_mode", None) if "page" in params and "page_size" in params: query_params = { @@ -236,7 +236,6 @@ class MemoryClient: "client.get_all", self, { - "api_version": "v2", "keys": list(kwargs.keys()), "sync_type": "sync", }, @@ -249,13 +248,13 @@ class MemoryClient: return result @api_error_handler - def search(self, query: str, **kwargs) -> Dict[str, Any]: + def search(self, query: str, options: Optional[SearchMemoryOptions] = None, **kwargs) -> Dict[str, Any]: """Search memories based on a query. Args: query: The search query string. - **kwargs: Additional parameters such as user_id, agent_id, app_id, - top_k, filters. + options: Typed options for the search operation (SearchMemoryOptions). + **kwargs: Additional parameters such as filters, top_k, rerank. Returns: A dictionary containing search results in v1.1 format: {"results": [...]} @@ -268,11 +267,9 @@ class MemoryClient: NetworkError: If network connectivity issues occur. MemoryNotFoundError: If the memory doesn't exist (for updates/deletes). """ - payload = {"query": query} + kwargs = {**(options.model_dump(exclude_unset=True) if options else {}), **kwargs} params = self._prepare_params(kwargs) - params.pop("async_mode", None) - - payload.update(params) + payload = {"query": query, **params} response = self.client.post("/v2/memories/search/", json=payload) response.raise_for_status() @@ -282,7 +279,6 @@ class MemoryClient: "client.search", self, { - "api_version": "v2", "keys": list(kwargs.keys()), "sync_type": "sync", }, @@ -298,36 +294,32 @@ class MemoryClient: def update( self, memory_id: str, - text: Optional[str] = None, - metadata: Optional[Dict[str, Any]] = None, - timestamp: Optional[Union[int, float, str]] = None, + options: Optional[UpdateMemoryOptions] = None, + **kwargs, ) -> Dict[str, Any]: - """ - Update a memory by ID. + """Update a memory by ID. Args: - memory_id (str): Memory ID. - text (str, optional): New content to update the memory with. - metadata (dict, optional): Metadata to update in the memory. - timestamp (int, float, or str, optional): Unix epoch timestamp or ISO 8601 string. + memory_id: The ID of the memory to update. + options: Typed options (UpdateMemoryOptions) with text, metadata, + and/or timestamp fields. + **kwargs: Alternatively pass text, metadata, timestamp as keyword args. Returns: Dict[str, Any]: The response from the server. - Example: - >>> client.update(memory_id="mem_123", text="Likes to play tennis on weekends") - >>> client.update(memory_id="mem_123", timestamp="2025-01-15T12:00:00Z") - """ - if text is None and metadata is None and timestamp is None: - raise ValueError("At least one of text, metadata, or timestamp must be provided for update.") + Raises: + ValueError: If none of text, metadata, or timestamp are provided. - payload = {} - if text is not None: - payload["text"] = text - if metadata is not None: - payload["metadata"] = metadata - if timestamp is not None: - payload["timestamp"] = timestamp + Example: + >>> client.update("mem_123", UpdateMemoryOptions(text="Updated text")) + >>> client.update("mem_123", text="Updated text") + """ + payload = {**(options.model_dump(exclude_unset=True) if options else {}), **kwargs} + payload = {k: v for k, v in payload.items() if v is not None} + + if not payload: + raise ValueError("At least one of text, metadata, or timestamp must be provided for update.") capture_client_event("client.update", self, {"memory_id": memory_id, "sync_type": "sync"}) params = self._prepare_params() @@ -360,10 +352,11 @@ class MemoryClient: return response.json() @api_error_handler - def delete_all(self, **kwargs) -> Dict[str, str]: + def delete_all(self, options: Optional[DeleteAllMemoryOptions] = None, **kwargs) -> Dict[str, str]: """Delete all memories, with optional filtering. Args: + options: Typed options for the delete_all operation (DeleteAllMemoryOptions). **kwargs: Optional parameters for filtering (user_id, agent_id, app_id). @@ -378,6 +371,7 @@ class MemoryClient: NetworkError: If network connectivity issues occur. MemoryNotFoundError: If the memory doesn't exist (for updates/deletes). """ + kwargs = {**(options.model_dump(exclude_unset=True) if options else {}), **kwargs} params = self._prepare_params(kwargs) response = self.client.delete("/v1/memories/", params=params) response.raise_for_status() @@ -667,13 +661,11 @@ class MemoryClient: @api_error_handler def update_project( self, + options: Optional[ProjectUpdateOptions] = None, custom_instructions: Optional[str] = None, custom_categories: Optional[List[str]] = None, retrieval_criteria: Optional[List[Dict[str, Any]]] = None, enable_graph: Optional[bool] = None, - version: Optional[str] = None, - inclusion_prompt: Optional[str] = None, - exclusion_prompt: Optional[str] = None, memory_depth: Optional[str] = None, usecase_setting: Optional[str] = None, multilingual: Optional[bool] = None, @@ -681,67 +673,53 @@ class MemoryClient: """Update the project settings. Args: - custom_instructions: New instructions for the project - custom_categories: New categories for the project - retrieval_criteria: New retrieval criteria for the project - enable_graph: Enable or disable the graph for the project - version: Version of the project - inclusion_prompt: Inclusion prompt for the project - exclusion_prompt: Exclusion prompt for the project - memory_depth: Memory depth for the project - usecase_setting: Usecase setting for the project - multilingual: Whether to use the input language for memory storage and retrieval + options: Typed options for the update operation (ProjectUpdateOptions). + custom_instructions: New instructions for the project. + custom_categories: New categories for the project. + retrieval_criteria: New retrieval criteria for the project. + enable_graph: Enable or disable the graph for the project. + memory_depth: Memory depth for the project. + usecase_setting: Usecase setting for the project. + multilingual: Whether to use the input language for memory storage and retrieval. Returns: Dictionary containing the API response. Raises: - ValidationError: If the input data is invalid. - AuthenticationError: If authentication fails. - RateLimitError: If rate limits are exceeded. - MemoryQuotaExceededError: If memory quota is exceeded. - NetworkError: If network connectivity issues occur. - MemoryNotFoundError: If the memory doesn't exist (for updates/deletes). - ValueError: If org_id or project_id are not set. + ValueError: If org_id or project_id are not set, or no update fields provided. """ logger.warning( - "update_project() method is going to be deprecated in version v1.0 of the package. Please use the client.project.update() method instead." + "update_project() method is going to be deprecated in version v1.0 of the package. " + "Please use the client.project.update() method instead." ) if not (self.org_id and self.project_id): raise ValueError("org_id and project_id must be set to update instructions or categories") - if ( - custom_instructions is None - and custom_categories is None - and retrieval_criteria is None - and enable_graph is None - and version is None - and inclusion_prompt is None - and exclusion_prompt is None - and memory_depth is None - and usecase_setting is None - and multilingual is None - ): + kwargs = { + **(options.model_dump(exclude_unset=True) if options else {}), + **{ + k: v + for k, v in { + "custom_instructions": custom_instructions, + "custom_categories": custom_categories, + "retrieval_criteria": retrieval_criteria, + "enable_graph": enable_graph, + "memory_depth": memory_depth, + "usecase_setting": usecase_setting, + "multilingual": multilingual, + }.items() + if v is not None + }, + } + + if not kwargs: raise ValueError( "Currently we only support updating custom_instructions or " "custom_categories or retrieval_criteria, so you must " "provide at least one of them" ) - payload = self._prepare_params( - { - "custom_instructions": custom_instructions, - "custom_categories": custom_categories, - "retrieval_criteria": retrieval_criteria, - "enable_graph": enable_graph, - "version": version, - "inclusion_prompt": inclusion_prompt, - "exclusion_prompt": exclusion_prompt, - "memory_depth": memory_depth, - "usecase_setting": usecase_setting, - "multilingual": multilingual, - } - ) + payload = self._prepare_params(kwargs) response = self.client.patch( f"/api/v1/orgs/organizations/{self.org_id}/projects/{self.project_id}/", json=payload, @@ -750,19 +728,7 @@ class MemoryClient: capture_client_event( "client.update_project", self, - { - "custom_instructions": custom_instructions, - "custom_categories": custom_categories, - "retrieval_criteria": retrieval_criteria, - "enable_graph": enable_graph, - "version": version, - "inclusion_prompt": inclusion_prompt, - "exclusion_prompt": exclusion_prompt, - "memory_depth": memory_depth, - "usecase_setting": usecase_setting, - "multilingual": multilingual, - "sync_type": "sync", - }, + {**kwargs, "sync_type": "sync"}, ) return response.json() @@ -937,20 +903,12 @@ class MemoryClient: Returns: A dictionary containing the prepared parameters. - - Raises: - ValueError: If either org_id or project_id is provided but not both. """ if kwargs is None: kwargs = {} - # Add org_id and project_id if both are available - if self.org_id and self.project_id: - kwargs["org_id"] = self.org_id - kwargs["project_id"] = self.project_id - elif self.org_id or self.project_id: - raise ValueError("Please provide both org_id and project_id") + # org_id and project_id are resolved from API key — not injected into params return {k: v for k, v in kwargs.items() if v is not None} @@ -966,8 +924,6 @@ class AsyncMemoryClient: self, api_key: Optional[str] = None, host: Optional[str] = None, - org_id: Optional[str] = None, - project_id: Optional[str] = None, client: Optional[httpx.AsyncClient] = None, ): """Initialize the AsyncMemoryClient. @@ -978,8 +934,6 @@ class AsyncMemoryClient: environment variable. host: The base URL for the Mem0 API. Defaults to "https://api.mem0.ai". - org_id: The ID of the organization. - project_id: The ID of the project. client: A custom httpx.AsyncClient instance. If provided, it will be used instead of creating a new one. Note that base_url and headers will be set/overridden as needed. @@ -989,8 +943,8 @@ class AsyncMemoryClient: """ self.api_key = api_key or os.getenv("MEM0_API_KEY") self.host = host or "https://api.mem0.ai" - self.org_id = org_id - self.project_id = project_id + self.org_id = None + self.project_id = None self.user_id = get_user_id() if not self.api_key: @@ -1085,20 +1039,12 @@ class AsyncMemoryClient: Returns: A dictionary containing the prepared parameters. - - Raises: - ValueError: If either org_id or project_id is provided but not both. """ if kwargs is None: kwargs = {} - # Add org_id and project_id if both are available - if self.org_id and self.project_id: - kwargs["org_id"] = self.org_id - kwargs["project_id"] = self.project_id - elif self.org_id or self.project_id: - raise ValueError("Please provide both org_id and project_id") + # org_id and project_id are resolved from API key — not injected into params return {k: v for k, v in kwargs.items() if v is not None} @@ -1109,25 +1055,41 @@ class AsyncMemoryClient: await self.async_client.aclose() @api_error_handler - async def add(self, messages, **kwargs) -> Dict[str, Any]: + async def add(self, messages, options: Optional[AddMemoryOptions] = None, **kwargs) -> Dict[str, Any]: + """Add a new memory. + + Args: + messages: A list of message dictionaries, a single message dictionary, + or a string. If a string is provided, it will be converted to + a user message. + options: Typed options for the add operation (AddMemoryOptions). + **kwargs: Additional parameters such as user_id, agent_id, app_id, + metadata, filters. + + Returns: + A dictionary containing the API response in v1.1 format. + + Raises: + ValidationError: If the input data is invalid. + AuthenticationError: If authentication fails. + RateLimitError: If rate limits are exceeded. + MemoryQuotaExceededError: If memory quota is exceeded. + NetworkError: If network connectivity issues occur. + MemoryNotFoundError: If the memory doesn't exist (for updates/deletes). + """ + kwargs = {**(options.model_dump(exclude_unset=True) if options else {}), **kwargs} # Handle different message input formats (align with OSS behavior) if isinstance(messages, str): messages = [{"role": "user", "content": messages}] elif isinstance(messages, dict): messages = [messages] elif not isinstance(messages, list): - raise ValueError( - f"messages must be str, dict, or list[dict], got {type(messages).__name__}" - ) - - kwargs = self._prepare_params(kwargs) - - # Set async_mode to True by default, but allow user override - if "async_mode" not in kwargs: - kwargs["async_mode"] = True + raise ValueError(f"messages must be str, dict, or list[dict], got {type(messages).__name__}") # Force v1.1 format for all add operations kwargs["output_format"] = "v1.1" + + kwargs = self._prepare_params(kwargs) payload = self._prepare_payload(messages, kwargs) response = await self.async_client.post("/v1/memories/", json=payload) response.raise_for_status() @@ -1145,9 +1107,26 @@ class AsyncMemoryClient: return response.json() @api_error_handler - async def get_all(self, **kwargs) -> Dict[str, Any]: + async def get_all(self, options: Optional[GetAllMemoryOptions] = None, **kwargs) -> Dict[str, Any]: + """Retrieve all memories, with optional filtering. + + Args: + options: Typed options for the get_all operation (GetAllMemoryOptions). + **kwargs: Optional parameters for filtering (filters, page, page_size). + + Returns: + A dictionary containing memories in v1.1 format: {"results": [...]} + + Raises: + ValidationError: If the input data is invalid. + AuthenticationError: If authentication fails. + RateLimitError: If rate limits are exceeded. + MemoryQuotaExceededError: If memory quota is exceeded. + NetworkError: If network connectivity issues occur. + MemoryNotFoundError: If the memory doesn't exist (for updates/deletes). + """ + kwargs = {**(options.model_dump(exclude_unset=True) if options else {}), **kwargs} params = self._prepare_params(kwargs) - params.pop("async_mode", None) if "page" in params and "page_size" in params: query_params = { @@ -1164,7 +1143,6 @@ class AsyncMemoryClient: "client.get_all", self, { - "api_version": "v2", "keys": list(kwargs.keys()), "sync_type": "async", }, @@ -1177,12 +1155,28 @@ class AsyncMemoryClient: return result @api_error_handler - async def search(self, query: str, **kwargs) -> Dict[str, Any]: - payload = {"query": query} - params = self._prepare_params(kwargs) - params.pop("async_mode", None) + async def search(self, query: str, options: Optional[SearchMemoryOptions] = None, **kwargs) -> Dict[str, Any]: + """Search memories based on a query. - payload.update(params) + Args: + query: The search query string. + options: Typed options for the search operation (SearchMemoryOptions). + **kwargs: Additional parameters such as filters, top_k, rerank. + + Returns: + A dictionary containing search results in v1.1 format: {"results": [...]} + + Raises: + ValidationError: If the input data is invalid. + AuthenticationError: If authentication fails. + RateLimitError: If rate limits are exceeded. + MemoryQuotaExceededError: If memory quota is exceeded. + NetworkError: If network connectivity issues occur. + MemoryNotFoundError: If the memory doesn't exist (for updates/deletes). + """ + kwargs = {**(options.model_dump(exclude_unset=True) if options else {}), **kwargs} + params = self._prepare_params(kwargs) + payload = {"query": query, **params} response = await self.async_client.post("/v2/memories/search/", json=payload) response.raise_for_status() @@ -1192,7 +1186,6 @@ class AsyncMemoryClient: "client.search", self, { - "api_version": "v2", "keys": list(kwargs.keys()), "sync_type": "async", }, @@ -1208,36 +1201,32 @@ class AsyncMemoryClient: async def update( self, memory_id: str, - text: Optional[str] = None, - metadata: Optional[Dict[str, Any]] = None, - timestamp: Optional[Union[int, float, str]] = None, + options: Optional[UpdateMemoryOptions] = None, + **kwargs, ) -> Dict[str, Any]: - """ - Update a memory by ID asynchronously. + """Update a memory by ID asynchronously. Args: - memory_id (str): Memory ID. - text (str, optional): New content to update the memory with. - metadata (dict, optional): Metadata to update in the memory. - timestamp (int, float, or str, optional): Unix epoch timestamp or ISO 8601 string. + memory_id: The ID of the memory to update. + options: Typed options (UpdateMemoryOptions) with text, metadata, + and/or timestamp fields. + **kwargs: Alternatively pass text, metadata, timestamp as keyword args. Returns: Dict[str, Any]: The response from the server. - Example: - >>> await client.update(memory_id="mem_123", text="Likes to play tennis on weekends") - >>> await client.update(memory_id="mem_123", timestamp="2025-01-15T12:00:00Z") - """ - if text is None and metadata is None and timestamp is None: - raise ValueError("At least one of text, metadata, or timestamp must be provided for update.") + Raises: + ValueError: If none of text, metadata, or timestamp are provided. - payload = {} - if text is not None: - payload["text"] = text - if metadata is not None: - payload["metadata"] = metadata - if timestamp is not None: - payload["timestamp"] = timestamp + Example: + >>> await client.update("mem_123", UpdateMemoryOptions(text="Updated text")) + >>> await client.update("mem_123", text="Updated text") + """ + payload = {**(options.model_dump(exclude_unset=True) if options else {}), **kwargs} + payload = {k: v for k, v in payload.items() if v is not None} + + if not payload: + raise ValueError("At least one of text, metadata, or timestamp must be provided for update.") capture_client_event("client.update", self, {"memory_id": memory_id, "sync_type": "async"}) params = self._prepare_params() @@ -1270,10 +1259,11 @@ class AsyncMemoryClient: return response.json() @api_error_handler - async def delete_all(self, **kwargs) -> Dict[str, str]: + async def delete_all(self, options: Optional[DeleteAllMemoryOptions] = None, **kwargs) -> Dict[str, str]: """Delete all memories, with optional filtering. Args: + options: Typed options for the delete_all operation (DeleteAllMemoryOptions). **kwargs: Optional parameters for filtering (user_id, agent_id, app_id). Returns: @@ -1287,6 +1277,7 @@ class AsyncMemoryClient: NetworkError: If network connectivity issues occur. MemoryNotFoundError: If the memory doesn't exist (for updates/deletes). """ + kwargs = {**(options.model_dump(exclude_unset=True) if options else {}), **kwargs} params = self._prepare_params(kwargs) response = await self.async_client.delete("/v1/memories/", params=params) response.raise_for_status() @@ -1554,63 +1545,65 @@ class AsyncMemoryClient: @api_error_handler async def update_project( self, + options: Optional[ProjectUpdateOptions] = None, custom_instructions: Optional[str] = None, custom_categories: Optional[List[str]] = None, retrieval_criteria: Optional[List[Dict[str, Any]]] = None, enable_graph: Optional[bool] = None, - version: Optional[str] = None, + memory_depth: Optional[str] = None, + usecase_setting: Optional[str] = None, multilingual: Optional[bool] = None, ) -> Dict[str, Any]: """Update the project settings. Args: - custom_instructions: New instructions for the project - custom_categories: New categories for the project - retrieval_criteria: New retrieval criteria for the project - enable_graph: Enable or disable the graph for the project - version: Version of the project - multilingual: Whether to use the input language for memory storage and retrieval + options: Typed options for the update operation (ProjectUpdateOptions). + custom_instructions: New instructions for the project. + custom_categories: New categories for the project. + retrieval_criteria: New retrieval criteria for the project. + enable_graph: Enable or disable the graph for the project. + memory_depth: Memory depth for the project. + usecase_setting: Usecase setting for the project. + multilingual: Whether to use the input language for memory storage and retrieval. Returns: Dictionary containing the API response. Raises: - ValidationError: If the input data is invalid. - AuthenticationError: If authentication fails. - RateLimitError: If rate limits are exceeded. - MemoryQuotaExceededError: If memory quota is exceeded. - NetworkError: If network connectivity issues occur. - MemoryNotFoundError: If the memory doesn't exist (for updates/deletes). - ValueError: If org_id or project_id are not set. + ValueError: If org_id or project_id are not set, or no update fields provided. """ logger.warning( - "update_project() method is going to be deprecated in version v1.0 of the package. Please use the client.project.update() method instead." + "update_project() method is going to be deprecated in version v1.0 of the package. " + "Please use the client.project.update() method instead." ) if not (self.org_id and self.project_id): raise ValueError("org_id and project_id must be set to update instructions or categories") - if ( - custom_instructions is None - and custom_categories is None - and retrieval_criteria is None - and enable_graph is None - and version is None - and multilingual is None - ): + kwargs = { + **(options.model_dump(exclude_unset=True) if options else {}), + **{ + k: v + for k, v in { + "custom_instructions": custom_instructions, + "custom_categories": custom_categories, + "retrieval_criteria": retrieval_criteria, + "enable_graph": enable_graph, + "memory_depth": memory_depth, + "usecase_setting": usecase_setting, + "multilingual": multilingual, + }.items() + if v is not None + }, + } + + if not kwargs: raise ValueError( - "Currently we only support updating custom_instructions or custom_categories or retrieval_criteria, so you must provide at least one of them" + "Currently we only support updating custom_instructions or " + "custom_categories or retrieval_criteria, so you must " + "provide at least one of them" ) - payload = self._prepare_params( - { - "custom_instructions": custom_instructions, - "custom_categories": custom_categories, - "retrieval_criteria": retrieval_criteria, - "enable_graph": enable_graph, - "version": version, - "multilingual": multilingual, - } - ) + payload = self._prepare_params(kwargs) response = await self.async_client.patch( f"/api/v1/orgs/organizations/{self.org_id}/projects/{self.project_id}/", json=payload, @@ -1619,15 +1612,7 @@ class AsyncMemoryClient: capture_client_event( "client.update_project", self, - { - "custom_instructions": custom_instructions, - "custom_categories": custom_categories, - "retrieval_criteria": retrieval_criteria, - "enable_graph": enable_graph, - "version": version, - "multilingual": multilingual, - "sync_type": "async", - }, + {**kwargs, "sync_type": "async"}, ) return response.json() diff --git a/mem0/client/types.py b/mem0/client/types.py new file mode 100644 index 000000000..e55b78713 --- /dev/null +++ b/mem0/client/types.py @@ -0,0 +1,103 @@ +"""Pydantic option models for MemoryClient methods. + +These models provide IDE autocompletion, runtime validation, and type safety. +Methods accept both typed options and **kwargs for backward compatibility. +""" + +from typing import Any, Dict, List, Optional, Union + +from pydantic import BaseModel, Field + + +class EntityOptions(BaseModel): + """Identity options for add/delete operations (top-level entity IDs).""" + + user_id: Optional[str] = Field(default=None, description="The user ID to associate with the memory") + agent_id: Optional[str] = Field(default=None, description="The agent ID to associate with the memory") + app_id: Optional[str] = Field(default=None, description="The app ID to associate with the memory") + run_id: Optional[str] = Field(default=None, description="The run ID to associate with the memory") + + +class AddMemoryOptions(EntityOptions): + """Options for the add() method.""" + + metadata: Optional[Dict[str, Any]] = Field(default=None, description="Additional metadata for the memory") + infer: Optional[bool] = Field(default=None, description="Whether to infer memories from the input") + custom_categories: Optional[List[Dict[str, Any]]] = Field( + default=None, description="Custom categories for memory classification" + ) + custom_instructions: Optional[str] = Field(default=None, description="Custom instructions for fact extraction") + timestamp: Optional[int] = Field(default=None, description="Unix timestamp for the memory") + structured_data_schema: Optional[Dict[str, Any]] = Field( + default=None, description="Schema for structured data extraction" + ) + enable_graph: Optional[bool] = Field(default=None, description="Whether to enable graph memory for this operation") + + +class SearchMemoryOptions(BaseModel): + """Options for the search() method. + + Identity fields (user_id, agent_id, etc.) must be passed inside the + ``filters`` dict — the v2 API does not accept them at the top level. + """ + + filters: Optional[Dict[str, Any]] = Field( + default=None, description="Filters for the search (e.g. {'user_id': '...'})" + ) + metadata: Optional[Dict[str, Any]] = Field(default=None, description="Additional metadata for the search") + top_k: Optional[int] = Field(default=None, description="Number of results to return") + rerank: Optional[bool] = Field(default=None, description="Whether to rerank results") + threshold: Optional[float] = Field(default=None, description="Minimum similarity score threshold") + fields: Optional[List[str]] = Field(default=None, description="Fields to include in the response") + categories: Optional[List[str]] = Field(default=None, description="Categories to filter by") + enable_graph: Optional[bool] = Field(default=None, description="Whether to enable graph memory for this search") + + +class GetAllMemoryOptions(BaseModel): + """Options for the get_all() method. + + Identity fields (user_id, agent_id, etc.) must be passed inside the + ``filters`` dict — the v2 API does not accept them at the top level. + """ + + filters: Optional[Dict[str, Any]] = Field( + default=None, description="Filters for retrieval (e.g. {'user_id': '...'})" + ) + page: Optional[int] = Field(default=None, description="Page number for pagination") + page_size: Optional[int] = Field(default=None, description="Number of items per page") + start_date: Optional[str] = Field( + default=None, description="Filter memories created on or after this date (ISO 8601)" + ) + end_date: Optional[str] = Field( + default=None, description="Filter memories created on or before this date (ISO 8601)" + ) + categories: Optional[List[str]] = Field(default=None, description="Categories to filter by") + enable_graph: Optional[bool] = Field(default=None, description="Whether to enable graph memory for retrieval") + + +class DeleteAllMemoryOptions(EntityOptions): + """Options for the delete_all() method.""" + + pass + + +class UpdateMemoryOptions(BaseModel): + """Options for the update() method.""" + + text: Optional[str] = Field(default=None, description="New text content for the memory") + metadata: Optional[Dict[str, Any]] = Field(default=None, description="Updated metadata") + timestamp: Optional[Union[int, float, str]] = Field(default=None, description="Updated timestamp") + + +class ProjectUpdateOptions(BaseModel): + """Options for project update operations.""" + + custom_instructions: Optional[str] = Field(default=None, description="Custom instructions for fact extraction") + custom_categories: Optional[List[Dict[str, Any]]] = Field( + default=None, description="Custom categories for classification" + ) + enable_graph: Optional[bool] = Field(default=None, description="Whether to enable graph memory") + memory_depth: Optional[str] = Field(default=None, description="Memory depth configuration") + usecase_setting: Optional[Any] = Field(default=None, description="Use case specific settings") + multilingual: Optional[bool] = Field(default=None, description="Whether to enable multilingual support") + retrieval_criteria: Optional[List[Any]] = Field(default=None, description="Criteria for memory retrieval") diff --git a/mem0/configs/base.py b/mem0/configs/base.py index dd0dd9df4..9c379bb95 100644 --- a/mem0/configs/base.py +++ b/mem0/configs/base.py @@ -3,11 +3,11 @@ from typing import Any, Dict, Optional from pydantic import BaseModel, Field +from mem0.configs.rerankers.config import RerankerConfig from mem0.embeddings.configs import EmbedderConfig from mem0.graphs.configs import GraphStoreConfig from mem0.llms.configs import LlmConfig from mem0.vector_stores.configs import VectorStoreConfig -from mem0.configs.rerankers.config import RerankerConfig # Set up the directory path home_dir = os.path.expanduser("~") @@ -56,8 +56,8 @@ class MemoryConfig(BaseModel): description="The version of the API", default="v1.1", ) - custom_fact_extraction_prompt: Optional[str] = Field( - description="Custom prompt for the fact extraction", + custom_instructions: Optional[str] = Field( + description="Custom instructions for fact extraction", default=None, ) custom_update_memory_prompt: Optional[str] = Field( diff --git a/mem0/configs/vector_stores/elasticsearch.py b/mem0/configs/vector_stores/elasticsearch.py index 6044383cc..0e9faac3f 100644 --- a/mem0/configs/vector_stores/elasticsearch.py +++ b/mem0/configs/vector_stores/elasticsearch.py @@ -17,7 +17,7 @@ class ElasticsearchConfig(BaseModel): use_ssl: bool = Field(True, description="Use SSL for connection") auto_create_index: bool = Field(True, description="Automatically create index during initialization") custom_search_query: Optional[Callable[[List[float], int, Optional[Dict]], Dict]] = Field( - None, description="Custom search query function. Parameters: (query, limit, filters) -> Dict" + None, description="Custom search query function. Parameters: (query, top_k, filters) -> Dict" ) headers: Optional[Dict[str, str]] = Field(None, description="Custom headers to include in requests") diff --git a/mem0/graphs/neptune/base.py b/mem0/graphs/neptune/base.py index 59aa9f8e6..20499deaf 100644 --- a/mem0/graphs/neptune/base.py +++ b/mem0/graphs/neptune/base.py @@ -346,14 +346,14 @@ class NeptuneBase(ABC): ): pass - def search(self, query, filters, limit=100): + def search(self, query, filters, top_k=100): """ Search for memories and related graph data. Args: query (str): Query to search for. filters (dict): A dictionary containing filters to be applied during the search. - limit (int): The maximum number of nodes and relationships to retrieve. Defaults to 100. + top_k (int): The maximum number of nodes and relationships to retrieve. Defaults to 100. Returns: dict: A dictionary containing: @@ -438,13 +438,13 @@ class NeptuneBase(ABC): """ pass - def get_all(self, filters, limit=100): + def get_all(self, filters, top_k=100): """ Retrieves all nodes and relationships from the graph database based on filtering criteria. Args: filters (dict): A dictionary containing filters to be applied during the retrieval. - limit (int): The maximum number of nodes and relationships to retrieve. Defaults to 100. + top_k (int): The maximum number of nodes and relationships to retrieve. Defaults to 100. Returns: list: A list of dictionaries, each containing: - 'contexts': The base data store response for each memory. @@ -452,7 +452,7 @@ class NeptuneBase(ABC): """ # return all nodes and relationships - query, params = self._get_all_cypher(filters, limit) + query, params = self._get_all_cypher(filters, top_k) results = self.graph.query(query, params=params) final_results = [] @@ -470,13 +470,13 @@ class NeptuneBase(ABC): return final_results @abstractmethod - def _get_all_cypher(self, filters, limit): + def _get_all_cypher(self, filters, top_k): """ Returns the OpenCypher query and parameters to get all edges/nodes in the memory store """ pass - def _search_graph_db(self, node_list, filters, limit=100): + def _search_graph_db(self, node_list, filters, top_k=100): """ Search similar nodes among and their respective incoming and outgoing relations. """ @@ -484,14 +484,14 @@ class NeptuneBase(ABC): for node in node_list: n_embedding = self.embedding_model.embed(node) - cypher_query, params = self._search_graph_db_cypher(n_embedding, filters, limit) + cypher_query, params = self._search_graph_db_cypher(n_embedding, filters, top_k) ans = self.graph.query(cypher_query, params=params) result_relations.extend(ans) return result_relations @abstractmethod - def _search_graph_db_cypher(self, n_embedding, filters, limit): + def _search_graph_db_cypher(self, n_embedding, filters, top_k): """ Returns the OpenCypher query and parameters to search for similar nodes in the memory store """ diff --git a/mem0/graphs/neptune/neptunedb.py b/mem0/graphs/neptune/neptunedb.py index e5ebf87a0..44a20aa7b 100644 --- a/mem0/graphs/neptune/neptunedb.py +++ b/mem0/graphs/neptune/neptunedb.py @@ -380,7 +380,7 @@ class MemoryGraph(NeptuneBase): source_nodes = self.vector_store.search( query="", vectors=source_embedding, - limit=self.vector_store_limit, + top_k=self.vector_store_limit, filters={"user_id": user_id}, ) @@ -413,7 +413,7 @@ class MemoryGraph(NeptuneBase): destination_nodes = self.vector_store.search( query="", vectors=destination_embedding, - limit=self.vector_store_limit, + top_k=self.vector_store_limit, filters={"user_id": user_id}, ) @@ -455,12 +455,12 @@ class MemoryGraph(NeptuneBase): logger.debug(f"delete_all query={cypher}") return cypher, params - def _get_all_cypher(self, filters, limit): + def _get_all_cypher(self, filters, top_k): """ Returns the OpenCypher query and parameters to get all edges/nodes in the memory store :param filters: search filters - :param limit: return limit + :param top_k: return limit :return: str, dict """ @@ -469,16 +469,16 @@ class MemoryGraph(NeptuneBase): RETURN n.name AS source, type(r) AS relationship, m.name AS target LIMIT $limit """ - params = {"user_id": filters["user_id"], "limit": limit} + params = {"user_id": filters["user_id"], "limit": top_k} return cypher, params - def _search_graph_db_cypher(self, n_embedding, filters, limit): + def _search_graph_db_cypher(self, n_embedding, filters, top_k): """ Returns the OpenCypher query and parameters to search for similar nodes in the memory store :param n_embedding: node vector :param filters: search filters - :param limit: return limit + :param top_k: return limit :return: str, dict """ @@ -486,7 +486,7 @@ class MemoryGraph(NeptuneBase): search_nodes = self.vector_store.search( query="", vectors=n_embedding, - limit=self.vector_store_limit, + top_k=self.vector_store_limit, filters=filters, ) @@ -504,7 +504,7 @@ class MemoryGraph(NeptuneBase): params = { "n_ids": ids, "user_id": filters["user_id"], - "limit": limit, + "limit": top_k, } logger.debug(f"_search_graph_db\n query={cypher_query}") diff --git a/mem0/graphs/neptune/neptunegraph.py b/mem0/graphs/neptune/neptunegraph.py index 866ed372f..d458b436c 100644 --- a/mem0/graphs/neptune/neptunegraph.py +++ b/mem0/graphs/neptune/neptunegraph.py @@ -3,8 +3,8 @@ import logging from .base import NeptuneBase try: - from langchain_aws import NeptuneAnalyticsGraph from botocore.config import Config + from langchain_aws import NeptuneAnalyticsGraph except ImportError: raise ImportError("langchain_aws is not installed. Please install it using 'make install_all'.") @@ -412,12 +412,12 @@ class MemoryGraph(NeptuneBase): logger.debug(f"delete_all query={cypher}") return cypher, params - def _get_all_cypher(self, filters, limit): + def _get_all_cypher(self, filters, top_k): """ Returns the OpenCypher query and parameters to get all edges/nodes in the memory store :param filters: search filters - :param limit: return limit + :param top_k: return limit :return: str, dict """ @@ -426,16 +426,16 @@ class MemoryGraph(NeptuneBase): RETURN n.name AS source, type(r) AS relationship, m.name AS target LIMIT $limit """ - params = {"user_id": filters["user_id"], "limit": limit} + params = {"user_id": filters["user_id"], "limit": top_k} return cypher, params - def _search_graph_db_cypher(self, n_embedding, filters, limit): + def _search_graph_db_cypher(self, n_embedding, filters, top_k): """ Returns the OpenCypher query and parameters to search for similar nodes in the memory store :param n_embedding: node vector :param filters: search filters - :param limit: return limit + :param top_k: return limit :return: str, dict """ @@ -468,7 +468,7 @@ class MemoryGraph(NeptuneBase): "n_embedding": n_embedding, "threshold": self.threshold, "user_id": filters["user_id"], - "limit": limit, + "limit": top_k, } logger.debug(f"_search_graph_db\n query={cypher_query}") diff --git a/mem0/memory/apache_age_memory.py b/mem0/memory/apache_age_memory.py index e3c69f1a4..721345533 100644 --- a/mem0/memory/apache_age_memory.py +++ b/mem0/memory/apache_age_memory.py @@ -213,14 +213,14 @@ class MemoryGraph: return {"deleted_entities": deleted_entities, "added_entities": added_entities} - def search(self, query, filters, limit=100): + def search(self, query, filters, top_k=100): """ Search for memories and related graph data. Args: query (str): Query to search for. filters (dict): A dictionary containing filters to be applied during the search. - limit (int): The maximum number of nodes and relationships to retrieve. Defaults to 100. + top_k (int): The maximum number of nodes and relationships to retrieve. Defaults to 100. Returns: list: A list of dicts with keys "source", "relationship", "destination". @@ -286,13 +286,13 @@ class MemoryGraph: ) self.ag.commit() - def get_all(self, filters, limit=100): + def get_all(self, filters, top_k=100): """ Retrieves all nodes and relationships from the graph database based on optional filtering criteria. Args: filters (dict): A dictionary containing filters to be applied during the retrieval. - limit (int): The maximum number of nodes and relationships to retrieve. Defaults to 100. + top_k (int): The maximum number of nodes and relationships to retrieve. Defaults to 100. Returns: list: A list of dictionaries, each containing: - 'source': The source node name. @@ -308,7 +308,7 @@ class MemoryGraph: where_parts.extend(["n.run_id = %s", "m.run_id = %s"]) params.extend([filters["run_id"], filters["run_id"]]) where_clause = " AND ".join(where_parts) - params.append(limit) + params.append(top_k) results = self._exec_cypher( f"MATCH (n)-[r]->(m) WHERE {where_clause} " @@ -408,7 +408,7 @@ class MemoryGraph: # -- graph DB operations --------------------------------------------------- - def _search_graph_db(self, node_list, filters, limit=100): + def _search_graph_db(self, node_list, filters, top_k=100): """Search similar nodes and their respective incoming and outgoing relations.""" result_relations = [] @@ -430,7 +430,7 @@ class MemoryGraph: rel_where = " AND ".join(rel_where_parts) # For each similar node, fetch its relationships - for sn in similar_nodes[:limit]: + for sn in similar_nodes[:top_k]: node_name = sn["name"] similarity = sn["similarity"] diff --git a/mem0/memory/graph_memory.py b/mem0/memory/graph_memory.py index 80a3b905a..2dabc2392 100644 --- a/mem0/memory/graph_memory.py +++ b/mem0/memory/graph_memory.py @@ -93,14 +93,14 @@ class MemoryGraph: return {"deleted_entities": deleted_entities, "added_entities": added_entities} - def search(self, query, filters, limit=100): + def search(self, query, filters, top_k=100): """ Search for memories and related graph data. Args: query (str): Query to search for. filters (dict): A dictionary containing filters to be applied during the search. - limit (int): The maximum number of nodes and relationships to retrieve. Defaults to 100. + top_k (int): The maximum number of nodes and relationships to retrieve. Defaults to 100. Returns: dict: A dictionary containing: @@ -171,18 +171,18 @@ class MemoryGraph: params["run_id"] = filters["run_id"] self.graph.query(cypher, params=params) - def get_all(self, filters, limit=100): + def get_all(self, filters, top_k=100): """ Retrieves all nodes and relationships from the graph database based on optional filtering criteria. Args: filters (dict): A dictionary containing filters to be applied during the retrieval. - limit (int): The maximum number of nodes and relationships to retrieve. Defaults to 100. + top_k (int): The maximum number of nodes and relationships to retrieve. Defaults to 100. Returns: list: A list of dictionaries, each containing: - 'contexts': The base data store response for each memory. - 'entities': A list of strings representing the nodes and relationships """ - params = {"user_id": filters["user_id"], "limit": limit} + params = {"user_id": filters["user_id"], "limit": top_k} # Build node properties based on filters node_props = ["user_id: $user_id"] @@ -291,7 +291,7 @@ class MemoryGraph: logger.debug(f"Extracted entities: {entities}") return entities - def _search_graph_db(self, node_list, filters, limit=100): + def _search_graph_db(self, node_list, filters, top_k=100): """Search similar nodes among and their respective incoming and outgoing relations.""" result_relations = [] @@ -332,7 +332,7 @@ class MemoryGraph: "n_embedding": n_embedding, "threshold": self.threshold, "user_id": filters["user_id"], - "limit": limit, + "limit": top_k, } if filters.get("agent_id"): params["agent_id"] = filters["agent_id"] diff --git a/mem0/memory/kuzu_memory.py b/mem0/memory/kuzu_memory.py index 0a9f1a4a1..d798bf6c4 100644 --- a/mem0/memory/kuzu_memory.py +++ b/mem0/memory/kuzu_memory.py @@ -113,14 +113,14 @@ class MemoryGraph: return {"deleted_entities": deleted_entities, "added_entities": added_entities} - def search(self, query, filters, limit=5): + def search(self, query, filters, top_k=5): """ Search for memories and related graph data. Args: query (str): Query to search for. filters (dict): A dictionary containing filters to be applied during the search. - limit (int): The maximum number of nodes and relationships to retrieve. Defaults to 100. + top_k (int): The maximum number of nodes and relationships to retrieve. Defaults to 100. Returns: dict: A dictionary containing: @@ -139,7 +139,7 @@ class MemoryGraph: bm25 = BM25Okapi(search_outputs_sequence) tokenized_query = query.split(" ") - reranked_results = bm25.get_top_n(tokenized_query, search_outputs_sequence, n=limit) + reranked_results = bm25.get_top_n(tokenized_query, search_outputs_sequence, n=top_k) search_results = [] for item in reranked_results: @@ -191,12 +191,12 @@ class MemoryGraph: params["run_id"] = filters["run_id"] self.kuzu_execute(cypher, parameters=params) - def get_all(self, filters, limit=100): + def get_all(self, filters, top_k=100): """ Retrieves all nodes and relationships from the graph database based on optional filtering criteria. Args: filters (dict): A dictionary containing filters to be applied during the retrieval. - limit (int): The maximum number of nodes and relationships to retrieve. Defaults to 100. + top_k (int): The maximum number of nodes and relationships to retrieve. Defaults to 100. Returns: list: A list of dictionaries, each containing: - 'contexts': The base data store response for each memory. @@ -205,7 +205,7 @@ class MemoryGraph: params = { "user_id": filters["user_id"], - "limit": limit, + "limit": top_k, } # Build node properties based on filters node_props = ["user_id: $user_id"] @@ -316,14 +316,14 @@ class MemoryGraph: logger.debug(f"Extracted entities: {entities}") return entities - def _search_graph_db(self, node_list, filters, limit=100, threshold=None): + def _search_graph_db(self, node_list, filters, top_k=100, threshold=None): """Search similar nodes among and their respective incoming and outgoing relations.""" result_relations = [] params = { "threshold": threshold if threshold else self.threshold, "user_id": filters["user_id"], - "limit": limit, + "limit": top_k, } # Build node properties for filtering node_props = ["user_id: $user_id"] @@ -364,7 +364,7 @@ class MemoryGraph: parameters=params)) # Kuzu does not support sort/limit over unions. Do it manually for now. - result_relations.extend(sorted(results, key=lambda x: x["similarity"], reverse=True)[:limit]) + result_relations.extend(sorted(results, key=lambda x: x["similarity"], reverse=True)[:top_k]) return result_relations diff --git a/mem0/memory/main.py b/mem0/memory/main.py index 582e56d88..c18b093a0 100644 --- a/mem0/memory/main.py +++ b/mem0/memory/main.py @@ -258,7 +258,7 @@ class Memory(MemoryBase): def __init__(self, config: MemoryConfig = MemoryConfig()): self.config = config - self.custom_fact_extraction_prompt = self.config.custom_fact_extraction_prompt + self.custom_instructions = self.config.custom_instructions self.custom_update_memory_prompt = self.config.custom_update_memory_prompt self.embedding_model = EmbedderFactory.create( self.config.embedder.provider, @@ -525,8 +525,8 @@ class Memory(MemoryBase): parsed_messages = parse_messages(messages) - if self.config.custom_fact_extraction_prompt: - system_prompt = self.config.custom_fact_extraction_prompt + if self.config.custom_instructions: + system_prompt = self.config.custom_instructions user_prompt = f"Input:\n{parsed_messages}" else: # Determine if this should use agent memory extraction based on agent_id presence @@ -582,7 +582,7 @@ class Memory(MemoryBase): existing_memories = self.vector_store.search( query=new_mem, vectors=messages_embeddings, - limit=5, + top_k=5, filters=search_filters, ) for mem in existing_memories: @@ -785,7 +785,7 @@ class Memory(MemoryBase): agent_id: Optional[str] = None, run_id: Optional[str] = None, filters: Optional[Dict[str, Any]] = None, - limit: int = 100, + top_k: int = 100, ): """ List all memories. @@ -797,7 +797,7 @@ class Memory(MemoryBase): filters (dict, optional): Additional custom key-value filters to apply to the search. These are merged with the ID-based scoping filters. For example, `filters={"actor_id": "some_user"}`. - limit (int, optional): The maximum number of memories to return. Defaults to 100. + top_k (int, optional): The maximum number of memories to return. Defaults to 100. Returns: dict: A dictionary containing a list of memories under the "results" key, @@ -815,13 +815,13 @@ class Memory(MemoryBase): keys, encoded_ids = process_telemetry_filters(effective_filters) capture_event( - "mem0.get_all", self, {"limit": limit, "keys": keys, "encoded_ids": encoded_ids, "sync_type": "sync"} + "mem0.get_all", self, {"top_k": top_k, "keys": keys, "encoded_ids": encoded_ids, "sync_type": "sync"} ) with concurrent.futures.ThreadPoolExecutor() as executor: - future_memories = executor.submit(self._get_all_from_vector_store, effective_filters, limit) + future_memories = executor.submit(self._get_all_from_vector_store, effective_filters, top_k) future_graph_entities = ( - executor.submit(self.graph.get_all, effective_filters, limit) if self.enable_graph else None + executor.submit(self.graph.get_all, effective_filters, top_k) if self.enable_graph else None ) concurrent.futures.wait( @@ -836,8 +836,8 @@ class Memory(MemoryBase): return {"results": all_memories_result} - def _get_all_from_vector_store(self, filters, limit): - memories_result = self.vector_store.list(filters=filters, limit=limit) + def _get_all_from_vector_store(self, filters, top_k): + memories_result = self.vector_store.list(filters=filters, top_k=top_k) # Handle different vector store return formats by inspecting first element if isinstance(memories_result, (tuple, list)) and len(memories_result) > 0: @@ -890,7 +890,7 @@ class Memory(MemoryBase): user_id: Optional[str] = None, agent_id: Optional[str] = None, run_id: Optional[str] = None, - limit: int = 100, + top_k: int = 100, filters: Optional[Dict[str, Any]] = None, threshold: Optional[float] = None, rerank: bool = True, @@ -902,7 +902,7 @@ class Memory(MemoryBase): user_id (str, optional): ID of the user to search for. Defaults to None. agent_id (str, optional): ID of the agent to search for. Defaults to None. run_id (str, optional): ID of the run to search for. Defaults to None. - limit (int, optional): Limit the number of results. Defaults to 100. + top_k (int, optional): Maximum number of results to return. Defaults to 100. filters (dict, optional): Legacy filters to apply to the search. Defaults to None. threshold (float, optional): Minimum score for a memory to be included in the results. Defaults to None. filters (dict, optional): Enhanced metadata filtering with operators: @@ -947,7 +947,7 @@ class Memory(MemoryBase): "mem0.search", self, { - "limit": limit, + "top_k": top_k, "version": self.api_version, "keys": keys, "encoded_ids": encoded_ids, @@ -958,9 +958,9 @@ class Memory(MemoryBase): ) with concurrent.futures.ThreadPoolExecutor() as executor: - future_memories = executor.submit(self._search_vector_store, query, effective_filters, limit, threshold) + future_memories = executor.submit(self._search_vector_store, query, effective_filters, top_k, threshold) future_graph_entities = ( - executor.submit(self.graph.search, query, effective_filters, limit) if self.enable_graph else None + executor.submit(self.graph.search, query, effective_filters, top_k) if self.enable_graph else None ) concurrent.futures.wait( @@ -973,7 +973,7 @@ class Memory(MemoryBase): # Apply reranking if enabled and reranker is available if rerank and self.reranker and original_memories: try: - reranked_memories = self.reranker.rerank(query, original_memories, limit) + reranked_memories = self.reranker.rerank(query, original_memories, top_k) original_memories = reranked_memories except Exception as e: logger.warning(f"Reranking failed, using original results: {e}") @@ -1079,9 +1079,9 @@ class Memory(MemoryBase): return True return False - def _search_vector_store(self, query, filters, limit, threshold: Optional[float] = None): + def _search_vector_store(self, query, filters, top_k, threshold: Optional[float] = None): embeddings = self.embedding_model.embed(query, "search") - memories = self.vector_store.search(query=query, vectors=embeddings, limit=limit, filters=filters) + memories = self.vector_store.search(query=query, vectors=embeddings, top_k=top_k, filters=filters) promoted_payload_keys = [ "user_id", @@ -1643,8 +1643,8 @@ class AsyncMemory(MemoryBase): return returned_memories parsed_messages = parse_messages(messages) - if self.config.custom_fact_extraction_prompt: - system_prompt = self.config.custom_fact_extraction_prompt + if self.config.custom_instructions: + system_prompt = self.config.custom_instructions user_prompt = f"Input:\n{parsed_messages}" else: # Determine if this should use agent memory extraction based on agent_id presence @@ -1699,7 +1699,7 @@ class AsyncMemory(MemoryBase): self.vector_store.search, query=new_mem_content, vectors=embeddings, - limit=5, + top_k=5, filters=search_filters, ) return [{"id": mem.id, "text": mem.payload.get("data", "")} for mem in existing_mems] @@ -1922,7 +1922,7 @@ class AsyncMemory(MemoryBase): agent_id: Optional[str] = None, run_id: Optional[str] = None, filters: Optional[Dict[str, Any]] = None, - limit: int = 100, + top_k: int = 100, ): """ List all memories. @@ -1934,7 +1934,7 @@ class AsyncMemory(MemoryBase): filters (dict, optional): Additional custom key-value filters to apply to the search. These are merged with the ID-based scoping filters. For example, `filters={"actor_id": "some_user"}`. - limit (int, optional): The maximum number of memories to return. Defaults to 100. + top_k (int, optional): The maximum number of memories to return. Defaults to 100. Returns: dict: A dictionary containing a list of memories under the "results" key, @@ -1955,19 +1955,19 @@ class AsyncMemory(MemoryBase): keys, encoded_ids = process_telemetry_filters(effective_filters) capture_event( - "mem0.get_all", self, {"limit": limit, "keys": keys, "encoded_ids": encoded_ids, "sync_type": "async"} + "mem0.get_all", self, {"top_k": top_k, "keys": keys, "encoded_ids": encoded_ids, "sync_type": "async"} ) - vector_store_task = asyncio.create_task(self._get_all_from_vector_store(effective_filters, limit)) + vector_store_task = asyncio.create_task(self._get_all_from_vector_store(effective_filters, top_k)) graph_task = None if self.enable_graph: graph_get_all = getattr(self.graph, "get_all", None) if callable(graph_get_all): if asyncio.iscoroutinefunction(graph_get_all): - graph_task = asyncio.create_task(graph_get_all(effective_filters, limit)) + graph_task = asyncio.create_task(graph_get_all(effective_filters, top_k)) else: - graph_task = asyncio.create_task(asyncio.to_thread(graph_get_all, effective_filters, limit)) + graph_task = asyncio.create_task(asyncio.to_thread(graph_get_all, effective_filters, top_k)) results_dict = {} if graph_task: @@ -1978,8 +1978,8 @@ class AsyncMemory(MemoryBase): return results_dict - async def _get_all_from_vector_store(self, filters, limit): - memories_result = await asyncio.to_thread(self.vector_store.list, filters=filters, limit=limit) + async def _get_all_from_vector_store(self, filters, top_k): + memories_result = await asyncio.to_thread(self.vector_store.list, filters=filters, top_k=top_k) # Handle different vector store return formats by inspecting first element if isinstance(memories_result, (tuple, list)) and len(memories_result) > 0: @@ -2032,7 +2032,7 @@ class AsyncMemory(MemoryBase): user_id: Optional[str] = None, agent_id: Optional[str] = None, run_id: Optional[str] = None, - limit: int = 100, + top_k: int = 100, filters: Optional[Dict[str, Any]] = None, threshold: Optional[float] = None, metadata_filters: Optional[Dict[str, Any]] = None, @@ -2045,7 +2045,7 @@ class AsyncMemory(MemoryBase): user_id (str, optional): ID of the user to search for. Defaults to None. agent_id (str, optional): ID of the agent to search for. Defaults to None. run_id (str, optional): ID of the run to search for. Defaults to None. - limit (int, optional): Limit the number of results. Defaults to 100. + top_k (int, optional): Maximum number of results to return. Defaults to 100. filters (dict, optional): Legacy filters to apply to the search. Defaults to None. threshold (float, optional): Minimum score for a memory to be included in the results. Defaults to None. filters (dict, optional): Enhanced metadata filtering with operators: @@ -2091,7 +2091,7 @@ class AsyncMemory(MemoryBase): "mem0.search", self, { - "limit": limit, + "top_k": top_k, "version": self.api_version, "keys": keys, "encoded_ids": encoded_ids, @@ -2101,14 +2101,14 @@ class AsyncMemory(MemoryBase): }, ) - vector_store_task = asyncio.create_task(self._search_vector_store(query, effective_filters, limit, threshold)) + vector_store_task = asyncio.create_task(self._search_vector_store(query, effective_filters, top_k, threshold)) graph_task = None if self.enable_graph: if hasattr(self.graph.search, "__await__"): # Check if graph search is async - graph_task = asyncio.create_task(self.graph.search(query, effective_filters, limit)) + graph_task = asyncio.create_task(self.graph.search(query, effective_filters, top_k)) else: - graph_task = asyncio.create_task(asyncio.to_thread(self.graph.search, query, effective_filters, limit)) + graph_task = asyncio.create_task(asyncio.to_thread(self.graph.search, query, effective_filters, top_k)) if graph_task: original_memories, graph_entities = await asyncio.gather(vector_store_task, graph_task) @@ -2121,7 +2121,7 @@ class AsyncMemory(MemoryBase): try: # Run reranking in thread pool to avoid blocking async loop reranked_memories = await asyncio.to_thread( - self.reranker.rerank, query, original_memories, limit + self.reranker.rerank, query, original_memories, top_k ) original_memories = reranked_memories except Exception as e: @@ -2228,10 +2228,10 @@ class AsyncMemory(MemoryBase): return True return False - async def _search_vector_store(self, query, filters, limit, threshold: Optional[float] = None): + async def _search_vector_store(self, query, filters, top_k, threshold: Optional[float] = None): embeddings = await asyncio.to_thread(self.embedding_model.embed, query, "search") memories = await asyncio.to_thread( - self.vector_store.search, query=query, vectors=embeddings, limit=limit, filters=filters + self.vector_store.search, query=query, vectors=embeddings, top_k=top_k, filters=filters ) promoted_payload_keys = [ diff --git a/mem0/memory/memgraph_memory.py b/mem0/memory/memgraph_memory.py index 3a29b9e71..d0e871100 100644 --- a/mem0/memory/memgraph_memory.py +++ b/mem0/memory/memgraph_memory.py @@ -98,14 +98,14 @@ class MemoryGraph: return {"deleted_entities": deleted_entities, "added_entities": added_entities} - def search(self, query, filters, limit=100): + def search(self, query, filters, top_k=100): """ Search for memories and related graph data. Args: query (str): Query to search for. filters (dict): A dictionary containing filters to be applied during the search. - limit (int): The maximum number of nodes and relationships to retrieve. Defaults to 100. + top_k (int): The maximum number of nodes and relationships to retrieve. Defaults to 100. Returns: dict: A dictionary containing: @@ -172,14 +172,14 @@ class MemoryGraph: params = {"user_id": filters["user_id"]} self.graph.query(cypher, params=params) - def get_all(self, filters, limit=100): + def get_all(self, filters, top_k=100): """ Retrieves all nodes and relationships from the graph database based on optional filtering criteria. Args: filters (dict): A dictionary containing filters to be applied during the retrieval. Supports 'user_id' (required) and 'agent_id' (optional). - limit (int): The maximum number of nodes and relationships to retrieve. Defaults to 100. + top_k (int): The maximum number of nodes and relationships to retrieve. Defaults to 100. Returns: list: A list of dictionaries, each containing: - 'source': The source node name. @@ -193,14 +193,14 @@ class MemoryGraph: RETURN n.name AS source, type(r) AS relationship, m.name AS target LIMIT $limit """ - params = {"user_id": filters["user_id"], "agent_id": filters["agent_id"], "limit": limit} + params = {"user_id": filters["user_id"], "agent_id": filters["agent_id"], "limit": top_k} else: query = """ MATCH (n:Entity {user_id: $user_id})-[r]->(m:Entity {user_id: $user_id}) RETURN n.name AS source, type(r) AS relationship, m.name AS target LIMIT $limit """ - params = {"user_id": filters["user_id"], "limit": limit} + params = {"user_id": filters["user_id"], "limit": top_k} results = self.graph.query(query, params=params) @@ -293,7 +293,7 @@ class MemoryGraph: logger.debug(f"Extracted entities: {entities}") return entities - def _search_graph_db(self, node_list, filters, limit=100): + def _search_graph_db(self, node_list, filters, top_k=100): """Search similar nodes among and their respective incoming and outgoing relations.""" result_relations = [] @@ -324,7 +324,7 @@ class MemoryGraph: "threshold": self.threshold, "user_id": filters["user_id"], "agent_id": filters["agent_id"], - "limit": limit, + "limit": top_k, } else: cypher_query = """ @@ -348,7 +348,7 @@ class MemoryGraph: "n_embedding": n_embedding, "threshold": self.threshold, "user_id": filters["user_id"], - "limit": limit, + "limit": top_k, } ans = self.graph.query(cypher_query, params=params) diff --git a/mem0/memory/utils.py b/mem0/memory/utils.py index 61e3863e3..f03341efe 100644 --- a/mem0/memory/utils.py +++ b/mem0/memory/utils.py @@ -38,7 +38,7 @@ def ensure_json_instruction(system_prompt, user_prompt): OpenAI's API requires the word 'json' to appear in the messages when response_format is set to {"type": "json_object"}. When users provide a - custom_fact_extraction_prompt that doesn't include 'json', this causes a + custom_instructions that doesn't include 'json', this causes a 400 error. This function appends a JSON format instruction to the system prompt if 'json' is not already present in either prompt. diff --git a/mem0/proxy/main.py b/mem0/proxy/main.py index 4baaf5ec2..922f88614 100644 --- a/mem0/proxy/main.py +++ b/mem0/proxy/main.py @@ -59,7 +59,7 @@ class Completions: run_id: Optional[str] = None, metadata: Optional[dict] = None, filters: Optional[dict] = None, - limit: Optional[int] = 10, + top_k: Optional[int] = 10, # LLM arguments timeout: Optional[Union[float, str, httpx.Timeout]] = None, temperature: Optional[float] = None, @@ -103,7 +103,7 @@ class Completions: prepared_messages = self._prepare_messages(messages) if prepared_messages[-1]["role"] == "user": self._async_add_to_memory(messages, user_id, agent_id, run_id, metadata, filters) - relevant_memories = self._fetch_relevant_memories(messages, user_id, agent_id, run_id, filters, limit) + relevant_memories = self._fetch_relevant_memories(messages, user_id, agent_id, run_id, filters, top_k) logger.debug(f"Retrieved {len(relevant_memories)} relevant memories") prepared_messages[-1]["content"] = self._format_query_with_memories(messages, relevant_memories) @@ -163,7 +163,7 @@ class Completions: threading.Thread(target=add_task, daemon=True).start() - def _fetch_relevant_memories(self, messages, user_id, agent_id, run_id, filters, limit): + def _fetch_relevant_memories(self, messages, user_id, agent_id, run_id, filters, top_k): # Currently, only pass the last 6 messages to the search API to prevent long query message_input = [f"{message['role']}: {message['content']}" for message in messages][-6:] # TODO: Make it better by summarizing the past conversation @@ -173,7 +173,7 @@ class Completions: agent_id=agent_id, run_id=run_id, filters=filters, - limit=limit, + top_k=top_k, ) def _format_query_with_memories(self, messages, relevant_memories): diff --git a/mem0/vector_stores/azure_ai_search.py b/mem0/vector_stores/azure_ai_search.py index 6165efc6b..da2ea1333 100644 --- a/mem0/vector_stores/azure_ai_search.py +++ b/mem0/vector_stores/azure_ai_search.py @@ -205,14 +205,14 @@ class AzureAISearch(VectorStoreBase): filter_expression = " and ".join(filter_conditions) return filter_expression - def search(self, query, vectors, limit=5, filters=None): + def search(self, query, vectors, top_k=5, filters=None): """ Search for similar vectors. Args: query (str): Query. vectors (List[float]): Query vector. - limit (int, optional): Number of results to return. Defaults to 5. + top_k (int, optional): Number of results to return. Defaults to 5. filters (Dict, optional): Filters to apply to the search. Defaults to None. Returns: @@ -222,13 +222,13 @@ class AzureAISearch(VectorStoreBase): if filters: filter_expression = self._build_filter_expression(filters) - vector_query = VectorizedQuery(vector=vectors, k_nearest_neighbors=limit, fields="vector") + vector_query = VectorizedQuery(vector=vectors, k_nearest_neighbors=top_k, fields="vector") if self.hybrid_search: search_results = self.search_client.search( search_text=query, vector_queries=[vector_query], filter=filter_expression, - top=limit, + top=top_k, vector_filter_mode=self.vector_filter_mode, search_fields=["payload"], ) @@ -236,7 +236,7 @@ class AzureAISearch(VectorStoreBase): search_results = self.search_client.search( vector_queries=[vector_query], filter=filter_expression, - top=limit, + top=top_k, vector_filter_mode=self.vector_filter_mode, ) @@ -327,13 +327,13 @@ class AzureAISearch(VectorStoreBase): index = self.index_client.get_index(self.index_name) return {"name": index.name, "fields": index.fields} - def list(self, filters=None, limit=100): + def list(self, filters=None, top_k=100): """ List all vectors in the index. Args: filters (dict, optional): Filters to apply to the list. - limit (int, optional): Number of vectors to return. Defaults to 100. + top_k (int, optional): Number of vectors to return. Defaults to 100. Returns: List[OutputData]: List of vectors. @@ -342,7 +342,7 @@ class AzureAISearch(VectorStoreBase): if filters: filter_expression = self._build_filter_expression(filters) - search_results = self.search_client.search(search_text="*", filter=filter_expression, top=limit) + search_results = self.search_client.search(search_text="*", filter=filter_expression, top=top_k) results = [] for result in search_results: payload = json.loads(extract_json(result["payload"])) diff --git a/mem0/vector_stores/azure_mysql.py b/mem0/vector_stores/azure_mysql.py index 2d9ab373b..7c8204753 100644 --- a/mem0/vector_stores/azure_mysql.py +++ b/mem0/vector_stores/azure_mysql.py @@ -7,8 +7,8 @@ from pydantic import BaseModel try: import pymysql - from pymysql.cursors import DictCursor from dbutils.pooled_db import PooledDB + from pymysql.cursors import DictCursor except ImportError: raise ImportError( "Azure MySQL vector store requires PyMySQL and DBUtils. " @@ -243,7 +243,7 @@ class AzureMySQL(VectorStoreBase): self, query: str, vectors: List[float], - limit: int = 5, + top_k: int = 5, filters: Optional[Dict] = None, ) -> List[OutputData]: """ @@ -252,7 +252,7 @@ class AzureMySQL(VectorStoreBase): Args: query (str): Query string (not used in vector search) vectors (List[float]): Query vector - limit (int): Number of results to return + top_k (int): Number of results to return filters (Dict, optional): Filters to apply to the search Returns: @@ -291,9 +291,9 @@ class AzureMySQL(VectorStoreBase): distance = 1 - similarity scored_results.append((row['id'], distance, row['payload'])) - # Sort by distance and limit + # Sort by distance and apply limit scored_results.sort(key=lambda x: x[1]) - scored_results = scored_results[:limit] + scored_results = scored_results[:top_k] return [ OutputData(id=r[0], score=float(r[1]), payload=json.loads(r[2]) if isinstance(r[2], str) else r[2]) @@ -406,14 +406,14 @@ class AzureMySQL(VectorStoreBase): def list( self, filters: Optional[Dict] = None, - limit: int = 100 + top_k: int = 100 ) -> List[List[OutputData]]: """ List all vectors in the collection. Args: filters (Dict, optional): Filters to apply - limit (int): Number of vectors to return + top_k (int): Number of vectors to return Returns: List[List[OutputData]]: List of vectors @@ -436,7 +436,7 @@ class AzureMySQL(VectorStoreBase): {filter_clause} LIMIT %s """, - (*filter_params, limit) + (*filter_params, top_k) ) results = cur.fetchall() diff --git a/mem0/vector_stores/baidu.py b/mem0/vector_stores/baidu.py index 2c211abe9..f9b2ce54b 100644 --- a/mem0/vector_stores/baidu.py +++ b/mem0/vector_stores/baidu.py @@ -185,14 +185,14 @@ class BaiduDB(VectorStoreBase): row = Row(id=idx, vector=vector, metadata=metadata) self._table.upsert(rows=[row]) - def search(self, query: str, vectors: list, limit: int = 5, filters: dict = None) -> list: + def search(self, query: str, vectors: list, top_k: int = 5, filters: dict = None) -> list: """ Search for similar vectors. Args: query (str): Query string. vectors (List[float]): Query vector. - limit (int, optional): Number of results to return. Defaults to 5. + top_k (int, optional): Number of results to return. Defaults to 5. filters (Dict, optional): Filters to apply to the search. Defaults to None. Returns: @@ -207,7 +207,7 @@ class BaiduDB(VectorStoreBase): request = VectorTopkSearchRequest( vector_field="vector", vector=FloatVector(vectors), - limit=limit, + limit=top_k, filter=search_filter, config=VectorSearchConfig(ef=200), ) @@ -313,20 +313,20 @@ class BaiduDB(VectorStoreBase): """ return self._table.stats() - def list(self, filters: dict = None, limit: int = 100) -> list: + def list(self, filters: dict = None, top_k: int = 100) -> list: """ List all vectors in the table. Args: filters (Dict, optional): Filters to apply to the list. - limit (int, optional): Number of vectors to return. Defaults to 100. + top_k (int, optional): Number of vectors to return. Defaults to 100. Returns: List[OutputData]: List of vectors. """ projections = ["id", "metadata"] list_filter = self._create_filter(filters) if filters else None - result = self._table.select(filter=list_filter, projections=projections, limit=limit) + result = self._table.select(filter=list_filter, projections=projections, limit=top_k) memories = [] for row in result.rows: diff --git a/mem0/vector_stores/base.py b/mem0/vector_stores/base.py index 3e22499d7..77378ee18 100644 --- a/mem0/vector_stores/base.py +++ b/mem0/vector_stores/base.py @@ -13,7 +13,7 @@ class VectorStoreBase(ABC): pass @abstractmethod - def search(self, query, vectors, limit=5, filters=None): + def search(self, query, vectors, top_k=5, filters=None): """Search for similar vectors.""" pass @@ -48,7 +48,7 @@ class VectorStoreBase(ABC): pass @abstractmethod - def list(self, filters=None, limit=None): + def list(self, filters=None, top_k=None): """List all memories.""" pass diff --git a/mem0/vector_stores/cassandra.py b/mem0/vector_stores/cassandra.py index 24e4fea88..c6579a29f 100644 --- a/mem0/vector_stores/cassandra.py +++ b/mem0/vector_stores/cassandra.py @@ -7,8 +7,8 @@ import numpy as np from pydantic import BaseModel try: - from cassandra.cluster import Cluster from cassandra.auth import PlainTextAuthProvider + from cassandra.cluster import Cluster except ImportError: raise ImportError( "Apache Cassandra vector store requires cassandra-driver. " @@ -214,7 +214,7 @@ class CassandraDB(VectorStoreBase): self, query: str, vectors: List[float], - limit: int = 5, + top_k: int = 5, filters: Optional[Dict] = None, ) -> List[OutputData]: """ @@ -223,7 +223,7 @@ class CassandraDB(VectorStoreBase): Args: query (str): Query string (not used in vector search) vectors (List[float]): Query vector - limit (int): Number of results to return + top_k (int): Number of results to return filters (Dict, optional): Filters to apply to the search Returns: @@ -263,9 +263,9 @@ class CassandraDB(VectorStoreBase): scored_results.append((row.id, distance, row.payload)) - # Sort by distance and limit + # Sort by distance and apply limit scored_results.sort(key=lambda x: x[1]) - scored_results = scored_results[:limit] + scored_results = scored_results[:top_k] return [ OutputData( @@ -427,14 +427,14 @@ class CassandraDB(VectorStoreBase): def list( self, filters: Optional[Dict] = None, - limit: int = 100 + top_k: int = 100 ) -> List[List[OutputData]]: """ List all vectors in the collection. Args: filters (Dict, optional): Filters to apply - limit (int): Number of vectors to return + top_k (int): Number of vectors to return Returns: List[List[OutputData]]: List of vectors @@ -443,7 +443,7 @@ class CassandraDB(VectorStoreBase): query = f""" SELECT id, vector, payload FROM {self.keyspace}.{self.collection_name} - LIMIT {limit} + LIMIT {top_k} """ rows = self.session.execute(query) diff --git a/mem0/vector_stores/chroma.py b/mem0/vector_stores/chroma.py index 63818a5ba..476d5f0d2 100644 --- a/mem0/vector_stores/chroma.py +++ b/mem0/vector_stores/chroma.py @@ -141,7 +141,7 @@ class ChromaDB(VectorStoreBase): self.collection.add(ids=ids, embeddings=vectors, metadatas=payloads) def search( - self, query: str, vectors: List[list], limit: int = 5, filters: Optional[Dict] = None + self, query: str, vectors: List[list], top_k: int = 5, filters: Optional[Dict] = None ) -> List[OutputData]: """ Search for similar vectors. @@ -149,14 +149,14 @@ class ChromaDB(VectorStoreBase): Args: query (str): Query. vectors (List[list]): List of vectors to search. - limit (int, optional): Number of results to return. Defaults to 5. + top_k (int, optional): Number of results to return. Defaults to 5. filters (Optional[Dict], optional): Filters to apply to the search. Defaults to None. Returns: List[OutputData]: Search results. """ where_clause = self._generate_where_clause(filters) if filters else None - results = self.collection.query(query_embeddings=vectors, where=where_clause, n_results=limit) + results = self.collection.query(query_embeddings=vectors, where=where_clause, n_results=top_k) final_results = self._parse_output(results) return final_results @@ -222,19 +222,19 @@ class ChromaDB(VectorStoreBase): """ return self.client.get_collection(name=self.collection_name) - def list(self, filters: Optional[Dict] = None, limit: int = 100) -> List[OutputData]: + def list(self, filters: Optional[Dict] = None, top_k: int = 100) -> List[OutputData]: """ List all vectors in a collection. Args: filters (Optional[Dict], optional): Filters to apply to the list. Defaults to None. - limit (int, optional): Number of vectors to return. Defaults to 100. + top_k (int, optional): Number of vectors to return. Defaults to 100. Returns: List[OutputData]: List of vectors. """ where_clause = self._generate_where_clause(filters) if filters else None - results = self.collection.get(where=where_clause, limit=limit) + results = self.collection.get(where=where_clause, limit=top_k) return [self._parse_output(results)] def reset(self): diff --git a/mem0/vector_stores/databricks.py b/mem0/vector_stores/databricks.py index a8cd8ec86..595e125d5 100644 --- a/mem0/vector_stores/databricks.py +++ b/mem0/vector_stores/databricks.py @@ -2,20 +2,28 @@ import json import logging import re import uuid -from typing import Optional, List -from datetime import datetime, date -from databricks.sdk.service.catalog import ColumnInfo, ColumnTypeName, TableType, DataSourceFormat -from databricks.sdk.service.catalog import TableConstraint, PrimaryKeyConstraint +from datetime import date, datetime +from typing import List, Optional + from databricks.sdk import WorkspaceClient +from databricks.sdk.service.catalog import ( + ColumnInfo, + ColumnTypeName, + DataSourceFormat, + PrimaryKeyConstraint, + TableConstraint, + TableType, +) from databricks.sdk.service.sql import StatementParameterListItem from databricks.sdk.service.vectorsearch import ( - VectorIndexType, DeltaSyncVectorIndexSpecRequest, DirectAccessVectorIndexSpec, EmbeddingSourceColumn, EmbeddingVectorColumn, + VectorIndexType, ) from pydantic import BaseModel + from mem0.memory.utils import extract_json from mem0.vector_stores.base import VectorStoreBase @@ -446,14 +454,14 @@ class Databricks(VectorStoreBase): logger.error(f"Insert operation failed: {e}") raise - def search(self, query: str, vectors: list, limit: int = 5, filters: dict = None) -> List[MemoryResult]: + def search(self, query: str, vectors: list, top_k: int = 5, filters: dict = None) -> List[MemoryResult]: """ Search for similar vectors or text using the Databricks Vector Search index. Args: query (str): Search query text (for text-based search). vectors (list): Query vector (for vector-based search). - limit (int): Maximum number of results. + top_k (int): Maximum number of results. filters (dict): Filters to apply. Returns: @@ -468,7 +476,7 @@ class Databricks(VectorStoreBase): query_kwargs = { "index_name": self.fully_qualified_index_name, "columns": self.column_names, - "num_results": limit, + "num_results": top_k, "query_type": self.query_type, "filters_json": filters_json, } @@ -719,20 +727,20 @@ class Databricks(VectorStoreBase): logger.error(f"Failed to get info for index '{name or self.index_name}': {e}") raise - def list(self, filters: dict = None, limit: int = None) -> list[MemoryResult]: + def list(self, filters: dict = None, top_k: int = None) -> list[MemoryResult]: """ List all recent created memories from the vector store. Args: filters (dict, optional): Filters to apply. - limit (int, optional): Maximum number of results. + top_k (int, optional): Maximum number of results. Returns: List containing list of MemoryResult objects. """ try: filters_json = json.dumps(filters) if filters else None - num_results = limit or 100 + num_results = top_k or 100 columns = self.column_names # Use query_text for Delta Sync with model endpoint, query_vector otherwise query_kwargs = { diff --git a/mem0/vector_stores/elasticsearch.py b/mem0/vector_stores/elasticsearch.py index b73eedcdd..99fd61d5c 100644 --- a/mem0/vector_stores/elasticsearch.py +++ b/mem0/vector_stores/elasticsearch.py @@ -129,7 +129,7 @@ class ElasticsearchDB(VectorStoreBase): return results def search( - self, query: str, vectors: List[float], limit: int = 5, filters: Optional[Dict] = None + self, query: str, vectors: List[float], top_k: int = 5, filters: Optional[Dict] = None ) -> List[OutputData]: """ Search with two options: @@ -137,10 +137,10 @@ class ElasticsearchDB(VectorStoreBase): 2. Use KNN search on vectors with pre-filtering if no custom search query is provided """ if self.custom_search_query: - search_query = self.custom_search_query(vectors, limit, filters) + search_query = self.custom_search_query(vectors, top_k, filters) else: search_query = { - "knn": {"field": "vector", "query_vector": vectors, "k": limit, "num_candidates": limit * 2} + "knn": {"field": "vector", "query_vector": vectors, "k": top_k, "num_candidates": top_k * 2} } if filters: filter_conditions = [] @@ -203,7 +203,7 @@ class ElasticsearchDB(VectorStoreBase): """Get information about a collection (index).""" return self.client.indices.get(index=name) - def list(self, filters: Optional[Dict] = None, limit: Optional[int] = None) -> List[List[OutputData]]: + def list(self, filters: Optional[Dict] = None, top_k: Optional[int] = None) -> List[List[OutputData]]: """List all memories.""" query: Dict[str, Any] = {"query": {"match_all": {}}} @@ -213,8 +213,8 @@ class ElasticsearchDB(VectorStoreBase): filter_conditions.append({"term": {f"metadata.{key}": value}}) query["query"] = {"bool": {"must": filter_conditions}} - if limit: - query["size"] = limit + if top_k: + query["size"] = top_k response = self.client.search(index=self.collection_name, body=query) diff --git a/mem0/vector_stores/faiss.py b/mem0/vector_stores/faiss.py index 03865c0ac..4fddc89bd 100644 --- a/mem0/vector_stores/faiss.py +++ b/mem0/vector_stores/faiss.py @@ -2,14 +2,13 @@ import logging import os import pickle import uuid +import warnings from pathlib import Path from typing import Dict, List, Optional import numpy as np from pydantic import BaseModel -import warnings - try: # Suppress SWIG deprecation warnings from FAISS warnings.filterwarnings("ignore", category=DeprecationWarning, message=".*SwigPy.*") @@ -115,23 +114,23 @@ class FAISS(VectorStoreBase): except Exception as e: logger.warning(f"Failed to save FAISS index: {e}") - def _parse_output(self, scores, ids, limit=None) -> List[OutputData]: + def _parse_output(self, scores, ids, top_k=None) -> List[OutputData]: """ Parse the output data. Args: scores: Similarity scores from FAISS. ids: Indices from FAISS. - limit: Maximum number of results to return. + top_k: Maximum number of results to return. Returns: List[OutputData]: Parsed output data. """ - if limit is None: - limit = len(ids) + if top_k is None: + top_k = len(ids) results = [] - for i in range(min(len(ids), limit)): + for i in range(min(len(ids), top_k)): if ids[i] == -1: # FAISS returns -1 for empty results continue @@ -225,7 +224,7 @@ class FAISS(VectorStoreBase): logger.info(f"Inserted {len(vectors)} vectors into collection {self.collection_name}") def search( - self, query: str, vectors: List[list], limit: int = 5, filters: Optional[Dict] = None + self, query: str, vectors: List[list], top_k: int = 5, filters: Optional[Dict] = None ) -> List[OutputData]: """ Search for similar vectors. @@ -233,7 +232,7 @@ class FAISS(VectorStoreBase): Args: query (str): Query (not used, kept for API compatibility). vectors (List[list]): List of vectors to search. - limit (int, optional): Number of results to return. Defaults to 5. + top_k (int, optional): Number of results to return. Defaults to 5. filters (Optional[Dict], optional): Filters to apply to the search. Defaults to None. Returns: @@ -250,19 +249,19 @@ class FAISS(VectorStoreBase): if self.normalize_L2 and self.distance_strategy.lower() == "euclidean": faiss.normalize_L2(query_vectors) - fetch_k = limit * 2 if filters else limit + fetch_k = top_k * 2 if filters else top_k scores, indices = self.index.search(query_vectors, fetch_k) - results = self._parse_output(scores[0], indices[0], limit) + results = self._parse_output(scores[0], indices[0], top_k) if filters: filtered_results = [] for result in results: if self._apply_filters(result.payload, filters): filtered_results.append(result) - if len(filtered_results) >= limit: + if len(filtered_results) >= top_k: break - results = filtered_results[:limit] + results = filtered_results[:top_k] return results @@ -450,13 +449,13 @@ class FAISS(VectorStoreBase): "distance": self.distance_strategy, } - def list(self, filters: Optional[Dict] = None, limit: int = 100) -> List[OutputData]: + def list(self, filters: Optional[Dict] = None, top_k: int = 100) -> List[OutputData]: """ List all vectors in a collection. Args: filters (Optional[Dict], optional): Filters to apply to the list. Defaults to None. - limit (int, optional): Number of vectors to return. Defaults to 100. + top_k (int, optional): Number of vectors to return. Defaults to 100. Returns: List[OutputData]: List of vectors. @@ -482,7 +481,7 @@ class FAISS(VectorStoreBase): ) count += 1 - if count >= limit: + if count >= top_k: break return [results] diff --git a/mem0/vector_stores/langchain.py b/mem0/vector_stores/langchain.py index 844c5d414..451c21283 100644 --- a/mem0/vector_stores/langchain.py +++ b/mem0/vector_stores/langchain.py @@ -91,15 +91,15 @@ class Langchain(VectorStoreBase): texts = [payload.get("data", "") for payload in payloads] if payloads else [""] * len(vectors) self.client.add_texts(texts=texts, metadatas=payloads, ids=ids) - def search(self, query: str, vectors: List[List[float]], limit: int = 5, filters: Optional[Dict] = None): + def search(self, query: str, vectors: List[List[float]], top_k: int = 5, filters: Optional[Dict] = None): """ Search for similar vectors in LangChain. """ # For each vector, perform a similarity search if filters: - results = self.client.similarity_search_by_vector(embedding=vectors, k=limit, filter=filters) + results = self.client.similarity_search_by_vector(embedding=vectors, k=top_k, filter=filters) else: - results = self.client.similarity_search_by_vector(embedding=vectors, k=limit) + results = self.client.similarity_search_by_vector(embedding=vectors, k=top_k) final_results = self._parse_output(results) return final_results @@ -152,7 +152,7 @@ class Langchain(VectorStoreBase): """ return {"name": self.collection_name} - def list(self, filters=None, limit=None): + def list(self, filters=None, top_k=None): """ List all vectors in a collection. """ @@ -164,7 +164,7 @@ class Langchain(VectorStoreBase): # Handle all filters, not just user_id where_clause = filters - result = self.client._collection.get(where=where_clause, limit=limit) + result = self.client._collection.get(where=where_clause, limit=top_k) # Convert the result to the expected format if result and isinstance(result, dict): diff --git a/mem0/vector_stores/milvus.py b/mem0/vector_stores/milvus.py index 4a0cd7961..83787aeb9 100644 --- a/mem0/vector_stores/milvus.py +++ b/mem0/vector_stores/milvus.py @@ -139,14 +139,14 @@ class MilvusDB(VectorStoreBase): return memory - def search(self, query: str, vectors: list, limit: int = 5, filters: dict = None) -> list: + def search(self, query: str, vectors: list, top_k: int = 5, filters: dict = None) -> list: """ Search for similar vectors. Args: query (str): Query. vectors (List[float]): Query vector. - limit (int, optional): Number of results to return. Defaults to 5. + top_k (int, optional): Number of results to return. Defaults to 5. filters (Dict, optional): Filters to apply to the search. Defaults to None. Returns: @@ -156,7 +156,7 @@ class MilvusDB(VectorStoreBase): hits = self.client.search( collection_name=self.collection_name, data=[vectors], - limit=limit, + limit=top_k, filter=query_filter, output_fields=["*"], ) @@ -235,19 +235,19 @@ class MilvusDB(VectorStoreBase): """ return self.client.get_collection_stats(collection_name=self.collection_name) - def list(self, filters: dict = None, limit: int = 100) -> list: + def list(self, filters: dict = None, top_k: int = 100) -> list: """ List all vectors in a collection. Args: filters (Dict, optional): Filters to apply to the list. - limit (int, optional): Number of vectors to return. Defaults to 100. + top_k (int, optional): Number of vectors to return. Defaults to 100. Returns: List[OutputData]: List of vectors. """ query_filter = self._create_filter(filters) if filters else None - result = self.client.query(collection_name=self.collection_name, filter=query_filter, limit=limit) + result = self.client.query(collection_name=self.collection_name, filter=query_filter, limit=top_k) memories = [] for data in result: obj = OutputData(id=data.get("id"), score=None, payload=data.get("metadata")) diff --git a/mem0/vector_stores/mongodb.py b/mem0/vector_stores/mongodb.py index 5fae0ff36..e470bd503 100644 --- a/mem0/vector_stores/mongodb.py +++ b/mem0/vector_stores/mongodb.py @@ -113,14 +113,14 @@ class MongoDB(VectorStoreBase): except PyMongoError as e: logger.error(f"Error inserting data: {e}") - def search(self, query: str, vectors: List[float], limit=5, filters: Optional[Dict] = None) -> List[OutputData]: + def search(self, query: str, vectors: List[float], top_k=5, filters: Optional[Dict] = None) -> List[OutputData]: """ Search for similar vectors using the vector search index. Args: query (str): Query string vectors (List[float]): Query vector. - limit (int, optional): Number of results to return. Defaults to 5. + top_k (int, optional): Number of results to return. Defaults to 5. filters (Dict, optional): Filters to apply to the search. Returns: @@ -139,8 +139,8 @@ class MongoDB(VectorStoreBase): { "$vectorSearch": { "index": self.index_name, - "limit": limit, - "numCandidates": min(limit * 20, 10000), + "limit": top_k, + "numCandidates": min(top_k * 20, 10000), "queryVector": vectors, "path": "embedding", } @@ -271,13 +271,13 @@ class MongoDB(VectorStoreBase): logger.error(f"Error getting collection info: {e}") return {} - def list(self, filters: Optional[Dict] = None, limit: int = 100) -> List[OutputData]: + def list(self, filters: Optional[Dict] = None, top_k: int = 100) -> List[OutputData]: """ List vectors in the collection. Args: filters (Dict, optional): Filters to apply to the list. - limit (int, optional): Number of vectors to return. + top_k (int, optional): Number of vectors to return. Returns: List[OutputData]: List of vectors. @@ -292,7 +292,7 @@ class MongoDB(VectorStoreBase): if filter_conditions: query = {"$and": filter_conditions} - cursor = self.collection.find(query).limit(limit) + cursor = self.collection.find(query).limit(top_k) results = [OutputData(id=str(doc["_id"]), score=None, payload=doc.get("payload")) for doc in cursor] logger.info(f"Retrieved {len(results)} documents from collection '{self.collection_name}'.") return results diff --git a/mem0/vector_stores/neptune_analytics.py b/mem0/vector_stores/neptune_analytics.py index e05e09033..584c72555 100644 --- a/mem0/vector_stores/neptune_analytics.py +++ b/mem0/vector_stores/neptune_analytics.py @@ -133,7 +133,7 @@ class NeptuneAnalyticsVector(VectorStoreBase): def search( - self, query: str, vectors: List[float], limit: int = 5, filters: Optional[Dict] = None + self, query: str, vectors: List[float], top_k: int = 5, filters: Optional[Dict] = None ) -> List[OutputData]: """ Search for similar vectors using embedding similarity. @@ -144,7 +144,7 @@ class NeptuneAnalyticsVector(VectorStoreBase): Args: query (str): Search query text (unused in vector search). vectors (List[float]): Query embedding vector. - limit (int, optional): Maximum number of results to return. Defaults to 5. + top_k (int, optional): Maximum number of results to return. Defaults to 5. filters (Optional[Dict]): Optional filters to apply to search results. Returns: @@ -159,7 +159,7 @@ class NeptuneAnalyticsVector(VectorStoreBase): query_string = f""" CALL neptune.algo.vectors.topKByEmbeddingWithFiltering({{ - topK: {limit}, + topK: {top_k}, embedding: {vectors} {filter_clause} }} @@ -309,7 +309,7 @@ class NeptuneAnalyticsVector(VectorStoreBase): pass - def list(self, filters: Optional[Dict] = None, limit: int = 100) -> List[OutputData]: + def list(self, filters: Optional[Dict] = None, top_k: int = 100) -> List[OutputData]: """ List all vectors in the collection with optional filtering. @@ -317,7 +317,7 @@ class NeptuneAnalyticsVector(VectorStoreBase): Args: filters (Optional[Dict]): Optional filters to apply based on metadata. - limit (int, optional): Maximum number of vectors to return. Defaults to 100. + top_k (int, optional): Maximum number of vectors to return. Defaults to 100. Returns: List[OutputData]: List of vectors with their metadata. @@ -325,7 +325,7 @@ class NeptuneAnalyticsVector(VectorStoreBase): where_clause = self._get_where_clause(filters) if filters else "" para = { - "limit": limit, + "limit": top_k, } query_string = f""" MATCH (n :{self.collection_name}) diff --git a/mem0/vector_stores/opensearch.py b/mem0/vector_stores/opensearch.py index aee256640..9eb4ca6db 100644 --- a/mem0/vector_stores/opensearch.py +++ b/mem0/vector_stores/opensearch.py @@ -158,7 +158,7 @@ class OpenSearchDB(VectorStoreBase): return results def search( - self, query: str, vectors: List[float], limit: int = 5, filters: Optional[Dict] = None + self, query: str, vectors: List[float], top_k: int = 5, filters: Optional[Dict] = None ) -> List[OutputData]: """Search for similar vectors using OpenSearch k-NN search with optional filters.""" @@ -167,13 +167,13 @@ class OpenSearchDB(VectorStoreBase): "knn": { "vector_field": { "vector": vectors, - "k": limit * 2, + "k": top_k * 2, } } } # Start building the full query - query_body = {"size": limit * 2, "query": None} + query_body = {"size": top_k * 2, "query": None} # Prepare filter conditions if applicable filter_clauses = [] @@ -196,7 +196,7 @@ class OpenSearchDB(VectorStoreBase): hits = response["hits"]["hits"] results = [ OutputData(id=hit["_source"].get("id"), score=hit["_score"], payload=hit["_source"].get("payload", {})) - for hit in hits[:limit] # Ensure we don't exceed limit + for hit in hits[:top_k] # Ensure we don't exceed top_k ] return results except Exception as e: @@ -284,7 +284,7 @@ class OpenSearchDB(VectorStoreBase): """Get information about a collection (index).""" return self.client.indices.get(index=name) - def list(self, filters: Optional[Dict] = None, limit: Optional[int] = None) -> List[OutputData]: + def list(self, filters: Optional[Dict] = None, top_k: Optional[int] = None) -> List[OutputData]: try: """List all memories with optional filters.""" query: Dict = {"query": {"match_all": {}}} @@ -299,8 +299,8 @@ class OpenSearchDB(VectorStoreBase): if filter_clauses: query["query"] = {"bool": {"filter": filter_clauses}} - if limit: - query["size"] = limit + if top_k: + query["size"] = top_k response = self.client.search(index=self.collection_name, body=query) hits = response["hits"]["hits"] diff --git a/mem0/vector_stores/pgvector.py b/mem0/vector_stores/pgvector.py index e2d020a66..44f34b7ed 100644 --- a/mem0/vector_stores/pgvector.py +++ b/mem0/vector_stores/pgvector.py @@ -203,7 +203,7 @@ class PGVector(VectorStoreBase): self, query: str, vectors: list[float], - limit: Optional[int] = 5, + top_k: Optional[int] = 5, filters: Optional[dict] = None, ) -> List[OutputData]: """ @@ -212,7 +212,7 @@ class PGVector(VectorStoreBase): Args: query (str): Query. vectors (List[float]): Query vector. - limit (int, optional): Number of results to return. Defaults to 5. + top_k (int, optional): Number of results to return. Defaults to 5. filters (Dict, optional): Filters to apply to the search. Defaults to None. Returns: @@ -237,7 +237,7 @@ class PGVector(VectorStoreBase): ORDER BY distance LIMIT %s """, - (vectors, *filter_params, limit), + (vectors, *filter_params, top_k), ) results = cur.fetchall() @@ -350,14 +350,14 @@ class PGVector(VectorStoreBase): def list( self, filters: Optional[dict] = None, - limit: Optional[int] = 100 + top_k: Optional[int] = 100 ) -> List[OutputData]: """ List all vectors in a collection. Args: filters (Dict, optional): Filters to apply to the list. - limit (int, optional): Number of vectors to return. Defaults to 100. + top_k (int, optional): Number of vectors to return. Defaults to 100. Returns: List[OutputData]: List of vectors. @@ -380,7 +380,7 @@ class PGVector(VectorStoreBase): """ with self._get_cursor() as cur: - cur.execute(query, (*filter_params, limit)) + cur.execute(query, (*filter_params, top_k)) results = cur.fetchall() return [[OutputData(id=str(r[0]), score=None, payload=r[2]) for r in results]] diff --git a/mem0/vector_stores/pinecone.py b/mem0/vector_stores/pinecone.py index 08ccf8bc6..0de07096d 100644 --- a/mem0/vector_stores/pinecone.py +++ b/mem0/vector_stores/pinecone.py @@ -204,7 +204,7 @@ class PineconeDB(VectorStoreBase): return pinecone_filter def search( - self, query: str, vectors: List[float], limit: int = 5, filters: Optional[Dict] = None + self, query: str, vectors: List[float], top_k: int = 5, filters: Optional[Dict] = None ) -> List[OutputData]: """ Search for similar vectors. @@ -212,7 +212,7 @@ class PineconeDB(VectorStoreBase): Args: query (str): Query. vectors (list): List of vectors to search. - limit (int, optional): Number of results to return. Defaults to 5. + top_k (int, optional): Number of results to return. Defaults to 5. filters (dict, optional): Filters to apply to the search. Defaults to None. Returns: @@ -222,7 +222,7 @@ class PineconeDB(VectorStoreBase): query_params = { "vector": vectors, - "top_k": limit, + "top_k": top_k, "include_metadata": True, "include_values": False, } @@ -320,13 +320,13 @@ class PineconeDB(VectorStoreBase): """ return self.client.describe_index(self.collection_name) - def list(self, filters: Optional[Dict] = None, limit: int = 100) -> List[OutputData]: + def list(self, filters: Optional[Dict] = None, top_k: int = 100) -> List[OutputData]: """ List vectors in an index with optional filtering. Args: filters (dict, optional): Filters to apply to the list. Defaults to None. - limit (int, optional): Number of vectors to return. Defaults to 100. + top_k (int, optional): Number of vectors to return. Defaults to 100. Returns: dict: List of vectors with their metadata. @@ -340,7 +340,7 @@ class PineconeDB(VectorStoreBase): query_params = { "vector": zero_vector, - "top_k": limit, + "top_k": top_k, "include_metadata": True, "include_values": True, } diff --git a/mem0/vector_stores/qdrant.py b/mem0/vector_stores/qdrant.py index 3241722de..d0bac7f42 100644 --- a/mem0/vector_stores/qdrant.py +++ b/mem0/vector_stores/qdrant.py @@ -308,14 +308,14 @@ class Qdrant(VectorStoreBase): must_not=must_not or None, ) - def search(self, query: str, vectors: list, limit: int = 5, filters: dict = None) -> list: + def search(self, query: str, vectors: list, top_k: int = 5, filters: dict = None) -> list: """ Search for similar vectors. Args: query (str): Query. vectors (list): Query vector. - limit (int, optional): Number of results to return. Defaults to 5. + top_k (int, optional): Number of results to return. Defaults to 5. filters (dict, optional): Filters to apply to the search. Defaults to None. Returns: @@ -326,7 +326,7 @@ class Qdrant(VectorStoreBase): collection_name=self.collection_name, query=vectors, query_filter=query_filter, - limit=limit, + limit=top_k, ) return hits.points @@ -404,13 +404,13 @@ class Qdrant(VectorStoreBase): """ return self.client.get_collection(collection_name=self.collection_name) - def list(self, filters: dict = None, limit: int = 100) -> list: + def list(self, filters: dict = None, top_k: int = 100) -> list: """ List all vectors in a collection. Args: filters (dict, optional): Filters to apply to the list. Defaults to None. - limit (int, optional): Number of vectors to return. Defaults to 100. + top_k (int, optional): Number of vectors to return. Defaults to 100. Returns: list: List of vectors. @@ -419,7 +419,7 @@ class Qdrant(VectorStoreBase): result = self.client.scroll( collection_name=self.collection_name, scroll_filter=query_filter, - limit=limit, + limit=top_k, with_payload=True, with_vectors=False, ) diff --git a/mem0/vector_stores/redis.py b/mem0/vector_stores/redis.py index 6e2544a0a..904904dac 100644 --- a/mem0/vector_stores/redis.py +++ b/mem0/vector_stores/redis.py @@ -141,7 +141,7 @@ class RedisDB(VectorStoreBase): data.append(entry) self.index.load(data, id_field="memory_id") - def search(self, query: str, vectors: list, limit: int = 5, filters: dict = None): + def search(self, query: str, vectors: list, top_k: int = 5, filters: dict = None): conditions = [Tag(key) == value for key, value in filters.items() if value is not None] filter = reduce(lambda x, y: x & y, conditions) @@ -150,7 +150,7 @@ class RedisDB(VectorStoreBase): vector_field_name="embedding", return_fields=["memory_id", "hash", "agent_id", "run_id", "user_id", "memory", "metadata", "created_at"], filter_expression=filter, - num_results=limit, + num_results=top_k, ) results = self.index.query(v) @@ -254,15 +254,15 @@ class RedisDB(VectorStoreBase): # Recreate the index with the same parameters self.create_col(collection_name, self.embedding_model_dims) - def list(self, filters: dict = None, limit: int = None) -> list: + def list(self, filters: dict = None, top_k: int = None) -> list: """ List all recent created memories from the vector store. """ conditions = [Tag(key) == value for key, value in filters.items() if value is not None] filter = reduce(lambda x, y: x & y, conditions) query = Query(str(filter)).sort_by("created_at", asc=False) - if limit is not None: - query = Query(str(filter)).sort_by("created_at", asc=False).paging(0, limit) + if top_k is not None: + query = Query(str(filter)).sort_by("created_at", asc=False).paging(0, top_k) results = self.index.search(query) return [ diff --git a/mem0/vector_stores/s3_vectors.py b/mem0/vector_stores/s3_vectors.py index f6504c379..c44ce15c8 100644 --- a/mem0/vector_stores/s3_vectors.py +++ b/mem0/vector_stores/s3_vectors.py @@ -99,12 +99,12 @@ class S3Vectors(VectorStoreBase): vectors=vectors_to_put, ) - def search(self, query, vectors, limit=5, filters=None): + def search(self, query, vectors, top_k=5, filters=None): params = { "vectorBucketName": self.vector_bucket_name, "indexName": self.collection_name, "queryVector": {"float32": vectors}, - "topK": limit, + "topK": top_k, "returnMetadata": True, "returnDistance": True, } @@ -149,7 +149,7 @@ class S3Vectors(VectorStoreBase): response = self.client.get_index(vectorBucketName=self.vector_bucket_name, indexName=self.collection_name) return response.get("index", {}) - def list(self, filters=None, limit=None): + def list(self, filters=None, top_k=None): # Note: list_vectors does not support metadata filtering. if filters: logger.warning("S3 Vectors `list` does not support metadata filtering. Ignoring filters.") @@ -160,8 +160,8 @@ class S3Vectors(VectorStoreBase): "returnData": False, "returnMetadata": True, } - if limit: - params["maxResults"] = limit + if top_k: + params["maxResults"] = top_k paginator = self.client.get_paginator("list_vectors") pages = paginator.paginate(**params) diff --git a/mem0/vector_stores/supabase.py b/mem0/vector_stores/supabase.py index e55a979cb..79d2257b5 100644 --- a/mem0/vector_stores/supabase.py +++ b/mem0/vector_stores/supabase.py @@ -116,7 +116,7 @@ class Supabase(VectorStoreBase): self.collection.upsert(records) def search( - self, query: str, vectors: List[float], limit: int = 5, filters: Optional[dict] = None + self, query: str, vectors: List[float], top_k: int = 5, filters: Optional[dict] = None ) -> List[OutputData]: """ Search for similar vectors. @@ -124,7 +124,7 @@ class Supabase(VectorStoreBase): Args: query (str): Query. vectors (List[float]): Query vector. - limit (int, optional): Number of results to return. Defaults to 5. + top_k (int, optional): Number of results to return. Defaults to 5. filters (Dict, optional): Filters to apply to the search. Defaults to None. Returns: @@ -132,7 +132,7 @@ class Supabase(VectorStoreBase): """ filters = self._preprocess_filters(filters) results = self.collection.query( - data=vectors, limit=limit, filters=filters, include_metadata=True, include_value=True + data=vectors, limit=top_k, filters=filters, include_metadata=True, include_value=True ) return [OutputData(id=str(result[0]), score=float(result[1]), payload=result[2]) for result in results] @@ -209,13 +209,13 @@ class Supabase(VectorStoreBase): "index": {"method": info.index_method, "metric": info.distance_metric}, } - def list(self, filters: Optional[dict] = None, limit: int = 100) -> List[OutputData]: + def list(self, filters: Optional[dict] = None, top_k: int = 100) -> List[OutputData]: """ List vectors in the collection. Args: filters (Dict, optional): Filters to apply - limit (int, optional): Maximum number of results to return. Defaults to 100. + top_k (int, optional): Maximum number of results to return. Defaults to 100. Returns: List[OutputData]: List of vectors @@ -223,7 +223,7 @@ class Supabase(VectorStoreBase): filters = self._preprocess_filters(filters) query = [0] * self.embedding_model_dims ids = self.collection.query( - data=query, limit=limit, filters=filters, include_metadata=True, include_value=False + data=query, limit=top_k, filters=filters, include_metadata=True, include_value=False ) ids = [id[0] for id in ids] records = self.collection.fetch(ids=ids) diff --git a/mem0/vector_stores/turbopuffer.py b/mem0/vector_stores/turbopuffer.py index a6708746f..bef4b40b2 100644 --- a/mem0/vector_stores/turbopuffer.py +++ b/mem0/vector_stores/turbopuffer.py @@ -157,7 +157,7 @@ class TurbopufferDB(VectorStoreBase): return ("And", tuple(conditions)) def search( - self, query: str, vectors: List[float], limit: int = 5, filters: Optional[Dict] = None + self, query: str, vectors: List[float], top_k: int = 5, filters: Optional[Dict] = None ) -> List[OutputData]: """ Search for similar vectors. @@ -165,7 +165,7 @@ class TurbopufferDB(VectorStoreBase): Args: query (str): Query text (unused in vector search, kept for interface consistency). vectors (list): Query vector to search with. - limit (int, optional): Number of results to return. Defaults to 5. + top_k (int, optional): Number of results to return. Defaults to 5. filters (dict, optional): Filters to apply to the search. Defaults to None. Returns: @@ -173,7 +173,7 @@ class TurbopufferDB(VectorStoreBase): """ query_params = { "rank_by": ("vector", "ANN", vectors), - "top_k": limit, + "top_k": top_k, "include_attributes": True, } @@ -290,20 +290,20 @@ class TurbopufferDB(VectorStoreBase): except Exception: return {"name": self.collection_name} - def list(self, filters: Optional[Dict] = None, limit: int = 100) -> list: + def list(self, filters: Optional[Dict] = None, top_k: int = 100) -> list: """ List vectors in the namespace with optional filtering. Args: filters (dict, optional): Filters to apply. Defaults to None. - limit (int, optional): Number of vectors to return. Defaults to 100. + top_k (int, optional): Number of vectors to return. Defaults to 100. Returns: list: Wrapped list of OutputData objects ([[results]]). """ query_params = { "rank_by": ("vector", "ANN", [0.0] * self.embedding_model_dims), - "top_k": limit, + "top_k": top_k, "include_attributes": True, } diff --git a/mem0/vector_stores/upstash_vector.py b/mem0/vector_stores/upstash_vector.py index 82dc0f441..19d154993 100644 --- a/mem0/vector_stores/upstash_vector.py +++ b/mem0/vector_stores/upstash_vector.py @@ -98,7 +98,7 @@ class UpstashVector(VectorStoreBase): self, query: str, vectors: List[list], - limit: int = 5, + top_k: int = 5, filters: Optional[Dict] = None, ) -> List[OutputData]: """ @@ -106,7 +106,7 @@ class UpstashVector(VectorStoreBase): Args: query (list): Query vector. - limit (int, optional): Number of results to return. Defaults to 5. + top_k (int, optional): Number of results to return. Defaults to 5. filters (Dict, optional): Filters to apply to the search. Returns: @@ -120,7 +120,7 @@ class UpstashVector(VectorStoreBase): if self.enable_embeddings: response = self.client.query( data=query, - top_k=limit, + top_k=top_k, filter=filters_str or "", include_metadata=True, namespace=self.collection_name, @@ -129,7 +129,7 @@ class UpstashVector(VectorStoreBase): queries = [ { "vector": v, - "top_k": limit, + "top_k": top_k, "filter": filters_str or "", "include_metadata": True, "namespace": self.collection_name, @@ -205,12 +205,12 @@ class UpstashVector(VectorStoreBase): return None return OutputData(id=vector.id, score=None, payload=vector.metadata) - def list(self, filters: Optional[Dict] = None, limit: int = 100) -> List[List[OutputData]]: + def list(self, filters: Optional[Dict] = None, top_k: int = 100) -> List[List[OutputData]]: """ List all memories. Args: filters (Dict, optional): Filters to apply to the search. Defaults to None. - limit (int, optional): Number of results to return. Defaults to 100. + top_k (int, optional): Number of results to return. Defaults to 100. Returns: List[OutputData]: Search results. """ @@ -233,7 +233,7 @@ class UpstashVector(VectorStoreBase): ) with query: while True: - if len(results) >= limit: + if len(results) >= top_k: break res = query.fetch_next(100) if not res: diff --git a/mem0/vector_stores/valkey.py b/mem0/vector_stores/valkey.py index 273aaba11..4dc0dc78d 100644 --- a/mem0/vector_stores/valkey.py +++ b/mem0/vector_stores/valkey.py @@ -410,14 +410,14 @@ class ValkeyDB(VectorStoreBase): return memory_results - def search(self, query: str, vectors: list, limit: int = 5, filters: dict = None, ef_runtime: int = None): + def search(self, query: str, vectors: list, top_k: int = 5, filters: dict = None, ef_runtime: int = None): """ Search for similar vectors in the index. Args: query (str): The search query. vectors (list): The vector to search for. - limit (int, optional): Maximum number of results to return. Defaults to 5. + top_k (int, optional): Maximum number of results to return. Defaults to 5. filters (dict, optional): Filters to apply to the search. Defaults to None. ef_runtime (int, optional): HNSW ef_runtime parameter for this query. Only used with HNSW index. Defaults to None. @@ -429,10 +429,10 @@ class ValkeyDB(VectorStoreBase): # Build the KNN part with optional EF_RUNTIME for HNSW if self.index_type == "hnsw" and ef_runtime is not None: - knn_part = f"[KNN {limit} @embedding $vec_param EF_RUNTIME {ef_runtime} AS vector_score]" + knn_part = f"[KNN {top_k} @embedding $vec_param EF_RUNTIME {ef_runtime} AS vector_score]" else: # For FLAT indexes or when ef_runtime is None, use basic KNN - knn_part = f"[KNN {limit} @embedding $vec_param AS vector_score]" + knn_part = f"[KNN {top_k} @embedding $vec_param AS vector_score]" # Build the complete query q = self._build_search_query(knn_part, filters) @@ -768,7 +768,7 @@ class ValkeyDB(VectorStoreBase): return q - def list(self, filters: dict = None, limit: int = None) -> list: + def list(self, filters: dict = None, top_k: int = None) -> list: """ List all recent created memories from the vector store. @@ -778,7 +778,7 @@ class ValkeyDB(VectorStoreBase): Values are used as-is without validation - wildcards, special characters, lists, etc. are passed through literally to Valkey search. Multiple filters are combined with AND logic. - limit (int, optional): Maximum number of results to return. Defaults to 1000 + top_k (int, optional): Maximum number of results to return. Defaults to 1000 if not specified. Returns: @@ -789,10 +789,10 @@ class ValkeyDB(VectorStoreBase): # Since Valkey search requires vector format, use a dummy vector search # that returns all documents by using a zero vector and large K dummy_vector = [0.0] * self.embedding_model_dims - search_limit = limit if limit is not None else 1000 # Large default + search_limit = top_k if top_k is not None else 1000 # Large default # Use the existing search method which handles filters properly - search_results = self.search("", dummy_vector, limit=search_limit, filters=filters) + search_results = self.search("", dummy_vector, top_k=search_limit, filters=filters) # Convert search results to list format (match Redis format) class MemoryResult: diff --git a/mem0/vector_stores/vertex_ai_vector_search.py b/mem0/vector_stores/vertex_ai_vector_search.py index 9e2a9a5c4..d9fb6e782 100644 --- a/mem0/vector_stores/vertex_ai_vector_search.py +++ b/mem0/vector_stores/vertex_ai_vector_search.py @@ -5,7 +5,9 @@ from typing import Any, Dict, List, Optional, Tuple import google.api_core.exceptions from google.cloud import aiplatform, aiplatform_v1 -from google.cloud.aiplatform.matching_engine.matching_engine_index_endpoint import Namespace +from google.cloud.aiplatform.matching_engine.matching_engine_index_endpoint import ( + Namespace, +) from google.oauth2 import service_account from pydantic import BaseModel @@ -206,20 +208,20 @@ class GoogleMatchingEngine(VectorStoreBase): raise def search( - self, query: str, vectors: List[float], limit: int = 5, filters: Optional[Dict] = None + self, query: str, vectors: List[float], top_k: int = 5, filters: Optional[Dict] = None ) -> List[OutputData]: """ Search for similar vectors. Args: query (str): Query. vectors (List[float]): Query vector. - limit (int, optional): Number of results to return. Defaults to 5. + top_k (int, optional): Number of results to return. Defaults to 5. filters (Optional[Dict], optional): Filters to apply to the search. Defaults to None. Returns: List[OutputData]: Search results (unwrapped) """ logger.debug("Starting search") - logger.debug("Limit: %d, Filters: %s", limit, filters) + logger.debug("Limit: %d, Filters: %s", top_k, filters) try: filter_namespaces = [] @@ -241,7 +243,7 @@ class GoogleMatchingEngine(VectorStoreBase): response = self.index_endpoint.find_neighbors( deployed_index_id=self.deployment_index_id, queries=[vectors], - num_neighbors=limit, + num_neighbors=top_k, filter=filter_namespaces if filter_namespaces else None, return_full_datapoint=True, ) @@ -453,12 +455,12 @@ class GoogleMatchingEngine(VectorStoreBase): "region": self.region, } - def list(self, filters: Optional[Dict] = None, limit: Optional[int] = None) -> List[List[OutputData]]: + def list(self, filters: Optional[Dict] = None, top_k: Optional[int] = None) -> List[List[OutputData]]: """List vectors matching the given filters. Args: filters: Optional filters to apply - limit: Optional maximum number of results to return + top_k: Optional maximum number of results to return Returns: List[List[OutputData]]: List of matching vectors wrapped in an extra array @@ -466,17 +468,17 @@ class GoogleMatchingEngine(VectorStoreBase): """ logger.debug("Starting list operation") logger.debug("Filters: %s", filters) - logger.debug("Limit: %s", limit) + logger.debug("Limit: %s", top_k) try: # Use a zero vector for the search dimension = 768 # This should be configurable based on the model zero_vector = [0.0] * dimension - # Use a large limit if none specified - search_limit = limit if limit is not None else 10000 + # Use a large top_k if none specified + search_limit = top_k if top_k is not None else 10000 - results = self.search(query=zero_vector, limit=search_limit, filters=filters) + results = self.search(query=zero_vector, top_k=search_limit, filters=filters) logger.debug("Found %d results", len(results)) return [results] # Wrap in extra array to match interface @@ -609,7 +611,7 @@ class GoogleMatchingEngine(VectorStoreBase): logger.debug("Filter: %s", filter) embedding = self.embedder.embed_query(query) - results = self.search(query=embedding, limit=k, filters=filter) + results = self.search(query=embedding, top_k=k, filters=filter) docs_and_scores = [ (Document(page_content=result.payload.get("text", ""), metadata=result.payload), result.score) diff --git a/mem0/vector_stores/weaviate.py b/mem0/vector_stores/weaviate.py index cb1ed6d3a..b76bf90b3 100644 --- a/mem0/vector_stores/weaviate.py +++ b/mem0/vector_stores/weaviate.py @@ -179,7 +179,7 @@ class Weaviate(VectorStoreBase): batch.add_object(collection=self.collection_name, properties=data_object, uuid=object_id, vector=vector) def search( - self, query: str, vectors: List[float], limit: int = 5, filters: Optional[Dict] = None + self, query: str, vectors: List[float], top_k: int = 5, filters: Optional[Dict] = None ) -> List[OutputData]: """ Search for similar vectors. @@ -194,7 +194,7 @@ class Weaviate(VectorStoreBase): response = collection.query.hybrid( query="", vector=vectors, - limit=limit, + limit=top_k, filters=combined_filter, return_properties=["hash", "created_at", "updated_at", "user_id", "agent_id", "run_id", "data", "category"], return_metadata=MetadataQuery(score=True), @@ -313,7 +313,7 @@ class Weaviate(VectorStoreBase): return schema return None - def list(self, filters=None, limit=100) -> List[OutputData]: + def list(self, filters=None, top_k=100) -> List[OutputData]: """ List all vectors in a collection. """ @@ -325,7 +325,7 @@ class Weaviate(VectorStoreBase): filter_conditions.append(Filter.by_property(key).equal(value)) combined_filter = Filter.all_of(filter_conditions) if filter_conditions else None response = collection.query.fetch_objects( - limit=limit, + limit=top_k, filters=combined_filter, return_properties=["hash", "created_at", "updated_at", "user_id", "agent_id", "run_id", "data", "category"], ) diff --git a/server/main.py b/server/main.py index 3be534dcd..961cb2f2a 100644 --- a/server/main.py +++ b/server/main.py @@ -135,7 +135,7 @@ class SearchRequest(BaseModel): run_id: Optional[str] = None agent_id: Optional[str] = None filters: Optional[Dict[str, Any]] = None - limit: Optional[int] = Field(None, description="Maximum number of results to return.") + top_k: Optional[int] = Field(None, description="Maximum number of results to return.") threshold: Optional[float] = Field(None, description="Minimum similarity score for results.") diff --git a/skills/mem0/client/differences.md b/skills/mem0/client/differences.md index d4849115d..37344b277 100644 --- a/skills/mem0/client/differences.md +++ b/skills/mem0/client/differences.md @@ -99,7 +99,7 @@ These methods exist in Python but not TypeScript: | `vector_store` | `vectorStore` | | `graph_store` | `graphStore` | | `history_db_path` | `historyDbPath` | -| `custom_fact_extraction_prompt` | `customPrompt` | +| `custom_instructions` | `customInstructions` | | `enable_graph` | `enableGraph` | ## OSS Scope Parameter Naming diff --git a/skills/mem0/client/python.md b/skills/mem0/client/python.md index 280373f99..e9a5fe002 100644 --- a/skills/mem0/client/python.md +++ b/skills/mem0/client/python.md @@ -390,7 +390,7 @@ config = { } }, "history_db_path": "history.db", # SQLite path for change history - "custom_fact_extraction_prompt": "...", # Custom LLM prompt for extraction + "custom_instructions": "...", # Custom LLM prompt for extraction "custom_update_memory_prompt": "...", # Custom LLM prompt for updates "enable_graph": False, # Enable graph memory } diff --git a/tests/memory/test_apache_age_e2e.py b/tests/memory/test_apache_age_e2e.py index 5919136d6..c55ea74a0 100644 --- a/tests/memory/test_apache_age_e2e.py +++ b/tests/memory/test_apache_age_e2e.py @@ -16,10 +16,10 @@ Run: import json import os -import pytest from unittest.mock import MagicMock, patch import age +import pytest from mem0.memory.apache_age_memory import MemoryGraph # noqa: E402 @@ -386,7 +386,7 @@ class TestPublicAPICRUD: ) self.mg.ag.commit() - results = self.mg.get_all({"user_id": "u1"}, limit=3) + results = self.mg.get_all({"user_id": "u1"}, top_k=3) assert len(results) == 3 def test_delete_all_removes_user_data(self): diff --git a/tests/memory/test_apache_age_memory.py b/tests/memory/test_apache_age_memory.py index e0823cb4d..c0cc5e086 100644 --- a/tests/memory/test_apache_age_memory.py +++ b/tests/memory/test_apache_age_memory.py @@ -175,7 +175,7 @@ class TestGetAll: {"source": "alice", "relationship": "KNOWS", "target": "bob"}, {"source": "alice", "relationship": "LIKES", "target": "hiking"}, ]) - results = instance.get_all({"user_id": "u1"}, limit=10) + results = instance.get_all({"user_id": "u1"}, top_k=10) assert len(results) == 2 assert results[0]["source"] == "alice" assert results[0]["relationship"] == "KNOWS" @@ -187,7 +187,7 @@ class TestGetAll: instance._exec_cypher = MagicMock(return_value=[ {"source": "n0", "relationship": "R", "target": "m0"}, ]) - instance.get_all({"user_id": "u1"}, limit=3) + instance.get_all({"user_id": "u1"}, top_k=3) # Verify limit was passed as a parameter to the query cypher_stmt = instance._exec_cypher.call_args[0][0] assert "LIMIT %s" in cypher_stmt diff --git a/tests/memory/test_graph_memory_soft_delete.py b/tests/memory/test_graph_memory_soft_delete.py index 111345851..7b17f5c5d 100644 --- a/tests/memory/test_graph_memory_soft_delete.py +++ b/tests/memory/test_graph_memory_soft_delete.py @@ -88,7 +88,7 @@ class TestSearchExcludesSoftDeleted: def test_get_all_filters_soft_deleted(self): mg = _create_graph_memory() - mg.get_all(filters={"user_id": "user1"}, limit=10) + mg.get_all(filters={"user_id": "user1"}, top_k=10) cypher = mg.graph.query.call_args[0][0] assert "r.valid IS NULL OR r.valid = true" in cypher @@ -271,7 +271,7 @@ class TestSoftDeleteWithFilters: def test_get_all_with_agent_id_filters_soft_deleted(self): mg = _create_graph_memory() - mg.get_all(filters={"user_id": "user1", "agent_id": "agent1"}, limit=10) + mg.get_all(filters={"user_id": "user1", "agent_id": "agent1"}, top_k=10) cypher = mg.graph.query.call_args[0][0] assert "r.valid IS NULL OR r.valid = true" in cypher diff --git a/tests/memory/test_main.py b/tests/memory/test_main.py index 15507e2be..c4aaea953 100644 --- a/tests/memory/test_main.py +++ b/tests/memory/test_main.py @@ -35,7 +35,7 @@ class TestAddToVectorStoreErrors: memory = Memory() memory.config = mocker.MagicMock() - memory.config.custom_fact_extraction_prompt = None + memory.config.custom_instructions = None memory.config.custom_update_memory_prompt = None memory.api_version = "v1.1" @@ -139,7 +139,7 @@ class TestAsyncAddToVectorStoreErrors: memory = AsyncMemory() memory.config = mocker.MagicMock() - memory.config.custom_fact_extraction_prompt = None + memory.config.custom_instructions = None memory.config.custom_update_memory_prompt = None memory.api_version = "v1.1" @@ -187,7 +187,7 @@ def _build_memory_instance(mocker, memory_cls): mocker.patch("mem0.memory.main.MEM0_TELEMETRY", False) memory = memory_cls() memory.config = mocker.MagicMock() - memory.config.custom_fact_extraction_prompt = None + memory.config.custom_instructions = None memory.config.custom_update_memory_prompt = None memory.api_version = "v1.1" memory.vector_store = mocker.MagicMock() @@ -309,8 +309,8 @@ def test_create_then_search_and_get_all_return_same_timestamps(mocker): memory.vector_store.list.return_value = [[mem_result]] # Step 3: Call search and get_all, compare timestamps - search_results = memory._search_vector_store("pizza", filters={"user_id": "alice"}, limit=10, threshold=None) - get_all_results = memory._get_all_from_vector_store(filters={"user_id": "alice"}, limit=100) + search_results = memory._search_vector_store("pizza", filters={"user_id": "alice"}, top_k=10, threshold=None) + get_all_results = memory._get_all_from_vector_store(filters={"user_id": "alice"}, top_k=100) search_item = search_results[0] get_all_item = get_all_results[0] @@ -376,8 +376,8 @@ def test_search_and_get_all_consistent_after_update(mocker): memory.vector_store.search.return_value = [mem_result] memory.vector_store.list.return_value = [[mem_result]] - search_results = memory._search_vector_store("pizza", filters={"user_id": "alice"}, limit=10, threshold=None) - get_all_results = memory._get_all_from_vector_store(filters={"user_id": "alice"}, limit=100) + search_results = memory._search_vector_store("pizza", filters={"user_id": "alice"}, top_k=10, threshold=None) + get_all_results = memory._get_all_from_vector_store(filters={"user_id": "alice"}, top_k=100) assert search_results[0]["created_at"] == get_all_results[0]["created_at"] assert search_results[0]["updated_at"] == get_all_results[0]["updated_at"] diff --git a/tests/memory/test_neptune_analytics_memory.py b/tests/memory/test_neptune_analytics_memory.py index 46d8ec90d..dbec73fd8 100644 --- a/tests/memory/test_neptune_analytics_memory.py +++ b/tests/memory/test_neptune_analytics_memory.py @@ -1,8 +1,10 @@ import unittest from unittest.mock import MagicMock, patch + import pytest -from mem0.graphs.neptune.neptunegraph import MemoryGraph + from mem0.graphs.neptune.base import NeptuneBase +from mem0.graphs.neptune.neptunegraph import MemoryGraph class TestNeptuneMemory(unittest.TestCase): @@ -134,7 +136,7 @@ class TestNeptuneMemory(unittest.TestCase): mock_bm25_instance.get_top_n.return_value = reranked_results # Call the search method - result = self.memory_graph.search("Find Alice", self.test_filters, limit=5) + result = self.memory_graph.search("Find Alice", self.test_filters, top_k=5) # Verify the method calls self.memory_graph._retrieve_nodes_from_data.assert_called_once_with("Find Alice", self.test_filters) @@ -162,7 +164,7 @@ class TestNeptuneMemory(unittest.TestCase): self.mock_graph.query.return_value = mock_query_result # Call the get_all method - result = self.memory_graph.get_all(self.test_filters, limit=10) + result = self.memory_graph.get_all(self.test_filters, top_k=10) # Verify the method calls self.memory_graph._get_all_cypher.assert_called_once_with(self.test_filters, 10) @@ -256,7 +258,7 @@ class TestNeptuneMemory(unittest.TestCase): self.mock_graph.query.side_effect = [mock_query_result1, mock_query_result2] # Call the _search_graph_db method - result = self.memory_graph._search_graph_db(node_list, self.test_filters, limit=10) + result = self.memory_graph._search_graph_db(node_list, self.test_filters, top_k=10) # Verify the method calls self.assertEqual(self.mock_embedding_model.embed.call_count, 2) diff --git a/tests/memory/test_neptune_memory.py b/tests/memory/test_neptune_memory.py index c491b99af..ad28b6e34 100644 --- a/tests/memory/test_neptune_memory.py +++ b/tests/memory/test_neptune_memory.py @@ -190,7 +190,7 @@ class TestNeptuneMemory(unittest.TestCase): mock_bm25_instance.get_top_n.return_value = reranked_results # Call the search method - result = self.memory_graph.search("Find Alice", self.test_filters, limit=5) + result = self.memory_graph.search("Find Alice", self.test_filters, top_k=5) # Verify the method calls self.memory_graph._retrieve_nodes_from_data.assert_called_once_with("Find Alice", self.test_filters) @@ -218,7 +218,7 @@ class TestNeptuneMemory(unittest.TestCase): self.mock_graph.query.return_value = mock_query_result # Call the get_all method - result = self.memory_graph.get_all(self.test_filters, limit=10) + result = self.memory_graph.get_all(self.test_filters, top_k=10) # Verify the method calls self.memory_graph._get_all_cypher.assert_called_once_with(self.test_filters, 10) @@ -331,7 +331,7 @@ class TestNeptuneMemory(unittest.TestCase): self.mock_graph.query.side_effect = [mock_query_result1, mock_query_result2] # Call the _search_graph_db method - result = self.memory_graph._search_graph_db(node_list, self.test_filters, limit=10) + result = self.memory_graph._search_graph_db(node_list, self.test_filters, top_k=10) # Verify the method calls self.assertEqual(self.mock_embedding_model.embed.call_count, 2) diff --git a/tests/test_main.py b/tests/test_main.py index 37cb11015..2a12ce55e 100644 --- a/tests/test_main.py +++ b/tests/test_main.py @@ -60,7 +60,7 @@ def memory_custom_instance(): config = MemoryConfig( version="v1.1", - custom_fact_extraction_prompt="custom prompt extracting memory in json format", + custom_instructions="custom prompt extracting memory in json format", custom_update_memory_prompt="custom prompt determining memory update", ) config.graph_store.config = {"some_config": "value"} @@ -156,7 +156,7 @@ def test_search(memory_instance, version, enable_graph): assert result["results"][0]["score"] == 0.9 memory_instance.vector_store.search.assert_called_once_with( - query="test query", vectors=[0.1, 0.2, 0.3], limit=100, filters={"user_id": "test_user"} + query="test query", vectors=[0.1, 0.2, 0.3], top_k=100, filters={"user_id": "test_user"} ) memory_instance.embedding_model.embed.assert_called_once_with("test query", "search") @@ -286,7 +286,7 @@ def test_get_all(memory_instance, version, enable_graph, expected_result): else: assert "relations" not in result - memory_instance.vector_store.list.assert_called_once_with(filters={"user_id": "test_user"}, limit=100) + memory_instance.vector_store.list.assert_called_once_with(filters={"user_id": "test_user"}, top_k=100) if enable_graph: memory_instance.graph.get_all.assert_called_once_with({"user_id": "test_user"}, 100) @@ -314,7 +314,7 @@ def test_custom_prompts(memory_custom_instance): memory_custom_instance.llm.generate_response.assert_any_call( messages=[ - {"role": "system", "content": memory_custom_instance.config.custom_fact_extraction_prompt}, + {"role": "system", "content": memory_custom_instance.config.custom_instructions}, {"role": "user", "content": f"Input:\n{mock_parse_messages.return_value}"}, ], response_format={"type": "json_object"}, diff --git a/tests/vector_stores/test_azure_ai_search.py b/tests/vector_stores/test_azure_ai_search.py index f9815be29..1478d90c4 100644 --- a/tests/vector_stores/test_azure_ai_search.py +++ b/tests/vector_stores/test_azure_ai_search.py @@ -533,7 +533,7 @@ def test_search_basic(azure_ai_search_instance): # Search with a vector query_text = "test query" # Add a query string query_vector = [0.1, 0.2, 0.3] - results = instance.search(query_text, query_vector, limit=5) # Pass the query string + results = instance.search(query_text, query_vector, top_k=5) # Pass the query string # Verify search was called correctly mock_search_client.search.assert_called_once() diff --git a/tests/vector_stores/test_azure_mysql.py b/tests/vector_stores/test_azure_mysql.py index 1b6dd6ef3..05b61dc4c 100644 --- a/tests/vector_stores/test_azure_mysql.py +++ b/tests/vector_stores/test_azure_mysql.py @@ -1,7 +1,8 @@ import json -import pytest from unittest.mock import Mock, patch +import pytest + from mem0.vector_stores.azure_mysql import AzureMySQL, OutputData @@ -118,7 +119,7 @@ def test_search(azure_mysql_instance): ]) query_vector = [0.2, 0.3, 0.4] - results = azure_mysql_instance.search(query="test", vectors=query_vector, limit=5) + results = azure_mysql_instance.search(query="test", vectors=query_vector, top_k=5) assert isinstance(results, list) assert cursor.execute.called @@ -218,7 +219,7 @@ def test_list(azure_mysql_instance): } ]) - results = azure_mysql_instance.list(limit=10) + results = azure_mysql_instance.list(top_k=10) assert isinstance(results, list) assert len(results) > 0 diff --git a/tests/vector_stores/test_baidu.py b/tests/vector_stores/test_baidu.py index 981c7790a..cd615feb7 100644 --- a/tests/vector_stores/test_baidu.py +++ b/tests/vector_stores/test_baidu.py @@ -112,7 +112,7 @@ def test_search(mochow_instance, mock_mochow_client): mochow_instance._table.vector_search.return_value = mock_search_results vectors = [0.1, 0.2, 0.3] - results = mochow_instance.search(query="test", vectors=vectors, limit=2) + results = mochow_instance.search(query="test", vectors=vectors, top_k=2) # Verify search was called with correct parameters mochow_instance._table.vector_search.assert_called_once() @@ -143,7 +143,7 @@ def test_search_with_filters(mochow_instance, mock_mochow_client): vectors = [0.1, 0.2, 0.3] filters = {"user_id": "user123", "agent_id": "agent456"} - mochow_instance.search(query="test", vectors=vectors, limit=2, filters=filters) + mochow_instance.search(query="test", vectors=vectors, top_k=2, filters=filters) # Verify search was called with filter call_args = mochow_instance._table.vector_search.call_args @@ -196,7 +196,7 @@ def test_list(mochow_instance, mock_mochow_client): mock_result.rows = [{"id": "id1", "metadata": {"name": "vector1"}}, {"id": "id2", "metadata": {"name": "vector2"}}] mochow_instance._table.select.return_value = mock_result - results = mochow_instance.list(limit=2) + results = mochow_instance.list(top_k=2) mochow_instance._table.select.assert_called_once_with(filter=None, projections=["id", "metadata"], limit=2) diff --git a/tests/vector_stores/test_cassandra.py b/tests/vector_stores/test_cassandra.py index 3194e4a8d..edb1207dc 100644 --- a/tests/vector_stores/test_cassandra.py +++ b/tests/vector_stores/test_cassandra.py @@ -1,7 +1,8 @@ import json -import pytest from unittest.mock import Mock, patch +import pytest + from mem0.vector_stores.cassandra import CassandraDB, OutputData @@ -106,7 +107,7 @@ def test_search(cassandra_instance): cassandra_instance.session.execute = Mock(return_value=[mock_row1, mock_row2]) query_vector = [0.2, 0.3, 0.4] - results = cassandra_instance.search(query="test", vectors=query_vector, limit=5) + results = cassandra_instance.search(query="test", vectors=query_vector, top_k=5) assert isinstance(results, list) assert len(results) <= 5 @@ -213,7 +214,7 @@ def test_list(cassandra_instance): cassandra_instance.session.execute = Mock(return_value=[mock_row]) - results = cassandra_instance.list(limit=10) + results = cassandra_instance.list(top_k=10) assert isinstance(results, list) assert len(results) > 0 @@ -264,7 +265,7 @@ def test_search_with_filters(cassandra_instance): results = cassandra_instance.search( query="test", vectors=query_vector, - limit=5, + top_k=5, filters={"category": "A"} ) diff --git a/tests/vector_stores/test_chroma.py b/tests/vector_stores/test_chroma.py index 57c16d4d9..e8d2e3c28 100644 --- a/tests/vector_stores/test_chroma.py +++ b/tests/vector_stores/test_chroma.py @@ -2,8 +2,8 @@ from unittest.mock import Mock, patch import pytest -from mem0.vector_stores.chroma import ChromaDB from mem0.configs.vector_stores.chroma import ChromaDbConfig +from mem0.vector_stores.chroma import ChromaDB @pytest.fixture @@ -39,7 +39,7 @@ def test_search_vectors(chromadb_instance, mock_chromadb_client): chromadb_instance.collection.query.return_value = mock_result vectors = [[0.1, 0.2, 0.3]] - results = chromadb_instance.search(query="", vectors=vectors, limit=2) + results = chromadb_instance.search(query="", vectors=vectors, top_k=2) chromadb_instance.collection.query.assert_called_once_with(query_embeddings=vectors, where=None, n_results=2) @@ -60,7 +60,7 @@ def test_search_vectors_with_filters(chromadb_instance, mock_chromadb_client): vectors = [[0.1, 0.2, 0.3]] filters = {"user_id": "alice", "agent_id": "agent1", "run_id": "run1"} - results = chromadb_instance.search(query="", vectors=vectors, limit=2, filters=filters) + results = chromadb_instance.search(query="", vectors=vectors, top_k=2, filters=filters) # Verify that _generate_where_clause was called with the filters expected_where = {"$and": [{"user_id": {"$eq": "alice"}}, {"agent_id": {"$eq": "agent1"}}, {"run_id": {"$eq": "run1"}}]} @@ -86,7 +86,7 @@ def test_search_vectors_with_single_filter(chromadb_instance, mock_chromadb_clie vectors = [[0.1, 0.2, 0.3]] filters = {"user_id": "alice"} - results = chromadb_instance.search(query="", vectors=vectors, limit=2, filters=filters) + results = chromadb_instance.search(query="", vectors=vectors, top_k=2, filters=filters) # Verify that single filter is passed with $eq operator expected_where = {"user_id": {"$eq": "alice"}} @@ -108,7 +108,7 @@ def test_search_vectors_with_no_filters(chromadb_instance, mock_chromadb_client) chromadb_instance.collection.query.return_value = mock_result vectors = [[0.1, 0.2, 0.3]] - results = chromadb_instance.search(query="", vectors=vectors, limit=2, filters=None) + results = chromadb_instance.search(query="", vectors=vectors, top_k=2, filters=None) chromadb_instance.collection.query.assert_called_once_with( query_embeddings=vectors, where=None, n_results=2 @@ -162,7 +162,7 @@ def test_list_vectors(chromadb_instance): } chromadb_instance.collection.get.return_value = mock_result - results = chromadb_instance.list(limit=2) + results = chromadb_instance.list(top_k=2) chromadb_instance.collection.get.assert_called_once_with(where=None, limit=2) @@ -181,7 +181,7 @@ def test_list_vectors_with_filters(chromadb_instance): chromadb_instance.collection.get.return_value = mock_result filters = {"user_id": "alice", "agent_id": "agent1", "run_id": "run1"} - results = chromadb_instance.list(filters=filters, limit=2) + results = chromadb_instance.list(filters=filters, top_k=2) # Verify that _generate_where_clause was called with the filters expected_where = {"$and": [{"user_id": {"$eq": "alice"}}, {"agent_id": {"$eq": "agent1"}}, {"run_id": {"$eq": "run1"}}]} @@ -203,7 +203,7 @@ def test_list_vectors_with_single_filter(chromadb_instance): chromadb_instance.collection.get.return_value = mock_result filters = {"user_id": "alice"} - results = chromadb_instance.list(filters=filters, limit=2) + results = chromadb_instance.list(filters=filters, top_k=2) # Verify that single filter is passed with $eq operator expected_where = {"user_id": {"$eq": "alice"}} diff --git a/tests/vector_stores/test_databricks.py b/tests/vector_stores/test_databricks.py index 93c0324d5..1a970d7a3 100644 --- a/tests/vector_stores/test_databricks.py +++ b/tests/vector_stores/test_databricks.py @@ -5,9 +5,15 @@ import pytest pytest.importorskip("databricks", reason="databricks-sdk package not installed") -from databricks.sdk.service.vectorsearch import VectorIndexType, QueryVectorIndexResponse, ResultManifest, ResultData, ColumnInfo -from mem0.vector_stores.databricks import Databricks +from databricks.sdk.service.vectorsearch import ( + ColumnInfo, + QueryVectorIndexResponse, + ResultData, + ResultManifest, + VectorIndexType, +) +from mem0.vector_stores.databricks import Databricks # ---------------------- Fixtures ---------------------- # @@ -185,7 +191,7 @@ def test_search_delta_sync_text(db_instance_delta, mock_workspace_client): mock_workspace_client.vector_search_indexes.query_index.return_value = SimpleNamespace( result=SimpleNamespace(data_array=[row]) ) - results = db_instance_delta.search(query="hello", vectors=None, limit=1) + results = db_instance_delta.search(query="hello", vectors=None, top_k=1) mock_workspace_client.vector_search_indexes.query_index.assert_called_once() assert len(results) == 1 assert results[0].id == "id1" @@ -210,7 +216,7 @@ def test_search_direct_access_vector(db_instance_direct, mock_workspace_client): mock_workspace_client.vector_search_indexes.query_index.return_value = SimpleNamespace( result=SimpleNamespace(data_array=[row]) ) - results = db_instance_direct.search(query="", vectors=[0.1, 0.2, 0.3, 0.4], limit=1) + results = db_instance_direct.search(query="", vectors=[0.1, 0.2, 0.3, 0.4], top_k=1) assert len(results) == 1 assert results[0].id == "id2" assert results[0].score == 0.77 @@ -235,7 +241,7 @@ def test_search_delta_sync_self_managed_vectors(mock_workspace_client): mock_workspace_client.vector_search_indexes.query_index.return_value = SimpleNamespace( result=SimpleNamespace(data_array=[]) ) - inst.search(query="ignored", vectors=[0.1, 0.2, 0.3, 0.4], limit=5) + inst.search(query="ignored", vectors=[0.1, 0.2, 0.3, 0.4], top_k=5) call_kwargs = mock_workspace_client.vector_search_indexes.query_index.call_args.kwargs assert "query_vector" in call_kwargs assert "query_text" not in call_kwargs @@ -423,7 +429,7 @@ def test_list_memories(db_instance_delta, mock_workspace_client): ] ) ) - res = db_instance_delta.list(limit=1) + res = db_instance_delta.list(top_k=1) assert isinstance(res, list) assert len(res[0]) == 1 assert res[0][0].id == "id-get" @@ -453,7 +459,7 @@ def test_list_memories_direct_access(db_instance_direct, mock_workspace_client): ] ) ) - res = db_instance_direct.list(limit=5) + res = db_instance_direct.list(top_k=5) assert isinstance(res, list) assert len(res[0]) == 1 assert res[0][0].id == "id-da-list" @@ -516,7 +522,7 @@ def test_list_memories_delta_sync_self_managed(mock_workspace_client): mock_workspace_client.vector_search_indexes.query_index.return_value = SimpleNamespace( result=SimpleNamespace(data_array=[]) ) - inst.list(limit=5) + inst.list(top_k=5) call_kwargs = mock_workspace_client.vector_search_indexes.query_index.call_args.kwargs assert "query_vector" in call_kwargs assert "query_text" not in call_kwargs @@ -527,7 +533,7 @@ def test_list_memories_default_limit(db_instance_delta, mock_workspace_client): mock_workspace_client.vector_search_indexes.query_index.return_value = SimpleNamespace( result=SimpleNamespace(data_array=[]) ) - db_instance_delta.list(limit=None) + db_instance_delta.list(top_k=None) call_kwargs = mock_workspace_client.vector_search_indexes.query_index.call_args.kwargs assert call_kwargs["num_results"] == 100 @@ -624,8 +630,8 @@ def test_reset(db_instance_delta, mock_workspace_client): def test_e2e_config_to_factory_delta_sync(mock_workspace_client): """End-to-end: VectorStoreConfig validates docs-correct params, factory creates Databricks instance.""" - from mem0.vector_stores.configs import VectorStoreConfig from mem0.utils.factory import VectorStoreFactory + from mem0.vector_stores.configs import VectorStoreConfig # Step 1: Config validation (simulates what Memory.from_config does) vs_config = VectorStoreConfig( @@ -655,8 +661,8 @@ def test_e2e_config_to_factory_delta_sync(mock_workspace_client): def test_e2e_config_to_factory_direct_access(mock_workspace_client): """End-to-end: DIRECT_ACCESS via config → factory creates correct instance.""" - from mem0.vector_stores.configs import VectorStoreConfig from mem0.utils.factory import VectorStoreFactory + from mem0.vector_stores.configs import VectorStoreConfig mock_workspace_client.tables.exists.return_value = SimpleNamespace(table_exists=True) @@ -699,8 +705,8 @@ def test_e2e_old_docs_config_rejected(): def test_e2e_crud_lifecycle_delta_sync(mock_workspace_client): """End-to-end CRUD lifecycle: insert → search → get → list → update → delete.""" - from mem0.vector_stores.configs import VectorStoreConfig from mem0.utils.factory import VectorStoreFactory + from mem0.vector_stores.configs import VectorStoreConfig vs_config = VectorStoreConfig( provider="databricks", @@ -736,7 +742,7 @@ def test_e2e_crud_lifecycle_delta_sync(mock_workspace_client): data_array=[["mem-001", "h1", None, None, "u1", "test memory", None, None, None, 0.95]] ) ) - results = db.search(query="test", vectors=None, limit=5) + results = db.search(query="test", vectors=None, top_k=5) assert len(results) == 1 assert results[0].id == "mem-001" assert results[0].payload["data"] == "test memory" @@ -767,7 +773,7 @@ def test_e2e_crud_lifecycle_delta_sync(mock_workspace_client): data_array=[["mem-001", "h1", None, None, "u1", "test memory", None, None, None]] ) ) - listed = db.list(filters={"user_id": "u1"}, limit=10) + listed = db.list(filters={"user_id": "u1"}, top_k=10) assert len(listed[0]) == 1 assert listed[0][0].id == "mem-001" list_kwargs = mock_workspace_client.vector_search_indexes.query_index.call_args.kwargs @@ -797,8 +803,8 @@ def test_e2e_crud_lifecycle_delta_sync(mock_workspace_client): def test_e2e_crud_lifecycle_direct_access(mock_workspace_client): """End-to-end CRUD lifecycle for DIRECT_ACCESS: insert → search → get → list.""" - from mem0.vector_stores.configs import VectorStoreConfig from mem0.utils.factory import VectorStoreFactory + from mem0.vector_stores.configs import VectorStoreConfig mock_workspace_client.tables.exists.return_value = SimpleNamespace(table_exists=True) @@ -838,7 +844,7 @@ def test_e2e_crud_lifecycle_direct_access(mock_workspace_client): data_array=[["mem-da-001", "h1", None, None, "u1", "direct memory", None, None, None, [0.1, 0.2, 0.3, 0.4], 0.9]] ) ) - results = db.search(query="", vectors=[0.1, 0.2, 0.3, 0.4], limit=5) + results = db.search(query="", vectors=[0.1, 0.2, 0.3, 0.4], top_k=5) assert len(results) == 1 search_kwargs = mock_workspace_client.vector_search_indexes.query_index.call_args.kwargs assert "query_vector" in search_kwargs @@ -867,7 +873,7 @@ def test_e2e_crud_lifecycle_direct_access(mock_workspace_client): data_array=[["mem-da-001", "h1", None, None, "u1", "direct memory", None, None, None, [0.1, 0.2, 0.3, 0.4]]] ) ) - db.list(limit=5) + db.list(top_k=5) list_kwargs = mock_workspace_client.vector_search_indexes.query_index.call_args.kwargs assert "query_vector" in list_kwargs assert "query_text" not in list_kwargs diff --git a/tests/vector_stores/test_elasticsearch.py b/tests/vector_stores/test_elasticsearch.py index db7a82e14..61ff9efbd 100644 --- a/tests/vector_stores/test_elasticsearch.py +++ b/tests/vector_stores/test_elasticsearch.py @@ -9,8 +9,8 @@ try: except ImportError: raise ImportError("Elasticsearch requires extra dependencies. Install with `pip install elasticsearch`") from None -from mem0.vector_stores.elasticsearch import ElasticsearchDB, OutputData from mem0.configs.vector_stores.elasticsearch import ElasticsearchConfig +from mem0.vector_stores.elasticsearch import ElasticsearchDB, OutputData class TestElasticsearchDB(unittest.TestCase): @@ -189,7 +189,7 @@ class TestElasticsearchDB(unittest.TestCase): # Perform search vectors = [[0.1] * 1536] - results = self.es_db.search(query="", vectors=vectors, limit=5) + results = self.es_db.search(query="", vectors=vectors, top_k=5) # Verify search call self.client_mock.search.assert_called_once() @@ -221,7 +221,7 @@ class TestElasticsearchDB(unittest.TestCase): vectors = [[0.1] * 1536] limit = 5 filters = {"key1": "value1"} - self.es_db.search(query="", vectors=vectors, limit=limit, filters=filters) + self.es_db.search(query="", vectors=vectors, top_k=limit, filters=filters) # Verify custom search query function was called self.es_db.custom_search_query.assert_called_once_with(vectors, limit, filters) @@ -272,7 +272,7 @@ class TestElasticsearchDB(unittest.TestCase): self.client_mock.search.return_value = mock_response # Perform list operation - results = self.es_db.list(limit=10) + results = self.es_db.list(top_k=10) # Verify search call self.client_mock.search.assert_called_once() diff --git a/tests/vector_stores/test_faiss.py b/tests/vector_stores/test_faiss.py index 120b65caf..dc730f4b8 100644 --- a/tests/vector_stores/test_faiss.py +++ b/tests/vector_stores/test_faiss.py @@ -101,7 +101,7 @@ def test_search(faiss_instance, mock_faiss_index): with patch.object(faiss_instance, "_parse_output", return_value=expected_results): # Call search - results = faiss_instance.search(query="test query", vectors=query_vector, limit=2) + results = faiss_instance.search(query="test query", vectors=query_vector, top_k=2) # Verify numpy.array was called (but we don't check exact call arguments since it's complex) assert mock_np_array.called @@ -148,7 +148,7 @@ def test_search_with_filters(faiss_instance, mock_faiss_index): with patch.object(faiss_instance, "_apply_filters", side_effect=lambda p, f: p.get("category") == "A"): # Call search with filters results = faiss_instance.search( - query="test query", vectors=query_vector, limit=2, filters={"category": "A"} + query="test query", vectors=query_vector, top_k=2, filters={"category": "A"} ) # Verify numpy.array was called @@ -241,7 +241,7 @@ def test_list(faiss_instance): assert len(results[0]) == 3 # Test listing with a limit - results = faiss_instance.list(limit=2) + results = faiss_instance.list(top_k=2) assert len(results[0]) == 2 # Test listing with filters diff --git a/tests/vector_stores/test_langchain_vector_store.py b/tests/vector_stores/test_langchain_vector_store.py index 6dd85d042..fc32adc58 100644 --- a/tests/vector_stores/test_langchain_vector_store.py +++ b/tests/vector_stores/test_langchain_vector_store.py @@ -48,7 +48,7 @@ def test_search_vectors(langchain_instance): # Test search without filters vectors = [[0.1, 0.2, 0.3]] - results = langchain_instance.search(query="", vectors=vectors, limit=2) + results = langchain_instance.search(query="", vectors=vectors, top_k=2) langchain_instance.client.similarity_search_by_vector.assert_called_once_with(embedding=vectors, k=2) @@ -60,7 +60,7 @@ def test_search_vectors(langchain_instance): # Test search with filters filters = {"name": "vector1"} - langchain_instance.search(query="", vectors=vectors, limit=2, filters=filters) + langchain_instance.search(query="", vectors=vectors, top_k=2, filters=filters) langchain_instance.client.similarity_search_by_vector.assert_called_with(embedding=vectors, k=2, filter=filters) @@ -75,7 +75,7 @@ def test_search_vectors_with_agent_id_run_id_filters(langchain_instance): vectors = [[0.1, 0.2, 0.3]] filters = {"user_id": "alice", "agent_id": "agent1", "run_id": "run1"} - results = langchain_instance.search(query="", vectors=vectors, limit=2, filters=filters) + results = langchain_instance.search(query="", vectors=vectors, top_k=2, filters=filters) # Verify that filters were passed to the underlying vector store langchain_instance.client.similarity_search_by_vector.assert_called_once_with( @@ -96,7 +96,7 @@ def test_search_vectors_with_single_filter(langchain_instance): vectors = [[0.1, 0.2, 0.3]] filters = {"user_id": "alice"} - results = langchain_instance.search(query="", vectors=vectors, limit=2, filters=filters) + results = langchain_instance.search(query="", vectors=vectors, top_k=2, filters=filters) # Verify that filters were passed to the underlying vector store langchain_instance.client.similarity_search_by_vector.assert_called_once_with( @@ -114,7 +114,7 @@ def test_search_vectors_with_no_filters(langchain_instance): langchain_instance.client.similarity_search_by_vector.return_value = mock_docs vectors = [[0.1, 0.2, 0.3]] - results = langchain_instance.search(query="", vectors=vectors, limit=2, filters=None) + results = langchain_instance.search(query="", vectors=vectors, top_k=2, filters=None) # Verify that no filters were passed to the underlying vector store langchain_instance.client.similarity_search_by_vector.assert_called_once_with( @@ -155,7 +155,7 @@ def test_list_with_filters(langchain_instance): langchain_instance.client._collection = mock_collection filters = {"user_id": "alice", "agent_id": "agent1", "run_id": "run1"} - results = langchain_instance.list(filters=filters, limit=10) + results = langchain_instance.list(filters=filters, top_k=10) # Verify that the collection.get method was called with the correct filters mock_collection.get.assert_called_once_with(where=filters, limit=10) @@ -180,7 +180,7 @@ def test_list_with_single_filter(langchain_instance): langchain_instance.client._collection = mock_collection filters = {"user_id": "alice"} - results = langchain_instance.list(filters=filters, limit=10) + results = langchain_instance.list(filters=filters, top_k=10) # Verify that the collection.get method was called with the correct filter mock_collection.get.assert_called_once_with(where=filters, limit=10) @@ -202,7 +202,7 @@ def test_list_with_no_filters(langchain_instance): } langchain_instance.client._collection = mock_collection - results = langchain_instance.list(filters=None, limit=10) + results = langchain_instance.list(filters=None, top_k=10) # Verify that the collection.get method was called with no filters mock_collection.get.assert_called_once_with(where=None, limit=10) @@ -220,7 +220,7 @@ def test_list_with_exception(langchain_instance): mock_collection.get.side_effect = Exception("Test exception") langchain_instance.client._collection = mock_collection - results = langchain_instance.list(filters={"user_id": "alice"}, limit=10) + results = langchain_instance.list(filters={"user_id": "alice"}, top_k=10) # Verify that an empty list is returned when an exception occurs assert results == [] diff --git a/tests/vector_stores/test_milvus.py b/tests/vector_stores/test_milvus.py index 976056a25..637411a2e 100644 --- a/tests/vector_stores/test_milvus.py +++ b/tests/vector_stores/test_milvus.py @@ -8,10 +8,12 @@ These tests verify: 4. Update/upsert operations """ -import pytest from unittest.mock import MagicMock, patch -from mem0.vector_stores.milvus import MilvusDB + +import pytest + from mem0.configs.vector_stores.milvus import MetricType +from mem0.vector_stores.milvus import MilvusDB class TestMilvusDB: @@ -118,7 +120,7 @@ class TestMilvusDB: results = milvus_db.search( query="test query", vectors=query_vector, - limit=5, + top_k=5, filters=filters ) @@ -199,7 +201,7 @@ class TestMilvusDB: {"id": "mem2", "metadata": {"user_id": "alice"}} ] - results = milvus_db.list(filters={"user_id": "alice"}, limit=10) + results = milvus_db.list(filters={"user_id": "alice"}, top_k=10) # Verify query was called with filter call_args = mock_milvus_client.query.call_args diff --git a/tests/vector_stores/test_mongodb.py b/tests/vector_stores/test_mongodb.py index 33a8a7ad0..215f7f7c0 100644 --- a/tests/vector_stores/test_mongodb.py +++ b/tests/vector_stores/test_mongodb.py @@ -86,7 +86,7 @@ def test_search(mongo_vector_fixture): ] mock_collection.list_search_indexes.return_value = ["test_collection_vector_index"] - results = mongo_vector.search("query_str", query_vector, limit=2) + results = mongo_vector.search("query_str", query_vector, top_k=2) mock_collection.list_search_indexes.assert_called_with(name="test_collection_vector_index") mock_collection.aggregate.assert_called_once_with( [ @@ -119,7 +119,7 @@ def test_search_with_filters(mongo_vector_fixture): mock_collection.list_search_indexes.return_value = ["test_collection_vector_index"] filters = {"user_id": "alice", "agent_id": "agent1", "run_id": "run1"} - results = mongo_vector.search("query_str", query_vector, limit=2, filters=filters) + results = mongo_vector.search("query_str", query_vector, top_k=2, filters=filters) # Verify that the aggregation pipeline includes the filter stage mock_collection.aggregate.assert_called_once() @@ -153,7 +153,7 @@ def test_search_with_single_filter(mongo_vector_fixture): mock_collection.list_search_indexes.return_value = ["test_collection_vector_index"] filters = {"user_id": "alice"} - results = mongo_vector.search("query_str", query_vector, limit=2, filters=filters) + results = mongo_vector.search("query_str", query_vector, top_k=2, filters=filters) # Verify that the aggregation pipeline includes the filter stage mock_collection.aggregate.assert_called_once() @@ -177,7 +177,7 @@ def test_search_with_no_filters(mongo_vector_fixture): ] mock_collection.list_search_indexes.return_value = ["test_collection_vector_index"] - results = mongo_vector.search("query_str", query_vector, limit=2, filters=None) + results = mongo_vector.search("query_str", query_vector, top_k=2, filters=None) # Verify that the aggregation pipeline does not include the filter stage mock_collection.aggregate.assert_called_once() @@ -315,7 +315,7 @@ def test_list(mongo_vector_fixture): {"_id": "id2", "payload": {"key": "value2"}}, ] - results = mongo_vector.list(limit=2) + results = mongo_vector.list(top_k=2) mock_collection.find.assert_called_once_with({}) mock_cursor.limit.assert_called_once_with(2) @@ -334,7 +334,7 @@ def test_list_with_filters(mongo_vector_fixture): ] filters = {"user_id": "alice", "agent_id": "agent1", "run_id": "run1"} - results = mongo_vector.list(filters=filters, limit=2) + results = mongo_vector.list(filters=filters, top_k=2) # Verify that the find method was called with the correct query expected_query = { @@ -363,7 +363,7 @@ def test_list_with_single_filter(mongo_vector_fixture): ] filters = {"user_id": "alice"} - results = mongo_vector.list(filters=filters, limit=2) + results = mongo_vector.list(filters=filters, top_k=2) # Verify that the find method was called with the correct query expected_query = { @@ -387,7 +387,7 @@ def test_list_with_no_filters(mongo_vector_fixture): {"_id": "id1", "payload": {"key": "value1"}}, ] - results = mongo_vector.list(filters=None, limit=2) + results = mongo_vector.list(filters=None, top_k=2) # Verify that the find method was called with empty query mock_collection.find.assert_called_once_with({}) diff --git a/tests/vector_stores/test_neptune_analytics.py b/tests/vector_stores/test_neptune_analytics.py index a643cc26e..0fde89ae0 100644 --- a/tests/vector_stores/test_neptune_analytics.py +++ b/tests/vector_stores/test_neptune_analytics.py @@ -123,7 +123,7 @@ class TestNeptuneAnalyticsOperations: payloads=SAMPLE_PAYLOADS ) - result = na_instance.search(query="", vectors=VECTOR_1, limit=1) + result = na_instance.search(query="", vectors=VECTOR_1, top_k=1) assert len(result) == 1 assert "label" not in result[0].payload diff --git a/tests/vector_stores/test_opensearch.py b/tests/vector_stores/test_opensearch.py index 626fdf048..09046c367 100644 --- a/tests/vector_stores/test_opensearch.py +++ b/tests/vector_stores/test_opensearch.py @@ -229,7 +229,7 @@ class TestOpenSearchDB(unittest.TestCase): } self.client_mock.search.return_value = mock_response vectors = [[0.1] * 1536] - results = self.os_db.search(query="", vectors=vectors, limit=5) + results = self.os_db.search(query="", vectors=vectors, top_k=5) self.client_mock.search.assert_called_once() search_args = self.client_mock.search.call_args[1] self.assertEqual(search_args["index"], "test_collection") @@ -345,7 +345,7 @@ class TestOpenSearchDB(unittest.TestCase): def test_search_error_logs_with_exc_info(self, mock_logger): """Search error logging should include exc_info for full stack trace.""" self.client_mock.search.side_effect = Exception("Search failed") - results = self.os_db.search(query="", vectors=[[0.1] * 1536], limit=5) + results = self.os_db.search(query="", vectors=[[0.1] * 1536], top_k=5) self.assertEqual(results, []) mock_logger.error.assert_called_once() call_kwargs = mock_logger.error.call_args diff --git a/tests/vector_stores/test_pgvector.py b/tests/vector_stores/test_pgvector.py index 436c9708c..7479f3e23 100644 --- a/tests/vector_stores/test_pgvector.py +++ b/tests/vector_stores/test_pgvector.py @@ -432,7 +432,7 @@ class TestPGVector(unittest.TestCase): maxconn=4 ) - results = pgvector.search("test query", [0.1, 0.2, 0.3], limit=2) + results = pgvector.search("test query", [0.1, 0.2, 0.3], top_k=2) # Verify the _get_cursor context manager was called mock_get_cursor.assert_called() @@ -481,7 +481,7 @@ class TestPGVector(unittest.TestCase): maxconn=4 ) - results = pgvector.search("test query", [0.1, 0.2, 0.3], limit=2) + results = pgvector.search("test query", [0.1, 0.2, 0.3], top_k=2) # Verify the _get_cursor context manager was called mock_get_cursor.assert_called() @@ -1031,7 +1031,7 @@ class TestPGVector(unittest.TestCase): maxconn=4 ) - results = pgvector.list(limit=2) + results = pgvector.list(top_k=2) # Verify the _get_cursor context manager was called mock_get_cursor.assert_called() @@ -1079,7 +1079,7 @@ class TestPGVector(unittest.TestCase): maxconn=4 ) - results = pgvector.list(limit=2) + results = pgvector.list(top_k=2) # Verify the _get_cursor context manager was called mock_get_cursor.assert_called() @@ -1127,7 +1127,7 @@ class TestPGVector(unittest.TestCase): ) filters = {"user_id": "alice", "agent_id": "agent1", "run_id": "run1"} - results = pgvector.search("test query", [0.1, 0.2, 0.3], limit=2, filters=filters) + results = pgvector.search("test query", [0.1, 0.2, 0.3], top_k=2, filters=filters) # Verify the _get_cursor context manager was called mock_get_cursor.assert_called() @@ -1177,7 +1177,7 @@ class TestPGVector(unittest.TestCase): ) filters = {"user_id": "alice", "agent_id": "agent1", "run_id": "run1"} - results = pgvector.search("test query", [0.1, 0.2, 0.3], limit=2, filters=filters) + results = pgvector.search("test query", [0.1, 0.2, 0.3], top_k=2, filters=filters) # Verify the _get_cursor context manager was called mock_get_cursor.assert_called() @@ -1227,7 +1227,7 @@ class TestPGVector(unittest.TestCase): ) filters = {"user_id": "alice"} - results = pgvector.search("test query", [0.1, 0.2, 0.3], limit=2, filters=filters) + results = pgvector.search("test query", [0.1, 0.2, 0.3], top_k=2, filters=filters) # Verify the _get_cursor context manager was called mock_get_cursor.assert_called() @@ -1275,7 +1275,7 @@ class TestPGVector(unittest.TestCase): ) filters = {"user_id": "alice"} - results = pgvector.search("test query", [0.1, 0.2, 0.3], limit=2, filters=filters) + results = pgvector.search("test query", [0.1, 0.2, 0.3], top_k=2, filters=filters) # Verify the _get_cursor context manager was called mock_get_cursor.assert_called() @@ -1323,7 +1323,7 @@ class TestPGVector(unittest.TestCase): maxconn=4 ) - results = pgvector.search("test query", [0.1, 0.2, 0.3], limit=2, filters=None) + results = pgvector.search("test query", [0.1, 0.2, 0.3], top_k=2, filters=None) # Verify the _get_cursor context manager was called mock_get_cursor.assert_called() @@ -1372,7 +1372,7 @@ class TestPGVector(unittest.TestCase): maxconn=4 ) - results = pgvector.search("test query", [0.1, 0.2, 0.3], limit=2, filters=None) + results = pgvector.search("test query", [0.1, 0.2, 0.3], top_k=2, filters=None) # Verify the _get_cursor context manager was called mock_get_cursor.assert_called() @@ -1421,7 +1421,7 @@ class TestPGVector(unittest.TestCase): ) filters = {"user_id": "alice", "agent_id": "agent1"} - results = pgvector.list(filters=filters, limit=2) + results = pgvector.list(filters=filters, top_k=2) # Verify the _get_cursor context manager was called mock_get_cursor.assert_called() @@ -1470,7 +1470,7 @@ class TestPGVector(unittest.TestCase): ) filters = {"user_id": "alice", "agent_id": "agent1"} - results = pgvector.list(filters=filters, limit=2) + results = pgvector.list(filters=filters, top_k=2) # Verify the _get_cursor context manager was called mock_get_cursor.assert_called() @@ -1519,7 +1519,7 @@ class TestPGVector(unittest.TestCase): ) filters = {"user_id": "alice"} - results = pgvector.list(filters=filters, limit=2) + results = pgvector.list(filters=filters, top_k=2) # Verify the _get_cursor context manager was called mock_get_cursor.assert_called() @@ -1567,7 +1567,7 @@ class TestPGVector(unittest.TestCase): ) filters = {"user_id": "alice"} - results = pgvector.list(filters=filters, limit=2) + results = pgvector.list(filters=filters, top_k=2) # Verify the _get_cursor context manager was called mock_get_cursor.assert_called() @@ -1615,7 +1615,7 @@ class TestPGVector(unittest.TestCase): maxconn=4 ) - results = pgvector.list(filters=None, limit=2) + results = pgvector.list(filters=None, top_k=2) # Verify the _get_cursor context manager was called mock_get_cursor.assert_called() @@ -1663,7 +1663,7 @@ class TestPGVector(unittest.TestCase): maxconn=4 ) - results = pgvector.list(filters=None, limit=2) + results = pgvector.list(filters=None, top_k=2) # Verify the _get_cursor context manager was called mock_get_cursor.assert_called() diff --git a/tests/vector_stores/test_pinecone.py b/tests/vector_stores/test_pinecone.py index fb7043398..89194aea1 100644 --- a/tests/vector_stores/test_pinecone.py +++ b/tests/vector_stores/test_pinecone.py @@ -80,7 +80,7 @@ def test_insert_vectors(pinecone_db): def test_search_vectors(pinecone_db): pinecone_db.index.query.return_value.matches = [{"id": "id1", "score": 0.9, "metadata": {"name": "vector1"}}] - results = pinecone_db.search("test query", [0.1] * 128, limit=1) + results = pinecone_db.search("test query", [0.1] * 128, top_k=1) pinecone_db.index.query.assert_called_with( vector=[0.1] * 128, top_k=1, diff --git a/tests/vector_stores/test_qdrant.py b/tests/vector_stores/test_qdrant.py index 8cca2eb7f..5d10d2139 100644 --- a/tests/vector_stores/test_qdrant.py +++ b/tests/vector_stores/test_qdrant.py @@ -84,7 +84,7 @@ class TestQdrant(unittest.TestCase): mock_point = MagicMock(id=str(uuid.uuid4()), score=0.95, payload={"key": "value"}) self.client_mock.query_points.return_value = MagicMock(points=[mock_point]) - results = self.qdrant.search(query="", vectors=vectors, limit=1) + results = self.qdrant.search(query="", vectors=vectors, top_k=1) self.client_mock.query_points.assert_called_once_with( collection_name="test_collection", @@ -108,7 +108,7 @@ class TestQdrant(unittest.TestCase): self.client_mock.query_points.return_value = MagicMock(points=[mock_point]) filters = {"user_id": "alice", "agent_id": "agent1", "run_id": "run1"} - results = self.qdrant.search(query="", vectors=vectors, limit=1, filters=filters) + results = self.qdrant.search(query="", vectors=vectors, top_k=1, filters=filters) # Verify that _create_filter was called and query_filter was passed self.client_mock.query_points.assert_called_once() @@ -138,7 +138,7 @@ class TestQdrant(unittest.TestCase): self.client_mock.query_points.return_value = MagicMock(points=[mock_point]) filters = {"user_id": "alice"} - results = self.qdrant.search(query="", vectors=vectors, limit=1, filters=filters) + results = self.qdrant.search(query="", vectors=vectors, top_k=1, filters=filters) # Verify that a Filter object was created with single condition call_args = self.client_mock.query_points.call_args[1] @@ -155,7 +155,7 @@ class TestQdrant(unittest.TestCase): mock_point = MagicMock(id=str(uuid.uuid4()), score=0.95, payload={"key": "value"}) self.client_mock.query_points.return_value = MagicMock(points=[mock_point]) - results = self.qdrant.search(query="", vectors=vectors, limit=1, filters=None) + results = self.qdrant.search(query="", vectors=vectors, top_k=1, filters=None) call_args = self.client_mock.query_points.call_args[1] self.assertIsNone(call_args["query_filter"]) @@ -298,7 +298,7 @@ class TestQdrant(unittest.TestCase): self.client_mock.scroll.return_value = [mock_point] filters = {"user_id": "alice", "agent_id": "agent1", "run_id": "run1"} - results = self.qdrant.list(filters=filters, limit=10) + results = self.qdrant.list(filters=filters, top_k=10) # Verify that _create_filter was called and scroll_filter was passed self.client_mock.scroll.assert_called_once() @@ -327,7 +327,7 @@ class TestQdrant(unittest.TestCase): self.client_mock.scroll.return_value = [mock_point] filters = {"user_id": "alice"} - results = self.qdrant.list(filters=filters, limit=10) + results = self.qdrant.list(filters=filters, top_k=10) # Verify that a Filter object was created with single condition call_args = self.client_mock.scroll.call_args[1] @@ -344,7 +344,7 @@ class TestQdrant(unittest.TestCase): mock_point = MagicMock(id=str(uuid.uuid4()), score=0.95, payload={"key": "value"}) self.client_mock.scroll.return_value = [mock_point] - results = self.qdrant.list(filters=None, limit=10) + results = self.qdrant.list(filters=None, top_k=10) call_args = self.client_mock.scroll.call_args[1] self.assertIsNone(call_args["scroll_filter"]) diff --git a/tests/vector_stores/test_s3_vectors.py b/tests/vector_stores/test_s3_vectors.py index e8141e2f5..3ad69cb3a 100644 --- a/tests/vector_stores/test_s3_vectors.py +++ b/tests/vector_stores/test_s3_vectors.py @@ -1,7 +1,7 @@ -from mem0.configs.vector_stores.s3_vectors import S3VectorsConfig import pytest from botocore.exceptions import ClientError +from mem0.configs.vector_stores.s3_vectors import S3VectorsConfig from mem0.memory.main import Memory from mem0.vector_stores.s3_vectors import S3Vectors @@ -152,7 +152,7 @@ def test_search(mock_boto_client): embedding_model_dims=EMBEDDING_DIMS, ) query_vector = [0.1, 0.2] - results = store.search(query="test", vectors=query_vector, limit=1) + results = store.search(query="test", vectors=query_vector, top_k=1) mock_boto_client.query_vectors.assert_called_once() assert len(results) == 1 diff --git a/tests/vector_stores/test_supabase.py b/tests/vector_stores/test_supabase.py index e051ccf1b..b63fb92d6 100644 --- a/tests/vector_stores/test_supabase.py +++ b/tests/vector_stores/test_supabase.py @@ -67,7 +67,7 @@ def test_search_vectors(supabase_instance, mock_collection): vectors = [[0.1, 0.2, 0.3]] filters = {"category": "test"} - results = supabase_instance.search(query="", vectors=vectors, limit=2, filters=filters) + results = supabase_instance.search(query="", vectors=vectors, top_k=2, filters=filters) mock_collection.query.assert_called_once_with( data=vectors, limit=2, filters={"category": {"$eq": "test"}}, include_metadata=True, include_value=True @@ -118,7 +118,7 @@ def test_list_vectors(supabase_instance, mock_collection): mock_collection.query.return_value = mock_query_results mock_collection.fetch.return_value = mock_fetch_results - results = supabase_instance.list(limit=2, filters={"category": "test"}) + results = supabase_instance.list(top_k=2, filters={"category": "test"}) assert len(results[0]) == 2 assert results[0][0].id == "id1" diff --git a/tests/vector_stores/test_turbopuffer.py b/tests/vector_stores/test_turbopuffer.py index 7e677ed60..d1915b94e 100644 --- a/tests/vector_stores/test_turbopuffer.py +++ b/tests/vector_stores/test_turbopuffer.py @@ -289,7 +289,7 @@ class TestSearch: ] db.namespace.query.return_value = mock_response - results = db.search("test query", [0.1, 0.2, 0.3, 0.4], limit=2) + results = db.search("test query", [0.1, 0.2, 0.3, 0.4], top_k=2) db.namespace.query.assert_called_once_with( rank_by=("vector", "ANN", [0.1, 0.2, 0.3, 0.4]), @@ -306,7 +306,7 @@ class TestSearch: db.namespace.query.return_value = mock_response results = db.search( - "query", [0.1, 0.2, 0.3, 0.4], limit=5, filters={"user_id": "u1"} + "query", [0.1, 0.2, 0.3, 0.4], top_k=5, filters={"user_id": "u1"} ) call_kwargs = db.namespace.query.call_args[1] @@ -515,7 +515,7 @@ class TestList: mock_response.rows = [_make_row("id1", dist=0.1, data="hello")] db.namespace.query.return_value = mock_response - db.list(filters={"user_id": "u1"}, limit=50) + db.list(filters={"user_id": "u1"}, top_k=50) call_kwargs = db.namespace.query.call_args[1] assert call_kwargs["filters"] == ("user_id", "Eq", "u1") diff --git a/tests/vector_stores/test_upstash_vector.py b/tests/vector_stores/test_upstash_vector.py index ad3423415..b64357047 100644 --- a/tests/vector_stores/test_upstash_vector.py +++ b/tests/vector_stores/test_upstash_vector.py @@ -60,7 +60,7 @@ def test_search_vectors(upstash_instance, mock_index): results = upstash_instance.search( query="hello world", vectors=vectors, - limit=2, + top_k=2, filters={"age": 30, "name": "John"}, ) @@ -132,7 +132,7 @@ def test_list_vectors(upstash_instance): filters = {"age": 30, "name": "John"} print("filters", filters) - [results] = upstash_instance.list(filters=filters, limit=15) + [results] = upstash_instance.list(filters=filters, top_k=15) upstash_instance.client.info.return_value = { "dimension": 10, @@ -193,7 +193,7 @@ def test_search_vectors_with_embeddings(upstash_instance_with_embeddings, mock_i results = upstash_instance_with_embeddings.search( query="hello world", vectors=[], - limit=2, + top_k=2, filters={"age": 30, "name": "John"}, ) @@ -302,7 +302,7 @@ def test_search_vectors_empty_filters(upstash_instance): results = upstash_instance.search( query="hello world", vectors=vectors, - limit=1, + top_k=1, filters=None, ) diff --git a/tests/vector_stores/test_valkey.py b/tests/vector_stores/test_valkey.py index 65fa89fbb..7269830c8 100644 --- a/tests/vector_stores/test_valkey.py +++ b/tests/vector_stores/test_valkey.py @@ -59,7 +59,7 @@ def test_search_filter_syntax(valkey_db, mock_valkey_client): valkey_db.search( query="test query", vectors=np.random.rand(1536).tolist(), - limit=5, + top_k=5, filters={"user_id": "test_user"}, ) @@ -72,7 +72,7 @@ def test_search_filter_syntax(valkey_db, mock_valkey_client): valkey_db.search( query="test query", vectors=np.random.rand(1536).tolist(), - limit=5, + top_k=5, filters={"user_id": "test_user", "agent_id": "test_agent"}, ) @@ -104,7 +104,7 @@ def test_search_without_filters(valkey_db, mock_valkey_client): results = valkey_db.search( query="test query", vectors=np.random.rand(1536).tolist(), - limit=5, + top_k=5, ) # Check that the search was called with the correct syntax @@ -405,7 +405,7 @@ def test_list(valkey_db, mock_valkey_client): mock_ft.search.return_value = mock_results # Call list - results = valkey_db.list(filters={"user_id": "test_user"}, limit=10) + results = valkey_db.list(filters={"user_id": "test_user"}, top_k=10) # Check that search was called with the correct arguments mock_ft.search.assert_called_once() @@ -437,7 +437,7 @@ def test_search_error_handling(valkey_db, mock_valkey_client): valkey_db.search( query="test query", vectors=np.random.rand(1536).tolist(), - limit=5, + top_k=5, filters={"user_id": "test_user"}, ) @@ -776,7 +776,7 @@ def test_search_with_invalid_metadata(valkey_db, mock_valkey_client): mock_valkey_client.ft.return_value.search.return_value = mock_result # Should handle invalid JSON gracefully - results = valkey_db.search(query="test query", vectors=np.random.rand(1536).tolist(), limit=5) + results = valkey_db.search(query="test query", vectors=np.random.rand(1536).tolist(), top_k=5) assert len(results) == 1 @@ -790,7 +790,7 @@ def test_search_with_hnsw_ef_runtime(valkey_db, mock_valkey_client): mock_result.docs = [] mock_valkey_client.ft.return_value.search.return_value = mock_result - valkey_db.search(query="test query", vectors=np.random.rand(1536).tolist(), limit=5) + valkey_db.search(query="test query", vectors=np.random.rand(1536).tolist(), top_k=5) # Verify the search was called assert mock_valkey_client.ft.return_value.search.called diff --git a/tests/vector_stores/test_vertex_ai_vector_search.py b/tests/vector_stores/test_vertex_ai_vector_search.py index d0d1f4c9b..661d27ac0 100644 --- a/tests/vector_stores/test_vertex_ai_vector_search.py +++ b/tests/vector_stores/test_vertex_ai_vector_search.py @@ -100,7 +100,7 @@ def test_search_vectors(vector_store, mock_vertex_ai): mock_vertex_ai["endpoint"].find_neighbors.return_value = [[mock_neighbor]] - results = vector_store.search(query="", vectors=vectors, filters=filters, limit=1) + results = vector_store.search(query="", vectors=vectors, filters=filters, top_k=1) mock_vertex_ai["endpoint"].find_neighbors.assert_called_once_with( deployed_index_id=vector_store.deployment_index_id, diff --git a/tests/vector_stores/test_weaviate.py b/tests/vector_stores/test_weaviate.py index 680ec8455..76bcff043 100644 --- a/tests/vector_stores/test_weaviate.py +++ b/tests/vector_stores/test_weaviate.py @@ -1,10 +1,10 @@ import os -import uuid -import httpx import unittest +import uuid from unittest.mock import MagicMock, patch import dotenv +import httpx import weaviate from weaviate.exceptions import UnexpectedStatusCodeException @@ -137,7 +137,7 @@ class TestWeaviateDB(unittest.TestCase): mock_hybrid.return_value = mock_response vectors = [[0.1] * 1536] - results = self.weaviate_db.search(query="", vectors=vectors, limit=5) + results = self.weaviate_db.search(query="", vectors=vectors, top_k=5) mock_hybrid.assert_called_once() @@ -170,7 +170,7 @@ class TestWeaviateDB(unittest.TestCase): self.client_mock.collections.get.return_value.query.fetch_objects = mock_fetch mock_fetch.return_value = mock_response - results = self.weaviate_db.list(limit=10) + results = self.weaviate_db.list(top_k=10) mock_fetch.assert_called_once()