|
9 | 9 | from collections.abc import AsyncIterator |
10 | 10 |
|
11 | 11 | import pytest |
| 12 | +from websockets.exceptions import PayloadTooBig |
12 | 13 |
|
13 | 14 | from hypeman.lib import ( |
14 | 15 | CopyCallbacks, |
@@ -62,9 +63,9 @@ def __exit__( |
62 | 63 | @dataclass |
63 | 64 | class FakeConnector: |
64 | 65 | connections: deque[FakeWebSocket] |
65 | | - calls: list[tuple[str, dict[str, str], int | None]] = field(default_factory=lambda: []) |
| 66 | + calls: list[tuple[str, dict[str, str], int]] = field(default_factory=lambda: []) |
66 | 67 |
|
67 | | - def __call__(self, url: str, *, additional_headers: dict[str, str], max_size: int | None) -> FakeWebSocket: |
| 68 | + def __call__(self, url: str, *, additional_headers: dict[str, str], max_size: int) -> FakeWebSocket: |
68 | 69 | self.calls.append((url, additional_headers, max_size)) |
69 | 70 | return self.connections.popleft() |
70 | 71 |
|
@@ -102,9 +103,9 @@ async def __aexit__( |
102 | 103 | @dataclass |
103 | 104 | class FakeAsyncConnector: |
104 | 105 | connections: deque[FakeAsyncWebSocket] |
105 | | - calls: list[tuple[str, dict[str, str], int | None]] = field(default_factory=lambda: []) |
| 106 | + calls: list[tuple[str, dict[str, str], int]] = field(default_factory=lambda: []) |
106 | 107 |
|
107 | | - def __call__(self, url: str, *, additional_headers: dict[str, str], max_size: int | None) -> FakeAsyncWebSocket: |
| 108 | + def __call__(self, url: str, *, additional_headers: dict[str, str], max_size: int) -> FakeAsyncWebSocket: |
108 | 109 | self.calls.append((url, additional_headers, max_size)) |
109 | 110 | return self.connections.popleft() |
110 | 111 |
|
@@ -155,7 +156,7 @@ def test_exec_uses_client_auth_url_and_all_request_dimensions() -> None: |
155 | 156 | assert result.output == b"outerr" |
156 | 157 | assert result.exit_code == 0 |
157 | 158 | assert connector.calls == [ |
158 | | - ("wss://example.test/api/instances/inst_123/exec", {"Authorization": "Bearer secret"}, None) |
| 159 | + ("wss://example.test/api/instances/inst_123/exec", {"Authorization": "Bearer secret"}, 2**20) |
159 | 160 | ] |
160 | 161 | assert json.loads(str(websocket.sent[0])) == { |
161 | 162 | "command": ["sh", "-lc", "echo hi"], |
@@ -196,6 +197,31 @@ def test_exec_never_retries_after_dispatch() -> None: |
196 | 197 | assert len(websocket.sent) == 1 |
197 | 198 |
|
198 | 199 |
|
| 200 | +@pytest.mark.parametrize(("tty", "resize"), [(False, [(24, 80)]), (True, [(0, 80)]), (True, [(24, -1)])]) |
| 201 | +def test_exec_rejects_invalid_resize_before_connect(tty: bool, resize: list[tuple[int, int]]) -> None: |
| 202 | + connector = FakeConnector(deque()) |
| 203 | + with pytest.raises(ValueError, match="resize"): |
| 204 | + exec(FakeClient(), "inst", ["true"], tty=tty, resize=resize, connector=connector) |
| 205 | + assert not connector.calls |
| 206 | + |
| 207 | + |
| 208 | +def test_exec_rejects_oversized_inbound_message() -> None: |
| 209 | + oversized = PayloadTooBig(2**20 + 1, 2**20) |
| 210 | + connector = FakeConnector(deque([FakeWebSocket(deque([oversized]))])) |
| 211 | + with pytest.raises(ExecProtocolError, match="before an exitCode") as exc_info: |
| 212 | + exec(FakeClient(), "inst", ["true"], connector=connector) |
| 213 | + assert isinstance(exc_info.value.__cause__, PayloadTooBig) |
| 214 | + assert connector.calls[0][2] == 2**20 |
| 215 | + |
| 216 | + |
| 217 | +@pytest.mark.asyncio |
| 218 | +async def test_exec_async_rejects_invalid_resize_before_connect() -> None: |
| 219 | + connector = FakeAsyncConnector(deque()) |
| 220 | + with pytest.raises(ValueError, match="resize"): |
| 221 | + await exec_async(FakeClient(), "inst", ["true"], tty=False, resize=[(24, 80)], connector=connector) |
| 222 | + assert not connector.calls |
| 223 | + |
| 224 | + |
199 | 225 | @pytest.mark.asyncio |
200 | 226 | async def test_exec_async_supports_streaming_stdin() -> None: |
201 | 227 | websocket = FakeAsyncWebSocket(deque([b"done", '{"exitCode":7}'])) |
@@ -248,6 +274,7 @@ def test_cp_upload_file_preserves_mode_and_reports_progress(tmp_path: Path) -> N |
248 | 274 | expected_request["gid"] = source_stat.st_gid |
249 | 275 | assert request == expected_request |
250 | 276 | assert websocket.sent[1:] == [b"payload", '{"type":"end"}'] |
| 277 | + assert connector.calls[0][2] == 2**20 |
251 | 278 | assert events == [ |
252 | 279 | ("start", (str(source), 7)), |
253 | 280 | ("progress", 7), |
@@ -420,27 +447,31 @@ async def test_cp_async_upload_and_download(tmp_path: Path) -> None: |
420 | 447 | source = tmp_path / "source" |
421 | 448 | source.write_bytes(b"async") |
422 | 449 | upload_socket = FakeAsyncWebSocket(deque([upload_result(5)])) |
| 450 | + upload_connector = FakeAsyncConnector(deque([upload_socket])) |
423 | 451 | await cp_to_instance_async( |
424 | 452 | FakeClient(), |
425 | 453 | "inst", |
426 | 454 | source, |
427 | 455 | "/guest/source", |
428 | | - connector=FakeAsyncConnector(deque([upload_socket])), |
| 456 | + connector=upload_connector, |
429 | 457 | ) |
430 | 458 | assert upload_socket.sent[1] == b"async" |
| 459 | + assert upload_connector.calls[0][2] == 2**20 |
431 | 460 |
|
432 | 461 | download_socket = FakeAsyncWebSocket( |
433 | 462 | deque([file_header("result", size=5), b"async", text_frame({"type": "end", "final": True})]) |
434 | 463 | ) |
435 | 464 | destination = tmp_path / "dest" |
| 465 | + download_connector = FakeAsyncConnector(deque([download_socket])) |
436 | 466 | await cp_from_instance_async( |
437 | 467 | FakeClient(), |
438 | 468 | "inst", |
439 | 469 | "/guest/result", |
440 | 470 | destination, |
441 | | - connector=FakeAsyncConnector(deque([download_socket])), |
| 471 | + connector=download_connector, |
442 | 472 | ) |
443 | 473 | assert (destination / "result").read_bytes() == b"async" |
| 474 | + assert download_connector.calls[0][2] == 2**20 |
444 | 475 |
|
445 | 476 |
|
446 | 477 | def test_invalid_instance_id_is_rejected_before_connect() -> None: |
|
0 commit comments