You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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的分区机制,通过控制分区数量和每个分区内的线程数,直接锁定总并发数。

具体逻辑:

  1. 将DataFrame重分区为8个分区(对应目标并发数);
  2. 每个分区内用单线程的线程池处理请求,确保单分区仅占用1个连接;
  3. 总并发数=分区数×每个分区线程数=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。

具体逻辑:

  1. 为每个executor的Python进程初始化一个全局Session,配置连接池最大连接数为1;
  2. 将Spark的任务并行度(如spark.sql.shuffle.partitions)设为8,确保同时运行的任务数为8;
  3. 每个任务对应一个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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.04 12:03:29