mirror of
https://gitflic.ru/project/openide/openide.git
synced 2026-09-27 10:03:11 +07:00
Provide special type for replace usages on dataclasses (PY-28957)
This commit is contained in:
@@ -15,7 +15,7 @@ import com.jetbrains.python.psi.types.*
|
||||
class PyDataclassesTypeProvider : PyTypeProviderBase() {
|
||||
|
||||
override fun getReferenceExpressionType(referenceExpression: PyReferenceExpression, context: TypeEvalContext): PyType? {
|
||||
return getDataclassTypeForCallee(referenceExpression, context)
|
||||
return getDataclassTypeForCallee(referenceExpression, context) ?: getDataclassesReplaceType(referenceExpression, context)
|
||||
}
|
||||
|
||||
override fun getParameterType(param: PyNamedParameter, func: PyFunction, context: TypeEvalContext): Ref<PyType>? {
|
||||
@@ -55,13 +55,46 @@ class PyDataclassesTypeProvider : PyTypeProviderBase() {
|
||||
.firstOrNull { it != null }
|
||||
}
|
||||
|
||||
private fun getDataclassesReplaceType(referenceExpression: PyReferenceExpression, context: TypeEvalContext): PyCallableType? {
|
||||
val call = PyCallExpressionNavigator.getPyCallExpressionByCallee(referenceExpression) ?: return null
|
||||
val callee = call.callee as? PyReferenceExpression ?: return null
|
||||
|
||||
val resolveContext = PyResolveContext.noImplicits().withTypeEvalContext(context)
|
||||
val resolvedCallee = PyUtil.multiResolveTopPriority(callee.getReference(resolveContext)).singleOrNull()
|
||||
|
||||
if (resolvedCallee is PyCallable && resolvedCallee.qualifiedName == "dataclasses.replace") {
|
||||
val obj = call.getArgument(0, "obj", PyTypedElement::class.java) ?: return null
|
||||
val objType = context.getType(obj) as? PyClassType ?: return null
|
||||
if (objType.isDefinition) return null
|
||||
|
||||
val dataclassType = getDataclassTypeForClass(objType.pyClass, context) ?: return null
|
||||
val dataclassParameters = dataclassType.getParameters(context) ?: return null
|
||||
|
||||
val parameters = mutableListOf<PyCallableParameter>()
|
||||
val elementGenerator = PyElementGenerator.getInstance(referenceExpression.project)
|
||||
|
||||
parameters.add(PyCallableParameterImpl.nonPsi("obj", objType))
|
||||
parameters.add(PyCallableParameterImpl.psi(elementGenerator.createSingleStarParameter()))
|
||||
|
||||
val ellipsis = elementGenerator.createEllipsis()
|
||||
|
||||
dataclassParameters.forEach {
|
||||
parameters.add(PyCallableParameterImpl.nonPsi(it.name, it.getType(context), it.defaultValue ?: ellipsis))
|
||||
}
|
||||
|
||||
return PyCallableTypeImpl(parameters, dataclassType.getReturnType(context))
|
||||
}
|
||||
|
||||
return null
|
||||
}
|
||||
|
||||
private fun getDataclassTypeForClass(cls: PyClass, context: TypeEvalContext): PyCallableType? {
|
||||
val dataclassParameters = parseDataclassParameters(cls, context)
|
||||
if (dataclassParameters == null || !dataclassParameters.init) {
|
||||
return null
|
||||
}
|
||||
|
||||
val parameters = ArrayList<PyCallableParameter>()
|
||||
val parameters = mutableListOf<PyCallableParameter>()
|
||||
val ellipsis = PyElementGenerator.getInstance(cls.project).createEllipsis()
|
||||
|
||||
cls.processClassLevelDeclarations { element, _ ->
|
||||
@@ -101,7 +134,7 @@ class PyDataclassesTypeProvider : PyTypeProviderBase() {
|
||||
|
||||
private fun getTypeForParameter(element: PyTargetExpression, context: TypeEvalContext): PyType? {
|
||||
val type = context.getType(element)
|
||||
if (type is PyCollectionType && type is PyClassType && type.classQName == DATACLASSES_INITVAR_TYPE) {
|
||||
if (type is PyCollectionType && type.classQName == DATACLASSES_INITVAR_TYPE) {
|
||||
return type.elementTypes.firstOrNull()
|
||||
}
|
||||
return type
|
||||
|
||||
@@ -0,0 +1,45 @@
|
||||
from dataclasses import dataclass, field, InitVar, replace
|
||||
|
||||
|
||||
@dataclass
|
||||
class A:
|
||||
a: int
|
||||
b: str = "str"
|
||||
|
||||
|
||||
replace(A(1))
|
||||
replace(A(1), a=1)
|
||||
replace(A(1), a=1, b="abc")
|
||||
replace(A(1), a=1, b="abc", <warning descr="Unexpected argument">c=2</warning>)
|
||||
|
||||
|
||||
@dataclass
|
||||
class B:
|
||||
a: int
|
||||
b: str = field(default="str", init=False)
|
||||
|
||||
|
||||
replace(B(1))
|
||||
replace(B(1), a=1)
|
||||
replace(B(1), a=1, <warning descr="Unexpected argument">b="abc"</warning>)
|
||||
replace(B(1), a=1, <warning descr="Unexpected argument">b="abc"</warning>, <warning descr="Unexpected argument">c=2</warning>)
|
||||
|
||||
|
||||
@dataclass
|
||||
class C:
|
||||
a: int
|
||||
b: InitVar[str] = "str"
|
||||
|
||||
|
||||
replace(C(1))
|
||||
replace(C(1), a=1)
|
||||
replace(C(1), a=1, b="abc")
|
||||
replace(C(1), a=1, b="abc", <warning descr="Unexpected argument">c=2</warning>)
|
||||
|
||||
|
||||
class D:
|
||||
pass
|
||||
|
||||
|
||||
replace(D())
|
||||
replace(D(), a=1, b=2)
|
||||
+20
@@ -0,0 +1,20 @@
|
||||
class _InitVarMeta(type):
|
||||
def __getitem__(self, params):
|
||||
return self
|
||||
|
||||
class InitVar(metaclass=_InitVarMeta):
|
||||
pass
|
||||
|
||||
|
||||
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
|
||||
|
||||
|
||||
def replace(obj, **changes):
|
||||
pass
|
||||
@@ -0,0 +1,42 @@
|
||||
from dataclasses import dataclass, field, InitVar, replace
|
||||
|
||||
|
||||
@dataclass
|
||||
class A:
|
||||
a: int
|
||||
b: str = "str"
|
||||
|
||||
|
||||
replace(A(1))
|
||||
replace(A(1), a=1, b="abc")
|
||||
replace(A(1), <warning descr="Expected type 'int', got 'str' instead">a="str"</warning>, <warning descr="Expected type 'str', got 'int' instead">b=1</warning>)
|
||||
|
||||
|
||||
@dataclass
|
||||
class B:
|
||||
a: int
|
||||
b: str = field(default="str", init=False)
|
||||
|
||||
|
||||
replace(B(1))
|
||||
replace(B(1), a=1)
|
||||
replace(B(1), <warning descr="Expected type 'int', got 'str' instead">a="str"</warning>)
|
||||
|
||||
|
||||
@dataclass
|
||||
class C:
|
||||
a: int
|
||||
b: InitVar[str] = "str"
|
||||
|
||||
|
||||
replace(C(1))
|
||||
replace(C(1), a=1, b="str")
|
||||
replace(C(1), <warning descr="Expected type 'int', got 'str' instead">a="str"</warning>, <warning descr="Expected type 'str', got 'int' instead">b=1</warning>)
|
||||
|
||||
|
||||
class D:
|
||||
pass
|
||||
|
||||
|
||||
replace(D())
|
||||
replace(D(), a=1, b=2)
|
||||
@@ -0,0 +1,20 @@
|
||||
class _InitVarMeta(type):
|
||||
def __getitem__(self, params):
|
||||
return self
|
||||
|
||||
class InitVar(metaclass=_InitVarMeta):
|
||||
pass
|
||||
|
||||
|
||||
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
|
||||
|
||||
|
||||
def replace(obj, **changes):
|
||||
pass
|
||||
@@ -0,0 +1,35 @@
|
||||
from dataclasses import dataclass, field, InitVar, replace
|
||||
|
||||
|
||||
@dataclass
|
||||
class A:
|
||||
a: int
|
||||
b: str = "str"
|
||||
|
||||
|
||||
replace(A(1), <arg1>)
|
||||
|
||||
|
||||
@dataclass
|
||||
class B:
|
||||
a: int
|
||||
b: str = field(default="str", init=False)
|
||||
|
||||
|
||||
replace(B(1), <arg2>)
|
||||
|
||||
|
||||
@dataclass
|
||||
class C:
|
||||
a: int
|
||||
b: InitVar[str] = "str"
|
||||
|
||||
|
||||
replace(C(1), <arg3>)
|
||||
|
||||
|
||||
class D:
|
||||
pass
|
||||
|
||||
|
||||
replace(D(), <arg4>)
|
||||
@@ -0,0 +1,20 @@
|
||||
class _InitVarMeta(type):
|
||||
def __getitem__(self, params):
|
||||
return self
|
||||
|
||||
class InitVar(metaclass=_InitVarMeta):
|
||||
pass
|
||||
|
||||
|
||||
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
|
||||
|
||||
|
||||
def replace(obj, **changes):
|
||||
pass
|
||||
@@ -28,7 +28,6 @@ import com.intellij.psi.PsiElement;
|
||||
import com.intellij.psi.PsiFile;
|
||||
import com.intellij.util.ArrayUtil;
|
||||
import com.intellij.util.Function;
|
||||
import java.util.HashSet;
|
||||
import com.jetbrains.python.fixtures.LightMarkedTestCase;
|
||||
import com.jetbrains.python.psi.LanguageLevel;
|
||||
import com.jetbrains.python.psi.PyArgumentList;
|
||||
@@ -752,10 +751,25 @@ public class PyParameterInfoTest extends LightMarkedTestCase {
|
||||
);
|
||||
}
|
||||
|
||||
// PY-28957
|
||||
public void testDataclassesReplace() {
|
||||
runWithLanguageLevel(
|
||||
LanguageLevel.PYTHON37,
|
||||
() -> {
|
||||
final Map<String, PsiElement> marks = loadMultiFileTest(4);
|
||||
|
||||
feignCtrlP(marks.get("<arg1>").getTextOffset()).check("obj: A, *, a: int=..., b: str=\"str\"", ArrayUtil.EMPTY_STRING_ARRAY);
|
||||
feignCtrlP(marks.get("<arg2>").getTextOffset()).check("obj: B, *, a: int=...", ArrayUtil.EMPTY_STRING_ARRAY);
|
||||
feignCtrlP(marks.get("<arg3>").getTextOffset()).check("obj: C, *, a: int=..., b: str=\"str\"", ArrayUtil.EMPTY_STRING_ARRAY);
|
||||
feignCtrlP(marks.get("<arg4>").getTextOffset()).check("obj, **changes", new String[]{"**changes"});
|
||||
}
|
||||
);
|
||||
}
|
||||
|
||||
// EA-102450
|
||||
public void testKeywordOnlyWithFilledPositional() {
|
||||
runWithLanguageLevel(
|
||||
LanguageLevel.PYTHON30,
|
||||
LanguageLevel.PYTHON34,
|
||||
() -> {
|
||||
final Map<String, PsiElement> test = loadTest(4);
|
||||
|
||||
|
||||
@@ -296,4 +296,9 @@ public class Py3TypeCheckerInspectionTest extends PyInspectionTestCase {
|
||||
public void testInitializingDataclass() {
|
||||
runWithLanguageLevel(LanguageLevel.PYTHON37, () -> super.doMultiFileTest());
|
||||
}
|
||||
|
||||
// PY-28957
|
||||
public void testDataclassesReplace() {
|
||||
runWithLanguageLevel(LanguageLevel.PYTHON37, () -> super.doMultiFileTest());
|
||||
}
|
||||
}
|
||||
|
||||
@@ -334,6 +334,11 @@ public class PyArgumentListInspectionTest extends PyInspectionTestCase {
|
||||
runWithLanguageLevel(LanguageLevel.PYTHON37, this::doMultiFileTest);
|
||||
}
|
||||
|
||||
// PY-28957
|
||||
public void testDataclassesReplace() {
|
||||
runWithLanguageLevel(LanguageLevel.PYTHON37, this::doMultiFileTest);
|
||||
}
|
||||
|
||||
public void testInitializingImportedTypingNamedTupleInheritor() {
|
||||
runWithLanguageLevel(LanguageLevel.PYTHON37, this::doMultiFileTest);
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user