PySpark中如何限制对API端点的并行连接数量?
问题:PySpark中限制API调用并发连接数
假设启动PySpark会话时配置了32个executor,但希望仅以8个并发连接调用API端点,该如何实现?示例代码如下:
data = [("A", "url1"), ("B", "url2"), ("C", "url3")] columns = ["col1", "col2"] df = spark.createDataFrame(data, columns) def api_udf(url): response = requests.get(url) if response.status_code == 200: return response.json() else: return None api_udf_spark = udf(api_udf, StringType()) # Q. How to limit concurrent connections when applying UDF? result_df = df.withColumn("api_response", api_udf_spark(df["col2"])) result_df.show(truncate=False)
解决方案
可以通过两种核心方式精确控制API调用的并发数,适配32个executor的场景:
方法1:使用mapPartitions结合线程池(推荐)
利用Spark的分区机制,通过控制分区数量和每个分区内的线程数,直接锁定总并发数。
具体逻辑:
- 将DataFrame重分区为8个分区(对应目标并发数);
- 每个分区内用单线程的线程池处理请求,确保单分区仅占用1个连接;
- 总并发数=分区数×每个分区线程数=8×1=8,刚好符合需求。
代码示例:
from pyspark.sql import Row import requests from concurrent.futures import ThreadPoolExecutor def process_partition(partition): # 每个分区用1个线程处理,控制单分区并发数 with ThreadPoolExecutor(max_workers=1) as executor: results = [] for row in partition: url = row.col2 # 提交请求并同步获取结果 response = executor.submit(requests.get, url).result() resp_json = response.json() if response.status_code == 200 else None results.append(Row(col1=row.col1, col2=row.col2, api_response=resp_json)) return results # 重分区为8个,锁定总并发数 df_repartitioned = df.repartition(8) # 通过mapPartitions批量处理每个分区数据 result_df = df_repartitioned.rdd.mapPartitions(process_partition).toDF() result_df.show(truncate=False)
如果需要灵活调整,比如将分区数设为4,只需将线程池max_workers改为2,总并发数仍保持为8。
方法2:通过requests连接池+调整Spark并行度
借助requests的连接池限制每个executor的连接数,同时调整Spark任务并行度,将总并发控制在8。
具体逻辑:
- 为每个executor的Python进程初始化一个全局
Session,配置连接池最大连接数为1; - 将Spark的任务并行度(如
spark.sql.shuffle.partitions)设为8,确保同时运行的任务数为8; - 每个任务对应一个executor的连接,总并发数=任务数×每个任务连接数=8×1=8。
代码示例:
from pyspark.sql.functions import udf from pyspark.sql.types import StringType import requests from requests.adapters import HTTPAdapter # 每个executor的Python进程维护一个全局Session对象 session = None def init_session(): global session if session is None: session = requests.Session() # 配置每个host的最大连接数为1 adapter = HTTPAdapter(pool_connections=1, pool_maxsize=1) session.mount('http://', adapter) session.mount('https://', adapter) def api_udf(url): init_session() response = session.get(url) # 将JSON转为字符串适配StringType返回类型 return str(response.json()) if response.status_code == 200 else None api_udf_spark = udf(api_udf, StringType()) # 需在Spark会话初始化时配置任务并行度 # spark.conf.set("spark.sql.shuffle.partitions", "8") result_df = df.withColumn("api_response", api_udf_spark(df["col2"])) result_df.show(truncate=False)
注意:此方法必须配合调整Spark的任务并行度配置,确保同时运行的任务数不超过8,避免连接数超出限制。
内容的提问来源于stack exchange,提问作者The Wanderer
相关产品推荐
相关产品推荐

