feat(ts-sdk): add Azure MySQL vector store (#5827)

Co-authored-by: kartik-mem0 <kartik.labhshetwar@mem0.ai>
This commit is contained in:
Rod Boev
2026-07-06 13:47:43 -04:00
committed by GitHub
parent 2cc060fd76
commit c944bed460
8 changed files with 756 additions and 2 deletions
+7 -1
View File
@@ -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"
+72
View File
@@ -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
+1
View File
@@ -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";
+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 { 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":
@@ -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<void>;
constructor(config: AzureMySQLConfig) {
this.collectionName = validateIdentifier(
config.collectionName || "memories",
"collectionName",
);
this.config = config;
}
private col(): string {
return `\`${this.collectionName}\``;
}
async initialize(): Promise<void> {
if (!this._initPromise) {
this._initPromise = this._doInitialize();
}
return this._initPromise;
}
private async _doInitialize(): Promise<void> {
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<string, any> | 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<string, any>[],
): Promise<void> {
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<VectorStoreResult[]> {
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<RowDataPacket[]>(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<VectorStoreResult[] | null> {
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<RowDataPacket[]>(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<VectorStoreResult | null> {
await this.initialize();
const [rows] = await this.pool!.execute<RowDataPacket[]>(
`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<string, any>,
): Promise<void> {
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<void> {
await this.initialize();
await this.pool!.execute(`DELETE FROM ${this.col()} WHERE id = ?`, [
vectorId,
]);
}
async deleteCol(): Promise<void> {
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<RowDataPacket[]>(listSql, [
...params,
topK,
]);
const [countRows] = await this.pool!.execute<RowDataPacket[]>(
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<string> {
await this.initialize();
const [rows] = await this.pool!.execute<RowDataPacket[]>(
"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<void> {
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<void> {
if (this.pool) {
await this.pool.end();
}
}
}
@@ -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"],
@@ -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) {
+1
View File
@@ -28,6 +28,7 @@ const external = [
"fastembed",
"compromise",
"natural",
"mysql2",
"@turbopuffer/turbopuffer",
];