diff --git a/python/psi-api/src/com/jetbrains/python/psi/PyAnnotation.java b/python/psi-api/src/com/jetbrains/python/psi/PyAnnotation.java index a6182fb2eb65..71dc7286d84b 100644 --- a/python/psi-api/src/com/jetbrains/python/psi/PyAnnotation.java +++ b/python/psi-api/src/com/jetbrains/python/psi/PyAnnotation.java @@ -22,9 +22,7 @@ import org.jetbrains.annotations.Nullable; /** * @author yole */ -public interface PyAnnotation extends PyElement, StubBasedPsiElement { - PyExpression getValue(); - +public interface PyAnnotation extends PyTypedElement, StubBasedPsiElement { @Nullable - PyClass resolveToClass(); + PyExpression getValue(); } diff --git a/python/psi-api/src/com/jetbrains/python/psi/PyTargetExpression.java b/python/psi-api/src/com/jetbrains/python/psi/PyTargetExpression.java index 0b67d85d3dc3..399b8f8a7bac 100644 --- a/python/psi-api/src/com/jetbrains/python/psi/PyTargetExpression.java +++ b/python/psi-api/src/com/jetbrains/python/psi/PyTargetExpression.java @@ -28,7 +28,7 @@ import org.jetbrains.annotations.Nullable; * @author yole */ public interface PyTargetExpression extends PyQualifiedExpression, PsiNamedElement, PsiNameIdentifierOwner, PyDocStringOwner, - StubBasedPsiElement { + PyQualifiedNameOwner, StubBasedPsiElement { PyTargetExpression[] EMPTY_ARRAY = new PyTargetExpression[0]; /** diff --git a/python/src/META-INF/python-core.xml b/python/src/META-INF/python-core.xml index 292237aac6ac..7adbc46c189c 100644 --- a/python/src/META-INF/python-core.xml +++ b/python/src/META-INF/python-core.xml @@ -589,6 +589,9 @@ + + + diff --git a/python/src/com/jetbrains/python/codeInsight/PyTypingTypeProvider.java b/python/src/com/jetbrains/python/codeInsight/PyTypingTypeProvider.java new file mode 100644 index 000000000000..f03befac6c14 --- /dev/null +++ b/python/src/com/jetbrains/python/codeInsight/PyTypingTypeProvider.java @@ -0,0 +1,375 @@ +/* + * Copyright 2000-2014 JetBrains s.r.o. + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package com.jetbrains.python.codeInsight; + +import com.google.common.collect.ImmutableMap; +import com.google.common.collect.ImmutableSet; +import com.intellij.openapi.project.Project; +import com.intellij.psi.PsiElement; +import com.intellij.psi.PsiPolyVariantReference; +import com.intellij.psi.util.QualifiedName; +import com.jetbrains.python.PyNames; +import com.jetbrains.python.psi.*; +import com.jetbrains.python.psi.impl.PyExpressionCodeFragmentImpl; +import com.jetbrains.python.psi.impl.PyPsiUtils; +import com.jetbrains.python.psi.resolve.PyResolveContext; +import com.jetbrains.python.psi.types.*; +import org.jetbrains.annotations.NotNull; +import org.jetbrains.annotations.Nullable; + +import java.util.ArrayList; +import java.util.Collections; +import java.util.List; + +/** + * @author vlan + */ +public class PyTypingTypeProvider extends PyTypeProviderBase { + private static ImmutableMap BUILTIN_COLLECTIONS = ImmutableMap.builder() + .put("typing.List", "list") + .put("typing.Dict", "dict") + .put("typing.Set", PyNames.SET) + .put("typing.Tuple", PyNames.TUPLE) + .build(); + + private static ImmutableSet GENERIC_CLASSES = ImmutableSet.builder() + .add("typing.Generic") + .add("typing.AbstractGeneric") + .add("typing.Protocol") + .build(); + + public PyType getParameterType(@NotNull PyNamedParameter param, @NotNull PyFunction func, @NotNull TypeEvalContext context) { + final PyAnnotation annotation = param.getAnnotation(); + if (annotation != null) { + // XXX: Requires switching from stub to AST + final PyExpression value = annotation.getValue(); + if (value != null) { + return getTypingType(value, context); + } + } + return null; + } + + @Nullable + @Override + public PyType getReturnType(@NotNull Callable callable, @NotNull TypeEvalContext context) { + if (callable instanceof PyFunction) { + final PyFunction function = (PyFunction)callable; + final PyAnnotation annotation = function.getAnnotation(); + if (annotation != null) { + // XXX: Requires switching from stub to AST + final PyExpression value = annotation.getValue(); + if (value != null) { + return getTypingType(value, context); + } + } + final PyType constructorType = getGenericConstructorType(function, context); + if (constructorType != null) { + return constructorType; + } + } + return null; + } + + @Nullable + private static PyType getGenericConstructorType(@NotNull PyFunction function, @NotNull TypeEvalContext context) { + if (PyUtil.isInit(function)) { + final PyClass cls = function.getContainingClass(); + if (cls != null) { + final List genericTypes = collectGenericTypes(cls, context); + + final PyType elementType; + if (genericTypes.size() == 1) { + elementType = genericTypes.get(0); + } + else if (genericTypes.size() > 1) { + elementType = PyTupleType.create(cls, genericTypes.toArray(new PyType[genericTypes.size()])); + } + else { + elementType = null; + } + + if (elementType != null) { + return new PyCollectionTypeImpl(cls, false, elementType); + } + } + } + return null; + } + + @NotNull + private static List collectGenericTypes(@NotNull PyClass cls, @NotNull TypeEvalContext context) { + boolean isGeneric = false; + for (PyClass ancestor : cls.getAncestorClasses(context)) { + if (GENERIC_CLASSES.contains(ancestor.getQualifiedName())) { + isGeneric = true; + break; + } + } + if (isGeneric) { + final ArrayList results = new ArrayList(); + // XXX: Requires switching from stub to AST + for (PyExpression expr : cls.getSuperClassExpressions()) { + if (expr instanceof PySubscriptionExpression) { + final PyExpression indexExpr = ((PySubscriptionExpression)expr).getIndexExpression(); + if (indexExpr != null) { + final PyGenericType genericType = getGenericType(indexExpr, context); + if (genericType != null) { + results.add(genericType); + } + } + } + } + return results; + } + return Collections.emptyList(); + } + + @Nullable + private static PyType getTypingType(@NotNull PyExpression expression, @NotNull TypeEvalContext context) { + final PyType unionType = getUnionType(expression, context); + if (unionType != null) { + return unionType; + } + final PyType parameterizedType = getParameterizedType(expression, context); + if (parameterizedType != null) { + return parameterizedType; + } + final PyType builtinCollection = getBuiltinCollection(expression, context); + if (builtinCollection != null) { + return builtinCollection; + } + final PyType genericType = getGenericType(expression, context); + if (genericType != null) { + return genericType; + } + final PyType functionType = getFunctionType(expression, context); + if (functionType != null) { + return functionType; + } + final PyType stringBasedType = getStringBasedType(expression, context); + if (stringBasedType != null) { + return stringBasedType; + } + return null; + } + + @Nullable + private static PyType getStringBasedType(@NotNull PyExpression expression, @NotNull TypeEvalContext context) { + if (expression instanceof PyStringLiteralExpression) { + // XXX: Requires switching from stub to AST + final String contents = ((PyStringLiteralExpression)expression).getStringValue(); + final Project project = expression.getProject(); + final PyExpressionCodeFragmentImpl codeFragment = new PyExpressionCodeFragmentImpl(project, "dummy.py", contents, false); + codeFragment.setContext(expression.getContainingFile()); + final PsiElement element = codeFragment.getFirstChild(); + if (element instanceof PyExpressionStatement) { + final PyExpression dummyExpr = ((PyExpressionStatement)element).getExpression(); + return getType(dummyExpr, context); + } + } + return null; + } + + @Nullable + private static PyType getFunctionType(@NotNull PyExpression expression, @NotNull TypeEvalContext context) { + if (expression instanceof PySubscriptionExpression) { + final PySubscriptionExpression subscriptionExpr = (PySubscriptionExpression)expression; + final PyExpression operand = subscriptionExpr.getOperand(); + final String operandName = resolveToQualifiedName(operand, context); + if ("typing.Function".equals(operandName)) { + final PyExpression indexExpr = subscriptionExpr.getIndexExpression(); + if (indexExpr instanceof PyTupleExpression) { + final PyTupleExpression tupleExpr = (PyTupleExpression)indexExpr; + final PyExpression[] elements = tupleExpr.getElements(); + if (elements.length == 2) { + final PyExpression parametersExpr = elements[0]; + if (parametersExpr instanceof PyListLiteralExpression) { + final List parameters = new ArrayList(); + final PyListLiteralExpression listExpr = (PyListLiteralExpression)parametersExpr; + for (PyExpression argExpr : listExpr.getElements()) { + parameters.add(new PyCallableParameterImpl(null, getType(argExpr, context))); + } + final PyExpression returnTypeExpr = elements[1]; + final PyType returnType = getType(returnTypeExpr, context); + return new PyCallableTypeImpl(parameters, returnType); + } + } + } + } + } + return null; + } + + @Nullable + private static PyType getUnionType(@NotNull PyExpression expression, @NotNull TypeEvalContext context) { + if (expression instanceof PySubscriptionExpression) { + final PySubscriptionExpression subscriptionExpr = (PySubscriptionExpression)expression; + final PyExpression operand = subscriptionExpr.getOperand(); + final String operandName = resolveToQualifiedName(operand, context); + if ("typing.Union".equals(operandName)) { + return PyUnionType.union(getIndexTypes(subscriptionExpr, context)); + } + } + return null; + } + + @Nullable + private static PyGenericType getGenericType(@NotNull PyExpression expression, @NotNull TypeEvalContext context) { + final PsiElement resolved = resolve(expression, context); + if (resolved instanceof PyTargetExpression) { + final PyTargetExpression targetExpr = (PyTargetExpression)resolved; + final QualifiedName calleeName = targetExpr.getCalleeName(); + if (calleeName != null && "typevar".equals(calleeName.toString())) { + // XXX: Requires switching from stub to AST + final PyExpression assigned = targetExpr.findAssignedValue(); + if (assigned instanceof PyCallExpression) { + final PyCallExpression assignedCall = (PyCallExpression)assigned; + final PyExpression callee = assignedCall.getCallee(); + if (callee != null) { + final String calleeQName = resolveToQualifiedName(callee, context); + if ("typing.typevar".equals(calleeQName)) { + final PyExpression[] arguments = assignedCall.getArguments(); + if (arguments.length > 0) { + final PyExpression firstArgument = arguments[0]; + if (firstArgument instanceof PyStringLiteralExpression) { + final String name = ((PyStringLiteralExpression)firstArgument).getStringValue(); + if (name != null) { + return new PyGenericType(name, getGenericTypeBound(arguments, context)); + } + } + } + } + + } + } + } + } + return null; + } + + @Nullable + private static PyType getGenericTypeBound(@NotNull PyExpression[] typeVarArguments, @NotNull TypeEvalContext context) { + final List types = new ArrayList(); + if (typeVarArguments.length > 1) { + final PyExpression secondArgument = typeVarArguments[1]; + if (secondArgument instanceof PyKeywordArgument) { + final PyKeywordArgument valuesArgument = (PyKeywordArgument)secondArgument; + final PyExpression valueExpr = PyPsiUtils.flattenParens(valuesArgument.getValueExpression()); + if (valueExpr instanceof PyTupleExpression) { + final PyTupleExpression tupleExpr = (PyTupleExpression)valueExpr; + for (PyExpression expr : tupleExpr.getElements()) { + types.add(getType(expr, context)); + } + } + } + } + return PyUnionType.union(types); + } + + @NotNull + private static List getIndexTypes(@NotNull PySubscriptionExpression expression, @NotNull TypeEvalContext context) { + final List types = new ArrayList(); + final PyExpression indexExpr = expression.getIndexExpression(); + if (indexExpr instanceof PyTupleExpression) { + final PyTupleExpression tupleExpr = (PyTupleExpression)indexExpr; + for (PyExpression expr : tupleExpr.getElements()) { + types.add(getType(expr, context)); + } + } + return types; + } + + @Nullable + private static PyType getParameterizedType(@NotNull PyExpression expression, @NotNull TypeEvalContext context) { + if (expression instanceof PySubscriptionExpression) { + final PySubscriptionExpression subscriptionExpr = (PySubscriptionExpression)expression; + final PyExpression operand = subscriptionExpr.getOperand(); + final PyExpression indexExpr = subscriptionExpr.getIndexExpression(); + final PyType operandType = getType(operand, context); + if (operandType instanceof PyClassType) { + final PyClass cls = ((PyClassType)operandType).getPyClass(); + if (PyNames.TUPLE.equals(cls.getQualifiedName())) { + final List indexTypes = getIndexTypes(subscriptionExpr, context); + return PyTupleType.create(expression, indexTypes.toArray(new PyType[indexTypes.size()])); + } + else if (indexExpr != null) { + final PyType indexType = context.getType(indexExpr); + return new PyCollectionTypeImpl(cls, false, indexType); + } + } + } + return null; + } + + @Nullable + private static PyType getBuiltinCollection(@NotNull PyExpression expression, @NotNull TypeEvalContext context) { + final String collectionName = resolveToQualifiedName(expression, context); + final String builtinName = BUILTIN_COLLECTIONS.get(collectionName); + return builtinName != null ? PyTypeParser.getTypeByName(expression, builtinName) : null; + } + + @Nullable + private static PyType getType(@NotNull PyExpression expression, @NotNull TypeEvalContext context) { + // It is possible to replace PyAnnotation.getType() with this implementation + final PyType typingType = getTypingType(expression, context); + if (typingType != null) { + return typingType; + } + final PyType type = context.getType(expression); + if (type instanceof PyClassLikeType) { + final PyClassLikeType classType = (PyClassLikeType)type; + if (classType.isDefinition()) { + return classType.toInstance(); + } + } + else if (type instanceof PyNoneType) { + return type; + } + return null; + } + + @Nullable + private static PsiElement resolve(@NotNull PyExpression expression, @NotNull TypeEvalContext context) { + if (expression instanceof PyReferenceOwner) { + final PyReferenceOwner referenceOwner = (PyReferenceOwner)expression; + final PyResolveContext resolveContext = PyResolveContext.noImplicits().withTypeEvalContext(context); + final PsiPolyVariantReference reference = referenceOwner.getReference(resolveContext); + final PsiElement element = reference.resolve(); + if (element instanceof PyFunction) { + final PyFunction function = (PyFunction)element; + if (PyUtil.isInit(function)) { + final PyClass cls = function.getContainingClass(); + if (cls != null) { + return cls; + } + } + } + return element; + } + return null; + } + + @Nullable + private static String resolveToQualifiedName(@NotNull PyExpression expression, @NotNull TypeEvalContext context) { + final PsiElement element = resolve(expression, context); + if (element instanceof PyQualifiedNameOwner) { + final PyQualifiedNameOwner qualifiedNameOwner = (PyQualifiedNameOwner)element; + return qualifiedNameOwner.getQualifiedName(); + } + return null; + } +} diff --git a/python/src/com/jetbrains/python/psi/impl/PyAnnotationImpl.java b/python/src/com/jetbrains/python/psi/impl/PyAnnotationImpl.java index f3c806dac7c5..f8a4d482c2f2 100644 --- a/python/src/com/jetbrains/python/psi/impl/PyAnnotationImpl.java +++ b/python/src/com/jetbrains/python/psi/impl/PyAnnotationImpl.java @@ -16,14 +16,16 @@ package com.jetbrains.python.psi.impl; import com.intellij.lang.ASTNode; -import com.intellij.psi.PsiElement; -import com.intellij.psi.PsiPolyVariantReference; import com.jetbrains.python.PyElementTypes; import com.jetbrains.python.psi.PyAnnotation; -import com.jetbrains.python.psi.PyClass; import com.jetbrains.python.psi.PyExpression; -import com.jetbrains.python.psi.PyReferenceExpression; import com.jetbrains.python.psi.stubs.PyAnnotationStub; +import com.jetbrains.python.psi.types.PyClassLikeType; +import com.jetbrains.python.psi.types.PyNoneType; +import com.jetbrains.python.psi.types.PyType; +import com.jetbrains.python.psi.types.TypeEvalContext; +import org.jetbrains.annotations.NotNull; +import org.jetbrains.annotations.Nullable; /** * @author yole @@ -37,19 +39,26 @@ public class PyAnnotationImpl extends PyBaseElementImpl implem super(stub, PyElementTypes.ANNOTATION); } + @Nullable @Override public PyExpression getValue() { return findChildByClass(PyExpression.class); } + @Nullable @Override - public PyClass resolveToClass() { - PyExpression expr = getValue(); - if (expr instanceof PyReferenceExpression) { - final PsiPolyVariantReference reference = ((PyReferenceExpression)expr).getReference(); - final PsiElement result = reference.resolve(); - if (result instanceof PyClass) { - return (PyClass) result; + public PyType getType(@NotNull TypeEvalContext context, @NotNull TypeEvalContext.Key key) { + final PyExpression value = getValue(); + if (value != null) { + final PyType type = context.getType(value); + if (type instanceof PyClassLikeType) { + final PyClassLikeType classType = (PyClassLikeType)type; + if (classType.isDefinition()) { + return classType.toInstance(); + } + } + else if (type instanceof PyNoneType) { + return type; } } return null; diff --git a/python/src/com/jetbrains/python/psi/impl/PyClassImpl.java b/python/src/com/jetbrains/python/psi/impl/PyClassImpl.java index 7538fbae2fd2..bb2104fc0c77 100644 --- a/python/src/com/jetbrains/python/psi/impl/PyClassImpl.java +++ b/python/src/com/jetbrains/python/psi/impl/PyClassImpl.java @@ -20,7 +20,6 @@ import com.intellij.lang.ASTNode; import com.intellij.openapi.util.Comparing; import com.intellij.openapi.util.Key; import com.intellij.openapi.util.NotNullLazyValue; -import com.intellij.openapi.vfs.VirtualFile; import com.intellij.psi.*; import com.intellij.psi.scope.PsiScopeProcessor; import com.intellij.psi.search.LocalSearchScope; @@ -187,6 +186,11 @@ public class PyClassImpl extends PyPresentableElementImpl implement } } } + // Heuristic: unfold Foo[Bar] to Foo for subscription expressions for superclasses + else if (expression instanceof PySubscriptionExpression) { + final PySubscriptionExpression subscriptionExpr = (PySubscriptionExpression)expression; + return subscriptionExpr.getOperand(); + } return expression; } @@ -237,27 +241,7 @@ public class PyClassImpl extends PyPresentableElementImpl implement @Nullable public String getQualifiedName() { - String name = getName(); - final PyClassStub stub = getStub(); - PsiElement ancestor = stub != null ? stub.getParentStub().getPsi() : getParent(); - while (!(ancestor instanceof PsiFile)) { - if (ancestor == null) return name; // can this happen? - if (ancestor instanceof PyClass) { - name = ((PyClass)ancestor).getName() + "." + name; - } - ancestor = stub != null ? ((StubBasedPsiElement)ancestor).getStub().getParentStub().getPsi() : ancestor.getParent(); - } - - PsiFile psiFile = ((PsiFile)ancestor).getOriginalFile(); - final PyFile builtins = PyBuiltinCache.getInstance(this).getBuiltinsFile(); - if (!psiFile.equals(builtins)) { - VirtualFile vFile = psiFile.getVirtualFile(); - if (vFile != null) { - final String packageName = QualifiedNameFinder.findShortestImportableName(this, vFile); - return packageName + "." + name; - } - } - return name; + return QualifiedNameFinder.getQualifiedName(this); } @Override diff --git a/python/src/com/jetbrains/python/psi/impl/PyFunctionImpl.java b/python/src/com/jetbrains/python/psi/impl/PyFunctionImpl.java index 2531b37bf8b2..86f84abb7fb6 100644 --- a/python/src/com/jetbrains/python/psi/impl/PyFunctionImpl.java +++ b/python/src/com/jetbrains/python/psi/impl/PyFunctionImpl.java @@ -184,11 +184,11 @@ public class PyFunctionImpl extends PyPresentableElementImpl imp } } if (context.maySwitchToAST(this) && LanguageLevel.forElement(this).isAtLeast(LanguageLevel.PYTHON30)) { - PyAnnotation anno = getAnnotation(); - if (anno != null) { - PyClass pyClass = anno.resolveToClass(); - if (pyClass != null) { - return new PyClassTypeImpl(pyClass, false); + final PyAnnotation annotation = getAnnotation(); + if (annotation != null) { + final PyType type = context.getType(annotation); + if (type != null) { + return type; } } } @@ -571,7 +571,7 @@ public class PyFunctionImpl extends PyPresentableElementImpl imp @Override public PyAnnotation getAnnotation() { - return findChildByClass(PyAnnotation.class); + return getStubOrPsiChild(PyElementTypes.ANNOTATION); } @NotNull @@ -699,21 +699,6 @@ public class PyFunctionImpl extends PyPresentableElementImpl imp @Nullable @Override public String getQualifiedName() { - String name = getName(); - if (name == null) { - return null; - } - PyClass containingClass = getContainingClass(); - if (containingClass != null) { - return containingClass.getQualifiedName() + "." + name; - } - if (PsiTreeUtil.getStubOrPsiParent(this) instanceof PyFile) { - VirtualFile virtualFile = getContainingFile().getVirtualFile(); - if (virtualFile != null) { - final String packageName = QualifiedNameFinder.findShortestImportableName(this, virtualFile); - return packageName + "." + name; - } - } - return null; + return QualifiedNameFinder.getQualifiedName(this); } } diff --git a/python/src/com/jetbrains/python/psi/impl/PyNamedParameterImpl.java b/python/src/com/jetbrains/python/psi/impl/PyNamedParameterImpl.java index bf792fef0dae..cc9178c07f9e 100644 --- a/python/src/com/jetbrains/python/psi/impl/PyNamedParameterImpl.java +++ b/python/src/com/jetbrains/python/psi/impl/PyNamedParameterImpl.java @@ -181,11 +181,11 @@ public class PyNamedParameterImpl extends PyPresentableElementImpl= 2) { - final PyNamedParameter param = parameters[1].getAsNamed(); - if (arg != null && param != null) { - final Map arguments = new LinkedHashMap(); - arguments.put(arg, param); - final AnalyzeCallResults results = new AnalyzeCallResults(callable, receiver, arguments); - if (firstResults == null) { - firstResults = results; - } - if (match(context.getType(param), context.getType(arg), context)) { - return results; - } + resolveResult = ref.multiResolve(false); + AnalyzeCallResults firstResults = null; + for (ResolveResult result : resolveResult) { + final PsiElement resolved = result.getElement(); + if (resolved instanceof PyTypedElement) { + final PyTypedElement typedElement = (PyTypedElement)resolved; + final PyType type = context.getType(typedElement); + if (!(type instanceof PyFunctionType)) { + return null; + } + final Callable callable = ((PyFunctionType)type).getCallable(); + final boolean isRight = PyNames.isRightOperatorName(typedElement.getName()); + final PyExpression arg = isRight ? expr.getLeftExpression() : expr.getRightExpression(); + final PyExpression receiver = isRight ? expr.getRightExpression() : expr.getLeftExpression(); + final PyParameter[] parameters = callable.getParameterList().getParameters(); + if (parameters.length >= 2) { + final PyNamedParameter param = parameters[1].getAsNamed(); + if (arg != null && param != null) { + final Map arguments = new LinkedHashMap(); + arguments.put(arg, param); + final AnalyzeCallResults results = new AnalyzeCallResults(callable, receiver, arguments); + if (firstResults == null) { + firstResults = results; + } + if (match(context.getType(param), context.getType(arg), context)) { + return results; } } } } - if (firstResults != null) { - return firstResults; - } + } + if (firstResults != null) { + return firstResults; } return null; } @@ -471,22 +469,20 @@ public class PyTypeChecker { public static AnalyzeCallResults analyzeCall(@NotNull PySubscriptionExpression expr, @NotNull TypeEvalContext context) { final PsiReference ref = expr.getReference(PyResolveContext.noImplicits().withTypeEvalContext(context)); final PsiElement resolved; - if (ref != null) { - resolved = ref.resolve(); - if (resolved instanceof PyTypedElement) { - final PyType type = context.getType((PyTypedElement)resolved); - if (type instanceof PyFunctionType) { - final Callable callable = ((PyFunctionType)type).getCallable(); - final PyParameter[] parameters = callable.getParameterList().getParameters(); - if (parameters.length == 2) { - final PyNamedParameter param = parameters[1].getAsNamed(); - if (param != null) { - final Map arguments = new LinkedHashMap(); - final PyExpression arg = expr.getIndexExpression(); - if (arg != null) { - arguments.put(arg, param); - return new AnalyzeCallResults(callable, expr.getOperand(), arguments); - } + resolved = ref.resolve(); + if (resolved instanceof PyTypedElement) { + final PyType type = context.getType((PyTypedElement)resolved); + if (type instanceof PyFunctionType) { + final Callable callable = ((PyFunctionType)type).getCallable(); + final PyParameter[] parameters = callable.getParameterList().getParameters(); + if (parameters.length == 2) { + final PyNamedParameter param = parameters[1].getAsNamed(); + if (param != null) { + final Map arguments = new LinkedHashMap(); + final PyExpression arg = expr.getIndexExpression(); + if (arg != null) { + arguments.put(arg, param); + return new AnalyzeCallResults(callable, expr.getOperand(), arguments); } } } diff --git a/python/testData/typing/typing.py b/python/testData/typing/typing.py new file mode 100644 index 000000000000..1fbd5fc8081b --- /dev/null +++ b/python/testData/typing/typing.py @@ -0,0 +1,582 @@ +"""Static type checking helpers""" + +from abc import ABCMeta, abstractmethod, abstractproperty +import inspect +import sys +import re + + +__all__ = [ + # Type system related + 'AbstractGeneric', + 'AbstractGenericMeta', + 'Any', + 'AnyStr', + 'Dict', + 'Function', + 'Generic', + 'GenericMeta', + 'IO', + 'List', + 'Match', + 'Pattern', + 'Protocol', + 'Set', + 'Tuple', + 'Undefined', + 'Union', + 'cast', + 'forwardref', + 'overload', + 'typevar', + # Protocols and abstract base classes + 'Container', + 'Iterable', + 'Iterator', + 'Sequence', + 'Sized', + 'AbstractSet', + 'Mapping', + 'BinaryIO', + 'TextIO', +] + + +def builtinclass(cls): + """Mark a class as a built-in/extension class for type checking.""" + return cls + + +def ducktype(type): + """Return a duck type declaration decorator. + + The decorator only affects type checking. + """ + def decorator(cls): + return cls + return decorator + + +def disjointclass(type): + """Return a disjoint class declaration decorator. + + The decorator only affects type checking. + """ + def decorator(cls): + return cls + return decorator + + +class GenericMeta(type): + """Metaclass for generic classes that support indexing by types.""" + + def __getitem__(self, args): + # Just ignore args; they are for compile-time checks only. + return self + + +class Generic(metaclass=GenericMeta): + """Base class for generic classes.""" + + +class AbstractGenericMeta(ABCMeta): + """Metaclass for abstract generic classes that support type indexing. + + This is used for both protocols and ordinary abstract classes. + """ + + def __new__(mcls, name, bases, namespace): + cls = super().__new__(mcls, name, bases, namespace) + # 'Protocol' must be an explicit base class in order for a class to + # be a protocol. + cls._is_protocol = name == 'Protocol' or Protocol in bases + return cls + + def __getitem__(self, args): + # Just ignore args; they are for compile-time checks only. + return self + + +class Protocol(metaclass=AbstractGenericMeta): + """Base class for protocol classes.""" + + @classmethod + def __subclasshook__(cls, c): + if not cls._is_protocol: + # No structural checks since this isn't a protocol. + return NotImplemented + + if cls is Protocol: + # Every class is a subclass of the empty protocol. + return True + + # Find all attributes defined in the protocol. + attrs = cls._get_protocol_attrs() + + for attr in attrs: + if not any(attr in d.__dict__ for d in c.__mro__): + return NotImplemented + return True + + @classmethod + def _get_protocol_attrs(cls): + # Get all Protocol base classes. + protocol_bases = [] + for c in cls.__mro__: + if getattr(c, '_is_protocol', False) and c.__name__ != 'Protocol': + protocol_bases.append(c) + + # Get attributes included in protocol. + attrs = set() + for base in protocol_bases: + for attr in base.__dict__.keys(): + # Include attributes not defined in any non-protocol bases. + for c in cls.__mro__: + if (c is not base and attr in c.__dict__ and + not getattr(c, '_is_protocol', False)): + break + else: + if (not attr.startswith('_abc_') and + attr != '__abstractmethods__' and + attr != '_is_protocol' and + attr != '__dict__' and + attr != '_get_protocol_attrs' and + attr != '__module__'): + attrs.add(attr) + + return attrs + + +class AbstractGeneric(metaclass=AbstractGenericMeta): + """Base class for abstract generic classes.""" + + +class TypeAlias: + """Class for defining generic aliases for library types.""" + + def __init__(self, target_type): + self.target_type = target_type + + def __getitem__(self, typeargs): + return self.target_type + + +Traceback = object() # TODO proper type object + + +# Define aliases for built-in types that support indexing. +List = TypeAlias(list) +Dict = TypeAlias(dict) +Set = TypeAlias(set) +Tuple = TypeAlias(tuple) +Function = TypeAlias(callable) +Pattern = TypeAlias(type(re.compile(''))) +Match = TypeAlias(type(re.match('', ''))) + +def union(x): return x + +Union = TypeAlias(union) + +class typevar: + def __init__(self, name, *, values=None): + self.name = name + self.values = values + + +# Predefined type variables. +AnyStr = typevar('AnyStr', values=(str, bytes)) + + +class forwardref: + def __init__(self, name): + self.name = name + + +def Any(x): + """The Any type; can also be used to cast a value to type Any.""" + return x + +def cast(type, object): + """Cast a value to a type. + + This only affects static checking; simply return object at runtime. + """ + return object + + +def overload(func): + """Function decorator for defining overloaded functions.""" + frame = sys._getframe(1) + locals = frame.f_locals + # See if there is a previous overload variant available. Also verify + # that the existing function really is overloaded: otherwise, replace + # the definition. The latter is actually important if we want to reload + # a library module such as genericpath with a custom one that uses + # overloading in the implementation. + if func.__name__ in locals and hasattr(locals[func.__name__], 'dispatch'): + orig_func = locals[func.__name__] + + def wrapper(*args, **kwargs): + ret, ok = orig_func.dispatch(*args, **kwargs) + if ok: + return ret + return func(*args, **kwargs) + wrapper.isoverload = True + wrapper.dispatch = make_dispatcher(func, orig_func.dispatch) + wrapper.next = orig_func + wrapper.__name__ = func.__name__ + if hasattr(func, '__isabstractmethod__'): + # Note that we can't reliably check that abstractmethod is + # used consistently across overload variants, so we let a + # static checker do it. + wrapper.__isabstractmethod__ = func.__isabstractmethod__ + return wrapper + else: + # Return the initial overload variant. + func.isoverload = True + func.dispatch = make_dispatcher(func) + func.next = None + return func + + +def is_erased_type(t): + return t is Any or isinstance(t, typevar) + + +def make_dispatcher(func, previous=None): + """Create argument dispatcher for an overloaded function. + + Also handle chaining of multiple overload variants. + """ + (args, varargs, varkw, defaults, + kwonlyargs, kwonlydefaults, annotations) = inspect.getfullargspec(func) + + argtypes = [] + for arg in args: + ann = annotations.get(arg) + if isinstance(ann, forwardref): + ann = ann.name + if is_erased_type(ann): + ann = None + elif isinstance(ann, str): + # The annotation is a string => evaluate it lazily when the + # overloaded function is first called. + frame = sys._getframe(2) + t = [None] + ann_str = ann + def check(x): + if not t[0]: + # Evaluate string in the context of the overload caller. + t[0] = eval(ann_str, frame.f_globals, frame.f_locals) + if is_erased_type(t[0]): + # Anything goes. + t[0] = object + if isinstance(t[0], type): + return isinstance(x, t[0]) + else: + return t[0](x) + ann = check + argtypes.append(ann) + + maxargs = len(argtypes) + minargs = maxargs + if defaults: + minargs = len(argtypes) - len(defaults) + + def dispatch(*args, **kwargs): + if previous: + ret, ok = previous(*args, **kwargs) + if ok: + return ret, ok + + nargs = len(args) + if nargs < minargs or nargs > maxargs: + # Invalid argument count. + return None, False + + for i in range(nargs): + argtype = argtypes[i] + if argtype: + if isinstance(argtype, type): + if not isinstance(args[i], argtype): + break + else: + if not argtype(args[i]): + break + else: + return func(*args, **kwargs), True + return None, False + return dispatch + + +class Undefined: + """Class that represents an undefined value with a specified type. + + At runtime the name Undefined is bound to an instance of this + class. The intent is that any operation on an Undefined object + raises an exception, including use in a boolean context. Some + operations cannot be disallowed: Undefined can be used as an + operand of 'is', and it can be assigned to variables and stored in + containers. + + 'Undefined' makes it possible to declare the static type of a + variable even if there is no useful default value to initialize it + with: + + from typing import Undefined + x = Undefined(int) + y = Undefined # type: int + + The latter form can be used if efficiency is of utmost importance, + since it saves a call operation and potentially additional + operations needed to evaluate a type expression. Undefined(x) + just evaluates to Undefined, ignoring the argument value. + """ + + def __repr__(self): + return '' + + def __setattr__(self, attr, value): + raise AttributeError("'Undefined' object has no attribute '%s'" % attr) + + def __eq__(self, other): + raise TypeError("'Undefined' object cannot be compared") + + def __call__(self, type): + return self + + def __bool__(self): + raise TypeError("'Undefined' object is not valid as a boolean") + + +Undefined = Undefined() + + +# Abstract classes + + +T = typevar('T') +KT = typevar('KT') +VT = typevar('VT') + + +class SupportsInt(Protocol): + @abstractmethod + def __int__(self) -> int: pass + + +class SupportsFloat(Protocol): + @abstractmethod + def __float__(self) -> float: pass + + +class SupportsAbs(Protocol[T]): + @abstractmethod + def __abs__(self) -> T: pass + + +class SupportsRound(Protocol[T]): + @abstractmethod + def __round__(self, ndigits: int = 0) -> T: pass + + +class Reversible(Protocol[T]): + @abstractmethod + def __reversed__(self) -> 'Iterator[T]': pass + + +class Sized(Protocol): + @abstractmethod + def __len__(self) -> int: pass + + +class Container(Protocol[T]): + @abstractmethod + def __contains__(self, x) -> bool: pass + + +class Iterable(Protocol[T]): + @abstractmethod + def __iter__(self) -> 'Iterator[T]': pass + + +class Iterator(Iterable[T], Protocol[T]): + @abstractmethod + def __next__(self) -> T: pass + + +class Sequence(Sized, Iterable[T], Container[T], AbstractGeneric[T]): + @abstractmethod + @overload + def __getitem__(self, i: int) -> T: pass + + @abstractmethod + @overload + def __getitem__(self, s: slice) -> 'Sequence[T]': pass + + @abstractmethod + def __reversed__(self, s: slice) -> Iterator[T]: pass + + @abstractmethod + def index(self, x) -> int: pass + + @abstractmethod + def count(self, x) -> int: pass + + +for t in list, tuple, str, bytes, range: + Sequence.register(t) + + +class AbstractSet(Sized, Iterable[T], AbstractGeneric[T]): + @abstractmethod + def __contains__(self, x: object) -> bool: pass + @abstractmethod + def __and__(self, s: 'AbstractSet[T]') -> 'AbstractSet[T]': pass + @abstractmethod + def __or__(self, s: 'AbstractSet[T]') -> 'AbstractSet[T]': pass + @abstractmethod + def __sub__(self, s: 'AbstractSet[T]') -> 'AbstractSet[T]': pass + @abstractmethod + def __xor__(self, s: 'AbstractSet[T]') -> 'AbstractSet[T]': pass + @abstractmethod + def isdisjoint(self, s: 'AbstractSet[T]') -> bool: pass + + +for t in set, frozenset, type({}.keys()), type({}.items()): + AbstractSet.register(t) + + +class Mapping(Sized, Iterable[KT], AbstractGeneric[KT, VT]): + @abstractmethod + def __getitem__(self, k: KT) -> VT: pass + @abstractmethod + def __setitem__(self, k: KT, v: VT) -> None: pass + @abstractmethod + def __delitem__(self, v: KT) -> None: pass + @abstractmethod + def __contains__(self, o: object) -> bool: pass + + @abstractmethod + def clear(self) -> None: pass + @abstractmethod + def copy(self) -> 'Mapping[KT, VT]': pass + @overload + @abstractmethod + def get(self, k: KT) -> VT: pass + @overload + @abstractmethod + def get(self, k: KT, default: VT) -> VT: pass + @overload + @abstractmethod + def pop(self, k: KT) -> VT: pass + @overload + @abstractmethod + def pop(self, k: KT, default: VT) -> VT: pass + @abstractmethod + def popitem(self) -> Tuple[KT, VT]: pass + @overload + @abstractmethod + def setdefault(self, k: KT) -> VT: pass + @overload + @abstractmethod + def setdefault(self, k: KT, default: VT) -> VT: pass + + @overload + @abstractmethod + def update(self, m: 'Mapping[KT, VT]') -> None: pass + @overload + @abstractmethod + def update(self, m: Iterable[Tuple[KT, VT]]) -> None: pass + + @abstractmethod + def keys(self) -> AbstractSet[KT]: pass + @abstractmethod + def values(self) -> AbstractSet[VT]: pass + @abstractmethod + def items(self) -> AbstractSet[Tuple[KT, VT]]: pass + + +# TODO Consider more types: os.environ, etc. However, these add dependencies. +Mapping.register(dict) + + +# Note that the BinaryIO and TextIO classes must be in sync with typing module +# stubs. + + +class IO(AbstractGeneric[AnyStr]): + @abstractproperty + def mode(self) -> str: pass + @abstractproperty + def name(self) -> str: pass + @abstractmethod + def close(self) -> None: pass + @abstractmethod + def closed(self) -> bool: pass + @abstractmethod + def fileno(self) -> int: pass + @abstractmethod + def flush(self) -> None: pass + @abstractmethod + def isatty(self) -> bool: pass + @abstractmethod + def read(self, n: int = -1) -> AnyStr: pass + @abstractmethod + def readable(self) -> bool: pass + @abstractmethod + def readline(self, limit: int = -1) -> AnyStr: pass + @abstractmethod + def readlines(self, hint: int = -1) -> List[AnyStr]: pass + @abstractmethod + def seek(self, offset: int, whence: int = 0) -> int: pass + @abstractmethod + def seekable(self) -> bool: pass + @abstractmethod + def tell(self) -> int: pass + @abstractmethod + def truncate(self, size: int = None) -> int: pass + @abstractmethod + def writable(self) -> bool: pass + @abstractmethod + def write(self, s: AnyStr) -> int: pass + @abstractmethod + def writelines(self, lines: List[AnyStr]) -> None: pass + + @abstractmethod + def __enter__(self) -> 'IO[AnyStr]': pass + @abstractmethod + def __exit__(self, type, value, traceback) -> None: pass + + +class BinaryIO(IO[bytes]): + @overload + @abstractmethod + def write(self, s: bytes) -> int: pass + @overload + @abstractmethod + def write(self, s: bytearray) -> int: pass + + @abstractmethod + def __enter__(self) -> 'BinaryIO': pass + + +class TextIO(IO[str]): + @abstractproperty + def buffer(self) -> BinaryIO: pass + @abstractproperty + def encoding(self) -> str: pass + @abstractproperty + def errors(self) -> str: pass + @abstractproperty + def line_buffering(self) -> bool: pass + @abstractproperty + def newlines(self) -> Any: pass + @abstractmethod + def __enter__(self) -> 'TextIO': pass + + +# TODO Register IO/TextIO/BinaryIO as the base class of file-like types. + + +del t diff --git a/python/testSrc/com/jetbrains/python/PyTypeTest.java b/python/testSrc/com/jetbrains/python/PyTypeTest.java index b0e41bcffddc..b6ed1090edad 100644 --- a/python/testSrc/com/jetbrains/python/PyTypeTest.java +++ b/python/testSrc/com/jetbrains/python/PyTypeTest.java @@ -835,6 +835,37 @@ public class PyTypeTest extends PyTestCase { "expr = (1,) + (True, 'spam') + ()"); } + public void testConstructorUnification() { + doTest("C[int]", + "class C(object):\n" + + " def __init__(self, x):\n" + + " '''\n" + + " :type x: T\n" + + " :rtype: C[T]\n" + + " '''\n" + + " pass\n" + + "\n" + + "expr = C(10)\n"); + } + + public void testGenericClassMethodUnification() { + doTest("int", + "class C(object):\n" + + " def __init__(self, x):\n" + + " '''\n" + + " :type x: T\n" + + " :rtype: C[T]\n" + + " '''\n" + + " pass\n" + + " def foo(self):\n" + + " '''\n" + + " :rtype: T\n" + + " '''\n" + + " pass\n" + + "\n" + + "expr = C(10).foo()\n"); + } + private static TypeEvalContext getTypeEvalContext(@NotNull PyExpression element) { return TypeEvalContext.userInitiated(element.getContainingFile()).withTracing(); } diff --git a/python/testSrc/com/jetbrains/python/PyTypingTest.java b/python/testSrc/com/jetbrains/python/PyTypingTest.java new file mode 100644 index 000000000000..9513795cb5b7 --- /dev/null +++ b/python/testSrc/com/jetbrains/python/PyTypingTest.java @@ -0,0 +1,267 @@ +/* + * Copyright 2000-2014 JetBrains s.r.o. + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package com.jetbrains.python; + +import com.intellij.testFramework.LightProjectDescriptor; +import com.jetbrains.python.documentation.PythonDocumentationProvider; +import com.jetbrains.python.fixtures.PyTestCase; +import com.jetbrains.python.psi.LanguageLevel; +import com.jetbrains.python.psi.PyExpression; +import com.jetbrains.python.psi.types.PyType; +import com.jetbrains.python.psi.types.TypeEvalContext; +import org.jetbrains.annotations.NotNull; +import org.jetbrains.annotations.Nullable; + +/** + * Tests for a type system based on mypy's typing module. + * + * @author vlan + */ +public class PyTypingTest extends PyTestCase { + @Nullable + @Override + protected LightProjectDescriptor getProjectDescriptor() { + return ourPy3Descriptor; + } + + @Override + public void setUp() throws Exception { + super.setUp(); + setLanguageLevel(LanguageLevel.PYTHON32); + } + + @Override + public void tearDown() throws Exception { + setLanguageLevel(null); + super.tearDown(); + } + + public void testClassType() { + doTest("Foo", + "class Foo:" + + " pass\n" + + "\n" + + "def f(expr: Foo):\n" + + " pass\n"); + } + + public void testClassReturnType() { + doTest("Foo", + "class Foo:" + + " pass\n" + + "\n" + + "def f() -> Foo:\n" + + " pass\n" + + "\n" + + "expr = f()\n"); + + } + + public void testNoneType() { + doTest("None", + "def f(expr: None):\n" + + " pass\n"); + } + + public void testNoneReturnType() { + doTest("None", + "def f() -> None:\n" + + " return 0\n" + + "expr = f()\n"); + } + + public void testUnionType() { + doTest("int | str", + "from typing import Union\n" + + "\n" + + "def f(expr: Union[int, str]):\n" + + " pass\n"); + } + + public void testBuiltinList() { + doTest("list", + "from typing import List\n" + + "\n" + + "def f(expr: List):\n" + + " pass\n"); + } + + public void testBuiltinListWithParameter() { + doTest("list[int]", + "from typing import List\n" + + "\n" + + "def f(expr: List[int]):\n" + + " pass\n"); + } + + public void testBuiltinDictWithParameters() { + doTest("dict[str, int]", + "from typing import Dict\n" + + "\n" + + "def f(expr: Dict[str, int]):\n" + + " pass\n"); + } + + public void testBuiltinTuple() { + doTest("tuple", + "from typing import Tuple\n" + + "\n" + + "def f(expr: Tuple):\n" + + " pass\n"); + } + + public void testBuiltinTupleWithParameters() { + doTest("(int, str)", + "from typing import Tuple\n" + + "\n" + + "def f(expr: Tuple[int, str]):\n" + + " pass\n"); + } + + public void testAnyType() { + doTest("unknown", + "from typing import Any\n" + + "\n" + + "def f(expr: Any):\n" + + " pass\n"); + } + + public void testGenericType() { + doTest("A", + "from typing import typevar\n" + + "\n" + + "T = typevar('A')\n" + + "\n" + + "def f(expr: T):\n" + + " pass\n"); + } + + public void testGenericBoundedType() { + doTest("T <= int | str", + "from typing import typevar\n" + + "\n" + + "T = typevar('T', values=(int, str))\n" + + "\n" + + "def f(expr: T):\n" + + " pass\n"); + } + + public void testParameterizedClass() { + doTest("C[int]", + "from typing import Generic, typevar\n" + + "\n" + + "T = typevar('T')\n" + + "\n" + + "class C(Generic[T]):\n" + + " def __init__(self, x: T):\n" + + " pass\n" + + "\n" + + "expr = C(10)\n"); + } + + public void testParameterizedClassMethod() { + doTest("int", + "from typing import Generic, typevar\n" + + "\n" + + "T = typevar('T')\n" + + "\n" + + "class C(Generic[T]):\n" + + " def __init__(self, x: T):\n" + + " pass\n" + + " def foo(self) -> T:\n" + + " pass\n" + + "\n" + + "expr = C(10).foo()\n"); + } + + public void testParameterizedClassInheritance() { + doTest("int", + "from typing import Generic, typevar\n" + + "\n" + + "T = typevar('T')\n" + + "\n" + + "class B(Generic[T]):\n" + + " def foo(self) -> T:\n" + + " pass\n" + + "class C(B[T]):\n" + + " def __init__(self, x: T):\n" + + " pass\n" + + "\n" + + "expr = C(10).foo()\n"); + } + + public void testAnyStrUnification() { + doTest("bytes", + "from typing import AnyStr\n" + + "\n" + + "def foo(x: AnyStr) -> AnyStr:\n" + + " pass\n" + + "\n" + + "expr = foo(b'bar')\n"); + } + + public void testAnyStrForUnknown() { + doTest("str | bytes", + "from typing import AnyStr\n" + + "\n" + + "def foo(x: AnyStr) -> AnyStr:\n" + + " pass\n" + + "\n" + + "def bar(x):\n" + + " expr = foo(x)\n"); + } + + public void testFunctionType() { + doTest("(int, str) -> str", + "from typing import Function\n" + + "\n" + + "def foo(expr: Function[[int, str], str]):\n" + + " pass\n"); + } + + public void testTypeInStringLiteral() { + doTest("C", + "class C:\n" + + " def foo(self, expr: 'C'):\n" + + " pass\n"); + } + + public void testQualifiedTypeInStringLiteral() { + doTest("str", + "import typing\n" + + "\n" + + "def foo(x: 'typing.AnyStr') -> typing.AnyStr:\n" + + " pass\n" + + "\n" + + "expr = foo('bar')\n"); + } + + private void doTest(@NotNull String expectedType, @NotNull String text) { + myFixture.copyDirectoryToProject("typing", ""); + myFixture.configureByText(PythonFileType.INSTANCE, text); + final PyExpression expr = myFixture.findElementByText("expr", PyExpression.class); + final TypeEvalContext codeAnalysis = TypeEvalContext.codeAnalysis(expr.getContainingFile()); + final TypeEvalContext userInitiated = TypeEvalContext.userInitiated(expr.getContainingFile()).withTracing(); + assertType(expectedType, expr, codeAnalysis, "code analysis"); + assertType(expectedType, expr, userInitiated, "user initiated"); + } + + private static void assertType(String expectedType, PyExpression expr, TypeEvalContext context, String contextName) { + final PyType actual = context.getType(expr); + final String actualType = PythonDocumentationProvider.getTypeName(actual, context); + assertEquals("Failed in " + contextName + " context", expectedType, actualType); + } +}