From 9ef644b95e41a32b80e87ab9f53d3cfbc6a52ef6 Mon Sep 17 00:00:00 2001
From: Faizan Habib <91795555+faizan842@users.noreply.github.com>
Date: Fri, 17 Oct 2025 23:48:49 +0530
Subject: [PATCH] Add Apache Cassandra vector store support (#3578)
---
docs/components/vectordbs/dbs/cassandra.mdx | 181 +++++++
docs/docs.json | 3 +-
mem0/configs/vector_stores/cassandra.py | 77 +++
mem0/utils/factory.py | 1 +
mem0/vector_stores/cassandra.py | 496 ++++++++++++++++++++
mem0/vector_stores/configs.py | 1 +
pyproject.toml | 1 +
tests/vector_stores/test_cassandra.py | 316 +++++++++++++
8 files changed, 1075 insertions(+), 1 deletion(-)
create mode 100644 docs/components/vectordbs/dbs/cassandra.mdx
create mode 100644 mem0/configs/vector_stores/cassandra.py
create mode 100644 mem0/vector_stores/cassandra.py
create mode 100644 tests/vector_stores/test_cassandra.py
diff --git a/docs/components/vectordbs/dbs/cassandra.mdx b/docs/components/vectordbs/dbs/cassandra.mdx
new file mode 100644
index 000000000..658ff239d
--- /dev/null
+++ b/docs/components/vectordbs/dbs/cassandra.mdx
@@ -0,0 +1,181 @@
+---
+title: Apache Cassandra
+---
+
+[Apache Cassandra](https://cassandra.apache.org/) is a highly scalable, distributed NoSQL database designed for handling large amounts of data across many commodity servers with no single point of failure. It supports vector storage for semantic search capabilities in AI applications and can scale to massive datasets with linear performance improvements.
+
+### Usage
+
+```python
+import os
+from mem0 import Memory
+
+os.environ["OPENAI_API_KEY"] = "sk-xx"
+
+config = {
+ "vector_store": {
+ "provider": "cassandra",
+ "config": {
+ "contact_points": ["127.0.0.1"],
+ "port": 9042,
+ "username": "cassandra",
+ "password": "cassandra",
+ "keyspace": "mem0",
+ "collection_name": "memories",
+ }
+ }
+}
+
+m = Memory.from_config(config)
+messages = [
+ {"role": "user", "content": "I'm planning to watch a movie tonight. Any recommendations?"},
+ {"role": "assistant", "content": "How about 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"})
+```
+
+#### Using DataStax Astra DB
+
+For managed Cassandra with DataStax Astra DB:
+
+```python
+config = {
+ "vector_store": {
+ "provider": "cassandra",
+ "config": {
+ "contact_points": ["dummy"], # Not used with secure connect bundle
+ "username": "token",
+ "password": "AstraCS:...", # Your Astra DB application token
+ "keyspace": "mem0",
+ "collection_name": "memories",
+ "secure_connect_bundle": "/path/to/secure-connect-bundle.zip"
+ }
+ }
+}
+```
+
+
+When using DataStax Astra DB, provide the secure connect bundle path. The contact_points parameter is ignored when a secure connect bundle is provided.
+
+
+### Config
+
+Here are the parameters available for configuring Apache Cassandra:
+
+| Parameter | Description | Default Value |
+| --- | --- | --- |
+| `contact_points` | List of contact point IP addresses | Required |
+| `port` | Cassandra port | `9042` |
+| `username` | Database username | `None` |
+| `password` | Database password | `None` |
+| `keyspace` | Keyspace name | `"mem0"` |
+| `collection_name` | Table name for storing vectors | `"memories"` |
+| `embedding_model_dims` | Dimensions of embedding vectors | `1536` |
+| `secure_connect_bundle` | Path to Astra DB secure connect bundle | `None` |
+| `protocol_version` | CQL protocol version | `4` |
+| `load_balancing_policy` | Custom load balancing policy | `None` |
+
+### Setup
+
+#### Option 1: Local Cassandra Setup using Docker:
+
+```bash
+# Pull and run Cassandra container
+docker run --name mem0-cassandra \
+ -p 9042:9042 \
+ -e CASSANDRA_CLUSTER_NAME="Mem0Cluster" \
+ -d cassandra:latest
+
+# Wait for Cassandra to start (may take 1-2 minutes)
+docker exec -it mem0-cassandra cqlsh
+
+# Create keyspace
+CREATE KEYSPACE IF NOT EXISTS mem0
+WITH replication = {'class': 'SimpleStrategy', 'replication_factor': 1};
+```
+
+#### Option 2: DataStax Astra DB (Managed Cloud):
+
+1. Sign up at [DataStax Astra](https://astra.datastax.com/)
+2. Create a new database
+3. Download the secure connect bundle
+4. Generate an application token
+
+
+For production deployments, use DataStax Astra DB for fully managed Cassandra with automatic scaling, backups, and security.
+
+
+#### Option 3: Install Cassandra Locally:
+
+**Ubuntu/Debian:**
+```bash
+# Add Apache Cassandra repository
+echo "deb https://downloads.apache.org/cassandra/debian 40x main" | sudo tee -a /etc/apt/sources.list.d/cassandra.sources.list
+curl https://downloads.apache.org/cassandra/KEYS | sudo apt-key add -
+
+# Install Cassandra
+sudo apt-get update
+sudo apt-get install cassandra
+
+# Start Cassandra
+sudo systemctl start cassandra
+
+# Verify installation
+nodetool status
+```
+
+**macOS:**
+```bash
+# Using Homebrew
+brew install cassandra
+
+# Start Cassandra
+brew services start cassandra
+
+# Connect to CQL shell
+cqlsh
+```
+
+### Python Client Installation
+
+Install the required Python package:
+
+```bash
+pip install cassandra-driver
+```
+
+### Performance Considerations
+
+- **Replication Factor**: For production, use replication factor of at least 3
+- **Consistency Level**: Balance between consistency and performance (QUORUM recommended)
+- **Partitioning**: Cassandra automatically distributes data across nodes
+- **Scaling**: Add nodes to linearly increase capacity and performance
+
+### Advanced Configuration
+
+```python
+from cassandra.policies import DCAwareRoundRobinPolicy
+
+config = {
+ "vector_store": {
+ "provider": "cassandra",
+ "config": {
+ "contact_points": ["node1.example.com", "node2.example.com", "node3.example.com"],
+ "port": 9042,
+ "username": "mem0_user",
+ "password": "secure_password",
+ "keyspace": "mem0_prod",
+ "collection_name": "memories",
+ "protocol_version": 4,
+ "load_balancing_policy": DCAwareRoundRobinPolicy(local_dc='DC1')
+ }
+ }
+}
+```
+
+
+For production use, configure appropriate replication strategies and consistency levels based on your availability and consistency requirements.
+
+
diff --git a/docs/docs.json b/docs/docs.json
index bb0ad7ff9..472673c21 100644
--- a/docs/docs.json
+++ b/docs/docs.json
@@ -59,7 +59,7 @@
"icon": "star",
"pages": [
"platform/features/platform-overview",
- "platform/features/v2-memory-filters",
+ "platform/features/v2-memory-filters",
"platform/features/contextual-add",
"platform/features/async-client",
"platform/features/async-mode-default-change",
@@ -153,6 +153,7 @@
"pages": [
"components/vectordbs/dbs/qdrant",
"components/vectordbs/dbs/chroma",
+ "components/vectordbs/dbs/cassandra",
"components/vectordbs/dbs/pgvector",
"components/vectordbs/dbs/milvus",
"components/vectordbs/dbs/pinecone",
diff --git a/mem0/configs/vector_stores/cassandra.py b/mem0/configs/vector_stores/cassandra.py
new file mode 100644
index 000000000..40e629a81
--- /dev/null
+++ b/mem0/configs/vector_stores/cassandra.py
@@ -0,0 +1,77 @@
+from typing import Any, Dict, List, Optional
+
+from pydantic import BaseModel, Field, model_validator
+
+
+class CassandraConfig(BaseModel):
+ """Configuration for Apache Cassandra vector database."""
+
+ contact_points: List[str] = Field(
+ ...,
+ description="List of contact point addresses (e.g., ['127.0.0.1', '127.0.0.2'])"
+ )
+ port: int = Field(9042, description="Cassandra port")
+ username: Optional[str] = Field(None, description="Database username")
+ password: Optional[str] = Field(None, description="Database password")
+ keyspace: str = Field("mem0", description="Keyspace name")
+ collection_name: str = Field("memories", description="Table name")
+ embedding_model_dims: int = Field(1536, description="Dimensions of the embedding model")
+ secure_connect_bundle: Optional[str] = Field(
+ None,
+ description="Path to secure connect bundle for DataStax Astra DB"
+ )
+ protocol_version: int = Field(4, description="CQL protocol version")
+ load_balancing_policy: Optional[Any] = Field(
+ None,
+ description="Custom load balancing policy object"
+ )
+
+ @model_validator(mode="before")
+ @classmethod
+ def check_auth(cls, values: Dict[str, Any]) -> Dict[str, Any]:
+ """Validate authentication parameters."""
+ username = values.get("username")
+ password = values.get("password")
+
+ # Both username and password must be provided together or not at all
+ if (username and not password) or (password and not username):
+ raise ValueError(
+ "Both 'username' and 'password' must be provided together for authentication"
+ )
+
+ return values
+
+ @model_validator(mode="before")
+ @classmethod
+ def check_connection_config(cls, values: Dict[str, Any]) -> Dict[str, Any]:
+ """Validate connection configuration."""
+ secure_connect_bundle = values.get("secure_connect_bundle")
+ contact_points = values.get("contact_points")
+
+ # Either secure_connect_bundle or contact_points must be provided
+ if not secure_connect_bundle and not contact_points:
+ raise ValueError(
+ "Either 'contact_points' or 'secure_connect_bundle' must be provided"
+ )
+
+ return values
+
+ @model_validator(mode="before")
+ @classmethod
+ def validate_extra_fields(cls, values: Dict[str, Any]) -> Dict[str, Any]:
+ """Validate that no extra fields are provided."""
+ allowed_fields = set(cls.model_fields.keys())
+ input_fields = set(values.keys())
+ extra_fields = input_fields - allowed_fields
+
+ if extra_fields:
+ raise ValueError(
+ f"Extra fields not allowed: {', '.join(extra_fields)}. "
+ f"Please input only the following fields: {', '.join(allowed_fields)}"
+ )
+
+ return values
+
+ class Config:
+ arbitrary_types_allowed = True
+
diff --git a/mem0/utils/factory.py b/mem0/utils/factory.py
index 534a9322b..ab3fc77a3 100644
--- a/mem0/utils/factory.py
+++ b/mem0/utils/factory.py
@@ -184,6 +184,7 @@ class VectorStoreFactory:
"langchain": "mem0.vector_stores.langchain.Langchain",
"s3_vectors": "mem0.vector_stores.s3_vectors.S3Vectors",
"baidu": "mem0.vector_stores.baidu.BaiduDB",
+ "cassandra": "mem0.vector_stores.cassandra.CassandraDB",
"neptune": "mem0.vector_stores.neptune_analytics.NeptuneAnalyticsVector",
}
diff --git a/mem0/vector_stores/cassandra.py b/mem0/vector_stores/cassandra.py
new file mode 100644
index 000000000..24e4fea88
--- /dev/null
+++ b/mem0/vector_stores/cassandra.py
@@ -0,0 +1,496 @@
+import json
+import logging
+import uuid
+from typing import Any, Dict, List, Optional
+
+import numpy as np
+from pydantic import BaseModel
+
+try:
+ from cassandra.cluster import Cluster
+ from cassandra.auth import PlainTextAuthProvider
+except ImportError:
+ raise ImportError(
+ "Apache Cassandra vector store requires cassandra-driver. "
+ "Please install it using 'pip install cassandra-driver'"
+ )
+
+from mem0.vector_stores.base import VectorStoreBase
+
+logger = logging.getLogger(__name__)
+
+
+class OutputData(BaseModel):
+ id: Optional[str]
+ score: Optional[float]
+ payload: Optional[dict]
+
+
+class CassandraDB(VectorStoreBase):
+ def __init__(
+ self,
+ contact_points: List[str],
+ port: int = 9042,
+ username: Optional[str] = None,
+ password: Optional[str] = None,
+ keyspace: str = "mem0",
+ collection_name: str = "memories",
+ embedding_model_dims: int = 1536,
+ secure_connect_bundle: Optional[str] = None,
+ protocol_version: int = 4,
+ load_balancing_policy: Optional[Any] = None,
+ ):
+ """
+ Initialize the Apache Cassandra vector store.
+
+ Args:
+ contact_points (List[str]): List of contact point addresses (e.g., ['127.0.0.1'])
+ port (int): Cassandra port (default: 9042)
+ username (str, optional): Database username
+ password (str, optional): Database password
+ keyspace (str): Keyspace name (default: "mem0")
+ collection_name (str): Table name (default: "memories")
+ embedding_model_dims (int): Dimension of the embedding vector (default: 1536)
+ secure_connect_bundle (str, optional): Path to secure connect bundle for Astra DB
+ protocol_version (int): CQL protocol version (default: 4)
+ load_balancing_policy (Any, optional): Custom load balancing policy
+ """
+ self.contact_points = contact_points
+ self.port = port
+ self.username = username
+ self.password = password
+ self.keyspace = keyspace
+ self.collection_name = collection_name
+ self.embedding_model_dims = embedding_model_dims
+ self.secure_connect_bundle = secure_connect_bundle
+ self.protocol_version = protocol_version
+ self.load_balancing_policy = load_balancing_policy
+
+ # Initialize connection
+ self.cluster = None
+ self.session = None
+ self._setup_connection()
+
+ # Create keyspace and table if they don't exist
+ self._create_keyspace()
+ self._create_table()
+
+ def _setup_connection(self):
+ """Setup Cassandra cluster connection."""
+ try:
+ # Setup authentication
+ auth_provider = None
+ if self.username and self.password:
+ auth_provider = PlainTextAuthProvider(
+ username=self.username,
+ password=self.password
+ )
+
+ # Connect to Astra DB using secure connect bundle
+ if self.secure_connect_bundle:
+ self.cluster = Cluster(
+ cloud={'secure_connect_bundle': self.secure_connect_bundle},
+ auth_provider=auth_provider,
+ protocol_version=self.protocol_version
+ )
+ else:
+ # Connect to standard Cassandra cluster
+ cluster_kwargs = {
+ 'contact_points': self.contact_points,
+ 'port': self.port,
+ 'protocol_version': self.protocol_version
+ }
+
+ if auth_provider:
+ cluster_kwargs['auth_provider'] = auth_provider
+
+ if self.load_balancing_policy:
+ cluster_kwargs['load_balancing_policy'] = self.load_balancing_policy
+
+ self.cluster = Cluster(**cluster_kwargs)
+
+ self.session = self.cluster.connect()
+ logger.info("Successfully connected to Cassandra cluster")
+ except Exception as e:
+ logger.error(f"Failed to connect to Cassandra: {e}")
+ raise
+
+ def _create_keyspace(self):
+ """Create keyspace if it doesn't exist."""
+ try:
+ # Use SimpleStrategy for single datacenter, NetworkTopologyStrategy for production
+ query = f"""
+ CREATE KEYSPACE IF NOT EXISTS {self.keyspace}
+ WITH replication = {{'class': 'SimpleStrategy', 'replication_factor': 1}}
+ """
+ self.session.execute(query)
+ self.session.set_keyspace(self.keyspace)
+ logger.info(f"Keyspace '{self.keyspace}' is ready")
+ except Exception as e:
+ logger.error(f"Failed to create keyspace: {e}")
+ raise
+
+ def _create_table(self):
+ """Create table with vector column if it doesn't exist."""
+ try:
+ # Create table with vector stored as list and payload as text (JSON)
+ query = f"""
+ CREATE TABLE IF NOT EXISTS {self.keyspace}.{self.collection_name} (
+ id text PRIMARY KEY,
+ vector list,
+ payload text
+ )
+ """
+ self.session.execute(query)
+ logger.info(f"Table '{self.collection_name}' is ready")
+ except Exception as e:
+ logger.error(f"Failed to create table: {e}")
+ raise
+
+ def create_col(self, name: str = None, vector_size: int = None, distance: str = "cosine"):
+ """
+ Create a new collection (table in Cassandra).
+
+ Args:
+ name (str, optional): Collection name (uses self.collection_name if not provided)
+ vector_size (int, optional): Vector dimension (uses self.embedding_model_dims if not provided)
+ distance (str): Distance metric (cosine, euclidean, dot_product)
+ """
+ table_name = name or self.collection_name
+ dims = vector_size or self.embedding_model_dims
+
+ try:
+ query = f"""
+ CREATE TABLE IF NOT EXISTS {self.keyspace}.{table_name} (
+ id text PRIMARY KEY,
+ vector list,
+ payload text
+ )
+ """
+ self.session.execute(query)
+ logger.info(f"Created collection '{table_name}' with vector dimension {dims}")
+ except Exception as e:
+ logger.error(f"Failed to create collection: {e}")
+ raise
+
+ def insert(
+ self,
+ vectors: List[List[float]],
+ payloads: Optional[List[Dict]] = None,
+ ids: Optional[List[str]] = None
+ ):
+ """
+ Insert vectors into the collection.
+
+ Args:
+ vectors (List[List[float]]): List of vectors to insert
+ payloads (List[Dict], optional): List of payloads corresponding to vectors
+ ids (List[str], optional): List of IDs corresponding to vectors
+ """
+ logger.info(f"Inserting {len(vectors)} vectors into collection {self.collection_name}")
+
+ if payloads is None:
+ payloads = [{}] * len(vectors)
+ if ids is None:
+ ids = [str(uuid.uuid4()) for _ in range(len(vectors))]
+
+ try:
+ query = f"""
+ INSERT INTO {self.keyspace}.{self.collection_name} (id, vector, payload)
+ VALUES (?, ?, ?)
+ """
+ prepared = self.session.prepare(query)
+
+ for vector, payload, vec_id in zip(vectors, payloads, ids):
+ self.session.execute(
+ prepared,
+ (vec_id, vector, json.dumps(payload))
+ )
+ except Exception as e:
+ logger.error(f"Failed to insert vectors: {e}")
+ raise
+
+ def search(
+ self,
+ query: str,
+ vectors: List[float],
+ limit: int = 5,
+ filters: Optional[Dict] = None,
+ ) -> List[OutputData]:
+ """
+ Search for similar vectors using cosine similarity.
+
+ Args:
+ query (str): Query string (not used in vector search)
+ vectors (List[float]): Query vector
+ limit (int): Number of results to return
+ filters (Dict, optional): Filters to apply to the search
+
+ Returns:
+ List[OutputData]: Search results
+ """
+ try:
+ # Fetch all vectors (in production, you'd want pagination or filtering)
+ query_cql = f"""
+ SELECT id, vector, payload
+ FROM {self.keyspace}.{self.collection_name}
+ """
+ rows = self.session.execute(query_cql)
+
+ # Calculate cosine similarity in Python
+ query_vec = np.array(vectors)
+ scored_results = []
+
+ for row in rows:
+ if not row.vector:
+ continue
+
+ vec = np.array(row.vector)
+
+ # Cosine similarity
+ similarity = np.dot(query_vec, vec) / (np.linalg.norm(query_vec) * np.linalg.norm(vec))
+ distance = 1 - similarity
+
+ # Apply filters if provided
+ if filters:
+ try:
+ payload = json.loads(row.payload) if row.payload else {}
+ match = all(payload.get(k) == v for k, v in filters.items())
+ if not match:
+ continue
+ except json.JSONDecodeError:
+ continue
+
+ scored_results.append((row.id, distance, row.payload))
+
+ # Sort by distance and limit
+ scored_results.sort(key=lambda x: x[1])
+ scored_results = scored_results[:limit]
+
+ return [
+ OutputData(
+ id=r[0],
+ score=float(r[1]),
+ payload=json.loads(r[2]) if r[2] else {}
+ )
+ for r in scored_results
+ ]
+ except Exception as e:
+ logger.error(f"Search failed: {e}")
+ raise
+
+ def delete(self, vector_id: str):
+ """
+ Delete a vector by ID.
+
+ Args:
+ vector_id (str): ID of the vector to delete
+ """
+ try:
+ query = f"""
+ DELETE FROM {self.keyspace}.{self.collection_name}
+ WHERE id = ?
+ """
+ prepared = self.session.prepare(query)
+ self.session.execute(prepared, (vector_id,))
+ logger.info(f"Deleted vector with id: {vector_id}")
+ except Exception as e:
+ logger.error(f"Failed to delete vector: {e}")
+ raise
+
+ def update(
+ self,
+ vector_id: str,
+ vector: Optional[List[float]] = None,
+ payload: Optional[Dict] = None,
+ ):
+ """
+ Update a vector and its payload.
+
+ Args:
+ vector_id (str): ID of the vector to update
+ vector (List[float], optional): Updated vector
+ payload (Dict, optional): Updated payload
+ """
+ try:
+ if vector is not None:
+ query = f"""
+ UPDATE {self.keyspace}.{self.collection_name}
+ SET vector = ?
+ WHERE id = ?
+ """
+ prepared = self.session.prepare(query)
+ self.session.execute(prepared, (vector, vector_id))
+
+ if payload is not None:
+ query = f"""
+ UPDATE {self.keyspace}.{self.collection_name}
+ SET payload = ?
+ WHERE id = ?
+ """
+ prepared = self.session.prepare(query)
+ self.session.execute(prepared, (json.dumps(payload), vector_id))
+
+ logger.info(f"Updated vector with id: {vector_id}")
+ except Exception as e:
+ logger.error(f"Failed to update vector: {e}")
+ raise
+
+ def get(self, vector_id: str) -> Optional[OutputData]:
+ """
+ Retrieve a vector by ID.
+
+ Args:
+ vector_id (str): ID of the vector to retrieve
+
+ Returns:
+ OutputData: Retrieved vector or None if not found
+ """
+ try:
+ query = f"""
+ SELECT id, vector, payload
+ FROM {self.keyspace}.{self.collection_name}
+ WHERE id = ?
+ """
+ prepared = self.session.prepare(query)
+ row = self.session.execute(prepared, (vector_id,)).one()
+
+ if not row:
+ return None
+
+ return OutputData(
+ id=row.id,
+ score=None,
+ payload=json.loads(row.payload) if row.payload else {}
+ )
+ except Exception as e:
+ logger.error(f"Failed to get vector: {e}")
+ return None
+
+ def list_cols(self) -> List[str]:
+ """
+ List all collections (tables in the keyspace).
+
+ Returns:
+ List[str]: List of collection names
+ """
+ try:
+ query = f"""
+ SELECT table_name
+ FROM system_schema.tables
+ WHERE keyspace_name = '{self.keyspace}'
+ """
+ rows = self.session.execute(query)
+ return [row.table_name for row in rows]
+ except Exception as e:
+ logger.error(f"Failed to list collections: {e}")
+ return []
+
+ def delete_col(self):
+ """Delete the collection (table)."""
+ try:
+ query = f"""
+ DROP TABLE IF EXISTS {self.keyspace}.{self.collection_name}
+ """
+ self.session.execute(query)
+ logger.info(f"Deleted collection '{self.collection_name}'")
+ except Exception as e:
+ logger.error(f"Failed to delete collection: {e}")
+ raise
+
+ def col_info(self) -> Dict[str, Any]:
+ """
+ Get information about the collection.
+
+ Returns:
+ Dict[str, Any]: Collection information
+ """
+ try:
+ # Get row count (approximate)
+ query = f"""
+ SELECT COUNT(*) as count
+ FROM {self.keyspace}.{self.collection_name}
+ """
+ row = self.session.execute(query).one()
+ count = row.count if row else 0
+
+ return {
+ "name": self.collection_name,
+ "keyspace": self.keyspace,
+ "count": count,
+ "vector_dims": self.embedding_model_dims
+ }
+ except Exception as e:
+ logger.error(f"Failed to get collection info: {e}")
+ return {}
+
+ def list(
+ self,
+ filters: Optional[Dict] = None,
+ limit: int = 100
+ ) -> List[List[OutputData]]:
+ """
+ List all vectors in the collection.
+
+ Args:
+ filters (Dict, optional): Filters to apply
+ limit (int): Number of vectors to return
+
+ Returns:
+ List[List[OutputData]]: List of vectors
+ """
+ try:
+ query = f"""
+ SELECT id, vector, payload
+ FROM {self.keyspace}.{self.collection_name}
+ LIMIT {limit}
+ """
+ rows = self.session.execute(query)
+
+ results = []
+ for row in rows:
+ # Apply filters if provided
+ if filters:
+ try:
+ payload = json.loads(row.payload) if row.payload else {}
+ match = all(payload.get(k) == v for k, v in filters.items())
+ if not match:
+ continue
+ except json.JSONDecodeError:
+ continue
+
+ results.append(
+ OutputData(
+ id=row.id,
+ score=None,
+ payload=json.loads(row.payload) if row.payload else {}
+ )
+ )
+
+ return [results]
+ except Exception as e:
+ logger.error(f"Failed to list vectors: {e}")
+ return [[]]
+
+ def reset(self):
+ """Reset the collection by truncating it."""
+ try:
+ logger.warning(f"Resetting collection {self.collection_name}...")
+ query = f"""
+ TRUNCATE TABLE {self.keyspace}.{self.collection_name}
+ """
+ self.session.execute(query)
+ logger.info(f"Collection '{self.collection_name}' has been reset")
+ except Exception as e:
+ logger.error(f"Failed to reset collection: {e}")
+ raise
+
+ def __del__(self):
+ """Close the cluster connection when the object is deleted."""
+ try:
+ if self.cluster:
+ self.cluster.shutdown()
+ logger.info("Cassandra cluster connection closed")
+ except Exception:
+ pass
+
diff --git a/mem0/vector_stores/configs.py b/mem0/vector_stores/configs.py
index 42edf53a6..d08bae37a 100644
--- a/mem0/vector_stores/configs.py
+++ b/mem0/vector_stores/configs.py
@@ -18,6 +18,7 @@ class VectorStoreConfig(BaseModel):
"mongodb": "MongoDBConfig",
"milvus": "MilvusDBConfig",
"baidu": "BaiduDBConfig",
+ "cassandra": "CassandraConfig",
"neptune": "NeptuneAnalyticsConfig",
"upstash_vector": "UpstashVectorConfig",
"azure_ai_search": "AzureAISearchConfig",
diff --git a/pyproject.toml b/pyproject.toml
index 69415e13a..d60cc3346 100644
--- a/pyproject.toml
+++ b/pyproject.toml
@@ -35,6 +35,7 @@ graph = [
vector_stores = [
"vecs>=0.4.0",
"chromadb>=0.4.24",
+ "cassandra-driver>=3.29.0",
"weaviate-client>=4.4.0,<4.15.0",
"pinecone<=7.3.0",
"pinecone-text>=0.10.0",
diff --git a/tests/vector_stores/test_cassandra.py b/tests/vector_stores/test_cassandra.py
new file mode 100644
index 000000000..3194e4a8d
--- /dev/null
+++ b/tests/vector_stores/test_cassandra.py
@@ -0,0 +1,316 @@
+import json
+import pytest
+from unittest.mock import Mock, patch
+
+from mem0.vector_stores.cassandra import CassandraDB, OutputData
+
+
+@pytest.fixture
+def mock_session():
+ """Create a mock Cassandra session."""
+ session = Mock()
+ session.execute = Mock(return_value=Mock())
+ session.prepare = Mock(return_value=Mock())
+ session.set_keyspace = Mock()
+ return session
+
+
+@pytest.fixture
+def mock_cluster(mock_session):
+ """Create a mock Cassandra cluster."""
+ cluster = Mock()
+ cluster.connect = Mock(return_value=mock_session)
+ cluster.shutdown = Mock()
+ return cluster
+
+
+@pytest.fixture
+def cassandra_instance(mock_cluster, mock_session):
+ """Create a CassandraDB instance with mocked cluster."""
+ with patch('mem0.vector_stores.cassandra.Cluster') as mock_cluster_class:
+ mock_cluster_class.return_value = mock_cluster
+
+ instance = CassandraDB(
+ contact_points=['127.0.0.1'],
+ port=9042,
+ username='testuser',
+ password='testpass',
+ keyspace='test_keyspace',
+ collection_name='test_collection',
+ embedding_model_dims=128,
+ )
+ instance.session = mock_session
+ return instance
+
+
+def test_cassandra_init(mock_cluster, mock_session):
+ """Test CassandraDB initialization."""
+ with patch('mem0.vector_stores.cassandra.Cluster') as mock_cluster_class:
+ mock_cluster_class.return_value = mock_cluster
+
+ instance = CassandraDB(
+ contact_points=['127.0.0.1'],
+ port=9042,
+ username='testuser',
+ password='testpass',
+ keyspace='test_keyspace',
+ collection_name='test_collection',
+ embedding_model_dims=128,
+ )
+
+ assert instance.contact_points == ['127.0.0.1']
+ assert instance.port == 9042
+ assert instance.username == 'testuser'
+ assert instance.keyspace == 'test_keyspace'
+ assert instance.collection_name == 'test_collection'
+ assert instance.embedding_model_dims == 128
+
+
+def test_create_col(cassandra_instance):
+ """Test collection creation."""
+ cassandra_instance.create_col(name="new_collection", vector_size=256)
+
+ # Verify that execute was called (table creation)
+ assert cassandra_instance.session.execute.called
+
+
+def test_insert(cassandra_instance):
+ """Test vector insertion."""
+ vectors = [[0.1, 0.2, 0.3], [0.4, 0.5, 0.6]]
+ payloads = [{"text": "test1"}, {"text": "test2"}]
+ ids = ["id1", "id2"]
+
+ # Mock prepared statement
+ mock_prepared = Mock()
+ cassandra_instance.session.prepare = Mock(return_value=mock_prepared)
+
+ cassandra_instance.insert(vectors=vectors, payloads=payloads, ids=ids)
+
+ assert cassandra_instance.session.prepare.called
+ assert cassandra_instance.session.execute.called
+
+
+def test_search(cassandra_instance):
+ """Test vector search."""
+ # Mock the database response
+ mock_row1 = Mock()
+ mock_row1.id = 'id1'
+ mock_row1.vector = [0.1, 0.2, 0.3]
+ mock_row1.payload = json.dumps({"text": "test1"})
+
+ mock_row2 = Mock()
+ mock_row2.id = 'id2'
+ mock_row2.vector = [0.4, 0.5, 0.6]
+ mock_row2.payload = json.dumps({"text": "test2"})
+
+ cassandra_instance.session.execute = Mock(return_value=[mock_row1, mock_row2])
+
+ query_vector = [0.2, 0.3, 0.4]
+ results = cassandra_instance.search(query="test", vectors=query_vector, limit=5)
+
+ assert isinstance(results, list)
+ assert len(results) <= 5
+ assert cassandra_instance.session.execute.called
+
+
+def test_delete(cassandra_instance):
+ """Test vector deletion."""
+ mock_prepared = Mock()
+ cassandra_instance.session.prepare = Mock(return_value=mock_prepared)
+
+ cassandra_instance.delete(vector_id="test_id")
+
+ assert cassandra_instance.session.prepare.called
+ assert cassandra_instance.session.execute.called
+
+
+def test_update(cassandra_instance):
+ """Test vector update."""
+ new_vector = [0.7, 0.8, 0.9]
+ new_payload = {"text": "updated"}
+
+ mock_prepared = Mock()
+ cassandra_instance.session.prepare = Mock(return_value=mock_prepared)
+
+ cassandra_instance.update(vector_id="test_id", vector=new_vector, payload=new_payload)
+
+ assert cassandra_instance.session.prepare.called
+ assert cassandra_instance.session.execute.called
+
+
+def test_get(cassandra_instance):
+ """Test retrieving a vector by ID."""
+ # Mock the database response
+ mock_row = Mock()
+ mock_row.id = 'test_id'
+ mock_row.vector = [0.1, 0.2, 0.3]
+ mock_row.payload = json.dumps({"text": "test"})
+
+ mock_result = Mock()
+ mock_result.one = Mock(return_value=mock_row)
+
+ mock_prepared = Mock()
+ cassandra_instance.session.prepare = Mock(return_value=mock_prepared)
+ cassandra_instance.session.execute = Mock(return_value=mock_result)
+
+ result = cassandra_instance.get(vector_id="test_id")
+
+ assert result is not None
+ assert isinstance(result, OutputData)
+ assert result.id == "test_id"
+
+
+def test_list_cols(cassandra_instance):
+ """Test listing collections."""
+ # Mock the database response
+ mock_row1 = Mock()
+ mock_row1.table_name = "collection1"
+
+ mock_row2 = Mock()
+ mock_row2.table_name = "collection2"
+
+ cassandra_instance.session.execute = Mock(return_value=[mock_row1, mock_row2])
+
+ collections = cassandra_instance.list_cols()
+
+ assert isinstance(collections, list)
+ assert len(collections) == 2
+ assert "collection1" in collections
+
+
+def test_delete_col(cassandra_instance):
+ """Test collection deletion."""
+ cassandra_instance.delete_col()
+
+ assert cassandra_instance.session.execute.called
+
+
+def test_col_info(cassandra_instance):
+ """Test getting collection information."""
+ # Mock the database response
+ mock_row = Mock()
+ mock_row.count = 100
+
+ mock_result = Mock()
+ mock_result.one = Mock(return_value=mock_row)
+
+ cassandra_instance.session.execute = Mock(return_value=mock_result)
+
+ info = cassandra_instance.col_info()
+
+ assert isinstance(info, dict)
+ assert 'name' in info
+ assert 'keyspace' in info
+
+
+def test_list(cassandra_instance):
+ """Test listing vectors."""
+ # Mock the database response
+ mock_row = Mock()
+ mock_row.id = 'id1'
+ mock_row.vector = [0.1, 0.2, 0.3]
+ mock_row.payload = json.dumps({"text": "test1"})
+
+ cassandra_instance.session.execute = Mock(return_value=[mock_row])
+
+ results = cassandra_instance.list(limit=10)
+
+ assert isinstance(results, list)
+ assert len(results) > 0
+
+
+def test_reset(cassandra_instance):
+ """Test resetting the collection."""
+ cassandra_instance.reset()
+
+ assert cassandra_instance.session.execute.called
+
+
+def test_astra_db_connection(mock_cluster, mock_session):
+ """Test connection with DataStax Astra DB secure connect bundle."""
+ with patch('mem0.vector_stores.cassandra.Cluster') as mock_cluster_class:
+ mock_cluster_class.return_value = mock_cluster
+
+ instance = CassandraDB(
+ contact_points=['127.0.0.1'],
+ port=9042,
+ username='testuser',
+ password='testpass',
+ keyspace='test_keyspace',
+ collection_name='test_collection',
+ embedding_model_dims=128,
+ secure_connect_bundle='/path/to/bundle.zip'
+ )
+
+ assert instance.secure_connect_bundle == '/path/to/bundle.zip'
+
+
+def test_search_with_filters(cassandra_instance):
+ """Test vector search with filters."""
+ # Mock the database response
+ mock_row1 = Mock()
+ mock_row1.id = 'id1'
+ mock_row1.vector = [0.1, 0.2, 0.3]
+ mock_row1.payload = json.dumps({"text": "test1", "category": "A"})
+
+ mock_row2 = Mock()
+ mock_row2.id = 'id2'
+ mock_row2.vector = [0.4, 0.5, 0.6]
+ mock_row2.payload = json.dumps({"text": "test2", "category": "B"})
+
+ cassandra_instance.session.execute = Mock(return_value=[mock_row1, mock_row2])
+
+ query_vector = [0.2, 0.3, 0.4]
+ results = cassandra_instance.search(
+ query="test",
+ vectors=query_vector,
+ limit=5,
+ filters={"category": "A"}
+ )
+
+ assert isinstance(results, list)
+ # Should only return filtered results
+ for result in results:
+ assert result.payload.get("category") == "A"
+
+
+def test_output_data_model():
+ """Test OutputData model."""
+ data = OutputData(
+ id="test_id",
+ score=0.95,
+ payload={"text": "test"}
+ )
+
+ assert data.id == "test_id"
+ assert data.score == 0.95
+ assert data.payload == {"text": "test"}
+
+
+def test_insert_without_ids(cassandra_instance):
+ """Test vector insertion without providing IDs."""
+ vectors = [[0.1, 0.2, 0.3], [0.4, 0.5, 0.6]]
+ payloads = [{"text": "test1"}, {"text": "test2"}]
+
+ mock_prepared = Mock()
+ cassandra_instance.session.prepare = Mock(return_value=mock_prepared)
+
+ cassandra_instance.insert(vectors=vectors, payloads=payloads)
+
+ assert cassandra_instance.session.prepare.called
+ assert cassandra_instance.session.execute.called
+
+
+def test_insert_without_payloads(cassandra_instance):
+ """Test vector insertion without providing payloads."""
+ vectors = [[0.1, 0.2, 0.3], [0.4, 0.5, 0.6]]
+ ids = ["id1", "id2"]
+
+ mock_prepared = Mock()
+ cassandra_instance.session.prepare = Mock(return_value=mock_prepared)
+
+ cassandra_instance.insert(vectors=vectors, ids=ids)
+
+ assert cassandra_instance.session.prepare.called
+ assert cassandra_instance.session.execute.called
+