Make EnumValueOptimizer a CodeRewriterPass

Bug: b/284304606
Change-Id: Iaed455751d5a299c5948a395e9c227a657efea03
diff --git a/src/main/java/com/android/tools/r8/ir/analysis/proto/GeneratedMessageLiteBuilderShrinker.java b/src/main/java/com/android/tools/r8/ir/analysis/proto/GeneratedMessageLiteBuilderShrinker.java
index c6658c4..a5e1ecb 100644
--- a/src/main/java/com/android/tools/r8/ir/analysis/proto/GeneratedMessageLiteBuilderShrinker.java
+++ b/src/main/java/com/android/tools/r8/ir/analysis/proto/GeneratedMessageLiteBuilderShrinker.java
@@ -332,7 +332,7 @@
     // Run the enum optimization to optimize all Enum.ordinal() invocations. This is required to
     // get rid of the enum switch in dynamicMethod().
     if (enumValueOptimizer != null) {
-      enumValueOptimizer.rewriteConstantEnumMethodCalls(code);
+      enumValueOptimizer.run(code.context(), code);
     }
   }
 
diff --git a/src/main/java/com/android/tools/r8/ir/conversion/IRConverter.java b/src/main/java/com/android/tools/r8/ir/conversion/IRConverter.java
index 7aae701..957a34f 100644
--- a/src/main/java/com/android/tools/r8/ir/conversion/IRConverter.java
+++ b/src/main/java/com/android/tools/r8/ir/conversion/IRConverter.java
@@ -738,9 +738,7 @@
 
     if (enumValueOptimizer != null) {
       assert appView.enableWholeProgramOptimizations();
-      timing.begin("Rewrite constant enum methods");
-      enumValueOptimizer.rewriteConstantEnumMethodCalls(code);
-      timing.end();
+      enumValueOptimizer.run(context, code, timing);
     }
 
     timing.begin("Rewrite array length");
diff --git a/src/main/java/com/android/tools/r8/ir/conversion/passes/ArrayConstructionSimplifier.java b/src/main/java/com/android/tools/r8/ir/conversion/passes/ArrayConstructionSimplifier.java
index bc8f8fb..bbbb807 100644
--- a/src/main/java/com/android/tools/r8/ir/conversion/passes/ArrayConstructionSimplifier.java
+++ b/src/main/java/com/android/tools/r8/ir/conversion/passes/ArrayConstructionSimplifier.java
@@ -7,7 +7,6 @@
 import com.android.tools.r8.graph.AppInfo;
 import com.android.tools.r8.graph.AppView;
 import com.android.tools.r8.graph.DexClass;
-import com.android.tools.r8.graph.DexItemFactory;
 import com.android.tools.r8.graph.DexType;
 import com.android.tools.r8.graph.ProgramMethod;
 import com.android.tools.r8.ir.analysis.type.ArrayTypeElement;
@@ -84,20 +83,17 @@
  */
 public class ArrayConstructionSimplifier extends CodeRewriterPass<AppInfo> {
 
-  private final DexItemFactory dexItemFactory;
-
   public ArrayConstructionSimplifier(AppView<?> appView) {
     super(appView);
-    this.dexItemFactory = appView.dexItemFactory();
   }
 
   @Override
-  String getTimingId() {
+  protected String getTimingId() {
     return "ArrayConstructionSimplifier";
   }
 
   @Override
-  void rewriteCode(ProgramMethod method, IRCode code) {
+  protected void rewriteCode(ProgramMethod method, IRCode code) {
     WorkList<BasicBlock> worklist = WorkList.newIdentityWorkList(code.blocks);
     while (worklist.hasNext()) {
       BasicBlock block = worklist.next();
@@ -106,7 +102,7 @@
   }
 
   @Override
-  boolean shouldRewriteCode(ProgramMethod method, IRCode code) {
+  protected boolean shouldRewriteCode(ProgramMethod method, IRCode code) {
     return appView.options().isGeneratingDex();
   }
 
diff --git a/src/main/java/com/android/tools/r8/ir/conversion/passes/BinopRewriter.java b/src/main/java/com/android/tools/r8/ir/conversion/passes/BinopRewriter.java
index fd1c703..7c47790 100644
--- a/src/main/java/com/android/tools/r8/ir/conversion/passes/BinopRewriter.java
+++ b/src/main/java/com/android/tools/r8/ir/conversion/passes/BinopRewriter.java
@@ -239,12 +239,12 @@
   }
 
   @Override
-  String getTimingId() {
+  protected String getTimingId() {
     return "BinopRewriter";
   }
 
   @Override
-  boolean shouldRewriteCode(ProgramMethod method, IRCode code) {
+  protected boolean shouldRewriteCode(ProgramMethod method, IRCode code) {
     return true;
   }
 
diff --git a/src/main/java/com/android/tools/r8/ir/conversion/passes/CodeRewriterPass.java b/src/main/java/com/android/tools/r8/ir/conversion/passes/CodeRewriterPass.java
index aea343b..50a469d 100644
--- a/src/main/java/com/android/tools/r8/ir/conversion/passes/CodeRewriterPass.java
+++ b/src/main/java/com/android/tools/r8/ir/conversion/passes/CodeRewriterPass.java
@@ -14,11 +14,11 @@
 
 public abstract class CodeRewriterPass<T extends AppInfo> {
 
-  final AppView<?> appView;
-  final DexItemFactory dexItemFactory;
-  final InternalOptions options;
+  protected final AppView<?> appView;
+  protected final DexItemFactory dexItemFactory;
+  protected final InternalOptions options;
 
-  CodeRewriterPass(AppView<?> appView) {
+  protected CodeRewriterPass(AppView<?> appView) {
     this.appView = appView;
     this.dexItemFactory = appView.dexItemFactory();
     this.options = appView.options();
@@ -39,9 +39,9 @@
     }
   }
 
-  abstract String getTimingId();
+  protected abstract String getTimingId();
 
-  abstract void rewriteCode(ProgramMethod method, IRCode code);
+  protected abstract void rewriteCode(ProgramMethod method, IRCode code);
 
-  abstract boolean shouldRewriteCode(ProgramMethod method, IRCode code);
+  protected abstract boolean shouldRewriteCode(ProgramMethod method, IRCode code);
 }
diff --git a/src/main/java/com/android/tools/r8/ir/conversion/passes/CommonSubexpressionElimination.java b/src/main/java/com/android/tools/r8/ir/conversion/passes/CommonSubexpressionElimination.java
index fb168cd..036bd0f 100644
--- a/src/main/java/com/android/tools/r8/ir/conversion/passes/CommonSubexpressionElimination.java
+++ b/src/main/java/com/android/tools/r8/ir/conversion/passes/CommonSubexpressionElimination.java
@@ -35,12 +35,12 @@
   }
 
   @Override
-  boolean shouldRewriteCode(ProgramMethod method, IRCode code) {
+  protected boolean shouldRewriteCode(ProgramMethod method, IRCode code) {
     return true;
   }
 
   @Override
-  void rewriteCode(ProgramMethod method, IRCode code) {
+  protected void rewriteCode(ProgramMethod method, IRCode code) {
     int noCandidate = code.reserveMarkingColor();
     if (hasCSECandidate(code, noCandidate)) {
       final ListMultimap<Wrapper<Instruction>, Value> instructionToValue =
diff --git a/src/main/java/com/android/tools/r8/ir/conversion/passes/DexConstantOptimizer.java b/src/main/java/com/android/tools/r8/ir/conversion/passes/DexConstantOptimizer.java
index d170c0b..7a34b63 100644
--- a/src/main/java/com/android/tools/r8/ir/conversion/passes/DexConstantOptimizer.java
+++ b/src/main/java/com/android/tools/r8/ir/conversion/passes/DexConstantOptimizer.java
@@ -59,18 +59,18 @@
   }
 
   @Override
-  String getTimingId() {
+  protected String getTimingId() {
     return "DexConstantOptimizer";
   }
 
   @Override
-  void rewriteCode(ProgramMethod method, IRCode code) {
+  protected void rewriteCode(ProgramMethod method, IRCode code) {
     useDedicatedConstantForLitInstruction(code);
     shortenLiveRanges(code, constantCanonicalizer);
   }
 
   @Override
-  boolean shouldRewriteCode(ProgramMethod method, IRCode code) {
+  protected boolean shouldRewriteCode(ProgramMethod method, IRCode code) {
     return true;
   }
 
diff --git a/src/main/java/com/android/tools/r8/ir/conversion/passes/NaturalIntLoopRemover.java b/src/main/java/com/android/tools/r8/ir/conversion/passes/NaturalIntLoopRemover.java
index 16e9415..298020d 100644
--- a/src/main/java/com/android/tools/r8/ir/conversion/passes/NaturalIntLoopRemover.java
+++ b/src/main/java/com/android/tools/r8/ir/conversion/passes/NaturalIntLoopRemover.java
@@ -36,12 +36,12 @@
   }
 
   @Override
-  String getTimingId() {
+  protected String getTimingId() {
     return "NaturalIntLoopRemover";
   }
 
   @Override
-  void rewriteCode(ProgramMethod method, IRCode code) {
+  protected void rewriteCode(ProgramMethod method, IRCode code) {
     boolean loopRemoved = false;
     for (BasicBlock comparisonBlockCandidate : code.blocks) {
       if (isComparisonBlock(comparisonBlockCandidate)) {
@@ -56,7 +56,7 @@
   }
 
   @Override
-  boolean shouldRewriteCode(ProgramMethod method, IRCode code) {
+  protected boolean shouldRewriteCode(ProgramMethod method, IRCode code) {
     return appView.options().enableLoopUnrolling;
   }
 
diff --git a/src/main/java/com/android/tools/r8/ir/conversion/passes/ParentConstructorHoistingCodeRewriter.java b/src/main/java/com/android/tools/r8/ir/conversion/passes/ParentConstructorHoistingCodeRewriter.java
index 06e410b..685926d 100644
--- a/src/main/java/com/android/tools/r8/ir/conversion/passes/ParentConstructorHoistingCodeRewriter.java
+++ b/src/main/java/com/android/tools/r8/ir/conversion/passes/ParentConstructorHoistingCodeRewriter.java
@@ -43,12 +43,12 @@
   }
 
   @Override
-  String getTimingId() {
+  protected String getTimingId() {
     return "Parent constructor hoisting pass";
   }
 
   @Override
-  void rewriteCode(ProgramMethod method, IRCode code) {
+  protected void rewriteCode(ProgramMethod method, IRCode code) {
     for (InvokeDirect invoke : getOrComputeSideEffectFreeConstructorCalls(code)) {
       hoistSideEffectFreeConstructorCall(code, invoke);
     }
@@ -136,9 +136,8 @@
 
   /** Only run this when the rewriting may actually enable more constructor inlining. */
   @Override
-  boolean shouldRewriteCode(ProgramMethod method, IRCode code) {
+  protected boolean shouldRewriteCode(ProgramMethod method, IRCode code) {
     if (!appView.hasClassHierarchy()) {
-      assert !appView.enableWholeProgramOptimizations();
       return false;
     }
     if (!method.getDefinition().isInstanceInitializer()
diff --git a/src/main/java/com/android/tools/r8/ir/conversion/passes/SplitBranch.java b/src/main/java/com/android/tools/r8/ir/conversion/passes/SplitBranch.java
index 8da5119..9290a77 100644
--- a/src/main/java/com/android/tools/r8/ir/conversion/passes/SplitBranch.java
+++ b/src/main/java/com/android/tools/r8/ir/conversion/passes/SplitBranch.java
@@ -34,12 +34,12 @@
   }
 
   @Override
-  String getTimingId() {
+  protected String getTimingId() {
     return "SplitBranch";
   }
 
   @Override
-  boolean shouldRewriteCode(ProgramMethod method, IRCode code) {
+  protected boolean shouldRewriteCode(ProgramMethod method, IRCode code) {
     return true;
   }
 
@@ -54,7 +54,7 @@
    * known boolean values.
    */
   @Override
-  void rewriteCode(ProgramMethod method, IRCode code) {
+  protected void rewriteCode(ProgramMethod method, IRCode code) {
     List<BasicBlock> candidates = computeCandidates(code);
     if (candidates.isEmpty()) {
       return;
diff --git a/src/main/java/com/android/tools/r8/ir/conversion/passes/TrivialGotosCollapser.java b/src/main/java/com/android/tools/r8/ir/conversion/passes/TrivialGotosCollapser.java
index 7c9b4d0..fef138e 100644
--- a/src/main/java/com/android/tools/r8/ir/conversion/passes/TrivialGotosCollapser.java
+++ b/src/main/java/com/android/tools/r8/ir/conversion/passes/TrivialGotosCollapser.java
@@ -29,12 +29,12 @@
   }
 
   @Override
-  String getTimingId() {
+  protected String getTimingId() {
     return "TrivialGotosCollapser";
   }
 
   @Override
-  void rewriteCode(ProgramMethod method, IRCode code) {
+  protected void rewriteCode(ProgramMethod method, IRCode code) {
     assert code.isConsistentGraph(appView);
     List<BasicBlock> blocksToRemove = new ArrayList<>();
     // Rewrite all non-fallthrough targets to the end of trivial goto chains and remove
@@ -77,7 +77,7 @@
   }
 
   @Override
-  boolean shouldRewriteCode(ProgramMethod method, IRCode code) {
+  protected boolean shouldRewriteCode(ProgramMethod method, IRCode code) {
     return true;
   }
 
diff --git a/src/main/java/com/android/tools/r8/ir/optimize/enums/EnumValueOptimizer.java b/src/main/java/com/android/tools/r8/ir/optimize/enums/EnumValueOptimizer.java
index 9a88b2e..3fb3242 100644
--- a/src/main/java/com/android/tools/r8/ir/optimize/enums/EnumValueOptimizer.java
+++ b/src/main/java/com/android/tools/r8/ir/optimize/enums/EnumValueOptimizer.java
@@ -14,6 +14,7 @@
 import com.android.tools.r8.graph.DexItemFactory;
 import com.android.tools.r8.graph.DexMethod;
 import com.android.tools.r8.graph.DexType;
+import com.android.tools.r8.graph.ProgramMethod;
 import com.android.tools.r8.ir.analysis.type.ClassTypeElement;
 import com.android.tools.r8.ir.analysis.type.TypeAnalysis;
 import com.android.tools.r8.ir.analysis.type.TypeElement;
@@ -33,6 +34,7 @@
 import com.android.tools.r8.ir.code.InvokeVirtual;
 import com.android.tools.r8.ir.code.StaticGet;
 import com.android.tools.r8.ir.code.Value;
+import com.android.tools.r8.ir.conversion.passes.CodeRewriterPass;
 import com.android.tools.r8.ir.optimize.info.FieldOptimizationInfo;
 import com.android.tools.r8.shaking.AppInfoWithLiveness;
 import com.android.tools.r8.utils.ArrayUtils;
@@ -47,21 +49,31 @@
 import java.util.Arrays;
 import java.util.Set;
 
-public class EnumValueOptimizer {
-
-  private final AppView<AppInfoWithLiveness> appView;
-  private final DexItemFactory factory;
+public class EnumValueOptimizer extends CodeRewriterPass<AppInfoWithLiveness> {
 
   public EnumValueOptimizer(AppView<AppInfoWithLiveness> appView) {
-    this.appView = appView;
-    this.factory = appView.dexItemFactory();
+    super(appView);
+  }
+
+  @Override
+  protected String getTimingId() {
+    return "EnumValueOptimizer";
+  }
+
+  @Override
+  protected void rewriteCode(ProgramMethod method, IRCode code) {
+    rewriteConstantEnumMethodCalls(code);
+  }
+
+  @Override
+  protected boolean shouldRewriteCode(ProgramMethod method, IRCode code) {
+    return code.metadata().mayHaveInvokeMethodWithReceiver();
   }
 
   @SuppressWarnings("ConstantConditions")
-  public void rewriteConstantEnumMethodCalls(IRCode code) {
+  private void rewriteConstantEnumMethodCalls(IRCode code) {
     IRMetadata metadata = code.metadata();
-    if (!metadata.mayHaveInvokeMethodWithReceiver()
-        && !(metadata.mayHaveInvokeStatic() && metadata.mayHaveArrayLength())) {
+    if (!metadata.mayHaveInvokeMethodWithReceiver()) {
       return;
     }
 
@@ -76,14 +88,15 @@
         if (!receiver.getType().isClassType()
             || !appView
                 .appInfo()
-                .isSubtype(receiver.getType().asClassType().getClassType(), factory.enumType)) {
+                .isSubtype(
+                    receiver.getType().asClassType().getClassType(), dexItemFactory.enumType)) {
           continue;
         }
 
         DexMethod invokedMethod = methodWithReceiver.getInvokedMethod();
-        boolean isOrdinalInvoke = invokedMethod.match(factory.enumMembers.ordinalMethod);
-        boolean isNameInvoke = invokedMethod.match(factory.enumMembers.nameMethod);
-        boolean isToStringInvoke = invokedMethod.match(factory.enumMembers.toString);
+        boolean isOrdinalInvoke = invokedMethod.match(dexItemFactory.enumMembers.ordinalMethod);
+        boolean isNameInvoke = invokedMethod.match(dexItemFactory.enumMembers.nameMethod);
+        boolean isToStringInvoke = invokedMethod.match(dexItemFactory.enumMembers.toString);
         if (!isOrdinalInvoke && !isNameInvoke && !isToStringInvoke) {
           continue;
         }
@@ -161,9 +174,10 @@
             appView
                 .appInfo()
                 .resolveMethodOnClassLegacy(
-                    enumFieldType.getClassType(), factory.objectMembers.toString)
+                    enumFieldType.getClassType(), dexItemFactory.objectMembers.toString)
                 .getSingleTarget();
-        if (singleTarget != null && singleTarget.getReference() != factory.enumMembers.toString) {
+        if (singleTarget != null
+            && singleTarget.getReference() != dexItemFactory.enumMembers.toString) {
           continue;
         }
 
@@ -334,7 +348,7 @@
                       enumInstanceField,
                       appView
                           .appInfo()
-                          .resolveField(factory.enumMembers.ordinalField, code.context())
+                          .resolveField(dexItemFactory.enumMembers.ordinalField, code.context())
                           .getResolvedField());
         }
         if (ordinalValue == null) {
@@ -350,14 +364,14 @@
   private SingleStringValue getNameValue(
       IRCode code, AbstractValue abstractValue, boolean neverNull) {
     AbstractValue ordinalValue =
-        getEnumFieldValue(code, abstractValue, factory.enumMembers.nameField, neverNull);
+        getEnumFieldValue(code, abstractValue, dexItemFactory.enumMembers.nameField, neverNull);
     return ordinalValue == null ? null : ordinalValue.asSingleStringValue();
   }
 
   private SingleNumberValue getOrdinalValue(
       IRCode code, AbstractValue abstractValue, boolean neverNull) {
     AbstractValue ordinalValue =
-        getEnumFieldValue(code, abstractValue, factory.enumMembers.ordinalField, neverNull);
+        getEnumFieldValue(code, abstractValue, dexItemFactory.enumMembers.ordinalField, neverNull);
     return ordinalValue == null ? null : ordinalValue.asSingleNumberValue();
   }