annotide 0.1.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.
- annotide/__init__.py +22 -0
- annotide/cli.py +214 -0
- annotide/client.py +660 -0
- annotide/errors.py +58 -0
- annotide/mcp_server.py +213 -0
- annotide/models.py +2683 -0
- annotide/py.typed +0 -0
- annotide-0.1.0.dist-info/METADATA +18 -0
- annotide-0.1.0.dist-info/RECORD +14 -0
- annotide-0.1.0.dist-info/WHEEL +5 -0
- annotide-0.1.0.dist-info/entry_points.txt +2 -0
- annotide-0.1.0.dist-info/licenses/LICENSE +201 -0
- annotide-0.1.0.dist-info/licenses/NOTICE +7 -0
- annotide-0.1.0.dist-info/top_level.txt +1 -0
annotide/client.py
ADDED
|
@@ -0,0 +1,660 @@
|
|
|
1
|
+
"""Synchronous client for the Annotide REST API (API-3).
|
|
2
|
+
|
|
3
|
+
Thin by design: one method per endpoint a script needs, plus a few helpers
|
|
4
|
+
(`export`, `import_file`, `take_snapshot`, `wait_for_job`) that chain them.
|
|
5
|
+
Responses are the API's JSON as plain dicts, typed with the generated
|
|
6
|
+
`annotide.models` TypedDicts.
|
|
7
|
+
"""
|
|
8
|
+
|
|
9
|
+
from __future__ import annotations
|
|
10
|
+
|
|
11
|
+
import json as jsonlib
|
|
12
|
+
import os
|
|
13
|
+
import time
|
|
14
|
+
import uuid
|
|
15
|
+
from collections.abc import Callable, Iterator, Mapping
|
|
16
|
+
from pathlib import Path
|
|
17
|
+
from typing import IO, Any, Literal, cast
|
|
18
|
+
from urllib.parse import urljoin, urlsplit
|
|
19
|
+
|
|
20
|
+
import httpx
|
|
21
|
+
|
|
22
|
+
from annotide.errors import AnnotationError, ApiError, JobFailedError, JobTimeoutError
|
|
23
|
+
from annotide.models import (
|
|
24
|
+
AnnotationRead,
|
|
25
|
+
DatasetFilter,
|
|
26
|
+
ImportStatus,
|
|
27
|
+
ItemRead,
|
|
28
|
+
JobRead,
|
|
29
|
+
JobStatusOutput,
|
|
30
|
+
JobTypeOutput,
|
|
31
|
+
LabelSchemaVersionRead,
|
|
32
|
+
ModelVersionCreate,
|
|
33
|
+
ModelVersionRead,
|
|
34
|
+
ProjectRead,
|
|
35
|
+
ProjectStats,
|
|
36
|
+
SnapshotRead,
|
|
37
|
+
TaskRead,
|
|
38
|
+
UserRead,
|
|
39
|
+
)
|
|
40
|
+
|
|
41
|
+
__all__ = ["Client"]
|
|
42
|
+
|
|
43
|
+
URL_ENV = "ANNOTIDE_URL"
|
|
44
|
+
API_KEY_ENV = "ANNOTIDE_API_KEY"
|
|
45
|
+
|
|
46
|
+
TERMINAL_JOB_STATUSES: frozenset[str] = frozenset({"succeeded", "failed", "cancelled"})
|
|
47
|
+
#: Gateway errors: the request may not have reached the app. Retried only
|
|
48
|
+
#: when a repeat is harmless (a read, or a create carrying an Idempotency-Key).
|
|
49
|
+
RETRY_STATUSES = frozenset({502, 503, 504})
|
|
50
|
+
IDEMPOTENT_METHODS = frozenset({"GET", "HEAD", "PUT", "DELETE"})
|
|
51
|
+
MAX_RETRY_AFTER = 60.0
|
|
52
|
+
PAGE_SIZE = 100
|
|
53
|
+
|
|
54
|
+
Split = Literal["train", "val", "test"]
|
|
55
|
+
|
|
56
|
+
|
|
57
|
+
def _drop_none(values: Mapping[str, Any]) -> dict[str, Any]:
|
|
58
|
+
return {key: value for key, value in values.items() if value is not None}
|
|
59
|
+
|
|
60
|
+
|
|
61
|
+
def _retry_after(response: httpx.Response) -> float:
|
|
62
|
+
try:
|
|
63
|
+
seconds = float(response.headers.get("Retry-After", "1"))
|
|
64
|
+
except ValueError:
|
|
65
|
+
seconds = 1.0
|
|
66
|
+
return min(max(seconds, 0.0), MAX_RETRY_AFTER)
|
|
67
|
+
|
|
68
|
+
|
|
69
|
+
def _backoff(attempt: int) -> float:
|
|
70
|
+
return float(min(0.5 * 2**attempt, 8.0))
|
|
71
|
+
|
|
72
|
+
|
|
73
|
+
def _problem_reasons(problem: Mapping[str, Any], limit: int = 5) -> str:
|
|
74
|
+
"""`loc: msg` for the first few field errors or violations, or ""."""
|
|
75
|
+
reasons: list[str] = []
|
|
76
|
+
for key in ("errors", "violations"):
|
|
77
|
+
entries = problem.get(key)
|
|
78
|
+
if not isinstance(entries, list):
|
|
79
|
+
continue
|
|
80
|
+
for entry in entries:
|
|
81
|
+
if isinstance(entry, dict):
|
|
82
|
+
where = ".".join(str(part) for part in entry.get("loc", []) if part != "body")
|
|
83
|
+
message = str(entry.get("msg") or entry.get("message") or entry)
|
|
84
|
+
reasons.append(f"{where}: {message}" if where else message)
|
|
85
|
+
else:
|
|
86
|
+
reasons.append(str(entry))
|
|
87
|
+
if len(reasons) > limit:
|
|
88
|
+
reasons = [*reasons[:limit], f"and {len(reasons) - limit} more"]
|
|
89
|
+
return "; ".join(reasons)
|
|
90
|
+
|
|
91
|
+
|
|
92
|
+
def _api_error(response: httpx.Response, method: str, path: str) -> ApiError:
|
|
93
|
+
try:
|
|
94
|
+
problem = response.json()
|
|
95
|
+
except ValueError:
|
|
96
|
+
problem = None
|
|
97
|
+
if not isinstance(problem, dict):
|
|
98
|
+
return ApiError(
|
|
99
|
+
response.status_code,
|
|
100
|
+
response.reason_phrase or "Error",
|
|
101
|
+
response.text or None,
|
|
102
|
+
method=method,
|
|
103
|
+
path=path,
|
|
104
|
+
)
|
|
105
|
+
detail = problem.get("detail")
|
|
106
|
+
if not isinstance(detail, str | None):
|
|
107
|
+
# FastAPI's 422 carries a list of validation errors in `detail`.
|
|
108
|
+
detail = str(detail)
|
|
109
|
+
# A 422 names each bad field in `errors` (and QA-6 rule breaks in
|
|
110
|
+
# `violations`): the part a caller, or an agent, needs to fix the request.
|
|
111
|
+
reasons = _problem_reasons(problem)
|
|
112
|
+
if reasons:
|
|
113
|
+
detail = f"{detail}: {reasons}" if detail else reasons
|
|
114
|
+
return ApiError(
|
|
115
|
+
int(problem.get("status") or response.status_code),
|
|
116
|
+
str(problem.get("title") or response.reason_phrase or "Error"),
|
|
117
|
+
detail,
|
|
118
|
+
type=str(problem.get("type") or "about:blank"),
|
|
119
|
+
method=method,
|
|
120
|
+
path=path,
|
|
121
|
+
)
|
|
122
|
+
|
|
123
|
+
|
|
124
|
+
class Client:
|
|
125
|
+
"""A connection to one platform install, authenticated with an API key.
|
|
126
|
+
|
|
127
|
+
`base_url` is the site root (`https://annotate.example.com`), not the
|
|
128
|
+
`/api/v1` prefix; both it and `api_key` fall back to the
|
|
129
|
+
`ANNOTIDE_URL` / `ANNOTIDE_API_KEY` environment variables. Mint a key
|
|
130
|
+
under Settings → API keys; a CI job should use a service account's key.
|
|
131
|
+
"""
|
|
132
|
+
|
|
133
|
+
def __init__(
|
|
134
|
+
self,
|
|
135
|
+
base_url: str | None = None,
|
|
136
|
+
api_key: str | None = None,
|
|
137
|
+
*,
|
|
138
|
+
timeout: float = 30.0,
|
|
139
|
+
max_retries: int = 3,
|
|
140
|
+
transport: httpx.BaseTransport | None = None,
|
|
141
|
+
sleep: Callable[[float], None] = time.sleep,
|
|
142
|
+
) -> None:
|
|
143
|
+
base_url = base_url or os.environ.get(URL_ENV)
|
|
144
|
+
api_key = api_key or os.environ.get(API_KEY_ENV)
|
|
145
|
+
if not base_url:
|
|
146
|
+
raise AnnotationError(f"no base URL: pass base_url or set {URL_ENV}")
|
|
147
|
+
if not api_key:
|
|
148
|
+
raise AnnotationError(f"no API key: pass api_key or set {API_KEY_ENV}")
|
|
149
|
+
self._root = base_url.rstrip("/") + "/"
|
|
150
|
+
self._max_retries = max_retries
|
|
151
|
+
self._sleep = sleep
|
|
152
|
+
self._http = httpx.Client(
|
|
153
|
+
base_url=urljoin(self._root, "api/v1/"),
|
|
154
|
+
headers={"Authorization": f"Bearer {api_key}", "Accept": "application/json"},
|
|
155
|
+
timeout=timeout,
|
|
156
|
+
transport=transport,
|
|
157
|
+
)
|
|
158
|
+
# Signed download URLs carry their own authorisation: never send the
|
|
159
|
+
# API key to the storage host.
|
|
160
|
+
self._storage = httpx.Client(timeout=timeout, transport=transport, follow_redirects=True)
|
|
161
|
+
|
|
162
|
+
def close(self) -> None:
|
|
163
|
+
self._http.close()
|
|
164
|
+
self._storage.close()
|
|
165
|
+
|
|
166
|
+
def __enter__(self) -> Client:
|
|
167
|
+
return self
|
|
168
|
+
|
|
169
|
+
def __exit__(self, *_: object) -> None:
|
|
170
|
+
self.close()
|
|
171
|
+
|
|
172
|
+
# ------------------------------------------------------------------
|
|
173
|
+
# Transport
|
|
174
|
+
# ------------------------------------------------------------------
|
|
175
|
+
|
|
176
|
+
def _send(
|
|
177
|
+
self,
|
|
178
|
+
method: str,
|
|
179
|
+
path: str,
|
|
180
|
+
*,
|
|
181
|
+
params: Mapping[str, Any] | None = None,
|
|
182
|
+
json: Any = None,
|
|
183
|
+
data: Mapping[str, Any] | None = None,
|
|
184
|
+
files: Mapping[str, tuple[str, IO[bytes], str]] | None = None,
|
|
185
|
+
idempotency_key: str | None = None,
|
|
186
|
+
) -> httpx.Response:
|
|
187
|
+
headers = {"Idempotency-Key": idempotency_key} if idempotency_key else None
|
|
188
|
+
repeatable = method in IDEMPOTENT_METHODS or idempotency_key is not None
|
|
189
|
+
attempt = 0
|
|
190
|
+
while True:
|
|
191
|
+
for _, stream, _ in (files or {}).values():
|
|
192
|
+
stream.seek(0)
|
|
193
|
+
try:
|
|
194
|
+
response = self._http.request(
|
|
195
|
+
method,
|
|
196
|
+
path.lstrip("/"),
|
|
197
|
+
params=_drop_none(params or {}),
|
|
198
|
+
json=json,
|
|
199
|
+
data=data,
|
|
200
|
+
files=files,
|
|
201
|
+
headers=headers,
|
|
202
|
+
)
|
|
203
|
+
except httpx.TransportError as exc:
|
|
204
|
+
# A failed connect never reached the server, so it is always
|
|
205
|
+
# safe to repeat; anything later only when repeatable.
|
|
206
|
+
safe = repeatable or isinstance(exc, httpx.ConnectError)
|
|
207
|
+
if not safe or attempt >= self._max_retries:
|
|
208
|
+
raise AnnotationError(f"{method} {path}: {exc}") from exc
|
|
209
|
+
self._sleep(_backoff(attempt))
|
|
210
|
+
attempt += 1
|
|
211
|
+
continue
|
|
212
|
+
if attempt < self._max_retries:
|
|
213
|
+
# The rate limiter answers before any handler runs, so a 429
|
|
214
|
+
# is safe to repeat whatever the method.
|
|
215
|
+
if response.status_code == 429:
|
|
216
|
+
self._sleep(_retry_after(response))
|
|
217
|
+
attempt += 1
|
|
218
|
+
continue
|
|
219
|
+
if response.status_code in RETRY_STATUSES and repeatable:
|
|
220
|
+
self._sleep(_backoff(attempt))
|
|
221
|
+
attempt += 1
|
|
222
|
+
continue
|
|
223
|
+
if response.is_error:
|
|
224
|
+
raise _api_error(response, method, path)
|
|
225
|
+
return response
|
|
226
|
+
|
|
227
|
+
def request(
|
|
228
|
+
self,
|
|
229
|
+
method: str,
|
|
230
|
+
path: str,
|
|
231
|
+
*,
|
|
232
|
+
params: Mapping[str, Any] | None = None,
|
|
233
|
+
json: Any = None,
|
|
234
|
+
idempotency_key: str | None = None,
|
|
235
|
+
) -> Any:
|
|
236
|
+
"""Call any endpoint under `/api/v1` and return its JSON (None for 204).
|
|
237
|
+
|
|
238
|
+
The escape hatch for endpoints without a method here; errors, retries
|
|
239
|
+
and authentication behave as for every other call.
|
|
240
|
+
"""
|
|
241
|
+
response = self._send(
|
|
242
|
+
method.upper(), path, params=params, json=json, idempotency_key=idempotency_key
|
|
243
|
+
)
|
|
244
|
+
if response.status_code == 204 or not response.content:
|
|
245
|
+
return None
|
|
246
|
+
return response.json()
|
|
247
|
+
|
|
248
|
+
def _create(self, path: str, body: Any, idempotency_key: str | None) -> Any:
|
|
249
|
+
"""POST a creation. The key makes a retried create answer the first row."""
|
|
250
|
+
return self.request("POST", path, json=body, idempotency_key=idempotency_key or _new_key())
|
|
251
|
+
|
|
252
|
+
def _paginate(self, path: str, params: Mapping[str, Any] | None = None) -> Iterator[Any]:
|
|
253
|
+
cursor: str | None = None
|
|
254
|
+
while True:
|
|
255
|
+
page = self.request(
|
|
256
|
+
"GET", path, params={**(params or {}), "limit": PAGE_SIZE, "cursor": cursor}
|
|
257
|
+
)
|
|
258
|
+
yield from page["items"]
|
|
259
|
+
cursor = page.get("next_cursor")
|
|
260
|
+
if not cursor:
|
|
261
|
+
return
|
|
262
|
+
|
|
263
|
+
# ------------------------------------------------------------------
|
|
264
|
+
# Identity
|
|
265
|
+
# ------------------------------------------------------------------
|
|
266
|
+
|
|
267
|
+
def me(self) -> UserRead:
|
|
268
|
+
"""The user (or service account) the API key acts as."""
|
|
269
|
+
return cast(UserRead, self.request("GET", "auth/me"))
|
|
270
|
+
|
|
271
|
+
# ------------------------------------------------------------------
|
|
272
|
+
# Projects and items
|
|
273
|
+
# ------------------------------------------------------------------
|
|
274
|
+
|
|
275
|
+
def list_projects(self) -> Iterator[ProjectRead]:
|
|
276
|
+
"""Every project the caller is a member of, following the cursor."""
|
|
277
|
+
return cast(Iterator[ProjectRead], self._paginate("projects"))
|
|
278
|
+
|
|
279
|
+
def get_project(self, project_id: str) -> ProjectRead:
|
|
280
|
+
return cast(ProjectRead, self.request("GET", f"projects/{project_id}"))
|
|
281
|
+
|
|
282
|
+
def create_project(
|
|
283
|
+
self, name: str, *, idempotency_key: str | None = None, **fields: Any
|
|
284
|
+
) -> ProjectRead:
|
|
285
|
+
"""Create a project; `fields` are the other `ProjectCreate` keys."""
|
|
286
|
+
body = {"name": name, **fields}
|
|
287
|
+
return cast(ProjectRead, self._create("projects", body, idempotency_key))
|
|
288
|
+
|
|
289
|
+
def get_stats(self, project_id: str) -> ProjectStats:
|
|
290
|
+
"""Dashboard numbers: items, tasks, annotations, review, class balance."""
|
|
291
|
+
return cast(ProjectStats, self.request("GET", f"projects/{project_id}/stats"))
|
|
292
|
+
|
|
293
|
+
def list_items(
|
|
294
|
+
self,
|
|
295
|
+
project_id: str,
|
|
296
|
+
*,
|
|
297
|
+
status: str | None = None,
|
|
298
|
+
media_type: str | None = None,
|
|
299
|
+
q: str | None = None,
|
|
300
|
+
) -> Iterator[ItemRead]:
|
|
301
|
+
"""The project's items, each with a short-lived signed `media_url`."""
|
|
302
|
+
params = {"status": status, "media_type": media_type, "q": q}
|
|
303
|
+
return cast(Iterator[ItemRead], self._paginate(f"projects/{project_id}/items", params))
|
|
304
|
+
|
|
305
|
+
def get_item(self, item_id: str) -> ItemRead:
|
|
306
|
+
return cast(ItemRead, self.request("GET", f"items/{item_id}"))
|
|
307
|
+
|
|
308
|
+
def list_annotations(self, item_id: str) -> list[AnnotationRead]:
|
|
309
|
+
"""Every annotation version of the item, newest first."""
|
|
310
|
+
return cast(list[AnnotationRead], self.request("GET", f"items/{item_id}/annotations"))
|
|
311
|
+
|
|
312
|
+
def download_media(self, item: ItemRead, *, max_bytes: int | None = None) -> bytes:
|
|
313
|
+
"""The item's media, fetched from storage on its signed `media_url` (ARC-3).
|
|
314
|
+
|
|
315
|
+
`max_bytes` refuses anything larger before reading it all.
|
|
316
|
+
"""
|
|
317
|
+
url = item.get("media_url")
|
|
318
|
+
if not url:
|
|
319
|
+
raise AnnotationError(f"item {item['id']} has no media URL")
|
|
320
|
+
with self._storage.stream("GET", urljoin(self._root, str(url))) as response:
|
|
321
|
+
if response.is_error:
|
|
322
|
+
response.read()
|
|
323
|
+
raise _api_error(response, "GET", "media download")
|
|
324
|
+
chunks: list[bytes] = []
|
|
325
|
+
size = 0
|
|
326
|
+
for chunk in response.iter_bytes():
|
|
327
|
+
size += len(chunk)
|
|
328
|
+
if max_bytes is not None and size > max_bytes:
|
|
329
|
+
raise AnnotationError(f"item {item['id']} is larger than {max_bytes} bytes")
|
|
330
|
+
chunks.append(chunk)
|
|
331
|
+
return b"".join(chunks)
|
|
332
|
+
|
|
333
|
+
def list_schema_versions(self, project_id: str) -> list[LabelSchemaVersionRead]:
|
|
334
|
+
"""The project's label schema versions, newest first."""
|
|
335
|
+
return cast(
|
|
336
|
+
list[LabelSchemaVersionRead], self.request("GET", f"projects/{project_id}/schemas")
|
|
337
|
+
)
|
|
338
|
+
|
|
339
|
+
def create_prelabel(
|
|
340
|
+
self,
|
|
341
|
+
item_id: str,
|
|
342
|
+
*,
|
|
343
|
+
model_version_id: str,
|
|
344
|
+
result: Mapping[str, Any],
|
|
345
|
+
label_schema_version_id: str | None = None,
|
|
346
|
+
) -> AnnotationRead:
|
|
347
|
+
"""Post a pre-label as an external producer (API-8): a model-authored draft."""
|
|
348
|
+
body = _drop_none(
|
|
349
|
+
{
|
|
350
|
+
"model_version_id": model_version_id,
|
|
351
|
+
"result": dict(result),
|
|
352
|
+
"label_schema_version_id": label_schema_version_id,
|
|
353
|
+
}
|
|
354
|
+
)
|
|
355
|
+
return cast(AnnotationRead, self.request("POST", f"items/{item_id}/prelabels", json=body))
|
|
356
|
+
|
|
357
|
+
# ------------------------------------------------------------------
|
|
358
|
+
# Tasks (WF-2, WF-3)
|
|
359
|
+
# ------------------------------------------------------------------
|
|
360
|
+
|
|
361
|
+
def claim_task(self, project_id: str, *, task_type: str = "annotate") -> TaskRead | None:
|
|
362
|
+
"""Claim the next open task (taking its lock), or None when the queue is empty."""
|
|
363
|
+
task = self.request(
|
|
364
|
+
"POST", "tasks/next", params={"project_id": project_id, "type": task_type}
|
|
365
|
+
)
|
|
366
|
+
return cast(TaskRead | None, task)
|
|
367
|
+
|
|
368
|
+
def release_task(self, task_id: str) -> TaskRead:
|
|
369
|
+
"""Give a claimed task back to the queue."""
|
|
370
|
+
return cast(TaskRead, self.request("POST", f"tasks/{task_id}/release"))
|
|
371
|
+
|
|
372
|
+
def scan(self, project_id: str, *, idempotency_key: str | None = None) -> JobRead:
|
|
373
|
+
"""Queue a scan of the project's source storage for new items."""
|
|
374
|
+
return cast(JobRead, self._create(f"projects/{project_id}/scan", {}, idempotency_key))
|
|
375
|
+
|
|
376
|
+
# ------------------------------------------------------------------
|
|
377
|
+
# Snapshots
|
|
378
|
+
# ------------------------------------------------------------------
|
|
379
|
+
|
|
380
|
+
def list_snapshots(self, project_id: str) -> Iterator[SnapshotRead]:
|
|
381
|
+
return cast(Iterator[SnapshotRead], self._paginate(f"projects/{project_id}/snapshots"))
|
|
382
|
+
|
|
383
|
+
def get_snapshot(self, project_id: str, snapshot_id: str) -> SnapshotRead:
|
|
384
|
+
return cast(
|
|
385
|
+
SnapshotRead, self.request("GET", f"projects/{project_id}/snapshots/{snapshot_id}")
|
|
386
|
+
)
|
|
387
|
+
|
|
388
|
+
def create_snapshot(
|
|
389
|
+
self,
|
|
390
|
+
project_id: str,
|
|
391
|
+
name: str,
|
|
392
|
+
*,
|
|
393
|
+
filter: DatasetFilter | None = None, # noqa: A002 - the API's field name
|
|
394
|
+
split: Mapping[str, Any] | None = None,
|
|
395
|
+
label_schema_version_id: str | None = None,
|
|
396
|
+
idempotency_key: str | None = None,
|
|
397
|
+
) -> JobRead:
|
|
398
|
+
"""Queue a snapshot job (EXP-1). `split` is `{train, val, test, seed, group_by}`."""
|
|
399
|
+
body = _drop_none(
|
|
400
|
+
{
|
|
401
|
+
"name": name,
|
|
402
|
+
"filter": filter,
|
|
403
|
+
"split": split,
|
|
404
|
+
"label_schema_version_id": label_schema_version_id,
|
|
405
|
+
}
|
|
406
|
+
)
|
|
407
|
+
return cast(
|
|
408
|
+
JobRead, self._create(f"projects/{project_id}/snapshots", body, idempotency_key)
|
|
409
|
+
)
|
|
410
|
+
|
|
411
|
+
def take_snapshot(
|
|
412
|
+
self,
|
|
413
|
+
project_id: str,
|
|
414
|
+
name: str,
|
|
415
|
+
*,
|
|
416
|
+
filter: DatasetFilter | None = None, # noqa: A002 - the API's field name
|
|
417
|
+
split: Mapping[str, Any] | None = None,
|
|
418
|
+
label_schema_version_id: str | None = None,
|
|
419
|
+
timeout: float = 600.0,
|
|
420
|
+
) -> SnapshotRead:
|
|
421
|
+
"""Create a snapshot, wait for it, and return the frozen snapshot."""
|
|
422
|
+
job = self.create_snapshot(
|
|
423
|
+
project_id,
|
|
424
|
+
name,
|
|
425
|
+
filter=filter,
|
|
426
|
+
split=split,
|
|
427
|
+
label_schema_version_id=label_schema_version_id,
|
|
428
|
+
)
|
|
429
|
+
done = self.wait_for_job(job["id"], timeout=timeout)
|
|
430
|
+
snapshot_id = (done["result"] or {}).get("snapshot_id")
|
|
431
|
+
if not snapshot_id:
|
|
432
|
+
raise AnnotationError(f"snapshot job {done['id']} returned no snapshot_id")
|
|
433
|
+
return self.get_snapshot(project_id, str(snapshot_id))
|
|
434
|
+
|
|
435
|
+
# ------------------------------------------------------------------
|
|
436
|
+
# Exports and imports
|
|
437
|
+
# ------------------------------------------------------------------
|
|
438
|
+
|
|
439
|
+
def create_export(
|
|
440
|
+
self,
|
|
441
|
+
project_id: str,
|
|
442
|
+
format: str, # noqa: A002 - the API's field name
|
|
443
|
+
*,
|
|
444
|
+
snapshot_id: str | None = None,
|
|
445
|
+
split: Split | None = None,
|
|
446
|
+
filter: DatasetFilter | None = None, # noqa: A002 - the API's field name
|
|
447
|
+
label_schema_version_id: str | None = None,
|
|
448
|
+
idempotency_key: str | None = None,
|
|
449
|
+
) -> JobRead:
|
|
450
|
+
"""Queue an export (EXP-5): `coco`, `yolo` or `native`.
|
|
451
|
+
|
|
452
|
+
Export a snapshot for a reproducible dataset; `split` (needs
|
|
453
|
+
`snapshot_id`) exports one partition of a split snapshot.
|
|
454
|
+
"""
|
|
455
|
+
body = _drop_none(
|
|
456
|
+
{
|
|
457
|
+
"format": format,
|
|
458
|
+
"snapshot_id": snapshot_id,
|
|
459
|
+
"split": split,
|
|
460
|
+
"filter": filter,
|
|
461
|
+
"label_schema_version_id": label_schema_version_id,
|
|
462
|
+
}
|
|
463
|
+
)
|
|
464
|
+
return cast(JobRead, self._create(f"projects/{project_id}/exports", body, idempotency_key))
|
|
465
|
+
|
|
466
|
+
def download_export(self, job_id: str, dest: str | Path) -> Path:
|
|
467
|
+
"""Stream a succeeded export's archive to `dest` (a file, or a directory).
|
|
468
|
+
|
|
469
|
+
The archive is written to `<dest>.part` first and renamed when
|
|
470
|
+
complete, so an interrupted download never leaves a truncated file.
|
|
471
|
+
"""
|
|
472
|
+
link = self.request("GET", f"jobs/{job_id}/download")
|
|
473
|
+
url = urljoin(self._root, str(link["url"]))
|
|
474
|
+
target = Path(dest)
|
|
475
|
+
if target.is_dir():
|
|
476
|
+
name = Path(urlsplit(url).path).name or f"{job_id}.zip"
|
|
477
|
+
target = target / name
|
|
478
|
+
partial = target.with_name(target.name + ".part")
|
|
479
|
+
with self._storage.stream("GET", url) as response:
|
|
480
|
+
if response.is_error:
|
|
481
|
+
response.read()
|
|
482
|
+
raise _api_error(response, "GET", "export download")
|
|
483
|
+
with partial.open("wb") as out:
|
|
484
|
+
for chunk in response.iter_bytes():
|
|
485
|
+
out.write(chunk)
|
|
486
|
+
partial.replace(target)
|
|
487
|
+
return target
|
|
488
|
+
|
|
489
|
+
def export(
|
|
490
|
+
self,
|
|
491
|
+
project_id: str,
|
|
492
|
+
format: str, # noqa: A002 - the API's field name
|
|
493
|
+
dest: str | Path,
|
|
494
|
+
*,
|
|
495
|
+
snapshot_id: str | None = None,
|
|
496
|
+
split: Split | None = None,
|
|
497
|
+
filter: DatasetFilter | None = None, # noqa: A002 - the API's field name
|
|
498
|
+
timeout: float = 1800.0,
|
|
499
|
+
) -> Path:
|
|
500
|
+
"""Queue an export, wait for it and download the archive to `dest`."""
|
|
501
|
+
job = self.create_export(
|
|
502
|
+
project_id, format, snapshot_id=snapshot_id, split=split, filter=filter
|
|
503
|
+
)
|
|
504
|
+
self.wait_for_job(job["id"], timeout=timeout)
|
|
505
|
+
return self.download_export(job["id"], dest)
|
|
506
|
+
|
|
507
|
+
def create_import(
|
|
508
|
+
self,
|
|
509
|
+
project_id: str,
|
|
510
|
+
format: str, # noqa: A002 - the API's field name
|
|
511
|
+
path: str,
|
|
512
|
+
*,
|
|
513
|
+
connector_id: str | None = None,
|
|
514
|
+
class_mapping: Mapping[str, str | None] | None = None,
|
|
515
|
+
status: ImportStatus | None = None,
|
|
516
|
+
dry_run: bool = False,
|
|
517
|
+
idempotency_key: str | None = None,
|
|
518
|
+
) -> JobRead:
|
|
519
|
+
"""Queue an import (EXP-6) of a file already in storage, at `path`.
|
|
520
|
+
|
|
521
|
+
Formats: `coco`, `yolo`, `voc`, `cvat`, `label_studio`. A `dry_run`
|
|
522
|
+
reports what would be imported without writing anything.
|
|
523
|
+
"""
|
|
524
|
+
body = _drop_none(
|
|
525
|
+
{
|
|
526
|
+
"format": format,
|
|
527
|
+
"path": path,
|
|
528
|
+
"connector_id": connector_id,
|
|
529
|
+
"class_mapping": dict(class_mapping) if class_mapping is not None else None,
|
|
530
|
+
"status": status,
|
|
531
|
+
"dry_run": dry_run,
|
|
532
|
+
}
|
|
533
|
+
)
|
|
534
|
+
return cast(JobRead, self._create(f"projects/{project_id}/imports", body, idempotency_key))
|
|
535
|
+
|
|
536
|
+
def upload_import(
|
|
537
|
+
self,
|
|
538
|
+
project_id: str,
|
|
539
|
+
file: str | Path,
|
|
540
|
+
format: str, # noqa: A002 - the API's field name
|
|
541
|
+
*,
|
|
542
|
+
class_mapping: Mapping[str, str | None] | None = None,
|
|
543
|
+
status: ImportStatus | None = None,
|
|
544
|
+
dry_run: bool = False,
|
|
545
|
+
) -> JobRead:
|
|
546
|
+
"""Upload a local annotation file and queue its import (EXP-6).
|
|
547
|
+
|
|
548
|
+
This endpoint takes no `Idempotency-Key`, so a gateway error is not
|
|
549
|
+
retried (a repeat could import twice); a 429 or a failed connect is.
|
|
550
|
+
"""
|
|
551
|
+
source = Path(file)
|
|
552
|
+
data = _drop_none(
|
|
553
|
+
{
|
|
554
|
+
"format": format,
|
|
555
|
+
"status": status,
|
|
556
|
+
"dry_run": "true" if dry_run else "false",
|
|
557
|
+
"class_mapping": jsonlib.dumps(dict(class_mapping)) if class_mapping else None,
|
|
558
|
+
}
|
|
559
|
+
)
|
|
560
|
+
with source.open("rb") as stream:
|
|
561
|
+
response = self._send(
|
|
562
|
+
"POST",
|
|
563
|
+
f"projects/{project_id}/imports/upload",
|
|
564
|
+
data=data,
|
|
565
|
+
files={"file": (source.name, stream, "application/octet-stream")},
|
|
566
|
+
)
|
|
567
|
+
return cast(JobRead, response.json())
|
|
568
|
+
|
|
569
|
+
def import_file(
|
|
570
|
+
self,
|
|
571
|
+
project_id: str,
|
|
572
|
+
file: str | Path,
|
|
573
|
+
format: str, # noqa: A002 - the API's field name
|
|
574
|
+
*,
|
|
575
|
+
class_mapping: Mapping[str, str | None] | None = None,
|
|
576
|
+
status: ImportStatus | None = None,
|
|
577
|
+
dry_run: bool = False,
|
|
578
|
+
timeout: float = 1800.0,
|
|
579
|
+
) -> JobRead:
|
|
580
|
+
"""Upload and import a local file, wait, and return the finished job.
|
|
581
|
+
|
|
582
|
+
The job's `result` carries the import tally (items matched, shapes
|
|
583
|
+
written, unmapped classes).
|
|
584
|
+
"""
|
|
585
|
+
job = self.upload_import(
|
|
586
|
+
project_id,
|
|
587
|
+
file,
|
|
588
|
+
format,
|
|
589
|
+
class_mapping=class_mapping,
|
|
590
|
+
status=status,
|
|
591
|
+
dry_run=dry_run,
|
|
592
|
+
)
|
|
593
|
+
return self.wait_for_job(job["id"], timeout=timeout)
|
|
594
|
+
|
|
595
|
+
# ------------------------------------------------------------------
|
|
596
|
+
# Models (BYOM-2, EXP-8)
|
|
597
|
+
# ------------------------------------------------------------------
|
|
598
|
+
|
|
599
|
+
def create_model_version(self, model_id: str, body: ModelVersionCreate) -> ModelVersionRead:
|
|
600
|
+
"""Register a trained version of a model.
|
|
601
|
+
|
|
602
|
+
With `snapshot_id` and `snapshot_digest` the version records what it
|
|
603
|
+
was trained on (EXP-8); the platform answers 409 when the digest is
|
|
604
|
+
not that snapshot's. The route takes no `Idempotency-Key`, so a
|
|
605
|
+
gateway error is not retried.
|
|
606
|
+
"""
|
|
607
|
+
return cast(
|
|
608
|
+
ModelVersionRead, self.request("POST", f"models/{model_id}/versions", json=body)
|
|
609
|
+
)
|
|
610
|
+
|
|
611
|
+
# ------------------------------------------------------------------
|
|
612
|
+
# Jobs
|
|
613
|
+
# ------------------------------------------------------------------
|
|
614
|
+
|
|
615
|
+
def list_jobs(
|
|
616
|
+
self,
|
|
617
|
+
project_id: str,
|
|
618
|
+
*,
|
|
619
|
+
status: JobStatusOutput | None = None,
|
|
620
|
+
type: JobTypeOutput | None = None, # noqa: A002 - the API's field name
|
|
621
|
+
) -> Iterator[JobRead]:
|
|
622
|
+
params = {"status": status, "type": type}
|
|
623
|
+
return cast(Iterator[JobRead], self._paginate(f"projects/{project_id}/jobs", params))
|
|
624
|
+
|
|
625
|
+
def get_job(self, job_id: str) -> JobRead:
|
|
626
|
+
return cast(JobRead, self.request("GET", f"jobs/{job_id}"))
|
|
627
|
+
|
|
628
|
+
def cancel_job(self, job_id: str) -> JobRead:
|
|
629
|
+
return cast(JobRead, self.request("POST", f"jobs/{job_id}/cancel"))
|
|
630
|
+
|
|
631
|
+
def retry_job(self, job_id: str) -> JobRead:
|
|
632
|
+
return cast(JobRead, self.request("POST", f"jobs/{job_id}/retry"))
|
|
633
|
+
|
|
634
|
+
def wait_for_job(
|
|
635
|
+
self,
|
|
636
|
+
job_id: str,
|
|
637
|
+
*,
|
|
638
|
+
timeout: float = 600.0,
|
|
639
|
+
interval: float = 2.0,
|
|
640
|
+
raise_on_failure: bool = True,
|
|
641
|
+
) -> JobRead:
|
|
642
|
+
"""Poll until the job is `succeeded`, `failed` or `cancelled`.
|
|
643
|
+
|
|
644
|
+
Raises `JobFailedError` for a job that did not succeed (unless
|
|
645
|
+
`raise_on_failure` is false) and `JobTimeoutError` when `timeout` passes.
|
|
646
|
+
"""
|
|
647
|
+
deadline = time.monotonic() + timeout
|
|
648
|
+
while True:
|
|
649
|
+
job = self.get_job(job_id)
|
|
650
|
+
if job["status"] in TERMINAL_JOB_STATUSES:
|
|
651
|
+
if raise_on_failure and job["status"] != "succeeded":
|
|
652
|
+
raise JobFailedError(job)
|
|
653
|
+
return job
|
|
654
|
+
if time.monotonic() >= deadline:
|
|
655
|
+
raise JobTimeoutError(job, timeout)
|
|
656
|
+
self._sleep(interval)
|
|
657
|
+
|
|
658
|
+
|
|
659
|
+
def _new_key() -> str:
|
|
660
|
+
return uuid.uuid4().hex
|