Guard against recursive phis using seen set

This fixes the imprecision caught by TypePropagationThroughPhiTest,
since we now no longer backtrack when hitting `value != initialValue`,
but rather when hitting `!seenPhis.add(phi)`.

Bug: b/309575527
Change-Id: Idd421ad82fa0652755b453ceed4570056de00975
diff --git a/src/main/java/com/android/tools/r8/optimize/argumentpropagation/ArgumentPropagatorCodeScanner.java b/src/main/java/com/android/tools/r8/optimize/argumentpropagation/ArgumentPropagatorCodeScanner.java
index bea4ad0..e0a11a3 100644
--- a/src/main/java/com/android/tools/r8/optimize/argumentpropagation/ArgumentPropagatorCodeScanner.java
+++ b/src/main/java/com/android/tools/r8/optimize/argumentpropagation/ArgumentPropagatorCodeScanner.java
@@ -231,7 +231,7 @@
 
     // State used by internalComputeNonReceiverValueState.
     private Map<Value, NonEmptyValueState> cache = new IdentityHashMap<>();
-    private Value initialValue;
+    private Set<Phi> seenPhis = Sets.newIdentityHashSet();
     private DexType staticType;
     private ProgramMember<?, ?> target;
 
@@ -313,15 +313,14 @@
       }
 
       assert this.cache.isEmpty();
-      assert this.initialValue == null;
+      assert this.seenPhis.isEmpty();
       assert this.staticType == null;
       assert this.target == null;
-      this.initialValue = value;
       this.staticType = staticType;
       this.target = target;
       NonEmptyValueState result = internalGetOrComputeNonReceiverValueState(value);
       this.cache.clear();
-      this.initialValue = null;
+      this.seenPhis.clear();
       this.staticType = null;
       this.target = null;
       return result;
@@ -332,10 +331,13 @@
     }
 
     private NonEmptyValueState internalComputeNonReceiverValueState(Value value) {
-      assert value == initialValue || initialValue.getAliasedValue().isPhi();
-
       if (value.isPhi()) {
-        return computePhiState(value.asPhi());
+        // In presence of recursive phis we fall back to computing the state of the phi from the
+        // phi type instead of from its operands.
+        Phi phi = value.asPhi();
+        if (seenPhis.add(phi)) {
+          return computePhiState(phi);
+        }
       }
 
       // If the current value is an argument of the declaring method, then we have no information
@@ -404,10 +406,6 @@
     //  same value multiple times.
     // TODO(b/302281503): Canonicalize computed in flow.
     private InFlow computeInFlow(Value value) {
-      if (value != initialValue) {
-        assert initialValue.getAliasedValue().isPhi();
-        return computeBaseInFlow(value);
-      }
       Value valueRoot = value.getAliasedValue(aliasedValueConfiguration);
       if (valueRoot.isArgument()) {
         MethodParameter inParameter =
@@ -479,8 +477,11 @@
         return null;
       }
       NonEmptyValueState leftValue = internalGetOrComputeNonReceiverValueState(phi.getOperand(0));
+      if (leftValue.isUnknown() || leftValue.asConcrete().hasNonBaseInFlow()) {
+        return null;
+      }
       NonEmptyValueState rightValue = internalGetOrComputeNonReceiverValueState(phi.getOperand(1));
-      if (leftValue.isUnknown() && rightValue.isUnknown()) {
+      if (rightValue.isUnknown() || rightValue.asConcrete().hasNonBaseInFlow()) {
         return null;
       }
       IfThenElseAbstractFunction result =
@@ -601,7 +602,6 @@
     }
 
     private NonEmptyValueState computeInFlowState(Value value) {
-      assert value == initialValue || initialValue.getAliasedValue().isPhi();
       InFlow inFlow = computeInFlow(value);
       if (inFlow != null && !inFlow.isUnknown()) {
         assert inFlow.isBaseInFlow()
diff --git a/src/main/java/com/android/tools/r8/optimize/argumentpropagation/codescanner/ConcreteValueState.java b/src/main/java/com/android/tools/r8/optimize/argumentpropagation/codescanner/ConcreteValueState.java
index 06ee793..a7ab2a7 100644
--- a/src/main/java/com/android/tools/r8/optimize/argumentpropagation/codescanner/ConcreteValueState.java
+++ b/src/main/java/com/android/tools/r8/optimize/argumentpropagation/codescanner/ConcreteValueState.java
@@ -97,6 +97,15 @@
     return !inFlow.isEmpty();
   }
 
+  public boolean hasNonBaseInFlow() {
+    for (InFlow inFlow : inFlow) {
+      if (!inFlow.isBaseInFlow()) {
+        return true;
+      }
+    }
+    return false;
+  }
+
   public Set<InFlow> getInFlow() {
     assert inFlow.isEmpty() || inFlow instanceof HashSet<?>;
     return inFlow;
diff --git a/src/test/java8/ir/com/android/tools/r8/ir/optimize/intlongarithmetic/ComposeBinopTest.java b/src/test/java8/ir/com/android/tools/r8/ir/optimize/intlongarithmetic/ComposeBinopTest.java
index 9989e20..b8ef372 100644
--- a/src/test/java8/ir/com/android/tools/r8/ir/optimize/intlongarithmetic/ComposeBinopTest.java
+++ b/src/test/java8/ir/com/android/tools/r8/ir/optimize/intlongarithmetic/ComposeBinopTest.java
@@ -7,6 +7,7 @@
 import static org.junit.Assert.assertEquals;
 import static org.junit.Assert.assertTrue;
 
+import com.android.tools.r8.KeepConstantArguments;
 import com.android.tools.r8.NeverInline;
 import com.android.tools.r8.TestBase;
 import com.android.tools.r8.TestParameters;
@@ -130,7 +131,7 @@
   }
 
   @Test
-  public void testD8() throws Exception {
+  public void testRuntime() throws Exception {
     testForRuntime(parameters)
         .addProgramClasses(Main.class)
         .run(parameters.getRuntime(), Main.class)
@@ -142,6 +143,7 @@
     testForR8(parameters.getBackend())
         .addProgramClasses(Main.class)
         .addKeepMainRule(Main.class)
+        .enableConstantArgumentAnnotations()
         .enableInliningAnnotations()
         .setMinApi(parameters)
         .compile()
@@ -239,6 +241,7 @@
       composeTests(i3);
     }
 
+    @KeepConstantArguments
     @NeverInline
     private static void bitSameInput(int a) {
       // a & a => a, a | a => a.
@@ -249,6 +252,7 @@
       System.out.println(a - a);
     }
 
+    @KeepConstantArguments
     @NeverInline
     private static void shareShiftCstRight(int a, int b) {
       // (x shift: val) | (y shift: val) => (x | y) shift: val.
@@ -261,6 +265,7 @@
       System.out.println((a >>> 3) & (b >>> 3));
     }
 
+    @KeepConstantArguments
     @NeverInline
     private static void shareShiftCstLeft(int a) {
       // (x shift: val) | (y shift: val) => (x | y) shift: val.
@@ -273,6 +278,7 @@
       System.out.println((116 >>> a) & (227 >>> a));
     }
 
+    @KeepConstantArguments
     @NeverInline
     private static void shareShiftVar(int a, int b, int c) {
       // (x shift: val) | (y shift: val) => (x | y) shift: val.
@@ -285,6 +291,7 @@
       System.out.println((a >>> c) & (b >>> c));
     }
 
+    @KeepConstantArguments
     @NeverInline
     private static void andOrCompositionCst(int a) {
       // For all permutations of & and |, represented by &| and |&.
@@ -296,6 +303,7 @@
       System.out.println((a | 0b10101010) | (a | 0b11110000));
     }
 
+    @KeepConstantArguments
     @NeverInline
     private static void andOrCompositionVar(int a, int b, int c) {
       // For all permutations of & and |.
@@ -306,6 +314,7 @@
       System.out.println((a | b) | (a | c));
     }
 
+    @KeepConstantArguments
     @NeverInline
     private static void composeTests(int dirty) {
       // This is rewritten to ((0b0001111111111110 & dirty) >> 3).
diff --git a/src/test/java8/optimize/com/android/tools/r8/optimize/argumentpropagation/TypePropagationThroughPhiTest.java b/src/test/java8/optimize/com/android/tools/r8/optimize/argumentpropagation/TypePropagationThroughPhiTest.java
index 89adea8..dc2e16d 100644
--- a/src/test/java8/optimize/com/android/tools/r8/optimize/argumentpropagation/TypePropagationThroughPhiTest.java
+++ b/src/test/java8/optimize/com/android/tools/r8/optimize/argumentpropagation/TypePropagationThroughPhiTest.java
@@ -40,9 +40,8 @@
               assertEquals(
                   "java.lang.String",
                   mainClass.uniqueMethodWithOriginalName("foo").getParameter(0).getTypeName());
-              // TODO(b/309575527): Should be String.
               assertEquals(
-                  "java.lang.CharSequence",
+                  "java.lang.String",
                   mainClass.uniqueMethodWithOriginalName("bar").getParameter(0).getTypeName());
             });
   }