Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
24 changes: 24 additions & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
1 change: 1 addition & 0 deletions src/archastro/platform/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
58 changes: 56 additions & 2 deletions src/archastro/platform/runtime/http_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,16 +7,58 @@
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

import httpx
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``.
Expand Down Expand Up @@ -72,14 +114,16 @@ 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
self._get_access_token = get_access_token
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

Expand Down Expand Up @@ -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)}"

Expand All @@ -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(
Expand Down Expand Up @@ -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"}
Expand All @@ -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()
Expand Down Expand Up @@ -347,14 +395,16 @@ 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
self._get_access_token = get_access_token
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()

Expand Down Expand Up @@ -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)}"

Expand All @@ -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(
Expand Down Expand Up @@ -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"}
Expand All @@ -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()
Expand Down
Loading
Loading