以下是一个完整的 PySpark 解决方案,用于根据输入的表名和过滤条件生成 INSERT 语句。代码考虑了不同数据类型的格式化、NULL 处理、单引号转义以及列名的安全引用,并提供了两种使用方式(收集到驱动端和分布式写入文件)。
函数实现
frompyspark.sqlimportDataFramefrompyspark.sql.typesimport(StringType,IntegerType,LongType,ShortType,ByteType,FloatType,DoubleType,DecimalType,BooleanType,DateType,TimestampType)defgenerate_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(当使用分布式写入模式时,结果直接写入文件)。 """iftarget_tableisNone:target_table=table_name# 读取源表并应用过滤df=spark.table(table_name).filter(filter_condition)columns=df.columns schema=df.schemadefformat_value(value,data_type):"""将 PySpark 值转换为 SQL 字面量字符串"""ifvalueisNone:return"NULL"# 字符串类型:单引号包裹,内部单引号转义为两个单引号ifisinstance(data_type,StringType):return"'"+str(value).replace("'","''")+"'"# 数值类型:直接转为字符串elifisinstance(data_type,(IntegerType,LongType,ShortType,ByteType,FloatType,DoubleType,DecimalType)):returnstr(value)# 布尔类型:使用 TRUE/FALSE(可根据目标数据库调整为 1/0)elifisinstance(data_type,BooleanType):return"TRUE"ifvalueelse"FALSE"# 日期和时间戳:转为字符串并加单引号(ISO 格式)elifisinstance(data_type,(DateType,TimestampType)):return"'"+str(value)+"'"# 其他类型(如二进制、复杂类型)默认按字符串处理,并给出警告else:print(f"警告:未处理的类型{data_type},按字符串处理")return"'"+str(value).replace("'","''")+"'"# 构建列名列表,使用反引号避免特殊字符问题column_list=", ".join([f"`{c}`"forcincolumns])# 使用迭代器逐行生成语句,避免一次性加载所有数据到内存defrow_to_insert(row):values=[]forcol_name,fieldinzip(columns,schema.fields):val=row[col_name]values.append(format_value(val,field.dataType))values_str=", ".join(values)returnf"INSERT INTO `{target_table}` ({column_list}) VALUES ({values_str});"# 方式一:收集到驱动端返回列表(适用于小数据量)# return [row_to_insert(row) for row in df.collect()]# 方式二:使用 toLocalIterator 生成器,逐条返回(可配合外部循环写入文件)defgenerate():forrowindf.toLocalIterator():yieldrow_to_insert(row)returngenerate()使用示例
# 生成 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 条fori,stmtinenumerate(insert_gen):ifi>=5:breakprint(stmt)# 示例:将所有语句写入文件withopen("/dbfs/tmp/insert_statements.sql","w")asf:forstmtingenerate_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 语句并写入文件,适合大规模数据。
frompyspark.sql.functionsimportcol,lit,concat,when,isnull,format_stringfrompyspark.sql.typesimportStringTypedefgenerate_insert_statements_distributed(table_name:str,filter_condition:str,target_table:str=None,output_path:str="/dbfs/tmp/insert_statements"):iftarget_tableisNone:target_table=table_name df=spark.table(table_name).filter(filter_condition)columns=df.columns# 为每个列构建 SQL 字面量表达式value_exprs=[]forcincolumns:col_type=df.schema[c].dataType col_expr=col(c)ifisinstance(col_type,StringType):# 字符串:添加单引号并转义expr=concat(lit("'"),col_expr.cast(StringType()),lit("'"))expr=expr.replace("'","''")# 注意:此方法在 Spark 中不可用,需使用 regexp_replaceexpr=concat(lit("'"),regexp_replace(col_expr.cast(StringType()),"'","''"),lit("'"))elifisinstance(col_type,(IntegerType,LongType,ShortType,ByteType,FloatType,DoubleType,DecimalType)):expr=col_expr.cast(StringType())elifisinstance(col_type,BooleanType):expr=when(col_expr,lit("TRUE")).otherwise(lit("FALSE"))elifisinstance(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}`"forcincolumns])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 语句的完整实现。根据数据量大小和具体需求,可以选择简单模式(收集到驱动端)或分布式模式(直接写入文件)。使用时请根据目标数据库的语法调整数据类型格式化和转义规则。