diff --git a/python/python-psi-impl/resources/messages/PyPsiBundle.properties b/python/python-psi-impl/resources/messages/PyPsiBundle.properties
index f593731fc7a4..a8ae335b8d2d 100644
--- a/python/python-psi-impl/resources/messages/PyPsiBundle.properties
+++ b/python/python-psi-impl/resources/messages/PyPsiBundle.properties
@@ -1180,6 +1180,8 @@ INSP.type.hints.invalid.type.argument=Invalid type argument
INSP.type.hints.generic.type.alias.is.not.generic.or.already.parameterized=Type alias is not generic or already specialized
INSP.type.hints.default.type.must.be.type.expression=Default type must be a type expression
INSP.type.hints.invalid.type.expression=Invalid type expression
+INSP.type.hints.default.type.do.not.match.constraints=Default type of TypeVar must be one of the constraint types
+INSP.type.hints.default.type.do.not.match.bounds=Default type of TypeVar is not a subtype of the bound
INSP.type.hints.illegal.callable.format='Callable' must be used as 'Callable[[arg, ...], result]'
INSP.type.hints.illegal.first.parameter='Callable' first parameter must be a parameter expression
INSP.type.hints.parameters.to.generic.types.must.be.types=Parameters to generic types must be types
diff --git a/python/python-psi-impl/src/com/jetbrains/python/codeInsight/typing/PyTypingTypeProvider.java b/python/python-psi-impl/src/com/jetbrains/python/codeInsight/typing/PyTypingTypeProvider.java
index 5602b5b5c5b2..12133a838d77 100644
--- a/python/python-psi-impl/src/com/jetbrains/python/codeInsight/typing/PyTypingTypeProvider.java
+++ b/python/python-psi-impl/src/com/jetbrains/python/codeInsight/typing/PyTypingTypeProvider.java
@@ -1559,6 +1559,12 @@ public final class PyTypingTypeProvider extends PyTypeProviderWithCustomContext<
return null;
}
+ @ApiStatus.Internal
+ public static @Nullable PyTypeParameterType getTypeParameterTypeFromTypeParameter(@NotNull PyTypeParameter typeParameter,
+ @NotNull TypeEvalContext context) {
+ return staticWithCustomContext(context, c -> getTypeParameterTypeFromTypeParameter(typeParameter, c));
+ }
+
private static @Nullable PyTypeParameterType getTypeParameterTypeFromTypeParameter(@NotNull PsiElement element, @NotNull Context context) {
if (element instanceof PyTypeParameter typeParameter) {
String name = typeParameter.getName();
diff --git a/python/python-psi-impl/src/com/jetbrains/python/inspections/PyTypeHintsInspection.kt b/python/python-psi-impl/src/com/jetbrains/python/inspections/PyTypeHintsInspection.kt
index 645a5956b760..5b3ce5c1eb91 100644
--- a/python/python-psi-impl/src/com/jetbrains/python/inspections/PyTypeHintsInspection.kt
+++ b/python/python-psi-impl/src/com/jetbrains/python/inspections/PyTypeHintsInspection.kt
@@ -127,7 +127,11 @@ class PyTypeHintsInspection : PyInspection() {
val defaultExpression = typeParameter.defaultExpression
if (defaultExpression == null) return
when(typeParameter.kind) {
- PyAstTypeParameter.Kind.TypeVar -> checkTypeVarDefaultType(defaultExpression)
+ PyAstTypeParameter.Kind.TypeVar -> {
+ val typeVarType = PyTypingTypeProvider.getTypeParameterTypeFromTypeParameter(typeParameter, myTypeEvalContext) as? PyTypeVarType
+ ?: return
+ checkTypeVarDefaultType(defaultExpression, typeVarType)
+ }
PyAstTypeParameter.Kind.ParamSpec -> checkParamSpecDefaultValue(defaultExpression)
PyAstTypeParameter.Kind.TypeVarTuple -> checkTypeVarTupleDefaultValue(defaultExpression, typeParameter)
}
@@ -470,9 +474,12 @@ class PyTypeHintsInspection : PyInspection() {
ProblemHighlightType.GENERIC_ERROR)
}
- default?.let { checkTypeVarDefaultType(it) }
-
- // TODO match bounds and constraints
+ default?.let {
+ val type = Ref.deref(PyTypingTypeProvider.getType(call, myTypeEvalContext))
+ if (type is PyTypeVarType) {
+ checkTypeVarDefaultType(it, type)
+ }
+ }
constraints.asSequence().plus(bound).forEach {
if (it != null) {
@@ -490,19 +497,18 @@ class PyTypeHintsInspection : PyInspection() {
}
}
- private fun checkTypeVarDefaultType(defaultExpression: PyExpression) {
- val type = Ref.deref(PyTypingTypeProvider.getType(defaultExpression, myTypeEvalContext))
- when (type) {
- is PyParamSpecType -> registerProblem(defaultExpression, PyPsiBundle.message("INSP.type.hints.cannot.be.used.in.default.type.of.type.var", "ParamSpec"))
- is PyTypeVarTupleType -> registerProblem(defaultExpression, PyPsiBundle.message("INSP.type.hints.cannot.be.used.in.default.type.of.type.var", "TypeVarTuple"))
+ private fun checkTypeVarDefaultType(defaultExpression: PyExpression, typeVarType: PyTypeVarType) {
+ val typeRef = typeVarType.defaultType
+ if (typeRef == null) {
+ registerProblem(defaultExpression, PyPsiBundle.message("INSP.type.hints.default.type.must.be.type.expression"))
+ return
}
- checkIsCorrectTypeExpression(defaultExpression)
- }
-
- private fun checkIsCorrectTypeExpression(expression: PyExpression) {
- if (PyTypingTypeProvider.getType(expression, myTypeEvalContext) == null) {
- registerProblem(expression, PyPsiBundle.message("INSP.type.hints.default.type.must.be.type.expression"))
+ val defaultType = typeRef.get()
+ when (defaultType) {
+ is PyParamSpecType -> registerProblem(defaultExpression, PyPsiBundle.message("INSP.type.hints.cannot.be.used.in.default.type.of.type.var", "ParamSpec"))
+ is PyTypeVarTupleType -> registerProblem(defaultExpression, PyPsiBundle.message("INSP.type.hints.cannot.be.used.in.default.type.of.type.var", "TypeVarTuple"))
+ else -> validateTypeVarDefaultType(typeVarType, defaultType, defaultExpression)
}
}
@@ -520,7 +526,9 @@ class PyTypeHintsInspection : PyInspection() {
if (defaultExpression is PyEllipsisLiteralExpression) return
if (defaultExpression is PyListLiteralExpression) {
defaultExpression.elements.forEach {
- checkIsCorrectTypeExpression(it)
+ if (PyTypingTypeProvider.getType(it, myTypeEvalContext) == null) {
+ registerProblem(it, PyPsiBundle.message("INSP.type.hints.default.type.must.be.type.expression"))
+ }
}
return
}
@@ -1368,6 +1376,29 @@ class PyTypeHintsInspection : PyInspection() {
.mapNotNull { it.qualifiedName }
.any { names.contains(it) }
}
+
+ private fun validateTypeVarDefaultType(typeVarType: PyTypeVarType, defaultType: PyType?, defaultExpression: PyExpression) {
+ val defaultTypes = when (defaultType) {
+ is PyTypeVarType -> defaultType.constraints.ifEmpty {
+ val objectType = PyBuiltinCache.getInstance(defaultExpression).objectType ?: return
+ listOf(defaultType.bound ?: objectType)
+ }
+ else -> listOf(defaultType)
+ }
+
+ when {
+ typeVarType.bound != null -> {
+ if (!defaultTypes.all { PyTypeChecker.match (typeVarType.bound, it, myTypeEvalContext) }) {
+ registerProblem(defaultExpression, PyPsiBundle.message("INSP.type.hints.default.type.do.not.match.bounds"))
+ }
+ }
+ typeVarType.constraints.isNotEmpty() -> {
+ if (!typeVarType.constraints.containsAll(defaultTypes)) {
+ registerProblem(defaultExpression, PyPsiBundle.message("INSP.type.hints.default.type.do.not.match.constraints"))
+ }
+ }
+ }
+ }
}
companion object {
diff --git a/python/testSrc/com/jetbrains/python/inspections/PyTypeHintsInspectionTest.java b/python/testSrc/com/jetbrains/python/inspections/PyTypeHintsInspectionTest.java
index b8c7d5b585ac..e8e3145f06e8 100644
--- a/python/testSrc/com/jetbrains/python/inspections/PyTypeHintsInspectionTest.java
+++ b/python/testSrc/com/jetbrains/python/inspections/PyTypeHintsInspectionTest.java
@@ -2648,6 +2648,169 @@ public class PyTypeHintsInspectionTest extends PyInspectionTestCase {
""");
}
+ // PY-76870
+ public void testTypeVarDefaultCanBeSubclassOfBound() {
+ doTestByText("""
+ from typing import TypeVar, List
+
+ T1 = TypeVar('T1', bound=int, default=bool)
+ """);
+ }
+
+ // PY-76870
+ public void testTypeVarDefaultCanNotBeSubclassOfConstraint() {
+ doTestByText("""
+ from typing import TypeVar, List
+
+ T1 = TypeVar('T1', int, str, default=bool)
+ """);
+ }
+
+ // PY-76870
+ public void testTypeVarDefaultTypeMatchesConstraints() {
+ doTestByText("""
+ from typing import TypeVar, List
+
+ # Default type matches one of the constraints
+ T1 = TypeVar('T1', str, int, default=str)
+ T2 = TypeVar('T2', str, int, default=int)
+
+ # Default type doesn't match any of the constraints
+ T3 = TypeVar('T3', str, int, default=bool)
+ T4 = TypeVar('T4', str, int, default=List[int])
+ """);
+ }
+
+ // PY-76870
+ public void testTypeVarDefaultTypeReferringToTypeVarMatchesBound() {
+ doTestByText("""
+ from typing import TypeVar
+
+ Y1 = TypeVar("Y1", bound=int)
+ Invalid = TypeVar("Invalid", float, str, default=Y1)
+ """);
+ }
+
+ // PY-76870
+ public void testTypeVarDefaultTypeReferringToTypeVarMatchesConstraints() {
+ doTestByText("""
+ from typing import TypeVar
+
+ Y1 = TypeVar("Y1", int, str)
+ AlsoOk2 = TypeVar("AlsoOk2", int, str, bool, default=Y1) # OK
+ AlsoInvalid2 = TypeVar("AlsoInvalid2", bool, complex, default=Y1)
+ """);
+ }
+
+ // PY-76870
+ public void testTypeVarDefaultTypeReferringToTypeVarWithoutConstraints() {
+ doTestByText("""
+ from typing import TypeVar
+ T = TypeVar("T")
+ Invalid = TypeVar("Invalid", str, int, default=T)
+ """);
+ }
+
+ // PY-76870
+ public void testTypeVarDefaultBoundMatchedAgainstBound() {
+ doTestByText("""
+ from typing import TypeVar
+
+ X1 = TypeVar("X1", bound=int)
+ Ok1 = TypeVar("Ok1", default=X1, bound=float)
+ """);
+ }
+
+ // PY-76870
+ public void testTypeVarDefaultBoundNotMatchedAgainstConstraints() {
+ doTestByText("""
+ from typing import TypeVar
+
+ Y3 = TypeVar("Y3", bound=int)
+ Invalid3 = TypeVar("Invalid3", str, complex, default=Y3)
+ """);
+ }
+
+ // PY-76870
+ public void testTypeVarDefaultBoundNotMatchedAgainstBound() {
+ doTestByText("""
+ from typing import TypeVar
+
+ X1 = TypeVar("X1", bound=int)
+ Invalid1 = TypeVar("Invalid1", default=X1, bound=str)
+ """);
+ }
+
+ // PY-76870
+ public void testTypeVarDefaultConstraintsNotMatchedAgainstBound() {
+ doTestByText("""
+ from typing import TypeVar
+
+ Y4 = TypeVar("Y4", int, str)
+ Invalid4 = TypeVar("Invalid4", bound=str, default=Y4)
+ """);
+ }
+
+ // PY-76870
+ public void testTypeVarDefaultChecksNewSyntax() {
+ doTestByText("""
+ # NOT OK
+ def foo1[T1: int = str](): ...
+ def foo2[T1: (int, bool) = str](): ...
+ def foo3[T1: int, T2: str = T1](): ...
+ def foo4[T1: (int, bool), T2: str = T1](): ...
+ def foo5[T1: (int, bool), T2: (int, str) = T1](): ...
+ def foo6[T1: (int, str, float), T2: (int, float) = T1](): ...
+ def foo7[T1: (int, str) = bool](): ...
+
+ # OK
+ def bar1[T1: int = bool](): ...
+ def bar2[T1: (int, str) = str](): ...
+ def bar3[T1: bool, T2: int = T1](): ...
+ def bar4[T1: (int, str), T2: (str, int) = T1](): ...
+ def bar5[T1: (int, str), T2: (str, int, float) = T1](): ...
+ """);
+ }
+
+ // PY-76870
+ public void testTypeVarDefaultTypeMatchedWithObject() {
+ doTestByText("""
+ from typing import TypeVar, Any
+
+ T = TypeVar('T')
+ T1 = TypeVar('T1', bound=object, default=T)
+ T2 = TypeVar('T2', int, object, default=T)
+ """);
+ }
+
+ // PY-76870
+ public void testTypeVarDefaultAnyInConstraintsAndBound() {
+ doTestByText("""
+ from typing import TypeVar, Any
+ T1 = TypeVar("T1", int, str)
+ Ok1 = TypeVar("Ok1", int, str, Any, default=T1)
+ T2 = TypeVar("T2", int, str, Any)
+ Ok2 = TypeVar("Ok2", int, str, Any, default=T2)
+ T3 = TypeVar("T3", bound=Any)
+ Ok3 = TypeVar("Ok3", bound=Any, default=T3)
+ T4 = TypeVar("T4", bound=str)
+ Ok4 = TypeVar("Ok4", bound=Any, default=T4)
+ T5 = TypeVar("T5", bound=Any)
+ Ok5 = TypeVar("Ok5", bound=Any, default=T5)
+
+ Y1 = TypeVar("Y1", int, str, Any)
+ NotOk1 = TypeVar("NotOk1", int, str, default=Y1)
+ Y2 = TypeVar("Y2", bound=Any)
+ NotOk2 = TypeVar("NotOk2", int, str, default=Y2)
+ Y3 = TypeVar("Y3", bound=str)
+ NotOk3 = TypeVar("NotOk3", int, Any, default=Y3)
+ Y4 = TypeVar("Y4", str, Any)
+ NotOk4 = TypeVar("NotOk4", int, Any, default=Y4)
+ Y5 = TypeVar("Y5", bound=Any)
+ NotOk5 = TypeVar("NotOk5", int, str, Any, default=Y5)
+ """);
+ }
+
@NotNull
@Override
protected Class extends PyInspection> getInspectionClass() {