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()