Provide special type for replace usages on dataclasses (PY-28957)

This commit is contained in:
Semyon Proshev
2018-03-20 14:20:25 +03:00
parent 610c1b70c3
commit f5550fbc81
10 changed files with 244 additions and 5 deletions
@@ -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)
@@ -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);
}