77import ctypes .wintypes
88import os
99import struct
10+ from collections .abc import Iterator
1011from typing import TYPE_CHECKING
1112
1213from cuda .pathfinder ._dynamic_libs .load_dl_common import LoadedDL
2223
2324# Set up kernel32 functions with proper types
2425kernel32 = ctypes .windll .kernel32 # type: ignore[attr-defined]
26+ psapi = ctypes .windll .psapi # type: ignore[attr-defined]
2527
2628# GetModuleHandleW
2729kernel32 .GetModuleHandleW .argtypes = [ctypes .wintypes .LPCWSTR ]
2830kernel32 .GetModuleHandleW .restype = ctypes .wintypes .HMODULE
2931
32+ # GetCurrentProcess
33+ kernel32 .GetCurrentProcess .argtypes = []
34+ kernel32 .GetCurrentProcess .restype = ctypes .wintypes .HANDLE
35+
3036# LoadLibraryExW
3137kernel32 .LoadLibraryExW .argtypes = [
3238 ctypes .wintypes .LPCWSTR , # lpLibFileName
4753kernel32 .AddDllDirectory .argtypes = [ctypes .wintypes .LPCWSTR ]
4854kernel32 .AddDllDirectory .restype = ctypes .c_void_p # DLL_DIRECTORY_COOKIE
4955
56+ # EnumProcessModules
57+ psapi .EnumProcessModules .argtypes = [
58+ ctypes .wintypes .HANDLE ,
59+ ctypes .POINTER (ctypes .wintypes .HMODULE ),
60+ ctypes .wintypes .DWORD ,
61+ ctypes .POINTER (ctypes .wintypes .DWORD ),
62+ ]
63+ psapi .EnumProcessModules .restype = ctypes .wintypes .BOOL
64+
5065
5166def ctypes_handle_to_unsigned_int (handle : ctypes .wintypes .HMODULE ) -> int :
5267 """Convert ctypes HMODULE to unsigned int."""
@@ -101,6 +116,41 @@ def abs_path_for_dynamic_library(libname: str, handle: ctypes.wintypes.HMODULE)
101116 return buffer .value
102117
103118
119+ def _iter_loaded_module_handles () -> Iterator [ctypes .wintypes .HMODULE ]:
120+ process_handle = kernel32 .GetCurrentProcess ()
121+ capacity = 64
122+ module_size = ctypes .sizeof (ctypes .wintypes .HMODULE )
123+ while True :
124+ module_handles = (ctypes .wintypes .HMODULE * capacity )()
125+ needed = ctypes .wintypes .DWORD ()
126+ ok = psapi .EnumProcessModules (
127+ process_handle ,
128+ module_handles ,
129+ ctypes .sizeof (module_handles ),
130+ ctypes .byref (needed ),
131+ )
132+ if not ok :
133+ error_code = ctypes .GetLastError () # type: ignore[attr-defined]
134+ raise RuntimeError (f"EnumProcessModules failed (error code: { error_code } )" )
135+ count = needed .value // module_size
136+ if count <= capacity :
137+ for raw_handle in module_handles [:count ]:
138+ if raw_handle is None :
139+ continue
140+ yield ctypes .wintypes .HMODULE (int (raw_handle ))
141+ return
142+ capacity = count
143+
144+
145+ def _find_loaded_module (dll_names : tuple [str , ...]) -> tuple [ctypes .wintypes .HMODULE , str ] | None :
146+ wanted = {dll_name .casefold () for dll_name in dll_names }
147+ for handle in _iter_loaded_module_handles ():
148+ abs_path = abs_path_for_dynamic_library ("loaded module" , handle )
149+ if os .path .basename (abs_path ).casefold () in wanted :
150+ return handle , abs_path
151+ return None
152+
153+
104154def check_if_already_loaded_from_elsewhere (desc : LibDescriptor , have_abs_path : bool ) -> LoadedDL | None :
105155 for dll_name in desc .windows_dlls :
106156 handle = kernel32 .GetModuleHandleW (dll_name )
@@ -112,6 +162,14 @@ def check_if_already_loaded_from_elsewhere(desc: LibDescriptor, have_abs_path: b
112162 # activate it even if the library was already loaded from elsewhere.
113163 add_dll_directory (abs_path )
114164 return LoadedDL (abs_path , True , ctypes_handle_to_unsigned_int (handle ), "was-already-loaded-from-elsewhere" )
165+ # Observed on newer Windows CUPTI builds: GetModuleHandleW(basename)
166+ # can miss an already loaded DLL, so fall back to enumerating loaded modules.
167+ loaded = _find_loaded_module (desc .windows_dlls )
168+ if loaded is not None :
169+ handle , abs_path = loaded
170+ if have_abs_path and desc .requires_add_dll_directory :
171+ add_dll_directory (abs_path )
172+ return LoadedDL (abs_path , True , ctypes_handle_to_unsigned_int (handle ), "was-already-loaded-from-elsewhere" )
115173 return None
116174
117175
0 commit comments