diff --git a/python/src/com/jetbrains/python/codeInsight/stdlib/PyDataclassTypeProvider.kt b/python/src/com/jetbrains/python/codeInsight/stdlib/PyDataclassTypeProvider.kt index 62b4a77a2817..21152e214476 100644 --- a/python/src/com/jetbrains/python/codeInsight/stdlib/PyDataclassTypeProvider.kt +++ b/python/src/com/jetbrains/python/codeInsight/stdlib/PyDataclassTypeProvider.kt @@ -117,6 +117,7 @@ class PyDataclassTypeProvider : PyTypeProviderBase() { var seenInit = false val keywordOnly = linkedSetOf() var seenKeywordOnlyClass = false + val seenNames = mutableSetOf() 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? { + 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, diff --git a/python/testData/paramInfo/InitializingDataclassHierarchy/a.py b/python/testData/paramInfo/InitializingDataclassHierarchy/a.py index ece4da7dbfc5..207b6b6b480f 100644 --- a/python/testData/paramInfo/InitializingDataclassHierarchy/a.py +++ b/python/testData/paramInfo/InitializingDataclassHierarchy/a.py @@ -41,17 +41,4 @@ class A4: class B4(A4): b: str -B4() - - -@dataclass -class A5: - x: Any = 15.0 - y: int = 0 - -@dataclass -class B5(A5): - z: int = 10 - x: int = 15 - -B5() \ No newline at end of file +B4() \ No newline at end of file diff --git a/python/testData/paramInfo/InitializingDataclassOverridingField/a.py b/python/testData/paramInfo/InitializingDataclassOverridingField/a.py new file mode 100644 index 000000000000..c5542b247a21 --- /dev/null +++ b/python/testData/paramInfo/InitializingDataclassOverridingField/a.py @@ -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() + + +@dataclass +class A2: + a: int + aa: int + +@dataclass +class B2(A2): + aa: int = field(init=False) + + def __post_init__(self): + self.aa = 44 + +B2() \ No newline at end of file diff --git a/python/testData/paramInfo/InitializingDataclassOverridingField/dataclasses.py b/python/testData/paramInfo/InitializingDataclassOverridingField/dataclasses.py new file mode 100644 index 000000000000..44ee41d023c3 --- /dev/null +++ b/python/testData/paramInfo/InitializingDataclassOverridingField/dataclasses.py @@ -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 diff --git a/python/testSrc/com/jetbrains/python/PyParameterInfoTest.java b/python/testSrc/com/jetbrains/python/PyParameterInfoTest.java index 669316ee0b47..0094d99580fa 100644 --- a/python/testSrc/com/jetbrains/python/PyParameterInfoTest.java +++ b/python/testSrc/com/jetbrains/python/PyParameterInfoTest.java @@ -755,7 +755,7 @@ public class PyParameterInfoTest extends LightMarkedTestCase { runWithLanguageLevel( LanguageLevel.PYTHON37, () -> { - final Map marks = loadMultiFileTest(5); + final Map marks = loadMultiFileTest(4); feignCtrlP(marks.get("").getTextOffset()).check("a: int, b: str", new String[]{"a: int, "}); feignCtrlP(marks.get("").getTextOffset()).check("a: int, b: str", new String[]{"a: int, "}); @@ -764,8 +764,6 @@ public class PyParameterInfoTest extends LightMarkedTestCase { feignCtrlP(marks.get("").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("").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 marks = loadMultiFileTest(2); + + feignCtrlP(marks.get("").getTextOffset()).check("x: int=15, y: int=0, z: int=10", new String[]{"x: int=15, "}); + feignCtrlP(marks.get("").getTextOffset()).check("a: int", new String[]{"a: int"}); + } + ); + } + // PY-26354 public void testInitializingAttrsUsingPep526() { runWithLanguageLevel(