mapPartitions: connection pool паттерн для JDBC и HTTP клиентов

Фундаментальный разбор mapPartitions в PySpark: антипаттерн row-by-row с JDBC/HTTP, ошибка NotSerializableException, connection pool на уровне Executor, паттерн Batch Insert для JDBC, Rate Limiting для HTTP API, управление ресурсами через try/finally, ленивые итераторы vs list(), контроль числа партиций и мониторинг в Spark UI.

optimization

1. Проблема построчной обработки: антипаттерн .map() с внешними системами

Многие production пайплайны в Data Engineering требуют обращения к внешним системам: загрузки enrichment-данных из PostgreSQL, обогащения записей через REST API, записи результатов в ClickHouse или Elasticsearch, ML-инференса через gRPC эндпоинт. На первый взгляд задача выглядит просто — возьми каждую строку и обратись к внешней системе. Но именно здесь начинаются серьёзные архитектурные проблемы.

Почему row-by-row с внешними системами — катастрофа

Рассмотрим то, что делает большинство начинающих инженеров:

import psycopg2
from pyspark.sql import SparkSession, functions as F
from pyspark.sql.types import StringType

spark = SparkSession.builder.getOrCreate()

# Антипаттерн: JDBC-соединение в .map() / .withColumn() + UDF
@udf(returnType=StringType())
def enrich_with_db_naive(user_id: int) -> str:
    """
    ❌ АНТИПАТТЕРН: создаём соединение на КАЖДУЮ строку!
    При 1M строк создаём 1M TCP соединений!
    """
    conn = psycopg2.connect(
        host="postgres", database="users",
        user="spark", password="secret"
    )  # 3-way TCP handshake + TLS + auth = ~10-50 ms
    cur = conn.cursor()
    cur.execute("SELECT tier FROM users WHERE id = %s", (user_id,))
    result = cur.fetchone()
    conn.close()  # ещё overhead на закрытие
    return result[0] if result else "unknown"

df = spark.table("bronze.events")
enriched = df.withColumn("tier", enrich_with_db_naive(F.col("user_id")))
enriched.write.parquet("hdfs://cluster/silver/events/")

Посмотрим что происходит при 10 млн строк:

Схема показывает разницу: при map() создаётся 10 миллионов соединений, при mapPartitions() с 200 партициями — только 200. Разница в нагрузке на PostgreSQL — в 50 000 раз!

Реальные последствия антипаттерна

  • PostgreSQL упадёт или начнёт throttling. Дефолтный max_connections = 100. Один Spark-кластер с 50 Executor'ами × 200 партиций × 1 соединение на строку = сотни тысяч попыток подключения в секунду.
  • Время выполнения растёт катастрофически. Создание JDBC соединения занимает 10-50 мс. При 10M строк: 10M × 30 мс = 83 часа только на соединения!
  • OOM на воркерах. Каждое открытое соединение занимает RAM. Тысячи одновременных соединений исчерпают память Executor'а.

2. Философия mapPartitions: пакетная обработка данных

mapPartitions() переводит логику с уровня строк на уровень физических разделов (Partitions). Spark вызывает вашу функцию один раз на партицию, передавая итератор всех строк.

Анатомия mapPartitions

Диаграмма показывает ключевую разницу: функция вызывается один раз, соединение создаётся один раз, а итератор обрабатывается в цикле. Это и есть паттерн mapPartitions.

Сравнение API: map vs mapPartitions

from pyspark.sql import SparkSession

spark = SparkSession.builder.getOrCreate()
sc = spark.sparkContext

# map: функция вызывается для КАЖДОГО элемента
# f(x) вызывается N раз
rdd = sc.parallelize([1, 2, 3, 4, 5])
result = rdd.map(lambda x: x * 2)  # f вызывается 5 раз

# mapPartitions: функция вызывается для каждой ПАРТИЦИИ
# f(iterator) вызывается num_partitions раз
def process_partition(iterator):
    # Здесь инициализация тяжёлых объектов
    # Они создаются ОДИН РАЗ на партицию
    expensive_client = create_db_connection()  # один раз!

    for item in iterator:
        yield process_item(item, expensive_client)

    expensive_client.close()  # один раз!

result = rdd.mapPartitions(process_partition)

# DataFrame API: используем df.rdd.mapPartitions() или foreachPartition()
df_result = df.rdd.mapPartitions(process_partition).toDF(schema)

Когда mapPartitions необходим (не опционален)

mapPartitions — это не просто оптимизация. В некоторых сценариях это единственный правильный подход:

Сценарий Почему map() не работает Почему mapPartitions работает
JDBC queries 1M соединений перегрузят БД 200 соединений — управляемо
HTTP API с rate limiting 1M requests/sec — 429 errors 200 connections × батчи
ML model inference 1M model.load() — нет смысла Загружаем модель 1 раз на партицию
Redis batch lookup 1M отдельных GET pipeline.execute() батчами
Запись в Kafka 1M отдельных produce() Один Producer на партицию

3. Анатомия несериализуемых объектов: ошибка NotSerializableException

Перед тем как написать правильный код, важно понять типичную ошибку, которую совершают при попытке передать JDBC-соединение из Driver'а на Executor'ы.

Почему нельзя создать соединение на Driver'е

# ❌ ОШИБКА: создаём соединение на Driver'е
import psycopg2

# Это выполняется на Driver'е
conn = psycopg2.connect(host="postgres", database="users",
                        user="spark", password="secret")

def process_row_bad(row):
    # Spark пытается сериализовать 'conn' как часть closure!
    # conn — это объект с открытым TCP-сокетом
    # TCP-сокеты не могут быть сериализованы!
    cur = conn.cursor()  # ← conn захвачен в closure
    cur.execute("SELECT tier FROM users WHERE id = %s", (row.user_id,))
    return cur.fetchone()[0]

df.rdd.map(process_row_bad)
# Ошибка:
# PicklingError: Can't pickle local object
# Или: java.io.NotSerializableException: psycopg2.extensions.connection

Правильное решение: создавать соединение внутри mapPartitions

# ✅ ПРАВИЛЬНО: соединение создаётся ВНУТРИ функции на Executor'е
def process_partition_correct(rows_iterator):
    """
    Эта функция выполняется на Executor'е.
    conn создаётся локально — не сериализуется, не передаётся по сети.
    """
    import psycopg2  # импорт ВНУТРИ функции — это важно для PySpark!

    # Соединение создаётся на Executor'е, в памяти воркера
    # Оно НЕ передаётся по сети — никакой сериализации!
    conn = psycopg2.connect(
        host="postgres", database="users",
        user="spark", password="secret"
    )
    conn.autocommit = True
    cur = conn.cursor()

    try:
        for row in rows_iterator:
            cur.execute(
                "SELECT tier, segment FROM users WHERE id = %s",
                (row.user_id,)
            )
            result = cur.fetchone()
            if result:
                # yield возвращает обогащённую строку
                yield (row.user_id, row.amount, result[0], result[1])
            else:
                yield (row.user_id, row.amount, "unknown", "default")
    finally:
        # ОБЯЗАТЕЛЬНО: закрываем в finally!
        cur.close()
        conn.close()

# Применяем к DataFrame
from pyspark.sql.types import StructType, StructField, LongType, DoubleType, StringType

output_schema = StructType([
    StructField("user_id",  LongType()),
    StructField("amount",   DoubleType()),
    StructField("tier",     StringType()),
    StructField("segment",  StringType()),
])

result_df = df.rdd.mapPartitions(process_partition_correct).toDF(output_schema)

4. Реализация JDBC Connection Pool паттерна

Для production использования нужно идти дальше простого «один коннект на партицию». Если несколько Task'ов выполняются на одном Executor'е (при executor-cores = 5: пять одновременных Task'ов), каждый создаст своё соединение. Можно оптимизировать это через паттерн Singleton на уровне Python Worker.

Паттерн Lazy Singleton для Connection Pool

import threading

# Connection Pool хранится как module-level переменная
# Python Workers на одном Executor'е МОГУТ разделять глобальное состояние
# (но это зависит от конфигурации PySpark Workers)
_connection_pool = None
_pool_lock = threading.Lock()


def get_connection_pool():
    """
    Ленивая инициализация пула соединений.

    На одном Executor'е может работать несколько Python Workers (pyspark.daemon).
    Синглтон гарантирует создание пула только один раз на Worker-процесс.

    ВАЖНО: в PySpark каждый Worker — отдельный Python-процесс.
    Поэтому пул создаётся один раз на Python Worker-процесс,
    а не один раз на весь Executor JVM.
    """
    global _connection_pool
    if _connection_pool is None:
        with _pool_lock:
            if _connection_pool is None:
                from psycopg2 import pool
                _connection_pool = pool.ThreadedConnectionPool(
                    minconn=1,      # минимум соединений в пуле
                    maxconn=5,      # максимум соединений
                    host=os.getenv("PG_HOST", "postgres"),
                    database=os.getenv("PG_DB", "users"),
                    user=os.getenv("PG_USER", "spark"),
                    password=os.getenv("PG_PASSWORD", ""),
                    connect_timeout=10,
                )
    return _connection_pool


def process_partition_with_pool(rows_iterator):
    """
    Использует Connection Pool для переиспользования соединений
    между Task'ами на одном Worker-процессе.
    """
    import os

    pool = get_connection_pool()
    conn = pool.getconn()  # берём соединение из пула

    try:
        conn.autocommit = True
        cur = conn.cursor()

        for row in rows_iterator:
            cur.execute(
                "SELECT tier, segment, credit_limit "
                "FROM users WHERE id = %s",
                (row.user_id,)
            )
            db_row = cur.fetchone()
            if db_row:
                yield row + db_row  # Python tuple unpacking
            else:
                yield row + ("unknown", "default", 0)

        cur.close()

    except Exception as e:
        conn.rollback()
        raise
    finally:
        pool.putconn(conn)  # возвращаем соединение в пул

Полный production-ready пример

import os
from pyspark.sql import SparkSession, functions as F
from pyspark.sql.types import StructType, StructField, LongType, DoubleType, StringType, IntegerType


def create_enrichment_pipeline(spark: SparkSession):
    """
    Production-ready пайплайн обогащения событий данными из PostgreSQL.
    Использует mapPartitions с Connection Pool.
    """

    # ── Конфигурация ──────────────────────────────────────────────────
    PG_HOST = os.getenv("PG_HOST", "postgres.internal")
    PG_DB = os.getenv("PG_DB", "users_db")
    PG_USER = os.getenv("PG_USER", "spark_reader")
    PG_PASSWORD = os.getenv("PG_PASSWORD", "")

    # ── Функция обработки партиции ────────────────────────────────────
    def enrich_partition(rows_iterator):
        """
        Обогащает строки партиции данными из PostgreSQL.
        Вызывается ОДИН РАЗ на партицию.
        """
        import psycopg2
        import psycopg2.extras

        # Используем переменные из outer scope (безопасно для строк/примитивов)
        conn = psycopg2.connect(
            host=PG_HOST,
            database=PG_DB,
            user=PG_USER,
            password=PG_PASSWORD,
            connect_timeout=30,
            options="-c statement_timeout=10000"  # 10 сек таймаут на запрос
        )

        try:
            with conn.cursor(cursor_factory=psycopg2.extras.RealDictCursor) as cur:
                for row in rows_iterator:
                    cur.execute(
                        """
                        SELECT u.tier, u.segment, u.credit_limit, c.country_name
                        FROM users u
                        LEFT JOIN countries c ON u.country_code = c.code
                        WHERE u.id = %s
                        """,
                        (row.user_id,)
                    )
                    user_data = cur.fetchone()

                    if user_data:
                        yield (
                            row.event_id,
                            row.user_id,
                            row.amount,
                            row.event_type,
                            user_data["tier"],
                            user_data["segment"],
                            user_data["credit_limit"],
                            user_data["country_name"],
                        )
                    else:
                        yield (
                            row.event_id,
                            row.user_id,
                            row.amount,
                            row.event_type,
                            "free",
                            "unknown",
                            0,
                            "Unknown",
                        )

        except Exception as e:
            # Логируем ошибку с номером партиции для диагностики
            import traceback
            print(f"ERROR in partition processing: {e}")
            print(traceback.format_exc())
            raise
        finally:
            conn.close()

    # ── Схема результата ──────────────────────────────────────────────
    enriched_schema = StructType([
        StructField("event_id",     StringType()),
        StructField("user_id",      LongType()),
        StructField("amount",       DoubleType()),
        StructField("event_type",   StringType()),
        StructField("user_tier",    StringType()),
        StructField("user_segment", StringType()),
        StructField("credit_limit", IntegerType()),
        StructField("country_name", StringType()),
    ])

    # ── Основной пайплайн ─────────────────────────────────────────────
    events_df = spark.table("bronze.events")

    # Контролируем число параллельных соединений к PostgreSQL
    # 50 партиций = max 50 одновременных соединений
    events_repartitioned = events_df.repartition(50)

    enriched_df = (
        events_repartitioned.rdd
        .mapPartitions(enrich_partition)
        .toDF(enriched_schema)
    )

    return enriched_df

5. Паттерн Batch Insert для JDBC: пакетная запись

При записи данных в PostgreSQL через mapPartitions важно не только открыть одно соединение, но и использовать батч-вставку (executemany / execute_batch) вместо построчной записи.

Почему одиночные INSERT медленные

Каждый cursor.execute(INSERT ...) — это отдельный сетевой round-trip к СУБД. При 50 000 строк в партиции это 50 000 round-trips. При latency 1 мс = 50 секунд только на ожидание сети.

execute_values: оптимальная пакетная вставка

def write_partition_to_postgres(rows_iterator):
    """
    Пакетная запись строк партиции в PostgreSQL.
    Использует execute_values для максимальной производительности.
    """
    import psycopg2
    import psycopg2.extras

    BATCH_SIZE = 5000  # строк на один INSERT запрос

    conn = psycopg2.connect(
        host=os.getenv("PG_HOST"),
        database=os.getenv("PG_DB"),
        user=os.getenv("PG_USER"),
        password=os.getenv("PG_PASSWORD"),
    )
    conn.autocommit = False  # Транзакция на весь батч!

    INSERT_SQL = """
        INSERT INTO silver.enriched_events
            (event_id, user_id, amount, event_type, processed_at)
        VALUES %s
        ON CONFLICT (event_id) DO UPDATE SET
            amount = EXCLUDED.amount,
            processed_at = EXCLUDED.processed_at
    """

    inserted_total = 0
    batch = []

    try:
        with conn.cursor() as cur:
            for row in rows_iterator:
                # Собираем батч
                batch.append((
                    row.event_id,
                    row.user_id,
                    row.amount,
                    row.event_type,
                    row.processed_at,
                ))

                # Отправляем батч когда он заполнен
                if len(batch) >= BATCH_SIZE:
                    psycopg2.extras.execute_values(
                        cur, INSERT_SQL, batch,
                        template=None,
                        page_size=1000  # размер одного VALUES (...) chunk
                    )
                    inserted_total += len(batch)
                    batch = []

            # Отправляем остаток (последний неполный батч)
            if batch:
                psycopg2.extras.execute_values(cur, INSERT_SQL, batch)
                inserted_total += len(batch)

        conn.commit()
        print(f"Partition: inserted {inserted_total} rows")

    except Exception as e:
        conn.rollback()
        raise
    finally:
        conn.close()

    # mapPartitions ожидает генератор!
    # При записи (side effect) возвращаем пустой итератор
    return iter([])  # или yield ничего не возвращаем


# Используем foreachPartition для чистых side-effect операций
df.foreachPartition(write_partition_to_postgres)
# foreachPartition = mapPartitions для side effects (не возвращает данные)

Сравнение производительности вставки

Метод 100K строк Сетевых запросов Время
Одиночные INSERT в map() 100,000 100,000 ~100 сек
Одиночные INSERT в mapPartitions() 100,000 100,000 (но 1 коннект) ~80 сек
execute_values batch=5K 100,000 20 ~2 сек
COPY FROM (максимум) 100,000 1 ~0.5 сек

6. HTTP API: паттерн Session + Rate Limiting

Интеграция с REST API требует особой осторожности: внешние сервисы обычно защищены Rate Limiting.

Session Reuse: Connection Pooling для HTTP

import requests
from requests.adapters import HTTPAdapter
from urllib3.util.retry import Retry


def create_robust_http_session(
    max_retries: int = 3,
    backoff_factor: float = 0.5,
    pool_connections: int = 10,
    pool_maxsize: int = 20,
) -> requests.Session:
    """
    Создаёт Session с Connection Pooling и автоматическим retry.

    pool_connections: число пулов соединений (обычно = число хостов)
    pool_maxsize: максимум соединений в одном пуле
    backoff_factor: задержка между retry = backoff_factor × {1, 2, 4, 8...} сек
    """
    session = requests.Session()

    retry_strategy = Retry(
        total=max_retries,
        backoff_factor=backoff_factor,
        status_forcelist=[429, 500, 502, 503, 504],  # коды для retry
        allowed_methods=["GET", "POST"],
        respect_retry_after_header=True,  # соблюдаем Retry-After заголовок
    )

    adapter = HTTPAdapter(
        max_retries=retry_strategy,
        pool_connections=pool_connections,
        pool_maxsize=pool_maxsize,
    )

    session.mount("http://", adapter)
    session.mount("https://", adapter)

    return session


def enrich_via_api(rows_iterator):
    """
    Обогащение данных через внешний REST API.
    Использует Session для Connection Reuse и Bulk API где возможно.
    """
    import time

    API_URL = os.getenv("ENRICHMENT_API_URL", "https://api.example.com")
    API_KEY = os.getenv("ENRICHMENT_API_KEY", "")
    BATCH_SIZE = 100   # строк на один API запрос (Bulk API)
    RPS_LIMIT = 50     # запросов в секунду (Rate Limit)

    session = create_robust_http_session()
    session.headers.update({
        "Authorization": f"Bearer {API_KEY}",
        "Content-Type": "application/json",
    })

    batch = []
    batch_map = {}   # user_id → row (для сопоставления ответов)
    request_count = 0
    window_start = time.time()

    def send_batch_and_yield():
        """Отправляет текущий батч и возвращает обогащённые строки."""
        nonlocal request_count, window_start

        # Rate limiting: не превышаем RPS_LIMIT запросов в секунду
        request_count += 1
        elapsed = time.time() - window_start
        if elapsed < 1.0 and request_count >= RPS_LIMIT:
            sleep_time = 1.0 - elapsed
            time.sleep(sleep_time)
            request_count = 0
            window_start = time.time()

        # Отправляем Bulk API запрос
        response = session.post(
            f"{API_URL}/v1/users/bulk",
            json={"ids": list(batch_map.keys())},
            timeout=30,
        )
        response.raise_for_status()

        api_results = response.json().get("users", {})

        # Сопоставляем ответы с исходными строками
        for user_id, original_row in batch_map.items():
            user_data = api_results.get(str(user_id), {})
            yield (
                original_row.event_id,
                user_id,
                original_row.amount,
                user_data.get("tier", "free"),
                user_data.get("country", "unknown"),
            )

    try:
        for row in rows_iterator:
            batch.append(row)
            batch_map[row.user_id] = row

            if len(batch) >= BATCH_SIZE:
                yield from send_batch_and_yield()
                batch = []
                batch_map = {}

        # Последний неполный батч
        if batch:
            yield from send_batch_and_yield()

    finally:
        session.close()

7. Управление ресурсами: try/finally и очистка

Утечка соединений (Connection Leak) — критическая проблема в production. Если Task упадёт с ошибкой до явного conn.close() — соединение зависнет.

Паттерн правильной очистки ресурсов

def safe_partition_processing(rows_iterator):
    """
    Образцово-показательный паттерн управления ресурсами.
    Гарантирует освобождение ресурсов даже при ошибках.
    """
    import psycopg2
    from contextlib import contextmanager

    # Контекстный менеджер для автоматической очистки
    @contextmanager
    def managed_connection():
        conn = None
        try:
            conn = psycopg2.connect(
                host=os.getenv("PG_HOST"),
                database=os.getenv("PG_DB"),
                user=os.getenv("PG_USER"),
                password=os.getenv("PG_PASSWORD"),
            )
            yield conn
        except psycopg2.OperationalError as e:
            # Ошибка подключения — все строки партиции пропускаем с дефолтами
            print(f"DB connection failed: {e}")
            raise  # re-raise → Task упадёт → Spark повторит Task
        finally:
            if conn is not None and not conn.closed:
                try:
                    conn.close()
                except Exception:
                    pass  # игнорируем ошибки при закрытии

    # Используем контекстный менеджер
    with managed_connection() as conn:
        with conn.cursor() as cur:
            for row in rows_iterator:
                try:
                    cur.execute(
                        "SELECT tier FROM users WHERE id = %s",
                        (row.user_id,)
                    )
                    result = cur.fetchone()
                    yield row + ((result[0] if result else "unknown"),)

                except psycopg2.Error as e:
                    # Строчный уровень: ошибка одной строки не прерывает партицию
                    print(f"Error processing user_id={row.user_id}: {e}")
                    yield row + ("error",)  # дефолтное значение при ошибке строки

Уровни обработки ошибок

Важно понимать разницу между:

  • Ошибка соединения (conn.connect() failed): должна прокинуть исключение → Spark повторит Task
  • Ошибка отдельной строки (fetchone() вернул None): обрабатываем внутри, партиция продолжается
  • Timeout (statement_timeout exceeded): зависит от контекста, часто re-raise

8. Контроль числа партиций: квотирование нагрузки

Число активных Tasks = число параллельных соединений к внешней системе.

Как число партиций влияет на нагрузку

from pyspark.sql import SparkSession, functions as F

spark = SparkSession.builder.getOrCreate()
df = spark.table("bronze.events")

# Проверяем сколько партиций сейчас
current_partitions = df.rdd.getNumPartitions()
print(f"Текущих партиций: {current_partitions}")
# Например: 800 партиций → 800 одновременных соединений к PostgreSQL!

# ── Уменьшаем число партиций для контроля нагрузки ───────────────────

# coalesce: уменьшение БЕЗ полного shuffle (быстрее)
# Использовать когда: текущих партиций >> нужных, нет data skew
df_controlled = df.coalesce(50)  # max 50 одновременных соединений

# repartition: пересортировка через full shuffle (медленнее, но равномернее)
# Использовать когда: нужна точная равномерность размеров партиций
df_even = df.repartition(50)

print(f"После coalesce: {df_controlled.rdd.getNumPartitions()}")  # 50


def calculate_optimal_partitions(
    total_rows: int,
    target_rows_per_partition: int = 50_000,
    max_parallel_connections: int = 50,
    min_partitions: int = 10,
) -> int:
    """
    Рассчитывает оптимальное число партиций для mapPartitions с внешней системой.

    target_rows_per_partition: сколько строк обрабатывать за одно соединение
    max_parallel_connections: максимум одновременных соединений к БД/API
    """
    # По объёму данных
    partitions_by_data = max(min_partitions,
                             total_rows // target_rows_per_partition)

    # Ограничиваем нагрузку на внешнюю систему
    optimal = min(partitions_by_data, max_parallel_connections)

    return optimal


# Используем калькулятор
total = df.count()
n_partitions = calculate_optimal_partitions(
    total_rows=total,
    target_rows_per_partition=50_000,
    max_parallel_connections=30,  # PostgreSQL max_connections / 3 (треть под Spark)
)
print(f"Рекомендуемое число партиций: {n_partitions}")

df_final = df.coalesce(n_partitions) if df.rdd.getNumPartitions() > n_partitions \
           else df.repartition(n_partitions)

9. Ленивая природа итераторов: антипаттерн list(iterator)

Это самая коварная ошибка при работе с mapPartitions.

Почему list() убивает преимущества mapPartitions

# ❌ АНТИПАТТЕРН: материализация всей партиции в RAM
def process_partition_bad(rows_iterator):
    # Загружаем ВСЮ партицию в список!
    # Если партиция = 5 GB → OOM!
    all_rows = list(rows_iterator)  # ← ОПАСНО!

    # Теперь обрабатываем список
    results = []
    for row in all_rows:
        results.append(process_row(row))
    return results

# ✅ ПРАВИЛЬНО: стриминговая обработка через генератор
def process_partition_correct(rows_iterator):
    """
    Данные никогда полностью не загружаются в RAM.
    Каждая строка обрабатывается и освобождается из памяти.
    """
    # Инициализируем один раз
    conn = create_connection()

    try:
        # Обрабатываем ПОТОКОВО через генератор
        for row in rows_iterator:
            result = process_single_row(row, conn)
            yield result  # ← yield вместо return/append
    finally:
        conn.close()

Паттерн батчевой обработки без OOM

Когда нужно обрабатывать данные батчами (для Bulk API или executemany), используем генератор батчей:

def chunked_iterator(iterator, chunk_size: int):
    """
    Генератор: разбивает итератор на батчи заданного размера.
    Никогда не загружает весь итератор в память!
    """
    batch = []
    for item in iterator:
        batch.append(item)
        if len(batch) >= chunk_size:
            yield batch
            batch = []
    if batch:  # последний неполный батч
        yield batch


def process_partition_batched(rows_iterator):
    """
    Обработка батчами без материализации всей партиции.
    Каждый батч = chunk_size строк в RAM, не вся партиция.
    """
    session = create_http_session()

    try:
        for batch in chunked_iterator(rows_iterator, chunk_size=500):
            # В памяти только 500 строк одновременно!
            ids = [row.user_id for row in batch]

            response = session.post(
                "https://api.example.com/bulk",
                json={"ids": ids},
                timeout=10,
            )
            response.raise_for_status()

            api_data = {item["id"]: item for item in response.json()["users"]}

            for row in batch:
                user_info = api_data.get(row.user_id, {})
                yield (row.event_id, row.user_id, row.amount,
                       user_info.get("tier", "free"))
    finally:
        session.close()

10. Мониторинг в Spark UI и Production-паттерны

Что смотреть в Spark UI при mapPartitions

Вкладка Stages → конкретный Stage:

До mapPartitions (map + conn на строку):
  Duration: 2h 45min
  Task Time: очень высокий (в основном ожидание сети)
  Shuffle Read: 0 MB (нет shuffle)
  Executor Deserialize Time: высокий (сериализация замыканий)

После mapPartitions с Connection Pool:
  Duration: 8 min
  Task Time: значительно меньше
  Partition count: 50 (управляемо)
  Stragglers: нет (равномерная нагрузка)

На что обращать внимание:

  • Task Duration = время выполнения одного Task. При mapPartitions с JDBC должно быть примерно одинаковым для всех Task'ов (нет Stragglers).
  • Executor Deserialize Time = время десериализации Task'а. Должно быть минимальным — мы не передаём соединения по сети.
  • GC Time = при правильной реализации (без list()) должен быть низким.

Полная диагностика pipeline'а через метрики

from pyspark.sql import SparkSession
import time


def benchmark_mappartitions(
    spark: SparkSession,
    df,
    partition_function,
    label: str = "benchmark"
) -> dict:
    """
    Запускает функцию и собирает метрики производительности.
    """
    sc = spark.sparkContext

    t0 = time.time()

    # Запускаем action (count) для материализации
    result_rdd = df.rdd.mapPartitions(partition_function)
    result_count = result_rdd.count()

    elapsed = time.time() - t0

    # Получаем метрики из SparkContext
    status = sc.statusTracker()
    completed_stages = [
        s for s in status.getActiveStageIds()
    ]

    return {
        "label": label,
        "total_rows": result_count,
        "elapsed_seconds": elapsed,
        "rows_per_second": result_count / elapsed if elapsed > 0 else 0,
        "num_partitions": df.rdd.getNumPartitions(),
        "parallel_connections": df.rdd.getNumPartitions(),
    }

Итоговый production-ready шаблон

def create_mappartitions_pipeline(
    spark: SparkSession,
    source_table: str,
    target_table: str,
    max_db_connections: int = 30,
    batch_size: int = 5000,
) -> None:
    """
    Универсальный шаблон production mapPartitions pipeline:
    1. Читаем источник
    2. Контролируем число партиций (= соединений)
    3. Обогащаем через mapPartitions
    4. Записываем результат
    """
    from pyspark.sql.types import StructType, StructField, LongType, StringType

    # ── 1. Источник ───────────────────────────────────────────────────
    df = spark.table(source_table)
    original_partitions = df.rdd.getNumPartitions()
    print(f"Исходных партиций: {original_partitions}")

    # ── 2. Контроль нагрузки на внешнюю систему ───────────────────────
    if original_partitions > max_db_connections:
        df = df.coalesce(max_db_connections)
        print(f"Сокращено до {max_db_connections} партиций")
    elif original_partitions < 10:
        df = df.repartition(max_db_connections)
        print(f"Увеличено до {max_db_connections} партиций")

    # ── 3. Функция обработки партиции ────────────────────────────────
    def process(rows_iterator):
        import psycopg2

        conn = psycopg2.connect(
            host=os.getenv("PG_HOST"),
            database=os.getenv("PG_DB"),
            user=os.getenv("PG_USER"),
            password=os.getenv("PG_PASSWORD"),
        )

        batch = []
        results = []

        try:
            with conn.cursor() as cur:
                for row in rows_iterator:
                    batch.append(row.user_id)

                    if len(batch) >= batch_size:
                        # Bulk query вместо одиночных!
                        cur.execute(
                            "SELECT id, tier FROM users WHERE id = ANY(%s)",
                            (batch,)
                        )
                        tiers = {r[0]: r[1] for r in cur.fetchall()}

                        for user_id in batch:
                            yield (user_id, tiers.get(user_id, "free"))

                        batch = []

                # Последний неполный батч
                if batch:
                    cur.execute(
                        "SELECT id, tier FROM users WHERE id = ANY(%s)",
                        (batch,)
                    )
                    tiers = {r[0]: r[1] for r in cur.fetchall()}
                    for user_id in batch:
                        yield (user_id, tiers.get(user_id, "free"))

        finally:
            conn.close()

    # ── 4. Применяем и записываем ─────────────────────────────────────
    result_schema = StructType([
        StructField("user_id", LongType()),
        StructField("tier",    StringType()),
    ])

    result_df = df.rdd.mapPartitions(process).toDF(result_schema)

    result_df.write \
        .mode("overwrite") \
        .saveAsTable(target_table)

    print(f"Pipeline завершён: {target_table}")

Итоги: когда использовать mapPartitions

mapPartitions обязателен когда:

  • Инициализация клиента (JDBC, HTTP, gRPC, Redis) стоит дорого
  • Внешняя система имеет ограничение на число соединений
  • Нужна пакетная отправка данных (executemany, Bulk API)
  • Загружаете ML-модель для инференса (один раз на партицию)
  • Соединяетесь с Key-Value хранилищем (Redis, Cassandra)

Три золотых правила:

  1. Создавайте соединение внутри функции (на Executor'е), никогда на Driver'е
  2. Используйте try/finally для гарантированного закрытия соединений
  3. Никогда не делайте list(iterator) — обрабатывайте данные потоково через генераторы

Контролируйте нагрузку: coalesce(N) перед mapPartitions, где N = максимально допустимое число одновременных соединений к целевой системе.