PySpark根据输入的表名和过滤条件生成 INSERT 语句

以下是一个完整的 PySpark 解决方案,用于根据输入的表名和过滤条件生成 INSERT 语句。代码考虑了不同数据类型的格式化、NULL 处理、单引号转义以及列名的安全引用,并提供了两种使用方式(收集到驱动端和分布式写入文件)。


函数实现

python 复制代码
from pyspark.sql import DataFrame
from pyspark.sql.types import (
    StringType, IntegerType, LongType, ShortType, ByteType,
    FloatType, DoubleType, DecimalType, BooleanType,
    DateType, TimestampType
)

def generate_insert_statements(
    table_name: str,
    filter_condition: str,
    target_table: str = None,
    batch_size: int = 10000
):
    """
    根据源表和过滤条件生成 INSERT 语句。

    参数:
        table_name (str): 源表名(可包含数据库名,如 `db.table`)。
        filter_condition (str): 过滤条件(SQL WHERE 子句,不含 WHERE 关键字)。
        target_table (str, optional): 目标表名,若不提供则使用源表名。
        batch_size (int, optional): 分批处理大小,用于控制内存占用(仅在分布式写入时使用)。

    返回:
        list: 包含 INSERT 语句的列表(当使用 collect 模式时)。
        或 None(当使用分布式写入模式时,结果直接写入文件)。
    """
    if target_table is None:
        target_table = table_name

    # 读取源表并应用过滤
    df = spark.table(table_name).filter(filter_condition)
    columns = df.columns
    schema = df.schema

    def format_value(value, data_type):
        """将 PySpark 值转换为 SQL 字面量字符串"""
        if value is None:
            return "NULL"
        # 字符串类型:单引号包裹,内部单引号转义为两个单引号
        if isinstance(data_type, StringType):
            return "'" + str(value).replace("'", "''") + "'"
        # 数值类型:直接转为字符串
        elif isinstance(data_type, (IntegerType, LongType, ShortType, ByteType,
                                    FloatType, DoubleType, DecimalType)):
            return str(value)
        # 布尔类型:使用 TRUE/FALSE(可根据目标数据库调整为 1/0)
        elif isinstance(data_type, BooleanType):
            return "TRUE" if value else "FALSE"
        # 日期和时间戳:转为字符串并加单引号(ISO 格式)
        elif isinstance(data_type, (DateType, TimestampType)):
            return "'" + str(value) + "'"
        # 其他类型(如二进制、复杂类型)默认按字符串处理,并给出警告
        else:
            print(f"警告:未处理的类型 {data_type},按字符串处理")
            return "'" + str(value).replace("'", "''") + "'"

    # 构建列名列表,使用反引号避免特殊字符问题
    column_list = ", ".join([f"`{c}`" for c in columns])

    # 使用迭代器逐行生成语句,避免一次性加载所有数据到内存
    def row_to_insert(row):
        values = []
        for col_name, field in zip(columns, schema.fields):
            val = row[col_name]
            values.append(format_value(val, field.dataType))
        values_str = ", ".join(values)
        return f"INSERT INTO `{target_table}` ({column_list}) VALUES ({values_str});"

    # 方式一:收集到驱动端返回列表(适用于小数据量)
    # return [row_to_insert(row) for row in df.collect()]

    # 方式二:使用 toLocalIterator 生成器,逐条返回(可配合外部循环写入文件)
    def generate():
        for row in df.toLocalIterator():
            yield row_to_insert(row)

    return generate()

使用示例

python 复制代码
# 生成 INSERT 语句(返回生成器,可迭代处理)
insert_gen = generate_insert_statements(
    table_name="sales_db.orders",
    filter_condition="order_date >= '2024-01-01' AND status = 'completed'",
    target_table="archive_db.orders"
)

# 示例:打印前 5 条
for i, stmt in enumerate(insert_gen):
    if i >= 5:
        break
    print(stmt)

# 示例:将所有语句写入文件
with open("/dbfs/tmp/insert_statements.sql", "w") as f:
    for stmt in generate_insert_statements("sales_db.orders", "status = 'completed'"):
        f.write(stmt + "\n")

注意事项

  1. 大数据量
    · 使用 collect() 会将所有数据加载到驱动节点,可能导致内存溢出。
    · 推荐使用 toLocalIterator() 生成器逐条处理,或使用分布式写入文件(见下文"分布式生成方式")。
  2. 复杂数据类型
    · 对于数组、结构体、映射等复杂类型,上述代码会按字符串处理,可能不符合目标数据库语法。
    · 建议针对具体类型扩展 format_value 函数(如序列化为 JSON 或使用目标数据库支持的格式)。
  3. SQL 方言差异
    · 布尔值 TRUE/FALSE 在 MySQL、PostgreSQL 中通用,但 SQL Server 使用 1/0,Oracle 使用 1/0 或 'Y'/'N'。
    · 日期和时间戳格式可能需要调整(如 Oracle 的 TO_DATE)。
  4. 转义与安全性
    · 字符串内部单引号已处理为两个单引号,符合 SQL 标准。
    · 表名和列名使用反引号包裹,避免与数据库关键字冲突。
  5. 性能优化
    · 若数据量极大,建议直接在 Spark 中构建 INSERT 语句字符串列,然后使用 df.write.text() 分布式写出,避免驱动程序成为瓶颈。

分布式生成方式(可选)

以下方法利用 Spark 的分布式能力,直接在各个分区生成 INSERT 语句并写入文件,适合大规模数据。

python 复制代码
from pyspark.sql.functions import col, lit, concat, when, isnull, format_string
from pyspark.sql.types import StringType

def generate_insert_statements_distributed(
    table_name: str,
    filter_condition: str,
    target_table: str = None,
    output_path: str = "/dbfs/tmp/insert_statements"
):
    if target_table is None:
        target_table = table_name

    df = spark.table(table_name).filter(filter_condition)
    columns = df.columns

    # 为每个列构建 SQL 字面量表达式
    value_exprs = []
    for c in columns:
        col_type = df.schema[c].dataType
        col_expr = col(c)
        if isinstance(col_type, StringType):
            # 字符串:添加单引号并转义
            expr = concat(lit("'"), col_expr.cast(StringType()), lit("'"))
            expr = expr.replace("'", "''")  # 注意:此方法在 Spark 中不可用,需使用 regexp_replace
            expr = concat(lit("'"), 
                          regexp_replace(col_expr.cast(StringType()), "'", "''"), 
                          lit("'"))
        elif isinstance(col_type, (IntegerType, LongType, ShortType, ByteType,
                                   FloatType, DoubleType, DecimalType)):
            expr = col_expr.cast(StringType())
        elif isinstance(col_type, BooleanType):
            expr = when(col_expr, lit("TRUE")).otherwise(lit("FALSE"))
        elif isinstance(col_type, (DateType, TimestampType)):
            expr = concat(lit("'"), col_expr.cast(StringType()), lit("'"))
        else:
            expr = concat(lit("'"), col_expr.cast(StringType()), lit("'"))
        # 处理 NULL:使用 coalesce 将 NULL 替换为字符串 'NULL'(不加引号)
        expr = when(isnull(col_expr), lit("NULL")).otherwise(expr)
        value_exprs.append(expr)

    # 构建完整的 INSERT 语句列
    column_list_str = ", ".join([f"`{c}`" for c in columns])
    values_expr = concat(lit("("), 
                         concat_ws(", ", *value_exprs), 
                         lit(")"))
    insert_stmt_col = concat(
        lit(f"INSERT INTO `{target_table}` ({column_list_str}) VALUES "),
        values_expr,
        lit(";")
    )

    # 选择该列并写入文本文件(每个分区一个文件)
    df.select(insert_stmt_col.alias("insert_sql")).write.mode("overwrite").text(output_path)

说明:

· 使用 regexp_replace 转义字符串中的单引号。

· 使用 when(isnull(col), ...) 区分 NULL 和其他值。

· 结果通过 write.text() 分布写入指定路径,每个分区生成一个文件,避免驱动端内存压力。


总结

以上代码提供了从 PySpark 表生成 INSERT 语句的完整实现。根据数据量大小和具体需求,可以选择简单模式(收集到驱动端)或分布式模式(直接写入文件)。使用时请根据目标数据库的语法调整数据类型格式化和转义规则。

相关推荐
GodSure091421 分钟前
Java单一职责原则SRP详解
java·python·单一职责原则
Mojitocean31 分钟前
Win开发环境配置(持续更新)
开发语言·python
千里码aicood39 分钟前
flask基于数字人的储粮知识问答原型系统研究与实现
后端·python·flask
E_ICEBLUE1 小时前
Python 实现 Markdown 转 Word、PDF:从转换到页面设置
python·pdf·word·markdown·格式转换
哒咩哒咩1291 小时前
Agent 智能体开发全攻略:从 ReAct 到企业级架构
python·langchain·fastapi
H愚公移山H1 小时前
Tengine2.4.1 + OpenSSL1.1.1w 全平台编译踩坑手册(CentOS7 x86_64|Linux x64|ARM64|2026实战复盘)
python
码云骑士1 小时前
120-视频理解-大模型视频分析-抽帧-自动字幕摘要-高光片段
python·音视频
xfan_me1 小时前
手机在网状态接口-空号查询-空号过滤API
数据库·人工智能·python·智能手机
测试秃头怪1 小时前
Postman中变量的使用
自动化测试·软件测试·python·测试工具·测试用例·接口测试·postman