From a3894b7fbc9e2ac6c555dbb74a9fb3f1defac25c Mon Sep 17 00:00:00 2001 From: TJaniF Date: Wed, 10 Apr 2024 15:37:22 +0200 Subject: [PATCH 01/10] probably a minor crime against python --- airflow/models/taskinstance.py | 39 ++++++++++++++++++++++------------ 1 file changed, 26 insertions(+), 13 deletions(-) diff --git a/airflow/models/taskinstance.py b/airflow/models/taskinstance.py index d52a71c5b2e16..d82fc82e118a5 100644 --- a/airflow/models/taskinstance.py +++ b/airflow/models/taskinstance.py @@ -397,7 +397,7 @@ def _creator_note(val): return TaskInstanceNote(*val) -def _execute_task(task_instance: TaskInstance | TaskInstancePydantic, context: Context, task_orig: Operator): +def _execute_task(task_instance: TaskInstance | TaskInstancePydantic, context: Context, task_orig: Operator, jinja_env=None): """ Execute Task (optionally with a Timeout) and push Xcom results. @@ -433,7 +433,7 @@ def _execute_task(task_instance: TaskInstance | TaskInstancePydantic, context: C if execute_callable.__name__ == "execute": execute_callable_kwargs[f"{task_to_execute.__class__.__name__}__sentinel"] = _sentinel - def _execute_callable(context: Context, **execute_callable_kwargs): + def _execute_callable(context: Context, jinja_env=None, **execute_callable_kwargs): try: # Print a marker for log grouping of details before task execution log.info("::endgroup::") @@ -453,6 +453,16 @@ def _execute_callable(context: Context, **execute_callable_kwargs): # Print a marker post execution for internals of post task processing log.info("::group::Post task execution logs") + # DAG authors define map_index_template at the task level + if jinja_env is not None and (template := context.get("map_index_template")) is not None: + rendered_map_index = task_instance.task.rendered_map_index = jinja_env.from_string(template).render(context) + task_instance.task.log.info("Map index rendered as %s", rendered_map_index) + else: + rendered_map_index = None + + task_instance.rendered_map_index = rendered_map_index + + # If a timeout is specified for the task, make it fail # if it goes beyond if task_to_execute.execution_timeout: @@ -470,12 +480,12 @@ def _execute_callable(context: Context, **execute_callable_kwargs): raise AirflowTaskTimeout() # Run task in timeout wrapper with timeout(timeout_seconds): - result = _execute_callable(context=context, **execute_callable_kwargs) + result = _execute_callable(context=context, **execute_callable_kwargs, jinja_env=jinja_env) except AirflowTaskTimeout: task_to_execute.on_kill() raise else: - result = _execute_callable(context=context, **execute_callable_kwargs) + result = _execute_callable(context=context, **execute_callable_kwargs, jinja_env=jinja_env) cm = nullcontext() if InternalApiConfig.get_use_internal_api() else create_session() with cm as session_or_null: if task_to_execute.do_xcom_push: @@ -501,7 +511,13 @@ def _execute_callable(context: Context, **execute_callable_kwargs): _record_task_map_for_downstreams( task_instance=task_instance, task=task_orig, value=xcom_value, session=session_or_null ) - return result + + print(jinja_env) + print(context.get("map_index_template")) + + + + return result, task_instance.rendered_map_index def _refresh_from_db( @@ -2716,29 +2732,26 @@ def signal_handler(signum, frame): # Execute the task with set_current_context(context): - result = self._execute_task(context, task_orig) + result, rendered_map_index = self._execute_task(context, task_orig, jinja_env=jinja_env) + + self.rendered_map_index = rendered_map_index # Run post_execute callback self.task.post_execute(context=context, result=result) - # DAG authors define map_index_template at the task level - if jinja_env is not None and (template := context.get("map_index_template")) is not None: - rendered_map_index = self.rendered_map_index = jinja_env.from_string(template).render(context) - self.log.info("Map index rendered as %s", rendered_map_index) - Stats.incr(f"operator_successes_{self.task.task_type}", tags=self.stats_tags) # Same metric with tagging Stats.incr("operator_successes", tags={**self.stats_tags, "task_type": self.task.task_type}) Stats.incr("ti_successes", tags=self.stats_tags) - def _execute_task(self, context: Context, task_orig: Operator): + def _execute_task(self, context: Context, task_orig: Operator, jinja_env=None): """ Execute Task (optionally with a Timeout) and push Xcom results. :param context: Jinja2 context :param task_orig: origin task """ - return _execute_task(self, context, task_orig) + return _execute_task(self, context, task_orig, jinja_env) @provide_session def defer_task(self, session: Session, defer: TaskDeferred) -> None: From 6735b7d7fd861b94a6d43d99a70b01bc1d91155d Mon Sep 17 00:00:00 2001 From: TJaniF Date: Wed, 10 Apr 2024 15:58:50 +0200 Subject: [PATCH 02/10] remove print statements --- airflow/models/taskinstance.py | 5 ----- 1 file changed, 5 deletions(-) diff --git a/airflow/models/taskinstance.py b/airflow/models/taskinstance.py index d82fc82e118a5..ca8c05d3fffae 100644 --- a/airflow/models/taskinstance.py +++ b/airflow/models/taskinstance.py @@ -512,11 +512,6 @@ def _execute_callable(context: Context, jinja_env=None, **execute_callable_kwarg task_instance=task_instance, task=task_orig, value=xcom_value, session=session_or_null ) - print(jinja_env) - print(context.get("map_index_template")) - - - return result, task_instance.rendered_map_index From 4673ac5833ddb6945e390ccfcfae9afa348e0664 Mon Sep 17 00:00:00 2001 From: TJaniF Date: Wed, 10 Apr 2024 17:04:12 +0200 Subject: [PATCH 03/10] unit tests and pre-commit --- airflow/models/taskinstance.py | 11 ++++++----- tests/models/test_mappedoperator.py | 26 ++++++++++++++++++++++++++ 2 files changed, 32 insertions(+), 5 deletions(-) diff --git a/airflow/models/taskinstance.py b/airflow/models/taskinstance.py index ca8c05d3fffae..7e35a3d2d1bc9 100644 --- a/airflow/models/taskinstance.py +++ b/airflow/models/taskinstance.py @@ -397,7 +397,9 @@ def _creator_note(val): return TaskInstanceNote(*val) -def _execute_task(task_instance: TaskInstance | TaskInstancePydantic, context: Context, task_orig: Operator, jinja_env=None): +def _execute_task( + task_instance: TaskInstance | TaskInstancePydantic, context: Context, task_orig: Operator, jinja_env=None +): """ Execute Task (optionally with a Timeout) and push Xcom results. @@ -455,14 +457,13 @@ def _execute_callable(context: Context, jinja_env=None, **execute_callable_kwarg # DAG authors define map_index_template at the task level if jinja_env is not None and (template := context.get("map_index_template")) is not None: - rendered_map_index = task_instance.task.rendered_map_index = jinja_env.from_string(template).render(context) - task_instance.task.log.info("Map index rendered as %s", rendered_map_index) - else: + rendered_map_index = jinja_env.from_string(template).render(context) + log.info("Map index rendered as %s", rendered_map_index) + else: rendered_map_index = None task_instance.rendered_map_index = rendered_map_index - # If a timeout is specified for the task, make it fail # if it goes beyond if task_to_execute.execution_timeout: diff --git a/tests/models/test_mappedoperator.py b/tests/models/test_mappedoperator.py index 64304cf3069ab..7a880d475cddb 100644 --- a/tests/models/test_mappedoperator.py +++ b/tests/models/test_mappedoperator.py @@ -635,6 +635,30 @@ def task1(map_name): return task1.expand(map_name=map_names) +def _create_named_map_index_renders_on_failure_classic(*, task_id, map_names, template): + class HasMapName(BaseOperator): + def __init__(self, *, map_name: str, **kwargs): + super().__init__(**kwargs) + self.map_name = map_name + raise AirflowSkipException("Imagine this task failed!") + + return HasMapName.partial(task_id=task_id, map_index_template=template).expand( + map_name=map_names, + ) + + +def _create_named_map_index_renders_on_failure_taskflow(*, task_id, map_names, template): + from airflow.operators.python import get_current_context + + @task(task_id=task_id, map_index_template=template) + def task1(map_name): + context = get_current_context() + context["map_name"] = map_name + raise AirflowSkipException("Imagine this task failed!") + + return task1.expand(map_name=map_names) + + @pytest.mark.parametrize( "template, expected_rendered_names", [ @@ -649,6 +673,8 @@ def task1(map_name): [ pytest.param(_create_mapped_with_name_template_classic, id="classic"), pytest.param(_create_mapped_with_name_template_taskflow, id="taskflow"), + pytest.param(_create_named_map_index_renders_on_failure_classic, id="classic-failure"), + pytest.param(_create_named_map_index_renders_on_failure_taskflow, id="taskflow-failure"), ], ) def test_expand_mapped_task_instance_with_named_index( From bee7cfe0048b0870d2d20838c50084a9552694cf Mon Sep 17 00:00:00 2001 From: TJaniF Date: Thu, 11 Apr 2024 13:23:34 +0200 Subject: [PATCH 04/10] render template regardless of task outcome --- airflow/models/taskinstance.py | 48 ++++++++++++++++++++-------------- 1 file changed, 28 insertions(+), 20 deletions(-) diff --git a/airflow/models/taskinstance.py b/airflow/models/taskinstance.py index 727f4d94bb1f0..eaf8161f0af37 100644 --- a/airflow/models/taskinstance.py +++ b/airflow/models/taskinstance.py @@ -435,7 +435,7 @@ def _execute_task( if execute_callable.__name__ == "execute": execute_callable_kwargs[f"{task_to_execute.__class__.__name__}__sentinel"] = _sentinel - def _execute_callable(context: Context, jinja_env=None, **execute_callable_kwargs): + def _execute_callable(context: Context, **execute_callable_kwargs): try: # Print a marker for log grouping of details before task execution log.info("::endgroup::") @@ -455,15 +455,6 @@ def _execute_callable(context: Context, jinja_env=None, **execute_callable_kwarg # Print a marker post execution for internals of post task processing log.info("::group::Post task execution logs") - # DAG authors define map_index_template at the task level - if jinja_env is not None and (template := context.get("map_index_template")) is not None: - rendered_map_index = jinja_env.from_string(template).render(context) - log.info("Map index rendered as %s", rendered_map_index) - else: - rendered_map_index = None - - task_instance.rendered_map_index = rendered_map_index - # If a timeout is specified for the task, make it fail # if it goes beyond if task_to_execute.execution_timeout: @@ -481,12 +472,12 @@ def _execute_callable(context: Context, jinja_env=None, **execute_callable_kwarg raise AirflowTaskTimeout() # Run task in timeout wrapper with timeout(timeout_seconds): - result = _execute_callable(context=context, **execute_callable_kwargs, jinja_env=jinja_env) + result = _execute_callable(context=context, **execute_callable_kwargs) except AirflowTaskTimeout: task_to_execute.on_kill() raise else: - result = _execute_callable(context=context, **execute_callable_kwargs, jinja_env=jinja_env) + result = _execute_callable(context=context, **execute_callable_kwargs) cm = nullcontext() if InternalApiConfig.get_use_internal_api() else create_session() with cm as session_or_null: if task_to_execute.do_xcom_push: @@ -513,7 +504,7 @@ def _execute_callable(context: Context, jinja_env=None, **execute_callable_kwarg task_instance=task_instance, task=task_orig, value=xcom_value, session=session_or_null ) - return result, task_instance.rendered_map_index + return result def _refresh_from_db( @@ -2653,6 +2644,16 @@ def _register_dataset_changes(self, *, events: DatasetEventAccessors, session: S session=session, ) + def _render_map_index(self, context, jinja_env=None): + # DAG authors define map_index_template at the task level + if jinja_env is not None and (template := context.get("map_index_template")) is not None: + rendered_map_index = jinja_env.from_string(template).render(context) + log.info("Map index rendered as %s", rendered_map_index) + else: + rendered_map_index = None + + return rendered_map_index + def _execute_task_with_callbacks(self, context: Context, test_mode: bool = False, *, session: Session): """Prepare Task for Execution.""" if TYPE_CHECKING: @@ -2725,11 +2726,18 @@ def signal_handler(signum, frame): previous_state=TaskInstanceState.QUEUED, task_instance=self, session=session ) - # Execute the task - with set_current_context(context): - result, rendered_map_index = self._execute_task(context, task_orig, jinja_env=jinja_env) - - self.rendered_map_index = rendered_map_index + try: + # Execute the task + with set_current_context(context): + result = self._execute_task(context, task_orig) + except Exception: + # If the task failed, swallow rendering error so it doesn't mask the main error. + with contextlib.suppress(jinja2.TemplateSyntaxError, jinja2.UndefinedError): + self.rendered_map_index = self._render_map_index(context, jinja_env=jinja_env) + raise + else: + # If the task succeeded, render normally to let rendering error bubble up. + self.rendered_map_index = self._render_map_index(context, jinja_env=jinja_env) # Run post_execute callback self.task.post_execute(context=context, result=result) @@ -2739,14 +2747,14 @@ def signal_handler(signum, frame): Stats.incr("operator_successes", tags={**self.stats_tags, "task_type": self.task.task_type}) Stats.incr("ti_successes", tags=self.stats_tags) - def _execute_task(self, context: Context, task_orig: Operator, jinja_env=None): + def _execute_task(self, context: Context, task_orig: Operator): """ Execute Task (optionally with a Timeout) and push Xcom results. :param context: Jinja2 context :param task_orig: origin task """ - return _execute_task(self, context, task_orig, jinja_env) + return _execute_task(self, context, task_orig) @provide_session def defer_task(self, session: Session, defer: TaskDeferred) -> None: From 63c3da2df6ee6f379af591d5fca5be89b32c5a7e Mon Sep 17 00:00:00 2001 From: TJaniF Date: Thu, 11 Apr 2024 13:31:11 +0200 Subject: [PATCH 05/10] cleanup --- airflow/models/taskinstance.py | 5 +---- 1 file changed, 1 insertion(+), 4 deletions(-) diff --git a/airflow/models/taskinstance.py b/airflow/models/taskinstance.py index eaf8161f0af37..1ddae9116642c 100644 --- a/airflow/models/taskinstance.py +++ b/airflow/models/taskinstance.py @@ -397,9 +397,7 @@ def _creator_note(val): return TaskInstanceNote(*val) -def _execute_task( - task_instance: TaskInstance | TaskInstancePydantic, context: Context, task_orig: Operator, jinja_env=None -): +def _execute_task(task_instance: TaskInstance | TaskInstancePydantic, context: Context, task_orig: Operator): """ Execute Task (optionally with a Timeout) and push Xcom results. @@ -503,7 +501,6 @@ def _execute_callable(context: Context, **execute_callable_kwargs): _record_task_map_for_downstreams( task_instance=task_instance, task=task_orig, value=xcom_value, session=session_or_null ) - return result From e94d4b12d97c09aafe8f769dcdc3f204d48e6bf6 Mon Sep 17 00:00:00 2001 From: TJaniF Date: Thu, 11 Apr 2024 17:24:53 +0200 Subject: [PATCH 06/10] transform indices to strings --- tests/models/test_mappedoperator.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/tests/models/test_mappedoperator.py b/tests/models/test_mappedoperator.py index b49f8cd49308a..427e741749c22 100644 --- a/tests/models/test_mappedoperator.py +++ b/tests/models/test_mappedoperator.py @@ -704,6 +704,8 @@ def test_expand_mapped_task_instance_with_named_index( .order_by(TaskInstance.map_index) ).all() + indices = [str(index) for index in indices] + assert indices == expected_rendered_names From 1cdee6e02b4cef6af6c576ea12088afc9e58bc92 Mon Sep 17 00:00:00 2001 From: TJaniF Date: Thu, 11 Apr 2024 18:04:55 +0200 Subject: [PATCH 07/10] attempt 2 --- tests/models/test_mappedoperator.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/models/test_mappedoperator.py b/tests/models/test_mappedoperator.py index 427e741749c22..37ac574cdc2d8 100644 --- a/tests/models/test_mappedoperator.py +++ b/tests/models/test_mappedoperator.py @@ -704,7 +704,7 @@ def test_expand_mapped_task_instance_with_named_index( .order_by(TaskInstance.map_index) ).all() - indices = [str(index) for index in indices] + indices = [str(index) if index is not None else None for index in indices] assert indices == expected_rendered_names From b45d35de1c8573dbc0cda06f1a4196c87b972d4c Mon Sep 17 00:00:00 2001 From: Tzu-ping Chung Date: Fri, 12 Apr 2024 16:11:34 +0800 Subject: [PATCH 08/10] Properly fail a classic operator --- tests/models/test_mappedoperator.py | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/tests/models/test_mappedoperator.py b/tests/models/test_mappedoperator.py index 37ac574cdc2d8..f2609f0aa2b18 100644 --- a/tests/models/test_mappedoperator.py +++ b/tests/models/test_mappedoperator.py @@ -640,6 +640,9 @@ class HasMapName(BaseOperator): def __init__(self, *, map_name: str, **kwargs): super().__init__(**kwargs) self.map_name = map_name + + def execute(self, context): + context["map_name"] = self.map_name raise AirflowSkipException("Imagine this task failed!") return HasMapName.partial(task_id=task_id, map_index_template=template).expand( @@ -704,8 +707,6 @@ def test_expand_mapped_task_instance_with_named_index( .order_by(TaskInstance.map_index) ).all() - indices = [str(index) if index is not None else None for index in indices] - assert indices == expected_rendered_names From 194924c1062da2aa595f73d64a0ca035a976ad5f Mon Sep 17 00:00:00 2001 From: Tzu-ping Chung Date: Fri, 12 Apr 2024 16:12:33 +0800 Subject: [PATCH 09/10] Ensure context is set to render map index name --- airflow/models/taskinstance.py | 39 ++++++++++++++++------------------ 1 file changed, 18 insertions(+), 21 deletions(-) diff --git a/airflow/models/taskinstance.py b/airflow/models/taskinstance.py index 1ddae9116642c..797010ff463f7 100644 --- a/airflow/models/taskinstance.py +++ b/airflow/models/taskinstance.py @@ -2641,16 +2641,6 @@ def _register_dataset_changes(self, *, events: DatasetEventAccessors, session: S session=session, ) - def _render_map_index(self, context, jinja_env=None): - # DAG authors define map_index_template at the task level - if jinja_env is not None and (template := context.get("map_index_template")) is not None: - rendered_map_index = jinja_env.from_string(template).render(context) - log.info("Map index rendered as %s", rendered_map_index) - else: - rendered_map_index = None - - return rendered_map_index - def _execute_task_with_callbacks(self, context: Context, test_mode: bool = False, *, session: Session): """Prepare Task for Execution.""" if TYPE_CHECKING: @@ -2723,18 +2713,25 @@ def signal_handler(signum, frame): previous_state=TaskInstanceState.QUEUED, task_instance=self, session=session ) - try: - # Execute the task - with set_current_context(context): + def _render_map_index(context: Context, *, jinja_env: jinja2.Environment | None) -> str | None: + """Render named map index if the DAG author defined map_index_template at the task level.""" + if jinja_env is None or (template := context.get("map_index_template")) is None: + return None + rendered_map_index = jinja_env.from_string(template).render(context) + log.info("Map index rendered as %s", rendered_map_index) + return rendered_map_index + + # Execute the task. + with set_current_context(context): + try: result = self._execute_task(context, task_orig) - except Exception: - # If the task failed, swallow rendering error so it doesn't mask the main error. - with contextlib.suppress(jinja2.TemplateSyntaxError, jinja2.UndefinedError): - self.rendered_map_index = self._render_map_index(context, jinja_env=jinja_env) - raise - else: - # If the task succeeded, render normally to let rendering error bubble up. - self.rendered_map_index = self._render_map_index(context, jinja_env=jinja_env) + except Exception: + # If the task failed, swallow rendering error so it doesn't mask the main error. + with contextlib.suppress(jinja2.TemplateSyntaxError, jinja2.UndefinedError): + self.rendered_map_index = _render_map_index(context, jinja_env=jinja_env) + raise + else: # If the task succeeded, render normally to let rendering error bubble up. + self.rendered_map_index = _render_map_index(context, jinja_env=jinja_env) # Run post_execute callback self.task.post_execute(context=context, result=result) From 2ea6ad72bef974048408cb5a82130489f19cacbb Mon Sep 17 00:00:00 2001 From: TJaniF Date: Fri, 12 Apr 2024 13:11:47 +0200 Subject: [PATCH 10/10] change log level of rendered map index to debug --- airflow/models/taskinstance.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/airflow/models/taskinstance.py b/airflow/models/taskinstance.py index 2edb1106e4fc4..3eb8820a5f979 100644 --- a/airflow/models/taskinstance.py +++ b/airflow/models/taskinstance.py @@ -2719,7 +2719,7 @@ def _render_map_index(context: Context, *, jinja_env: jinja2.Environment | None) if jinja_env is None or (template := context.get("map_index_template")) is None: return None rendered_map_index = jinja_env.from_string(template).render(context) - log.info("Map index rendered as %s", rendered_map_index) + log.debug("Map index rendered as %s", rendered_map_index) return rendered_map_index # Execute the task.