import z3, json, subprocess, sys, time

MASK = (1 << 64) - 1

def forward_symbolic(state0, state1):
    s1, s0 = state0, state1
    s1 ^= s1 << 23
    s1 ^= z3.LShR(s1, 17)
    s1 ^= s0
    s1 ^= z3.LShR(s0, 26)
    return s0, s1

def undo_xor_right(value, bits):
    result = value
    for _ in range(64 // bits + 1):
        result = value ^ (result >> bits)
    return result & MASK

def undo_xor_left(value, bits):
    result = value
    for _ in range(64 // bits + 1):
        result = (value ^ (result << bits)) & MASK
    return result

def backward(state0, state1):
    previous1 = state0
    mixed = state1 ^ previous1 ^ (previous1 >> 26)
    mixed = undo_xor_right(mixed, 17)
    previous0 = undo_xor_left(mixed, 23)
    return previous0, previous1

def to_double(state0):
    return (state0 >> 11) / (1 << 53)

def predict(observed, count):
    solver = z3.Solver()
    initial0, initial1 = z3.BitVecs("initial0 initial1", 64)
    state0, state1 = initial0, initial1
    for value in reversed(observed):
        state0, state1 = forward_symbolic(state0, state1)
        solver.add(z3.LShR(state0, 11) == int(value * (1 << 53)))
    if solver.check() != z3.sat:
        return None
    model = solver.model()
    state0, state1 = model[initial0].as_long(), model[initial1].as_long()
    upcoming = []
    for _ in range(count):
        upcoming.append(to_double(state0))
        state0, state1 = backward(state0, state1)
    return upcoming

if __name__ == "__main__":
    shown = int(sys.argv[1]) if len(sys.argv) > 1 else 5
    trials = int(sys.argv[2]) if len(sys.argv) > 2 else 1
    results = []
    for _ in range(trials):
        values = json.loads(subprocess.check_output(["node", "-e", "console.log(JSON.stringify(Array.from({length: 30}, Math.random)))"]))
        observed, hidden = values[:shown], values[shown : shown + 10]
        started = time.time()
        guess = predict(observed, 10)
        elapsed = time.time() - started
        exact = 0 if guess is None else sum(1 for g, h in zip(guess, hidden) if g == h)
        results.append({"solved": guess is not None, "exact": exact, "seconds": round(elapsed, 3)})
    print(json.dumps({"shown": shown, "trials": trials, "all10": sum(r["exact"] == 10 for r in results), "unsat": sum(not r["solved"] for r in results), "medianSeconds": sorted(r["seconds"] for r in results)[trials // 2]}))
