-
Notifications
You must be signed in to change notification settings - Fork 154
Expand file tree
/
Copy pathtest_preview_loader.py
More file actions
68 lines (51 loc) · 1.9 KB
/
Copy pathtest_preview_loader.py
File metadata and controls
68 lines (51 loc) · 1.9 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
# SPDX-FileCopyrightText: Copyright (c) <2026> NVIDIA CORPORATION & AFFILIATES. All rights reserved.
#
# SPDX-License-Identifier: Apache-2.0
import types
import pytest
from cuda.tile import _preview_loader
_mock_type = type("PreviewForeignCall", (), {})
_preview = types.SimpleNamespace(PreviewForeignCall=_mock_type)
def _missing_import(name):
raise ModuleNotFoundError(name="cuda.tile_preview")
@pytest.mark.parametrize("mock_import,expected", [
(_missing_import, None),
(lambda name: _preview, _mock_type),
])
def test_import_preview_package(monkeypatch, mock_import, expected):
calls = []
def counted_import(name):
calls.append(name)
return mock_import(name)
get_type = _preview_loader.get_preview_foreign_call_type
get_type.cache_clear()
monkeypatch.setattr(_preview_loader.importlib, "import_module", counted_import)
try:
assert get_type() is expected
assert get_type() is expected # cached
assert len(calls) == 1
finally:
get_type.cache_clear()
def test_import_fails_on_dependency(monkeypatch):
error = ModuleNotFoundError(name="preview_dependency")
def fail_import(name):
raise error
get_type = _preview_loader.get_preview_foreign_call_type
get_type.cache_clear()
monkeypatch.setattr(_preview_loader.importlib, "import_module", fail_import)
try:
with pytest.raises(ModuleNotFoundError) as exc_info:
get_type()
assert exc_info.value is error
finally:
get_type.cache_clear()
def test_import_fails_when_type_missing(monkeypatch):
preview = types.SimpleNamespace()
get_type = _preview_loader.get_preview_foreign_call_type
get_type.cache_clear()
monkeypatch.setattr(_preview_loader.importlib, "import_module", lambda name: preview)
try:
with pytest.raises(AttributeError):
get_type()
finally:
get_type.cache_clear()