@@ -18,12 +18,16 @@ def restore_optional_import_state():
1818 saved_nvvm_attempted = _program ._nvvm_import_attempted
1919 saved_driver = _linker ._driver
2020 saved_inited = _linker ._inited
21+ saved_nvjitlink = _linker ._nvjitlink
22+ saved_nvjitlink_version = _linker ._nvjitlink_version
2123 saved_use_nvjitlink = _linker ._use_nvjitlink_backend
2224
2325 _program ._nvvm_module = None
2426 _program ._nvvm_import_attempted = False
2527 _linker ._driver = None
2628 _linker ._inited = False
29+ _linker ._nvjitlink = None
30+ _linker ._nvjitlink_version = None
2731 _linker ._use_nvjitlink_backend = None
2832
2933 yield
@@ -32,6 +36,8 @@ def restore_optional_import_state():
3236 _program ._nvvm_import_attempted = saved_nvvm_attempted
3337 _linker ._driver = saved_driver
3438 _linker ._inited = saved_inited
39+ _linker ._nvjitlink = saved_nvjitlink
40+ _linker ._nvjitlink_version = saved_nvjitlink_version
3541 _linker ._use_nvjitlink_backend = saved_use_nvjitlink
3642
3743
@@ -168,12 +174,22 @@ def fake__optional_cuda_import(modname, probe_function=None):
168174 assert _linker ._use_nvjitlink_backend is False
169175
170176
171- @pytest .mark .agent_authored (model = "grok-4.5 " )
177+ @pytest .mark .agent_authored (model = "gpt-5.6 " )
172178def test_decide_nvjitlink_or_driver_selects_nvjitlink_when_version_symbol_present (monkeypatch ):
179+ version_calls = 0
180+
181+ class FakeModule :
182+ def version (self ):
183+ nonlocal version_calls
184+ version_calls += 1
185+ return (13 , 4 )
186+
187+ nvjitlink_module = FakeModule ()
188+
173189 def fake__optional_cuda_import (modname , probe_function = None ):
174190 assert modname == "cuda.bindings.nvjitlink"
175191 assert probe_function is None
176- return object ()
192+ return nvjitlink_module
177193
178194 monkeypatch .setattr (_linker , "_optional_cuda_import" , fake__optional_cuda_import )
179195 monkeypatch .setattr (_linker , "_nvjitlink_has_version_symbol" , lambda _nvjitlink : True )
@@ -182,21 +198,24 @@ def fake__optional_cuda_import(modname, probe_function=None):
182198
183199 assert use_driver_backend is False
184200 assert _linker ._use_nvjitlink_backend is True
201+ assert _linker ._nvjitlink is nvjitlink_module
202+ assert _linker ._nvjitlink_version == (13 , 4 )
203+ assert version_calls == 1
185204
186205
187- @pytest .mark .agent_authored (model = "grok-4.5 " )
188- def test_decide_nvjitlink_or_driver_does_not_call_version (monkeypatch ):
189- """Regression guard for #2408: must not call module.version()."""
206+ @pytest .mark .agent_authored (model = "gpt-5.6 " )
207+ def test_decide_nvjitlink_or_driver_does_not_call_version_when_symbol_missing (monkeypatch ):
208+ """Regression guard for #2408: old nvJitLink must not call module.version()."""
190209 called = {"version" : False , "inspect" : False }
191210
192211 class FakeModule :
193212 def version (self ):
194213 called ["version" ] = True
195- raise AssertionError ("module.version() must not be used for nvJitLink probing " )
214+ raise AssertionError ("module.version() must not be called when its symbol is missing " )
196215
197216 def fake_has_version (_nvjitlink ):
198217 called ["inspect" ] = True
199- return True
218+ return False
200219
201220 def fake__optional_cuda_import (modname , probe_function = None ):
202221 assert modname == "cuda.bindings.nvjitlink"
@@ -206,6 +225,9 @@ def fake__optional_cuda_import(modname, probe_function=None):
206225 monkeypatch .setattr (_linker , "_optional_cuda_import" , fake__optional_cuda_import )
207226 monkeypatch .setattr (_linker , "_nvjitlink_has_version_symbol" , fake_has_version )
208227
209- assert _linker ._decide_nvjitlink_or_driver () is False
228+ with pytest .warns (RuntimeWarning , match = "too old \\ (<12.3\\ )" ):
229+ assert _linker ._decide_nvjitlink_or_driver () is True
210230 assert called ["inspect" ] is True
211231 assert called ["version" ] is False
232+ assert _linker ._nvjitlink is None
233+ assert _linker ._nvjitlink_version is None
0 commit comments