Skip to content

Commit dae53fb

Browse files
kaxilodaneau-astro
authored andcommitted
Remove callbacks from DAG default_args when serializating it (apache#57397)
1 parent 08469ec commit dae53fb

2 files changed

Lines changed: 79 additions & 0 deletions

File tree

airflow-core/src/airflow/serialization/serialized_objects.py

Lines changed: 17 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2507,6 +2507,23 @@ def serialize_dag(cls, dag: DAG) -> dict:
25072507
serialized_dag["has_on_success_callback"] = True
25082508
if dag.has_on_failure_callback:
25092509
serialized_dag["has_on_failure_callback"] = True
2510+
2511+
# TODO: Move this logic to a better place -- ideally before serializing contents of default_args.
2512+
# There is some duplication with this and SerializedBaseOperator.partial_kwargs serialization.
2513+
# Ideally default_args goes through same logic as fields of SerializedBaseOperator.
2514+
if serialized_dag.get("default_args", {}):
2515+
default_args_dict = serialized_dag["default_args"][Encoding.VAR]
2516+
callbacks_to_remove = []
2517+
for k, v in list(default_args_dict.items()):
2518+
if k in [
2519+
f"on_{x}_callback" for x in ("execute", "failure", "success", "retry", "skipped")
2520+
]:
2521+
if bool(v):
2522+
default_args_dict[f"has_{k}"] = True
2523+
callbacks_to_remove.append(k)
2524+
for k in callbacks_to_remove:
2525+
del default_args_dict[k]
2526+
25102527
return serialized_dag
25112528
except SerializationError:
25122529
raise

airflow-core/tests/unit/serialization/test_dag_serialization.py

Lines changed: 62 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -4291,3 +4291,65 @@ def test_partial_kwargs_end_to_end_deserialization(self):
42914291
assert "owner" in deserialized_task.partial_kwargs
42924292
assert deserialized_task.partial_kwargs["retry_delay"] == timedelta(seconds=600)
42934293
assert deserialized_task.partial_kwargs["owner"] == "custom_owner"
4294+
4295+
4296+
@pytest.mark.parametrize(
4297+
["callbacks", "expected_has_flags", "absent_keys"],
4298+
[
4299+
pytest.param(
4300+
{
4301+
"on_failure_callback": lambda ctx: None,
4302+
"on_success_callback": lambda ctx: None,
4303+
"on_retry_callback": lambda ctx: None,
4304+
},
4305+
["has_on_failure_callback", "has_on_success_callback", "has_on_retry_callback"],
4306+
["on_failure_callback", "on_success_callback", "on_retry_callback"],
4307+
id="multiple_callbacks",
4308+
),
4309+
pytest.param(
4310+
{"on_failure_callback": lambda ctx: None},
4311+
["has_on_failure_callback"],
4312+
["on_failure_callback", "has_on_success_callback", "on_success_callback"],
4313+
id="single_callback",
4314+
),
4315+
pytest.param(
4316+
{"on_failure_callback": lambda ctx: None, "on_execute_callback": None},
4317+
["has_on_failure_callback"],
4318+
["on_failure_callback", "has_on_execute_callback", "on_execute_callback"],
4319+
id="callback_with_none",
4320+
),
4321+
pytest.param(
4322+
{},
4323+
[],
4324+
[
4325+
"has_on_execute_callback",
4326+
"has_on_failure_callback",
4327+
"has_on_success_callback",
4328+
"has_on_retry_callback",
4329+
"has_on_skipped_callback",
4330+
],
4331+
id="no_callbacks",
4332+
),
4333+
],
4334+
)
4335+
def test_dag_default_args_callbacks_serialization(callbacks, expected_has_flags, absent_keys):
4336+
"""Test callbacks in DAG default_args are serialized as boolean flags."""
4337+
default_args = {"owner": "test_owner", "retries": 2, **callbacks}
4338+
4339+
with DAG(dag_id="test_default_args_callbacks", default_args=default_args) as dag:
4340+
BashOperator(task_id="task1", bash_command="echo 1", dag=dag)
4341+
4342+
serialized_dag_dict = SerializedDAG.serialize_dag(dag)
4343+
default_args_dict = serialized_dag_dict["default_args"][Encoding.VAR]
4344+
4345+
for flag in expected_has_flags:
4346+
assert default_args_dict.get(flag) is True
4347+
4348+
for key in absent_keys:
4349+
assert key not in default_args_dict
4350+
4351+
assert default_args_dict["owner"] == "test_owner"
4352+
assert default_args_dict["retries"] == 2
4353+
4354+
deserialized_dag = SerializedDAG.deserialize_dag(serialized_dag_dict)
4355+
assert deserialized_dag.dag_id == "test_default_args_callbacks"

0 commit comments

Comments
 (0)