diff --git a/python/src/com/jetbrains/python/codeInsight/stdlib/PyStdlibTypeProvider.java b/python/src/com/jetbrains/python/codeInsight/stdlib/PyStdlibTypeProvider.java index de5c258191fd..aed41c170012 100644 --- a/python/src/com/jetbrains/python/codeInsight/stdlib/PyStdlibTypeProvider.java +++ b/python/src/com/jetbrains/python/codeInsight/stdlib/PyStdlibTypeProvider.java @@ -148,7 +148,7 @@ public class PyStdlibTypeProvider extends PyTypeProviderBase { final String qname = function.getQualifiedName(); if (qname != null) { if (OPEN_FUNCTIONS.contains(qname) && callSite instanceof PyCallExpression) { - return getOpenFunctionType(qname, PyCallExpressionHelper.mapArguments(callSite, function, context), callSite); + return getOpenFunctionType(qname, PyCallExpressionHelper.mapArguments(callSite, function, context).getMappedParameters(), callSite); } else if ("tuple.__init__".equals(qname) && callSite instanceof PyCallExpression) { return getTupleInitializationType((PyCallExpression)callSite, context); diff --git a/python/src/com/jetbrains/python/inspections/PyTypeCheckerInspection.java b/python/src/com/jetbrains/python/inspections/PyTypeCheckerInspection.java index d4f3a6e664d7..bb4bdf98d238 100644 --- a/python/src/com/jetbrains/python/inspections/PyTypeCheckerInspection.java +++ b/python/src/com/jetbrains/python/inspections/PyTypeCheckerInspection.java @@ -15,17 +15,12 @@ */ package com.jetbrains.python.inspections; -import com.google.common.collect.Sets; import com.intellij.codeInspection.LocalInspectionToolSession; -import com.intellij.codeInspection.ProblemHighlightType; import com.intellij.codeInspection.ProblemsHolder; import com.intellij.openapi.diagnostic.Logger; import com.intellij.openapi.util.Key; -import com.intellij.openapi.util.Pair; -import com.intellij.openapi.util.text.StringUtil; import com.intellij.psi.PsiElementVisitor; import com.intellij.util.containers.ContainerUtil; -import com.intellij.util.containers.hash.LinkedHashMap; import com.jetbrains.python.PyNames; import com.jetbrains.python.codeInsight.controlflow.ScopeOwner; import com.jetbrains.python.codeInsight.dataflow.scope.ScopeUtil; @@ -33,12 +28,15 @@ import com.jetbrains.python.codeInsight.typing.PyTypingTypeProvider; import com.jetbrains.python.documentation.PythonDocumentationProvider; import com.jetbrains.python.inspections.quickfix.PyMakeFunctionReturnTypeQuickFix; import com.jetbrains.python.psi.*; +import com.jetbrains.python.psi.impl.PyCallExpressionHelper; import com.jetbrains.python.psi.types.*; +import one.util.streamex.StreamEx; import org.jetbrains.annotations.Nls; import org.jetbrains.annotations.NotNull; import org.jetbrains.annotations.Nullable; import java.util.*; +import java.util.stream.Collectors; import static com.jetbrains.python.psi.PyUtil.as; @@ -182,18 +180,20 @@ public class PyTypeCheckerInspection extends PyInspection { } private void checkCallSite(@NotNull PyCallSiteExpression callSite) { - final List resultsSet = PyTypeChecker.analyzeCallSite(callSite, myTypeEvalContext); - final List>> problemsSet = - new ArrayList<>(); - for (PyTypeChecker.AnalyzeCallResults results : resultsSet) { - problemsSet.add(checkMapping(results.getReceiver(), results.getArguments())); + final List calleesResults = StreamEx + .of(PyTypeChecker.analyzeCallSite(callSite, myTypeEvalContext)) + .filter(Visitor::callDoesNotHaveUnmappedArgumentsAndUnfilledParameters) + .map(this::analyzeCallee) + .toList(); + + if (matchedCalleeResultsExist(calleesResults)) return; + + if (calleesResults.size() == 1) { + PyTypeCheckerInspectionProblemRegistrar.registerSingleCalleeProblem(this, calleesResults.get(0), myTypeEvalContext); } - if (!problemsSet.isEmpty()) { - final Map> minProblems = Collections.min(problemsSet, - Comparator.comparingInt(Map::size)); - for (Map.Entry> entry : minProblems.entrySet()) { - registerProblem(entry.getKey(), entry.getValue().getFirst(), entry.getValue().getSecond()); - } + else if (!calleesResults.isEmpty()) { + PyTypeCheckerInspectionProblemRegistrar + .registerMultiCalleeProblem(this, callSite, getArgumentTypes(calleesResults), calleesResults, myTypeEvalContext); } } @@ -210,86 +210,84 @@ public class PyTypeCheckerInspection extends PyInspection { } } + private static boolean callDoesNotHaveUnmappedArgumentsAndUnfilledParameters(@NotNull PyTypeChecker.AnalyzeCallResults callResults) { + final PyCallExpressionHelper.ArgumentMappingResults mapping = callResults.getMapping(); + return mapping.getUnmappedArguments().isEmpty() && mapping.getUnmappedParameters().isEmpty(); + } + @NotNull - private Map> checkMapping(@Nullable PyExpression receiver, - @NotNull Map mapping) { - final Map> problems = - new HashMap<>(); - final Map substitutions = new LinkedHashMap<>(); - boolean genericsCollected = false; - for (Map.Entry entry : mapping.entrySet()) { - final PyNamedParameter param = entry.getValue(); - final PyExpression arg = entry.getKey(); - final PyType expectedArgType = param.getArgumentType(myTypeEvalContext); - if (expectedArgType == null) { - continue; - } - final PyType actualArgType = myTypeEvalContext.getType(arg); - if (!genericsCollected) { - substitutions.putAll(PyTypeChecker.unifyReceiver(receiver, myTypeEvalContext)); - genericsCollected = true; - } - final Pair problem = checkTypes(expectedArgType, actualArgType, myTypeEvalContext, substitutions); - if (problem != null) { - problems.put(arg, problem); - } + private AnalyzeCalleeResults analyzeCallee(@NotNull PyTypeChecker.AnalyzeCallResults results) { + final List result = new ArrayList<>(); + final PyExpression receiver = results.getReceiver(); + + Map substitutions = null; + + for (Map.Entry entry : results.getMapping().getMappedParameters().entrySet()) { + final AnalyzeArgumentResult argumentResult = + analyzeArgument(receiver, entry.getValue(), entry.getKey(), substitutions); + + substitutions = argumentResult.mySubstitutions; + + result.add(argumentResult); } - return problems; + + return new AnalyzeCalleeResults(results.getCallable(), result); } - @Nullable - private static Pair checkTypes(@Nullable PyType expected, - @Nullable PyType actual, - @NotNull TypeEvalContext context, - @NotNull Map substitutions) { - if (actual != null && expected != null) { - if (!PyTypeChecker.match(expected, actual, context, substitutions)) { - final String expectedName = PythonDocumentationProvider.getTypeName(expected, context); - String quotedExpectedName = String.format("'%s'", expectedName); - final boolean hasGenerics = PyTypeChecker.hasGenerics(expected, context); - ProblemHighlightType highlightType = ProblemHighlightType.GENERIC_ERROR_OR_WARNING; - if (hasGenerics) { - final PyType substitute = PyTypeChecker.substitute(expected, substitutions, context); - if (substitute != null) { - quotedExpectedName = String.format("'%s' (matched generic type '%s')", - PythonDocumentationProvider.getTypeName(substitute, context), - expectedName); - highlightType = ProblemHighlightType.WEAK_WARNING; - } - } - final String actualName = PythonDocumentationProvider.getTypeName(actual, context); - String msg = String.format("Expected type %s, got '%s' instead", quotedExpectedName, actualName); - if (expected instanceof PyStructuralType) { - final Set expectedAttributes = ((PyStructuralType)expected).getAttributeNames(); - final Set actualAttributes = getAttributes(actual, context); - if (actualAttributes != null) { - final Sets.SetView missingAttributes = Sets.difference(expectedAttributes, actualAttributes); - if (missingAttributes.size() == 1) { - msg = String.format("Type '%s' doesn't have expected attribute '%s'", actualName, missingAttributes.iterator().next()); - } - else { - msg = String.format("Type '%s' doesn't have expected attributes %s", - actualName, - StringUtil.join(missingAttributes, s -> String.format("'%s'", s), ", ")); - } - } - } - return Pair.create(msg, highlightType); - } - } - return null; + private static boolean matchedCalleeResultsExist(@NotNull List calleesResults) { + return calleesResults + .stream() + .anyMatch(calleeResults -> calleeResults.getResults().stream().allMatch(AnalyzeArgumentResult::isMatched)); } - } - @Nullable - private static Set getAttributes(@NotNull PyType type, @NotNull TypeEvalContext context) { - if (type instanceof PyStructuralType) { - return ((PyStructuralType)type).getAttributeNames(); + @NotNull + private static List getArgumentTypes(@NotNull List calleesResults) { + return calleesResults + .stream() + .map(AnalyzeCalleeResults::getResults) + .max(Comparator.comparingInt(List::size)) + .orElse(Collections.emptyList()) + .stream() + .map(AnalyzeArgumentResult::getActualType) + .collect(Collectors.toList()); } - else if (type instanceof PyClassLikeType) { - return ((PyClassLikeType)type).getMemberNames(true, context); + + /** + * @param receiver call receiver + * @param parameter callee parameter + * @param argument passed argument + * @param substitutions generics substitutions + * @return an object that contains expected argument type, expected argument type after substitution, actual argument type, + * flag with result of matching actual type against expected one and generics substitutions + * Note: generics substitutions are not recalculated if they were calculated before + */ + @NotNull + private AnalyzeArgumentResult analyzeArgument(@Nullable PyExpression receiver, + @NotNull PyNamedParameter parameter, + @NotNull PyExpression argument, + @Nullable Map substitutions) { + final PyType expectedArgumentType = parameter.getArgumentType(myTypeEvalContext); + final PyType actualArgumentType = myTypeEvalContext.getType(argument); + + if (expectedArgumentType == null) { + return new AnalyzeArgumentResult(argument, null, null, actualArgumentType, true, substitutions); + } + else { + substitutions = substitutions != null ? substitutions : PyTypeChecker.unifyReceiver(receiver, myTypeEvalContext); + + final PyType expectedTypeAfterSubstitution = PyTypeChecker.hasGenerics(expectedArgumentType, myTypeEvalContext) + ? PyTypeChecker.substitute(expectedArgumentType, substitutions, myTypeEvalContext) + : null; + return new AnalyzeArgumentResult( + argument, + expectedArgumentType, + expectedTypeAfterSubstitution, + actualArgumentType, + actualArgumentType == null || PyTypeChecker.match(expectedArgumentType, actualArgumentType, myTypeEvalContext, substitutions), + substitutions + ); + } } - return null; } @Override @@ -309,4 +307,87 @@ public class PyTypeCheckerInspection extends PyInspection { public String getDisplayName() { return "Type checker"; } + + static class AnalyzeCalleeResults { + + @NotNull + private final PyCallable myCallable; + + @NotNull + private final List myResults; + + public AnalyzeCalleeResults(@NotNull PyCallable callable, + @NotNull List results) { + myCallable = callable; + myResults = results; + } + + @NotNull + public PyCallable getCallable() { + return myCallable; + } + + @NotNull + public List getResults() { + return myResults; + } + } + + static class AnalyzeArgumentResult { + + @NotNull + private final PyExpression myArgument; + + @Nullable + private final PyType myExpectedType; + + @Nullable + private final PyType myExpectedTypeAfterSubstitution; + + @Nullable + private final PyType myActualType; + + private final boolean myIsMatched; + + @Nullable + private final Map mySubstitutions; + + public AnalyzeArgumentResult(@NotNull PyExpression argument, + @Nullable PyType expectedType, + @Nullable PyType expectedTypeAfterSubstitution, + @Nullable PyType actualType, + boolean isMatched, + @Nullable Map substitutions) { + myArgument = argument; + myExpectedType = expectedType; + myExpectedTypeAfterSubstitution = expectedTypeAfterSubstitution; + myActualType = actualType; + myIsMatched = isMatched; + mySubstitutions = substitutions; + } + + @NotNull + public PyExpression getArgument() { + return myArgument; + } + + @Nullable + public PyType getExpectedType() { + return myExpectedType; + } + + @Nullable + public PyType getExpectedTypeAfterSubstitution() { + return myExpectedTypeAfterSubstitution; + } + + @Nullable + public PyType getActualType() { + return myActualType; + } + + public boolean isMatched() { + return myIsMatched; + } + } } diff --git a/python/src/com/jetbrains/python/inspections/PyTypeCheckerInspectionProblemRegistrar.java b/python/src/com/jetbrains/python/inspections/PyTypeCheckerInspectionProblemRegistrar.java new file mode 100644 index 000000000000..036b124314f6 --- /dev/null +++ b/python/src/com/jetbrains/python/inspections/PyTypeCheckerInspectionProblemRegistrar.java @@ -0,0 +1,195 @@ +/* + * Copyright 2000-2017 JetBrains s.r.o. + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package com.jetbrains.python.inspections; + +import com.google.common.collect.Sets; +import com.intellij.codeInspection.ProblemHighlightType; +import com.intellij.openapi.util.text.StringUtil; +import com.intellij.psi.PsiElement; +import com.intellij.util.ObjectUtils; +import com.intellij.xml.util.XmlStringUtil; +import com.jetbrains.python.documentation.PythonDocumentationProvider; +import com.jetbrains.python.psi.PyBinaryExpression; +import com.jetbrains.python.psi.PyCallExpression; +import com.jetbrains.python.psi.PyCallSiteExpression; +import com.jetbrains.python.psi.PySubscriptionExpression; +import com.jetbrains.python.psi.types.PyClassLikeType; +import com.jetbrains.python.psi.types.PyStructuralType; +import com.jetbrains.python.psi.types.PyType; +import com.jetbrains.python.psi.types.TypeEvalContext; +import org.jetbrains.annotations.NotNull; +import org.jetbrains.annotations.Nullable; + +import java.util.List; +import java.util.Set; +import java.util.stream.Collectors; + +class PyTypeCheckerInspectionProblemRegistrar { + + static void registerSingleCalleeProblem(@NotNull PyInspectionVisitor visitor, + @NotNull PyTypeCheckerInspection.AnalyzeCalleeResults calleeResults, + @NotNull TypeEvalContext context) { + for (PyTypeCheckerInspection.AnalyzeArgumentResult argumentResult : calleeResults.getResults()) { + if (argumentResult.isMatched()) continue; + + visitor.registerProblem(argumentResult.getArgument(), + getSingleCalleeProblemMessage(argumentResult, context), + getSingleCalleeHighlightType(argumentResult.getExpectedTypeAfterSubstitution())); + } + } + + static void registerMultiCalleeProblem(@NotNull PyInspectionVisitor visitor, + @NotNull PyCallSiteExpression callSite, + @NotNull List argumentTypes, + @NotNull List calleesResults, + @NotNull TypeEvalContext context) { + visitor.registerProblem(getMultiCalleeElementToHighlight(callSite), + getMultiCalleeProblemMessage(argumentTypes, calleesResults, context), + getMultiCalleeHighlightType(calleesResults)); + } + + @NotNull + private static String getSingleCalleeProblemMessage(@NotNull PyTypeCheckerInspection.AnalyzeArgumentResult argumentResult, + @NotNull TypeEvalContext context) { + final PyType actualType = argumentResult.getActualType(); + final PyType expectedType = argumentResult.getExpectedType(); + + assert actualType != null; // see PyTypeCheckerInspection.Visitor.analyzeArgument() + assert expectedType != null; // see PyTypeCheckerInspection.Visitor.analyzeArgument() + + final String actualTypeName = PythonDocumentationProvider.getTypeName(actualType, context); + + if (expectedType instanceof PyStructuralType) { + final Set expectedAttributes = ((PyStructuralType)expectedType).getAttributeNames(); + final Set actualAttributes = getAttributes(actualType, context); + + if (actualAttributes != null) { + final Sets.SetView missingAttributes = Sets.difference(expectedAttributes, actualAttributes); + if (missingAttributes.size() == 1) { + return String.format("Type '%s' doesn't have expected attribute '%s'", actualTypeName, missingAttributes.iterator().next()); + } + else { + return String.format("Type '%s' doesn't have expected attributes %s", + actualTypeName, + StringUtil.join(missingAttributes, s -> String.format("'%s'", s), ", ")); + } + } + } + + final String expectedTypeRepresentation = getSingleCalleeExpectedTypeRepresentation(expectedType, + argumentResult.getExpectedTypeAfterSubstitution(), + context); + + return String.format("Expected type %s, got '%s' instead", expectedTypeRepresentation, actualTypeName); + } + + @NotNull + private static ProblemHighlightType getSingleCalleeHighlightType(@Nullable PyType expectedTypeAfterSubstitution) { + return expectedTypeAfterSubstitution == null ? ProblemHighlightType.GENERIC_ERROR_OR_WARNING : ProblemHighlightType.WEAK_WARNING; + } + + @NotNull + private static PsiElement getMultiCalleeElementToHighlight(@NotNull PyCallSiteExpression callSite) { + if (callSite instanceof PyCallExpression) { + return ObjectUtils.notNull(((PyCallExpression)callSite).getArgumentList(), callSite); + } + else if (callSite instanceof PyBinaryExpression) { + return ObjectUtils.notNull(((PyBinaryExpression)callSite).getPsiOperator(), callSite); + } + else if (callSite instanceof PySubscriptionExpression) { + return ObjectUtils.notNull(((PySubscriptionExpression)callSite).getIndexExpression(), callSite); + } + else { + return callSite; + } + } + + @NotNull + private static String getMultiCalleeProblemMessage(@NotNull List argumentTypes, + @NotNull List calleesResults, + @NotNull TypeEvalContext context) { + return XmlStringUtil.wrapInHtml("Unexpected type(s):
" + + XmlStringUtil.escapeString(getMultiCalleeActualTypesRepresentation(argumentTypes, context)) + "
" + + "Possible types:
" + + XmlStringUtil.escapeString(getMultiCalleePossibleExpectedTypesRepresentation(calleesResults, context))); + } + + /** + * @param calleesResults results of analyzing arguments passed to callees + * @return {@link ProblemHighlightType#WEAK_WARNING} if all expected types were substituted for all callees, + * {@link ProblemHighlightType#GENERIC_ERROR_OR_WARNING} otherwise. + */ + @NotNull + private static ProblemHighlightType getMultiCalleeHighlightType(@NotNull List calleesResults) { + final boolean allExpectedTypesWereSubstituted = calleesResults + .stream() + .flatMap(calleeResults -> calleeResults.getResults().stream()) + .allMatch(argumentResult -> argumentResult.getExpectedTypeAfterSubstitution() != null); + + return allExpectedTypesWereSubstituted ? ProblemHighlightType.WEAK_WARNING : ProblemHighlightType.GENERIC_ERROR_OR_WARNING; + } + + @Nullable + private static Set getAttributes(@NotNull PyType type, @NotNull TypeEvalContext context) { + if (type instanceof PyStructuralType) { + return ((PyStructuralType)type).getAttributeNames(); + } + else if (type instanceof PyClassLikeType) { + return ((PyClassLikeType)type).getMemberNames(true, context); + } + return null; + } + + @NotNull + private static String getSingleCalleeExpectedTypeRepresentation(@NotNull PyType expectedType, + @Nullable PyType expectedTypeAfterSubstitution, + @NotNull TypeEvalContext context) { + final String expectedTypeName = PythonDocumentationProvider.getTypeName(expectedType, context); + + return expectedTypeAfterSubstitution == null + ? String.format("'%s'", expectedTypeName) + : String.format("'%s' (matched generic type '%s')", + PythonDocumentationProvider.getTypeName(expectedTypeAfterSubstitution, context), + expectedTypeName); + } + + @NotNull + private static String getMultiCalleeActualTypesRepresentation(@NotNull List argumentTypes, @NotNull TypeEvalContext context) { + return argumentTypes + .stream() + .map(type -> PythonDocumentationProvider.getTypeName(type, context)) + .collect(Collectors.joining(", ", "(", ")")); + } + + @NotNull + private static String getMultiCalleePossibleExpectedTypesRepresentation(@NotNull List calleesResults, + @NotNull TypeEvalContext context) { + return calleesResults + .stream() + .map(calleeResult -> getMultiCalleeExpectedTypesRepresentation(calleeResult.getResults(), context)) + .collect(Collectors.joining("
")); + } + + @NotNull + private static String getMultiCalleeExpectedTypesRepresentation(@NotNull List calleeResults, + @NotNull TypeEvalContext context) { + return calleeResults + .stream() + .map(argumentResult -> ObjectUtils.chooseNotNull(argumentResult.getExpectedTypeAfterSubstitution(), argumentResult.getExpectedType())) + .map(type -> PythonDocumentationProvider.getTypeName(type, context)) + .collect(Collectors.joining(", ", "(", ")")); + } +} diff --git a/python/src/com/jetbrains/python/psi/impl/PyBinaryExpressionImpl.java b/python/src/com/jetbrains/python/psi/impl/PyBinaryExpressionImpl.java index 0c9cb953681b..4976a1e98fa3 100644 --- a/python/src/com/jetbrains/python/psi/impl/PyBinaryExpressionImpl.java +++ b/python/src/com/jetbrains/python/psi/impl/PyBinaryExpressionImpl.java @@ -1,5 +1,5 @@ /* - * Copyright 2000-2016 JetBrains s.r.o. + * Copyright 2000-2017 JetBrains s.r.o. * * Licensed under the Apache License, Version 2.0 (the "License"); * you may not use this file except in compliance with the License. @@ -143,7 +143,7 @@ public class PyBinaryExpressionImpl extends PyElementImpl implements PyBinaryExp final List matchedTypes = new ArrayList<>(); for (PyTypeChecker.AnalyzeCallResults result : results) { boolean matched = true; - for (Map.Entry entry : result.getArguments().entrySet()) { + for (Map.Entry entry : result.getMapping().getMappedParameters().entrySet()) { final PyExpression argument = entry.getKey(); final PyNamedParameter parameter = entry.getValue(); if (parameter.isPositionalContainer() || parameter.isKeywordContainer()) { diff --git a/python/src/com/jetbrains/python/psi/impl/PyCallExpressionHelper.java b/python/src/com/jetbrains/python/psi/impl/PyCallExpressionHelper.java index 19f903b800f0..3696428e74b4 100644 --- a/python/src/com/jetbrains/python/psi/impl/PyCallExpressionHelper.java +++ b/python/src/com/jetbrains/python/psi/impl/PyCallExpressionHelper.java @@ -645,20 +645,20 @@ public class PyCallExpressionHelper { } @NotNull - public static Map mapArguments(@NotNull PyCallSiteExpression callSite, - @NotNull PyCallable callable, - @NotNull List parameters, - @NotNull TypeEvalContext context) { + public static ArgumentMappingResults mapArguments(@NotNull PyCallSiteExpression callSite, + @NotNull PyCallable callable, + @NotNull List parameters, + @NotNull TypeEvalContext context) { final List arguments = PyTypeChecker.getArguments(callSite, callable); final PyResolveContext resolveContext = PyResolveContext.noImplicits().withTypeEvalContext(context); final List explicitParameters = PyTypeChecker.filterExplicitParameters(parameters, callable, callSite, resolveContext); - return analyzeArguments(arguments, explicitParameters).getMappedParameters(); + return analyzeArguments(arguments, explicitParameters); } @NotNull - public static Map mapArguments(@NotNull PyCallSiteExpression callSite, - @NotNull PyCallable callable, - @NotNull TypeEvalContext context) { + public static ArgumentMappingResults mapArguments(@NotNull PyCallSiteExpression callSite, + @NotNull PyCallable callable, + @NotNull TypeEvalContext context) { final List parameters = PyUtil.getParameters(callable, context); return mapArguments(callSite, callable, parameters, context); } diff --git a/python/src/com/jetbrains/python/psi/impl/PyFunctionImpl.java b/python/src/com/jetbrains/python/psi/impl/PyFunctionImpl.java index ad770d04d2e0..f88f88c6ab5b 100644 --- a/python/src/com/jetbrains/python/psi/impl/PyFunctionImpl.java +++ b/python/src/com/jetbrains/python/psi/impl/PyFunctionImpl.java @@ -1,5 +1,5 @@ /* - * Copyright 2000-2016 JetBrains s.r.o. + * Copyright 2000-2017 JetBrains s.r.o. * * Licensed under the Apache License, Version 2.0 (the "License"); * you may not use this file except in compliance with the License. @@ -224,7 +224,9 @@ public class PyFunctionImpl extends PyBaseElementImpl implements } final PyExpression receiver = PyTypeChecker.getReceiver(callSite, this); - final Map mapping = PyCallExpressionHelper.mapArguments(callSite, this, context); + final Map mapping = + PyCallExpressionHelper.mapArguments(callSite, this, context).getMappedParameters(); + return getCallType(receiver, mapping, context); } diff --git a/python/src/com/jetbrains/python/psi/types/PyTypeChecker.java b/python/src/com/jetbrains/python/psi/types/PyTypeChecker.java index ed4c4329f79a..88ba5d387bca 100644 --- a/python/src/com/jetbrains/python/psi/types/PyTypeChecker.java +++ b/python/src/com/jetbrains/python/psi/types/PyTypeChecker.java @@ -619,7 +619,7 @@ public class PyTypeChecker { final List results = new ArrayList<>(); for (PyCallable callable : multiResolveCallee(callSite, context)) { final PyExpression receiver = getReceiver(callSite, callable); - final Map mapping = PyCallExpressionHelper.mapArguments(callSite, callable, context); + final PyCallExpressionHelper.ArgumentMappingResults mapping = PyCallExpressionHelper.mapArguments(callSite, callable, context); results.add(new AnalyzeCallResults(callable, receiver, mapping)); } return results; @@ -817,13 +817,13 @@ public class PyTypeChecker { public static class AnalyzeCallResults { @NotNull private final PyCallable myCallable; @Nullable private final PyExpression myReceiver; - @NotNull private final Map myArguments; + @NotNull private final PyCallExpressionHelper.ArgumentMappingResults myMapping; public AnalyzeCallResults(@NotNull PyCallable callable, @Nullable PyExpression receiver, - @NotNull Map arguments) { + @NotNull PyCallExpressionHelper.ArgumentMappingResults mapping) { myCallable = callable; myReceiver = receiver; - myArguments = arguments; + myMapping = mapping; } @NotNull @@ -837,8 +837,8 @@ public class PyTypeChecker { } @NotNull - public Map getArguments() { - return myArguments; + public PyCallExpressionHelper.ArgumentMappingResults getMapping() { + return myMapping; } } } diff --git a/python/src/com/jetbrains/python/pyi/PyiTypeProvider.java b/python/src/com/jetbrains/python/pyi/PyiTypeProvider.java index cce91f1f540f..03207e585b5e 100644 --- a/python/src/com/jetbrains/python/pyi/PyiTypeProvider.java +++ b/python/src/com/jetbrains/python/pyi/PyiTypeProvider.java @@ -1,5 +1,5 @@ /* - * Copyright 2000-2016 JetBrains s.r.o. + * Copyright 2000-2017 JetBrains s.r.o. * * Licensed under the Apache License, Version 2.0 (the "License"); * you may not use this file except in compliance with the License. @@ -124,11 +124,11 @@ public class PyiTypeProvider extends PyTypeProviderBase { } final PyExpression receiver = PyTypeChecker.getReceiver(callSite, overload); - final Map mapping = mapArguments(callSite, overload, context); + final PyCallExpressionHelper.ArgumentMappingResults mapping = mapArguments(callSite, overload, context); if (mapping == null) { continue; } - final Map substitutions = PyTypeChecker.unifyGenericCall(receiver, mapping, context); + final Map substitutions = PyTypeChecker.unifyGenericCall(receiver, mapping.getMappedParameters(), context); final PyType unifiedType = substitutions != null ? PyTypeChecker.substitute(returnType, substitutions, context) : null; if (unifiedType != null) { @@ -210,15 +210,17 @@ public class PyiTypeProvider extends PyTypeProviderBase { } @Nullable - private static Map mapArguments(@NotNull PyCallSiteExpression callSite, - @NotNull PyFunction function, - @NotNull TypeEvalContext context) { + private static PyCallExpressionHelper.ArgumentMappingResults mapArguments(@NotNull PyCallSiteExpression callSite, + @NotNull PyFunction function, + @NotNull TypeEvalContext context) { final List parameters = Arrays.asList(function.getParameterList().getParameters()); - final Map map = PyCallExpressionHelper.mapArguments(callSite, function, parameters, context); + final PyCallExpressionHelper.ArgumentMappingResults mapping = + PyCallExpressionHelper.mapArguments(callSite, function, parameters, context); + final PyCallExpression callExpr = as(callSite, PyCallExpression.class); - if (callExpr != null && callExpr.getArguments().length != map.size()) { + if (callExpr != null && callExpr.getArguments().length != mapping.getMappedParameters().size()) { return null; } - return map; + return mapping; } } diff --git a/python/testData/inspections/PyTypeCheckerInspection/BuiltinNumeric.py b/python/testData/inspections/PyTypeCheckerInspection/BuiltinNumeric.py index 3646e9542072..6278c8f82bec 100644 --- a/python/testData/inspections/PyTypeCheckerInspection/BuiltinNumeric.py +++ b/python/testData/inspections/PyTypeCheckerInspection/BuiltinNumeric.py @@ -5,6 +5,6 @@ def test(): float(False) complex(False) divmod(False, False) - divmod('foo', u'bar') + divmod('foo', u'bar') pow(False, True) - round(False, 'foo') + round(False, 'foo') diff --git a/python/testData/inspections/PyTypeCheckerInspection/BuiltinsPy3.py b/python/testData/inspections/PyTypeCheckerInspection/BuiltinsPy3.py index 5b66044a2a13..b1b849fccd46 100644 --- a/python/testData/inspections/PyTypeCheckerInspection/BuiltinsPy3.py +++ b/python/testData/inspections/PyTypeCheckerInspection/BuiltinsPy3.py @@ -13,4 +13,4 @@ def test_numerics(): divmod(False, False) divmod(b'foo', 'bar') pow(False, True) - round(False, 'foo') + round(False, 'foo') diff --git a/python/testData/inspections/PyTypeCheckerInspection/SecondFormIter.py b/python/testData/inspections/PyTypeCheckerInspection/SecondFormIter.py index a331ff954278..96b10610f8a2 100644 --- a/python/testData/inspections/PyTypeCheckerInspection/SecondFormIter.py +++ b/python/testData/inspections/PyTypeCheckerInspection/SecondFormIter.py @@ -7,7 +7,7 @@ def test_second_form(): def test_second_form_fail(): - for chunk in iter(10, ''): + for chunk in iter(10, ''): pass diff --git a/python/testData/pyi/inspections/overloadedGenerics/OverloadedGenerics.py b/python/testData/pyi/inspections/overloadedGenerics/OverloadedGenerics.py index aa87709d8784..8f28a38c15ec 100644 --- a/python/testData/pyi/inspections/overloadedGenerics/OverloadedGenerics.py +++ b/python/testData/pyi/inspections/overloadedGenerics/OverloadedGenerics.py @@ -1,6 +1,6 @@ from m1 import g, Gen g(Gen(10).get(10, 10)) -g(Gen(10).get(10, 'foo')) -g(Gen('foo').get(10, 10)) +g(Gen(10).get(10, 'foo')) +g(Gen('foo').get(10, 10)) g(Gen('foo').get(10, 'foo')) diff --git a/python/testData/pyi/inspections/overloads/Overloads.py b/python/testData/pyi/inspections/overloads/Overloads.py index c4015235a3c3..9a6ba08983ee 100644 --- a/python/testData/pyi/inspections/overloads/Overloads.py +++ b/python/testData/pyi/inspections/overloads/Overloads.py @@ -4,7 +4,7 @@ from m1 import f, g, C, stub_only def test_overloaded_function(x): g(f(10)) g(f('foo')) - g(f({1: 2})) + g(f({1: 2})) g(f(x)) @@ -12,7 +12,7 @@ def test_overloaded_subscription_operator_parameters(): c = C() print(c[10]) print(c['foo']) - print(c[{1: 2}]) + print(c[{1: 2}]) def test_overloaded_binary_operator_parameters(): @@ -26,4 +26,4 @@ def test_stub_only_function(x): g(stub_only(10)) g(stub_only('foo')) g(stub_only(x)) - g(stub_only({1: 2})) + g(stub_only({1: 2})) diff --git a/python/testData/pyi/inspections/overloadsWithDifferentNumberOfParameters/OverloadsWithDifferentNumberOfParameters.py b/python/testData/pyi/inspections/overloadsWithDifferentNumberOfParameters/OverloadsWithDifferentNumberOfParameters.py new file mode 100644 index 000000000000..5df28a83977c --- /dev/null +++ b/python/testData/pyi/inspections/overloadsWithDifferentNumberOfParameters/OverloadsWithDifferentNumberOfParameters.py @@ -0,0 +1,37 @@ +from m1 import f, g, h + + +def test_different_number_of_parameters(): + f(5) + f("a") + + f(5, "a") + f("a", "b") + f(5, 6) + f("a", 5) + + +def test_same_number_of_parameters_but_one_is_default(): + g(5) + g("a") + + g("a", False) + g(5, 6) + g(False, 5) + + g(5, "a") + g("a", "b") + g(5, 6) + g("a", 5) + + +def test_different_number_of_parameters_one_is_default(): + h(5) + h(lambda x: x) + + h("a") + + h("a", False) + h(5, False) # fail + h("a", 5) # fail + h(False, "a") # fail \ No newline at end of file diff --git a/python/testData/pyi/inspections/overloadsWithDifferentNumberOfParameters/m1.pyi b/python/testData/pyi/inspections/overloadsWithDifferentNumberOfParameters/m1.pyi new file mode 100644 index 000000000000..13f8bf935349 --- /dev/null +++ b/python/testData/pyi/inspections/overloadsWithDifferentNumberOfParameters/m1.pyi @@ -0,0 +1,18 @@ +from typing import overload + +@overload +def f(i: int) -> None: ... +@overload +def f(i: int, s: str) -> None: ... + + +@overload +def g(i: int, b: bool = ...) -> None: ... +@overload +def g(i: int, s: str) -> None: ... + + +@overload +def h(i: int) -> None: ... +@overload +def h(i: str, b: bool = ...) -> None: ... \ No newline at end of file diff --git a/python/testSrc/com/jetbrains/python/pyi/PyiInspectionsTest.java b/python/testSrc/com/jetbrains/python/pyi/PyiInspectionsTest.java index 9f412d4f9fa3..2d05cd95dd63 100644 --- a/python/testSrc/com/jetbrains/python/pyi/PyiInspectionsTest.java +++ b/python/testSrc/com/jetbrains/python/pyi/PyiInspectionsTest.java @@ -1,5 +1,5 @@ /* - * Copyright 2000-2015 JetBrains s.r.o. + * Copyright 2000-2017 JetBrains s.r.o. * * Licensed under the Apache License, Version 2.0 (the "License"); * you may not use this file except in compliance with the License. @@ -63,6 +63,10 @@ public class PyiInspectionsTest extends PyTestCase { doPyTest(PyTypeCheckerInspection.class); } + public void testOverloadsWithDifferentNumberOfParameters() { + doPyTest(PyTypeCheckerInspection.class); + } + public void testOverloadedGenerics() { doPyTest(PyTypeCheckerInspection.class); }