Skip to content
Merged
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
123 changes: 83 additions & 40 deletions contextual_orchestrator/server.py
Comment thread
seonghobae marked this conversation as resolved.

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

📝 Info: Disconnected streams logged as successful 200

When a client disconnects mid-stream, _stream_route_completion now returns cleanly, so do_POST still records the analytics event at server.py:5270-5280 with status_code: 200 and response_streamed: True. Aborted deliveries are counted as fully successful streamed responses.

(Refers to this code)

Open in Devin Review

Was this helpful? React with 👍 or 👎 to provide feedback.

Original file line number Diff line number Diff line change
Expand Up @@ -5950,6 +5950,26 @@ def _send_error(
) -> None:
self._send(_error_payload(code, message, {"request_id": uuid.uuid4().hex, **(detail or {})}), status)

def _write_response(self, writer: Callable[[], None]) -> bool:
"""Run a response-writing callback, swallowing a dead-peer disconnect.

A client that gave up waiting (e.g. on a slow orchestration run)
closes its end of the socket before this thread finishes writing.
The write then raises BrokenPipeError/ConnectionError/OSError --
there is nothing left to deliver, so this is not a server error.
Without this guard, that exception propagates out of do_POST's
try block into its own `except Exception: self._send_error(...)`
handler, which calls back into a send method on the same closed
socket and raises again -- uncaught this time, crashing the
request-handling thread (visible as a second, unhandled
BrokenPipeError in server logs after the first).
"""
try:
writer()
return True
except (BrokenPipeError, ConnectionError, OSError):
return False

def _send(
self,
payload: dict[str, Any],
Expand All @@ -5958,46 +5978,64 @@ def _send(
extra_headers: dict[str, str] | None = None,
) -> None:
raw = json.dumps(payload, ensure_ascii=False).encode("utf-8")
self.send_response(status)
self.send_header("content-type", "application/json; charset=utf-8")
self.send_header("content-length", str(len(raw)))
self._send_security_headers()
for name, value in (extra_headers or {}).items():
self.send_header(name, value)
self.end_headers()
self.wfile.write(raw)

def _write() -> None:
self.send_response(status)
self.send_header("content-type", "application/json; charset=utf-8")
self.send_header("content-length", str(len(raw)))
self._send_security_headers()
for name, value in (extra_headers or {}).items():
self.send_header(name, value)
self.end_headers()
self.wfile.write(raw)

self._write_response(_write)

def _send_text(self, payload: str, content_type: str, status: int = 200) -> None:
raw = payload.encode("utf-8")
self.send_response(status)
self.send_header("content-type", content_type)
self.send_header("content-length", str(len(raw)))
self._send_security_headers()
self.end_headers()
self.wfile.write(raw)

def _write() -> None:
self.send_response(status)
self.send_header("content-type", content_type)
self.send_header("content-length", str(len(raw)))
self._send_security_headers()
self.end_headers()
self.wfile.write(raw)

self._write_response(_write)

def _send_sse(self, body: str, status: int = 200) -> None:
raw = body.encode("utf-8")
self.send_response(status)
self.send_header("content-type", "text/event-stream; charset=utf-8")
self.send_header("cache-control", "no-cache")
self.send_header("content-length", str(len(raw)))
self._send_security_headers()
self.end_headers()
self.wfile.write(raw)

def _begin_sse(self) -> None:

def _write() -> None:
self.send_response(status)
self.send_header("content-type", "text/event-stream; charset=utf-8")
self.send_header("cache-control", "no-cache")
self.send_header("content-length", str(len(raw)))
self._send_security_headers()
self.end_headers()
self.wfile.write(raw)

self._write_response(_write)

def _begin_sse(self) -> bool:
# Incremental SSE: no content-length; the connection close delimits the body.
self.send_response(200)
self.send_header("content-type", "text/event-stream; charset=utf-8")
self.send_header("cache-control", "no-cache")
self.send_header("connection", "close")
self._send_security_headers()
self.end_headers()
def _write() -> None:
self.send_response(200)
self.send_header("content-type", "text/event-stream; charset=utf-8")
self.send_header("cache-control", "no-cache")
self.send_header("connection", "close")
self._send_security_headers()
self.end_headers()

def _write_sse(self, frame: str) -> None:
self.wfile.write(frame.encode("utf-8"))
self.wfile.flush()
return self._write_response(_write)

def _write_sse(self, frame: str) -> bool:
def _write() -> None:
self.wfile.write(frame.encode("utf-8"))
self.wfile.flush()

return self._write_response(_write)

def _stream_route_completion(self, orchestrator: Any, security: Any, messages: Any, model_name: str) -> None:
"""Pipe a worker's live deltas out as OpenAI chat.completion.chunk SSE frames."""
Expand All @@ -6017,12 +6055,14 @@ def frame(delta: dict[str, Any], finish: str | None = None) -> str:

security.acquire_run_slot()
try:
self._begin_sse()
self._write_sse(frame({"role": "assistant"}))
if not self._begin_sse() or not self._write_sse(frame({"role": "assistant"})):
return
try:
for delta in orchestrator.stream_route(messages, workflow_run_id=run_id):
self._write_sse(frame({"content": delta}))
self._write_sse(frame({}, finish="stop"))
if not self._write_sse(frame({"content": delta})):
return
Comment on lines 6061 to +6063

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

📝 Info: Generator closed by refcount, not explicit close

On disconnect, the for loop returns without explicitly closing the stream_route generator. CPython refcounting closes it promptly so upstream work stops, but this relies on refcount semantics rather than an explicit close().

Open in Devin Review

Was this helpful? React with 👍 or 👎 to provide feedback.

if not self._write_sse(frame({}, finish="stop")):
return
except ToolFallbackStoppedError as exc:
detail = {
"request_id": uuid.uuid4().hex,
Expand All @@ -6033,12 +6073,15 @@ def frame(delta: dict[str, Any], finish: str | None = None) -> str:
TOOL_FALLBACK_STOPPED_MESSAGE,
detail,
)
self._write_sse(
if not self._write_sse(
f"data: {json.dumps(payload, ensure_ascii=False)}\n\n"
)
self._write_sse(frame({}, finish="error"))
):
return
if not self._write_sse(frame({}, finish="error")):
return
except Exception: # noqa: BLE001 - headers already sent; surface as a terminal error frame
self._write_sse(frame({}, finish="error"))
if not self._write_sse(frame({}, finish="error")):
return
self._write_sse("data: [DONE]\n\n")
finally:
security.release_run_slot()
Expand Down
117 changes: 117 additions & 0 deletions tests/test_http_response_write_disconnect_safety.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,117 @@
"""A disconnected client during response write must not crash the handler thread."""

from __future__ import annotations

from pathlib import Path
import sys

sys.path.insert(0, str(Path(__file__).resolve().parents[1]))

from contextual_orchestrator import ModelAgent, TaskOrchestrator # noqa: E402
from contextual_orchestrator.server import SecurityConfig, build_server # noqa: E402

_TEST_AUTH_TOKEN = "http_response_write_disconnect_safety_token" # noqa: S105


def build() -> TaskOrchestrator:
return TaskOrchestrator(
[ModelAgent("general_agent", "mock-planner", tags=("reasoning", "writing"))]
)


def test_write_response_swallows_a_broken_pipe_from_a_disconnected_client() -> None:
"""``_write_response`` must absorb BrokenPipeError/ConnectionError/OSError.

A client that gives up waiting on a slow response closes its socket
before the write completes. Before this guard existed, that exception
propagated out of do_POST's own `except Exception` handler into a
second call to the same closed socket, which raised again -- uncaught
-- and crashed the request-handling thread. The fix point is
`_write_response`; every `_send*`/`_begin_sse`/`_write_sse` method
routes through it, so testing it directly covers all of them.
"""
server = build_server(build(), port=0, security=SecurityConfig(auth_token=_TEST_AUTH_TOKEN))
try:
handler_cls = server.RequestHandlerClass

def _disconnected_write() -> None:
raise BrokenPipeError("simulated client disconnect mid-write")

# Must return quietly, not raise -- self is unused by the method,
# so a dummy stands in for a real connected-handler instance.
assert handler_cls._write_response(object(), _disconnected_write) is False

def _reset_by_peer() -> None:
raise ConnectionResetError("simulated peer reset")

assert handler_cls._write_response(object(), _reset_by_peer) is False
finally:
server.server_close()


def test_stream_stops_consuming_and_releases_slot_after_disconnect() -> None:
"""A dead SSE peer must stop paid upstream work and release concurrency."""
server = build_server(build(), port=0, security=SecurityConfig(auth_token=_TEST_AUTH_TOKEN))
yielded: list[str] = []

class Orchestrator:
def stream_route(self, messages, workflow_run_id):
for delta in ("first", "second"):
yielded.append(delta)
yield delta

class Security:
acquired = 0
released = 0

def acquire_run_slot(self):
self.acquired += 1

def release_run_slot(self):
self.released += 1

class Handler:
writes = 0

def _begin_sse(self):
return True

def _write_sse(self, _frame):
self.writes += 1
return self.writes < 2

try:
security = Security()
handler = Handler()
server.RequestHandlerClass._stream_route_completion(
handler, Orchestrator(), security, [], "model-group"
)
assert yielded == ["first"]
assert security.acquired == security.released == 1
finally:
server.server_close()


def test_write_response_still_propagates_unrelated_errors() -> None:
"""Only disconnect-shaped errors are swallowed; real bugs must still surface."""
server = build_server(build(), port=0, security=SecurityConfig(auth_token=_TEST_AUTH_TOKEN))
try:
handler_cls = server.RequestHandlerClass

def _broken_serializer() -> None:
raise TypeError("payload was not JSON-serializable")

try:
handler_cls._write_response(object(), _broken_serializer)
except TypeError:
pass
else:
raise AssertionError("expected TypeError to propagate, not be swallowed")
finally:
server.server_close()


if __name__ == "__main__":
test_write_response_swallows_a_broken_pipe_from_a_disconnected_client()
test_write_response_still_propagates_unrelated_errors()
print("ok")
Loading