How to properly mock functions parallelized with Dask?

Viewed 172

I'm trying to write tests for Dask parallelized code. Here's a minimal reproducing example to give an idea of what I'm trying to do:

from distributed import Client
from unittest.mock import patch
from scratch_import import foo

def test():
    client=Client()
    fut = client.submit(foo, a=1, b=2)
    ans = client.gather(fut)
    assert ans == 3

@patch("scratch_import.foo")
def test_patch(mock_foo):
    mock_foo.return_value = 1
    client=Client()
    fut = client.submit(foo, a=1, b=2)
    ans = client.gather(fut)
    assert ans == 1

which imports foo() from "scratch_import.py"

def foo(a, b):
    return a + b

The issue I keep running into and don't understand is that the patch is applied all the way up until dask submits the function to the dask workers, which seems to strip the patch in favor of the original function definition. What causes this behavior and is it expected? Is there a proper way to patch/mock parallelized function calls?

0 Answers
Related