【问题标题】:Mocking SparkSession for unit testing模拟 SparkSession 进行单元测试
【发布时间】:2018-09-04 03:45:42
【问题描述】:

我的 spark 应用程序中有一个从 MySQL 数据库加载数据的方法。该方法看起来像这样。

trait DataManager {

val session: SparkSession

def loadFromDatabase(input: Input): DataFrame = {
            session.read.jdbc(input.jdbcUrl, s"(${input.selectQuery}) T0",
              input.columnName, 0L, input.maxId, input.parallelism, input.connectionProperties)
    }
}

该方法除了执行jdbc 方法并从数据库中加载数据之外什么都不做。我该如何测试这种方法?标准方法是创建对象session 的模拟,它是SparkSession 的一个实例。但由于 SparkSession 有一个私有构造函数,我无法使用 ScalaMock 模拟它。

这里的主要问题是我的函数是一个纯副作用函数(副作用是从关系数据库中提取数据),鉴于我在模拟 SparkSession 时遇到问题,我如何对这个函数进行单元测试。

那么有什么方法可以模拟 SparkSession 或任何其他比模拟测试此方法更好的方法吗?

【问题讨论】:

  • @himanshuIIITian 这不是那个问题的重复。我的问题非常具体到一个用例,我的方法只从数据库加载数据,如果可能的话,我如何使用模拟或任何其他方法对其进行测试。您链接的问题没有谈论如何模拟它或如何处理非常具体的场景..
  • 好的!我认为它与它相似。为混乱道歉。
  • 你到底想测试什么?如果查询可以执行?老实说,我不会测试这个方法,因为它不包含你实现的任何逻辑(没有冒犯)。您只是在运行 spark 提供的一些逻辑 - 这应该在他们这边进行测试。如果您仍然想对此进行测试,您可以使用嵌入式数据库。
  • 还是我把你的问题弄错了,你问谁来创建一个 spark-session?给你:SparkSession.builder().getOrCreate()

标签: scala unit-testing apache-spark mocking scalamock


【解决方案1】:

您可以使用 mockito scala 来模拟 SparkSession,如 this article 所示。

【讨论】:

    【解决方案2】:

    在您的情况下,我建议不要模拟 SparkSession。这或多或少会模拟整个功能(无论如何您都可以这样做)。如果您想测试此功能,我的建议是运行嵌入式数据库(如H2)并使用真正的 SparkSession。为此,您需要将 SparkSession 提供给您的DataManager

    未经测试的草图:

    您的代码:

    class DataManager (session: SparkSession) {
             def loadFromDatabase(input: Input): DataFrame = {
                session.read.jdbc(input.jdbcUrl, s"(${input.selectQuery}) T0",
                input.columnName, 0L, input.maxId, input.parallelism, input.connectionProperties)
             }
        }
    

    你的测试用例:

    class DataManagerTest extends FunSuite with BeforeAndAfter {
      override def beforeAll() {
        Connection conn = DriverManager.getConnection("jdbc:h2:~/test", "sa", "");
        // your insert statements goes here
        conn.close()
      }
    
      test ("should load data from database") {
        val dm = DataManager(SparkSession.builder().getOrCreate())
        val input = Input(jdbcUrl = "jdbc:h2:~/test", selectQuery="SELECT whateveryounedd FROM whereeveryouputit ")
        val expectedData = dm.loadFromDatabase(input)
        assert(//expectedData)
      }
    }
    

    【讨论】:

      猜你喜欢
      • 2013-11-12
      • 2018-11-21
      • 2018-10-19
      • 2013-07-24
      • 1970-01-01
      • 1970-01-01
      • 1970-01-01
      • 2018-04-04
      • 2020-04-06
      相关资源
      最近更新 更多