以下是一个完整的 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")
注意事项
- 大数据量
· 使用 collect() 会将所有数据加载到驱动节点,可能导致内存溢出。
· 推荐使用 toLocalIterator() 生成器逐条处理,或使用分布式写入文件(见下文"分布式生成方式")。 - 复杂数据类型
· 对于数组、结构体、映射等复杂类型,上述代码会按字符串处理,可能不符合目标数据库语法。
· 建议针对具体类型扩展 format_value 函数(如序列化为 JSON 或使用目标数据库支持的格式)。 - SQL 方言差异
· 布尔值 TRUE/FALSE 在 MySQL、PostgreSQL 中通用,但 SQL Server 使用 1/0,Oracle 使用 1/0 或 'Y'/'N'。
· 日期和时间戳格式可能需要调整(如 Oracle 的 TO_DATE)。 - 转义与安全性
· 字符串内部单引号已处理为两个单引号,符合 SQL 标准。
· 表名和列名使用反引号包裹,避免与数据库关键字冲突。 - 性能优化
· 若数据量极大,建议直接在 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 语句的完整实现。根据数据量大小和具体需求,可以选择简单模式(收集到驱动端)或分布式模式(直接写入文件)。使用时请根据目标数据库的语法调整数据类型格式化和转义规则。