Source code for airflow.example_dags.example_task_loops
# Licensed to the Apache Software Foundation (ASF) under one
# or more contributor license agreements. See the NOTICE file
# distributed with this work for additional information
# regarding copyright ownership. The ASF licenses this file
# to you under the Apache License, Version 2.0 (the
# "License"); you may not use this file except in compliance
# with the License. You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing,
# software distributed under the License is distributed on an
# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
# KIND, either express or implied. See the License for the
# specific language governing permissions and limitations
# under the License.
"""Repeat task groups with fixed counts, runtime conditions, and mapped tasks."""
from __future__ import annotations
# [START refine_estimate]
from airflow.sdk import dag, task, task_group
@dag(schedule=None, catchup=False, tags=["example"])
[docs]
def refine_estimate():
@task_group
def refine():
@task
def improve(*, loop):
estimate = 1.0 if loop.index == 0 else loop.previous["estimate"]
return (estimate + 2.0 / estimate) / 2.0
@task
def evaluate(estimate):
return {"estimate": estimate, "error": abs(estimate * estimate - 2.0)}
evaluate(improve())
def accurate_enough(*, loop):
return loop.result["error"] < 0.000001
@task
def finished():
print("Refinement finished.")
refinement = refine.loop(max_iterations=10, until=accurate_enough)
refinement >> finished()
refine_estimate()
# [END refine_estimate]
# [START fixed_loop]
@dag(schedule=None, catchup=False, tags=["example"])
[docs]
def fixed_task_loop():
@task_group
def accumulate():
@task
def increment(*, loop):
return (0 if loop.index == 0 else loop.previous) + 1
increment()
accumulate.loop(max_iterations=3)
fixed_task_loop()
# [END fixed_loop]
# [START mapped_loop]
@dag(schedule=None, catchup=False, tags=["example"])
[docs]
def mapped_task_loop():
@task_group
def process_batch():
@task
def choose_items(*, loop):
return [1, 2] if loop.index == 0 else [value + 1 for value in loop.previous]
@task
def process(value, *, loop, ti):
print(f"Iteration {loop.index}, mapped position {ti.map_index}")
return value + loop.index
process.expand(value=choose_items())
def batch_ready(*, loop):
return min(loop.result) >= 2
process_batch.loop(max_iterations=3, until=batch_ready)
mapped_task_loop()
# [END mapped_loop]
# [START partial_override_loop]
@dag(schedule=None, catchup=False, tags=["example"])
[docs]
def partial_override_loop():
@task_group
def accumulate(increment_by):
@task
def increment(increment_by, *, loop):
return (0 if loop.index == 0 else loop.previous) + increment_by
increment(increment_by)
accumulate.override(group_id="accumulate_more").partial(increment_by=5).loop(max_iterations=3)
partial_override_loop()
# [END partial_override_loop]