You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

Apache Airflow 3.1.7 DAG编写难题:Operator与Task数据传递

Apache Airflow 3.1.7 DAG 数据管道问题解决

需求概述

需要实现以下数据管道流程:

  • 基于data_interval从MS SQL数据库获取行数据;
  • 在Python函数中对获取的数据排序并处理;
  • 对每一行数据执行额外SQL查询补充信息;
  • 发送处理完成的数据。

核心问题

  1. 混合使用Operator(如SQLExecuteQueryOperator)与TaskFlow任务时,数据传递容易出错;
  2. 动态展开的handle_row任务无法正确与后续submit任务建立依赖关系。

问题代码示例

import logging
from datetime import datetime, timedelta
from zoneinfo import ZoneInfo

# The DAG object; we'll need this to instantiate a DAG
from airflow.sdk import dag, task, chain
from airflow.providers.standard.operators.python import PythonOperator, BranchPythonOperator, ShortCircuitOperator
from airflow.providers.standard.operators.empty import EmptyOperator
from airflow.providers.common.sql.operators.sql import SQLExecuteQueryOperator
from airflow.timetables.interval import CronDataIntervalTimetable

def output_processor_as_dict(results, descriptions):
    columns = [x[0] for x in descriptions[0]]
    return [ [dict(zip(columns, row)) for row in results[0]] ]

@dag(
    "troublesome_dag",
    default_args={
        "depends_on_past": False,
        "retries": 3,
        "retry_delay": timedelta(minutes=1)
    },
    schedule=CronDataIntervalTimetable("*/10 * * * *", timezone="Europe/Copenhagen"),
    start_date=datetime(2026,1,1),
    catchup=False,
    tags=["db"],
)
def trouble_dag():
    fetch_sql = SQLExecuteQueryOperator(
        task_id = 'fetch_processes',
        conn_id = 'conn_id',
        show_return_value_in_logs=True,
        requires_result_fetch = True,
        do_xcom_push=True,
        output_processor=output_processor_as_dict,
        sql = r"""SELECT TOP (50) * FROM tbl;""",
    )

    @task(show_return_value_in_logs=True)
    def make_diff_times(results: list, **kwargs):
        results.sort(key=lambda x: f'{x["id"]}')
            previous = None
            for row in results:
                if previous is not None and previous["group"] == row["group"]
                    # do something to row
        return results

    @task(show_return_value_in_logs=True)    
    def handle_row(row, **kwargs):
        if row["station"] == 'update':
            # query another database for more data
        return row

    submit = GrpcOperator(task_id='submit', ...)

    rows = make_diff_times(fetch_sql.output)  ## working as expected
    handle_process.expand(row=rows)           ## working as expected
    handle_process >> submit                  ## ???
 
trouble_dag()   

解决方案

1. 修复Operator与TaskFlow之间的数据传递

SQLExecuteQueryOperator的output_processor返回了嵌套列表([ [dict,...] ]),导致TaskFlow任务接收的数据格式异常,需调整输出处理器:

def output_processor_as_dict(results, descriptions):
    columns = [x[0] for x in descriptions[0]]
    # 直接返回字典列表,去除不必要的嵌套
    return [dict(zip(columns, row)) for row in results[0]]

同时修正make_diff_times的语法错误(缩进、缺少冒号)并优化排序逻辑:

@task(show_return_value_in_logs=True)
def make_diff_times(results: list, **kwargs):
    # 直接按id字段排序,无需转字符串
    results.sort(key=lambda x: x["id"])
    previous = None
    for row in results:
        if previous is not None and previous["group"] == row["group"]:
            # 在这里添加行处理逻辑,比如计算时间差
            row["time_diff"] = row["timestamp"] - previous["timestamp"]
        previous = row
    return results

2. 建立动态任务与submit的依赖关系

使用.expand()会生成多个并行的handle_row任务,要让所有任务完成后再执行submit,只需将展开后的任务集合赋值给变量,再与submit建立依赖:

方式1:使用TaskFlow的submit任务(推荐)

@task(show_return_value_in_logs=True)
def submit_task(processed_rows):
    # 实现Grpc发送逻辑
    import grpc
    # 替换为你的Grpc客户端实现
    from your_proto_def import service_pb2, service_pb2_grpc
    with grpc.insecure_channel('grpc-server:50051') as channel:
        stub = service_pb2_grpc.DataServiceStub(channel)
        # 根据proto定义构造请求
        request = service_pb2.SubmitRequest(rows=processed_rows)
        response = stub.SendData(request)
    logging.info(f"Grpc submit response: {response}")

# 任务串联
processed_list = make_diff_times(fetch_sql.output)
all_processed_rows = handle_row.expand(row=processed_list)
# 所有动态任务完成后执行submit
submit_task(all_processed_rows)

方式2:使用GrpcOperator

如果必须使用Operator,需要通过XCom收集所有handle_row的输出,再在GrpcOperator中获取:

submit = GrpcOperator(
    task_id='submit',
    # 通过XCom拉取所有handle_row的返回值
    data="{{ ti.xcom_pull(task_ids='handle_row', key='return_value') }}"
)

all_processed_rows = handle_row.expand(row=processed_list)
all_processed_rows >> submit

3. 补充:基于data_interval查询数据

原SQL未使用data_interval,需添加Airflow模板变量实现区间查询(适配MS SQL格式):

sql = r"""SELECT TOP (50) * FROM tbl 
          WHERE created_at BETWEEN '{{ data_interval_start }}' AND '{{ data_interval_end }}';"""

完整修正代码

import logging
from datetime import datetime, timedelta
from zoneinfo import ZoneInfo

from airflow.sdk import dag, task
from airflow.providers.common.sql.operators.sql import SQLExecuteQueryOperator
from airflow.timetables.interval import CronDataIntervalTimetable
# 替换为你的GrpcOperator导入路径
from your_provider.operators.grpc import GrpcOperator

def output_processor_as_dict(results, descriptions):
    columns = [x[0] for x in descriptions[0]]
    return [dict(zip(columns, row)) for row in results[0]]

@dag(
    "troublesome_dag",
    default_args={
        "depends_on_past": False,
        "retries": 3,
        "retry_delay": timedelta(minutes=1)
    },
    schedule=CronDataIntervalTimetable("*/10 * * * *", timezone="Europe/Copenhagen"),
    start_date=datetime(2026,1,1),
    catchup=False,
    tags=["db"],
)
def trouble_dag():
    fetch_sql = SQLExecuteQueryOperator(
        task_id = 'fetch_processes',
        conn_id = 'conn_id',
        show_return_value_in_logs=True,
        requires_result_fetch = True,
        do_xcom_push=True,
        output_processor=output_processor_as_dict,
        sql = r"""SELECT TOP (50) * FROM tbl 
                  WHERE created_at BETWEEN '{{ data_interval_start }}' AND '{{ data_interval_end }}';""",
    )

    @task(show_return_value_in_logs=True)
    def make_diff_times(results: list, **kwargs):
        results.sort(key=lambda x: x["id"])
        previous = None
        for row in results:
            if previous is not None and previous["group"] == row["group"]:
                row["time_diff"] = row["timestamp"] - previous["timestamp"]
            previous = row
        return results

    @task(show_return_value_in_logs=True)    
    def handle_row(row, **kwargs):
        if row["station"] == 'update':
            # 使用MsSqlHook执行额外查询
            from airflow.providers.microsoft.mssql.hooks.mssql import MsSqlHook
            hook = MsSqlHook(conn_id="another_ms_sql_conn")
            extra_data = hook.get_first("SELECT * FROM another_tbl WHERE id = %s", parameters=(row["id"],))
            if extra_data:
                # 假设查询返回的是元组,需要与字段名映射,这里简化处理
                extra_columns = ["extra_col1", "extra_col2"]
                row.update(dict(zip(extra_columns, extra_data)))
        return row

    @task(show_return_value_in_logs=True)
    def submit_task(processed_rows):
        # Grpc发送逻辑示例
        import grpc
        from your_proto import data_pb2, data_pb2_grpc
        with grpc.insecure_channel('grpc-service:50051') as channel:
            stub = data_pb2_grpc.DataSubmitStub(channel)
            # 构造请求消息
            rows_msg = [data_pb2.RowData(**row) for row in processed_rows]
            request = data_pb2.SubmitRequest(rows=rows_msg)
            response = stub.Submit(request)
        logging.info(f"Submit succeeded: {response.success}")

    # 任务依赖
    processed_list = make_diff_times(fetch_sql.output)
    all_processed = handle_row.expand(row=processed_list)
    submit_task(all_processed)

trouble_dag()

内容的提问来源于stack exchange,提问作者MrGumble

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.11 13:05:54