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
53 changes: 28 additions & 25 deletions mypy/checker.py
Original file line number Diff line number Diff line change
Expand Up @@ -1453,8 +1453,9 @@ def check_func_def(
self, defn: FuncItem, typ: CallableType, name: str | None, allow_empty: bool = False
) -> None:
"""Type check a function definition."""
if defn.type_args:
self.check_typevar_defaults(typ.variables, defn)
# Expand type variables with value restrictions to ordinary types.
self.check_typevar_defaults(typ.variables)
expanded = self.expand_typevars(defn, typ)
original_typ = typ
for item, typ in expanded:
Expand Down Expand Up @@ -1508,11 +1509,12 @@ def check_func_def(
not in {"__init__", "__new__", "__post_init__", "__replace__"}
and not is_private(defn.name) # private methods are not inherited
and (i != 0 or not found_self)
and not isinstance(defn, LambdaExpr)
):
ctx: Context = arg_type
if ctx.line < 0:
ctx = typ
self.fail(message_registry.FUNCTION_PARAMETER_CANNOT_BE_COVARIANT, ctx)
self.fail(
message_registry.FUNCTION_PARAMETER_CANNOT_BE_COVARIANT,
defn.arguments[i],
)
# Need to store arguments again for the expanded item.
store_argument_type(item, i, typ, self.named_generic_type)

Expand Down Expand Up @@ -1724,23 +1726,24 @@ def check_funcdef_item(
self.check_setattr_method(typ, defn)

# Refuse contravariant return type variable
if isinstance(typ.ret_type, TypeVarType):
if typ.ret_type.variance == CONTRAVARIANT:
self.fail(message_registry.RETURN_TYPE_CANNOT_BE_CONTRAVARIANT, typ.ret_type)
self.check_unbound_return_typevar(typ)
elif isinstance(original_typ.ret_type, TypeVarType) and original_typ.ret_type.values:
# Since type vars with values are expanded, the return type is changed
# to a raw value. This is a hack to get it back.
self.check_unbound_return_typevar(original_typ)
if not isinstance(item, LambdaExpr):
if isinstance(typ.ret_type, TypeVarType):
if typ.ret_type.variance == CONTRAVARIANT:
self.fail(message_registry.RETURN_TYPE_CANNOT_BE_CONTRAVARIANT, defn)
self.check_unbound_return_typevar(typ, defn)
elif isinstance(original_typ.ret_type, TypeVarType) and original_typ.ret_type.values:
# Since type vars with values are expanded, the return type is changed
# to a raw value. This is a hack to get it back.
self.check_unbound_return_typevar(original_typ, defn)

# Check that Generator functions have the appropriate return type.
if defn.is_generator:
if defn.is_async_generator:
if not self.is_async_generator_return_type(typ.ret_type):
self.fail(message_registry.INVALID_RETURN_TYPE_FOR_ASYNC_GENERATOR, typ)
self.fail(message_registry.INVALID_RETURN_TYPE_FOR_ASYNC_GENERATOR, defn)
else:
if not self.is_generator_return_type(typ.ret_type, defn.is_coroutine):
self.fail(message_registry.INVALID_RETURN_TYPE_FOR_GENERATOR, typ)
self.fail(message_registry.INVALID_RETURN_TYPE_FOR_GENERATOR, defn)

def require_correct_self_argument(self, func: Type, defn: FuncDef) -> bool:
func = get_proper_type(func)
Expand Down Expand Up @@ -1825,15 +1828,15 @@ def is_var_redefined_in_outer_context(self, v: Var, after_line: int) -> bool:
return True
return False

def check_unbound_return_typevar(self, typ: CallableType) -> None:
def check_unbound_return_typevar(self, typ: CallableType, context: Context) -> None:
"""Fails when the return typevar is not defined in arguments."""
if isinstance(typ.ret_type, TypeVarType) and typ.ret_type in typ.variables:
arg_type_visitor = CollectArgTypeVarTypes()
for argtype in typ.arg_types:
argtype.accept(arg_type_visitor)

if typ.ret_type not in arg_type_visitor.arg_types:
self.fail(message_registry.UNBOUND_TYPEVAR, typ.ret_type, code=TYPE_VAR)
self.fail(message_registry.UNBOUND_TYPEVAR, context, code=TYPE_VAR)
upper_bound = get_proper_type(typ.ret_type.upper_bound)
if not (
isinstance(upper_bound, Instance)
Expand All @@ -1842,7 +1845,7 @@ def check_unbound_return_typevar(self, typ: CallableType) -> None:
self.note(
"Consider using the upper bound "
f"{format_type(typ.ret_type.upper_bound, self.options)} instead",
context=typ.ret_type,
context=context,
code=TYPE_VAR,
)

Expand Down Expand Up @@ -2883,8 +2886,8 @@ def visit_class_def(self, defn: ClassDef) -> None:
context=defn,
code=codes.TYPE_VAR,
)
if typ.defn.type_vars:
self.check_typevar_defaults(typ.defn.type_vars)
if typ.defn.type_args:
self.check_typevar_defaults(typ.defn.type_vars, typ.defn)

if typ.is_protocol and typ.defn.type_vars:
self.check_protocol_variance(defn)
Expand Down Expand Up @@ -2948,14 +2951,14 @@ def check_init_subclass(self, defn: ClassDef) -> None:
# all other bases have already been checked.
break

def check_typevar_defaults(self, tvars: Sequence[TypeVarLikeType]) -> None:
def check_typevar_defaults(self, tvars: Sequence[TypeVarLikeType], context: Context) -> None:
for tv in tvars:
if not (isinstance(tv, TypeVarType) and tv.has_default()):
continue
if not is_subtype(tv.default, tv.upper_bound):
self.fail("TypeVar default must be a subtype of the bound type", tv)
self.fail("TypeVar default must be a subtype of the bound type", context)
if tv.values and not any(is_same_type(tv.default, value) for value in tv.values):
self.fail("TypeVar default must be one of the constraint types", tv)
self.fail("TypeVar default must be one of the constraint types", context)

def check_enum(self, defn: ClassDef) -> None:
assert defn.info.is_enum
Expand Down Expand Up @@ -6244,8 +6247,8 @@ def check_and_remove_capture_conflicts(
del type_map[expr]

def visit_type_alias_stmt(self, o: TypeAliasStmt) -> None:
if o.alias_node:
self.check_typevar_defaults(o.alias_node.alias_tvars)
if o.alias_node and o.type_args:
self.check_typevar_defaults(o.alias_node.alias_tvars, o)

with self.msg.filter_errors():
self.expr_checker.accept(o.value)
Expand Down
8 changes: 1 addition & 7 deletions mypy/checkexpr.py
Original file line number Diff line number Diff line change
Expand Up @@ -6450,14 +6450,8 @@ def visit_yield_from_expr(self, e: YieldFromExpr, allow_none_return: bool = Fals
elif self.chk.type_is_iterable(subexpr_type):
if is_async_def(subexpr_type) and not has_coroutine_decorator(return_type):
self.chk.msg.yield_from_invalid_operand_type(subexpr_type, e)

any_type = AnyType(TypeOfAny.special_form)
generic_generator_type = self.chk.named_generic_type(
"typing.Generator", [any_type, any_type, any_type]
)
generic_generator_type.set_line(e)
iter_type, _ = self.check_method_call_by_name(
"__iter__", subexpr_type, [], [], context=generic_generator_type
"__iter__", subexpr_type, [], [], context=e
)
else:
if not (is_async_def(subexpr_type) and has_coroutine_decorator(return_type)):
Expand Down
13 changes: 13 additions & 0 deletions test-data/unit/check-classes.test
Original file line number Diff line number Diff line change
Expand Up @@ -9776,3 +9776,16 @@ class C:

reveal_type(C.x) # N: Revealed type is "builtins.int | None"
[builtins fixtures/classmethod.pyi]

[case testTypeVarDefaultErrorLocation]
from typing import Generic

from lib import T

# some spaces

class C(Generic[T]): ...

[file lib.py]
from typing import TypeVar
T = TypeVar("T", default=int, bound=str) # E: TypeVar default must be a subtype of the bound type
23 changes: 16 additions & 7 deletions test-data/unit/check-functions.test
Original file line number Diff line number Diff line change
Expand Up @@ -2197,21 +2197,16 @@ class A(Generic[t]):
[out]
main:6: error: Cannot use a covariant type variable as a parameter

[case testRejectCovariantArgumentInLambda]
[case testAllowCovariantArgumentInLambda]
from typing import TypeVar, Generic, Callable

t = TypeVar('t', covariant=True)
class Thing(Generic[t]):
def chain(self, func: Callable[[t], None]) -> None: pass
def end(self) -> None:
return self.chain( # Note that lambda args have no line numbers
return self.chain(
lambda _: None)
[builtins fixtures/bool.pyi]
[out]
main:8: error: Cannot use a covariant type variable as a parameter

[case testRejectCovariantArgumentInLambdaSplitLine]
from typing import TypeVar, Generic, Callable

[case testRejectContravariantReturnType]
# flags: --no-strict-optional
Expand Down Expand Up @@ -3948,3 +3943,17 @@ convert4("hello", 3.15) # E: Missing positional arguments "third", "fourth" in
convert4(b'', "hello", 3.15) # E: Missing positional argument "fourth" in call to "convert4" \
# E: Argument 1 to "convert4" has incompatible type "bytes"; expected "int"
[builtins fixtures/primitives.pyi]

[case testGenericContextLambdaNoError]
from lib import takes_lambda
takes_lambda(lambda x: 1)

[file lib.py]
from typing import Any, Callable, TypeVar

# extra spaces

Ex = TypeVar("Ex", covariant=True)

def takes_lambda(func: Callable[[Ex], Any]) -> None:
pass
Loading