provide correct type for Python 3 no-args super call (PY-1330)

This commit is contained in:
Dmitry Jemerov
2010-07-19 16:40:25 +04:00
parent c0aca39007
commit f343a3bfd2
4 changed files with 41 additions and 8 deletions
@@ -138,6 +138,7 @@ public class PyCallExpressionImpl extends PyElementImpl implements PyCallExpress
if (must_be_super == PyBuiltinCache.getInstance(this).getClass(PyNames.SUPER)) {
PyArgumentList arglist = getArgumentList();
if (arglist != null) {
final PyClass containingClass = PsiTreeUtil.getParentOfType(this, PyClass.class);
PyExpression[] args = arglist.getArguments();
if (args.length > 1) {
PyExpression first_arg = args[0];
@@ -150,7 +151,7 @@ public class PyCallExpressionImpl extends PyElementImpl implements PyCallExpress
if (element instanceof PyParameter) {
final PyParameterList parameterList = PsiTreeUtil.getParentOfType(element, PyParameterList.class);
if (parameterList != null && element == parameterList.getParameters() [0]) {
return getSuperCallType(context, PsiTreeUtil.getParentOfType(this, PyClass.class), args[1]);
return getSuperCallType(context, containingClass, args[1]);
}
}
}
@@ -161,6 +162,9 @@ public class PyCallExpressionImpl extends PyElementImpl implements PyCallExpress
}
}
}
else if (((PyFile)getContainingFile()).getLanguageLevel().isPy3K()) {
return getFirstSuperClassType(containingClass);
}
}
}
}
@@ -177,10 +181,7 @@ public class PyCallExpressionImpl extends PyElementImpl implements PyCallExpress
PyClass second_class = ((PyClassType)second_type).getPyClass();
assert second_class != null;
if (first_class == second_class) {
final PyClass[] supers = first_class.getSuperClasses();
if (supers.length > 0) {
return new PyClassType(supers[0], false);
}
return getFirstSuperClassType(first_class);
}
if (second_class.isSubclass(first_class)) {
// TODO: super(Foo, Bar) is a superclass of Foo directly preceding Bar in MRO
@@ -190,4 +191,13 @@ public class PyCallExpressionImpl extends PyElementImpl implements PyCallExpress
}
return null;
}
private static PyType getFirstSuperClassType(PyClass first_class) {
// TODO handle __mro__ here
final PyClass[] supers = first_class.getSuperClasses();
if (supers.length > 0) {
return new PyClassType(supers[0], false);
}
return null;
}
}
+10
View File
@@ -0,0 +1,10 @@
class A(object):
def foo(self):
print "foo"
class B(A):
def foo(self):
super().foo()
# <ref>
B().foo()
@@ -9,6 +9,7 @@ import com.intellij.psi.util.PsiTreeUtil;
import com.jetbrains.python.fixtures.PyResolveTestCase;
import com.jetbrains.python.psi.*;
import com.jetbrains.python.psi.impl.PyBuiltinCache;
import com.jetbrains.python.psi.impl.PythonLanguageLevelPusher;
import com.jetbrains.python.psi.resolve.ImportedResolveResult;
public class PyResolveTest extends PyResolveTestCase {
@@ -250,6 +251,17 @@ public class PyResolveTest extends PyResolveTestCase {
assertEquals("A", ((PyFunction) targetElement).getContainingClass().getName());
}
public void testSuperPy3k() { // PY-1330
PythonLanguageLevelPusher.setForcedLanguageLevel(myFixture.getProject(), LanguageLevel.PYTHON30);
try {
final PyFunction pyFunction = assertResolvesTo(PyFunction.class, "foo");
assertEquals("A", pyFunction.getContainingClass().getName());
}
finally {
PythonLanguageLevelPusher.setForcedLanguageLevel(myFixture.getProject(), null);
}
}
public void testStackOverflow() {
PsiElement targetElement = resolve();
assertNull(targetElement);
@@ -47,11 +47,11 @@ public abstract class PyResolveTestCase extends PyLightFixtureTestCase {
protected abstract PsiElement doResolve() throws Exception;
protected <T extends PsiElement> void assertResolvesTo(final Class<T> aClass, final String name) {
assertResolvesTo(aClass, name, null);
protected <T extends PsiElement> T assertResolvesTo(final Class<T> aClass, final String name) {
return assertResolvesTo(aClass, name, null);
}
protected <T extends PsiElement> void assertResolvesTo(final Class<T> aClass,
protected <T extends PsiElement> T assertResolvesTo(final Class<T> aClass,
final String name,
String containingFilePath) {
final PsiElement element;
@@ -66,6 +66,7 @@ public abstract class PyResolveTestCase extends PyLightFixtureTestCase {
if (containingFilePath != null) {
assertEquals(containingFilePath, element.getContainingFile().getVirtualFile().getPath());
}
return (T)element;
}
protected int findMarkerOffset(final PsiFile psiFile) {