320 lines
10 KiB
Python
320 lines
10 KiB
Python
"""Absolute-deadline contracts for the killable HTTP control helper."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import base64
|
|
import io
|
|
import json
|
|
import subprocess
|
|
import threading
|
|
import time
|
|
import urllib.request
|
|
from http.server import BaseHTTPRequestHandler, HTTPServer
|
|
|
|
import pytest
|
|
|
|
from testing.tests.test_hermes_scm_broker_support import _load
|
|
|
|
|
|
class _Handler(BaseHTTPRequestHandler):
|
|
def _respond(self):
|
|
if self.path == "/slow":
|
|
time.sleep(10)
|
|
if self.path == "/redirect":
|
|
self.send_response(302)
|
|
self.send_header("Location", "http://127.0.0.1:1/evil")
|
|
self.send_header("Content-Length", "0")
|
|
self.end_headers()
|
|
return
|
|
status = 404 if self.path == "/missing" else 200
|
|
body = b"x" * 64 if self.path == "/large" else b'{"ok":true}'
|
|
self.send_response(status)
|
|
self.send_header("Content-Type", "application/json")
|
|
self.send_header("Content-Length", str(len(body)))
|
|
self.end_headers()
|
|
self.wfile.write(body)
|
|
|
|
do_GET = _respond
|
|
do_POST = _respond
|
|
|
|
def log_message(self, *_args):
|
|
return
|
|
|
|
|
|
@pytest.fixture(scope="module")
|
|
def server():
|
|
value = HTTPServer(("127.0.0.1", 0), _Handler)
|
|
thread = threading.Thread(target=value.serve_forever, daemon=True)
|
|
thread.start()
|
|
yield f"http://127.0.0.1:{value.server_port}"
|
|
value.shutdown()
|
|
|
|
|
|
def _control(url: str, **overrides):
|
|
value = {
|
|
"url": url,
|
|
"method": "GET",
|
|
"headers": {"Accept": "application/json"},
|
|
"body": "",
|
|
"maximum": 1024,
|
|
"timeout": 10,
|
|
}
|
|
value.update(overrides)
|
|
return value
|
|
|
|
|
|
def test_child_exchange_round_trips_and_reports_http_errors(server):
|
|
module = _load("deadline_http")
|
|
result = module._child_exchange(_control(server + "/ok"))
|
|
assert (result.status, result.content_type) == (200, "application/json")
|
|
assert result.body == b'{"ok":true}'
|
|
|
|
posted = module._child_exchange(
|
|
_control(
|
|
server + "/ok",
|
|
method="POST",
|
|
body=base64.b64encode(b'{"a":1}').decode(),
|
|
)
|
|
)
|
|
assert posted.status == 200
|
|
|
|
missing = module._child_exchange(_control(server + "/missing"))
|
|
assert missing.status == 404
|
|
|
|
|
|
def test_child_exchange_rejects_redirects_and_oversized_responses(server):
|
|
module = _load("deadline_http")
|
|
with pytest.raises(Exception):
|
|
module._child_exchange(_control(server + "/redirect"))
|
|
with pytest.raises(ValueError, match="too large"):
|
|
module._child_exchange(_control(server + "/large", maximum=8))
|
|
|
|
|
|
@pytest.mark.parametrize(
|
|
"overrides",
|
|
[
|
|
{"url": 1},
|
|
{"method": None},
|
|
{"headers": ["not", "dict"]},
|
|
{"headers": {"Accept": 5}},
|
|
{"body": b"raw"},
|
|
{"maximum": "big"},
|
|
{"maximum": -1},
|
|
{"maximum": 4 * 1024 * 1024},
|
|
{"timeout": "soon"},
|
|
{"timeout": 0},
|
|
{"timeout": 121},
|
|
],
|
|
)
|
|
def test_child_exchange_rejects_invalid_control_documents(overrides):
|
|
module = _load("deadline_http")
|
|
control = _control("http://127.0.0.1:1/")
|
|
control.update(overrides)
|
|
with pytest.raises(ValueError, match="invalid control"):
|
|
module._child_exchange(control)
|
|
|
|
|
|
def test_child_exchange_rejects_bad_or_oversized_encoded_bodies():
|
|
module = _load("deadline_http")
|
|
with pytest.raises(ValueError):
|
|
module._child_exchange(_control("http://127.0.0.1:1/", body="!!not-b64!!"))
|
|
huge = base64.b64encode(b"x" * (module.MAX_CONTROL_BYTES + 1)).decode()
|
|
with pytest.raises(ValueError, match="request too large"):
|
|
module._child_exchange(_control("http://127.0.0.1:1/", body=huge))
|
|
|
|
|
|
def test_read_response_enforces_deadline_and_size():
|
|
module = _load("deadline_http")
|
|
|
|
class Reader:
|
|
def __init__(self, body: bytes):
|
|
self.stream = io.BytesIO(body)
|
|
|
|
def read(self, limit: int):
|
|
return self.stream.read(limit)
|
|
|
|
with pytest.raises(TimeoutError):
|
|
module._read_response(Reader(b"abc"), 10, time.monotonic() - 1)
|
|
with pytest.raises(ValueError, match="too large"):
|
|
module._read_response(Reader(b"abcdef"), 4, time.monotonic() + 5)
|
|
value = module._read_response(Reader(b"abc"), 10, time.monotonic() + 5)
|
|
assert value == b"abc"
|
|
|
|
|
|
def _run_child_main(module, monkeypatch, raw: bytes) -> dict:
|
|
stdin = type("Stdin", (), {"buffer": io.BytesIO(raw)})()
|
|
stdout = io.StringIO()
|
|
monkeypatch.setattr(module.sys, "stdin", stdin)
|
|
monkeypatch.setattr(module.sys, "stdout", stdout)
|
|
code = module._child_main()
|
|
return code, json.loads(stdout.getvalue())
|
|
|
|
|
|
def test_child_main_reports_success_and_failure_documents(monkeypatch, server):
|
|
module = _load("deadline_http")
|
|
code, output = _run_child_main(
|
|
module, monkeypatch, json.dumps(_control(server + "/ok")).encode()
|
|
)
|
|
assert code == 0 and output["ok"] is True
|
|
assert base64.b64decode(output["body"]) == b'{"ok":true}'
|
|
|
|
for raw in (
|
|
b"x" * (module.MAX_CONTROL_BYTES + 1),
|
|
b"not json",
|
|
b"[1,2]",
|
|
):
|
|
code, output = _run_child_main(module, monkeypatch, raw)
|
|
assert code == 1 and output == {"ok": False}
|
|
|
|
|
|
class _FakeProcess:
|
|
def __init__(self, output: bytes, returncode: int = 0):
|
|
self.args = ["fake-child"]
|
|
self.pid = 0
|
|
self.returncode = returncode
|
|
self.output = output
|
|
self.stdin_payload = None
|
|
|
|
def communicate(self, payload=None, timeout=None):
|
|
self.stdin_payload = payload
|
|
return self.output, b""
|
|
|
|
|
|
def _fake_popen(output: bytes, returncode: int = 0):
|
|
def popen(*_args, **_kwargs):
|
|
return _FakeProcess(output, returncode)
|
|
|
|
return popen
|
|
|
|
|
|
def _success_output(body: bytes, **overrides) -> bytes:
|
|
value = {
|
|
"ok": True,
|
|
"status": 200,
|
|
"content_type": "application/json",
|
|
"body": base64.b64encode(body).decode("ascii"),
|
|
}
|
|
value.update(overrides)
|
|
return json.dumps(value).encode()
|
|
|
|
|
|
def test_exchange_validates_bounds_and_request_body():
|
|
module = _load("deadline_http")
|
|
request = urllib.request.Request("https://scm.bstein.dev/api/v1/x")
|
|
with pytest.raises(module.PolicyError, match="bounds are invalid"):
|
|
module.exchange(request, maximum=1, timeout=0)
|
|
with pytest.raises(module.PolicyError, match="bounds are invalid"):
|
|
module.exchange(request, maximum=-1, timeout=10)
|
|
request.data = "not-bytes"
|
|
with pytest.raises(module.PolicyError, match="safe size"):
|
|
module.exchange(request, maximum=1, timeout=10)
|
|
request.data = b"x" * (module.MAX_CONTROL_BYTES + 1)
|
|
with pytest.raises(module.PolicyError, match="safe size"):
|
|
module.exchange(request, maximum=1, timeout=10)
|
|
|
|
|
|
def test_exchange_returns_result_and_enforces_response_limit():
|
|
module = _load("deadline_http")
|
|
request = urllib.request.Request(
|
|
"https://scm.bstein.dev/api/v1/x", data=b"{}", method="POST"
|
|
)
|
|
result = module.exchange(
|
|
request, maximum=64, timeout=10, popen=_fake_popen(_success_output(b"body"))
|
|
)
|
|
assert (result.status, result.content_type, result.body) == (
|
|
200,
|
|
"application/json",
|
|
b"body",
|
|
)
|
|
with pytest.raises(module.PolicyError, match="safe size"):
|
|
module.exchange(
|
|
request,
|
|
maximum=2,
|
|
timeout=10,
|
|
popen=_fake_popen(_success_output(b"body")),
|
|
)
|
|
|
|
|
|
@pytest.mark.parametrize(
|
|
("output", "returncode", "match"),
|
|
[
|
|
(b"", 1, "request failed"),
|
|
(b"x" * (4 * 1024 * 1024 + 1), 0, "request failed"),
|
|
(b"not json", 0, "invalid evidence"),
|
|
(b'{"ok":false}', 0, "invalid evidence"),
|
|
(b"[1]", 0, "invalid evidence"),
|
|
(b'{"ok":true}', 0, "invalid evidence"),
|
|
(b'{"ok":true,"status":"200","content_type":"a","body":""}', 0, "invalid evidence"),
|
|
(b'{"ok":true,"status":200,"content_type":"a","body":"!!"}', 0, "invalid evidence"),
|
|
],
|
|
ids=[
|
|
"empty-error",
|
|
"oversized-output",
|
|
"invalid-json",
|
|
"false-result",
|
|
"array-result",
|
|
"missing-fields",
|
|
"empty-body",
|
|
"invalid-body",
|
|
],
|
|
)
|
|
def test_exchange_rejects_invalid_helper_evidence(output, returncode, match):
|
|
module = _load("deadline_http")
|
|
request = urllib.request.Request("https://scm.bstein.dev/api/v1/x")
|
|
with pytest.raises(module.PolicyError, match=match):
|
|
module.exchange(
|
|
request, maximum=64, timeout=10, popen=_fake_popen(output, returncode)
|
|
)
|
|
|
|
|
|
def test_exchange_kills_the_helper_when_the_deadline_expires(server):
|
|
module = _load("deadline_http")
|
|
request = urllib.request.Request(server + "/slow")
|
|
started = time.monotonic()
|
|
with pytest.raises(module.PolicyError, match="deadline exceeded"):
|
|
module.exchange(request, maximum=64, timeout=1.0)
|
|
assert time.monotonic() - started < 5
|
|
|
|
|
|
def test_exchange_fails_before_spawn_wait_when_time_is_exhausted(monkeypatch):
|
|
module = _load("deadline_http")
|
|
request = urllib.request.Request("https://scm.bstein.dev/api/v1/x")
|
|
killed = []
|
|
monkeypatch.setattr(module.os, "killpg", lambda pid, sig: killed.append((pid, sig)))
|
|
clock = iter([0.0, 1_000_000.0, 1_000_000.0, 1_000_000.0])
|
|
monkeypatch.setattr(module.time, "monotonic", lambda: next(clock))
|
|
with pytest.raises(module.PolicyError, match="deadline exceeded"):
|
|
module.exchange(
|
|
request, maximum=64, timeout=10, popen=_fake_popen(_success_output(b""))
|
|
)
|
|
assert killed
|
|
|
|
|
|
def test_terminate_tolerates_already_finished_helpers():
|
|
module = _load("deadline_http")
|
|
process = subprocess.Popen(
|
|
["/bin/true"], start_new_session=True, stdout=subprocess.DEVNULL
|
|
)
|
|
process.wait()
|
|
module._terminate(process)
|
|
assert process.returncode == 0
|
|
|
|
|
|
def test_child_entrypoint_round_trip_through_real_subprocess(server):
|
|
module = _load("deadline_http")
|
|
request = urllib.request.Request(server + "/ok")
|
|
result = module.exchange(request, maximum=1024, timeout=20)
|
|
assert result.status == 200 and result.body == b'{"ok":true}'
|
|
|
|
|
|
def test_open_bounded_adapts_result_to_the_urlopen_interface():
|
|
module = _load("deadline_http")
|
|
request = urllib.request.Request("https://scm.bstein.dev/api/v1/x")
|
|
with module.open_bounded(
|
|
request, maximum=64, timeout=10, popen=_fake_popen(_success_output(b"body"))
|
|
) as response:
|
|
assert response.status == 200
|
|
assert response.headers.get_content_type() == "application/json"
|
|
assert response.read(2) == b"bo"
|
|
assert response.read() == b"dy"
|