Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 8 additions & 2 deletions cuda_core/cuda/core/_cpp/resource_handles.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -1524,13 +1524,13 @@ DevicePtrHandle deviceptr_import_ipc(const MemoryPoolHandle& h_pool, const void*
ExportDataKey key;
std::memcpy(&key.data, data, sizeof(key.data));

GILReleaseGuard gil;
std::lock_guard<std::mutex> lock(ipc_import_mutex);

if (auto h = ipc_ptr_cache.lookup(key)) {
return h;
}

GILReleaseGuard gil;
CUdeviceptr ptr;
if (CUDA_SUCCESS != (err = p_cuMemPoolImportPointer(&ptr, *h_pool, data))) {
return {};
Expand All @@ -1545,8 +1545,14 @@ DevicePtrHandle deviceptr_import_ipc(const MemoryPoolHandle& h_pool, const void*
auto box = std::shared_ptr<DevicePtrBox>(
new DevicePtrBox{ptr, std::move(ds)},
[h_pool, key](DevicePtrBox* b) {
ipc_ptr_cache.unregister_handle(key);
// Release the GIL first (the GIL is the outermost lock), then hold
// the mutex across unregister + free. A concurrent import that finds
// this entry expired must wait until the mapping is gone; otherwise
// it re-imports the same allocation and the first cuMemFreeAsync
// unmaps it for both (nvbug 5570902).
GILReleaseGuard gil;
std::lock_guard<std::mutex> lock(ipc_import_mutex);
ipc_ptr_cache.unregister_handle(key);
const DeallocationStream& stream = b->deallocation;
cleanup_in_context(
deallocation_context(stream), "cuMemFreeAsync",
Expand Down
87 changes: 87 additions & 0 deletions cuda_core/tests/memory_ipc/test_ipc_concurrent_import.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,87 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0

"""Regression test for #2840: concurrent IPC imports must not deadlock.

``deviceptr_import_ipc()`` took ``ipc_import_mutex`` while still holding the
GIL and released the GIL only inside the critical section. Because destruction
runs in reverse declaration order, the importing thread reacquired the GIL
before dropping the mutex, while a second thread blocked on the mutex holding
the GIL.

Every import here succeeds and nothing is reported, so a hang indicates the
lock ordering rather than anything on an error path. A deadlocked child is
detected by the parent's join timeout; the per-directory conftest timeout is
the final backstop.
"""

import contextlib
import multiprocessing as mp
import threading

import pytest
from helpers.child_processes import child_timeout_sec, kill_subprocesses

from cuda.core import Buffer, Device

CHILD_TIMEOUT_SEC = child_timeout_sec()
NBYTES = 64
THREADS = 4

# these tests spawn new processes and files which fails for very many threads
pytestmark = pytest.mark.parallel_threads_limit(4)


def child_main(queue):
device = Device()
device.set_current()
mr = queue.get()
descriptor = queue.get()

# One descriptor for every thread is enough: the mutex is taken before the
# pointer cache is consulted, so a thread that would have been a cache hit
# still blocks on the lock while holding the GIL.
barrier = threading.Barrier(THREADS)

def importer():
# A current context, so the import succeeds and nothing is reported.
Device().set_current()
barrier.wait()
Buffer.from_ipc_descriptor(mr, descriptor, stream=device.default_stream).close()

threads = [threading.Thread(target=importer) for _ in range(THREADS)]
for thread in threads:
thread.start()
for thread in threads:
thread.join()

device.sync()


class TestIpcConcurrentImport:
"""Importing one descriptor from several threads at once must not hang."""

@pytest.fixture(autouse=True)
def _set_start_method(self):
# Ensure spawn is used for multiprocessing
with contextlib.suppress(RuntimeError):
mp.set_start_method("spawn", force=True)

def test_main(self, ipc_device, ipc_memory_resource):
ipc_device.set_current()
mr = ipc_memory_resource

stream = ipc_device.default_stream
buffer = mr.allocate(NBYTES, stream=stream)
stream.sync()

queue = mp.Queue()
process = mp.Process(target=child_main, args=(queue,))
process.start()
queue.put(mr)
queue.put(buffer.ipc_descriptor)

process.join(timeout=CHILD_TIMEOUT_SEC)
survivors = kill_subprocesses(process)
assert not survivors, "concurrent importers deadlocked (see #2840)"
assert process.exitcode == 0, f"child process failed with exit code {process.exitcode}"
Loading