From 1752941fec89320846db55aca7a4bf1b943cad2e Mon Sep 17 00:00:00 2001 From: Shahar Epstein <60007259+shahar1@users.noreply.github.com> Date: Fri, 15 May 2026 08:36:17 +0300 Subject: [PATCH 1/9] Enable ruff B008 (function-call-in-default-argument) and fix violations Adds [tool.ruff.lint.flake8-bugbear] extend-immutable-calls to exempt FastAPI DI callables (Depends, Query, Path, Body, Security) and the stateless cryptography SHA256 descriptor from B008, then fixes all remaining violations where conf.get*() or mutable objects were evaluated once at import time rather than at call time: - task-sdk: sensor timeout and Resources cpus/ram/disk/gpus now read from config at instantiation - providers/google: s3_to_gcs deferrable default, Vertex AI Ray head_node_type mutable default - shared/observability: SafeStatsdLogger, SafeDogStatsdLogger, and SafeOtelLogger metrics_validator/metric_tags_validator mutable defaults (PatternAllowListValidator shared across instances) - openlineage system test: setup_jinja() mutable Jinja Environment default - api_fastapi/common/parameters: LimitFilter.depends conf.getint moved to module-level constant (_FALLBACK_PAGE_LIMIT) to preserve OpenAPI schema default while making the intent explicit - providers/fab tests: types.SimpleNamespace mutable default in conftest --- .../src/airflow/api_fastapi/common/parameters.py | 6 ++++-- .../fab/auth_manager/api_fastapi/conftest.py | 4 +++- .../google/cloud/hooks/vertex_ai/ray.py | 4 ++-- .../google/cloud/operators/vertex_ai/ray.py | 4 ++-- .../google/cloud/transfers/s3_to_gcs.py | 8 ++++++-- .../tests/system/openlineage/operator.py | 4 ++-- pyproject.toml | 13 +++++++++++++ .../observability/metrics/datadog_logger.py | 12 ++++++++---- .../observability/metrics/otel_logger.py | 6 ++++-- .../observability/metrics/statsd_logger.py | 12 ++++++++---- task-sdk/src/airflow/sdk/bases/sensor.py | 4 +++- .../sdk/definitions/operator_resources.py | 16 ++++++++-------- 12 files changed, 63 insertions(+), 30 deletions(-) diff --git a/airflow-core/src/airflow/api_fastapi/common/parameters.py b/airflow-core/src/airflow/api_fastapi/common/parameters.py index a93ec040e1cec..7aa2eb222c64a 100644 --- a/airflow-core/src/airflow/api_fastapi/common/parameters.py +++ b/airflow-core/src/airflow/api_fastapi/common/parameters.py @@ -77,6 +77,8 @@ T = TypeVar("T") +_FALLBACK_PAGE_LIMIT: int = conf.getint("api", "fallback_page_limit") + class BaseParam(OrmClause[T], ABC): """Base class for path or query parameters with ORM transformation.""" @@ -106,7 +108,7 @@ def to_orm(self, select: Select) -> Select: return select.limit(self.value) @classmethod - def depends(cls, limit: NonNegativeInt = conf.getint("api", "fallback_page_limit")) -> LimitFilter: + def depends(cls, limit: NonNegativeInt = _FALLBACK_PAGE_LIMIT) -> LimitFilter: return cls().set_value(min(limit, conf.getint("api", "maximum_page_limit"))) @@ -611,7 +613,7 @@ def inner( order_by: list[str] = Query( default=default_list, description=f"Attributes to order by, multi criteria sort is supported. Prefix with `-` for descending order. " - f"Supported attributes: `{', '.join(all_attrs) if all_attrs else self.get_primary_key_string()}`", + f"Supported attributes: `{', '.join(all_attrs) if all_attrs else self.get_primary_key_string()}`", # noqa: B008 ), ) -> SortParam: return self.set_value(order_by) diff --git a/providers/fab/tests/unit/fab/auth_manager/api_fastapi/conftest.py b/providers/fab/tests/unit/fab/auth_manager/api_fastapi/conftest.py index fb884b8ede7f4..af2615f7581ba 100644 --- a/providers/fab/tests/unit/fab/auth_manager/api_fastapi/conftest.py +++ b/providers/fab/tests/unit/fab/auth_manager/api_fastapi/conftest.py @@ -63,7 +63,9 @@ def _use(mapping: dict): @pytest.fixture def as_user(override_deps): @contextmanager - def _as(u=types.SimpleNamespace(id=1, username="tester")): + def _as(u=None): + if u is None: + u = types.SimpleNamespace(id=1, username="tester") with override_deps({get_user_dep: lambda: u}): yield u diff --git a/providers/google/src/airflow/providers/google/cloud/hooks/vertex_ai/ray.py b/providers/google/src/airflow/providers/google/cloud/hooks/vertex_ai/ray.py index de722ed595926..9d85480414530 100644 --- a/providers/google/src/airflow/providers/google/cloud/hooks/vertex_ai/ray.py +++ b/providers/google/src/airflow/providers/google/cloud/hooks/vertex_ai/ray.py @@ -66,7 +66,7 @@ def create_ray_cluster( self, project_id: str, location: str, - head_node_type: resources.Resources = resources.Resources(), + head_node_type: resources.Resources | None = None, python_version: str = "3.10", ray_version: str = "2.33", network: str | None = None, @@ -115,7 +115,7 @@ def create_ray_cluster( """ aiplatform.init(project=project_id, location=location, credentials=self.get_credentials()) cluster_path = vertex_ray.create_ray_cluster( - head_node_type=head_node_type, + head_node_type=head_node_type or resources.Resources(), python_version=python_version, ray_version=ray_version, network=network, diff --git a/providers/google/src/airflow/providers/google/cloud/operators/vertex_ai/ray.py b/providers/google/src/airflow/providers/google/cloud/operators/vertex_ai/ray.py index 4d6723977afe2..103d6d60dabba 100644 --- a/providers/google/src/airflow/providers/google/cloud/operators/vertex_ai/ray.py +++ b/providers/google/src/airflow/providers/google/cloud/operators/vertex_ai/ray.py @@ -140,7 +140,7 @@ def __init__( self, python_version: str, ray_version: Literal["2.9.3", "2.33", "2.42"], - head_node_type: resources.Resources = resources.Resources(), + head_node_type: resources.Resources | None = None, network: str | None = None, service_account: str | None = None, cluster_name: str | None = None, @@ -155,7 +155,7 @@ def __init__( **kwargs, ) -> None: super().__init__(*args, **kwargs) - self.head_node_type = head_node_type + self.head_node_type = head_node_type if head_node_type is not None else resources.Resources() self.python_version = python_version self.ray_version = ray_version self.network = network diff --git a/providers/google/src/airflow/providers/google/cloud/transfers/s3_to_gcs.py b/providers/google/src/airflow/providers/google/cloud/transfers/s3_to_gcs.py index 861769d57680c..3e6c1e7837a19 100644 --- a/providers/google/src/airflow/providers/google/cloud/transfers/s3_to_gcs.py +++ b/providers/google/src/airflow/providers/google/cloud/transfers/s3_to_gcs.py @@ -165,7 +165,7 @@ def __init__( replace=False, gzip=False, google_impersonation_chain: str | Sequence[str] | None = None, - deferrable=conf.getboolean("operators", "default_deferrable", fallback=False), + deferrable=None, poll_interval: int = 10, return_gcs_uris: bool = False, **kwargs, @@ -178,7 +178,11 @@ def __init__( self.verify = verify self.gzip = gzip self.google_impersonation_chain = google_impersonation_chain - self.deferrable = deferrable + self.deferrable = ( + deferrable + if deferrable is not None + else conf.getboolean("operators", "default_deferrable", fallback=False) + ) if poll_interval <= 0: raise ValueError("Invalid value for poll_interval. Expected value greater than 0") self.poll_interval = poll_interval diff --git a/providers/openlineage/tests/system/openlineage/operator.py b/providers/openlineage/tests/system/openlineage/operator.py index 0bc8aa7c003b5..b240b1602a9fe 100644 --- a/providers/openlineage/tests/system/openlineage/operator.py +++ b/providers/openlineage/tests/system/openlineage/operator.py @@ -197,7 +197,7 @@ def __init__( self, event_templates: dict[str, dict] | None = None, file_path: str | None = None, - env: Environment = setup_jinja(), + env: Environment | None = None, allow_duplicate_events_regex: str | None = None, clear_variables: bool = True, **kwargs, @@ -205,7 +205,7 @@ def __init__( super().__init__(**kwargs) self.event_templates = event_templates self.file_path = file_path - self.env = env + self.env = env if env is not None else setup_jinja() self.allow_duplicate_events_regex = allow_duplicate_events_regex self.clear_variables = clear_variables if self.event_templates and self.file_path: diff --git a/pyproject.toml b/pyproject.toml index fda1174169ee3..640d202fb2c2d 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -655,6 +655,7 @@ extend-select = [ "B004", # Checks for use of hasattr(x, "__call__") and replaces it with callable(x) "B006", # Checks for uses of mutable objects as function argument defaults. "B007", # Checks for unused variables in the loop + "B008", # Do not perform function call in argument defaults (use extend-immutable-calls for FastAPI DI) "B012", # Checks for `break`, `continue`, and `return` statements in `finally` blocks "B017", # Checks for pytest.raises context managers that catch Exception or BaseException. "B019", # Use of functools.lru_cache or functools.cache on methods can lead to memory leaks @@ -703,6 +704,18 @@ unfixable = [ "PT022", ] +[tool.ruff.lint.flake8-bugbear] +# FastAPI dependency injection uses function calls in argument defaults intentionally. +# SHA256 is a stateless algorithm descriptor (cryptography library). +extend-immutable-calls = [ + "fastapi.Body", + "fastapi.Depends", + "fastapi.Query", + "fastapi.Path", + "fastapi.Security", + "cryptography.hazmat.primitives.hashes.SHA256", +] + [tool.ruff.format] docstring-code-format = true diff --git a/shared/observability/src/airflow_shared/observability/metrics/datadog_logger.py b/shared/observability/src/airflow_shared/observability/metrics/datadog_logger.py index e6daf53b793c9..a018226a3c8c1 100644 --- a/shared/observability/src/airflow_shared/observability/metrics/datadog_logger.py +++ b/shared/observability/src/airflow_shared/observability/metrics/datadog_logger.py @@ -45,16 +45,20 @@ class SafeDogStatsdLogger: def __init__( self, dogstatsd_client: DogStatsd, - metrics_validator: ListValidator = PatternAllowListValidator(), + metrics_validator: ListValidator | None = None, metrics_tags: bool = False, - metric_tags_validator: ListValidator = PatternAllowListValidator(), + metric_tags_validator: ListValidator | None = None, stat_name_handler: Callable[[str], str] | None = None, statsd_influxdb_enabled: bool = False, ) -> None: self.dogstatsd = dogstatsd_client - self.metrics_validator = metrics_validator + self.metrics_validator = ( + metrics_validator if metrics_validator is not None else PatternAllowListValidator() + ) self.metrics_tags = metrics_tags - self.metric_tags_validator = metric_tags_validator + self.metric_tags_validator = ( + metric_tags_validator if metric_tags_validator is not None else PatternAllowListValidator() + ) self.stat_name_handler = stat_name_handler self.statsd_influxdb_enabled = statsd_influxdb_enabled diff --git a/shared/observability/src/airflow_shared/observability/metrics/otel_logger.py b/shared/observability/src/airflow_shared/observability/metrics/otel_logger.py index 5aaa77741f0e5..95c62617d8444 100644 --- a/shared/observability/src/airflow_shared/observability/metrics/otel_logger.py +++ b/shared/observability/src/airflow_shared/observability/metrics/otel_logger.py @@ -175,13 +175,15 @@ def __init__( self, otel_provider, prefix: str = DEFAULT_METRIC_NAME_PREFIX, - metrics_validator: ListValidator = PatternAllowListValidator(), + metrics_validator: ListValidator | None = None, stat_name_handler: Callable[[str], str] | None = None, statsd_influxdb_enabled: bool = False, ): self.otel: Callable = otel_provider self.prefix: str = prefix - self.metrics_validator = metrics_validator + self.metrics_validator = ( + metrics_validator if metrics_validator is not None else PatternAllowListValidator() + ) self.meter = otel_provider.get_meter(__name__) self.metrics_map = MetricsMap(self.meter) self.stat_name_handler = stat_name_handler diff --git a/shared/observability/src/airflow_shared/observability/metrics/statsd_logger.py b/shared/observability/src/airflow_shared/observability/metrics/statsd_logger.py index 7e8d29f3a26d8..6fa3828f11e19 100644 --- a/shared/observability/src/airflow_shared/observability/metrics/statsd_logger.py +++ b/shared/observability/src/airflow_shared/observability/metrics/statsd_logger.py @@ -67,16 +67,20 @@ class SafeStatsdLogger: def __init__( self, statsd_client: StatsClient, - metrics_validator: ListValidator = PatternAllowListValidator(), + metrics_validator: ListValidator | None = None, influxdb_tags_enabled: bool = False, - metric_tags_validator: ListValidator = PatternAllowListValidator(), + metric_tags_validator: ListValidator | None = None, stat_name_handler: Callable[[str], str] | None = None, statsd_influxdb_enabled: bool = False, ) -> None: self.statsd = statsd_client - self.metrics_validator = metrics_validator + self.metrics_validator = ( + metrics_validator if metrics_validator is not None else PatternAllowListValidator() + ) self.influxdb_tags_enabled = influxdb_tags_enabled - self.metric_tags_validator = metric_tags_validator + self.metric_tags_validator = ( + metric_tags_validator if metric_tags_validator is not None else PatternAllowListValidator() + ) self.stat_name_handler = stat_name_handler self.statsd_influxdb_enabled = statsd_influxdb_enabled diff --git a/task-sdk/src/airflow/sdk/bases/sensor.py b/task-sdk/src/airflow/sdk/bases/sensor.py index 3a877dd98aff5..3f1f0842e8982 100644 --- a/task-sdk/src/airflow/sdk/bases/sensor.py +++ b/task-sdk/src/airflow/sdk/bases/sensor.py @@ -116,7 +116,7 @@ def __init__( self, *, poke_interval: timedelta | float = 60, - timeout: timedelta | float = conf.getfloat("sensors", "default_timeout"), + timeout: timedelta | float | None = None, soft_fail: bool = False, mode: str = "poke", exponential_backoff: bool = False, @@ -128,6 +128,8 @@ def __init__( super().__init__(**kwargs) self.poke_interval = self._coerce_poke_interval(poke_interval).total_seconds() self.soft_fail = soft_fail + if timeout is None: + timeout = conf.getfloat("sensors", "default_timeout") self.timeout: int | float = self._coerce_timeout(timeout).total_seconds() self.mode = mode self.exponential_backoff = exponential_backoff diff --git a/task-sdk/src/airflow/sdk/definitions/operator_resources.py b/task-sdk/src/airflow/sdk/definitions/operator_resources.py index d6cbf10039d00..d8a0d02a1b440 100644 --- a/task-sdk/src/airflow/sdk/definitions/operator_resources.py +++ b/task-sdk/src/airflow/sdk/definitions/operator_resources.py @@ -125,15 +125,15 @@ class Resources: def __init__( self, - cpus=conf.getint("operators", "default_cpus"), - ram=conf.getint("operators", "default_ram"), - disk=conf.getint("operators", "default_disk"), - gpus=conf.getint("operators", "default_gpus"), + cpus=None, + ram=None, + disk=None, + gpus=None, ): - self.cpus = CpuResource(cpus) - self.ram = RamResource(ram) - self.disk = DiskResource(disk) - self.gpus = GpuResource(gpus) + self.cpus = CpuResource(cpus if cpus is not None else conf.getint("operators", "default_cpus")) + self.ram = RamResource(ram if ram is not None else conf.getint("operators", "default_ram")) + self.disk = DiskResource(disk if disk is not None else conf.getint("operators", "default_disk")) + self.gpus = GpuResource(gpus if gpus is not None else conf.getint("operators", "default_gpus")) def __eq__(self, other: object) -> bool: if not isinstance(other, self.__class__): From 13323d3cebae9659016e239eeafc11e3985d7485 Mon Sep 17 00:00:00 2001 From: Shahar Epstein <60007259+shahar1@users.noreply.github.com> Date: Fri, 15 May 2026 09:06:31 +0300 Subject: [PATCH 2/9] Address Copilot review feedback on B008 PR - Make RayHook.create_ray_cluster use `is not None` for head_node_type fallback so an explicit falsy value is forwarded instead of being replaced; matches the operator's behavior. - Add regression tests for the conf-default-at-instantiation behavior: * Sensor `timeout` reads `sensors.default_timeout` at construction * `Resources()` reads `operators.default_*` at construction and preserves explicit 0 instead of falling back to config * `S3ToGCSOperator.deferrable` reads `operators.default_deferrable` at construction; explicit `False` overrides a truthy config value --- .../google/cloud/hooks/vertex_ai/ray.py | 2 +- .../google/cloud/transfers/test_s3_to_gcs.py | 41 +++++++++++++++++++ task-sdk/tests/task_sdk/bases/test_sensor.py | 14 +++++++ .../definitions/test_operator_resources.py | 38 +++++++++++++++++ 4 files changed, 94 insertions(+), 1 deletion(-) diff --git a/providers/google/src/airflow/providers/google/cloud/hooks/vertex_ai/ray.py b/providers/google/src/airflow/providers/google/cloud/hooks/vertex_ai/ray.py index 9d85480414530..6328c0f33febc 100644 --- a/providers/google/src/airflow/providers/google/cloud/hooks/vertex_ai/ray.py +++ b/providers/google/src/airflow/providers/google/cloud/hooks/vertex_ai/ray.py @@ -115,7 +115,7 @@ def create_ray_cluster( """ aiplatform.init(project=project_id, location=location, credentials=self.get_credentials()) cluster_path = vertex_ray.create_ray_cluster( - head_node_type=head_node_type or resources.Resources(), + head_node_type=head_node_type if head_node_type is not None else resources.Resources(), python_version=python_version, ray_version=ray_version, network=network, diff --git a/providers/google/tests/unit/google/cloud/transfers/test_s3_to_gcs.py b/providers/google/tests/unit/google/cloud/transfers/test_s3_to_gcs.py index 19b3c1117289b..389356ef1de6f 100644 --- a/providers/google/tests/unit/google/cloud/transfers/test_s3_to_gcs.py +++ b/providers/google/tests/unit/google/cloud/transfers/test_s3_to_gcs.py @@ -32,6 +32,8 @@ ) from airflow.utils.timezone import utcnow +from tests_common.test_utils.config import conf_vars + PROJECT_ID = "test-project-id" TASK_ID = "test-s3-gcs-operator" S3_BUCKET = "test-bucket" @@ -103,6 +105,45 @@ def test_init(self): assert operator.poll_interval == POLL_INTERVAL assert operator.return_gcs_uris is True + def test_deferrable_default_read_from_conf_at_instantiation(self): + """When ``deferrable`` is omitted, the operator should read + ``operators.default_deferrable`` at instantiation time (not at module import). + """ + with conf_vars({("operators", "default_deferrable"): "True"}): + operator = S3ToGCSOperator( + task_id=TASK_ID, + bucket=S3_BUCKET, + prefix=S3_PREFIX, + dest_gcs=GCS_PATH_PREFIX, + return_gcs_uris=True, + ) + assert operator.deferrable is True + + with conf_vars({("operators", "default_deferrable"): "False"}): + operator = S3ToGCSOperator( + task_id=TASK_ID + "_2", + bucket=S3_BUCKET, + prefix=S3_PREFIX, + dest_gcs=GCS_PATH_PREFIX, + return_gcs_uris=True, + ) + assert operator.deferrable is False + + def test_deferrable_explicit_false_overrides_truthy_conf(self): + """An explicit ``deferrable=False`` must override a truthy + ``operators.default_deferrable`` config — only ``None`` falls back to config. + """ + with conf_vars({("operators", "default_deferrable"): "True"}): + operator = S3ToGCSOperator( + task_id=TASK_ID, + bucket=S3_BUCKET, + prefix=S3_PREFIX, + dest_gcs=GCS_PATH_PREFIX, + deferrable=False, + return_gcs_uris=True, + ) + assert operator.deferrable is False + @mock.patch("airflow.providers.google.cloud.transfers.s3_to_gcs.S3Hook") @mock.patch("airflow.providers.google.cloud.transfers.s3_to_gcs.GCSHook") def test_execute(self, gcs_mock_hook, s3_mock_hook): diff --git a/task-sdk/tests/task_sdk/bases/test_sensor.py b/task-sdk/tests/task_sdk/bases/test_sensor.py index 8a6fbcc15acf2..5e7f588a1551a 100644 --- a/task-sdk/tests/task_sdk/bases/test_sensor.py +++ b/task-sdk/tests/task_sdk/bases/test_sensor.py @@ -41,6 +41,8 @@ from airflow.sdk.execution_time.comms import RescheduleTask, TaskRescheduleStartDate from airflow.sdk.timezone import datetime +from tests_common.test_utils.config import conf_vars + if TYPE_CHECKING: from airflow.sdk.definitions.context import Context @@ -358,6 +360,18 @@ def test_sensor_with_invalid_timeout(self): task_id="test_sensor_task_3", return_value=None, poke_interval=10, timeout=positive_timeout ) + def test_sensor_timeout_default_read_from_conf_at_instantiation(self): + """When ``timeout`` is not supplied, it should be read from ``sensors.default_timeout`` + at instantiation time (not at module import time). + """ + with conf_vars({("sensors", "default_timeout"): "12345"}): + sensor = DummySensor(task_id="test_sensor_default_timeout", return_value=None, poke_interval=10) + assert sensor.timeout == 12345 + + with conf_vars({("sensors", "default_timeout"): "67"}): + sensor = DummySensor(task_id="test_sensor_default_timeout_2", return_value=None, poke_interval=10) + assert sensor.timeout == 67 + def test_sensor_with_exponential_backoff_off(self): sensor = DummySensor( task_id=SENSOR_OP, return_value=None, poke_interval=5, timeout=60, exponential_backoff=False diff --git a/task-sdk/tests/task_sdk/definitions/test_operator_resources.py b/task-sdk/tests/task_sdk/definitions/test_operator_resources.py index 9e0875cf076e8..389ab8701f63f 100644 --- a/task-sdk/tests/task_sdk/definitions/test_operator_resources.py +++ b/task-sdk/tests/task_sdk/definitions/test_operator_resources.py @@ -19,6 +19,8 @@ from airflow.sdk.definitions.operator_resources import Resources +from tests_common.test_utils.config import conf_vars + class TestResources: def test_resource_eq(self): @@ -41,3 +43,39 @@ def test_to_dict(self): "disk": {"name": "Disk", "qty": 1024, "units_str": "MB"}, "gpus": {"name": "GPU", "qty": 1, "units_str": "gpu(s)"}, } + + def test_defaults_read_from_conf_at_instantiation(self): + """When fields are omitted, ``Resources`` should read defaults from the ``operators`` + section at instantiation time (not at module import time). + """ + with conf_vars( + { + ("operators", "default_cpus"): "7", + ("operators", "default_ram"): "5120", + ("operators", "default_disk"): "8192", + ("operators", "default_gpus"): "3", + } + ): + r = Resources() + assert r.cpus.qty == 7 + assert r.ram.qty == 5120 + assert r.disk.qty == 8192 + assert r.gpus.qty == 3 + + def test_falsy_zero_values_are_preserved(self): + """Explicit ``0`` for a resource must not be replaced by the config default — + only ``None`` (the sentinel for "not supplied") should fall back to config. + """ + with conf_vars( + { + ("operators", "default_cpus"): "4", + ("operators", "default_ram"): "2048", + ("operators", "default_disk"): "1024", + ("operators", "default_gpus"): "2", + } + ): + r = Resources(cpus=0, ram=0, disk=0, gpus=0) + assert r.cpus.qty == 0 + assert r.ram.qty == 0 + assert r.disk.qty == 0 + assert r.gpus.qty == 0 From 81b1a4b5139b7b6cca9617e3eafc9fe70d7b3cdb Mon Sep 17 00:00:00 2001 From: Shahar Epstein <60007259+shahar1@users.noreply.github.com> Date: Fri, 15 May 2026 09:23:24 +0300 Subject: [PATCH 3/9] Add regression tests for Ray cluster default head_node_type fix Addresses remaining Copilot review comments on #66979: - Hook test: verify calling create_ray_cluster with head_node_type=None passes a fresh Resources() instance to vertex_ray.create_ray_cluster - New operator test file for vertex_ai/ray.py: verify omitting head_node_type uses a fresh Resources() per instance (not shared), and that execute() forwards a Resources() instance to the hook --- .../google/cloud/hooks/vertex_ai/test_ray.py | 17 ++++ .../cloud/operators/vertex_ai/test_ray.py | 99 +++++++++++++++++++ 2 files changed, 116 insertions(+) create mode 100644 providers/google/tests/unit/google/cloud/operators/vertex_ai/test_ray.py diff --git a/providers/google/tests/unit/google/cloud/hooks/vertex_ai/test_ray.py b/providers/google/tests/unit/google/cloud/hooks/vertex_ai/test_ray.py index d3d071c4b3ffd..26205939a85e6 100644 --- a/providers/google/tests/unit/google/cloud/hooks/vertex_ai/test_ray.py +++ b/providers/google/tests/unit/google/cloud/hooks/vertex_ai/test_ray.py @@ -94,6 +94,23 @@ def test_create_ray_cluster(self, mock_aiplatform_init, mock_create_ray_cluster) labels=None, ) + @mock.patch(RAY_STRING.format("vertex_ray.create_ray_cluster")) + @mock.patch(RAY_STRING.format("aiplatform.init")) + def test_create_ray_cluster_default_head_node_type( + self, mock_aiplatform_init, mock_create_ray_cluster + ) -> None: + self.hook.create_ray_cluster( + project_id=TEST_PROJECT_ID, + location=TEST_LOCATION, + head_node_type=None, + python_version=TEST_PYTHON_VERSION, + ray_version=TEST_RAY_VERSION, + cluster_name=TEST_CLUSTER_NAME, + ) + mock_aiplatform_init.assert_called_once() + call_kwargs = mock_create_ray_cluster.call_args.kwargs + assert isinstance(call_kwargs["head_node_type"], Resources) + @mock.patch(RAY_STRING.format("vertex_ray.delete_ray_cluster")) @mock.patch(RAY_STRING.format("aiplatform.init")) @mock.patch(RAY_STRING.format("PersistentResourceServiceClient.persistent_resource_path")) diff --git a/providers/google/tests/unit/google/cloud/operators/vertex_ai/test_ray.py b/providers/google/tests/unit/google/cloud/operators/vertex_ai/test_ray.py new file mode 100644 index 0000000000000..889c81849e08b --- /dev/null +++ b/providers/google/tests/unit/google/cloud/operators/vertex_ai/test_ray.py @@ -0,0 +1,99 @@ +# +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +from __future__ import annotations + +from unittest import mock + +import pytest + +pytest.importorskip("google.cloud.aiplatform.vertex_ray.util.resources") +from google.cloud.aiplatform.vertex_ray.util.resources import Resources + +from airflow.providers.google.cloud.operators.vertex_ai.ray import CreateRayClusterOperator + +TEST_GCP_CONN_ID = "test-gcp-conn-id" +TEST_LOCATION = "us-central1" +TEST_PROJECT_ID = "test-project-id" +TEST_PYTHON_VERSION = "3.10" +TEST_RAY_VERSION = "2.33" +TEST_CLUSTER_NAME = "test-cluster-name" + +VERTEX_AI_RAY_OP_PATH = "airflow.providers.google.cloud.operators.vertex_ai.ray.{}" + + +class TestCreateRayClusterOperator: + @mock.patch(VERTEX_AI_RAY_OP_PATH.format("RayHook")) + def test_create_ray_cluster_with_explicit_head_node_type(self, mock_hook_cls): + explicit_head = Resources() + op = CreateRayClusterOperator( + task_id="test-task", + project_id=TEST_PROJECT_ID, + location=TEST_LOCATION, + python_version=TEST_PYTHON_VERSION, + ray_version=TEST_RAY_VERSION, + head_node_type=explicit_head, + gcp_conn_id=TEST_GCP_CONN_ID, + ) + assert op.head_node_type is explicit_head + + @mock.patch(VERTEX_AI_RAY_OP_PATH.format("RayHook")) + def test_create_ray_cluster_default_head_node_type_is_fresh_resources(self, mock_hook_cls): + op1 = CreateRayClusterOperator( + task_id="test-task-1", + project_id=TEST_PROJECT_ID, + location=TEST_LOCATION, + python_version=TEST_PYTHON_VERSION, + ray_version=TEST_RAY_VERSION, + gcp_conn_id=TEST_GCP_CONN_ID, + ) + op2 = CreateRayClusterOperator( + task_id="test-task-2", + project_id=TEST_PROJECT_ID, + location=TEST_LOCATION, + python_version=TEST_PYTHON_VERSION, + ray_version=TEST_RAY_VERSION, + gcp_conn_id=TEST_GCP_CONN_ID, + ) + assert isinstance(op1.head_node_type, Resources) + assert isinstance(op2.head_node_type, Resources) + assert op1.head_node_type is not op2.head_node_type + + @mock.patch(VERTEX_AI_RAY_OP_PATH.format("VertexAIRayClusterLink")) + @mock.patch(VERTEX_AI_RAY_OP_PATH.format("RayHook")) + def test_execute_without_head_node_type_passes_default_resources(self, mock_hook_cls, mock_link): + mock_hook = mock_hook_cls.return_value + mock_hook.create_ray_cluster.return_value = ( + f"projects/{TEST_PROJECT_ID}/locations/{TEST_LOCATION}/persistentResources/{TEST_CLUSTER_NAME}" + ) + mock_hook.extract_cluster_id.return_value = TEST_CLUSTER_NAME + + op = CreateRayClusterOperator( + task_id="test-task", + project_id=TEST_PROJECT_ID, + location=TEST_LOCATION, + python_version=TEST_PYTHON_VERSION, + ray_version=TEST_RAY_VERSION, + gcp_conn_id=TEST_GCP_CONN_ID, + ) + + ti_mock = mock.MagicMock() + context = {"ti": ti_mock, "task": mock.MagicMock()} + op.execute(context=context) + + call_kwargs = mock_hook.create_ray_cluster.call_args.kwargs + assert isinstance(call_kwargs["head_node_type"], Resources) From d011b8386bb0fdc1dacc4d3bf9bf43fc7d143262 Mon Sep 17 00:00:00 2001 From: Shahar Epstein <60007259+shahar1@users.noreply.github.com> Date: Fri, 15 May 2026 19:15:47 +0300 Subject: [PATCH 4/9] Address Wei Lee review comments - parameters.py: compute Query() before inner() instead of suppressing B008 with noqa - cleaner and avoids the linter exception - s3_to_gcs.py: revert deferrable=None pattern; restore the canonical conf.getboolean() default enforced by check_deferrable_default checker; add type annotation so B008 doesn't flag the unannotated call --- .../airflow/api_fastapi/common/parameters.py | 14 +++---- .../google/cloud/transfers/s3_to_gcs.py | 8 +--- .../google/cloud/transfers/test_s3_to_gcs.py | 41 ------------------- 3 files changed, 9 insertions(+), 54 deletions(-) diff --git a/airflow-core/src/airflow/api_fastapi/common/parameters.py b/airflow-core/src/airflow/api_fastapi/common/parameters.py index 7aa2eb222c64a..dfd0cd473cbf6 100644 --- a/airflow-core/src/airflow/api_fastapi/common/parameters.py +++ b/airflow-core/src/airflow/api_fastapi/common/parameters.py @@ -609,13 +609,13 @@ def dynamic_depends(self, default: str | Sequence[str] | None = None) -> Callabl else: default_list = list(default) - def inner( - order_by: list[str] = Query( - default=default_list, - description=f"Attributes to order by, multi criteria sort is supported. Prefix with `-` for descending order. " - f"Supported attributes: `{', '.join(all_attrs) if all_attrs else self.get_primary_key_string()}`", # noqa: B008 - ), - ) -> SortParam: + _order_by_query = Query( + default=default_list, + description=f"Attributes to order by, multi criteria sort is supported. Prefix with `-` for descending order. " + f"Supported attributes: `{', '.join(all_attrs) if all_attrs else self.get_primary_key_string()}`", + ) + + def inner(order_by: list[str] = _order_by_query) -> SortParam: return self.set_value(order_by) return inner diff --git a/providers/google/src/airflow/providers/google/cloud/transfers/s3_to_gcs.py b/providers/google/src/airflow/providers/google/cloud/transfers/s3_to_gcs.py index 3e6c1e7837a19..b7df46631ddd3 100644 --- a/providers/google/src/airflow/providers/google/cloud/transfers/s3_to_gcs.py +++ b/providers/google/src/airflow/providers/google/cloud/transfers/s3_to_gcs.py @@ -165,7 +165,7 @@ def __init__( replace=False, gzip=False, google_impersonation_chain: str | Sequence[str] | None = None, - deferrable=None, + deferrable: bool = conf.getboolean("operators", "default_deferrable", fallback=False), poll_interval: int = 10, return_gcs_uris: bool = False, **kwargs, @@ -178,11 +178,7 @@ def __init__( self.verify = verify self.gzip = gzip self.google_impersonation_chain = google_impersonation_chain - self.deferrable = ( - deferrable - if deferrable is not None - else conf.getboolean("operators", "default_deferrable", fallback=False) - ) + self.deferrable = deferrable if poll_interval <= 0: raise ValueError("Invalid value for poll_interval. Expected value greater than 0") self.poll_interval = poll_interval diff --git a/providers/google/tests/unit/google/cloud/transfers/test_s3_to_gcs.py b/providers/google/tests/unit/google/cloud/transfers/test_s3_to_gcs.py index 389356ef1de6f..19b3c1117289b 100644 --- a/providers/google/tests/unit/google/cloud/transfers/test_s3_to_gcs.py +++ b/providers/google/tests/unit/google/cloud/transfers/test_s3_to_gcs.py @@ -32,8 +32,6 @@ ) from airflow.utils.timezone import utcnow -from tests_common.test_utils.config import conf_vars - PROJECT_ID = "test-project-id" TASK_ID = "test-s3-gcs-operator" S3_BUCKET = "test-bucket" @@ -105,45 +103,6 @@ def test_init(self): assert operator.poll_interval == POLL_INTERVAL assert operator.return_gcs_uris is True - def test_deferrable_default_read_from_conf_at_instantiation(self): - """When ``deferrable`` is omitted, the operator should read - ``operators.default_deferrable`` at instantiation time (not at module import). - """ - with conf_vars({("operators", "default_deferrable"): "True"}): - operator = S3ToGCSOperator( - task_id=TASK_ID, - bucket=S3_BUCKET, - prefix=S3_PREFIX, - dest_gcs=GCS_PATH_PREFIX, - return_gcs_uris=True, - ) - assert operator.deferrable is True - - with conf_vars({("operators", "default_deferrable"): "False"}): - operator = S3ToGCSOperator( - task_id=TASK_ID + "_2", - bucket=S3_BUCKET, - prefix=S3_PREFIX, - dest_gcs=GCS_PATH_PREFIX, - return_gcs_uris=True, - ) - assert operator.deferrable is False - - def test_deferrable_explicit_false_overrides_truthy_conf(self): - """An explicit ``deferrable=False`` must override a truthy - ``operators.default_deferrable`` config — only ``None`` falls back to config. - """ - with conf_vars({("operators", "default_deferrable"): "True"}): - operator = S3ToGCSOperator( - task_id=TASK_ID, - bucket=S3_BUCKET, - prefix=S3_PREFIX, - dest_gcs=GCS_PATH_PREFIX, - deferrable=False, - return_gcs_uris=True, - ) - assert operator.deferrable is False - @mock.patch("airflow.providers.google.cloud.transfers.s3_to_gcs.S3Hook") @mock.patch("airflow.providers.google.cloud.transfers.s3_to_gcs.GCSHook") def test_execute(self, gcs_mock_hook, s3_mock_hook): From 5b855efff0bdea9b46027a9123d9d456bad8785b Mon Sep 17 00:00:00 2001 From: Shahar Epstein <60007259+shahar1@users.noreply.github.com> Date: Sat, 16 May 2026 15:43:22 +0300 Subject: [PATCH 5/9] Update providers/google/src/airflow/providers/google/cloud/hooks/vertex_ai/ray.py Co-authored-by: Jens Scheffler <95105677+jscheffl@users.noreply.github.com> --- .../src/airflow/providers/google/cloud/hooks/vertex_ai/ray.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/providers/google/src/airflow/providers/google/cloud/hooks/vertex_ai/ray.py b/providers/google/src/airflow/providers/google/cloud/hooks/vertex_ai/ray.py index 6328c0f33febc..9d85480414530 100644 --- a/providers/google/src/airflow/providers/google/cloud/hooks/vertex_ai/ray.py +++ b/providers/google/src/airflow/providers/google/cloud/hooks/vertex_ai/ray.py @@ -115,7 +115,7 @@ def create_ray_cluster( """ aiplatform.init(project=project_id, location=location, credentials=self.get_credentials()) cluster_path = vertex_ray.create_ray_cluster( - head_node_type=head_node_type if head_node_type is not None else resources.Resources(), + head_node_type=head_node_type or resources.Resources(), python_version=python_version, ray_version=ray_version, network=network, From 62c2be0b478944c3f9243e08bd166b919b3fe31c Mon Sep 17 00:00:00 2001 From: Shahar Epstein <60007259+shahar1@users.noreply.github.com> Date: Sat, 16 May 2026 15:51:34 +0300 Subject: [PATCH 6/9] Update providers/openlineage/tests/system/openlineage/operator.py Co-authored-by: Jens Scheffler <95105677+jscheffl@users.noreply.github.com> --- providers/openlineage/tests/system/openlineage/operator.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/providers/openlineage/tests/system/openlineage/operator.py b/providers/openlineage/tests/system/openlineage/operator.py index b240b1602a9fe..51ec3fd67ece0 100644 --- a/providers/openlineage/tests/system/openlineage/operator.py +++ b/providers/openlineage/tests/system/openlineage/operator.py @@ -205,7 +205,7 @@ def __init__( super().__init__(**kwargs) self.event_templates = event_templates self.file_path = file_path - self.env = env if env is not None else setup_jinja() + self.env = env or setup_jinja() self.allow_duplicate_events_regex = allow_duplicate_events_regex self.clear_variables = clear_variables if self.event_templates and self.file_path: From 4cc68574bd15000d1a743b6ffe6a9001de9c08ba Mon Sep 17 00:00:00 2001 From: Shahar Epstein <60007259+shahar1@users.noreply.github.com> Date: Sat, 16 May 2026 15:51:46 +0300 Subject: [PATCH 7/9] Update shared/observability/src/airflow_shared/observability/metrics/datadog_logger.py Co-authored-by: Jens Scheffler <95105677+jscheffl@users.noreply.github.com> --- .../airflow_shared/observability/metrics/datadog_logger.py | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) diff --git a/shared/observability/src/airflow_shared/observability/metrics/datadog_logger.py b/shared/observability/src/airflow_shared/observability/metrics/datadog_logger.py index a018226a3c8c1..b4a8d0e075f7c 100644 --- a/shared/observability/src/airflow_shared/observability/metrics/datadog_logger.py +++ b/shared/observability/src/airflow_shared/observability/metrics/datadog_logger.py @@ -52,9 +52,7 @@ def __init__( statsd_influxdb_enabled: bool = False, ) -> None: self.dogstatsd = dogstatsd_client - self.metrics_validator = ( - metrics_validator if metrics_validator is not None else PatternAllowListValidator() - ) + self.metrics_validator = metrics_validator or PatternAllowListValidator() self.metrics_tags = metrics_tags self.metric_tags_validator = ( metric_tags_validator if metric_tags_validator is not None else PatternAllowListValidator() From 58589034b74713576b300cf3f761dd30967ae7e7 Mon Sep 17 00:00:00 2001 From: Shahar Epstein <60007259+shahar1@users.noreply.github.com> Date: Sat, 16 May 2026 18:04:15 +0300 Subject: [PATCH 8/9] Update providers/google/src/airflow/providers/google/cloud/operators/vertex_ai/ray.py Co-authored-by: Jens Scheffler <95105677+jscheffl@users.noreply.github.com> --- .../airflow/providers/google/cloud/operators/vertex_ai/ray.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/providers/google/src/airflow/providers/google/cloud/operators/vertex_ai/ray.py b/providers/google/src/airflow/providers/google/cloud/operators/vertex_ai/ray.py index 103d6d60dabba..0a90637f1037c 100644 --- a/providers/google/src/airflow/providers/google/cloud/operators/vertex_ai/ray.py +++ b/providers/google/src/airflow/providers/google/cloud/operators/vertex_ai/ray.py @@ -155,7 +155,7 @@ def __init__( **kwargs, ) -> None: super().__init__(*args, **kwargs) - self.head_node_type = head_node_type if head_node_type is not None else resources.Resources() + self.head_node_type = head_node_type or resources.Resources() self.python_version = python_version self.ray_version = ray_version self.network = network From bec9beaba4a16dfe09b2ecd881923f515c159b9c Mon Sep 17 00:00:00 2001 From: Shahar Epstein <60007259+shahar1@users.noreply.github.com> Date: Sat, 16 May 2026 18:10:55 +0300 Subject: [PATCH 9/9] Use `or` instead of `is not None else` for validator defaults MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Address Jens's review comment on PR #66979 — the ListValidator instances assigned to `metrics_validator` / `metric_tags_validator` are never falsy except when `None`, so the shorter `or` form is equivalent and easier to read. --- .../observability/metrics/datadog_logger.py | 4 +--- .../airflow_shared/observability/metrics/otel_logger.py | 4 +--- .../airflow_shared/observability/metrics/statsd_logger.py | 8 ++------ 3 files changed, 4 insertions(+), 12 deletions(-) diff --git a/shared/observability/src/airflow_shared/observability/metrics/datadog_logger.py b/shared/observability/src/airflow_shared/observability/metrics/datadog_logger.py index b4a8d0e075f7c..129354c3a4381 100644 --- a/shared/observability/src/airflow_shared/observability/metrics/datadog_logger.py +++ b/shared/observability/src/airflow_shared/observability/metrics/datadog_logger.py @@ -54,9 +54,7 @@ def __init__( self.dogstatsd = dogstatsd_client self.metrics_validator = metrics_validator or PatternAllowListValidator() self.metrics_tags = metrics_tags - self.metric_tags_validator = ( - metric_tags_validator if metric_tags_validator is not None else PatternAllowListValidator() - ) + self.metric_tags_validator = metric_tags_validator or PatternAllowListValidator() self.stat_name_handler = stat_name_handler self.statsd_influxdb_enabled = statsd_influxdb_enabled diff --git a/shared/observability/src/airflow_shared/observability/metrics/otel_logger.py b/shared/observability/src/airflow_shared/observability/metrics/otel_logger.py index 95c62617d8444..8d25b23372a10 100644 --- a/shared/observability/src/airflow_shared/observability/metrics/otel_logger.py +++ b/shared/observability/src/airflow_shared/observability/metrics/otel_logger.py @@ -181,9 +181,7 @@ def __init__( ): self.otel: Callable = otel_provider self.prefix: str = prefix - self.metrics_validator = ( - metrics_validator if metrics_validator is not None else PatternAllowListValidator() - ) + self.metrics_validator = metrics_validator or PatternAllowListValidator() self.meter = otel_provider.get_meter(__name__) self.metrics_map = MetricsMap(self.meter) self.stat_name_handler = stat_name_handler diff --git a/shared/observability/src/airflow_shared/observability/metrics/statsd_logger.py b/shared/observability/src/airflow_shared/observability/metrics/statsd_logger.py index 6fa3828f11e19..3500a04dc7be3 100644 --- a/shared/observability/src/airflow_shared/observability/metrics/statsd_logger.py +++ b/shared/observability/src/airflow_shared/observability/metrics/statsd_logger.py @@ -74,13 +74,9 @@ def __init__( statsd_influxdb_enabled: bool = False, ) -> None: self.statsd = statsd_client - self.metrics_validator = ( - metrics_validator if metrics_validator is not None else PatternAllowListValidator() - ) + self.metrics_validator = metrics_validator or PatternAllowListValidator() self.influxdb_tags_enabled = influxdb_tags_enabled - self.metric_tags_validator = ( - metric_tags_validator if metric_tags_validator is not None else PatternAllowListValidator() - ) + self.metric_tags_validator = metric_tags_validator or PatternAllowListValidator() self.stat_name_handler = stat_name_handler self.statsd_influxdb_enabled = statsd_influxdb_enabled