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