Merge branch 'main' into mem0-1.0.0
This commit is contained in:
@@ -13,7 +13,7 @@ install:
|
||||
install_all:
|
||||
pip install ruff==0.6.9 groq together boto3 litellm ollama chromadb weaviate weaviate-client sentence_transformers vertexai \
|
||||
google-generativeai elasticsearch opensearch-py vecs "pinecone<7.0.0" pinecone-text faiss-cpu langchain-community \
|
||||
upstash-vector azure-search-documents langchain-memgraph langchain-neo4j langchain-aws rank-bm25 pymochow pymongo psycopg kuzu databricks-sdk
|
||||
upstash-vector azure-search-documents langchain-memgraph langchain-neo4j langchain-aws rank-bm25 pymochow pymongo psycopg kuzu databricks-sdk valkey
|
||||
|
||||
# Format code with ruff
|
||||
format:
|
||||
|
||||
@@ -70,3 +70,39 @@ The v2 search API is powerful and flexible, allowing for more precise memory ret
|
||||
)
|
||||
```
|
||||
</CodeGroup>
|
||||
|
||||
<CodeGroup>
|
||||
```python Categories Filter Examples
|
||||
# Example 1: Using 'contains' for partial matching
|
||||
finance_memories = m.search(
|
||||
query="What are my financial goals?",
|
||||
version="v2",
|
||||
filters={
|
||||
"AND": [
|
||||
{ "user_id": "alice" },
|
||||
{
|
||||
"categories": {
|
||||
"contains": "finance"
|
||||
}
|
||||
}
|
||||
]
|
||||
},
|
||||
)
|
||||
|
||||
# Example 2: Using 'in' for exact matching
|
||||
personal_memories = m.search(
|
||||
query="What personal information do you have?",
|
||||
version="v2",
|
||||
filters={
|
||||
"AND": [
|
||||
{ "user_id": "alice" },
|
||||
{
|
||||
"categories": {
|
||||
"in": ["personal_information"]
|
||||
}
|
||||
}
|
||||
]
|
||||
},
|
||||
)
|
||||
```
|
||||
</CodeGroup>
|
||||
|
||||
@@ -7,6 +7,49 @@ mode: "wide"
|
||||
<Tabs>
|
||||
<Tab title="Python">
|
||||
|
||||
<Update label="2025-09-03" description="v0.1.117">
|
||||
|
||||
**New Features & Updates:**
|
||||
- **OpenMemory:**
|
||||
- Added memory export / import feature
|
||||
- Added vector store integrations: Weaviate, FAISS, PGVector, Chroma, Redis, Elasticsearch, Milvus
|
||||
- Added `export_openmemory.sh` migration script
|
||||
- **Vector Stores:**
|
||||
- Added Amazon S3 Vectors support
|
||||
- Added Databricks Mosaic AI vector store support
|
||||
- Added support for OpenAI Store
|
||||
- **Graph Memory:** Added support for graph memory using Kuzu
|
||||
- **Azure:** Added Azure Identity for Azure OpenAI and Azure AI Search authentication
|
||||
- **Elasticsearch:** Added headers configuration support
|
||||
|
||||
**Improvements:**
|
||||
- Added custom connection client to enable connecting to local containers for Weaviate
|
||||
- Updated configuration AWS Bedrock
|
||||
- Fixed dependency issues and tests; updated docstrings
|
||||
- **Documentation:**
|
||||
- Fixed Graph Docs page missing in sidebar
|
||||
- Updated integration documentation
|
||||
- Added version param in Search V2 API documentation
|
||||
- Updated Databricks documentation and refactored docs
|
||||
- Updated favicon logo
|
||||
- Fixed typos and Typescript docs
|
||||
|
||||
**Bug Fixes:**
|
||||
- Baidu: Added missing provider for Baidu vector DB
|
||||
- MongoDB: Replaced `query_vector` args in search method
|
||||
- Fixed new memory mistaken for current
|
||||
- AsyncMemory._add_to_vector_store: handled edge case when no facts found
|
||||
- Fixed missing commas in Kuzu graph INSERT queries
|
||||
- Fixed inconsistent created and updated properties for Graph
|
||||
- Fixed missing `app_id` on client for Neptune Analytics
|
||||
- Correctly pick AWS region from environment variable
|
||||
- Fixed Ollama model existence check
|
||||
|
||||
**Refactoring:**
|
||||
- **PGVector:** Use internal connection pools and context managers
|
||||
|
||||
</Update>
|
||||
|
||||
<Update label="2025-08-14" description="v0.1.116">
|
||||
|
||||
**New Features & Updates:**
|
||||
@@ -568,6 +611,11 @@ mode: "wide"
|
||||
|
||||
<Tab title="TypeScript">
|
||||
|
||||
<Update label="2025-09-04" description="v2.1.38">
|
||||
**New Features:**
|
||||
- **Client:** Added `metadata` param to `update` method.
|
||||
</Update>
|
||||
|
||||
<Update label="2025-08-04" description="v2.1.37">
|
||||
**New Features:**
|
||||
- **OSS:** Added `RedisCloud` search module check
|
||||
@@ -1039,6 +1087,11 @@ mode: "wide"
|
||||
|
||||
<Tab title="Vercel AI SDK">
|
||||
|
||||
<Update label="2025-09-03" description="v2.0.2">
|
||||
**Bug Fix:**
|
||||
- **Vercel AI SDK:** Fixed streaming response in the AI SDK.
|
||||
</Update>
|
||||
|
||||
<Update label="2025-08-05" description="v2.0.1">
|
||||
**New Features:**
|
||||
- **Vercel AI SDK:** Added a new param `host` to the config.
|
||||
|
||||
@@ -8,7 +8,7 @@ iconType: "solid"
|
||||
|
||||
The `config` is defined as an object with two main keys:
|
||||
- `vector_store`: Specifies the vector database provider and its configuration
|
||||
- `provider`: The name of the vector database (e.g., "chroma", "pgvector", "qdrant", "milvus", "upstash_vector", "azure_ai_search", "vertex_ai_vector_search")
|
||||
- `provider`: The name of the vector database (e.g., "chroma", "pgvector", "qdrant", "milvus", "upstash_vector", "azure_ai_search", "vertex_ai_vector_search", "valkey")
|
||||
- `config`: A nested dictionary containing provider-specific settings
|
||||
|
||||
|
||||
|
||||
@@ -1,7 +1,9 @@
|
||||
[Chroma](https://www.trychroma.com/) is an AI-native open-source vector database that simplifies building LLM apps by providing tools for storing, embedding, and searching embeddings with a focus on simplicity and speed.
|
||||
[Chroma](https://www.trychroma.com/) is an AI-native open-source vector database that simplifies building LLM apps by providing tools for storing, embedding, and searching embeddings with a focus on simplicity and speed. It supports both local deployment and cloud hosting through ChromaDB Cloud.
|
||||
|
||||
### Usage
|
||||
|
||||
#### Local Installation
|
||||
|
||||
```python
|
||||
import os
|
||||
from mem0 import Memory
|
||||
@@ -14,6 +16,9 @@ config = {
|
||||
"config": {
|
||||
"collection_name": "test",
|
||||
"path": "db",
|
||||
# Optional: ChromaDB Cloud configuration
|
||||
# "api_key": "your-chroma-cloud-api-key",
|
||||
# "tenant": "your-chroma-cloud-tenant-id",
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -39,3 +44,5 @@ Here are the parameters available for configuring Chroma:
|
||||
| `path` | Path for the Chroma database | `db` |
|
||||
| `host` | The host where the Chroma server is running | `None` |
|
||||
| `port` | The port where the Chroma server is running | `None` |
|
||||
| `api_key` | ChromaDB Cloud API key (for cloud usage) | `None` |
|
||||
| `tenant` | ChromaDB Cloud tenant ID (for cloud usage) | `None` |
|
||||
@@ -0,0 +1,49 @@
|
||||
# Valkey Vector Store
|
||||
|
||||
[Valkey](https://valkey.io/) is an open source (BSD) high-performance key/value datastore that supports a variety of workloads and rich datastructures including vector search.
|
||||
|
||||
## Installation
|
||||
|
||||
```bash
|
||||
pip install mem0ai[vector_stores]
|
||||
```
|
||||
|
||||
## Usage
|
||||
|
||||
```python
|
||||
config = {
|
||||
"vector_store": {
|
||||
"provider": "valkey",
|
||||
"config": {
|
||||
"collection_name": "test",
|
||||
"valkey_url": "valkey://localhost:6379",
|
||||
"embedding_model_dims": 1536,
|
||||
"index_type": "flat"
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
m = Memory.from_config(config)
|
||||
messages = [
|
||||
{"role": "user", "content": "I'm planning to watch a movie tonight. Any recommendations?"},
|
||||
{"role": "assistant", "content": "How about a thriller movies? They can be quite engaging."},
|
||||
{"role": "user", "content": "I'm not a big fan of thriller movies but I love sci-fi movies."},
|
||||
{"role": "assistant", "content": "Got it! I'll avoid thriller recommendations and suggest sci-fi movies in the future."}
|
||||
]
|
||||
m.add(messages, user_id="alice", metadata={"category": "movies"})
|
||||
```
|
||||
|
||||
## Parameters
|
||||
|
||||
Let's see the available parameters for the `valkey` config:
|
||||
|
||||
| Parameter | Description | Default Value |
|
||||
| --- | --- | --- |
|
||||
| `collection_name` | The name of the collection to store the vectors | `mem0` |
|
||||
| `valkey_url` | Connection URL for the Valkey server | `valkey://localhost:6379` |
|
||||
| `embedding_model_dims` | Dimensions of the embedding model | `1536` |
|
||||
| `index_type` | Vector index algorithm (`hnsw` or `flat`) | `hnsw` |
|
||||
| `hnsw_m` | Number of bi-directional links for HNSW | `16` |
|
||||
| `hnsw_ef_construction` | Size of dynamic candidate list for HNSW | `200` |
|
||||
| `hnsw_ef_runtime` | Size of dynamic candidate list for search | `10` |
|
||||
| `distance_metric` | Distance metric for vector similarity | `cosine` |
|
||||
@@ -11,7 +11,7 @@ Mem0 includes built-in support for various popular databases. Memory can utilize
|
||||
See the list of supported vector databases below.
|
||||
|
||||
<Note>
|
||||
The following vector databases are supported in the Python implementation. The TypeScript implementation currently only supports Qdrant, Redis,Vectorize and in-memory vector database.
|
||||
The following vector databases are supported in the Python implementation. The TypeScript implementation currently only supports Qdrant, Redis, Valkey, Vectorize and in-memory vector database.
|
||||
</Note>
|
||||
|
||||
<CardGroup cols={3}>
|
||||
@@ -24,6 +24,7 @@ See the list of supported vector databases below.
|
||||
<Card title="MongoDB" href="/components/vectordbs/dbs/mongodb"></Card>
|
||||
<Card title="Azure" href="/components/vectordbs/dbs/azure"></Card>
|
||||
<Card title="Redis" href="/components/vectordbs/dbs/redis"></Card>
|
||||
<Card title="Valkey" href="/components/vectordbs/dbs/valkey"></Card>
|
||||
<Card title="Elasticsearch" href="/components/vectordbs/dbs/elasticsearch"></Card>
|
||||
<Card title="OpenSearch" href="/components/vectordbs/dbs/opensearch"></Card>
|
||||
<Card title="Supabase" href="/components/vectordbs/dbs/supabase"></Card>
|
||||
|
||||
+5
-3
@@ -1,8 +1,8 @@
|
||||
{
|
||||
"$schema": "https://mintlify.com/docs.json",
|
||||
"theme": "maple",
|
||||
"name": "Mem0",
|
||||
"description": "Mem0 is a self-improving memory layer for LLM applications, enabling personalized AI experiences that save costs and delight users.",
|
||||
"theme": "maple",
|
||||
"colors": {
|
||||
"primary": "#6c60f0",
|
||||
"light": "#E6FFA2",
|
||||
@@ -46,7 +46,7 @@
|
||||
},
|
||||
{
|
||||
"group": "Platform",
|
||||
"icon": "cogs",
|
||||
"icon": "globe",
|
||||
"pages": [
|
||||
"platform/overview",
|
||||
"platform/quickstart",
|
||||
@@ -58,6 +58,7 @@
|
||||
"platform/features/platform-overview",
|
||||
"platform/features/contextual-add",
|
||||
"platform/features/async-client",
|
||||
"platform/features/graph-memory",
|
||||
"platform/features/advanced-retrieval",
|
||||
"platform/features/criteria-retrieval",
|
||||
"platform/features/multimodal-support",
|
||||
@@ -151,6 +152,7 @@
|
||||
"components/vectordbs/dbs/mongodb",
|
||||
"components/vectordbs/dbs/azure",
|
||||
"components/vectordbs/dbs/redis",
|
||||
"components/vectordbs/dbs/valkey",
|
||||
"components/vectordbs/dbs/elasticsearch",
|
||||
"components/vectordbs/dbs/opensearch",
|
||||
"components/vectordbs/dbs/supabase",
|
||||
@@ -380,7 +382,7 @@
|
||||
"background": {
|
||||
"color": {
|
||||
"light": "#fff",
|
||||
"dark": "#0f1117"
|
||||
"dark": "#09090b"
|
||||
}
|
||||
},
|
||||
"navbar": {
|
||||
|
||||
@@ -24,8 +24,7 @@ Before you begin, make sure you have:
|
||||
|
||||
Installed Google ADK and Mem0 SDK:
|
||||
```bash
|
||||
pip install google-adk
|
||||
pip install mem0ai
|
||||
pip install google-adk mem0ai python-dotenv
|
||||
```
|
||||
|
||||
## Code Breakdown
|
||||
@@ -35,21 +34,25 @@ Let's get started and understand the different components required in building a
|
||||
```python
|
||||
# Import dependencies
|
||||
import os
|
||||
import asyncio
|
||||
from google.adk.agents import Agent
|
||||
from google.adk.sessions import InMemorySessionService
|
||||
from google.adk.runners import Runner
|
||||
from google.adk.sessions import InMemorySessionService
|
||||
from google.genai import types
|
||||
from mem0 import MemoryClient
|
||||
from dotenv import load_dotenv
|
||||
|
||||
# Set up API keys (replace with your actual keys)
|
||||
os.environ["GOOGLE_API_KEY"] = "your-google-api-key"
|
||||
os.environ["MEM0_API_KEY"] = "your-mem0-api-key"
|
||||
load_dotenv()
|
||||
|
||||
# Set up environment variables
|
||||
# os.environ["GOOGLE_API_KEY"] = "your-google-api-key"
|
||||
# os.environ["MEM0_API_KEY"] = "your-mem0-api-key"
|
||||
|
||||
# Define a global user ID for simplicity
|
||||
USER_ID = "Alex"
|
||||
|
||||
# Initialize Mem0 client
|
||||
mem0_client = MemoryClient()
|
||||
mem0 = MemoryClient()
|
||||
```
|
||||
|
||||
## Define Memory Tools
|
||||
|
||||
@@ -17,7 +17,7 @@ Before setting up Mem0 with AgentOps, ensure you have:
|
||||
|
||||
1. Installed the required packages:
|
||||
```bash
|
||||
pip install mem0ai agentops
|
||||
pip install mem0ai agentops python-dotenv
|
||||
```
|
||||
|
||||
2. Valid API keys:
|
||||
@@ -37,7 +37,9 @@ import asyncio
|
||||
import logging
|
||||
from dotenv import load_dotenv
|
||||
import agentops
|
||||
import openai
|
||||
|
||||
load_dotenv()
|
||||
#Set up environment variables for API keys
|
||||
os.environ["AGENTOPS_API_KEY"] = os.getenv("AGENTOPS_API_KEY")
|
||||
os.environ["OPENAI_API_KEY"] = os.getenv("OPENAI_API_KEY")
|
||||
|
||||
@@ -18,7 +18,7 @@ Before setting up Mem0 with Agno, ensure you have:
|
||||
|
||||
1. Installed the required packages:
|
||||
```bash
|
||||
pip install agno mem0ai
|
||||
pip install agno mem0ai python-dotenv
|
||||
```
|
||||
|
||||
2. Valid API keys:
|
||||
@@ -81,7 +81,7 @@ agent = Agent(
|
||||
|
||||
def chat_user(
|
||||
user_input: Optional[str] = None,
|
||||
user_id: str = "user_123",
|
||||
user_id: str = "alex",
|
||||
image_path: Optional[str] = None
|
||||
) -> str:
|
||||
"""
|
||||
@@ -120,13 +120,13 @@ def chat_user(
|
||||
})
|
||||
|
||||
# Store messages in memory
|
||||
client.add(messages, user_id=user_id)
|
||||
client.add(messages, user_id=user_id, output_format='v1.1')
|
||||
print("✅ Image and text stored in memory.")
|
||||
|
||||
if user_input:
|
||||
# Search for relevant memories
|
||||
memories = client.search(user_input, user_id=user_id)
|
||||
memory_context = "\n".join(f"- {m['memory']}" for m in memories.get('results', []))
|
||||
memories = client.search(user_input, user_id=user_id, output_format='v1.1')
|
||||
memory_context = "\n".join(f"- {m['memory']}" for m in memories['results'])
|
||||
|
||||
# Construct the prompt
|
||||
prompt = f"""
|
||||
@@ -150,7 +150,8 @@ User question:
|
||||
response = agent.run(prompt)
|
||||
|
||||
# Store the interaction in memory
|
||||
client.add(f"User: {user_input}\nAssistant: {response.content}", user_id=user_id)
|
||||
interaction_message = [{"role": "user", "content": f"User: {user_input}\nAssistant: {response.content}"}]
|
||||
client.add(interaction_message, user_id=user_id, output_format='v1.1')
|
||||
return response.content
|
||||
|
||||
return "No user input or image provided."
|
||||
@@ -159,9 +160,9 @@ User question:
|
||||
# Example Usage
|
||||
if __name__ == "__main__":
|
||||
response = chat_user(
|
||||
"This is the picture of what I brought with me in the trip to Bahamas",
|
||||
"I like to travel and my favorite destination is London",
|
||||
image_path="travel_items.jpeg",
|
||||
user_id="user_123"
|
||||
user_id="alex"
|
||||
)
|
||||
print(response)
|
||||
```
|
||||
|
||||
@@ -1,3 +1,7 @@
|
||||
---
|
||||
title: AutoGen
|
||||
---
|
||||
|
||||
Build conversational AI agents with memory capabilities. This integration combines AutoGen for creating AI agents with Mem0 for memory management, enabling context-aware and personalized interactions.
|
||||
|
||||
## Overview
|
||||
@@ -10,7 +14,7 @@ In this guide, we'll explore an example of creating a conversational AI system w
|
||||
Install necessary libraries:
|
||||
|
||||
```bash
|
||||
pip install pyautogen mem0ai openai
|
||||
pip install autogen mem0ai openai python-dotenv
|
||||
```
|
||||
|
||||
First, we'll import the necessary libraries and set up our configurations.
|
||||
@@ -22,15 +26,18 @@ import os
|
||||
from autogen import ConversableAgent
|
||||
from mem0 import MemoryClient
|
||||
from openai import OpenAI
|
||||
from dotenv import load_dotenv
|
||||
|
||||
load_dotenv()
|
||||
|
||||
# Configuration
|
||||
OPENAI_API_KEY = 'sk-xxx' # Replace with your actual OpenAI API key
|
||||
MEM0_API_KEY = 'your-mem0-key' # Replace with your actual Mem0 API key from https://app.mem0.ai
|
||||
USER_ID = "customer_service_bot"
|
||||
# OPENAI_API_KEY = 'sk-xxx' # Replace with your actual OpenAI API key
|
||||
# MEM0_API_KEY = 'your-mem0-key' # Replace with your actual Mem0 API key from https://app.mem0.ai
|
||||
USER_ID = "alice"
|
||||
|
||||
# Set up OpenAI API key
|
||||
os.environ['OPENAI_API_KEY'] = OPENAI_API_KEY
|
||||
os.environ['MEM0_API_KEY'] = MEM0_API_KEY
|
||||
OPENAI_API_KEY = os.environ.get('OPENAI_API_KEY')
|
||||
# os.environ['MEM0_API_KEY'] = MEM0_API_KEY
|
||||
|
||||
# Initialize Mem0 and AutoGen agents
|
||||
memory_client = MemoryClient()
|
||||
@@ -55,7 +62,7 @@ conversation = [
|
||||
{"role": "assistant", "content": "Thank you for the information. Let's troubleshoot this issue..."}
|
||||
]
|
||||
|
||||
memory_client.add(messages=conversation, user_id=USER_ID)
|
||||
memory_client.add(messages=conversation, user_id=USER_ID, output_format="v1.1")
|
||||
print("Conversation added to memory.")
|
||||
```
|
||||
|
||||
@@ -65,7 +72,7 @@ Create a function to get context-aware responses based on user's question and pr
|
||||
|
||||
```python
|
||||
def get_context_aware_response(question):
|
||||
relevant_memories = memory_client.search(question, user_id=USER_ID)
|
||||
relevant_memories = memory_client.search(question, user_id=USER_ID, output_format='v1.1')
|
||||
context = "\n".join([m["memory"] for m in relevant_memories.get('results', [])])
|
||||
|
||||
prompt = f"""Answer the user question considering the previous interactions:
|
||||
@@ -97,7 +104,7 @@ manager = ConversableAgent(
|
||||
)
|
||||
|
||||
def escalate_to_manager(question):
|
||||
relevant_memories = memory_client.search(question, user_id=USER_ID)
|
||||
relevant_memories = memory_client.search(question, user_id=USER_ID, output_format='v1.1')
|
||||
context = "\n".join([m["memory"] for m in relevant_memories.get('results', [])])
|
||||
|
||||
prompt = f"""
|
||||
|
||||
@@ -16,7 +16,7 @@ In this guide, we'll build a voice agent that:
|
||||
Install necessary libraries:
|
||||
|
||||
```bash
|
||||
pip install elevenlabs mem0 python-dotenv
|
||||
pip install elevenlabs mem0ai python-dotenv
|
||||
```
|
||||
|
||||
Configure your environment variables:
|
||||
|
||||
@@ -17,7 +17,7 @@ Before setting up Mem0 with Google ADK, ensure you have:
|
||||
|
||||
1. Installed the required packages:
|
||||
```bash
|
||||
pip install google-adk mem0ai
|
||||
pip install google-adk mem0ai python-dotenv
|
||||
```
|
||||
|
||||
2. Valid API keys:
|
||||
@@ -30,15 +30,19 @@ The following example demonstrates how to create a Google ADK agent with Mem0 me
|
||||
|
||||
```python
|
||||
import os
|
||||
import asyncio
|
||||
from google.adk.agents import Agent
|
||||
from google.adk.runners import Runner
|
||||
from google.adk.sessions import InMemorySessionService
|
||||
from google.genai import types
|
||||
from mem0 import MemoryClient
|
||||
from dotenv import load_dotenv
|
||||
|
||||
load_dotenv()
|
||||
|
||||
# Set up environment variables
|
||||
os.environ["GOOGLE_API_KEY"] = "your-google-api-key"
|
||||
os.environ["MEM0_API_KEY"] = "your-mem0-api-key"
|
||||
# os.environ["GOOGLE_API_KEY"] = "your-google-api-key"
|
||||
# os.environ["MEM0_API_KEY"] = "your-mem0-api-key"
|
||||
|
||||
# Initialize Mem0 client
|
||||
mem0 = MemoryClient()
|
||||
@@ -46,17 +50,18 @@ mem0 = MemoryClient()
|
||||
# Define memory function tools
|
||||
def search_memory(query: str, user_id: str) -> dict:
|
||||
"""Search through past conversations and memories"""
|
||||
memories = mem0.search(query, user_id=user_id)
|
||||
memories = mem0.search(query, user_id=user_id, output_format='v1.1')
|
||||
if memories.get('results', []):
|
||||
memory_context = "\n".join([f"- {mem['memory']}" for mem in memories.get('results', [])])
|
||||
memory_list = memories['results']
|
||||
memory_context = "\n".join([f"- {mem['memory']}" for mem in memory_list])
|
||||
return {"status": "success", "memories": memory_context}
|
||||
return {"status": "no_memories", "message": "No relevant memories found"}
|
||||
|
||||
def save_memory(content: str, user_id: str) -> dict:
|
||||
"""Save important information to memory"""
|
||||
try:
|
||||
mem0.add([{"role": "user", "content": content}], user_id=user_id)
|
||||
return {"status": "success", "message": "Information saved to memory"}
|
||||
result = mem0.add([{"role": "user", "content": content}], user_id=user_id, output_format='v1.1')
|
||||
return {"status": "success", "message": "Information saved to memory", "result": result}
|
||||
except Exception as e:
|
||||
return {"status": "error", "message": f"Failed to save memory: {str(e)}"}
|
||||
|
||||
@@ -72,7 +77,7 @@ personal_assistant = Agent(
|
||||
tools=[search_memory, save_memory]
|
||||
)
|
||||
|
||||
def chat_with_agent(user_input: str, user_id: str) -> str:
|
||||
async def chat_with_agent(user_input: str, user_id: str) -> str:
|
||||
"""
|
||||
Handle user input with automatic memory integration.
|
||||
|
||||
@@ -85,7 +90,7 @@ def chat_with_agent(user_input: str, user_id: str) -> str:
|
||||
"""
|
||||
# Set up session and runner
|
||||
session_service = InMemorySessionService()
|
||||
session = session_service.create_session(
|
||||
session = await session_service.create_session(
|
||||
app_name="memory_assistant",
|
||||
user_id=user_id,
|
||||
session_id=f"session_{user_id}"
|
||||
@@ -107,10 +112,10 @@ def chat_with_agent(user_input: str, user_id: str) -> str:
|
||||
|
||||
# Example usage
|
||||
if __name__ == "__main__":
|
||||
response = chat_with_agent(
|
||||
response = asyncio.run(chat_with_agent(
|
||||
"I love Italian food and I'm planning a trip to Rome next month",
|
||||
user_id="alice"
|
||||
)
|
||||
))
|
||||
print(response)
|
||||
```
|
||||
|
||||
|
||||
@@ -16,7 +16,7 @@ In this guide, we'll create a Travel Agent AI that:
|
||||
Install necessary libraries:
|
||||
|
||||
```bash
|
||||
pip install langchain langchain_openai mem0ai
|
||||
pip install langchain langchain_openai mem0ai python-dotenv
|
||||
```
|
||||
|
||||
Import required modules and set up configurations:
|
||||
@@ -30,10 +30,13 @@ from langchain_openai import ChatOpenAI
|
||||
from langchain_core.messages import SystemMessage, HumanMessage, AIMessage
|
||||
from langchain_core.prompts import ChatPromptTemplate, MessagesPlaceholder
|
||||
from mem0 import MemoryClient
|
||||
from dotenv import load_dotenv
|
||||
|
||||
load_dotenv()
|
||||
|
||||
# Configuration
|
||||
os.environ["OPENAI_API_KEY"] = "your-openai-api-key"
|
||||
os.environ["MEM0_API_KEY"] = "your-mem0-api-key"
|
||||
# os.environ["OPENAI_API_KEY"] = "your-openai-api-key"
|
||||
# os.environ["MEM0_API_KEY"] = "your-mem0-api-key"
|
||||
|
||||
# Initialize LangChain and Mem0
|
||||
llm = ChatOpenAI(model="gpt-4o-mini")
|
||||
@@ -61,8 +64,11 @@ Create functions to handle context retrieval, response generation, and addition
|
||||
```python
|
||||
def retrieve_context(query: str, user_id: str) -> List[Dict]:
|
||||
"""Retrieve relevant context from Mem0"""
|
||||
memories = mem0.search(query, user_id=user_id)
|
||||
serialized_memories = ' '.join([mem["memory"] for mem in memories.get('results', [])])
|
||||
try:
|
||||
memories = mem0.search(query, user_id=user_id, output_format='v1.1')
|
||||
memory_list = memories['results']
|
||||
|
||||
serialized_memories = ' '.join([mem["memory"] for mem in memory_list])
|
||||
context = [
|
||||
{
|
||||
"role": "system",
|
||||
@@ -74,6 +80,10 @@ def retrieve_context(query: str, user_id: str) -> List[Dict]:
|
||||
}
|
||||
]
|
||||
return context
|
||||
except Exception as e:
|
||||
print(f"Error retrieving memories: {e}")
|
||||
# Return empty context if there's an error
|
||||
return [{"role": "user", "content": query}]
|
||||
|
||||
def generate_response(input: str, context: List[Dict]) -> str:
|
||||
"""Generate a response using the language model"""
|
||||
@@ -86,6 +96,7 @@ def generate_response(input: str, context: List[Dict]) -> str:
|
||||
|
||||
def save_interaction(user_id: str, user_input: str, assistant_response: str):
|
||||
"""Save the interaction to Mem0"""
|
||||
try:
|
||||
interaction = [
|
||||
{
|
||||
"role": "user",
|
||||
@@ -96,7 +107,10 @@ def save_interaction(user_id: str, user_input: str, assistant_response: str):
|
||||
"content": assistant_response
|
||||
}
|
||||
]
|
||||
mem0.add(interaction, user_id=user_id)
|
||||
result = mem0.add(interaction, user_id=user_id, output_format='v1.1')
|
||||
print(f"Memory saved successfully: {len(result.get('results', []))} memories added")
|
||||
except Exception as e:
|
||||
print(f"Error saving interaction: {e}")
|
||||
```
|
||||
|
||||
## Create Chat Turn Function
|
||||
@@ -124,7 +138,7 @@ Set up the main program loop for user interaction:
|
||||
```python
|
||||
if __name__ == "__main__":
|
||||
print("Welcome to your personal Travel Agent Planner! How can I assist you with your travel plans today?")
|
||||
user_id = "john"
|
||||
user_id = "alice"
|
||||
|
||||
while True:
|
||||
user_input = input("You: ")
|
||||
|
||||
@@ -16,7 +16,7 @@ In this guide, we'll create a Customer Support AI Agent that:
|
||||
Install necessary libraries:
|
||||
|
||||
```bash
|
||||
pip install langgraph langchain-openai mem0ai
|
||||
pip install langgraph langchain-openai mem0ai python-dotenv
|
||||
```
|
||||
|
||||
|
||||
@@ -31,14 +31,17 @@ from langgraph.graph.message import add_messages
|
||||
from langchain_openai import ChatOpenAI
|
||||
from mem0 import MemoryClient
|
||||
from langchain_core.messages import SystemMessage, HumanMessage, AIMessage
|
||||
from dotenv import load_dotenv
|
||||
|
||||
load_dotenv()
|
||||
|
||||
# Configuration
|
||||
OPENAI_API_KEY = 'sk-xxx' # Replace with your actual OpenAI API key
|
||||
MEM0_API_KEY = 'your-mem0-key' # Replace with your actual Mem0 API key
|
||||
# OPENAI_API_KEY = 'sk-xxx' # Replace with your actual OpenAI API key
|
||||
# MEM0_API_KEY = 'your-mem0-key' # Replace with your actual Mem0 API key
|
||||
|
||||
# Initialize LangChain and Mem0
|
||||
llm = ChatOpenAI(model="gpt-4", api_key=OPENAI_API_KEY)
|
||||
mem0 = MemoryClient(api_key=MEM0_API_KEY)
|
||||
llm = ChatOpenAI(model="gpt-4")
|
||||
mem0 = MemoryClient()
|
||||
```
|
||||
|
||||
## Define State and Graph
|
||||
@@ -62,11 +65,15 @@ def chatbot(state: State):
|
||||
messages = state["messages"]
|
||||
user_id = state["mem0_user_id"]
|
||||
|
||||
try:
|
||||
# Retrieve relevant memories
|
||||
memories = mem0.search(messages[-1].content, user_id=user_id)
|
||||
memories = mem0.search(messages[-1].content, user_id=user_id, output_format='v1.1')
|
||||
|
||||
# Handle dict response format
|
||||
memory_list = memories['results']
|
||||
|
||||
context = "Relevant information from previous conversations:\n"
|
||||
for memory in memories.get('results', []):
|
||||
for memory in memory_list:
|
||||
context += f"- {memory['memory']}\n"
|
||||
|
||||
system_message = SystemMessage(content=f"""You are a helpful customer support assistant. Use the provided context to personalize your responses and remember user preferences and past interactions.
|
||||
@@ -76,7 +83,28 @@ def chatbot(state: State):
|
||||
response = llm.invoke(full_messages)
|
||||
|
||||
# Store the interaction in Mem0
|
||||
mem0.add(f"User: {messages[-1].content}\nAssistant: {response.content}", user_id=user_id)
|
||||
try:
|
||||
interaction = [
|
||||
{
|
||||
"role": "user",
|
||||
"content": messages[-1].content
|
||||
},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": response.content
|
||||
}
|
||||
]
|
||||
result = mem0.add(interaction, user_id=user_id, output_format='v1.1')
|
||||
print(f"Memory saved: {len(result.get('results', []))} memories added")
|
||||
except Exception as e:
|
||||
print(f"Error saving memory: {e}")
|
||||
|
||||
return {"messages": [response]}
|
||||
|
||||
except Exception as e:
|
||||
print(f"Error in chatbot: {e}")
|
||||
# Fallback response without memory context
|
||||
response = llm.invoke(messages)
|
||||
return {"messages": [response]}
|
||||
```
|
||||
|
||||
@@ -115,7 +143,7 @@ Set up the main program loop for user interaction:
|
||||
```python
|
||||
if __name__ == "__main__":
|
||||
print("Welcome to Customer Support! How can I assist you today?")
|
||||
mem0_user_id = "customer_123" # You can generate or retrieve this based on your user management system
|
||||
mem0_user_id = "alice" # You can generate or retrieve this based on your user management system
|
||||
while True:
|
||||
user_input = input("You: ")
|
||||
if user_input.lower() in ['quit', 'exit', 'bye']:
|
||||
|
||||
@@ -13,7 +13,7 @@ LlamaIndex supports Mem0 as a [memory store](https://llamahub.ai/l/memory/llama-
|
||||
To install the required package, run:
|
||||
|
||||
```bash
|
||||
pip install llama-index-core llama-index-memory-mem0
|
||||
pip install llama-index-core llama-index-memory-mem0 python-dotenv
|
||||
```
|
||||
|
||||
### Setup with Mem0 Platform
|
||||
@@ -25,18 +25,23 @@ Set your Mem0 Platform API key as an environment variable. You can replace `<you
|
||||
</Note>
|
||||
|
||||
```python
|
||||
os.environ["MEM0_API_KEY"] = "<your-mem0-api-key>"
|
||||
from dotenv import load_dotenv
|
||||
import os
|
||||
|
||||
load_dotenv()
|
||||
|
||||
# os.environ["MEM0_API_KEY"] = "<your-mem0-api-key>"
|
||||
```
|
||||
|
||||
Import the necessary modules and create a Mem0Memory instance:
|
||||
```python
|
||||
from llama_index.memory.mem0 import Mem0Memory
|
||||
|
||||
context = {"user_id": "user_1"}
|
||||
context = {"user_id": "alice"}
|
||||
memory_from_client = Mem0Memory.from_client(
|
||||
context=context,
|
||||
api_key="<your-mem0-api-key>",
|
||||
search_msg_limit=4, # optional, default is 5
|
||||
output_format='v1.1', # Remove deprecation warnings
|
||||
)
|
||||
```
|
||||
|
||||
@@ -44,8 +49,8 @@ Context is used to identify the user, agent or the conversation in the Mem0. It
|
||||
|
||||
```python
|
||||
context = {
|
||||
"user_id": "user_1",
|
||||
"agent_id": "agent_1",
|
||||
"user_id": "alice",
|
||||
"agent_id": "llama_agent_1",
|
||||
"run_id": "run_1",
|
||||
}
|
||||
```
|
||||
@@ -98,17 +103,20 @@ memory_from_config = Mem0Memory.from_config(
|
||||
context=context,
|
||||
config=config,
|
||||
search_msg_limit=4, # optional, default is 5
|
||||
output_format='v1.1', # Remove deprecation warnings
|
||||
)
|
||||
```
|
||||
|
||||
Initialize the LLM
|
||||
|
||||
```python
|
||||
import os
|
||||
from llama_index.llms.openai import OpenAI
|
||||
from dotenv import load_dotenv
|
||||
|
||||
os.environ["OPENAI_API_KEY"] = "<your-openai-api-key>"
|
||||
llm = OpenAI(model="gpt-4o")
|
||||
load_dotenv()
|
||||
|
||||
# os.environ["OPENAI_API_KEY"] = "<your-openai-api-key>"
|
||||
llm = OpenAI(model="gpt-4o-mini")
|
||||
```
|
||||
|
||||
### SimpleChatEngine
|
||||
@@ -122,7 +130,7 @@ agent = SimpleChatEngine.from_defaults(
|
||||
)
|
||||
|
||||
# Start the chat
|
||||
response = agent.chat("Hi, My name is Mayank")
|
||||
response = agent.chat("Hi, My name is Alice")
|
||||
print(response)
|
||||
```
|
||||
Now we will learn how to use Mem0 with FunctionCalling and ReAct agents.
|
||||
@@ -165,7 +173,7 @@ agent = FunctionCallingAgent.from_tools(
|
||||
)
|
||||
|
||||
# Start the chat
|
||||
response = agent.chat("Hi, My name is Mayank")
|
||||
response = agent.chat("Hi, My name is Alice")
|
||||
print(response)
|
||||
```
|
||||
|
||||
@@ -182,7 +190,7 @@ agent = ReActAgent.from_tools(
|
||||
)
|
||||
|
||||
# Start the chat
|
||||
response = agent.chat("Hi, My name is Mayank")
|
||||
response = agent.chat("Hi, My name is Alice")
|
||||
print(response)
|
||||
```
|
||||
|
||||
|
||||
@@ -92,7 +92,7 @@ async def websocket_endpoint(websocket: WebSocket):
|
||||
await websocket.accept()
|
||||
|
||||
# Basic setup with minimal configuration
|
||||
user_id = "user123"
|
||||
user_id = "alice"
|
||||
|
||||
# WebSocket transport
|
||||
transport = FastAPIWebsocketTransport(
|
||||
|
||||
Binary file not shown.
|
Before Width: | Height: | Size: 66 KiB After Width: | Height: | Size: 75 KiB |
+22
-6
@@ -1076,9 +1076,14 @@
|
||||
"run_id": {"type": "string"},
|
||||
"created_at": {"type": "string", "format": "date-time"},
|
||||
"updated_at": {"type": "string", "format": "date-time"},
|
||||
"categories": {"type": "array", "items": {"type": "string"}},
|
||||
"categories": {"type": "object", "properties": {
|
||||
"in": {"type": "array", "items": {"type": "string"}}
|
||||
}},
|
||||
"metadata": {"type": "object"},
|
||||
"keywords": {"type": "string"}
|
||||
"keywords": {"type": "object", "properties": {
|
||||
"contains": {"type": "string"},
|
||||
"icontains": {"type": "string"}
|
||||
}}
|
||||
},
|
||||
"additionalProperties": {
|
||||
"type": "object",
|
||||
@@ -1094,7 +1099,7 @@
|
||||
}
|
||||
}
|
||||
},
|
||||
"description": "Filters to apply to the memories. Available fields are: user_id, agent_id, app_id, run_id, created_at, updated_at, categories, keywords. Supports logical operators (AND, OR) and comparison operators (in, gte, lte, gt, lt, ne, contains, icontains)",
|
||||
"description": "Filters to apply to the memories. Available fields are: user_id, agent_id, app_id, run_id, created_at, updated_at, categories, keywords. Supports logical operators (AND, OR) and comparison operators (in, gte, lte, gt, lt, ne, contains, icontains). For categories field, use 'contains' for partial matching (e.g., {\"categories\": {\"contains\": \"finance\"}}) or 'in' for exact matching (e.g., {\"categories\": {\"in\": [\"personal_information\"]}}).",
|
||||
"style": "deepObject",
|
||||
"explode": true
|
||||
},
|
||||
@@ -5074,10 +5079,16 @@
|
||||
"type": "string",
|
||||
"description": "The query to search for in the memory."
|
||||
},
|
||||
"version": {
|
||||
"title": "Version",
|
||||
"type": "string",
|
||||
"default": "v2",
|
||||
"description": "The version of the memory to use. This should always be v2."
|
||||
},
|
||||
"filters": {
|
||||
"title": "Filters",
|
||||
"type": "object",
|
||||
"description": "A dictionary of filters to apply to the search. Available fields are: user_id, agent_id, app_id, run_id, created_at, updated_at, categories, keywords. Supports logical operators (AND, OR) and comparison operators (in, gte, lte, gt, lt, ne, contains, icontains).",
|
||||
"description": "A dictionary of filters to apply to the search. Available fields are: user_id, agent_id, app_id, run_id, created_at, updated_at, categories, keywords. Supports logical operators (AND, OR) and comparison operators (in, gte, lte, gt, lt, ne, contains, icontains). For categories field, use 'contains' for partial matching (e.g., {\"categories\": {\"contains\": \"finance\"}}) or 'in' for exact matching (e.g., {\"categories\": {\"in\": [\"personal_information\"]}}).",
|
||||
"properties": {
|
||||
"user_id": {"type": "string"},
|
||||
"agent_id": {"type": "string"},
|
||||
@@ -5085,8 +5096,13 @@
|
||||
"run_id": {"type": "string"},
|
||||
"created_at": {"type": "string", "format": "date-time"},
|
||||
"updated_at": {"type": "string", "format": "date-time"},
|
||||
"text": {"type": "string"},
|
||||
"categories": {"type": "array", "items": {"type": "string"}},
|
||||
"keywords": {"type": "object", "properties": {
|
||||
"contains": {"type": "string"},
|
||||
"icontains": {"type": "string"}
|
||||
}},
|
||||
"categories": {"type": "object", "properties": {
|
||||
"in": {"type": "array", "items": {"type": "string"}}
|
||||
}},
|
||||
"metadata": {"type": "object"}
|
||||
},
|
||||
"additionalProperties": {
|
||||
|
||||
@@ -285,6 +285,10 @@ curl -X POST "https://api.mem0.ai/v1/memories/" \
|
||||
|
||||
Our advanced search allows you to set custom search filters. You can filter by user_id, agent_id, app_id, run_id, created_at, updated_at, categories, and text. The filters support logical operators (AND, OR) and comparison operators (in, gte, lte, gt, lt, ne, contains, icontains, `*`). The wildcard character (`*`) matches everything for a specific field.
|
||||
|
||||
For the **categories** field specifically:
|
||||
- Use `contains` for partial matching (e.g., `{"categories": {"contains": "finance"}}`)
|
||||
- Use `in` for exact matching (e.g., `{"categories": {"in": ["personal_information"]}}`).
|
||||
|
||||
Here you need to define `version` as `v2` in the search method.
|
||||
|
||||
#### Example 1: Search using user_id and agent_id filters
|
||||
@@ -407,17 +411,32 @@ curl -X POST "https://api.mem0.ai/v1/memories/search/?version=v2" \
|
||||
```
|
||||
</CodeGroup>
|
||||
|
||||
#### Example 3: Search using metadata and categories
|
||||
#### Example 3: Search using categories filters
|
||||
|
||||
<CodeGroup>
|
||||
```python Python
|
||||
query = "What do you know about me?"
|
||||
# Example 3a: Using 'contains' for partial matching
|
||||
query = "What are my financial goals?"
|
||||
filters = {
|
||||
"AND": [
|
||||
{"metadata": {"food": "vegan"}},
|
||||
{ "user_id": "alice" },
|
||||
{
|
||||
"categories": {
|
||||
"contains": "food_preferences"
|
||||
"contains": "finance"
|
||||
}
|
||||
}
|
||||
]
|
||||
}
|
||||
client.search(query, version="v2", filters=filters)
|
||||
|
||||
# Example 3b: Using 'in' for exact matching
|
||||
query = "What personal information do you have?"
|
||||
filters = {
|
||||
"AND": [
|
||||
{ "user_id": "alice" },
|
||||
{
|
||||
"categories": {
|
||||
"in": ["personal_information"]
|
||||
}
|
||||
}
|
||||
]
|
||||
@@ -426,39 +445,72 @@ client.search(query, version="v2", filters=filters)
|
||||
```
|
||||
|
||||
```javascript JavaScript
|
||||
const query = "What do you know about me?";
|
||||
const filters = {
|
||||
// Example 3a: Using 'contains' for partial matching
|
||||
const query1 = "What are my financial goals?";
|
||||
const filters1 = {
|
||||
"AND": [
|
||||
{"metadata": {"food": "vegan"}},
|
||||
{ "user_id": "alice" },
|
||||
{
|
||||
"categories": {
|
||||
"contains": "food_preferences"
|
||||
"contains": "finance"
|
||||
}
|
||||
}
|
||||
]
|
||||
};
|
||||
|
||||
client.search(query, { version: "v2", filters })
|
||||
client.search(query1, { version: "v2", filters: filters1 })
|
||||
.then(results => console.log(results))
|
||||
.catch(error => console.error(error));
|
||||
|
||||
// Example 3b: Using 'in' for exact matching
|
||||
const query2 = "What personal information do you have?";
|
||||
const filters2 = {
|
||||
"AND": [
|
||||
{ "user_id": "alice" },
|
||||
{
|
||||
"categories": {
|
||||
"in": ["personal_information"]
|
||||
}
|
||||
}
|
||||
]
|
||||
};
|
||||
|
||||
client.search(query2, { version: "v2", filters: filters2 })
|
||||
.then(results => console.log(results))
|
||||
.catch(error => console.error(error));
|
||||
```
|
||||
|
||||
```bash cURL
|
||||
# Example 3a: Using 'contains' for partial matching
|
||||
curl -X POST "https://api.mem0.ai/v1/memories/search/?version=v2" \
|
||||
-H "Authorization: Token your-api-key" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"query": "What do you know about me?",
|
||||
"query": "What are my financial goals?",
|
||||
"filters": {
|
||||
"AND": [
|
||||
{
|
||||
"metadata": {
|
||||
"food": "vegan"
|
||||
}
|
||||
},
|
||||
{ "user_id": "alice" },
|
||||
{
|
||||
"categories": {
|
||||
"contains": "food_preferences"
|
||||
"contains": "finance"
|
||||
}
|
||||
}
|
||||
]
|
||||
}
|
||||
}'
|
||||
|
||||
# Example 3b: Using 'in' for exact matching
|
||||
curl -X POST "https://api.mem0.ai/v1/memories/search/?version=v2" \
|
||||
-H "Authorization: Token your-api-key" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"query": "What personal information do you have?",
|
||||
"filters": {
|
||||
"AND": [
|
||||
{ "user_id": "alice" },
|
||||
{
|
||||
"categories": {
|
||||
"in": ["personal_information"]
|
||||
}
|
||||
}
|
||||
]
|
||||
@@ -693,6 +745,10 @@ curl -X GET "https://api.mem0.ai/v1/memories/?user_id=alex&keywords=to play&page
|
||||
|
||||
Our advanced retrieval allows you to set custom filters when fetching memories. You can filter by user_id, agent_id, app_id, run_id, created_at, updated_at, categories, and keywords. The filters support logical operators (AND, OR) and comparison operators (in, gte, lte, gt, lt, ne, contains, icontains, `*`). The wildcard character (`*`) matches everything for a specific field.
|
||||
|
||||
For the **categories** field specifically:
|
||||
- Use `contains` for partial matching (e.g., `{"categories": {"contains": "finance"}}`)
|
||||
- Use `in` for exact matching (e.g., `{"categories": {"in": ["personal_information"]}}`).
|
||||
|
||||
Here you need to define `version` as `v2` in the get_all method.
|
||||
|
||||
<CodeGroup>
|
||||
|
||||
@@ -9,8 +9,6 @@ description: "Advanced memory search with keyword expansion, intelligent reranki
|
||||
|
||||
Advanced Retrieval gives you precise control over how memories are found and ranked. While basic search uses semantic similarity, these advanced options help you find exactly what you need, when you need it.
|
||||
|
||||
<CodeGroup>
|
||||
|
||||
## Search Enhancement Options
|
||||
|
||||
### Keyword Search
|
||||
|
||||
@@ -36,7 +36,7 @@ Learn about the key features and capabilities that make Mem0 a powerful platform
|
||||
<Card title="Memory Export" icon="file-export" href="memory-export">
|
||||
Export memories in structured formats using customizable Pydantic schemas.
|
||||
</Card>
|
||||
<Card title="Graph Memory" icon="graph" href="graph-memory">
|
||||
<Card title="Graph Memory" icon="circle-nodes" href="graph-memory">
|
||||
Add memories in the form of nodes and edges in a graph database and search for related memories.
|
||||
</Card>
|
||||
</CardGroup>
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
{
|
||||
"name": "mem0ai",
|
||||
"version": "2.1.37",
|
||||
"version": "2.1.38",
|
||||
"description": "The Memory Layer For Your AI Apps",
|
||||
"main": "./dist/index.js",
|
||||
"module": "./dist/index.mjs",
|
||||
|
||||
@@ -251,11 +251,19 @@ export default class MemoryClient {
|
||||
return response;
|
||||
}
|
||||
|
||||
async update(memoryId: string, message: string): Promise<Array<Memory>> {
|
||||
async update(
|
||||
memoryId: string,
|
||||
{ text, metadata }: { text?: string; metadata?: Record<string, any> },
|
||||
): Promise<Array<Memory>> {
|
||||
if (text === undefined && metadata === undefined) {
|
||||
throw new Error("Either text or metadata must be provided for update.");
|
||||
}
|
||||
|
||||
if (this.telemetryId === "") await this.ping();
|
||||
this._validateOrgProject();
|
||||
const payload = {
|
||||
text: message,
|
||||
text: text,
|
||||
metadata: metadata,
|
||||
};
|
||||
|
||||
const payloadKeys = Object.keys(payload);
|
||||
|
||||
@@ -89,7 +89,7 @@ export class MemoryGraph {
|
||||
|
||||
this.llm = LLMFactory.create(this.llmProvider, this.config.llm.config);
|
||||
this.structuredLlm = LLMFactory.create(
|
||||
"openai_structured",
|
||||
this.llmProvider,
|
||||
this.config.llm.config,
|
||||
);
|
||||
this.threshold = 0.7;
|
||||
|
||||
@@ -3,7 +3,7 @@ from abc import ABC, abstractmethod
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
import httpx
|
||||
from pydantic import BaseModel, Field
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
|
||||
from mem0.client.utils import api_error_handler
|
||||
from mem0.memory.telemetry import capture_client_event
|
||||
@@ -20,9 +20,7 @@ class ProjectConfig(BaseModel):
|
||||
project_id: Optional[str] = Field(default=None, description="Project ID")
|
||||
user_email: Optional[str] = Field(default=None, description="User email")
|
||||
|
||||
class Config:
|
||||
validate_assignment = True
|
||||
extra = "forbid"
|
||||
model_config = ConfigDict(validate_assignment=True, extra="forbid")
|
||||
|
||||
|
||||
class BaseProject(ABC):
|
||||
|
||||
@@ -28,6 +28,7 @@ class OpenAIConfig(BaseLlmConfig):
|
||||
openrouter_base_url: Optional[str] = None,
|
||||
site_url: Optional[str] = None,
|
||||
app_name: Optional[str] = None,
|
||||
store: bool = False,
|
||||
# Response monitoring callback
|
||||
response_callback: Optional[Callable[[Any, dict, dict], None]] = None,
|
||||
):
|
||||
@@ -72,5 +73,7 @@ class OpenAIConfig(BaseLlmConfig):
|
||||
self.openrouter_base_url = openrouter_base_url
|
||||
self.site_url = site_url
|
||||
self.app_name = app_name
|
||||
self.store = store
|
||||
|
||||
# Response monitoring
|
||||
self.response_callback = response_callback
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
from pydantic import BaseModel, Field, model_validator
|
||||
from pydantic import BaseModel, ConfigDict, Field, model_validator
|
||||
|
||||
|
||||
class AzureAISearchConfig(BaseModel):
|
||||
@@ -54,6 +54,4 @@ class AzureAISearchConfig(BaseModel):
|
||||
|
||||
return values
|
||||
|
||||
model_config = {
|
||||
"arbitrary_types_allowed": True,
|
||||
}
|
||||
model_config = ConfigDict(arbitrary_types_allowed=True)
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
from typing import Any, Dict
|
||||
|
||||
from pydantic import BaseModel, Field, model_validator
|
||||
from pydantic import BaseModel, ConfigDict, Field, model_validator
|
||||
|
||||
|
||||
class BaiduDBConfig(BaseModel):
|
||||
@@ -24,6 +24,4 @@ class BaiduDBConfig(BaseModel):
|
||||
)
|
||||
return values
|
||||
|
||||
model_config = {
|
||||
"arbitrary_types_allowed": True,
|
||||
}
|
||||
model_config = ConfigDict(arbitrary_types_allowed=True)
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
from typing import Any, ClassVar, Dict, Optional
|
||||
|
||||
from pydantic import BaseModel, Field, model_validator
|
||||
from pydantic import BaseModel, ConfigDict, Field, model_validator
|
||||
|
||||
|
||||
class ChromaDbConfig(BaseModel):
|
||||
@@ -10,17 +10,37 @@ class ChromaDbConfig(BaseModel):
|
||||
raise ImportError("The 'chromadb' library is required. Please install it using 'pip install chromadb'.")
|
||||
Client: ClassVar[type] = Client
|
||||
|
||||
collection_name: str = Field("mem0", description="Default name for the collection")
|
||||
collection_name: str = Field("mem0", description="Default name for the collection/database")
|
||||
client: Optional[Client] = Field(None, description="Existing ChromaDB client instance")
|
||||
path: Optional[str] = Field(None, description="Path to the database directory")
|
||||
host: Optional[str] = Field(None, description="Database connection remote host")
|
||||
port: Optional[int] = Field(None, description="Database connection remote port")
|
||||
# ChromaDB Cloud configuration
|
||||
api_key: Optional[str] = Field(None, description="ChromaDB Cloud API key")
|
||||
tenant: Optional[str] = Field(None, description="ChromaDB Cloud tenant ID")
|
||||
|
||||
@model_validator(mode="before")
|
||||
def check_host_port_or_path(cls, values):
|
||||
def check_connection_config(cls, values):
|
||||
host, port, path = values.get("host"), values.get("port"), values.get("path")
|
||||
if not path and not (host and port):
|
||||
raise ValueError("Either 'host' and 'port' or 'path' must be provided.")
|
||||
api_key, tenant = values.get("api_key"), values.get("tenant")
|
||||
|
||||
# Check if cloud configuration is provided
|
||||
cloud_config = bool(api_key and tenant)
|
||||
|
||||
# If cloud configuration is provided, remove any default path that might have been added
|
||||
if cloud_config and path == "/tmp/chroma":
|
||||
values.pop("path", None)
|
||||
return values
|
||||
|
||||
# Check if local/server configuration is provided (excluding default tmp path for cloud config)
|
||||
local_config = bool(path and path != "/tmp/chroma") or bool(host and port)
|
||||
|
||||
if not cloud_config and not local_config:
|
||||
raise ValueError("Either ChromaDB Cloud configuration (api_key, tenant) or local configuration (path or host/port) must be provided.")
|
||||
|
||||
if cloud_config and local_config:
|
||||
raise ValueError("Cannot specify both cloud configuration and local configuration. Choose one.")
|
||||
|
||||
return values
|
||||
|
||||
@model_validator(mode="before")
|
||||
@@ -35,6 +55,4 @@ class ChromaDbConfig(BaseModel):
|
||||
)
|
||||
return values
|
||||
|
||||
model_config = {
|
||||
"arbitrary_types_allowed": True,
|
||||
}
|
||||
model_config = ConfigDict(arbitrary_types_allowed=True)
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
from pydantic import BaseModel, Field, model_validator
|
||||
from pydantic import BaseModel, ConfigDict, Field, model_validator
|
||||
|
||||
from databricks.sdk.service.vectorsearch import EndpointType, VectorIndexType, PipelineType
|
||||
|
||||
@@ -58,6 +58,4 @@ class DatabricksConfig(BaseModel):
|
||||
|
||||
return self
|
||||
|
||||
model_config = {
|
||||
"arbitrary_types_allowed": True,
|
||||
}
|
||||
model_config = ConfigDict(arbitrary_types_allowed=True)
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
from pydantic import BaseModel, Field, model_validator
|
||||
from pydantic import BaseModel, ConfigDict, Field, model_validator
|
||||
|
||||
|
||||
class FAISSConfig(BaseModel):
|
||||
@@ -34,6 +34,4 @@ class FAISSConfig(BaseModel):
|
||||
)
|
||||
return values
|
||||
|
||||
model_config = {
|
||||
"arbitrary_types_allowed": True,
|
||||
}
|
||||
model_config = ConfigDict(arbitrary_types_allowed=True)
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
from typing import Any, ClassVar, Dict
|
||||
|
||||
from pydantic import BaseModel, Field, model_validator
|
||||
from pydantic import BaseModel, ConfigDict, Field, model_validator
|
||||
|
||||
|
||||
class LangchainConfig(BaseModel):
|
||||
@@ -27,6 +27,4 @@ class LangchainConfig(BaseModel):
|
||||
)
|
||||
return values
|
||||
|
||||
model_config = {
|
||||
"arbitrary_types_allowed": True,
|
||||
}
|
||||
model_config = ConfigDict(arbitrary_types_allowed=True)
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
from enum import Enum
|
||||
from typing import Any, Dict
|
||||
|
||||
from pydantic import BaseModel, Field, model_validator
|
||||
from pydantic import BaseModel, ConfigDict, Field, model_validator
|
||||
|
||||
|
||||
class MetricType(str, Enum):
|
||||
@@ -39,6 +39,4 @@ class MilvusDBConfig(BaseModel):
|
||||
)
|
||||
return values
|
||||
|
||||
model_config = {
|
||||
"arbitrary_types_allowed": True,
|
||||
}
|
||||
model_config = ConfigDict(arbitrary_types_allowed=True)
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
import os
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
from pydantic import BaseModel, Field, model_validator
|
||||
from pydantic import BaseModel, ConfigDict, Field, model_validator
|
||||
|
||||
|
||||
class PineconeConfig(BaseModel):
|
||||
@@ -52,6 +52,4 @@ class PineconeConfig(BaseModel):
|
||||
)
|
||||
return values
|
||||
|
||||
model_config = {
|
||||
"arbitrary_types_allowed": True,
|
||||
}
|
||||
model_config = ConfigDict(arbitrary_types_allowed=True)
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
from typing import Any, ClassVar, Dict, Optional
|
||||
|
||||
from pydantic import BaseModel, Field, model_validator
|
||||
from pydantic import BaseModel, ConfigDict, Field, model_validator
|
||||
|
||||
|
||||
class QdrantConfig(BaseModel):
|
||||
@@ -44,6 +44,4 @@ class QdrantConfig(BaseModel):
|
||||
)
|
||||
return values
|
||||
|
||||
model_config = {
|
||||
"arbitrary_types_allowed": True,
|
||||
}
|
||||
model_config = ConfigDict(arbitrary_types_allowed=True)
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
from typing import Any, Dict
|
||||
|
||||
from pydantic import BaseModel, Field, model_validator
|
||||
from pydantic import BaseModel, ConfigDict, Field, model_validator
|
||||
|
||||
|
||||
# TODO: Upgrade to latest pydantic version
|
||||
@@ -21,6 +21,4 @@ class RedisDBConfig(BaseModel):
|
||||
)
|
||||
return values
|
||||
|
||||
model_config = {
|
||||
"arbitrary_types_allowed": True,
|
||||
}
|
||||
model_config = ConfigDict(arbitrary_types_allowed=True)
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
from pydantic import BaseModel, Field, model_validator
|
||||
from pydantic import BaseModel, ConfigDict, Field, model_validator
|
||||
|
||||
|
||||
class S3VectorsConfig(BaseModel):
|
||||
@@ -29,6 +29,4 @@ class S3VectorsConfig(BaseModel):
|
||||
)
|
||||
return values
|
||||
|
||||
model_config = {
|
||||
"arbitrary_types_allowed": True,
|
||||
}
|
||||
model_config = ConfigDict(arbitrary_types_allowed=True)
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
import os
|
||||
from typing import Any, ClassVar, Dict, Optional
|
||||
|
||||
from pydantic import BaseModel, Field, model_validator
|
||||
from pydantic import BaseModel, ConfigDict, Field, model_validator
|
||||
|
||||
try:
|
||||
from upstash_vector import Index
|
||||
@@ -31,6 +31,4 @@ class UpstashVectorConfig(BaseModel):
|
||||
raise ValueError("Either a client or URL and token must be provided.")
|
||||
return values
|
||||
|
||||
model_config = {
|
||||
"arbitrary_types_allowed": True,
|
||||
}
|
||||
model_config = ConfigDict(arbitrary_types_allowed=True)
|
||||
|
||||
@@ -0,0 +1,15 @@
|
||||
from pydantic import BaseModel
|
||||
|
||||
|
||||
class ValkeyConfig(BaseModel):
|
||||
"""Configuration for Valkey vector store."""
|
||||
|
||||
valkey_url: str
|
||||
collection_name: str
|
||||
embedding_model_dims: int
|
||||
timezone: str = "UTC"
|
||||
index_type: str = "hnsw" # Default to HNSW, can be 'hnsw' or 'flat'
|
||||
# HNSW specific parameters with recommended defaults
|
||||
hnsw_m: int = 16 # Number of connections per layer (default from Valkey docs)
|
||||
hnsw_ef_construction: int = 200 # Search width during construction
|
||||
hnsw_ef_runtime: int = 10 # Search width during queries
|
||||
@@ -1,6 +1,6 @@
|
||||
from typing import Dict, Optional
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
|
||||
|
||||
class GoogleMatchingEngineConfig(BaseModel):
|
||||
@@ -15,7 +15,7 @@ class GoogleMatchingEngineConfig(BaseModel):
|
||||
service_account_json: Optional[Dict] = Field(None, description="Service account credentials as dictionary (alternative to credentials_path)")
|
||||
vector_search_api_endpoint: Optional[str] = Field(None, description="Vector search API endpoint")
|
||||
|
||||
model_config = {"extra": "forbid"}
|
||||
model_config = ConfigDict(extra="forbid")
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
from typing import Any, ClassVar, Dict, Optional
|
||||
|
||||
from pydantic import BaseModel, Field, model_validator
|
||||
from pydantic import BaseModel, ConfigDict, Field, model_validator
|
||||
|
||||
|
||||
class WeaviateConfig(BaseModel):
|
||||
@@ -38,6 +38,4 @@ class WeaviateConfig(BaseModel):
|
||||
|
||||
return values
|
||||
|
||||
model_config = {
|
||||
"arbitrary_types_allowed": True,
|
||||
}
|
||||
model_config = ConfigDict(arbitrary_types_allowed=True)
|
||||
|
||||
@@ -50,7 +50,7 @@ Entity Consistency:
|
||||
- Ensure that relationships are coherent and logically align with the context of the message.
|
||||
- Maintain consistent naming for entities across the extracted data.
|
||||
|
||||
Strive to construct a coherent and easily understandable knowledge graph by eshtablishing all the relationships among the entities and adherence to the user’s context.
|
||||
Strive to construct a coherent and easily understandable knowledge graph by establishing all the relationships among the entities and adherence to the user’s context.
|
||||
|
||||
Adhere strictly to these guidelines to ensure high-quality knowledge graph extraction."""
|
||||
|
||||
|
||||
+86
-27
@@ -136,23 +136,27 @@ class AWSBedrockLLM(LLMBase):
|
||||
else:
|
||||
self._format_messages = self._format_messages_generic
|
||||
|
||||
def _format_messages_anthropic(self, messages: List[Dict[str, str]]) -> List[Dict[str, Any]]:
|
||||
def _format_messages_anthropic(self, messages: List[Dict[str, str]]) -> tuple[List[Dict[str, Any]], Optional[str]]:
|
||||
"""Format messages for Anthropic models."""
|
||||
formatted_messages = []
|
||||
system_message = None
|
||||
|
||||
for message in messages:
|
||||
role = message["role"]
|
||||
content = message["content"]
|
||||
|
||||
if role == "system":
|
||||
# Anthropic doesn't support system messages, prepend to first user message
|
||||
continue
|
||||
# Anthropic supports system messages as a separate parameter
|
||||
# see: https://docs.anthropic.com/en/docs/build-with-claude/prompt-engineering/system-prompts
|
||||
system_message = content
|
||||
elif role == "user":
|
||||
formatted_messages.append({"role": "user", "content": [{"type": "text", "text": content}]})
|
||||
# Use Converse API format
|
||||
formatted_messages.append({"role": "user", "content": [{"text": content}]})
|
||||
elif role == "assistant":
|
||||
formatted_messages.append({"role": "assistant", "content": [{"type": "text", "text": content}]})
|
||||
# Use Converse API format
|
||||
formatted_messages.append({"role": "assistant", "content": [{"text": content}]})
|
||||
|
||||
return formatted_messages
|
||||
return formatted_messages, system_message
|
||||
|
||||
def _format_messages_cohere(self, messages: List[Dict[str, str]]) -> str:
|
||||
"""Format messages for Cohere models."""
|
||||
@@ -451,48 +455,103 @@ class AWSBedrockLLM(LLMBase):
|
||||
logger.error(f"Failed to generate response: {e}")
|
||||
raise RuntimeError(f"Failed to generate response: {e}")
|
||||
|
||||
@staticmethod
|
||||
def _convert_tools_to_converse_format(tools: List[Dict]) -> List[Dict]:
|
||||
"""Convert OpenAI-style tools to Converse API format."""
|
||||
if not tools:
|
||||
return []
|
||||
|
||||
converse_tools = []
|
||||
for tool in tools:
|
||||
if tool.get("type") == "function" and "function" in tool:
|
||||
func = tool["function"]
|
||||
converse_tool = {
|
||||
"toolSpec": {
|
||||
"name": func["name"],
|
||||
"description": func.get("description", ""),
|
||||
"inputSchema": {
|
||||
"json": func.get("parameters", {})
|
||||
}
|
||||
}
|
||||
}
|
||||
converse_tools.append(converse_tool)
|
||||
|
||||
return converse_tools
|
||||
|
||||
def _generate_with_tools(self, messages: List[Dict[str, str]], tools: List[Dict], stream: bool = False) -> Dict[str, Any]:
|
||||
"""Generate response with tool calling support."""
|
||||
"""Generate response with tool calling support using correct message format."""
|
||||
# Format messages for tool-enabled models
|
||||
system_message = None
|
||||
if self.provider == "anthropic":
|
||||
formatted_messages = self._format_messages_anthropic(messages)
|
||||
formatted_messages, system_message = self._format_messages_anthropic(messages)
|
||||
elif self.provider == "amazon":
|
||||
formatted_messages = self._format_messages_amazon(messages)
|
||||
else:
|
||||
formatted_messages = [{"role": "user", "content": messages[-1]["content"]}]
|
||||
formatted_messages = [{"role": "user", "content": [{"text": messages[-1]["content"]}]}]
|
||||
|
||||
# Prepare inference configuration
|
||||
inference_config = {
|
||||
"temperature": self.model_config.get("temperature", 0.1),
|
||||
# Prepare tool configuration in Converse API format
|
||||
tool_config = None
|
||||
if tools:
|
||||
converse_tools = self._convert_tools_to_converse_format(tools)
|
||||
if converse_tools:
|
||||
tool_config = {"tools": converse_tools}
|
||||
|
||||
# Prepare converse parameters
|
||||
converse_params = {
|
||||
"modelId": self.config.model,
|
||||
"messages": formatted_messages,
|
||||
"inferenceConfig": {
|
||||
"maxTokens": self.model_config.get("max_tokens", 2000),
|
||||
"temperature": self.model_config.get("temperature", 0.1),
|
||||
"topP": self.model_config.get("top_p", 0.9),
|
||||
}
|
||||
}
|
||||
|
||||
# Prepare tools configuration
|
||||
tools_config = {"tools": self._convert_tool_format(tools)}
|
||||
# Add system message if present (for Anthropic)
|
||||
if system_message:
|
||||
converse_params["system"] = [{"text": system_message}]
|
||||
|
||||
# Add tool config if present
|
||||
if tool_config:
|
||||
converse_params["toolConfig"] = tool_config
|
||||
|
||||
# Make API call
|
||||
response = self.client.converse(
|
||||
modelId=self.config.model,
|
||||
messages=formatted_messages,
|
||||
inferenceConfig=inference_config,
|
||||
toolConfig=tools_config,
|
||||
)
|
||||
response = self.client.converse(**converse_params)
|
||||
|
||||
return self._parse_response(response, tools)
|
||||
|
||||
def _generate_standard(self, messages: List[Dict[str, str]], stream: bool = False) -> str:
|
||||
"""Generate standard text response."""
|
||||
# Format messages according to provider
|
||||
"""Generate standard text response using Converse API for Anthropic models."""
|
||||
# For Anthropic models, always use Converse API
|
||||
if self.provider == "anthropic":
|
||||
formatted_messages = self._format_messages_anthropic(messages)
|
||||
input_body = {
|
||||
formatted_messages, system_message = self._format_messages_anthropic(messages)
|
||||
|
||||
# Prepare converse parameters
|
||||
converse_params = {
|
||||
"modelId": self.config.model,
|
||||
"messages": formatted_messages,
|
||||
"max_tokens": self.model_config.get("max_tokens", 2000),
|
||||
"inferenceConfig": {
|
||||
"maxTokens": self.model_config.get("max_tokens", 2000),
|
||||
"temperature": self.model_config.get("temperature", 0.1),
|
||||
"top_p": self.model_config.get("top_p", 0.9),
|
||||
"anthropic_version": "bedrock-2023-05-31",
|
||||
"topP": self.model_config.get("top_p", 0.9),
|
||||
}
|
||||
}
|
||||
|
||||
# Add system message if present
|
||||
if system_message:
|
||||
converse_params["system"] = [{"text": system_message}]
|
||||
|
||||
# Use converse API for Anthropic models
|
||||
response = self.client.converse(**converse_params)
|
||||
|
||||
# Parse Converse API response
|
||||
if hasattr(response, 'output') and hasattr(response.output, 'message'):
|
||||
return response.output.message.content[0].text
|
||||
elif 'output' in response and 'message' in response['output']:
|
||||
return response['output']['message']['content'][0]['text']
|
||||
else:
|
||||
return str(response)
|
||||
|
||||
elif self.provider == "amazon" and "nova" in self.config.model.lower():
|
||||
# Nova models use converse API even without tools
|
||||
formatted_messages = self._format_messages_amazon(messages)
|
||||
|
||||
@@ -91,6 +91,14 @@ class OllamaLLM(LLMBase):
|
||||
"messages": messages,
|
||||
}
|
||||
|
||||
# Handle JSON response format by modifying the system prompt
|
||||
if response_format and response_format.get("type") == "json_object":
|
||||
# Add JSON format instruction to the last message or create a system message
|
||||
if messages and messages[-1]["role"] == "user":
|
||||
messages[-1]["content"] += "\n\nPlease respond with valid JSON only."
|
||||
else:
|
||||
messages.append({"role": "user", "content": "Please respond with valid JSON only."})
|
||||
|
||||
# Add options for Ollama (temperature, num_predict, top_p)
|
||||
options = {
|
||||
"temperature": self.config.temperature,
|
||||
|
||||
@@ -124,6 +124,12 @@ class OpenAILLM(LLMBase):
|
||||
|
||||
params.update(**openrouter_params)
|
||||
|
||||
else:
|
||||
openai_specific_generation_params = ["store"]
|
||||
for param in openai_specific_generation_params:
|
||||
if hasattr(self.config, param):
|
||||
params[param] = getattr(self.config, param)
|
||||
|
||||
if response_format:
|
||||
params["response_format"] = response_format
|
||||
if tools: # TODO: Remove tools if no issues found with new memory addition logic
|
||||
|
||||
@@ -7,6 +7,7 @@ import logging
|
||||
import os
|
||||
import uuid
|
||||
import warnings
|
||||
|
||||
from copy import deepcopy
|
||||
from datetime import datetime
|
||||
from typing import Any, Dict, Optional
|
||||
@@ -39,6 +40,9 @@ from mem0.utils.factory import (
|
||||
RerankerFactory,
|
||||
)
|
||||
|
||||
# Suppress SWIG deprecation warnings globally
|
||||
warnings.filterwarnings("ignore", category=DeprecationWarning, message=".*SwigPy.*")
|
||||
warnings.filterwarnings("ignore", category=DeprecationWarning, message=".*swigvarlink.*")
|
||||
|
||||
def _build_filters_and_metadata(
|
||||
*, # Enforce keyword-only arguments
|
||||
@@ -416,6 +420,10 @@ class Memory(MemoryBase):
|
||||
response = ""
|
||||
|
||||
try:
|
||||
if not response or not response.strip():
|
||||
logger.warning("Empty response from LLM, no memories to extract")
|
||||
new_memories_with_actions = {}
|
||||
else:
|
||||
response = remove_code_blocks(response)
|
||||
new_memories_with_actions = json.loads(response)
|
||||
except Exception as e:
|
||||
@@ -1412,6 +1420,10 @@ class AsyncMemory(MemoryBase):
|
||||
logger.error(f"Error in new memory actions response: {e}")
|
||||
response = ""
|
||||
try:
|
||||
if not response or not response.strip():
|
||||
logger.warning("Empty response from LLM, no memories to extract")
|
||||
new_memories_with_actions = {}
|
||||
else:
|
||||
response = remove_code_blocks(response)
|
||||
new_memories_with_actions = json.loads(response)
|
||||
except Exception as e:
|
||||
|
||||
@@ -170,6 +170,7 @@ class VectorStoreFactory:
|
||||
"pinecone": "mem0.vector_stores.pinecone.PineconeDB",
|
||||
"mongodb": "mem0.vector_stores.mongodb.MongoDB",
|
||||
"redis": "mem0.vector_stores.redis.RedisDB",
|
||||
"valkey": "mem0.vector_stores.valkey.ValkeyDB",
|
||||
"databricks": "mem0.vector_stores.databricks.Databricks",
|
||||
"elasticsearch": "mem0.vector_stores.elasticsearch.ElasticsearchDB",
|
||||
"vertex_ai_vector_search": "mem0.vector_stores.vertex_ai_vector_search.GoogleMatchingEngine",
|
||||
@@ -179,6 +180,7 @@ class VectorStoreFactory:
|
||||
"faiss": "mem0.vector_stores.faiss.FAISS",
|
||||
"langchain": "mem0.vector_stores.langchain.Langchain",
|
||||
"s3_vectors": "mem0.vector_stores.s3_vectors.S3Vectors",
|
||||
"baidu": "mem0.vector_stores.baidu.BaiduDB",
|
||||
}
|
||||
|
||||
@classmethod
|
||||
|
||||
@@ -28,6 +28,8 @@ class ChromaDB(VectorStoreBase):
|
||||
host: Optional[str] = None,
|
||||
port: Optional[int] = None,
|
||||
path: Optional[str] = None,
|
||||
api_key: Optional[str] = None,
|
||||
tenant: Optional[str] = None,
|
||||
):
|
||||
"""
|
||||
Initialize the Chromadb vector store.
|
||||
@@ -38,10 +40,21 @@ class ChromaDB(VectorStoreBase):
|
||||
host (str, optional): Host address for chromadb server. Defaults to None.
|
||||
port (int, optional): Port for chromadb server. Defaults to None.
|
||||
path (str, optional): Path for local chromadb database. Defaults to None.
|
||||
api_key (str, optional): ChromaDB Cloud API key. Defaults to None.
|
||||
tenant (str, optional): ChromaDB Cloud tenant ID. Defaults to None.
|
||||
"""
|
||||
if client:
|
||||
self.client = client
|
||||
elif api_key and tenant:
|
||||
# Initialize ChromaDB Cloud client
|
||||
logger.info("Initializing ChromaDB Cloud client")
|
||||
self.client = chromadb.CloudClient(
|
||||
api_key=api_key,
|
||||
tenant=tenant,
|
||||
database="mem0" # Use fixed database name for cloud
|
||||
)
|
||||
else:
|
||||
# Initialize local or server client
|
||||
self.settings = Settings(anonymized_telemetry=False)
|
||||
|
||||
if host and port:
|
||||
|
||||
@@ -21,6 +21,7 @@ class VectorStoreConfig(BaseModel):
|
||||
"upstash_vector": "UpstashVectorConfig",
|
||||
"azure_ai_search": "AzureAISearchConfig",
|
||||
"redis": "RedisDBConfig",
|
||||
"valkey": "ValkeyConfig",
|
||||
"databricks": "DatabricksConfig",
|
||||
"elasticsearch": "ElasticsearchConfig",
|
||||
"vertex_ai_vector_search": "GoogleMatchingEngineConfig",
|
||||
|
||||
@@ -465,7 +465,7 @@ class Databricks(VectorStoreBase):
|
||||
|
||||
# Parse results
|
||||
result_data = sdk_results.result if hasattr(sdk_results, "result") else sdk_results
|
||||
data_array = result_data.data_array if hasattr(result_data, "data_array") else []
|
||||
data_array = result_data.data_array if getattr(result_data, "data_array", None) else []
|
||||
|
||||
memory_results = []
|
||||
for row in data_array:
|
||||
@@ -708,7 +708,7 @@ class Databricks(VectorStoreBase):
|
||||
pass
|
||||
memory_id = row_dict.get("memory_id") or row_dict.get("id")
|
||||
memory_results.append(MemoryResult(id=memory_id, payload=payload))
|
||||
return memory_results
|
||||
return [memory_results]
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to list memories: {e}")
|
||||
return []
|
||||
|
||||
@@ -8,7 +8,13 @@ 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.*")
|
||||
warnings.filterwarnings("ignore", category=DeprecationWarning, message=".*swigvarlink.*")
|
||||
|
||||
logging.getLogger("faiss").setLevel(logging.WARNING)
|
||||
logging.getLogger("faiss.loader").setLevel(logging.WARNING)
|
||||
|
||||
|
||||
@@ -0,0 +1,824 @@
|
||||
import json
|
||||
import logging
|
||||
from datetime import datetime
|
||||
from typing import Dict
|
||||
|
||||
import numpy as np
|
||||
import pytz
|
||||
import valkey
|
||||
from pydantic import BaseModel
|
||||
from valkey.exceptions import ResponseError
|
||||
|
||||
from mem0.memory.utils import extract_json
|
||||
from mem0.vector_stores.base import VectorStoreBase
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Default fields for the Valkey index
|
||||
DEFAULT_FIELDS = [
|
||||
{"name": "memory_id", "type": "tag"},
|
||||
{"name": "hash", "type": "tag"},
|
||||
{"name": "agent_id", "type": "tag"},
|
||||
{"name": "run_id", "type": "tag"},
|
||||
{"name": "user_id", "type": "tag"},
|
||||
{"name": "memory", "type": "tag"}, # Using TAG instead of TEXT for Valkey compatibility
|
||||
{"name": "metadata", "type": "tag"}, # Using TAG instead of TEXT for Valkey compatibility
|
||||
{"name": "created_at", "type": "numeric"},
|
||||
{"name": "updated_at", "type": "numeric"},
|
||||
{
|
||||
"name": "embedding",
|
||||
"type": "vector",
|
||||
"attrs": {"distance_metric": "cosine", "algorithm": "flat", "datatype": "float32"},
|
||||
},
|
||||
]
|
||||
|
||||
excluded_keys = {"user_id", "agent_id", "run_id", "hash", "data", "created_at", "updated_at"}
|
||||
|
||||
|
||||
class OutputData(BaseModel):
|
||||
id: str
|
||||
score: float
|
||||
payload: Dict
|
||||
|
||||
|
||||
class ValkeyDB(VectorStoreBase):
|
||||
def __init__(
|
||||
self,
|
||||
valkey_url: str,
|
||||
collection_name: str,
|
||||
embedding_model_dims: int,
|
||||
timezone: str = "UTC",
|
||||
index_type: str = "hnsw",
|
||||
hnsw_m: int = 16,
|
||||
hnsw_ef_construction: int = 200,
|
||||
hnsw_ef_runtime: int = 10,
|
||||
):
|
||||
"""
|
||||
Initialize the Valkey vector store.
|
||||
|
||||
Args:
|
||||
valkey_url (str): Valkey URL.
|
||||
collection_name (str): Collection name.
|
||||
embedding_model_dims (int): Embedding model dimensions.
|
||||
timezone (str, optional): Timezone for timestamps. Defaults to "UTC".
|
||||
index_type (str, optional): Index type ('hnsw' or 'flat'). Defaults to "hnsw".
|
||||
hnsw_m (int, optional): HNSW M parameter (connections per node). Defaults to 16.
|
||||
hnsw_ef_construction (int, optional): HNSW ef_construction parameter. Defaults to 200.
|
||||
hnsw_ef_runtime (int, optional): HNSW ef_runtime parameter. Defaults to 10.
|
||||
"""
|
||||
self.embedding_model_dims = embedding_model_dims
|
||||
self.collection_name = collection_name
|
||||
self.prefix = f"mem0:{collection_name}"
|
||||
self.timezone = timezone
|
||||
self.index_type = index_type.lower()
|
||||
self.hnsw_m = hnsw_m
|
||||
self.hnsw_ef_construction = hnsw_ef_construction
|
||||
self.hnsw_ef_runtime = hnsw_ef_runtime
|
||||
|
||||
# Validate index type
|
||||
if self.index_type not in ["hnsw", "flat"]:
|
||||
raise ValueError(f"Invalid index_type: {index_type}. Must be 'hnsw' or 'flat'")
|
||||
|
||||
# Connect to Valkey
|
||||
try:
|
||||
self.client = valkey.from_url(valkey_url)
|
||||
logger.debug(f"Successfully connected to Valkey at {valkey_url}")
|
||||
except Exception as e:
|
||||
logger.exception(f"Failed to connect to Valkey at {valkey_url}: {e}")
|
||||
raise
|
||||
|
||||
# Create the index schema
|
||||
self._create_index(embedding_model_dims)
|
||||
|
||||
def _build_index_schema(self, collection_name, embedding_dims, distance_metric, prefix):
|
||||
"""
|
||||
Build the FT.CREATE command for index creation.
|
||||
|
||||
Args:
|
||||
collection_name (str): Name of the collection/index
|
||||
embedding_dims (int): Vector embedding dimensions
|
||||
distance_metric (str): Distance metric (e.g., "COSINE", "L2", "IP")
|
||||
prefix (str): Key prefix for the index
|
||||
|
||||
Returns:
|
||||
list: Complete FT.CREATE command as list of arguments
|
||||
"""
|
||||
# Build the vector field configuration based on index type
|
||||
if self.index_type == "hnsw":
|
||||
vector_config = [
|
||||
"embedding",
|
||||
"VECTOR",
|
||||
"HNSW",
|
||||
"12", # Attribute count: TYPE, FLOAT32, DIM, dims, DISTANCE_METRIC, metric, M, m, EF_CONSTRUCTION, ef_construction, EF_RUNTIME, ef_runtime
|
||||
"TYPE",
|
||||
"FLOAT32",
|
||||
"DIM",
|
||||
str(embedding_dims),
|
||||
"DISTANCE_METRIC",
|
||||
distance_metric,
|
||||
"M",
|
||||
str(self.hnsw_m),
|
||||
"EF_CONSTRUCTION",
|
||||
str(self.hnsw_ef_construction),
|
||||
"EF_RUNTIME",
|
||||
str(self.hnsw_ef_runtime),
|
||||
]
|
||||
elif self.index_type == "flat":
|
||||
vector_config = [
|
||||
"embedding",
|
||||
"VECTOR",
|
||||
"FLAT",
|
||||
"6", # Attribute count: TYPE, FLOAT32, DIM, dims, DISTANCE_METRIC, metric
|
||||
"TYPE",
|
||||
"FLOAT32",
|
||||
"DIM",
|
||||
str(embedding_dims),
|
||||
"DISTANCE_METRIC",
|
||||
distance_metric,
|
||||
]
|
||||
else:
|
||||
# This should never happen due to constructor validation, but be defensive
|
||||
raise ValueError(f"Unsupported index_type: {self.index_type}. Must be 'hnsw' or 'flat'")
|
||||
|
||||
# Build the complete command (comma is default separator for TAG fields)
|
||||
cmd = [
|
||||
"FT.CREATE",
|
||||
collection_name,
|
||||
"ON",
|
||||
"HASH",
|
||||
"PREFIX",
|
||||
"1",
|
||||
prefix,
|
||||
"SCHEMA",
|
||||
"memory_id",
|
||||
"TAG",
|
||||
"hash",
|
||||
"TAG",
|
||||
"agent_id",
|
||||
"TAG",
|
||||
"run_id",
|
||||
"TAG",
|
||||
"user_id",
|
||||
"TAG",
|
||||
"memory",
|
||||
"TAG",
|
||||
"metadata",
|
||||
"TAG",
|
||||
"created_at",
|
||||
"NUMERIC",
|
||||
"updated_at",
|
||||
"NUMERIC",
|
||||
] + vector_config
|
||||
|
||||
return cmd
|
||||
|
||||
def _create_index(self, embedding_model_dims):
|
||||
"""
|
||||
Create the search index with the specified schema.
|
||||
|
||||
Args:
|
||||
embedding_model_dims (int): Dimensions for the vector embeddings.
|
||||
|
||||
Raises:
|
||||
ValueError: If the search module is not available.
|
||||
Exception: For other errors during index creation.
|
||||
"""
|
||||
# Check if the search module is available
|
||||
try:
|
||||
# Try to execute a search command
|
||||
self.client.execute_command("FT._LIST")
|
||||
except ResponseError as e:
|
||||
if "unknown command" in str(e).lower():
|
||||
raise ValueError(
|
||||
"Valkey search module is not available. Please ensure Valkey is running with the search module enabled. "
|
||||
"The search module can be loaded using the --loadmodule option with the valkey-search library. "
|
||||
"For installation and setup instructions, refer to the Valkey Search documentation."
|
||||
)
|
||||
else:
|
||||
logger.exception(f"Error checking search module: {e}")
|
||||
raise
|
||||
|
||||
# Check if the index already exists
|
||||
try:
|
||||
self.client.ft(self.collection_name).info()
|
||||
return
|
||||
except ResponseError as e:
|
||||
if "not found" not in str(e).lower():
|
||||
logger.exception(f"Error checking index existence: {e}")
|
||||
raise
|
||||
|
||||
# Build and execute the index creation command
|
||||
cmd = self._build_index_schema(
|
||||
self.collection_name,
|
||||
embedding_model_dims,
|
||||
"COSINE", # Fixed distance metric for initialization
|
||||
self.prefix,
|
||||
)
|
||||
|
||||
try:
|
||||
self.client.execute_command(*cmd)
|
||||
logger.info(f"Successfully created {self.index_type.upper()} index {self.collection_name}")
|
||||
except Exception as e:
|
||||
logger.exception(f"Error creating index {self.collection_name}: {e}")
|
||||
raise
|
||||
|
||||
def create_col(self, name=None, vector_size=None, distance=None):
|
||||
"""
|
||||
Create a new collection (index) in Valkey.
|
||||
|
||||
Args:
|
||||
name (str, optional): Name for the collection. Defaults to None, which uses the current collection_name.
|
||||
vector_size (int, optional): Size of the vector embeddings. Defaults to None, which uses the current embedding_model_dims.
|
||||
distance (str, optional): Distance metric to use. Defaults to None, which uses 'cosine'.
|
||||
|
||||
Returns:
|
||||
The created index object.
|
||||
"""
|
||||
# Use provided parameters or fall back to instance attributes
|
||||
collection_name = name or self.collection_name
|
||||
embedding_dims = vector_size or self.embedding_model_dims
|
||||
distance_metric = distance or "COSINE"
|
||||
prefix = f"mem0:{collection_name}"
|
||||
|
||||
# Try to drop the index if it exists (cleanup before creation)
|
||||
self._drop_index(collection_name, log_level="silent")
|
||||
|
||||
# Build and execute the index creation command
|
||||
cmd = self._build_index_schema(
|
||||
collection_name,
|
||||
embedding_dims,
|
||||
distance_metric, # Configurable distance metric
|
||||
prefix,
|
||||
)
|
||||
|
||||
try:
|
||||
self.client.execute_command(*cmd)
|
||||
logger.info(f"Successfully created {self.index_type.upper()} index {collection_name}")
|
||||
|
||||
# Update instance attributes if creating a new collection
|
||||
if name:
|
||||
self.collection_name = collection_name
|
||||
self.prefix = prefix
|
||||
|
||||
return self.client.ft(collection_name)
|
||||
except Exception as e:
|
||||
logger.exception(f"Error creating collection {collection_name}: {e}")
|
||||
raise
|
||||
|
||||
def insert(self, vectors: list, payloads: list = None, ids: list = None):
|
||||
"""
|
||||
Insert vectors and their payloads into the index.
|
||||
|
||||
Args:
|
||||
vectors (list): List of vectors to insert.
|
||||
payloads (list, optional): List of payloads corresponding to the vectors.
|
||||
ids (list, optional): List of IDs for the vectors.
|
||||
"""
|
||||
for vector, payload, id in zip(vectors, payloads, ids):
|
||||
try:
|
||||
# Create the key for the hash
|
||||
key = f"{self.prefix}:{id}"
|
||||
|
||||
# Check for required fields and provide defaults if missing
|
||||
if "data" not in payload:
|
||||
# Silently use default value for missing 'data' field
|
||||
pass
|
||||
|
||||
# Ensure created_at is present
|
||||
if "created_at" not in payload:
|
||||
payload["created_at"] = datetime.now(pytz.timezone(self.timezone)).isoformat()
|
||||
|
||||
# Prepare the hash data
|
||||
hash_data = {
|
||||
"memory_id": id,
|
||||
"hash": payload.get("hash", f"hash_{id}"), # Use a default hash if not provided
|
||||
"memory": payload.get("data", f"data_{id}"), # Use a default data if not provided
|
||||
"created_at": int(datetime.fromisoformat(payload["created_at"]).timestamp()),
|
||||
"embedding": np.array(vector, dtype=np.float32).tobytes(),
|
||||
}
|
||||
|
||||
# Add optional fields
|
||||
for field in ["agent_id", "run_id", "user_id"]:
|
||||
if field in payload:
|
||||
hash_data[field] = payload[field]
|
||||
|
||||
# Add metadata
|
||||
hash_data["metadata"] = json.dumps({k: v for k, v in payload.items() if k not in excluded_keys})
|
||||
|
||||
# Store in Valkey
|
||||
self.client.hset(key, mapping=hash_data)
|
||||
logger.debug(f"Successfully inserted vector with ID {id}")
|
||||
except KeyError as e:
|
||||
logger.error(f"Error inserting vector with ID {id}: Missing required field {e}")
|
||||
except Exception as e:
|
||||
logger.exception(f"Error inserting vector with ID {id}: {e}")
|
||||
raise
|
||||
|
||||
def _build_search_query(self, knn_part, filters=None):
|
||||
"""
|
||||
Build a search query string with filters.
|
||||
|
||||
Args:
|
||||
knn_part (str): The KNN part of the query.
|
||||
filters (dict, optional): Filters to apply to the search. Each key-value pair
|
||||
becomes a tag filter (@key:{value}). None values are ignored.
|
||||
Values are used as-is (no validation) - wildcards, lists, etc. are
|
||||
passed through literally to Valkey search. Multiple filters are
|
||||
combined with AND logic (space-separated).
|
||||
|
||||
Returns:
|
||||
str: The complete search query string in format "filter_expr =>[KNN...]"
|
||||
or "*=>[KNN...]" if no valid filters.
|
||||
"""
|
||||
# No filters, just use the KNN search
|
||||
if not filters or not any(value is not None for key, value in filters.items()):
|
||||
return f"*=>{knn_part}"
|
||||
|
||||
# Build filter expression
|
||||
filter_parts = []
|
||||
for key, value in filters.items():
|
||||
if value is not None:
|
||||
# Use the correct filter syntax for Valkey
|
||||
filter_parts.append(f"@{key}:{{{value}}}")
|
||||
|
||||
# No valid filter parts
|
||||
if not filter_parts:
|
||||
return f"*=>{knn_part}"
|
||||
|
||||
# Combine filter parts with proper syntax
|
||||
filter_expr = " ".join(filter_parts)
|
||||
return f"{filter_expr} =>{knn_part}"
|
||||
|
||||
def _execute_search(self, query, params):
|
||||
"""
|
||||
Execute a search query.
|
||||
|
||||
Args:
|
||||
query (str): The search query to execute.
|
||||
params (dict): The query parameters.
|
||||
|
||||
Returns:
|
||||
The search results.
|
||||
"""
|
||||
try:
|
||||
return self.client.ft(self.collection_name).search(query, query_params=params)
|
||||
except ResponseError as e:
|
||||
logger.error(f"Search failed with query '{query}': {e}")
|
||||
raise
|
||||
|
||||
def _process_search_results(self, results):
|
||||
"""
|
||||
Process search results into OutputData objects.
|
||||
|
||||
Args:
|
||||
results: The search results from Valkey.
|
||||
|
||||
Returns:
|
||||
list: List of OutputData objects.
|
||||
"""
|
||||
memory_results = []
|
||||
for doc in results.docs:
|
||||
# Extract the score
|
||||
score = float(doc.vector_score) if hasattr(doc, "vector_score") else None
|
||||
|
||||
# Create the payload
|
||||
payload = {
|
||||
"hash": doc.hash,
|
||||
"data": doc.memory,
|
||||
"created_at": self._format_timestamp(int(doc.created_at), self.timezone),
|
||||
}
|
||||
|
||||
# Add updated_at if available
|
||||
if hasattr(doc, "updated_at"):
|
||||
payload["updated_at"] = self._format_timestamp(int(doc.updated_at), self.timezone)
|
||||
|
||||
# Add optional fields
|
||||
for field in ["agent_id", "run_id", "user_id"]:
|
||||
if hasattr(doc, field):
|
||||
payload[field] = getattr(doc, field)
|
||||
|
||||
# Add metadata
|
||||
if hasattr(doc, "metadata"):
|
||||
try:
|
||||
metadata = json.loads(extract_json(doc.metadata))
|
||||
payload.update(metadata)
|
||||
except (json.JSONDecodeError, TypeError) as e:
|
||||
logger.warning(f"Failed to parse metadata: {e}")
|
||||
|
||||
# Create the result
|
||||
memory_results.append(OutputData(id=doc.memory_id, score=score, payload=payload))
|
||||
|
||||
return memory_results
|
||||
|
||||
def search(self, query: str, vectors: list, limit: 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.
|
||||
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.
|
||||
|
||||
Returns:
|
||||
list: List of OutputData objects.
|
||||
"""
|
||||
# Convert the vector to bytes
|
||||
vector_bytes = np.array(vectors, dtype=np.float32).tobytes()
|
||||
|
||||
# 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]"
|
||||
else:
|
||||
# For FLAT indexes or when ef_runtime is None, use basic KNN
|
||||
knn_part = f"[KNN {limit} @embedding $vec_param AS vector_score]"
|
||||
|
||||
# Build the complete query
|
||||
q = self._build_search_query(knn_part, filters)
|
||||
|
||||
# Log the query for debugging (only in debug mode)
|
||||
logger.debug(f"Valkey search query: {q}")
|
||||
|
||||
# Set up the query parameters
|
||||
params = {"vec_param": vector_bytes}
|
||||
|
||||
# Execute the search
|
||||
results = self._execute_search(q, params)
|
||||
|
||||
# Process the results
|
||||
return self._process_search_results(results)
|
||||
|
||||
def delete(self, vector_id):
|
||||
"""
|
||||
Delete a vector from the index.
|
||||
|
||||
Args:
|
||||
vector_id (str): ID of the vector to delete.
|
||||
"""
|
||||
try:
|
||||
key = f"{self.prefix}:{vector_id}"
|
||||
self.client.delete(key)
|
||||
logger.debug(f"Successfully deleted vector with ID {vector_id}")
|
||||
except Exception as e:
|
||||
logger.exception(f"Error deleting vector with ID {vector_id}: {e}")
|
||||
raise
|
||||
|
||||
def update(self, vector_id=None, vector=None, payload=None):
|
||||
"""
|
||||
Update a vector in the index.
|
||||
|
||||
Args:
|
||||
vector_id (str): ID of the vector to update.
|
||||
vector (list, optional): New vector data.
|
||||
payload (dict, optional): New payload data.
|
||||
"""
|
||||
try:
|
||||
key = f"{self.prefix}:{vector_id}"
|
||||
|
||||
# Check for required fields and provide defaults if missing
|
||||
if "data" not in payload:
|
||||
# Silently use default value for missing 'data' field
|
||||
pass
|
||||
|
||||
# Ensure created_at is present
|
||||
if "created_at" not in payload:
|
||||
payload["created_at"] = datetime.now(pytz.timezone(self.timezone)).isoformat()
|
||||
|
||||
# Prepare the hash data
|
||||
hash_data = {
|
||||
"memory_id": vector_id,
|
||||
"hash": payload.get("hash", f"hash_{vector_id}"), # Use a default hash if not provided
|
||||
"memory": payload.get("data", f"data_{vector_id}"), # Use a default data if not provided
|
||||
"created_at": int(datetime.fromisoformat(payload["created_at"]).timestamp()),
|
||||
"embedding": np.array(vector, dtype=np.float32).tobytes(),
|
||||
}
|
||||
|
||||
# Add updated_at if available
|
||||
if "updated_at" in payload:
|
||||
hash_data["updated_at"] = int(datetime.fromisoformat(payload["updated_at"]).timestamp())
|
||||
|
||||
# Add optional fields
|
||||
for field in ["agent_id", "run_id", "user_id"]:
|
||||
if field in payload:
|
||||
hash_data[field] = payload[field]
|
||||
|
||||
# Add metadata
|
||||
hash_data["metadata"] = json.dumps({k: v for k, v in payload.items() if k not in excluded_keys})
|
||||
|
||||
# Update in Valkey
|
||||
self.client.hset(key, mapping=hash_data)
|
||||
logger.debug(f"Successfully updated vector with ID {vector_id}")
|
||||
except KeyError as e:
|
||||
logger.error(f"Error updating vector with ID {vector_id}: Missing required field {e}")
|
||||
except Exception as e:
|
||||
logger.exception(f"Error updating vector with ID {vector_id}: {e}")
|
||||
raise
|
||||
|
||||
def _format_timestamp(self, timestamp, timezone=None):
|
||||
"""
|
||||
Format a timestamp with the specified timezone.
|
||||
|
||||
Args:
|
||||
timestamp (int): The timestamp to format.
|
||||
timezone (str, optional): The timezone to use. Defaults to UTC.
|
||||
|
||||
Returns:
|
||||
str: The formatted timestamp.
|
||||
"""
|
||||
# Use UTC as default timezone if not specified
|
||||
tz = pytz.timezone(timezone or "UTC")
|
||||
return datetime.fromtimestamp(timestamp, tz=tz).isoformat(timespec="microseconds")
|
||||
|
||||
def _process_document_fields(self, result, vector_id):
|
||||
"""
|
||||
Process document fields from a Valkey hash result.
|
||||
|
||||
Args:
|
||||
result (dict): The hash result from Valkey.
|
||||
vector_id (str): The vector ID.
|
||||
|
||||
Returns:
|
||||
dict: The processed payload.
|
||||
str: The memory ID.
|
||||
"""
|
||||
# Create the payload with error handling
|
||||
payload = {}
|
||||
|
||||
# Convert bytes to string for text fields
|
||||
for k in result:
|
||||
if k not in ["embedding"]:
|
||||
if isinstance(result[k], bytes):
|
||||
try:
|
||||
result[k] = result[k].decode("utf-8")
|
||||
except UnicodeDecodeError:
|
||||
# If decoding fails, keep the bytes
|
||||
pass
|
||||
|
||||
# Add required fields with error handling
|
||||
for field in ["hash", "memory", "created_at"]:
|
||||
if field in result:
|
||||
if field == "created_at":
|
||||
try:
|
||||
payload[field] = self._format_timestamp(int(result[field]), self.timezone)
|
||||
except (ValueError, TypeError):
|
||||
payload[field] = result[field]
|
||||
else:
|
||||
payload[field] = result[field]
|
||||
else:
|
||||
# Use default values for missing fields
|
||||
if field == "hash":
|
||||
payload[field] = "unknown"
|
||||
elif field == "memory":
|
||||
payload[field] = "unknown"
|
||||
elif field == "created_at":
|
||||
payload[field] = self._format_timestamp(
|
||||
int(datetime.now(tz=pytz.timezone(self.timezone)).timestamp()), self.timezone
|
||||
)
|
||||
|
||||
# Rename memory to data for consistency
|
||||
if "memory" in payload:
|
||||
payload["data"] = payload.pop("memory")
|
||||
|
||||
# Add updated_at if available
|
||||
if "updated_at" in result:
|
||||
try:
|
||||
payload["updated_at"] = self._format_timestamp(int(result["updated_at"]), self.timezone)
|
||||
except (ValueError, TypeError):
|
||||
payload["updated_at"] = result["updated_at"]
|
||||
|
||||
# Add optional fields
|
||||
for field in ["agent_id", "run_id", "user_id"]:
|
||||
if field in result:
|
||||
payload[field] = result[field]
|
||||
|
||||
# Add metadata
|
||||
if "metadata" in result:
|
||||
try:
|
||||
metadata = json.loads(extract_json(result["metadata"]))
|
||||
payload.update(metadata)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
logger.warning(f"Failed to parse metadata: {result.get('metadata')}")
|
||||
|
||||
# Use memory_id from result if available, otherwise use vector_id
|
||||
memory_id = result.get("memory_id", vector_id)
|
||||
|
||||
return payload, memory_id
|
||||
|
||||
def _convert_bytes(self, data):
|
||||
"""Convert bytes data back to string"""
|
||||
if isinstance(data, bytes):
|
||||
try:
|
||||
return data.decode("utf-8")
|
||||
except UnicodeDecodeError:
|
||||
return data
|
||||
if isinstance(data, dict):
|
||||
return {self._convert_bytes(key): self._convert_bytes(value) for key, value in data.items()}
|
||||
if isinstance(data, list):
|
||||
return [self._convert_bytes(item) for item in data]
|
||||
if isinstance(data, tuple):
|
||||
return tuple(self._convert_bytes(item) for item in data)
|
||||
return data
|
||||
|
||||
def get(self, vector_id):
|
||||
"""
|
||||
Get a vector by ID.
|
||||
|
||||
Args:
|
||||
vector_id (str): ID of the vector to get.
|
||||
|
||||
Returns:
|
||||
OutputData: The retrieved vector.
|
||||
"""
|
||||
try:
|
||||
key = f"{self.prefix}:{vector_id}"
|
||||
result = self.client.hgetall(key)
|
||||
|
||||
if not result:
|
||||
raise KeyError(f"Vector with ID {vector_id} not found")
|
||||
|
||||
# Convert bytes keys/values to strings
|
||||
result = self._convert_bytes(result)
|
||||
|
||||
logger.debug(f"Retrieved result keys: {result.keys()}")
|
||||
|
||||
# Process the document fields
|
||||
payload, memory_id = self._process_document_fields(result, vector_id)
|
||||
|
||||
return OutputData(id=memory_id, payload=payload, score=0.0)
|
||||
except KeyError:
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.exception(f"Error getting vector with ID {vector_id}: {e}")
|
||||
raise
|
||||
|
||||
def list_cols(self):
|
||||
"""
|
||||
List all collections (indices) in Valkey.
|
||||
|
||||
Returns:
|
||||
list: List of collection names.
|
||||
"""
|
||||
try:
|
||||
# Use the FT._LIST command to list all indices
|
||||
return self.client.execute_command("FT._LIST")
|
||||
except Exception as e:
|
||||
logger.exception(f"Error listing collections: {e}")
|
||||
raise
|
||||
|
||||
def _drop_index(self, collection_name, log_level="error"):
|
||||
"""
|
||||
Drop an index by name using the documented FT.DROPINDEX command.
|
||||
|
||||
Args:
|
||||
collection_name (str): Name of the index to drop.
|
||||
log_level (str): Logging level for missing index ("silent", "info", "error").
|
||||
"""
|
||||
try:
|
||||
self.client.execute_command("FT.DROPINDEX", collection_name)
|
||||
logger.info(f"Successfully deleted index {collection_name}")
|
||||
return True
|
||||
except ResponseError as e:
|
||||
if "Unknown index name" in str(e):
|
||||
# Index doesn't exist - handle based on context
|
||||
if log_level == "silent":
|
||||
pass # No logging in situations where this is expected such as initial index creation
|
||||
elif log_level == "info":
|
||||
logger.info(f"Index {collection_name} doesn't exist, skipping deletion")
|
||||
return False
|
||||
else:
|
||||
# Real error - always log and raise
|
||||
logger.error(f"Error deleting index {collection_name}: {e}")
|
||||
raise
|
||||
except Exception as e:
|
||||
# Non-ResponseError exceptions - always log and raise
|
||||
logger.error(f"Error deleting index {collection_name}: {e}")
|
||||
raise
|
||||
|
||||
def delete_col(self):
|
||||
"""
|
||||
Delete the current collection (index).
|
||||
"""
|
||||
return self._drop_index(self.collection_name, log_level="info")
|
||||
|
||||
def col_info(self, name=None):
|
||||
"""
|
||||
Get information about a collection (index).
|
||||
|
||||
Args:
|
||||
name (str, optional): Name of the collection. Defaults to None, which uses the current collection_name.
|
||||
|
||||
Returns:
|
||||
dict: Information about the collection.
|
||||
"""
|
||||
try:
|
||||
collection_name = name or self.collection_name
|
||||
return self.client.ft(collection_name).info()
|
||||
except Exception as e:
|
||||
logger.exception(f"Error getting collection info for {collection_name}: {e}")
|
||||
raise
|
||||
|
||||
def reset(self):
|
||||
"""
|
||||
Reset the index by deleting and recreating it.
|
||||
"""
|
||||
try:
|
||||
collection_name = self.collection_name
|
||||
logger.warning(f"Resetting index {collection_name}...")
|
||||
|
||||
# Delete the index
|
||||
self.delete_col()
|
||||
|
||||
# Recreate the index
|
||||
self._create_index(self.embedding_model_dims)
|
||||
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.exception(f"Error resetting index {self.collection_name}: {e}")
|
||||
raise
|
||||
|
||||
def _build_list_query(self, filters=None):
|
||||
"""
|
||||
Build a query for listing vectors.
|
||||
|
||||
Args:
|
||||
filters (dict, optional): Filters to apply to the list. Each key-value pair
|
||||
becomes a tag filter (@key:{value}). None values are ignored.
|
||||
Values are used as-is (no validation) - wildcards, lists, etc. are
|
||||
passed through literally to Valkey search.
|
||||
|
||||
Returns:
|
||||
str: The query string. Returns "*" if no valid filters provided.
|
||||
"""
|
||||
# Default query
|
||||
q = "*"
|
||||
|
||||
# Add filters if provided
|
||||
if filters and any(value is not None for key, value in filters.items()):
|
||||
filter_conditions = []
|
||||
for key, value in filters.items():
|
||||
if value is not None:
|
||||
filter_conditions.append(f"@{key}:{{{value}}}")
|
||||
|
||||
if filter_conditions:
|
||||
q = " ".join(filter_conditions)
|
||||
|
||||
return q
|
||||
|
||||
def list(self, filters: dict = None, limit: int = None) -> list:
|
||||
"""
|
||||
List all recent created memories from the vector store.
|
||||
|
||||
Args:
|
||||
filters (dict, optional): Filters to apply to the list. Each key-value pair
|
||||
becomes a tag filter (@key:{value}). None values are ignored.
|
||||
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
|
||||
if not specified.
|
||||
|
||||
Returns:
|
||||
list: Nested list format [[MemoryResult(), ...]] matching Redis implementation.
|
||||
Each MemoryResult contains id and payload with hash, data, timestamps, etc.
|
||||
"""
|
||||
try:
|
||||
# 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
|
||||
|
||||
# Use the existing search method which handles filters properly
|
||||
search_results = self.search("", dummy_vector, limit=search_limit, filters=filters)
|
||||
|
||||
# Convert search results to list format (match Redis format)
|
||||
class MemoryResult:
|
||||
def __init__(self, id: str, payload: dict, score: float = None):
|
||||
self.id = id
|
||||
self.payload = payload
|
||||
self.score = score
|
||||
|
||||
memory_results = []
|
||||
for result in search_results:
|
||||
# Create payload in the expected format
|
||||
payload = {
|
||||
"hash": result.payload.get("hash", ""),
|
||||
"data": result.payload.get("data", ""),
|
||||
"created_at": result.payload.get("created_at"),
|
||||
"updated_at": result.payload.get("updated_at"),
|
||||
}
|
||||
|
||||
# Add metadata (exclude system fields)
|
||||
for key, value in result.payload.items():
|
||||
if key not in ["data", "hash", "created_at", "updated_at"]:
|
||||
payload[key] = value
|
||||
|
||||
# Create MemoryResult object (matching Redis format)
|
||||
memory_results.append(MemoryResult(id=result.id, payload=payload))
|
||||
|
||||
# Return nested list format like Redis
|
||||
return [memory_results]
|
||||
|
||||
except Exception as e:
|
||||
logger.exception(f"Error in list method: {e}")
|
||||
return [[]] # Return empty result on error
|
||||
@@ -2,7 +2,7 @@ from datetime import datetime
|
||||
from typing import List, Optional
|
||||
from uuid import UUID
|
||||
|
||||
from pydantic import BaseModel, Field, validator
|
||||
from pydantic import BaseModel, ConfigDict, Field, validator
|
||||
|
||||
|
||||
class MemoryBase(BaseModel):
|
||||
@@ -33,8 +33,7 @@ class Memory(MemoryBase):
|
||||
categories: Optional[List[Category]] = None
|
||||
app: App
|
||||
|
||||
class Config:
|
||||
from_attributes = True
|
||||
model_config = ConfigDict(from_attributes=True)
|
||||
|
||||
class MemoryUpdate(BaseModel):
|
||||
content: Optional[str] = None
|
||||
|
||||
+3
-2
@@ -14,7 +14,7 @@ requires-python = ">=3.9,<4.0"
|
||||
dependencies = [
|
||||
"qdrant-client>=1.9.1",
|
||||
"pydantic>=2.7.3",
|
||||
"openai>=1.90.0,<1.100.0",
|
||||
"openai>=1.90.0,<1.110.0",
|
||||
"posthog>=3.5.0",
|
||||
"pytz>=2024.1",
|
||||
"sqlalchemy>=2.0.31",
|
||||
@@ -42,6 +42,7 @@ vector_stores = [
|
||||
"psycopg-pool>=3.2.6,<4.0.0",
|
||||
"pymongo>=4.13.2",
|
||||
"pymochow>=2.2.9",
|
||||
"valkey>=6.0.0",
|
||||
"databricks-sdk>=0.63.0",
|
||||
"azure-identity>=1.24.0",
|
||||
"redis>=5.0.0,<6.0.0",
|
||||
@@ -53,7 +54,7 @@ llms = [
|
||||
"groq>=0.3.0",
|
||||
"together>=0.2.10",
|
||||
"litellm>=1.74.0",
|
||||
"openai>=1.90.0,<1.100.0",
|
||||
"openai>=1.90.0,<1.110.0",
|
||||
"ollama>=0.1.0",
|
||||
"vertexai>=0.1.0",
|
||||
"google-generativeai>=0.3.0",
|
||||
|
||||
@@ -55,7 +55,7 @@ def test_generate_response_without_tools(mock_openai_client):
|
||||
response = llm.generate_response(messages)
|
||||
|
||||
mock_openai_client.chat.completions.create.assert_called_once_with(
|
||||
model="gpt-4o", messages=messages, temperature=0.7, max_tokens=100, top_p=1.0
|
||||
model="gpt-4o", messages=messages, temperature=0.7, max_tokens=100, top_p=1.0, store=False
|
||||
)
|
||||
assert response == "I'm doing well, thank you for asking!"
|
||||
|
||||
@@ -97,7 +97,7 @@ def test_generate_response_with_tools(mock_openai_client):
|
||||
response = llm.generate_response(messages, tools=tools)
|
||||
|
||||
mock_openai_client.chat.completions.create.assert_called_once_with(
|
||||
model="gpt-4o", messages=messages, temperature=0.7, max_tokens=100, top_p=1.0, tools=tools, tool_choice="auto"
|
||||
model="gpt-4o", messages=messages, temperature=0.7, max_tokens=100, top_p=1.0, tools=tools, tool_choice="auto", store=False
|
||||
)
|
||||
|
||||
assert response["content"] == "I've added the memory for you."
|
||||
|
||||
@@ -66,7 +66,7 @@ class TestAddToVectorStoreErrors:
|
||||
mock_memory.llm.generate_response.side_effect = ['{"facts": ["test fact"]}', ""]
|
||||
|
||||
# Execute
|
||||
with caplog.at_level(logging.ERROR):
|
||||
with caplog.at_level(logging.WARNING):
|
||||
result = mock_memory._add_to_vector_store(
|
||||
messages=[{"role": "user", "content": "test"}], metadata={}, filters={}, infer=True
|
||||
)
|
||||
@@ -74,7 +74,7 @@ class TestAddToVectorStoreErrors:
|
||||
# Verify
|
||||
assert mock_memory.llm.generate_response.call_count == 2
|
||||
assert result == [] # Should return empty list when no memories processed
|
||||
assert "Invalid JSON response" in caplog.text
|
||||
assert "Empty response from LLM, no memories to extract" in caplog.text
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@@ -117,11 +117,11 @@ class TestAsyncAddToVectorStoreErrors:
|
||||
mock_capture_event = mocker.MagicMock()
|
||||
mocker.patch("mem0.memory.main.capture_event", mock_capture_event)
|
||||
|
||||
with caplog.at_level(logging.ERROR):
|
||||
with caplog.at_level(logging.WARNING):
|
||||
result = await mock_async_memory._add_to_vector_store(
|
||||
messages=[{"role": "user", "content": "test"}], metadata={}, effective_filters={}, infer=True
|
||||
)
|
||||
|
||||
assert result == []
|
||||
assert "Invalid JSON response" in caplog.text
|
||||
assert "Empty response from LLM, no memories to extract" in caplog.text
|
||||
assert mock_capture_event.call_count == 1
|
||||
|
||||
@@ -300,8 +300,9 @@ def test_list_memories(db_instance_delta, mock_workspace_client):
|
||||
result=SimpleNamespace(data_array=[row])
|
||||
)
|
||||
res = db_instance_delta.list(limit=1)
|
||||
assert len(res) == 1
|
||||
assert res[0].id == "id3"
|
||||
assert isinstance(res, list)
|
||||
assert len(res[0]) == 1
|
||||
assert res[0][0].id == "id3"
|
||||
|
||||
|
||||
# ---------------------- Reset Tests ---------------------- #
|
||||
|
||||
@@ -0,0 +1,862 @@
|
||||
import json
|
||||
from datetime import datetime
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
import pytz
|
||||
from valkey.exceptions import ResponseError
|
||||
|
||||
from mem0.vector_stores.valkey import ValkeyDB
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_valkey_client():
|
||||
"""Create a mock Valkey client."""
|
||||
with patch("valkey.from_url") as mock_client:
|
||||
# Mock the ft method
|
||||
mock_ft = MagicMock()
|
||||
mock_client.return_value.ft = MagicMock(return_value=mock_ft)
|
||||
mock_client.return_value.execute_command = MagicMock()
|
||||
mock_client.return_value.hset = MagicMock()
|
||||
mock_client.return_value.hgetall = MagicMock()
|
||||
mock_client.return_value.delete = MagicMock()
|
||||
yield mock_client.return_value
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def valkey_db(mock_valkey_client):
|
||||
"""Create a ValkeyDB instance with a mock client."""
|
||||
# Initialize the ValkeyDB with test parameters
|
||||
valkey_db = ValkeyDB(
|
||||
valkey_url="valkey://localhost:6379",
|
||||
collection_name="test_collection",
|
||||
embedding_model_dims=1536,
|
||||
)
|
||||
# Replace the client with our mock
|
||||
valkey_db.client = mock_valkey_client
|
||||
return valkey_db
|
||||
|
||||
|
||||
def test_search_filter_syntax(valkey_db, mock_valkey_client):
|
||||
"""Test that the search filter syntax is correctly formatted for Valkey."""
|
||||
# Mock search results
|
||||
mock_doc = MagicMock()
|
||||
mock_doc.memory_id = "test_id"
|
||||
mock_doc.hash = "test_hash"
|
||||
mock_doc.memory = "test_data"
|
||||
mock_doc.created_at = str(int(datetime.now().timestamp()))
|
||||
mock_doc.metadata = json.dumps({"key": "value"})
|
||||
mock_doc.vector_score = "0.5"
|
||||
|
||||
mock_results = MagicMock()
|
||||
mock_results.docs = [mock_doc]
|
||||
|
||||
mock_ft = mock_valkey_client.ft.return_value
|
||||
mock_ft.search.return_value = mock_results
|
||||
|
||||
# Test with user_id filter
|
||||
valkey_db.search(
|
||||
query="test query",
|
||||
vectors=np.random.rand(1536).tolist(),
|
||||
limit=5,
|
||||
filters={"user_id": "test_user"},
|
||||
)
|
||||
|
||||
# Check that the search was called with the correct filter syntax
|
||||
args, kwargs = mock_ft.search.call_args
|
||||
assert "@user_id:{test_user}" in args[0]
|
||||
assert "=>[KNN" in args[0]
|
||||
|
||||
# Test with multiple filters
|
||||
valkey_db.search(
|
||||
query="test query",
|
||||
vectors=np.random.rand(1536).tolist(),
|
||||
limit=5,
|
||||
filters={"user_id": "test_user", "agent_id": "test_agent"},
|
||||
)
|
||||
|
||||
# Check that the search was called with the correct filter syntax
|
||||
args, kwargs = mock_ft.search.call_args
|
||||
assert "@user_id:{test_user}" in args[0]
|
||||
assert "@agent_id:{test_agent}" in args[0]
|
||||
assert "=>[KNN" in args[0]
|
||||
|
||||
|
||||
def test_search_without_filters(valkey_db, mock_valkey_client):
|
||||
"""Test search without filters."""
|
||||
# Mock search results
|
||||
mock_doc = MagicMock()
|
||||
mock_doc.memory_id = "test_id"
|
||||
mock_doc.hash = "test_hash"
|
||||
mock_doc.memory = "test_data"
|
||||
mock_doc.created_at = str(int(datetime.now().timestamp()))
|
||||
mock_doc.metadata = json.dumps({"key": "value"})
|
||||
mock_doc.vector_score = "0.5"
|
||||
|
||||
mock_results = MagicMock()
|
||||
mock_results.docs = [mock_doc]
|
||||
|
||||
mock_ft = mock_valkey_client.ft.return_value
|
||||
mock_ft.search.return_value = mock_results
|
||||
|
||||
# Test without filters
|
||||
results = valkey_db.search(
|
||||
query="test query",
|
||||
vectors=np.random.rand(1536).tolist(),
|
||||
limit=5,
|
||||
)
|
||||
|
||||
# Check that the search was called with the correct syntax
|
||||
args, kwargs = mock_ft.search.call_args
|
||||
assert "*=>[KNN" in args[0]
|
||||
|
||||
# Check that results are processed correctly
|
||||
assert len(results) == 1
|
||||
assert results[0].id == "test_id"
|
||||
assert results[0].payload["hash"] == "test_hash"
|
||||
assert results[0].payload["data"] == "test_data"
|
||||
assert "created_at" in results[0].payload
|
||||
|
||||
|
||||
def test_insert(valkey_db, mock_valkey_client):
|
||||
"""Test inserting vectors."""
|
||||
# Prepare test data
|
||||
vectors = [np.random.rand(1536).tolist()]
|
||||
payloads = [{"hash": "test_hash", "data": "test_data", "user_id": "test_user"}]
|
||||
ids = ["test_id"]
|
||||
|
||||
# Call insert
|
||||
valkey_db.insert(vectors=vectors, payloads=payloads, ids=ids)
|
||||
|
||||
# Check that hset was called with the correct arguments
|
||||
mock_valkey_client.hset.assert_called_once()
|
||||
args, kwargs = mock_valkey_client.hset.call_args
|
||||
assert args[0] == "mem0:test_collection:test_id"
|
||||
assert "memory_id" in kwargs["mapping"]
|
||||
assert kwargs["mapping"]["memory_id"] == "test_id"
|
||||
assert kwargs["mapping"]["hash"] == "test_hash"
|
||||
assert kwargs["mapping"]["memory"] == "test_data"
|
||||
assert kwargs["mapping"]["user_id"] == "test_user"
|
||||
assert "created_at" in kwargs["mapping"]
|
||||
assert "embedding" in kwargs["mapping"]
|
||||
|
||||
|
||||
def test_insert_handles_missing_created_at(valkey_db, mock_valkey_client):
|
||||
"""Test inserting vectors with missing created_at field."""
|
||||
# Prepare test data
|
||||
vectors = [np.random.rand(1536).tolist()]
|
||||
payloads = [{"hash": "test_hash", "data": "test_data"}] # No created_at
|
||||
ids = ["test_id"]
|
||||
|
||||
# Call insert
|
||||
valkey_db.insert(vectors=vectors, payloads=payloads, ids=ids)
|
||||
|
||||
# Check that hset was called with the correct arguments
|
||||
mock_valkey_client.hset.assert_called_once()
|
||||
args, kwargs = mock_valkey_client.hset.call_args
|
||||
assert "created_at" in kwargs["mapping"] # Should be added automatically
|
||||
|
||||
|
||||
def test_delete(valkey_db, mock_valkey_client):
|
||||
"""Test deleting a vector."""
|
||||
# Call delete
|
||||
valkey_db.delete("test_id")
|
||||
|
||||
# Check that delete was called with the correct key
|
||||
mock_valkey_client.delete.assert_called_once_with("mem0:test_collection:test_id")
|
||||
|
||||
|
||||
def test_update(valkey_db, mock_valkey_client):
|
||||
"""Test updating a vector."""
|
||||
# Prepare test data
|
||||
vector = np.random.rand(1536).tolist()
|
||||
payload = {
|
||||
"hash": "test_hash",
|
||||
"data": "updated_data",
|
||||
"created_at": datetime.now(pytz.timezone("UTC")).isoformat(),
|
||||
"user_id": "test_user",
|
||||
}
|
||||
|
||||
# Call update
|
||||
valkey_db.update(vector_id="test_id", vector=vector, payload=payload)
|
||||
|
||||
# Check that hset was called with the correct arguments
|
||||
mock_valkey_client.hset.assert_called_once()
|
||||
args, kwargs = mock_valkey_client.hset.call_args
|
||||
assert args[0] == "mem0:test_collection:test_id"
|
||||
assert kwargs["mapping"]["memory_id"] == "test_id"
|
||||
assert kwargs["mapping"]["memory"] == "updated_data"
|
||||
|
||||
|
||||
def test_update_handles_missing_created_at(valkey_db, mock_valkey_client):
|
||||
"""Test updating vectors with missing created_at field."""
|
||||
# Prepare test data
|
||||
vector = np.random.rand(1536).tolist()
|
||||
payload = {"hash": "test_hash", "data": "updated_data"} # No created_at
|
||||
|
||||
# Call update
|
||||
valkey_db.update(vector_id="test_id", vector=vector, payload=payload)
|
||||
|
||||
# Check that hset was called with the correct arguments
|
||||
mock_valkey_client.hset.assert_called_once()
|
||||
args, kwargs = mock_valkey_client.hset.call_args
|
||||
assert "created_at" in kwargs["mapping"] # Should be added automatically
|
||||
|
||||
|
||||
def test_get(valkey_db, mock_valkey_client):
|
||||
"""Test getting a vector."""
|
||||
# Mock hgetall to return a vector
|
||||
mock_valkey_client.hgetall.return_value = {
|
||||
"memory_id": "test_id",
|
||||
"hash": "test_hash",
|
||||
"memory": "test_data",
|
||||
"created_at": str(int(datetime.now().timestamp())),
|
||||
"metadata": json.dumps({"key": "value"}),
|
||||
"user_id": "test_user",
|
||||
}
|
||||
|
||||
# Call get
|
||||
result = valkey_db.get("test_id")
|
||||
|
||||
# Check that hgetall was called with the correct key
|
||||
mock_valkey_client.hgetall.assert_called_once_with("mem0:test_collection:test_id")
|
||||
|
||||
# Check the result
|
||||
assert result.id == "test_id"
|
||||
assert result.payload["hash"] == "test_hash"
|
||||
assert result.payload["data"] == "test_data"
|
||||
assert "created_at" in result.payload
|
||||
assert result.payload["key"] == "value" # From metadata
|
||||
assert result.payload["user_id"] == "test_user"
|
||||
|
||||
|
||||
def test_get_not_found(valkey_db, mock_valkey_client):
|
||||
"""Test getting a vector that doesn't exist."""
|
||||
# Mock hgetall to return empty dict (not found)
|
||||
mock_valkey_client.hgetall.return_value = {}
|
||||
|
||||
# Call get should raise KeyError
|
||||
with pytest.raises(KeyError, match="Vector with ID test_id not found"):
|
||||
valkey_db.get("test_id")
|
||||
|
||||
|
||||
def test_list_cols(valkey_db, mock_valkey_client):
|
||||
"""Test listing collections."""
|
||||
# Reset the mock to clear previous calls
|
||||
mock_valkey_client.execute_command.reset_mock()
|
||||
|
||||
# Mock execute_command to return list of indices
|
||||
mock_valkey_client.execute_command.return_value = ["test_collection", "another_collection"]
|
||||
|
||||
# Call list_cols
|
||||
result = valkey_db.list_cols()
|
||||
|
||||
# Check that execute_command was called with the correct command
|
||||
mock_valkey_client.execute_command.assert_called_with("FT._LIST")
|
||||
|
||||
# Check the result
|
||||
assert result == ["test_collection", "another_collection"]
|
||||
|
||||
|
||||
def test_delete_col(valkey_db, mock_valkey_client):
|
||||
"""Test deleting a collection."""
|
||||
# Reset the mock to clear previous calls
|
||||
mock_valkey_client.execute_command.reset_mock()
|
||||
|
||||
# Test successful deletion
|
||||
result = valkey_db.delete_col()
|
||||
assert result is True
|
||||
|
||||
# Check that execute_command was called with the correct command
|
||||
mock_valkey_client.execute_command.assert_called_once_with("FT.DROPINDEX", "test_collection")
|
||||
|
||||
# Test error handling - real errors should still raise
|
||||
mock_valkey_client.execute_command.side_effect = ResponseError("Error dropping index")
|
||||
with pytest.raises(ResponseError, match="Error dropping index"):
|
||||
valkey_db.delete_col()
|
||||
|
||||
# Test idempotent behavior - "Unknown index name" should return False, not raise
|
||||
mock_valkey_client.execute_command.side_effect = ResponseError("Unknown index name")
|
||||
result = valkey_db.delete_col()
|
||||
assert result is False
|
||||
|
||||
|
||||
def test_context_aware_logging(valkey_db, mock_valkey_client):
|
||||
"""Test that _drop_index handles different log levels correctly."""
|
||||
# Mock "Unknown index name" error
|
||||
mock_valkey_client.execute_command.side_effect = ResponseError("Unknown index name")
|
||||
|
||||
# Test silent mode - should not log anything (we can't easily test log output, but ensure no exception)
|
||||
result = valkey_db._drop_index("test_collection", log_level="silent")
|
||||
assert result is False
|
||||
|
||||
# Test info mode - should not raise exception
|
||||
result = valkey_db._drop_index("test_collection", log_level="info")
|
||||
assert result is False
|
||||
|
||||
# Test default mode - should not raise exception
|
||||
result = valkey_db._drop_index("test_collection")
|
||||
assert result is False
|
||||
|
||||
|
||||
def test_col_info(valkey_db, mock_valkey_client):
|
||||
"""Test getting collection info."""
|
||||
# Mock ft().info() to return index info
|
||||
mock_ft = mock_valkey_client.ft.return_value
|
||||
|
||||
# Reset the mock to clear previous calls
|
||||
mock_ft.info.reset_mock()
|
||||
|
||||
mock_ft.info.return_value = {"index_name": "test_collection", "num_docs": 100}
|
||||
|
||||
# Call col_info
|
||||
result = valkey_db.col_info()
|
||||
|
||||
# Check that ft().info() was called
|
||||
assert mock_ft.info.called
|
||||
|
||||
# Check the result
|
||||
assert result["index_name"] == "test_collection"
|
||||
assert result["num_docs"] == 100
|
||||
|
||||
|
||||
def test_create_col(valkey_db, mock_valkey_client):
|
||||
"""Test creating a new collection."""
|
||||
# Call create_col
|
||||
valkey_db.create_col(name="new_collection", vector_size=768, distance="IP")
|
||||
|
||||
# Check that execute_command was called to create the index
|
||||
assert mock_valkey_client.execute_command.called
|
||||
args = mock_valkey_client.execute_command.call_args[0]
|
||||
assert args[0] == "FT.CREATE"
|
||||
assert args[1] == "new_collection"
|
||||
|
||||
# Check that the distance metric was set correctly
|
||||
distance_metric_index = args.index("DISTANCE_METRIC")
|
||||
assert args[distance_metric_index + 1] == "IP"
|
||||
|
||||
# Check that the vector size was set correctly
|
||||
dim_index = args.index("DIM")
|
||||
assert args[dim_index + 1] == "768"
|
||||
|
||||
|
||||
def test_list(valkey_db, mock_valkey_client):
|
||||
"""Test listing vectors."""
|
||||
# Mock search results
|
||||
mock_doc = MagicMock()
|
||||
mock_doc.memory_id = "test_id"
|
||||
mock_doc.hash = "test_hash"
|
||||
mock_doc.memory = "test_data"
|
||||
mock_doc.created_at = str(int(datetime.now().timestamp()))
|
||||
mock_doc.metadata = json.dumps({"key": "value"})
|
||||
mock_doc.vector_score = "0.5" # Add missing vector_score
|
||||
|
||||
mock_results = MagicMock()
|
||||
mock_results.docs = [mock_doc]
|
||||
|
||||
mock_ft = mock_valkey_client.ft.return_value
|
||||
mock_ft.search.return_value = mock_results
|
||||
|
||||
# Call list
|
||||
results = valkey_db.list(filters={"user_id": "test_user"}, limit=10)
|
||||
|
||||
# Check that search was called with the correct arguments
|
||||
mock_ft.search.assert_called_once()
|
||||
args, kwargs = mock_ft.search.call_args
|
||||
# Now expects full search query with KNN part due to dummy vector approach
|
||||
assert "@user_id:{test_user}" in args[0]
|
||||
assert "=>[KNN" in args[0]
|
||||
# Verify the results format
|
||||
assert len(results) == 1
|
||||
assert len(results[0]) == 1
|
||||
assert results[0][0].id == "test_id"
|
||||
|
||||
# Check the results
|
||||
assert len(results) == 1 # One list of results
|
||||
assert len(results[0]) == 1 # One result in the list
|
||||
assert results[0][0].id == "test_id"
|
||||
assert results[0][0].payload["hash"] == "test_hash"
|
||||
assert results[0][0].payload["data"] == "test_data"
|
||||
|
||||
|
||||
def test_search_error_handling(valkey_db, mock_valkey_client):
|
||||
"""Test search error handling when query fails."""
|
||||
# Mock search to fail with an error
|
||||
mock_ft = mock_valkey_client.ft.return_value
|
||||
mock_ft.search.side_effect = ResponseError("Invalid filter expression")
|
||||
|
||||
# Call search should raise the error
|
||||
with pytest.raises(ResponseError, match="Invalid filter expression"):
|
||||
valkey_db.search(
|
||||
query="test query",
|
||||
vectors=np.random.rand(1536).tolist(),
|
||||
limit=5,
|
||||
filters={"user_id": "test_user"},
|
||||
)
|
||||
|
||||
# Check that search was called once
|
||||
assert mock_ft.search.call_count == 1
|
||||
|
||||
|
||||
def test_drop_index_error_handling(valkey_db, mock_valkey_client):
|
||||
"""Test error handling when dropping an index."""
|
||||
# Reset the mock to clear previous calls
|
||||
mock_valkey_client.execute_command.reset_mock()
|
||||
|
||||
# Test 1: Real error (not "Unknown index name") should raise
|
||||
mock_valkey_client.execute_command.side_effect = ResponseError("Error dropping index")
|
||||
with pytest.raises(ResponseError, match="Error dropping index"):
|
||||
valkey_db._drop_index("test_collection")
|
||||
|
||||
# Test 2: "Unknown index name" with default log_level should return False
|
||||
mock_valkey_client.execute_command.side_effect = ResponseError("Unknown index name")
|
||||
result = valkey_db._drop_index("test_collection")
|
||||
assert result is False
|
||||
|
||||
# Test 3: "Unknown index name" with silent log_level should return False
|
||||
mock_valkey_client.execute_command.side_effect = ResponseError("Unknown index name")
|
||||
result = valkey_db._drop_index("test_collection", log_level="silent")
|
||||
assert result is False
|
||||
|
||||
# Test 4: "Unknown index name" with info log_level should return False
|
||||
mock_valkey_client.execute_command.side_effect = ResponseError("Unknown index name")
|
||||
result = valkey_db._drop_index("test_collection", log_level="info")
|
||||
assert result is False
|
||||
|
||||
# Test 5: Successful deletion should return True
|
||||
mock_valkey_client.execute_command.side_effect = None # Reset to success
|
||||
result = valkey_db._drop_index("test_collection")
|
||||
assert result is True
|
||||
|
||||
|
||||
def test_reset(valkey_db, mock_valkey_client):
|
||||
"""Test resetting an index."""
|
||||
# Mock delete_col and _create_index
|
||||
with (
|
||||
patch.object(valkey_db, "delete_col", return_value=True) as mock_delete_col,
|
||||
patch.object(valkey_db, "_create_index") as mock_create_index,
|
||||
):
|
||||
# Call reset
|
||||
result = valkey_db.reset()
|
||||
|
||||
# Check that delete_col and _create_index were called
|
||||
mock_delete_col.assert_called_once()
|
||||
mock_create_index.assert_called_once_with(1536)
|
||||
|
||||
# Check the result
|
||||
assert result is True
|
||||
|
||||
|
||||
def test_build_list_query(valkey_db):
|
||||
"""Test building a list query with and without filters."""
|
||||
# Test without filters
|
||||
query = valkey_db._build_list_query(None)
|
||||
assert query == "*"
|
||||
|
||||
# Test with empty filters
|
||||
query = valkey_db._build_list_query({})
|
||||
assert query == "*"
|
||||
|
||||
# Test with filters
|
||||
query = valkey_db._build_list_query({"user_id": "test_user"})
|
||||
assert query == "@user_id:{test_user}"
|
||||
|
||||
# Test with multiple filters
|
||||
query = valkey_db._build_list_query({"user_id": "test_user", "agent_id": "test_agent"})
|
||||
assert "@user_id:{test_user}" in query
|
||||
assert "@agent_id:{test_agent}" in query
|
||||
|
||||
|
||||
def test_process_document_fields(valkey_db):
|
||||
"""Test processing document fields from hash results."""
|
||||
# Create a mock result with all fields
|
||||
result = {
|
||||
"memory_id": "test_id",
|
||||
"hash": "test_hash",
|
||||
"memory": "test_data",
|
||||
"created_at": "1625097600", # 2021-07-01 00:00:00 UTC
|
||||
"updated_at": "1625184000", # 2021-07-02 00:00:00 UTC
|
||||
"user_id": "test_user",
|
||||
"agent_id": "test_agent",
|
||||
"metadata": json.dumps({"key": "value"}),
|
||||
}
|
||||
|
||||
# Process the document fields
|
||||
payload, memory_id = valkey_db._process_document_fields(result, "default_id")
|
||||
|
||||
# Check the results
|
||||
assert memory_id == "test_id"
|
||||
assert payload["hash"] == "test_hash"
|
||||
assert payload["data"] == "test_data" # memory renamed to data
|
||||
assert "created_at" in payload
|
||||
assert "updated_at" in payload
|
||||
assert payload["user_id"] == "test_user"
|
||||
assert payload["agent_id"] == "test_agent"
|
||||
assert payload["key"] == "value" # From metadata
|
||||
|
||||
# Test with missing fields
|
||||
result = {
|
||||
# No memory_id
|
||||
"hash": "test_hash",
|
||||
# No memory
|
||||
# No created_at
|
||||
}
|
||||
|
||||
# Process the document fields
|
||||
payload, memory_id = valkey_db._process_document_fields(result, "default_id")
|
||||
|
||||
# Check the results
|
||||
assert memory_id == "default_id" # Should use default_id
|
||||
assert payload["hash"] == "test_hash"
|
||||
assert "data" in payload # Should have default value
|
||||
assert "created_at" in payload # Should have default value
|
||||
|
||||
|
||||
def test_init_connection_error():
|
||||
"""Test that initialization handles connection errors."""
|
||||
# Mock the from_url to raise an exception
|
||||
with patch("valkey.from_url") as mock_from_url:
|
||||
mock_from_url.side_effect = Exception("Connection failed")
|
||||
|
||||
# Initialize ValkeyDB should raise the exception
|
||||
with pytest.raises(Exception, match="Connection failed"):
|
||||
ValkeyDB(
|
||||
valkey_url="valkey://localhost:6379",
|
||||
collection_name="test_collection",
|
||||
embedding_model_dims=1536,
|
||||
)
|
||||
|
||||
|
||||
def test_build_search_query(valkey_db):
|
||||
"""Test building search queries with different filter scenarios."""
|
||||
# Test with no filters
|
||||
knn_part = "[KNN 5 @embedding $vec_param AS vector_score]"
|
||||
query = valkey_db._build_search_query(knn_part)
|
||||
assert query == f"*=>{knn_part}"
|
||||
|
||||
# Test with empty filters
|
||||
query = valkey_db._build_search_query(knn_part, {})
|
||||
assert query == f"*=>{knn_part}"
|
||||
|
||||
# Test with None values in filters
|
||||
query = valkey_db._build_search_query(knn_part, {"user_id": None})
|
||||
assert query == f"*=>{knn_part}"
|
||||
|
||||
# Test with single filter
|
||||
query = valkey_db._build_search_query(knn_part, {"user_id": "test_user"})
|
||||
assert query == f"@user_id:{{test_user}} =>{knn_part}"
|
||||
|
||||
# Test with multiple filters
|
||||
query = valkey_db._build_search_query(knn_part, {"user_id": "test_user", "agent_id": "test_agent"})
|
||||
assert "@user_id:{test_user}" in query
|
||||
assert "@agent_id:{test_agent}" in query
|
||||
assert f"=>{knn_part}" in query
|
||||
|
||||
|
||||
def test_get_error_handling(valkey_db, mock_valkey_client):
|
||||
"""Test error handling in the get method."""
|
||||
# Mock hgetall to raise an exception
|
||||
mock_valkey_client.hgetall.side_effect = Exception("Unexpected error")
|
||||
|
||||
# Call get should raise the exception
|
||||
with pytest.raises(Exception, match="Unexpected error"):
|
||||
valkey_db.get("test_id")
|
||||
|
||||
|
||||
def test_list_error_handling(valkey_db, mock_valkey_client):
|
||||
"""Test error handling in the list method."""
|
||||
# Mock search to raise an exception
|
||||
mock_ft = mock_valkey_client.ft.return_value
|
||||
mock_ft.search.side_effect = Exception("Unexpected error")
|
||||
|
||||
# Call list should return empty result on error
|
||||
results = valkey_db.list(filters={"user_id": "test_user"})
|
||||
|
||||
# Check that the result is an empty list
|
||||
assert results == [[]]
|
||||
|
||||
|
||||
def test_create_index_other_error():
|
||||
"""Test that initialization handles other errors during index creation."""
|
||||
# Mock the execute_command to raise a different error
|
||||
with patch("valkey.from_url") as mock_client:
|
||||
mock_client.return_value.execute_command.side_effect = ResponseError("Some other error")
|
||||
mock_client.return_value.ft = MagicMock()
|
||||
mock_client.return_value.ft.return_value.info.side_effect = ResponseError("not found")
|
||||
|
||||
# Initialize ValkeyDB should raise the exception
|
||||
with pytest.raises(ResponseError, match="Some other error"):
|
||||
ValkeyDB(
|
||||
valkey_url="valkey://localhost:6379",
|
||||
collection_name="test_collection",
|
||||
embedding_model_dims=1536,
|
||||
)
|
||||
|
||||
|
||||
def test_create_col_error(valkey_db, mock_valkey_client):
|
||||
"""Test error handling in create_col method."""
|
||||
# Mock execute_command to raise an exception
|
||||
mock_valkey_client.execute_command.side_effect = Exception("Failed to create index")
|
||||
|
||||
# Call create_col should raise the exception
|
||||
with pytest.raises(Exception, match="Failed to create index"):
|
||||
valkey_db.create_col(name="new_collection", vector_size=768)
|
||||
|
||||
|
||||
def test_list_cols_error(valkey_db, mock_valkey_client):
|
||||
"""Test error handling in list_cols method."""
|
||||
# Reset the mock to clear previous calls
|
||||
mock_valkey_client.execute_command.reset_mock()
|
||||
|
||||
# Mock execute_command to raise an exception
|
||||
mock_valkey_client.execute_command.side_effect = Exception("Failed to list indices")
|
||||
|
||||
# Call list_cols should raise the exception
|
||||
with pytest.raises(Exception, match="Failed to list indices"):
|
||||
valkey_db.list_cols()
|
||||
|
||||
|
||||
def test_col_info_error(valkey_db, mock_valkey_client):
|
||||
"""Test error handling in col_info method."""
|
||||
# Mock ft().info() to raise an exception
|
||||
mock_ft = mock_valkey_client.ft.return_value
|
||||
mock_ft.info.side_effect = Exception("Failed to get index info")
|
||||
|
||||
# Call col_info should raise the exception
|
||||
with pytest.raises(Exception, match="Failed to get index info"):
|
||||
valkey_db.col_info()
|
||||
|
||||
|
||||
# Additional tests to improve coverage
|
||||
|
||||
|
||||
def test_invalid_index_type():
|
||||
"""Test validation of invalid index type."""
|
||||
with pytest.raises(ValueError, match="Invalid index_type: invalid. Must be 'hnsw' or 'flat'"):
|
||||
ValkeyDB(
|
||||
valkey_url="valkey://localhost:6379",
|
||||
collection_name="test_collection",
|
||||
embedding_model_dims=1536,
|
||||
index_type="invalid",
|
||||
)
|
||||
|
||||
|
||||
def test_index_existence_check_error(mock_valkey_client):
|
||||
"""Test error handling when checking index existence."""
|
||||
# Mock ft().info() to raise a ResponseError that's not "not found"
|
||||
mock_ft = MagicMock()
|
||||
mock_ft.info.side_effect = ResponseError("Some other error")
|
||||
mock_valkey_client.ft.return_value = mock_ft
|
||||
|
||||
with patch("valkey.from_url", return_value=mock_valkey_client):
|
||||
with pytest.raises(ResponseError):
|
||||
ValkeyDB(
|
||||
valkey_url="valkey://localhost:6379",
|
||||
collection_name="test_collection",
|
||||
embedding_model_dims=1536,
|
||||
)
|
||||
|
||||
|
||||
def test_flat_index_creation(mock_valkey_client):
|
||||
"""Test creation of FLAT index type."""
|
||||
mock_ft = MagicMock()
|
||||
# Mock the info method to raise ResponseError with "not found" to trigger index creation
|
||||
mock_ft.info.side_effect = ResponseError("Index not found")
|
||||
mock_valkey_client.ft.return_value = mock_ft
|
||||
|
||||
with patch("valkey.from_url", return_value=mock_valkey_client):
|
||||
# Mock the execute_command to avoid the actual exception
|
||||
mock_valkey_client.execute_command.return_value = None
|
||||
|
||||
ValkeyDB(
|
||||
valkey_url="valkey://localhost:6379",
|
||||
collection_name="test_collection",
|
||||
embedding_model_dims=1536,
|
||||
index_type="flat",
|
||||
)
|
||||
|
||||
# Verify that execute_command was called (index creation)
|
||||
assert mock_valkey_client.execute_command.called
|
||||
|
||||
|
||||
def test_index_creation_error(mock_valkey_client):
|
||||
"""Test error handling during index creation."""
|
||||
mock_ft = MagicMock()
|
||||
mock_ft.info.side_effect = ResponseError("Unknown index name") # Index doesn't exist
|
||||
mock_valkey_client.ft.return_value = mock_ft
|
||||
mock_valkey_client.execute_command.side_effect = Exception("Failed to create index")
|
||||
|
||||
with patch("valkey.from_url", return_value=mock_valkey_client):
|
||||
with pytest.raises(Exception, match="Failed to create index"):
|
||||
ValkeyDB(
|
||||
valkey_url="valkey://localhost:6379",
|
||||
collection_name="test_collection",
|
||||
embedding_model_dims=1536,
|
||||
)
|
||||
|
||||
|
||||
def test_insert_missing_required_field(valkey_db, mock_valkey_client):
|
||||
"""Test error handling when inserting vector with missing required field."""
|
||||
# Mock hset to raise KeyError (missing required field)
|
||||
mock_valkey_client.hset.side_effect = KeyError("missing_field")
|
||||
|
||||
# This should not raise an exception but should log the error
|
||||
valkey_db.insert(vectors=[np.random.rand(1536).tolist()], payloads=[{"memory": "test"}], ids=["test_id"])
|
||||
|
||||
|
||||
def test_insert_general_error(valkey_db, mock_valkey_client):
|
||||
"""Test error handling for general exceptions during insert."""
|
||||
# Mock hset to raise a general exception
|
||||
mock_valkey_client.hset.side_effect = Exception("Database error")
|
||||
|
||||
with pytest.raises(Exception, match="Database error"):
|
||||
valkey_db.insert(vectors=[np.random.rand(1536).tolist()], payloads=[{"memory": "test"}], ids=["test_id"])
|
||||
|
||||
|
||||
def test_search_with_invalid_metadata(valkey_db, mock_valkey_client):
|
||||
"""Test search with invalid JSON metadata."""
|
||||
# Mock search results with invalid JSON metadata
|
||||
mock_doc = MagicMock()
|
||||
mock_doc.memory_id = "test_id"
|
||||
mock_doc.hash = "test_hash"
|
||||
mock_doc.memory = "test_data"
|
||||
mock_doc.created_at = str(int(datetime.now().timestamp()))
|
||||
mock_doc.metadata = "invalid_json" # Invalid JSON
|
||||
mock_doc.vector_score = "0.5"
|
||||
|
||||
mock_result = MagicMock()
|
||||
mock_result.docs = [mock_doc]
|
||||
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)
|
||||
|
||||
assert len(results) == 1
|
||||
|
||||
|
||||
def test_search_with_hnsw_ef_runtime(valkey_db, mock_valkey_client):
|
||||
"""Test search with HNSW ef_runtime parameter."""
|
||||
valkey_db.index_type = "hnsw"
|
||||
valkey_db.hnsw_ef_runtime = 20
|
||||
|
||||
mock_result = MagicMock()
|
||||
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)
|
||||
|
||||
# Verify the search was called
|
||||
assert mock_valkey_client.ft.return_value.search.called
|
||||
|
||||
|
||||
def test_delete_error(valkey_db, mock_valkey_client):
|
||||
"""Test error handling during vector deletion."""
|
||||
mock_valkey_client.delete.side_effect = Exception("Delete failed")
|
||||
|
||||
with pytest.raises(Exception, match="Delete failed"):
|
||||
valkey_db.delete("test_id")
|
||||
|
||||
|
||||
def test_update_missing_required_field(valkey_db, mock_valkey_client):
|
||||
"""Test error handling when updating vector with missing required field."""
|
||||
mock_valkey_client.hset.side_effect = KeyError("missing_field")
|
||||
|
||||
# This should not raise an exception but should log the error
|
||||
valkey_db.update(vector_id="test_id", vector=np.random.rand(1536).tolist(), payload={"memory": "updated"})
|
||||
|
||||
|
||||
def test_update_general_error(valkey_db, mock_valkey_client):
|
||||
"""Test error handling for general exceptions during update."""
|
||||
mock_valkey_client.hset.side_effect = Exception("Update failed")
|
||||
|
||||
with pytest.raises(Exception, match="Update failed"):
|
||||
valkey_db.update(vector_id="test_id", vector=np.random.rand(1536).tolist(), payload={"memory": "updated"})
|
||||
|
||||
|
||||
def test_get_with_binary_data_and_unicode_error(valkey_db, mock_valkey_client):
|
||||
"""Test get method with binary data that fails UTF-8 decoding."""
|
||||
# Mock result with binary data that can't be decoded
|
||||
mock_result = {
|
||||
"memory_id": "test_id",
|
||||
"hash": b"\xff\xfe", # Invalid UTF-8 bytes
|
||||
"memory": "test_memory",
|
||||
"created_at": "1234567890",
|
||||
"updated_at": "invalid_timestamp",
|
||||
"metadata": "{}",
|
||||
"embedding": b"binary_embedding_data",
|
||||
}
|
||||
mock_valkey_client.hgetall.return_value = mock_result
|
||||
|
||||
result = valkey_db.get("test_id")
|
||||
|
||||
# Should handle binary data gracefully
|
||||
assert result.id == "test_id"
|
||||
assert result.payload["data"] == "test_memory"
|
||||
|
||||
|
||||
def test_get_with_invalid_timestamps(valkey_db, mock_valkey_client):
|
||||
"""Test get method with invalid timestamp values."""
|
||||
mock_result = {
|
||||
"memory_id": "test_id",
|
||||
"hash": "test_hash",
|
||||
"memory": "test_memory",
|
||||
"created_at": "invalid_timestamp",
|
||||
"updated_at": "also_invalid",
|
||||
"metadata": "{}",
|
||||
"embedding": b"binary_data",
|
||||
}
|
||||
mock_valkey_client.hgetall.return_value = mock_result
|
||||
|
||||
result = valkey_db.get("test_id")
|
||||
|
||||
# Should handle invalid timestamps gracefully
|
||||
assert result.id == "test_id"
|
||||
assert "created_at" in result.payload
|
||||
|
||||
|
||||
def test_get_with_invalid_metadata_json(valkey_db, mock_valkey_client):
|
||||
"""Test get method with invalid JSON metadata."""
|
||||
mock_result = {
|
||||
"memory_id": "test_id",
|
||||
"hash": "test_hash",
|
||||
"memory": "test_memory",
|
||||
"created_at": "1234567890",
|
||||
"updated_at": "1234567890",
|
||||
"metadata": "invalid_json{", # Invalid JSON
|
||||
"embedding": b"binary_data",
|
||||
}
|
||||
mock_valkey_client.hgetall.return_value = mock_result
|
||||
|
||||
result = valkey_db.get("test_id")
|
||||
|
||||
# Should handle invalid JSON gracefully
|
||||
assert result.id == "test_id"
|
||||
|
||||
|
||||
def test_list_with_missing_fields_and_defaults(valkey_db, mock_valkey_client):
|
||||
"""Test list method with documents missing various fields."""
|
||||
# Mock search results with missing fields but valid timestamps
|
||||
mock_doc1 = MagicMock()
|
||||
mock_doc1.memory_id = "fallback_id"
|
||||
mock_doc1.hash = "test_hash" # Provide valid hash
|
||||
mock_doc1.memory = "test_memory" # Provide valid memory
|
||||
mock_doc1.created_at = str(int(datetime.now().timestamp())) # Valid timestamp
|
||||
mock_doc1.updated_at = str(int(datetime.now().timestamp())) # Valid timestamp
|
||||
mock_doc1.metadata = json.dumps({"key": "value"}) # Valid JSON
|
||||
mock_doc1.vector_score = "0.5"
|
||||
|
||||
mock_result = MagicMock()
|
||||
mock_result.docs = [mock_doc1]
|
||||
mock_valkey_client.ft.return_value.search.return_value = mock_result
|
||||
|
||||
results = valkey_db.list()
|
||||
|
||||
# Should handle the search-based list approach
|
||||
assert len(results) == 1
|
||||
inner_results = results[0]
|
||||
assert len(inner_results) == 1
|
||||
result = inner_results[0]
|
||||
assert result.id == "fallback_id"
|
||||
assert "hash" in result.payload
|
||||
assert "data" in result.payload # memory is renamed to data
|
||||
Reference in New Issue
Block a user