|
| 1 | +"""Exercise stream deadlines with real sockets, including HTTPX body reads.""" |
| 2 | + |
| 3 | +import json |
| 4 | +import threading |
| 5 | +from functools import partial |
| 6 | +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer |
| 7 | +from types import SimpleNamespace |
| 8 | + |
| 9 | +import anyio |
| 10 | +import httpx |
| 11 | +import pytest |
| 12 | + |
| 13 | +from hyperbrowser.client.managers.async_manager.sandboxes import ( |
| 14 | + sandbox_transport as async_transport, |
| 15 | +) |
| 16 | +from hyperbrowser.client.managers.async_manager.sandboxes.sandbox_processes import ( |
| 17 | + SandboxProcessesApi as AsyncProcesses, |
| 18 | +) |
| 19 | +from hyperbrowser.client.managers.sync_manager.sandboxes import ( |
| 20 | + sandbox_transport as sync_transport, |
| 21 | +) |
| 22 | +from hyperbrowser.client.managers.sync_manager.sandboxes.sandbox_processes import ( |
| 23 | + SandboxProcessesApi as SyncProcesses, |
| 24 | +) |
| 25 | +from hyperbrowser.exceptions import HyperbrowserError |
| 26 | +from hyperbrowser.sandbox_common import RuntimeConnection |
| 27 | + |
| 28 | + |
| 29 | +@pytest.fixture |
| 30 | +def stream_server(): |
| 31 | + servers = [] |
| 32 | + |
| 33 | + def start(respond): |
| 34 | + stopped = threading.Event() |
| 35 | + requests = [] |
| 36 | + |
| 37 | + class Handler(BaseHTTPRequestHandler): |
| 38 | + def log_message(self, *args): |
| 39 | + pass |
| 40 | + |
| 41 | + def do_GET(self): |
| 42 | + requests.append((self.command, self.headers.get("Authorization"))) |
| 43 | + self.rfile.read(int(self.headers.get("Content-Length", 0))) |
| 44 | + try: |
| 45 | + respond(self, stopped) |
| 46 | + except (BrokenPipeError, ConnectionResetError): |
| 47 | + pass # Expected when a timeout closes the client socket. |
| 48 | + |
| 49 | + do_POST = do_GET |
| 50 | + |
| 51 | + def headers_for(self, status=200, content_type="text/event-stream"): |
| 52 | + self.send_response(status) |
| 53 | + self.send_header("Content-Type", content_type) |
| 54 | + self.end_headers() |
| 55 | + |
| 56 | + def event(self, name, data): |
| 57 | + self.wfile.write( |
| 58 | + ("event: " + name + "\ndata: " + json.dumps(data) + "\n\n").encode() |
| 59 | + ) |
| 60 | + self.wfile.flush() |
| 61 | + |
| 62 | + def started(self): |
| 63 | + self.headers_for() |
| 64 | + self.event( |
| 65 | + "started", |
| 66 | + { |
| 67 | + "id": "p1", |
| 68 | + "status": "running", |
| 69 | + "command": "quiet", |
| 70 | + "cwd": "/tmp", |
| 71 | + "started_at": 1, |
| 72 | + }, |
| 73 | + ) |
| 74 | + |
| 75 | + def finished(self): |
| 76 | + self.event( |
| 77 | + "output", |
| 78 | + { |
| 79 | + "seq": 1, |
| 80 | + "stream": "stdout", |
| 81 | + "data": "finished", |
| 82 | + "timestamp": 2, |
| 83 | + }, |
| 84 | + ) |
| 85 | + self.event( |
| 86 | + "done", |
| 87 | + { |
| 88 | + "id": "p1", |
| 89 | + "status": "exited", |
| 90 | + "exit_code": 0, |
| 91 | + "started_at": 1, |
| 92 | + "completed_at": 2, |
| 93 | + "last_seq": 1, |
| 94 | + }, |
| 95 | + ) |
| 96 | + |
| 97 | + server = ThreadingHTTPServer(("127.0.0.1", 0), Handler) |
| 98 | + thread = threading.Thread( |
| 99 | + target=partial(server.serve_forever, poll_interval=0.02), daemon=True |
| 100 | + ) |
| 101 | + thread.start() |
| 102 | + servers.append((server, thread, stopped)) |
| 103 | + return SimpleNamespace( |
| 104 | + url="http://127.0.0.1:{}".format(server.server_port), requests=requests |
| 105 | + ) |
| 106 | + |
| 107 | + yield start |
| 108 | + for server, thread, stopped in servers: |
| 109 | + stopped.set() |
| 110 | + server.shutdown() |
| 111 | + server.server_close() |
| 112 | + thread.join() |
| 113 | + |
| 114 | + |
| 115 | +@pytest.fixture(params=["sync", "async"]) |
| 116 | +def client(request, monkeypatch): |
| 117 | + asynchronous = request.param == "async" |
| 118 | + module = async_transport if asynchronous else sync_transport |
| 119 | + |
| 120 | + async def call(function, *args, **kwargs): |
| 121 | + if asynchronous: |
| 122 | + return await function(*args, **kwargs) |
| 123 | + return await anyio.to_thread.run_sync(partial(function, *args, **kwargs)) |
| 124 | + |
| 125 | + def create(server, request_timeout, idle_timeout): |
| 126 | + monkeypatch.setattr(module, "PROCESS_STREAM_IDLE_TIMEOUT_SECONDS", idle_timeout) |
| 127 | + refreshes = [] |
| 128 | + |
| 129 | + def resolve(refresh): |
| 130 | + refreshes.append(refresh) |
| 131 | + return RuntimeConnection( |
| 132 | + sandbox_id="test", |
| 133 | + base_url=server.url, |
| 134 | + token="fresh" if refresh else "old", |
| 135 | + ) |
| 136 | + |
| 137 | + async def async_resolve(refresh): |
| 138 | + return resolve(refresh) |
| 139 | + |
| 140 | + transport = module.RuntimeTransport( |
| 141 | + async_resolve if asynchronous else resolve, timeout=request_timeout |
| 142 | + ) |
| 143 | + api = (AsyncProcesses if asynchronous else SyncProcesses)(transport) |
| 144 | + return SimpleNamespace( |
| 145 | + api=api, transport=transport, call=call, refreshes=refreshes |
| 146 | + ) |
| 147 | + |
| 148 | + return create |
| 149 | + |
| 150 | + |
| 151 | +@pytest.mark.anyio |
| 152 | +async def test_process_stream_accepts_line_terminators( |
| 153 | + stream_server, client, monkeypatch |
| 154 | +): |
| 155 | + # HTTPX 0.23, our minimum supported version, keeps the newline on each line. |
| 156 | + iter_lines = httpx.Response.iter_lines |
| 157 | + aiter_lines = httpx.Response.aiter_lines |
| 158 | + |
| 159 | + def legacy_lines(response): |
| 160 | + for line in iter_lines(response): |
| 161 | + yield line + "\n" |
| 162 | + |
| 163 | + async def async_legacy_lines(response): |
| 164 | + async for line in aiter_lines(response): |
| 165 | + yield line + "\n" |
| 166 | + |
| 167 | + monkeypatch.setattr(httpx.Response, "iter_lines", legacy_lines) |
| 168 | + monkeypatch.setattr(httpx.Response, "aiter_lines", async_legacy_lines) |
| 169 | + |
| 170 | + def respond(handler, stopped): |
| 171 | + handler.started() |
| 172 | + handler.finished() |
| 173 | + |
| 174 | + c = client(stream_server(respond), request_timeout=1, idle_timeout=1) |
| 175 | + handle = await c.call(c.api.start, "quiet") |
| 176 | + try: |
| 177 | + result = await c.call(handle.wait, timeout_sec=3) |
| 178 | + assert (result.stdout, result.exit_code) == ("finished", 0) |
| 179 | + finally: |
| 180 | + await c.call(handle.disconnect) |
| 181 | + |
| 182 | + |
| 183 | +@pytest.mark.anyio |
| 184 | +@pytest.mark.parametrize("refresh", [False, True]) |
| 185 | +async def test_quiet_process_and_heartbeats_outlive_request_timeout( |
| 186 | + stream_server, client, refresh |
| 187 | +): |
| 188 | + def respond(handler, stopped): |
| 189 | + if refresh and handler.headers["Authorization"] == "Bearer old": |
| 190 | + handler.headers_for(401, "application/json") |
| 191 | + handler.wfile.write(b'{"error":"expired"}') |
| 192 | + return |
| 193 | + handler.started() |
| 194 | + # Longer than the ordinary request timeout, shorter than stream idle. |
| 195 | + if stopped.wait(0.6): |
| 196 | + return |
| 197 | + # Total duration exceeds stream idle too: each heartbeat resets it. |
| 198 | + for _ in range(4): |
| 199 | + handler.event("keepalive", {}) |
| 200 | + if stopped.wait(0.25): |
| 201 | + return |
| 202 | + handler.finished() |
| 203 | + |
| 204 | + server = stream_server(respond) |
| 205 | + c = client(server, request_timeout=0.25, idle_timeout=1.0) |
| 206 | + handle = await c.call(c.api.start, "quiet") |
| 207 | + try: |
| 208 | + result = await c.call(handle.wait, timeout_sec=5) |
| 209 | + assert (result.stdout, result.exit_code) == ("finished", 0) |
| 210 | + assert c.refreshes == ([False, True] if refresh else [False]) |
| 211 | + assert len(server.requests) == (2 if refresh else 1) |
| 212 | + finally: |
| 213 | + await c.call(handle.disconnect) |
| 214 | + |
| 215 | + |
| 216 | +@pytest.mark.anyio |
| 217 | +async def test_missing_heartbeats_fail_without_reexecuting(stream_server, client): |
| 218 | + def respond(handler, stopped): |
| 219 | + handler.started() |
| 220 | + stopped.wait(10) |
| 221 | + |
| 222 | + server = stream_server(respond) |
| 223 | + c = client(server, request_timeout=5, idle_timeout=0.25) |
| 224 | + handle = await c.call(c.api.start, "quiet") |
| 225 | + try: |
| 226 | + with pytest.raises(HyperbrowserError) as exc: |
| 227 | + await c.call(handle.wait, timeout_sec=2) |
| 228 | + assert exc.value.code == "incomplete_output" |
| 229 | + assert exc.value.details["process_id"] == "p1" |
| 230 | + assert not exc.value.retryable |
| 231 | + assert len(server.requests) == 1 |
| 232 | + finally: |
| 233 | + await c.call(handle.disconnect) |
| 234 | + |
| 235 | + |
| 236 | +@pytest.mark.anyio |
| 237 | +async def test_response_headers_keep_ordinary_request_timeout(stream_server, client): |
| 238 | + def respond(handler, stopped): |
| 239 | + stopped.wait(2) |
| 240 | + |
| 241 | + server = stream_server(respond) |
| 242 | + c = client(server, request_timeout=0.25, idle_timeout=5) |
| 243 | + with pytest.raises(HyperbrowserError) as exc: |
| 244 | + await c.call(c.api.start, "quiet") |
| 245 | + assert isinstance(exc.value.original_error, httpx.ReadTimeout) |
| 246 | + assert len(server.requests) == 1 |
| 247 | + |
| 248 | + |
| 249 | +@pytest.mark.anyio |
| 250 | +async def test_json_body_keeps_ordinary_request_timeout(stream_server, client): |
| 251 | + def respond(handler, stopped): |
| 252 | + handler.headers_for(200, "application/json") |
| 253 | + stopped.wait(2) |
| 254 | + |
| 255 | + server = stream_server(respond) |
| 256 | + c = client(server, request_timeout=0.25, idle_timeout=5) |
| 257 | + with pytest.raises(HyperbrowserError) as exc: |
| 258 | + await c.call(c.transport.request_json, "/sandbox/processes/p1") |
| 259 | + assert isinstance(exc.value.original_error, httpx.ReadTimeout) |
| 260 | + assert len(server.requests) == 1 |
0 commit comments