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",