diff --git a/docs/components/vectordbs/dbs/cassandra.mdx b/docs/components/vectordbs/dbs/cassandra.mdx index 02613f61a..4cbe7d602 100644 --- a/docs/components/vectordbs/dbs/cassandra.mdx +++ b/docs/components/vectordbs/dbs/cassandra.mdx @@ -7,7 +7,8 @@ description: "Use Apache Cassandra as a distributed vector store in Mem0 with se ### Usage -```python + +```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" } }); +``` + + #### Using DataStax Astra DB For managed Cassandra with DataStax Astra DB: -```python + +```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', + }, + }, +}; +``` + + -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. ### 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` | + +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. + + ### 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 + +```bash Python pip install cassandra-driver ``` +```bash TypeScript +npm install cassandra-driver +``` + + ### Performance Considerations - **Replication Factor**: For production, use replication factor of at least 3 @@ -156,7 +215,8 @@ pip install cassandra-driver ### Advanced Configuration -```python + +```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, + }, + }, +}; +``` + + For production use, configure appropriate replication strategies and consistency levels based on your availability and consistency requirements. diff --git a/mem0-ts/package.json b/mem0-ts/package.json index 7e45904ac..fd5d653f8 100644 --- a/mem0-ts/package.json +++ b/mem0-ts/package.json @@ -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", diff --git a/mem0-ts/pnpm-lock.yaml b/mem0-ts/pnpm-lock.yaml index 449cc2dc7..05b260a13 100644 --- a/mem0-ts/pnpm-lock.yaml +++ b/mem0-ts/pnpm-lock.yaml @@ -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: {} diff --git a/mem0-ts/src/oss/src/index.ts b/mem0-ts/src/oss/src/index.ts index c1d2a8995..41c15cbe8 100644 --- a/mem0-ts/src/oss/src/index.ts +++ b/mem0-ts/src/oss/src/index.ts @@ -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"; diff --git a/mem0-ts/src/oss/src/utils/factory.ts b/mem0-ts/src/oss/src/utils/factory.ts index cfd935838..956cd665d 100644 --- a/mem0-ts/src/oss/src/utils/factory.ts +++ b/mem0-ts/src/oss/src/utils/factory.ts @@ -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": diff --git a/mem0-ts/src/oss/src/vector_stores/cassandra.ts b/mem0-ts/src/oss/src/vector_stores/cassandra.ts new file mode 100644 index 000000000..fb314e507 --- /dev/null +++ b/mem0-ts/src/oss/src/vector_stores/cassandra.ts @@ -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; + execute( + query: string, + params?: any[], + options?: Record, + ): Promise<{ rows?: any[]; pageState?: string | null }>; +} + +interface CassandraVector { + id: string; + vector: number[]; + payload: Record; +} + +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; + + 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 { + if (!this._initPromise) { + this._initPromise = this._doInitialize(); + } + return this._initPromise; + } + + private async _doInitialize(): Promise { + 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, + 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[], + ): Promise { + 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 { + return null; + } + + async search( + query: number[], + topK: number = 5, + filters?: SearchFilters, + ): Promise { + 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 { + 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, + ): Promise { + 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 { + await this.initialize(); + + await this.client!.execute( + ` + DELETE FROM ${this.keyspace}.${this.collectionName} + WHERE id = ? + `, + [vectorId], + { prepare: true }, + ); + } + + async deleteCol(): Promise { + 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, + 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 { + 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 { + 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 = {}; + + 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 { + 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, + 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 = { + $and: "AND", + $or: "OR", + $not: "NOT", + }; + const normalized: Record = {}; + 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, + fetchSize: number, + ): Promise { + 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; + } + } +} diff --git a/mem0-ts/src/oss/tests/factory.unit.test.ts b/mem0-ts/src/oss/tests/factory.unit.test.ts index 85f6858d6..7f896607a 100644 --- a/mem0-ts/src/oss/tests/factory.unit.test.ts +++ b/mem0-ts/src/oss/tests/factory.unit.test.ts @@ -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) => { diff --git a/mem0-ts/src/oss/tests/vector-stores-compat.test.ts b/mem0-ts/src/oss/tests/vector-stores-compat.test.ts index 5c8225c1f..e2573c5c6 100644 --- a/mem0-ts/src/oss/tests/vector-stores-compat.test.ts +++ b/mem0-ts/src/oss/tests/vector-stores-compat.test.ts @@ -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 = {}, + ) => { + 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 = {}, + ) => { + 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 // ─────────────────────────────────────────────────────────────────────────── diff --git a/mem0-ts/tsup.config.ts b/mem0-ts/tsup.config.ts index 11d5bf342..3d70b13d9 100644 --- a/mem0-ts/tsup.config.ts +++ b/mem0-ts/tsup.config.ts @@ -10,6 +10,7 @@ const external = [ "pg", "zod", "better-sqlite3", + "cassandra-driver", "@pinecone-database/pinecone", "@qdrant/js-client-rest", "redis",