Fixed type evaluation for right Python operators, including 'str.__rmul__'

This commit is contained in:
Andrey Vlasovskikh
2011-12-21 20:05:56 +04:00
parent 04d26eae08
commit 7b16af897a
4 changed files with 99 additions and 86 deletions
@@ -3,7 +3,6 @@ package com.jetbrains.python.psi.impl;
import com.intellij.lang.ASTNode;
import com.intellij.psi.PsiElement;
import com.intellij.psi.PsiPolyVariantReference;
import com.intellij.psi.PsiReference;
import com.intellij.psi.tree.IElementType;
import com.intellij.psi.util.PsiTreeUtil;
import com.intellij.util.IncorrectOperationException;
@@ -109,27 +108,16 @@ public class PyBinaryExpressionImpl extends PyElementImpl implements PyBinaryExp
}
public PyType getType(@NotNull TypeEvalContext context) {
final PyExpression lhs = getLeftExpression();
final PyExpression rhs = getRightExpression();
final PsiElement operator = getPsiOperator();
if (lhs != null && rhs != null && operator != null) {
final PsiReference ref = getReference(PyResolveContext.noImplicits().withTypeEvalContext(context));
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();
final PyTypeChecker.AnalyzeCallResults results = PyTypeChecker.analyzeCall(this, context);
if (results != null) {
final PyType type = results.getFunction().getReturnType(context, this);
if (!PyTypeChecker.isUnknown(type) && !(type instanceof PyNoneType)) {
return type;
}
}
if (PyNames.COMPARISON_OPERATORS.contains(getReferencedName())) {
return PyBuiltinCache.getInstance(this).getBoolType();
}
return null;
}
@@ -295,82 +295,97 @@ public class PyTypeChecker {
}
}
@Nullable
public static AnalyzeCallResults analyzeCall(@NotNull PyCallExpression call, @NotNull TypeEvalContext context) {
final PyArgumentList args = call.getArgumentList();
if (args != null) {
final CallArgumentsMapping mapping = args.analyzeCall(PyResolveContext.noImplicits().withTypeEvalContext(context));
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) {
final PyExpression callee = call.getCallee();
final PyExpression receiver = callee instanceof PyQualifiedExpression ? ((PyQualifiedExpression)callee).getQualifier() : null;
return new AnalyzeCallResults((PyFunction)callable, receiver, arguments);
}
}
}
return null;
}
@Nullable
public static AnalyzeCallResults analyzeCall(@NotNull PyBinaryExpression expr, @NotNull TypeEvalContext context) {
final PsiPolyVariantReference ref = expr.getReference(PyResolveContext.noImplicits().withTypeEvalContext(context));
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;
}
}
return null;
}
@Nullable
public static AnalyzeCallResults analyzeCall(@NotNull PySubscriptionExpression expr, @NotNull TypeEvalContext context) {
final PsiReference ref = expr.getReference(PyResolveContext.noImplicits().withTypeEvalContext(context));
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;
}
@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);
}
}
}
return analyzeCall((PyCallExpression)parent, context);
}
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;
}
}
return analyzeCall((PyBinaryExpression)callSite, context);
}
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 analyzeCall((PySubscriptionExpression)callSite, context);
}
return null;
}
@@ -34,7 +34,7 @@ i = "test"
'%(name)s' % {'name': i} #ok
'%s' % i #ok
'%f' % i #Unexpected type
'%f' % (2 * i + 5) #ok
'%f' % (2 * 3 + 5) #ok
s = "%s" % "a".upper() #ok
x = ['a', 'b', 'c']
print "%d: %s" % (len(x), ", ".join(x)) #ok
@@ -342,6 +342,16 @@ def test_right_operators():
]
def test_string_integer():
print('foo' + 'bar')
print(2 + 3)
print('foo' + <warning descr="Expected type 'one of (str, unicode)', got 'int' instead">3</warning>)
print(3 + <warning descr="Expected type 'one of (int, long, float, complex)', got 'str' instead">'foo'</warning>)
print('foo' + 'bar' * 3)
print('foo' + 3 * 'bar')
print('foo' + <warning descr="Expected type 'one of (str, unicode)', got 'int' instead">2 * 3</warning>)
def test_isinstance_implicit_self_types():
x = 1
if isinstance(x, unicode):