Source code for apache_airflow_providers_imednet.sensors
"""Airflow sensors for iMednet operations."""
from __future__ import annotations
from collections.abc import Sequence
from typing import Any
try: # pragma: no cover - optional Airflow dependency
try:
from airflow.sdk.bases.sensor import BaseSensorOperator
except ImportError:
from airflow.sensors.base import BaseSensorOperator
except (ImportError, ModuleNotFoundError): # pragma: no cover - placeholder fallback
class BaseSensorOperator: # type: ignore
"""Fallback BaseSensorOperator."""
template_fields: Sequence[str] = ()
def __init__(self, *args: Any, **kwargs: Any) -> None: # pragma: no cover
"""Initialize fallback BaseSensorOperator."""
from imednet import ImednetSDK
from ._airflow_compat import AirflowException, Context
from .hooks import ImednetHook
[docs]class ImednetJobSensor(BaseSensorOperator):
"""Poll iMednet for job completion."""
template_fields: Sequence[str] = ("study_key", "batch_id")
[docs] def __init__(
self,
*,
study_key: str,
batch_id: str,
imednet_conn_id: str = "imednet_default",
poke_interval: float = 60,
**kwargs: Any,
) -> None:
"""Initialize the job sensor.
:param study_key: The study key identifier.
:param batch_id: The batch or job identifier to monitor.
:param imednet_conn_id: Airflow connection ID to use for credentials.
:param poke_interval: Seconds between polling attempts.
:param kwargs: Additional Airflow BaseSensorOperator arguments.
"""
super().__init__(poke_interval=poke_interval, **kwargs)
self.study_key = study_key
self.batch_id = batch_id
self.imednet_conn_id = imednet_conn_id
def _get_sdk(self) -> ImednetSDK:
"""Get the Imednet SDK client."""
return ImednetHook(self.imednet_conn_id).get_sdk_client()
[docs] def poke(self, context: Context) -> bool:
"""Check the status of the job."""
from imednet.spi.models import JobStatus
from imednet.spi.utils import JobFailedError, evaluate_job_state
sdk = self._get_sdk()
job: JobStatus = sdk.jobs.get(self.study_key, self.batch_id)
try:
return evaluate_job_state(job)
except JobFailedError as e:
raise AirflowException(str(e)) from e
__all__ = ["ImednetJobSensor"]