Fixed code insight for returning 'self' in base class methods (PY-10977, PY-11413)

If a method of a base class returns a value of this class, then it is
cast to a derived class this method has been invoked on. This cast is
considered safe, since it cannot result in false positives.
This commit is contained in:
Andrey Vlasovskikh
2014-01-16 16:49:05 +04:00
parent b56b6cc01b
commit e2db1eaa5f
5 changed files with 72 additions and 18 deletions
@@ -174,29 +174,28 @@ public class PyFunctionImpl extends PyPresentableElementImpl<PyFunctionStub> imp
@Nullable
@Override
public PyType getReturnType(@NotNull TypeEvalContext context, @Nullable PyQualifiedExpression callSite) {
final PyType type = getGenericReturnType(context, callSite);
PyType type = getGenericReturnType(context, callSite);
if (callSite == null) {
return type;
}
final PyTypeChecker.AnalyzeCallResults results = PyTypeChecker.analyzeCallSite(callSite, context);
if (PyTypeChecker.hasGenerics(type, context)) {
if (results != null) {
final Map<PyGenericType, PyType> substitutions = PyTypeChecker.unifyGenericCall(this, results.getReceiver(), results.getArguments(),
context);
if (substitutions != null) {
return PyTypeChecker.substitute(type, substitutions, context);
}
type = substitutions != null ? PyTypeChecker.substitute(type, substitutions, context) : null;
}
return null;
else {
type = null;
}
}
if (results != null) {
type = replaceSelf(type, results.getReceiver(), context);
}
if (results != null && isDynamicallyEvaluated(results.getArguments().values(), context)) {
return PyUnionType.createWeakType(type);
}
else {
return type;
}
return type;
}
@Nullable
@@ -205,16 +204,36 @@ public class PyFunctionImpl extends PyPresentableElementImpl<PyFunctionStub> imp
*/
public PyType getReturnTypeWithoutCallSite(@NotNull TypeEvalContext context,
@Nullable PyExpression receiver) {
final PyType type = getGenericReturnType(context, null);
PyType type = getGenericReturnType(context, null);
if (PyTypeChecker.hasGenerics(type, context)) {
final Map<PyGenericType, PyType> substitutions =
PyTypeChecker.unifyGenericCall(this, receiver, Maps.<PyExpression, PyNamedParameter>newHashMap(), context);
final Map<PyGenericType, PyType> substitutions = PyTypeChecker.unifyGenericCall(this, receiver,
Maps.<PyExpression, PyNamedParameter>newHashMap(),
context);
if (substitutions != null) {
return PyTypeChecker.substitute(type, substitutions, context);
type = PyTypeChecker.substitute(type, substitutions, context);
}
else {
type = null;
}
return null;
}
return type;
return replaceSelf(type, receiver, context);
}
@Nullable
private PyType replaceSelf(@Nullable PyType returnType, @Nullable PyExpression receiver, @NotNull TypeEvalContext context) {
if (receiver != null) {
// TODO: Currently we substitute only simple subclass types, but we could handle union and collection types as well
if (returnType instanceof PyClassType) {
final PyClassType returnClassType = (PyClassType)returnType;
if (returnClassType.getPyClass() == getContainingClass()) {
final PyType receiverType = context.getType(receiver);
if (receiverType instanceof PyClassType && PyTypeChecker.match(returnType, receiverType, context)) {
return receiverType;
}
}
}
}
return returnType;
}
private static boolean isDynamicallyEvaluated(@NotNull Collection<PyNamedParameter> parameters, @NotNull TypeEvalContext context) {
@@ -241,8 +241,8 @@ public class PyTargetExpressionImpl extends PyPresentableElementImpl<PyTargetExp
if (exprType instanceof PyClassType) {
final PyClass cls = ((PyClassType)exprType).getPyClass();
final PyFunction enter = cls.findMethodByName(PyNames.ENTER, true);
if (enter != null) {
final PyType enterType = enter.getReturnType(context, null);
if (enter instanceof PyFunctionImpl) {
final PyType enterType = ((PyFunctionImpl)enter).getReturnTypeWithoutCallSite(context, expression);
if (enterType != null) {
return enterType;
}
@@ -0,0 +1,12 @@
class C(object):
def __enter__(self):
return self
class D(C):
def foo(self):
pass
with D() as cm:
cm.foo() # pass
@@ -0,0 +1,13 @@
class C(object):
def get_self(self):
return self
class D(C):
def foo(self):
pass
d = D()
print(d.foo())
print(d.get_self().foo()) # pass
@@ -326,6 +326,16 @@ public class PyUnresolvedReferencesInspectionTest extends PyTestCase {
doTest();
}
// PY-10977
public void testContextManagerSubclass() {
doTest();
}
// PY-11413
public void testReturnSelfInSuperClass() {
doTest();
}
private void doTest() {
myFixture.configureByFile(TEST_DIRECTORY + getTestName(true) + ".py");
myFixture.enableInspections(PyUnresolvedReferencesInspection.class);