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..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 @@ -16,22 +16,33 @@ # under the License. from __future__ import annotations +import asyncio import contextlib +import inspect import warnings -from collections.abc import Callable, Generator, Iterable, Mapping, MutableMapping, Sequence -from contextlib import closing, contextmanager, suppress +from collections.abc import ( + 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 deprecated import deprecated +from asgiref.sync import sync_to_async +from deprecated import deprecated # type: ignore[import-untyped] from methodtools import lru_cache from more_itertools import chunked try: - from sqlalchemy import create_engine, inspect + from sqlalchemy import create_engine, inspect as sa_inspect from sqlalchemy.engine import make_url from sqlalchemy.exc import ArgumentError, NoSuchModuleError except ImportError: @@ -64,6 +75,8 @@ T = TypeVar("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'.""" @@ -342,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: @@ -1092,3 +1105,173 @@ 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, func: Callable, *args, **kwargs) -> Any: + 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): # type: ignore[override] + 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): + 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 + 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: ... + + async def run_async( + 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, + ) -> tuple | list | list[tuple] | list[list[tuple] | tuple] | 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 + _last_description = 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 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: ... 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/operators/sql.py b/providers/common/sql/src/airflow/providers/common/sql/operators/sql.py index ab9569ee18f9c..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 @@ -18,12 +18,14 @@ 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 airflow import XComArg +from airflow.configuration import conf from airflow.models import SkipMixin from airflow.providers.common.compat.sdk import ( AirflowException, @@ -34,6 +36,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 +323,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 +355,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 +368,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 +396,62 @@ 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.") + 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.") + return f"{module}:{qualname}" + 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.deferrable: + if self.handler: + handler_path = self._get_handler_import_path() + 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.""" 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..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 @@ -17,20 +17,29 @@ # 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: from collections.abc import AsyncIterator from typing import Any +from collections.abc import ( + Iterable, + Mapping, +) -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 +59,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 +103,117 @@ 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 | Iterable[str], + conn_id: str, + autocommit: bool, + split_statements: bool, + return_last: bool, + parameters: Iterable[Any] | Mapping[str, Any] | None = None, + handler_path: str | None = None, + ): + 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, qualname = self.handler_path.rsplit(":", 1) + if module_path in sys.modules: + module = await sync_to_async(importlib.reload)(sys.modules[module_path]) + 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: + """ + Return 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)() + 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) + + 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"}) + + 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