fix(server): catch client disconnect errors in response write path
This commit is contained in:
16
server.py
16
server.py
@@ -164,7 +164,7 @@ def _build_csp_report_only_policy() -> str:
|
||||
|
||||
from api.auth import check_auth
|
||||
from api.config import HOST, PORT, STATE_DIR, SESSION_DIR, DEFAULT_WORKSPACE
|
||||
from api.helpers import j, get_profile_cookie
|
||||
from api.helpers import j, get_profile_cookie, _CLIENT_DISCONNECT_ERRORS
|
||||
from api.profiles import set_request_profile, clear_request_profile
|
||||
from api.routes import handle_delete, handle_get, handle_patch, handle_post, handle_put
|
||||
from api.startup import auto_install_agent_deps, fix_credential_permissions
|
||||
@@ -317,8 +317,13 @@ class Handler(BaseHTTPRequestHandler):
|
||||
# reconnect races; do not convert it into a misleading server 500.
|
||||
return
|
||||
except Exception as e:
|
||||
if isinstance(e, _CLIENT_DISCONNECT_ERRORS):
|
||||
return
|
||||
print(f'[webui] ERROR {self.command} {self.path}\n' + traceback.format_exc(), flush=True)
|
||||
return j(self, {'error': 'Internal server error'}, status=500)
|
||||
try:
|
||||
j(self, {'error': 'Internal server error'}, status=500)
|
||||
except Exception:
|
||||
pass
|
||||
finally:
|
||||
clear_request_profile()
|
||||
|
||||
@@ -348,8 +353,13 @@ class Handler(BaseHTTPRequestHandler):
|
||||
# reconnect races; do not convert it into a misleading server 500.
|
||||
return
|
||||
except Exception as e:
|
||||
if isinstance(e, _CLIENT_DISCONNECT_ERRORS):
|
||||
return
|
||||
print(f'[webui] ERROR {self.command} {self.path}\n' + traceback.format_exc(), flush=True)
|
||||
return j(self, {'error': 'Internal server error'}, status=500)
|
||||
try:
|
||||
j(self, {'error': 'Internal server error'}, status=500)
|
||||
except Exception:
|
||||
pass
|
||||
finally:
|
||||
clear_request_profile()
|
||||
|
||||
|
||||
279
tests/test_broken_pipe_cascade.py
Normal file
279
tests/test_broken_pipe_cascade.py
Normal file
@@ -0,0 +1,279 @@
|
||||
"""
|
||||
Test for BrokenPipeError + SSL BAD_LENGTH cascading failure fix.
|
||||
|
||||
When a client disconnects mid-response:
|
||||
1. First write raises BrokenPipeError (or ssl.SSLError on TLS)
|
||||
2. Exception handler tries to send 500 JSON through same broken socket
|
||||
3. Without the fix, this produces a second exception and noisy traceback
|
||||
|
||||
The fix:
|
||||
- helpers._safe_write() catches _CLIENT_DISCONNECT_ERRORS (including ssl.SSLError)
|
||||
- server.py exception handlers wrap their 500-response j() calls in try/except
|
||||
"""
|
||||
import ssl
|
||||
import unittest
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import sys
|
||||
import os
|
||||
|
||||
# Ensure project root is on path
|
||||
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
|
||||
from api.helpers import j, t, _safe_write, _CLIENT_DISCONNECT_ERRORS
|
||||
|
||||
|
||||
class MockHandler:
|
||||
"""Minimal mock of BaseHTTPRequestHandler for testing response helpers."""
|
||||
|
||||
def __init__(self, write_raises=None, end_headers_raises=None):
|
||||
self._write_raises = write_raises
|
||||
self._end_headers_raises = end_headers_raises
|
||||
self.headers = {}
|
||||
self._wfile = MagicMock()
|
||||
self._sent_headers = []
|
||||
self._response_status = None
|
||||
|
||||
@property
|
||||
def wfile(self):
|
||||
return self._wfile
|
||||
|
||||
def send_response(self, status):
|
||||
self._response_status = status
|
||||
|
||||
def send_header(self, key, value):
|
||||
self._sent_headers.append((key, value))
|
||||
|
||||
def end_headers(self):
|
||||
if self._end_headers_raises:
|
||||
raise self._end_headers_raises
|
||||
|
||||
def _wfile_write(self, data):
|
||||
if self._write_raises:
|
||||
raise self._write_raises
|
||||
return len(data)
|
||||
|
||||
|
||||
class TestSafeWrite(unittest.TestCase):
|
||||
"""Test _safe_write swallows client disconnect errors silently."""
|
||||
|
||||
def _make_handler(self, end_headers_raises=None, write_raises=None):
|
||||
handler = MockHandler(
|
||||
end_headers_raises=end_headers_raises,
|
||||
write_raises=write_raises,
|
||||
)
|
||||
handler.wfile.write = handler._wfile_write
|
||||
return handler
|
||||
|
||||
def test_safe_write_success(self):
|
||||
handler = self._make_handler()
|
||||
_safe_write(handler, b"hello")
|
||||
self.assertEqual(handler._response_status, None) # send_response not called by _safe_write
|
||||
|
||||
def test_safe_write_broken_pipe(self):
|
||||
handler = self._make_handler(write_raises=BrokenPipeError())
|
||||
# Should NOT raise
|
||||
_safe_write(handler, b"hello")
|
||||
|
||||
def test_safe_write_connection_reset(self):
|
||||
handler = self._make_handler(write_raises=ConnectionResetError())
|
||||
_safe_write(handler, b"hello")
|
||||
|
||||
def test_safe_write_connection_aborted(self):
|
||||
handler = self._make_handler(write_raises=ConnectionAbortedError())
|
||||
_safe_write(handler, b"hello")
|
||||
|
||||
def test_safe_write_timeout(self):
|
||||
handler = self._make_handler(write_raises=TimeoutError())
|
||||
_safe_write(handler, b"hello")
|
||||
|
||||
def test_safe_write_ssl_bad_length(self):
|
||||
handler = self._make_handler(
|
||||
write_raises=ssl.SSLError("[BAD_LENGTH] write failed")
|
||||
)
|
||||
_safe_write(handler, b"hello")
|
||||
|
||||
def test_safe_write_end_headers_broken_pipe(self):
|
||||
handler = self._make_handler(end_headers_raises=BrokenPipeError())
|
||||
_safe_write(handler, b"hello")
|
||||
|
||||
def test_safe_write_end_headers_ssl_error(self):
|
||||
handler = self._make_handler(
|
||||
end_headers_raises=ssl.SSLError("[BAD_LENGTH]")
|
||||
)
|
||||
_safe_write(handler, b"hello")
|
||||
|
||||
def test_safe_write_other_error_propagates(self):
|
||||
handler = self._make_handler(write_raises=ValueError("unexpected"))
|
||||
with self.assertRaises(ValueError):
|
||||
_safe_write(handler, b"hello")
|
||||
|
||||
|
||||
class TestJsonHelper(unittest.TestCase):
|
||||
"""Test j() helper with disconnect errors."""
|
||||
|
||||
def _make_handler(self, end_headers_raises=None, write_raises=None):
|
||||
handler = MockHandler(
|
||||
end_headers_raises=end_headers_raises,
|
||||
write_raises=write_raises,
|
||||
)
|
||||
handler.wfile.write = handler._wfile_write
|
||||
return handler
|
||||
|
||||
def test_j_success(self):
|
||||
handler = self._make_handler()
|
||||
j(handler, {"ok": True}, status=200)
|
||||
self.assertEqual(handler._response_status, 200)
|
||||
self.assertTrue(any(h[0] == "Content-Type" for h in handler._sent_headers))
|
||||
|
||||
def test_j_broken_pipe_on_write(self):
|
||||
handler = self._make_handler(write_raises=BrokenPipeError())
|
||||
# Should NOT raise — headers sent, write fails silently
|
||||
j(handler, {"ok": True}, status=200)
|
||||
self.assertEqual(handler._response_status, 200)
|
||||
|
||||
def test_j_ssl_error_on_write(self):
|
||||
handler = self._make_handler(
|
||||
write_raises=ssl.SSLError("[BAD_LENGTH] write failed")
|
||||
)
|
||||
j(handler, {"ok": True}, status=200)
|
||||
self.assertEqual(handler._response_status, 200)
|
||||
|
||||
def test_j_broken_pipe_on_end_headers(self):
|
||||
handler = self._make_handler(end_headers_raises=BrokenPipeError())
|
||||
j(handler, {"ok": True}, status=200)
|
||||
self.assertEqual(handler._response_status, 200)
|
||||
|
||||
|
||||
class TestTextHelper(unittest.TestCase):
|
||||
"""Test t() helper with disconnect errors."""
|
||||
|
||||
def _make_handler(self, end_headers_raises=None, write_raises=None):
|
||||
handler = MockHandler(
|
||||
end_headers_raises=end_headers_raises,
|
||||
write_raises=write_raises,
|
||||
)
|
||||
handler.wfile.write = handler._wfile_write
|
||||
return handler
|
||||
|
||||
def test_t_success(self):
|
||||
handler = self._make_handler()
|
||||
t(handler, "hello", status=200)
|
||||
self.assertEqual(handler._response_status, 200)
|
||||
|
||||
def test_t_broken_pipe(self):
|
||||
handler = self._make_handler(write_raises=BrokenPipeError())
|
||||
t(handler, "hello", status=200)
|
||||
self.assertEqual(handler._response_status, 200)
|
||||
|
||||
def test_t_ssl_error(self):
|
||||
handler = self._make_handler(
|
||||
write_raises=ssl.SSLError("[BAD_LENGTH]")
|
||||
)
|
||||
t(handler, "hello", status=200)
|
||||
self.assertEqual(handler._response_status, 200)
|
||||
|
||||
|
||||
class TestClientDisconnectErrorsTuple(unittest.TestCase):
|
||||
"""Verify _CLIENT_DISCONNECT_ERRORS includes all expected types."""
|
||||
|
||||
def test_includes_broken_pipe(self):
|
||||
self.assertIn(BrokenPipeError, _CLIENT_DISCONNECT_ERRORS)
|
||||
|
||||
def test_includes_connection_reset(self):
|
||||
self.assertIn(ConnectionResetError, _CLIENT_DISCONNECT_ERRORS)
|
||||
|
||||
def test_includes_connection_aborted(self):
|
||||
self.assertIn(ConnectionAbortedError, _CLIENT_DISCONNECT_ERRORS)
|
||||
|
||||
def test_includes_timeout(self):
|
||||
self.assertIn(TimeoutError, _CLIENT_DISCONNECT_ERRORS)
|
||||
|
||||
def test_includes_ssl_error(self):
|
||||
self.assertIn(ssl.SSLError, _CLIENT_DISCONNECT_ERRORS)
|
||||
|
||||
def test_excludes_broad_oserror(self):
|
||||
"""OSError is too broad — it masks real errors like file-not-found."""
|
||||
self.assertNotIn(OSError, _CLIENT_DISCONNECT_ERRORS)
|
||||
|
||||
|
||||
class TestServerDisconnectHandling(unittest.TestCase):
|
||||
"""Test server.py skips 500 response when client disconnects."""
|
||||
|
||||
def _make_handler(self, route_raises=None):
|
||||
"""Build a Handler with mocked socket and route."""
|
||||
from server import Handler
|
||||
handler = Handler.__new__(Handler)
|
||||
handler.command = "GET"
|
||||
handler.path = "/api/test"
|
||||
handler._req_t0 = 0.0
|
||||
handler.headers = {}
|
||||
handler.wfile = MagicMock()
|
||||
handler.wfile.write = MagicMock()
|
||||
handler.send_response = MagicMock()
|
||||
handler.send_header = MagicMock()
|
||||
handler.end_headers = MagicMock()
|
||||
handler._route_raises = route_raises
|
||||
return handler
|
||||
|
||||
def test_do_get_skips_500_on_broken_pipe(self):
|
||||
from server import Handler
|
||||
handler = self._make_handler(route_raises=BrokenPipeError())
|
||||
|
||||
def _fake_handle_get(self, parsed):
|
||||
raise self._route_raises
|
||||
|
||||
# Patch handle_get to raise BrokenPipeError
|
||||
import server as _server_mod
|
||||
orig_handle_get = _server_mod.handle_get
|
||||
_server_mod.handle_get = _fake_handle_get
|
||||
try:
|
||||
Handler.do_GET(handler)
|
||||
finally:
|
||||
_server_mod.handle_get = orig_handle_get
|
||||
|
||||
# send_response should NEVER be called for the 500 — client is gone
|
||||
handler.send_response.assert_not_called()
|
||||
|
||||
def test_handle_write_skips_500_on_connection_reset(self):
|
||||
from server import Handler
|
||||
handler = self._make_handler(route_raises=ConnectionResetError())
|
||||
handler.command = "POST"
|
||||
|
||||
def _fake_route(self, parsed):
|
||||
raise self._route_raises
|
||||
|
||||
import server as _server_mod
|
||||
orig_check_auth = _server_mod.check_auth
|
||||
_server_mod.check_auth = lambda h, p: True
|
||||
try:
|
||||
Handler._handle_write(handler, _fake_route)
|
||||
finally:
|
||||
_server_mod.check_auth = orig_check_auth
|
||||
|
||||
handler.send_response.assert_not_called()
|
||||
|
||||
def test_do_get_sends_500_on_real_error(self):
|
||||
from server import Handler
|
||||
handler = self._make_handler(route_raises=ValueError("real bug"))
|
||||
|
||||
def _fake_handle_get(self, parsed):
|
||||
raise self._route_raises
|
||||
|
||||
import server as _server_mod
|
||||
orig_handle_get = _server_mod.handle_get
|
||||
orig_check_auth = _server_mod.check_auth
|
||||
_server_mod.handle_get = _fake_handle_get
|
||||
_server_mod.check_auth = lambda h, p: True
|
||||
try:
|
||||
Handler.do_GET(handler)
|
||||
finally:
|
||||
_server_mod.handle_get = orig_handle_get
|
||||
_server_mod.check_auth = orig_check_auth
|
||||
|
||||
# Should send 500 for real errors
|
||||
handler.send_response.assert_called_once_with(500)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user