Pandas read_sql设chunksize时如何从指定批次断点续传
问题背景
- 使用Pandas读取大体量SQL表并导出为
.csv文件,通过设置chunksize参数分批写入时,频繁出现数据库连接中断问题。 - 需要实现断点续传能力:从中断的批次位置继续导出,不需要重新加载已经保存完成的前置分块。
- 约束条件:所有待导出的SQL表均无
ID列,无法通过主键范围做数据过滤。 - 原有尝试:通过统计已保存的分块文件数量确定已完成批次号,尝试用
next方法跳过迭代器中已保存的分块,该方案未达到预期效果。
原有实现代码如下:
# Read tables to save tables = pd.read_csv('../data/to_extract.csv') # Check which tables and batch have been saved already all_files = os.listdir('../data') batch_size = 10000 def save_chunk(chunk, db, table_name, batch_no): chunk.to_csv(f'../data/{db}.{table_name}_{batch_no:04d}.csv.zip', compression={'method': 'zip', 'archive_name': f'{db}.{table_name}_{batch_no:04d}.csv'}, index=False, ) def get_and_save_data(row): table_name = row['TABLE_NAME'] db = row['TABLE_SCHEMA'] batch_no = len( [i for i in all_files if i.startswith(f"{db}.{table_name}")]) iterator = pd.read_sql_query(f"SELECT * FROM {db}.{table_name}", cnxn, chunksize=batch_size) nb_chunk_to_get = int(np.floor(row.CURRENT_ROWS / batch_size) - batch_no) if batch_no > 0: chunk = next((x for i, x in enumerate( iterator) if i == batch_no), None) ## Here I try to skip to the batch I want save_chunk(chunk, rename_dict, db, table_name, batch_no) batch_no += 1 for chunk in tqdm(iterator, total=nb_chunk_to_get, desc=f"{db}.{table_name}"): save_chunk(chunk, rename_dict, db, table_name, batch_no) batch_no += 1 rows_iter = (row for _, row in tables.iterrows()) with ThreadPoolExecutor(max_workers=2) as pool: tqdm(pool.map(get_and_save_data, rows_iter), total=len(tables), desc='overall')
原有方案失效原因
pd.read_sql_query返回的分块迭代器基于数据库服务端游标实现,逐批从数据库拉取数据。通过枚举迭代器跳过前N个块的逻辑,本质还是会把前N个块的所有数据从数据库传输到客户端,遇到连接不稳定的场景,拉取前置块的过程中就会触发连接中断,根本无法到达目标批次,没有实现真正的跳过。- 多线程场景下全局共用同一个
cnxn数据库连接,不同线程的游标操作会互相干扰,本身就是连接频繁中断的核心诱因之一。 - 已完成批次计数逻辑存在缺陷:如果分块写入中途失败(比如压缩过程中断、连接断开),残留的损坏文件也会被统计为已完成,导致偏移量计算错误。
修复方案
核心逻辑是改用数据库服务端分页直接跳过已导出的行数,不需要在客户端拉取、遍历前置分块,从根源上减少无效传输和连接占用。同时调整连接管理、文件校验逻辑,提升稳定性。
注意:不同数据库分页语法有差异,以下示例以SQL Server/PostgreSQL支持的
OFFSET ... FETCH NEXT语法为例,MySQL可替换为LIMIT {offset}, {batch_size}。
修正后的实现代码:
import os import zipfile import numpy as np import pandas as pd from tqdm import tqdm from concurrent.futures import ThreadPoolExecutor # 替换为实际使用的数据库驱动连接方法,例如sqlalchemy.create_engine from your_db_module import create_db_connection tables = pd.read_csv('../data/to_extract.csv') batch_size = 10000 save_dir = '../data' def is_valid_zip(file_path): """校验压缩分块文件是否完整,避免损坏文件被误判为已完成""" try: with zipfile.ZipFile(file_path, 'r') as zf: return zf.testzip() is None except Exception: return False def save_chunk(chunk, db, table_name, batch_no): file_path = f'{save_dir}/{db}.{table_name}_{batch_no:04d}.csv.zip' # 先写临时文件,写入完成后再重命名为正式文件名,避免中途中断产生损坏文件 temp_path = file_path + '.tmp' chunk.to_csv(temp_path, compression={'method': 'zip', 'archive_name': f'{db}.{table_name}_{batch_no:04d}.csv'}, index=False, ) # 校验文件完整性后重命名 if is_valid_zip(temp_path): os.rename(temp_path, file_path) else: os.remove(temp_path) raise IOError(f"Batch {batch_no} of {db}.{table_name} save failed, file corrupted") def get_and_save_data(row): table_name = row['TABLE_NAME'] db = row['TABLE_SCHEMA'] # 每个线程单独创建数据库连接,避免多线程共用连接冲突 cnxn = create_db_connection() total_rows = row.CURRENT_ROWS # 统计有效已完成批次数量 existing_files = [i for i in os.listdir(save_dir) if i.startswith(f"{db}.{table_name}") and i.endswith('.csv.zip')] valid_batch_nos = [] for f in existing_files: batch_str = f.split('_')[-1].split('.')[0] try: bn = int(batch_str) if is_valid_zip(os.path.join(save_dir, f)): valid_batch_nos.append(bn) else: # 删除损坏文件 os.remove(os.path.join(save_dir, f)) except ValueError: continue # 已完成批次从0开始连续计数,取最大连续编号作为起始点 batch_no = 0 while batch_no in valid_batch_nos: batch_no += 1 total_batch = int(np.ceil(total_rows / batch_size)) # 从断点位置开始逐批拉取,服务端直接跳过已导出的行 for current_bn in tqdm(range(batch_no, total_batch), initial=batch_no, total=total_batch, desc=f"{db}.{table_name}"): offset = current_bn * batch_size # 服务端分页查询,直接拉取目标批次数据,不需要传输前置块 sql = f""" SELECT * FROM {db}.{table_name} OFFSET {offset} ROWS FETCH NEXT {batch_size} ROWS ONLY """ chunk = pd.read_sql_query(sql, cnxn) if len(chunk) == 0: break save_chunk(chunk, db, table_name, current_bn) cnxn.close() rows_iter = (row for _, row in tables.iterrows()) with ThreadPoolExecutor(max_workers=2) as pool: list(tqdm(pool.map(get_and_save_data, rows_iter), total=len(tables), desc='overall'))
额外优化点
- 临时文件写入机制:所有分块先写
.tmp后缀临时文件,校验完整性后再重命名为正式文件,彻底避免中断产生的损坏文件被误统计。 - 连接隔离:每个工作线程独立创建数据库连接,从根源上避免多线程共用连接导致的状态混乱、连接中断问题。
- 批次连续性校验:即使存在编号不连续的残留文件,也会从第一个缺失的批次开始导出,不会出现漏数、重复数问题。
内容的提问来源于stack exchange,提问作者Roger
相关产品推荐
相关产品推荐

