#!/usr/bin/env python3
"""Run gcra.lua on a real redis-server and compare with rl_sim.gcra_schedule().

    REDIS_SERVER=/path/to/redis-server python3 test_redis_gcra.py

Starts a throwaway server on 127.0.0.1:6391 (no persistence). Simulated time
is scaled by 1000 before it is sent, so the PX expiry (which Redis measures
in real time) never fires during the test; decisions are scale invariant.
"""
import os
import socket
import subprocess
import threading
import time

import rl_sim as R

EPOCH = 1_758_000_000_000_000        # epoch-sized microsecond offset
SCALE = 1000
PORT = 6391


class Conn:
    def __init__(self):
        self.s = socket.create_connection(("127.0.0.1", PORT))
        self.f = self.s.makefile("rb")

    def cmd(self, *args):
        out = [b"*%d\r\n" % len(args)]
        for a in args:
            a = str(a).encode()
            out.append(b"$%d\r\n%s\r\n" % (len(a), a))
        self.s.sendall(b"".join(out))
        return self.read()

    def read(self):
        line = self.f.readline()
        kind, rest = line[:1], line[1:-2]
        if kind in (b"+", b"-"):
            if kind == b"-":
                raise RuntimeError(rest.decode())
            return rest.decode()
        if kind == b":":
            return int(rest)
        if kind == b"$":
            n = int(rest)
            return None if n < 0 else self.f.read(n + 2)[:-2].decode()
        return [self.read() for _ in range(int(rest))]


def admitted_by_script(conn, sha, key, arrivals):
    res = []
    for t in arrivals:
        r = conn.cmd("EVALSHA", sha, 1, key, R.T * SCALE, R.TAU * SCALE, EPOCH + t * SCALE)
        res.append(r[0] == 1)
    return res


def burst(sha, key, threads=8, per_thread=50, script=True):
    """All threads fire at the same simulated instant; count admissions."""
    now, count, lock = EPOCH, [0], threading.Lock()

    def worker():
        c = Conn()
        for _ in range(per_thread):
            if script:
                ok = c.cmd("EVALSHA", sha, 1, key, R.T * SCALE, R.TAU * SCALE, now)[0] == 1
            else:                                   # racy read-modify-write
                v = c.cmd("GET", key)
                tat = int(v) if v else now
                ok = tat - R.TAU * SCALE <= now
                if ok:
                    c.cmd("SET", key, max(tat, now) + R.T * SCALE)
            if ok:
                with lock:
                    count[0] += 1

    ts = [threading.Thread(target=worker) for _ in range(threads)]
    for t in ts:
        t.start()
    for t in ts:
        t.join()
    return count[0]


def main():
    server = os.environ.get("REDIS_SERVER", "redis-server")
    p = subprocess.Popen([server, "--port", str(PORT), "--bind", "127.0.0.1",
                          "--save", "", "--appendonly", "no"],
                         stdout=subprocess.DEVNULL, env=dict(os.environ, LC_ALL="C"))
    try:
        for _ in range(50):
            try:
                conn = Conn()
                break
            except OSError:
                time.sleep(0.1)
        print("redis", conn.cmd("INFO", "server").split("redis_version:")[1].split()[0])
        src = open(os.path.join(os.path.dirname(__file__), "gcra.lua")).read()
        sha = conn.cmd("SCRIPT", "LOAD", src)
        bad = total = 0
        for s in range(1, 6):
            arr = R.equivalence_trace(s)
            want = [d is not None for d in R.gcra_schedule(arr)]
            got = admitted_by_script(conn, sha, f"eq:{s}", arr)
            bad += sum(a != b for a, b in zip(got, want))
            total += len(arr)
        print(f"gcra.lua vs gcra_schedule: {total} requests, {bad} mismatches")
        runs = [burst(sha, f"atomic:{i}") for i in range(5)]
        print(f"8 threads x 50 at one instant, script: admitted {runs}")
        runs = [burst(sha, f"racy:{i}", script=False) for i in range(5)]
        print(f"8 threads x 50 at one instant, GET then SET: admitted {runs}")
    finally:
        p.terminate()
        p.wait()


if __name__ == "__main__":
    main()
