mirror of
https://gitflic.ru/project/openide/openide.git
synced 2026-09-27 10:03:11 +07:00
Initial support for generic types in Python
This commit is contained in:
@@ -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);
|
||||
|
||||
Reference in New Issue
Block a user