#!/usr/bin/env python3
"""sic wire protocol v2 daemon — read one frame, exec command, forward stdin."""
import os
import struct
import sys


def read_exactly(fd: int, n: int) -> bytes:
    """Read exactly n bytes from fd, or raise EOFError."""
    buf = bytearray()
    while len(buf) < n:
        chunk = os.read(fd, n - len(buf))
        if not chunk:
            raise EOFError(f"expected {n} bytes, got {len(buf)}")
        buf.extend(chunk)
    return bytes(buf)


def parse_netstring(data: bytes) -> tuple[bytes, bytes]:
    """Parse a DJB netstring from data. Returns (content, rest)."""
    colon = data.find(b":")
    if colon < 0:
        raise ValueError("netstring missing colon")
    # Enforce non-canonical lengths: must be all digits
    if not data[:colon].isdigit():
        raise ValueError("netstring length not all digits")
    ns_len = int(data[:colon])
    # Need at least colon+1 + ns_len + 1 for trailing comma
    if len(data) < colon + 1 + ns_len + 1:
        raise ValueError("netstring truncated")
    content = data[colon + 1 : colon + 1 + ns_len]
    if data[colon + 1 + ns_len : colon + 1 + ns_len + 1] != b",":
        raise ValueError("netstring missing trailing comma")
    # Require exact consumption: no rest after the comma
    rest = data[colon + 1 + ns_len + 1 :]
    if rest:
        raise ValueError("netstring has trailing data after comma")
    return content, rest


def write_all(fd: int, data: bytes) -> None:
    """Write all bytes to fd, looping on short writes."""
    while data:
        n = os.write(fd, data)
        data = data[n:]


def main() -> None:
    # Read 5-byte preamble: magic + len32
    try:
        preamble = read_exactly(sys.stdin.fileno(), 5)
    except EOFError:
        sys.stderr.write("sicd: EOF reading preamble\n")
        sys.exit(1)

    if preamble[0:1] != b"\x00":
        sys.stderr.write("sicd: missing magic byte 0x00\n")
        sys.exit(1)

    ns_len = struct.unpack(">I", preamble[1:5])[0]

    # Read the netstring
    try:
        ns_data = read_exactly(sys.stdin.fileno(), ns_len)
    except EOFError:
        sys.stderr.write("sicd: EOF reading netstring body\n")
        sys.exit(1)

    try:
        content, _ = parse_netstring(ns_data)
    except ValueError as e:
        sys.stderr.write(f"sicd: {e}\n")
        sys.exit(1)

    # Split at first NUL
    nul_pos = content.find(b"\x00")
    if nul_pos < 0:
        sys.stderr.write("sicd: missing NUL separator between command and payload\n")
        sys.exit(1)

    command_bytes = content[:nul_pos]
    payload = content[nul_pos + 1 :]

    # Decode command, validate ASCII
    try:
        command_str = command_bytes.decode("ascii")
    except UnicodeDecodeError:
        sys.stderr.write("sicd: command contains non-ASCII bytes\n")
        sys.exit(1)

    # Split on ASCII space only (not .split())
    argv = command_str.split(" ")

    # Create pipe for child's stdin
    r_fd, w_fd = os.pipe()

    pid = os.fork()
    if pid == 0:
        # Child
        os.close(w_fd)
        os.dup2(r_fd, sys.stdin.fileno())
        os.close(r_fd)
        try:
            os.execvp(argv[0], argv)
        except FileNotFoundError:
            sys.stderr.write(f"sicd: command not found: {argv[0]}\n")
            os._exit(1)
        except Exception as e:
            sys.stderr.write(f"sicd: exec failed: {e}\n")
            os._exit(1)
    else:
        # Parent
        os.close(r_fd)
        # Write payload first
        try:
            write_all(w_fd, payload)
        except BrokenPipeError:
            # Child died before we could write payload; reap and inherit its status
            os.close(w_fd)
            _, status = os.waitpid(pid, 0)
            if os.WIFEXITED(status):
                sys.exit(os.WEXITSTATUS(status))
            else:
                sys.exit(1)
        # Then pump remaining stdin to child
        try:
            while True:
                chunk = os.read(sys.stdin.fileno(), 65536)
                if not chunk:
                    break
                write_all(w_fd, chunk)
        except BrokenPipeError:
            pass
        except OSError:
            pass
        os.close(w_fd)
        # Wait for child
        _, status = os.waitpid(pid, 0)
        if os.WIFEXITED(status):
            sys.exit(os.WEXITSTATUS(status))
        else:
            sys.exit(1)


if __name__ == "__main__":
    main()
