diff --git a/python/src/com/jetbrains/python/codeInsight/stdlib/PyStdlibTypeProvider.java b/python/src/com/jetbrains/python/codeInsight/stdlib/PyStdlibTypeProvider.java index e49b062ebd6b..b4d75e321417 100644 --- a/python/src/com/jetbrains/python/codeInsight/stdlib/PyStdlibTypeProvider.java +++ b/python/src/com/jetbrains/python/codeInsight/stdlib/PyStdlibTypeProvider.java @@ -1,5 +1,6 @@ package com.jetbrains.python.codeInsight.stdlib; +import com.intellij.openapi.extensions.Extensions; import com.intellij.openapi.util.Ref; import com.intellij.openapi.vfs.VirtualFile; import com.intellij.psi.PsiElement; @@ -8,6 +9,7 @@ import com.jetbrains.python.documentation.StructuredDocString; 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.*; @@ -26,6 +28,16 @@ public class PyStdlibTypeProvider extends PyTypeProviderBase { private Properties myStdlibTypes2 = new Properties(); private Properties myStdlibTypes3 = new Properties(); + @Nullable + public static PyStdlibTypeProvider getInstance() { + for (PyTypeProvider typeProvider : Extensions.getExtensions(PyTypeProvider.EP_NAME)) { + if (typeProvider instanceof PyStdlibTypeProvider) { + return (PyStdlibTypeProvider)typeProvider; + } + } + return null; + } + @Override public PyType getReferenceType(@NotNull PsiElement referenceTarget, @NotNull TypeEvalContext context, @Nullable PsiElement anchor) { if (referenceTarget instanceof PyFunction && @@ -56,6 +68,19 @@ public class PyStdlibTypeProvider extends PyTypeProviderBase { return null; } + @Nullable + public PyType getConstructorType(@NotNull PyClass cls) { + final String classQName = cls.getQualifiedName(); + if (classQName != null) { + final PyQualifiedName canonicalQName = ResolveImportUtil.restoreStdlibCanonicalPath(PyQualifiedName.fromDottedString(classQName)); + if (canonicalQName != null) { + final PyQualifiedName qname = canonicalQName.append(PyNames.INIT); + return getReturnTypeByQName(qname.toString(), cls); + } + } + return null; + } + @Nullable private PyType getReturnTypeByQName(@NotNull String qname, @NotNull PsiElement anchor) { final LanguageLevel level = LanguageLevel.forElement(anchor); diff --git a/python/src/com/jetbrains/python/codeInsight/stdlib/StdlibTypes2.properties b/python/src/com/jetbrains/python/codeInsight/stdlib/StdlibTypes2.properties index e90f1addfc06..39db2b4e1e72 100644 --- a/python/src/com/jetbrains/python/codeInsight/stdlib/StdlibTypes2.properties +++ b/python/src/com/jetbrains/python/codeInsight/stdlib/StdlibTypes2.properties @@ -34,17 +34,20 @@ __builtin__.divmod = \ :rtype: (int or long or float, int or long or float) \n\ __builtin__.enumerate.__init__ = \ - :type iterable: collections.Iterable \n\ + :type iterable: collections.Iterable of T \n\ :type start: int or long \n\ - :rtype: enumerate of (int, unknown) \n\ + :rtype: enumerate of (int, T) \n\ __builtin__.enumerate.__new__ = \ - :rtype: enumerate of (int, unknown) \n\ + :rtype: enumerate of (int, T) \n\ + +__builtin__.enumerate.next = \ + :rtype: (int, T) \n\ __builtin__.filter = \ :type function_or_none: collections.Callable or None \n\ - :type sequence: collections.Iterable \n\ - :rtype: list \n\ + :type sequence: collections.Iterable of T \n\ + :rtype: list of T \n\ __builtin__.filter.1 = \ :type sequence: bytes \n\ @@ -73,7 +76,8 @@ __builtin__.hash = \ :rtype: int \n\ __builtin__.iter = \ - :rtype: collections.Iterable \n\ + :type source: collections.Iterable of T \n\ + :rtype: collections.Iterator of T \n\ __builtin__.len = \ :type p_object: collections.Sized \n\ @@ -1050,6 +1054,10 @@ __builtin__.tuple.__getitem__ = \ ## 5.8 Mapping types +__builtin__.dict.__init__ = \ + :type seq: collections.Iterable of (T, V) \n\ + :rtype: dict of (T, V) \n\ + __builtin__.dict.__getitem__ = \ :type y: object \n\ :rtype: unknown \n\ @@ -1076,7 +1084,7 @@ __builtin__.dict.has_key = \ :rtype: bool \n\ __builtin__.dict.items = \ - :rtype: list of (unknown, unknown) \n\ + :rtype: list of (T, V) \n\ __builtin__.dict.iteritems = \ :rtype: collections.Iterable of (unknown, unknown) \n\ @@ -1088,10 +1096,10 @@ __builtin__.dict.itervalues = \ :rtype: collections.Iterable \n\ __builtin__.dict.keys = \ - :rtype: list \n\ + :rtype: list of T \n\ __builtin__.dict.values = \ - :rtype: list \n\ + :rtype: list of V \n\ ## 5.9. File objects @@ -1334,6 +1342,16 @@ datetime.datetime.__rsub__ = \ :type y: datetime.datetime \n\ :rtype: datetime.timedelta \n\ + +## 8.3. collections + +collections.Iterator.__init__ = \ + :rtype: collections.Iterator of T \n\ + +collections.Iterator.next = \ + :rtype: T \n\ + + ## 9.4. decimal decimal.Decimal.as_tuple = \ diff --git a/python/src/com/jetbrains/python/inspections/PyTypeCheckerInspection.java b/python/src/com/jetbrains/python/inspections/PyTypeCheckerInspection.java index a92f26cb7a49..0c4f1ded6b61 100644 --- a/python/src/com/jetbrains/python/inspections/PyTypeCheckerInspection.java +++ b/python/src/com/jetbrains/python/inspections/PyTypeCheckerInspection.java @@ -4,6 +4,7 @@ import com.intellij.codeInspection.LocalInspectionToolSession; 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.jetbrains.python.documentation.PythonDocumentationProvider; @@ -13,6 +14,7 @@ import org.jetbrains.annotations.Nls; import org.jetbrains.annotations.NotNull; import org.jetbrains.annotations.Nullable; +import java.util.HashMap; import java.util.Map; /** @@ -40,9 +42,17 @@ 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 Map mapped = res.getPlainMappedParams(); - for (Map.Entry entry : mapped.entrySet()) { + 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)); + } + } + for (Map.Entry entry : res.getPlainMappedParams().entrySet()) { final PyNamedParameter p = entry.getValue(); if (p.isPositionalContainer() || p.isKeywordContainer()) { // TODO: Support *args, **kwargs @@ -50,7 +60,7 @@ public class PyTypeCheckerInspection extends PyInspection { } final PyType argType = entry.getKey().getType(myTypeEvalContext); final PyType paramType = p.getType(myTypeEvalContext); - checkTypes(paramType, argType, entry.getKey(), myTypeEvalContext, true); + checkTypes(paramType, argType, entry.getKey(), myTypeEvalContext, substitutions, true); } } } @@ -107,7 +117,7 @@ public class PyTypeCheckerInspection extends PyInspection { if (p != null) { final PyType argType = argument.getType(myTypeEvalContext); final PyType paramType = p.getType(myTypeEvalContext); - return checkTypes(paramType, argType, argument, myTypeEvalContext, registerProblem); + return checkTypes(paramType, argType, argument, myTypeEvalContext, new HashMap(), registerProblem); } } } @@ -115,12 +125,24 @@ public class PyTypeCheckerInspection extends PyInspection { } @Nullable - private String checkTypes(PyType superType, PyType subType, PsiElement node, TypeEvalContext context, boolean registerPoblem) { + private String checkTypes(@Nullable PyType superType, @Nullable PyType subType, @Nullable PsiElement node, + @NotNull TypeEvalContext context, @NotNull Map substitutions, boolean registerPoblem) { if (subType != null && superType != null) { - if (!PyTypeChecker.match(superType, subType, context)) { - final String msg = String.format("Expected type '%s', got '%s' instead", - PythonDocumentationProvider.getTypeName(superType, context), - PythonDocumentationProvider.getTypeName(subType, myTypeEvalContext)); + if (!PyTypeChecker.match(superType, subType, context, substitutions)) { + final String superName = PythonDocumentationProvider.getTypeName(superType, context); + String expected = String.format("'%s'", superName); + if (PyTypeChecker.hasGenerics(superType, context)) { + final Ref subst = PyTypeChecker.substitute(superType, substitutions, context); + if (subst != null) { + expected = String.format("'%s' (matched generic type '%s')", + PythonDocumentationProvider.getTypeName(subst.get(), context), + superName); + + } + } + final String msg = String.format("Expected type %s, got '%s' instead", + expected, + PythonDocumentationProvider.getTypeName(subType, context)); if (registerPoblem) { registerProblem(node, msg); } diff --git a/python/src/com/jetbrains/python/psi/PyBinaryExpression.java b/python/src/com/jetbrains/python/psi/PyBinaryExpression.java index 1279844ddc8a..ce6a94eb7ec1 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 PyQualifiedExpression, PyReferenceOwner { +public interface PyBinaryExpression extends PyReferenceExpression { 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 aa87322295f2..2b5c2e7bc856 100644 --- a/python/src/com/jetbrains/python/psi/impl/PyBinaryExpressionImpl.java +++ b/python/src/com/jetbrains/python/psi/impl/PyBinaryExpressionImpl.java @@ -11,6 +11,8 @@ 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; @@ -32,6 +34,7 @@ public class PyBinaryExpressionImpl extends PyElementImpl implements PyBinaryExp pyVisitor.visitPyBinaryExpression(this); } + @Nullable public PyExpression getLeftExpression() { return PsiTreeUtil.getChildOfType(this, PyExpression.class); } @@ -67,6 +70,7 @@ public class PyBinaryExpressionImpl extends PyElementImpl implements PyBinaryExp return buf.toString().equals(chars); } + @Nullable public PyExpression getOppositeExpression(PyExpression expression) throws IllegalArgumentException { PyExpression right = getRightExpression(); PyExpression left = getLeftExpression(); @@ -86,7 +90,7 @@ public class PyBinaryExpressionImpl extends PyElementImpl implements PyBinaryExp if (left == child.getPsi() && right != null) { replace(right); } - else if (right == child.getPsi()) { + else if (right == child.getPsi() && left != null) { replace(left); } else { @@ -94,19 +98,27 @@ public class PyBinaryExpressionImpl extends PyElementImpl implements PyBinaryExp } } + @NotNull @Override - public PsiReference getReference() { + public QualifiedResolveResult followAssignmentsChain(PyResolveContext resolveContext) { + return new PyReferenceExpressionImpl.QualifiedResolveResultEmpty(); + } + + @Override + public PyQualifiedName asQualifiedName() { + return PyQualifiedName.fromReferenceChain(PyResolveUtil.unwindQualifiers(this)); + } + + @NotNull + @Override + public PsiPolyVariantReference getReference() { return getReference(PyResolveContext.noImplicits()); } - @Nullable + @NotNull @Override public PsiPolyVariantReference getReference(PyResolveContext context) { - final PyElementType t = getOperator(); - if (t != null && t.getSpecialMethodName() != null) { - return new PyOperatorReferenceImpl(this, context); - } - return null; + return new PyOperatorReferenceImpl(this, context); } public PyType getType(@NotNull TypeEvalContext context) { @@ -115,20 +127,18 @@ public class PyBinaryExpressionImpl extends PyElementImpl implements PyBinaryExp final PsiElement operator = getPsiOperator(); if (lhs != null && rhs != null && operator != null) { final PsiReference ref = getReference(PyResolveContext.noImplicits().withTypeEvalContext(context)); - if (ref != null) { - final PsiElement resolved = ref.resolve(); - if (resolved instanceof Callable) { - context.trace("Binary operator reference resolved to %s", resolved); - final PyType res = ((Callable)resolved).getReturnType(context, null); - context.trace("Binary operator reference resolve type is %s", res); - if (!PyTypeChecker.isUnknown(res) && !(res instanceof PyNoneType)) { - return res; - } - } - else { - context.trace("Failed to resolve binary operator reference %s", this); + final PsiElement resolved = ref.resolve(); + if (resolved instanceof Callable) { + context.trace("Binary operator reference resolved to %s", resolved); + final PyType res = ((Callable)resolved).getReturnType(context, this); + context.trace("Binary operator reference resolve type is %s", res); + if (!PyTypeChecker.isUnknown(res) && !(res instanceof PyNoneType)) { + return res; } } + else { + context.trace("Failed to resolve binary operator reference %s", this); + } if (PyNames.COMPARISON_OPERATORS.contains(getReferencedName())) { return PyBuiltinCache.getInstance(this).getBoolType(); } diff --git a/python/src/com/jetbrains/python/psi/impl/PyFunctionImpl.java b/python/src/com/jetbrains/python/psi/impl/PyFunctionImpl.java index fa0452ee37c6..caa9f02e015c 100644 --- a/python/src/com/jetbrains/python/psi/impl/PyFunctionImpl.java +++ b/python/src/com/jetbrains/python/psi/impl/PyFunctionImpl.java @@ -3,6 +3,7 @@ package com.jetbrains.python.psi.impl; import com.intellij.lang.ASTNode; import com.intellij.openapi.extensions.Extensions; 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; @@ -129,7 +130,28 @@ public class PyFunctionImpl extends PyPresentableElementImpl imp } @Nullable - public PyType getReturnType(TypeEvalContext typeEvalContext, @Nullable PyReferenceExpression callSite) { + public PyType getReturnType(@NotNull TypeEvalContext context, @Nullable PyReferenceExpression callSite) { + final PyType type = getGenericReturnType(context, callSite); + 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(); + } + } + } + } + return null; + } + return type; + } + + @Nullable + private PyType getGenericReturnType(TypeEvalContext typeEvalContext, @Nullable PyReferenceExpression callSite) { if (typeEvalContext.maySwitchToAST(this)) { PyAnnotation anno = getAnnotation(); if (anno != null) { @@ -161,7 +183,7 @@ public class PyFunctionImpl extends PyPresentableElementImpl imp @Nullable private PyType getYieldStatementType(@NotNull final TypeEvalContext context) { - PyType elementType = null; + Ref elementType = null; final PyBuiltinCache cache = PyBuiltinCache.getInstance(this); final PyClass listClass = cache.getClass("list"); final PyStatementList statements = getStatementList(); @@ -175,16 +197,16 @@ public class PyFunctionImpl extends PyPresentableElementImpl imp }); final int n = types.size(); if (n == 1) { - elementType = types.iterator().next(); + elementType = Ref.create(types.iterator().next()); } else if (n > 0) { - elementType = PyUnionType.union(types); + elementType = Ref.create(PyUnionType.union(types)); } } if (elementType != null) { final PyType it = PyTypeParser.getTypeByName(this, "Iterator"); if (it instanceof PyClassType) { - return new PyCollectionTypeImpl(((PyClassType)it).getPyClass(), false, elementType); + return new PyCollectionTypeImpl(((PyClassType)it).getPyClass(), false, elementType.get()); } } return null; diff --git a/python/src/com/jetbrains/python/psi/impl/PyNamedParameterImpl.java b/python/src/com/jetbrains/python/psi/impl/PyNamedParameterImpl.java index 589a7a541035..fe0d8dacd8f3 100644 --- a/python/src/com/jetbrains/python/psi/impl/PyNamedParameterImpl.java +++ b/python/src/com/jetbrains/python/psi/impl/PyNamedParameterImpl.java @@ -13,13 +13,11 @@ import com.jetbrains.python.PyElementTypes; import com.jetbrains.python.PyNames; import com.jetbrains.python.PyTokenTypes; import com.jetbrains.python.PythonDialectsTokenSetProvider; +import com.jetbrains.python.codeInsight.stdlib.PyStdlibTypeProvider; import com.jetbrains.python.documentation.StructuredDocString; import com.jetbrains.python.psi.*; import com.jetbrains.python.psi.stubs.PyNamedParameterStub; -import com.jetbrains.python.psi.types.PyClassType; -import com.jetbrains.python.psi.types.PyType; -import com.jetbrains.python.psi.types.PyTypeParser; -import com.jetbrains.python.psi.types.TypeEvalContext; +import com.jetbrains.python.psi.types.*; import org.jetbrains.annotations.NotNull; import org.jetbrains.annotations.Nullable; @@ -161,6 +159,20 @@ public class PyNamedParameterImpl extends PyPresentableElementImpl resolveMember(String name, + @Nullable PyExpression location, + AccessDirection direction, + PyResolveContext resolveContext) { + return null; + } + + @Override + public Object[] getCompletionVariants(String completionPrefix, PyExpression location, ProcessingContext context) { + return new Object[0]; + } + + @NotNull + @Override + public String getName() { + return myName; + } + + @Override + public boolean isBuiltin(TypeEvalContext context) { + return false; + } + + @Override + public boolean equals(Object o) { + if (this == o) { + return true; + } + if (o == null || getClass() != o.getClass()) { + return false; + } + final PyGenericType type = (PyGenericType)o; + return myName.equals(type.myName); + } + + @Override + public int hashCode() { + return myName.hashCode(); + } + + @NotNull + @Override + public String toString() { + return "PyGenericType: " + myName; + } +} diff --git a/python/src/com/jetbrains/python/psi/types/PyTupleType.java b/python/src/com/jetbrains/python/psi/types/PyTupleType.java index 74843233b2f9..f070167e317d 100644 --- a/python/src/com/jetbrains/python/psi/types/PyTupleType.java +++ b/python/src/com/jetbrains/python/psi/types/PyTupleType.java @@ -6,6 +6,7 @@ import com.intellij.util.Function; import com.jetbrains.python.psi.PyExpression; import com.jetbrains.python.psi.impl.PyBuiltinCache; import com.jetbrains.python.psi.impl.PyConstantExpressionEvaluator; +import org.jetbrains.annotations.Nullable; import java.util.Arrays; @@ -26,7 +27,8 @@ public class PyTupleType extends PyClassType implements PySubscriptableType { } public String getName() { - return "tuple(" + StringUtil.join(myElementTypes, new Function() { + return "(" + StringUtil.join(myElementTypes, new Function() { + @Nullable public String fun(PyType pyType) { return pyType == null ? "unknown" : pyType.getName(); } diff --git a/python/src/com/jetbrains/python/psi/types/PyTypeChecker.java b/python/src/com/jetbrains/python/psi/types/PyTypeChecker.java index 94d6ccd2f825..8460f5db9df5 100644 --- a/python/src/com/jetbrains/python/psi/types/PyTypeChecker.java +++ b/python/src/com/jetbrains/python/psi/types/PyTypeChecker.java @@ -1,10 +1,16 @@ package com.jetbrains.python.psi.types; +import com.intellij.openapi.extensions.Extensions; +import com.intellij.openapi.util.Ref; import com.jetbrains.python.PyNames; -import com.jetbrains.python.psi.PyClass; +import com.jetbrains.python.codeInsight.stdlib.PyStdlibTypeProvider; +import com.jetbrains.python.psi.*; +import com.jetbrains.python.psi.impl.PyTypeProvider; +import com.jetbrains.python.psi.resolve.PyResolveContext; +import org.jetbrains.annotations.NotNull; import org.jetbrains.annotations.Nullable; -import java.util.List; +import java.util.*; /** * @author vlan @@ -13,11 +19,17 @@ public class PyTypeChecker { private PyTypeChecker() { } - public static boolean match(@Nullable PyType expected, @Nullable PyType actual, TypeEvalContext context) { - return match(expected, actual, context, true); + public static boolean match(@Nullable PyType expected, @Nullable PyType actual, @NotNull TypeEvalContext context) { + return match(expected, actual, context, null, true); } - private static boolean match(@Nullable PyType expected, @Nullable PyType actual, TypeEvalContext context, boolean resolveReferences) { + public static boolean match(@Nullable PyType expected, @Nullable PyType actual, @NotNull TypeEvalContext context, + @Nullable Map substitutions) { + return match(expected, actual, context, substitutions, true); + } + + private static boolean match(@Nullable PyType expected, @Nullable PyType actual, @NotNull TypeEvalContext context, + @Nullable Map substitutions, boolean resolveReferences) { // TODO: subscriptable types?, module types?, etc. if (expected == null || actual == null) { return true; @@ -32,17 +44,17 @@ public class PyTypeChecker { return true; } if (expected instanceof PyTypeReference) { - return match(((PyTypeReference)expected).resolve(null, context), actual, context, resolveReferences); + return match(((PyTypeReference)expected).resolve(null, context), actual, context, substitutions, resolveReferences); } if (actual instanceof PyTypeReference) { - return match(expected, ((PyTypeReference)actual).resolve(null, context), context, false); + return match(expected, ((PyTypeReference)actual).resolve(null, context), context, substitutions, false); } if (isUnknown(actual)) { return true; } if (actual instanceof PyUnionType) { for (PyType m : ((PyUnionType)actual).getMembers()) { - if (!match(expected, m, context, resolveReferences)) { + if (!match(expected, m, context, substitutions, resolveReferences)) { return false; } } @@ -50,7 +62,7 @@ public class PyTypeChecker { } if (expected instanceof PyUnionType) { for (PyType t : ((PyUnionType)expected).getMembers()) { - if (match(t, actual, context, resolveReferences)) { + if (match(t, actual, context, substitutions, resolveReferences)) { return true; } } @@ -65,7 +77,7 @@ public class PyTypeChecker { } final PyType superElementType = ((PyCollectionType)expected).getElementType(context); final PyType subElementType = ((PyCollectionType)actual).getElementType(context); - return match(superElementType, subElementType, context, resolveReferences); + return match(superElementType, subElementType, context, substitutions, resolveReferences); } else if (expected instanceof PyTupleType && actual instanceof PyTupleType) { final PyTupleType superTupleType = (PyTupleType)expected; @@ -75,7 +87,7 @@ public class PyTypeChecker { } else { for (int i = 0; i < superTupleType.getElementCount(); i++) { - if (!match(superTupleType.getElementType(i), subTupleType.getElementType(i), context, resolveReferences)) { + if (!match(superTupleType.getElementType(i), subTupleType.getElementType(i), context, substitutions, resolveReferences)) { return false; } } @@ -92,6 +104,17 @@ public class PyTypeChecker { if (expected.equals(actual)) { return true; } + if (expected instanceof PyGenericType && substitutions != null) { + final PyGenericType generic = (PyGenericType)expected; + final PyType subst = substitutions.get(generic); + if (subst != null) { + return match(subst, actual, context, substitutions, resolveReferences); + } + else { + substitutions.put(generic, actual); + return true; + } + } final String superName = expected.getName(); final String subName = actual.getName(); // TODO: No inheritance check for builtin numerics at this moment @@ -125,6 +148,149 @@ public class PyTypeChecker { return false; } + public static boolean hasGenerics(@Nullable PyType type, @NotNull TypeEvalContext context) { + final Set collected = new HashSet(); + collectGenerics(type, context, collected); + return !collected.isEmpty(); + } + + private static void collectGenerics(@Nullable PyType type, @NotNull TypeEvalContext context, @NotNull Set collected) { + if (type instanceof PyGenericType) { + collected.add((PyGenericType)type); + } + else if (type instanceof PyUnionType) { + final PyUnionType union = (PyUnionType)type; + for (PyType t : union.getMembers()) { + collectGenerics(t, context, collected); + } + } + else if (type instanceof PyCollectionType) { + final PyCollectionType collection = (PyCollectionType)type; + collectGenerics(collection.getElementType(context), context, collected); + } + else if (type instanceof PyTupleType) { + final PyTupleType tuple = (PyTupleType)type; + final int n = tuple.getElementCount(); + for (int i = 0; i < n; i++) { + collectGenerics(tuple.getElementType(i), context, collected); + } + } + } + + @Nullable + public static Ref substitute(@Nullable PyType type, @NotNull Map substitutions, + @NotNull TypeEvalContext context) { + if (hasGenerics(type, context)) { + if (type instanceof PyGenericType) { + final PyType subst = substitutions.get((PyGenericType)type); + return subst != null ? Ref.create(subst) : null; + } + else if (type instanceof PyUnionType) { + final PyUnionType union = (PyUnionType)type; + final List results = new ArrayList(); + for (PyType t : union.getMembers()) { + final Ref subst = substitute(t, substitutions, context); + if (subst == null) { + return null; + } + results.add(subst.get()); + } + return Ref.create(PyUnionType.union(results)); + } + else if (type instanceof PyCollectionTypeImpl) { + final PyCollectionTypeImpl collection = (PyCollectionTypeImpl)type; + final PyType elem = collection.getElementType(context); + final Ref subst = substitute(elem, substitutions, context); + if (subst == null) { + return null; + } + final PyType result = new PyCollectionTypeImpl(collection.getPyClass(), collection.isDefinition(), subst.get()); + return Ref.create(result); + } + else if (type instanceof PyTupleType) { + final PyTupleType tuple = (PyTupleType)type; + final int n = tuple.getElementCount(); + final List results = new ArrayList(); + for (int i = 0; i < n; i++) { + final Ref subst = substitute(tuple.getElementType(i), substitutions, context); + if (subst == null) { + return null; + } + results.add(subst.get()); + } + final PyType result = new PyTupleType((PyTupleType)type, results.toArray(new PyType[results.size()])); + return Ref.create(result); + } + } + return Ref.create(type); + } + + @Nullable + public static Map unifyGenericCall(@NotNull PyFunction function, @NotNull PyCallExpression call, + @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 PyNamedParameter p = entry.getValue(); + final String name = p.getName(); + if (p.isPositionalContainer() || p.isKeywordContainer() || name == null) { + continue; + } + final PyType argType = entry.getKey().getType(context); + final PyType paramType = p.getType(context); + if (!match(paramType, argType, context, substitutions)) { + return null; + } + } + return substitutions; + } + + @NotNull + public static Map collectCallGenerics(@NotNull PyFunction function, @NotNull PyCallExpression call, + @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); + } + 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); + } + } + } + } + } + } + return substitutions; + } + private static boolean matchClasses(@Nullable PyClass superClass, @Nullable PyClass subClass) { if (superClass == null || subClass == null || subClass.isSubclass(superClass) || PyABCUtil.isSubclass(subClass, superClass)) { return true; diff --git a/python/src/com/jetbrains/python/psi/types/PyTypeParser.java b/python/src/com/jetbrains/python/psi/types/PyTypeParser.java index 53a684e84766..8a95c948ac4d 100644 --- a/python/src/com/jetbrains/python/psi/types/PyTypeParser.java +++ b/python/src/com/jetbrains/python/psi/types/PyTypeParser.java @@ -138,6 +138,9 @@ public class PyTypeParser { types.put(whole, t); return t; } + if (type.length() == 1 && 'T' <= type.charAt(0) && type.charAt(0) <= 'Z') { + return new PyGenericType(type); + } final Matcher m = PARAMETRIZED_CLASS.matcher(type); if (m.matches()) { final PyType objType = parseObjectType(anchor, m.group(1), builtinCache, types, fullRanges, offset + m.start(1)); @@ -229,8 +232,10 @@ public class PyTypeParser { PyClassType dict = PyBuiltinCache.getInstance(anchor).getDictType(); if (dict != null) { if (m.matches()) { + PyType from = parse(anchor, m.group(1), types, fullRanges, offset + m.start(1)); PyType to = parse(anchor, m.group(2), types, fullRanges, offset + m.start(2)); - return new PyCollectionTypeImpl(dict.getPyClass(), false, to); + final PyType p = new PyTupleType(anchor, new PyType[] {from, to}); + return new PyCollectionTypeImpl(dict.getPyClass(), false, p); } return dict; } diff --git a/python/testSrc/com/jetbrains/python/PyTypeParserTest.java b/python/testSrc/com/jetbrains/python/PyTypeParserTest.java index f966b0da1c8b..327634fb1c9d 100644 --- a/python/testSrc/com/jetbrains/python/PyTypeParserTest.java +++ b/python/testSrc/com/jetbrains/python/PyTypeParserTest.java @@ -62,6 +62,14 @@ public class PyTypeParserTest extends PyTestCase { assertEquals(7, result.getTypes().values().size()); } + public void testGenericType() { + myFixture.configureByFile("typeParser/typeParser.py"); + final PyType type = PyTypeParser.getTypeByName(myFixture.getFile(), "T"); + assertNotNull(type); + assertInstanceOf(type, PyGenericType.class); + assertEquals("T", type.getName()); + } + // PY-4223 public void testSphinxFormattedType() { myFixture.configureByFile("typeParser/typeParser.py"); diff --git a/python/testSrc/com/jetbrains/python/PyTypeTest.java b/python/testSrc/com/jetbrains/python/PyTypeTest.java index cc1e0c2510f3..e46f925dded0 100644 --- a/python/testSrc/com/jetbrains/python/PyTypeTest.java +++ b/python/testSrc/com/jetbrains/python/PyTypeTest.java @@ -224,7 +224,7 @@ public class PyTypeTest extends PyTestCase { assertTrue(PyTypeChecker.isUnknown(t)); doTest("int", text); } - + public void testReturnTypeReferenceEquality() { final String text = "def foo(x): return x.bar\n" + "def xyzzy(a, x):\n" + @@ -301,6 +301,36 @@ public class PyTypeTest extends PyTestCase { assertFalse(actual.isBuiltin(context)); } + public void testGenericConcrete() { + PyExpression expr = parseExpr("def f(x):\n" + + " '''\n" + + " :type x: T\n" + + " :rtype: T\n" + + " '''\n" + + " return x\n" + + "\n" + + "expr = f(1)\n"); + TypeEvalContext context = TypeEvalContext.slow().withTracing(); + PyType actual = expr.getType(context); + assertNotNull(actual); + assertEquals("int", actual.getName()); + } + + public void testGenericConcreteMismatch() { + PyExpression expr = parseExpr("def f(x, y):\n" + + " '''\n" + + " :type x: T\n" + + " :rtype: T\n" + + " '''\n" + + " return x\n" + + "\n" + + "expr = f(1)\n"); + TypeEvalContext context = TypeEvalContext.slow().withTracing(); + PyType actual = expr.getType(context); + assertNotNull(actual); + assertEquals("int", actual.getName()); + } + private PyExpression parseExpr(String text) { myFixture.configureByText(PythonFileType.INSTANCE, text); return myFixture.findElementByText("expr", PyExpression.class);