From c944bed4602a2edcc0a62a009f887054139a90d2 Mon Sep 17 00:00:00 2001 From: Rod Boev Date: Mon, 6 Jul 2026 13:47:43 -0400 Subject: [PATCH] feat(ts-sdk): add Azure MySQL vector store (#5827) Co-authored-by: kartik-mem0 --- mem0-ts/package.json | 8 +- mem0-ts/pnpm-lock.yaml | 72 ++++ mem0-ts/src/oss/src/index.ts | 1 + mem0-ts/src/oss/src/utils/factory.ts | 3 + .../src/oss/src/vector_stores/azure_mysql.ts | 370 ++++++++++++++++++ mem0-ts/src/oss/tests/factory.unit.test.ts | 6 + .../oss/tests/vector-stores-compat.test.ts | 297 +++++++++++++- mem0-ts/tsup.config.ts | 1 + 8 files changed, 756 insertions(+), 2 deletions(-) create mode 100644 mem0-ts/src/oss/src/vector_stores/azure_mysql.ts diff --git a/mem0-ts/package.json b/mem0-ts/package.json index 717a23a07..04a0bcaf4 100644 --- a/mem0-ts/package.json +++ b/mem0-ts/package.json @@ -132,7 +132,13 @@ "redis": "^4.6.13", "iovalkey": "^0.3.3", "compromise": "^14.0.0", - "natural": "^8.0.1" + "natural": "^8.0.1", + "mysql2": "^3.0.0" + }, + "peerDependenciesMeta": { + "mysql2": { + "optional": true + } }, "engines": { "node": ">=18" diff --git a/mem0-ts/pnpm-lock.yaml b/mem0-ts/pnpm-lock.yaml index 1da2a8096..5fe30d89c 100644 --- a/mem0-ts/pnpm-lock.yaml +++ b/mem0-ts/pnpm-lock.yaml @@ -95,6 +95,9 @@ importers: groq-sdk: specifier: 0.3.0 version: 0.3.0 + mysql2: + specifier: ^3.0.0 + version: 3.22.5(@types/node@22.19.21) natural: specifier: ^8.0.1 version: 8.1.1 @@ -1445,6 +1448,10 @@ packages: asynckit@0.4.0: resolution: {integrity: sha512-Oei9OH4tRh0YqU3GxhX79dM/mwVgvbZJaSNaRk+bshkj0S5cfHcgYakreBjrHwatXKbz+IoIdYLxrKim2MjW0Q==} + aws-ssl-profiles@1.1.2: + resolution: {integrity: sha512-NZKeq9AfyQvEeNlN0zSYAaWrmBffJh3IELMZfRpJVWgrpEbtEpnjvzqBPf+mxoI287JohRDoa+/nsfqqiZmF6g==} + engines: {node: '>= 6.0.0'} + axios@1.17.0: resolution: {integrity: sha512-J8SwNxprqqpbfenehxWYXE7CW+wM1BB4w3+N+g+/Wx40xM4rsLrfPmHHxSWIxJLYDgSY/HqlFPIYb2/S3rxafw==} @@ -2002,6 +2009,9 @@ packages: gearhash-jit@1.0.2: resolution: {integrity: sha512-UhzJL4KXSdqAKepy/tZwmi2Rcy0YMmtiC4DQS4SURCuIWdh8ECZtnXK2ePRMLigfB61hRKdLK/Vgg2bSw73izQ==} + generate-function@2.3.1: + resolution: {integrity: sha512-eeB5GfMNeevm/GRYq20ShmsaGcmI81kIX2K9XQx5miC8KdHaC6Jm0qQ8ZNeGOi7wYB8OsdxKs+Y2oVuTFuVwKQ==} + generic-pool@3.9.0: resolution: {integrity: sha512-hymDOu5B53XvN4QT9dBmZxPX4CWhBPPLguTZ9MMFeFa/Kg0xWVfylOVNlJji/E7yTZWFd/q9GO5TxDLq156D7g==} engines: {node: '>= 4'} @@ -2146,6 +2156,10 @@ packages: resolution: {integrity: sha512-1dhVQZXhcHje7798IVM+xoo/1ZdVfzOMIc8/rgVSijRK38EDqOJoGula9N/8ZI5RD8QTxNQtK/Gozpr+qUqRRA==} engines: {node: '>=20.0.0'} + iconv-lite@0.7.3: + resolution: {integrity: sha512-IKXpvIzjnC9XTAUbVBcMfGS0EPaIXtW6v+zr+RRp+hqULEpo0owZax6wyRwPOJbWbzjYspQwusTsfVr0ifh4uQ==} + engines: {node: '>=0.10.0'} + ieee754@1.2.1: resolution: {integrity: sha512-dcyqhDvX1C46lXZcVqCpK+FtMRQVdIMN6/Df5js2zouUsqG7I6sFxitIC+7KYK29KdXOLHdu9zL4sFnoVQnqaA==} @@ -2219,6 +2233,9 @@ packages: resolution: {integrity: sha512-41Cifkg6e8TylSpdtTpeLVMqvSBEVzTttHvERD741+pnZ8ANv0004MRL43QKPDlK9cGvNp6NZWZUBlbGXYxxng==} engines: {node: '>=0.12.0'} + is-property@1.0.2: + resolution: {integrity: sha512-Ks/IoX00TtClbGQr4TWXemAnktAQvYB7HzcCxDGqEZU6oCmb2INHuOoKxbtR+HFkmYWBKv/dOZtGRiAjDhj92g==} + is-stream@2.0.1: resolution: {integrity: sha512-hFoiJiTl63nn+kstHGBtewWSKnQLpyb155KHheA1l39uvtO9nWIop1p3udqPcUd/xbF1VLMO4n7OI6p7RbngDg==} engines: {node: '>=8'} @@ -2532,6 +2549,10 @@ packages: lru-cache@5.1.1: resolution: {integrity: sha512-KpNARQA3Iwv+jTA0utUVVbrh+Jlrr1Fv0e56GGzAFOXN7dk/FviaDW8LHmK52DlcH4WP2n6gI8vN1aesBFgo9w==} + lru.min@1.1.4: + resolution: {integrity: sha512-DqC6n3QQ77zdFpCMASA1a3Jlb64Hv2N2DciFGkO/4L9+q/IpIAuRlKOvCXabtRW6cQf8usbmM6BE/TOPysCdIA==} + engines: {bun: '>=1.0.0', deno: '>=1.30.0', node: '>=8.0.0'} + magic-string@0.30.21: resolution: {integrity: sha512-vd2F4YUyEXKGcLHoq+TEyCjxueSeHnFxyyjNp80yg0XV4vUhnDer/lvvlqM/arB5bXQN5K2/3oinyCRyx8T2CQ==} @@ -2685,9 +2706,19 @@ packages: resolution: {integrity: sha512-71ippSywq5Yb7/tVYyGbkBggbU8H3u5Rz56fH60jGFgr8uHwxs+aSKeqmluIVzM0m0kB7xQjKS6qPfd0b2ZoqQ==} hasBin: true + mysql2@3.22.5: + resolution: {integrity: sha512-95uZ2TrPWAZdwpB3vvvDbmEMcNG8yIeNCyu6GUcr/QnWEE/wXm7+mhOCsdQfWQDTV7qYT/PDUZ4U4UPP4AsXqQ==} + engines: {node: '>= 8.0'} + peerDependencies: + '@types/node': '>= 8' + mz@2.7.0: resolution: {integrity: sha512-z81GNO7nnYMEhrGh9LeymoE4+Yr0Wn5McHIZMK5cfQCl+NDX08sCZgUc9/6MHni9IWuFLm1Z3HTCXu2z9fN62Q==} + named-placeholders@1.1.6: + resolution: {integrity: sha512-Tz09sEL2EEuv5fFowm419c1+a/jSMiBjI9gHxVLrVdbUkkNUUfjsVYs9pVZu5oCon/kmRh9TfLEObFtkVxmY0w==} + engines: {node: '>=8.0.0'} + napi-build-utils@2.0.0: resolution: {integrity: sha512-GEbrYkbfF7MoNaoh2iGG84Mnf/WZfB0GdGEsM8wz7Expx/LlWf5U8t9nvJKXSp3qr5IsEbK04cBGhol/KwOsWA==} @@ -3141,6 +3172,9 @@ packages: resolution: {integrity: sha512-b3rppTKm9T+PsVCBEOUR46GWI7fdOs00VKZ1+9c1EWDaDMvjQc6tUwuFyIprgGgTcWoVHSKrU8H31ZHA2e0RHA==} engines: {node: '>=10'} + safer-buffer@2.1.2: + resolution: {integrity: sha512-YZo3K82SD7Riyi0E1EQPojLz7kpepnSQI9IyPbHHg1XXXevb5dJI7tpyN2ADxGcQbHG7vcyRHk0cbwqcQriUtg==} + semver-compare@1.0.0: resolution: {integrity: sha512-YM3/ITh2MJ5MtzaM429anh+x2jiLVjqILF4m4oyQB18W7Ggea7BfqdH/wGMK7dDiMghv/6WG7znWMwUDzJiXow==} @@ -3225,6 +3259,10 @@ packages: sprintf-js@1.1.3: resolution: {integrity: sha512-Oo+0REFV59/rz3gfJNKQiBlwfHaSESl1pcGyABQsnnIfWOFt6JNj5gCog2U6MLZ//IGYD+nA8nI+mTShREReaA==} + sql-escaper@1.3.3: + resolution: {integrity: sha512-BsTCV265VpTp8tm1wyIm1xqQCS+Q9NHx2Sr+WcnUrgLrQ6yiDIvHYJV5gHxsj1lMBy2zm5twLaZao8Jd+S8JJw==} + engines: {bun: '>=1.0.0', deno: '>=2.0.0', node: '>=12.0.0'} + stack-utils@2.0.6: resolution: {integrity: sha512-XlkWvfIm6RmsWtNJx+uqtKLS8eqFbxUg0ZzLXqY0caEy9l7hruX8IpiDnjsLavoBgqCCR71TqWO8MaXYheJ3RQ==} engines: {node: '>=10'} @@ -5259,6 +5297,8 @@ snapshots: asynckit@0.4.0: {} + aws-ssl-profiles@1.1.2: {} + axios@1.17.0: dependencies: follow-redirects: 1.16.0 @@ -5851,6 +5891,10 @@ snapshots: gearhash-jit@1.0.2: {} + generate-function@2.3.1: + dependencies: + is-property: 1.0.2 + generic-pool@3.9.0: {} gensync@1.0.0-beta.2: {} @@ -6047,6 +6091,10 @@ snapshots: iceberg-js@0.8.1: {} + iconv-lite@0.7.3: + dependencies: + safer-buffer: 2.1.2 + ieee754@1.2.1: {} ignore-by-default@1.0.1: {} @@ -6111,6 +6159,8 @@ snapshots: is-number@7.0.0: {} + is-property@1.0.2: {} + is-stream@2.0.1: {} is-wsl@3.1.1: @@ -6584,6 +6634,8 @@ snapshots: dependencies: yallist: 3.1.1 + lru.min@1.1.4: {} + magic-string@0.30.21: dependencies: '@jridgewell/sourcemap-codec': 1.5.5 @@ -6711,12 +6763,28 @@ snapshots: mustache@4.2.0: {} + mysql2@3.22.5(@types/node@22.19.21): + dependencies: + '@types/node': 22.19.21 + aws-ssl-profiles: 1.1.2 + denque: 2.1.0 + generate-function: 2.3.1 + iconv-lite: 0.7.3 + long: 5.3.2 + lru.min: 1.1.4 + named-placeholders: 1.1.6 + sql-escaper: 1.3.3 + mz@2.7.0: dependencies: any-promise: 1.3.0 object-assign: 4.1.1 thenify-all: 1.6.0 + named-placeholders@1.1.6: + dependencies: + lru.min: 1.1.4 + napi-build-utils@2.0.0: {} natural-compare@1.4.0: {} @@ -7219,6 +7287,8 @@ snapshots: safe-stable-stringify@2.5.0: {} + safer-buffer@2.1.2: {} + semver-compare@1.0.0: {} semver@6.3.1: {} @@ -7288,6 +7358,8 @@ snapshots: sprintf-js@1.1.3: {} + sql-escaper@1.3.3: {} + stack-utils@2.0.6: dependencies: escape-string-regexp: 2.0.0 diff --git a/mem0-ts/src/oss/src/index.ts b/mem0-ts/src/oss/src/index.ts index 842a0bc2c..e3e15ab40 100644 --- a/mem0-ts/src/oss/src/index.ts +++ b/mem0-ts/src/oss/src/index.ts @@ -32,6 +32,7 @@ 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/azure_mysql"; export * from "./vector_stores/cassandra"; export * from "./vector_stores/s3_vectors"; export * from "./vector_stores/vertex_ai_vector_search"; diff --git a/mem0-ts/src/oss/src/utils/factory.ts b/mem0-ts/src/oss/src/utils/factory.ts index 7caea89f3..79197c0d7 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 { AzureMySQLDB } from "../vector_stores/azure_mysql"; import { VertexAIVectorSearch } from "../vector_stores/vertex_ai_vector_search"; import { CassandraDB } from "../vector_stores/cassandra"; import { PineconeDB } from "../vector_stores/pinecone"; @@ -135,6 +136,8 @@ export class VectorStoreFactory { return new VertexAIVectorSearch(config as any); case "pgvector": return new PGVector(config as any); + case "azure_mysql": + return new AzureMySQLDB(config as any); case "cassandra": return new CassandraDB(config as any); case "pinecone": diff --git a/mem0-ts/src/oss/src/vector_stores/azure_mysql.ts b/mem0-ts/src/oss/src/vector_stores/azure_mysql.ts new file mode 100644 index 000000000..0b5bdf210 --- /dev/null +++ b/mem0-ts/src/oss/src/vector_stores/azure_mysql.ts @@ -0,0 +1,370 @@ +import { createPool } from "mysql2/promise"; +import type { Pool, RowDataPacket } from "mysql2/promise"; +import { VectorStore } from "./base"; +import { SearchFilters, VectorStoreConfig, VectorStoreResult } from "../types"; + +const SAFE_IDENTIFIER_RE = /^[a-zA-Z_][a-zA-Z0-9_]{0,127}$/; + +function validateIdentifier( + name: string, + label: string = "identifier", +): 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; +} + +function cosineSimilarity(a: number[], b: number[]): number { + let dot = 0; + let normA = 0; + let normB = 0; + for (let i = 0; i < a.length; i++) { + dot += a[i] * b[i]; + normA += a[i] * a[i]; + normB += b[i] * b[i]; + } + const denom = Math.sqrt(normA) * Math.sqrt(normB); + return denom === 0 ? 0 : dot / denom; +} + +interface AzureMySQLConfig extends VectorStoreConfig { + host: string; + port?: number; + user: string; + password?: string; + database: string; + collectionName: string; + embeddingModelDims: number; + useAzureCredential?: boolean; + sslCa?: string; + sslDisabled?: boolean; + maxConn?: number; +} + +export class AzureMySQLDB implements VectorStore { + private pool?: Pool; + private readonly collectionName: string; + private readonly config: AzureMySQLConfig; + private _initPromise?: Promise; + + constructor(config: AzureMySQLConfig) { + this.collectionName = validateIdentifier( + config.collectionName || "memories", + "collectionName", + ); + this.config = config; + } + + private col(): string { + return `\`${this.collectionName}\``; + } + + async initialize(): Promise { + if (!this._initPromise) { + this._initPromise = this._doInitialize(); + } + return this._initPromise; + } + + private async _doInitialize(): Promise { + let password = this.config.password; + + if (this.config.useAzureCredential) { + try { + const { DefaultAzureCredential } = await import("@azure/identity"); + const credential = new DefaultAzureCredential(); + const token = await credential.getToken( + "https://ossrdbms-aad.database.windows.net/.default", + ); + password = token.token; + } catch (err) { + throw new Error(`Azure credential authentication failed: ${err}`); + } + } + + const ssl: Record | undefined = this.config.sslDisabled + ? undefined + : { + rejectUnauthorized: true, + ...(this.config.sslCa ? { ca: this.config.sslCa } : {}), + }; + + this.pool = createPool({ + host: this.config.host, + port: this.config.port ?? 3306, + user: this.config.user, + password, + database: this.config.database, + ssl, + connectionLimit: this.config.maxConn ?? 5, + waitForConnections: true, + ...(this.config.useAzureCredential + ? { + authPlugins: { + mysql_clear_password: () => () => + Buffer.from(`${password ?? ""}\0`), + }, + } + : {}), + }); + + await this.pool.execute(` + CREATE TABLE IF NOT EXISTS ${this.col()} ( + id VARCHAR(255) PRIMARY KEY, + vector JSON, + payload JSON, + text_lemmatized VARCHAR(1000) GENERATED ALWAYS AS + (CAST(JSON_UNQUOTE(JSON_EXTRACT(payload, '$.textLemmatized')) AS CHAR(1000))) STORED + ) + `); + + try { + await this.pool.execute( + `CREATE FULLTEXT INDEX ft_text_lemmatized ON ${this.col()} (text_lemmatized)`, + ); + } catch { + // Index may already exist or FULLTEXT may be unsupported; continue silently. + } + + await this.pool.execute(` + CREATE TABLE IF NOT EXISTS memory_migrations ( + id INT PRIMARY KEY, + user_id VARCHAR(255) NOT NULL + ) + `); + } + + async insert( + vectors: number[][], + ids: string[], + payloads: Record[], + ): Promise { + await this.initialize(); + const conn = await this.pool!.getConnection(); + try { + await conn.beginTransaction(); + for (let i = 0; i < vectors.length; i++) { + await conn.execute( + `INSERT INTO ${this.col()} (id, vector, payload) VALUES (?, ?, ?) AS new + ON DUPLICATE KEY UPDATE vector = new.vector, payload = new.payload`, + [ids[i], JSON.stringify(vectors[i]), JSON.stringify(payloads[i])], + ); + } + await conn.commit(); + } catch (err) { + await conn.rollback(); + throw err; + } finally { + conn.release(); + } + } + + async search( + query: number[], + topK: number = 5, + filters?: SearchFilters, + ): Promise { + await this.initialize(); + + const conditions: string[] = []; + const params: any[] = []; + + if (filters) { + for (const [k, v] of Object.entries(filters)) { + conditions.push("JSON_EXTRACT(payload, ?) = ?"); + params.push(`$.${k}`, JSON.stringify(v)); + } + } + + const whereClause = + conditions.length > 0 ? "WHERE " + conditions.join(" AND ") : ""; + const sql = `SELECT id, vector, payload FROM ${this.col()} ${whereClause}`; + const [rows] = await this.pool!.execute(sql, params); + + const scored: Array<{ id: string; score: number; payload: any }> = []; + for (const row of rows) { + const vec: number[] = + typeof row.vector === "string" ? JSON.parse(row.vector) : row.vector; + const score = cosineSimilarity(query, vec); + const payload = + typeof row.payload === "string" ? JSON.parse(row.payload) : row.payload; + scored.push({ id: row.id as string, score, payload }); + } + + scored.sort((a, b) => b.score - a.score); + return scored.slice(0, topK).map((r) => ({ + id: r.id, + payload: r.payload, + score: r.score, + })); + } + + async keywordSearch( + query: string, + topK: number = 5, + filters?: SearchFilters, + ): Promise { + try { + await this.initialize(); + + const conditions: string[] = []; + const params: any[] = [query, query]; + + if (filters) { + for (const [k, v] of Object.entries(filters)) { + conditions.push("JSON_EXTRACT(payload, ?) = ?"); + params.push(`$.${k}`, JSON.stringify(v)); + } + } + + const filterClause = + conditions.length > 0 ? " AND " + conditions.join(" AND ") : ""; + const sql = ` + SELECT id, payload, + MATCH(text_lemmatized) AGAINST(? IN NATURAL LANGUAGE MODE) AS score + FROM ${this.col()} + WHERE MATCH(text_lemmatized) AGAINST(? IN NATURAL LANGUAGE MODE) + ${filterClause} + ORDER BY score DESC + LIMIT ? + `; + params.push(topK); + + const [rows] = await this.pool!.execute(sql, params); + return rows.map((row) => ({ + id: row.id as string, + payload: + typeof row.payload === "string" + ? JSON.parse(row.payload) + : row.payload, + score: Number(row.score), + })); + } catch { + return null; + } + } + + async get(vectorId: string): Promise { + await this.initialize(); + const [rows] = await this.pool!.execute( + `SELECT id, payload FROM ${this.col()} WHERE id = ?`, + [vectorId], + ); + if (rows.length === 0) return null; + const row = rows[0]; + return { + id: row.id as string, + payload: + typeof row.payload === "string" ? JSON.parse(row.payload) : row.payload, + }; + } + + async update( + vectorId: string, + vector: number[], + payload: Record, + ): Promise { + await this.initialize(); + const conn = await this.pool!.getConnection(); + try { + await conn.beginTransaction(); + await conn.execute( + `UPDATE ${this.col()} SET vector = ?, payload = ? WHERE id = ?`, + [JSON.stringify(vector), JSON.stringify(payload), vectorId], + ); + await conn.commit(); + } catch (err) { + await conn.rollback(); + throw err; + } finally { + conn.release(); + } + } + + async delete(vectorId: string): Promise { + await this.initialize(); + await this.pool!.execute(`DELETE FROM ${this.col()} WHERE id = ?`, [ + vectorId, + ]); + } + + async deleteCol(): Promise { + await this.initialize(); + await this.pool!.execute(`DROP TABLE IF EXISTS ${this.col()}`); + } + + async list( + filters?: SearchFilters, + topK: number = 100, + ): Promise<[VectorStoreResult[], number]> { + await this.initialize(); + + const conditions: string[] = []; + const params: any[] = []; + + if (filters) { + for (const [k, v] of Object.entries(filters)) { + conditions.push("JSON_EXTRACT(payload, ?) = ?"); + params.push(`$.${k}`, JSON.stringify(v)); + } + } + + const whereClause = + conditions.length > 0 ? "WHERE " + conditions.join(" AND ") : ""; + const listSql = `SELECT id, payload FROM ${this.col()} ${whereClause} LIMIT ?`; + const countSql = `SELECT COUNT(*) AS cnt FROM ${this.col()} ${whereClause}`; + + const [rows] = await this.pool!.execute(listSql, [ + ...params, + topK, + ]); + const [countRows] = await this.pool!.execute( + countSql, + params, + ); + + const results: VectorStoreResult[] = rows.map((row) => ({ + id: row.id as string, + payload: + typeof row.payload === "string" ? JSON.parse(row.payload) : row.payload, + })); + + return [results, Number(countRows[0].cnt)]; + } + + async getUserId(): Promise { + await this.initialize(); + const [rows] = await this.pool!.execute( + "SELECT user_id FROM memory_migrations WHERE id = 1", + ); + if (rows.length > 0) { + return rows[0].user_id as string; + } + const randomUserId = + Math.random().toString(36).substring(2, 15) + + Math.random().toString(36).substring(2, 15); + await this.pool!.execute( + "INSERT INTO memory_migrations (id, user_id) VALUES (1, ?) AS new ON DUPLICATE KEY UPDATE user_id = new.user_id", + [randomUserId], + ); + return randomUserId; + } + + async setUserId(userId: string): Promise { + await this.initialize(); + await this.pool!.execute( + "INSERT INTO memory_migrations (id, user_id) VALUES (1, ?) AS new ON DUPLICATE KEY UPDATE user_id = new.user_id", + [userId], + ); + } + + async close(): Promise { + if (this.pool) { + await this.pool.end(); + } + } +} diff --git a/mem0-ts/src/oss/tests/factory.unit.test.ts b/mem0-ts/src/oss/tests/factory.unit.test.ts index 7f896607a..72b18c4ee 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/azure_mysql", () => ({ + AzureMySQLDB: jest + .fn() + .mockImplementation((config) => ({ type: "azure_mysql", config })), +})); jest.mock("../src/vector_stores/cassandra", () => ({ CassandraDB: jest .fn() @@ -294,6 +299,7 @@ describe("VectorStoreFactory", () => { ["vectorize"], ["azure-ai-search"], ["pgvector"], + ["azure_mysql"], ["cassandra"], ["s3-vectors"], ["s3_vectors"], 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 e2573c5c6..87d1057fa 100644 --- a/mem0-ts/src/oss/tests/vector-stores-compat.test.ts +++ b/mem0-ts/src/oss/tests/vector-stores-compat.test.ts @@ -1716,7 +1716,302 @@ describe("LangchainVectorStore – backward compat", () => { }); // ─────────────────────────────────────────────────────────────────────────── -// 8. Memory class — ensure it works with each provider via mocked factories +// 8. AzureMySQL — mock mysql2 pool, test interface + idempotent init + CRUD +// ─────────────────────────────────────────────────────────────────────────── +describe("AzureMySQL – backward compat with mocked client", () => { + let AzureMySQLDB: any; + let mockPool: any; + + beforeEach(() => { + jest.resetModules(); + + const rows = new Map< + string, + { id: string; vector: string; payload: string } + >(); + let userId: string | null = null; + + mockPool = { + execute: jest + .fn() + .mockImplementation(async (sql: string, params?: any[]) => { + const q = sql.trim().toUpperCase(); + + if ( + q.startsWith("CREATE TABLE") || + q.startsWith("CREATE FULLTEXT") || + q.startsWith("DROP TABLE") + ) { + return [{ affectedRows: 0 }, []]; + } + + // INSERT into main table (ON DUPLICATE KEY) + if (q.startsWith("INSERT INTO `") && q.includes("ON DUPLICATE KEY")) { + const [id, vector, payload] = params!; + rows.set(id, { id, vector, payload }); + return [{ affectedRows: 1 }, []]; + } + + // INSERT into memory_migrations + if (q.startsWith("INSERT INTO MEMORY_MIGRATIONS")) { + userId = params![0]; + return [{ affectedRows: 1 }, []]; + } + + // SELECT id, payload FROM table WHERE id = ? (single-row get by PK) + if ( + q.startsWith("SELECT ID, PAYLOAD FROM") && + q.includes("WHERE ID = ?") + ) { + const row = rows.get(params![0]); + return [row ? [row] : [], []]; + } + + // SELECT id, vector, payload FROM table (search with optional filters) + if (q.startsWith("SELECT ID, VECTOR, PAYLOAD FROM")) { + return [[...rows.values()], []]; + } + + if (q.includes("MATCH(TEXT_LEMMATIZED) AGAINST")) { + const term = String(params?.[0] ?? "").toLowerCase(); + const limit = Number(params?.[params!.length - 1] ?? rows.size); + const matched = [...rows.values()] + .filter((row) => { + const payload = JSON.parse(row.payload); + return String(payload.textLemmatized ?? "") + .toLowerCase() + .includes(term); + }) + .slice(0, limit) + .map((row) => ({ ...row, score: 1 })); + return [matched, []]; + } + + // SELECT id, payload FROM table (list with LIMIT) + if (q.startsWith("SELECT ID, PAYLOAD FROM")) { + return [ + [...rows.values()].slice(0, params![params!.length - 1]), + [], + ]; + } + + // SELECT COUNT(*) + if (q.startsWith("SELECT COUNT(*)")) { + return [[{ cnt: rows.size }], []]; + } + + // UPDATE + if (q.startsWith("UPDATE `")) { + const [vector, payload, id] = params!; + if (rows.has(id)) { + rows.set(id, { id, vector, payload }); + } + return [{ affectedRows: 1 }, []]; + } + + // DELETE + if (q.startsWith("DELETE FROM `")) { + rows.delete(params![0]); + return [{ affectedRows: 1 }, []]; + } + + // SELECT user_id FROM memory_migrations + if (q.startsWith("SELECT USER_ID FROM MEMORY_MIGRATIONS")) { + return [userId ? [{ user_id: userId }] : [], []]; + } + + return [[], []]; + }), + getConnection: jest.fn().mockImplementation(async () => ({ + beginTransaction: jest.fn().mockResolvedValue(undefined), + execute: jest + .fn() + .mockImplementation(async (sql: string, params?: any[]) => { + return mockPool.execute(sql, params); + }), + commit: jest.fn().mockResolvedValue(undefined), + rollback: jest.fn().mockResolvedValue(undefined), + release: jest.fn(), + })), + end: jest.fn().mockResolvedValue(undefined), + }; + + jest.doMock("mysql2/promise", () => ({ + createPool: jest.fn().mockReturnValue(mockPool), + })); + + AzureMySQLDB = require("../src/vector_stores/azure_mysql").AzureMySQLDB; + }); + + afterEach(() => { + jest.restoreAllMocks(); + jest.resetModules(); + }); + + it("implements full VectorStore interface", () => { + const store = new AzureMySQLDB({ + host: "localhost", + user: "test", + database: "testdb", + collectionName: "memories", + embeddingModelDims: 4, + }); + expect(typeof store.insert).toBe("function"); + expect(typeof store.search).toBe("function"); + expect(typeof store.keywordSearch).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", async () => { + const mysql2 = require("mysql2/promise"); + const store = new AzureMySQLDB({ + host: "localhost", + user: "test", + database: "testdb", + collectionName: "memories", + embeddingModelDims: 4, + }); + + const p1 = store.initialize(); + const p2 = store.initialize(); + const p3 = store.initialize(); + await Promise.all([p1, p2, p3]); + + // createPool called only once despite 3 initialize() calls + expect(mysql2.createPool).toHaveBeenCalledTimes(1); + }); + + it("full CRUD cycle", async () => { + const store = new AzureMySQLDB({ + host: "localhost", + user: "test", + database: "testdb", + collectionName: "memories", + embeddingModelDims: 4, + }); + await store.initialize(); + + const vec1 = [1, 0, 0, 0]; + const vec2 = [0, 1, 0, 0]; + + // Insert + await store.insert( + [vec1, vec2], + ["id-1", "id-2"], + [{ data: "alpha" }, { data: "beta" }], + ); + + // Get + const item = await store.get("id-1"); + expect(item).not.toBeNull(); + expect(item!.id).toBe("id-1"); + + // Search — vec1 should rank first + const results = await store.search(vec1, 2); + expect(results.length).toBeGreaterThan(0); + + // Update + await store.update("id-1", [0, 0, 1, 0], { data: "updated" }); + + // List + const [listed, count] = await store.list(); + expect(listed.length).toBeGreaterThan(0); + expect(count).toBeGreaterThan(0); + + // Delete + await store.delete("id-2"); + + // DeleteCol + await store.deleteCol(); + }); + + it("keywordSearch matches textLemmatized payloads", async () => { + const store = new AzureMySQLDB({ + host: "localhost", + user: "test", + database: "testdb", + collectionName: "memories", + embeddingModelDims: 4, + }); + await store.initialize(); + + await store.insert( + [[1, 0, 0, 0]], + ["id-1"], + [ + { + data: "alpha", + textLemmatized: "alpha normalized", + }, + ], + ); + + const results = await store.keywordSearch("normalized", 5); + expect(results).not.toBeNull(); + expect(results![0].id).toBe("id-1"); + }); + + it("getUserId and setUserId roundtrip", async () => { + const store = new AzureMySQLDB({ + host: "localhost", + user: "test", + database: "testdb", + collectionName: "memories", + embeddingModelDims: 4, + }); + await store.initialize(); + + await store.setUserId("custom-user"); + const retrieved = await store.getUserId(); + expect(retrieved).toBe("custom-user"); + }); + + it("uses mysql_clear_password semantics for Azure tokens", async () => { + jest.doMock("@azure/identity", () => ({ + DefaultAzureCredential: jest.fn().mockImplementation(() => ({ + getToken: jest.fn().mockResolvedValue({ token: "aad-token" }), + })), + })); + + const mysql2 = require("mysql2/promise"); + const store = new AzureMySQLDB({ + host: "localhost", + user: "test", + database: "testdb", + collectionName: "memories", + embeddingModelDims: 4, + useAzureCredential: true, + }); + await store.initialize(); + + const poolConfig = mysql2.createPool.mock.calls[0][0]; + const plugin = poolConfig.authPlugins.mysql_clear_password(); + expect(plugin()).toEqual(Buffer.from("aad-token\0")); + }); + + it("rejects invalid collectionName at construction", () => { + expect( + () => + new AzureMySQLDB({ + host: "localhost", + user: "test", + database: "testdb", + collectionName: "drop--table", + embeddingModelDims: 4, + }), + ).toThrow("Invalid collectionName"); + }); +}); + +// ─────────────────────────────────────────────────────────────────────────── +// 9. Memory class — ensure it works with each provider via mocked factories // ─────────────────────────────────────────────────────────────────────────── describe("Memory class – backward compat with all providers", () => { function createMockEmbedder(dims: number) { diff --git a/mem0-ts/tsup.config.ts b/mem0-ts/tsup.config.ts index 6cd4ed914..ae8d290cc 100644 --- a/mem0-ts/tsup.config.ts +++ b/mem0-ts/tsup.config.ts @@ -28,6 +28,7 @@ const external = [ "fastembed", "compromise", "natural", + "mysql2", "@turbopuffer/turbopuffer", ];