Skip to content

Commit bbd7904

Browse files
authored
Fix protocol subtyping in case of restricted self-types (#22106)
Fixes #19341 Fixes #20061 This is a bit hacky, but an alternative would be to change return type of a lot of functions in `checkmember.py`. Also although this is a niche use case, it seems to be important for `numpy` and `pandas`.
1 parent 7907393 commit bbd7904

3 files changed

Lines changed: 50 additions & 18 deletions

File tree

‎mypy/checkmember.py‎

Lines changed: 22 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -101,6 +101,7 @@ def __init__(
101101
rvalue: Expression | None = None,
102102
suppress_errors: bool = False,
103103
preserve_type_var_ids: bool = False,
104+
attribute_error: bool = False,
104105
) -> None:
105106
self.is_lvalue = is_lvalue
106107
self.is_super = is_super
@@ -121,6 +122,13 @@ def __init__(
121122
# It is needed to avoid infinite recursion in cases involving self-referential
122123
# generic methods, see find_member() for details. Do not use for other purposes!
123124
self.preserve_type_var_ids = preserve_type_var_ids
125+
# This is again used for protocol subtype checks. Normally protocol checks first
126+
# quickly reject an attribute access by checking if there is a relevant symbol.
127+
# This doesn't work in case of restricted self-types. We cannot rely on error
128+
# messages, as those may be emitted in cases where an attribute access is valid
129+
# for the purposes of structural subtyping. So instead we use this mutable attribute
130+
# to indicate such failures reliably.
131+
self.attribute_error = attribute_error
124132

125133
def named_type(self, name: str) -> Instance:
126134
return self.chk.named_type(name)
@@ -152,6 +160,7 @@ def copy_modified(
152160
rvalue=self.rvalue,
153161
suppress_errors=self.suppress_errors,
154162
preserve_type_var_ids=self.preserve_type_var_ids,
163+
attribute_error=self.attribute_error,
155164
)
156165
if self_type is not None:
157166
mx.self_type = self_type
@@ -375,9 +384,7 @@ def analyze_instance_member_access(
375384
if isinstance(method, (FuncDef, OverloadedFuncDef)) and method.is_trivial_self:
376385
signature = bind_self_fast(signature, mx.self_type)
377386
else:
378-
signature = check_self_arg(
379-
signature, mx.self_type, method.is_class, mx.context, name, mx.msg
380-
)
387+
signature = check_self_arg(signature, method.is_class, name, mx)
381388
signature = bind_self(signature, mx.self_type, is_classmethod=method.is_class)
382389
typ = map_instance_to_supertype(typ, method.info)
383390
member_type = expand_type_by_instance(signature, typ)
@@ -501,6 +508,9 @@ def analyze_union_member_access(name: str, typ: UnionType, mx: MemberContext) ->
501508
# Self types should be bound to every individual item of a union.
502509
item_mx = mx.copy_modified(self_type=subtype)
503510
results.append(_analyze_member_access(name, subtype, item_mx))
511+
if item_mx.attribute_error:
512+
# Record an error if there is an error on at least one union item.
513+
mx.attribute_error = True
504514
return make_simplified_union(results)
505515

506516

@@ -973,7 +983,7 @@ def expand_and_bind_callable(
973983
if is_trivial_self:
974984
typ = bind_self_fast(typ, mx.self_type)
975985
else:
976-
typ = check_self_arg(typ, mx.self_type, var.is_classmethod, mx.context, name, mx.msg)
986+
typ = check_self_arg(typ, var.is_classmethod, name, mx)
977987
typ = bind_self(typ, mx.self_type, var.is_classmethod)
978988
expanded = expand_type_by_instance(typ, itype)
979989
freeze_all_type_vars(expanded)
@@ -1050,12 +1060,7 @@ def expand_self_type_if_needed(
10501060

10511061

10521062
def check_self_arg(
1053-
functype: FunctionLike,
1054-
dispatched_arg_type: Type,
1055-
is_classmethod: bool,
1056-
context: Context,
1057-
name: str,
1058-
msg: MessageBuilder,
1063+
functype: FunctionLike, is_classmethod: bool, name: str, mx: MemberContext
10591064
) -> FunctionLike:
10601065
"""Check that an instance has a valid type for a method with annotated 'self'.
10611066
@@ -1070,14 +1075,15 @@ def f(self: S) -> T: ...
10701075
if not items:
10711076
return functype
10721077
new_items = []
1078+
dispatched_arg_type = mx.self_type
10731079
if is_classmethod:
10741080
dispatched_arg_type = TypeType.make_normalized(dispatched_arg_type)
10751081
p_dispatched_arg_type = get_proper_type(dispatched_arg_type)
10761082

10771083
for item in items:
10781084
if not item.arg_types or item.arg_kinds[0] not in (ARG_POS, ARG_STAR):
10791085
# No positional first (self) argument (*args is okay).
1080-
msg.no_formal_self(name, item, context)
1086+
mx.msg.no_formal_self(name, item, mx.context)
10811087
# This is pretty bad, so just return the original signature if
10821088
# there is at least one such error.
10831089
return functype
@@ -1119,9 +1125,10 @@ def f(self: S) -> T: ...
11191125
raise NotImplementedError
11201126
if not new_items:
11211127
# Choose first item for the message (it may be not very helpful for overloads).
1122-
msg.incompatible_self_argument(
1123-
name, dispatched_arg_type, items[0], is_classmethod, context
1128+
mx.msg.incompatible_self_argument(
1129+
name, dispatched_arg_type, items[0], is_classmethod, mx.context
11241130
)
1131+
mx.attribute_error = True
11251132
return functype
11261133
if len(new_items) == 1:
11271134
return new_items[0]
@@ -1287,7 +1294,7 @@ def analyze_class_attribute_access(
12871294
and not is_trivial_self
12881295
and not t.bound()
12891296
):
1290-
t = check_self_arg(t, mx.self_type, False, mx.context, name, mx.msg)
1297+
t = check_self_arg(t, False, name, mx)
12911298
t = add_class_tvars(
12921299
t,
12931300
isuper,
@@ -1492,7 +1499,7 @@ def analyze_decorator_or_funcbase_access(
14921499
typ = mx.chk.function_type(defn)
14931500
if isinstance(defn, (FuncDef, OverloadedFuncDef)) and defn.is_trivial_self:
14941501
return bind_self_fast(typ, mx.self_type)
1495-
typ = check_self_arg(typ, mx.self_type, defn.is_class, mx.context, name, mx.msg)
1502+
typ = check_self_arg(typ, defn.is_class, name, mx)
14961503
return bind_self(typ, original_type=mx.self_type, is_classmethod=defn.is_class)
14971504

14981505

‎mypy/subtypes.py‎

Lines changed: 12 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -1406,7 +1406,12 @@ def get_protocol_member(
14061406
# if constructor signature didn't match, this can cause many false negatives.
14071407
return None
14081408

1409-
subtype = find_member(member, left, original_left, class_obj=class_obj, is_lvalue=is_lvalue)
1409+
# We only reject attributes that give an error on subtype. This is not very principled,
1410+
# but this simplifies logic because we apply subtype as a self-type for supertype
1411+
# attribute access.
1412+
subtype = find_member(
1413+
member, left, original_left, class_obj=class_obj, is_lvalue=is_lvalue, strict_attrs=True
1414+
)
14101415
if isinstance(subtype, PartialType):
14111416
subtype = (
14121417
NoneType()
@@ -1426,6 +1431,7 @@ def find_member(
14261431
is_operator: bool = False,
14271432
class_obj: bool = False,
14281433
is_lvalue: bool = False,
1434+
strict_attrs: bool = False,
14291435
) -> Type | None:
14301436
type_checker = checker_state.type_checker
14311437
if type_checker is None:
@@ -1488,9 +1494,12 @@ def find_member(
14881494
with type_checker.msg.filter_errors(filter_deprecated=True):
14891495
if class_obj:
14901496
fallback = itype.type.metaclass_type or mx.named_type("builtins.type")
1491-
return analyze_class_attribute_access(itype, name, mx, mcs_fallback=fallback)
1497+
result = analyze_class_attribute_access(itype, name, mx, mcs_fallback=fallback)
14921498
else:
1493-
return analyze_instance_member_access(name, itype, mx, info)
1499+
result = analyze_instance_member_access(name, itype, mx, info)
1500+
if strict_attrs and mx.attribute_error:
1501+
return None
1502+
return result
14941503

14951504

14961505
def find_member_simple(

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

Lines changed: 16 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -4803,3 +4803,19 @@ class C:
48034803
# This should not cause infinite recursion.
48044804
x: P[int] = C()
48054805
[builtins fixtures/tuple.pyi]
4806+
4807+
[case testLateSelfArgCheckRejectsProtoImpl]
4808+
from typing import Generic, Protocol, TypeVar
4809+
4810+
T = TypeVar("T")
4811+
4812+
class SupportsFoo(Protocol):
4813+
def foo(self) -> None: ...
4814+
4815+
class Bad(Generic[T]):
4816+
x: T
4817+
def foo(self: Bad[int]) -> None: ...
4818+
4819+
bad: Bad[str] = Bad()
4820+
fail: SupportsFoo = bad # E: Incompatible types in assignment (expression has type "Bad[str]", variable has type "SupportsFoo")
4821+
[builtins fixtures/tuple.pyi]

0 commit comments

Comments
 (0)