File size: 3,774 Bytes
95456ed
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
c4b6f85
95456ed
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
import ast
import sys
import unittest
from pathlib import Path

SCANDL2_ROOT = Path(__file__).resolve().parents[1]
PROJECT_ROOT = SCANDL2_ROOT.parent
EXCLUDED_DIRS = {"__pycache__", ".git", "tests"}
EXPECTED_IMPORT_FAILURES = {
    ("ScanDL2/app.py", "import gradio as gr"),
    (
        "ScanDL2/create_data.py",
        "from ScanDL2.scandl_module.scripts.sp_load_celer_zuco import load_emtec, process_emtec",
    ),
    (
        "ScanDL2/create_data.py",
        "from ScanDL2.scandl_module.scripts.sp_load_celer_zuco import load_bsc, process_bsc",
    ),
    (
        "ScanDL2/fix_dur_module/train_seq2seq.py",
        "from ScanDL2.CONSTANTS import (\n    COMPLETE_FIXDUR_MODULE_TRAIN_PATH_BSC",
    ),
    (
        "ScanDL2/scandl_module/scripts/sp_run_train.py",
        "from ScanDL2.CONSTANTS import (\n    COMPLETE_SCANDL_MODULE_TRAIN_PATH_BSC",
    ),
    ("ScanDL2/scandl_module/original_scandl/utils/logger.py", "import tensorflow as tf"),
    (
        "ScanDL2/scandl_module/original_scandl/utils/logger.py",
        "from tensorflow.python import pywrap_tensorflow",
    ),
    (
        "ScanDL2/scandl_module/original_scandl/utils/logger.py",
        "from tensorflow.core.util import event_pb2",
    ),
    (
        "ScanDL2/scandl_module/original_scandl/utils/logger.py",
        "from tensorflow.python.util import compat",
    ),
}

if str(PROJECT_ROOT) not in sys.path:
    sys.path.insert(0, str(PROJECT_ROOT))


class ScanDL2ImportTests(unittest.TestCase):
    def test_import_lines(self):
        failures = []

        for file_path in _python_files(SCANDL2_ROOT):
            source = file_path.read_text()
            tree = ast.parse(source, filename=str(file_path))

            for node in ast.walk(tree):
                if not isinstance(node, (ast.Import, ast.ImportFrom)):
                    continue
                if isinstance(node, ast.ImportFrom) and node.module == "__future__":
                    continue

                import_line = ast.get_source_segment(source, node)
                try:
                    exec(
                        compile(import_line, str(file_path), "exec"),
                        _import_globals(file_path),
                    )
                except Exception as exc:
                    relative_path = str(file_path.relative_to(PROJECT_ROOT))
                    if _is_expected_failure(relative_path, import_line):
                        continue
                    failures.append(
                        f"{relative_path}:{node.lineno}\n"
                        f"{import_line}\n"
                        f"{type(exc).__name__}: {exc}"
                    )

        if failures:
            self.fail("Failed import line(s):\n\n" + "\n\n".join(failures))


def _python_files(root):
    for file_path in root.rglob("*.py"):
        if any(part in EXCLUDED_DIRS for part in file_path.parts):
            continue
        yield file_path


def _import_globals(file_path):
    module_path = file_path.relative_to(PROJECT_ROOT).with_suffix("")
    module_parts = module_path.parts

    if module_parts[-1] == "__init__":
        module_name = ".".join(module_parts[:-1])
        package = module_name
    else:
        module_name = ".".join(module_parts)
        package = ".".join(module_parts[:-1])

    return {
        "__name__": module_name,
        "__package__": package,
    }


def _is_expected_failure(relative_path, import_line):
    return any(
        relative_path == expected_path and import_line.startswith(expected_import)
        for expected_path, expected_import in EXPECTED_IMPORT_FAILURES
    )


if __name__ == "__main__":
    unittest.main()