Fix pgvector sslmode handling for PostgreSQL URIs (#5308)

This commit is contained in:
Ren yiwei
2026-06-04 22:02:47 +08:00
committed by GitHub
parent 95b6f95f7b
commit 866888df41
2 changed files with 55 additions and 12 deletions
+33 -3
View File
@@ -4,7 +4,7 @@ import unittest
import uuid
from unittest.mock import MagicMock, patch
from mem0.vector_stores.pgvector import PGVector, _build_filter_conditions
from mem0.vector_stores.pgvector import PGVector, _build_filter_conditions, _with_sslmode
class TestPGVector(unittest.TestCase):
@@ -2084,8 +2084,8 @@ class TestPGVector(unittest.TestCase):
connection_string=connection_string
)
# Verify ConnectionPool was called with the connection string including sslmode
expected_conn_string = f"{connection_string} sslmode=require"
# Verify ConnectionPool was called with sslmode as a URI query parameter
expected_conn_string = f"{connection_string}?sslmode=require"
mock_connection_pool.assert_called_with(
conninfo=expected_conn_string,
min_size=1,
@@ -2095,6 +2095,36 @@ class TestPGVector(unittest.TestCase):
self.assertEqual(pgvector.collection_name, "test_collection")
self.assertEqual(pgvector.embedding_model_dims, 3)
def test_with_sslmode_appends_uri_query_param(self):
"""Test sslmode is appended to URI connection strings without corrupting dbname."""
connection_string = _with_sslmode(
"postgresql://user:pass@localhost:5432/db?connect_timeout=10",
"require",
)
self.assertEqual(
connection_string,
"postgresql://user:pass@localhost:5432/db?connect_timeout=10&sslmode=require",
)
def test_with_sslmode_replaces_existing_uri_sslmode(self):
"""Test existing URI sslmode values are replaced rather than duplicated."""
connection_string = _with_sslmode(
"postgresql://user:pass@localhost:5432/db?sslmode=prefer&connect_timeout=10",
"require",
)
self.assertEqual(
connection_string,
"postgresql://user:pass@localhost:5432/db?connect_timeout=10&sslmode=require",
)
def test_with_sslmode_preserves_keyword_conninfo_format(self):
"""Test non-URI conninfo strings keep PostgreSQL keyword syntax."""
connection_string = _with_sslmode("dbname=test user=postgres sslmode=prefer", "require")
self.assertEqual(connection_string, "dbname=test user=postgres sslmode=require")
# Enhanced Test for Index Creation with DiskANN
@patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 3)
@patch('mem0.vector_stores.pgvector.ConnectionPool')