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

vatsrahul1001 pushed a commit to branch v3-3-test
in repository https://gitbox.apache.org/repos/asf/airflow.git


The following commit(s) were added to refs/heads/v3-3-test by this push:
     new 5b892a83e94 [v3-3-test] Preserve custom operator defaults in mapped 
tasks (#72589) (#72828)
5b892a83e94 is described below

commit 5b892a83e94e068c72610c1514875ff59c8c381d
Author: github-actions[bot] 
<41898282+github-actions[bot]@users.noreply.github.com>
AuthorDate: Thu Sep 10 16:27:06 2026 +0530

    [v3-3-test] Preserve custom operator defaults in mapped tasks (#72589) 
(#72828)
    
    * Preserve custom operator defaults in mapped tasks
    
    * Preserve mapping defaults for custom operators
    
    * Apply suggestion from @uranusjr
    
    
    
    ---------
    (cherry picked from commit c2c8dce0e47c0dd0105158cec0e9ee7917722f77)
    
    Co-authored-by: Eitan Shalev 
<[email protected]>
    Co-authored-by: Tzu-ping Chung <[email protected]>
    Co-authored-by: Rahul Vats <[email protected]>
---
 task-sdk/src/airflow/sdk/bases/operator.py         | 19 +++++++++++++++++-
 .../task_sdk/definitions/test_mappedoperator.py    | 23 +++++++++++++++++++++-
 2 files changed, 40 insertions(+), 2 deletions(-)

diff --git a/task-sdk/src/airflow/sdk/bases/operator.py 
b/task-sdk/src/airflow/sdk/bases/operator.py
index 97da8869686..940fee84b87 100644
--- a/task-sdk/src/airflow/sdk/bases/operator.py
+++ b/task-sdk/src/airflow/sdk/bases/operator.py
@@ -270,6 +270,21 @@ OPERATOR_DEFAULTS: dict[str, Any] = {
 }
 
 
+def _get_operator_defaults(operator_class: type[BaseOperator]) -> dict[str, 
Any]:
+    """Get default values for an operator class's BaseOperator arguments."""
+    operator_defaults = OPERATOR_DEFAULTS.copy()
+    operator_classes = operator_class.mro()
+    for operator_base in reversed(operator_classes[: 
operator_classes.index(BaseOperator)]):
+        if (init := operator_base.__dict__.get("__init__")) is None:
+            continue
+        operator_defaults.update(
+            (name, parameter.default)
+            for name, parameter in inspect.signature(init).parameters.items()
+            if name in OPERATOR_DEFAULTS and parameter.default is not 
inspect.Parameter.empty
+        )
+    return operator_defaults
+
+
 # This is what handles the actual mapping.
 
 if TYPE_CHECKING:
@@ -372,7 +387,9 @@ else:
         )
 
         # Fill fields not provided by the user with default values.
-        partial_kwargs.update((k, v) for k, v in OPERATOR_DEFAULTS.items() if 
k not in partial_kwargs)
+        partial_kwargs.update(
+            (k, v) for k, v in _get_operator_defaults(operator_class).items() 
if k not in partial_kwargs
+        )
 
         # Post-process arguments. Should be kept in sync with 
_TaskDecorator.expand().
         if "task_concurrency" in kwargs:  # Reject deprecated option.
diff --git a/task-sdk/tests/task_sdk/definitions/test_mappedoperator.py 
b/task-sdk/tests/task_sdk/definitions/test_mappedoperator.py
index 93c5cc19aed..546a39c55ff 100644
--- a/task-sdk/tests/task_sdk/definitions/test_mappedoperator.py
+++ b/task-sdk/tests/task_sdk/definitions/test_mappedoperator.py
@@ -25,7 +25,7 @@ from unittest import mock
 import pendulum
 import pytest
 
-from airflow.sdk import TaskInstanceState, TriggerRule
+from airflow.sdk import ExceptionRetryPolicy, TaskInstanceState, TriggerRule
 from airflow.sdk.bases.operator import BaseOperator
 from airflow.sdk.bases.xcom import BaseXCom
 from airflow.sdk.definitions.dag import DAG
@@ -115,6 +115,27 @@ def test_task_mapping_override_default_args():
     assert mapped.owner == "airflow"
 
 
+def test_mapped_task_preserves_custom_base_operator_default():
+    retry_policy = ExceptionRetryPolicy(rules=[])
+
+    class CustomRetryOperator(BaseOperator):
+        def __init__(self, *, value: str, retry_policy=retry_policy, **kwargs):
+            super().__init__(retry_policy=retry_policy, **kwargs)
+            self.value = value
+
+        def execute(self, context):
+            pass
+
+    direct = CustomRetryOperator(task_id="direct", value="direct")
+    mapped = 
CustomRetryOperator.partial(task_id="mapped").expand(value=["mapped"])
+    unmapped = mapped.unmap({"value": "mapped"})
+
+    assert direct.retry_policy is retry_policy
+    assert mapped.partial_kwargs["retry_policy"] is retry_policy
+    assert mapped.partial_kwargs["inlets"] == []
+    assert unmapped.retry_policy is retry_policy
+
+
 def test_map_unknown_arg_raises():
     with pytest.raises(TypeError, match=r"argument 'file'"):
         BaseOperator.partial(task_id="a").expand(file=[1, 2, {"a": "b"}])

Reply via email to