From 77c5aa73d74528b68c2a4f348feb6b70b9091ad2 Mon Sep 17 00:00:00 2001 From: Shashwati Date: Tue, 7 Jul 2026 18:44:26 +0530 Subject: [PATCH] Fix pgvector register_vector to support psycopg3 (#69443) Signed-off-by: Shashwati --- .../providers/pgvector/operators/pgvector.py | 9 +++-- .../unit/pgvector/operators/test_pgvector.py | 34 +++++++++++++------ 2 files changed, 30 insertions(+), 13 deletions(-) diff --git a/providers/pgvector/src/airflow/providers/pgvector/operators/pgvector.py b/providers/pgvector/src/airflow/providers/pgvector/operators/pgvector.py index ae48a377ea7a7..4342123001f96 100644 --- a/providers/pgvector/src/airflow/providers/pgvector/operators/pgvector.py +++ b/providers/pgvector/src/airflow/providers/pgvector/operators/pgvector.py @@ -17,9 +17,8 @@ # under the License. from __future__ import annotations -from pgvector.psycopg2 import register_vector - from airflow.providers.common.sql.operators.sql import SQLExecuteQueryOperator +from airflow.providers.postgres.hooks.postgres import USE_PSYCOPG3 class PgVectorIngestOperator(SQLExecuteQueryOperator): @@ -42,6 +41,12 @@ def __init__(self, *args, **kwargs) -> None: def _register_vector(self) -> None: """Register the vector type with your connection.""" conn = self.get_db_hook().get_conn() + # ``PostgresHook`` selects psycopg2 or psycopg3 depending on the runtime environment, so the + # matching pgvector registration helper must be used for the connection object it returns. + if USE_PSYCOPG3: + from pgvector.psycopg import register_vector + else: + from pgvector.psycopg2 import register_vector register_vector(conn) def execute(self, context): diff --git a/providers/pgvector/tests/unit/pgvector/operators/test_pgvector.py b/providers/pgvector/tests/unit/pgvector/operators/test_pgvector.py index 6fe144d5bee99..761ec2a9353f5 100644 --- a/providers/pgvector/tests/unit/pgvector/operators/test_pgvector.py +++ b/providers/pgvector/tests/unit/pgvector/operators/test_pgvector.py @@ -16,7 +16,7 @@ # under the License. from __future__ import annotations -from unittest.mock import Mock, patch +from unittest.mock import MagicMock, Mock, patch import pytest @@ -32,25 +32,37 @@ def pg_vector_ingest_operator(): ) -@patch("airflow.providers.pgvector.operators.pgvector.register_vector") +@patch("airflow.providers.pgvector.operators.pgvector.USE_PSYCOPG3", False) @patch("airflow.providers.pgvector.operators.pgvector.PgVectorIngestOperator.get_db_hook") -def test_register_vector(mock_get_db_hook, mock_register_vector, pg_vector_ingest_operator): - # Create a mock database connection +def test_register_vector_psycopg2(mock_get_db_hook, pg_vector_ingest_operator): + """With psycopg2, the psycopg2 flavour of ``register_vector`` is used.""" mock_db_hook = Mock() mock_get_db_hook.return_value = mock_db_hook + mock_register_vector = MagicMock() - pg_vector_ingest_operator._register_vector() - mock_register_vector.assert_called_with(mock_db_hook.get_conn()) + with patch.dict("sys.modules", {"pgvector.psycopg2": MagicMock(register_vector=mock_register_vector)}): + pg_vector_ingest_operator._register_vector() + mock_register_vector.assert_called_once_with(mock_db_hook.get_conn()) -@patch("airflow.providers.pgvector.operators.pgvector.register_vector") -@patch("airflow.providers.pgvector.operators.pgvector.SQLExecuteQueryOperator.execute") + +@patch("airflow.providers.pgvector.operators.pgvector.USE_PSYCOPG3", True) @patch("airflow.providers.pgvector.operators.pgvector.PgVectorIngestOperator.get_db_hook") -def test_execute( - mock_get_db_hook, mock_execute_query_operator_execute, mock_register_vector, pg_vector_ingest_operator -): +def test_register_vector_psycopg3(mock_get_db_hook, pg_vector_ingest_operator): + """With psycopg3, the psycopg flavour of ``register_vector`` is used.""" mock_db_hook = Mock() mock_get_db_hook.return_value = mock_db_hook + mock_register_vector = MagicMock() + + with patch.dict("sys.modules", {"pgvector.psycopg": MagicMock(register_vector=mock_register_vector)}): + pg_vector_ingest_operator._register_vector() + mock_register_vector.assert_called_once_with(mock_db_hook.get_conn()) + + +@patch("airflow.providers.pgvector.operators.pgvector.PgVectorIngestOperator._register_vector") +@patch("airflow.providers.pgvector.operators.pgvector.SQLExecuteQueryOperator.execute") +def test_execute(mock_execute_query_operator_execute, mock_register_vector, pg_vector_ingest_operator): pg_vector_ingest_operator.execute(None) + mock_register_vector.assert_called_once() mock_execute_query_operator_execute.assert_called_once()