Apache Airflow 3.1.7 DAG编写难题:Operator与Task数据传递
Apache Airflow 3.1.7 DAG 数据管道问题解决
需求概述
需要实现以下数据管道流程:
- 基于
data_interval从MS SQL数据库获取行数据; - 在Python函数中对获取的数据排序并处理;
- 对每一行数据执行额外SQL查询补充信息;
- 发送处理完成的数据。
核心问题
- 混合使用Operator(如
SQLExecuteQueryOperator)与TaskFlow任务时,数据传递容易出错; - 动态展开的
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
相关产品推荐
相关产品推荐

