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
3 changes: 3 additions & 0 deletions docs/WORKFLOWS.md
Original file line number Diff line number Diff line change
Expand Up @@ -77,6 +77,8 @@ contract.

Credential handling: a caller-supplied `Authorization` header is always used as-is and never overwritten. With no header, `NEBULA_API_TOKEN` is injected only when the call is on the Nebula url (no custom `spec.api.url`) — the Nebula token is never sent to a custom endpoint. A Nebula-path call with no token available fails closed.

`spec.api.retries` (default `0`, at most `10`) sets how many times a transient failure is retried before the task fails. A transient failure is a connection error or an HTTP status of 5xx, 408, or 429; other 4xx statuses are never retried. Retries back off exponentially: the first waits 1s and each later one doubles, capped at 60s. When a retryable response carries a `Retry-After` header (seconds or an HTTP date), that wait is used instead, also capped at 60s. Each retry logs a warning with the attempt count and the wait. A cancelled task stops retrying immediately.

```yaml
spec:
taskType: api
Expand All @@ -89,6 +91,7 @@ spec:
messages:
- role: user
content: Hello
retries: 3
response:
parse_json: true
```
Expand Down
158 changes: 152 additions & 6 deletions src/worker/executors/api_executor.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,9 @@
import email.utils
import logging
import math
import os
import threading
from datetime import UTC, datetime
from pathlib import Path
from typing import Any, ClassVar

Expand All @@ -11,13 +14,25 @@
from shared.tasks.task_type import TaskType
from shared.utils.redact import is_credential_key

from .base_executor import ExecutionError, Executor, ExecutorTask
from .base_executor import ExecutionError, Executor, ExecutorTask, TaskCancelledError

logger = logging.getLogger(__name__)

# Cache key: (base_url, timeout_seconds, verify_tls, follow_redirects)
_ClientKey = tuple[str, float, bool, bool]

# Base delay between retry attempts, doubled each retry.
_RETRY_BACKOFF_SEC = 1.0
# Upper bound on any single retry wait.
_RETRY_BACKOFF_MAX_SEC = 60.0
# Upper bound on spec.api.retries.
_MAX_RETRIES = 10


def _is_retryable_status(status_code: int) -> bool:
"""Whether an HTTP status is transient and worth retrying."""
return status_code >= 500 or status_code in (408, 429)


class APIExecutor(Executor):
"""Performs a single HTTP request defined by task YAML.
Expand All @@ -36,6 +51,24 @@ class APIExecutor(Executor):
_clients: ClassVar[dict[_ClientKey, httpx.Client]] = {}
_clients_lock: ClassVar[threading.Lock] = threading.Lock()

def __init__(self, *args: Any, **kwargs: Any) -> None:
super().__init__(*args, **kwargs)
self._cancel_event = threading.Event()
self._cancel_lock = threading.Lock()
self._active_task_id: str | None = None
self._pending_cancelled_ids: set[str] = set()

def cancel(self, task_id: str) -> None:
Comment thread
kaiitunnz marked this conversation as resolved.
"""Signal the executor to abort the current request and any retries."""
with self._cancel_lock:
if self._active_task_id is None:
# No run in flight: record the id so a pre-start cancel lands.
self._pending_cancelled_ids.add(task_id)
self._cancel_event.set()
elif self._active_task_id == task_id:
self._cancel_event.set()
# A cancellation for a different task than the active one is ignored.

@classmethod
def _base_url(cls, url: str) -> str:
"""Extract scheme + host + port from a URL for pool keying."""
Expand Down Expand Up @@ -78,6 +111,98 @@ def _get_client(
)
return client

def _request_with_retries(
Comment thread
kaiitunnz marked this conversation as resolved.
self,
client: httpx.Client,
method: str,
url: str,
headers: dict[str, Any],
params: dict[str, Any] | None,
request_kwargs: dict[str, Any],
retries: int,
) -> httpx.Response:
"""Issue the request, retrying transient failures up to ``retries`` times.

A retryable failure is a connection error or a transient HTTP status
(5xx, 408, 429). Non-retryable failures and a cancelled task stop the
Comment thread
kaiitunnz marked this conversation as resolved.
loop immediately. The final attempt's failure propagates to the caller.
"""
attempt = 0
while True:
if self._cancel_event.is_set():
raise TaskCancelledError("API request cancelled")
try:
resp = client.request(
method,
url,
headers=headers,
params=params,
**request_kwargs,
)
except httpx.RequestError as exc:
if attempt >= retries:
raise
attempt += 1
delay = self._backoff_delay(attempt)
logger.warning(
"API request failed (attempt %d/%d): %s; retrying in %.1fs",
attempt,
retries,
exc,
delay,
)
self._wait_for_backoff(delay)
continue
if resp.is_error and _is_retryable_status(resp.status_code):
if attempt >= retries:
return resp
attempt += 1
delay = self._backoff_delay(attempt, resp)
logger.warning(
"API request returned %s (attempt %d/%d); retrying in %.1fs",
resp.status_code,
attempt,
retries,
delay,
)
self._wait_for_backoff(delay)
continue
return resp

def _backoff_delay(self, attempt: int, resp: httpx.Response | None = None) -> float:
"""Return the wait before the next attempt, honouring Retry-After."""
if resp is not None:
retry_after = self._retry_after_seconds(resp)
if retry_after is not None:
return min(retry_after, _RETRY_BACKOFF_MAX_SEC)
return min(_RETRY_BACKOFF_SEC * (2 ** (attempt - 1)), _RETRY_BACKOFF_MAX_SEC)

@staticmethod
def _retry_after_seconds(resp: httpx.Response) -> float | None:
"""Parse the Retry-After header as seconds or an HTTP date."""
value = resp.headers.get("Retry-After")
if value is None:
return None
try:
seconds = float(value)
except ValueError:
pass
else:
if math.isfinite(seconds) and seconds >= 0:
return seconds
try:
retry_at = email.utils.parsedate_to_datetime(value)
except (TypeError, ValueError):
return None
if retry_at.tzinfo is None:
retry_at = retry_at.replace(tzinfo=UTC)
return max(0.0, (retry_at - datetime.now(UTC)).total_seconds())

def _wait_for_backoff(self, delay: float) -> None:
"""Wait out the retry backoff, aborting early if the task is cancelled."""
if self._cancel_event.wait(delay):
raise TaskCancelledError("API request cancelled")

@classmethod
def close_all_clients(cls) -> None:
"""Close and discard all cached HTTP clients."""
Expand All @@ -95,6 +220,19 @@ def cleanup_after_run(self) -> None:
self.close_all_clients()

def run(self, task: ExecutorTask, out_dir: Path) -> APIResult:
with self._cancel_lock:
self._active_task_id = task.task_id
# Event is set iff this task's id was pending; other ids are stale.
if task.task_id not in self._pending_cancelled_ids:
self._cancel_event.clear()
self._pending_cancelled_ids.clear()
try:
return self._run(task, out_dir)
finally:
with self._cancel_lock:
self._active_task_id = None

def _run(self, task: ExecutorTask, out_dir: Path) -> APIResult:
spec = self.require_spec(task, ApiSpecStrict)
api_cfg = spec.api or {}
if not isinstance(api_cfg, dict):
Expand Down Expand Up @@ -164,15 +302,23 @@ def run(self, task: ExecutorTask, out_dir: Path) -> APIResult:
raise_for_status = bool(response_cfg.get("raise_for_status", True))
max_body_bytes = int(response_cfg.get("max_body_bytes", 200000))

retries = api_cfg.get("retries", 0)
Comment thread
kaiitunnz marked this conversation as resolved.
if not isinstance(retries, int) or isinstance(retries, bool) or retries < 0:
raise ExecutionError("spec.api.retries must be a non-negative integer")
if retries > _MAX_RETRIES:
raise ExecutionError(f"spec.api.retries must be at most {_MAX_RETRIES}")

try:
base = self._base_url(str(url))
client = self._get_client(base, timeout, verify_tls, follow_redirects)
resp = client.request(
resp = self._request_with_retries(
client,
method,
str(url),
headers=headers,
params=params,
**request_kwargs,
headers,
params,
request_kwargs,
retries,
)
except httpx.RequestError as exc:
raise ExecutionError(f"API request failed: {exc}", retryable=True) from exc
Expand Down Expand Up @@ -225,7 +371,7 @@ def run(self, task: ExecutorTask, out_dir: Path) -> APIResult:
message = f"API request returned status {resp.status_code}"
if body_text:
message = f"{message}: {body_text[:200]}"
retryable = resp.status_code >= 500 or resp.status_code in (408, 429)
retryable = _is_retryable_status(resp.status_code)
raise ExecutionError(message, retryable=retryable)

return result
Loading
Loading