from pyspark.sql import SparkSession
from pyspark.sql.functions import col, coalesce, trim, when, lit, sum
from pyspark.sql.types import StringType, NumericType
# 在 Databricks 中,spark 会话通常已经存在,无需重新创建
# 如果需要显式创建,使用:
# spark = SparkSession.builder.getOrCreate()
# 配置参数
database_name = "your_database" # 替换为实际数据库名
result_list = []
# 获取数据库中的所有表和视图(包括 Delta 表)
tables = spark.catalog.listTables(database_name)
for table in tables:
table_name = table.name
full_table_name = f"{database_name}.{table_name}"
try:
# 读取表或视图(Delta 表和普通表均可)
df = spark.table(full_table_name)
# 快速判断是否为空表,避免不必要的缓存
total_count = df.count()
if total_count == 0:
continue
# 缓存数据以便多次引用(实际上我们只需扫描一次,缓存并非必需)
df.cache()
# 构建所有字段的聚合表达式(一次性计算)
agg_exprs = []
field_meta = [] # 用于记录每个字段的类型和名称
for field in df.schema.fields:
col_name = field.name
col_type = field.dataType
if isinstance(col_type, StringType):
# 字符串类型:统计 null 或 trim 后为空字符串
modified_col = trim(coalesce(col(col_name), lit("")))
condition = (modified_col == lit(""))
count_expr = sum(when(condition, 1).otherwise(0)).alias(f"cnt_{col_name}")
elif isinstance(col_type, NumericType):
# 数值类型:统计 null 或零值
modified_col = coalesce(col(col_name), lit(0))
condition = (modified_col == lit(0))
count_expr = sum(when(condition, 1).otherwise(0)).alias(f"cnt_{col_name}")
else:
# 其他类型:仅统计 null
condition = col(col_name).isNull()
count_expr = sum(when(condition, 1).otherwise(0)).alias(f"cnt_{col_name}")
agg_exprs.append(count_expr)
field_meta.append((col_name, str(col_type)))
# 执行一次聚合,获取所有字段的统计值
stats_row = df.agg(*agg_exprs).collect()[0]
# 整理结果
for col_name, col_type in field_meta:
stat_count = stats_row[f"cnt_{col_name}"]
percentage = round((stat_count / total_count) * 100, 2) if total_count > 0 else 0.0
result_list.append((
database_name,
table_name,
col_name,
col_type,
stat_count,
total_count,
float(percentage)
))
df.unpersist() # 释放缓存
except Exception as e:
print(f"Error processing table {table_name}: {str(e)}")
continue
# 创建结果 DataFrame
result_columns = [
"database_name",
"table_name",
"column_name",
"column_type",
"stat_count",
"total_rows",
"percentage"
]
result_df = spark.createDataFrame(result_list, result_columns)
# 显示结果
result_df.show(truncate=False)
# 可选:将结果保存到 Delta 表
# result_df.write.format("delta").mode("overwrite").saveAsTable("your_audit_table")