"""Unit tests for the host control server (issue #468, chunk 1). Mostly exercises the pure `dispatch()` (socket-free, like the orchestrator server tests), plus a real-socket round-trip through `BrokerClient` that proves the full sign -> POST -> verify -> act seam over HTTP. """ from __future__ import annotations import io import json import secrets import threading import unittest import urllib.error import urllib.request from unittest.mock import MagicMock, patch from bot_bottle.orchestrator.broker import ( BrokerAuthError, LaunchBroker, LaunchRequest, StubBroker, sign_request, ) from bot_bottle.orchestrator.broker_client import BrokerClient from bot_bottle.orchestrator.host_server import ( MAX_BODY_BYTES, Handler, HostControlServer, broker_secret_from_env, dispatch, main, make_host_server, ) def _body(obj: object) -> bytes: return json.dumps(obj).encode() class _RaisingBroker(LaunchBroker): """A broker whose backend launch always fails — exercises the 502 path (an operational backend failure, distinct from a fail-closed provenance 401).""" def _launch(self, req: LaunchRequest) -> None: raise RuntimeError("docker down") def _teardown(self, req: LaunchRequest) -> None: raise RuntimeError("docker down") class TestDispatch(unittest.TestCase): def setUp(self) -> None: self.secret = secrets.token_bytes(16) self.broker = StubBroker(self.secret) def _token(self, **kwargs: object) -> str: return sign_request(LaunchRequest(**kwargs), self.secret) # type: ignore[arg-type] def test_health(self) -> None: status, payload = dispatch(self.broker, "GET", "/health", b"") self.assertEqual(200, status) self.assertEqual("ok", payload["status"]) def test_broker_launch_verifies_and_acts(self) -> None: token = self._token( op="launch", bottle_id="b1", source_ip="10.243.0.1", image_ref="img", slot=2, ) status, payload = dispatch(self.broker, "POST", "/broker", _body({"token": token})) self.assertEqual(200, status) self.assertEqual("launch", payload["op"]) self.assertEqual("b1", payload["bottle_id"]) self.assertEqual("img", payload["image_ref"]) self.assertEqual(2, payload["slot"]) self.assertEqual(["b1"], [r.bottle_id for r in self.broker.launched]) def test_broker_teardown_acts(self) -> None: token = self._token(op="teardown", bottle_id="b1") status, _ = dispatch(self.broker, "POST", "/broker", _body({"token": token})) self.assertEqual(200, status) self.assertEqual(["b1"], [r.bottle_id for r in self.broker.torn_down]) def test_forged_token_is_401_and_nothing_acted(self) -> None: forged = sign_request( LaunchRequest(op="launch", bottle_id="b1"), secrets.token_bytes(16)) status, payload = dispatch(self.broker, "POST", "/broker", _body({"token": forged})) self.assertEqual(401, status) self.assertIn("broker auth failed", str(payload["error"])) self.assertEqual([], self.broker.launched) # fail-closed: never launched def test_backend_failure_is_502(self) -> None: broker = _RaisingBroker(self.secret) token = self._token(op="launch", bottle_id="b1", image_ref="img") status, payload = dispatch(broker, "POST", "/broker", _body({"token": token})) self.assertEqual(502, status) self.assertIn("backend launch failed", str(payload["error"])) def test_missing_token_is_400(self) -> None: status, _ = dispatch(self.broker, "POST", "/broker", _body({})) self.assertEqual(400, status) def test_bad_json_is_400(self) -> None: status, _ = dispatch(self.broker, "POST", "/broker", b"{not json") self.assertEqual(400, status) def test_empty_body_is_missing_token_400(self) -> None: # Empty body parses to {} (no token) → 400, never reaching the broker. status, _ = dispatch(self.broker, "POST", "/broker", b"") self.assertEqual(400, status) self.assertEqual([], self.broker.launched) def test_non_object_body_is_400(self) -> None: status, _ = dispatch(self.broker, "POST", "/broker", b"[1, 2]") self.assertEqual(400, status) def test_unknown_route_404(self) -> None: status, _ = dispatch(self.broker, "GET", "/nope", b"") self.assertEqual(404, status) def test_trailing_slash_normalized(self) -> None: status, _ = dispatch(self.broker, "GET", "/health/", b"") self.assertEqual(200, status) class TestBrokerSecretFromEnv(unittest.TestCase): def test_reads_hex_secret(self) -> None: s = secrets.token_bytes(16) self.assertEqual(s, broker_secret_from_env({"BOT_BOTTLE_BROKER_SECRET": s.hex()})) def test_unset_is_none(self) -> None: self.assertIsNone(broker_secret_from_env({})) def test_invalid_hex_is_none(self) -> None: self.assertIsNone(broker_secret_from_env({"BOT_BOTTLE_BROKER_SECRET": "not-hex"})) class TestSeamRoundTrip(unittest.TestCase): """The whole point of chunk 1: a request signed by the orchestrator side is POSTed to a real host control server, verified there, and acted on — over HTTP, not an in-process call.""" def _serve(self, broker: LaunchBroker) -> BrokerClient: server = make_host_server(broker, "127.0.0.1", 0) self.addCleanup(server.server_close) threading.Thread(target=server.serve_forever, daemon=True).start() self.addCleanup(server.shutdown) host, port = server.server_address[0], server.server_address[1] return BrokerClient(f"http://{host}:{port}") def test_sign_post_verify_act_over_http(self) -> None: secret = secrets.token_bytes(16) broker = StubBroker(secret) client = self._serve(broker) req = LaunchRequest( op="launch", bottle_id="b1", source_ip="10.0.0.1", image_ref="img", slot=1) got = client.submit(sign_request(req, secret)) self.assertEqual(req, got) # the controller echoes the verified request self.assertEqual(["b1"], [r.bottle_id for r in broker.launched]) def test_forged_token_raises_broker_auth_error_over_http(self) -> None: secret = secrets.token_bytes(16) broker = StubBroker(secret) client = self._serve(broker) forged = sign_request( LaunchRequest(op="launch", bottle_id="b1"), secrets.token_bytes(16)) with self.assertRaises(BrokerAuthError): client.submit(forged) self.assertEqual([], broker.launched) # fail-closed across the wire class TestRequestLimits(unittest.TestCase): """The privileged listener must not let a caller that can merely reach the socket (no signed token) exhaust it via an oversized declared body.""" def _base(self) -> str: self.broker = StubBroker(secrets.token_bytes(16)) server = make_host_server(self.broker, "127.0.0.1", 0) self.addCleanup(server.server_close) threading.Thread(target=server.serve_forever, daemon=True).start() self.addCleanup(server.shutdown) host, port = server.server_address[0], server.server_address[1] return f"http://{host}:{port}" def test_oversized_body_is_rejected_before_acting(self) -> None: base = self._base() big = b"x" * (MAX_BODY_BYTES + 1) req = urllib.request.Request( f"{base}/broker", data=big, method="POST", headers={"Content-Type": "application/json"}) with self.assertRaises(urllib.error.HTTPError) as cm: urllib.request.urlopen(req, timeout=5) self.assertEqual(413, cm.exception.code) self.assertEqual([], self.broker.launched) # never reached the broker class TestServeUnit(unittest.TestCase): """Drive `Handler._serve` directly (no socket). The real per-request handler runs in a daemon thread whose coverage/trace data is lost, so the bounded-body and error paths are exercised here in the main thread instead.""" def _handler(self, broker: LaunchBroker, headers: dict[str, str], body: bytes = b"") -> tuple[Handler, MagicMock]: server = HostControlServer.__new__(HostControlServer) server.broker = broker h = Handler.__new__(Handler) h.server = server h.headers = headers # type: ignore[assignment] — dict is a valid .get() stand-in h.path = "/broker" h.rfile = io.BytesIO(body) h.wfile = io.BytesIO() send_response = MagicMock() h.send_response = send_response # type: ignore[method-assign] h.send_header = MagicMock() # type: ignore[method-assign] h.end_headers = MagicMock() # type: ignore[method-assign] return h, send_response def test_oversized_content_length_is_413(self) -> None: broker = StubBroker(secrets.token_bytes(16)) h, send_response = self._handler(broker, {"Content-Length": str(MAX_BODY_BYTES + 1)}) h.do_POST() # exercises do_POST -> _serve send_response.assert_called_once_with(413) self.assertEqual([], broker.launched) # rejected before the broker def test_invalid_content_length_is_400(self) -> None: h, send_response = self._handler(StubBroker(secrets.token_bytes(16)), {"Content-Length": "not-a-number"}) h._serve("POST") send_response.assert_called_once_with(400) def test_valid_request_dispatches_200(self) -> None: secret = secrets.token_bytes(16) broker = StubBroker(secret) body = _body({"token": sign_request( LaunchRequest(op="teardown", bottle_id="b1"), secret)}) h, send_response = self._handler(broker, {"Content-Length": str(len(body))}, body) h._serve("POST") send_response.assert_called_once_with(200) self.assertEqual(["b1"], [r.bottle_id for r in broker.torn_down]) def test_dispatch_exception_becomes_500(self) -> None: # dispatch is total, but the handler still guards it: a raised dispatch # returns 500 rather than dropping the connection. h, send_response = self._handler( StubBroker(secrets.token_bytes(16)), {"Content-Length": "0"}) with patch("bot_bottle.orchestrator.host_server.dispatch", side_effect=RuntimeError("boom")): h._serve("POST") send_response.assert_called_once_with(500) def test_health_over_do_get(self) -> None: h, send_response = self._handler(StubBroker(secrets.token_bytes(16)), {}) h.path = "/health" h.do_GET() send_response.assert_called_once_with(200) class TestMain(unittest.TestCase): def test_fail_closed_without_secret(self) -> None: with patch("bot_bottle.orchestrator.host_server.broker_secret_from_env", return_value=None): self.assertEqual(2, main(["--port", "0"])) def test_serves_then_shuts_down_cleanly(self) -> None: fake = MagicMock() fake.server_address = ("127.0.0.1", 0) fake.serve_forever.side_effect = KeyboardInterrupt with patch("bot_bottle.orchestrator.host_server.broker_secret_from_env", return_value=b"k"), \ patch("bot_bottle.orchestrator.host_server.make_host_server", return_value=fake): self.assertEqual(0, main(["--port", "0"])) fake.serve_forever.assert_called_once() fake.server_close.assert_called_once() if __name__ == "__main__": unittest.main()