diff --git a/CHANGES.md b/CHANGES.md index a6b0b443..3d33727d 100644 --- a/CHANGES.md +++ b/CHANGES.md @@ -1,6 +1,15 @@ In development ============== +- Fix pickling of Python 3.14 lazy annotation (`__annotate__`) functions that + have to be pickled by value, for instance when `functools.update_wrapper` + copies the `__annotate__` function of a method onto a wrapper instance. Such + functions close over the namespace of the class they were defined in, which + can hold unpicklable objects such as `_abc_impl`. Their annotations are now + evaluated at pickling time instead, as is already the case for the + annotations of dynamic functions. ([issue #585]( + https://github.com/cloudpipe/cloudpickle/issues/585)) + - Make pickling of functions depending on globals in notebook more deterministic. ([PR#560](https://github.com/cloudpipe/cloudpickle/pull/560)) diff --git a/cloudpickle/cloudpickle.py b/cloudpickle/cloudpickle.py index 08882306..1e5c400e 100644 --- a/cloudpickle/cloudpickle.py +++ b/cloudpickle/cloudpickle.py @@ -754,6 +754,50 @@ def _function_getstate(func): return state, slotstate +def _make_eager_annotate(annotations): + """Rebuild a PEP 649 ``__annotate__`` function from evaluated annotations. + + The original annotate function is a closure over the namespace in which the + annotated object was defined (for methods, the class namespace). That + namespace can hold unpicklable objects such as ``_abc_impl``, so when the + annotate function has to be pickled by value we snapshot the annotations + instead, mirroring what ``_function_getstate`` already does for the + ``__annotations__`` of dynamic functions. + """ + + def __annotate__(format, /): + from annotationlib import Format, annotations_to_string + + if format == Format.STRING: + return annotations_to_string(annotations) + elif format in (Format.VALUE, Format.FORWARDREF): + return dict(annotations) + raise NotImplementedError(format) + + return __annotate__ + + +def _eager_annotate_reduce(func): + """Reducer for PEP 649 ``__annotate__`` functions pickled by value. + + Returns NotImplemented when the annotations cannot be evaluated eagerly + (e.g. they contain forward references), in which case the generic dynamic + function reducer is used instead. + """ + if sys.version_info < (3, 14) or func.__name__ != "__annotate__": + return NotImplemented + + from annotationlib import Format, call_annotate_function + + try: + annotations = call_annotate_function(func, Format.VALUE) + except Exception: + return NotImplemented + if not isinstance(annotations, dict): + return NotImplemented + return _make_eager_annotate, (annotations,) + + def _class_getstate(obj): clsdict = _extract_class_dict(obj) clsdict.pop("__weakref__", None) @@ -1286,8 +1330,10 @@ def _function_reduce(self, obj): """ if _should_pickle_by_reference(obj): return NotImplemented - else: - return self._dynamic_function_reduce(obj) + reduce = _eager_annotate_reduce(obj) + if reduce is not NotImplemented: + return reduce + return self._dynamic_function_reduce(obj) def _function_getnewargs(self, func): code = func.__code__ diff --git a/tests/cloudpickle_test.py b/tests/cloudpickle_test.py index e2097d1c..1702551c 100644 --- a/tests/cloudpickle_test.py +++ b/tests/cloudpickle_test.py @@ -2707,6 +2707,99 @@ class C(abc.ABC): c2 = C2() assert isinstance(c2, C2) + @pytest.mark.skipif( + sys.version_info < (3, 14), + reason="functools.update_wrapper copies __annotate__ starting in Python 3.14", + ) + def test_update_wrapper_with_annotated_abc_method(self): + # see https://github.com/cloudpipe/cloudpickle/issues/585 + # functools.update_wrapper copies the (lazy) __annotate__ function of + # the wrapped callable onto the wrapper instance on Python 3.14+. For a + # method, that function closes over the namespace of the defining + # class, which for an ABC holds an unpicklable _abc_impl object. + class FuncWrapper: + def __init__(self, function): + self.function = function + functools.update_wrapper(self, self.function) + + def __call__(self, *args, **kwargs): + return self.function(*args, **kwargs) + + class AbstractClass(abc.ABC): + a: int + + def method(self, arg: str) -> str: + return arg.upper() + + wrapped = FuncWrapper(AbstractClass().method) + + wrapped_clone = pickle_depickle(wrapped, protocol=self.protocol) + + assert wrapped_clone("abc") == "ABC" + assert wrapped_clone.__name__ == "method" + assert wrapped_clone.__wrapped__.__annotations__ == {"arg": str, "return": str} + if sys.version_info >= (3, 14): + import annotationlib + + # The annotations copied onto the wrapper survive the roundtrip, + # evaluated eagerly at pickling time. + assert annotationlib.get_annotations(wrapped_clone) == { + "arg": str, + "return": str, + } + assert annotationlib.get_annotations( + wrapped_clone, format=annotationlib.Format.STRING + ) == {"arg": "str", "return": "str"} + + @pytest.mark.skipif( + sys.version_info < (3, 14), + reason="PEP 649 lazy annotations require Python 3.14+", + ) + def test_pickle_annotate_function_of_method(self): + # A PEP 649 __annotate__ function that has to be pickled by value is + # snapshotted into an eagerly evaluated one, see issue #585. + import annotationlib + + class AbstractClass(abc.ABC): + a: int + + def method(self, arg: str) -> "str": + return arg.upper() + + annotate_clone = pickle_depickle( + AbstractClass.method.__annotate__, protocol=self.protocol + ) + assert annotate_clone(annotationlib.Format.VALUE) == { + "arg": str, + "return": "str", + } + assert annotate_clone(annotationlib.Format.STRING) == { + "arg": "str", + "return": "str", + } + + @pytest.mark.skipif( + sys.version_info < (3, 14), + reason="PEP 649 lazy annotations require Python 3.14+", + ) + def test_unresolvable_annotations_stay_lazy(self): + # When the annotations cannot be evaluated eagerly, the annotate + # function is pickled by value as any other function. + import annotationlib + + class FuncWrapper: + def __init__(self, function): + self.function = function + functools.update_wrapper(self, self.function) + + def func(arg: "UndefinedName") -> int: # noqa: F821 + return 0 + + clone = pickle_depickle(FuncWrapper(func), protocol=self.protocol) + assert annotationlib.get_annotations( + clone, format=annotationlib.Format.STRING + ) == {"arg": "UndefinedName", "return": "int"} + def test_function_annotations(self): def f(a: int) -> str: pass