#!/usr/bin/env python3
"""
Reproduction harness for three SavedDataStorage bugs in the vanilla Minecraft: Java Edition
dedicated server. Drives an unmodified server.jar over stdin and watches world/data/ from outside.
Nothing in the game is patched or instrumented from the inside.

  torn-write    kill -9 while a data/*.dat file is being rewritten in place. The file is left
                truncated, fails to load on the next start and is then overwritten with defaults.
  lost-write    one failed write (file made read-only for a single save) is never retried, not by
                later saves and not by a clean /stop, although the fault is gone by then.
  ticket-churn  chunk_tickets.dat of every dimension is rewritten on every save of an idle server.

Needs Python 3.8+, Java 25 for 26.x, nothing else. Optional: strace on Linux (--strace).

  python3 saveddata_repro.py --accept-eula torn-write
  python3 saveddata_repro.py --accept-eula --jar server.jar --java /path/to/java all
"""
import argparse, gzip, hashlib, json, os, platform, random, re, shutil, signal, string, subprocess
import sys, threading, time, urllib.request

MANIFEST = "https://piston-meta.mojang.com/mc/game/version_manifest_v2.json"
IS_WIN = platform.system() == "Windows"
T0 = time.monotonic()


def log(msg=""):
    print(f"[{time.monotonic() - T0:8.3f}] {msg}", flush=True)


def fetch_jar(version, dest):
    if os.path.exists(dest):
        return dest
    log(f"downloading server.jar for {version}")
    man = json.load(urllib.request.urlopen(MANIFEST, timeout=60))
    if version == "latest":
        version = man["latest"]["release"]
    meta_url = next(v["url"] for v in man["versions"] if v["id"] == version)
    dl = json.load(urllib.request.urlopen(meta_url, timeout=60))["downloads"]["server"]
    data = urllib.request.urlopen(dl["url"], timeout=300).read()
    if hashlib.sha1(data).hexdigest() != dl["sha1"]:
        sys.exit("sha1 mismatch on downloaded server.jar")
    with open(dest, "wb") as f:
        f.write(data)
    return dest


class Server:
    def __init__(self, a, workdir, tag):
        self.a, self.dir, self.tag = a, workdir, tag
        self.lines, self.cv, self.proc = [], threading.Condition(), None

    def start(self):
        cmd = [self.a.java, f"-Xmx{self.a.heap}", "-jar", self.a.jar, "nogui"]
        if self.a.strace:
            cmd = ["strace", "-f", "-qq", "-tt", "-o", os.path.join(self.dir, f"strace-{self.tag}.log"),
                   "-e", "trace=openat,rename,renameat,renameat2,unlink,unlinkat"] + cmd
        kw = {} if IS_WIN else {"start_new_session": True}
        self.proc = subprocess.Popen(cmd, cwd=self.dir, stdin=subprocess.PIPE, stdout=subprocess.PIPE,
                                     stderr=subprocess.STDOUT, text=True, bufsize=1, errors="replace", **kw)
        threading.Thread(target=self._pump, daemon=True).start()
        self.wait_for(r"Done \(", 600)
        log(f"server up (pid {self.proc.pid}, run '{self.tag}')")
        return self

    def _pump(self):
        with open(os.path.join(self.dir, f"console-{self.tag}.log"), "w") as out:
            for line in self.proc.stdout:
                out.write(line[:2000].rstrip("\n") + "\n")
                with self.cv:
                    self.lines.append(line[:2000].rstrip("\n"))
                    self.cv.notify_all()

    def mark(self):
        with self.cv:
            return len(self.lines)

    def wait_for(self, pattern, timeout=120, since=0):
        rx, end = re.compile(pattern), time.monotonic() + timeout
        with self.cv:
            i = since
            while True:
                while i < len(self.lines):
                    if rx.search(self.lines[i]):
                        return self.lines[i]
                    i += 1
                left = end - time.monotonic()
                if left <= 0 or (self.proc.poll() is not None and i >= len(self.lines)):
                    raise TimeoutError(f"never saw /{pattern}/ in server output")
                self.cv.wait(min(left, 0.5))

    def grep(self, pattern, since=0):
        rx = re.compile(pattern)
        with self.cv:
            return [l for l in self.lines[since:] if rx.search(l)]

    def cmd(self, line):
        self.proc.stdin.write(line + "\n")
        self.proc.stdin.flush()

    def run(self, line, expect, timeout=120):
        m = self.mark()
        self.cmd(line)
        return self.wait_for(expect, timeout, m)

    def save(self, flush=True):
        return self.run("save-all flush" if flush else "save-all", r"Saved the game", 600)

    def kill9(self):
        if IS_WIN:
            self.proc.kill()
        else:
            os.killpg(os.getpgid(self.proc.pid), signal.SIGKILL)
        self.proc.wait(30)

    def stop(self):
        self.cmd("stop")
        self.proc.wait(300)


class Watcher(threading.Thread):
    """Polls stat() on a set of files and records every change of (size, inode). Optionally fires a
    callback the first time `victim` is seen smaller than kill_below bytes."""

    def __init__(self, paths, victim=None, kill_below=None, on_hit=None):
        super().__init__(daemon=True)
        self.paths, self.victim, self.kill_below, self.on_hit = paths, victim, kill_below, on_hit
        self.events, self.stop_flag, self.fired = [], threading.Event(), None

    @staticmethod
    def snap(p):
        try:
            s = os.stat(p)
            return (s.st_size, s.st_ino)
        except OSError:
            return (None, None)

    def run(self):
        last = {p: self.snap(p) for p in self.paths}
        t0 = time.monotonic()
        for p, v in last.items():
            self.events.append((0.0, p, v))
        while not self.stop_flag.is_set():
            for p in self.paths:
                v = self.snap(p)
                if v != last[p]:
                    t = (time.monotonic() - t0) * 1000
                    self.events.append((t, p, v))
                    last[p] = v
                    if (self.on_hit and not self.fired and p == self.victim and v[0] is not None
                            and v[0] < self.kill_below):
                        self.fired = t
                        self.on_hit()
                        return
            time.sleep(0.0002)

    def finish(self):
        self.stop_flag.set()
        self.join(5)
        return self.events

    def report(self, root):
        for p in self.paths:
            ev = [(t, v) for t, q, v in self.events if q == p]
            sizes = [v[0] for _, v in ev if v[0] is not None]
            inodes = {v[1] for _, v in ev if v[1] is not None}
            changes = ev[1:]
            window = (changes[-1][0] - changes[0][0]) if len(changes) > 1 else 0.0
            log(f"  {os.path.relpath(p, root):<40} start={ev[0][1][0]} min_seen={min(sizes) if sizes else None} "
                f"end={ev[-1][1][0]} size_changes={len(changes)} inodes={len(inodes)} "
                f"{'REPLACED (new inode)' if len(inodes) > 1 else 'IN PLACE (same inode)'} window={window:.1f}ms")


def gzip_state(path):
    size = os.path.getsize(path)
    if size == 0:
        return "EMPTY (0 bytes)"
    try:
        n = 0
        with gzip.open(path, "rb") as f:
            while True:
                b = f.read(1 << 20)
                if not b:
                    break
                n += len(b)
        return f"valid gzip, {size} bytes on disk, {n} bytes of NBT"
    except Exception as e:
        return f"CORRUPT, {size} bytes on disk: {type(e).__name__}: {e}"


def prepare(a, name):
    d = os.path.join(a.workdir, name)
    if os.path.exists(d):
        shutil.rmtree(d)
    os.makedirs(d)
    with open(os.path.join(d, "eula.txt"), "w") as f:
        f.write("eula=true\n")
    with open(os.path.join(d, "server.properties"), "w") as f:
        f.write("\n".join([
            "server-ip=127.0.0.1", f"server-port={random.randint(20000, 40000)}",
            "level-type=minecraft\\:flat", "generate-structures=false",
            "pause-when-empty-seconds=-1", "max-tick-time=-1", "view-distance=2", "simulation-distance=2",
        ]) + "\n")
    return d


def find_data(d, filename):
    for root, _, files in os.walk(os.path.join(d, "world")):
        if filename in files and os.sep + "data" in root:
            return os.path.join(root, filename)
    sys.exit(f"{filename} not found under {d}/world")


def show_strace(d, tag, names):
    p = os.path.join(d, f"strace-{tag}.log")
    if not os.path.exists(p):
        return
    log(f"  strace ({os.path.basename(p)}), calls touching {', '.join(names)}:")
    shown = 0
    for line in open(p, errors="replace"):
        if any(n in line for n in names) and ("O_WRONLY" in line or "O_RDWR" in line or "rename" in line or "unlink" in line):
            log("    " + line.strip()[:260])
            shown += 1
            if shown >= 24:
                break


def probe(s, command, label):
    """Run a command and log whatever the server answers. Tolerant: output format is not parsed."""
    m = s.mark()
    s.cmd(command)
    try:
        line = s.wait_for(r"\]: ", 5, m)
        log(f"  /{command} ({label}) -> " + line.split("]: ", 1)[-1][:100])
    except TimeoutError:
        log(f"  /{command} ({label}) -> no answer")


def torn_write(a):
    log("=== torn-write ===")
    d = prepare(a, "torn-write")
    s = Server(a, d, "1-populate").start()
    rnd = random.Random(1)
    log(f"filling the scoreboard: {a.objectives} objectives with {a.name_kb} KB display names")
    for i in range(a.objectives):
        name = "".join(rnd.choices(string.ascii_lowercase + string.digits, k=a.name_kb * 1024))
        s.cmd(f'scoreboard objectives add bulk{i} dummy "{name}"')
    s.run("scoreboard objectives add marker dummy", r"Created new objective \[marker\]", 600)
    s.save()
    victim = find_data(d, "scoreboard.dat")
    level = os.path.join(d, "world", "level.dat")
    base = os.path.getsize(victim)
    log(f"baseline {os.path.relpath(victim, d)}: {gzip_state(victim)}")

    log("step A: one more save with no kill, watching scoreboard.dat against level.dat")
    s.run("scoreboard objectives add dirty1 dummy", r"Created new objective \[dirty1\]")
    w = Watcher([victim, level])
    w.start()
    s.save()
    time.sleep(0.3)
    w.finish()
    w.report(d)

    probe(s, "time query gametime", "before kill")
    verdicts = []
    for attempt in range(1, a.attempts + 1):
        log(f"step B (attempt {attempt}): dirty the scoreboard, save-all, SIGKILL as soon as the file shrinks")
        s.run(f"scoreboard objectives add dirty{attempt + 1} dummy", r"Created new objective")
        w = Watcher([victim], victim, base // 2, s.kill9)
        w.start()
        s.cmd("save-all")
        w.join(60)
        if not w.fired:
            w.finish()
            log("  write window missed, trying again")
            continue
        log(f"  killed {w.fired:.1f} ms after the watcher armed; size at kill was "
            f"{[v[0] for _, _, v in w.events][-1]} of {base} bytes")
        state = gzip_state(victim)
        log(f"  scoreboard.dat after kill: {state}")
        shutil.copy(victim, os.path.join(d, f"evidence-scoreboard-after-kill-{attempt}.dat"))
        show_strace(d, s.tag, ["scoreboard.dat", "level"])

        s = Server(a, d, f"2-restart-{attempt}").start()
        errs = s.grep(r"Error loading saved data|Failed to parse saved data")
        for e in errs:
            log("  server log: " + e[:220])
        listing = s.run("scoreboard objectives list", r"There are no objectives|There are \d+ objective")
        log("  /scoreboard objectives list -> " + listing.split("]: ", 1)[-1][:80])
        probe(s, "time query gametime", "after restart")
        s.stop()
        log(f"  scoreboard.dat after clean stop: {gzip_state(victim)}")
        lost = "no objectives" in listing
        verdicts.append(lost)
        log(f"  RESULT: {'REPRODUCED, all scoreboard data lost' if lost else 'not reproduced'}")
        break
    return bool(verdicts and verdicts[-1])


def lost_write(a):
    log("=== lost-write ===")
    if not IS_WIN and os.geteuid() == 0:
        log("  skipped: root ignores file permissions, run as a normal user")
        return None
    d = prepare(a, "lost-write")
    s = Server(a, d, "1").start()
    s.run("scoreboard objectives add before dummy", r"Created new objective")
    s.save()
    victim = find_data(d, "scoreboard.dat")
    m0 = os.stat(victim).st_mtime_ns
    log(f"baseline saved: {gzip_state(victim)}")

    os.chmod(victim, 0o444)
    log("scoreboard.dat made read-only, adding objective 'after' and saving")
    s.run("scoreboard objectives add after dummy", r"Created new objective")
    mark = s.mark()
    s.save()
    time.sleep(0.5)
    for e in s.grep(r"Could not save data|AccessDenied|FileSystemException", mark)[:2]:
        log("  server log: " + e[:200])
    failed = bool(s.grep(r"Could not save data", mark))

    os.chmod(victim, 0o644)
    log("fault cleared (file writable again), saving again")
    s.save()
    time.sleep(0.5)
    m1 = os.stat(victim).st_mtime_ns
    log(f"  scoreboard.dat rewritten by the next save-all: {'yes' if m1 != m0 else 'NO'}")
    s.stop()
    m2 = os.stat(victim).st_mtime_ns
    log(f"  scoreboard.dat rewritten by clean /stop:        {'yes' if m2 != m0 else 'NO'}")

    s = Server(a, d, "2-restart").start()
    listing = s.run("scoreboard objectives list", r"There are no objectives|There are \d+ objective")
    log("  after restart: " + listing.split("]: ", 1)[-1][:120])
    s.stop()
    lost = failed and "[after]" not in listing
    log(f"  RESULT: {'REPRODUCED, objective [after] was never written' if lost else 'not reproduced'}")
    return lost


def ticket_churn(a):
    log("=== ticket-churn ===")
    d = prepare(a, "ticket-churn")
    s = Server(a, d, "1").start()
    s.save()
    files = []
    for root, _, fs in os.walk(os.path.join(d, "world")):
        if os.sep + "data" in root:
            files += [os.path.join(root, f) for f in fs if f.endswith(".dat")]
    def digest(f):
        return hashlib.sha256(open(f, "rb").read()).hexdigest()

    last = {f: (os.stat(f).st_mtime_ns, digest(f)) for f in files}
    rewrites = {f: 0 for f in files}
    identical = {f: 0 for f in files}
    log(f"idle server, no players, {a.saves} x save-all flush, {a.gap}s apart")
    for _ in range(a.saves):
        time.sleep(a.gap)
        s.save()
        time.sleep(0.3)
        for f in files:
            m, h = os.stat(f).st_mtime_ns, digest(f)
            if m != last[f][0]:
                rewrites[f] += 1
                identical[f] += h == last[f][1]
                last[f] = (m, h)
    s.stop()
    for f in sorted(files):
        log(f"  {os.path.relpath(f, os.path.join(d, 'world')):<68} {os.path.getsize(f):>5} B  "
            f"rewritten {rewrites[f]}/{a.saves}, byte-identical {identical[f]}/{a.saves}")
    churn = [f for f in files if identical[f] == a.saves]
    log(f"  RESULT: {'REPRODUCED, ' + str(len(churn)) + ' file(s) rewritten with identical bytes on every save' if churn else 'not reproduced'}")
    return bool(churn)


def main():
    ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
    ap.add_argument("scenario", choices=["torn-write", "lost-write", "ticket-churn", "all"])
    ap.add_argument("--accept-eula", action="store_true", help="you agree to https://aka.ms/MinecraftEULA (writes eula.txt)")
    ap.add_argument("--version", default="26.3", help="version to download if --jar is not given (or 'latest')")
    ap.add_argument("--jar", help="path to an existing vanilla server.jar")
    ap.add_argument("--java", default=os.path.join(os.environ["JAVA_HOME"], "bin", "java") if os.environ.get("JAVA_HOME") else "java")
    ap.add_argument("--heap", default="2G")
    ap.add_argument("--workdir", default="repro-work")
    ap.add_argument("--objectives", type=int, default=200, help="torn-write: number of bulk objectives")
    ap.add_argument("--name-kb", type=int, default=30, help="torn-write: display name size per objective, KB")
    ap.add_argument("--attempts", type=int, default=5)
    ap.add_argument("--saves", type=int, default=5, help="ticket-churn: number of saves")
    ap.add_argument("--gap", type=float, default=3.0, help="ticket-churn: seconds between saves")
    ap.add_argument("--strace", action="store_true", help="Linux: run the JVM under strace and print the open/rename calls")
    a = ap.parse_args()
    if not a.accept_eula:
        sys.exit("The server will not start without the Minecraft EULA. Read https://aka.ms/MinecraftEULA and pass --accept-eula.")
    a.workdir = os.path.abspath(a.workdir)
    os.makedirs(a.workdir, exist_ok=True)
    a.jar = os.path.abspath(a.jar) if a.jar else fetch_jar(a.version, os.path.join(a.workdir, f"server-{a.version}.jar"))
    ver = subprocess.run([a.java, "-version"], capture_output=True, text=True).stderr.splitlines()[0]
    log(f"{platform.platform()} | {ver} | {os.path.basename(a.jar)}")
    todo = {"torn-write": [torn_write], "lost-write": [lost_write], "ticket-churn": [ticket_churn],
            "all": [torn_write, lost_write, ticket_churn]}[a.scenario]
    results = {f.__name__: f(a) for f in todo}
    log("summary: " + ", ".join(f"{k}={'REPRODUCED' if v else 'skipped' if v is None else 'not reproduced'}" for k, v in results.items()))
    sys.exit(0 if all(v is not False for v in results.values()) else 1)


if __name__ == "__main__":
    main()
