diff --git a/providers/openai/docs/changelog.rst b/providers/openai/docs/changelog.rst index 3dda0fb893860..cf4f0d37d2a48 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 09618b9048e6e..cca4ef1977acd 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 f9ba07a849784..761f29c4fd2e3 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 @@ 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 @@ 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. @@ -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) @@ -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( diff --git a/providers/openai/src/airflow/providers/openai/operators/openai.py b/providers/openai/src/airflow/providers/openai/operators/openai.py index d0234b17b0f30..fe03fe1bf407e 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 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 @@ 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: + """ + 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 49fd0900cc022..2b41b3f6ea284 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 @@ 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 @@ 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." @@ -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, } @@ -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, } @@ -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, } @@ -140,7 +144,8 @@ 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, } ) @@ -148,9 +153,17 @@ async def run(self) -> AsyncIterator[TriggerEvent]: 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 cc370911777cd..6b56df91f5583 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.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) + + +@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 diff --git a/providers/openai/tests/unit/openai/operators/test_openai.py b/providers/openai/tests/unit/openai/operators/test_openai.py index 7ec22be98cc15..7e6cdc5c68443 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_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 @@ def test_failed_event_raises(self, event): 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 fabaad35343f0..b9c6a14e44667 100644 --- a/providers/openai/tests/unit/openai/test_exceptions.py +++ b/providers/openai/tests/unit/openai/test_exceptions.py @@ -21,7 +21,11 @@ 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 @@ [ 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 0a2c4322a5ed6..d72620ed3c501 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 @@ def test_rejects_both_timeout_and_end_time(self): @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 @@ async def test_openai_batch_for_terminal_status( ) expected_result = { "status": mock_status, + "termination_reason": mock_termination_reason, "message": mock_message, "batch_id": self.BATCH_ID, } @@ -187,6 +194,7 @@ async def test_openai_batch_for_timeout(self, mock_monotonic, mock_batch, mock_b 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 @@ async def test_openai_batch_yields_single_terminal_event(self, mock_batch): TriggerEvent( { "status": "success", + "termination_reason": "completed", "message": f"Batch {self.BATCH_ID} has completed successfully.", "batch_id": self.BATCH_ID, } @@ -254,6 +263,7 @@ async def test_openai_batch_for_unexpected_error(self, mock_batch): ) expected_result = { "status": "error", + "termination_reason": "polling_error", "message": "'float' object has no attribute 'status'", "batch_id": self.BATCH_ID, } @@ -261,3 +271,26 @@ async def test_openai_batch_for_unexpected_error(self, mock_batch): 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()