import z3, struct, subprocess, json, sys

def run_node(expression):
    return json.loads(subprocess.check_output(["node", "-e", f"console.log(JSON.stringify({expression}))"]))

def top53(value):
    return int(value * (1 << 53))

def step(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 to_double(state0):
    return (state0 >> 11) / (1 << 53)

def step_int(state0, state1):
    mask = (1 << 64) - 1
    s1, s0 = state0, state1
    s1 ^= (s1 << 23) & mask
    s1 ^= s1 >> 17
    s1 ^= s0
    s1 ^= s0 >> 26
    return s0, s1 & mask

values = run_node("Array.from({length: 10}, Math.random)")
observed, hidden = values[:5], values[5:]
solver = z3.Solver()
state0, state1 = z3.BitVecs("state0 state1", 64)
for value in reversed(observed):
    state0, state1 = step(state0, state1)
    solver.add(z3.LShR(state0, 11) == top53(value))
if solver.check() != z3.sat:
    sys.exit("unsat")
model = solver.model()
start0 = model[z3.BitVec("state0", 64)].as_long()
start1 = model[z3.BitVec("state1", 64)].as_long()
a, b = start0, start1
for _ in observed:
    a, b = step_int(a, b)
predicted = []
for _ in hidden:
    a, b = step_int(a, b)
print("observed", len(observed))
print("hidden  ", hidden)
