feat(vector-store): Add Valkey vector store support (#3272)

This commit is contained in:
Swarnaprakash Udayakumar
2025-09-09 15:31:53 -07:00
committed by GitHub
parent e64488b598
commit e3f0277cb9
11 changed files with 1758 additions and 3 deletions
+1 -1
View File
@@ -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:
+1 -1
View File
@@ -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
+49
View File
@@ -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` |
+2 -1
View File
@@ -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>
+1
View File
@@ -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",
+15
View File
@@ -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
View File
@@ -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",
+1
View File
@@ -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",
+824
View File
@@ -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
+1
View File
@@ -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",
+862
View File
@@ -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