Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
33 changes: 25 additions & 8 deletions src/nooa/runtime/stream_wrappers.py
Original file line number Diff line number Diff line change
Expand Up @@ -44,25 +44,42 @@ def __init__(
self.mode = getattr(original, "mode", "w")

def write(self, data: str) -> int:
"""Write to contextvar buffer if set, otherwise to original stream."""
"""Write to the task buffer, falling back if a stale buffer was closed."""
buffer = self._buffer_var.get()
if buffer is not None:
return buffer.write(data)
try:
return buffer.write(data)
except ValueError:
# Background tasks inherit contextvars. They can outlive the
# execution capture that installed this buffer, leaving a
# closed StringIO in their copied context. Preserve logging
# and exception reporting by routing those late writes to the
# process stream instead.
if not getattr(buffer, "closed", False):
raise
return self._original.write(data)

def writelines(self, lines: list[str]) -> None:
"""Write multiple lines."""
"""Write multiple lines, falling back if a stale buffer was closed."""
buffer = self._buffer_var.get()
if buffer is not None:
buffer.writelines(lines)
else:
self._original.writelines(lines)
try:
buffer.writelines(lines)
return
except ValueError:
if not getattr(buffer, "closed", False):
raise
self._original.writelines(lines)

def flush(self) -> None:
"""Flush both buffer and original stream."""
"""Flush the active buffer when usable, then the original stream."""
buffer = self._buffer_var.get()
if buffer is not None:
buffer.flush()
try:
buffer.flush()
except ValueError:
if not getattr(buffer, "closed", False):
raise
self._original.flush()

def fileno(self) -> int:
Expand Down
34 changes: 34 additions & 0 deletions tests/runtime/test_stream_wrappers.py
Original file line number Diff line number Diff line change
Expand Up @@ -68,6 +68,17 @@ def test_write_returns_length_from_original(self):
n = stream.write("xyz")
assert n == 3

def test_write_falls_back_to_original_when_inherited_buffer_is_closed(self):
stream, buf_var, original = make_stream()
buf = io.StringIO()
token = buf_var.set(buf)
buf.close()
try:
assert stream.write("late log") == len("late log")
assert original.getvalue() == "late log"
finally:
buf_var.reset(token)


class TestContextVarStreamWritelines:
"""Tests for ContextVarStream.writelines() routing."""
Expand All @@ -88,6 +99,17 @@ def test_writelines_to_original_when_no_buffer(self):
stream.writelines(["x", "y", "z"])
assert original.getvalue() == "xyz"

def test_writelines_falls_back_to_original_when_buffer_is_closed(self):
stream, buf_var, original = make_stream()
buf = io.StringIO()
token = buf_var.set(buf)
buf.close()
try:
stream.writelines(["late", " log"])
assert original.getvalue() == "late log"
finally:
buf_var.reset(token)


class TestContextVarStreamFlush:
"""Tests for ContextVarStream.flush() behavior with and without a buffer."""
Expand Down Expand Up @@ -118,6 +140,18 @@ def test_flush_without_buffer_flushes_original(self):
stream.flush()
mock_original.flush.assert_called_once()

def test_flush_ignores_closed_inherited_buffer(self):
stream, buf_var, original = make_stream()
buf = io.StringIO()
token = buf_var.set(buf)
buf.close()
try:
stream.flush() # must not mask late exception reporting
original.write("still usable")
assert original.getvalue() == "still usable"
finally:
buf_var.reset(token)


class TestContextVarStreamFileno:
"""Tests for ContextVarStream.fileno() delegation to the original stream."""
Expand Down
Loading