From bae7f56999bf3fdc38cf57ad11dd4c849d61c681 Mon Sep 17 00:00:00 2001 From: Ryan Heard Date: Thu, 8 Oct 2026 09:15:49 -0400 Subject: [PATCH] Discard walrus narrowing from repeated inference of call arguments --- mypy/binder.py | 40 +++++++++-- mypy/checkexpr.py | 102 +++++++++++++++++++++-------- test-data/unit/check-python38.test | 97 +++++++++++++++++++++++++++ 3 files changed, 206 insertions(+), 33 deletions(-) diff --git a/mypy/binder.py b/mypy/binder.py index 9984781733cd4..d7de0c8ebe6df 100644 --- a/mypy/binder.py +++ b/mypy/binder.py @@ -72,12 +72,15 @@ class Frame: types -- the concept predates literal types. """ - def __init__(self, id: int, conditional_frame: bool = False) -> None: + def __init__(self, id: int, conditional_frame: bool = False, discard: bool = False) -> None: self.id = id self.types: dict[Key, CurrentType] = {} self.unreachable = False self.conditional_frame = conditional_frame self.suppress_unreachable_warnings = False + # For a frame that will be discarded: types removed from outer frames while + # this frame is on the stack, to be restored when it is popped. + self.removed_types: list[tuple[Frame, Key, CurrentType]] | None = [] if discard else None def __repr__(self) -> str: return f"Frame({self.id}, {self.types}, {self.unreachable}, {self.conditional_frame})" @@ -123,7 +126,7 @@ def __enter__(self) -> Frame: if self.try_frame: self.binder.try_frames.add(len(self.binder.frames) - 1) - new_frame = self.binder.push_frame(self.conditional_frame) + new_frame = self.binder.push_frame(self.conditional_frame, discard=self.discard) if self.try_frame: # An exception may occur immediately self.binder.allow_jump(-1) @@ -208,6 +211,9 @@ def __init__(self, options: Options) -> None: # expression caches when needed. self.version = 0 + # Indices (in self.frames) of the frames that will be discarded when popped. + self.discard_frames: list[int] = [] + def _get_id(self) -> int: self.next_id += 1 return self.next_id @@ -220,9 +226,14 @@ def _add_dependencies(self, key: Key, value: Key | None = None) -> None: for elt in subkeys(key): self._add_dependencies(elt, value) - def push_frame(self, conditional_frame: bool = False) -> Frame: - """Push a new frame into the binder.""" - f = Frame(self._get_id(), conditional_frame) + def push_frame(self, conditional_frame: bool = False, discard: bool = False) -> Frame: + """Push a new frame into the binder. + + If discard is True, the frame must be popped with discard=True. + """ + f = Frame(self._get_id(), conditional_frame, discard) + if discard: + self.discard_frames.append(len(self.frames)) self.frames.append(f) self.options_on_return.append([]) return f @@ -295,8 +306,13 @@ def cleanse(self, expr: Expression) -> None: def _cleanse_key(self, key: Key) -> None: """Remove all references to a key from the binder.""" - for frame in self.frames: + for i, frame in enumerate(self.frames): if key in frame.types: + if self.discard_frames and i < self.discard_frames[-1]: + # Restore the type when the innermost discarded frame is popped. + removed = self.frames[self.discard_frames[-1]].removed_types + assert removed is not None + removed.append((frame, key, frame.types[key])) del frame.types[key] def update_from_options(self, frames: list[Frame]) -> bool: @@ -438,8 +454,18 @@ def pop_frame(self, can_skip: bool, fall_through: int, *, discard: bool = False) options = self.options_on_return.pop() if discard: + assert self.discard_frames and self.discard_frames[-1] == len(self.frames) + self.discard_frames.pop() + removed = result.removed_types + assert removed is not None + for frame, key, current in reversed(removed): + frame.types[key] = current + if result.types or removed or result.unreachable: + # Types visible through the frame stack are now different again. + self.version += 1 self.last_pop_changed = False return result + assert result.removed_types is None, "Frame was pushed with discard=True" if can_skip: options.insert(0, self.frames[-1]) @@ -607,6 +633,8 @@ def frame_context( If discard is True, then this is a temporary throw-away frame (used e.g. for isolation) and its effect will be discarded on pop. + This includes types removed from outer frames while it was pushed, + which are restored. After the context manager exits, self.last_pop_changed indicates whether any types changed in the newly-topmost frame as a result diff --git a/mypy/checkexpr.py b/mypy/checkexpr.py index 3e0772c302378..0ccbf4f5a58e6 100644 --- a/mypy/checkexpr.py +++ b/mypy/checkexpr.py @@ -577,7 +577,7 @@ def visit_call_expr_inner(self, e: CallExpr, allow_none_return: bool = False) -> e.arg_names, e.callee.arg_kinds, e.callee.arg_names, - lambda i: self.accept(e.args[i]), + lambda i: self.accept_and_discard_narrowing(e.args[i]), ) arg_types = [ @@ -1214,7 +1214,7 @@ def try_infer_partial_value_type_from_call( and methodname in self.item_args[typename] and e.arg_kinds == [ARG_POS] ): - item_type = self.accept(e.args[0]) + item_type = self.accept_and_discard_narrowing(e.args[0]) if mypy.checker.is_valid_inferred_type(item_type, self.chk.options): return self.chk.named_generic_type(typename, [item_type]) elif ( @@ -1222,7 +1222,7 @@ def try_infer_partial_value_type_from_call( and methodname in self.container_args[typename] and e.arg_kinds == [ARG_POS] ): - arg_type = get_proper_type(self.accept(e.args[0])) + arg_type = get_proper_type(self.accept_and_discard_narrowing(e.args[0])) if isinstance(arg_type, Instance): arg_typename = arg_type.type.fullname if arg_typename in self.container_args[typename][methodname]: @@ -1325,7 +1325,7 @@ def apply_signature_hook( arg_names, callee.arg_kinds, callee.arg_names, - lambda i: self.accept(args[i]), + lambda i: self.accept_and_discard_narrowing(args[i]), ) formal_arg_exprs: list[list[Expression]] = [[] for _ in range(num_formals)] for formal, actuals in enumerate(formal_to_actual): @@ -1782,7 +1782,7 @@ def check_callable_call( arg_names, callee.arg_kinds, callee.arg_names, - lambda i: self.accept(args[i]), + lambda i: self.accept_and_discard_narrowing(args[i]), ) if callee.special_sig == "tuple" and len(args) == 1: @@ -1823,7 +1823,7 @@ def check_callable_call( arg_names, callee.arg_kinds, callee.arg_names, - lambda i: self.accept(args[i]), + lambda i: self.accept_and_discard_narrowing(args[i]), ) param_spec = callee.param_spec() @@ -2028,6 +2028,25 @@ def infer_arg_types_in_empty_context(self, args: list[Expression]) -> list[Type] res.append(arg_type) return res + def accept_and_discard_narrowing(self, node: Expression) -> Type: + """Infer the type of an expression that will be inferred again later. + + Narrowing done here (by assignment expressions) is discarded, as it would + otherwise affect expressions before the assignment in the later inference. + """ + with self.chk.binder.frame_context(can_skip=False, discard=True): + return self.accept(node) + + def apply_argument_narrowing(self, args: list[Expression]) -> None: + """Apply narrowing from assignment expressions in arguments. + + This is needed after the arguments were only inferred in discarded binder + frames. Infer them once more for the side effects; the types and errors + from the inference that was used are already recorded. + """ + with self.msg.filter_errors(filter_revealed_type=True), self.chk.local_type_map: + self.infer_arg_types_in_empty_context(args) + def infer_more_unions_for_recursive_type(self, type_context: Type) -> bool: """Adjust type inference of unions if type context has a recursive type. @@ -2187,8 +2206,13 @@ def infer_function_type_arguments( # Disable type errors during type inference. There may be errors # due to partial available context information at this time, but # these errors can be safely ignored as the arguments will be - # inferred again later. - with self.msg.filter_errors(): + # inferred again later. For the same reason, discard any narrowing + # done here (by assignment expressions), as it would otherwise affect + # arguments that come before the assignment when they are inferred again. + with ( + self.msg.filter_errors(), + self.chk.binder.frame_context(can_skip=False, discard=True), + ): arg_types = self.infer_arg_types_in_context( callee_type, args, arg_kinds, formal_to_actual ) @@ -2259,7 +2283,7 @@ def infer_function_type_arguments( arg_names, callee_type.arg_kinds, callee_type.arg_names, - lambda a: self.accept(args[a]), + lambda a: self.accept_and_discard_narrowing(args[a]), ) # If the regular two-phase inference didn't work, try inferring type # variables while allowing for polymorphic solutions, i.e. for solutions @@ -2347,11 +2371,12 @@ def infer_function_type_arguments_pass2( arg_names, callee_type.arg_kinds, callee_type.arg_names, - lambda a: self.accept(args[a]), + lambda a: self.accept_and_discard_narrowing(args[a]), ) - # Same as during first pass, disable type errors (we still have partial context). - with self.msg.filter_errors(): + # Same as during first pass, disable type errors (we still have partial context) + # and discard narrowing. + with self.msg.filter_errors(), self.chk.binder.frame_context(can_skip=False, discard=True): arg_types = self.infer_arg_types_in_context( callee_type, args, arg_kinds, formal_to_actual ) @@ -2851,7 +2876,13 @@ def check_overload_call( """Checks a call to an overloaded function.""" # Normalize unpacked kwargs before checking the call. callee = callee.with_unpacked_kwargs() - arg_types = self.infer_arg_types_in_empty_context(args) + # The arguments are inferred several times below, against different items. + # This is done in discarded binder frames, so that narrowing from assignment + # expressions in the arguments can't affect arguments before the assignment + # in the next inference. The narrowing is applied once the result is known. + assignment_expression_effect = self.chk.assignment_expression_effect + with self.chk.binder.frame_context(can_skip=False, discard=True): + arg_types = self.infer_arg_types_in_empty_context(args) # Step 1: Filter call targets to remove ones where the argument counts don't match plausible_targets = self.plausible_overload_call_targets( arg_types, arg_kinds, arg_names, callee @@ -2924,6 +2955,9 @@ def check_overload_call( unioned_result = None else: inferred_result = None + if unioned_result is not None or inferred_result is not None: + if assignment_expression_effect != self.chk.assignment_expression_effect: + self.apply_argument_narrowing(args) if unioned_result is not None: if inferred_types is not None: for inferred_type in inferred_types: @@ -3071,7 +3105,10 @@ def infer_overload_return_type( for typ in plausible_targets: assert self.msg is self.chk.msg - with self.msg.filter_errors(filter_revealed_type=True) as w: + with ( + self.msg.filter_errors(filter_revealed_type=True) as w, + self.chk.binder.frame_context(can_skip=False, discard=True), + ): with self.chk.local_type_map as m: ret_type, infer_type = self.check_call( callee=typ, @@ -3118,15 +3155,17 @@ def infer_overload_return_type( self.chk.store_types(type_maps[0]) return erase_type(return_types[0]), erase_type(inferred_types[0]) else: - return self.check_call( - callee=AnyType(TypeOfAny.special_form), - args=args, - arg_kinds=arg_kinds, - arg_names=arg_names, - context=context, - callable_name=callable_name, - object_type=object_type, - ) + # The caller applies the narrowing, as for other results. + with self.chk.binder.frame_context(can_skip=False, discard=True): + return self.check_call( + callee=AnyType(TypeOfAny.special_form), + args=args, + arg_kinds=arg_kinds, + arg_names=arg_names, + context=context, + callable_name=callable_name, + object_type=object_type, + ) else: # Success! No ambiguity; return the first match. self.chk.store_types(type_maps[0]) @@ -3434,11 +3473,20 @@ def check_union_call( arg_names: Sequence[str | None] | None, context: Context, ) -> tuple[Type, Type]: + items = callee.relevant_items() + results: list[tuple[Type, Type]] = [] with self.msg.disable_type_names(): - results = [ - self.check_call(subtype, args, arg_kinds, context, arg_names) - for subtype in callee.relevant_items() - ] + for i, subtype in enumerate(items): + if i < len(items) - 1: + # Check the arguments against each item from the same binder state: + # discard narrowing from assignment expressions, except for the last + # item. + with self.chk.binder.frame_context(can_skip=False, discard=True): + results.append( + self.check_call(subtype, args, arg_kinds, context, arg_names) + ) + else: + results.append(self.check_call(subtype, args, arg_kinds, context, arg_names)) return (make_simplified_union([res[0] for res in results]), callee) diff --git a/test-data/unit/check-python38.test b/test-data/unit/check-python38.test index d2ede55840cfe..83604fbe647ac 100644 --- a/test-data/unit/check-python38.test +++ b/test-data/unit/check-python38.test @@ -831,3 +831,100 @@ def fn() -> tuple[int, int]: # E: Unsupported operand types for + ("None" and "int") \ # N: Left operand is of type "int | None" [builtins fixtures/dict.pyi] + +[case testWalrusInLaterArgumentOfGenericCall] +from typing import Optional, TypeVar + +T = TypeVar("T") + +class Node: + next: Optional[Node] + +def pair(a: T, b: object) -> T: ... + +def f(node: Optional[Node]) -> None: + while node is not None: + reveal_type(pair(node, (node := node.next))) # N: Revealed type is "__main__.Node" + reveal_type(node) # N: Revealed type is "__main__.Node | None" + +[case testWalrusInLaterArgumentOfOverloadedCall] +from typing import Optional, Union, overload + +@overload +def f(x: int, y: object) -> int: ... +@overload +def f(x: None, y: object) -> str: ... +def f(x: Optional[int], y: object) -> Union[int, str]: ... + +def g(x: Optional[int]) -> None: + if x is not None: + reveal_type(f(x, (x := None))) # N: Revealed type is "builtins.int" + reveal_type(x) # N: Revealed type is "None" + +[case testWalrusInArgumentOfOverloadedCallNarrowsAfterCall] +from typing import Any, Optional, Union, overload + +class C: ... + +@overload +def f(x: int, y: object) -> int: ... +@overload +def f(x: str, y: object) -> str: ... +def f(x: Union[int, str], y: object) -> Union[int, str]: ... + +def union_math(u: Union[int, str], y: Optional[int]) -> None: + reveal_type(f(u, (y := 1))) # N: Revealed type is "builtins.int | builtins.str" + reveal_type(y) # N: Revealed type is "builtins.int" + +def any_arg(a: Any, y: Optional[int]) -> None: + reveal_type(f(a, (y := 1))) # N: Revealed type is "Any" + reveal_type(y) # N: Revealed type is "builtins.int" + +def no_match(y: Optional[int]) -> None: + f(C(), (y := 1)) # E: No overload variant of "f" matches argument types "C", "int" \ + # N: Possible overload variants: \ + # N: def f(x: int, y: object) -> int \ + # N: def f(x: str, y: object) -> str + reveal_type(y) # N: Revealed type is "builtins.int" + +[case testWalrusInLaterArgumentKeepsAttributeNarrowing] +from typing import Optional, TypeVar + +T = TypeVar("T") + +class A: + x: Optional[int] + +def pair(a: T, b: object) -> T: ... + +def f(a: A, b: A) -> None: + if a.x is not None: + reveal_type(pair(a.x, (a := b))) # N: Revealed type is "builtins.int" + reveal_type(a.x) # N: Revealed type is "builtins.int | None" + +[case testWalrusInLaterArgumentOtherCalls] +from typing import Callable, Optional, TypeVar, Union + +T = TypeVar("T") + +class Node: + next: Optional[Node] + +def star(a: T, *rest: object) -> T: ... + +def star_arg(node: Optional[Node]) -> None: + while node is not None: + reveal_type(star(node, *[(node := node.next)])) # N: Revealed type is "__main__.Node" + +def partial_type(node: Optional[Node]) -> None: + out = [] + while node is not None: + out.append(star(node, (node := node.next))) + reveal_type(out) # N: Revealed type is "builtins.list[__main__.Node]" + +def union_callee( + c: Union[Callable[[Node, object], int], Callable[[Node, object], str]], node: Optional[Node] +) -> None: + if node is not None: + reveal_type(c(node, (node := node.next))) # N: Revealed type is "builtins.int | builtins.str" +[builtins fixtures/list.pyi]