diff --git a/mem0-ts/src/oss/src/vector_stores/pgvector.ts b/mem0-ts/src/oss/src/vector_stores/pgvector.ts index ac957b168..97393b9fe 100644 --- a/mem0-ts/src/oss/src/vector_stores/pgvector.ts +++ b/mem0-ts/src/oss/src/vector_stores/pgvector.ts @@ -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; 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 { if (!this._initPromise) { this._initPromise = this._doInitialize(); @@ -102,31 +133,28 @@ export class PGVector implements VectorStore { } private async createDatabase(dbName: string): Promise { - // 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 { - // 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 { 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 { 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 { 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, ): Promise { - 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 { - 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 { - await this.client.query(`DROP TABLE IF EXISTS ${this.collectionName}`); + await this.client.query(`DROP TABLE IF EXISTS ${this.col()}`); } private async listCols(): Promise { @@ -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} `; diff --git a/mem0-ts/src/oss/tests/pgvector.unit.test.ts b/mem0-ts/src/oss/tests/pgvector.unit.test.ts index ef4827ced..d0d2b7cca 100644 --- a/mem0-ts/src/oss/tests/pgvector.unit.test.ts +++ b/mem0-ts/src/oss/tests/pgvector.unit.test.ts @@ -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); - } }); }); diff --git a/mem0/reranker/llm_reranker.py b/mem0/reranker/llm_reranker.py index a474ea2d8..16d7f47c3 100644 --- a/mem0/reranker/llm_reranker.py +++ b/mem0/reranker/llm_reranker.py @@ -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 diff --git a/mem0/vector_stores/azure_mysql.py b/mem0/vector_stores/azure_mysql.py index 40c2d97d0..9391d6bb9 100644 --- a/mem0/vector_stores/azure_mysql.py +++ b/mem0/vector_stores/azure_mysql.py @@ -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: diff --git a/mem0/vector_stores/cassandra.py b/mem0/vector_stores/cassandra.py index c6579a29f..3916e103e 100644 --- a/mem0/vector_stores/cassandra.py +++ b/mem0/vector_stores/cassandra.py @@ -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}") diff --git a/mem0/vector_stores/pgvector.py b/mem0/vector_stores/pgvector.py index eda8fd5a4..8642766f6 100644 --- a/mem0/vector_stores/pgvector.py +++ b/mem0/vector_stores/pgvector.py @@ -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]] diff --git a/tests/rerankers/test_llm_reranker_rerank.py b/tests/rerankers/test_llm_reranker_rerank.py index 9e16f149a..607d99b1b 100644 --- a/tests/rerankers/test_llm_reranker_rerank.py +++ b/tests/rerankers/test_llm_reranker_rerank.py @@ -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 diff --git a/tests/vector_stores/test_pgvector.py b/tests/vector_stores/test_pgvector.py index 7479f3e23..faa2029bc 100644 --- a/tests/vector_stores/test_pgvector.py +++ b/tests/vector_stores/test_pgvector.py @@ -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)