diff --git a/README.md b/README.md index f5cba64..4f8919e 100644 --- a/README.md +++ b/README.md @@ -144,6 +144,30 @@ async def main() -> None: asyncio.run(main()) ``` +### Request Timeouts + +Each HTTP request times out after 30 seconds by default. Generated resource +methods take no timeout argument, so the override is scoped instead: +`request_timeout(seconds)` applies to every request sent inside the block, in +the current thread or asyncio task (and code that copies its context, such as +`asyncio.to_thread`). A streaming call reads the value when you start iterating +it, so iterate inside the block. + +```python +from archastro.platform import request_timeout + +with request_timeout(5.0): + docs = client.knowledge_documents.list(source=["cso_..."]) +``` + +The value is a per-request httpx timeout, not a total for the block. After a +401, the token refresh request and the retried request each get the same value, +so a single call can take up to three times it. To hold a chain of calls to one +deadline, pass each call the time that remains. A value +that is not positive raises `TimeoutError` without sending the request. +`HttpClient` and `SyncHttpClient` also accept `timeout=` to change the default +(`DEFAULT_TIMEOUT_S`, 30 seconds). + ## Examples - [`examples/org_system_user_token`](examples/org_system_user_token) — run the diff --git a/src/archastro/platform/__init__.py b/src/archastro/platform/__init__.py index 4495032..2f81077 100644 --- a/src/archastro/platform/__init__.py +++ b/src/archastro/platform/__init__.py @@ -6,6 +6,7 @@ from .auth import AsyncAuthClient, AuthClient, AuthTokens # noqa: F401 from .client import AsyncPlatformClient, PlatformClient # noqa: F401 +from .runtime.http_client import DEFAULT_TIMEOUT_S, request_timeout # noqa: F401 from .v1 import ( V1, # noqa: F401 AsyncV1, # noqa: F401 diff --git a/src/archastro/platform/runtime/http_client.py b/src/archastro/platform/runtime/http_client.py index 27c52cb..b194252 100644 --- a/src/archastro/platform/runtime/http_client.py +++ b/src/archastro/platform/runtime/http_client.py @@ -7,6 +7,8 @@ import json import threading from collections.abc import AsyncIterator, Callable, Coroutine, Iterator +from contextlib import contextmanager +from contextvars import ContextVar from functools import cache from typing import Any, TypeVar, overload @@ -14,9 +16,49 @@ from pydantic import TypeAdapter DEFAULT_API_PREFIX = "/api/v1" +DEFAULT_TIMEOUT_S = 30.0 T = TypeVar("T") +_request_timeout: ContextVar[float | None] = ContextVar("archastro_request_timeout", default=None) + + +@contextmanager +def request_timeout(seconds: float) -> Iterator[None]: + """Override the timeout of every request issued inside this block. + + Applies to all SDK clients in the current context (a thread, or an + asyncio task and anything that copies its context such as + ``asyncio.to_thread``), including requests made by generated resource + methods, which take no timeout argument of their own. The value is read + when a request is sent; a stream reads it when iteration starts, so start + reading inside the block. It is the httpx timeout for each HTTP request, not + a total for the block: after a 401, the token refresh request and the retry + each get the same value again, so one call can spend up to three times it. + Blocks nest; the innermost wins. + + A value that is not positive (including NaN) raises :class:`TimeoutError` + before any request is sent, so a caller passing down the remainder of an + overall deadline fails fast once that deadline has passed. + """ + token = _request_timeout.set(float(seconds)) + try: + yield + finally: + _request_timeout.reset(token) + + +def _resolve_timeout(method: str, path: str) -> Any: + """The per-request timeout override, or httpx's client-default sentinel.""" + seconds = _request_timeout.get() + if seconds is None: + return httpx.USE_CLIENT_DEFAULT + if not seconds > 0: + raise TimeoutError( + f"request timeout budget exhausted before {method} {path} (timeout={seconds:.3f}s)" + ) + return seconds + def _encode_query(query: dict[str, Any] | None) -> dict[str, Any] | None: """Drop unset parameters and encode list values as ``key[]=item``. @@ -72,6 +114,7 @@ def __init__( path_prefix: str | None = None, default_headers: dict[str, str] | None = None, refresh_only: bool = False, + timeout: float = DEFAULT_TIMEOUT_S, ): self._base_url = base_url.rstrip("/") self._access_token = access_token @@ -79,7 +122,8 @@ def __init__( self._on_refresh_token = on_refresh_token self._path_prefix = path_prefix self._default_headers = default_headers or {} - self._client = httpx.AsyncClient(timeout=30.0) + self._timeout = timeout + self._client = httpx.AsyncClient(timeout=timeout) self._refresh_task: asyncio.Task[str] | None = None self._refresh_only = refresh_only @@ -110,6 +154,7 @@ async def _do_fetch( headers: dict[str, str] | None = None, query: dict[str, Any] | None = None, ) -> httpx.Response: + timeout = _resolve_timeout(method, path) token = self._get_token() url = f"{self._base_url}{self._transform_path(path)}" @@ -130,6 +175,7 @@ async def _do_fetch( json=body if body is not None and method not in ("GET", "HEAD") else None, headers=req_headers, params=params, + timeout=timeout, ) async def _execute( @@ -287,6 +333,7 @@ async def stream_sse( f"Refresh-only HTTP client cannot make requests outside {auth_prefix}" ) + timeout = _resolve_timeout(method, path) url = f"{self._base_url}{self._transform_path(path)}" sends_body = body is not None and method not in ("GET", "HEAD") req_headers = {**self._default_headers, "Accept": "text/event-stream"} @@ -305,6 +352,7 @@ async def stream_sse( json=body if sends_body else None, headers=req_headers, params=params, + timeout=timeout, ) as response: if response.status_code >= 400: await response.aread() @@ -347,6 +395,7 @@ def __init__( path_prefix: str | None = None, default_headers: dict[str, str] | None = None, refresh_only: bool = False, + timeout: float = DEFAULT_TIMEOUT_S, ): self._base_url = base_url.rstrip("/") self._access_token = access_token @@ -354,7 +403,8 @@ def __init__( self._on_refresh_token = on_refresh_token self._path_prefix = path_prefix self._default_headers = default_headers or {} - self._client = httpx.Client(timeout=30.0) + self._timeout = timeout + self._client = httpx.Client(timeout=timeout) self._refresh_only = refresh_only self._refresh_lock = threading.Lock() @@ -385,6 +435,7 @@ def _do_fetch( headers: dict[str, str] | None = None, query: dict[str, Any] | None = None, ) -> httpx.Response: + timeout = _resolve_timeout(method, path) token = self._get_token() url = f"{self._base_url}{self._transform_path(path)}" @@ -405,6 +456,7 @@ def _do_fetch( json=body if body is not None and method not in ("GET", "HEAD") else None, headers=req_headers, params=params, + timeout=timeout, ) def _execute( @@ -544,6 +596,7 @@ def stream_sse_sync( f"Refresh-only HTTP client cannot make requests outside {auth_prefix}" ) + timeout = _resolve_timeout(method, path) url = f"{self._base_url}{self._transform_path(path)}" sends_body = body is not None and method not in ("GET", "HEAD") req_headers = {**self._default_headers, "Accept": "text/event-stream"} @@ -562,6 +615,7 @@ def stream_sse_sync( json=body if sends_body else None, headers=req_headers, params=params, + timeout=timeout, ) as response: if response.status_code >= 400: response.read() diff --git a/tests/test_http_client.py b/tests/test_http_client.py index ee6183d..c6f802f 100644 --- a/tests/test_http_client.py +++ b/tests/test_http_client.py @@ -7,7 +7,12 @@ import pytest from pydantic import BaseModel, ValidationError -from archastro.platform.runtime.http_client import ApiError, HttpClient, SyncHttpClient +from archastro.platform.runtime.http_client import ( + ApiError, + HttpClient, + SyncHttpClient, + request_timeout, +) def _mock_response(status: int, body: dict | None = None) -> httpx.Response: @@ -287,6 +292,7 @@ def test_sync_client_sends_auth_headers_query_and_json_body(): "Authorization": "Bearer sat_test", }, params={"limit": 10}, + timeout=httpx.USE_CLIENT_DEFAULT, ) @@ -575,6 +581,7 @@ async def test_async_client_sends_auth_headers_query_and_json_body(): "Authorization": "Bearer sat_test", }, params={"limit": 10}, + timeout=httpx.USE_CLIENT_DEFAULT, ) @@ -594,6 +601,7 @@ async def test_async_client_drops_body_and_keeps_path_for_get_outside_api_prefix json=None, headers={"Content-Type": "application/json"}, params=None, + timeout=httpx.USE_CLIENT_DEFAULT, ) @@ -943,3 +951,296 @@ def handler(request: httpx.Request) -> httpx.Response: ) assert seen[0].url.query == b"q=notes&page_size=25" + + +# --- request timeout contract ------------------------------------------------ +# +# Callers that run under an overall budget (the benchmark harness passes each +# platform call the *remaining* seconds of its step) need every request to +# honor a timeout they choose, including requests issued by generated resource +# methods that expose no timeout argument. These tests read the timeout httpx +# actually attaches to the outgoing request, not the arguments handed to it. + + +def _timeouts(request: httpx.Request) -> dict[str, float]: + return request.extensions["timeout"] + + +def _all(seconds: float) -> dict[str, float]: + return {"connect": seconds, "read": seconds, "write": seconds, "pool": seconds} + + +def _recording_sync_client(**kwargs) -> tuple[SyncHttpClient, list[httpx.Request]]: + seen: list[httpx.Request] = [] + + def handler(request: httpx.Request) -> httpx.Response: + seen.append(request) + return httpx.Response(200, json={}) + + client = SyncHttpClient(base_url="https://api.test", **kwargs) + client._client = httpx.Client(transport=httpx.MockTransport(handler), timeout=client._timeout) + return client, seen + + +def _recording_async_client(**kwargs) -> tuple[HttpClient, list[httpx.Request]]: + seen: list[httpx.Request] = [] + + async def handler(request: httpx.Request) -> httpx.Response: + seen.append(request) + return httpx.Response(200, json={}) + + client = HttpClient(base_url="https://api.test", **kwargs) + client._client = httpx.AsyncClient( + transport=httpx.MockTransport(handler), timeout=client._timeout + ) + return client, seen + + +def test_sync_default_request_timeout_is_thirty_seconds(): + client, seen = _recording_sync_client() + + client.request("/api/v1/things") + + assert _timeouts(seen[0]) == _all(30.0) + + +def test_sync_constructor_timeout_sets_the_client_default(): + client = SyncHttpClient(base_url="https://api.test", timeout=7.5) + + assert client._client.timeout == httpx.Timeout(7.5) + + +async def test_async_constructor_timeout_sets_the_client_default(): + client = HttpClient(base_url="https://api.test", timeout=7.5) + + assert client._client.timeout == httpx.Timeout(7.5) + + +def test_sync_request_timeout_overrides_only_inside_its_block(): + client, seen = _recording_sync_client(timeout=30.0) + + with request_timeout(4.25): + client.request("/api/v1/things") + client.request_raw("/api/v1/things/content") + client.request("/api/v1/things") + + assert [_timeouts(r) for r in seen] == [_all(4.25), _all(4.25), _all(30.0)] + + +async def test_async_request_timeout_overrides_only_inside_its_block(): + client, seen = _recording_async_client(timeout=30.0) + + with request_timeout(4.25): + await client.request("/api/v1/things") + await client.request_raw("/api/v1/things/content") + await client.request("/api/v1/things") + + assert [_timeouts(r) for r in seen] == [_all(4.25), _all(4.25), _all(30.0)] + + +def test_nested_request_timeout_innermost_wins_then_restores(): + client, seen = _recording_sync_client() + + with request_timeout(10.0): + with request_timeout(2.0): + client.request("/api/v1/inner") + client.request("/api/v1/outer") + + assert [_timeouts(r) for r in seen] == [_all(2.0), _all(10.0)] + + +def _refreshing_transport(seen: list[httpx.Request]): + """First /things call is a 401; /auth/refresh issues a token; retry succeeds.""" + calls = {"things": 0} + + def answer(request: httpx.Request) -> httpx.Response: + seen.append(request) + if request.url.path == "/api/v1/auth/refresh": + return httpx.Response(200, json={"access_token": "fresh"}) + calls["things"] += 1 + if calls["things"] == 1: + return httpx.Response(401, json={}) + return httpx.Response(200, json={"ok": True}) + + return answer + + +def test_sync_refresh_request_and_retry_use_the_same_override(): + # The refresh is sent by a separate refresh-only client, the way the + # generated `with_credentials` builder wires it, so this pins that the + # override reaches that client too and cannot fall back to its 30s default. + seen: list[httpx.Request] = [] + transport = httpx.MockTransport(_refreshing_transport(seen)) + refresh_http = SyncHttpClient(base_url="https://api.test", refresh_only=True) + refresh_http._client = httpx.Client(transport=transport) + client = SyncHttpClient( + base_url="https://api.test", + access_token="expired", + on_refresh_token=lambda: refresh_http.request("/api/v1/auth/refresh", method="POST")[ + "access_token" + ], + ) + client._client = httpx.Client(transport=transport) + + with request_timeout(3.0): + assert client.request("/api/v1/things") == {"ok": True} + + assert [(r.url.path, _timeouts(r)) for r in seen] == [ + ("/api/v1/things", _all(3.0)), + ("/api/v1/auth/refresh", _all(3.0)), + ("/api/v1/things", _all(3.0)), + ] + + +async def test_async_refresh_request_and_retry_use_the_same_override(): + # The async refresh runs in its own asyncio task; tasks copy the caller's + # context, which is what carries the override into the refresh request. + seen: list[httpx.Request] = [] + transport = _refreshing_transport(seen) + + async def answer(request: httpx.Request) -> httpx.Response: + return transport(request) + + refresh_http = HttpClient(base_url="https://api.test", refresh_only=True) + refresh_http._client = httpx.AsyncClient(transport=httpx.MockTransport(answer)) + + async def refresh() -> str: + tokens = await refresh_http.request("/api/v1/auth/refresh", method="POST") + return tokens["access_token"] + + client = HttpClient( + base_url="https://api.test", access_token="expired", on_refresh_token=refresh + ) + client._client = httpx.AsyncClient(transport=httpx.MockTransport(answer)) + + with request_timeout(3.0): + assert await client.request("/api/v1/things") == {"ok": True} + + assert [(r.url.path, _timeouts(r)) for r in seen] == [ + ("/api/v1/things", _all(3.0)), + ("/api/v1/auth/refresh", _all(3.0)), + ("/api/v1/things", _all(3.0)), + ] + + +def test_request_timeout_does_not_leak_into_other_threads(): + import threading + + client, seen = _recording_sync_client(timeout=30.0) + entered = threading.Event() + release = threading.Event() + + def hold_override(): + with request_timeout(1.0): + entered.set() + release.wait(5) + + holder = threading.Thread(target=hold_override) + holder.start() + try: + assert entered.wait(5) + client.request("/api/v1/things") + finally: + release.set() + holder.join() + + assert _timeouts(seen[0]) == _all(30.0) + + +@pytest.mark.parametrize("budget", [0.0, -0.5, float("nan")]) +def test_sync_expired_budget_raises_timeout_error_without_sending(budget): + client, seen = _recording_sync_client() + + with request_timeout(budget): + with pytest.raises(TimeoutError, match="GET /api/v1/things"): + client.request("/api/v1/things") + with pytest.raises(TimeoutError): + client.request_raw("/api/v1/things") + with pytest.raises(TimeoutError): + list(client.stream_sse_sync("/api/v1/things/stream")) + + assert seen == [] + + +@pytest.mark.parametrize("budget", [0.0, -0.5, float("nan")]) +async def test_async_expired_budget_raises_timeout_error_without_sending(budget): + client, seen = _recording_async_client() + + with request_timeout(budget): + with pytest.raises(TimeoutError, match="POST /api/v1/things"): + await client.request("/api/v1/things", method="POST", body={}) + with pytest.raises(TimeoutError): + await client.request_raw("/api/v1/things") + with pytest.raises(TimeoutError): + [ev async for ev in client.stream_sse("/api/v1/things/stream")] + + assert seen == [] + + +def test_sync_stream_honors_request_timeout(): + seen: list[httpx.Request] = [] + + def handler(request: httpx.Request) -> httpx.Response: + seen.append(request) + return httpx.Response(200, content=b'event: done\ndata: {"ok": true}\n\n') + + client = SyncHttpClient(base_url="https://api.test") + client._client = httpx.Client(transport=httpx.MockTransport(handler)) + + with request_timeout(6.0): + events = list(client.stream_sse_sync("/api/v1/echo/stream")) + + assert events == [{"event": "done", "data": {"ok": True}}] + assert _timeouts(seen[0]) == _all(6.0) + + +async def test_async_stream_honors_request_timeout(): + seen: list[httpx.Request] = [] + + async def handler(request: httpx.Request) -> httpx.Response: + seen.append(request) + return httpx.Response(200, content=b'event: done\ndata: {"ok": true}\n\n') + + client = HttpClient(base_url="https://api.test") + client._client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + + with request_timeout(6.0): + events = [ev async for ev in client.stream_sse("/api/v1/echo/stream")] + + assert events == [{"event": "done", "data": {"ok": True}}] + assert _timeouts(seen[0]) == _all(6.0) + + +def test_generated_resource_methods_honor_request_timeout(): + # Imported from the package root: that is the public path the README + # teaches, and a regeneration that drops the re-export fails here. + from archastro.platform import DEFAULT_TIMEOUT_S, PlatformClient + from archastro.platform import request_timeout as public_request_timeout + + assert public_request_timeout is request_timeout + assert DEFAULT_TIMEOUT_S == 30.0 + + seen: list[httpx.Request] = [] + + def handler(request: httpx.Request) -> httpx.Response: + seen.append(request) + return httpx.Response( + 200, + json={ + "data": [], + "has_next": False, + "has_prev": False, + "page": 1, + "page_size": 10, + "total_entries": 0, + "total_pages": 0, + }, + ) + + client = PlatformClient.with_secret_key("sk_test", base_url="https://api.test") + client._http._client = httpx.Client(transport=httpx.MockTransport(handler)) + + with public_request_timeout(2.5): + client.knowledge_documents.list(source=["cso_a"], page_size=10) + + assert _timeouts(seen[0]) == _all(2.5)