"""What the client does when the engine is slow, busy, or briefly gone.""" import httpx import pytest from fluksio.sdk.client import WAIT_TOLERANCE, ApiError, Client, RunHandle def a_client(handler, monkeypatch, **kwargs): """A client whose transport is a function, and whose backoff costs nothing.""" monkeypatch.setattr("fluksio.sdk.client.time.sleep", lambda _seconds: None) transport = httpx.MockTransport(handler) http = httpx.Client(transport=transport, base_url="http://engine") return Client(url="http://engine", token="t", http=http, **kwargs) def test_idempotent_get_is_tried_again(monkeypatch): calls = [] def handler(request): calls.append(request) if len(calls) < 3: raise httpx.ReadTimeout("too slow", request=request) return httpx.Response(200, json={"id": "r1", "status": "ok"}) client = a_client(handler, monkeypatch) assert client.run("r1")["status"] == "ok" assert len(calls) == 3 def test_a_busy_engine_is_asked_again(monkeypatch): codes = iter([503, 503, 200]) def handler(request): code = next(codes) return httpx.Response(code, json={"id": "r1"} if code == 200 else {}) client = a_client(handler, monkeypatch) assert client.run("r1")["id"] == "r1" def test_giving_up_raises_what_it_last_saw(monkeypatch): def handler(request): raise httpx.ReadTimeout("too slow", request=request) client = a_client(handler, monkeypatch, retries=2) with pytest.raises(httpx.ReadTimeout): client.run("r1") def test_a_write_is_not_repeated(monkeypatch): calls = [] def handler(request): calls.append(request) raise httpx.ReadTimeout("too slow", request=request) client = a_client(handler, monkeypatch) with pytest.raises(httpx.ReadTimeout): client.put_source("train", "fit", "code") assert len(calls) == 1, "a source write means something different twice" def test_submit_carries_one_key_across_its_retries(monkeypatch): bodies = [] def handler(request): bodies.append(httpx.Response(200, content=request.content).json()) if len(bodies) < 3: raise httpx.ConnectError("no route", request=request) return httpx.Response(202, json={"id": "r1", "status": "queued"}) client = a_client(handler, monkeypatch) assert client.submit("train", {"lr": 0.1}).id == "r1" keys = {body["idempotency_key"] for body in bodies} assert len(keys) == 1, "a retry must not read as a second run" assert len(next(iter(keys))) == 32 def test_a_sweep_keys_every_entry(monkeypatch): seen = {} def handler(request): body = httpx.Response(200, content=request.content).json() seen["runs"] = body["runs"] return httpx.Response(202, json=[{"id": "r1"}, {"id": "r2"}]) client = a_client(handler, monkeypatch) client.sweep("train", [{"params": {"lr": 0.1}}, {"params": {"lr": 0.2}}]) keys = [entry["idempotency_key"] for entry in seen["runs"]] assert len(set(keys)) == 2 def test_waiting_survives_a_few_bad_answers(monkeypatch): answers = iter( [503] * (WAIT_TOLERANCE - 1) + [200] # then the run is finished ) def handler(request): code = next(answers) if code != 200: return httpx.Response(code, json={}) return httpx.Response(200, json={"id": "r1", "status": "ok"}) client = a_client(handler, monkeypatch, retries=0) handle = RunHandle(client, "r1", {"status": "running"}) assert handle.wait(poll=0).status == "ok" def test_waiting_gives_up_eventually(monkeypatch): def handler(request): return httpx.Response(503, json={}) client = a_client(handler, monkeypatch, retries=0) handle = RunHandle(client, "r1", {"status": "running"}) with pytest.raises(ApiError): handle.wait(poll=0) def test_a_run_that_is_gone_stops_the_wait_at_once(monkeypatch): calls = [] def handler(request): calls.append(request) return httpx.Response(404, json={"detail": "no such run"}) client = a_client(handler, monkeypatch, retries=0) handle = RunHandle(client, "r1", {"status": "running"}) with pytest.raises(ApiError) as caught: handle.wait(poll=0) assert caught.value.status == 404 assert len(calls) == 1, "a 404 is an answer, not a blip"