Add Apache Cassandra vector store support (#3578)

This commit is contained in:
Faizan Habib
2025-10-17 23:48:49 +05:30
committed by GitHub
parent 7afbaae7a3
commit 9ef644b95e
8 changed files with 1075 additions and 1 deletions
+181
View File
@@ -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"
}
}
}
```
<Note>
When using DataStax Astra DB, provide the secure connect bundle path. The contact_points parameter is ignored when a secure connect bundle is provided.
</Note>
### 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
<Tip>
For production deployments, use DataStax Astra DB for fully managed Cassandra with automatic scaling, backups, and security.
</Tip>
#### 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')
}
}
}
```
<Warning>
For production use, configure appropriate replication strategies and consistency levels based on your availability and consistency requirements.
</Warning>
+2 -1
View File
@@ -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",
+77
View File
@@ -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
+1
View File
@@ -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",
}
+496
View File
@@ -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<float> and payload as text (JSON)
query = f"""
CREATE TABLE IF NOT EXISTS {self.keyspace}.{self.collection_name} (
id text PRIMARY KEY,
vector list<float>,
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<float>,
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
+1
View File
@@ -18,6 +18,7 @@ class VectorStoreConfig(BaseModel):
"mongodb": "MongoDBConfig",
"milvus": "MilvusDBConfig",
"baidu": "BaiduDBConfig",
"cassandra": "CassandraConfig",
"neptune": "NeptuneAnalyticsConfig",
"upstash_vector": "UpstashVectorConfig",
"azure_ai_search": "AzureAISearchConfig",
+1
View File
@@ -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",
+316
View File
@@ -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