feat(vector-store): Add Valkey vector store support (#3272)
This commit is contained in:
committed by
GitHub
parent
e64488b598
commit
e3f0277cb9
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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>
|
||||
|
||||
@@ -152,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",
|
||||
|
||||
@@ -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
|
||||
@@ -165,6 +165,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",
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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
|
||||
@@ -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",
|
||||
|
||||
@@ -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