From 91c56f8e5896bd91307ac0cd97755943002a6d0c Mon Sep 17 00:00:00 2001 From: Karen Braganza Date: Tue, 3 Feb 2026 12:55:30 -0500 Subject: [PATCH 01/14] First draft of DbApiHookAsync --- .../airflow/providers/common/sql/hooks/sql.py | 180 +++++++++++++++++- 1 file changed, 178 insertions(+), 2 deletions(-) diff --git a/providers/common/sql/src/airflow/providers/common/sql/hooks/sql.py b/providers/common/sql/src/airflow/providers/common/sql/hooks/sql.py index 3b439947596c0..237cd6aab6fc1 100644 --- a/providers/common/sql/src/airflow/providers/common/sql/hooks/sql.py +++ b/providers/common/sql/src/airflow/providers/common/sql/hooks/sql.py @@ -16,16 +16,27 @@ # under the License. from __future__ import annotations +import asyncio import contextlib import warnings -from collections.abc import Callable, Generator, Iterable, Mapping, MutableMapping, Sequence -from contextlib import closing, contextmanager, suppress +from collections.abc import ( + AsyncIterator, + Awaitable, + Callable, + Generator, + Iterable, + Mapping, + MutableMapping, + Sequence, +) +from contextlib import asynccontextmanager, closing, contextmanager, suppress from datetime import datetime from functools import cached_property from typing import TYPE_CHECKING, Any, Literal, Protocol, TypeVar, cast, overload from urllib.parse import urlparse import sqlparse +from asgiref.sync import sync_to_async from deprecated import deprecated from methodtools import lru_cache from more_itertools import chunked @@ -64,6 +75,7 @@ T = TypeVar("T") +HANDLER = Callable[[Any], T | Awaitable[T]] SQL_PLACEHOLDERS = frozenset({"%s", "?"}) WARNING_MESSAGE = """Import of {} from the 'airflow.providers.common.sql.hooks' module is deprecated and will be removed in the future. Please import it from 'airflow.providers.common.sql.hooks.handlers'.""" @@ -1092,3 +1104,167 @@ def get_db_log_messages(self, conn) -> None: :param conn: Connection object """ + + +class DbApiHookAsync(DbApiHook): + """Abstract base class for asynchronous sql hooks.""" + + # Override to provide the connection name. + conn_name_attr: str + # Override to have a default connection id for a particular dbHook + default_conn_name = "default_conn_id" + # Override if this db doesn't support semicolons in SQL queries + strip_semicolon = False + # Override if this db supports autocommit. + supports_autocommit = False + # Override if this db supports execute many. + supports_executemany = False + # Override with the object that exposes the connect method + connector: ConnectorProtocol | None = None + # Override with db-specific query to check connection + _test_connection_sql = "select 1" + # Default SQL placeholder + _placeholder: str = "%s" + _dialects: MutableMapping[str, MutableMapping] = resolve_dialects() + _resolve_target_fields = conf.getboolean("core", "dbapihook_resolve_target_fields", fallback=False) + + def __init__(self, *args, schema=None, log_sql=True, **kwargs): + super().__init__(*args, schema=schema, log_sql=log_sql, **kwargs) + self._conn_lock = asyncio.Lock() + + async def get_conn(self) -> Any: + async with self._conn_lock: + if not self._connection: + self._connection = await sync_to_async(self.get_connection)(self.get_conn_id()) + db = self._connection + if self.connector is None: + raise RuntimeError(f"{type(self).__name__} didn't have `self.connector` set!") + host = db.host or "" + login = db.login or "" + schema = db.schema or "" + return await self.connector.connect( + host=host, port=cast("int", db.port), username=login, schema=schema + ) + + async def _maybe_await(self, result: Any) -> Any: + return await result if inspect.isawaitable(result) else result + + @asynccontextmanager + async def _create_autocommit_connection(self, autocommit: bool = False): + conn = await self.get_conn() + try: + if self.supports_autocommit: + self.set_autocommit(conn, autocommit) + yield conn + finally: + close = getattr(conn, "aclose", None) or conn.close + await self._maybe_await(close()) + + @asynccontextmanager + async def _get_cursor(self, conn: Any) -> AsyncIterator[Any]: + cur_or_cm = await self._maybe_await(conn.cursor()) + if hasattr(cur_or_cm, "__aenter__") and hasattr(cur_or_cm, "__aexit__"): + async with cur_or_cm as cur: + yield cur + return + + cur = await self._maybe_await(cur_or_cm) + try: + yield cur + finally: + close = getattr(cur, "aclose", None) or getattr(cur, "close", None) + if close: + await self._maybe_await(close()) + + async def _run_command(self, cur, sql_statement, parameters): + """Run a statement using an already open cursor.""" + if self.log_sql: + self.log.info("Running statement: %s, parameters: %s", sql_statement, parameters) + + if parameters: + # If we're using psycopg3, we might need to handle parameters differently + if isinstance(parameters, list): + parameters = tuple(parameters) + await self._maybe_await(cur.execute(sql_statement, parameters)) + else: + await self._maybe_await(cur.execute(sql_statement)) + + # According to PEP 249, this is -1 when query result is not applicable. + if cur.rowcount >= 0: + self.log.info("Rows affected: %s", cur.rowcount) + + @overload + def run( + self, + sql: str | Iterable[str], + autocommit: bool = ..., + parameters: Iterable | Mapping[str, Any] | None = ..., + handler: None = ..., + split_statements: bool = ..., + return_last: bool = ..., + ) -> None: ... + + @overload + def run( + self, + sql: str | Iterable[str], + autocommit: bool = ..., + parameters: Iterable | Mapping[str, Any] | None = ..., + handler: HANDLER = ..., + split_statements: bool = ..., + return_last: bool = ..., + ) -> T | list[T] | None: ... + + async def run( + self, + sql: str | Iterable[str], + autocommit: bool = False, + parameters: Iterable | Mapping[str, Any] | None = None, + handler: HANDLER | None = None, + split_statements: bool = False, + return_last: bool = True, + ) -> T | list[T] | None: + self.descriptions = [] + + if isinstance(sql, str): + if split_statements: + sql_list: Iterable[str] = self.split_sql_string( + sql=sql, + strip_semicolon=self.strip_semicolon, + ) + else: + sql_list = [sql] if sql.strip() else [] + else: + sql_list = sql + + if sql_list: + self.log.debug("Executing following statements against DB: %s", sql_list) + else: + raise ValueError("List of SQL statements is empty") + _last_result = None + async with self._create_autocommit_connection(autocommit) as conn: + async with self._get_cursor(conn) as cur: + results = [] + for sql_statement in sql_list: + await self._run_command(cur, sql_statement, parameters) + + if handler is not None: + handled = await self._maybe_await(handler(cur)) + result = self._make_common_data_structure(handled) + if handlers.return_single_query_results(sql, return_last, split_statements): + _last_result = result + _last_description = cur.description + else: + results.append(result) + self.descriptions.append(cur.description) + # If autocommit was set to False or db does not support autocommit, we do a manual commit. + if not self.get_autocommit(conn): + await self._maybe_await(conn.commit()) + # Logs all database messages or errors sent to the client + await self._maybe_await(self.get_db_log_messages(conn)) + if handler is None: + return None + if handlers.return_single_query_results(sql, return_last, split_statements): + self.descriptions = [_last_description] + return _last_result + return results From 3f488e80fe62f68b4e785f1a161073273057d7fb Mon Sep 17 00:00:00 2001 From: Karen Braganza Date: Tue, 3 Feb 2026 14:18:05 -0500 Subject: [PATCH 02/14] First draft of SQLExecuteQueryTrigger --- .../airflow/providers/common/sql/hooks/sql.py | 9 +- .../common/sql/operators/generic_transfer.py | 6 +- .../providers/common/sql/triggers/sql.py | 107 +++++++++++++++++- .../providers/common/sql/triggers/sql.pyi | 2 +- .../unit/common/sql/triggers/test_sql.py | 6 +- 5 files changed, 117 insertions(+), 13 deletions(-) diff --git a/providers/common/sql/src/airflow/providers/common/sql/hooks/sql.py b/providers/common/sql/src/airflow/providers/common/sql/hooks/sql.py index 237cd6aab6fc1..7bc268606bf5d 100644 --- a/providers/common/sql/src/airflow/providers/common/sql/hooks/sql.py +++ b/providers/common/sql/src/airflow/providers/common/sql/hooks/sql.py @@ -1146,8 +1146,13 @@ async def get_conn(self) -> Any: host=host, port=cast("int", db.port), username=login, schema=schema ) - async def _maybe_await(self, result: Any) -> Any: - return await result if inspect.isawaitable(result) else result + async def _maybe_await(self, func: Callable, *args, **kwargs) -> Any: + result = func(*args, **kwargs) + + if inspect.isawaitable(result): + return await result + + return await sync_to_async(func)(*args, **kwargs) @asynccontextmanager async def _create_autocommit_connection(self, autocommit: bool = False): diff --git a/providers/common/sql/src/airflow/providers/common/sql/operators/generic_transfer.py b/providers/common/sql/src/airflow/providers/common/sql/operators/generic_transfer.py index 6c5e149cd3450..cb87ca408ab45 100644 --- a/providers/common/sql/src/airflow/providers/common/sql/operators/generic_transfer.py +++ b/providers/common/sql/src/airflow/providers/common/sql/operators/generic_transfer.py @@ -23,7 +23,7 @@ from airflow.providers.common.compat.sdk import AirflowException, BaseHook, BaseOperator from airflow.providers.common.sql.hooks.sql import DbApiHook -from airflow.providers.common.sql.triggers.sql import SQLExecuteQueryTrigger +from airflow.providers.common.sql.triggers.sql import SQLGenericTransferTrigger if TYPE_CHECKING: import jinja2 @@ -147,7 +147,7 @@ def execute(self, context: Context): if self.page_size and isinstance(self.sql, str): self.defer( - trigger=SQLExecuteQueryTrigger( + trigger=SQLGenericTransferTrigger( conn_id=self.source_conn_id, hook_params=self.source_hook_params, sql=self.get_paginated_sql(0), @@ -207,7 +207,7 @@ def execute_complete( ) self.defer( - trigger=SQLExecuteQueryTrigger( + trigger=SQLGenericTransferTrigger( conn_id=self.source_conn_id, hook_params=self.source_hook_params, sql=self.get_paginated_sql(offset), diff --git a/providers/common/sql/src/airflow/providers/common/sql/triggers/sql.py b/providers/common/sql/src/airflow/providers/common/sql/triggers/sql.py index c1188e0e1b348..0e021606cbb51 100644 --- a/providers/common/sql/src/airflow/providers/common/sql/triggers/sql.py +++ b/providers/common/sql/src/airflow/providers/common/sql/triggers/sql.py @@ -17,10 +17,14 @@ # under the License. from __future__ import annotations +import importlib +import sys from typing import TYPE_CHECKING +from asgiref.sync import sync_to_async + from airflow.providers.common.compat.sdk import AirflowException, BaseHook -from airflow.providers.common.sql.hooks.sql import DbApiHook +from airflow.providers.common.sql.hooks.sql import DbApiHook, DbApiHookAsync from airflow.triggers.base import BaseTrigger, TriggerEvent if TYPE_CHECKING: @@ -28,9 +32,9 @@ from typing import Any -class SQLExecuteQueryTrigger(BaseTrigger): +class SQLGenericTransferTrigger(BaseTrigger): """ - A trigger that executes SQL code in async mode. + A SQL trigger that executes SQL to get records in async mode. :param sql: the sql statement to be executed (str) or a list of sql statements to execute :param conn_id: the connection ID used to connect to the database @@ -50,7 +54,7 @@ def __init__( self.hook_params = hook_params def serialize(self) -> tuple[str, dict[str, Any]]: - """Serialize the SQLExecuteQueryTrigger arguments and classpath.""" + """Serialize the SQLGenericTransferTrigger arguments and classpath.""" return ( f"{self.__class__.__module__}.{self.__class__.__name__}", { @@ -94,3 +98,98 @@ async def run(self) -> AsyncIterator[TriggerEvent]: except Exception as e: self.log.exception("An error occurred: %s", e) yield TriggerEvent({"status": "failure", "message": str(e)}) + + +class SQLExecuteQueryTrigger(BaseTrigger): + """ + A SQL trigger that executes SQL code in async mode. + + :param sql: the sql statement to be executed (str) or a list of sql statements to execute + :param conn_id: the connection ID used to connect to the database + :param hook_params: hook parameters + """ + + def __init__( + self, + sql: str | list[str], + conn_id: str, + autocommit, + parameters, + handler_path, + split_statements, + return_last, + ): + super().__init__() + self.sql = sql + self.conn_id = conn_id + self.autocommit = (autocommit,) + self.parameters = (parameters,) + self.handler_path = (handler_path,) + self.split_statements = (split_statements,) + self.return_last = return_last + + def serialize(self) -> tuple[str, dict[str, Any]]: + """Serialize the SQLExecuteQueryTrigger arguments and classpath.""" + return ( + f"{self.__class__.__module__}.{self.__class__.__name__}", + { + "sql": self.sql, + "conn_id": self.conn_id, + "autocommit": self.autocommit, + "parameters": self.parameters, + "handler_path": self.handler_path, + "split_statements": self.split_statements, + "return_last": self.return_last, + }, + ) + + async def _import_from_handler_path(self): + """Import the handler callable from the path provided by the user.""" + module_path, func_name = self.handler_path.rsplit(".", 1) + if module_path in sys.modules: + module = await sync_to_async(importlib.reload)(sys.modules[module_path]) + module = await sync_to_async(importlib.import_module)(module_path) + return getattr(module, func_name) + + async def get_hook(self) -> DbApiHookAsync: + """ + Return DbApiHookAsync. + + :return: DbApiHookAsync for this connection + """ + connection = sync_to_async(BaseHook.get_connection)(self.conn_id) + hook = sync_to_async(connection.get_hook) + if not isinstance(hook, DbApiHookAsync): + raise AirflowException( + f"You are trying to use `common-sql` with {hook.__class__.__name__}," + " but its provider does not support it. Please upgrade the provider to a version that" + " supports `common-sql`. The hook class should be a subclass of DbApiHookAsync" + f" Got {hook.__class__.__name__} hook with class hierarchy: {hook.__class__.mro()}" + ) + return hook + + async def run(self) -> AsyncIterator[TriggerEvent]: + try: + hook = await self.get_hook() + handler = None + if self.handler_path: + handler = await self._import_from_handler_path() + + self.log.info("Extracting data from %s", self.conn_id) + self.log.info("Executing: \n %s", self.sql) + + results = await hook.run( + sql=self.sql, + autocommit=self.autocommit, + parameters=self.parameters, + handler=handler, + split_statements=self.split_statements, + return_last=self.return_last, + ) + + self.log.info("Executing query from %s done!", self.conn_id) + self.log.debug("results: %s", results) + yield TriggerEvent({"status": "success", "results": results}) + except Exception as e: + self.log.exception("An error occurred: %s", e) + yield TriggerEvent({"status": "failure", "message": str(e)}) diff --git a/providers/common/sql/src/airflow/providers/common/sql/triggers/sql.pyi b/providers/common/sql/src/airflow/providers/common/sql/triggers/sql.pyi index 529da972c7dc2..fece1e7d56678 100644 --- a/providers/common/sql/src/airflow/providers/common/sql/triggers/sql.pyi +++ b/providers/common/sql/src/airflow/providers/common/sql/triggers/sql.pyi @@ -34,7 +34,7 @@ from typing import Any from airflow.triggers.base import BaseTrigger as BaseTrigger, TriggerEvent as TriggerEvent -class SQLExecuteQueryTrigger(BaseTrigger): +class SQLGenericTransferTrigger(BaseTrigger): def __init__( self, sql: str | list[str], conn_id: str, hook_params: dict | None = None, **kwargs ) -> None: ... diff --git a/providers/common/sql/tests/unit/common/sql/triggers/test_sql.py b/providers/common/sql/tests/unit/common/sql/triggers/test_sql.py index e09fba590be73..9b5a003ae0dac 100644 --- a/providers/common/sql/tests/unit/common/sql/triggers/test_sql.py +++ b/providers/common/sql/tests/unit/common/sql/triggers/test_sql.py @@ -20,7 +20,7 @@ from airflow.models.connection import Connection from airflow.providers.common.sql.hooks.sql import DbApiHook -from airflow.providers.common.sql.triggers.sql import SQLExecuteQueryTrigger +from airflow.providers.common.sql.triggers.sql import SQLGenericTransferTrigger from airflow.triggers.base import TriggerEvent try: @@ -35,7 +35,7 @@ from tests_common.test_utils.operators.run_deferrable import run_trigger -class TestSQLExecuteQueryTrigger: +class TestSQLGenericTransferTrigger: @mock.patch(f"{BASEHOOK_PATCH_PATH}.get_connection") def test_run(self, mock_get_connection): data = [(1, "Alice"), (2, "Bob")] @@ -45,7 +45,7 @@ def test_run(self, mock_get_connection): mock_get_connection.return_value = mock_connection mock_connection.get_hook.side_effect = lambda hook_params: mock_hook - trigger = SQLExecuteQueryTrigger(sql="SELECT * FROM users;", conn_id="test_conn_id") + trigger = SQLGenericTransferTrigger(sql="SELECT * FROM users;", conn_id="test_conn_id") actual = run_trigger(trigger) assert len(actual) == 1 From fa4a172c440ed90c9536a3d276107926fb847fb8 Mon Sep 17 00:00:00 2001 From: Karen Braganza Date: Tue, 3 Feb 2026 15:04:08 -0500 Subject: [PATCH 03/14] First draft of deferrable SQLExecuteQueryOperator --- .../providers/common/sql/operators/sql.py | 83 +++++++++++++------ 1 file changed, 59 insertions(+), 24 deletions(-) diff --git a/providers/common/sql/src/airflow/providers/common/sql/operators/sql.py b/providers/common/sql/src/airflow/providers/common/sql/operators/sql.py index ab9569ee18f9c..da44aa53538b1 100644 --- a/providers/common/sql/src/airflow/providers/common/sql/operators/sql.py +++ b/providers/common/sql/src/airflow/providers/common/sql/operators/sql.py @@ -23,7 +23,10 @@ from functools import cached_property from typing import TYPE_CHECKING, Any, ClassVar, NoReturn, SupportsAbs +from sqlalchemy import inspect + from airflow import XComArg +from airflow.configuration import conf from airflow.models import SkipMixin from airflow.providers.common.compat.sdk import ( AirflowException, @@ -34,6 +37,7 @@ ) from airflow.providers.common.sql.hooks.handlers import fetch_all_handler, return_single_query_results from airflow.providers.common.sql.hooks.sql import DbApiHook +from airflow.providers.common.sql.triggers.sql import SQLExecuteQueryTrigger from airflow.utils.helpers import merge_dicts if TYPE_CHECKING: @@ -320,6 +324,7 @@ class SQLExecuteQueryOperator(BaseSQLOperator): :param requires_result_fetch: (optional) if True, ensures that query results are fetched before completing execution. If `do_xcom_push` is True, results are fetched automatically, making this parameter redundant. (default: False). + :param deferrable: (optional) Run operator in the deferrable mode. .. seealso:: For more information on how to use this operator, take a look at the guide: @@ -351,6 +356,7 @@ def __init__( return_last: bool = True, show_return_value_in_logs: bool = False, requires_result_fetch: bool = False, + deferrable: bool = conf.getboolean("operators", "default_deferrable", fallback=False), **kwargs, ) -> None: super().__init__(conn_id=conn_id, database=database, **kwargs) @@ -363,6 +369,7 @@ def __init__( self.return_last = return_last self.show_return_value_in_logs = show_return_value_in_logs self.requires_result_fetch = requires_result_fetch + self.deferrable = deferrable def _process_output( self, results: list[Any], descriptions: list[Sequence[Sequence] | None] @@ -390,33 +397,61 @@ def _process_output( def _should_run_output_processing(self) -> bool: return self.do_xcom_push + def _get_handler_import_path(self) -> str: + if not inspect.isfunction(self.handler): + raise ValueError("The handler must be a function object.") + qualname = self.handler.__qualname__ + if "" in qualname or "" in qualname: + raise ValueError("The handler must not be a nested/local or lambda function.") + module = self.handler.__module__ + top_name = qualname.split(".", 1)[0] + return f"{module}.{top_name}" + def execute(self, context): self.log.info("Executing: %s", self.sql) - hook = self.get_db_hook() - if self.split_statements is not None: - extra_kwargs = {"split_statements": self.split_statements} + handler_path = None + if self.handler: + handler_path = self._get_handler_import_path() + if self.deferrable: + self.defer( + timeout=self.timeout, + trigger=SQLExecuteQueryTrigger( + sql=self.sql, + autocommit=self.autocommit, + parameters=self.parameters, + handler_path=handler_path + if self._should_run_output_processing() or self.requires_result_fetch + else None, + split_statements=self.split_statements, + return_last=self.return_last, + ), + ) else: - extra_kwargs = {} - output = hook.run( - sql=self.sql, - autocommit=self.autocommit, - parameters=self.parameters, - handler=self.handler - if self._should_run_output_processing() or self.requires_result_fetch - else None, - return_last=self.return_last, - **extra_kwargs, - ) - if not self._should_run_output_processing(): - return None - if return_single_query_results(self.sql, self.return_last, self.split_statements): - # For simplicity, we pass always list as input to _process_output, regardless if - # single query results are going to be returned, and we return the first element - # of the list in this case from the (always) list returned by _process_output - return self._process_output([output], hook.descriptions)[-1] - result = self._process_output(output, hook.descriptions) - self.log.info("result: %s", result) - return result + hook = self.get_db_hook() + if self.split_statements is not None: + extra_kwargs = {"split_statements": self.split_statements} + else: + extra_kwargs = {} + output = hook.run( + sql=self.sql, + autocommit=self.autocommit, + parameters=self.parameters, + handler=self.handler + if self._should_run_output_processing() or self.requires_result_fetch + else None, + return_last=self.return_last, + **extra_kwargs, + ) + if not self._should_run_output_processing(): + return None + if return_single_query_results(self.sql, self.return_last, self.split_statements): + # For simplicity, we pass always list as input to _process_output, regardless if + # single query results are going to be returned, and we return the first element + # of the list in this case from the (always) list returned by _process_output + return self._process_output([output], hook.descriptions)[-1] + result = self._process_output(output, hook.descriptions) + self.log.info("result: %s", result) + return result def prepare_template(self) -> None: """Parse template file for attribute parameters.""" From 9d156d8d0a5c604b220a484898221668bd20a54c Mon Sep 17 00:00:00 2001 From: Karen Braganza Date: Tue, 10 Feb 2026 09:55:00 -0500 Subject: [PATCH 04/14] Prek changes --- .../airflow/providers/common/sql/hooks/sql.py | 42 ++++++++++--------- .../providers/common/sql/operators/sql.py | 3 +- .../providers/common/sql/triggers/sql.py | 4 +- 3 files changed, 25 insertions(+), 24 deletions(-) diff --git a/providers/common/sql/src/airflow/providers/common/sql/hooks/sql.py b/providers/common/sql/src/airflow/providers/common/sql/hooks/sql.py index 7bc268606bf5d..dfa16ee5ce9ea 100644 --- a/providers/common/sql/src/airflow/providers/common/sql/hooks/sql.py +++ b/providers/common/sql/src/airflow/providers/common/sql/hooks/sql.py @@ -18,9 +18,9 @@ import asyncio import contextlib +import inspect import warnings from collections.abc import ( - AsyncIterator, Awaitable, Callable, Generator, @@ -42,7 +42,7 @@ from more_itertools import chunked try: - from sqlalchemy import create_engine, inspect + from sqlalchemy import create_engine from sqlalchemy.engine import make_url from sqlalchemy.exc import ArgumentError, NoSuchModuleError except ImportError: @@ -75,7 +75,8 @@ T = TypeVar("T") -HANDLER = Callable[[Any], T | Awaitable[T]] +HANDLER = Callable[[Any], Any | Awaitable[Any]] +ROW = tuple[Any, ...] SQL_PLACEHOLDERS = frozenset({"%s", "?"}) WARNING_MESSAGE = """Import of {} from the 'airflow.providers.common.sql.hooks' module is deprecated and will be removed in the future. Please import it from 'airflow.providers.common.sql.hooks.handlers'.""" @@ -1147,15 +1148,15 @@ async def get_conn(self) -> Any: ) async def _maybe_await(self, func: Callable, *args, **kwargs) -> Any: - result = func(*args, **kwargs) - - if inspect.isawaitable(result): - return await result - - return await sync_to_async(func)(*args, **kwargs) + if inspect.iscoroutinefunction(func): + result = func(*args, **kwargs) + if inspect.isawaitable(result): + return await result + else: + return await sync_to_async(func)(*args, **kwargs) @asynccontextmanager - async def _create_autocommit_connection(self, autocommit: bool = False): + async def _create_autocommit_connection(self, autocommit: bool = False): # type: ignore[override] conn = await self.get_conn() try: if self.supports_autocommit: @@ -1166,8 +1167,8 @@ async def _create_autocommit_connection(self, autocommit: bool = False): await self._maybe_await(close()) @asynccontextmanager - async def _get_cursor(self, conn: Any) -> AsyncIterator[Any]: - cur_or_cm = await self._maybe_await(conn.cursor()) + async def _get_cursor(self, conn): + cur_or_cm = await self._maybe_await(conn.cursor) if hasattr(cur_or_cm, "__aenter__") and hasattr(cur_or_cm, "__aexit__"): async with cur_or_cm as cur: yield cur @@ -1179,7 +1180,7 @@ async def _get_cursor(self, conn: Any) -> AsyncIterator[Any]: finally: close = getattr(cur, "aclose", None) or getattr(cur, "close", None) if close: - await self._maybe_await(close()) + await self._maybe_await(close) async def _run_command(self, cur, sql_statement, parameters): """Run a statement using an already open cursor.""" @@ -1190,9 +1191,9 @@ async def _run_command(self, cur, sql_statement, parameters): # If we're using psycopg3, we might need to handle parameters differently if isinstance(parameters, list): parameters = tuple(parameters) - await self._maybe_await(cur.execute(sql_statement, parameters)) + await self._maybe_await(cur.execute, sql_statement, parameters) else: - await self._maybe_await(cur.execute(sql_statement)) + await self._maybe_await(cur.execute, sql_statement) # According to PEP 249, this is -1 when query result is not applicable. if cur.rowcount >= 0: @@ -1218,7 +1219,7 @@ def run( handler: HANDLER = ..., split_statements: bool = ..., return_last: bool = ..., - ) -> T | list[T] | None: ... + ) -> tuple | list | list[tuple] | list[list[tuple] | tuple] | None: ... async def run( self, @@ -1228,7 +1229,7 @@ async def run( handler: HANDLER | None = None, split_statements: bool = False, return_last: bool = True, - ) -> T | list[T] | None: + ) -> tuple | list | list[tuple] | list[list[tuple] | tuple] | None: self.descriptions = [] if isinstance(sql, str): @@ -1247,6 +1248,7 @@ async def run( else: raise ValueError("List of SQL statements is empty") _last_result = None + _last_description = None async with self._create_autocommit_connection(autocommit) as conn: async with self._get_cursor(conn) as cur: results = [] @@ -1254,7 +1256,7 @@ async def run( await self._run_command(cur, sql_statement, parameters) if handler is not None: - handled = await self._maybe_await(handler(cur)) + handled = await self._maybe_await(handler, cur) result = self._make_common_data_structure(handled) if handlers.return_single_query_results(sql, return_last, split_statements): _last_result = result @@ -1264,9 +1266,9 @@ async def run( self.descriptions.append(cur.description) # If autocommit was set to False or db does not support autocommit, we do a manual commit. if not self.get_autocommit(conn): - await self._maybe_await(conn.commit()) + await self._maybe_await(conn.commit) # Logs all database messages or errors sent to the client - await self._maybe_await(self.get_db_log_messages(conn)) + await self._maybe_await(self.get_db_log_messages, conn) if handler is None: return None if handlers.return_single_query_results(sql, return_last, split_statements): diff --git a/providers/common/sql/src/airflow/providers/common/sql/operators/sql.py b/providers/common/sql/src/airflow/providers/common/sql/operators/sql.py index da44aa53538b1..698285bdcdc34 100644 --- a/providers/common/sql/src/airflow/providers/common/sql/operators/sql.py +++ b/providers/common/sql/src/airflow/providers/common/sql/operators/sql.py @@ -18,13 +18,12 @@ from __future__ import annotations import ast +import inspect import re from collections.abc import Callable, Iterable, Mapping, Sequence from functools import cached_property from typing import TYPE_CHECKING, Any, ClassVar, NoReturn, SupportsAbs -from sqlalchemy import inspect - from airflow import XComArg from airflow.configuration import conf from airflow.models import SkipMixin diff --git a/providers/common/sql/src/airflow/providers/common/sql/triggers/sql.py b/providers/common/sql/src/airflow/providers/common/sql/triggers/sql.py index 0e021606cbb51..a8a7eb2825163 100644 --- a/providers/common/sql/src/airflow/providers/common/sql/triggers/sql.py +++ b/providers/common/sql/src/airflow/providers/common/sql/triggers/sql.py @@ -157,8 +157,8 @@ async def get_hook(self) -> DbApiHookAsync: :return: DbApiHookAsync for this connection """ - connection = sync_to_async(BaseHook.get_connection)(self.conn_id) - hook = sync_to_async(connection.get_hook) + connection = await sync_to_async(BaseHook.get_connection)(self.conn_id) + hook = await sync_to_async(connection.get_hook) if not isinstance(hook, DbApiHookAsync): raise AirflowException( f"You are trying to use `common-sql` with {hook.__class__.__name__}," From ecd3d1ae087b567eb4ed1a57cda60b24953a927e Mon Sep 17 00:00:00 2001 From: Karen Braganza Date: Tue, 10 Feb 2026 11:11:34 -0500 Subject: [PATCH 05/14] mypy changes --- .../src/airflow/providers/common/sql/hooks/sql.py | 12 ++++++------ 1 file changed, 6 insertions(+), 6 deletions(-) diff --git a/providers/common/sql/src/airflow/providers/common/sql/hooks/sql.py b/providers/common/sql/src/airflow/providers/common/sql/hooks/sql.py index dfa16ee5ce9ea..eea79498e60f1 100644 --- a/providers/common/sql/src/airflow/providers/common/sql/hooks/sql.py +++ b/providers/common/sql/src/airflow/providers/common/sql/hooks/sql.py @@ -37,12 +37,12 @@ import sqlparse from asgiref.sync import sync_to_async -from deprecated import deprecated +from deprecated import deprecated # type: ignore[import-untyped] from methodtools import lru_cache from more_itertools import chunked try: - from sqlalchemy import create_engine + from sqlalchemy import create_engine, inspect as sa_inspect from sqlalchemy.engine import make_url from sqlalchemy.exc import ArgumentError, NoSuchModuleError except ImportError: @@ -355,7 +355,7 @@ def inspector(self) -> Inspector: "SQLAlchemy is required for database inspection. " "Install it with: pip install 'apache-airflow-providers-common-sql[sqlalchemy]'" ) - return inspect(self.get_sqlalchemy_engine()) + return sa_inspect(self.get_sqlalchemy_engine()) @cached_property def dialect_name(self) -> str: @@ -1200,7 +1200,7 @@ async def _run_command(self, cur, sql_statement, parameters): self.log.info("Rows affected: %s", cur.rowcount) @overload - def run( + async def run_async( self, sql: str | Iterable[str], autocommit: bool = ..., @@ -1211,7 +1211,7 @@ def run( ) -> None: ... @overload - def run( + async def run_async( self, sql: str | Iterable[str], autocommit: bool = ..., @@ -1221,7 +1221,7 @@ def run( return_last: bool = ..., ) -> tuple | list | list[tuple] | list[list[tuple] | tuple] | None: ... - async def run( + async def run_async( self, sql: str | Iterable[str], autocommit: bool = False, From f35879b8a12066a1441ca727d4043378ed7c2921 Mon Sep 17 00:00:00 2001 From: Karen Braganza Date: Tue, 10 Feb 2026 11:28:45 -0500 Subject: [PATCH 06/14] mypy changes --- .../sql/src/airflow/providers/common/sql/triggers/sql.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/providers/common/sql/src/airflow/providers/common/sql/triggers/sql.py b/providers/common/sql/src/airflow/providers/common/sql/triggers/sql.py index a8a7eb2825163..f2e5947234090 100644 --- a/providers/common/sql/src/airflow/providers/common/sql/triggers/sql.py +++ b/providers/common/sql/src/airflow/providers/common/sql/triggers/sql.py @@ -158,7 +158,7 @@ async def get_hook(self) -> DbApiHookAsync: :return: DbApiHookAsync for this connection """ connection = await sync_to_async(BaseHook.get_connection)(self.conn_id) - hook = await sync_to_async(connection.get_hook) + hook = await sync_to_async(connection.get_hook)() if not isinstance(hook, DbApiHookAsync): raise AirflowException( f"You are trying to use `common-sql` with {hook.__class__.__name__}," @@ -178,7 +178,7 @@ async def run(self) -> AsyncIterator[TriggerEvent]: self.log.info("Extracting data from %s", self.conn_id) self.log.info("Executing: \n %s", self.sql) - results = await hook.run( + results = await hook.run_async( sql=self.sql, autocommit=self.autocommit, parameters=self.parameters, From fdd676decca95662c92e65405515a932b7c8842e Mon Sep 17 00:00:00 2001 From: Karen Braganza Date: Tue, 10 Feb 2026 11:48:55 -0500 Subject: [PATCH 07/14] mypy changes --- .../airflow/providers/common/sql/triggers/sql.py | 15 ++++++++++----- 1 file changed, 10 insertions(+), 5 deletions(-) diff --git a/providers/common/sql/src/airflow/providers/common/sql/triggers/sql.py b/providers/common/sql/src/airflow/providers/common/sql/triggers/sql.py index f2e5947234090..784ec907a2bba 100644 --- a/providers/common/sql/src/airflow/providers/common/sql/triggers/sql.py +++ b/providers/common/sql/src/airflow/providers/common/sql/triggers/sql.py @@ -31,6 +31,11 @@ from collections.abc import AsyncIterator from typing import Any +from collections.abc import ( + Iterable, + Mapping, +) + class SQLGenericTransferTrigger(BaseTrigger): """ @@ -113,11 +118,11 @@ def __init__( self, sql: str | list[str], conn_id: str, - autocommit, - parameters, - handler_path, - split_statements, - return_last, + autocommit: bool | None = None, + parameters: Iterable | Mapping[str, Any] | None = None, + handler_path: str | None = None, + split_statements: bool | None = None, + return_last: bool | None = None, ): super().__init__() self.sql = sql From c0ff9bdb42f5b26ac831ff13cb498d40aafc4363 Mon Sep 17 00:00:00 2001 From: Karen Braganza Date: Tue, 10 Feb 2026 11:58:35 -0500 Subject: [PATCH 08/14] mypy changes --- .../common/sql/src/airflow/providers/common/sql/triggers/sql.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/providers/common/sql/src/airflow/providers/common/sql/triggers/sql.py b/providers/common/sql/src/airflow/providers/common/sql/triggers/sql.py index 784ec907a2bba..7d2869cb0a686 100644 --- a/providers/common/sql/src/airflow/providers/common/sql/triggers/sql.py +++ b/providers/common/sql/src/airflow/providers/common/sql/triggers/sql.py @@ -116,7 +116,7 @@ class SQLExecuteQueryTrigger(BaseTrigger): def __init__( self, - sql: str | list[str], + sql: str | Iterable[str], conn_id: str, autocommit: bool | None = None, parameters: Iterable | Mapping[str, Any] | None = None, From 12cab64bc8f629eea3e654be4f0f4b308e08bc99 Mon Sep 17 00:00:00 2001 From: Karen Braganza Date: Tue, 10 Feb 2026 12:10:39 -0500 Subject: [PATCH 09/14] mypy changes --- .../common/sql/src/airflow/providers/common/sql/triggers/sql.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/providers/common/sql/src/airflow/providers/common/sql/triggers/sql.py b/providers/common/sql/src/airflow/providers/common/sql/triggers/sql.py index 7d2869cb0a686..a3b377bd238da 100644 --- a/providers/common/sql/src/airflow/providers/common/sql/triggers/sql.py +++ b/providers/common/sql/src/airflow/providers/common/sql/triggers/sql.py @@ -119,7 +119,7 @@ def __init__( sql: str | Iterable[str], conn_id: str, autocommit: bool | None = None, - parameters: Iterable | Mapping[str, Any] | None = None, + parameters: Iterable[Any] | Mapping[str, Any] | None = None, handler_path: str | None = None, split_statements: bool | None = None, return_last: bool | None = None, From 0d991ef9dcf251e73cfbc2c4c59a5d97a60ebe42 Mon Sep 17 00:00:00 2001 From: Karen Braganza Date: Tue, 10 Feb 2026 12:21:07 -0500 Subject: [PATCH 10/14] mypy changes --- .../sql/src/airflow/providers/common/sql/triggers/sql.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/providers/common/sql/src/airflow/providers/common/sql/triggers/sql.py b/providers/common/sql/src/airflow/providers/common/sql/triggers/sql.py index a3b377bd238da..6f2b27c22f45c 100644 --- a/providers/common/sql/src/airflow/providers/common/sql/triggers/sql.py +++ b/providers/common/sql/src/airflow/providers/common/sql/triggers/sql.py @@ -118,11 +118,11 @@ def __init__( self, sql: str | Iterable[str], conn_id: str, - autocommit: bool | None = None, + autocommit: bool, + split_statements: bool, + return_last: bool, parameters: Iterable[Any] | Mapping[str, Any] | None = None, handler_path: str | None = None, - split_statements: bool | None = None, - return_last: bool | None = None, ): super().__init__() self.sql = sql From 8c6b14b29a718e1438056502ef9156585201ee33 Mon Sep 17 00:00:00 2001 From: Karen Braganza Date: Tue, 10 Feb 2026 12:30:27 -0500 Subject: [PATCH 11/14] mypy changes --- .../sql/src/airflow/providers/common/sql/triggers/sql.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/providers/common/sql/src/airflow/providers/common/sql/triggers/sql.py b/providers/common/sql/src/airflow/providers/common/sql/triggers/sql.py index 6f2b27c22f45c..8dc8731797214 100644 --- a/providers/common/sql/src/airflow/providers/common/sql/triggers/sql.py +++ b/providers/common/sql/src/airflow/providers/common/sql/triggers/sql.py @@ -127,10 +127,10 @@ def __init__( super().__init__() self.sql = sql self.conn_id = conn_id - self.autocommit = (autocommit,) - self.parameters = (parameters,) - self.handler_path = (handler_path,) - self.split_statements = (split_statements,) + self.autocommit = autocommit + self.parameters = parameters + self.handler_path = handler_path + self.split_statements = split_statements self.return_last = return_last def serialize(self) -> tuple[str, dict[str, Any]]: From 30e699c78bfab31cf937e8e0c7538fbc6ceb4235 Mon Sep 17 00:00:00 2001 From: Karen Braganza Date: Tue, 10 Feb 2026 12:43:11 -0500 Subject: [PATCH 12/14] mypy changes --- .../providers/common/sql/triggers/sql.py | 37 +++++++++++++------ 1 file changed, 26 insertions(+), 11 deletions(-) diff --git a/providers/common/sql/src/airflow/providers/common/sql/triggers/sql.py b/providers/common/sql/src/airflow/providers/common/sql/triggers/sql.py index 8dc8731797214..b642492d77277 100644 --- a/providers/common/sql/src/airflow/providers/common/sql/triggers/sql.py +++ b/providers/common/sql/src/airflow/providers/common/sql/triggers/sql.py @@ -183,18 +183,33 @@ async def run(self) -> AsyncIterator[TriggerEvent]: self.log.info("Extracting data from %s", self.conn_id) self.log.info("Executing: \n %s", self.sql) - results = await hook.run_async( - sql=self.sql, - autocommit=self.autocommit, - parameters=self.parameters, - handler=handler, - split_statements=self.split_statements, - return_last=self.return_last, - ) + if handler: + results = await hook.run_async( + sql=self.sql, + autocommit=self.autocommit, + parameters=self.parameters, + handler=handler, + split_statements=self.split_statements, + return_last=self.return_last, + ) + + self.log.info("Executing query from %s done!", self.conn_id) + self.log.debug("results: %s", results) + yield TriggerEvent({"status": "success", "results": results}) + + else: + await hook.run_async( + sql=self.sql, + autocommit=self.autocommit, + parameters=self.parameters, + handler=handler, + split_statements=self.split_statements, + return_last=self.return_last, + ) + + self.log.info("Executing query from %s done!", self.conn_id) + yield TriggerEvent({"status": "success"}) - self.log.info("Executing query from %s done!", self.conn_id) - self.log.debug("results: %s", results) - yield TriggerEvent({"status": "success", "results": results}) except Exception as e: self.log.exception("An error occurred: %s", e) yield TriggerEvent({"status": "failure", "message": str(e)}) From f10d6f2e1eb1288fb073651c61522daad34590e0 Mon Sep 17 00:00:00 2001 From: Karen Braganza Date: Tue, 10 Feb 2026 13:26:31 -0500 Subject: [PATCH 13/14] Make changes to use a callable in the trigger --- .../airflow/providers/common/sql/operators/sql.py | 13 +++++++------ .../airflow/providers/common/sql/triggers/sql.py | 10 +++++++--- 2 files changed, 14 insertions(+), 9 deletions(-) diff --git a/providers/common/sql/src/airflow/providers/common/sql/operators/sql.py b/providers/common/sql/src/airflow/providers/common/sql/operators/sql.py index 698285bdcdc34..76ba290addacd 100644 --- a/providers/common/sql/src/airflow/providers/common/sql/operators/sql.py +++ b/providers/common/sql/src/airflow/providers/common/sql/operators/sql.py @@ -399,19 +399,20 @@ def _should_run_output_processing(self) -> bool: def _get_handler_import_path(self) -> str: if not inspect.isfunction(self.handler): raise ValueError("The handler must be a function object.") - qualname = self.handler.__qualname__ + module = getattr(self.handler, "__module__", None) + qualname = getattr(self.handler, "__qualname__", None) + if not module or not qualname: + raise ValueError("handler must have __module__ and __qualname__") if "" in qualname or "" in qualname: raise ValueError("The handler must not be a nested/local or lambda function.") - module = self.handler.__module__ - top_name = qualname.split(".", 1)[0] - return f"{module}.{top_name}" + return f"{module}:{qualname}" def execute(self, context): self.log.info("Executing: %s", self.sql) handler_path = None - if self.handler: - handler_path = self._get_handler_import_path() if self.deferrable: + if self.handler: + handler_path = self._get_handler_import_path() self.defer( timeout=self.timeout, trigger=SQLExecuteQueryTrigger( diff --git a/providers/common/sql/src/airflow/providers/common/sql/triggers/sql.py b/providers/common/sql/src/airflow/providers/common/sql/triggers/sql.py index b642492d77277..7813e95346dbf 100644 --- a/providers/common/sql/src/airflow/providers/common/sql/triggers/sql.py +++ b/providers/common/sql/src/airflow/providers/common/sql/triggers/sql.py @@ -150,11 +150,15 @@ def serialize(self) -> tuple[str, dict[str, Any]]: async def _import_from_handler_path(self): """Import the handler callable from the path provided by the user.""" - module_path, func_name = self.handler_path.rsplit(".", 1) + module_path, qualname = self.handler_path.rsplit(":", 1) if module_path in sys.modules: module = await sync_to_async(importlib.reload)(sys.modules[module_path]) - module = await sync_to_async(importlib.import_module)(module_path) - return getattr(module, func_name) + else: + module = await sync_to_async(importlib.import_module)(module_path) + obj = module + for part in qualname.split("."): + obj = getattr(obj, part) + return obj async def get_hook(self) -> DbApiHookAsync: """ From 9cf520d71ba8017eb4154c54718d0eaad580b6b5 Mon Sep 17 00:00:00 2001 From: Karen Braganza Date: Tue, 10 Feb 2026 13:40:59 -0500 Subject: [PATCH 14/14] Add stub for mypy --- .../providers/common/sql/hooks/sql.pyi | 34 ++++++++++++++++++- 1 file changed, 33 insertions(+), 1 deletion(-) diff --git a/providers/common/sql/src/airflow/providers/common/sql/hooks/sql.pyi b/providers/common/sql/src/airflow/providers/common/sql/hooks/sql.pyi index 931ccd1c78059..986ab23b71ddb 100644 --- a/providers/common/sql/src/airflow/providers/common/sql/hooks/sql.pyi +++ b/providers/common/sql/src/airflow/providers/common/sql/hooks/sql.pyi @@ -32,7 +32,7 @@ Definition of the public interface for airflow.providers.common.sql.src.airflow.providers.common.sql.hooks.sql. """ -from collections.abc import Callable, Generator, Iterable, Mapping, MutableMapping, Sequence +from collections.abc import Awaitable, Callable, Generator, Iterable, Mapping, MutableMapping, Sequence from functools import cached_property as cached_property from typing import Any, Literal, Protocol, TypeVar, overload @@ -48,6 +48,7 @@ from airflow.providers.openlineage.extractors import OperatorLineage as Operator from airflow.providers.openlineage.sqlparser import DatabaseInfo as DatabaseInfo T = TypeVar("T") +HANDLER = Callable[[Any], Any | Awaitable[Any]] SQL_PLACEHOLDERS: Incomplete WARNING_MESSAGE: str @@ -205,3 +206,34 @@ class DbApiHook(BaseHook): split_statements: bool = ..., return_last: bool = ..., ) -> tuple | list | list[tuple] | list[list[tuple] | tuple] | None: ... + +class DbApiHookAsync(DbApiHook): + conn_name_attr: str + default_conn_name: str + strip_semicolon: bool + supports_autocommit: bool + supports_executemany: bool + connector: ConnectorProtocol | None + + def __init__(self, *args, schema: str | None = ..., log_sql: bool = ..., **kwargs: Any) -> None: ... + async def get_conn(self) -> Any: ... + @overload + async def run_async( + self, + sql: str | Iterable[str], + autocommit: bool = ..., + parameters: Iterable | Mapping[str, Any] | None = ..., + handler: None = ..., + split_statements: bool = ..., + return_last: bool = ..., + ) -> None: ... + @overload + async def run_async( + self, + sql: str | Iterable[str], + autocommit: bool = ..., + parameters: Iterable | Mapping[str, Any] | None = ..., + handler: HANDLER = ..., + split_statements: bool = ..., + return_last: bool = ..., + ) -> tuple | list | list[tuple] | list[list[tuple] | tuple] | None: ...