feat(ts-sdk): add Azure MySQL vector store (#5827)
Co-authored-by: kartik-mem0 <kartik.labhshetwar@mem0.ai>
This commit is contained in:
@@ -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"
|
||||
|
||||
Generated
+72
@@ -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
|
||||
|
||||
@@ -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";
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -28,6 +28,7 @@ const external = [
|
||||
"fastembed",
|
||||
"compromise",
|
||||
"natural",
|
||||
"mysql2",
|
||||
"@turbopuffer/turbopuffer",
|
||||
];
|
||||
|
||||
|
||||
Reference in New Issue
Block a user