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/__init__.py +137 -0
- arex/_poll.py +126 -0
- arex/_transport.py +303 -0
- arex/_validate.py +156 -0
- arex/_wire.py +126 -0
- arex/async_client.py +133 -0
- arex/client.py +197 -0
- arex/errors.py +334 -0
- arex/models.py +414 -0
- arex/py.typed +0 -0
- arex/research.py +363 -0
- arex_python_sdk-0.2.0.dist-info/METADATA +55 -0
- arex_python_sdk-0.2.0.dist-info/RECORD +15 -0
- arex_python_sdk-0.2.0.dist-info/WHEEL +4 -0
- arex_python_sdk-0.2.0.dist-info/licenses/LICENSE +203 -0
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
|
+
)
|