fix: sql injection, prompt injection (#4997)

Co-authored-by: kartik-mem0 <kartik.labhshetwar@mem0.ai>
This commit is contained in:
Harsh Vardhan Gupta
2026-04-29 00:51:16 +05:30
committed by GitHub
parent b66cf0f272
commit 1b95c99db4
8 changed files with 233 additions and 150 deletions
+62 -33
View File
@@ -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}
`;
+5 -7
View File
@@ -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);
}
});
});
+37 -21
View File
@@ -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
+14 -2
View File
@@ -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:
+19 -9
View File
@@ -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}")
+60 -49
View File
@@ -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]]
+9 -7
View File
@@ -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
+27 -22
View File
@@ -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)