feat(vector-stores): add Cassandra provider to TypeScript OSS SDK (#5823)

Co-authored-by: kartik-mem0 <kartik.labhshetwar@mem0.ai>
This commit is contained in:
Rod Boev
2026-07-06 12:27:32 -04:00
committed by GitHub
parent fec7cdf118
commit 7fb3feb5cd
9 changed files with 1183 additions and 7 deletions
+89 -7
View File
@@ -7,7 +7,8 @@ description: "Use Apache Cassandra as a distributed vector store in Mem0 with se
### Usage
```python
<CodeGroup>
```python Python
import os
from mem0 import Memory
@@ -37,11 +38,43 @@ messages = [
m.add(messages, user_id="alice", metadata={"category": "movies"})
```
```typescript TypeScript
import { Memory } from 'mem0ai/oss';
// Set OPENAI_API_KEY in your environment for the default embedder
const config = {
vectorStore: {
provider: 'cassandra',
config: {
contactPoints: ['127.0.0.1'],
localDataCenter: 'datacenter1', // required with contactPoints; "datacenter1" is the default for a single-node cluster
port: 9042,
username: 'cassandra',
password: 'cassandra',
keyspace: 'mem0',
collectionName: 'memories',
},
},
};
const memory = new Memory(config);
const 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."}
]
await memory.add(messages, { userId: "alice", metadata: { category: "movies" } });
```
</CodeGroup>
#### Using DataStax Astra DB
For managed Cassandra with DataStax Astra DB:
```python
<CodeGroup>
```python Python
config = {
"vector_store": {
"provider": "cassandra",
@@ -57,8 +90,24 @@ config = {
}
```
```typescript TypeScript
const config = {
vectorStore: {
provider: 'cassandra',
config: {
username: 'token',
password: 'AstraCS:...', // Your Astra DB application token
keyspace: 'mem0',
collectionName: 'memories',
secureConnectBundle: '/path/to/secure-connect-bundle.zip',
},
},
};
```
</CodeGroup>
<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.
When using DataStax Astra DB, provide the secure connect bundle path. Contact points and `localDataCenter` are not needed when a secure connect bundle is provided.
</Note>
### Config
@@ -78,6 +127,10 @@ Here are the parameters available for configuring Apache Cassandra:
| `protocol_version` | CQL protocol version | `4` |
| `load_balancing_policy` | Custom load balancing policy | `None` |
<Note>
The TypeScript SDK uses camelCase keys: `contactPoints`, `collectionName`, `embeddingModelDims`, `secureConnectBundle`, `protocolVersion`, and `loadBalancingPolicy`. It also requires `localDataCenter` (for example, `datacenter1`) when you connect with `contactPoints` instead of a secure connect bundle. The Node.js driver needs this to route queries; it has no default.
</Note>
### Setup
#### Option 1: Local Cassandra Setup using Docker:
@@ -139,14 +192,20 @@ brew services start cassandra
cqlsh
```
### Python Client Installation
### Client Installation
Install the required Python package:
Install the driver for your SDK:
```bash
<CodeGroup>
```bash Python
pip install cassandra-driver
```
```bash TypeScript
npm install cassandra-driver
```
</CodeGroup>
### Performance Considerations
- **Replication Factor**: For production, use replication factor of at least 3
@@ -156,7 +215,8 @@ pip install cassandra-driver
### Advanced Configuration
```python
<CodeGroup>
```python Python
from cassandra.policies import DCAwareRoundRobinPolicy
config = {
@@ -176,6 +236,28 @@ config = {
}
```
```typescript TypeScript
// The Node.js driver routes to localDataCenter by default, so set it to your
// primary DC for datacenter-aware routing. Pass loadBalancingPolicy only when
// you need a custom policy from the cassandra-driver package.
const config = {
vectorStore: {
provider: 'cassandra',
config: {
contactPoints: ['node1.example.com', 'node2.example.com', 'node3.example.com'],
localDataCenter: 'DC1',
port: 9042,
username: 'mem0_user',
password: 'secure_password',
keyspace: 'mem0_prod',
collectionName: 'memories',
protocolVersion: 4,
},
},
};
```
</CodeGroup>
<Warning>
For production use, configure appropriate replication strategies and consistency levels based on your availability and consistency requirements.
</Warning>
+1
View File
@@ -121,6 +121,7 @@
"@types/jest": "29.5.14",
"@types/pg": "8.11.0",
"better-sqlite3": "^12.6.2",
"cassandra-driver": "4.8.0",
"cloudflare": "^4.2.0",
"fastembed": "^2.1.0",
"groq-sdk": "0.3.0",
+24
View File
@@ -74,6 +74,9 @@ importers:
better-sqlite3:
specifier: ^12.6.2
version: 12.10.0
cassandra-driver:
specifier: 4.8.0
version: 4.8.0
cloudflare:
specifier: ^4.2.0
version: 4.5.0
@@ -1348,6 +1351,10 @@ packages:
engines: {node: '>=0.4.0'}
hasBin: true
adm-zip@0.5.17:
resolution: {integrity: sha512-+Ut8d9LLqwEvHHJl1+PIHqoyDxFgVN847JTVM3Izi3xHDWPE4UtzzXysMZQs64DMcrJfBeS/uoEP4AD3HQHnQQ==}
engines: {node: '>=12.0'}
afinn-165-financialmarketnews@3.0.0:
resolution: {integrity: sha512-0g9A1S3ZomFIGDTzZ0t6xmv4AuokBvBmpes8htiyHpH7N4xDmvSQL6UxL/Zcs2ypRb3VwgCscaD8Q3zEawKYhw==}
@@ -1559,6 +1566,10 @@ packages:
caniuse-lite@1.0.30001799:
resolution: {integrity: sha512-hG1bReV+OUU+MOqK4t/ZWI0tZOyz3rqS9XuhOUz1cIcbwBKjOyJEJuw9ER5JuNyqxNk8u/JUVbGibBOL1yrjFw==}
cassandra-driver@4.8.0:
resolution: {integrity: sha512-HritfMGq9V7SuESeSodHvArs0mLuMk7uh+7hQK2lqdvXrvm50aWxb4RPxkK3mPDdsgHjJ427xNRFITMH2ei+Sw==}
engines: {node: '>=18'}
chalk@4.1.2:
resolution: {integrity: sha512-oKnbhFyRIXpUuez8iBMmyEa4nbj4IOQyuhc/wy9kY7/WVPcwIO9VA668Pu8RkO7+0G76SLROeyw9CpQ061i4mA==}
engines: {node: '>=10'}
@@ -2466,6 +2477,9 @@ packages:
lodash.once@4.1.1:
resolution: {integrity: sha512-Sb487aTOCr9drQVL8pIxOzVhafOjZN9UU54hiN8PU3uAiSV7lx1yYNpbNmex2PK6dSJoNTSJUUswT651yww3Mg==}
long@5.2.5:
resolution: {integrity: sha512-e0r9YBBgNCq1D1o5Dp8FMH0N5hsFtXDBiVa0qoJPHpakvZkmDKPRoGffZJII/XsHvj9An9blm+cRJ01yQqU+Dw==}
long@5.3.2:
resolution: {integrity: sha512-mNAgZ1GmyNhD7AuqnTG3/VQ26o760+ZYBPKjPvugO8+nLbYfX6TVpJPseBvopbdY+qpZ/lKUnmEc1LeZYS3QAA==}
@@ -5096,6 +5110,8 @@ snapshots:
acorn@8.16.0: {}
adm-zip@0.5.17: {}
afinn-165-financialmarketnews@3.0.0: {}
afinn-165@2.0.2: {}
@@ -5315,6 +5331,12 @@ snapshots:
caniuse-lite@1.0.30001799: {}
cassandra-driver@4.8.0:
dependencies:
'@types/node': 18.19.130
adm-zip: 0.5.17
long: 5.2.5
chalk@4.1.2:
dependencies:
ansi-styles: 4.3.0
@@ -6412,6 +6434,8 @@ snapshots:
lodash.once@4.1.1: {}
long@5.2.5: {}
long@5.3.2: {}
lru-cache@10.4.3: {}
+1
View File
@@ -32,5 +32,6 @@ export * from "./vector_stores/langchain";
export * from "./vector_stores/vectorize";
export * from "./vector_stores/azure_ai_search";
export * from "./vector_stores/pgvector";
export * from "./vector_stores/cassandra";
export * from "./vector_stores/s3_vectors";
export * from "./utils/factory";
+3
View File
@@ -42,6 +42,7 @@ import { LangchainEmbedder } from "../embeddings/langchain";
import { LangchainVectorStore } from "../vector_stores/langchain";
import { AzureAISearch } from "../vector_stores/azure_ai_search";
import { PGVector } from "../vector_stores/pgvector";
import { CassandraDB } from "../vector_stores/cassandra";
import { PineconeDB } from "../vector_stores/pinecone";
import { S3Vectors } from "../vector_stores/s3_vectors";
@@ -130,6 +131,8 @@ export class VectorStoreFactory {
return new AzureAISearch(config as any);
case "pgvector":
return new PGVector(config as any);
case "cassandra":
return new CassandraDB(config as any);
case "pinecone":
return new PineconeDB(config as any);
case "s3-vectors":
@@ -0,0 +1,598 @@
import cassandra from "cassandra-driver";
import { VectorStore } from "./base";
import { SearchFilters, VectorStoreConfig, VectorStoreResult } from "../types";
const MIGRATION_ROW_ID = "mem0-user";
const SAFE_IDENTIFIER_RE = /^[A-Za-z_][A-Za-z0-9_]{0,127}$/;
interface CassandraConfig extends VectorStoreConfig {
contactPoints?: string[];
port?: number;
username?: string;
password?: string;
keyspace?: string;
collectionName?: string;
embeddingModelDims?: number;
secureConnectBundle?: string;
localDataCenter?: string;
protocolVersion?: number;
loadBalancingPolicy?: any;
client?: CassandraClientLike;
driver?: typeof cassandra;
}
interface CassandraClientLike {
connect?(): Promise<void>;
execute(
query: string,
params?: any[],
options?: Record<string, any>,
): Promise<{ rows?: any[]; pageState?: string | null }>;
}
interface CassandraVector {
id: string;
vector: number[];
payload: Record<string, any>;
}
export class CassandraDB implements VectorStore {
private static readonly PAGE_SIZE = 500;
private readonly driver: typeof cassandra;
private readonly contactPoints?: string[];
private readonly port: number;
private readonly username?: string;
private readonly password?: string;
private readonly keyspace: string;
private readonly collectionName: string;
private readonly dimension: number;
private readonly secureConnectBundle?: string;
private readonly localDataCenter?: string;
private readonly protocolVersion?: number;
private readonly loadBalancingPolicy?: any;
private client?: CassandraClientLike;
private _initPromise?: Promise<void>;
constructor(config: CassandraConfig) {
this.driver = config.driver || cassandra;
this.contactPoints = config.contactPoints;
this.port = config.port || 9042;
this.username = config.username;
this.password = config.password;
this.keyspace = this.validateIdentifier(
config.keyspace || "mem0",
"keyspace",
);
this.collectionName = this.validateIdentifier(
config.collectionName || "memories",
"collectionName",
);
this.dimension = config.embeddingModelDims || config.dimension || 1536;
this.secureConnectBundle = config.secureConnectBundle;
this.localDataCenter = config.localDataCenter;
this.protocolVersion = config.protocolVersion;
this.loadBalancingPolicy = config.loadBalancingPolicy;
this.client = config.client;
this.initialize().catch(console.error);
}
async initialize(): Promise<void> {
if (!this._initPromise) {
this._initPromise = this._doInitialize();
}
return this._initPromise;
}
private async _doInitialize(): Promise<void> {
if (!this.client) {
this.client = this.createClient();
}
if (typeof this.client.connect === "function") {
await this.client.connect();
}
await this.client.execute(`
CREATE KEYSPACE IF NOT EXISTS ${this.keyspace}
WITH replication = {'class': 'SimpleStrategy', 'replication_factor': 1}
`);
await this.client.execute(`
CREATE TABLE IF NOT EXISTS ${this.keyspace}.${this.collectionName} (
id text PRIMARY KEY,
vector list<float>,
payload text
)
`);
await this.client.execute(`
CREATE TABLE IF NOT EXISTS ${this.keyspace}.memory_migrations (
id text PRIMARY KEY,
user_id text
)
`);
}
async insert(
vectors: number[][],
ids: string[],
payloads: Record<string, any>[],
): Promise<void> {
await this.initialize();
this.assertBatchDimensions(vectors, "Vector");
const query = `
INSERT INTO ${this.keyspace}.${this.collectionName} (id, vector, payload)
VALUES (?, ?, ?)
`;
for (let index = 0; index < vectors.length; index += 1) {
await this.client!.execute(
query,
[ids[index], vectors[index], JSON.stringify(payloads[index] || {})],
{ prepare: true },
);
}
}
async keywordSearch(): Promise<null> {
return null;
}
async search(
query: number[],
topK: number = 5,
filters?: SearchFilters,
): Promise<VectorStoreResult[]> {
await this.initialize();
this.assertVectorDimension(query, "Query");
const scored: VectorStoreResult[] = [];
await this.scanRows(
`
SELECT id, vector, payload
FROM ${this.keyspace}.${this.collectionName}
`,
async (row) => {
const vector = this.normalizeVector(row.vector);
const payload = this.parsePayload(row.payload);
if (!vector || vector.length !== this.dimension) {
return;
}
const item: CassandraVector = {
id: String(row.id),
vector,
payload,
};
if (!this.filterVector(item, filters)) {
return;
}
this.pushTopResult(
scored,
{
id: item.id,
payload: item.payload,
score: this.cosineSimilarity(query, item.vector),
},
topK,
);
},
Math.max(topK, CassandraDB.PAGE_SIZE),
);
return scored;
}
async get(vectorId: string): Promise<VectorStoreResult | null> {
await this.initialize();
const result = await this.client!.execute(
`
SELECT id, payload
FROM ${this.keyspace}.${this.collectionName}
WHERE id = ?
`,
[vectorId],
{ prepare: true },
);
const row = result.rows?.[0];
if (!row) {
return null;
}
return {
id: String(row.id),
payload: this.parsePayload(row.payload),
};
}
async update(
vectorId: string,
vector: number[],
payload: Record<string, any>,
): Promise<void> {
await this.initialize();
this.assertVectorDimension(vector, "Vector");
await this.client!.execute(
`
INSERT INTO ${this.keyspace}.${this.collectionName} (id, vector, payload)
VALUES (?, ?, ?)
`,
[vectorId, vector, JSON.stringify(payload || {})],
{ prepare: true },
);
}
async delete(vectorId: string): Promise<void> {
await this.initialize();
await this.client!.execute(
`
DELETE FROM ${this.keyspace}.${this.collectionName}
WHERE id = ?
`,
[vectorId],
{ prepare: true },
);
}
async deleteCol(): Promise<void> {
await this.initialize();
await this.client!.execute(`
DROP TABLE IF EXISTS ${this.keyspace}.${this.collectionName}
`);
await this.client!.execute(`
CREATE TABLE IF NOT EXISTS ${this.keyspace}.${this.collectionName} (
id text PRIMARY KEY,
vector list<float>,
payload text
)
`);
}
async list(
filters?: SearchFilters,
topK: number = 100,
): Promise<[VectorStoreResult[], number]> {
await this.initialize();
const rows: VectorStoreResult[] = [];
let total = 0;
await this.scanRows(
`
SELECT id, payload
FROM ${this.keyspace}.${this.collectionName}
`,
async (row) => {
const item: CassandraVector = {
id: String(row.id),
vector: [],
payload: this.parsePayload(row.payload),
};
if (!this.filterVector(item, filters)) {
return;
}
total += 1;
if (rows.length < topK) {
rows.push({
id: item.id,
payload: item.payload,
});
}
},
CassandraDB.PAGE_SIZE,
);
return [rows, total];
}
async getUserId(): Promise<string> {
await this.initialize();
const result = await this.client!.execute(
`
SELECT user_id
FROM ${this.keyspace}.memory_migrations
WHERE id = ?
`,
[MIGRATION_ROW_ID],
{ prepare: true },
);
const existing = result.rows?.[0]?.user_id;
if (typeof existing === "string" && existing.length > 0) {
return existing;
}
const userId =
Math.random().toString(36).substring(2, 15) +
Math.random().toString(36).substring(2, 15);
await this.setUserId(userId);
return userId;
}
async setUserId(userId: string): Promise<void> {
await this.initialize();
await this.client!.execute(
`
INSERT INTO ${this.keyspace}.memory_migrations (id, user_id)
VALUES (?, ?)
`,
[MIGRATION_ROW_ID, userId],
{ prepare: true },
);
}
private createClient(): CassandraClientLike {
const clientConfig: Record<string, any> = {};
if (this.secureConnectBundle) {
clientConfig.cloud = {
secureConnectBundle: this.secureConnectBundle,
};
} else {
if (!this.contactPoints || this.contactPoints.length === 0) {
throw new Error(
"Cassandra vector store requires contactPoints when secureConnectBundle is not provided.",
);
}
if (!this.localDataCenter) {
throw new Error(
"Cassandra vector store requires localDataCenter when secureConnectBundle is not provided.",
);
}
clientConfig.contactPoints = this.contactPoints;
clientConfig.localDataCenter = this.localDataCenter;
clientConfig.protocolOptions = {
port: this.port,
};
}
if (this.protocolVersion !== undefined) {
clientConfig.protocolOptions = {
...(clientConfig.protocolOptions || {}),
maxVersion: this.protocolVersion,
};
}
if (this.loadBalancingPolicy) {
clientConfig.policies = {
loadBalancing: this.loadBalancingPolicy,
};
}
if (this.username && this.password) {
clientConfig.authProvider = new this.driver.auth.PlainTextAuthProvider(
this.username,
this.password,
);
}
return new this.driver.Client(clientConfig);
}
private validateIdentifier(name: string, label: string): string {
if (!SAFE_IDENTIFIER_RE.test(name)) {
throw new Error(
`Invalid ${label} '${name}': only letters, digits, and underscores are allowed, ` +
"must start with a letter or underscore, and be at most 128 characters.",
);
}
return name;
}
private cosineSimilarity(left: number[], right: number[]): number {
let dotProduct = 0;
let leftNorm = 0;
let rightNorm = 0;
for (let index = 0; index < left.length; index += 1) {
dotProduct += left[index] * right[index];
leftNorm += left[index] * left[index];
rightNorm += right[index] * right[index];
}
if (leftNorm === 0 || rightNorm === 0) {
return 0;
}
return dotProduct / (Math.sqrt(leftNorm) * Math.sqrt(rightNorm));
}
private normalizeVector(rawValue: any): number[] | undefined {
if (Array.isArray(rawValue)) {
return rawValue.map((value) => Number(value));
}
return undefined;
}
private parsePayload(rawValue: any): Record<string, any> {
if (typeof rawValue === "string") {
try {
const parsed = JSON.parse(rawValue);
if (parsed && typeof parsed === "object" && !Array.isArray(parsed)) {
return parsed;
}
} catch (error) {
return {};
}
}
if (rawValue && typeof rawValue === "object" && !Array.isArray(rawValue)) {
return rawValue;
}
return {};
}
private matchFieldCondition(
payload: Record<string, any>,
key: string,
value: any,
): boolean {
const payloadValue = payload[key];
if (typeof value !== "object" || value === null) {
if (value === "*") {
return true;
}
return payloadValue === value;
}
if (Array.isArray(value)) {
return value.includes(payloadValue);
}
if ("eq" in value) {
return payloadValue === value.eq;
}
if ("ne" in value) {
return payloadValue !== value.ne;
}
if ("gt" in value) {
return payloadValue > value.gt;
}
if ("gte" in value) {
return payloadValue >= value.gte;
}
if ("lt" in value) {
return payloadValue < value.lt;
}
if ("lte" in value) {
return payloadValue <= value.lte;
}
if ("in" in value) {
return Array.isArray(value.in) && value.in.includes(payloadValue);
}
if ("nin" in value) {
return !Array.isArray(value.nin) || !value.nin.includes(payloadValue);
}
if ("contains" in value) {
return (
typeof payloadValue === "string" &&
payloadValue.includes(value.contains)
);
}
if ("icontains" in value) {
return (
typeof payloadValue === "string" &&
payloadValue.toLowerCase().includes(value.icontains.toLowerCase())
);
}
return payloadValue === value;
}
private filterVector(
vector: CassandraVector,
filters?: SearchFilters,
): boolean {
if (!filters || Object.keys(filters).length === 0) {
return true;
}
const keyMap: Record<string, string> = {
$and: "AND",
$or: "OR",
$not: "NOT",
};
const normalized: Record<string, any> = {};
for (const [key, value] of Object.entries(filters)) {
const normalizedKey = keyMap[key] || key;
if (!(normalizedKey in normalized)) {
normalized[normalizedKey] = value;
}
}
for (const [key, value] of Object.entries(normalized)) {
if (key === "AND") {
if (!Array.isArray(value)) {
throw new Error(
`AND filter value must be a list of filter dicts, got ${typeof value}`,
);
}
if (
!value.every((entry: SearchFilters) =>
this.filterVector(vector, entry),
)
) {
return false;
}
} else if (key === "OR") {
if (!Array.isArray(value)) {
throw new Error(
`OR filter value must be a list of filter dicts, got ${typeof value}`,
);
}
if (
!value.some((entry: SearchFilters) =>
this.filterVector(vector, entry),
)
) {
return false;
}
} else if (key === "NOT") {
if (!Array.isArray(value)) {
throw new Error(
`NOT filter value must be a list of filter dicts, got ${typeof value}`,
);
}
if (
!value.every(
(entry: SearchFilters) => !this.filterVector(vector, entry),
)
) {
return false;
}
} else if (!this.matchFieldCondition(vector.payload, key, value)) {
return false;
}
}
return true;
}
private assertVectorDimension(vector: number[], label: string): void {
if (vector.length !== this.dimension) {
throw new Error(
`${label} dimension mismatch. Expected ${this.dimension}, got ${vector.length}`,
);
}
}
private assertBatchDimensions(vectors: number[][], label: string): void {
for (const vector of vectors) {
this.assertVectorDimension(vector, label);
}
}
private async scanRows(
query: string,
onRow: (row: any) => Promise<void>,
fetchSize: number,
): Promise<void> {
let pageState: string | undefined;
do {
const result = await this.client!.execute(query, [], {
autoPage: false,
fetchSize,
pageState,
});
for (const row of result.rows || []) {
await onRow(row);
}
pageState = result.pageState || undefined;
} while (pageState);
}
private pushTopResult(
results: VectorStoreResult[],
candidate: VectorStoreResult,
topK: number,
): void {
results.push(candidate);
results.sort((left, right) => (right.score || 0) - (left.score || 0));
if (results.length > topK) {
results.length = topK;
}
}
}
@@ -158,6 +158,11 @@ jest.mock("../src/vector_stores/pgvector", () => ({
.fn()
.mockImplementation((config) => ({ type: "pgvector", config })),
}));
jest.mock("../src/vector_stores/cassandra", () => ({
CassandraDB: jest
.fn()
.mockImplementation((config) => ({ type: "cassandra", config })),
}));
jest.mock("../src/vector_stores/s3_vectors", () => ({
S3Vectors: jest
.fn()
@@ -289,6 +294,7 @@ describe("VectorStoreFactory", () => {
["vectorize"],
["azure-ai-search"],
["pgvector"],
["cassandra"],
["s3-vectors"],
["s3_vectors"],
])("creates vector store for provider '%s'", (provider) => {
@@ -719,6 +719,466 @@ describe("AzureAISearch – backward compat with mocked client", () => {
});
});
// ───────────────────────────────────────────────────────────────────────────
// Cassandra — mock client, test interface + idempotent init
// ───────────────────────────────────────────────────────────────────────────
describe("Cassandra – backward compat with mocked client", () => {
let CassandraDB: any;
beforeEach(() => {
jest.resetModules();
jest.doMock("cassandra-driver", () => {
const rows = new Map<
string,
{ id: string; vector: number[]; payload: string }
>();
const memoryRows = () =>
Array.from(rows.entries())
.filter(([key]) => key.startsWith("memories:"))
.map(([, row]) => row);
class MockClient {
connect = jest.fn().mockResolvedValue(undefined);
execute = jest
.fn()
.mockImplementation(
async (
query: string,
params: any[] = [],
options: Record<string, any> = {},
) => {
const normalized = query.replace(/\s+/g, " ").trim();
if (normalized.startsWith("CREATE KEYSPACE IF NOT EXISTS")) {
return { rows: [] };
}
if (normalized.startsWith("CREATE TABLE IF NOT EXISTS")) {
return { rows: [] };
}
if (
normalized.startsWith(
"INSERT INTO mem0.memories (id, vector, payload) VALUES (?, ?, ?)",
)
) {
rows.set(`memories:${params[0]}`, {
id: params[0],
vector: params[1],
payload: params[2],
});
return { rows: [] };
}
if (
normalized.startsWith(
"SELECT id, payload FROM mem0.memories WHERE id = ?",
)
) {
const row = rows.get(`memories:${params[0]}`);
return {
rows: row ? [{ id: row.id, payload: row.payload }] : [],
};
}
if (
normalized.startsWith(
"SELECT id, vector, payload FROM mem0.memories",
)
) {
return {
rows: memoryRows().map((row) => ({
...row,
})),
};
}
if (
normalized.startsWith("SELECT id, payload FROM mem0.memories")
) {
return {
rows: memoryRows().map((row) => ({
id: row.id,
payload: row.payload,
})),
};
}
if (normalized.startsWith("DROP TABLE IF EXISTS mem0.memories")) {
for (const key of Array.from(rows.keys())) {
if (key.startsWith("memories:")) {
rows.delete(key);
}
}
return { rows: [] };
}
if (
normalized.startsWith("DELETE FROM mem0.memories WHERE id = ?")
) {
rows.delete(`memories:${params[0]}`);
return { rows: [] };
}
if (
normalized.startsWith(
"INSERT INTO mem0.memory_migrations (id, user_id) VALUES (?, ?)",
)
) {
rows.set(`migrations:${params[0]}`, {
id: params[0],
vector: [0],
payload: JSON.stringify({ user_id: params[1] }),
});
return { rows: [] };
}
if (
normalized.startsWith(
"SELECT user_id FROM mem0.memory_migrations WHERE id = ?",
)
) {
const row = rows.get(`migrations:${params[0]}`);
if (!row) {
return { rows: [] };
}
return {
rows: [{ user_id: JSON.parse(row.payload).user_id }],
};
}
throw new Error(
`Unexpected Cassandra query: ${normalized} prepare=${options.prepare}`,
);
},
);
}
return {
__esModule: true,
default: {
Client: jest.fn().mockImplementation(() => new MockClient()),
auth: {
PlainTextAuthProvider: jest
.fn()
.mockImplementation((username: string, password: string) => ({
username,
password,
})),
},
},
};
});
CassandraDB = require("../src/vector_stores/cassandra").CassandraDB;
});
afterEach(() => {
jest.restoreAllMocks();
jest.resetModules();
});
it("implements full VectorStore interface", () => {
const store = new CassandraDB({
client: {
execute: jest.fn().mockResolvedValue({ rows: [] }),
},
collectionName: "memories",
dimension: 3,
});
expect(typeof store.insert).toBe("function");
expect(typeof store.search).toBe("function");
expect(typeof store.get).toBe("function");
expect(typeof store.update).toBe("function");
expect(typeof store.delete).toBe("function");
expect(typeof store.deleteCol).toBe("function");
expect(typeof store.list).toBe("function");
expect(typeof store.getUserId).toBe("function");
expect(typeof store.setUserId).toBe("function");
expect(typeof store.initialize).toBe("function");
});
it("initialize() is idempotent (same promise returned)", async () => {
const cassandraDriver = require("cassandra-driver");
const store = new CassandraDB({
contactPoints: ["127.0.0.1"],
localDataCenter: "datacenter1",
collectionName: "memories",
dimension: 3,
});
const p1 = store.initialize();
const p2 = store.initialize();
const p3 = store.initialize();
await Promise.all([p1, p2, p3]);
const clientInstance = cassandraDriver.default.Client.mock.results[0].value;
expect(clientInstance.connect).toHaveBeenCalledTimes(1);
});
it("shapes Cassandra writes and normalizes search results", async () => {
const cassandraDriver = require("cassandra-driver");
const store = new CassandraDB({
contactPoints: ["127.0.0.1"],
localDataCenter: "datacenter1",
collectionName: "memories",
dimension: 3,
});
await store.initialize();
await store.insert(
[[1, 0, 0]],
["id-1"],
[{ user_id: "u1", topic: "alpha" }],
);
const clientInstance = cassandraDriver.default.Client.mock.results[0].value;
expect(clientInstance.execute).toHaveBeenCalledWith(
expect.stringContaining(
"INSERT INTO mem0.memories (id, vector, payload)",
),
["id-1", [1, 0, 0], JSON.stringify({ user_id: "u1", topic: "alpha" })],
{ prepare: true },
);
const results = await store.search([1, 0, 0], 5, { user_id: "u1" });
expect(results).toEqual([
{
id: "id-1",
payload: { user_id: "u1", topic: "alpha" },
score: 1,
},
]);
});
it("roundtrips migration user ids", async () => {
const store = new CassandraDB({
contactPoints: ["127.0.0.1"],
localDataCenter: "datacenter1",
collectionName: "memories",
dimension: 3,
});
await store.setUserId("custom-user");
expect(await store.getUserId()).toBe("custom-user");
});
it("supports get, update, delete, and list", async () => {
const store = new CassandraDB({
contactPoints: ["127.0.0.1"],
localDataCenter: "datacenter1",
collectionName: "memories",
dimension: 3,
});
await store.insert(
[
[1, 0, 0],
[0, 1, 0],
],
["id-1", "id-2"],
[
{ user_id: "u1", topic: "alpha" },
{ user_id: "u2", topic: "beta" },
],
);
expect(await store.get("missing")).toBeNull();
expect(await store.get("id-1")).toEqual({
id: "id-1",
payload: { user_id: "u1", topic: "alpha" },
});
await store.update("id-1", [0, 0, 1], {
user_id: "u1",
topic: "gamma",
});
expect(await store.get("id-1")).toEqual({
id: "id-1",
payload: { user_id: "u1", topic: "gamma" },
});
const [listed, count] = await store.list({ user_id: "u1" }, 10);
expect(count).toBe(1);
expect(listed).toEqual([
{
id: "id-1",
payload: { user_id: "u1", topic: "gamma" },
},
]);
await store.delete("id-2");
expect(await store.get("id-2")).toBeNull();
await store.deleteCol();
const [afterDrop, afterDropCount] = await store.list(undefined, 10);
expect(afterDrop).toEqual([]);
expect(afterDropCount).toBe(0);
});
it("scans paged search and list results", async () => {
const execute = jest
.fn()
.mockImplementation(
async (
query: string,
_params: any[] = [],
options: Record<string, any> = {},
) => {
const normalized = query.replace(/\s+/g, " ").trim();
if (normalized.startsWith("CREATE KEYSPACE IF NOT EXISTS")) {
return { rows: [] };
}
if (normalized.startsWith("CREATE TABLE IF NOT EXISTS")) {
return { rows: [] };
}
if (
normalized.startsWith(
"SELECT id, vector, payload FROM mem0.memories",
)
) {
if (!options.pageState) {
return {
rows: [
{
id: "id-1",
vector: [1, 0, 0],
payload: JSON.stringify({ user_id: "u1", topic: "alpha" }),
},
],
pageState: "page-2",
};
}
return {
rows: [
{
id: "id-2",
vector: [0, 1, 0],
payload: JSON.stringify({ user_id: "u2", topic: "beta" }),
},
],
pageState: null,
};
}
if (normalized.startsWith("SELECT id, payload FROM mem0.memories")) {
if (!options.pageState) {
return {
rows: [
{
id: "id-1",
payload: JSON.stringify({ user_id: "u1", topic: "alpha" }),
},
],
pageState: "page-2",
};
}
return {
rows: [
{
id: "id-2",
payload: JSON.stringify({ user_id: "u2", topic: "beta" }),
},
],
pageState: null,
};
}
if (
normalized.startsWith(
"SELECT user_id FROM mem0.memory_migrations WHERE id = ?",
)
) {
return { rows: [] };
}
throw new Error(`Unexpected Cassandra query: ${normalized}`);
},
);
const store = new CassandraDB({
client: { execute },
collectionName: "memories",
dimension: 3,
});
const searchResults = await store.search([1, 0, 0], 5);
expect(searchResults).toEqual([
{
id: "id-1",
payload: { user_id: "u1", topic: "alpha" },
score: 1,
},
{
id: "id-2",
payload: { user_id: "u2", topic: "beta" },
score: 0,
},
]);
const [listed, count] = await store.list(undefined, 10);
expect(count).toBe(2);
expect(listed).toEqual([
{
id: "id-1",
payload: { user_id: "u1", topic: "alpha" },
},
{
id: "id-2",
payload: { user_id: "u2", topic: "beta" },
},
]);
expect(execute).toHaveBeenCalledWith(
expect.stringContaining("SELECT id, vector, payload"),
[],
expect.objectContaining({
autoPage: false,
fetchSize: 500,
pageState: undefined,
}),
);
expect(execute).toHaveBeenCalledWith(
expect.stringContaining("SELECT id, vector, payload"),
[],
expect.objectContaining({
autoPage: false,
fetchSize: 500,
pageState: "page-2",
}),
);
expect(execute).toHaveBeenCalledWith(
expect.stringContaining("SELECT id, payload"),
[],
expect.objectContaining({
autoPage: false,
fetchSize: 500,
pageState: "page-2",
}),
);
});
it("rejects unsafe identifiers", () => {
expect(
() =>
new CassandraDB({
client: {
execute: jest.fn().mockResolvedValue({ rows: [] }),
},
keyspace: "bad-name",
collectionName: "memories",
dimension: 3,
}),
).toThrow("Invalid keyspace");
});
});
// ───────────────────────────────────────────────────────────────────────────
// 6. S3 Vectors — mock AWS client, test interface + init
// ───────────────────────────────────────────────────────────────────────────
+1
View File
@@ -10,6 +10,7 @@ const external = [
"pg",
"zod",
"better-sqlite3",
"cassandra-driver",
"@pinecone-database/pinecone",
"@qdrant/js-client-rest",
"redis",