diff --git a/providers/ssh/src/airflow/providers/ssh/hooks/ssh.py b/providers/ssh/src/airflow/providers/ssh/hooks/ssh.py index 5a5d1ab0fb933..43fca90dce57b 100644 --- a/providers/ssh/src/airflow/providers/ssh/hooks/ssh.py +++ b/providers/ssh/src/airflow/providers/ssh/hooks/ssh.py @@ -33,11 +33,7 @@ from tenacity import Retrying, stop_after_attempt, wait_fixed, wait_random from airflow.exceptions import AirflowException - -try: - from airflow.sdk import BaseHook -except ImportError: - from airflow.hooks.base import BaseHook # type: ignore[attr-defined,no-redef] +from airflow.providers.ssh.version_compat import BaseHook from airflow.utils.platform import getuser from airflow.utils.types import NOTSET, ArgNotSet diff --git a/providers/ssh/src/airflow/providers/ssh/version_compat.py b/providers/ssh/src/airflow/providers/ssh/version_compat.py index 4f8d5e32bca4a..f800a29bbb739 100644 --- a/providers/ssh/src/airflow/providers/ssh/version_compat.py +++ b/providers/ssh/src/airflow/providers/ssh/version_compat.py @@ -33,10 +33,16 @@ def get_base_airflow_version_tuple() -> tuple[int, int, int]: AIRFLOW_V_3_0_PLUS = get_base_airflow_version_tuple() >= (3, 0, 0) +AIRFLOW_V_3_1_PLUS: bool = get_base_airflow_version_tuple() >= (3, 1, 0) + +if AIRFLOW_V_3_1_PLUS: + from airflow.sdk import BaseHook +else: + from airflow.hooks.base import BaseHook # type: ignore[attr-defined,no-redef] if AIRFLOW_V_3_0_PLUS: from airflow.sdk import BaseOperator else: from airflow.models import BaseOperator # type: ignore[no-redef] -__all__ = ["AIRFLOW_V_3_0_PLUS", "BaseOperator"] +__all__ = ["AIRFLOW_V_3_0_PLUS", "AIRFLOW_V_3_1_PLUS", "BaseHook", "BaseOperator"]