This commit is contained in:
Douglas Creager
2025-03-26 13:13:46 -04:00
parent b7afaa219a
commit 5e74cf07fb
5 changed files with 192 additions and 30 deletions

View File

@@ -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

View File

@@ -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<Item = Argument<'a>> + '_ {
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<Argument<'a>> 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<Type<'db>>,
@@ -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<Item = (Argument<'a>, 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<Item = Type<'db>>) -> Self {
@@ -112,6 +165,37 @@ impl<'a, 'db> CallArgumentTypes<'a, 'db> {
pub(crate) fn iter(&self) -> impl Iterator<Item = (Argument<'a>, 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, '_> {

View File

@@ -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<ParameterForm>],
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,
}
}

View File

@@ -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

View File

@@ -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::<Box<_>>(),
);
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"));