如何Mock PySpark UDF?UDF测试断言失败问题求助
PySpark UDF测试断言失败及Mock解决方案
问题背景
有一个包含email和expected列的DataFrame,定义了如下clean_email PySpark UDF用于清洗邮箱地址:
@udf(returnType=StringType()) def clean_email(email): try: regex = r'\b[A-Za-z0-9._%+-]+@[A-Za-z0-9.-]+\.[A-Z|a-z]{2,}\b' replace={"%20":"" , "//":"" ,"/":""} for i in replace: if i in email: email=email.replace(i,"") if email is not None or '.jpg' in email or email.startswith('http'): if email.endswith('.') : email=email[:len(email)-1] return ''.join(e for e in email if (e.isalnum() or e in ['.', '@','-','_'])) else: return "" except Exception as x: print("Error Occured in email udf, Error: " + str(x))
编写测试代码对比expected列与UDF生成的Curated列:
df1=context.spark.read.option("header",True).csv("./test/input/11-udf-test/Book1.csv",schema=schema) df3=sorted(df1.select(col("expected"))).collect() df2=df1.withColumn("Curated", dataclean.clean_email(col("email"))) df4=sorted(df2.select(col("Curated"))).collect() assert df3== df4
测试时出现断言错误:
test/udf_test.py:32: AssertionError ====================================================== short test summary info ======================================================= FAILED test/udf_test.py::test_upper - AssertionError: assert [Row(expected...xpected=None)] == [Row(Curated=...t@gmail.com')] ==================================================== 1 failed in
一、修复UDF逻辑错误(解决断言失败)
断言失败的核心原因是clean_email UDF的逻辑存在两处关键错误:
1. 无效邮箱判断逻辑颠倒
原代码中if email is not None or '.jpg' in email or email.startswith('http')的逻辑完全错误:
- 需求应为:仅当邮箱不为空、不是图片链接、不是HTTP链接时,才执行清洗,否则返回空字符串
- 原逻辑会让无效邮箱(如http开头、jpg后缀)进入清洗分支,生成不符合预期的结果
修正后的条件判断:
if email is not None and '.jpg' not in email and not email.startswith('http'):
2. 空值处理漏洞
当email为None时,执行i in email会直接触发异常,需先判断email是否为None再执行替换操作。
完整修正后的UDF
@udf(returnType=StringType()) def clean_email(email): try: # 先过滤无效邮箱 if email is None or '.jpg' in email or email.startswith('http'): return "" # 执行特殊字符替换 replace = {"%20": "", "//": "", "/": ""} for key in replace: if key in email: email = email.replace(key, replace[key]) # 移除末尾多余的点号 if email.endswith('.'): email = email[:-1] # 保留合法邮箱字符 return ''.join(e for e in email if (e.isalnum() or e in ['.', '@', '-', '_'])) except Exception as x: print(f"Error Occured in email udf, Error: {str(x)}") return "" # 异常时返回空,避免数据报错
二、Mock PySpark UDF进行单元测试
为了脱离Spark环境快速测试,或在集成测试中替换UDF逻辑,可通过以下两种方式Mock:
方法1:抽离核心逻辑为普通函数,测试纯Python代码
将UDF的核心清洗逻辑抽成独立的普通函数,直接测试该函数,UDF仅作为包装层:
# 抽离核心清洗逻辑 def _clean_email_logic(email): if email is None or '.jpg' in email or email.startswith('http'): return "" replace = {"%20": "", "//": "", "/": ""} for key in replace: if key in email: email = email.replace(key, replace[key]) if email.endswith('.'): email = email[:-1] return ''.join(e for e in email if (e.isalnum() or e in ['.', '@', '-', '_'])) # UDF包装层 @udf(returnType=StringType()) def clean_email(email): try: return _clean_email_logic(email) except Exception as x: print(f"Error Occured in email udf, Error: {str(x)}") return ""
针对核心逻辑编写单元测试,无需启动Spark:
import unittest class TestEmailCleanLogic(unittest.TestCase): def test_valid_email(self): self.assertEqual(_clean_email_logic("test@example.com"), "test@example.com") self.assertEqual(_clean_email_logic("test%20name@example.com"), "testname@example.com") self.assertEqual(_clean_email_logic("test//name@example.com."), "testname@example.com") def test_invalid_email(self): self.assertEqual(_clean_email_logic("http://example.com"), "") self.assertEqual(_clean_email_logic("image.jpg"), "") self.assertEqual(_clean_email_logic(None), "") if __name__ == '__main__': unittest.main()
方法2:使用unittest.mock Patch UDF
如果需要在Spark测试中替换UDF实现,可通过mock.patch临时替换:
from unittest.mock import patch import pytest from pyspark.sql import SparkSession from pyspark.sql.functions import col @pytest.fixture(scope="session") def spark(): return SparkSession.builder.master("local[1]").getOrCreate() def test_clean_email_udf(spark): # 构造测试数据 test_data = [ ("test@example.com", "test@example.com"), ("test%20name@example.com.", "testname@example.com"), ("http://example.com", ""), (None, "") ] df = spark.createDataFrame(test_data, ["email", "expected"]) # Mock UDF的实现 with patch("dataclean.clean_email") as mock_udf: # 让Mock的UDF返回预期值 mock_udf.side_effect = lambda x: dict(test_data)[x] if x in dict(test_data) else "" df = df.withColumn("Curated", mock_udf(col("email"))) curated = sorted(df.select("Curated").collect()) expected = sorted(df.select("expected").collect()) assert curated == expected
内容的提问来源于stack exchange,提问作者Xi12
相关产品推荐
相关产品推荐

