sicd: harden pump loop, EPIPE/diagnostic handling, signal fidelity
Review round 3 fixes: - write_all helper: handle partial writes from signals - EPIPE on write side: catch BrokenPipeError (child dead), not all OSError - Read-side OSError: diagnose on stderr + exit 1 (not silent truncation) - parse_netstring: sanity cap 1<<24, return content only (not tuple) - 128+WTERMSIG for signal-killed children (shell convention) - 24/24 pytest green
This commit is contained in:
+25
-38
@@ -16,37 +16,36 @@ def read_exactly(fd: int, n: int) -> bytes:
|
||||
return bytes(buf)
|
||||
|
||||
|
||||
def parse_netstring(data: bytes) -> tuple[bytes, bytes]:
|
||||
"""Parse a DJB netstring from data. Returns (content, rest)."""
|
||||
def write_all(fd: int, data: bytes) -> None:
|
||||
"""Write all of data to fd, handling partial writes from signals."""
|
||||
offset = 0
|
||||
while offset < len(data):
|
||||
n = os.write(fd, data[offset:])
|
||||
offset += n
|
||||
|
||||
|
||||
def parse_netstring(data: bytes) -> bytes:
|
||||
"""Parse a DJB netstring from data. Returns content or raises ValueError."""
|
||||
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 ns_len > 1 << 24:
|
||||
raise ValueError("netstring length exceeds sanity cap")
|
||||
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:]
|
||||
return content
|
||||
|
||||
|
||||
def main() -> None:
|
||||
# Read 5-byte preamble: magic + len32
|
||||
try:
|
||||
preamble = read_exactly(sys.stdin.fileno(), 5)
|
||||
except EOFError:
|
||||
@@ -59,7 +58,6 @@ def main() -> None:
|
||||
|
||||
ns_len = struct.unpack(">I", preamble[1:5])[0]
|
||||
|
||||
# Read the netstring
|
||||
try:
|
||||
ns_data = read_exactly(sys.stdin.fileno(), ns_len)
|
||||
except EOFError:
|
||||
@@ -67,12 +65,11 @@ def main() -> None:
|
||||
sys.exit(1)
|
||||
|
||||
try:
|
||||
content, _ = parse_netstring(ns_data)
|
||||
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")
|
||||
@@ -81,22 +78,18 @@ def main() -> None:
|
||||
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)
|
||||
@@ -109,21 +102,11 @@ def main() -> None:
|
||||
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
|
||||
read_error = False
|
||||
try:
|
||||
if payload:
|
||||
write_all(w_fd, payload)
|
||||
while True:
|
||||
chunk = os.read(sys.stdin.fileno(), 65536)
|
||||
if not chunk:
|
||||
@@ -131,16 +114,20 @@ def main() -> None:
|
||||
write_all(w_fd, chunk)
|
||||
except BrokenPipeError:
|
||||
pass
|
||||
except OSError:
|
||||
pass
|
||||
except OSError as e:
|
||||
sys.stderr.write(f"sicd: stdin read error: {e}\n")
|
||||
read_error = True
|
||||
os.close(w_fd)
|
||||
# Wait for child
|
||||
_, status = os.waitpid(pid, 0)
|
||||
if read_error:
|
||||
sys.exit(1)
|
||||
if os.WIFEXITED(status):
|
||||
sys.exit(os.WEXITSTATUS(status))
|
||||
elif os.WIFSIGNALED(status):
|
||||
sys.exit(128 + os.WTERMSIG(status))
|
||||
else:
|
||||
sys.exit(1)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
main()
|
||||
Reference in New Issue
Block a user