From 06837ec267e3ca0b1734461d09919ff65cb1b1fe Mon Sep 17 00:00:00 2001 From: Tristan Simas Date: Tue, 29 Sep 2026 14:30:36 -0400 Subject: [PATCH] Keep decorator thread-local handles on their runtime owner --- pyproject.toml | 1 + src/arraybridge/decorators.py | 27 +++++++------ tests/test_durable_decorator_context.py | 50 +++++++++++++++++++++++++ 3 files changed, 66 insertions(+), 12 deletions(-) create mode 100644 tests/test_durable_decorator_context.py diff --git a/pyproject.toml b/pyproject.toml index 655eaf8..61d0176 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -75,6 +75,7 @@ all = [ ] dev = [ "pytest>=7.0", + "dill>=0.3.8", "pytest-cov>=4.0", "build>=1.0", "twine>=5.0", diff --git a/src/arraybridge/decorators.py b/src/arraybridge/decorators.py index e0e83fe..9612137 100644 --- a/src/arraybridge/decorators.py +++ b/src/arraybridge/decorators.py @@ -264,12 +264,22 @@ def _insertion_index(parameters: list[inspect.Parameter]) -> int: return len(parameters) -# Thread-local storage for GPU streams and contexts -_thread_gpu_contexts = threading.local() +class ThreadGPUContext: + """Runtime-owned thread-local streams keyed by framework/device identity. + Keep the runtime handle on its importable owner, not in decorator function + globals. A retained, unpublished decorated callable is serialized by value; + its durable closure must not pull a ``threading.local`` into history. + """ -class ThreadGPUContext: - """Thread-local streams keyed by framework-local device identity.""" + _contexts: ClassVar[threading.local] = threading.local() + + @classmethod + def current(cls) -> "ThreadGPUContext": + """Return this thread's runtime context without serializing its handle.""" + if not hasattr(cls._contexts, "context"): + cls._contexts.context = cls() + return cls._contexts.context def __init__(self): self._streams: dict[tuple[MemoryType, int], Any] = {} @@ -301,13 +311,6 @@ def stream_for( return device_id, self._streams[key] -def _get_thread_gpu_context(): - """Get or create thread-local GPU context.""" - if not hasattr(_thread_gpu_contexts, "context"): - _thread_gpu_contexts.context = ThreadGPUContext() - return _thread_gpu_contexts.context - - def memory_types( input_type: str | MemoryType, output_type: str | MemoryType, @@ -461,7 +464,7 @@ def gpu_wrapper(*args, **kwargs): # Check if GPU is available for this framework if framework is not None and mem_type.available_device_ids(framework): # Get thread-local context - ctx = _get_thread_gpu_context() + ctx = ThreadGPUContext.current() device_id, stream = ctx.stream_for(mem_type, framework) diff --git a/tests/test_durable_decorator_context.py b/tests/test_durable_decorator_context.py new file mode 100644 index 0000000..24385d5 --- /dev/null +++ b/tests/test_durable_decorator_context.py @@ -0,0 +1,50 @@ +"""Retained callable declarations exclude thread-local GPU lifecycle handles.""" + +import threading +from concurrent.futures import ThreadPoolExecutor + +import dill +import numpy as np +import pytest + +from arraybridge.decorators import ThreadGPUContext, cupy, numpy +from arraybridge.types import MemoryType + + +def test_context_is_stable_per_thread_and_never_shared_between_threads(): + current = ThreadGPUContext.current() + with ThreadPoolExecutor(max_workers=1) as executor: + other, repeated = executor.submit( + lambda: (ThreadGPUContext.current(), ThreadGPUContext.current()) + ).result() + assert ThreadGPUContext.current() is current + assert other is repeated + assert other is not current + + +@pytest.mark.parametrize("decorator", [numpy, cupy]) +def test_unpublished_decorated_callable_serializes_without_runtime_context(decorator): + # Like a replaced custom function retained by an undo snapshot, this + # declaration has no importable public alias. Dill must persist its body. + @decorator + def declared(image, scale=3): + return image * scale + + context = ThreadGPUContext.current() + key = (MemoryType.CUPY, -1) + handle = threading.local() + context._streams[key] = handle + try: + restored = dill.loads(dill.dumps(declared)) + assert restored.__name__ == declared.__name__ + assert restored.__wrapped__.__name__ == declared.__wrapped__.__name__ + # CPU NumPy invocation checks actual restored behavior. GPU behavior + # keeps its existing focused tests; this receipt starts no GPU runtime. + if decorator is numpy: + np.testing.assert_array_equal( + restored(np.array([1, 2], dtype=np.uint16), scale=4), [4, 8] + ) + assert ThreadGPUContext.current() is context + assert context._streams[key] is handle + finally: + del context._streams[key]