arex-python-sdk 0.2.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
arex/errors.py ADDED
@@ -0,0 +1,334 @@
1
+ """Error hierarchy raised by the AREX SDK."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import re
6
+ from collections.abc import Mapping
7
+ from datetime import datetime, timezone
8
+ from email.utils import parsedate_to_datetime
9
+ from typing import TYPE_CHECKING, Any, Literal
10
+
11
+ if TYPE_CHECKING:
12
+ from arex.models import ResearchRun
13
+
14
+ __all__ = [
15
+ "ArexApiError",
16
+ "ArexError",
17
+ "ArexNetworkError",
18
+ "ArexTimeoutError",
19
+ "AuthenticationError",
20
+ "CreditsExhaustedError",
21
+ "InvalidRequestError",
22
+ "NotFoundError",
23
+ "PermissionDeniedError",
24
+ "RateLimitError",
25
+ "ResearchRunCancelledError",
26
+ "ResearchRunFailedError",
27
+ "ServiceUnavailableError",
28
+ "SourceNotAvailableError",
29
+ "api_error_from_response",
30
+ ]
31
+
32
+ NetworkErrorCode = Literal["request_failed", "invalid_response"]
33
+ TimeoutPhase = Literal["request", "poll"]
34
+
35
+
36
+ class ArexError(Exception):
37
+ """Base class for every error raised by this SDK.
38
+
39
+ `run_id` names the research run being polled when the error surfaced, so
40
+ the caller can `wait(run_id)` again or `cancel(run_id)`. It is `None`
41
+ outside `wait`/`run`.
42
+ """
43
+
44
+ run_id: str | None = None
45
+
46
+
47
+ class ArexApiError(ArexError):
48
+ """An error response returned by the AREX API."""
49
+
50
+ def __init__(
51
+ self,
52
+ message: str,
53
+ *,
54
+ status: int,
55
+ code: str,
56
+ retryable: bool = False,
57
+ retry_after_ms: int | None = None,
58
+ request_id: str | None = None,
59
+ param: str | None = None,
60
+ details: Mapping[str, Any] | None = None,
61
+ ) -> None:
62
+ super().__init__(message)
63
+ self.status = status
64
+ self.code = code
65
+ self.retryable = retryable
66
+ self.retry_after_ms = retry_after_ms
67
+ self.request_id = request_id
68
+ self.param = param
69
+ self.details: Mapping[str, Any] = dict(details or {})
70
+
71
+
72
+ class AuthenticationError(ArexApiError):
73
+ """401 `unauthorized`."""
74
+
75
+
76
+ class PermissionDeniedError(ArexApiError):
77
+ """403 permission or account access denied."""
78
+
79
+
80
+ class InvalidRequestError(ArexApiError):
81
+ """Invalid request, blocked content, reused key, or client-side validation."""
82
+
83
+ def __init__(
84
+ self,
85
+ message: str,
86
+ *,
87
+ status: int = 0,
88
+ code: str = "invalid_request",
89
+ retryable: bool = False,
90
+ retry_after_ms: int | None = None,
91
+ request_id: str | None = None,
92
+ param: str | None = None,
93
+ details: Mapping[str, Any] | None = None,
94
+ ) -> None:
95
+ super().__init__(
96
+ message,
97
+ status=status,
98
+ code=code,
99
+ retryable=retryable,
100
+ retry_after_ms=retry_after_ms,
101
+ request_id=request_id,
102
+ param=param,
103
+ details=details,
104
+ )
105
+
106
+
107
+ class SourceNotAvailableError(ArexApiError):
108
+ """422 `source_not_available`, raised before the request is charged."""
109
+
110
+
111
+ class NotFoundError(ArexApiError):
112
+ """404 `not_found`."""
113
+
114
+
115
+ class RateLimitError(ArexApiError):
116
+ """Rate limits, outstanding runs, or an idempotent request still in progress."""
117
+
118
+
119
+ class CreditsExhaustedError(ArexApiError):
120
+ """429 `daily_credits_exhausted` / `weekly_credits_exhausted`."""
121
+
122
+ def __init__(
123
+ self,
124
+ message: str,
125
+ *,
126
+ status: int,
127
+ code: str,
128
+ retryable: bool = False,
129
+ retry_after_ms: int | None = None,
130
+ request_id: str | None = None,
131
+ param: str | None = None,
132
+ details: Mapping[str, Any] | None = None,
133
+ resource: str | None = None,
134
+ used: int | float | None = None,
135
+ limit: int | float | None = None,
136
+ requested: int | float | None = None,
137
+ reset_at: str | None = None,
138
+ ) -> None:
139
+ super().__init__(
140
+ message,
141
+ status=status,
142
+ code=code,
143
+ retryable=retryable,
144
+ retry_after_ms=retry_after_ms,
145
+ request_id=request_id,
146
+ param=param,
147
+ details=details,
148
+ )
149
+ self.resource = resource
150
+ self.used = used
151
+ self.limit = limit
152
+ self.requested = requested
153
+ self.reset_at = reset_at
154
+
155
+
156
+ class ServiceUnavailableError(ArexApiError):
157
+ """503 `temporarily_unavailable` / 500 `internal_error`."""
158
+
159
+
160
+ class ArexNetworkError(ArexError):
161
+ """The request never produced a usable response."""
162
+
163
+ def __init__(self, code: NetworkErrorCode, message: str | None = None) -> None:
164
+ super().__init__(message or _NETWORK_MESSAGES[code])
165
+ self.code = code
166
+
167
+
168
+ _NETWORK_MESSAGES: dict[str, str] = {
169
+ "request_failed": "The request to the AREX API failed.",
170
+ "invalid_response": "The AREX API returned a response the SDK cannot parse.",
171
+ }
172
+
173
+
174
+ class ArexTimeoutError(ArexError):
175
+ """A request or a research poll exceeded its deadline."""
176
+
177
+ def __init__(
178
+ self,
179
+ phase: TimeoutPhase,
180
+ timeout_ms: int,
181
+ run_id: str | None = None,
182
+ ) -> None:
183
+ subject = "Research run" if phase == "poll" else "Request"
184
+ suffix = f" (run {run_id})" if run_id else ""
185
+ super().__init__(f"{subject} timed out after {timeout_ms} ms{suffix}.")
186
+ self.phase = phase
187
+ self.timeout_ms = timeout_ms
188
+ self.run_id = run_id
189
+
190
+
191
+ class ResearchRunFailedError(ArexError):
192
+ """A research run reached the terminal `failed` status."""
193
+
194
+ def __init__(self, run: ResearchRun) -> None:
195
+ retryable = run.error.retryable if run.error else False
196
+ super().__init__(f"Research run {run.id} failed (retryable={retryable}).")
197
+ self.run = run
198
+ self.run_id = run.id
199
+
200
+
201
+ class ResearchRunCancelledError(ArexError):
202
+ """A research run reached the terminal `cancelled` status."""
203
+
204
+ def __init__(self, run: ResearchRun) -> None:
205
+ super().__init__(f"Research run {run.id} was cancelled.")
206
+ self.run = run
207
+ self.run_id = run.id
208
+
209
+
210
+ _ERROR_CLASSES: dict[str, type[ArexApiError]] = {
211
+ "invalid_request": InvalidRequestError,
212
+ "content_blocked": InvalidRequestError,
213
+ "idempotency_key_reused": InvalidRequestError,
214
+ "source_not_available": SourceNotAvailableError,
215
+ "not_found": NotFoundError,
216
+ "rate_limit_exceeded": RateLimitError,
217
+ "idempotency_key_in_progress": RateLimitError,
218
+ "outstanding_run_limit": RateLimitError,
219
+ "temporarily_unavailable": ServiceUnavailableError,
220
+ "internal_error": ServiceUnavailableError,
221
+ }
222
+
223
+
224
+ _CREDIT_CODES = frozenset({"daily_credits_exhausted", "weekly_credits_exhausted"})
225
+
226
+
227
+ def _header(headers: Mapping[str, str], name: str) -> str | None:
228
+ target = name.lower()
229
+ for key, value in headers.items():
230
+ if key.lower() == target:
231
+ return value
232
+ return None
233
+
234
+
235
+ # Longest wait the SDK will report. A server asking for more than this means
236
+ # "not soon", and clamping keeps the arithmetic and `sleep` calls in range.
237
+ _MAX_RETRY_AFTER_MS = 2**31 - 1
238
+ _DELTA_SECONDS = re.compile(r"[0-9]+(?:\.[0-9]+)?")
239
+
240
+
241
+ def _clamp_ms(ms: float) -> int:
242
+ return _MAX_RETRY_AFTER_MS if ms >= _MAX_RETRY_AFTER_MS else max(0, int(ms))
243
+
244
+
245
+ def retry_after_ms_from_headers(headers: Mapping[str, str]) -> int | None:
246
+ """`Retry-After` in milliseconds; accepts delta-seconds or an HTTP date.
247
+
248
+ Malformed values give `None`; huge ones clamp instead of overflowing.
249
+ """
250
+ raw = _header(headers, "Retry-After")
251
+ if raw is None:
252
+ return None
253
+ text = raw.strip()
254
+ if _DELTA_SECONDS.fullmatch(text):
255
+ # ASCII digits never fail `float`, but a long run of them is `inf`.
256
+ return _clamp_ms(float(text) * 1000)
257
+ try:
258
+ at = parsedate_to_datetime(text)
259
+ except (TypeError, ValueError, IndexError, OverflowError):
260
+ return None
261
+ if at.tzinfo is None:
262
+ # A date without a zone (`-0000`) is still GMT per RFC 9110.
263
+ at = at.replace(tzinfo=timezone.utc)
264
+ return _clamp_ms((at - datetime.now(timezone.utc)).total_seconds() * 1000)
265
+
266
+
267
+ def _optional_number(value: Any) -> int | float | None:
268
+ return (
269
+ value
270
+ if isinstance(value, (int, float)) and not isinstance(value, bool)
271
+ else None
272
+ )
273
+
274
+
275
+ def api_error_from_response(
276
+ status: int,
277
+ body: Any,
278
+ headers: Mapping[str, str],
279
+ ) -> ArexApiError:
280
+ """Map a wire error response onto the matching SDK error class."""
281
+ retry_after_ms = retry_after_ms_from_headers(headers)
282
+
283
+ if not isinstance(body, dict) or not isinstance(body.get("error"), dict):
284
+ raise ArexNetworkError("invalid_response")
285
+ details: dict[str, Any] = dict(body["error"])
286
+ code = details.get("code")
287
+ message = details.get("message")
288
+ retryable = details.get("retryable")
289
+ request_id = details.get("request_id")
290
+ param = details.get("param")
291
+ if (
292
+ not isinstance(code, str)
293
+ or not isinstance(message, str)
294
+ or not isinstance(retryable, bool)
295
+ or (request_id is not None and not isinstance(request_id, str))
296
+ or (param is not None and not isinstance(param, str))
297
+ ):
298
+ raise ArexNetworkError("invalid_response")
299
+
300
+ if code in _CREDIT_CODES:
301
+ resource = details.get("resource")
302
+ reset_at = details.get("reset_at")
303
+ return CreditsExhaustedError(
304
+ message,
305
+ status=status,
306
+ code=code,
307
+ retryable=retryable,
308
+ retry_after_ms=retry_after_ms,
309
+ request_id=request_id,
310
+ param=param,
311
+ details=details,
312
+ resource=resource if isinstance(resource, str) else None,
313
+ used=_optional_number(details.get("used")),
314
+ limit=_optional_number(details.get("limit")),
315
+ requested=_optional_number(details.get("requested")),
316
+ reset_at=reset_at if isinstance(reset_at, str) else None,
317
+ )
318
+
319
+ if status == 401:
320
+ cls: type[ArexApiError] = AuthenticationError
321
+ elif status == 403:
322
+ cls = PermissionDeniedError
323
+ else:
324
+ cls = _ERROR_CLASSES.get(code, ArexApiError)
325
+ return cls(
326
+ message,
327
+ status=status,
328
+ code=code,
329
+ retryable=retryable,
330
+ retry_after_ms=retry_after_ms,
331
+ request_id=request_id,
332
+ param=param,
333
+ details=details,
334
+ )