多线程调用API写入Spark DataFrame的数据丢失问题排查
问题描述
原本需要向表中导入20万条记录,单条循环插入耗时约5小时,于是改用Python多线程调用API,将返回结果累加至空Spark DataFrame后批量写入表。运行时所有API均返回状态码200,写入过程无报错,但最终表中记录数远少于20万(每次运行数量不同,比如1735、5000条),且已确认API返回有效数据,需排查数据丢失原因。
原实现代码:
def prepare_empty_df(schema, spark: SparkSession) -> DataFrame: empty_rdd = spark.sparkContext.emptyRDD() empty_df = spark.createDataFrame(empty_rdd, schema) return empty_df class RunApiCalls: def __init__(self, df: DataFrame=None): self.finalDf = df def do_some_transformations(df: DataFrame) -> DataFrame: return do_some_transformation_output_dataframe def get_json(self, spark, PARAMETER): try: token_headers = create_bearer_token() session = get_session() api_response = session.get(f'API_URL/?API_PARAMETER={PARAMETER}', headers=token_headers) print(f'API call: API_URL/?API_PARAMETER={PARAMETER} -> Status code: {api_response.status_code}') api_json_object = json.loads(api_response.text) string_data = json.dumps(api_json_object) json_df = spark.createDataFrame([(1, string_data)],["id","value"]) api_dataframe = do_some_transformations(json_df) self.finalDf = self.finalDf.unionAll(api_dataframe) except Exception as error: traceback.print_exc() def api_main(self, spark, batch_size, state_names) -> DataFrame: try: for i in range(0, len(state_names), batch_size): sub_list = state_names[i:i + batch_size] threads = [] for index in range(len(sub_list)): t = threading.Thread(target=self.get_json, name=str(index), args=(spark, sub_list[index])) threads.append(t) t.start() for index, thread in enumerate(threads): thread.join() print(f"All Threads completed for this sub_list{i}") return self.finalDf except Exception as e: traceback.print_exc() if __name__ == "__main__": spark = SparkSession.builder.appName('SOME_APP_NAME').getOrCreate() batch_size = 15 empty_df = prepare_empty_df(schema=schema, spark=spark) print('Created Empty Dataframe') api_param_list = get_list() print(f'api param list: {api_param_list}') api_call = RunApiCalls(df=empty_df) final_df = api_call.api_main(spark=spark, batch_size=batch_size, state_names=api_param_list) final_df.write.mode('append').saveAsTable("some_database.some_tablebname")
数据丢失原因分析
- DataFrame并发修改的竞态条件:Spark DataFrame是不可变对象,
unionAll操作会生成新的DataFrame。多线程环境下,多个线程同时读取同一个旧的self.finalDf,各自生成新的DataFrame后,最后只有一个线程的赋值会生效,其他线程的结果直接被覆盖。 - 无线程同步机制:代码中未对
self.finalDf的修改操作加锁,多个线程并发修改同一个实例变量,导致数据覆盖丢失,且每次运行的覆盖情况随机,所以最终记录数不稳定。 - SparkSession线程不安全:SparkSession并非线程安全组件,多线程共享同一个Session执行创建DataFrame等操作,可能引发未定义行为,间接导致数据丢失。
修复方案
用线程安全容器暂存数据
不要直接在多线程中修改Spark DataFrame,改用锁保护的列表暂存每个API返回的小DataFrame,所有线程执行完成后再统一合并:from threading import Lock class RunApiCalls: def __init__(self, df: DataFrame=None): self.finalDf = df self.df_list = [] self.lock = Lock() def get_json(self, spark, PARAMETER): try: token_headers = create_bearer_token() session = get_session() api_response = session.get(f'API_URL/?API_PARAMETER={PARAMETER}', headers=token_headers) print(f'API call: API_URL/?API_PARAMETER={PARAMETER} -> Status code: {api_response.status_code}') api_json_object = json.loads(api_response.text) string_data = json.dumps(api_json_object) json_df = spark.createDataFrame([(1, string_data)],["id","value"]) api_dataframe = do_some_transformations(json_df) with self.lock: self.df_list.append(api_dataframe) except Exception as error: traceback.print_exc() def api_main(self, spark, batch_size, state_names) -> DataFrame: try: for i in range(0, len(state_names), batch_size): sub_list = state_names[i:i + batch_size] threads = [] for index in range(len(sub_list)): t = threading.Thread(target=self.get_json, name=str(index), args=(spark, sub_list[index])) threads.append(t) t.start() for index, thread in enumerate(threads): thread.join() print(f"All Threads completed for this sub_list{i}") # 合并所有暂存的DataFrame if self.df_list: for df in self.df_list: self.finalDf = self.finalDf.unionAll(df) return self.finalDf except Exception as e: traceback.print_exc()改用Spark原生并行机制
Spark本身是分布式计算框架,建议用RDD的并行能力替代Python多线程,避免线程安全问题:def process_param(param): # 实现API调用和数据转换,返回符合schema的行数据 token_headers = create_bearer_token() session = get_session() api_response = session.get(f'API_URL/?API_PARAMETER={param}', headers=token_headers) api_json_object = json.loads(api_response.text) # 根据实际schema转换为行数据 return [transform_to_row(api_json_object)] if __name__ == "__main__": spark = SparkSession.builder.appName('SOME_APP_NAME').getOrCreate() api_param_list = get_list() # 并行处理所有参数 rdd = spark.sparkContext.parallelize(api_param_list) final_df = rdd.flatMap(process_param).toDF(schema) final_df.write.mode('append').saveAsTable("some_database.some_tablebname")避免多线程共享SparkSession
如果坚持使用Python多线程,需为每个线程创建独立的SparkSession,或对SparkSession的访问加锁,但更推荐使用Spark原生并行方式。
内容的提问来源于stack exchange,提问作者Metadata
相关产品推荐
相关产品推荐

