diff --git a/python/src/com/jetbrains/python/psi/impl/PyAssignmentStatementImpl.java b/python/src/com/jetbrains/python/psi/impl/PyAssignmentStatementImpl.java index 57664e8d97fb..c8ac9269fa0a 100644 --- a/python/src/com/jetbrains/python/psi/impl/PyAssignmentStatementImpl.java +++ b/python/src/com/jetbrains/python/psi/impl/PyAssignmentStatementImpl.java @@ -20,17 +20,14 @@ import com.intellij.lang.ASTNode; import com.intellij.openapi.util.Pair; import com.intellij.psi.PsiElement; import com.intellij.psi.PsiErrorElement; -import com.intellij.psi.PsiWhiteSpace; import com.intellij.psi.ResolveState; import com.intellij.psi.scope.PsiScopeProcessor; import com.intellij.psi.util.PsiTreeUtil; import com.intellij.util.SmartList; -import com.intellij.util.containers.HashMap; import com.jetbrains.python.PyElementTypes; import com.jetbrains.python.psi.*; import com.jetbrains.python.toolbox.FP; import com.jetbrains.python.toolbox.RepeatIterable; -import com.jetbrains.python.toolbox.RepeatIterator; import org.jetbrains.annotations.NotNull; import org.jetbrains.annotations.Nullable; @@ -98,21 +95,21 @@ public class PyAssignmentStatementImpl extends PyElementImpl implements PyAssign private static void mapToValues(PyExpression lhs, PyExpression rhs, List> map) { // cast for convenience PyTupleExpression lhs_tuple = null; - PyTargetExpression lhs_target = null; + PyExpression lhs_one = null; if (lhs instanceof PyTupleExpression) lhs_tuple = (PyTupleExpression)lhs; - else if (lhs instanceof PyTargetExpression) lhs_target = (PyTargetExpression)lhs; + else if (lhs != null) lhs_one = lhs; PyTupleExpression rhs_tuple = null; - PyTargetExpression rhs_target = null; + PyExpression rhs_one = null; if (rhs instanceof PyTupleExpression) rhs_tuple = (PyTupleExpression)rhs; - else if (rhs instanceof PyTargetExpression) rhs_target = (PyTargetExpression)rhs; + else if (rhs != null) rhs_one = rhs; // - if (lhs_target != null) { // single LHS, single RHS (direct mapping) or multiple RHS (packing) - map.add(new Pair(lhs_target, rhs)); + if (lhs_one != null) { // single LHS, single RHS (direct mapping) or multiple RHS (packing) + map.add(new Pair(lhs_one, rhs)); } - else if (lhs_tuple != null && rhs_target != null) { // multiple LHS, single RHS: unpacking - //for (PyExpression tuple_elt : lhs_tuple.getElements()) map.add(new Pair(tuple_elt, rhs_target)); - map.addAll(FP.zipList(lhs_tuple, new RepeatIterable(rhs_target))); + else if (lhs_tuple != null && rhs_one != null) { // multiple LHS, single RHS: unpacking + //for (PyExpression tuple_elt : lhs_tuple.getElements()) map.add(new Pair(tuple_elt, rhs_one)); + map.addAll(FP.zipList(lhs_tuple, new RepeatIterable(rhs_one))); } else if (lhs_tuple != null && rhs_tuple != null) { // multiple both sides: piecewise mapping map.addAll(FP.zipList(lhs_tuple, rhs_tuple, null, null)); diff --git a/python/testData/psi/assignment/Multiple.py b/python/testData/psi/assignment/Multiple.py new file mode 100644 index 000000000000..91c5d4360eb5 --- /dev/null +++ b/python/testData/psi/assignment/Multiple.py @@ -0,0 +1 @@ +a = b = c = 1 diff --git a/python/testData/psi/assignment/Simple.py b/python/testData/psi/assignment/Simple.py new file mode 100644 index 000000000000..680e39e574ad --- /dev/null +++ b/python/testData/psi/assignment/Simple.py @@ -0,0 +1 @@ +a = "foo" diff --git a/python/testData/psi/assignment/SubscribedSource.py b/python/testData/psi/assignment/SubscribedSource.py new file mode 100644 index 000000000000..52080e78587b --- /dev/null +++ b/python/testData/psi/assignment/SubscribedSource.py @@ -0,0 +1 @@ +a = foo['bar'] diff --git a/python/testData/psi/assignment/SubscribedTarget.py b/python/testData/psi/assignment/SubscribedTarget.py new file mode 100644 index 000000000000..81c5fb6a1f11 --- /dev/null +++ b/python/testData/psi/assignment/SubscribedTarget.py @@ -0,0 +1 @@ +foo['bar'] = 1 diff --git a/python/testData/psi/assignment/TupleMapped.py b/python/testData/psi/assignment/TupleMapped.py new file mode 100644 index 000000000000..4fd19fe221b4 --- /dev/null +++ b/python/testData/psi/assignment/TupleMapped.py @@ -0,0 +1 @@ +a, b = 1, 2 diff --git a/python/testData/psi/assignment/TuplePack.py b/python/testData/psi/assignment/TuplePack.py new file mode 100644 index 000000000000..e1b4b0e34d0e --- /dev/null +++ b/python/testData/psi/assignment/TuplePack.py @@ -0,0 +1 @@ +some_tuple = 1, 2 diff --git a/python/testData/psi/assignment/TupleUnpack.py b/python/testData/psi/assignment/TupleUnpack.py new file mode 100644 index 000000000000..a0c96320be5a --- /dev/null +++ b/python/testData/psi/assignment/TupleUnpack.py @@ -0,0 +1 @@ +a, b = some_tuple diff --git a/python/testSrc/com/jetbrains/python/PyAssignmentMappingTest.java b/python/testSrc/com/jetbrains/python/PyAssignmentMappingTest.java new file mode 100644 index 000000000000..29e265ee8a8d --- /dev/null +++ b/python/testSrc/com/jetbrains/python/PyAssignmentMappingTest.java @@ -0,0 +1,154 @@ +package com.jetbrains.python; + +import com.intellij.openapi.util.Pair; +import com.intellij.psi.PsiElement; +import com.jetbrains.python.psi.PyAssignmentStatement; +import com.jetbrains.python.psi.PyExpression; +import com.jetbrains.python.psi.PyTargetExpression; +import com.jetbrains.python.psi.PySubscriptionExpression; +import java.util.List; +import java.util.Map; + +/** + * Tests assignment mapping and tracking. + * User: dcheryasov + * Date: Dec 11, 2009 2:13:51 AM + */ +public class PyAssignmentMappingTest extends MarkedTestCase { + + public String getTestDataPath() { + return PythonTestUtil.getTestDataPath() + "/psi/assignment/"; + } + + + public void testSimple() throws Exception { + Map marks = loadTest(); + assertEquals(2, marks.size()); + PsiElement src = marks.get("").getParent(); // const -> expr; + PsiElement dst = marks.get("").getParent(); // ident -> target expr + assertTrue(dst instanceof PyTargetExpression); + PyAssignmentStatement stmt = (PyAssignmentStatement)dst.getParent(); + List> mapping = stmt.getTargetsToValuesMapping(); + assertEquals(1, mapping.size()); + Pair pair = mapping.get(0); + assertEquals(dst, pair.getFirst()); + assertEquals(src, pair.getSecond()); + } + + public void testSubscribedSource() throws Exception { + Map marks = loadTest(); + assertEquals(2, marks.size()); + PsiElement src = marks.get("").getParent().getParent(); // const -> ref foo -> subscr expr; + PsiElement dst = marks.get("").getParent(); // ident -> target expr + assertTrue(dst instanceof PyTargetExpression); + PyAssignmentStatement stmt = (PyAssignmentStatement)dst.getParent(); + List> mapping = stmt.getTargetsToValuesMapping(); + assertEquals(1, mapping.size()); + Pair pair = mapping.get(0); + assertEquals(dst, pair.getFirst()); + assertEquals(src, pair.getSecond()); + } + + public void testSubscribedTarget() throws Exception { + Map marks = loadTest(); + assertEquals(2, marks.size()); + PsiElement src = marks.get("").getParent(); // const -> expr; + PsiElement dst = marks.get("").getParent().getParent(); // ident -> target expr + assertTrue(dst instanceof PySubscriptionExpression); + PyAssignmentStatement stmt = (PyAssignmentStatement)src.getParent(); + List> mapping = stmt.getTargetsToValuesMapping(); + assertEquals(1, mapping.size()); + Pair pair = mapping.get(0); + assertEquals(dst, pair.getFirst()); + assertEquals(src, pair.getSecond()); + } + + + public void testMultiple() throws Exception { + Map marks = loadTest(); + final int TARGET_NUM = 3; + assertEquals(TARGET_NUM+1, marks.size()); + PsiElement src = marks.get("").getParent(); // const -> expr; + PsiElement[] dsts = new PsiElement[TARGET_NUM]; + for (int i=0; i").getParent(); // ident -> target expr + assertTrue(dst instanceof PyTargetExpression); + dsts[i] = dst; + } + PyAssignmentStatement stmt = (PyAssignmentStatement)src.getParent(); + List> mapping = stmt.getTargetsToValuesMapping(); + assertEquals(TARGET_NUM, mapping.size()); + for (int i=0; i pair = mapping.get(i); + assertEquals(dsts[i], pair.getFirst()); + assertEquals(src, pair.getSecond()); + } + } + + public void testTupleMapped() throws Exception { + Map marks = loadTest(); + final int PAIR_NUM = 2; + assertEquals(PAIR_NUM*2, marks.size()); + PsiElement[] srcs = new PsiElement[PAIR_NUM]; + PsiElement[] dsts = new PsiElement[PAIR_NUM]; + for (int i=0; i").getParent(); // ident -> target expr + assertTrue(dst instanceof PyTargetExpression); + dsts[i] = dst; + PsiElement src = marks.get("").getParent(); // ident -> target expr + assertTrue(src instanceof PyExpression); + srcs[i] = src; + } + PyAssignmentStatement stmt = (PyAssignmentStatement)srcs[0].getParent().getParent(); // tuple expr -> assignment + List> mapping = stmt.getTargetsToValuesMapping(); + assertEquals(PAIR_NUM, mapping.size()); + for (int i=0; i pair = mapping.get(i); + assertEquals(dsts[i], pair.getFirst()); + assertEquals(srcs[i], pair.getSecond()); + } + } + + public void testTuplePack() throws Exception { + Map marks = loadTest(); + final int SRC_NUM = 2; + assertEquals(SRC_NUM+1, marks.size()); + PsiElement[] srcs = new PsiElement[SRC_NUM]; + for (int i=0; i").getParent(); // ident -> target expr + assertTrue(src instanceof PyExpression); + srcs[i] = src; + } + PsiElement dst = marks.get("").getParent(); // ident -> target expr + PyAssignmentStatement stmt = (PyAssignmentStatement)dst.getParent(); + List> mapping = stmt.getTargetsToValuesMapping(); + assertEquals(1, mapping.size()); + Pair pair = mapping.get(0); + assertEquals(dst, pair.getFirst()); + for (PsiElement src : srcs) { + assertEquals(src.getParent(), pair.getSecond()); // numeric expr -> tuple + } + } + + + public void testTupleUnpack() throws Exception { + Map marks = loadTest(); + final int DST_NUM = 2; + assertEquals(DST_NUM+1, marks.size()); + PsiElement[] dsts = new PsiElement[DST_NUM]; + for (int i=0; i").getParent(); // ident -> target expr + assertTrue(dst instanceof PyTargetExpression); + dsts[i] = dst; + } + PsiElement src = marks.get("").getParent(); // ident -> target expr + PyAssignmentStatement stmt = (PyAssignmentStatement)src.getParent(); + List> mapping = stmt.getTargetsToValuesMapping(); + assertEquals(DST_NUM, mapping.size()); + for (int i=0; i pair = mapping.get(i); + assertEquals(dsts[i], pair.getFirst()); + assertEquals(src, pair.getSecond()); + } + } +} diff --git a/python/testSrc/com/jetbrains/python/PythonAllTestsSuite.java b/python/testSrc/com/jetbrains/python/PythonAllTestsSuite.java index 5ff3e4441982..9b00de7212c1 100644 --- a/python/testSrc/com/jetbrains/python/PythonAllTestsSuite.java +++ b/python/testSrc/com/jetbrains/python/PythonAllTestsSuite.java @@ -20,6 +20,7 @@ public class PythonAllTestsSuite { PyMultiFileResolveTest.class, PyResolveCalleeTest.class, PyToJavaResolveTest.class, + PyAssignmentTrackingTest.class, PythonCompletionTest.class, PyInheritorsSearchTest.class, PyParameterInfoTest.class,