diff --git a/Makefile b/Makefile
index 3c8287edd..14098f00c 100644
--- a/Makefile
+++ b/Makefile
@@ -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:
diff --git a/docs/components/vectordbs/config.mdx b/docs/components/vectordbs/config.mdx
index a36e55795..89d995d21 100644
--- a/docs/components/vectordbs/config.mdx
+++ b/docs/components/vectordbs/config.mdx
@@ -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
diff --git a/docs/components/vectordbs/dbs/valkey.mdx b/docs/components/vectordbs/dbs/valkey.mdx
new file mode 100644
index 000000000..3c6d72e84
--- /dev/null
+++ b/docs/components/vectordbs/dbs/valkey.mdx
@@ -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` |
diff --git a/docs/components/vectordbs/overview.mdx b/docs/components/vectordbs/overview.mdx
index 83b55d20c..ba504541c 100644
--- a/docs/components/vectordbs/overview.mdx
+++ b/docs/components/vectordbs/overview.mdx
@@ -11,7 +11,7 @@ Mem0 includes built-in support for various popular databases. Memory can utilize
See the list of supported vector databases below.
- 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.
@@ -24,6 +24,7 @@ See the list of supported vector databases below.
+
diff --git a/docs/docs.json b/docs/docs.json
index ebf1e2a1c..2001418e9 100644
--- a/docs/docs.json
+++ b/docs/docs.json
@@ -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",
diff --git a/mem0/configs/vector_stores/valkey.py b/mem0/configs/vector_stores/valkey.py
new file mode 100644
index 000000000..1c04049e6
--- /dev/null
+++ b/mem0/configs/vector_stores/valkey.py
@@ -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
diff --git a/mem0/utils/factory.py b/mem0/utils/factory.py
index 1f11da8e1..a27d16c69 100644
--- a/mem0/utils/factory.py
+++ b/mem0/utils/factory.py
@@ -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",
diff --git a/mem0/vector_stores/configs.py b/mem0/vector_stores/configs.py
index b2e14bbe0..f9570f749 100644
--- a/mem0/vector_stores/configs.py
+++ b/mem0/vector_stores/configs.py
@@ -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",
diff --git a/mem0/vector_stores/valkey.py b/mem0/vector_stores/valkey.py
new file mode 100644
index 000000000..c4539dcd2
--- /dev/null
+++ b/mem0/vector_stores/valkey.py
@@ -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
diff --git a/pyproject.toml b/pyproject.toml
index 386e828fa..537c24105 100644
--- a/pyproject.toml
+++ b/pyproject.toml
@@ -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",
diff --git a/tests/vector_stores/test_valkey.py b/tests/vector_stores/test_valkey.py
new file mode 100644
index 000000000..482f9b83d
--- /dev/null
+++ b/tests/vector_stores/test_valkey.py
@@ -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