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

Lee-W pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/airflow.git


The following commit(s) were added to refs/heads/main by this push:
     new 4b6606543fd Distinguish OpenAI batch timeout and cancellation in 
deferrable tasks (#72149)
4b6606543fd is described below

commit 4b6606543fdfd2fd1a9dae3d0be1cdb784eb0e2c
Author: Wei Lee <[email protected]>
AuthorDate: Tue Sep 29 17:09:33 2026 +0900

    Distinguish OpenAI batch timeout and cancellation in deferrable tasks 
(#72149)
---
 providers/openai/docs/changelog.rst                |  19 ++++
 .../src/airflow/providers/openai/exceptions.py     |  11 ++
 .../src/airflow/providers/openai/hooks/openai.py   |  50 ++++++++-
 .../airflow/providers/openai/operators/openai.py   |  63 ++++++++++--
 .../airflow/providers/openai/triggers/openai.py    |  19 +++-
 .../openai/tests/unit/openai/hooks/test_openai.py  |  36 +++++++
 .../tests/unit/openai/operators/test_openai.py     | 114 ++++++++++++++++++++-
 .../openai/tests/unit/openai/test_exceptions.py    |  15 ++-
 .../tests/unit/openai/triggers/test_openai.py      |  47 +++++++--
 9 files changed, 350 insertions(+), 24 deletions(-)

diff --git a/providers/openai/docs/changelog.rst 
b/providers/openai/docs/changelog.rst
index 3dda0fb8938..cf4f0d37d2a 100644
--- a/providers/openai/docs/changelog.rst
+++ b/providers/openai/docs/changelog.rst
@@ -20,6 +20,25 @@
 Changelog
 ---------
 
+.. warning::
+    A deferred ``OpenAITriggerBatchOperator`` that times out now raises 
``OpenAIBatchTimeout``
+    instead of ``OpenAIBatchJobException``, which is what 1.8.2 and earlier 
raised for the same
+    condition. ``OpenAIBatchTimeout`` is not a subclass of 
``OpenAIBatchJobException``, so an
+    ``on_failure_callback``, ``except`` clause, or retry rule keyed on
+    ``OpenAIBatchJobException`` no longer matches a deferred timeout. Catch or 
check for
+    ``OpenAIBatchTimeout`` as well to keep handling timeouts.
+
+    A cancelled batch now raises ``OpenAIBatchCancelled``, a subclass of
+    ``OpenAIBatchJobException``, so existing code that catches 
``OpenAIBatchJobException`` keeps
+    matching cancellations unchanged.
+
+.. note::
+    A deferred ``OpenAITriggerBatchOperator`` that times out now requests 
cancellation of the
+    batch, matching the non-deferrable path. Previously a deferred timeout 
only failed the
+    task and left the batch running (and billing) on OpenAI's side. 
Cancellation on OpenAI's
+    side is asynchronous, so the batch reports ``cancelling`` for a while 
before it settles as
+    ``cancelled``.
+
 2.0.0
 .....
 
diff --git a/providers/openai/src/airflow/providers/openai/exceptions.py 
b/providers/openai/src/airflow/providers/openai/exceptions.py
index 09618b9048e..cca4ef1977a 100644
--- a/providers/openai/src/airflow/providers/openai/exceptions.py
+++ b/providers/openai/src/airflow/providers/openai/exceptions.py
@@ -24,6 +24,17 @@ class OpenAIBatchJobException(AirflowException):
     """Raise when OpenAI Batch Job fails to start AFTER processing the 
request."""
 
 
+class OpenAIBatchCancelled(OpenAIBatchJobException):
+    """
+    Raise when an OpenAI Batch Job was cancelled.
+
+    Cancellation is a decision, not a failure, so it gets its own subclass: 
callers
+    that want to distinguish "someone cancelled this batch" from "the batch 
failed"
+    can catch this specifically, while existing handlers written against
+    ``OpenAIBatchJobException`` keep working unchanged.
+    """
+
+
 class OpenAIBatchTimeout(AirflowException):
     """Raise when OpenAI Batch Job times out."""
 
diff --git a/providers/openai/src/airflow/providers/openai/hooks/openai.py 
b/providers/openai/src/airflow/providers/openai/hooks/openai.py
index f9ba07a8497..761f29c4fd2 100644
--- a/providers/openai/src/airflow/providers/openai/hooks/openai.py
+++ b/providers/openai/src/airflow/providers/openai/hooks/openai.py
@@ -54,9 +54,10 @@ if TYPE_CHECKING:
     from openai.types.vector_stores import VectorStoreFile, 
VectorStoreFileBatch, VectorStoreFileDeleted
 from airflow.exceptions import AirflowProviderDeprecationWarning
 from airflow.providers.common.compat.module_loading import import_string
-from airflow.providers.common.compat.sdk import BaseHook
+from airflow.providers.common.compat.sdk import AirflowException, BaseHook
 from airflow.providers.openai.exceptions import (
     OpenAIAgentSessionError,
+    OpenAIBatchCancelled,
     OpenAIBatchJobException,
     OpenAIBatchTimeout,
     OpenAITriggerEventError,
@@ -97,6 +98,42 @@ class BatchStatus(str, Enum):
 TRIGGER_EVENT_STATUSES = frozenset({"success", "error", "cancelled"})
 
 
+class TerminationReason(str, Enum):
+    """Enum for the ``termination_reason`` field of a trigger's terminal 
event."""
+
+    TIMEOUT = "timeout"
+    COMPLETED = "completed"
+    CANCELLED = "cancelled"
+    FAILED = "failed"
+    EXPIRED = "expired"
+    UNEXPECTED_STATUS = "unexpected_status"
+    POLLING_ERROR = "polling_error"
+
+
+# Maps the trigger's ``termination_reason`` field to the exception 
``execute_complete``
+# should raise. Keyed on the reason field, never on the message text, so that a
+# rewording of the trigger's message never silently changes which exception a
+# downstream task can catch.
+_TERMINATION_REASON_EXCEPTIONS: dict[str, type[AirflowException]] = {
+    TerminationReason.TIMEOUT: OpenAIBatchTimeout,
+    TerminationReason.CANCELLED: OpenAIBatchCancelled,
+}
+
+
+def build_batch_error(message: str, termination_reason: str | None) -> 
AirflowException:
+    """
+    Build (but do not raise) the exception matching a trigger event's 
termination reason.
+
+    ``termination_reason`` is ``None`` when the event was produced by a trigger
+    serialized before this field existed (a rolling upgrade in flight); that 
case
+    falls back to ``OpenAIBatchJobException``, matching today's behavior.
+    """
+    if termination_reason is None:
+        return OpenAIBatchJobException(message)
+    exception_class = _TERMINATION_REASON_EXCEPTIONS.get(termination_reason, 
OpenAIBatchJobException)
+    return exception_class(message)
+
+
 def validate_execute_complete_event(event: dict[str, Any] | None = None) -> 
dict[str, Any]:
     """
     Validate the event a deferred task resumes with, returning it if 
well-formed.
@@ -685,7 +722,12 @@ class OpenAIHook(BaseHook):
         start = time.monotonic()
         while True:
             if start + timeout < time.monotonic():
-                self.cancel_batch(batch_id=batch_id)
+                try:
+                    self.cancel_batch(batch_id=batch_id)
+                except Exception as e:
+                    self.log.warning(
+                        "Failed to request cancellation of batch %s after 
timeout: %s", batch_id, e
+                    )
                 raise OpenAIBatchTimeout(f"Timeout: OpenAI Batch {batch_id} is 
not ready after {timeout}s")
             batch = self.get_batch(batch_id=batch_id)
 
@@ -697,10 +739,10 @@ class OpenAIHook(BaseHook):
             if batch.status == BatchStatus.FAILED:
                 raise OpenAIBatchJobException(f"Batch failed - \n{batch_id}")
             if batch.status in (BatchStatus.CANCELLED, BatchStatus.CANCELLING):
-                raise OpenAIBatchJobException(f"Batch failed - batch was 
cancelled:\n{batch_id}")
+                raise OpenAIBatchCancelled(f"Batch failed - batch was 
cancelled:\n{batch_id}")
             if batch.status == BatchStatus.EXPIRED:
                 raise OpenAIBatchJobException(
-                    f"Batch failed - batch couldn't be completed within the 
hour time window :\n{batch_id}"
+                    f"Batch failed - batch couldn't be completed within its 
completion window:\n{batch_id}"
                 )
 
             raise OpenAIBatchJobException(
diff --git a/providers/openai/src/airflow/providers/openai/operators/openai.py 
b/providers/openai/src/airflow/providers/openai/operators/openai.py
index d0234b17b0f..fe03fe1bf40 100644
--- a/providers/openai/src/airflow/providers/openai/operators/openai.py
+++ b/providers/openai/src/airflow/providers/openai/operators/openai.py
@@ -22,8 +22,12 @@ from functools import cached_property
 from typing import TYPE_CHECKING, Any, ClassVar
 
 from airflow.providers.common.compat.sdk import BaseOperator, conf
-from airflow.providers.openai.exceptions import OpenAIBatchJobException
-from airflow.providers.openai.hooks.openai import OpenAIHook, 
validate_execute_complete_event
+from airflow.providers.openai.hooks.openai import (
+    OpenAIHook,
+    TerminationReason,
+    build_batch_error,
+    validate_execute_complete_event,
+)
 from airflow.providers.openai.triggers.openai import OpenAIBatchTrigger
 
 if TYPE_CHECKING:
@@ -360,9 +364,10 @@ class OpenAITriggerBatchOperator(BaseOperator):
     :param deferrable: Optional. Run operator in the deferrable mode.
     :param wait_seconds: Optional. Number of seconds between checks. Only used 
when ``deferrable`` is False.
         Defaults to 3 seconds.
-    :param timeout: Optional. The amount of time, in seconds, to wait for the 
request to complete.
-        Applies in both deferrable and non-deferrable mode. Defaults to 24 
hours, which is the SLA for
-        OpenAI Batch API.
+    :param timeout: Optional. The number of seconds to wait for the batch to 
complete, in both
+        deferrable and non-deferrable mode. Defaults to 24 hours, the SLA for 
OpenAI Batch API.
+        In deferrable mode, if ``execution_timeout`` is set shorter than 
``timeout``, the task is
+        failed with ``TaskDeferralTimeout`` before the trigger times out, and 
the batch is not cancelled.
     :param wait_for_completion: Optional. Whether to wait for the batch to 
complete. If set to False, the operator
         will return immediately after triggering the batch. Defaults to True.
     :param metadata: Optional. A set of key-value pairs that can be attached 
to the batch. (templated)
@@ -455,17 +460,59 @@ class OpenAITriggerBatchOperator(BaseOperator):
         Invoke this callback when the trigger fires; return immediately.
 
         Relies on trigger to throw an exception, otherwise it assumes 
execution was
-        successful.
+        successful. The exception raised depends on the event's 
``termination_reason``:
+        :class:`~airflow.providers.openai.exceptions.OpenAIBatchTimeout` for a 
timeout
+        (matching the exception the synchronous path raises for the same 
condition),
+        :class:`~airflow.providers.openai.exceptions.OpenAIBatchCancelled` for 
a cancellation
+        (a subclass of 
:class:`~airflow.providers.openai.exceptions.OpenAIBatchJobException`),
+        and 
:class:`~airflow.providers.openai.exceptions.OpenAIBatchJobException` for any
+        other failure (including events from a trigger serialized before
+        ``termination_reason`` existed).
+
+        On a timeout, cancellation of the batch is requested before the 
timeout is raised
+        (see :meth:`_cancel_batch_quietly`). No other termination reason 
triggers
+        cancellation: a ``polling_error`` may be a transient, Airflow-side 
failure rather than
+        a real batch problem, and cancellation is irreversible, so it is left 
alone to run to
+        its own 24-hour completion window instead.
         """
         event = validate_execute_complete_event(event)
         if event["status"] != "success":
-            raise OpenAIBatchJobException(event["message"])
+            if event.get("termination_reason") == TerminationReason.TIMEOUT:
+                batch_id = event["batch_id"]
+                self.log.warning(
+                    "%s timed out waiting for batch %s; requesting 
cancellation.",
+                    self.task_id,
+                    batch_id,
+                )
+                self._cancel_batch_quietly(batch_id)
+            raise build_batch_error(event["message"], 
event.get("termination_reason"))
 
         self.log.info("%s completed successfully.", self.task_id)
         return event["batch_id"]
 
+    def _cancel_batch_quietly(self, batch_id: str) -> None:
+        """
+        Best-effort request to cancel a batch; never raises.
+
+        Takes ``batch_id`` as a parameter rather than reading 
``self.batch_id`` because it has
+        two callers with different sources for it: ``execute_complete``, after 
a deferred
+        timeout, passes the batch id carried by the trigger event, since it 
runs on a resumed
+        task instance where ``execute``'s assignment to ``self.batch_id`` 
never happened;
+        ``on_kill`` passes ``self.batch_id`` directly, already set by 
``execute`` on this same
+        operator instance.
+
+        Cancellation on OpenAI's side is asynchronous: the batch reports 
``cancelling`` for up
+        to 10 minutes before it settles as ``cancelled``, so this only 
requests cancellation. A
+        failure to cancel is logged, not raised, so it never masks the real 
failure reason
+        (the timeout, or the kill).
+        """
+        try:
+            self.hook.cancel_batch(batch_id)
+        except Exception as e:
+            self.log.warning("Failed to request cancellation of batch %s: %s", 
batch_id, e)
+
     def on_kill(self) -> None:
         """Cancel the batch if task is cancelled."""
         if self.batch_id:
             self.log.info("on_kill: cancel the OpenAI Batch %s", self.batch_id)
-            self.hook.cancel_batch(self.batch_id)
+            self._cancel_batch_quietly(self.batch_id)
diff --git a/providers/openai/src/airflow/providers/openai/triggers/openai.py 
b/providers/openai/src/airflow/providers/openai/triggers/openai.py
index 49fd0900cc0..2b41b3f6ea2 100644
--- a/providers/openai/src/airflow/providers/openai/triggers/openai.py
+++ b/providers/openai/src/airflow/providers/openai/triggers/openai.py
@@ -21,7 +21,7 @@ import time
 from collections.abc import AsyncIterator
 from typing import Any
 
-from airflow.providers.openai.hooks.openai import BatchStatus, OpenAIHook
+from airflow.providers.openai.hooks.openai import BatchStatus, OpenAIHook, 
TerminationReason
 from airflow.triggers.base import BaseTrigger, TriggerEvent
 
 
@@ -103,6 +103,7 @@ class OpenAIBatchTrigger(BaseTrigger):
                     yield TriggerEvent(
                         {
                             "status": "error",
+                            "termination_reason": TerminationReason.TIMEOUT,
                             "message": (
                                 f"Batch {self.batch_id} has not reached a 
terminal status after "
                                 f"{elapsed:.0f} seconds."
@@ -116,6 +117,7 @@ class OpenAIBatchTrigger(BaseTrigger):
                 yield TriggerEvent(
                     {
                         "status": "success",
+                        "termination_reason": TerminationReason.COMPLETED,
                         "message": f"Batch {self.batch_id} has completed 
successfully.",
                         "batch_id": self.batch_id,
                     }
@@ -124,6 +126,7 @@ class OpenAIBatchTrigger(BaseTrigger):
                 yield TriggerEvent(
                     {
                         "status": "cancelled",
+                        "termination_reason": TerminationReason.CANCELLED,
                         "message": f"Batch {self.batch_id} has been 
cancelled.",
                         "batch_id": self.batch_id,
                     }
@@ -132,6 +135,7 @@ class OpenAIBatchTrigger(BaseTrigger):
                 yield TriggerEvent(
                     {
                         "status": "error",
+                        "termination_reason": TerminationReason.FAILED,
                         "message": f"Batch failed:\n{self.batch_id}",
                         "batch_id": self.batch_id,
                     }
@@ -140,7 +144,8 @@ class OpenAIBatchTrigger(BaseTrigger):
                 yield TriggerEvent(
                     {
                         "status": "error",
-                        "message": f"Batch couldn't be completed within the 
hour time window :\n{self.batch_id}",
+                        "termination_reason": TerminationReason.EXPIRED,
+                        "message": f"Batch couldn't be completed within its 
completion window:\n{self.batch_id}",
                         "batch_id": self.batch_id,
                     }
                 )
@@ -148,9 +153,17 @@ class OpenAIBatchTrigger(BaseTrigger):
                 yield TriggerEvent(
                     {
                         "status": "error",
+                        "termination_reason": 
TerminationReason.UNEXPECTED_STATUS,
                         "message": f"Batch {self.batch_id} has failed.",
                         "batch_id": self.batch_id,
                     }
                 )
         except Exception as e:
-            yield TriggerEvent({"status": "error", "message": str(e), 
"batch_id": self.batch_id})
+            yield TriggerEvent(
+                {
+                    "status": "error",
+                    "termination_reason": TerminationReason.POLLING_ERROR,
+                    "message": str(e),
+                    "batch_id": self.batch_id,
+                }
+            )
diff --git a/providers/openai/tests/unit/openai/hooks/test_openai.py 
b/providers/openai/tests/unit/openai/hooks/test_openai.py
index cc370911777..6b56df91f55 100644
--- a/providers/openai/tests/unit/openai/hooks/test_openai.py
+++ b/providers/openai/tests/unit/openai/hooks/test_openai.py
@@ -41,6 +41,7 @@ from airflow.exceptions import 
AirflowProviderDeprecationWarning
 from airflow.models import Connection
 from airflow.providers.openai.exceptions import (
     OpenAIAgentSessionError,
+    OpenAIBatchCancelled,
     OpenAIBatchJobException,
     OpenAIBatchTimeout,
     OpenAITriggerEventError,
@@ -692,6 +693,41 @@ def 
test_wait_for_in_progress_batch_timeout(mock_openai_hook, mock_wip_batch):
     assert mock_openai_hook.conn.batches.cancel.call_count == 1
 
 
+def 
test_wait_for_in_progress_batch_timeout_cancel_failure_does_not_mask_timeout(
+    mock_openai_hook, mock_wip_batch, caplog
+):
+    """A cancellation failure inside the timeout branch must not replace 
``OpenAIBatchTimeout``
+    with the cancellation's own exception, and the failure must still be 
logged.
+    """
+    mock_openai_hook.conn.batches.retrieve.return_value = mock_wip_batch
+    mock_openai_hook.conn.batches.cancel.side_effect = RuntimeError("cancel 
failed")
+
+    with caplog.at_level("WARNING"):
+        with pytest.raises(OpenAIBatchTimeout, match="Timeout"):
+            mock_openai_hook.wait_for_batch(batch_id=BATCH_ID, 
wait_seconds=0.01, timeout=0.01)
+
+    assert mock_openai_hook.conn.batches.cancel.call_count == 1
+    assert any("Failed to request cancellation of batch" in message for 
message in caplog.messages)
+
+
[email protected]("status", ["cancelled", "cancelling"])
+def 
test_wait_for_cancelled_batch_raises_exact_cancelled_type(mock_openai_hook, 
status):
+    """``OpenAIBatchCancelled`` is a subclass of ``OpenAIBatchJobException``, 
so asserting
+    only the base class would stay green even if this raised the wrong (base) 
type. Assert
+    the exact type to prove the exception was actually narrowed.
+    """
+    mock_openai_hook.conn.batches.retrieve.return_value = create_batch(status)
+    with pytest.raises(OpenAIBatchCancelled):
+        mock_openai_hook.wait_for_batch(batch_id=BATCH_ID)
+
+
+def 
test_wait_for_expired_batch_message_does_not_mention_hour_window(mock_openai_hook):
+    mock_openai_hook.conn.batches.retrieve.return_value = 
create_batch("expired")
+    with pytest.raises(OpenAIBatchJobException, match="completion window") as 
exc_info:
+        mock_openai_hook.wait_for_batch(batch_id=BATCH_ID)
+    assert "hour time window" not in str(exc_info.value)
+
+
 def test_openai_hook_test_connection(mock_openai_hook):
     result, message = mock_openai_hook.test_connection()
     assert result is True
diff --git a/providers/openai/tests/unit/openai/operators/test_openai.py 
b/providers/openai/tests/unit/openai/operators/test_openai.py
index 7ec22be98cc..7e6cdc5c684 100644
--- a/providers/openai/tests/unit/openai/operators/test_openai.py
+++ b/providers/openai/tests/unit/openai/operators/test_openai.py
@@ -30,7 +30,12 @@ from openai.types.responses.response import IncompleteDetails
 from openai.types.responses.response_usage import InputTokensDetails, 
OutputTokensDetails
 
 from airflow.providers.common.compat.sdk import DAG, BaseOperator, Context, 
TaskDeferred, XComArg
-from airflow.providers.openai.exceptions import OpenAIBatchJobException, 
OpenAITriggerEventError
+from airflow.providers.openai.exceptions import (
+    OpenAIBatchCancelled,
+    OpenAIBatchJobException,
+    OpenAIBatchTimeout,
+    OpenAITriggerEventError,
+)
 from airflow.providers.openai.hooks.openai import OpenAIHook
 from airflow.providers.openai.operators.openai import (
     OpenAIEmbeddingOperator,
@@ -898,6 +903,29 @@ def 
test_openai_trigger_batch_operator_deferred_logs_active_knob(mock_log, mock_
     )
 
 
+def test_openai_trigger_batch_operator_on_kill_cancels_batch_quietly(caplog):
+    """on_kill()'s cancellation failure is logged, not raised."""
+    operator = OpenAITriggerBatchOperator(
+        task_id=TASK_ID,
+        conn_id=CONN_ID,
+        file_id=FILE_ID,
+        endpoint=BATCH_ENDPOINT,
+    )
+    operator.batch_id = BATCH_ID
+    mock_hook_instance = Mock(spec=OpenAIHook)
+    mock_hook_instance.cancel_batch.side_effect = RuntimeError("cancel failed")
+    operator.hook = mock_hook_instance
+
+    with caplog.at_level("WARNING"):
+        try:
+            operator.on_kill()
+        except Exception as e:
+            pytest.fail(f"on_kill() should not raise: {e}")
+
+    mock_hook_instance.cancel_batch.assert_called_once_with(BATCH_ID)
+    assert any("Failed to request cancellation of batch" in message for 
message in caplog.messages)
+
+
 class TestOpenAITriggerBatchOperatorExecuteComplete:
     def _operator(self):
         return OpenAITriggerBatchOperator(
@@ -935,3 +963,87 @@ class TestOpenAITriggerBatchOperatorExecuteComplete:
     def test_invalid_event_raises_instead_of_succeeding(self, event):
         with pytest.raises(OpenAITriggerEventError):
             self._operator().execute_complete(Context(), event)
+
+    @pytest.mark.parametrize(
+        ("termination_reason", "expected_exc", "cancel_expected"),
+        [
+            pytest.param("timeout", OpenAIBatchTimeout, True, id="timeout"),
+            pytest.param("cancelled", OpenAIBatchCancelled, False, 
id="cancelled"),
+            pytest.param("failed", OpenAIBatchJobException, False, 
id="failed"),
+            pytest.param("expired", OpenAIBatchJobException, False, 
id="expired"),
+            pytest.param("unexpected_status", OpenAIBatchJobException, False, 
id="unexpected-status"),
+            pytest.param("polling_error", OpenAIBatchJobException, False, 
id="polling-error"),
+            pytest.param(None, OpenAIBatchJobException, False, 
id="missing-reason"),
+        ],
+    )
+    def test_execute_complete_raises_exception_matching_termination_reason(
+        self, termination_reason, expected_exc, cancel_expected
+    ):
+        """Covers both which exception a termination reason maps to, and 
whether it triggers
+        cancellation, off a mocked hook so the assertions never depend on 
``_cancel_batch_quietly``
+        falling through to a real ``OpenAIHook`` looking up ``test_conn_id``.
+        """
+        operator = self._operator()
+        mock_hook_instance = Mock(spec=OpenAIHook)
+        operator.hook = mock_hook_instance
+        event = {"status": "error", "message": "boom", "batch_id": BATCH_ID}
+        if termination_reason is not None:
+            event["termination_reason"] = termination_reason
+
+        with pytest.raises(expected_exc, match="boom"):
+            operator.execute_complete(Context(), event)
+
+        if cancel_expected:
+            mock_hook_instance.cancel_batch.assert_called_once_with(BATCH_ID)
+        else:
+            mock_hook_instance.cancel_batch.assert_not_called()
+
+    @pytest.mark.parametrize("status", ["error", "cancelled"])
+    def test_execute_complete_missing_termination_reason_falls_back(self, 
status):
+        """A trigger serialized before ``termination_reason`` existed sends an 
event without
+        that key; ``execute_complete`` must fall back to 
``OpenAIBatchJobException`` exactly,
+        not raise ``KeyError``.
+        """
+        event = {"status": status, "message": "boom", "batch_id": BATCH_ID}
+        with pytest.raises(OpenAIBatchJobException, match="boom") as exc_info:
+            self._operator().execute_complete(Context(), event)
+        assert type(exc_info.value) is OpenAIBatchJobException
+
+    def test_timeout_requests_cancellation_using_event_batch_id(self):
+        """The resumed task is a fresh operator instance, so ``self.batch_id`` 
is ``None`` here.
+        Cancellation must use ``event["batch_id"]``; if this test is made to 
pass by
+        reading ``self.batch_id`` instead, it should fail again as soon as 
that read returns
+        ``None`` for a real resumed task.
+        """
+        operator = self._operator()
+        assert operator.batch_id is None
+        mock_hook_instance = Mock(spec=OpenAIHook)
+        operator.hook = mock_hook_instance
+        event = {
+            "status": "error",
+            "termination_reason": "timeout",
+            "message": "boom",
+            "batch_id": BATCH_ID,
+        }
+
+        with pytest.raises(OpenAIBatchTimeout):
+            operator.execute_complete(Context(), event)
+
+        mock_hook_instance.cancel_batch.assert_called_once_with(BATCH_ID)
+
+    def test_cancel_failure_does_not_mask_timeout(self):
+        operator = self._operator()
+        mock_hook_instance = Mock(spec=OpenAIHook)
+        mock_hook_instance.cancel_batch.side_effect = RuntimeError("cancel 
failed")
+        operator.hook = mock_hook_instance
+        event = {
+            "status": "error",
+            "termination_reason": "timeout",
+            "message": "boom",
+            "batch_id": BATCH_ID,
+        }
+
+        with pytest.raises(OpenAIBatchTimeout):
+            operator.execute_complete(Context(), event)
+
+        mock_hook_instance.cancel_batch.assert_called_once_with(BATCH_ID)
diff --git a/providers/openai/tests/unit/openai/test_exceptions.py 
b/providers/openai/tests/unit/openai/test_exceptions.py
index fabaad35343..b9c6a14e446 100644
--- a/providers/openai/tests/unit/openai/test_exceptions.py
+++ b/providers/openai/tests/unit/openai/test_exceptions.py
@@ -21,7 +21,11 @@ from unittest.mock import Mock
 
 import pytest
 
-from airflow.providers.openai.exceptions import OpenAIBatchJobException, 
OpenAIBatchTimeout
+from airflow.providers.openai.exceptions import (
+    OpenAIBatchCancelled,
+    OpenAIBatchJobException,
+    OpenAIBatchTimeout,
+)
 from airflow.providers.openai.hooks.openai import OpenAIHook
 
 
@@ -30,6 +34,7 @@ from airflow.providers.openai.hooks.openai import OpenAIHook
     [
         OpenAIBatchTimeout,
         OpenAIBatchJobException,
+        OpenAIBatchCancelled,
     ],
 )
 def test_wait_for_batch_raise_exception(exception_class):
@@ -38,3 +43,11 @@ def test_wait_for_batch_raise_exception(exception_class):
     hook = mock_hook_instance
     with pytest.raises(exception_class):
         hook.wait_for_batch(batch_id="batch_id")
+
+
+def test_batch_cancelled_is_subclass_of_batch_job_exception():
+    """Cancellation is deliberately a subclass, not a sibling, of the generic 
batch failure
+    exception: existing ``except OpenAIBatchJobException`` handlers must keep 
working
+    unchanged after cancellation gets its own exception type.
+    """
+    assert issubclass(OpenAIBatchCancelled, OpenAIBatchJobException)
diff --git a/providers/openai/tests/unit/openai/triggers/test_openai.py 
b/providers/openai/tests/unit/openai/triggers/test_openai.py
index 0a2c4322a5e..d72620ed3c5 100644
--- a/providers/openai/tests/unit/openai/triggers/test_openai.py
+++ b/providers/openai/tests/unit/openai/triggers/test_openai.py
@@ -119,22 +119,28 @@ class TestOpenAIBatchTrigger:
 
     @pytest.mark.asyncio
     @pytest.mark.parametrize(
-        ("mock_batch_status", "mock_status", "mock_message"),
+        ("mock_batch_status", "mock_status", "mock_termination_reason", 
"mock_message"),
         [
-            (str(BatchStatus.COMPLETED), "success", "Batch batch_id has 
completed successfully."),
-            (str(BatchStatus.CANCELLING), "cancelled", "Batch batch_id has 
been cancelled."),
-            (str(BatchStatus.CANCELLED), "cancelled", "Batch batch_id has been 
cancelled."),
-            (str(BatchStatus.FAILED), "error", "Batch failed:\nbatch_id"),
+            (
+                str(BatchStatus.COMPLETED),
+                "success",
+                "completed",
+                "Batch batch_id has completed successfully.",
+            ),
+            (str(BatchStatus.CANCELLING), "cancelled", "cancelled", "Batch 
batch_id has been cancelled."),
+            (str(BatchStatus.CANCELLED), "cancelled", "cancelled", "Batch 
batch_id has been cancelled."),
+            (str(BatchStatus.FAILED), "error", "failed", "Batch 
failed:\nbatch_id"),
             (
                 str(BatchStatus.EXPIRED),
                 "error",
-                "Batch couldn't be completed within the hour time window 
:\nbatch_id",
+                "expired",
+                "Batch couldn't be completed within its completion 
window:\nbatch_id",
             ),
         ],
     )
     @mock.patch("airflow.providers.openai.hooks.openai.OpenAIHook.get_batch")
     async def test_openai_batch_for_terminal_status(
-        self, mock_batch, mock_batch_status, mock_status, mock_message
+        self, mock_batch, mock_batch_status, mock_status, 
mock_termination_reason, mock_message
     ):
         """Assert that run trigger messages in case of job finished"""
         mock_batch.return_value = self.mock_get_batch(mock_batch_status)
@@ -146,6 +152,7 @@ class TestOpenAIBatchTrigger:
         )
         expected_result = {
             "status": mock_status,
+            "termination_reason": mock_termination_reason,
             "message": mock_message,
             "batch_id": self.BATCH_ID,
         }
@@ -187,6 +194,7 @@ class TestOpenAIBatchTrigger:
         await asyncio.sleep(0.1)
         event = task.result()
         assert event.payload["status"] == "error"
+        assert event.payload["termination_reason"] == "timeout"
         assert f"Batch {self.BATCH_ID} has not reached a terminal status 
after" in event.payload["message"]
         asyncio.get_event_loop().stop()
 
@@ -235,6 +243,7 @@ class TestOpenAIBatchTrigger:
             TriggerEvent(
                 {
                     "status": "success",
+                    "termination_reason": "completed",
                     "message": f"Batch {self.BATCH_ID} has completed 
successfully.",
                     "batch_id": self.BATCH_ID,
                 }
@@ -254,6 +263,7 @@ class TestOpenAIBatchTrigger:
         )
         expected_result = {
             "status": "error",
+            "termination_reason": "polling_error",
             "message": "'float' object has no attribute 'status'",
             "batch_id": self.BATCH_ID,
         }
@@ -261,3 +271,26 @@ class TestOpenAIBatchTrigger:
         await asyncio.sleep(0.1)
         assert TriggerEvent(expected_result) == task.result()
         asyncio.get_event_loop().stop()
+
+    @pytest.mark.asyncio
+    @mock.patch("airflow.providers.openai.hooks.openai.OpenAIHook.get_batch")
+    async def test_openai_batch_for_unexpected_status(self, mock_batch):
+        """A batch status outside the known terminal set falls into the 
`unexpected_status` branch."""
+        mock_batch.return_value = self.mock_get_batch("validating")
+        mock_batch.return_value.status = "some_future_status"
+        trigger = OpenAIBatchTrigger(
+            conn_id=self.CONN_ID,
+            batch_id=self.BATCH_ID,
+            poll_interval=self.POLL_INTERVAL,
+            timeout=self.TIMEOUT,
+        )
+        expected_result = {
+            "status": "error",
+            "termination_reason": "unexpected_status",
+            "message": f"Batch {self.BATCH_ID} has failed.",
+            "batch_id": self.BATCH_ID,
+        }
+        task = asyncio.create_task(trigger.run().__anext__())
+        await asyncio.sleep(0.1)
+        assert TriggerEvent(expected_result) == task.result()
+        asyncio.get_event_loop().stop()

Reply via email to