This is an automated email from the ASF dual-hosted git repository.

lidavidm pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/arrow-adbc.git


The following commit(s) were added to refs/heads/main by this push:
     new 818484573 fix(python/adbc_driver_postgresql): handle kwargs in dbapi 
connect (#2700)
818484573 is described below

commit 818484573384cc9bb2ebdc02c81cf1cda5b971ae
Author: Filip Wojciechowski <[email protected]>
AuthorDate: Tue Apr 15 00:04:37 2025 -0700

    fix(python/adbc_driver_postgresql): handle kwargs in dbapi connect (#2700)
    
    Closes #2696.
---
 .../adbc_driver_postgresql/__init__.py                     | 11 +++++++++--
 .../adbc_driver_postgresql/adbc_driver_postgresql/dbapi.py | 10 ++++++----
 python/adbc_driver_postgresql/tests/test_dbapi.py          | 14 ++++++++++++++
 3 files changed, 29 insertions(+), 6 deletions(-)

diff --git a/python/adbc_driver_postgresql/adbc_driver_postgresql/__init__.py 
b/python/adbc_driver_postgresql/adbc_driver_postgresql/__init__.py
index e2be5458f..750ba78a2 100644
--- a/python/adbc_driver_postgresql/adbc_driver_postgresql/__init__.py
+++ b/python/adbc_driver_postgresql/adbc_driver_postgresql/__init__.py
@@ -17,6 +17,7 @@
 
 import enum
 import functools
+import typing
 
 import adbc_driver_manager
 
@@ -53,9 +54,15 @@ class StatementOptions(enum.Enum):
     USE_COPY = "adbc.postgresql.use_copy"
 
 
-def connect(uri: str) -> adbc_driver_manager.AdbcDatabase:
+def connect(
+    uri: str,
+    db_kwargs: typing.Optional[typing.Dict[str, str]] = None,
+) -> adbc_driver_manager.AdbcDatabase:
     """Create a low level ADBC connection to PostgreSQL."""
-    return adbc_driver_manager.AdbcDatabase(driver=_driver_path(), uri=uri)
+    db_options = dict(db_kwargs or {})
+    db_options["driver"] = _driver_path()
+    db_options["uri"] = uri
+    return adbc_driver_manager.AdbcDatabase(**db_options)
 
 
 @functools.lru_cache
diff --git a/python/adbc_driver_postgresql/adbc_driver_postgresql/dbapi.py 
b/python/adbc_driver_postgresql/adbc_driver_postgresql/dbapi.py
index 88e309cb7..b5fbc5a9e 100644
--- a/python/adbc_driver_postgresql/adbc_driver_postgresql/dbapi.py
+++ b/python/adbc_driver_postgresql/adbc_driver_postgresql/dbapi.py
@@ -98,7 +98,7 @@ def connect(
     uri: str,
     db_kwargs: typing.Optional[typing.Dict[str, str]] = None,
     conn_kwargs: typing.Optional[typing.Dict[str, str]] = None,
-    **kwargs
+    **kwargs,
 ) -> "Connection":
     """
     Connect to PostgreSQL via ADBC.
@@ -118,9 +118,11 @@ def connect(
     conn = None
 
     try:
-        db = adbc_driver_postgresql.connect(uri)
-        conn = adbc_driver_manager.AdbcConnection(db)
-        return adbc_driver_manager.dbapi.Connection(db, conn, **kwargs)
+        db = adbc_driver_postgresql.connect(uri, db_kwargs=db_kwargs)
+        conn = adbc_driver_manager.AdbcConnection(db, **(conn_kwargs or {}))
+        return adbc_driver_manager.dbapi.Connection(
+            db, conn, conn_kwargs=conn_kwargs, **kwargs
+        )
     except Exception:
         if conn:
             conn.close()
diff --git a/python/adbc_driver_postgresql/tests/test_dbapi.py 
b/python/adbc_driver_postgresql/tests/test_dbapi.py
index ff31969b4..bcd9825b6 100644
--- a/python/adbc_driver_postgresql/tests/test_dbapi.py
+++ b/python/adbc_driver_postgresql/tests/test_dbapi.py
@@ -471,3 +471,17 @@ def test_txn_status(postgres: dbapi.Connection) -> None:
         assert status() == "active"
         postgres.rollback()
         assert status() == "intrans"
+
+
+def test_connect_conn_kwargs_db_schema(postgres_uri: str, postgres: 
dbapi.Connection):
+    """Verify current DB schema can be set via conn_kwargs."""
+    schema_key = "adbc.connection.db_schema"
+    schema_name = "dbapi_test_schema_via_option"
+
+    with postgres.cursor() as cur:
+        cur.execute(f"DROP SCHEMA IF EXISTS {schema_name} CASCADE")
+        cur.execute(f"CREATE SCHEMA {schema_name}")
+    postgres.commit()
+    with dbapi.connect(postgres_uri, conn_kwargs={schema_key: schema_name}) as 
conn:
+        option_value = conn.adbc_connection.get_option(schema_key)
+        assert option_value == schema_name

Reply via email to