Airflow中DataprocSubmitJobOperator如何通过XCom或其他方式传递任务数据?
解决Airflow DataprocSubmitJobOperator任务间数据传递的XCom KeyError问题
问题根源
你遇到的KeyError: 'ti'是因为DataprocSubmitJobOperator提交的PySpark任务运行在独立的Dataproc集群节点上,完全脱离Airflow的运行上下文——这些节点没有Airflow Task Instance(ti)的环境变量和上下文对象,所以直接在PySpark代码里调用ti.xcom_push必然报错。
可行解决方案
方案1:用中间存储传递数据(推荐,适合大数据量)
这是生产环境最常用的方式,把PySpark任务的输出写入中间存储(如GCS、HDFS、MongoDB/MySQL等),后续任务从该存储读取数据。
步骤1:修改获取客户列表的PySpark任务,将结果写入中间存储
from pyspark.sql import SparkSession # 初始化Spark会话 spark = SparkSession.builder.appName("FetchCustomers").getOrCreate() # 从Mongo读取客户数据 customers_df = spark.read.format("mongo") \ .option("uri", "mongodb://your-mongo-host:27017/db.collection") \ .load() # 将数据写入GCS(示例,也可以用HDFS/数据库) customers_df.write.mode("overwrite").parquet("gs://your-gcs-bucket/customers_data.parquet")
步骤2:在DAG中给后续任务传递存储路径参数
from airflow.providers.google.cloud.operators.dataproc import DataprocSubmitJobOperator from airflow.providers.google.cloud.dataproc import SparkJob run_dataproc_spark_insights = DataprocSubmitJobOperator( task_id="run_dataproc_spark_insights", region="your-gcp-region", project_id="your-gcp-project", job=SparkJob( main_python_file_uri="gs://your-gcs-bucket/insights_job.py", # 将存储路径作为参数传给后续PySpark任务 args=["gs://your-gcs-bucket/customers_data.parquet"] ), gcp_conn_id="google_cloud_default" ) # 设置任务依赖 run_dataproc_spark_getcutomers >> run_dataproc_spark_insights
步骤3:后续PySpark任务读取中间存储
在insights_job.py中读取传入的路径:
from pyspark.sql import SparkSession import sys spark = SparkSession.builder.appName("CustomerInsights").getOrCreate() # 获取DAG传递的路径参数 data_path = sys.argv[1] customers_df = spark.read.parquet(data_path) # 后续业务逻辑...
方案2:用XCom传递小数据量(仅适合少量数据,如客户ID列表)
如果数据量极小(比如几百个客户ID),可以通过捕获Dataproc作业的输出,再用Airflow的PythonOperator将数据推送到XCom。
步骤1:修改PySpark任务,将结果打印到标准输出
from pyspark.sql import SparkSession import json spark = SparkSession.builder.appName("FetchCustomers").getOrCreate() customers_df = spark.read.format("mongo").load() # 仅提取需要的小数据(比如前100个客户ID) customer_ids = [row.id for row in customers_df.select("id").limit(100).collect()] # 将数据转为JSON格式打印到stdout,Dataproc会捕获该输出 print(json.dumps(customer_ids))
步骤2:添加PythonOperator捕获输出并推XCom
from airflow.operators.python import PythonOperator from airflow.providers.google.cloud.hooks.dataproc import DataprocHook def extract_dataproc_output(**context): # 初始化Dataproc Hook hook = DataprocHook(gcp_conn_id="google_cloud_default") # 获取Dataproc作业的Job ID(DataprocSubmitJobOperator默认会把job_id推到XCom) job_id = context["ti"].xcom_pull(task_ids="run_dataproc_spark_getcutomers", key="job_id") # 获取作业的标准输出 job_output = hook.get_job_output(job_id=job_id, region="your-gcp-region") # 解析输出中的JSON数据 import json customer_ids = json.loads(job_output.strip()) # 将数据推送到XCom context["ti"].xcom_push(key="customer_ids", value=customer_ids) # 定义捕获输出的任务 pull_dataproc_output = PythonOperator( task_id="pull_dataproc_output", python_callable=extract_dataproc_output, provide_context=True, gcp_conn_id="google_cloud_default" )
步骤3:后续任务从XCom获取数据
run_dataproc_spark_insights = DataprocSubmitJobOperator( task_id="run_dataproc_spark_insights", region="your-gcp-region", project_id="your-gcp-project", job=SparkJob( main_python_file_uri="gs://your-gcs-bucket/insights_job.py", # 从XCom读取数据作为参数 args=["{{ ti.xcom_pull(task_ids='pull_dataproc_output', key='customer_ids') }}"] ), gcp_conn_id="google_cloud_default" ) # 设置任务依赖 run_dataproc_spark_getcutomers >> pull_dataproc_output >> run_dataproc_spark_insights
内容的提问来源于stack exchange,提问作者Karan Alang
相关产品推荐
相关产品推荐

