Skip to content
52 changes: 49 additions & 3 deletions tests/cloudpickle_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -44,6 +44,7 @@
from cloudpickle.cloudpickle import _find_module, _make_empty_cell, cell_set

from .testutils import subprocess_pickle_echo
from .testutils import assert_run_python_script


HAVE_WEAKSET = hasattr(weakref, 'WeakSet')
Expand Down Expand Up @@ -287,19 +288,21 @@ def some_method(self, x):
clone_class = pickle_depickle(SomeClass, protocol=self.protocol)
self.assertEqual(clone_class(1).one(), 1)
self.assertEqual(clone_class(5).some_method(41), 7)
clone_class = subprocess_pickle_echo(SomeClass)
clone_class = subprocess_pickle_echo(SomeClass, protocol=self.protocol)
self.assertEqual(clone_class(5).some_method(41), 7)

# pickle the class instances
self.assertEqual(pickle_depickle(SomeClass(1)).one(), 1)
self.assertEqual(pickle_depickle(SomeClass(5)).some_method(41), 7)
new_instance = subprocess_pickle_echo(SomeClass(5))
new_instance = subprocess_pickle_echo(SomeClass(5),
protocol=self.protocol)
self.assertEqual(new_instance.some_method(41), 7)

# pickle the method instances
self.assertEqual(pickle_depickle(SomeClass(1).one)(), 1)
self.assertEqual(pickle_depickle(SomeClass(5).some_method)(41), 7)
new_method = subprocess_pickle_echo(SomeClass(5).some_method)
new_method = subprocess_pickle_echo(SomeClass(5).some_method,
protocol=self.protocol)
self.assertEqual(new_method(41), 7)

def test_partial(self):
Expand Down Expand Up @@ -748,6 +751,49 @@ def test_builtin_type__new__(self):
for t in list, tuple, set, frozenset, dict, object:
self.assertTrue(pickle_depickle(t.__new__) is t.__new__)

def test_interactively_defined_function(self):
# Check that callables defined in the __main__ module of a Python
# script (or jupyter kernel) can be pickled / unpickled / executed.
code = """\
from testutils import subprocess_pickle_echo

CONSTANT = 42

class Foo(object):

def method(self, x):
return x


def f1():
return Foo


def f2(x):
return Foo().method(x)


def f3():
return Foo().method(CONSTANT)


cloned = subprocess_pickle_echo(lambda x: x**2, protocol={protocol})
assert cloned(3) == 9

cloned = subprocess_pickle_echo(Foo, protocol={protocol})
assert cloned().method(2) == Foo().method(2)

cloned = subprocess_pickle_echo(f1, protocol={protocol})
assert cloned()().method('a') == f1()().method('a')

cloned = subprocess_pickle_echo(f2, protocol={protocol})
assert cloned(2) == f2(2)

cloned = subprocess_pickle_echo(f3, protocol={protocol})
assert cloned() == f3()
""".format(protocol=self.protocol)
assert_run_python_script(code)

@pytest.mark.skipif(sys.version_info >= (3, 0),
reason="hardcoded pickle bytes for 2.7")
def test_function_pickle_compat_0_4_0(self):
Expand Down
49 changes: 44 additions & 5 deletions tests/testutils.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,11 @@
import sys
import os
import tempfile
from subprocess import Popen
from subprocess import check_output
from subprocess import PIPE
from subprocess import STDOUT
from subprocess import CalledProcessError

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Can you import all those names at once? :-)


from cloudpickle import dumps
from pickle import loads
Expand All @@ -16,7 +20,7 @@ class TimeoutExpired(Exception):
timeout_supported = False


def subprocess_pickle_echo(input_data):
def subprocess_pickle_echo(input_data, protocol=None):
"""Echo function with a child Python process

Pickle the input data into a buffer, send it to a subprocess via
Expand All @@ -27,8 +31,8 @@ def subprocess_pickle_echo(input_data):
[1, 'a', None]

"""
pickled_input_data = dumps(input_data)
cmd = [sys.executable, __file__]
pickled_input_data = dumps(input_data, protocol=protocol)
cmd = [sys.executable, __file__] # run then pickle_echo() in __main__
cwd = os.getcwd()
proc = Popen(cmd, stdin=PIPE, stdout=PIPE, stderr=PIPE, cwd=cwd)
try:
Expand All @@ -48,7 +52,7 @@ def subprocess_pickle_echo(input_data):
raise RuntimeError(message)


def pickle_echo(stream_in=None, stream_out=None):
def pickle_echo(stream_in=None, stream_out=None, protocol=None):
"""Read a pickle from stdin and pickle it back to stdout"""
if stream_in is None:
stream_in = sys.stdin
Expand All @@ -64,9 +68,44 @@ def pickle_echo(stream_in=None, stream_out=None):
input_bytes = stream_in.read()
stream_in.close()
unpickled_content = loads(input_bytes)
stream_out.write(dumps(unpickled_content))
stream_out.write(dumps(unpickled_content, protocol=protocol))
stream_out.close()


def assert_run_python_script(source_code, timeout=5):
"""Utility to help check pickleability of objects defined in __main__

The script provided in the source code should return 0 and not print
anything on stderr or stdout.
"""
fd, source_file = tempfile.mkstemp(suffix='_src_test_cloudpickle.py')
try:
with open(fd, 'wb') as f:
f.write(source_code.encode('utf-8'))

cmd = [sys.executable, source_file]
pythonpath = "{cwd}/tests:{cwd}".format(cwd=os.getcwd())

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Can't you use __file__ or something? getcwd sounds fragile.

kwargs = {
'cwd': os.getcwd(),
'stderr': STDOUT,
'env': {'PYTHONPATH': pythonpath},
}
if timeout_supported:
kwargs['timeout'] = timeout
try:
try:
out = check_output(cmd, **kwargs)
except CalledProcessError as e:
raise RuntimeError(u"script errored with output:\n%s"
% e.output.decode('utf-8'))
if out != b"":
raise AssertionError(out.decode('utf-8'))
except TimeoutExpired as e:
raise RuntimeError(u"script timeout, output so far:\n%s"
% e.output.decode('utf-8'))
finally:
os.unlink(source_file)


if __name__ == '__main__':
pickle_echo()