diff --git a/contextual_orchestrator/server.py b/contextual_orchestrator/server.py index e01483152..3aa547392 100644 --- a/contextual_orchestrator/server.py +++ b/contextual_orchestrator/server.py @@ -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 + 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() diff --git a/tests/test_http_response_write_disconnect_safety.py b/tests/test_http_response_write_disconnect_safety.py new file mode 100644 index 000000000..45177dcfa --- /dev/null +++ b/tests/test_http_response_write_disconnect_safety.py @@ -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")