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 4aef57e2f29338f23237932b8cb6b0f025dae543 Author: Wei Lee <[email protected]> AuthorDate: Sat Sep 19 17:33:26 2026 +0900 Share one TerminationReason enum between the trigger and its consumers The trigger wrote termination_reason as a bare string and the hook keyed its exception map on bare strings, so a typo on either side type-checked fine and only a test could catch it. Both sides, and the operator's timeout check, now go through TerminationReason, making a misspelled reason an attribute error. The enum subclasses str, so the value crossing the triggerer boundary is unchanged and the existing tests still pin it. --- .../openai/src/airflow/providers/openai/hooks/openai.py | 17 +++++++++++++++-- .../src/airflow/providers/openai/operators/openai.py | 3 ++- .../src/airflow/providers/openai/triggers/openai.py | 16 ++++++++-------- 3 files changed, 25 insertions(+), 11 deletions(-) diff --git a/providers/openai/src/airflow/providers/openai/hooks/openai.py b/providers/openai/src/airflow/providers/openai/hooks/openai.py index 4e6562f27a6..761f29c4fd2 100644 --- a/providers/openai/src/airflow/providers/openai/hooks/openai.py +++ b/providers/openai/src/airflow/providers/openai/hooks/openai.py @@ -97,13 +97,26 @@ class BatchStatus(str, Enum): #: Statuses the provider's trigger emits in its terminal event. 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]] = { - "timeout": OpenAIBatchTimeout, - "cancelled": OpenAIBatchCancelled, + TerminationReason.TIMEOUT: OpenAIBatchTimeout, + TerminationReason.CANCELLED: OpenAIBatchCancelled, } diff --git a/providers/openai/src/airflow/providers/openai/operators/openai.py b/providers/openai/src/airflow/providers/openai/operators/openai.py index a0c936ba690..68b116bb82f 100644 --- a/providers/openai/src/airflow/providers/openai/operators/openai.py +++ b/providers/openai/src/airflow/providers/openai/operators/openai.py @@ -24,6 +24,7 @@ from typing import TYPE_CHECKING, Any, ClassVar from airflow.providers.common.compat.sdk import BaseOperator, conf from airflow.providers.openai.hooks.openai import ( OpenAIHook, + TerminationReason, build_batch_error, validate_execute_complete_event, ) @@ -486,7 +487,7 @@ class OpenAITriggerBatchOperator(BaseOperator): """ event = validate_execute_complete_event(event) if event["status"] != "success": - if event.get("termination_reason") == "timeout": + if event.get("termination_reason") == TerminationReason.TIMEOUT: batch_id = event["batch_id"] self.log.warning( "%s timed out waiting for batch %s; requesting cancellation.", diff --git a/providers/openai/src/airflow/providers/openai/triggers/openai.py b/providers/openai/src/airflow/providers/openai/triggers/openai.py index 767aaa0c4bc..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,7 +103,7 @@ class OpenAIBatchTrigger(BaseTrigger): yield TriggerEvent( { "status": "error", - "termination_reason": "timeout", + "termination_reason": TerminationReason.TIMEOUT, "message": ( f"Batch {self.batch_id} has not reached a terminal status after " f"{elapsed:.0f} seconds." @@ -117,7 +117,7 @@ class OpenAIBatchTrigger(BaseTrigger): yield TriggerEvent( { "status": "success", - "termination_reason": "completed", + "termination_reason": TerminationReason.COMPLETED, "message": f"Batch {self.batch_id} has completed successfully.", "batch_id": self.batch_id, } @@ -126,7 +126,7 @@ class OpenAIBatchTrigger(BaseTrigger): yield TriggerEvent( { "status": "cancelled", - "termination_reason": "cancelled", + "termination_reason": TerminationReason.CANCELLED, "message": f"Batch {self.batch_id} has been cancelled.", "batch_id": self.batch_id, } @@ -135,7 +135,7 @@ class OpenAIBatchTrigger(BaseTrigger): yield TriggerEvent( { "status": "error", - "termination_reason": "failed", + "termination_reason": TerminationReason.FAILED, "message": f"Batch failed:\n{self.batch_id}", "batch_id": self.batch_id, } @@ -144,7 +144,7 @@ class OpenAIBatchTrigger(BaseTrigger): yield TriggerEvent( { "status": "error", - "termination_reason": "expired", + "termination_reason": TerminationReason.EXPIRED, "message": f"Batch couldn't be completed within its completion window:\n{self.batch_id}", "batch_id": self.batch_id, } @@ -153,7 +153,7 @@ class OpenAIBatchTrigger(BaseTrigger): yield TriggerEvent( { "status": "error", - "termination_reason": "unexpected_status", + "termination_reason": TerminationReason.UNEXPECTED_STATUS, "message": f"Batch {self.batch_id} has failed.", "batch_id": self.batch_id, } @@ -162,7 +162,7 @@ class OpenAIBatchTrigger(BaseTrigger): yield TriggerEvent( { "status": "error", - "termination_reason": "polling_error", + "termination_reason": TerminationReason.POLLING_ERROR, "message": str(e), "batch_id": self.batch_id, }
