Skip to content

Commit b5b279f

Browse files
authored
Find precise fallback for ambiguous Any overloads (#22060)
It is a common pattern (especially in numeric libraries) to have a fallback overload that has a relatively precise return type (i.e. not just `Any` or `list[Any]`). We should try to find and use that overload if there is an ambiguity caused by an argument that contains `Any`. This will likely cause some new errors, but this should limit the fallout from #22011
1 parent 25d86c6 commit b5b279f

5 files changed

Lines changed: 85 additions & 33 deletions

File tree

‎mypy/checker.py‎

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -222,6 +222,7 @@ def __init__(self) -> None:
222222
from mypy.state import state
223223
from mypy.subtypes import (
224224
find_member,
225+
has_any_type,
225226
infer_class_variances,
226227
is_callable_compatible,
227228
is_equivalent,
@@ -6011,7 +6012,7 @@ def check_untyped_after_decorator(self, typ: Type, func: FuncDef) -> None:
60116012
if not self.options.disallow_any_decorated or self.is_stub or self.current_node_deferred:
60126013
return
60136014

6014-
if mypy.checkexpr.has_any_type(typ):
6015+
if has_any_type(typ):
60156016
self.msg.untyped_decorated_function(typ, func)
60166017

60176018
def check_async_with_item(

‎mypy/checkexpr.py‎

Lines changed: 12 additions & 31 deletions
Original file line numberDiff line numberDiff line change
@@ -122,8 +122,10 @@
122122
from mypy.semanal_enum import ENUM_BASES
123123
from mypy.state import state
124124
from mypy.subtypes import (
125+
common_type,
125126
covers_at_runtime,
126127
find_member,
128+
has_any_type,
127129
is_same_type,
128130
is_subtype,
129131
merge_typevars_in_callables_by_name,
@@ -204,6 +206,7 @@
204206
has_recursive_types,
205207
has_type_vars,
206208
is_named_instance,
209+
remove_dups,
207210
split_with_prefix_and_suffix,
208211
)
209212
from mypy.types_utils import (
@@ -228,6 +231,9 @@
228231
# see https://github.com/python/mypy/pull/5255#discussion_r196896335 for discussion.
229232
MAX_UNIONS: Final = 5
230233

234+
# Maximum number or unique matched overload return types caused by Any
235+
# ambiguity where we try to find a precise fallback.
236+
MAX_PRECISE_OVERLOAD_FALLBACK: Final = 8
231237

232238
# Types considered safe for comparisons with --strict-equality due to known behaviour of __eq__.
233239
# NOTE: All these types are subtypes of AbstractSet.
@@ -3097,11 +3103,17 @@ def infer_overload_return_type(
30973103
if not matches:
30983104
return None
30993105
elif any_causes_overload_ambiguity(matches, return_types, arg_types, arg_kinds, arg_names):
3106+
return_types = remove_dups(return_types)
31003107
# An argument of type or containing the type 'Any' caused ambiguity.
31013108
# We try returning a precise type if we can. If not, we give up and just return 'Any'.
31023109
if all_same_types(return_types):
31033110
self.chk.store_types(type_maps[0])
31043111
return return_types[0], inferred_types[0]
3112+
elif len(return_types) < MAX_PRECISE_OVERLOAD_FALLBACK and (
3113+
common := common_type(return_types)
3114+
):
3115+
self.chk.store_types(type_maps[0])
3116+
return common, erase_type(inferred_types[0])
31053117
elif all_same_types([erase_type(typ) for typ in return_types]):
31063118
self.chk.store_types(type_maps[0])
31073119
return erase_type(return_types[0]), erase_type(inferred_types[0])
@@ -6660,37 +6672,6 @@ def try_parse_as_type_expression(self, maybe_type_expr: Expression) -> Type | No
66606672
return None
66616673

66626674

6663-
def has_any_type(t: Type, ignore_in_type_obj: bool = False) -> bool:
6664-
"""Whether t contains an Any type"""
6665-
return t.accept(HasAnyType(ignore_in_type_obj))
6666-
6667-
6668-
class HasAnyType(types.BoolTypeQuery):
6669-
def __init__(self, ignore_in_type_obj: bool) -> None:
6670-
super().__init__(types.ANY_STRATEGY)
6671-
self.ignore_in_type_obj = ignore_in_type_obj
6672-
6673-
def visit_any(self, t: AnyType) -> bool:
6674-
return t.type_of_any != TypeOfAny.special_form # special forms are not real Any types
6675-
6676-
def visit_callable_type(self, t: CallableType) -> bool:
6677-
if self.ignore_in_type_obj and t.is_type_obj():
6678-
return False
6679-
return super().visit_callable_type(t)
6680-
6681-
def visit_type_var(self, t: TypeVarType) -> bool:
6682-
default = [t.default] if t.has_default() else []
6683-
return self.query_types([t.upper_bound, *default] + t.values)
6684-
6685-
def visit_param_spec(self, t: ParamSpecType) -> bool:
6686-
default = [t.default] if t.has_default() else []
6687-
return self.query_types([t.upper_bound, *default, t.prefix])
6688-
6689-
def visit_type_var_tuple(self, t: TypeVarTupleType) -> bool:
6690-
default = [t.default] if t.has_default() else []
6691-
return self.query_types([t.upper_bound, *default])
6692-
6693-
66946675
def has_coroutine_decorator(t: Type) -> bool:
66956676
"""Whether t came from a function decorated with `@coroutine`."""
66966677
t = get_proper_type(t)

‎mypy/subtypes.py‎

Lines changed: 50 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -36,6 +36,7 @@
3636
)
3737
from mypy.options import Options
3838
from mypy.state import state
39+
from mypy.type_visitor import ANY_STRATEGY, BoolTypeQuery
3940
from mypy.types import (
4041
MAX_PROTOCOL_DEPTH,
4142
MYPYC_NATIVE_INT_NAMES,
@@ -2537,6 +2538,55 @@ def erase_return_self_types(typ: Type, self_type: Instance) -> Type:
25372538
return typ
25382539

25392540

2541+
def common_type(types: list[Type]) -> Type | None:
2542+
"""Return a type in the list that is both subtype and supertype of all other types.
2543+
2544+
If there are more than one such type, choose the one that has an Any component,
2545+
otherwise return None.
2546+
"""
2547+
candidates = []
2548+
for candidate in types:
2549+
if all(is_equivalent(candidate, other) for other in types):
2550+
candidates.append(candidate)
2551+
if len(candidates) == 1:
2552+
return candidates[0]
2553+
candidates = [c for c in candidates if has_any_type(c)]
2554+
if len(candidates) == 1:
2555+
return candidates[0]
2556+
return None
2557+
2558+
2559+
def has_any_type(t: Type, ignore_in_type_obj: bool = False) -> bool:
2560+
"""Whether t contains an Any type"""
2561+
return t.accept(HasAnyType(ignore_in_type_obj))
2562+
2563+
2564+
class HasAnyType(BoolTypeQuery):
2565+
def __init__(self, ignore_in_type_obj: bool) -> None:
2566+
super().__init__(ANY_STRATEGY)
2567+
self.ignore_in_type_obj = ignore_in_type_obj
2568+
2569+
def visit_any(self, t: AnyType) -> bool:
2570+
return t.type_of_any != TypeOfAny.special_form # special forms are not real Any types
2571+
2572+
def visit_callable_type(self, t: CallableType) -> bool:
2573+
if self.ignore_in_type_obj and t.is_type_obj():
2574+
return False
2575+
return super().visit_callable_type(t)
2576+
2577+
def visit_type_var(self, t: TypeVarType) -> bool:
2578+
default = [t.default] if t.has_default() else []
2579+
return self.query_types([t.upper_bound, *default] + t.values)
2580+
2581+
def visit_param_spec(self, t: ParamSpecType) -> bool:
2582+
default = [t.default] if t.has_default() else []
2583+
return self.query_types([t.upper_bound, *default, t.prefix])
2584+
2585+
def visit_type_var_tuple(self, t: TypeVarTupleType) -> bool:
2586+
default = [t.default] if t.has_default() else []
2587+
return self.query_types([t.upper_bound, *default])
2588+
2589+
25402590
def is_erased_instance(t: Instance) -> bool:
25412591
"""Is this an instance where all args are Any types?"""
25422592
if not t.args:

‎mypy/suggestions.py‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -34,7 +34,6 @@
3434

3535
from mypy.argmap import map_actuals_to_formals
3636
from mypy.build import Graph, State
37-
from mypy.checkexpr import has_any_type
3837
from mypy.find_sources import InvalidSourceList, SourceFinder
3938
from mypy.join import join_type_list
4039
from mypy.meet import meet_type_list
@@ -59,6 +58,7 @@
5958
from mypy.plugin import FunctionContext, MethodContext, Plugin
6059
from mypy.server.update import FineGrainedBuildManager
6160
from mypy.state import state
61+
from mypy.subtypes import has_any_type
6262
from mypy.traverser import TraverserVisitor
6363
from mypy.typeops import bind_self, make_simplified_union
6464
from mypy.types import (

‎test-data/unit/check-overloading.test‎

Lines changed: 20 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -6932,3 +6932,23 @@ class B(A):
69326932
def f(self, y: str) -> None: ... # This is currently allowed (note different name)
69336933
def f(self, *args, **kwargs) -> None: ...
69346934
[builtins fixtures/tuple.pyi]
6935+
6936+
[case testRespectPreciseOverloadFallbackIfPossible]
6937+
from typing import overload, Any
6938+
6939+
@overload
6940+
def f(x: list[int]) -> list[tuple[int, ...]]: ...
6941+
@overload
6942+
def f(x: list[Any]) -> list[tuple[Any, ...]]: ...
6943+
def f(x): pass
6944+
6945+
a: Any
6946+
la: list[Any]
6947+
li: list[int]
6948+
ls: list[str]
6949+
6950+
reveal_type(f(a)) # N: Revealed type is "builtins.list[builtins.tuple[Any, ...]]"
6951+
reveal_type(f(la)) # N: Revealed type is "builtins.list[builtins.tuple[Any, ...]]"
6952+
reveal_type(f(li)) # N: Revealed type is "builtins.list[builtins.tuple[builtins.int, ...]]"
6953+
reveal_type(f(ls)) # N: Revealed type is "builtins.list[builtins.tuple[Any, ...]]"
6954+
[builtins fixtures/tuple.pyi]

0 commit comments

Comments
 (0)