diff --git a/mypy/checker.py b/mypy/checker.py index 6a0e51d00913..69522410e5be 100644 --- a/mypy/checker.py +++ b/mypy/checker.py @@ -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: @@ -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) @@ -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) @@ -1825,7 +1828,7 @@ 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() @@ -1833,7 +1836,7 @@ def check_unbound_return_typevar(self, typ: CallableType) -> None: 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) @@ -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, ) @@ -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) @@ -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 @@ -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) diff --git a/mypy/checkexpr.py b/mypy/checkexpr.py index 04e2f7102478..3b28cbfe275c 100644 --- a/mypy/checkexpr.py +++ b/mypy/checkexpr.py @@ -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)): diff --git a/test-data/unit/check-classes.test b/test-data/unit/check-classes.test index dae74d9031e4..9f0d5462cba3 100644 --- a/test-data/unit/check-classes.test +++ b/test-data/unit/check-classes.test @@ -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 diff --git a/test-data/unit/check-functions.test b/test-data/unit/check-functions.test index 4cb58820d5bf..445b5ce26ee4 100644 --- a/test-data/unit/check-functions.test +++ b/test-data/unit/check-functions.test @@ -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 @@ -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