File size: 2,683 Bytes
c689a69
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
import ast
import sys
import unittest
from pathlib import Path

EYETTENTION_ROOT = Path(__file__).resolve().parents[1]
PROJECT_ROOT = EYETTENTION_ROOT.parent
EXCLUDED_DIRS = {"__pycache__", ".git", "tests"}
EXPECTED_IMPORT_FAILURES = {
    ("Eyettention/app.py", "import gradio as gr"),
}

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


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

        for file_path in _python_files(EYETTENTION_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()