diff --git a/crates/red_knot_python_semantic/resources/mdtest/generics/classes.md b/crates/red_knot_python_semantic/resources/mdtest/generics/classes.md index cc410a43d3..2d0f2b8050 100644 --- a/crates/red_knot_python_semantic/resources/mdtest/generics/classes.md +++ b/crates/red_knot_python_semantic/resources/mdtest/generics/classes.md @@ -53,7 +53,7 @@ class D(C[T]): ... (Examples `E` and `F` from above do not have analogues in the legacy syntax.) -## Inferring generic class parameters +## Specializing generic classes explicitly The type parameter can be specified explicitly: @@ -65,9 +65,50 @@ class C[T]: reveal_type(C[int]()) # revealed: C ``` +The specialization must match the generic types: + +```py +# error: [too-many-positional-arguments] "Too many positional arguments to explicit specialization of class `C`: expected 2, got 3" +reveal_type(C[int, int]()) # revealed: Unknown +``` + +If the type variable has an upper bound, the specialized type must satisfy that bound: + +```py +class Bounded[T: int]: + x: T + +# TODO: revealed: Bounded[int] +reveal_type(Bounded[int]()) # revealed: Bounded + +# TODO: error: [invalid-argument] +reveal_type(Bounded[str]()) # revealed: Bounded +``` + +If the type variable is constrained, the specialized type must satisfy those constraints: + +```py +class Constrained[T: (int, str)]: + x: T + +# TODO: revealed: Constrained[int] +reveal_type(Constrained[int]()) # revealed: Constrained + +# TODO: revealed: Constrained[str] +reveal_type(Constrained[str]()) # revealed: Constrained + +# error: [invalid-argument-type] +reveal_type(Constrained[object]()) # revealed: Unknown +``` + +## Inferring generic class parameters + We can infer the type parameter from a type context: ```py +class C[T]: + x: T + c: C[int] = C() # TODO: revealed: C[int] reveal_type(c) # revealed: C diff --git a/crates/red_knot_python_semantic/src/types/call/arguments.rs b/crates/red_knot_python_semantic/src/types/call/arguments.rs index cce0c81c0b..b37dc09875 100644 --- a/crates/red_knot_python_semantic/src/types/call/arguments.rs +++ b/crates/red_knot_python_semantic/src/types/call/arguments.rs @@ -1,7 +1,9 @@ +use std::borrow::Cow; use std::collections::VecDeque; use std::ops::{Deref, DerefMut}; use super::Type; +use crate::Db; /// Arguments for a single call, in source order. #[derive(Clone, Debug, Default)] @@ -32,6 +34,29 @@ impl<'a> CallArguments<'a> { pub(crate) fn iter(&self) -> impl Iterator> + '_ { self.0.iter().copied() } + + /// Unpacks any subscript tuple arguments into distinct arguments. + pub(crate) fn unpack_subscript_tuples(&self) -> Cow<'_, CallArguments<'a>> { + // If there are no subscript tuples, we can use the existing argument list as-is. + if self + .0 + .iter() + .all(|argument| !matches!(argument, Argument::PositionalSubscriptTuple(_))) + { + return Cow::Borrowed(self); + } + + let mut arguments = VecDeque::with_capacity(self.0.len()); + for argument in self.iter() { + match argument { + Argument::PositionalSubscriptTuple(count) => { + arguments.extend(std::iter::repeat_n(Argument::Positional, count)) + } + _ => arguments.push_back(argument), + } + } + Cow::Owned(CallArguments(arguments)) + } } impl<'a> FromIterator> for CallArguments<'a> { @@ -46,6 +71,8 @@ pub(crate) enum Argument<'a> { Synthetic, /// A positional argument. Positional, + /// A positional argument that is a packed tuple of multiple subscript expression arguments. + PositionalSubscriptTuple(usize), /// A starred positional argument (e.g. `*args`). Variadic, /// A keyword argument (e.g. `a=1`). @@ -54,7 +81,23 @@ pub(crate) enum Argument<'a> { Keywords, } +impl Argument<'_> { + pub(crate) fn subscript_argument<'db>( + db: &'db dyn Db, + slice_type: Type<'db>, + ) -> (Self, Type<'db>) { + match slice_type { + Type::Tuple(tuple) => ( + Argument::PositionalSubscriptTuple(tuple.len(db)), + slice_type, + ), + _ => (Argument::Positional, slice_type), + } + } +} + /// Arguments for a single call, in source order, along with inferred types for each argument. +#[derive(Clone)] pub(crate) struct CallArgumentTypes<'a, 'db> { arguments: CallArguments<'a>, types: VecDeque>, @@ -68,6 +111,16 @@ impl<'a, 'db> CallArgumentTypes<'a, 'db> { Self { arguments, types } } + /// Create a [`CallArgumentTypes`] from an iterator over non-variadic positional argument + /// types. + pub(crate) fn from_arguments( + arguments: impl IntoIterator, Type<'db>)>, + ) -> Self { + let (arguments, types): (VecDeque<_>, VecDeque<_>) = arguments.into_iter().collect(); + let arguments = CallArguments(arguments); + Self { arguments, types } + } + /// Create a [`CallArgumentTypes`] from an iterator over non-variadic positional argument /// types. pub(crate) fn positional(positional_tys: impl IntoIterator>) -> Self { @@ -112,6 +165,37 @@ impl<'a, 'db> CallArgumentTypes<'a, 'db> { pub(crate) fn iter(&self) -> impl Iterator, Type<'db>)> + '_ { self.arguments.iter().zip(self.types.iter().copied()) } + + /// Unpacks any subscript tuple arguments into distinct arguments. + pub(crate) fn unpack_subscript_tuples( + &self, + db: &'db dyn Db, + ) -> Cow<'_, CallArgumentTypes<'a, 'db>> { + // If there are no subscript tuples, we can use the existing argument list as-is. + if self + .arguments + .iter() + .all(|argument| !matches!(argument, Argument::PositionalSubscriptTuple(_))) + { + return Cow::Borrowed(self); + } + + let mut types = VecDeque::with_capacity(self.types.len()); + for (argument, ty) in self.iter() { + match (argument, ty) { + (Argument::PositionalSubscriptTuple(_), Type::Tuple(tuple)) => { + for ty in tuple.iter(db) { + types.push_back(ty); + } + } + _ => types.push_back(ty), + } + } + Cow::Owned(CallArgumentTypes { + arguments: self.arguments.unpack_subscript_tuples().into_owned(), + types, + }) + } } impl<'a> Deref for CallArgumentTypes<'a, '_> { diff --git a/crates/red_knot_python_semantic/src/types/call/bind.rs b/crates/red_knot_python_semantic/src/types/call/bind.rs index b82b8be099..a7bed1dc8a 100644 --- a/crates/red_knot_python_semantic/src/types/call/bind.rs +++ b/crates/red_knot_python_semantic/src/types/call/bind.rs @@ -3,6 +3,8 @@ //! [signatures][crate::types::signatures], we have to handle the fact that the callable might be a //! union of types, each of which might contain multiple overloads. +use std::borrow::Cow; + use smallvec::SmallVec; use super::{ @@ -547,10 +549,12 @@ impl<'db> CallableBinding<'db> { // two phases. // // [1] https://github.com/python/typing/pull/1839 + let callable_type = signature.callable_type; let overloads = signature .into_iter() .map(|signature| { Binding::match_parameters( + callable_type, signature, arguments, argument_forms, @@ -577,8 +581,9 @@ impl<'db> CallableBinding<'db> { // If this callable is a bound method, prepend the self instance onto the arguments list // before checking. argument_types.with_self(signature.bound_type, |argument_types| { + let callable_type = signature.callable_type; for (signature, overload) in signature.iter().zip(&mut self.overloads) { - overload.check_types(db, signature, argument_types); + overload.check_types(db, callable_type, signature, argument_types); } }); } @@ -716,11 +721,30 @@ pub(crate) struct Binding<'db> { impl<'db> Binding<'db> { fn match_parameters( + callable_type: Type<'db>, signature: &Signature<'db>, arguments: &CallArguments<'_>, argument_forms: &mut [Option], conflicting_forms: &mut [bool], ) -> Self { + // Special case: you explicitly specialize a generic class via a subscript expression, + // Class[T1, T2, ...]. This ends up calling a `__class_getitem__` method, defined on + // `type`, which specializes the class. Like all subscript expressions, multiple arguments + // are packed into a tuple before calling `__class_getitem__`. + // + // We would rather treat this as `__class_getitem__` taking in a distinct parameter for + // each type variable. This gives us better error messages if parameter matching fails, and + // makes it easier to type-check the arguments against any type parameter bounds or + // constraints. + let arguments = if matches!( + callable_type, + Type::Callable(CallableType::SpecializeClass(_)) + ) { + arguments.unpack_subscript_tuples() + } else { + Cow::Borrowed(arguments) + }; + let parameters = signature.parameters(); // The parameter that each argument is matched with. let mut argument_parameters = vec![None; arguments.len()]; @@ -743,7 +767,9 @@ impl<'db> Binding<'db> { }; for (argument_index, argument) in arguments.iter().enumerate() { let (index, parameter, positional) = match argument { - Argument::Positional | Argument::Synthetic => { + Argument::Positional + | Argument::PositionalSubscriptTuple(_) + | Argument::Synthetic => { if matches!(argument, Argument::Synthetic) { num_synthetic_args += 1; } @@ -840,9 +866,28 @@ impl<'db> Binding<'db> { fn check_types( &mut self, db: &'db dyn Db, + callable_type: Type<'db>, signature: &Signature<'db>, argument_types: &CallArgumentTypes<'_, 'db>, ) { + // Special case: you explicitly specialize a generic class via a subscript expression, + // Class[T1, T2, ...]. This ends up calling a `__class_getitem__` method, defined on + // `type`, which specializes the class. Like all subscript expressions, multiple arguments + // are packed into a tuple before calling `__class_getitem__`. + // + // We would rather treat this as `__class_getitem__` taking in a distinct parameter for + // each type variable. This gives us better error messages if parameter matching fails, and + // makes it easier to type-check the arguments against any type parameter bounds or + // constraints. + let argument_types = if matches!( + callable_type, + Type::Callable(CallableType::SpecializeClass(_)) + ) { + argument_types.unpack_subscript_tuples(db) + } else { + Cow::Borrowed(argument_types) + }; + let parameters = signature.parameters(); let mut num_synthetic_args = 0; let get_argument_index = |argument_index: usize, num_synthetic_args: usize| { @@ -954,6 +999,10 @@ impl<'db> CallableDescription<'db> { kind: "wrapper descriptor", name: "FunctionType.__get__", }), + Type::Callable(CallableType::SpecializeClass(class)) => Some(CallableDescription { + kind: "explicit specialization of class", + name: class.name(db), + }), _ => None, } } diff --git a/crates/red_knot_python_semantic/src/types/generics.rs b/crates/red_knot_python_semantic/src/types/generics.rs index 3defbf180d..6a449c1cc6 100644 --- a/crates/red_knot_python_semantic/src/types/generics.rs +++ b/crates/red_knot_python_semantic/src/types/generics.rs @@ -62,8 +62,9 @@ impl<'db> GenericContext<'db> { fn parameter_from_typevar(db: &'db dyn Db, typevar: &TypeVarInstance<'db>) -> Parameter<'db> { let mut parameter = Parameter::positional_only(Some(typevar.name(db).clone())); match typevar.bound_or_constraints(db) { - Some(TypeVarBoundOrConstraints::UpperBound(bound)) => { - parameter = parameter.with_annotated_type(bound); + Some(TypeVarBoundOrConstraints::UpperBound(_)) => { + // TODO: This should be TypeForm[bound] + parameter = parameter.with_annotated_type(Type::any()); } Some(TypeVarBoundOrConstraints::Constraints(constraints)) => { parameter = parameter diff --git a/crates/red_knot_python_semantic/src/types/infer.rs b/crates/red_knot_python_semantic/src/types/infer.rs index cb4e6cd3de..e413edf314 100644 --- a/crates/red_knot_python_semantic/src/types/infer.rs +++ b/crates/red_knot_python_semantic/src/types/infer.rs @@ -1974,7 +1974,7 @@ impl<'db> TypeInferenceBuilder<'db> { let tuple = TupleType::new( self.db(), elts.iter() - .map(|expr| self.infer_type_expression(expr)) + .map(|expr| self.infer_expression(expr)) .collect::>(), ); let constraints = TypeVarBoundOrConstraints::Constraints(tuple); @@ -5711,11 +5711,11 @@ impl<'db> TypeInferenceBuilder<'db> { // If the class defines `__getitem__`, return its return type. // // See: https://docs.python.org/3/reference/datamodel.html#class-getitem-versus-getitem - match value_ty.try_call_dunder( + let arguments = CallArgumentTypes::from_arguments([Argument::subscript_argument( self.db(), - "__getitem__", - CallArgumentTypes::positional([slice_ty]), - ) { + slice_ty, + )]); + match value_ty.try_call_dunder(self.db(), "__getitem__", arguments) { Ok(outcome) => return outcome.return_type(self.db()), Err(err @ CallDunderError::PossiblyUnbound { .. }) => { self.context.report_lint( @@ -5774,21 +5774,14 @@ impl<'db> TypeInferenceBuilder<'db> { ); } - match ty.try_call( - self.db(), - CallArgumentTypes::positional([value_ty, slice_ty]), - ) { + let arguments = CallArgumentTypes::from_arguments([ + (Argument::Synthetic, value_ty), + Argument::subscript_argument(self.db(), slice_ty), + ]); + match ty.try_call(self.db(), arguments) { Ok(bindings) => return bindings.return_type(self.db()), Err(CallError(_, bindings)) => { - self.context.report_lint( - &CALL_NON_CALLABLE, - value_node, - format_args!( - "Method `__class_getitem__` of type `{}` is not callable on object of type `{}`", - bindings.callable_type().display(self.db()), - value_ty.display(self.db()), - ), - ); + bindings.report_diagnostics(&self.context, value_node.into()); return bindings.return_type(self.db()); } } @@ -5799,12 +5792,6 @@ impl<'db> TypeInferenceBuilder<'db> { if class.is_known(self.db(), KnownClass::Type) { return KnownClass::GenericAlias.to_instance(self.db()); } - - if class.generic_context(self.db()).is_some() { - // TODO: specialize the generic class using these explicit type - // variable assignments - return value_ty; - } } report_non_subscriptable( @@ -7512,7 +7499,7 @@ mod tests { check_typevar("T", None, None, None); check_typevar("U", Some("A"), None, None); - check_typevar("V", None, Some(&["A", "B"]), None); + check_typevar("V", None, Some(&["Literal[A]", "Literal[B]"]), None); check_typevar("W", None, None, Some("A")); check_typevar("X", Some("A"), None, Some("A1"));