This is an automated email from the ASF dual-hosted git repository.

shahar1 pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/airflow.git


The following commit(s) were added to refs/heads/main by this push:
     new 83968fb101c Allow argument-provision checks of templated fields in 
operator __init__ (#70505)
83968fb101c is described below

commit 83968fb101c7cab9ce25e3fc975c0df8aac84a68
Author: Shahar Epstein <[email protected]>
AuthorDate: Tue Jul 28 09:08:57 2026 +0300

    Allow argument-provision checks of templated fields in operator __init__ 
(#70505)
    
    The rendering pipeline is one-way with respect to None: an unset field
    stays unset, but a field that was passed can render to None. A check that
    only asks which arguments the author supplied therefore has to run in the
    constructor, and the hook now permits it while still rejecting every read
    of the un-rendered value.
---
 airflow-core/docs/howto/custom-operator.rst        |  16 +++
 contributing-docs/05_pull_requests.rst             |  30 ++++++
 providers/.pre-commit-config.yaml                  |   4 +-
 scripts/ci/prek/validate_operators_init.py         |  81 +++++++++++---
 .../ci/prek/validate_operators_init_exemptions.txt |   8 +-
 .../tests/ci/prek/test_validate_operators_init.py  | 116 +++++++++++++++++++--
 6 files changed, 223 insertions(+), 32 deletions(-)

diff --git a/airflow-core/docs/howto/custom-operator.rst 
b/airflow-core/docs/howto/custom-operator.rst
index 858b1154c41..b5e14128ca0 100644
--- a/airflow-core/docs/howto/custom-operator.rst
+++ b/airflow-core/docs/howto/custom-operator.rst
@@ -336,6 +336,22 @@ Therefore, the following example is invalid:
             def __init__(self, foo) -> None:
                 self.foo = foo.lower()  # assignment should be only self.foo = 
foo
 
+4. Checking whether an argument was *passed* is allowed in the constructor — 
``execute()`` cannot
+tell a supplied field from a missing one, because a field can be ``None`` 
after rendering. Write it
+as ``is None`` / ``is not None``, never as a truthiness test. Anything that 
inspects the *value*
+still belongs in ``execute()``:
+
+.. code-block:: python
+
+        class HelloOperator(BaseOperator):
+            template_fields = ("foo", "bar")
+
+            def __init__(self, foo=None, bar=None) -> None:
+                if foo is None and bar is None:  # allowed: asks what was 
passed
+                    raise ValueError("Either 'foo' or 'bar' must be provided")
+                self.foo = foo
+                self.bar = bar
+
 When an operator inherits from a base operator and does not have a constructor 
defined on its own, the limitations above
 do not apply. However, the templated fields must be set properly in the parent 
according to those limitations.
 
diff --git a/contributing-docs/05_pull_requests.rst 
b/contributing-docs/05_pull_requests.rst
index 2b3c692283b..7ba8e8d5e0e 100644
--- a/contributing-docs/05_pull_requests.rst
+++ b/contributing-docs/05_pull_requests.rst
@@ -434,6 +434,36 @@ In such cases we can usually do something like this
         my_field = my_field or []
         self.my_field = my_field
 
+For such a mutable default the ``if my_field is None: my_field = []`` spelling 
is equally fine.
+
+Checks that only ask **whether an argument was passed** are the other 
exception, and they belong in
+the constructor — it is the only place that can answer the question. With
+``render_template_as_native_obj=True`` a field that *was* passed can render to
+``None``, so in ``execute`` a supplied argument is indistinguishable from a 
missing one. Write the
+check as ``is None`` / ``is not None``, never as a truthiness test — a 
supplied argument can itself
+be ``""``, ``0`` or an empty list:
+
+.. code-block:: python
+
+    from airflow.utils.helpers import exactly_one
+
+
+    def __init__(self, *, command: str | None = None, powershell: str | None = 
None, **kwargs):
+        if not exactly_one(command is not None, powershell is not None):
+            raise ValueError("Must provide exactly one of 'command' or 
'powershell'")
+        super().__init__(**kwargs)
+        self.command = command
+        self.powershell = powershell
+
+``self.field is None`` works the same way, after the assignment. Anything that 
looks at the *value*
+still has to move to ``execute``, and the operand has to be the field itself — 
indirection such as
+``all(v is None for v in [a, b])`` is still flagged.
+
+Converting an existing truthiness check is not neutral: a guard that raises 
when two arguments are
+*both* set only gets stricter, but ``if not field: raise`` gets looser — 
``""`` and ``0`` start
+passing, and that half is a value check. ``not exactly_one(a, b)`` carries 
both, since it also
+requires at least one.
+
 The reason for doing it is that we are working on a cleaning up our code to 
have
 `prek hook <../scripts/ci/prek/validate_operators_init.py>`_
 that will make sure all the cases where logic (such as validation and complex 
conversion)
diff --git a/providers/.pre-commit-config.yaml 
b/providers/.pre-commit-config.yaml
index ab6fef5a7f0..27d886f162a 100644
--- a/providers/.pre-commit-config.yaml
+++ b/providers/.pre-commit-config.yaml
@@ -52,8 +52,8 @@ repos:
         always_run: true
         pass_filenames: false
       - id: validate-operators-init
-        name: No templated field logic checks in operator __init__
-        description: Prevent templated field logic checks in operators' 
__init__
+        name: Validate templated fields in operator __init__
+        description: Require plain assignment of templated fields and forbid 
reading their un-rendered values
         language: python
         entry: ../scripts/ci/prek/validate_operators_init.py
         pass_filenames: true
diff --git a/scripts/ci/prek/validate_operators_init.py 
b/scripts/ci/prek/validate_operators_init.py
index b8893700c90..664d3e4e471 100755
--- a/scripts/ci/prek/validate_operators_init.py
+++ b/scripts/ci/prek/validate_operators_init.py
@@ -25,6 +25,7 @@ from __future__ import annotations
 
 import ast
 import sys
+from collections.abc import Iterator
 from pathlib import Path
 from typing import Any
 
@@ -179,7 +180,7 @@ def _handle_constructor_statement(
     if isinstance(ctor_stmt, ast.Assign):
         if isinstance(ctor_stmt.targets[0], ast.Attribute):
             for target in ctor_stmt.targets:
-                if isinstance(target, ast.Attribute) and target.attr in 
template_fields:
+                if isinstance(target, ast.Attribute) and _target_name(target) 
in template_fields:
                     if isinstance(ctor_stmt.value, ast.IfExp) and 
_is_value_preserving_ternary(
                         ctor_stmt.value, target.attr
                     ):
@@ -196,10 +197,10 @@ def _handle_constructor_statement(
                         )
         elif isinstance(ctor_stmt.targets[0], ast.Tuple) and 
isinstance(ctor_stmt.value, ast.Tuple):
             for target, value in zip(ctor_stmt.targets[0].elts, 
ctor_stmt.value.elts):
-                if isinstance(target, ast.Attribute):
+                if isinstance(target, ast.Attribute) and _target_name(target) 
in template_fields:
                     _handle_assigned_field(assigned_template_fields, 
invalid_assignments, target, value)
     elif isinstance(ctor_stmt, ast.AnnAssign):
-        if isinstance(ctor_stmt.target, ast.Attribute) and 
ctor_stmt.target.attr in template_fields:
+        if isinstance(ctor_stmt.target, ast.Attribute) and 
_target_name(ctor_stmt.target) in template_fields:
             _handle_assigned_field(
                 assigned_template_fields, invalid_assignments, 
ctor_stmt.target, ctor_stmt.value
             )
@@ -225,11 +226,11 @@ def _handle_assigned_field(
 
 def _target_name(target: ast.expr) -> str | None:
     """
-    Resolve an assignment target to the field name it binds.
+    Resolve an assignment target — or a comparison operand — to the field name 
it refers to.
 
-    :param target: The assignment target node.
-    :return: The attribute name for ``self.<name>`` targets, the identifier 
for bare-name
-        targets, or None for anything else.
+    :param target: The node to resolve.
+    :return: The attribute name for ``self.<name>`` nodes, the identifier for 
bare names,
+        or None for anything else.
     """
     if isinstance(target, ast.Attribute) and isinstance(target.value, 
ast.Name) and target.value.id == "self":
         return target.attr
@@ -254,6 +255,22 @@ def _is_super_init_call(node: ast.Call) -> bool:
     )
 
 
+def _is_none_check(node: ast.expr) -> bool:
+    """
+    Check whether an expression is an ``x is None`` / ``x is not None`` 
comparison.
+
+    :param node: The expression node to check.
+    :return: True if the node compares a single operand against the ``None`` 
literal by identity.
+    """
+    return (
+        isinstance(node, ast.Compare)
+        and len(node.ops) == 1
+        and isinstance(node.ops[0], (ast.Is, ast.IsNot))
+        and isinstance(node.comparators[0], ast.Constant)
+        and node.comparators[0].value is None
+    )
+
+
 def _is_value_preserving_ternary(value: ast.IfExp, field: str) -> bool:
     """
     Check whether a ternary keeps the field value intact when it is set.
@@ -274,11 +291,7 @@ def _is_value_preserving_ternary(value: ast.IfExp, field: 
str) -> bool:
         isinstance(test, ast.Compare)
         and isinstance(test.left, ast.Name)
         and test.left.id == field
-        and len(test.ops) == 1
-        and isinstance(test.ops[0], (ast.Is, ast.IsNot))
-        and len(test.comparators) == 1
-        and isinstance(test.comparators[0], ast.Constant)
-        and test.comparators[0].value is None
+        and _is_none_check(test)
     )
 
 
@@ -289,12 +302,13 @@ def _collect_sanctioned_uses(ctor: ast.FunctionDef, 
template_fields: list[str])
     Sanctioned patterns are the ones the project documents as safe in a 
constructor:
     ``self.field = field``, ``self.field = field or <default>``, the equivalent
     value-preserving ternaries, the local rebind ``field = field or 
<default>``,
-    tuple assignments pairing names one-to-one, and forwarding via
-    ``super().__init__(field=field)``.
+    tuple assignments pairing names one-to-one, forwarding via
+    ``super().__init__(field=field)``, and ``field is None`` / ``field is not 
None``
+    provision checks.
 
     :param ctor: The constructor function node.
     :param template_fields: The template fields of the class.
-    :return: Set of ``id()``s of Name nodes participating in sanctioned 
patterns.
+    :return: Set of ``id()``s of nodes participating in sanctioned patterns.
     """
     sanctioned: set[int] = set()
 
@@ -325,6 +339,9 @@ def _collect_sanctioned_uses(ctor: ast.FunctionDef, 
template_fields: list[str])
             for keyword in node.keywords:
                 if keyword.arg is not None and keyword.arg in template_fields:
                     mark(keyword.value, keyword.arg)
+        elif isinstance(node, ast.Compare) and _is_none_check(node):
+            # Reads whether the argument was passed, not its value — only 
__init__ can see that.
+            sanctioned.add(id(node.left))
     return sanctioned
 
 
@@ -375,14 +392,20 @@ def _check_constructor_field_logic(
             if node.id in template_fields and node.id in bound_names and 
id(node) not in sanctioned:
                 findings.setdefault(node.lineno, set()).add(node.id)
         elif isinstance(node, ast.Attribute) and isinstance(node.ctx, 
ast.Load):
-            if isinstance(node.value, ast.Name) and node.value.id == "self" 
and node.attr in template_fields:
+            if (
+                isinstance(node.value, ast.Name)
+                and node.value.id == "self"
+                and node.attr in template_fields
+                and id(node) not in sanctioned
+            ):
                 findings.setdefault(node.lineno, 
set()).add(f"self.{node.attr}")
 
     if findings:
         console.print(
             f"{class_node.name}'s constructor applies logic to template 
fields. Template fields "
             f"are rendered after the constructor runs, so validation or 
transformation here acts "
-            f"on the un-rendered Jinja expression and should move to 
execute():"
+            f"on the un-rendered Jinja expression and should move to execute() 
"
+            f"(see contributing-docs/05_pull_requests.rst):"
         )
         for lineno in sorted(findings):
             source = source_lines[lineno - 1].strip() if lineno <= 
len(source_lines) else ""
@@ -392,11 +415,31 @@ def _check_constructor_field_logic(
     return len(findings)
 
 
+def _iter_nested_statements(node: ast.AST) -> Iterator[ast.stmt]:
+    """
+    Yield the statements nested inside ``node``, excluding ``node`` itself.
+
+    Nested scopes are skipped: ``self.<field> = ...`` in a function or class 
defined in the
+    constructor either belongs to another object or runs after rendering.
+
+    :param node: The node to descend into.
+    :return: Iterator over the nested statements.
+    """
+    if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef, ast.ClassDef)):
+        return
+    for child in ast.iter_child_nodes(node):
+        if isinstance(child, ast.stmt):
+            yield child
+        yield from _iter_nested_statements(child)
+
+
 def _check_constructor_template_fields(class_node: ast.ClassDef, 
template_fields: list[str]) -> int:
     """
     This method checks a class's constructor for missing or invalid 
assignments of template fields.
     When there isn't a constructor - it assumes that the template fields are 
defined in the parent's
     constructor correctly.
+    Nested statements can only add invalid assignments: a branch may not run, 
so it cannot satisfy
+    the requirement that the field be assigned.
     TODO: Enhance this function to work with nested inheritance trees through 
dynamic imports.
 
     :param class_node: the AST node representing the class definition
@@ -418,6 +461,9 @@ def _check_constructor_template_fields(class_node: 
ast.ClassDef, template_fields
                 missing_assignments = _handle_constructor_statement(
                     template_fields, ctor_stmt, missing_assignments, 
invalid_assignments
                 )
+                for nested_stmt in _iter_nested_statements(ctor_stmt):
+                    _handle_parent_constructor_kwargs(template_fields, 
nested_stmt, [], invalid_assignments)
+                    _handle_constructor_statement(template_fields, 
nested_stmt, [], invalid_assignments)
 
     if init_flag and missing_assignments:
         count += len(missing_assignments)
@@ -428,6 +474,7 @@ def _check_constructor_template_fields(class_node: 
ast.ClassDef, template_fields
         )
         console.print(f"[red]{missing_assignments}[/red]")
 
+    invalid_assignments = sorted(set(invalid_assignments))
     if invalid_assignments:
         count += len(invalid_assignments)
         console.print(
diff --git a/scripts/ci/prek/validate_operators_init_exemptions.txt 
b/scripts/ci/prek/validate_operators_init_exemptions.txt
index 608ce3bda51..76ba180595a 100644
--- a/scripts/ci/prek/validate_operators_init_exemptions.txt
+++ b/scripts/ci/prek/validate_operators_init_exemptions.txt
@@ -3,10 +3,8 @@
 #
 # Format: <repo-relative-path>::<ClassName>  (one class per line)
 #
-# Fixing a class (moving template-field validation/transformation out of 
__init__ into
-# execute()) MUST remove its entry in the same PR — the hook fails on stale 
entries.
+# Fixing a class MUST remove its entry in the same PR — the hook fails on 
stale entries.
 # Burn-down tracked at https://github.com/apache/airflow/issues/70296
-providers/amazon/src/airflow/providers/amazon/aws/operators/emr.py::EmrAddStepsOperator
 
providers/amazon/src/airflow/providers/amazon/aws/operators/neptune.py::NeptuneStartDbClusterOperator
 
providers/amazon/src/airflow/providers/amazon/aws/operators/neptune.py::NeptuneStopDbClusterOperator
 
providers/amazon/src/airflow/providers/amazon/aws/operators/s3.py::S3DeleteObjectsOperator
@@ -21,7 +19,6 @@ 
providers/google/src/airflow/providers/google/cloud/operators/cloud_storage_tran
 
providers/google/src/airflow/providers/google/cloud/operators/dataproc.py::DataprocCreateClusterOperator
 
providers/google/src/airflow/providers/google/cloud/operators/dataproc.py::DataprocSubmitJobOperator
 
providers/google/src/airflow/providers/google/cloud/operators/functions.py::CloudFunctionDeployFunctionOperator
-providers/google/src/airflow/providers/google/cloud/operators/gcs.py::GCSDeleteObjectsOperator
 
providers/google/src/airflow/providers/google/cloud/operators/gcs.py::GCSFileTransformOperator
 
providers/google/src/airflow/providers/google/cloud/operators/gcs.py::GCSListObjectsOperator
 
providers/google/src/airflow/providers/google/cloud/sensors/bigquery_dts.py::BigQueryDataTransferServiceTransferRunSensor
@@ -30,10 +27,7 @@ 
providers/google/src/airflow/providers/google/cloud/transfers/azure_fileshare_to
 
providers/google/src/airflow/providers/google/cloud/transfers/bigquery_to_mssql.py::BigQueryToMsSqlOperator
 
providers/google/src/airflow/providers/google/cloud/transfers/gcs_to_bigquery.py::GCSToBigQueryOperator
 
providers/google/src/airflow/providers/google/cloud/transfers/gcs_to_gcs.py::GCSToGCSOperator
-providers/google/src/airflow/providers/google/cloud/transfers/gcs_to_local.py::GCSToLocalFilesystemOperator
 
providers/google/src/airflow/providers/google/marketing_platform/operators/campaign_manager.py::GoogleCampaignManagerDeleteReportOperator
 
providers/microsoft/azure/src/airflow/providers/microsoft/azure/transfers/gcs_to_wasb.py::GCSToAzureBlobStorageOperator
-providers/microsoft/azure/src/airflow/providers/microsoft/azure/transfers/oracle_to_azure_data_lake.py::OracleToAzureDataLakeOperator
 
providers/microsoft/psrp/src/airflow/providers/microsoft/psrp/operators/psrp.py::PsrpOperator
-providers/oracle/src/airflow/providers/oracle/transfers/oracle_to_oracle.py::OracleToOracleOperator
 
providers/standard/src/airflow/providers/standard/operators/bash.py::BashOperator
diff --git a/scripts/tests/ci/prek/test_validate_operators_init.py 
b/scripts/tests/ci/prek/test_validate_operators_init.py
index cc308796e4a..4546a173ca7 100644
--- a/scripts/tests/ci/prek/test_validate_operators_init.py
+++ b/scripts/tests/ci/prek/test_validate_operators_init.py
@@ -66,8 +66,8 @@ class TestConstructorFieldLogic:
             pytest.param("self._validate(foo)\nself.foo = foo", 1, 
id="validation-call"),
             pytest.param(
                 "if foo is not None:\n    self._validate(foo)\nself.foo = foo",
-                2,
-                id="nested-validation-call",
+                1,
+                id="validation-call-behind-provision-guard",
             ),
             pytest.param("self.foo = foo\nself.bar = foo.upper()", 1, 
id="derived-assignment"),
             pytest.param(
@@ -80,6 +80,52 @@ class TestConstructorFieldLogic:
                 2,
                 id="ctor-validation-raise",
             ),
+            pytest.param(
+                "if foo is None:\n    raise ValueError('foo is 
required')\nself.foo = foo",
+                0,
+                id="provision-check-raise",
+            ),
+            pytest.param(
+                "self.foo = foo\nif self.foo is not None:\n    self.bar = 1",
+                0,
+                id="provision-check-on-self-attribute",
+            ),
+            pytest.param(
+                "if foo is None:\n    raise ValueError('required')\n"
+                "if foo.startswith('x'):\n    raise 
ValueError('bad')\nself.foo = foo",
+                1,
+                id="provision-check-alongside-value-read",
+            ),
+            pytest.param(
+                "if foo.get('x') is None:\n    raise 
ValueError('bad')\nself.foo = foo",
+                1,
+                id="none-check-on-derived-value",
+            ),
+            pytest.param(
+                "if foo == None:\n    raise ValueError('required')\nself.foo = 
foo",
+                1,
+                id="equality-check-is-not-a-provision-check",
+            ),
+            pytest.param(
+                "if foo is False:\n    raise ValueError('bad')\nself.foo = 
foo",
+                1,
+                id="identity-check-against-a-non-none-constant",
+            ),
+            pytest.param(
+                "if foo is bar:\n    raise ValueError('bad')\nself.foo = foo",
+                1,
+                id="identity-check-against-a-non-constant",
+            ),
+            pytest.param(
+                "if foo is None is not bar:\n    raise 
ValueError('bad')\nself.foo = foo",
+                1,
+                id="chained-comparison-is-not-a-provision-check",
+            ),
+            pytest.param(
+                "if not exactly_one(foo is None, bar is None):\n    raise 
ValueError('x')\nself.foo = foo",
+                0,
+                id="provision-check-passed-to-a-helper",
+            ),
         ],
     )
     def test_flags_logic_but_not_sanctioned_patterns(self, ctor_body: str, 
expected: int):
@@ -148,15 +194,73 @@ class TestTemplateFieldExtraction:
 
 
 class TestValuePreservingTernaryAssignment:
-    def test_ternary_default_is_a_valid_assignment(self):
-        code = """
+    @pytest.mark.parametrize(
+        "value, expected",
+        [
+            pytest.param("conf if conf else {}", 0, id="value-preserving"),
+            # Neither a valid assignment nor a substitute for one: invalid, 
and still missing.
+            pytest.param('conf if conf == "x" else {}', 2, 
id="value-comparison"),
+        ],
+    )
+    def test_ternary_default_assignment(self, value: str, expected: int):
+        code = f"""
         class Op(BaseOperator):
             template_fields = ("conf",)
 
             def __init__(self, conf=None, **kwargs):
-                self.conf = conf if conf else {}
+                self.conf = {value}
         """
-        assert _check_constructor_template_fields(_first_class(code), 
["conf"]) == 0
+        assert _check_constructor_template_fields(_first_class(code), 
["conf"]) == expected
+
+    @pytest.mark.parametrize(
+        "nested",
+        [
+            pytest.param("if baz:\n        self.foo = derive(bar)", 
id="doubly-nested"),
+            pytest.param("super().__init__(foo=derive(bar), **kwargs)", 
id="super-kwarg"),
+            pytest.param("self.foo = derive(bar)\n    self.foo = derive(baz)", 
id="reported-once"),
+            pytest.param(
+                "try:\n        pass\n    except ValueError:\n        self.foo 
= derive(bar)",
+                id="except-handler",
+            ),
+        ],
+    )
+    def test_derived_assignment_nested_in_a_branch_is_invalid(self, nested: 
str):
+        code = _operator_code(f"self.foo = foo\nif bar:\n    {nested}")
+        assert _check_constructor_template_fields(_first_class(code), ["foo"]) 
== 1
+
+    @pytest.mark.parametrize(
+        "nested",
+        [
+            pytest.param("hook.foo = derive()", id="other-object-attribute"),
+            pytest.param("hook.foo: str = derive()", 
id="other-object-annotated"),
+            pytest.param("cfg.foo, cfg.bar = 1, 2", 
id="other-object-tuple-target"),
+        ],
+    )
+    def test_assignment_to_another_objects_attribute_is_not_the_field(self, 
nested: str):
+        code = _operator_code(f"self.foo = foo\nif bar:\n    {nested}")
+        assert _check_constructor_template_fields(_first_class(code), ["foo"]) 
== 0
+
+    @pytest.mark.parametrize(
+        "ctor_body",
+        [
+            pytest.param("if kwargs:\n    self.foo = foo", id="branch"),
+            pytest.param("if foo is None:\n    super().__init__(foo=foo, 
**kwargs)", id="super-call"),
+        ],
+    )
+    def test_conditional_assignment_does_not_satisfy_the_field(self, 
ctor_body: str):
+        assert 
_check_constructor_template_fields(_first_class(_operator_code(ctor_body)), 
["foo"]) == 1
+
+    @pytest.mark.parametrize(
+        "scope",
+        [
+            pytest.param("def on_kill():\n    self.foo = derive()", 
id="nested-function"),
+            pytest.param("async def on_kill():\n    self.foo = derive()", 
id="nested-async-function"),
+            pytest.param("class Inner:\n    self.foo = derive()", 
id="nested-class"),
+        ],
+    )
+    def test_assignment_in_a_nested_scope_belongs_to_that_scope(self, scope: 
str):
+        code = _operator_code(f"self.foo = foo\n{scope}")
+        assert _check_constructor_template_fields(_first_class(code), ["foo"]) 
== 0
 
 
 class TestExemptions:

Reply via email to