Skip to content
19 changes: 19 additions & 0 deletions providers/openai/docs/changelog.rst
Original file line number Diff line number Diff line change
Expand Up @@ -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
.....

Expand Down
11 changes: 11 additions & 0 deletions providers/openai/src/airflow/providers/openai/exceptions.py
Original file line number Diff line number Diff line change
Expand Up @@ -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."""

Expand Down
50 changes: 46 additions & 4 deletions providers/openai/src/airflow/providers/openai/hooks/openai.py
Original file line number Diff line number Diff line change
Expand Up @@ -54,9 +54,10 @@
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,
Expand Down Expand Up @@ -97,6 +98,42 @@ def is_in_progress(cls, status: str) -> bool:
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.
Expand Down Expand Up @@ -685,7 +722,12 @@ def wait_for_batch(self, batch_id: str, wait_seconds: float = 3, timeout: float
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)

Expand All @@ -697,10 +739,10 @@ def wait_for_batch(self, batch_id: str, wait_seconds: float = 3, timeout: float
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(
Expand Down
63 changes: 55 additions & 8 deletions providers/openai/src/airflow/providers/openai/operators/openai.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,8 +22,12 @@
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:
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -455,17 +460,59 @@ def execute_complete(self, context: Context, event: Any = None) -> str:
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:
Comment thread
Lee-W marked this conversation as resolved.
"""
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)
19 changes: 16 additions & 3 deletions providers/openai/src/airflow/providers/openai/triggers/openai.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,7 +21,7 @@
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


Expand Down Expand Up @@ -103,6 +103,7 @@ async def run(self) -> AsyncIterator[TriggerEvent]:
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."
Expand All @@ -116,6 +117,7 @@ async def run(self) -> AsyncIterator[TriggerEvent]:
yield TriggerEvent(
{
"status": "success",
"termination_reason": TerminationReason.COMPLETED,
"message": f"Batch {self.batch_id} has completed successfully.",
"batch_id": self.batch_id,
}
Expand All @@ -124,6 +126,7 @@ async def run(self) -> AsyncIterator[TriggerEvent]:
yield TriggerEvent(
{
"status": "cancelled",
"termination_reason": TerminationReason.CANCELLED,
"message": f"Batch {self.batch_id} has been cancelled.",
"batch_id": self.batch_id,
}
Expand All @@ -132,6 +135,7 @@ async def run(self) -> AsyncIterator[TriggerEvent]:
yield TriggerEvent(
{
"status": "error",
"termination_reason": TerminationReason.FAILED,
"message": f"Batch failed:\n{self.batch_id}",
"batch_id": self.batch_id,
}
Expand All @@ -140,17 +144,26 @@ async def run(self) -> AsyncIterator[TriggerEvent]:
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,
}
)
else:
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,
}
)
36 changes: 36 additions & 0 deletions providers/openai/tests/unit/openai/hooks/test_openai.py
Original file line number Diff line number Diff line change
Expand Up @@ -41,6 +41,7 @@
from airflow.models import Connection
from airflow.providers.openai.exceptions import (
OpenAIAgentSessionError,
OpenAIBatchCancelled,
OpenAIBatchJobException,
OpenAIBatchTimeout,
OpenAITriggerEventError,
Expand Down Expand Up @@ -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)


@pytest.mark.parametrize("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
Expand Down
Loading