import java.lang.reflect.Field;
import java.util.*;
import java.util.concurrent.ConcurrentHashMap;
import java.util.concurrent.atomic.AtomicInteger;

/*
 * Clock-free probes of java.util.concurrent.ConcurrentHashMap internals.
 * Needs: --add-opens java.base/java.util.concurrent=ALL-UNNAMED
 *
 *   bins     bin-length histogram vs Poisson
 *   reuse    fraction of nodes cloned when the table doubles
 *   treeify  when a colliding bin becomes a TreeBin
 *   race     lost updates: get+put vs merge; check-then-act vs putIfAbsent
 *   recurse  nested computeIfAbsent in the same bin / another bin
 */
public class ChmProbe {
    static final Field TABLE, NEXT, FIRST;
    static {
        try {
            TABLE = ConcurrentHashMap.class.getDeclaredField("table");
            TABLE.setAccessible(true);
            Class<?> node = Class.forName("java.util.concurrent.ConcurrentHashMap$Node");
            NEXT = node.getDeclaredField("next");
            NEXT.setAccessible(true);
            FIRST = Class.forName("java.util.concurrent.ConcurrentHashMap$TreeBin").getDeclaredField("first");
            FIRST.setAccessible(true);
        } catch (ReflectiveOperationException e) { throw new RuntimeException(e); }
    }

    static Object[] table(ConcurrentHashMap<?, ?> m) throws Exception { return (Object[]) TABLE.get(m); }

    static int spread(int h) { return (h ^ (h >>> 16)) & 0x7fffffff; }

    /* length of a plain list bin; -1 for TreeBin or other special nodes */
    static int binLen(Object head) throws Exception {
        if (head == null) return 0;
        if (!head.getClass().getSimpleName().equals("Node")) return -1;
        int n = 0;
        for (Object e = head; e != null; e = NEXT.get(e)) n++;
        return n;
    }

    static long[] histogram(ConcurrentHashMap<?, ?> m, int maxLen) throws Exception {
        long[] h = new long[maxLen + 2];
        for (Object b : table(m)) {
            int len = binLen(b);
            if (len < 0 || len > maxLen) h[maxLen + 1]++; else h[len]++;
        }
        return h;
    }

    static double poisson(double lam, int k) {
        double p = Math.exp(-lam);
        for (int i = 1; i <= k; i++) p *= lam / i;
        return p;
    }

    static void bins(long seed) throws Exception {
        final int T = 1 << 17, maxLen = 8;
        SplittableRandom rnd = new SplittableRandom(seed);
        ConcurrentHashMap<Integer, Integer> m = new ConcurrentHashMap<>();
        HashSet<Integer> seen = new HashSet<>();
        long[] mix = new long[maxLen + 2];
        int samples = 0;
        int lo = 3 * T / 8, hi = 3 * T / 4;           /* one doubling cycle at table size T */
        while (m.size() < hi - 1) {
            int k = rnd.nextInt();
            if (!seen.add(k)) continue;
            m.put(k, k);
            int n = m.size();
            if (n == T / 2) {
                long[] h = histogram(m, maxLen);
                System.out.printf("bins seed=%d table=%d n=%d lambda=%.4f%n", seed, table(m).length, n, (double) n / T);
                for (int i = 0; i <= maxLen; i++)
                    System.out.printf("  len=%d measured=%.8f poisson=%.8f%n", i, (double) h[i] / T, poisson(0.5, i));
                System.out.printf("  len>%d measured=%d bins%n", maxLen, h[maxLen + 1]);
            }
            if (n > lo && n % 1024 == 0 && table(m).length == T) {
                long[] h = histogram(m, maxLen);
                for (int i = 0; i < mix.length; i++) mix[i] += h[i];
                samples++;
            }
        }
        System.out.printf("cycle seed=%d table=%d samples=%d n in (%d,%d)%n", seed, T, samples, lo, hi);
        for (int i = 0; i <= maxLen; i++)
            System.out.printf("  len=%d measured=%.8f%n", i, (double) mix[i] / ((double) samples * T));
    }

    static void reuse(long seed, int T) throws Exception {
        SplittableRandom rnd = new SplittableRandom(seed);
        ConcurrentHashMap<Integer, Integer> m = new ConcurrentHashMap<>(T / 2);
        int threshold = T - (T >>> 2);
        while (m.size() < threshold - 1) m.put(rnd.nextInt(), 0);
        Object[] old = table(m);
        if (old.length != T) throw new IllegalStateException("table " + old.length);
        Set<Object> oldNodes = Collections.newSetFromMap(new IdentityHashMap<>());
        int treeBins = 0;
        for (Object b : old) {
            if (binLen(b) < 0) { treeBins++; b = FIRST.get(b); }
            for (Object e = b; e != null; e = NEXT.get(e)) oldNodes.add(e);
        }
        int k;
        do { k = rnd.nextInt(); } while (m.containsKey(k));
        m.put(k, 0);                                  /* reaches sizeCtl: triggers transfer */
        Object[] nt = table(m);
        long reused = 0;
        for (Object b : nt) {
            if (binLen(b) < 0) b = FIRST.get(b);
            for (Object e = b; e != null; e = NEXT.get(e))
                if (oldNodes.contains(e)) reused++;
        }
        long total = oldNodes.size();
        System.out.printf("reuse seed=%d T=%d->%d nodes=%d tree_bins=%d reused=%d cloned=%d cloned_frac=%.4f%n",
                seed, old.length, nt.length, total, treeBins, reused, total - reused, (double) (total - reused) / total);
    }

    static final class Collide implements Comparable<Collide> {
        final int id;
        Collide(int id) { this.id = id; }
        @Override public int hashCode() { return 42; }
        @Override public boolean equals(Object o) { return o instanceof Collide c && c.id == id; }
        @Override public int compareTo(Collide o) { return Integer.compare(id, o.id); }
    }

    static void treeify() throws Exception {
        ConcurrentHashMap<Collide, Integer> m = new ConcurrentHashMap<>();
        for (int i = 1; i <= 12; i++) {
            m.put(new Collide(i), i);
            Object[] t = table(m);
            Object head = t[spread(42) & (t.length - 1)];
            System.out.printf("treeify keys=%d table=%d bin=%s%n", i, t.length, head.getClass().getSimpleName());
        }
    }

    static void race(int run) throws Exception {
        final int N = 1_000_000;
        for (String mode : new String[] {"get+put", "merge"}) {
            ConcurrentHashMap<String, Integer> m = new ConcurrentHashMap<>();
            m.put("k", 0);
            Runnable r = () -> {
                for (int i = 0; i < N; i++) {
                    if (mode.equals("merge")) m.merge("k", 1, Integer::sum);
                    else m.put("k", m.get("k") + 1);
                }
            };
            Thread a = new Thread(r), b = new Thread(r);
            a.start(); b.start(); a.join(); b.join();
            System.out.printf("race run=%d mode=%s expected=%d final=%d lost=%d%n", run, mode, 2 * N, m.get("k"), 2 * N - m.get("k"));
        }
        final int K = 200_000;
        for (String mode : new String[] {"contains+put", "putIfAbsent"}) {
            ConcurrentHashMap<Integer, Integer> m = new ConcurrentHashMap<>();
            AtomicInteger wins = new AtomicInteger();
            Runnable r = () -> {
                for (int i = 0; i < K; i++) {
                    if (mode.equals("putIfAbsent")) {
                        if (m.putIfAbsent(i, 1) == null) wins.incrementAndGet();
                    } else if (!m.containsKey(i)) {
                        m.put(i, 1);
                        wins.incrementAndGet();
                    }
                }
            };
            Thread a = new Thread(r), b = new Thread(r);
            a.start(); b.start(); a.join(); b.join();
            System.out.printf("race run=%d mode=%s keys=%d winners=%d double_winners=%d%n", run, mode, K, wins.get(), wins.get() - K);
        }
    }

    static void recurse() {
        ConcurrentHashMap<Integer, Integer> m = new ConcurrentHashMap<>(16);
        m.put(-1, 0);                                  /* force lazy table initialization */
        int n;
        try { n = table(m).length; } catch (Exception e) { throw new RuntimeException(e); }
        int[][] pairs = {{1, 1 + n}, {1, 2}};
        String[] names = {"same-bin", "other-bin"};
        for (int p = 0; p < 2; p++) {
            int k1 = pairs[p][0], k2 = pairs[p][1];
            try {
                Integer v = m.computeIfAbsent(k1, x -> m.computeIfAbsent(k2, y -> 7) + 1);
                System.out.printf("recurse table=%d %s k1=%d k2=%d -> ok, value=%d, k2 present=%b%n", n, names[p], k1, k2, v, m.containsKey(k2));
            } catch (IllegalStateException e) {
                System.out.printf("recurse table=%d %s k1=%d k2=%d -> %s: %s%n", n, names[p], k1, k2, e.getClass().getSimpleName(), e.getMessage());
            }
        }
    }

    public static void main(String[] args) throws Exception {
        System.out.println("java.version=" + System.getProperty("java.version") + " cpus=" + Runtime.getRuntime().availableProcessors());
        switch (args[0]) {
            case "bins" -> { for (long s = 1; s <= 3; s++) bins(s); }
            case "reuse" -> { for (int T : new int[] {1 << 10, 1 << 16, 1 << 18}) for (long s = 1; s <= 3; s++) reuse(s, T); }
            case "treeify" -> treeify();
            case "race" -> { for (int r = 1; r <= 5; r++) race(r); }
            case "recurse" -> recurse();
            default -> throw new IllegalArgumentException(args[0]);
        }
    }
}
