Generic and overloaded types for Python operators

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