如何用Spark替代多线程处理API调用任务(Spark3+Python11)
基于Spark分布式能力解决大数量API调用OOM问题
核心思路
原方案把所有ID拉到Driver端执行多线程调用,数据量过大时直接导致Driver内存溢出。改用Spark的mapPartitions算子,将ID数据分布式到各个Executor节点的分区中,在每个分区内进行多线程API调用,每个节点仅处理部分数据,避免全量数据集中在Driver端。
具体实现代码
from pyspark.sql import SparkSession, Row from concurrent.futures import ThreadPoolExecutor from pyspark.sql.types import StructType, StructField, StringType, IntegerType # 需根据API实际返回字段调整 # 定义API返回数据对应的Schema,需与test_api返回的Row字段完全匹配 result_schema = StructType([ StructField("idName", StringType(), nullable=False), StructField("api_field1", StringType(), nullable=True), StructField("api_field2", IntegerType(), nullable=True) # 添加API返回的其他字段 ]) def test_api(id_val): # 原有API调用逻辑,调整为返回Row对象(而非DataFrame) # 示例:模拟API返回结果,实际替换为真实API调用代码 api_response = {"idName": id_val, "api_field1": "data_" + id_val, "api_field2": 123} return Row(**api_response) if api_response else None def process_partition(ids_iter): # 每个分区内启动多线程处理API调用 with ThreadPoolExecutor(max_workers=10) as executor: results = executor.map(test_api, ids_iter) # 过滤空结果,返回有效数据 return filter(None, results) def main(): spark = SparkSession.builder.appName("APItoDelta").getOrCreate() # STEP 1 - 读取ID数据,保留为DataFrame,不collect到Driver ids_df = spark.sql(SQL.IDS).select("idName") # STEP 2 - 分布式处理:每个分区内多线程调用API # 可选:根据数据量手动调整分区数,如repartition(100) result_rdd = ids_df.rdd.map(lambda row: row.idName).mapPartitions(process_partition) # 将RDD转换为结构化DataFrame final_df = spark.createDataFrame(result_rdd, schema=result_schema) # STEP 3 - 写入Delta Lake final_df.write.format("delta") \ .mode("overwrite") \ .option("overwriteSchema", "true") \ .save("/temp_tables/test_api3") if __name__ == "__main__": main()
关键细节说明
- 分区调整:可通过
ids_df.repartition(N)手动设置分区数,N需结合集群资源和API并发限制调整,避免单分区数据量过大或并发过高触发API限流。 - Schema匹配:必须明确指定返回DataFrame的Schema,Spark无法自动推断分布式处理后的结构,需与
test_api返回的Row字段完全对应。 - 线程数控制:每个分区内的线程数不宜过高,防止API服务触发限流或Executor节点资源耗尽,建议根据API的QPS限制动态调整。
- 空值过滤:在
process_partition中过滤空结果,避免无效数据进入最终DataFrame。
内容的提问来源于stack exchange,提问作者user3692015
相关产品推荐
相关产品推荐

