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"}])