Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
40 changes: 34 additions & 6 deletions mypy/binder.py
Original file line number Diff line number Diff line change
Expand Up @@ -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})"
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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
Expand All @@ -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
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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])
Expand Down Expand Up @@ -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
Expand Down
102 changes: 75 additions & 27 deletions mypy/checkexpr.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 = [
Expand Down Expand Up @@ -1214,15 +1214,15 @@ 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 (
typename in self.container_args
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]:
Expand Down Expand Up @@ -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):
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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()
Expand Down Expand Up @@ -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.

Expand Down Expand Up @@ -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
)
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
)
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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])
Expand Down Expand Up @@ -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)

Expand Down
Loading
Loading