Don't add an ancestor's field as a parameter if it was overridden with init=False (PY-35548)

GitOrigin-RevId: 75b1a1c3d1852b98a58d066d6380f04e5dd013bd
This commit is contained in:
Semyon Proshev
2019-12-05 14:31:31 +00:00
committed by intellij-monorepo-bot
parent 455acf1999
commit 64470fd40d
5 changed files with 80 additions and 38 deletions
@@ -117,6 +117,7 @@ class PyDataclassTypeProvider : PyTypeProviderBase() {
var seenInit = false
val keywordOnly = linkedSetOf<String>()
var seenKeywordOnlyClass = false
val seenNames = mutableSetOf<String>()
for (currentType in StreamEx.of(clsType).append(cls.getAncestorTypes(context))) {
if (currentType == null ||
@@ -140,22 +141,24 @@ class PyDataclassTypeProvider : PyTypeProviderBase() {
.asSequence()
.filterNot { PyTypingTypeProvider.isClassVar(it, context) }
.mapNotNull { fieldToParameter(current, it, parameters.type, ellipsis, context) }
.forEach { parameter ->
parameter.name?.let {
// note: attributes are visited from inheritors to ancestors, in reversed order for every of them
.filterNot { it.first in seenNames }
.forEach { (name, parameter) ->
// note: attributes are visited from inheritors to ancestors, in reversed order for every of them
if (seenKeywordOnlyClass && it !in collected) {
keywordOnly += it
}
if (seenKeywordOnlyClass && name !in collected) {
keywordOnly += name
}
if (parameters.type == PyDataclassParameters.Type.STD) {
// std: attribute that overrides ancestor's attribute does not change the order but updates type
collected[it] = collected.remove(it) ?: parameter
}
else if (!collected.containsKey(it)) {
// attrs: attribute that overrides ancestor's attribute changes the order
collected[it] = parameter
}
if (parameter == null) {
seenNames.add(name)
}
else if (parameters.type == PyDataclassParameters.Type.STD) {
// std: attribute that overrides ancestor's attribute does not change the order but updates type
collected[name] = collected.remove(name) ?: parameter
}
else if (!collected.containsKey(name)) {
// attrs: attribute that overrides ancestor's attribute changes the order
collected[name] = parameter
}
}
}
@@ -189,20 +192,24 @@ class PyDataclassTypeProvider : PyTypeProviderBase() {
field: PyTargetExpression,
dataclassType: PyDataclassParameters.Type,
ellipsis: PyNoneLiteralExpression,
context: TypeEvalContext): PyCallableParameter? {
context: TypeEvalContext): Pair<String, PyCallableParameter?>? {
val fieldName = field.name ?: return null
val stub = field.stub
val fieldStub = if (stub == null) PyDataclassFieldStubImpl.create(field) else stub.getCustomStub(PyDataclassFieldStub::class.java)
if (fieldStub != null && !fieldStub.initValue()) return null
if (fieldStub != null && !fieldStub.initValue()) return fieldName to null
if (fieldStub == null && field.annotationValue == null) return null // skip fields that are not annotated
val name =
field.name?.let {
val parameterName =
fieldName.let {
if (dataclassType == PyDataclassParameters.Type.ATTRS && PyUtil.getInitialUnderscores(it) == 1) it.substring(1) else it
}
return PyCallableParameterImpl.nonPsi(name,
getTypeForParameter(cls, field, dataclassType, context),
getDefaultValueForParameter(cls, field, fieldStub, dataclassType, ellipsis, context))
val parameter = PyCallableParameterImpl.nonPsi(parameterName,
getTypeForParameter(cls, field, dataclassType, context),
getDefaultValueForParameter(cls, field, fieldStub, dataclassType, ellipsis, context))
return parameterName to parameter
}
private fun getTypeForParameter(cls: PyClass,
@@ -41,17 +41,4 @@ class A4:
class B4(A4):
b: str
B4(<arg4>)
@dataclass
class A5:
x: Any = 15.0
y: int = 0
@dataclass
class B5(A5):
z: int = 10
x: int = 15
B5(<arg5>)
B4(<arg4>)
@@ -0,0 +1,29 @@
from dataclasses import dataclass, field
@dataclass
class A1:
x: Any = 15.0
y: int = 0
@dataclass
class B1(A1):
z: int = 10
x: int = 15
B1(<arg1>)
@dataclass
class A2:
a: int
aa: int
@dataclass
class B2(A2):
aa: int = field(init=False)
def __post_init__(self):
self.aa = 44
B2(<arg2>)
@@ -0,0 +1,8 @@
def dataclass(_cls=None, *, init=True, repr=True, eq=True, order=False,
unsafe_hash=False, frozen=False):
pass
def field(*, default=MISSING, default_factory=MISSING, init=True, repr=True,
hash=None, compare=True, metadata=None):
pass
@@ -755,7 +755,7 @@ public class PyParameterInfoTest extends LightMarkedTestCase {
runWithLanguageLevel(
LanguageLevel.PYTHON37,
() -> {
final Map<String, PsiElement> marks = loadMultiFileTest(5);
final Map<String, PsiElement> marks = loadMultiFileTest(4);
feignCtrlP(marks.get("<arg1>").getTextOffset()).check("a: int, b: str", new String[]{"a: int, "});
feignCtrlP(marks.get("<arg2>").getTextOffset()).check("a: int, b: str", new String[]{"a: int, "});
@@ -764,8 +764,6 @@ public class PyParameterInfoTest extends LightMarkedTestCase {
feignCtrlP(marks.get("<arg4>").getTextOffset()).check(Arrays.asList("self: object", "cls: object"),
Arrays.asList(ArrayUtilRt.EMPTY_STRING_ARRAY, ArrayUtilRt.EMPTY_STRING_ARRAY),
Arrays.asList(new String[]{"self: object"}, new String[]{"cls: object"}));
feignCtrlP(marks.get("<arg5>").getTextOffset()).check("x: int=15, y: int=0, z: int=10", new String[]{"x: int=15, "});
}
);
}
@@ -790,6 +788,19 @@ public class PyParameterInfoTest extends LightMarkedTestCase {
);
}
// PY-28506, PY-31762, PY-35548
public void testInitializingDataclassOverridingField() {
runWithLanguageLevel(
LanguageLevel.getLatest(),
() -> {
final Map<String, PsiElement> marks = loadMultiFileTest(2);
feignCtrlP(marks.get("<arg1>").getTextOffset()).check("x: int=15, y: int=0, z: int=10", new String[]{"x: int=15, "});
feignCtrlP(marks.get("<arg2>").getTextOffset()).check("a: int", new String[]{"a: int"});
}
);
}
// PY-26354
public void testInitializingAttrsUsingPep526() {
runWithLanguageLevel(