diff --git a/platform/util/src/com/intellij/openapi/util/RecursionManager.java b/platform/util/src/com/intellij/openapi/util/RecursionManager.java index 214e7c720416..d9b053acfbb6 100644 --- a/platform/util/src/com/intellij/openapi/util/RecursionManager.java +++ b/platform/util/src/com/intellij/openapi/util/RecursionManager.java @@ -20,6 +20,7 @@ import com.intellij.reference.SoftReference; import com.intellij.util.containers.SoftHashMap; import org.jetbrains.annotations.NonNls; import org.jetbrains.annotations.NotNull; +import org.jetbrains.annotations.Nullable; import java.util.ArrayList; import java.util.LinkedHashMap; @@ -49,34 +50,10 @@ import java.util.Map; public class RecursionManager { private static final Logger LOG = Logger.getInstance("#com.intellij.openapi.util.RecursionManager"); private static final Object NULL = new Object(); - private static final ThreadLocal ourStamp = new ThreadLocal() { + private static final ThreadLocal ourStack = new ThreadLocal() { @Override - protected Integer initialValue() { - return 0; - } - }; - private static final ThreadLocal ourMemoizationStamp = new ThreadLocal() { - @Override - protected Integer initialValue() { - return 0; - } - }; - private static final ThreadLocal ourDepth = new ThreadLocal() { - @Override - protected Integer initialValue() { - return 0; - } - }; - private static final ThreadLocal> ourProgress = new ThreadLocal>() { - @Override - protected LinkedHashMap initialValue() { - return new LinkedHashMap(); - } - }; - private static final ThreadLocal> ourIntermediateCache = new ThreadLocal>() { - @Override - protected Map initialValue() { - return new SoftHashMap(); + protected CalculationStack initialValue() { + return new CalculationStack(); } }; @@ -89,93 +66,46 @@ public class RecursionManager { @Override public T doPreventingRecursion(@NotNull Object key, boolean memoize, Computable computation) { MyKey realKey = new MyKey(id, key); - LinkedHashMap progressMap = ourProgress.get(); - if (progressMap.containsKey(realKey)) { - _prohibitResultCaching(key); + final CalculationStack stack = ourStack.get(); + if (stack.checkReentrancy(realKey)) { return null; } if (memoize) { - SoftReference reference = ourIntermediateCache.get().get(realKey); - if (reference != null) { - if (ourDepth.get() == 0) { - throw new AssertionError("Memoized values with empty stack"); - } - Object o = reference.get(); - if (o != null) { - //noinspection unchecked - return o == NULL ? null : (T)o; - } + Object o = stack.getMemoizedValue(realKey); + if (o != null) { + //noinspection unchecked + return o == NULL ? null : (T)o; } } - if (progressMap.isEmpty() && ourStamp.get() != 0) { - throw new AssertionError("Non-zero stamp with empty stack: " + ourStamp.get()); - } - - checkDepth("1"); - - int sizeBefore = progressMap.size(); - progressMap.put(realKey, ourStamp.get()); - ourDepth.set(ourDepth.get() + 1); - - checkDepth("2"); - - int sizeAfter = progressMap.size(); - if (sizeAfter != sizeBefore + 1) { - LOG.error("Key doesn't lead to the map size increase: " + sizeBefore + " " + sizeAfter + " " + key); - } - - int startStamp = ourMemoizationStamp.get(); + final int sizeBefore = stack.progressMap.size(); + stack.beforeComputation(realKey); + final int sizeAfter = stack.progressMap.size(); + int startStamp = stack.memoizationStamp; try { T result = computation.compute(); - if (memoize && ourMemoizationStamp.get() == startStamp) { - ourIntermediateCache.get().put(realKey, new SoftReference(result == null ? NULL : result)); + if (memoize) { + stack.maybeMemoize(realKey, result == null ? NULL : result, startStamp); } return result; } finally { - if (sizeAfter != progressMap.size()) { - LOG.error("Map size changed: " + progressMap.size() + " " + sizeAfter + " " + key); - } - - ourDepth.set(ourDepth.get() - 1); - Integer value = progressMap.remove(realKey); - - if (sizeBefore != progressMap.size()) { - LOG.error("Map size doesn't decrease: " + progressMap.size() + " " + sizeBefore + " " + key); - } - - if (value == null) { - throw new AssertionError(key + " has changed its equals/hashCode"); - } - - ourStamp.set(value); - if (value == 0) { - ourIntermediateCache.get().clear(); - } - else if (progressMap.isEmpty()) { - ourIntermediateCache.get().clear(); - throw new AssertionError("Non-zero stamp for empty progress map: " + key + ", " + value); - } else { - checkZero(); - } - - checkDepth("3"); + stack.afterComputation(realKey, sizeBefore, sizeAfter); } } @Override public StackStamp markStack() { - final Integer stamp = ourStamp.get(); + final int stamp = ourStack.get().reentrancyCount; return new StackStamp() { @Override public boolean mayCacheNow() { - return Comparing.equal(stamp, ourStamp.get()); + return stamp == ourStack.get().reentrancyCount; } }; } @@ -183,7 +113,7 @@ public class RecursionManager { @Override public List currentStack() { ArrayList result = new ArrayList(); - LinkedHashMap map = ourProgress.get(); + LinkedHashMap map = ourStack.get().progressMap; for (MyKey pair : map.keySet()) { if (pair.first.equals(id)) { result.add(pair.second); @@ -194,51 +124,134 @@ public class RecursionManager { @Override public void prohibitResultCaching(Object since) { - _prohibitResultCaching(since); - ourMemoizationStamp.set(ourMemoizationStamp.get() + 1); + ourStack.get()._prohibitResultCaching(since, id); + ourStack.get().memoizationStamp++; } - private void _prohibitResultCaching(Object since) { - int stamp = ourStamp.get() + 1; - ourStamp.set(stamp); - - checkZero(); - - boolean inLoop = false; - for (Map.Entry entry: ourProgress.get().entrySet()) { - if (inLoop) { - entry.setValue(stamp); - } - else if (entry.getKey().first.equals(id) && entry.getKey().second.equals(since)) { - inLoop = true; - } - } - - checkZero(); - } }; } - private static void checkDepth(String s) { - LinkedHashMap progressMap = ourProgress.get(); - Integer depth = ourDepth.get(); - if (depth != progressMap.size()) { - ourDepth.set(progressMap.size()); - throw new AssertionError("Inconsistent depth " + s + "; depth=" + depth + "; map=" + progressMap); - } - } - - private static void checkZero() { - LinkedHashMap progressMap = ourProgress.get(); - if (!progressMap.isEmpty() && progressMap.get(progressMap.keySet().iterator().next()) != 0) { - throw new AssertionError("Prisoner Zero has escaped"); - } - } - private static class MyKey extends Pair { public MyKey(String first, Object second) { super(first, second); } } + private static class CalculationStack { + private int reentrancyCount; + private int memoizationStamp; + private int depth; + private final LinkedHashMap progressMap = new LinkedHashMap(); + private final Map intermediateCache = new SoftHashMap(); + + boolean checkReentrancy(MyKey realKey) { + if (progressMap.containsKey(realKey)) { + _prohibitResultCaching(realKey.second, realKey.first); + + return true; + } + return false; + } + + @Nullable + Object getMemoizedValue(MyKey realKey) { + SoftReference reference = intermediateCache.get(realKey); + if (reference != null) { + if (depth == 0) { + throw new AssertionError("Memoized values with empty stack"); + } + return reference.get(); + } + return null; + } + + final void beforeComputation(MyKey realKey) { + if (progressMap.isEmpty() && reentrancyCount != 0) { + throw new AssertionError("Non-zero stamp with empty stack: " + reentrancyCount); + } + + checkDepth("1"); + + int sizeBefore = progressMap.size(); + progressMap.put(realKey, reentrancyCount); + depth++; + + checkDepth("2"); + + int sizeAfter = progressMap.size(); + if (sizeAfter != sizeBefore + 1) { + LOG.error("Key doesn't lead to the map size increase: " + sizeBefore + " " + sizeAfter + " " + realKey.second); + } + } + + final void maybeMemoize(MyKey realKey, @NotNull Object result, int startStamp) { + if (memoizationStamp == startStamp) { + intermediateCache.put(realKey, new SoftReference(result)); + } + } + + final void afterComputation(MyKey realKey, int sizeBefore, int sizeAfter) { + if (sizeAfter != progressMap.size()) { + LOG.error("Map size changed: " + progressMap.size() + " " + sizeAfter + " " + realKey.second); + } + + depth--; + Integer value = progressMap.remove(realKey); + + if (sizeBefore != progressMap.size()) { + LOG.error("Map size doesn't decrease: " + progressMap.size() + " " + sizeBefore + " " + realKey.second); + } + + if (value == null) { + throw new AssertionError(realKey.second + " has changed its equals/hashCode"); + } + + reentrancyCount = value; + if (value == 0) { + intermediateCache.clear(); + } + else if (progressMap.isEmpty()) { + intermediateCache.clear(); + throw new AssertionError("Non-zero stamp for empty progress map: " + realKey.second + ", " + value); + } else { + checkZero(); + } + + checkDepth("3"); + + } + + private void _prohibitResultCaching(Object since, String id) { + reentrancyCount++; + + checkZero(); + + boolean inLoop = false; + for (Map.Entry entry: progressMap.entrySet()) { + if (inLoop) { + entry.setValue(reentrancyCount); + } + else if (entry.getKey().first.equals(id) && entry.getKey().second.equals(since)) { + inLoop = true; + } + } + + checkZero(); + } + + private void checkDepth(String s) { + if (depth != progressMap.size()) { + depth = progressMap.size(); + throw new AssertionError("Inconsistent depth " + s + "; depth=" + depth + "; map=" + progressMap); + } + } + + private void checkZero() { + if (!progressMap.isEmpty() && progressMap.get(progressMap.keySet().iterator().next()) != 0) { + throw new AssertionError("Prisoner Zero has escaped"); + } + } + + } + }