From 7de0ae67600159ef0c694594430718ddeabde2e2 Mon Sep 17 00:00:00 2001 From: Sebastian Rittau Date: Wed, 30 Jun 2021 12:56:01 +0200 Subject: [PATCH 1/6] stubgen: Add support for yield statements --- mypy/stubgen.py | 34 ++++++++++++++++++++++++++- mypy/traverser.py | 46 +++++++++++++++++++++++++++++++++--- test-data/unit/stubgen.test | 47 +++++++++++++++++++++++++++++++++++++ 3 files changed, 123 insertions(+), 4 deletions(-) diff --git a/mypy/stubgen.py b/mypy/stubgen.py index 91f461b84c15d..9d4a7bd56da44 100755 --- a/mypy/stubgen.py +++ b/mypy/stubgen.py @@ -91,7 +91,7 @@ from mypy.find_sources import create_source_list, InvalidSourceList from mypy.build import build from mypy.errors import CompileError, Errors -from mypy.traverser import has_return_statement +from mypy.traverser import all_yield_expressions, has_return_statement, has_yield_expression from mypy.moduleinspect import ModuleInspect @@ -557,6 +557,13 @@ def visit_mypy_file(self, o: MypyFile) -> None: else: alias = '_' + t self.import_tracker.add_import_from("typing", [(t, alias)]) + abc_imports = ["Generator"] + for t in abc_imports: + if t not in self.defined_names: + alias = None + else: + alias = '_' + t + self.import_tracker.add_import_from("collections.abc", [(t, alias)]) super().visit_mypy_file(o) undefined_names = [name for name in self._all_ or [] if name not in self._toplevel_names] @@ -662,6 +669,23 @@ def visit_func_def(self, o: FuncDef, is_abstract: bool = False, # Always assume abstract methods return Any unless explicitly annotated. Also # some dunder methods should not have a None return type. retname = None # implicit Any + elif has_yield_expression(o): + self.add_abc_import('Generator') + yield_name = 'None' + send_name = 'None' + return_name = 'None' + for expr, in_assignment in all_yield_expressions(o): + if expr.expr is not None: + self.add_typing_import('Any') + yield_name = 'Any' + if in_assignment: + self.add_typing_import('Any') + send_name = 'Any' + if has_return_statement(o): + self.add_typing_import('Any') + return_name = 'Any' + generator_name = self.typing_name('Generator') + retname = f'{generator_name}[{yield_name}, {send_name}, {return_name}]' elif not has_return_statement(o) and not is_abstract: retname = 'None' retfield = '' @@ -1107,6 +1131,14 @@ def add_typing_import(self, name: str) -> None: name = self.typing_name(name) self.import_tracker.require_name(name) + def add_abc_import(self, name: str) -> None: + """Add a name to be imported from collections.abc, unless it's imported already. + + The import will be internal to the stub. + """ + name = self.typing_name(name) + self.import_tracker.require_name(name) + def add_import_line(self, line: str) -> None: """Add a line of text to the import section, unless it's already there.""" if line not in self._import_lines: diff --git a/mypy/traverser.py b/mypy/traverser.py index d0b656c7a77f2..af132e267d6b3 100644 --- a/mypy/traverser.py +++ b/mypy/traverser.py @@ -1,6 +1,6 @@ """Generic node traverser visitor""" -from typing import List +from typing import List, Tuple from mypy_extensions import mypyc_attr from mypy.visitor import NodeVisitor @@ -319,9 +319,22 @@ def has_return_statement(fdef: FuncBase) -> bool: return seeker.found -class ReturnCollector(TraverserVisitor): +class YieldSeeker(TraverserVisitor): + def __init__(self) -> None: + self.found = False + + def visit_yield_expr(self, o: YieldExpr) -> None: + self.found = True + + +def has_yield_expression(fdef: FuncBase) -> bool: + seeker = YieldSeeker() + fdef.accept(seeker) + return seeker.found + + +class FuncCollectorBase(TraverserVisitor): def __init__(self) -> None: - self.return_statements: List[ReturnStmt] = [] self.inside_func = False def visit_func_def(self, defn: FuncDef) -> None: @@ -330,6 +343,12 @@ def visit_func_def(self, defn: FuncDef) -> None: super().visit_func_def(defn) self.inside_func = False + +class ReturnCollector(FuncCollectorBase): + def __init__(self) -> None: + super().__init__() + self.return_statements: List[ReturnStmt] = [] + def visit_return_stmt(self, stmt: ReturnStmt) -> None: self.return_statements.append(stmt) @@ -338,3 +357,24 @@ def all_return_statements(node: Node) -> List[ReturnStmt]: v = ReturnCollector() node.accept(v) return v.return_statements + + +class YieldCollector(FuncCollectorBase): + def __init__(self) -> None: + super().__init__() + self.in_assignment = False + self.yield_expressions: List[Tuple[YieldExpr, bool]] = [] + + def visit_assignment_stmt(self, stmt) -> None: + self.in_assignment = True + super().visit_assignment_stmt(stmt) + self.in_assignment = False + + def visit_yield_expr(self, expr: YieldExpr) -> None: + self.yield_expressions.append((expr, self.in_assignment)) + + +def all_yield_expressions(node: Node) -> List[Tuple[YieldExpr, bool]]: + v = YieldCollector() + node.accept(v) + return v.yield_expressions diff --git a/test-data/unit/stubgen.test b/test-data/unit/stubgen.test index daee45a840825..9e3ee2a2870f1 100644 --- a/test-data/unit/stubgen.test +++ b/test-data/unit/stubgen.test @@ -944,6 +944,53 @@ def f(): ... [out] def f() -> None: ... +[case testFunctionYields] +def f(): + yield 123 +def g(): + x = yield +def h1(): + yield + return +def h2(): + yield + return "abc" +def all(): + x = yield 123 + return "abc" +[out] +from collections.abc import Generator +from typing import Any + +def f() -> Generator[Any, None, None]: ... +def g() -> Generator[None, Any, None]: ... +def h1() -> Generator[None, None, None]: ... +def h2() -> Generator[None, None, Any]: ... +def all() -> Generator[Any, Any, Any]: ... + +[case testFunctionYieldsNone] +def f(): + yield + +[out] +from collections.abc import Generator + +def f() -> Generator[None, None, None]: ... + +[case testGeneratorAlreadyDefined] +class Generator: + pass + +def f(): + yield 123 +[out] +from collections.abc import Generator as _Generator +from typing import Any + +class Generator: ... + +def f() -> _Generator[Any, None, None]: ... + [case testCallable] from typing import Callable From a08c7f2f36c90f311aea20da67fe99a32f4ec9de Mon Sep 17 00:00:00 2001 From: Sebastian Rittau Date: Wed, 30 Jun 2021 13:30:27 +0200 Subject: [PATCH 2/6] Add missing type annotation --- mypy/traverser.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/mypy/traverser.py b/mypy/traverser.py index af132e267d6b3..a5f993bd2fa52 100644 --- a/mypy/traverser.py +++ b/mypy/traverser.py @@ -365,7 +365,7 @@ def __init__(self) -> None: self.in_assignment = False self.yield_expressions: List[Tuple[YieldExpr, bool]] = [] - def visit_assignment_stmt(self, stmt) -> None: + def visit_assignment_stmt(self, stmt: AssignmentStmt) -> None: self.in_assignment = True super().visit_assignment_stmt(stmt) self.in_assignment = False From fcd528d9812f12f70d0fca4fa7bd13b3d4f89a5b Mon Sep 17 00:00:00 2001 From: Sebastian Rittau Date: Wed, 30 Jun 2021 13:41:00 +0200 Subject: [PATCH 3/6] Refactor known imports --- mypy/stubgen.py | 25 +++++++++++-------------- 1 file changed, 11 insertions(+), 14 deletions(-) diff --git a/mypy/stubgen.py b/mypy/stubgen.py index 9d4a7bd56da44..8da6037c7c0d3 100755 --- a/mypy/stubgen.py +++ b/mypy/stubgen.py @@ -550,20 +550,17 @@ def visit_mypy_file(self, o: MypyFile) -> None: self.path = o.path self.defined_names = find_defined_names(o) self.referenced_names = find_referenced_names(o) - typing_imports = ["Any", "Optional", "TypeVar"] - for t in typing_imports: - if t not in self.defined_names: - alias = None - else: - alias = '_' + t - self.import_tracker.add_import_from("typing", [(t, alias)]) - abc_imports = ["Generator"] - for t in abc_imports: - if t not in self.defined_names: - alias = None - else: - alias = '_' + t - self.import_tracker.add_import_from("collections.abc", [(t, alias)]) + known_imports = { + "typing": ["Any", "Optional", "TypeVar"], + "collections.abc": ["Generator"], + } + for pkg, imports in known_imports.items(): + for t in imports: + if t not in self.defined_names: + alias = None + else: + alias = '_' + t + self.import_tracker.add_import_from(pkg, [(t, alias)]) super().visit_mypy_file(o) undefined_names = [name for name in self._all_ or [] if name not in self._toplevel_names] From 6bca216f9fe36401befc06e0c1aed9c4422951cb Mon Sep 17 00:00:00 2001 From: Sebastian Rittau Date: Wed, 30 Jun 2021 13:42:45 +0200 Subject: [PATCH 4/6] Remove Optional from known_imports (unused) --- mypy/stubgen.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/mypy/stubgen.py b/mypy/stubgen.py index 8da6037c7c0d3..9a3046364ef91 100755 --- a/mypy/stubgen.py +++ b/mypy/stubgen.py @@ -551,7 +551,7 @@ def visit_mypy_file(self, o: MypyFile) -> None: self.defined_names = find_defined_names(o) self.referenced_names = find_referenced_names(o) known_imports = { - "typing": ["Any", "Optional", "TypeVar"], + "typing": ["Any", "TypeVar"], "collections.abc": ["Generator"], } for pkg, imports in known_imports.items(): From 3ddf133dc3ababd95a8422083e8f0843e57b0204 Mon Sep 17 00:00:00 2001 From: Sebastian Rittau Date: Wed, 30 Jun 2021 13:45:08 +0200 Subject: [PATCH 5/6] Slightly improve return type logic --- mypy/stubgen.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/mypy/stubgen.py b/mypy/stubgen.py index 9a3046364ef91..ff051cca11e7d 100755 --- a/mypy/stubgen.py +++ b/mypy/stubgen.py @@ -678,9 +678,9 @@ def visit_func_def(self, o: FuncDef, is_abstract: bool = False, if in_assignment: self.add_typing_import('Any') send_name = 'Any' - if has_return_statement(o): - self.add_typing_import('Any') - return_name = 'Any' + if has_return_statement(o): + self.add_typing_import('Any') + return_name = 'Any' generator_name = self.typing_name('Generator') retname = f'{generator_name}[{yield_name}, {send_name}, {return_name}]' elif not has_return_statement(o) and not is_abstract: From 3b6f7e218d0901584a8c09f754ba1823eb66c44e Mon Sep 17 00:00:00 2001 From: Sebastian Rittau Date: Wed, 30 Jun 2021 13:45:35 +0200 Subject: [PATCH 6/6] Support explicit "yield None" --- mypy/stubgen.py | 5 ++++- test-data/unit/stubgen.test | 3 +++ 2 files changed, 7 insertions(+), 1 deletion(-) diff --git a/mypy/stubgen.py b/mypy/stubgen.py index ff051cca11e7d..1959f71d12ed4 100755 --- a/mypy/stubgen.py +++ b/mypy/stubgen.py @@ -672,7 +672,7 @@ def visit_func_def(self, o: FuncDef, is_abstract: bool = False, send_name = 'None' return_name = 'None' for expr, in_assignment in all_yield_expressions(o): - if expr.expr is not None: + if expr.expr is not None and not self.is_none_expr(expr.expr): self.add_typing_import('Any') yield_name = 'Any' if in_assignment: @@ -693,6 +693,9 @@ def visit_func_def(self, o: FuncDef, is_abstract: bool = False, self.add("){}: ...\n".format(retfield)) self._state = FUNC + def is_none_expr(self, expr: Expression) -> bool: + return isinstance(expr, NameExpr) and expr.name == "None" + def visit_decorator(self, o: Decorator) -> None: if self.is_private_name(o.func.name, o.func.fullname): return diff --git a/test-data/unit/stubgen.test b/test-data/unit/stubgen.test index 9e3ee2a2870f1..b1995b57e7759 100644 --- a/test-data/unit/stubgen.test +++ b/test-data/unit/stubgen.test @@ -971,11 +971,14 @@ def all() -> Generator[Any, Any, Any]: ... [case testFunctionYieldsNone] def f(): yield +def g(): + yield None [out] from collections.abc import Generator def f() -> Generator[None, None, None]: ... +def g() -> Generator[None, None, None]: ... [case testGeneratorAlreadyDefined] class Generator: