labtasker-client 2.0.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.
labtasker/worker.py ADDED
@@ -0,0 +1,473 @@
1
+ from __future__ import annotations
2
+
3
+ import functools
4
+ import json
5
+ import logging
6
+ import math
7
+ import os
8
+ import secrets
9
+ import threading
10
+ import time
11
+ import traceback as traceback_module
12
+ from collections.abc import Callable
13
+ from typing import ParamSpec, TypeVar, cast
14
+
15
+ from labtasker.binding import CompiledBinding, compile_binding
16
+ from labtasker.client import Client
17
+ from labtasker.errors import (
18
+ APIError,
19
+ ConfigError,
20
+ FatalWorkerError,
21
+ TransientError,
22
+ TransportError,
23
+ )
24
+ from labtasker.execution import (
25
+ ExecutionContext,
26
+ RunControl,
27
+ _validate_force_stop_timeout,
28
+ activate_context,
29
+ active_context_present,
30
+ deactivate_context,
31
+ )
32
+ from labtasker.journal import LocalRunJournal
33
+ from labtasker.models import ClaimResponse, TaskInfo
34
+ from labtasker.tee import WorkerTee, configure_worker_logger
35
+ from labtasker.types import JSONValue
36
+ from labtasker.validation import RequestValidationError, validate_identifier
37
+
38
+ P = ParamSpec("P")
39
+ R = TypeVar("R")
40
+ HEARTBEAT_INTERVAL_SECONDS = 60.0
41
+ POLL_INTERVAL_SECONDS = 1.0
42
+ TERMINAL_BACKOFF_SECONDS = (0.25, 0.5, 1.0, 2.0, 5.0, 10.0, 30.0)
43
+ MAX_REQUEST_BYTES = 1024 * 1024
44
+ logger = logging.getLogger("labtasker.worker")
45
+
46
+
47
+ class Heartbeat:
48
+ def __init__(
49
+ self,
50
+ client: Client,
51
+ *,
52
+ queue: str,
53
+ task_id: str,
54
+ run_id: str,
55
+ control: RunControl,
56
+ ) -> None:
57
+ self._client = client
58
+ self._queue = queue
59
+ self._task_id = task_id
60
+ self._run_id = run_id
61
+ self._control = control
62
+ self._stop = threading.Event()
63
+ self._thread = threading.Thread(
64
+ target=self._run,
65
+ name=f"labtasker-heartbeat-{run_id}",
66
+ daemon=True,
67
+ )
68
+
69
+ def start(self) -> None:
70
+ self._thread.start()
71
+
72
+ def stop(self) -> None:
73
+ self._stop.set()
74
+ self._thread.join()
75
+
76
+ def _run(self) -> None:
77
+ while not self._stop.wait(HEARTBEAT_INTERVAL_SECONDS):
78
+ try:
79
+ self._client._heartbeat(
80
+ task_id=self._task_id,
81
+ run_id=self._run_id,
82
+ queue=self._queue,
83
+ )
84
+ except APIError as error:
85
+ if _retryable_api_error(error):
86
+ logger.warning("Heartbeat Server error; retrying: %s", error.message)
87
+ continue
88
+ if error.code == "run_finalized" and error.details.get("action") == "complete":
89
+ self._control.complete()
90
+ elif error.code in {"run_finalized", "stale_run"}:
91
+ self._control.revoke(str(error.details.get("action", error.code)))
92
+ else:
93
+ self._control.fail(error)
94
+ return
95
+ except TransportError as error:
96
+ logger.warning("Heartbeat transport error; retrying: %s", error)
97
+
98
+
99
+ def loop(
100
+ *,
101
+ route: str = "default",
102
+ queue: str | None = None,
103
+ idle_timeout: float = 300.0,
104
+ force_stop_timeout: float | None = None,
105
+ ) -> Callable[[Callable[P, R]], Callable[P, None]]:
106
+ normalized_route = validate_identifier(route, field="route")
107
+ normalized_idle_timeout = _validate_idle_timeout(idle_timeout)
108
+ normalized_force_stop_timeout = _validate_force_stop_timeout(force_stop_timeout)
109
+
110
+ def decorate(function: Callable[P, R]) -> Callable[P, None]:
111
+ binding = compile_binding(function)
112
+
113
+ @functools.wraps(function)
114
+ def run(*args: P.args, **kwargs: P.kwargs) -> None:
115
+ binding.validate_invocation(
116
+ cast(tuple[object, ...], args),
117
+ cast(dict[str, object], kwargs),
118
+ )
119
+ _run_python_worker(
120
+ binding,
121
+ cast(tuple[object, ...], args),
122
+ cast(dict[str, object], kwargs),
123
+ route=normalized_route,
124
+ queue=queue,
125
+ idle_timeout=normalized_idle_timeout,
126
+ force_stop_timeout=normalized_force_stop_timeout,
127
+ )
128
+
129
+ return run
130
+
131
+ return decorate
132
+
133
+
134
+ def _run_python_worker(
135
+ binding: CompiledBinding,
136
+ startup_args: tuple[object, ...],
137
+ startup_kwargs: dict[str, object],
138
+ *,
139
+ route: str,
140
+ queue: str | None,
141
+ idle_timeout: float,
142
+ force_stop_timeout: float | None,
143
+ ) -> None:
144
+ _guard_worker_topology()
145
+ with Client(queue=queue) as client, WorkerTee() as tee:
146
+ configure_worker_logger()
147
+ queue_name = client.configuration.queue
148
+ _preflight(client, queue_name)
149
+ idle_deadline: float | None = None
150
+ while True:
151
+ claim = client._claim(route=route, run_id=_generate_run_id(), queue=queue_name)
152
+ if claim is None:
153
+ now = time.monotonic()
154
+ if idle_deadline is None:
155
+ idle_deadline = now + idle_timeout
156
+ if now >= idle_deadline:
157
+ logger.info("Worker idle timeout reached; stopping normally.")
158
+ return
159
+ time.sleep(min(POLL_INTERVAL_SECONDS, idle_deadline - now))
160
+ continue
161
+ idle_deadline = None
162
+ logger.info(
163
+ "Claimed Task %s as run %s (attempt %d, route %s).",
164
+ claim.task.id,
165
+ claim.run_id,
166
+ claim.task.attempt,
167
+ route,
168
+ )
169
+ _run_python_claim(
170
+ client,
171
+ tee,
172
+ binding,
173
+ startup_args,
174
+ startup_kwargs,
175
+ claim=claim,
176
+ queue=queue_name,
177
+ route=route,
178
+ force_stop_timeout=force_stop_timeout,
179
+ )
180
+
181
+
182
+ def _run_python_claim(
183
+ client: Client,
184
+ tee: WorkerTee,
185
+ binding: CompiledBinding,
186
+ startup_args: tuple[object, ...],
187
+ startup_kwargs: dict[str, object],
188
+ *,
189
+ claim: ClaimResponse,
190
+ queue: str,
191
+ route: str,
192
+ force_stop_timeout: float | None,
193
+ ) -> None:
194
+ try:
195
+ journal = LocalRunJournal.create(
196
+ claim=claim,
197
+ endpoint=client.configuration.endpoint_dict(),
198
+ queue=queue,
199
+ route=route,
200
+ )
201
+ except Exception:
202
+ try:
203
+ client._unclaim(task_id=claim.task.id, run_id=claim.run_id, queue=queue)
204
+ except Exception:
205
+ logger.exception("Could not return Task after local journal setup failed.")
206
+ raise
207
+ control = RunControl(force_stop_timeout=force_stop_timeout, force_stop=_force_stop_process)
208
+
209
+ def report_complete(result: dict[str, JSONValue]) -> bool:
210
+ accepted = report_complete_until_resolved(
211
+ client,
212
+ queue=queue,
213
+ task_id=claim.task.id,
214
+ run_id=claim.run_id,
215
+ result=result,
216
+ )
217
+ if accepted:
218
+ control.complete()
219
+ else:
220
+ control.revoke("stale_run")
221
+ return accepted
222
+
223
+ info = TaskInfo(
224
+ **claim.task.model_dump(),
225
+ run_id=claim.run_id,
226
+ run_dir=journal.run_dir,
227
+ )
228
+ context = ExecutionContext(
229
+ info=info,
230
+ kind="python",
231
+ journal=journal,
232
+ reporter=report_complete,
233
+ control=control,
234
+ )
235
+ heartbeat = Heartbeat(
236
+ client,
237
+ queue=queue,
238
+ task_id=claim.task.id,
239
+ run_id=claim.run_id,
240
+ control=control,
241
+ )
242
+ activate_context(context)
243
+ heartbeat.start()
244
+ fatal: FatalWorkerError | None = None
245
+ try:
246
+ with tee.capture(journal.log_path):
247
+ try:
248
+ binding.invoke(claim.task.args, startup_args, startup_kwargs)
249
+ except FatalWorkerError as error:
250
+ fatal = error
251
+ logger.critical("Fatal Worker failure for Task %s.", claim.task.id, exc_info=True)
252
+ if control.active and not context.finished:
253
+ _report_failure(client, journal, claim, queue, error)
254
+ except TransientError as error:
255
+ logger.warning("%s: %s", type(error).__name__, error)
256
+ if control.active and not context.finished:
257
+ _report_unclaim(client, journal, claim, queue)
258
+ except Exception as error:
259
+ logger.exception("Task %s failed.", claim.task.id)
260
+ if control.active and not context.finished:
261
+ _report_failure(client, journal, claim, queue, error)
262
+ else:
263
+ if control.active and not context.finished:
264
+ _report_complete(client, journal, claim, queue, {})
265
+ except KeyboardInterrupt:
266
+ if control.active and not context.finished:
267
+ _best_effort_unclaim(client, claim, queue)
268
+ raise
269
+ finally:
270
+ control.executor_done()
271
+ heartbeat.stop()
272
+ deactivate_context(context)
273
+ if fatal is not None:
274
+ raise fatal
275
+ if control.fatal_error is not None:
276
+ raise control.fatal_error
277
+
278
+
279
+ def report_complete_until_resolved(
280
+ client: Client,
281
+ *,
282
+ queue: str,
283
+ task_id: str,
284
+ run_id: str,
285
+ result: dict[str, JSONValue],
286
+ ) -> bool:
287
+ return _report_until_resolved(
288
+ lambda: client._complete(
289
+ task_id=task_id,
290
+ run_id=run_id,
291
+ result=result,
292
+ queue=queue,
293
+ )
294
+ )
295
+
296
+
297
+ def _report_complete(
298
+ client: Client,
299
+ journal: LocalRunJournal,
300
+ claim: ClaimResponse,
301
+ queue: str,
302
+ result: dict[str, JSONValue],
303
+ ) -> None:
304
+ _journal_best_effort(lambda: journal.reporting("complete", result))
305
+ accepted = report_complete_until_resolved(
306
+ client,
307
+ queue=queue,
308
+ task_id=claim.task.id,
309
+ run_id=claim.run_id,
310
+ result=result,
311
+ )
312
+ _finish_journal(journal, accepted)
313
+
314
+
315
+ def _report_unclaim(
316
+ client: Client,
317
+ journal: LocalRunJournal,
318
+ claim: ClaimResponse,
319
+ queue: str,
320
+ ) -> None:
321
+ _journal_best_effort(lambda: journal.reporting("unclaim"))
322
+ accepted = _report_until_resolved(
323
+ lambda: client._unclaim(task_id=claim.task.id, run_id=claim.run_id, queue=queue)
324
+ )
325
+ _finish_journal(journal, accepted)
326
+
327
+
328
+ def _report_failure(
329
+ client: Client,
330
+ journal: LocalRunJournal,
331
+ claim: ClaimResponse,
332
+ queue: str,
333
+ error: Exception,
334
+ ) -> None:
335
+ error_type, message, traceback = _failure_diagnostic(error, claim.run_id)
336
+ payload: dict[str, JSONValue] = {
337
+ "type": error_type,
338
+ "message": message,
339
+ "traceback": traceback,
340
+ }
341
+ _journal_best_effort(lambda: journal.reporting("fail", payload))
342
+ accepted = _report_until_resolved(
343
+ lambda: client._fail(
344
+ task_id=claim.task.id,
345
+ run_id=claim.run_id,
346
+ error_type=error_type,
347
+ message=message,
348
+ traceback=traceback,
349
+ queue=queue,
350
+ )
351
+ )
352
+ _finish_journal(journal, accepted)
353
+
354
+
355
+ def _report_until_resolved(operation: Callable[[], None]) -> bool:
356
+ attempt = 0
357
+ while True:
358
+ try:
359
+ operation()
360
+ return True
361
+ except APIError as error:
362
+ if error.code in {"stale_run", "run_finalized"}:
363
+ return False
364
+ if not _retryable_api_error(error):
365
+ raise
366
+ logger.warning("Terminal report Server error; retrying: %s", error.message)
367
+ except TransportError as error:
368
+ logger.warning("Terminal report transport error; retrying: %s", error)
369
+ delay = TERMINAL_BACKOFF_SECONDS[min(attempt, len(TERMINAL_BACKOFF_SECONDS) - 1)]
370
+ attempt += 1
371
+ time.sleep(delay)
372
+
373
+
374
+ def _failure_diagnostic(error: Exception, run_id: str) -> tuple[str, str, str | None]:
375
+ error_type = _safe_diagnostic_text(type(error).__name__)
376
+ message = _safe_diagnostic_text(str(error))
377
+ formatted = _safe_diagnostic_text(
378
+ "".join(traceback_module.format_exception(type(error), error, error.__traceback__))
379
+ )
380
+ body = {
381
+ "run_id": run_id,
382
+ "error": {"type": error_type, "message": message, "traceback": formatted},
383
+ }
384
+ encoded = json.dumps(
385
+ body,
386
+ ensure_ascii=False,
387
+ allow_nan=False,
388
+ separators=(",", ":"),
389
+ ).encode("utf-8")
390
+ if len(encoded) > MAX_REQUEST_BYTES:
391
+ return (
392
+ error_type,
393
+ "Failure diagnostics exceeded the 1 MiB limit; see local run.log.",
394
+ None,
395
+ )
396
+ return error_type, message, formatted
397
+
398
+
399
+ def _safe_diagnostic_text(value: str) -> str:
400
+ return "".join(
401
+ "\N{REPLACEMENT CHARACTER}" if 0xD800 <= ord(character) <= 0xDFFF else character
402
+ for character in value
403
+ )
404
+
405
+
406
+ def _finish_journal(journal: LocalRunJournal, accepted: bool) -> None:
407
+ _journal_best_effort(journal.acknowledged if accepted else journal.revoked)
408
+
409
+
410
+ def _journal_best_effort(operation: Callable[[], None]) -> None:
411
+ try:
412
+ operation()
413
+ except Exception:
414
+ logger.warning("Could not update local run journal.", exc_info=True)
415
+
416
+
417
+ def _best_effort_unclaim(client: Client, claim: ClaimResponse, queue: str) -> None:
418
+ try:
419
+ client._unclaim(task_id=claim.task.id, run_id=claim.run_id, queue=queue)
420
+ except Exception:
421
+ logger.warning("Could not return interrupted Task; heartbeat recovery will apply.")
422
+
423
+
424
+ def _preflight(client: Client, queue: str) -> None:
425
+ client._health()
426
+ if queue not in {item.name for item in client.list_queues()}:
427
+ raise ConfigError(
428
+ "invalid_config",
429
+ f"Queue {queue!r} does not exist.",
430
+ {"queue": queue},
431
+ )
432
+
433
+
434
+ def _guard_worker_topology() -> None:
435
+ if active_context_present() or os.environ.get("LABTASKER_RUN_ID") is not None:
436
+ raise ConfigError(
437
+ "invalid_config",
438
+ "A nested Worker cannot start inside an active Labtasker execution.",
439
+ {},
440
+ )
441
+ world_size = os.environ.get("WORLD_SIZE")
442
+ rank_present = os.environ.get("RANK") is not None or os.environ.get("LOCAL_RANK") is not None
443
+ try:
444
+ distributed = world_size is not None and int(world_size) > 1
445
+ except ValueError:
446
+ distributed = False
447
+ if distributed and rank_present:
448
+ raise ConfigError(
449
+ "invalid_config",
450
+ "Start labtasker loop outside torchrun or Accelerate, not inside each rank.",
451
+ {"WORLD_SIZE": world_size},
452
+ )
453
+
454
+
455
+ def _validate_idle_timeout(value: float) -> float:
456
+ if isinstance(value, bool) or not isinstance(value, (int, float)):
457
+ raise RequestValidationError("idle_timeout must be a finite non-negative number.")
458
+ normalized = float(value)
459
+ if not math.isfinite(normalized) or normalized < 0:
460
+ raise RequestValidationError("idle_timeout must be a finite non-negative number.")
461
+ return normalized
462
+
463
+
464
+ def _retryable_api_error(error: APIError) -> bool:
465
+ return error.status_code >= 500 or error.code == "database_busy"
466
+
467
+
468
+ def _generate_run_id() -> str:
469
+ return f"r_{secrets.token_urlsafe(9)}"
470
+
471
+
472
+ def _force_stop_process() -> None:
473
+ os._exit(1)
@@ -0,0 +1,13 @@
1
+ Metadata-Version: 2.5
2
+ Name: labtasker-client
3
+ Version: 2.0.0
4
+ Summary: A small task queue for parallel model inference and evaluation
5
+ Project-URL: Homepage, https://github.com/luocfprime/labtasker
6
+ Project-URL: Repository, https://github.com/luocfprime/labtasker.git
7
+ Author-email: lcf <luocfprime@gmail.com>
8
+ License-Expression: Apache-2.0
9
+ License-File: LICENSE
10
+ Requires-Python: >=3.11
11
+ Requires-Dist: httpx<1,>=0.28
12
+ Requires-Dist: pydantic<3,>=2.10
13
+ Requires-Dist: typer<1,>=0.16
@@ -0,0 +1,25 @@
1
+ labtasker/__init__.py,sha256=GzS65tbsjSJW422ehFe_tjB0AfWQNkgvq7ETutK7DYI,1573
2
+ labtasker/__main__.py,sha256=P1DGd62jwyBW7M6Ns6H_mYXcv1mrCX_-AkXQ5y6Sr8Q,37
3
+ labtasker/api.py,sha256=VpxuDVQeNucac0CKtrC3u57z5TRsauD1ewtGr53RzqA,3019
4
+ labtasker/binding.py,sha256=MS46lqlPiADEAIFCHKedF4rDR43weHsVXhInbYRSmno,6331
5
+ labtasker/cli.py,sha256=XhrB8UKwdBb-mDhToBURDZPO5fSB9VCmWu9y6-0WjFc,16296
6
+ labtasker/client.py,sha256=sRzn3Chngq0psTiEINgWXRzA_9v2EsMPgD-nlyMRc10,24766
7
+ labtasker/command_template.py,sha256=ThOr_gkAxgQXx2bGkOyqHi5orQwTx4wmg1s1-261AbQ,5303
8
+ labtasker/command_worker.py,sha256=jwQpfRXrOPhRGrfvqItNukoYVU--W0e_Lu7I9toXLTk,15018
9
+ labtasker/config.py,sha256=g2L7S8BmL67qId8UjM3-0qNEqYUHdn9l2s8peKz5eCQ,6958
10
+ labtasker/errors.py,sha256=S5PTHR7c_Bpqux2Y-dvW8hnqYxzCj3BhpXVuKMHQVxc,1640
11
+ labtasker/execution.py,sha256=LwpNKYWwkGB28g0M9h07mS5cXIYDxjdUJfszWqO0hfo,13289
12
+ labtasker/journal.py,sha256=9UKfvF26zbAybbPr1grAQ_opSkxOX-e10d4SQKk1Mcc,10826
13
+ labtasker/local.py,sha256=JcIisjT-Lt6hB0eSAELWRSHt5AeGolCLoimBEKEgwTE,6052
14
+ labtasker/models.py,sha256=c5GQYzg5lyoSeDSwi0LKN_7MA2CPhvtAXTOmiVK37_I,6125
15
+ labtasker/paths.py,sha256=MG3iLlNy50_arQRA61R11LvQPKiHZWKtBP3yfuCt5-8,1084
16
+ labtasker/py.typed,sha256=AbpHGcgLb-kRsJGnwFEktk7uzpZOCcBY74-YBdrKVGs,1
17
+ labtasker/tee.py,sha256=3HUv0MFsZBYAXqr6z4SfWmE50vbiZIHCGkZCTePvE5I,4260
18
+ labtasker/types.py,sha256=hfhc5AnWs5j_VvVfnz1N4_W9rwHfGOByc_b2Uz4q_Kw,713
19
+ labtasker/validation.py,sha256=t6Z7pAxLURhd9yF7Z7qWDRXAkUOfItecRT6YENXJk5U,8282
20
+ labtasker/worker.py,sha256=Z7JdNV40HNQ4LQm0dbe7yW6rFL5tc2WbIRe5H4rmqA4,14928
21
+ labtasker_client-2.0.0.dist-info/METADATA,sha256=6LRRq8EKioaYuSNZfUblo6qJkiHr8_uIkf9HUcsOnCU,475
22
+ labtasker_client-2.0.0.dist-info/WHEEL,sha256=zOwg4jB6zX2kU910N-cMawjivD6tO8NEWvE12je1bVk,87
23
+ labtasker_client-2.0.0.dist-info/entry_points.txt,sha256=o8P3clM7HqgDpayeELDuF6PEvT0evVYA6zinWiJ5d98,48
24
+ labtasker_client-2.0.0.dist-info/licenses/LICENSE,sha256=xx0jnfkXJvxRnG63LTGOxlggYnIysveWIZ6H3PNdCrQ,11357
25
+ labtasker_client-2.0.0.dist-info/RECORD,,
@@ -0,0 +1,4 @@
1
+ Wheel-Version: 1.0
2
+ Generator: hatchling 1.32.0
3
+ Root-Is-Purelib: true
4
+ Tag: py3-none-any
@@ -0,0 +1,2 @@
1
+ [console_scripts]
2
+ labtasker = labtasker.cli:app