diff --git a/python/src/com/jetbrains/python/codeInsight/stdlib/PyStdlibTypeProvider.java b/python/src/com/jetbrains/python/codeInsight/stdlib/PyStdlibTypeProvider.java index b4d75e321417..3c11a9623403 100644 --- a/python/src/com/jetbrains/python/codeInsight/stdlib/PyStdlibTypeProvider.java +++ b/python/src/com/jetbrains/python/codeInsight/stdlib/PyStdlibTypeProvider.java @@ -10,7 +10,6 @@ import com.jetbrains.python.psi.*; import com.jetbrains.python.psi.impl.PyBuiltinCache; import com.jetbrains.python.psi.impl.PyQualifiedName; import com.jetbrains.python.psi.impl.PyTypeProvider; -import com.jetbrains.python.psi.resolve.PyResolveContext; import com.jetbrains.python.psi.resolve.ResolveImportUtil; import com.jetbrains.python.psi.types.*; import org.jetbrains.annotations.NotNull; @@ -51,13 +50,13 @@ public class PyStdlibTypeProvider extends PyTypeProviderBase { @Nullable @Override - public PyType getReturnType(@NotNull PyFunction function, @Nullable PyReferenceExpression callSite, @NotNull TypeEvalContext context) { + public PyType getReturnType(@NotNull PyFunction function, @Nullable PyQualifiedExpression callSite, @NotNull TypeEvalContext context) { final String qname = getQualifiedName(function, callSite); if (qname != null) { if (callSite != null) { - final PsiElement parent = callSite.getParent(); - if (parent instanceof PyCallExpression) { - final PyType overloaded = getOverloadedReturnTypeByQName((PyCallExpression)parent, qname, function, context); + PyTypeChecker.AnalyzeCallResults results = PyTypeChecker.analyzeCallSite(callSite, context); + if (results != null) { + final PyType overloaded = getOverloadedReturnTypeByQName(results.getArguments(), qname, function, context); if (overloaded != null) { return overloaded; } @@ -123,7 +122,9 @@ public class PyStdlibTypeProvider extends PyTypeProviderBase { } @Nullable - private PyType getOverloadedReturnTypeByQName(@NotNull PyCallExpression expr, @NotNull String qname, @NotNull PsiElement anchor, + private PyType getOverloadedReturnTypeByQName(@NotNull Map arguments, + @NotNull String qname, + @NotNull PsiElement anchor, @NotNull TypeEvalContext context) { int i = 1; PyType rtype; @@ -131,40 +132,35 @@ public class PyStdlibTypeProvider extends PyTypeProviderBase { final String overloadedQName = String.format("%s.%d", qname, i); rtype = getReturnTypeByQName(overloadedQName, anchor); if (rtype != null) { - final PyArgumentList args = expr.getArgumentList(); - if (args != null) { - final CallArgumentsMapping res = args.analyzeCall(PyResolveContext.noImplicits().withTypeEvalContext(context)); - final Map mapped = res.getPlainMappedParams(); - boolean matched = true; - for (Map.Entry entry : mapped.entrySet()) { - final PyNamedParameter p = entry.getValue(); - final String name = p.getName(); - if (p.isPositionalContainer() || p.isKeywordContainer() || name == null) { - continue; - } - PyType argType = entry.getKey().getType(context); - // Special case for the 'mode' argument of the 'open()' builtin - if (("__builtin__.open".equals(qname) || "io.open".equals(qname)) && "mode".equals(name)) { - final PyBuiltinCache cache = PyBuiltinCache.getInstance(anchor); - final LanguageLevel level = LanguageLevel.forElement(anchor); - argType = cache.getUnicodeType(level); - final PyExpression modeExpr = entry.getKey(); - if (modeExpr instanceof PyStringLiteralExpression) { - final String literal = ((PyStringLiteralExpression)modeExpr).getStringValue(); - if (literal.contains("b")) { - argType = cache.getBytesType(level); - } + boolean matched = true; + for (Map.Entry entry : arguments.entrySet()) { + final PyNamedParameter p = entry.getValue(); + final String name = p.getName(); + if (p.isPositionalContainer() || p.isKeywordContainer() || name == null) { + continue; + } + PyType argType = entry.getKey().getType(context); + // Special case for the 'mode' argument of the 'open()' builtin + if (("__builtin__.open".equals(qname) || "io.open".equals(qname)) && "mode".equals(name)) { + final PyBuiltinCache cache = PyBuiltinCache.getInstance(anchor); + final LanguageLevel level = LanguageLevel.forElement(anchor); + argType = cache.getUnicodeType(level); + final PyExpression modeExpr = entry.getKey(); + if (modeExpr instanceof PyStringLiteralExpression) { + final String literal = ((PyStringLiteralExpression)modeExpr).getStringValue(); + if (literal.contains("b")) { + argType = cache.getBytesType(level); } } - final PyType paramType = getParameterTypeByQName(overloadedQName, name, anchor); - if (!PyTypeChecker.match(paramType, argType, context)) { - matched = false; - } } - if (matched) { - return rtype; + final PyType paramType = getParameterTypeByQName(overloadedQName, name, anchor); + if (!PyTypeChecker.match(paramType, argType, context)) { + matched = false; } } + if (matched) { + return rtype; + } } i++; } while (rtype != null); diff --git a/python/src/com/jetbrains/python/codeInsight/stdlib/StdlibTypes2.properties b/python/src/com/jetbrains/python/codeInsight/stdlib/StdlibTypes2.properties index 39db2b4e1e72..abc81aa49d05 100644 --- a/python/src/com/jetbrains/python/codeInsight/stdlib/StdlibTypes2.properties +++ b/python/src/com/jetbrains/python/codeInsight/stdlib/StdlibTypes2.properties @@ -41,8 +41,11 @@ __builtin__.enumerate.__init__ = \ __builtin__.enumerate.__new__ = \ :rtype: enumerate of (int, T) \n\ +__builtin__.enumerate.__iter__ = \ + :rtype: enumerate of (int, T) \n\ + __builtin__.enumerate.next = \ - :rtype: (int, T) \n\ + :rtype: (int, unknown) \n\ __builtin__.filter = \ :type function_or_none: collections.Callable or None \n\ @@ -66,7 +69,7 @@ __builtin__.getattr = \ :rtype: unknown \n\ __builtin__.globals = \ - :rtype: dict from bytes to object + :rtype: dict of (bytes, object) __builtin__.hasattr = \ :type name: string \n\ @@ -84,7 +87,7 @@ __builtin__.len = \ :rtype: int \n\ __builtin__.locals = \ - :rtype: dict from bytes to object + :rtype: dict of (bytes, object) __builtin__.map = \ :type function: collections.Callable \n\ @@ -92,7 +95,8 @@ __builtin__.map = \ :rtype: list \n\ __builtin__.next = \ - :type iterator: collections.Iterator \n\ + :type iterator: collections.Iterator of T \n\ + :rtype: T \n\ __builtin__.open = \ :type name: string \n\ @@ -137,7 +141,7 @@ __builtin__.round = \ :rtype: float \n\ __builtin__.vars = \ - :rtype: dict from bytes to object + :rtype: dict of (bytes, object) ## 5.4. Numeric types @@ -971,25 +975,29 @@ __builtin__.unicode.zfill = \ :type width: int or long \n\ :rtype: unicode \n\ +__builtin__.list.__init__ = \ + :type seq: collections.Iterable of T \n\ + :rtype: list of T \n\ + __builtin__.list.__add__ = \ - :type y: list \n\ - :rtype: list \n\ + :type y: list of T \n\ + :rtype: list of T \n\ __builtin__.list.__mul__ = \ :type n: int or long \n\ - :rtype: list \n\ + :rtype: list of T \n\ __builtin__.list.__rmul__ = \ :type n: int or long \n\ - :rtype: list \n\ + :rtype: list of T \n\ __builtin__.list.__getitem__ = \ :type y: int \n\ - :rtype: unknown \n\ + :rtype: T \n\ __builtin__.list.__setitem__ = \ :type i: int \n\ - :type y: object \n\ + :type y: T \n\ :rtype: None \n\ __builtin__.list.__delitem__ = \ @@ -997,30 +1005,32 @@ __builtin__.list.__delitem__ = \ :rtype: None \n\ __builtin__.list.append = \ - :type p_object: object \n\ + :type p_object: T \n\ + :rtype: None \n\ __builtin__.list.extend = \ - :type iterable: collections.Iterable \n\ + :type iterable: collections.Iterable of T \n\ + :rtype: None \n\ __builtin__.list.count = \ - :type value: object \n\ + :type value: T \n\ :rtype: int \n\ __builtin__.list.index = \ - :type value: object \n\ + :type value: T \n\ :type start: bool or int or long or None \n\ :type stop: bool or int or long or None \n\ :rtype: int \n\ __builtin__.list.insert = \ :type index: bool or int or long \n\ - :type p_object: object \n\ + :type p_object: T \n\ __builtin__.list.pop = \ :type index: bool or int or long \n\ __builtin__.list.remove = \ - :type value: object \n\ + :type value: T \n\ __builtin__.list.sort = \ :type reverse: bool \n\ @@ -1059,41 +1069,41 @@ __builtin__.dict.__init__ = \ :rtype: dict of (T, V) \n\ __builtin__.dict.__getitem__ = \ - :type y: object \n\ - :rtype: unknown \n\ + :type y: T \n\ + :rtype: V \n\ __builtin__.dict.__setitem__ = \ - :type i: object \n\ - :type y: object \n\ + :type i: T \n\ + :type y: V \n\ :rtype: None \n\ __builtin__.dict.__delitem__ = \ - :type y: object \n\ + :type y: T \n\ :rtype: None \n\ __builtin__.dict.copy = \ - :rtype: dict \n\ + :rtype: dict of (T, V) \n\ __builtin__.dict.fromkeys = \ - :type S: collections.Iterable \n\ - :type v: object or None \n\ - :rtype: dict \n\ + :type S: collections.Iterable of T \n\ + :type v: V or None \n\ + :rtype: dict of (T, V) \n\ __builtin__.dict.has_key = \ - :type k: object \n\ + :type k: T \n\ :rtype: bool \n\ __builtin__.dict.items = \ :rtype: list of (T, V) \n\ __builtin__.dict.iteritems = \ - :rtype: collections.Iterable of (unknown, unknown) \n\ + :rtype: collections.Iterable of (T, V) \n\ __builtin__.dict.iterkeys = \ - :rtype: collections.Iterable \n\ + :rtype: collections.Iterable of T \n\ __builtin__.dict.itervalues = \ - :rtype: collections.Iterable \n\ + :rtype: collections.Iterable of V \n\ __builtin__.dict.keys = \ :rtype: list of T \n\ @@ -1194,6 +1204,14 @@ datetime.date.__sub__ = \ :type y: datetime.date or datetime.timedelta \n\ :rtype: datetime.date or datetime.timedelta \n\ +datetime.date.__sub__.1 = \ + :type y: datetime.date \n\ + :rtype: datetime.timedelta \n\ + +datetime.date.__sub__.2 = \ + :type y: datetime.timedelta \n\ + :rtype: datetime.date \n\ + datetime.date.__rsub__ = \ :type y: datetime.date \n\ :rtype: datetime.timedelta \n\ @@ -1218,6 +1236,18 @@ datetime.timedelta.__add__ = \ :type y: datetime.timedelta or datetime.date or datetime.datetime \n\ :rtype: datetime.timedelta or datetime.date or datetime.datetime \n\ +datetime.timedelta.__add__.1 = \ + :type y: datetime.timedelta \n\ + :rtype: datetime.timedelta \n\ + +datetime.timedelta.__add__.2 = \ + :type y: datetime.datetime \n\ + :rtype: datetime.datetime \n\ + +datetime.timedelta.__add__.3 = \ + :type y: datetime.date \n\ + :rtype: datetime.date \n\ + datetime.timedelta.__radd__ = \ :type y: datetime.timedelta \n\ :rtype: datetime.timedelta \n\ @@ -1338,6 +1368,14 @@ datetime.datetime.__sub__ = \ :type y: datetime.datetime or datetime.timedelta \n\ :rtype: datetime.datetime or datetime.timedelta \n\ +datetime.datetime.__sub__.1 = \ + :type y: datetime.datetime \n\ + :rtype: datetime.timedelta \n\ + +datetime.datetime.__sub__.2 = \ + :type y: datetime.timedelta \n\ + :rtype: datetime.datetime \n\ + datetime.datetime.__rsub__ = \ :type y: datetime.datetime \n\ :rtype: datetime.timedelta \n\ diff --git a/python/src/com/jetbrains/python/inspections/PyTypeCheckerInspection.java b/python/src/com/jetbrains/python/inspections/PyTypeCheckerInspection.java index 0c4f1ded6b61..3b0f8fa8886a 100644 --- a/python/src/com/jetbrains/python/inspections/PyTypeCheckerInspection.java +++ b/python/src/com/jetbrains/python/inspections/PyTypeCheckerInspection.java @@ -5,11 +5,14 @@ import com.intellij.codeInspection.ProblemsHolder; import com.intellij.openapi.diagnostic.Logger; import com.intellij.openapi.util.Key; import com.intellij.openapi.util.Ref; -import com.intellij.psi.*; -import com.jetbrains.python.PyNames; +import com.intellij.psi.PsiElement; +import com.intellij.psi.PsiElementVisitor; import com.jetbrains.python.documentation.PythonDocumentationProvider; import com.jetbrains.python.psi.*; -import com.jetbrains.python.psi.types.*; +import com.jetbrains.python.psi.types.PyGenericType; +import com.jetbrains.python.psi.types.PyType; +import com.jetbrains.python.psi.types.PyTypeChecker; +import com.jetbrains.python.psi.types.TypeEvalContext; import org.jetbrains.annotations.Nls; import org.jetbrains.annotations.NotNull; import org.jetbrains.annotations.Nullable; @@ -42,17 +45,30 @@ public class PyTypeCheckerInspection extends PyInspection { public void visitPyCallExpression(PyCallExpression node) { final PyArgumentList args = node.getArgumentList(); if (args != null) { - final Map substitutions = new HashMap(); - final CallArgumentsMapping res = args.analyzeCall(resolveWithoutImplicits()); - final PyCallExpression.PyMarkedCallee markedCallee = res.getMarkedCallee(); - if (markedCallee != null) { - final Callable callable = markedCallee.getCallable(); - final PyFunction function = callable.asMethod(); - if (function != null) { - substitutions.putAll(PyTypeChecker.collectCallGenerics(function, node, myTypeEvalContext)); - } + final PyExpression callee = node.getCallee(); + if (callee instanceof PyQualifiedExpression) { + checkCallSite((PyQualifiedExpression)callee); } - for (Map.Entry entry : res.getPlainMappedParams().entrySet()) { + } + } + + @Override + public void visitPyBinaryExpression(PyBinaryExpression node) { + checkCallSite(node); + } + + @Override + public void visitPySubscriptionExpression(PySubscriptionExpression node) { + // TODO: Support slice PySliceExpressions + checkCallSite(node); + } + + private void checkCallSite(@Nullable PyQualifiedExpression callSite) { + final Map substitutions = new HashMap(); + final PyTypeChecker.AnalyzeCallResults results = PyTypeChecker.analyzeCallSite(callSite, myTypeEvalContext); + if (results != null) { + substitutions.putAll(PyTypeChecker.collectCallGenerics(results.getFunction(), results.getReceiver(), myTypeEvalContext)); + for (Map.Entry entry : results.getArguments().entrySet()) { final PyNamedParameter p = entry.getValue(); if (p.isPositionalContainer() || p.isKeywordContainer()) { // TODO: Support *args, **kwargs @@ -65,65 +81,6 @@ public class PyTypeCheckerInspection extends PyInspection { } } - @Override - public void visitPyBinaryExpression(PyBinaryExpression node) { - final PsiPolyVariantReference ref = node.getReference(resolveWithoutImplicits()); - if (ref != null) { - final ResolveResult[] results = ref.multiResolve(false); - String error = null; - PyExpression arg = null; - for (ResolveResult result : results) { - final PsiElement resolved = result.getElement(); - if (resolved instanceof PyFunction) { - final PyFunction fun = (PyFunction)resolved; - PyExpression expr = PyNames.isRightOperatorName(fun.getName()) ? node.getLeftExpression() : node.getRightExpression(); - String msg = checkSingleArgumentFunction(fun, expr, false); - if (msg == null) { - return; - } - if (error == null) { - error = msg; - arg = expr; - } - } - else { - return; - } - } - if (error != null) { - registerProblem(arg, error); - } - } - } - - @Override - public void visitPySubscriptionExpression(PySubscriptionExpression node) { - // TODO: Support slice PySliceExpressions - final PsiReference ref = node.getReference(resolveWithoutImplicits()); - if (ref != null) { - final PsiElement resolved = ref.resolve(); - if (resolved instanceof PyFunction) { - checkSingleArgumentFunction((PyFunction)resolved, node.getIndexExpression(), true); - } - } - } - - @Nullable - private String checkSingleArgumentFunction(@NotNull PyFunction fun, @Nullable PyExpression argument, boolean registerProblem) { - if (argument != null) { - final PyParameter[] parameters = fun.getParameterList().getParameters(); - if (parameters.length == 2) { - final PyNamedParameter p = parameters[1].getAsNamed(); - if (p != null) { - final PyType argType = argument.getType(myTypeEvalContext); - final PyType paramType = p.getType(myTypeEvalContext); - return checkTypes(paramType, argType, argument, myTypeEvalContext, new HashMap(), registerProblem); - } - } - } - return null; - } - @Nullable private String checkTypes(@Nullable PyType superType, @Nullable PyType subType, @Nullable PsiElement node, @NotNull TypeEvalContext context, @NotNull Map substitutions, boolean registerPoblem) { diff --git a/python/src/com/jetbrains/python/inspections/PyUnresolvedReferencesInspection.java b/python/src/com/jetbrains/python/inspections/PyUnresolvedReferencesInspection.java index 1044bcda2e61..8b7c8f12017c 100644 --- a/python/src/com/jetbrains/python/inspections/PyUnresolvedReferencesInspection.java +++ b/python/src/com/jetbrains/python/inspections/PyUnresolvedReferencesInspection.java @@ -380,7 +380,9 @@ public class PyUnresolvedReferencesInspection extends PyInspection { boolean marked_qualified = false; if (element instanceof PyQualifiedExpression) { final PyQualifiedExpression qexpr = (PyQualifiedExpression)element; - if (myIgnoredIdentifiers.contains(ref_text) || PyNames.COMPARISON_OPERATORS.contains(qexpr.getReferencedName())) { + if (myIgnoredIdentifiers.contains(ref_text) || + PyNames.COMPARISON_OPERATORS.contains(qexpr.getReferencedName()) || + refname == null) { return; } final PyExpression qualifier = qexpr.getQualifier(); @@ -469,7 +471,7 @@ public class PyUnresolvedReferencesInspection extends PyInspection { } return false; } - + private static void addCreateMemberFromUsageFixes(PyType qtype, PsiReference reference, String refText, List actions) { PsiElement element = reference.getElement(); if (qtype instanceof PyClassType) { diff --git a/python/src/com/jetbrains/python/psi/Callable.java b/python/src/com/jetbrains/python/psi/Callable.java index 2ce0b7243185..51267edd7619 100644 --- a/python/src/com/jetbrains/python/psi/Callable.java +++ b/python/src/com/jetbrains/python/psi/Callable.java @@ -23,7 +23,7 @@ public interface Callable extends PyElement { * @return the type of returned value. */ @Nullable - PyType getReturnType(TypeEvalContext typeEvalContext, @Nullable PyReferenceExpression callSite); + PyType getReturnType(@NotNull TypeEvalContext context, @Nullable PyQualifiedExpression callSite); /** * @return a methods returns itself, non-method callables return null. diff --git a/python/src/com/jetbrains/python/psi/PyBinaryExpression.java b/python/src/com/jetbrains/python/psi/PyBinaryExpression.java index ce6a94eb7ec1..1279844ddc8a 100644 --- a/python/src/com/jetbrains/python/psi/PyBinaryExpression.java +++ b/python/src/com/jetbrains/python/psi/PyBinaryExpression.java @@ -6,7 +6,7 @@ import org.jetbrains.annotations.Nullable; /** * @author yole */ -public interface PyBinaryExpression extends PyReferenceExpression { +public interface PyBinaryExpression extends PyQualifiedExpression, PyReferenceOwner { PyExpression getLeftExpression(); @Nullable PyExpression getRightExpression(); diff --git a/python/src/com/jetbrains/python/psi/impl/PyBinaryExpressionImpl.java b/python/src/com/jetbrains/python/psi/impl/PyBinaryExpressionImpl.java index 2b5c2e7bc856..33b82f078913 100644 --- a/python/src/com/jetbrains/python/psi/impl/PyBinaryExpressionImpl.java +++ b/python/src/com/jetbrains/python/psi/impl/PyBinaryExpressionImpl.java @@ -11,8 +11,6 @@ import com.jetbrains.python.PyElementTypes; import com.jetbrains.python.PyNames; import com.jetbrains.python.psi.*; import com.jetbrains.python.psi.resolve.PyResolveContext; -import com.jetbrains.python.psi.resolve.PyResolveUtil; -import com.jetbrains.python.psi.resolve.QualifiedResolveResult; import com.jetbrains.python.psi.types.PyNoneType; import com.jetbrains.python.psi.types.PyType; import com.jetbrains.python.psi.types.PyTypeChecker; @@ -98,17 +96,6 @@ public class PyBinaryExpressionImpl extends PyElementImpl implements PyBinaryExp } } - @NotNull - @Override - public QualifiedResolveResult followAssignmentsChain(PyResolveContext resolveContext) { - return new PyReferenceExpressionImpl.QualifiedResolveResultEmpty(); - } - - @Override - public PyQualifiedName asQualifiedName() { - return PyQualifiedName.fromReferenceChain(PyResolveUtil.unwindQualifiers(this)); - } - @NotNull @Override public PsiPolyVariantReference getReference() { diff --git a/python/src/com/jetbrains/python/psi/impl/PyFunctionImpl.java b/python/src/com/jetbrains/python/psi/impl/PyFunctionImpl.java index caa9f02e015c..99203b5ce420 100644 --- a/python/src/com/jetbrains/python/psi/impl/PyFunctionImpl.java +++ b/python/src/com/jetbrains/python/psi/impl/PyFunctionImpl.java @@ -6,9 +6,7 @@ import com.intellij.openapi.util.Pair; import com.intellij.openapi.util.Ref; import com.intellij.openapi.util.text.StringUtil; import com.intellij.openapi.vfs.VirtualFile; -import com.intellij.psi.PsiElement; -import com.intellij.psi.PsiFile; -import com.intellij.psi.ResolveState; +import com.intellij.psi.*; import com.intellij.psi.scope.PsiScopeProcessor; import com.intellij.psi.search.LocalSearchScope; import com.intellij.psi.search.SearchScope; @@ -130,18 +128,21 @@ public class PyFunctionImpl extends PyPresentableElementImpl imp } @Nullable - public PyType getReturnType(@NotNull TypeEvalContext context, @Nullable PyReferenceExpression callSite) { + @Override + public PyType getReturnType(@NotNull TypeEvalContext context, @Nullable PyQualifiedExpression callSite) { final PyType type = getGenericReturnType(context, callSite); + if (callSite == null) { + return type; + } if (PyTypeChecker.hasGenerics(type, context)) { - if (callSite != null) { - final PsiElement parent = callSite.getParent(); - if (parent instanceof PyCallExpression) { - final Map substitutions = PyTypeChecker.unifyGenericCall(this, (PyCallExpression)parent, context); - if (substitutions != null) { - final Ref result = PyTypeChecker.substitute(type, substitutions, context); - if (result != null) { - return result.get(); - } + final PyTypeChecker.AnalyzeCallResults results = PyTypeChecker.analyzeCallSite(callSite, context); + if (results != null) { + final Map substitutions = PyTypeChecker.unifyGenericCall(this, results.getReceiver(), results.getArguments(), + context); + if (substitutions != null) { + final Ref result = PyTypeChecker.substitute(type, substitutions, context); + if (result != null) { + return result.get(); } } } @@ -151,7 +152,7 @@ public class PyFunctionImpl extends PyPresentableElementImpl imp } @Nullable - private PyType getGenericReturnType(TypeEvalContext typeEvalContext, @Nullable PyReferenceExpression callSite) { + private PyType getGenericReturnType(TypeEvalContext typeEvalContext, @Nullable PyQualifiedExpression callSite) { if (typeEvalContext.maySwitchToAST(this)) { PyAnnotation anno = getAnnotation(); if (anno != null) { diff --git a/python/src/com/jetbrains/python/psi/impl/PyLambdaExpressionImpl.java b/python/src/com/jetbrains/python/psi/impl/PyLambdaExpressionImpl.java index c2fd8e0a7ca7..4c06fc713709 100644 --- a/python/src/com/jetbrains/python/psi/impl/PyLambdaExpressionImpl.java +++ b/python/src/com/jetbrains/python/psi/impl/PyLambdaExpressionImpl.java @@ -34,7 +34,9 @@ public class PyLambdaExpressionImpl extends PyElementImpl implements PyLambdaExp return childToPsiNotNull(PyElementTypes.PARAMETER_LIST_SET, 0); } - public PyType getReturnType(TypeEvalContext context, @Nullable PyReferenceExpression callSite) { + @Nullable + @Override + public PyType getReturnType(@NotNull TypeEvalContext context, @Nullable PyQualifiedExpression callSite) { final PyExpression body = getBody(); if (body != null) return body.getType(context); else return null; diff --git a/python/src/com/jetbrains/python/psi/impl/PyPrefixExpressionImpl.java b/python/src/com/jetbrains/python/psi/impl/PyPrefixExpressionImpl.java index 20b125cccd72..10a5ddf3817f 100644 --- a/python/src/com/jetbrains/python/psi/impl/PyPrefixExpressionImpl.java +++ b/python/src/com/jetbrains/python/psi/impl/PyPrefixExpressionImpl.java @@ -68,7 +68,7 @@ public class PyPrefixExpressionImpl extends PyElementImpl implements PyPrefixExp if (ref != null) { final PsiElement resolved = ref.resolve(); if (resolved instanceof Callable) { - return ((Callable)resolved).getReturnType(context, null); + return ((Callable)resolved).getReturnType(context, this); } } return null; diff --git a/python/src/com/jetbrains/python/psi/impl/PySubscriptionExpressionImpl.java b/python/src/com/jetbrains/python/psi/impl/PySubscriptionExpressionImpl.java index eb053d4007e1..0a35123ede1f 100644 --- a/python/src/com/jetbrains/python/psi/impl/PySubscriptionExpressionImpl.java +++ b/python/src/com/jetbrains/python/psi/impl/PySubscriptionExpressionImpl.java @@ -44,7 +44,7 @@ public class PySubscriptionExpressionImpl extends PyElementImpl implements PySub if (ref != null) { final PsiElement resolved = ref.resolve(); if (resolved instanceof Callable) { - res = ((Callable)resolved).getReturnType(context, null); + res = ((Callable)resolved).getReturnType(context, this); } } if (PyTypeChecker.isUnknown(res) || res instanceof PyNoneType) { diff --git a/python/src/com/jetbrains/python/psi/impl/PyTypeProvider.java b/python/src/com/jetbrains/python/psi/impl/PyTypeProvider.java index 7853cdace40f..722130b6f197 100644 --- a/python/src/com/jetbrains/python/psi/impl/PyTypeProvider.java +++ b/python/src/com/jetbrains/python/psi/impl/PyTypeProvider.java @@ -2,10 +2,7 @@ package com.jetbrains.python.psi.impl; import com.intellij.openapi.extensions.ExtensionPointName; import com.intellij.psi.PsiElement; -import com.jetbrains.python.psi.PyClass; -import com.jetbrains.python.psi.PyFunction; -import com.jetbrains.python.psi.PyNamedParameter; -import com.jetbrains.python.psi.PyReferenceExpression; +import com.jetbrains.python.psi.*; import com.jetbrains.python.psi.types.PyType; import com.jetbrains.python.psi.types.TypeEvalContext; import org.jetbrains.annotations.NotNull; @@ -19,7 +16,7 @@ public interface PyTypeProvider { @Nullable PyType getReferenceExpressionType(PyReferenceExpression referenceExpression, TypeEvalContext context); - + @Nullable PyType getReferenceType(@NotNull PsiElement referenceTarget, TypeEvalContext context, @Nullable PsiElement anchor); @@ -27,7 +24,7 @@ public interface PyTypeProvider { PyType getParameterType(PyNamedParameter param, final PyFunction func, TypeEvalContext context); @Nullable - PyType getReturnType(PyFunction function, @Nullable PyReferenceExpression callSite, TypeEvalContext context); + PyType getReturnType(PyFunction function, @Nullable PyQualifiedExpression callSite, TypeEvalContext context); @Nullable PyType getIterationType(PyClass iterable); diff --git a/python/src/com/jetbrains/python/psi/types/PyTypeChecker.java b/python/src/com/jetbrains/python/psi/types/PyTypeChecker.java index 8460f5db9df5..f4a6f4bc2ff8 100644 --- a/python/src/com/jetbrains/python/psi/types/PyTypeChecker.java +++ b/python/src/com/jetbrains/python/psi/types/PyTypeChecker.java @@ -2,6 +2,10 @@ package com.jetbrains.python.psi.types; import com.intellij.openapi.extensions.Extensions; import com.intellij.openapi.util.Ref; +import com.intellij.psi.PsiElement; +import com.intellij.psi.PsiPolyVariantReference; +import com.intellij.psi.PsiReference; +import com.intellij.psi.ResolveResult; import com.jetbrains.python.PyNames; import com.jetbrains.python.codeInsight.stdlib.PyStdlibTypeProvider; import com.jetbrains.python.psi.*; @@ -226,15 +230,12 @@ public class PyTypeChecker { } @Nullable - public static Map unifyGenericCall(@NotNull PyFunction function, @NotNull PyCallExpression call, + public static Map unifyGenericCall(@NotNull PyFunction function, + @Nullable PyExpression receiver, + @NotNull Map arguments, @NotNull TypeEvalContext context) { - final PyArgumentList args = call.getArgumentList(); - if (args == null) { - return null; - } - final Map substitutions = collectCallGenerics(function, call, context); - final CallArgumentsMapping res = args.analyzeCall(PyResolveContext.noImplicits().withTypeEvalContext(context)); - for (Map.Entry entry : res.getPlainMappedParams().entrySet()) { + final Map substitutions = collectCallGenerics(function, receiver, context); + for (Map.Entry entry : arguments.entrySet()) { final PyNamedParameter p = entry.getValue(); final String name = p.getName(); if (p.isPositionalContainer() || p.isKeywordContainer() || name == null) { @@ -250,40 +251,33 @@ public class PyTypeChecker { } @NotNull - public static Map collectCallGenerics(@NotNull PyFunction function, @NotNull PyCallExpression call, - @NotNull TypeEvalContext context) { + public static Map collectCallGenerics(@NotNull PyFunction function, @Nullable PyExpression receiver, + @NotNull TypeEvalContext context) { final Map substitutions = new HashMap(); - final PyExpression callee = call.getCallee(); - if (callee instanceof PyReferenceExpression) { - final PyReferenceExpression expr = (PyReferenceExpression)callee; - final PyExpression qualifier = expr.getQualifier(); - if (qualifier != null) { - final PyType qualType = qualifier.getType(context); - // Collect generic params of object type - final Set generics = new HashSet(); - collectGenerics(qualType, context, generics); - for (PyGenericType t : generics) { - substitutions.put(t, t); + // Collect generic params of object type + final Set generics = new HashSet(); + final PyType qualifierType = receiver != null ? receiver.getType(context) : null; + collectGenerics(qualifierType, context, generics); + for (PyGenericType t : generics) { + substitutions.put(t, t); + } + final PyClass cls = function.getContainingClass(); + if (cls != null) { + final PyFunction init = cls.findInitOrNew(true); + // Unify generics in constructor + if (init != null) { + final PyType initType = init.getReturnType(context, null); + if (initType != null) { + match(initType, qualifierType, context, substitutions); } - final PyClass cls = function.getContainingClass(); - if (cls != null) { - final PyFunction init = cls.findInitOrNew(true); - // Unify generics in constructor - if (init != null) { - final PyType initType = init.getReturnType(context, null); - if (initType != null) { - match(initType, qualType, context, substitutions); - } - } - else { - // Unify generics in stdlib pseudo-constructor - final PyStdlibTypeProvider stdlib = PyStdlibTypeProvider.getInstance(); - if (stdlib != null) { - final PyType initType = stdlib.getConstructorType(cls); - if (initType != null) { - match(initType, qualType, context, substitutions); - } - } + } + else { + // Unify generics in stdlib pseudo-constructor + final PyStdlibTypeProvider stdlib = PyStdlibTypeProvider.getInstance(); + if (stdlib != null) { + final PyType initType = stdlib.getConstructorType(cls); + if (initType != null) { + match(initType, qualifierType, context, substitutions); } } } @@ -300,4 +294,112 @@ public class PyTypeChecker { return superName != null && superName.equals(subClass.getName()); } } + + @Nullable + public static AnalyzeCallResults analyzeCallSite(@Nullable PyQualifiedExpression callSite, @NotNull TypeEvalContext context) { + if (callSite == null) { + return null; + } + final PsiElement parent = callSite.getParent(); + final PyResolveContext resolveContext = PyResolveContext.noImplicits().withTypeEvalContext(context); + if (parent instanceof PyCallExpression) { + final PyExpression receiver = callSite.getQualifier(); + final PyCallExpression call = (PyCallExpression)parent; + final PyArgumentList args = call.getArgumentList(); + if (args != null) { + final CallArgumentsMapping mapping = args.analyzeCall(resolveContext); + final Map arguments = mapping.getPlainMappedParams(); + final PyCallExpression.PyMarkedCallee markedCallee = mapping.getMarkedCallee(); + if (markedCallee != null) { + final Callable callable = markedCallee.getCallable(); + if (callable instanceof PyFunction) { + return new AnalyzeCallResults((PyFunction)callable, receiver, arguments); + } + } + } + } + else if (callSite instanceof PyBinaryExpression) { + final PyBinaryExpression expr = (PyBinaryExpression)callSite; + final PsiPolyVariantReference ref = expr.getReference(resolveContext); + if (ref != null) { + final ResolveResult[] resolveResult = ref.multiResolve(false); + AnalyzeCallResults firstResults = null; + for (ResolveResult result : resolveResult) { + final PsiElement resolved = result.getElement(); + if (resolved instanceof PyFunction) { + final PyFunction function = (PyFunction)resolved; + final boolean isRight = PyNames.isRightOperatorName(function.getName()); + final PyExpression arg = isRight ? expr.getLeftExpression() : expr.getRightExpression(); + final PyExpression receiver = isRight ? expr.getRightExpression() : expr.getLeftExpression(); + final PyParameter[] parameters = function.getParameterList().getParameters(); + if (parameters.length == 2) { + final PyNamedParameter param = parameters[1].getAsNamed(); + if (arg != null && param != null) { + final Map arguments = new HashMap(); + arguments.put(arg, param); + final AnalyzeCallResults resutls = new AnalyzeCallResults(function, receiver, arguments); + if (firstResults == null) { + firstResults = resutls; + } + if (match(param.getType(context), arg.getType(context), context)) { + return resutls; + } + } + } + } + } + if (firstResults != null) { + return firstResults; + } + } + } + else if (callSite instanceof PySubscriptionExpression) { + final PySubscriptionExpression expr = (PySubscriptionExpression)callSite; + final PsiReference ref = expr.getReference(resolveContext); + if (ref != null) { + final PsiElement resolved = ref.resolve(); + if (resolved instanceof PyFunction) { + final PyFunction function = (PyFunction)resolved; + final PyParameter[] parameters = function.getParameterList().getParameters(); + if (parameters.length == 2) { + final PyNamedParameter param = parameters[1].getAsNamed(); + if (param != null) { + final Map arguments = new HashMap(); + arguments.put(expr.getIndexExpression(), param); + return new AnalyzeCallResults(function, expr.getOperand(), arguments); + } + } + } + } + } + return null; + } + + public static class AnalyzeCallResults { + @NotNull private final PyFunction myFunction; + @Nullable private final PyExpression myReceiver; + @NotNull private final Map myArguments; + + public AnalyzeCallResults(@NotNull PyFunction function, @Nullable PyExpression receiver, + @NotNull Map arguments) { + myFunction = function; + myReceiver = receiver; + myArguments = arguments; + } + + @NotNull + public PyFunction getFunction() { + return myFunction; + } + + @Nullable + public PyExpression getReceiver() { + return myReceiver; + } + + @NotNull + public Map getArguments() { + return myArguments; + } + } } diff --git a/python/src/com/jetbrains/python/psi/types/PyTypeProviderBase.java b/python/src/com/jetbrains/python/psi/types/PyTypeProviderBase.java index c016ba06dc44..208bb6128941 100644 --- a/python/src/com/jetbrains/python/psi/types/PyTypeProviderBase.java +++ b/python/src/com/jetbrains/python/psi/types/PyTypeProviderBase.java @@ -20,7 +20,7 @@ public class PyTypeProviderBase implements PyTypeProvider { protected interface ReturnTypeCallback { @Nullable - PyType getType(@Nullable PyReferenceExpression callSite, @Nullable PyType qualifierType, TypeEvalContext context); + PyType getType(@Nullable PyQualifiedExpression callSite, @Nullable PyType qualifierType, TypeEvalContext context); } private static class ReturnTypeDescriptor { @@ -31,7 +31,7 @@ public class PyTypeProviderBase implements PyTypeProvider { } @Nullable - public PyType get(PyFunction function, @Nullable PyReferenceExpression callSite, TypeEvalContext context) { + public PyType get(PyFunction function, @Nullable PyQualifiedExpression callSite, TypeEvalContext context) { PyClass containingClass = function.getContainingClass(); if (containingClass != null) { final ReturnTypeCallback typeCallback = myStringToReturnTypeMap.get(containingClass.getQualifiedName()); @@ -47,7 +47,7 @@ public class PyTypeProviderBase implements PyTypeProvider { private final ReturnTypeCallback mySelfTypeCallback = new ReturnTypeCallback() { @Override - public PyType getType(@Nullable PyReferenceExpression callSite, @Nullable PyType qualifierType, TypeEvalContext context) { + public PyType getType(@Nullable PyQualifiedExpression callSite, @Nullable PyType qualifierType, TypeEvalContext context) { if (qualifierType instanceof PyClassType) { return new PyClassType(((PyClassType)qualifierType).getPyClass(), false); } @@ -79,7 +79,7 @@ public class PyTypeProviderBase implements PyTypeProvider { } @Override - public PyType getReturnType(PyFunction function, @Nullable PyReferenceExpression callSite, TypeEvalContext context) { + public PyType getReturnType(PyFunction function, @Nullable PyQualifiedExpression callSite, TypeEvalContext context) { ReturnTypeDescriptor descriptor; synchronized (myMethodToReturnTypeMap) { descriptor = myMethodToReturnTypeMap.get(function.getName()); diff --git a/python/testData/inspections/PyCallingNonCallableInspection/tupleNonCallable.py b/python/testData/inspections/PyCallingNonCallableInspection/tupleNonCallable.py index e08bc00761b3..a4e04009336a 100644 --- a/python/testData/inspections/PyCallingNonCallableInspection/tupleNonCallable.py +++ b/python/testData/inspections/PyCallingNonCallableInspection/tupleNonCallable.py @@ -1 +1 @@ -(1,2)() +(1,2)() diff --git a/python/testData/inspections/PyTypeCheckerInspection/test.py b/python/testData/inspections/PyTypeCheckerInspection/test.py index 97722e71b75b..401bd1a58a30 100644 --- a/python/testData/inspections/PyTypeCheckerInspection/test.py +++ b/python/testData/inspections/PyTypeCheckerInspection/test.py @@ -84,7 +84,7 @@ def f_list_tuple(spam, eggs): def test_list_tuple(): f_list_tuple([1, 2, 3], - (False, 2, '')) + (False, 2, '')) def test_builtin_numerics(): diff --git a/python/testSrc/com/jetbrains/python/PyTypeParserTest.java b/python/testSrc/com/jetbrains/python/PyTypeParserTest.java index 327634fb1c9d..328bcc15680e 100644 --- a/python/testSrc/com/jetbrains/python/PyTypeParserTest.java +++ b/python/testSrc/com/jetbrains/python/PyTypeParserTest.java @@ -39,9 +39,15 @@ public class PyTypeParserTest extends PyTestCase { public void testDictType() { myFixture.configureByFile("typeParser/typeParser.py"); - final PyCollectionType type = (PyCollectionType) PyTypeParser.getTypeByName(myFixture.getFile(), "dict from string to MyObject"); + final PyCollectionType type = (PyCollectionType) PyTypeParser.getTypeByName(myFixture.getFile(), "dict from str to MyObject"); + assertNotNull(type); assertClassType(type, "dict"); - assertClassType(type.getElementType(TypeEvalContext.fast()), "MyObject"); + final PyType elementType = type.getElementType(TypeEvalContext.fast()); + assertInstanceOf(elementType, PyTupleType.class); + final PyTupleType tupleType = (PyTupleType)elementType; + assertEquals(2, tupleType.getElementCount()); + assertClassType(tupleType.getElementType(0), "str"); + assertClassType(tupleType.getElementType(1), "MyObject"); } public void testUnionType() {