import z3, json, sys, time
MASK = (1 << 64) - 1
values = json.load(open("/tmp/chrome-rand.json"))
def step(s0, s1, sym):
    a, b = s0, s1
    x = a
    x ^= (x << 23) if sym else (x << 23) & MASK
    x ^= z3.LShR(x, 17) if sym else x >> 17
    x ^= b
    x ^= z3.LShR(b, 26) if sym else b >> 26
    return b, x if sym else x & MASK
for shown in (2, 3, 4):
    solver = z3.Solver()
    i0, i1 = z3.BitVecs("i0 i1", 64)
    s0, s1 = i0, i1
    for v in values[:shown]:
        s0, s1 = step(s0, s1, True)
        solver.add(z3.LShR(s0 + s1, 11) == int(v * 2**53))
    t = time.time(); ok = solver.check() == z3.sat; dt = time.time() - t
    if not ok: print(shown, "unsat", round(dt, 2)); continue
    m = solver.model(); a, b = m[i0].as_long(), m[i1].as_long()
    out = []
    for _ in range(len(values)):
        a, b = step(a, b, False); out.append(((a + b) & MASK) >> 11)
    hits = sum(1 for k in range(shown, len(values)) if out[k] == int(values[k] * 2**53))
    print(shown, "sat", round(dt, 2), "s, predicted", hits, "of", len(values) - shown)
