diff --git a/docs/changelog.md b/docs/changelog.md index a7041eb..33482dd 100644 --- a/docs/changelog.md +++ b/docs/changelog.md @@ -56,6 +56,8 @@ this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0.htm - Crash on modules whose `__all__` is a set - An `__all__` with comments inside the literal is reported but not rewritten, so the comments are not lost; `# unexport:` markers inside string literals are ignored +- Names bound in only one branch of an `if` that can't be decided statically (e.g. + `if sys.platform == "win32":`) are no longer added to `__all__` ## [0.4.0] - 2022-11-05 diff --git a/docs/tutorials/useful-features.md b/docs/tutorials/useful-features.md index 6cfe025..7331988 100644 --- a/docs/tutorials/useful-features.md +++ b/docs/tutorials/useful-features.md @@ -42,6 +42,26 @@ T = TypeVar("T") # not added to __all__ PublicT = TypeVar("PublicT") # unexport: public ``` +## Conditional names + +A name defined in only one branch of an `if` whose outcome depends on the platform, the +Python version or anything else unexport can't know is not added to `__all__`: on the +other branch it doesn't exist, and `from module import *` would fail there. Names bound +in every branch are added as usual. Write 'unexport: public' as a comment to add one +anyway, or list it yourself. + +```python +import sys + +if sys.platform == "win32": + class WinRegistry: ... # not added + +if sys.version_info >= (3, 11): + Feature = ... # added: bound in both branches +else: + Feature = ... +``` + ## Names you list yourself Names that are already in `__all__` stay there as long as the module still binds them, diff --git a/src/unexport/relate.py b/src/unexport/relate.py index ae4ae52..353c2e7 100644 --- a/src/unexport/relate.py +++ b/src/unexport/relate.py @@ -8,6 +8,7 @@ "get_parents", "is_bare_annotation", "is_comprehension_target", + "is_conditional_only", "is_runtime_missing", "relate", ) @@ -101,3 +102,44 @@ def is_bare_annotation(node: ast.AST) -> bool: """Whether node is the target of an annotation without a value (``X: int``), which binds nothing at runtime.""" parent = getattr(node, "parent", None) return isinstance(parent, ast.AnnAssign) and parent.target is node and parent.value is None + + +def _binds(statements: list[ast.stmt], name: str) -> bool: + """Whether the statements bind name at this level (not inside functions or classes they define).""" + nodes: list[ast.AST] = list(statements) + while nodes: + node = nodes.pop() + if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef, ast.ClassDef)): + if node.name == name: + return True + continue + if isinstance(node, ast.Lambda): + continue + if isinstance(node, ast.Name) and isinstance(node.ctx, ast.Store) and node.id == name: + return True + if isinstance(node, (ast.Import, ast.ImportFrom)) and any( + (alias.asname or alias.name).split(".")[0] == name for alias in node.names + ): + return True + nodes.extend(ast.iter_child_nodes(node)) + return False + + +def is_conditional_only(node: ast.AST, name: str) -> bool: + """Whether name is bound only in one branch of an ``if`` whose outcome isn't known statically. + + E.g. ``if sys.platform == "win32": class WinOnly: ...`` without an + ``else`` that binds ``WinOnly``: on other platforms the name doesn't + exist, and listing it in ``__all__`` breaks ``from module import *``. + """ + child = node + for parent in get_parents(node): + if isinstance(parent, (ast.FunctionDef, ast.AsyncFunctionDef, ast.ClassDef, ast.Lambda)): + return False + if isinstance(parent, ast.If) and _truth_on_import(parent.test) is None: + if child in parent.body and not _binds(parent.orelse, name): + return True + if child in parent.orelse and not _binds(parent.body, name): + return True + child = parent + return False diff --git a/src/unexport/rule.py b/src/unexport/rule.py index 3691a5d..e453c33 100644 --- a/src/unexport/rule.py +++ b/src/unexport/rule.py @@ -8,7 +8,13 @@ from unexport import constants as C from unexport import typing as T -from unexport.relate import first_occurrence, is_bare_annotation, is_comprehension_target, is_runtime_missing +from unexport.relate import ( + first_occurrence, + is_bare_annotation, + is_comprehension_target, + is_conditional_only, + is_runtime_missing, +) __all__ = ("Rule",) @@ -177,3 +183,17 @@ def _rule_name_not_bare_annotation(node) -> bool: if hasattr(node, "add"): return node.add is True return not is_bare_annotation(node) + + +@Rule.register( # type: ignore + ( # type: ignore + ast.ClassDef, + ast.FunctionDef, + ast.AsyncFunctionDef, + ast.Name, + ) +) +def _rule_not_conditional_only(node) -> bool: + if hasattr(node, "add"): + return node.add is True + return not is_conditional_only(node, node.id if isinstance(node, ast.Name) else node.name) diff --git a/tests/test_analyzer.py b/tests/test_analyzer.py index 22835c3..61f9214 100644 --- a/tests/test_analyzer.py +++ b/tests/test_analyzer.py @@ -437,6 +437,39 @@ def test_walrus_in_lambda(self): """ self.assertListEqual(self.expected_all(source), ["Real"]) + def test_names_bound_in_one_branch_only(self): + source = """\ + import sys + + if sys.platform == "win32": + class WinOnly: ... + Both = 1 + elif sys.platform == "darwin": + Both = 2 + else: + def Both(): ... + + if sys.version_info >= (3, 11): + Feature = 1 + else: + from compat import Feature + + if sys.platform == "linux": + LinuxOnly = 1 # unexport: public + """ + self.assertListEqual(self.expected_all(source), ["Both", "Feature", "LinuxOnly"]) + + def test_listed_conditional_name_is_kept(self): + source = """\ + import sys + + __all__ = ["WinOnly"] + + if sys.platform == "win32": + class WinOnly: ... + """ + self.assertListEqual(self.expected_all(source), ["WinOnly"]) + def test_rebound_after_del(self): source = """\ VALUE = 1