-
Notifications
You must be signed in to change notification settings - Fork 1
fix: stop a disconnected client from crashing the response-write thread #828
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Changes from all commits
42b6e4c
0d71754
bde1975
ab154d2
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
|
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe 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, (Refers to this code) Was this helpful? React with 👍 or 👎 to provide feedback. |
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -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], | ||
|
|
@@ -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.""" | ||
|
|
@@ -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
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. 📝 Info: Generator closed by refcount, not explicit close On disconnect, the 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, | ||
|
|
@@ -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() | ||
|
|
||
| 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") |
Uh oh!
There was an error while loading. Please reload this page.