This is an automated email from the ASF dual-hosted git repository. Lee-W pushed a commit to branch openai-batch-termination-reason in repository https://gitbox.apache.org/repos/asf/airflow.git
commit 7a08e5a3feb9e8fb5975223c004a28ce9fe1f516 Author: Wei Lee <[email protected]> AuthorDate: Tue Aug 18 15:03:32 2026 +0800 Distinguish OpenAI batch timeout and cancellation in deferrable tasks A deferred OpenAITriggerBatchOperator reported every non-success outcome as the same OpenAIBatchJobException, so a task that merely timed out looked identical to a batch that failed, and neither could be handled separately. The synchronous path has always raised OpenAIBatchTimeout for the same condition. The trigger now reports why it stopped, and the resuming task picks the matching exception from that field rather than from the message text. --- providers/openai/docs/changelog.rst | 12 ++++++ .../src/airflow/providers/openai/exceptions.py | 11 +++++ .../src/airflow/providers/openai/hooks/openai.py | 30 ++++++++++++-- .../airflow/providers/openai/operators/openai.py | 21 ++++++++-- .../airflow/providers/openai/triggers/openai.py | 17 +++++++- .../openai/tests/unit/openai/hooks/test_openai.py | 19 +++++++++ .../tests/unit/openai/operators/test_openai.py | 41 ++++++++++++++++++- .../openai/tests/unit/openai/test_exceptions.py | 15 ++++++- .../tests/unit/openai/triggers/test_openai.py | 47 ++++++++++++++++++---- 9 files changed, 195 insertions(+), 18 deletions(-) diff --git a/providers/openai/docs/changelog.rst b/providers/openai/docs/changelog.rst index 3dda0fb8938..a5f332b3d19 100644 --- a/providers/openai/docs/changelog.rst +++ b/providers/openai/docs/changelog.rst @@ -51,6 +51,18 @@ Changelog metadata DB on every run, regardless of the operator's ``do_xcom_push`` setting. +.. note:: + A deferred ``OpenAITriggerBatchOperator`` that times out now raises + ``OpenAIBatchTimeout``, matching the exception the non-deferrable path has always raised + for the same condition. Previously every non-success outcome of a deferred batch raised + the same ``OpenAIBatchJobException``, so a timeout and a genuine batch failure could not + be told apart or handled separately. + + A cancelled batch now raises ``OpenAIBatchCancelled``, a subclass of + ``OpenAIBatchJobException``, so existing code that catches ``OpenAIBatchJobException`` + keeps working unchanged while callers that want to distinguish cancellation can catch the + subclass specifically. + Breaking changes ~~~~~~~~~~~~~~~~ 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..2ee588bf281 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, @@ -96,6 +97,29 @@ class BatchStatus(str, Enum): #: Statuses the provider's trigger emits in its terminal event. TRIGGER_EVENT_STATUSES = frozenset({"success", "error", "cancelled"}) +# 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]] = { + "timeout": OpenAIBatchTimeout, + "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]: """ @@ -697,10 +721,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..d1332a55086 100644 --- a/providers/openai/src/airflow/providers/openai/operators/openai.py +++ b/providers/openai/src/airflow/providers/openai/operators/openai.py @@ -22,8 +22,11 @@ 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, + build_batch_error, + validate_execute_complete_event, +) from airflow.providers.openai.triggers.openai import OpenAIBatchTrigger if TYPE_CHECKING: @@ -372,6 +375,13 @@ class OpenAITriggerBatchOperator(BaseOperator): :param poll_interval: Optional. Number of seconds between checks. Only used when ``deferrable`` is True. Defaults to 60 seconds. + When ``deferrable`` is True and the batch does not reach a terminal state, ``execute_complete`` + raises :class:`~airflow.providers.openai.exceptions.OpenAIBatchTimeout`, matching the exception + raised by the synchronous path for the same condition. A cancelled batch raises + :class:`~airflow.providers.openai.exceptions.OpenAIBatchCancelled` (a subclass of + :class:`~airflow.providers.openai.exceptions.OpenAIBatchJobException`), and any other failure + raises :class:`~airflow.providers.openai.exceptions.OpenAIBatchJobException`. + .. seealso:: For more information on how to use this operator, please take a look at the guide: :ref:`howto/operator:OpenAITriggerBatchOperator` @@ -455,11 +465,14 @@ 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``: + ``OpenAIBatchTimeout`` for a timeout, ``OpenAIBatchCancelled`` for a cancellation, + and ``OpenAIBatchJobException`` for any other failure (including events from a + trigger serialized before ``termination_reason`` existed). """ event = validate_execute_complete_event(event) if event["status"] != "success": - raise OpenAIBatchJobException(event["message"]) + raise build_batch_error(event["message"], event.get("termination_reason")) self.log.info("%s completed successfully.", self.task_id) return event["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..767aaa0c4bc 100644 --- a/providers/openai/src/airflow/providers/openai/triggers/openai.py +++ b/providers/openai/src/airflow/providers/openai/triggers/openai.py @@ -103,6 +103,7 @@ class OpenAIBatchTrigger(BaseTrigger): yield TriggerEvent( { "status": "error", + "termination_reason": "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": "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": "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": "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": "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": "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": "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..6548d95654b 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,24 @@ def test_wait_for_in_progress_batch_timeout(mock_openai_hook, mock_wip_batch): assert mock_openai_hook.conn.batches.cancel.call_count == 1 [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..66613601cb1 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, @@ -935,3 +940,37 @@ 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"), + [ + pytest.param("timeout", OpenAIBatchTimeout, id="timeout"), + pytest.param("cancelled", OpenAIBatchCancelled, id="cancelled"), + pytest.param("failed", OpenAIBatchJobException, id="failed"), + pytest.param("expired", OpenAIBatchJobException, id="expired"), + pytest.param("unexpected_status", OpenAIBatchJobException, id="unexpected-status"), + pytest.param("polling_error", OpenAIBatchJobException, id="polling-error"), + ], + ) + def test_execute_complete_raises_exception_matching_termination_reason( + self, termination_reason, expected_exc + ): + event = { + "status": "error", + "termination_reason": termination_reason, + "message": "boom", + "batch_id": BATCH_ID, + } + with pytest.raises(expected_exc, match="boom"): + self._operator().execute_complete(Context(), event) + + @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 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()
