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.
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)
Три золотых правила:
- Создавайте соединение внутри функции (на Executor'е), никогда на Driver'е
- Используйте
try/finallyдля гарантированного закрытия соединений - Никогда не делайте
list(iterator)— обрабатывайте данные потоково через генераторы
Контролируйте нагрузку: coalesce(N) перед mapPartitions, где N = максимально допустимое число одновременных соединений к целевой системе.