From 2be950efeb39af54da2641b8c5c683f4a8740580 Mon Sep 17 00:00:00 2001 From: Hussein Awala Date: Tue, 28 Feb 2023 22:18:44 +0100 Subject: [PATCH 1/5] Add a check for not templateable fields --- airflow/serialization/serialized_objects.py | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/airflow/serialization/serialized_objects.py b/airflow/serialization/serialized_objects.py index d5b823a0a2946..5a89521e51eef 100644 --- a/airflow/serialization/serialized_objects.py +++ b/airflow/serialization/serialized_objects.py @@ -702,6 +702,7 @@ def __init__(self, *args, **kwargs): def task_type(self) -> str: # Overwrites task_type of BaseOperator to use _task_type instead of # __class__.__name__. + return self._task_type @task_type.setter @@ -770,8 +771,12 @@ def _serialize_node(cls, op: BaseOperator | MappedOperator, include_deps: bool) # Store all template_fields as they are if there are JSON Serializable # If not, store them as strings + # And raise an exception if the field is not templateable + forbidden_fields = set(BaseOperator(task_id="base_task").__dict__.keys()) if op.template_fields: for template_field in op.template_fields: + if template_field in forbidden_fields: + raise AirflowException(f"Cannot use {template_field} as template field") value = getattr(op, template_field, None) if not cls._is_excluded(value, template_field, op): serialize_op[template_field] = serialize_template_field(value) From 694c0edde27b34621410282815949a3339970787 Mon Sep 17 00:00:00 2001 From: Hussein Awala Date: Wed, 1 Mar 2023 00:26:49 +0100 Subject: [PATCH 2/5] switch to __init__ func parameters to avoid creating a new ti in the dag --- airflow/serialization/serialized_objects.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/airflow/serialization/serialized_objects.py b/airflow/serialization/serialized_objects.py index 5a89521e51eef..80f594ebabe0a 100644 --- a/airflow/serialization/serialized_objects.py +++ b/airflow/serialization/serialized_objects.py @@ -20,6 +20,7 @@ import collections.abc import datetime import enum +import inspect import logging import warnings import weakref @@ -772,7 +773,7 @@ def _serialize_node(cls, op: BaseOperator | MappedOperator, include_deps: bool) # Store all template_fields as they are if there are JSON Serializable # If not, store them as strings # And raise an exception if the field is not templateable - forbidden_fields = set(BaseOperator(task_id="base_task").__dict__.keys()) + forbidden_fields = set(inspect.signature(BaseOperator.__init__).parameters.keys()) if op.template_fields: for template_field in op.template_fields: if template_field in forbidden_fields: From af412a5ad2c2386f1272e87ef26ba2e78bd477cd Mon Sep 17 00:00:00 2001 From: Hussein Awala Date: Wed, 1 Mar 2023 00:40:48 +0100 Subject: [PATCH 3/5] add a unit test --- tests/serialization/test_dag_serialization.py | 18 +++++++++++++++++- 1 file changed, 17 insertions(+), 1 deletion(-) diff --git a/tests/serialization/test_dag_serialization.py b/tests/serialization/test_dag_serialization.py index 640c173081183..d7151d41ddd86 100644 --- a/tests/serialization/test_dag_serialization.py +++ b/tests/serialization/test_dag_serialization.py @@ -38,7 +38,7 @@ import airflow from airflow.datasets import Dataset -from airflow.exceptions import SerializationError +from airflow.exceptions import AirflowException, SerializationError from airflow.hooks.base import BaseHook from airflow.kubernetes.pod_generator import PodGenerator from airflow.models import DAG, Connection, DagBag, Operator @@ -1934,6 +1934,22 @@ def test_params_serialize_default(self): assert param.description == "hello" assert param.schema == {"type": "string"} + def test_not_templateable_fields_in_serialized_dag( + self, + ): + """ + Test that when we use not templateable fields, an Airflow exception is raised. + """ + + class TestOperator(BaseOperator): + template_fields = ("execution_timeout",) + + dag = DAG("test_not_templateable_fields", start_date=datetime(2019, 8, 1)) + with dag: + TestOperator(task_id="test", execution_timeout=timedelta(seconds=10)) + with pytest.raises(AirflowException, match="Cannot use execution_timeout as template field"): + SerializedDAG.to_dict(dag) + def test_kubernetes_optional(): """Serialisation / deserialisation continues to work without kubernetes installed""" From 7c0f660b8a822387c1128d8b5a7c0ce5d55c6bc8 Mon Sep 17 00:00:00 2001 From: Hussein Awala Date: Wed, 1 Mar 2023 21:53:30 +0100 Subject: [PATCH 4/5] update exception message --- airflow/serialization/serialized_objects.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/airflow/serialization/serialized_objects.py b/airflow/serialization/serialized_objects.py index 4788afa74e18f..7200094547f0a 100644 --- a/airflow/serialization/serialized_objects.py +++ b/airflow/serialization/serialized_objects.py @@ -777,7 +777,7 @@ def _serialize_node(cls, op: BaseOperator | MappedOperator, include_deps: bool) if op.template_fields: for template_field in op.template_fields: if template_field in forbidden_fields: - raise AirflowException(f"Cannot use {template_field} as template field") + raise AirflowException(f"Cannot template BaseOperator fields: {template_field}") value = getattr(op, template_field, None) if not cls._is_excluded(value, template_field, op): serialize_op[template_field] = serialize_template_field(value) From 931259dc9dbcb8b7a247966cdf53fc5d2295dedd Mon Sep 17 00:00:00 2001 From: Hussein Awala Date: Wed, 1 Mar 2023 23:00:35 +0100 Subject: [PATCH 5/5] update test after updating the exception message --- tests/serialization/test_dag_serialization.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/serialization/test_dag_serialization.py b/tests/serialization/test_dag_serialization.py index d5ce8a1413e10..967c433bdc160 100644 --- a/tests/serialization/test_dag_serialization.py +++ b/tests/serialization/test_dag_serialization.py @@ -2029,7 +2029,7 @@ class TestOperator(BaseOperator): dag = DAG("test_not_templateable_fields", start_date=datetime(2019, 8, 1)) with dag: TestOperator(task_id="test", execution_timeout=timedelta(seconds=10)) - with pytest.raises(AirflowException, match="Cannot use execution_timeout as template field"): + with pytest.raises(AirflowException, match="Cannot template BaseOperator fields: execution_timeout"): SerializedDAG.to_dict(dag)