This is an automated email from the ASF dual-hosted git repository.
ashb pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/airflow.git
The following commit(s) were added to refs/heads/main by this push:
new f72e6c9befa Fix short-read handling in task-sdk IPC framing (#69253)
f72e6c9befa is described below
commit f72e6c9befa7494467fbc8ceb883148dd9cb61ec
Author: Andrew Chang <[email protected]>
AuthorDate: Fri Aug 14 17:06:13 2026 +0800
Fix short-read handling in task-sdk IPC framing (#69253)
Stream sockets are not required to return the full requested count from
recv(), but the 4-byte length prefix on the supervisor↔subprocess socket
was read with a single sock.recv(4). On a short read, the parsed length
is garbage, which either raises msgspec.DecodeError in the subprocess or
silently deadlocks the supervisor's selector reader (length_needed=0
makes the falsy payload check skip forever, so no frame is ever
dispatched to handle_requests).
Accumulate the header bytes on both sides, mirroring the payload
accumulation that was already correct. Add regression tests using a
socket that returns 2 bytes at a time to reliably reproduce the short
read.
---
task-sdk/src/airflow/sdk/execution_time/comms.py | 12 +++++-
.../src/airflow/sdk/execution_time/supervisor.py | 16 +++++---
.../tests/task_sdk/execution_time/test_comms.py | 39 ++++++++++++++++++
.../task_sdk/execution_time/test_supervisor.py | 48 ++++++++++++++++++++++
4 files changed, 108 insertions(+), 7 deletions(-)
diff --git a/task-sdk/src/airflow/sdk/execution_time/comms.py
b/task-sdk/src/airflow/sdk/execution_time/comms.py
index 5c4ae14dff1..4a26f56e297 100644
--- a/task-sdk/src/airflow/sdk/execution_time/comms.py
+++ b/task-sdk/src/airflow/sdk/execution_time/comms.py
@@ -331,10 +331,18 @@ class CommsDecoder(Generic[ReceiveMsgType, SendMsgType]):
else:
len_bytes = self.socket.recv(4)
- if len_bytes == b"":
+ if not len_bytes:
raise EOFError("Request socket closed before length")
- length = int.from_bytes(len_bytes, byteorder="big")
+ # Stream sockets may return fewer bytes than requested; accumulate the
header.
+ len_buf = bytearray(len_bytes)
+ while len(len_buf) < 4:
+ chunk = self.socket.recv(4 - len(len_buf))
+ if not chunk:
+ raise EOFError(f"Request socket closed mid-length after
{len(len_buf)} of 4 bytes")
+ len_buf.extend(chunk)
+
+ length = int.from_bytes(len_buf, byteorder="big")
buffer = bytearray(length)
mv = memoryview(buffer)
diff --git a/task-sdk/src/airflow/sdk/execution_time/supervisor.py
b/task-sdk/src/airflow/sdk/execution_time/supervisor.py
index 62a1e181e24..4c3ae38d21f 100644
--- a/task-sdk/src/airflow/sdk/execution_time/supervisor.py
+++ b/task-sdk/src/airflow/sdk/execution_time/supervisor.py
@@ -2337,6 +2337,9 @@ def length_prefixed_frame_reader(
gen: Generator[None, _RequestFrame, None], on_close: Callable[[socket],
None]
):
length_needed: int | None = None
+ # Accumulates the 4-byte length header across selector callbacks; stream
+ # sockets may return fewer than the requested 4 bytes in a single recv.
+ header_buffer = bytearray()
# This will hold our accumulated/partial binary frame if it doesn't come
in a single read
buffer: memoryview | None = None
# position in the buffer to store next read
@@ -2347,16 +2350,19 @@ def length_prefixed_frame_reader(
next(gen)
def cb(sock: socket):
- nonlocal buffer, length_needed, pos
+ nonlocal buffer, length_needed, pos, header_buffer
if length_needed is None:
- # Read the 32bit length of the frame
- bytes = sock.recv(4)
- if bytes == b"":
+ chunk = sock.recv(4 - len(header_buffer))
+ if not chunk:
return False
+ header_buffer.extend(chunk)
+ if len(header_buffer) < 4:
+ return True
- length_needed = int.from_bytes(bytes, byteorder="big")
+ length_needed = int.from_bytes(header_buffer, byteorder="big")
buffer = memoryview(bytearray(length_needed))
+ header_buffer = bytearray()
if length_needed and buffer:
n = sock.recv_into(buffer[pos:])
if n == 0:
diff --git a/task-sdk/tests/task_sdk/execution_time/test_comms.py
b/task-sdk/tests/task_sdk/execution_time/test_comms.py
index 05dd682d890..665e5418f5b 100644
--- a/task-sdk/tests/task_sdk/execution_time/test_comms.py
+++ b/task-sdk/tests/task_sdk/execution_time/test_comms.py
@@ -340,3 +340,42 @@ class TestCommsDecoder:
server2.join(timeout=2)
assert result is not None
+
+ def test_read_frame_recovers_from_short_read_on_header(self):
+ msg = VariableResult(key="k", value="v", type="VariableResult")
+ payload = msgspec.msgpack.encode(_ResponseFrame(0, msg.model_dump(),
None))
+ wire = len(payload).to_bytes(4, byteorder="big") + payload
+
+ class ChunkedSocket:
+ def __init__(self, data: bytes, chunk_size: int):
+ self._data = data
+ self._chunk_size = chunk_size
+ self._pos = 0
+
+ def setblocking(self, flag):
+ pass
+
+ def recv(self, n):
+ remaining = self._data[self._pos :]
+ if not remaining:
+ return b""
+ chunk = remaining[: min(n, self._chunk_size)]
+ self._pos += len(chunk)
+ return chunk
+
+ def recv_into(self, buf):
+ remaining = self._data[self._pos :]
+ if not remaining:
+ return 0
+ take = min(len(buf), self._chunk_size, len(remaining))
+ buf[:take] = remaining[:take]
+ self._pos += take
+ return take
+
+ sock = ChunkedSocket(wire, chunk_size=2)
+ decoder = CommsDecoder(socket=sock, log=None)
+
+ result = decoder._get_response()
+ assert isinstance(result, VariableResult)
+ assert result.key == "k"
+ assert result.value == "v"
diff --git a/task-sdk/tests/task_sdk/execution_time/test_supervisor.py
b/task-sdk/tests/task_sdk/execution_time/test_supervisor.py
index 909f4942ef8..c9a977c9793 100644
--- a/task-sdk/tests/task_sdk/execution_time/test_supervisor.py
+++ b/task-sdk/tests/task_sdk/execution_time/test_supervisor.py
@@ -4703,3 +4703,51 @@ class TestMakeBufferedSocketReader:
finally:
r.close()
w.close()
+
+
+class TestLengthPrefixedFrameReader:
+ def test_recovers_from_short_read_on_header(self):
+ received: list[_RequestFrame] = []
+
+ def collecting_gen():
+ while True:
+ frame = yield
+ received.append(frame)
+
+ payload = msgspec.msgpack.encode(_RequestFrame(id=42, body={"key":
"foo"}))
+ wire = len(payload).to_bytes(4, byteorder="big") + payload
+
+ class ChunkedSocket:
+ def __init__(self, data: bytes, chunk_size: int):
+ self._data = data
+ self._chunk_size = chunk_size
+ self._pos = 0
+
+ def recv(self, n):
+ remaining = self._data[self._pos :]
+ if not remaining:
+ return b""
+ chunk = remaining[: min(n, self._chunk_size)]
+ self._pos += len(chunk)
+ return chunk
+
+ def recv_into(self, buf):
+ remaining = self._data[self._pos :]
+ if not remaining:
+ return 0
+ take = min(len(buf), self._chunk_size, len(remaining))
+ buf[:take] = remaining[:take]
+ self._pos += take
+ return take
+
+ sock = ChunkedSocket(wire, chunk_size=2)
+ on_close = MagicMock()
+ cb, _ = supervisor.length_prefixed_frame_reader(collecting_gen(),
on_close=on_close)
+
+ for _ in range(len(wire) + 1):
+ if not cb(sock):
+ break
+ if received:
+ break
+
+ assert received == [_RequestFrame(id=42, body={"key": "foo"})]