diff --git a/gateway/sicd b/gateway/sicd index 7406587..53f7e6b 100755 --- a/gateway/sicd +++ b/gateway/sicd @@ -1,99 +1,128 @@ #!/usr/bin/env python3 -"""sicd — sic wire protocol daemon, reference implementation. - -Reads one wire frame from stdin, peels off the outer layer, and execs -the command with the payload + trailing stdin forwarded to the child. -""" - +"""sic wire protocol v2 daemon — read one frame, exec command, forward stdin.""" import os import struct import sys -def read_exact(n: int) -> bytes: - """Read exactly n bytes from stdin, or raise EOFError.""" - buf = b"" +def read_exactly(fd: int, n: int) -> bytes: + """Read exactly n bytes from fd, or raise EOFError.""" + buf = bytearray() while len(buf) < n: - chunk = sys.stdin.buffer.read(n - len(buf)) + chunk = os.read(fd, n - len(buf)) if not chunk: - raise EOFError("unexpected EOF") - buf += chunk - return buf + 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, return (content, rest).""" - # Find the colon separating length from content + """Parse a DJB netstring from data. Returns (content, rest).""" colon = data.find(b":") - if colon == -1: - raise ValueError("netstring: missing colon") - try: - length = int(data[:colon]) - except ValueError: - raise ValueError("netstring: invalid length") - content_start = colon + 1 - content_end = content_start + length - if content_end + 1 > len(data): - raise ValueError("netstring: truncated") - if data[content_end:content_end+1] != b",": - raise ValueError("netstring: missing trailing comma") - return data[content_start:content_end], data[content_end+1:] + 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 main() -> int: +def main() -> None: + # Read 5-byte preamble: magic + len32 try: - # Read preamble: magic byte + 4-byte big-endian length - preamble = read_exact(5) - if preamble[0:1] != b"\x00": - sys.stderr.write("sicd: missing magic byte\n") - return 1 - ns_len = struct.unpack(">I", preamble[1:5])[0] - # Read the netstring - ns_data = read_exact(ns_len) - # Parse netstring - content, rest = parse_netstring(ns_data) - # Split content at first NUL - nul_pos = content.find(b"\x00") - if nul_pos == -1: - sys.stderr.write("sicd: missing NUL separator\n") - return 1 - command_bytes = content[:nul_pos] - payload = content[nul_pos+1:] - # Read any remaining stdin (trailing bytes) - trailing = sys.stdin.buffer.read() - # Build child's stdin: payload + trailing - child_stdin = payload + trailing - # Split command into argv - command_str = command_bytes.decode("ascii", errors="replace") - argv = command_str.split() - if not argv: - sys.stderr.write("sicd: empty command\n") - return 1 - # Fork and exec - pid = os.fork() - if pid == 0: - # Child - # Redirect stdin to pipe with child_stdin - r, w = os.pipe() - os.write(w, child_stdin) - os.close(w) - os.dup2(r, 0) - os.close(r) - os.execvp(argv[0], argv) - # If exec fails - sys.stderr.write(f"sicd: exec failed: {argv[0]}\n") - os._exit(1) - else: - # Parent - _, status = os.waitpid(pid, 0) - if os.WIFEXITED(status): - return os.WEXITSTATUS(status) - else: - return 1 - except (EOFError, ValueError, OSError) as e: + 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") - return 1 + 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 + os.write(w_fd, payload) + # Then pump remaining stdin to child + try: + while True: + chunk = os.read(sys.stdin.fileno(), 65536) + if not chunk: + break + os.write(w_fd, chunk) + 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__": - sys.exit(main()) + main()