From 4d9e2515684fc8201e1dcb387a0a08e5842d0fa3 Mon Sep 17 00:00:00 2001 From: Mikhail Golubev Date: Fri, 3 Mar 2023 14:44:32 +0200 Subject: [PATCH] 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 --- .../typing/PyTypingTypeProvider.java | 144 +++++++----------- .../inspections/PyTypeCheckerInspection.java | 20 +-- .../MemberResolvedToStub/a.py | 2 +- .../ParamSpecArgsKwargsInImportedFile/mod.py | 6 + .../com/jetbrains/python/Py3TypeTest.java | 52 +++++++ 5 files changed, 118 insertions(+), 106 deletions(-) create mode 100644 python/testData/types/ParamSpecArgsKwargsInImportedFile/mod.py 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 e693d91f8b12..2269f6f339ad 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 @@ -234,34 +234,67 @@ public class PyTypingTypeProvider extends PyTypeProviderWithCustomContext getParameterType(@NotNull PyNamedParameter param, @NotNull PyFunction func, @NotNull Context context) { - final Ref 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 typeFromTypeComment = getParameterTypeFromTypeComment(param, context); - if (typeFromTypeComment != null) { - return typeFromTypeComment; - } - - final PyFunctionTypeAnnotation annotation = getFunctionTypeAnnotation(func); - if (annotation != null) { - PyParameterTypeList list = annotation.getParameterTypeList(); - List 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 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 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 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 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 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 getParameterTypeFromAnnotation(@NotNull PyNamedParameter parameter, @NotNull Context context) { - final Ref 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()); diff --git a/python/python-psi-impl/src/com/jetbrains/python/inspections/PyTypeCheckerInspection.java b/python/python-psi-impl/src/com/jetbrains/python/inspections/PyTypeCheckerInspection.java index ebff87d313a2..331005d69d31 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/inspections/PyTypeCheckerInspection.java +++ b/python/python-psi-impl/src/com/jetbrains/python/inspections/PyTypeCheckerInspection.java @@ -345,9 +345,7 @@ public class PyTypeCheckerInspection extends PyInspection { List keywordArguments = getArgumentsMappedToKeywordContainer(mappedParameters); List 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, diff --git a/python/testData/inspections/PyProtectedMemberInspection/MemberResolvedToStub/a.py b/python/testData/inspections/PyProtectedMemberInspection/MemberResolvedToStub/a.py index 29080666fdb1..4218ee456c98 100644 --- a/python/testData/inspections/PyProtectedMemberInspection/MemberResolvedToStub/a.py +++ b/python/testData/inspections/PyProtectedMemberInspection/MemberResolvedToStub/a.py @@ -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 diff --git a/python/testData/types/ParamSpecArgsKwargsInImportedFile/mod.py b/python/testData/types/ParamSpecArgsKwargsInImportedFile/mod.py new file mode 100644 index 000000000000..e78b1d7adf1b --- /dev/null +++ b/python/testData/types/ParamSpecArgsKwargsInImportedFile/mod.py @@ -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: + ... diff --git a/python/testSrc/com/jetbrains/python/Py3TypeTest.java b/python/testSrc/com/jetbrains/python/Py3TypeTest.java index 8115c2771d6d..19c1da8be28e 100644 --- a/python/testSrc/com/jetbrains/python/Py3TypeTest.java +++ b/python/testSrc/com/jetbrains/python/Py3TypeTest.java @@ -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",