Optimize rewriting of field collections

Change-Id: Id768a3bf91e87a456eeb368f8c0c35ef29cf1e2c
diff --git a/src/main/java/com/android/tools/r8/graph/AbstractAccessContexts.java b/src/main/java/com/android/tools/r8/graph/AbstractAccessContexts.java
index 4c63987..0229986 100644
--- a/src/main/java/com/android/tools/r8/graph/AbstractAccessContexts.java
+++ b/src/main/java/com/android/tools/r8/graph/AbstractAccessContexts.java
@@ -6,9 +6,11 @@
 
 import com.android.tools.r8.errors.Unreachable;
 import com.android.tools.r8.graph.lens.GraphLens;
+import com.android.tools.r8.utils.MapUtils;
 import com.android.tools.r8.utils.collections.ProgramMethodSet;
 import java.util.IdentityHashMap;
 import java.util.Map;
+import java.util.Map.Entry;
 import java.util.function.BiConsumer;
 import java.util.function.Consumer;
 import java.util.function.Predicate;
@@ -306,19 +308,50 @@
 
     @Override
     ConcreteAccessContexts rewrittenWithLens(DexDefinitionSupplier definitions, GraphLens lens) {
-      Map<DexField, ProgramMethodSet> newAccessesWithContexts = new IdentityHashMap<>();
-      accessesWithContexts.forEach(
-          (access, contexts) -> {
-            ProgramMethodSet newContexts =
-                newAccessesWithContexts.computeIfAbsent(
-                    lens.lookupField(access), ignore -> ProgramMethodSet.create());
-            for (ProgramMethod context : contexts) {
-              ProgramMethod newContext = lens.mapProgramMethod(context, definitions);
-              assert newContext != null : "Unable to map context: " + context.toSourceString();
-              newContexts.add(newContext);
-            }
-          });
-      return new ConcreteAccessContexts(newAccessesWithContexts);
+      Map<DexField, ProgramMethodSet> rewrittenAccessesWithContexts = null;
+      for (Entry<DexField, ProgramMethodSet> entry : accessesWithContexts.entrySet()) {
+        DexField field = entry.getKey();
+        DexField rewrittenField = lens.lookupField(field);
+
+        ProgramMethodSet contexts = entry.getValue();
+        ProgramMethodSet rewrittenContexts = contexts.rewrittenWithLens(definitions, lens);
+
+        if (rewrittenField == field && rewrittenContexts == contexts) {
+          if (rewrittenAccessesWithContexts == null) {
+            continue;
+          }
+        } else {
+          if (rewrittenAccessesWithContexts == null) {
+            rewrittenAccessesWithContexts = new IdentityHashMap<>(accessesWithContexts.size());
+            MapUtils.forEachUntilExclusive(
+                accessesWithContexts, rewrittenAccessesWithContexts::put, field);
+          }
+        }
+        merge(rewrittenAccessesWithContexts, rewrittenField, rewrittenContexts);
+      }
+      if (rewrittenAccessesWithContexts != null) {
+        rewrittenAccessesWithContexts =
+            MapUtils.trimCapacityOfIdentityHashMapIfSizeLessThan(
+                rewrittenAccessesWithContexts, accessesWithContexts.size());
+        return new ConcreteAccessContexts(rewrittenAccessesWithContexts);
+      } else {
+        return this;
+      }
+    }
+
+    private static void merge(
+        Map<DexField, ProgramMethodSet> accessesWithContexts,
+        DexField field,
+        ProgramMethodSet contexts) {
+      ProgramMethodSet existingContexts = accessesWithContexts.put(field, contexts);
+      if (existingContexts != null) {
+        if (existingContexts.size() <= contexts.size()) {
+          contexts.addAll(existingContexts);
+        } else {
+          accessesWithContexts.put(field, existingContexts);
+          existingContexts.addAll(contexts);
+        }
+      }
     }
 
     @Override
diff --git a/src/main/java/com/android/tools/r8/graph/FieldAccessInfoImpl.java b/src/main/java/com/android/tools/r8/graph/FieldAccessInfoImpl.java
index 1bc6ae6..4344d55 100644
--- a/src/main/java/com/android/tools/r8/graph/FieldAccessInfoImpl.java
+++ b/src/main/java/com/android/tools/r8/graph/FieldAccessInfoImpl.java
@@ -41,14 +41,25 @@
 
   // Maps every direct and indirect reference in a read-context to the set of methods in which that
   // reference appears.
-  private AbstractAccessContexts readsWithContexts = AbstractAccessContexts.empty();
+  private AbstractAccessContexts readsWithContexts;
 
   // Maps every direct and indirect reference in a write-context to the set of methods in which that
   // reference appears.
-  private AbstractAccessContexts writesWithContexts = AbstractAccessContexts.empty();
+  private AbstractAccessContexts writesWithContexts;
 
   public FieldAccessInfoImpl(DexField field) {
+    this(field, 0, AbstractAccessContexts.empty(), AbstractAccessContexts.empty());
+  }
+
+  public FieldAccessInfoImpl(
+      DexField field,
+      int flags,
+      AbstractAccessContexts readsWithContexts,
+      AbstractAccessContexts writesWithContexts) {
     this.field = field;
+    this.flags = flags;
+    this.readsWithContexts = readsWithContexts;
+    this.writesWithContexts = writesWithContexts;
   }
 
   void destroyAccessContexts() {
@@ -360,10 +371,28 @@
   public FieldAccessInfoImpl rewrittenWithLens(
       DexDefinitionSupplier definitions, GraphLens lens, Timing timing) {
     timing.begin("Rewrite FieldAccessInfoImpl");
-    FieldAccessInfoImpl rewritten = new FieldAccessInfoImpl(lens.lookupField(field));
-    rewritten.flags = flags;
-    rewritten.readsWithContexts = readsWithContexts.rewrittenWithLens(definitions, lens);
-    rewritten.writesWithContexts = writesWithContexts.rewrittenWithLens(definitions, lens);
+    AbstractAccessContexts rewrittenReadsWithContexts =
+        readsWithContexts.rewrittenWithLens(definitions, lens);
+    AbstractAccessContexts rewrittenWritesWithContexts =
+        writesWithContexts.rewrittenWithLens(definitions, lens);
+    FieldAccessInfoImpl rewritten;
+    if (lens.isIdentityLensForFields(GraphLens.getIdentityLens())) {
+      if (rewrittenReadsWithContexts == readsWithContexts
+          && rewrittenWritesWithContexts == writesWithContexts) {
+        rewritten = this;
+      } else {
+        rewritten =
+            new FieldAccessInfoImpl(
+                field, flags, rewrittenReadsWithContexts, rewrittenWritesWithContexts);
+      }
+    } else {
+      rewritten =
+          new FieldAccessInfoImpl(
+              lens.lookupField(field),
+              flags,
+              rewrittenReadsWithContexts,
+              rewrittenWritesWithContexts);
+    }
     timing.end();
     return rewritten;
   }
diff --git a/src/main/java/com/android/tools/r8/graph/lens/GraphLens.java b/src/main/java/com/android/tools/r8/graph/lens/GraphLens.java
index 43deea3..5c884ce 100644
--- a/src/main/java/com/android/tools/r8/graph/lens/GraphLens.java
+++ b/src/main/java/com/android/tools/r8/graph/lens/GraphLens.java
@@ -29,6 +29,7 @@
 import com.android.tools.r8.optimize.MemberRebindingIdentityLens;
 import com.android.tools.r8.optimize.MemberRebindingLens;
 import com.android.tools.r8.shaking.KeepInfoCollection;
+import com.android.tools.r8.utils.CollectionUtils;
 import com.android.tools.r8.utils.InternalOptions;
 import com.android.tools.r8.utils.ListUtils;
 import com.android.tools.r8.utils.SetUtils;
@@ -42,6 +43,7 @@
 import it.unimi.dsi.fastutil.objects.Object2BooleanArrayMap;
 import it.unimi.dsi.fastutil.objects.Object2BooleanMap;
 import java.util.ArrayList;
+import java.util.Collection;
 import java.util.IdentityHashMap;
 import java.util.List;
 import java.util.Map;
@@ -432,6 +434,8 @@
 
   public abstract boolean isIdentityLens();
 
+  public abstract boolean isIdentityLensForFields(GraphLens codeLens);
+
   public boolean isMemberRebindingLens() {
     return false;
   }
@@ -523,6 +527,51 @@
     return result;
   }
 
+  public Set<DexField> rewriteFields(Set<DexField> fields, Timing timing) {
+    timing.begin("Rewrite fields");
+    GraphLens appliedLens = getIdentityLens();
+    Set<DexField> rewrittenFields;
+    if (isIdentityLensForFields(appliedLens)) {
+      assert verifyIsIdentityLensForFields(fields, appliedLens);
+      rewrittenFields = fields;
+    } else {
+      rewrittenFields = null;
+      for (DexField field : fields) {
+        DexField rewrittenField = getRenamedFieldSignature(field, appliedLens);
+        // If rewrittenFields is non-null we have previously seen a change and need to record the
+        // field no matter what.
+        if (rewrittenFields != null) {
+          rewrittenFields.add(rewrittenField);
+          continue;
+        }
+        // If the field has not been rewritten then we can reuse the input set.
+        if (rewrittenField == field) {
+          continue;
+        }
+        // Otherwise add the rewritten field and all previous fields to a new set.
+        rewrittenFields = SetUtils.newIdentityHashSet(fields.size());
+        CollectionUtils.forEachUntilExclusive(fields, rewrittenFields::add, field);
+        rewrittenFields.add(rewrittenField);
+      }
+      if (rewrittenFields == null) {
+        rewrittenFields = fields;
+      } else {
+        rewrittenFields =
+            SetUtils.trimCapacityOfIdentityHashSetIfSizeLessThan(rewrittenFields, fields.size());
+      }
+    }
+    timing.end();
+    return rewrittenFields;
+  }
+
+  private boolean verifyIsIdentityLensForFields(
+      Collection<DexField> fields, GraphLens appliedLens) {
+    for (DexField field : fields) {
+      assert lookupField(field, appliedLens) == field;
+    }
+    return true;
+  }
+
   @SuppressWarnings("unchecked")
   public <T extends DexReference> T rewriteReference(T reference) {
     return rewriteReference(reference, null);
diff --git a/src/main/java/com/android/tools/r8/graph/lens/IdentityGraphLens.java b/src/main/java/com/android/tools/r8/graph/lens/IdentityGraphLens.java
index 2d6b49e..02c6a27 100644
--- a/src/main/java/com/android/tools/r8/graph/lens/IdentityGraphLens.java
+++ b/src/main/java/com/android/tools/r8/graph/lens/IdentityGraphLens.java
@@ -27,6 +27,11 @@
   }
 
   @Override
+  public boolean isIdentityLensForFields(GraphLens codeLens) {
+    return true;
+  }
+
+  @Override
   public boolean isNonIdentityLens() {
     return false;
   }
diff --git a/src/main/java/com/android/tools/r8/graph/lens/NonIdentityGraphLens.java b/src/main/java/com/android/tools/r8/graph/lens/NonIdentityGraphLens.java
index ac9d814..ccbacc3 100644
--- a/src/main/java/com/android/tools/r8/graph/lens/NonIdentityGraphLens.java
+++ b/src/main/java/com/android/tools/r8/graph/lens/NonIdentityGraphLens.java
@@ -190,6 +190,11 @@
   }
 
   @Override
+  public boolean isIdentityLensForFields(GraphLens codeLens) {
+    return this == codeLens;
+  }
+
+  @Override
   public final boolean isNonIdentityLens() {
     return true;
   }
diff --git a/src/main/java/com/android/tools/r8/optimize/MemberRebindingLens.java b/src/main/java/com/android/tools/r8/optimize/MemberRebindingLens.java
index d94e65d..15b6ae9 100644
--- a/src/main/java/com/android/tools/r8/optimize/MemberRebindingLens.java
+++ b/src/main/java/com/android/tools/r8/optimize/MemberRebindingLens.java
@@ -85,6 +85,14 @@
         .build();
   }
 
+  @Override
+  public boolean isIdentityLensForFields(GraphLens codeLens) {
+    if (this == codeLens) {
+      return true;
+    }
+    return getPrevious().isIdentityLensForFields(codeLens);
+  }
+
   public FieldRebindingIdentityLens toRewrittenFieldRebindingLens(
       AppView<? extends AppInfoWithClassHierarchy> appView,
       GraphLens lens,
diff --git a/src/main/java/com/android/tools/r8/shaking/AppInfoWithLiveness.java b/src/main/java/com/android/tools/r8/shaking/AppInfoWithLiveness.java
index 2f12103..05112d6 100644
--- a/src/main/java/com/android/tools/r8/shaking/AppInfoWithLiveness.java
+++ b/src/main/java/com/android/tools/r8/shaking/AppInfoWithLiveness.java
@@ -1186,7 +1186,7 @@
         lens.rewriteReferences(liveTypes),
         lens.rewriteReferences(targetedMethods),
         lens.rewriteReferences(failedMethodResolutionTargets),
-        lens.rewriteReferences(failedFieldResolutionTargets),
+        lens.rewriteFields(failedFieldResolutionTargets, timing),
         lens.rewriteReferences(bootstrapMethods),
         lens.rewriteReferences(virtualMethodsTargetedByInvokeDirect),
         lens.rewriteReferences(liveMethods),
diff --git a/src/main/java/com/android/tools/r8/utils/CollectionUtils.java b/src/main/java/com/android/tools/r8/utils/CollectionUtils.java
index 9a9d27c..5434871 100644
--- a/src/main/java/com/android/tools/r8/utils/CollectionUtils.java
+++ b/src/main/java/com/android/tools/r8/utils/CollectionUtils.java
@@ -12,6 +12,7 @@
 import java.util.Set;
 import java.util.function.Consumer;
 import java.util.function.Function;
+import java.util.function.Predicate;
 
 public class CollectionUtils {
 
@@ -42,6 +43,26 @@
     }
   }
 
+  public static <T> void forEachUntilExclusive(
+      Collection<T> collection, Consumer<T> consumer, T stoppingCriterion) {
+    for (T element : collection) {
+      if (element.equals(stoppingCriterion)) {
+        break;
+      }
+      consumer.accept(element);
+    }
+  }
+
+  public static <T> void forEachUntilExclusive(
+      Collection<T> collection, Consumer<T> consumer, Predicate<? super T> stoppingCriterion) {
+    for (T element : collection) {
+      if (stoppingCriterion.test(element)) {
+        break;
+      }
+      consumer.accept(element);
+    }
+  }
+
   public static <T extends Comparable<T>> Collection<T> sort(Collection<T> items) {
     ArrayList<T> ts = new ArrayList<>(items);
     Collections.sort(ts);
diff --git a/src/main/java/com/android/tools/r8/utils/MapUtils.java b/src/main/java/com/android/tools/r8/utils/MapUtils.java
index 5757218..ed49f66 100644
--- a/src/main/java/com/android/tools/r8/utils/MapUtils.java
+++ b/src/main/java/com/android/tools/r8/utils/MapUtils.java
@@ -12,6 +12,7 @@
 import java.util.IdentityHashMap;
 import java.util.Map;
 import java.util.Map.Entry;
+import java.util.function.BiConsumer;
 import java.util.function.BiFunction;
 import java.util.function.BiPredicate;
 import java.util.function.Consumer;
@@ -39,6 +40,17 @@
     return map.values().iterator().next();
   }
 
+  public static <K, V> void forEachUntilExclusive(
+      Map<K, V> map, BiConsumer<K, V> consumer, K stoppingCriterion) {
+    for (Entry<K, V> entry : map.entrySet()) {
+      K key = entry.getKey();
+      if (key.equals(stoppingCriterion)) {
+        break;
+      }
+      consumer.accept(key, entry.getValue());
+    }
+  }
+
   public static <T, R> Function<T, R> ignoreKey(Supplier<R> supplier) {
     return ignore -> supplier.get();
   }
@@ -133,6 +145,28 @@
     return true;
   }
 
+  public static <K, V> Map<K, V> trimCapacity(Map<K, V> map, IntFunction<Map<K, V>> mapFactory) {
+    Map<K, V> newMap = mapFactory.apply(map.size());
+    newMap.putAll(map);
+    return newMap;
+  }
+
+  public static <K, V> Map<K, V> trimCapacityIfSizeLessThan(
+      Map<K, V> map, IntFunction<Map<K, V>> mapFactory, int expectedSize) {
+    if (map.size() < expectedSize) {
+      return trimCapacity(map, mapFactory);
+    }
+    return map;
+  }
+
+  public static <K, V> Map<K, V> trimCapacityOfIdentityHashMapIfSizeLessThan(
+      Map<K, V> map, int expectedSize) {
+    if (map.size() < expectedSize) {
+      return trimCapacity(map, IdentityHashMap::new);
+    }
+    return map;
+  }
+
   public static <K, V> Map<K, V> unmodifiableForTesting(Map<K, V> map) {
     return InternalOptions.assertionsEnabled() ? Collections.unmodifiableMap(map) : map;
   }
diff --git a/src/main/java/com/android/tools/r8/utils/SetUtils.java b/src/main/java/com/android/tools/r8/utils/SetUtils.java
index a0ecdfc..05af36b 100644
--- a/src/main/java/com/android/tools/r8/utils/SetUtils.java
+++ b/src/main/java/com/android/tools/r8/utils/SetUtils.java
@@ -125,6 +125,16 @@
     return element;
   }
 
+  public static <T> Set<T> trimCapacityOfIdentityHashSetIfSizeLessThan(
+      Set<T> set, int expectedSize) {
+    if (set.size() < expectedSize) {
+      Set<T> newSet = SetUtils.newIdentityHashSet(set.size());
+      newSet.addAll(set);
+      return newSet;
+    }
+    return set;
+  }
+
   public static <T> Set<T> unionIdentityHashSet(Set<T> one, Set<T> other) {
     Set<T> union = Sets.newIdentityHashSet();
     union.addAll(one);
diff --git a/src/main/java/com/android/tools/r8/utils/collections/ProgramMethodSet.java b/src/main/java/com/android/tools/r8/utils/collections/ProgramMethodSet.java
index 612ace4..0ac1e7c 100644
--- a/src/main/java/com/android/tools/r8/utils/collections/ProgramMethodSet.java
+++ b/src/main/java/com/android/tools/r8/utils/collections/ProgramMethodSet.java
@@ -11,9 +11,12 @@
 import com.android.tools.r8.graph.ProgramMethod;
 import com.android.tools.r8.graph.PrunedItems;
 import com.android.tools.r8.graph.lens.GraphLens;
+import com.android.tools.r8.utils.CollectionUtils;
 import com.android.tools.r8.utils.ForEachable;
 import com.google.common.collect.ImmutableMap;
+import java.util.ArrayList;
 import java.util.IdentityHashMap;
+import java.util.List;
 import java.util.Map;
 import java.util.concurrent.ConcurrentHashMap;
 
@@ -84,17 +87,48 @@
   }
 
   public ProgramMethodSet rewrittenWithLens(DexDefinitionSupplier definitions, GraphLens lens) {
-    ProgramMethodSet rewritten = ProgramMethodSet.create(size());
-    forEach(
-        method -> {
-          ProgramMethod newMethod = lens.mapProgramMethod(method, definitions);
-          if (newMethod != null) {
-            assert !newMethod.getDefinition().isObsolete();
-            rewritten.add(newMethod);
+    List<ProgramMethod> elementsToRemove = null;
+    ProgramMethodSet rewrittenMethods = null;
+    for (ProgramMethod method : this) {
+      ProgramMethod rewrittenMethod = lens.mapProgramMethod(method, definitions);
+      if (rewrittenMethod == null) {
+        assert lens.isEnumUnboxerLens();
+        // If everything has been unchanged up until now, then record that we should remove this
+        // method.
+        if (rewrittenMethods == null) {
+          if (elementsToRemove == null) {
+            elementsToRemove = new ArrayList<>();
           }
-        });
-    rewritten.trimCapacityIfSizeLessThan(size());
-    return rewritten;
+          elementsToRemove.add(method);
+        }
+        continue;
+      }
+      if (rewrittenMethod == method) {
+        if (rewrittenMethods != null) {
+          rewrittenMethods.add(rewrittenMethod);
+        }
+      } else {
+        if (rewrittenMethods == null) {
+          rewrittenMethods = ProgramMethodSet.create(size());
+          CollectionUtils.<ProgramMethod>forEachUntilExclusive(
+              this, rewrittenMethods::add, method::isStructurallyEqualTo);
+          if (elementsToRemove != null) {
+            rewrittenMethods.removeAll(elementsToRemove);
+            elementsToRemove = null;
+          }
+        }
+        rewrittenMethods.add(rewrittenMethod);
+      }
+    }
+    if (rewrittenMethods != null) {
+      rewrittenMethods.trimCapacityIfSizeLessThan(size());
+      return rewrittenMethods;
+    } else {
+      if (elementsToRemove != null) {
+        removeAll(elementsToRemove);
+      }
+      return this;
+    }
   }
 
   public ProgramMethodSet withoutPrunedItems(PrunedItems prunedItems) {