PY-59241 Infer PyParamSpecType with proper scope for P.args/P.kwargs-annotated parameters

Previously, we inferred tuple[Any, ...] for such parameters and then replaced it
with a PyParamSpecType in PyTypeCheckerInspection specifically for its checks.
However, inside the inspection we didn't have enough context information to
properly detect the right type parameter scope.

Now, PyParamSpecType is created for vararg parameters annotated with "P.args" and
"P.kwargs" directly in PyTypingTypeProvider, uniformly with other possible usages
of ParamSpec and with accurate scope checks.

This also changes how we handle unrecognized parameter type hints. Previously,
we returned null for them, delegating to other type providers, which could lead
to subtle bugs. Now, when a type hint is present, but can't be identified,
we return Any, thus ignoring any further sources of type information.
It was triggered in PyProtectedMemberInspectionTest.testMemberResolvedToStub
in particular, where we ignored an annotation containing a broken
forward reference to a containing class, proceeding to a nearby .pyi stub.
This was most likely an honest mistake, so I've updated the test data.

GitOrigin-RevId: 02945e756c5f9a8097360c3bcf3f1c5f267c02e4
This commit is contained in:
Mikhail Golubev
2023-03-06 23:32:11 +00:00
committed by intellij-monorepo-bot
parent c25b2fec9c
commit 4d9e251568
5 changed files with 118 additions and 106 deletions
@@ -234,34 +234,67 @@ public class PyTypingTypeProvider extends PyTypeProviderWithCustomContext<PyTypi
@Override
@Nullable
public Ref<PyType> getParameterType(@NotNull PyNamedParameter param, @NotNull PyFunction func, @NotNull Context context) {
final Ref<PyType> typeFromAnnotation = getParameterTypeFromAnnotation(param, context);
if (typeFromAnnotation != null) {
return typeFromAnnotation;
@Nullable PyExpression typeHint = getAnnotationValue(param, context.myContext);
if (typeHint == null) {
String paramTypeCommentHint = param.getTypeCommentAnnotation();
if (paramTypeCommentHint != null) {
typeHint = toExpression(paramTypeCommentHint, param);
}
}
final Ref<PyType> typeFromTypeComment = getParameterTypeFromTypeComment(param, context);
if (typeFromTypeComment != null) {
return typeFromTypeComment;
}
final PyFunctionTypeAnnotation annotation = getFunctionTypeAnnotation(func);
if (annotation != null) {
PyParameterTypeList list = annotation.getParameterTypeList();
List<PyExpression> paramTypes = list.getParameterTypes();
if (paramTypes.size() == 1) {
final PyNoneLiteralExpression noneExpr = as(paramTypes.get(0), PyNoneLiteralExpression.class);
if (noneExpr != null && noneExpr.isEllipsis()) {
if (typeHint == null) {
PyFunctionTypeAnnotation annotation = getFunctionTypeAnnotation(func);
if (annotation != null) {
PyExpression funcTypeCommentParamHint = findParamTypeHintInFunctionTypeComment(annotation, param, func);
if (funcTypeCommentParamHint == null) {
return Ref.create();
}
}
final int startOffset = omitFirstParamInTypeComment(func, annotation) ? 1 : 0;
final List<PyParameter> funcParams = Arrays.asList(func.getParameterList().getParameters());
final int i = funcParams.indexOf(param) - startOffset;
if (i >= 0 && i < paramTypes.size()) {
return getParameterTypeFromFunctionComment(paramTypes.get(i), context);
typeHint = funcTypeCommentParamHint;
}
}
if (typeHint == null) {
return null;
}
if (typeHint instanceof PyReferenceExpression ref && ref.isQualified() && (
param.isPositionalContainer() && "args".equals(ref.getReferencedName()) ||
param.isKeywordContainer() && "kwargs".equals(ref.getReferencedName())
)) {
typeHint = Objects.requireNonNull(ref.getQualifier());
}
PyType type = Ref.deref(getType(typeHint, context));
if (param.isPositionalContainer() && !(type instanceof PyParamSpecType)) {
return Ref.create(PyTypeUtil.toPositionalContainerType(param, type));
}
if (param.isKeywordContainer() && !(type instanceof PyParamSpecType)) {
return Ref.create(PyTypeUtil.toKeywordContainerType(param, type));
}
if (PyNames.NONE.equals(param.getDefaultValueText())) {
return Ref.create(PyUnionType.union(type, PyNoneType.INSTANCE));
}
return Ref.create(type);
}
private static @Nullable PyExpression findParamTypeHintInFunctionTypeComment(@NotNull PyFunctionTypeAnnotation annotation,
@NotNull PyNamedParameter param,
@NotNull PyFunction func) {
List<PyExpression> paramTypes = annotation.getParameterTypeList().getParameterTypes();
if (paramTypes.size() == 1 && paramTypes.get(0) instanceof PyNoneLiteralExpression noneExpr && noneExpr.isEllipsis()) {
return null;
}
int startOffset = omitFirstParamInTypeComment(func, annotation) ? 1 : 0;
List<PyParameter> funcParams = Arrays.asList(func.getParameterList().getParameters());
int i = funcParams.indexOf(param) - startOffset;
if (i >= 0 && i < paramTypes.size()) {
PyExpression paramTypeHint = paramTypes.get(i);
if (paramTypeHint instanceof PyStarExpression starExpression) {
return starExpression.getExpression();
}
else if (paramTypeHint instanceof PyDoubleStarExpression doubleStarExpression) {
return doubleStarExpression.getExpression();
}
else {
return paramTypeHint;
}
}
return null;
}
@@ -269,73 +302,6 @@ public class PyTypingTypeProvider extends PyTypeProviderWithCustomContext<PyTypi
return type instanceof PyCollectionType && GENERATOR.equals(((PyClassLikeType)type).getClassQName());
}
@Nullable
private static Ref<PyType> getParameterTypeFromFunctionComment(@NotNull PyExpression expression, @NotNull Context context) {
final PyStarExpression starExpr = as(expression, PyStarExpression.class);
if (starExpr != null) {
final PyExpression inner = starExpr.getExpression();
if (inner != null) {
return Ref.create(PyTypeUtil.toPositionalContainerType(expression, Ref.deref(getType(inner, context))));
}
}
final PyDoubleStarExpression doubleStarExpr = as(expression, PyDoubleStarExpression.class);
if (doubleStarExpr != null) {
final PyExpression inner = doubleStarExpr.getExpression();
if (inner != null) {
return Ref.create(PyTypeUtil.toKeywordContainerType(expression, Ref.deref(getType(inner, context))));
}
}
return getType(expression, context);
}
@Nullable
private static Ref<PyType> getParameterTypeFromTypeComment(@NotNull PyNamedParameter parameter, @NotNull Context context) {
final String typeComment = parameter.getTypeCommentAnnotation();
if (typeComment != null) {
final PyType type = Ref.deref(getStringBasedType(typeComment, parameter, context));
if (parameter.isPositionalContainer()) {
return Ref.create(PyTypeUtil.toPositionalContainerType(parameter, type));
}
if (parameter.isKeywordContainer()) {
return Ref.create(PyTypeUtil.toKeywordContainerType(parameter, type));
}
return Ref.create(type);
}
return null;
}
@Nullable
private static Ref<PyType> getParameterTypeFromAnnotation(@NotNull PyNamedParameter parameter, @NotNull Context context) {
final Ref<PyType> annotationValueTypeRef = Optional
.ofNullable(getAnnotationValue(parameter, context.myContext))
.map(text -> getType(text, context))
.orElse(null);
if (annotationValueTypeRef != null) {
final PyType annotationValueType = annotationValueTypeRef.get();
if (parameter.isPositionalContainer()) {
return Ref.create(PyTypeUtil.toPositionalContainerType(parameter, annotationValueType));
}
if (parameter.isKeywordContainer()) {
return Ref.create(PyTypeUtil.toKeywordContainerType(parameter, annotationValueType));
}
if (PyNames.NONE.equals(parameter.getDefaultValueText())) {
return Ref.create(PyUnionType.union(annotationValueType, PyNoneType.INSTANCE));
}
return Ref.create(annotationValueType);
}
return null;
}
@NotNull
private static PyType createTypingGenericType(@NotNull PsiElement anchor) {
return new PyCustomType(GENERIC, null, false, true, PyBuiltinCache.getInstance(anchor).getObjectType());
@@ -345,9 +345,7 @@ public class PyTypeCheckerInspection extends PyInspection {
List<PyExpression> keywordArguments = getArgumentsMappedToKeywordContainer(mappedParameters);
List<PyExpression> allArguments = ContainerUtil.concat(positionalArguments, keywordArguments);
PyParamSpecType paramSpecType = getParamSpecTypeFromContainerParameters(keywordContainer,
positionalContainer,
as(callableType.getCallable(), PyFunction.class));
PyParamSpecType paramSpecType = getParamSpecTypeFromContainerParameters(keywordContainer, positionalContainer);
if (paramSpecType != null) {
analyzeParamSpec(paramSpecType, allArguments, substitutions, result, unmatchedArguments, unmatchedParameters);
}
@@ -431,21 +429,11 @@ public class PyTypeCheckerInspection extends PyInspection {
}
@Nullable
private static PyParamSpecType getParamSpecTypeFromContainerParameters(@Nullable PyCallableParameter positionalContainer,
@Nullable PyCallableParameter keywordContainer,
@Nullable PyFunction function) {
private PyParamSpecType getParamSpecTypeFromContainerParameters(@Nullable PyCallableParameter positionalContainer,
@Nullable PyCallableParameter keywordContainer) {
if (positionalContainer == null && keywordContainer == null) return null;
PyCallableParameter container = Objects.requireNonNullElse(positionalContainer, keywordContainer);
PyParameter parameter = container.getParameter();
if (!(parameter instanceof PyNamedParameter)) return null;
String annotationValue = ((PyNamedParameter)parameter).getAnnotationValue();
if (annotationValue == null ||
!(annotationValue.endsWith(".args") || annotationValue.endsWith(".kwargs")) ||
annotationValue.split("\\.").length != 2) return null;
String containerName = StringUtil.substringBeforeLast(annotationValue, ".");
return new PyParamSpecType(containerName).withScopeOwner(function);
return as(container.getType(myTypeEvalContext), PyParamSpecType.class);
}
private boolean matchParameterAndArgument(@Nullable PyType parameterType,
@@ -3,7 +3,7 @@ class Test:
def foo1(cls):
return cls._bar()
def foo2(self, t: Test):
def foo2(self, t):
return t._bar()
@staticmethod
@@ -0,0 +1,6 @@
from typing import Callable, ParamSpec
P = ParamSpec('P')
def func(c: Callable[P, int], *args: P.args, **kwargs: P.kwargs) -> None:
...
@@ -1200,6 +1200,58 @@ public class Py3TypeTest extends PyTestCase {
""");
}
public void testParamSpecArgsKwargsInAnnotations() {
doTest("(c: (ParamSpec(\"P\")) -> int, ParamSpec(\"P\"), ParamSpec(\"P\")) -> None", """
from typing import Callable, ParamSpec
P = ParamSpec('P')
def func(c: Callable[P, int], *args: P.args, **kwargs: P.kwargs) -> None:
...
expr = func
""");
}
public void testParamSpecArgsKwargsInTypeComments() {
doTest("(c: (ParamSpec(\"P\")) -> int, ParamSpec(\"P\"), ParamSpec(\"P\")) -> None", """
from typing import Callable, ParamSpec
P = ParamSpec('P')
def func(c, # type: Callable[P, int]
*args, # type: P.args
**kwargs, # type: P.kwargs
):
# type: (...) -> None
...
expr = func
""");
}
public void testParamSpecArgsKwargsInFunctionTypeComment() {
doTest("(c: (ParamSpec(\"P\")) -> int, ParamSpec(\"P\"), ParamSpec(\"P\")) -> None", """
from typing import Callable, ParamSpec
P = ParamSpec('P')
def func(c, *args, **kwargs):
# type: (Callable[P, int], *P.args, **P.kwargs) -> None
...
expr = func
""");
}
public void testParamSpecArgsKwargsInImportedFile() {
doMultiFileTest("(c: (ParamSpec(\"P\")) -> int, ParamSpec(\"P\"), ParamSpec(\"P\")) -> None", """
from mod import func
expr = func
""");
}
// PY-49935
public void testParamSpecSeveral() {
doTest("(y: int, x: str) -> bool",