Initial support for generic types in Python

This commit is contained in:
Andrey Vlasovskikh
2011-12-20 16:13:10 +04:00
parent a89dd1b950
commit 71bea78f8a
14 changed files with 453 additions and 63 deletions
@@ -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);
@@ -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 = \
@@ -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<PyGenericType, PyType> substitutions = new HashMap<PyGenericType, PyType>();
final CallArgumentsMapping res = args.analyzeCall(resolveWithoutImplicits());
final Map<PyExpression, PyNamedParameter> mapped = res.getPlainMappedParams();
for (Map.Entry<PyExpression, PyNamedParameter> 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<PyExpression, PyNamedParameter> 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<PyGenericType, PyType>(), 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<PyGenericType, PyType> 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<PyType> 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);
}
@@ -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();
@@ -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();
}
@@ -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<PyFunctionStub> 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<PyGenericType, PyType> substitutions = PyTypeChecker.unifyGenericCall(this, (PyCallExpression)parent, context);
if (substitutions != null) {
final Ref<PyType> 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<PyFunctionStub> imp
@Nullable
private PyType getYieldStatementType(@NotNull final TypeEvalContext context) {
PyType elementType = null;
Ref<PyType> 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<PyFunctionStub> 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;
@@ -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<PyNamedParame
// must be 'self' or 'cls'
final PyClass containingClass = func.getContainingClass();
if (containingClass != null) {
PyType initType = null;
final PyFunction init = containingClass.findInitOrNew(true);
if (init != null) {
initType = init.getReturnType(context, null);
}
else {
final PyStdlibTypeProvider stdlib = PyStdlibTypeProvider.getInstance();
if (stdlib != null) {
initType = stdlib.getConstructorType(containingClass);
}
}
if (initType != null && !(initType instanceof PyNoneType)) {
return initType;
}
return new PyClassType(containingClass, modifier == PyFunction.Modifier.CLASSMETHOD);
}
}
@@ -369,7 +369,7 @@ public class PyReferenceExpressionImpl extends PyElementImpl implements PyRefere
}
}
private static class QualifiedResolveResultEmpty implements QualifiedResolveResult {
public static class QualifiedResolveResultEmpty implements QualifiedResolveResult {
// a trivial implementation
public QualifiedResolveResultEmpty() {
@@ -0,0 +1,70 @@
package com.jetbrains.python.psi.types;
import com.intellij.util.ProcessingContext;
import com.jetbrains.python.psi.AccessDirection;
import com.jetbrains.python.psi.PyExpression;
import com.jetbrains.python.psi.resolve.PyResolveContext;
import com.jetbrains.python.psi.resolve.RatedResolveResult;
import org.jetbrains.annotations.NotNull;
import org.jetbrains.annotations.Nullable;
import java.util.List;
/**
* @author vlan
*/
public class PyGenericType implements PyType {
private final String myName;
public PyGenericType(@NotNull String name) {
this.myName = name;
}
@Nullable
@Override
public List<? extends RatedResolveResult> 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;
}
}
@@ -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<PyType, String>() {
return "(" + StringUtil.join(myElementTypes, new Function<PyType, String>() {
@Nullable
public String fun(PyType pyType) {
return pyType == null ? "unknown" : pyType.getName();
}
@@ -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<PyGenericType, PyType> substitutions) {
return match(expected, actual, context, substitutions, true);
}
private static boolean match(@Nullable PyType expected, @Nullable PyType actual, @NotNull TypeEvalContext context,
@Nullable Map<PyGenericType, PyType> 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<PyGenericType> collected = new HashSet<PyGenericType>();
collectGenerics(type, context, collected);
return !collected.isEmpty();
}
private static void collectGenerics(@Nullable PyType type, @NotNull TypeEvalContext context, @NotNull Set<PyGenericType> 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<PyType> substitute(@Nullable PyType type, @NotNull Map<PyGenericType, PyType> 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<PyType> results = new ArrayList<PyType>();
for (PyType t : union.getMembers()) {
final Ref<PyType> 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<PyType> 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<PyType> results = new ArrayList<PyType>();
for (int i = 0; i < n; i++) {
final Ref<PyType> 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<PyGenericType, PyType> unifyGenericCall(@NotNull PyFunction function, @NotNull PyCallExpression call,
@NotNull TypeEvalContext context) {
final PyArgumentList args = call.getArgumentList();
if (args == null) {
return null;
}
final Map<PyGenericType, PyType> substitutions = collectCallGenerics(function, call, context);
final CallArgumentsMapping res = args.analyzeCall(PyResolveContext.noImplicits().withTypeEvalContext(context));
for (Map.Entry<PyExpression, PyNamedParameter> 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<PyGenericType, PyType> collectCallGenerics(@NotNull PyFunction function, @NotNull PyCallExpression call,
@NotNull TypeEvalContext context) {
final Map<PyGenericType, PyType> substitutions = new HashMap<PyGenericType, PyType>();
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<PyGenericType> generics = new HashSet<PyGenericType>();
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;
@@ -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;
}
@@ -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");
@@ -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);