dabla commented on code in PR #62922:
URL: https://github.com/apache/airflow/pull/62922#discussion_r4093307676


##########
task-sdk/src/airflow/sdk/execution_time/executor.py:
##########
@@ -0,0 +1,399 @@
+#
+# Licensed to the Apache Software Foundation (ASF) under one
+# or more contributor license agreements.  See the NOTICE file
+# distributed with this work for additional information
+# regarding copyright ownership.  The ASF licenses this file
+# to you under the Apache License, Version 2.0 (the
+# "License"); you may not use this file except in compliance
+# with the License.  You may obtain a copy of the License at
+#
+#   http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing,
+# software distributed under the License is distributed on an
+# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, either express or implied.  See the License for the
+# specific language governing permissions and limitations
+# under the License.
+from __future__ import annotations
+
+import inspect
+import logging
+import os
+import threading
+import time
+from asyncio import (
+    FIRST_COMPLETED,
+    AbstractEventLoop,
+    CancelledError,
+    Future,
+    Semaphore,
+    Task,
+    TimeoutError as AsyncTimeoutError,
+    gather,
+    wait,
+    wait_for,
+    wrap_future,
+)
+from collections.abc import AsyncIterable, Callable, Iterator
+from concurrent.futures import Executor, ThreadPoolExecutor
+from contextlib import suppress
+from typing import TYPE_CHECKING, Any
+
+from airflow.sdk import BaseOperator, TaskInstanceState, timezone
+from airflow.sdk.definitions._internal.logging_mixin import LoggingMixin
+from airflow.sdk.exceptions import TaskDeferred
+from airflow.sdk.execution_time.context import set_current_context
+from airflow.sdk.execution_time.task_runner import (
+    _execute_async_task,
+    _execute_task,
+    _run_task_state_change_callbacks,
+)
+
+if TYPE_CHECKING:
+    from airflow.sdk import Context
+    from airflow.sdk.execution_time.task_runner import IndexedTaskInstance
+
+
+_log = logging.getLogger(__name__)
+
+
+class AsyncAwareExecutor(Executor):
+    """
+    Executes both sync and async functions concurrently.
+
+    Sync functions run in a ThreadPoolExecutor.
+    Async coroutines run on an asyncio event loop with a semaphore limit.
+
+    :param loop: Event loop used to schedule async tasks and coordinate mixed 
execution.
+    :param max_workers: Maximum concurrent workers used by both thread pool 
and async semaphore.
+    :param shutdown_timeout: Maximum time to wait, in seconds, for in-flight 
async tasks and
+        thread-pool workers to finish during ``shutdown(wait=True)``. Python 
threads cannot be
+        forcibly stopped, so a worker stuck in blocking user code (slow HTTP 
call, blocked C
+        extension, a deadlocked DB driver, ...) would otherwise hang 
``shutdown()`` forever.
+    """
+
+    def __init__(
+        self, loop: AbstractEventLoop, max_workers: int | None = None, 
shutdown_timeout: float = 10.0
+    ):
+        if max_workers is None:
+            max_workers = os.cpu_count() or 1
+        if max_workers <= 0:
+            raise ValueError("max_workers must be greater than 0")
+
+        self._loop = loop
+        self._max_workers = max_workers
+        self._shutdown_timeout = shutdown_timeout
+        self._semaphore = Semaphore(max_workers)
+        self._thread_pool = ThreadPoolExecutor(max_workers=max_workers)
+        self._async_tasks: set[Task[Any]] = set()
+        self._shutdown = False
+
+    def __enter__(self):
+        return self
+
+    def __exit__(self, exc_type, exc_val, exc_tb):
+        if exc_type is not None:
+            # On error path, cancel futures but still wait briefly for them to
+            # process CancelledError and release resources (e.g., threading
+            # locks). Without waiting, cancelled tasks that hold _thread_lock
+            # never execute their finally blocks, permanently leaking the lock
+            # and causing subsequent comms.send() calls to deadlock.
+            self.shutdown(wait=True, cancel_futures=True)
+        else:
+            self.shutdown(wait=True)
+
+    def shutdown(self, wait: bool = True, *, cancel_futures: bool = False) -> 
None:
+        if self._shutdown:
+            return
+
+        self._shutdown = True
+
+        if cancel_futures:
+            for task in list(self._async_tasks):
+                task.cancel()
+
+        if wait and self._async_tasks:
+            with suppress(TimeoutError, AsyncTimeoutError):
+                self._loop.run_until_complete(
+                    wait_for(
+                        gather(*self._async_tasks, return_exceptions=True),
+                        timeout=self._shutdown_timeout,
+                    )
+                )
+
+        # ThreadPoolExecutor.shutdown(wait=True) blocks until every worker 
thread
+        # finishes its current work item, with no way to bound that wait or 
forcibly
+        # stop a thread stuck in blocking user code (slow HTTP call, blocked C
+        # extension, a deadlocked DB driver, ...). Ask the pool to stop 
accepting new
+        # work (and cancel anything not yet started) up front, then bound how 
long we
+        # personally wait on the worker threads instead of blocking 
indefinitely.
+        self._thread_pool.shutdown(wait=False, cancel_futures=cancel_futures)
+
+        if wait:
+            threads = list(getattr(self._thread_pool, "_threads", ()))
+            deadline = time.monotonic() + self._shutdown_timeout
+            for thread in threads:
+                remaining = deadline - time.monotonic()
+                if remaining <= 0:
+                    break
+                thread.join(timeout=remaining)
+
+            stuck = [thread.name for thread in threads if thread.is_alive()]
+            if stuck:
+                _log.error(
+                    "%d worker thread(s) still running %.1fs after shutdown 
was requested; "
+                    "giving up waiting to avoid blocking indefinitely. This 
may leak resources. "
+                    "Affected threads: %s",
+                    len(stuck),
+                    self._shutdown_timeout,
+                    stuck,
+                )
+
+    def submit(self, func: Callable[..., Any] | Any, *args, **kwargs) -> 
Future[Any]:  # type: ignore[override]
+        """
+        Submit a callable for execution.
+
+        Always returns an asyncio.Future for consistency, whether the callable
+        is sync (run in thread pool) or async (run on the event loop).
+        """
+        if self._shutdown:
+            raise RuntimeError("cannot schedule new futures after shutdown")
+
+        if inspect.iscoroutine(func):
+            coro = func
+        elif inspect.iscoroutinefunction(func):
+            coro = func(*args, **kwargs)
+        else:
+            # Wrap thread pool future as asyncio.Future for consistent return 
type
+            return wrap_future(self._thread_pool.submit(func, *args, 
**kwargs), loop=self._loop)
+
+        async def guarded():
+            try:
+                async with self._semaphore:
+                    return await coro
+            except CancelledError:
+                # If cancellation occurs while waiting for the semaphore,
+                # the inner coroutine was never awaited. Close it to prevent
+                # "coroutine was never awaited" RuntimeWarning.
+                coro.close()
+                raise
+
+        task = self._loop.create_task(guarded())
+        self._async_tasks.add(task)
+        task.add_done_callback(self._async_tasks.discard)
+        return task
+
+    async def run_sync(self, func: Callable[..., Any], *args, **kwargs) -> Any:
+        """Run a sync callable in this executor's thread pool and await its 
result."""
+        future = self._thread_pool.submit(func, *args, **kwargs)
+        return await wrap_future(future, loop=self._loop)
+
+    def map(  # type: ignore[override]

Review Comment:
   Fixed in 9a5119f37d. Renamed to `imap_unordered`, after the 
`multiprocessing.Pool` method with the same contract, and the `type: ignore` is 
gone with it. The class still subclasses `Executor` but no longer overrides 
`map`, so anything holding a plain `Executor` keeps the submission-order 
guarantee, and completion-order streaming has to be asked for by name. A test 
pins that `AsyncAwareExecutor.map is Executor.map`.
   
   ---
   Drafted-by: Claude Fable 5.1; reviewed by @dabla before posting



##########
task-sdk/src/airflow/sdk/execution_time/executor.py:
##########
@@ -0,0 +1,441 @@
+#
+# Licensed to the Apache Software Foundation (ASF) under one
+# or more contributor license agreements.  See the NOTICE file
+# distributed with this work for additional information
+# regarding copyright ownership.  The ASF licenses this file
+# to you under the Apache License, Version 2.0 (the
+# "License"); you may not use this file except in compliance
+# with the License.  You may obtain a copy of the License at
+#
+#   http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing,
+# software distributed under the License is distributed on an
+# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, either express or implied.  See the License for the
+# specific language governing permissions and limitations
+# under the License.
+from __future__ import annotations
+
+import contextvars
+import inspect
+import logging
+import os
+import threading
+import time
+from asyncio import (
+    FIRST_COMPLETED,
+    AbstractEventLoop,
+    CancelledError,
+    Future,
+    Semaphore,
+    Task,
+    TimeoutError as AsyncTimeoutError,
+    gather,
+    wait,
+    wait_for,
+    wrap_future,
+)
+from collections.abc import Callable, Iterable, Iterator
+from concurrent.futures import Executor, ThreadPoolExecutor
+from contextlib import suppress
+from typing import TYPE_CHECKING, Any, cast
+
+from airflow.sdk import BaseAsyncOperator, BaseOperator, TaskInstanceState, 
timezone
+from airflow.sdk.bases.operator import ExecutorSafeguard
+from airflow.sdk.definitions._internal.logging_mixin import LoggingMixin
+from airflow.sdk.exceptions import TaskDeferred
+from airflow.sdk.execution_time.callback_runner import create_executable_runner
+from airflow.sdk.execution_time.context import context_get_outlet_events, 
set_current_context
+from airflow.sdk.execution_time.task_runner import (
+    RuntimeTaskInstance,
+    _execute_task,
+    _run_task_state_change_callbacks,
+)
+
+if TYPE_CHECKING:
+    from structlog.typing import FilteringBoundLogger as Logger
+
+    from airflow.sdk import Context
+    from airflow.sdk.execution_time.task_runner import IndexedTaskInstance
+
+
+_log = logging.getLogger(__name__)
+
+
+class AsyncAwareExecutor(Executor):
+    """
+    Executes both sync and async functions concurrently.
+
+    Sync functions run in a ThreadPoolExecutor.
+    Async coroutines run on an asyncio event loop with a semaphore limit.
+
+    :param loop: Event loop used to schedule async tasks and coordinate mixed 
execution.
+    :param max_workers: Maximum concurrent workers used by both thread pool 
and async semaphore.
+    :param shutdown_timeout: Maximum time to wait, in seconds, for in-flight 
async tasks and
+        thread-pool workers to finish during ``shutdown(wait=True)``. Python 
threads cannot be
+        forcibly stopped, so a worker stuck in blocking user code (slow HTTP 
call, blocked C
+        extension, a deadlocked DB driver, ...) would otherwise hang 
``shutdown()`` forever.
+    """
+
+    def __init__(
+        self, loop: AbstractEventLoop, max_workers: int | None = None, 
shutdown_timeout: float = 10.0
+    ):
+        if max_workers is None:
+            max_workers = os.cpu_count() or 1
+        if max_workers <= 0:
+            raise ValueError("max_workers must be greater than 0")
+
+        self._loop = loop
+        self._max_workers = max_workers
+        self._shutdown_timeout = shutdown_timeout
+        self._semaphore = Semaphore(max_workers)
+        self._thread_pool = ThreadPoolExecutor(max_workers=max_workers)
+        self._async_tasks: set[Task[Any]] = set()
+        self._shutdown = False
+
+    def __enter__(self):
+        return self
+
+    def __exit__(self, exc_type, exc_val, exc_tb):
+        if exc_type is not None:
+            # On error path, cancel futures but still wait briefly for them to
+            # process CancelledError and release resources (e.g., threading
+            # locks). Without waiting, cancelled tasks that hold _thread_lock
+            # never execute their finally blocks, permanently leaking the lock
+            # and causing subsequent comms.send() calls to deadlock.
+            self.shutdown(wait=True, cancel_futures=True)
+        else:
+            self.shutdown(wait=True)
+
+    def shutdown(self, wait: bool = True, *, cancel_futures: bool = False) -> 
None:
+        if self._shutdown:
+            return
+
+        self._shutdown = True
+
+        if cancel_futures:
+            for task in list(self._async_tasks):
+                task.cancel()
+
+        if wait and self._async_tasks:
+            with suppress(TimeoutError, AsyncTimeoutError):
+                self._loop.run_until_complete(
+                    wait_for(
+                        gather(*self._async_tasks, return_exceptions=True),
+                        timeout=self._shutdown_timeout,
+                    )
+                )
+
+        # ThreadPoolExecutor.shutdown(wait=True) blocks until every worker 
thread
+        # finishes its current work item, with no way to bound that wait or 
forcibly
+        # stop a thread stuck in blocking user code (slow HTTP call, blocked C
+        # extension, a deadlocked DB driver, ...). Ask the pool to stop 
accepting new
+        # work (and cancel anything not yet started) up front, then bound how 
long we
+        # personally wait on the worker threads instead of blocking 
indefinitely.
+        self._thread_pool.shutdown(wait=False, cancel_futures=cancel_futures)
+
+        if wait:
+            threads = list(getattr(self._thread_pool, "_threads", ()))
+            deadline = time.monotonic() + self._shutdown_timeout
+            for thread in threads:
+                remaining = deadline - time.monotonic()
+                if remaining <= 0:
+                    break
+                thread.join(timeout=remaining)
+
+            stuck = [thread.name for thread in threads if thread.is_alive()]
+            if stuck:
+                _log.error(
+                    "%d worker thread(s) still running %.1fs after shutdown 
was requested; "
+                    "giving up waiting to avoid blocking indefinitely. This 
may leak resources. "
+                    "Affected threads: %s",
+                    len(stuck),
+                    self._shutdown_timeout,
+                    stuck,
+                )
+
+    def submit(self, func: Callable[..., Any] | Any, *args, **kwargs) -> 
Future[Any]:  # type: ignore[override]
+        """
+        Submit a callable for execution.
+
+        Always returns an asyncio.Future for consistency, whether the callable
+        is sync (run in thread pool) or async (run on the event loop).
+        """
+        if self._shutdown:
+            raise RuntimeError("cannot schedule new futures after shutdown")
+
+        if inspect.iscoroutine(func):
+            coro = func
+        elif inspect.iscoroutinefunction(func):
+            coro = func(*args, **kwargs)
+        else:
+            # Wrap thread pool future as asyncio.Future for consistent return 
type
+            return wrap_future(self._thread_pool.submit(func, *args, 
**kwargs), loop=self._loop)
+
+        async def guarded():
+            try:
+                async with self._semaphore:
+                    return await coro
+            except CancelledError:
+                # If cancellation occurs while waiting for the semaphore,
+                # the inner coroutine was never awaited. Close it to prevent
+                # "coroutine was never awaited" RuntimeWarning.
+                coro.close()
+                raise
+
+        task = self._loop.create_task(guarded())
+        self._async_tasks.add(task)
+        task.add_done_callback(self._async_tasks.discard)
+        return task
+
+    async def run_sync(self, func: Callable[..., Any], *args, **kwargs) -> Any:
+        """Run a sync callable in this executor's thread pool and await its 
result."""
+        future = self._thread_pool.submit(func, *args, **kwargs)
+        return await wrap_future(future, loop=self._loop)
+
+    def map(

Review Comment:
   Fixed in 9a5119f37d, see the newer thread on the same method: it is now 
`imap_unordered` and no longer overrides `Executor.map`.
   
   ---
   Drafted-by: Claude Fable 5.1; reviewed by @dabla before posting



-- 
This is an automated message from the Apache Git Service.
To respond to the message, please log on to GitHub and use the
URL above to go to the specific comment.

To unsubscribe, e-mail: [email protected]

For queries about this service, please contact Infrastructure at:
[email protected]

Reply via email to