fix: sql injection, prompt injection (#4997)
Co-authored-by: kartik-mem0 <kartik.labhshetwar@mem0.ai>
This commit is contained in:
committed by
GitHub
parent
b66cf0f272
commit
1b95c99db4
@@ -1,9 +1,33 @@
|
||||
import type { Client as ClientType } from "pg";
|
||||
import pkg from "pg";
|
||||
const { Client } = pkg;
|
||||
const { Client, escapeIdentifier } = pkg;
|
||||
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 escapeFilterKey(key: string): string {
|
||||
if (!SAFE_IDENTIFIER_RE.test(key)) {
|
||||
throw new Error(
|
||||
`Invalid filter key '${key}': only letters, digits, and underscores are allowed.`,
|
||||
);
|
||||
}
|
||||
return key;
|
||||
}
|
||||
|
||||
interface PGVectorConfig extends VectorStoreConfig {
|
||||
dbname?: string;
|
||||
user: string;
|
||||
@@ -25,10 +49,13 @@ export class PGVector implements VectorStore {
|
||||
private _initPromise?: Promise<void>;
|
||||
|
||||
constructor(config: PGVectorConfig) {
|
||||
this.collectionName = config.collectionName || "memories";
|
||||
this.collectionName = validateIdentifier(
|
||||
config.collectionName || "memories",
|
||||
"collectionName",
|
||||
);
|
||||
this.useDiskann = config.diskann || false;
|
||||
this.useHnsw = config.hnsw || false;
|
||||
this.dbName = config.dbname || "vector_store";
|
||||
this.dbName = validateIdentifier(config.dbname || "vector_store", "dbname");
|
||||
this.config = config;
|
||||
|
||||
this.client = new Client({
|
||||
@@ -41,6 +68,10 @@ export class PGVector implements VectorStore {
|
||||
this.initialize().catch(console.error);
|
||||
}
|
||||
|
||||
private col(): string {
|
||||
return escapeIdentifier(this.collectionName);
|
||||
}
|
||||
|
||||
async initialize(): Promise<void> {
|
||||
if (!this._initPromise) {
|
||||
this._initPromise = this._doInitialize();
|
||||
@@ -102,31 +133,28 @@ export class PGVector implements VectorStore {
|
||||
}
|
||||
|
||||
private async createDatabase(dbName: string): Promise<void> {
|
||||
// Create database (cannot be parameterized)
|
||||
await this.client.query(`CREATE DATABASE ${dbName}`);
|
||||
await this.client.query(`CREATE DATABASE ${escapeIdentifier(dbName)}`);
|
||||
}
|
||||
|
||||
private async createCol(embeddingModelDims: number): Promise<void> {
|
||||
// Create the table
|
||||
const dims = Math.floor(embeddingModelDims);
|
||||
await this.client.query(`
|
||||
CREATE TABLE IF NOT EXISTS ${this.collectionName} (
|
||||
CREATE TABLE IF NOT EXISTS ${this.col()} (
|
||||
id UUID PRIMARY KEY,
|
||||
vector vector(${embeddingModelDims}),
|
||||
vector vector(${dims}),
|
||||
payload JSONB
|
||||
);
|
||||
`);
|
||||
|
||||
// Create indexes based on configuration
|
||||
if (this.useDiskann && embeddingModelDims < 2000) {
|
||||
try {
|
||||
// Check if vectorscale extension is available
|
||||
const result = await this.client.query(
|
||||
"SELECT * FROM pg_extension WHERE extname = 'vectorscale'",
|
||||
);
|
||||
if (result.rows.length > 0) {
|
||||
await this.client.query(`
|
||||
CREATE INDEX IF NOT EXISTS ${this.collectionName}_diskann_idx
|
||||
ON ${this.collectionName}
|
||||
CREATE INDEX IF NOT EXISTS ${escapeIdentifier(this.collectionName + "_diskann_idx")}
|
||||
ON ${this.col()}
|
||||
USING diskann (vector);
|
||||
`);
|
||||
}
|
||||
@@ -136,8 +164,8 @@ export class PGVector implements VectorStore {
|
||||
} else if (this.useHnsw) {
|
||||
try {
|
||||
await this.client.query(`
|
||||
CREATE INDEX IF NOT EXISTS ${this.collectionName}_hnsw_idx
|
||||
ON ${this.collectionName}
|
||||
CREATE INDEX IF NOT EXISTS ${escapeIdentifier(this.collectionName + "_hnsw_idx")}
|
||||
ON ${this.col()}
|
||||
USING hnsw (vector vector_cosine_ops);
|
||||
`);
|
||||
} catch (error) {
|
||||
@@ -153,16 +181,15 @@ export class PGVector implements VectorStore {
|
||||
): Promise<void> {
|
||||
const values = vectors.map((vector, i) => ({
|
||||
id: ids[i],
|
||||
vector: `[${vector.join(",")}]`, // Format vector as string with square brackets
|
||||
vector: `[${vector.join(",")}]`,
|
||||
payload: payloads[i],
|
||||
}));
|
||||
|
||||
const query = `
|
||||
INSERT INTO ${this.collectionName} (id, vector, payload)
|
||||
INSERT INTO ${this.col()} (id, vector, payload)
|
||||
VALUES ($1, $2::vector, $3::jsonb)
|
||||
`;
|
||||
|
||||
// Execute inserts in parallel using Promise.all
|
||||
await Promise.all(
|
||||
values.map((value) =>
|
||||
this.client.query(query, [value.id, value.vector, value.payload]),
|
||||
@@ -182,7 +209,8 @@ export class PGVector implements VectorStore {
|
||||
|
||||
if (filters) {
|
||||
for (const [key, value] of Object.entries(filters)) {
|
||||
filterConditions.push(`payload->>'${key}' = $${filterIndex}`);
|
||||
const safeKey = escapeFilterKey(key);
|
||||
filterConditions.push(`payload->>'${safeKey}' = $${filterIndex}`);
|
||||
filterValues.push(value);
|
||||
filterIndex++;
|
||||
}
|
||||
@@ -195,7 +223,7 @@ export class PGVector implements VectorStore {
|
||||
|
||||
const searchQuery = `
|
||||
SELECT id, ts_rank_cd(to_tsvector('simple', payload->>'textLemmatized'), plainto_tsquery('simple', $1)) AS score, payload
|
||||
FROM ${this.collectionName}
|
||||
FROM ${this.col()}
|
||||
WHERE to_tsvector('simple', payload->>'textLemmatized') @@ plainto_tsquery('simple', $1)
|
||||
${filterClause}
|
||||
ORDER BY score DESC
|
||||
@@ -221,13 +249,14 @@ export class PGVector implements VectorStore {
|
||||
filters?: SearchFilters,
|
||||
): Promise<VectorStoreResult[]> {
|
||||
const filterConditions: string[] = [];
|
||||
const queryVector = `[${query.join(",")}]`; // Format query vector as string with square brackets
|
||||
const queryVector = `[${query.join(",")}]`;
|
||||
const filterValues: any[] = [queryVector, topK];
|
||||
let filterIndex = 3;
|
||||
|
||||
if (filters) {
|
||||
for (const [key, value] of Object.entries(filters)) {
|
||||
filterConditions.push(`payload->>'${key}' = $${filterIndex}`);
|
||||
const safeKey = escapeFilterKey(key);
|
||||
filterConditions.push(`payload->>'${safeKey}' = $${filterIndex}`);
|
||||
filterValues.push(value);
|
||||
filterIndex++;
|
||||
}
|
||||
@@ -240,7 +269,7 @@ export class PGVector implements VectorStore {
|
||||
|
||||
const searchQuery = `
|
||||
SELECT id, vector <=> $1::vector AS distance, payload
|
||||
FROM ${this.collectionName}
|
||||
FROM ${this.col()}
|
||||
${filterClause}
|
||||
ORDER BY distance
|
||||
LIMIT $2
|
||||
@@ -257,7 +286,7 @@ export class PGVector implements VectorStore {
|
||||
|
||||
async get(vectorId: string): Promise<VectorStoreResult | null> {
|
||||
const result = await this.client.query(
|
||||
`SELECT id, payload FROM ${this.collectionName} WHERE id = $1`,
|
||||
`SELECT id, payload FROM ${this.col()} WHERE id = $1`,
|
||||
[vectorId],
|
||||
);
|
||||
|
||||
@@ -274,10 +303,10 @@ export class PGVector implements VectorStore {
|
||||
vector: number[],
|
||||
payload: Record<string, any>,
|
||||
): Promise<void> {
|
||||
const vectorStr = `[${vector.join(",")}]`; // Format vector as string with square brackets
|
||||
const vectorStr = `[${vector.join(",")}]`;
|
||||
await this.client.query(
|
||||
`
|
||||
UPDATE ${this.collectionName}
|
||||
UPDATE ${this.col()}
|
||||
SET vector = $1::vector, payload = $2::jsonb
|
||||
WHERE id = $3
|
||||
`,
|
||||
@@ -286,14 +315,13 @@ export class PGVector implements VectorStore {
|
||||
}
|
||||
|
||||
async delete(vectorId: string): Promise<void> {
|
||||
await this.client.query(
|
||||
`DELETE FROM ${this.collectionName} WHERE id = $1`,
|
||||
[vectorId],
|
||||
);
|
||||
await this.client.query(`DELETE FROM ${this.col()} WHERE id = $1`, [
|
||||
vectorId,
|
||||
]);
|
||||
}
|
||||
|
||||
async deleteCol(): Promise<void> {
|
||||
await this.client.query(`DROP TABLE IF EXISTS ${this.collectionName}`);
|
||||
await this.client.query(`DROP TABLE IF EXISTS ${this.col()}`);
|
||||
}
|
||||
|
||||
private async listCols(): Promise<string[]> {
|
||||
@@ -315,7 +343,8 @@ export class PGVector implements VectorStore {
|
||||
|
||||
if (filters) {
|
||||
for (const [key, value] of Object.entries(filters)) {
|
||||
filterConditions.push(`payload->>'${key}' = $${paramIndex}`);
|
||||
const safeKey = escapeFilterKey(key);
|
||||
filterConditions.push(`payload->>'${safeKey}' = $${paramIndex}`);
|
||||
filterValues.push(value);
|
||||
paramIndex++;
|
||||
}
|
||||
@@ -328,14 +357,14 @@ export class PGVector implements VectorStore {
|
||||
|
||||
const listQuery = `
|
||||
SELECT id, payload
|
||||
FROM ${this.collectionName}
|
||||
FROM ${this.col()}
|
||||
${filterClause}
|
||||
LIMIT $${paramIndex}
|
||||
`;
|
||||
|
||||
const countQuery = `
|
||||
SELECT COUNT(*)
|
||||
FROM ${this.collectionName}
|
||||
FROM ${this.col()}
|
||||
${filterClause}
|
||||
`;
|
||||
|
||||
|
||||
@@ -56,10 +56,13 @@ jest.mock("pg", () => {
|
||||
return client;
|
||||
});
|
||||
|
||||
const escapeIdentifier = (str: string) => `"${str.replace(/"/g, '""')}"`;
|
||||
|
||||
return {
|
||||
__esModule: true,
|
||||
default: { Client },
|
||||
default: { Client, escapeIdentifier },
|
||||
Client,
|
||||
escapeIdentifier,
|
||||
__mock: { Client, clients },
|
||||
};
|
||||
});
|
||||
@@ -73,7 +76,7 @@ describe("PGVector - search()", () => {
|
||||
pg.__mock.clients.length = 0;
|
||||
});
|
||||
|
||||
test("converts cosine distance into a clamped similarity score", async () => {
|
||||
test("returns similarity score (1 - distance) clamped to [0, 1]", async () => {
|
||||
const store = new PGVector({
|
||||
collectionName: "memories",
|
||||
user: "postgres",
|
||||
@@ -119,10 +122,5 @@ describe("PGVector - search()", () => {
|
||||
expect.stringContaining("vector <=> $1::vector AS distance"),
|
||||
["[1,0,0]", 4],
|
||||
);
|
||||
|
||||
for (const result of results) {
|
||||
expect(result.score).toBeGreaterThanOrEqual(0);
|
||||
expect(result.score).toBeLessThanOrEqual(1);
|
||||
}
|
||||
});
|
||||
});
|
||||
|
||||
@@ -58,24 +58,35 @@ class LLMReranker(BaseReranker):
|
||||
# Initialize LLM using the factory
|
||||
self.llm = LlmFactory.create(llm_provider, llm_config)
|
||||
|
||||
# Default scoring prompt
|
||||
self.scoring_prompt = getattr(self.config, 'scoring_prompt', None) or self._get_default_prompt()
|
||||
|
||||
def _get_default_prompt(self) -> str:
|
||||
"""Get the default scoring prompt template."""
|
||||
return """You are a relevance scoring assistant. Given a query and a document, you need to score how relevant the document is to the query.
|
||||
# Honor custom scoring_prompt from config if provided
|
||||
custom_prompt = getattr(self.config, 'scoring_prompt', None)
|
||||
if custom_prompt:
|
||||
import warnings
|
||||
warnings.warn(
|
||||
"LLMRerankerConfig.scoring_prompt is deprecated and will be removed in a future version. "
|
||||
"The prompt is now used as the system message.",
|
||||
DeprecationWarning,
|
||||
stacklevel=2,
|
||||
)
|
||||
self._system_prompt = custom_prompt
|
||||
else:
|
||||
self._system_prompt = self._SYSTEM_PROMPT
|
||||
|
||||
Score the relevance on a scale from 0.0 to 1.0, where:
|
||||
- 1.0 = Perfectly relevant and directly answers the query
|
||||
- 0.8-0.9 = Highly relevant with good information
|
||||
- 0.6-0.7 = Moderately relevant with some useful information
|
||||
- 0.4-0.5 = Slightly relevant with limited useful information
|
||||
- 0.0-0.3 = Not relevant or no useful information
|
||||
_SYSTEM_PROMPT = (
|
||||
"You are a relevance scoring assistant. "
|
||||
"Given a query and a document, score how relevant the document is to the query.\n\n"
|
||||
"Score the relevance on a scale from 0.0 to 1.0, where:\n"
|
||||
"- 1.0 = Perfectly relevant and directly answers the query\n"
|
||||
"- 0.8-0.9 = Highly relevant with good information\n"
|
||||
"- 0.6-0.7 = Moderately relevant with some useful information\n"
|
||||
"- 0.4-0.5 = Slightly relevant with limited useful information\n"
|
||||
"- 0.0-0.3 = Not relevant or no useful information\n\n"
|
||||
"Respond with only a single numerical score between 0.0 and 1.0. "
|
||||
"Do not include any explanation or additional text."
|
||||
)
|
||||
|
||||
Query: "{query}"
|
||||
Document: "{document}"
|
||||
|
||||
Provide only a single numerical score between 0.0 and 1.0. Do not include any explanation or additional text."""
|
||||
# Maximum character length for query and document inputs to prevent prompt flooding.
|
||||
_MAX_INPUT_LEN = 4000
|
||||
|
||||
def _extract_score(self, response_text: str) -> float:
|
||||
"""Extract numerical score from LLM response."""
|
||||
@@ -119,12 +130,17 @@ Provide only a single numerical score between 0.0 and 1.0. Do not include any ex
|
||||
doc_text = str(doc)
|
||||
|
||||
try:
|
||||
# Generate scoring prompt
|
||||
prompt = self.scoring_prompt.format(query=query, document=doc_text)
|
||||
|
||||
# Get LLM response
|
||||
# Truncate inputs to prevent prompt flooding, then send as separate
|
||||
# system/user messages so instructions cannot be overridden by user data.
|
||||
safe_query = query[: self._MAX_INPUT_LEN]
|
||||
safe_doc = doc_text[: self._MAX_INPUT_LEN]
|
||||
user_message = f"Query: {safe_query}\n\nDocument: {safe_doc}"
|
||||
|
||||
response = self.llm.generate_response(
|
||||
messages=[{"role": "user", "content": prompt}]
|
||||
messages=[
|
||||
{"role": "system", "content": self._system_prompt},
|
||||
{"role": "user", "content": user_message},
|
||||
]
|
||||
)
|
||||
|
||||
# Extract score from response
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
import json
|
||||
import logging
|
||||
import re
|
||||
from contextlib import contextmanager
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
@@ -25,6 +26,17 @@ from mem0.vector_stores.base import VectorStoreBase
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_SAFE_IDENTIFIER_RE = re.compile(r'^[a-zA-Z_][a-zA-Z0-9_]{0,127}$')
|
||||
|
||||
|
||||
def _validate_identifier(name: str, label: str = "identifier") -> str:
|
||||
if not _SAFE_IDENTIFIER_RE.match(name):
|
||||
raise ValueError(
|
||||
f"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
|
||||
|
||||
|
||||
class OutputData(BaseModel):
|
||||
id: Optional[str]
|
||||
@@ -72,7 +84,7 @@ class AzureMySQL(VectorStoreBase):
|
||||
self.user = user
|
||||
self.password = password
|
||||
self.database = database
|
||||
self.collection_name = collection_name
|
||||
self.collection_name = _validate_identifier(collection_name, "collection_name")
|
||||
self.embedding_model_dims = embedding_model_dims
|
||||
self.use_azure_credential = use_azure_credential
|
||||
self.ssl_ca = ssl_ca
|
||||
@@ -174,7 +186,7 @@ class AzureMySQL(VectorStoreBase):
|
||||
vector_size (int, optional): Vector dimension (uses self.embedding_model_dims if not provided)
|
||||
distance (str): Distance metric (cosine, euclidean, dot_product)
|
||||
"""
|
||||
table_name = name or self.collection_name
|
||||
table_name = _validate_identifier(name, "table_name") if name else self.collection_name
|
||||
dims = vector_size or self.embedding_model_dims
|
||||
|
||||
with self._get_cursor(commit=True) as cur:
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
import json
|
||||
import logging
|
||||
import re
|
||||
import uuid
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
@@ -19,6 +20,17 @@ from mem0.vector_stores.base import VectorStoreBase
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_SAFE_IDENTIFIER_RE = re.compile(r'^[a-zA-Z_][a-zA-Z0-9_]{0,127}$')
|
||||
|
||||
|
||||
def _validate_identifier(name: str, label: str = "identifier") -> str:
|
||||
if not _SAFE_IDENTIFIER_RE.match(name):
|
||||
raise ValueError(
|
||||
f"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
|
||||
|
||||
|
||||
class OutputData(BaseModel):
|
||||
id: Optional[str]
|
||||
@@ -59,8 +71,8 @@ class CassandraDB(VectorStoreBase):
|
||||
self.port = port
|
||||
self.username = username
|
||||
self.password = password
|
||||
self.keyspace = keyspace
|
||||
self.collection_name = collection_name
|
||||
self.keyspace = _validate_identifier(keyspace, "keyspace")
|
||||
self.collection_name = _validate_identifier(collection_name, "collection_name")
|
||||
self.embedding_model_dims = embedding_model_dims
|
||||
self.secure_connect_bundle = secure_connect_bundle
|
||||
self.protocol_version = protocol_version
|
||||
@@ -156,7 +168,7 @@ class CassandraDB(VectorStoreBase):
|
||||
vector_size (int, optional): Vector dimension (uses self.embedding_model_dims if not provided)
|
||||
distance (str): Distance metric (cosine, euclidean, dot_product)
|
||||
"""
|
||||
table_name = name or self.collection_name
|
||||
table_name = _validate_identifier(name, "table_name") if name else self.collection_name
|
||||
dims = vector_size or self.embedding_model_dims
|
||||
|
||||
try:
|
||||
@@ -375,12 +387,10 @@ class CassandraDB(VectorStoreBase):
|
||||
List[str]: List of collection names
|
||||
"""
|
||||
try:
|
||||
query = f"""
|
||||
SELECT table_name
|
||||
FROM system_schema.tables
|
||||
WHERE keyspace_name = '{self.keyspace}'
|
||||
"""
|
||||
rows = self.session.execute(query)
|
||||
prepared = self.session.prepare(
|
||||
"SELECT table_name FROM system_schema.tables WHERE keyspace_name = ?"
|
||||
)
|
||||
rows = self.session.execute(prepared, (self.keyspace,))
|
||||
return [row.table_name for row in rows]
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to list collections: {e}")
|
||||
|
||||
@@ -7,6 +7,7 @@ from pydantic import BaseModel
|
||||
|
||||
# Try to import psycopg (psycopg3) first, then fall back to psycopg2
|
||||
try:
|
||||
from psycopg import sql
|
||||
from psycopg.types.json import Json
|
||||
from psycopg_pool import ConnectionPool
|
||||
PSYCOPG_VERSION = 3
|
||||
@@ -14,6 +15,7 @@ try:
|
||||
logger.info("Using psycopg (psycopg3) with ConnectionPool for PostgreSQL connections")
|
||||
except ImportError:
|
||||
try:
|
||||
from psycopg2 import sql
|
||||
from psycopg2.extras import Json, execute_values
|
||||
from psycopg2.pool import ThreadedConnectionPool as ConnectionPool
|
||||
PSYCOPG_VERSION = 2
|
||||
@@ -144,6 +146,10 @@ class PGVector(VectorStoreBase):
|
||||
cur.close()
|
||||
self.connection_pool.putconn(conn)
|
||||
|
||||
def _col(self) -> "sql.Identifier":
|
||||
"""Return a safely-quoted SQL identifier for the collection table."""
|
||||
return sql.Identifier(self.collection_name)
|
||||
|
||||
def create_col(self) -> None:
|
||||
"""
|
||||
Create a new collection (table in PostgreSQL).
|
||||
@@ -152,39 +158,45 @@ class PGVector(VectorStoreBase):
|
||||
with self._get_cursor(commit=True) as cur:
|
||||
cur.execute("CREATE EXTENSION IF NOT EXISTS vector")
|
||||
cur.execute(
|
||||
f"""
|
||||
CREATE TABLE IF NOT EXISTS {self.collection_name} (
|
||||
sql.SQL("""
|
||||
CREATE TABLE IF NOT EXISTS {} (
|
||||
id UUID PRIMARY KEY,
|
||||
vector vector({self.embedding_model_dims}),
|
||||
vector vector({}),
|
||||
payload JSONB
|
||||
);
|
||||
"""
|
||||
""").format(self._col(), sql.Literal(self.embedding_model_dims))
|
||||
)
|
||||
if self.use_diskann and self.embedding_model_dims < 2000:
|
||||
cur.execute("SELECT * FROM pg_extension WHERE extname = 'vectorscale'")
|
||||
if cur.fetchone():
|
||||
# Create DiskANN index if extension is installed for faster search
|
||||
cur.execute(
|
||||
f"""
|
||||
CREATE INDEX IF NOT EXISTS {self.collection_name}_diskann_idx
|
||||
ON {self.collection_name}
|
||||
sql.SQL("""
|
||||
CREATE INDEX IF NOT EXISTS {} ON {}
|
||||
USING diskann (vector);
|
||||
"""
|
||||
""").format(
|
||||
sql.Identifier(f"{self.collection_name}_diskann_idx"),
|
||||
self._col(),
|
||||
)
|
||||
)
|
||||
elif self.use_hnsw:
|
||||
cur.execute(
|
||||
f"""
|
||||
CREATE INDEX IF NOT EXISTS {self.collection_name}_hnsw_idx
|
||||
ON {self.collection_name}
|
||||
sql.SQL("""
|
||||
CREATE INDEX IF NOT EXISTS {} ON {}
|
||||
USING hnsw (vector vector_cosine_ops)
|
||||
"""
|
||||
""").format(
|
||||
sql.Identifier(f"{self.collection_name}_hnsw_idx"),
|
||||
self._col(),
|
||||
)
|
||||
)
|
||||
cur.execute(
|
||||
f"""
|
||||
CREATE INDEX IF NOT EXISTS {self.collection_name}_text_lemmatized_idx
|
||||
ON {self.collection_name}
|
||||
sql.SQL("""
|
||||
CREATE INDEX IF NOT EXISTS {} ON {}
|
||||
USING gin(to_tsvector('simple', payload->>'text_lemmatized'));
|
||||
"""
|
||||
""").format(
|
||||
sql.Identifier(f"{self.collection_name}_text_lemmatized_idx"),
|
||||
self._col(),
|
||||
)
|
||||
)
|
||||
|
||||
def insert(self, vectors: list[list[float]], payloads=None, ids=None) -> None:
|
||||
@@ -195,14 +207,14 @@ class PGVector(VectorStoreBase):
|
||||
if PSYCOPG_VERSION == 3:
|
||||
with self._get_cursor(commit=True) as cur:
|
||||
cur.executemany(
|
||||
f"INSERT INTO {self.collection_name} (id, vector, payload) VALUES (%s, %s, %s)",
|
||||
sql.SQL("INSERT INTO {} (id, vector, payload) VALUES (%s, %s, %s)").format(self._col()),
|
||||
data,
|
||||
)
|
||||
else:
|
||||
with self._get_cursor(commit=True) as cur:
|
||||
execute_values(
|
||||
cur,
|
||||
f"INSERT INTO {self.collection_name} (id, vector, payload) VALUES %s",
|
||||
sql.SQL("INSERT INTO {} (id, vector, payload) VALUES %s").format(self._col()),
|
||||
data,
|
||||
)
|
||||
|
||||
@@ -233,17 +245,17 @@ class PGVector(VectorStoreBase):
|
||||
filter_conditions.append("payload->>%s = %s")
|
||||
filter_params.extend([k, str(v)])
|
||||
|
||||
filter_clause = "WHERE " + " AND ".join(filter_conditions) if filter_conditions else ""
|
||||
filter_clause = sql.SQL("WHERE " + " AND ".join(filter_conditions)) if filter_conditions else sql.SQL("")
|
||||
|
||||
with self._get_cursor() as cur:
|
||||
cur.execute(
|
||||
f"""
|
||||
sql.SQL("""
|
||||
SELECT id, vector <=> %s::vector AS distance, payload
|
||||
FROM {self.collection_name}
|
||||
{filter_clause}
|
||||
FROM {}
|
||||
{}
|
||||
ORDER BY distance
|
||||
LIMIT %s
|
||||
""",
|
||||
""").format(self._col(), filter_clause),
|
||||
(vectors, *filter_params, top_k),
|
||||
)
|
||||
|
||||
@@ -270,21 +282,19 @@ class PGVector(VectorStoreBase):
|
||||
filter_conditions.append("payload->>%s = %s")
|
||||
filter_params.extend([k, str(v)])
|
||||
|
||||
filter_clause = ""
|
||||
if filter_conditions:
|
||||
filter_clause = "AND " + " AND ".join(filter_conditions)
|
||||
filter_clause = sql.SQL("AND " + " AND ".join(filter_conditions)) if filter_conditions else sql.SQL("")
|
||||
|
||||
try:
|
||||
with self._get_cursor() as cur:
|
||||
cur.execute(
|
||||
f"""
|
||||
sql.SQL("""
|
||||
SELECT id, ts_rank_cd(to_tsvector('simple', payload->>'text_lemmatized'), plainto_tsquery('simple', %s)) AS score, payload
|
||||
FROM {self.collection_name}
|
||||
FROM {}
|
||||
WHERE to_tsvector('simple', payload->>'text_lemmatized') @@ plainto_tsquery('simple', %s)
|
||||
{filter_clause}
|
||||
{}
|
||||
ORDER BY score DESC
|
||||
LIMIT %s
|
||||
""",
|
||||
""").format(self._col(), filter_clause),
|
||||
(query, query, *filter_params, top_k),
|
||||
)
|
||||
|
||||
@@ -302,7 +312,7 @@ class PGVector(VectorStoreBase):
|
||||
vector_id (str): ID of the vector to delete.
|
||||
"""
|
||||
with self._get_cursor(commit=True) as cur:
|
||||
cur.execute(f"DELETE FROM {self.collection_name} WHERE id = %s", (vector_id,))
|
||||
cur.execute(sql.SQL("DELETE FROM {} WHERE id = %s").format(self._col()), (vector_id,))
|
||||
|
||||
def update(
|
||||
self,
|
||||
@@ -321,7 +331,7 @@ class PGVector(VectorStoreBase):
|
||||
with self._get_cursor(commit=True) as cur:
|
||||
if vector:
|
||||
cur.execute(
|
||||
f"UPDATE {self.collection_name} SET vector = %s WHERE id = %s",
|
||||
sql.SQL("UPDATE {} SET vector = %s WHERE id = %s").format(self._col()),
|
||||
(vector, vector_id),
|
||||
)
|
||||
if payload:
|
||||
@@ -329,13 +339,13 @@ class PGVector(VectorStoreBase):
|
||||
if PSYCOPG_VERSION == 3:
|
||||
# psycopg3 uses psycopg.types.json.Json
|
||||
cur.execute(
|
||||
f"UPDATE {self.collection_name} SET payload = %s WHERE id = %s",
|
||||
sql.SQL("UPDATE {} SET payload = %s WHERE id = %s").format(self._col()),
|
||||
(Json(payload), vector_id),
|
||||
)
|
||||
else:
|
||||
# psycopg2 uses psycopg2.extras.Json
|
||||
cur.execute(
|
||||
f"UPDATE {self.collection_name} SET payload = %s WHERE id = %s",
|
||||
sql.SQL("UPDATE {} SET payload = %s WHERE id = %s").format(self._col()),
|
||||
(Json(payload), vector_id),
|
||||
)
|
||||
|
||||
@@ -352,7 +362,7 @@ class PGVector(VectorStoreBase):
|
||||
"""
|
||||
with self._get_cursor() as cur:
|
||||
cur.execute(
|
||||
f"SELECT id, vector, payload FROM {self.collection_name} WHERE id = %s",
|
||||
sql.SQL("SELECT id, vector, payload FROM {} WHERE id = %s").format(self._col()),
|
||||
(vector_id,),
|
||||
)
|
||||
result = cur.fetchone()
|
||||
@@ -374,7 +384,7 @@ class PGVector(VectorStoreBase):
|
||||
def delete_col(self) -> None:
|
||||
"""Delete a collection."""
|
||||
with self._get_cursor(commit=True) as cur:
|
||||
cur.execute(f"DROP TABLE IF EXISTS {self.collection_name}")
|
||||
cur.execute(sql.SQL("DROP TABLE IF EXISTS {}").format(self._col()))
|
||||
|
||||
def col_info(self) -> dict[str, Any]:
|
||||
"""
|
||||
@@ -385,14 +395,14 @@ class PGVector(VectorStoreBase):
|
||||
"""
|
||||
with self._get_cursor() as cur:
|
||||
cur.execute(
|
||||
f"""
|
||||
sql.SQL("""
|
||||
SELECT
|
||||
table_name,
|
||||
(SELECT COUNT(*) FROM {self.collection_name}) as row_count,
|
||||
(SELECT pg_size_pretty(pg_total_relation_size('{self.collection_name}'))) as total_size
|
||||
(SELECT COUNT(*) FROM {}) as row_count,
|
||||
(SELECT pg_size_pretty(pg_total_relation_size({}::regclass))) as total_size
|
||||
FROM information_schema.tables
|
||||
WHERE table_schema = 'public' AND table_name = %s
|
||||
""",
|
||||
""").format(self._col(), sql.Literal(self.collection_name)),
|
||||
(self.collection_name,),
|
||||
)
|
||||
result = cur.fetchone()
|
||||
@@ -421,17 +431,18 @@ class PGVector(VectorStoreBase):
|
||||
filter_conditions.append("payload->>%s = %s")
|
||||
filter_params.extend([k, str(v)])
|
||||
|
||||
filter_clause = "WHERE " + " AND ".join(filter_conditions) if filter_conditions else ""
|
||||
|
||||
query = f"""
|
||||
SELECT id, vector, payload
|
||||
FROM {self.collection_name}
|
||||
{filter_clause}
|
||||
LIMIT %s
|
||||
"""
|
||||
filter_clause = sql.SQL("WHERE " + " AND ".join(filter_conditions)) if filter_conditions else sql.SQL("")
|
||||
|
||||
with self._get_cursor() as cur:
|
||||
cur.execute(query, (*filter_params, top_k))
|
||||
cur.execute(
|
||||
sql.SQL("""
|
||||
SELECT id, vector, payload
|
||||
FROM {}
|
||||
{}
|
||||
LIMIT %s
|
||||
""").format(self._col(), filter_clause),
|
||||
(*filter_params, top_k),
|
||||
)
|
||||
results = cur.fetchall()
|
||||
return [[OutputData(id=str(r[0]), score=None, payload=r[2]) for r in results]]
|
||||
|
||||
|
||||
@@ -79,8 +79,8 @@ class TestRerank:
|
||||
reranker = LLMReranker({"provider": "openai"})
|
||||
reranker.rerank("query", [{"text": "some text"}])
|
||||
|
||||
prompt_sent = mock_llm_instance.generate_response.call_args[1]["messages"][0]["content"]
|
||||
assert "some text" in prompt_sent
|
||||
user_msg = mock_llm_instance.generate_response.call_args[1]["messages"][1]["content"]
|
||||
assert "some text" in user_msg
|
||||
|
||||
def test_content_field_extraction(self, mock_llm):
|
||||
_, mock_llm_instance = mock_llm
|
||||
@@ -89,8 +89,8 @@ class TestRerank:
|
||||
reranker = LLMReranker({"provider": "openai"})
|
||||
reranker.rerank("query", [{"content": "some content"}])
|
||||
|
||||
prompt_sent = mock_llm_instance.generate_response.call_args[1]["messages"][0]["content"]
|
||||
assert "some content" in prompt_sent
|
||||
user_msg = mock_llm_instance.generate_response.call_args[1]["messages"][1]["content"]
|
||||
assert "some content" in user_msg
|
||||
|
||||
def test_fallback_score_on_llm_error(self, mock_llm):
|
||||
_, mock_llm_instance = mock_llm
|
||||
@@ -106,12 +106,14 @@ class TestRerank:
|
||||
_, mock_llm_instance = mock_llm
|
||||
mock_llm_instance.generate_response.return_value = "0.7"
|
||||
|
||||
custom_prompt = "Rate this: query={query} doc={document}"
|
||||
custom_prompt = "Rate relevance on a scale of 0.0 to 1.0."
|
||||
reranker = LLMReranker({"provider": "openai", "scoring_prompt": custom_prompt})
|
||||
reranker.rerank("my query", [{"memory": "my doc"}])
|
||||
|
||||
prompt_sent = mock_llm_instance.generate_response.call_args[1]["messages"][0]["content"]
|
||||
assert prompt_sent == "Rate this: query=my query doc=my doc"
|
||||
messages = mock_llm_instance.generate_response.call_args[1]["messages"]
|
||||
assert messages[0]["content"] == custom_prompt
|
||||
assert "my query" in messages[1]["content"]
|
||||
assert "my doc" in messages[1]["content"]
|
||||
|
||||
def test_original_doc_not_mutated(self, mock_llm):
|
||||
_, mock_llm_instance = mock_llm
|
||||
|
||||
@@ -127,10 +127,10 @@ class TestPGVector(unittest.TestCase):
|
||||
|
||||
# Verify vector extension and table creation
|
||||
self.mock_cursor.execute.assert_any_call("CREATE EXTENSION IF NOT EXISTS vector")
|
||||
table_creation_calls = [call for call in self.mock_cursor.execute.call_args_list
|
||||
if "CREATE TABLE IF NOT EXISTS test_collection" in str(call)]
|
||||
table_creation_calls = [call for call in self.mock_cursor.execute.call_args_list
|
||||
if "CREATE TABLE IF NOT EXISTS" in str(call) and "test_collection" in str(call)]
|
||||
self.assertTrue(len(table_creation_calls) > 0)
|
||||
|
||||
|
||||
# Verify pgvector instance properties
|
||||
self.assertEqual(pgvector.collection_name, "test_collection")
|
||||
self.assertEqual(pgvector.embedding_model_dims, 3)
|
||||
@@ -179,8 +179,8 @@ class TestPGVector(unittest.TestCase):
|
||||
|
||||
# Verify vector extension and table creation
|
||||
self.mock_cursor.execute.assert_any_call("CREATE EXTENSION IF NOT EXISTS vector")
|
||||
table_creation_calls = [call for call in self.mock_cursor.execute.call_args_list
|
||||
if "CREATE TABLE IF NOT EXISTS test_collection" in str(call)]
|
||||
table_creation_calls = [call for call in self.mock_cursor.execute.call_args_list
|
||||
if "CREATE TABLE IF NOT EXISTS" in str(call) and "test_collection" in str(call)]
|
||||
self.assertTrue(len(table_creation_calls) > 0)
|
||||
|
||||
# Verify pgvector instance properties
|
||||
@@ -233,7 +233,7 @@ class TestPGVector(unittest.TestCase):
|
||||
# Verify vector extension and table creation
|
||||
self.mock_cursor.execute.assert_any_call("CREATE EXTENSION IF NOT EXISTS vector")
|
||||
table_creation_calls = [call for call in self.mock_cursor.execute.call_args_list
|
||||
if "CREATE TABLE IF NOT EXISTS test_collection" in str(call)]
|
||||
if "CREATE TABLE IF NOT EXISTS" in str(call) and "test_collection" in str(call)]
|
||||
self.assertTrue(len(table_creation_calls) > 0)
|
||||
|
||||
# Verify pgvector instance properties
|
||||
@@ -277,7 +277,7 @@ class TestPGVector(unittest.TestCase):
|
||||
# Verify vector extension and table creation
|
||||
self.mock_cursor.execute.assert_any_call("CREATE EXTENSION IF NOT EXISTS vector")
|
||||
table_creation_calls = [call for call in self.mock_cursor.execute.call_args_list
|
||||
if "CREATE TABLE IF NOT EXISTS test_collection" in str(call)]
|
||||
if "CREATE TABLE IF NOT EXISTS" in str(call) and "test_collection" in str(call)]
|
||||
self.assertTrue(len(table_creation_calls) > 0)
|
||||
|
||||
# Verify pgvector instance properties
|
||||
@@ -319,7 +319,7 @@ class TestPGVector(unittest.TestCase):
|
||||
|
||||
# Verify insert query was executed (psycopg3 uses executemany)
|
||||
insert_calls = [call for call in self.mock_cursor.executemany.call_args_list
|
||||
if "INSERT INTO test_collection" in str(call)]
|
||||
if "INSERT INTO" in str(call) and "test_collection" in str(call)]
|
||||
self.assertTrue(len(insert_calls) > 0)
|
||||
|
||||
# Verify data format
|
||||
@@ -392,7 +392,9 @@ class TestPGVector(unittest.TestCase):
|
||||
mock_execute_values.assert_called_once()
|
||||
call_args = mock_execute_values.call_args
|
||||
|
||||
self.assertIn("INSERT INTO test_collection", call_args[0][1])
|
||||
mock_psycopg2.sql.SQL.assert_any_call(
|
||||
"INSERT INTO {} (id, vector, payload) VALUES %s"
|
||||
)
|
||||
|
||||
# The data argument should be a list of tuples, one per vector
|
||||
data_arg = call_args[0][2]
|
||||
@@ -400,6 +402,9 @@ class TestPGVector(unittest.TestCase):
|
||||
self.assertEqual(data_arg[0][0], self.test_ids[0])
|
||||
self.assertEqual(data_arg[1][0], self.test_ids[1])
|
||||
|
||||
# Restore the module after the sys.modules patch reverts
|
||||
importlib.reload(sys.modules['mem0.vector_stores.pgvector'])
|
||||
|
||||
@patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 3)
|
||||
@patch('mem0.vector_stores.pgvector.ConnectionPool')
|
||||
@patch.object(PGVector, '_get_cursor')
|
||||
@@ -534,7 +539,7 @@ class TestPGVector(unittest.TestCase):
|
||||
|
||||
# Verify delete query was executed
|
||||
delete_calls = [call for call in self.mock_cursor.execute.call_args_list
|
||||
if "DELETE FROM test_collection" in str(call)]
|
||||
if "DELETE FROM" in str(call) and "test_collection" in str(call)]
|
||||
self.assertTrue(len(delete_calls) > 0)
|
||||
|
||||
@patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 2)
|
||||
@@ -573,7 +578,7 @@ class TestPGVector(unittest.TestCase):
|
||||
|
||||
# Verify delete query was executed
|
||||
delete_calls = [call for call in self.mock_cursor.execute.call_args_list
|
||||
if "DELETE FROM test_collection" in str(call)]
|
||||
if "DELETE FROM" in str(call) and "test_collection" in str(call)]
|
||||
self.assertTrue(len(delete_calls) > 0)
|
||||
|
||||
@patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 3)
|
||||
@@ -615,7 +620,7 @@ class TestPGVector(unittest.TestCase):
|
||||
|
||||
# Verify update queries were executed
|
||||
update_calls = [call for call in self.mock_cursor.execute.call_args_list
|
||||
if "UPDATE test_collection" in str(call)]
|
||||
if "UPDATE" in str(call) and "test_collection" in str(call)]
|
||||
self.assertTrue(len(update_calls) > 0)
|
||||
|
||||
@patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 2)
|
||||
@@ -657,7 +662,7 @@ class TestPGVector(unittest.TestCase):
|
||||
|
||||
# Verify update queries were executed
|
||||
update_calls = [call for call in self.mock_cursor.execute.call_args_list
|
||||
if "UPDATE test_collection" in str(call)]
|
||||
if "UPDATE" in str(call) and "test_collection" in str(call)]
|
||||
self.assertTrue(len(update_calls) > 0)
|
||||
|
||||
@patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 3)
|
||||
@@ -867,7 +872,7 @@ class TestPGVector(unittest.TestCase):
|
||||
|
||||
# Verify delete_col query was executed
|
||||
delete_calls = [call for call in self.mock_cursor.execute.call_args_list
|
||||
if "DROP TABLE IF EXISTS test_collection" in str(call)]
|
||||
if "DROP TABLE IF EXISTS" in str(call) and "test_collection" in str(call)]
|
||||
self.assertTrue(len(delete_calls) > 0)
|
||||
|
||||
@patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 2)
|
||||
@@ -906,7 +911,7 @@ class TestPGVector(unittest.TestCase):
|
||||
|
||||
# Verify delete_col query was executed
|
||||
delete_calls = [call for call in self.mock_cursor.execute.call_args_list
|
||||
if "DROP TABLE IF EXISTS test_collection" in str(call)]
|
||||
if "DROP TABLE IF EXISTS" in str(call) and "test_collection" in str(call)]
|
||||
self.assertTrue(len(delete_calls) > 0)
|
||||
|
||||
@patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 3)
|
||||
@@ -1802,7 +1807,7 @@ class TestPGVector(unittest.TestCase):
|
||||
|
||||
# Verify the update query was executed
|
||||
update_calls = [call for call in self.mock_cursor.execute.call_args_list
|
||||
if "UPDATE test_collection SET payload" in str(call)]
|
||||
if "UPDATE" in str(call) and "test_collection" in str(call) and "SET payload" in str(call)]
|
||||
self.assertTrue(len(update_calls) > 0)
|
||||
|
||||
@patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 2)
|
||||
@@ -1843,7 +1848,7 @@ class TestPGVector(unittest.TestCase):
|
||||
|
||||
# Verify the update query was executed
|
||||
update_calls = [call for call in self.mock_cursor.execute.call_args_list
|
||||
if "UPDATE test_collection SET payload" in str(call)]
|
||||
if "UPDATE" in str(call) and "test_collection" in str(call) and "SET payload" in str(call)]
|
||||
self.assertTrue(len(update_calls) > 0)
|
||||
|
||||
@patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 2)
|
||||
@@ -1861,7 +1866,7 @@ class TestPGVector(unittest.TestCase):
|
||||
|
||||
# Only raise exception on the delete operation, not during setup
|
||||
def execute_side_effect(*args, **kwargs):
|
||||
if args and "DELETE FROM" in str(args[0]):
|
||||
if args and ("DELETE FROM" in str(args[0]) or "DELETE" in repr(args[0])):
|
||||
raise Exception("Database error")
|
||||
return MagicMock()
|
||||
mock_cursor.execute.side_effect = execute_side_effect
|
||||
@@ -2003,9 +2008,9 @@ class TestPGVector(unittest.TestCase):
|
||||
|
||||
# Verify only vector update query was executed (not payload)
|
||||
vector_update_calls = [call for call in self.mock_cursor.execute.call_args_list
|
||||
if "UPDATE test_collection SET vector" in str(call) and "payload" not in str(call)]
|
||||
if "UPDATE" in str(call) and "test_collection" in str(call) and "SET vector" in str(call) and "payload" not in str(call)]
|
||||
payload_update_calls = [call for call in self.mock_cursor.execute.call_args_list
|
||||
if "UPDATE test_collection SET payload" in str(call)]
|
||||
if "UPDATE" in str(call) and "test_collection" in str(call) and "SET payload" in str(call)]
|
||||
|
||||
self.assertTrue(len(vector_update_calls) > 0)
|
||||
self.assertEqual(len(payload_update_calls), 0)
|
||||
@@ -2045,9 +2050,9 @@ class TestPGVector(unittest.TestCase):
|
||||
|
||||
# Verify both vector and payload update queries were executed
|
||||
vector_update_calls = [call for call in self.mock_cursor.execute.call_args_list
|
||||
if "UPDATE test_collection SET vector" in str(call)]
|
||||
if "UPDATE" in str(call) and "test_collection" in str(call) and "SET vector" in str(call)]
|
||||
payload_update_calls = [call for call in self.mock_cursor.execute.call_args_list
|
||||
if "UPDATE test_collection SET payload" in str(call)]
|
||||
if "UPDATE" in str(call) and "test_collection" in str(call) and "SET payload" in str(call)]
|
||||
|
||||
self.assertTrue(len(vector_update_calls) > 0)
|
||||
self.assertTrue(len(payload_update_calls) > 0)
|
||||
|
||||
Reference in New Issue
Block a user