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.
@@ -0,0 +1,173 @@
1
+ """Compile the command Worker's intentionally small interpolation language.
2
+
3
+ Grammar::
4
+
5
+ path = segment ("." segment)*
6
+ segment = start continue*
7
+ start = ASCII letter | "_"
8
+ continue = start | ASCII digit
9
+
10
+ In text, ``%{{`` emits a literal ``%{`` and ``%{path}`` emits a path piece.
11
+ Every scanner transition advances the Unicode input by at least one code point.
12
+ The language is regular and fail-fast; it has no nesting or error recovery.
13
+ """
14
+
15
+ from __future__ import annotations
16
+
17
+ import json
18
+ from dataclasses import dataclass
19
+
20
+ from labtasker.paths import PathError, select_path
21
+ from labtasker.types import JSONValue
22
+
23
+
24
+ class TemplateSyntaxError(ValueError):
25
+ def __init__(self, message: str, *, element: int, column: int) -> None:
26
+ self.element = element
27
+ self.column = column
28
+ super().__init__(f"argv element {element}, column {column}: {message}")
29
+
30
+
31
+ class TemplateBindingError(ValueError):
32
+ pass
33
+
34
+
35
+ @dataclass(frozen=True, slots=True)
36
+ class _PathPiece:
37
+ segments: tuple[str, ...]
38
+
39
+
40
+ @dataclass(frozen=True, slots=True)
41
+ class CompiledTemplate:
42
+ pieces: tuple[str | _PathPiece, ...]
43
+
44
+ def resolve(self, args: dict[str, JSONValue]) -> str:
45
+ rendered: list[str] = []
46
+ for piece in self.pieces:
47
+ if isinstance(piece, str):
48
+ rendered.append(piece)
49
+ continue
50
+ try:
51
+ selected = select_path(args, piece.segments)
52
+ except PathError as error:
53
+ raise TemplateBindingError(str(error)) from error
54
+ rendered.append(_render_value(selected))
55
+ value = "".join(rendered)
56
+ if "\0" in value:
57
+ raise TemplateBindingError("resolved argv element contains NUL")
58
+ return value
59
+
60
+
61
+ def compile_argv(argv: list[str]) -> tuple[CompiledTemplate, ...]:
62
+ if not argv:
63
+ raise ValueError("command must contain at least one argv element")
64
+ return tuple(_compile_element(value, element=index) for index, value in enumerate(argv, 1))
65
+
66
+
67
+ def resolve_argv(
68
+ templates: tuple[CompiledTemplate, ...],
69
+ args: dict[str, JSONValue],
70
+ ) -> list[str]:
71
+ return [template.resolve(args) for template in templates]
72
+
73
+
74
+ def _compile_element(value: str, *, element: int) -> CompiledTemplate:
75
+ pieces: list[str | _PathPiece] = []
76
+ literal: list[str] = []
77
+ index = 0
78
+ while index < len(value):
79
+ previous = index
80
+ if value.startswith("%{{", index):
81
+ literal.append("%{")
82
+ index += 3
83
+ elif value.startswith("%{", index):
84
+ if literal:
85
+ pieces.append("".join(literal))
86
+ literal.clear()
87
+ opening = index
88
+ index += 2
89
+ segments, index = _scan_path(value, index, element=element, opening=opening)
90
+ pieces.append(_PathPiece(segments))
91
+ else:
92
+ if value[index] == "\0":
93
+ raise TemplateSyntaxError(
94
+ "literal argv text contains NUL",
95
+ element=element,
96
+ column=index + 1,
97
+ )
98
+ literal.append(value[index])
99
+ index += 1
100
+ if index <= previous: # pragma: no cover - implementation invariant
101
+ raise AssertionError("command template scanner did not advance")
102
+ if literal or not pieces:
103
+ pieces.append("".join(literal))
104
+ return CompiledTemplate(tuple(pieces))
105
+
106
+
107
+ def _scan_path(
108
+ value: str,
109
+ index: int,
110
+ *,
111
+ element: int,
112
+ opening: int,
113
+ ) -> tuple[tuple[str, ...], int]:
114
+ segments: list[str] = []
115
+ current: list[str] = []
116
+ expecting_start = True
117
+ while index < len(value):
118
+ character = value[index]
119
+ if character == "}":
120
+ if expecting_start:
121
+ raise TemplateSyntaxError(
122
+ "placeholder has an empty path segment",
123
+ element=element,
124
+ column=index + 1,
125
+ )
126
+ segments.append("".join(current))
127
+ return tuple(segments), index + 1
128
+ if expecting_start:
129
+ if _is_start(character):
130
+ current.append(character)
131
+ expecting_start = False
132
+ index += 1
133
+ continue
134
+ elif character == ".":
135
+ segments.append("".join(current))
136
+ current.clear()
137
+ expecting_start = True
138
+ index += 1
139
+ continue
140
+ elif _is_continue(character):
141
+ current.append(character)
142
+ index += 1
143
+ continue
144
+ raise TemplateSyntaxError(
145
+ f"invalid placeholder character {character!r}",
146
+ element=element,
147
+ column=index + 1,
148
+ )
149
+ raise TemplateSyntaxError(
150
+ "unterminated placeholder",
151
+ element=element,
152
+ column=opening + 1,
153
+ )
154
+
155
+
156
+ def _is_start(character: str) -> bool:
157
+ return character == "_" or "A" <= character <= "Z" or "a" <= character <= "z"
158
+
159
+
160
+ def _is_continue(character: str) -> bool:
161
+ return _is_start(character) or "0" <= character <= "9"
162
+
163
+
164
+ def _render_value(value: object) -> str:
165
+ if isinstance(value, str):
166
+ return value
167
+ return json.dumps(
168
+ value,
169
+ ensure_ascii=False,
170
+ allow_nan=False,
171
+ sort_keys=True,
172
+ separators=(",", ":"),
173
+ )
@@ -0,0 +1,492 @@
1
+ from __future__ import annotations
2
+
3
+ import errno
4
+ import logging
5
+ import os
6
+ import select
7
+ import signal
8
+ import subprocess
9
+ import sys
10
+ import threading
11
+ import time
12
+ from collections.abc import Iterator
13
+ from contextlib import contextmanager, suppress
14
+ from pathlib import Path
15
+ from typing import IO, Any
16
+
17
+ from labtasker.client import Client
18
+ from labtasker.command_template import (
19
+ CompiledTemplate,
20
+ TemplateBindingError,
21
+ compile_argv,
22
+ resolve_argv,
23
+ )
24
+ from labtasker.execution import RunControl, _validate_force_stop_timeout
25
+ from labtasker.journal import LocalRunJournal
26
+ from labtasker.models import ClaimResponse
27
+ from labtasker.tee import configure_worker_logger
28
+ from labtasker.types import JSONValue
29
+ from labtasker.validation import validate_identifier
30
+ from labtasker.worker import (
31
+ POLL_INTERVAL_SECONDS,
32
+ Heartbeat,
33
+ _best_effort_unclaim,
34
+ _finish_journal,
35
+ _generate_run_id,
36
+ _guard_worker_topology,
37
+ _journal_best_effort,
38
+ _preflight,
39
+ _report_until_resolved,
40
+ _safe_diagnostic_text,
41
+ _validate_idle_timeout,
42
+ )
43
+
44
+ logger = logging.getLogger("labtasker.command_worker")
45
+ _PLATFORM = sys.platform
46
+ _POSIX_PROCESS_GROUPS = os.name == "posix" and hasattr(os, "killpg")
47
+
48
+
49
+ def run_command_worker(
50
+ argv: list[str],
51
+ *,
52
+ route: str = "default",
53
+ queue: str | None = None,
54
+ idle_timeout: float = 300.0,
55
+ force_stop_timeout: float | None = None,
56
+ ) -> None:
57
+ templates = compile_argv(argv)
58
+ normalized_route = validate_identifier(route, field="route")
59
+ normalized_idle_timeout = _validate_idle_timeout(idle_timeout)
60
+ normalized_force_stop_timeout = _validate_force_stop_timeout(force_stop_timeout)
61
+ _guard_command_worker_platform()
62
+ _guard_worker_topology()
63
+ configure_worker_logger()
64
+ with Client(queue=queue) as client:
65
+ queue_name = client.configuration.queue
66
+ _preflight(client, queue_name)
67
+ idle_deadline: float | None = None
68
+ while True:
69
+ claim = client._claim(
70
+ route=normalized_route,
71
+ run_id=_generate_run_id(),
72
+ queue=queue_name,
73
+ )
74
+ if claim is None:
75
+ now = time.monotonic()
76
+ if idle_deadline is None:
77
+ idle_deadline = now + normalized_idle_timeout
78
+ if now >= idle_deadline:
79
+ logger.info("Worker idle timeout reached; stopping normally.")
80
+ return
81
+ time.sleep(min(POLL_INTERVAL_SECONDS, idle_deadline - now))
82
+ continue
83
+ idle_deadline = None
84
+ logger.info(
85
+ "Claimed Task %s as run %s (attempt %d, route %s).",
86
+ claim.task.id,
87
+ claim.run_id,
88
+ claim.task.attempt,
89
+ normalized_route,
90
+ )
91
+ _run_command_claim(
92
+ client,
93
+ templates,
94
+ claim=claim,
95
+ queue=queue_name,
96
+ route=normalized_route,
97
+ force_stop_timeout=normalized_force_stop_timeout,
98
+ )
99
+
100
+
101
+ def _guard_command_worker_platform() -> None:
102
+ if not _POSIX_PROCESS_GROUPS:
103
+ raise NotImplementedError(
104
+ "Command Workers require POSIX process-group support; "
105
+ f"platform {_PLATFORM!r} is not supported."
106
+ )
107
+
108
+
109
+ def _run_command_claim(
110
+ client: Client,
111
+ templates: tuple[CompiledTemplate, ...],
112
+ *,
113
+ claim: ClaimResponse,
114
+ queue: str,
115
+ route: str,
116
+ force_stop_timeout: float | None,
117
+ ) -> None:
118
+ try:
119
+ journal = LocalRunJournal.create(
120
+ claim=claim,
121
+ endpoint=client.configuration.endpoint_dict(),
122
+ queue=queue,
123
+ route=route,
124
+ )
125
+ except Exception:
126
+ _best_effort_unclaim(client, claim, queue)
127
+ raise
128
+
129
+ control = RunControl(force_stop_timeout=None, force_stop=lambda: None)
130
+ heartbeat = Heartbeat(
131
+ client,
132
+ queue=queue,
133
+ task_id=claim.task.id,
134
+ run_id=claim.run_id,
135
+ control=control,
136
+ )
137
+ process: subprocess.Popen[bytes] | None = None
138
+ heartbeat.start()
139
+ try:
140
+ try:
141
+ resolved = resolve_argv(templates, claim.task.args)
142
+ except TemplateBindingError as error:
143
+ _report_command_failure(client, journal, claim, queue, "TaskBindingError", str(error))
144
+ return
145
+ environment = _command_environment(client, claim, journal, queue, route)
146
+ try:
147
+ if _interactive_terminal():
148
+ process = _run_pty(
149
+ resolved,
150
+ environment,
151
+ journal.log_path,
152
+ control,
153
+ force_stop_timeout,
154
+ )
155
+ else:
156
+ process = _run_pipes(
157
+ resolved,
158
+ environment,
159
+ journal.log_path,
160
+ control,
161
+ force_stop_timeout,
162
+ )
163
+ except OSError as error:
164
+ _report_command_failure(
165
+ client,
166
+ journal,
167
+ claim,
168
+ queue,
169
+ type(error).__name__,
170
+ str(error),
171
+ )
172
+ return
173
+ if control.fatal_error is not None:
174
+ raise control.fatal_error
175
+ if control.revoked:
176
+ _journal_best_effort(journal.revoked)
177
+ return
178
+ try:
179
+ journal = LocalRunJournal.open(journal.run_dir)
180
+ except Exception:
181
+ logger.warning("Could not reload command child journal.", exc_info=True)
182
+ if journal.phase == "acknowledged" and journal.terminal_action == "complete":
183
+ return
184
+ if journal.phase == "reporting" and journal.terminal_action == "complete":
185
+ result = journal.read_result()
186
+ accepted = _report_command_complete(client, claim, queue, result)
187
+ _finish_journal(journal, accepted)
188
+ return
189
+ if control.completed:
190
+ return
191
+ if process.returncode == 0:
192
+ _journal_best_effort(lambda: journal.reporting("complete", {}))
193
+ accepted = _report_command_complete(client, claim, queue, {})
194
+ _finish_journal(journal, accepted)
195
+ return
196
+ message = _returncode_message(process.returncode)
197
+ _report_command_failure(client, journal, claim, queue, "CommandProcessError", message)
198
+ except KeyboardInterrupt:
199
+ if control.active:
200
+ _best_effort_unclaim(client, claim, queue)
201
+ raise
202
+ finally:
203
+ control.executor_done()
204
+ heartbeat.stop()
205
+
206
+
207
+ def _run_pipes(
208
+ argv: list[str],
209
+ environment: dict[str, str],
210
+ log_path: Path,
211
+ control: RunControl,
212
+ force_stop_timeout: float | None,
213
+ ) -> subprocess.Popen[bytes]:
214
+ process = subprocess.Popen(
215
+ argv,
216
+ stdin=subprocess.DEVNULL,
217
+ stdout=subprocess.PIPE,
218
+ stderr=subprocess.PIPE,
219
+ env=environment,
220
+ close_fds=True,
221
+ start_new_session=True,
222
+ )
223
+ assert process.stdout is not None
224
+ assert process.stderr is not None
225
+ lock = threading.Lock()
226
+ try:
227
+ with log_path.open("ab", buffering=0) as log:
228
+ stdout_thread = _start_drain(process.stdout, sys.stdout, log, lock, "stdout")
229
+ stderr_thread = _start_drain(process.stderr, sys.stderr, log, lock, "stderr")
230
+ _wait_process(process, control, force_stop_timeout)
231
+ stdout_thread.join()
232
+ stderr_thread.join()
233
+ except BaseException:
234
+ _terminate_process_group(process, force_stop_timeout)
235
+ raise
236
+ return process
237
+
238
+
239
+ def _start_drain(
240
+ source: IO[bytes],
241
+ destination: object,
242
+ log: IO[bytes],
243
+ lock: threading.Lock,
244
+ name: str,
245
+ ) -> threading.Thread:
246
+ def drain() -> None:
247
+ while True:
248
+ chunk = source.read(65536)
249
+ if not chunk:
250
+ return
251
+ with lock:
252
+ log.write(chunk)
253
+ _write_bytes(destination, chunk)
254
+
255
+ thread = threading.Thread(target=drain, name=f"labtasker-command-{name}", daemon=True)
256
+ thread.start()
257
+ return thread
258
+
259
+
260
+ def _run_pty(
261
+ argv: list[str],
262
+ environment: dict[str, str],
263
+ log_path: Path,
264
+ control: RunControl,
265
+ force_stop_timeout: float | None,
266
+ ) -> subprocess.Popen[bytes]:
267
+ import pty
268
+ import termios
269
+
270
+ master, slave = pty.openpty()
271
+ _copy_terminal_size(sys.stdin.fileno(), slave)
272
+ process = subprocess.Popen(
273
+ argv,
274
+ stdin=slave,
275
+ stdout=slave,
276
+ stderr=slave,
277
+ env=environment,
278
+ close_fds=True,
279
+ start_new_session=True,
280
+ )
281
+ os.close(slave)
282
+ try:
283
+ with log_path.open("ab", buffering=0) as log, _raw_terminal(sys.stdin.fileno()):
284
+ last_size: bytes | None = None
285
+ output_open = True
286
+ while output_open or process.poll() is None:
287
+ if control.revoked and process.poll() is None:
288
+ _terminate_process_group(process, force_stop_timeout)
289
+ size = _terminal_size(sys.stdin.fileno())
290
+ if size is not None and size != last_size:
291
+ try:
292
+ import fcntl
293
+
294
+ fcntl.ioctl(master, termios.TIOCSWINSZ, size)
295
+ except OSError:
296
+ pass
297
+ last_size = size
298
+ readers = [master]
299
+ if process.poll() is None:
300
+ readers.append(sys.stdin.fileno())
301
+ ready, _, _ = select.select(readers, [], [], 0.1)
302
+ if master in ready:
303
+ try:
304
+ chunk = os.read(master, 65536)
305
+ except OSError as error:
306
+ if error.errno != errno.EIO:
307
+ raise
308
+ chunk = b""
309
+ if chunk:
310
+ log.write(chunk)
311
+ _write_bytes(sys.stdout, chunk)
312
+ else:
313
+ output_open = False
314
+ if sys.stdin.fileno() in ready:
315
+ chunk = os.read(sys.stdin.fileno(), 65536)
316
+ if chunk:
317
+ os.write(master, chunk)
318
+ process.wait()
319
+ except BaseException:
320
+ _terminate_process_group(process, force_stop_timeout)
321
+ raise
322
+ finally:
323
+ os.close(master)
324
+ return process
325
+
326
+
327
+ def _wait_process(
328
+ process: subprocess.Popen[bytes],
329
+ control: RunControl,
330
+ force_stop_timeout: float | None,
331
+ ) -> None:
332
+ while process.poll() is None:
333
+ if control.revoked:
334
+ _terminate_process_group(process, force_stop_timeout)
335
+ return
336
+ with suppress(subprocess.TimeoutExpired):
337
+ process.wait(timeout=0.1)
338
+
339
+
340
+ def _terminate_process_group(
341
+ process: subprocess.Popen[bytes],
342
+ force_stop_timeout: float | None,
343
+ ) -> None:
344
+ if process.poll() is not None:
345
+ return
346
+ os.killpg(process.pid, signal.SIGTERM)
347
+ if force_stop_timeout is None:
348
+ process.wait()
349
+ return
350
+ try:
351
+ process.wait(timeout=force_stop_timeout)
352
+ except subprocess.TimeoutExpired:
353
+ os.killpg(process.pid, signal.SIGKILL)
354
+ process.wait()
355
+
356
+
357
+ def _command_environment(
358
+ client: Client,
359
+ claim: ClaimResponse,
360
+ journal: LocalRunJournal,
361
+ queue: str,
362
+ route: str,
363
+ ) -> dict[str, str]:
364
+ environment = dict(os.environ)
365
+ environment.update(
366
+ {
367
+ "LABTASKER_QUEUE": queue,
368
+ "LABTASKER_TASK_ID": claim.task.id,
369
+ "LABTASKER_RUN_ID": claim.run_id,
370
+ "LABTASKER_ROUTE": route,
371
+ "LABTASKER_RUN_DIR": str(journal.run_dir),
372
+ }
373
+ )
374
+ configuration = client.configuration
375
+ if configuration.local is None:
376
+ assert configuration.url is not None
377
+ environment["LABTASKER_URL"] = configuration.url
378
+ environment.pop("LABTASKER_SOCKET", None)
379
+ environment.pop("LABTASKER_LOCAL_DIRECTORY", None)
380
+ else:
381
+ environment["LABTASKER_SOCKET"] = str(configuration.local.socket)
382
+ environment["LABTASKER_LOCAL_DIRECTORY"] = str(configuration.local.directory)
383
+ environment.pop("LABTASKER_URL", None)
384
+ token = configuration.token
385
+ if token is None or configuration.local is not None:
386
+ environment.pop("LABTASKER_TOKEN", None)
387
+ else:
388
+ environment["LABTASKER_TOKEN"] = token
389
+ return environment
390
+
391
+
392
+ def _report_command_complete(
393
+ client: Client,
394
+ claim: ClaimResponse,
395
+ queue: str,
396
+ result: dict[str, JSONValue],
397
+ ) -> bool:
398
+ return _report_until_resolved(
399
+ lambda: client._complete(
400
+ task_id=claim.task.id,
401
+ run_id=claim.run_id,
402
+ result=result,
403
+ queue=queue,
404
+ )
405
+ )
406
+
407
+
408
+ def _report_command_failure(
409
+ client: Client,
410
+ journal: LocalRunJournal,
411
+ claim: ClaimResponse,
412
+ queue: str,
413
+ error_type: str,
414
+ message: str,
415
+ ) -> None:
416
+ error_type = _safe_diagnostic_text(error_type)
417
+ message = _safe_diagnostic_text(message)
418
+ logger.error("%s: %s", error_type, message)
419
+ payload: dict[str, JSONValue] = {
420
+ "type": error_type,
421
+ "message": message,
422
+ "traceback": None,
423
+ }
424
+ _journal_best_effort(lambda: journal.reporting("fail", payload))
425
+ accepted = _report_until_resolved(
426
+ lambda: client._fail(
427
+ task_id=claim.task.id,
428
+ run_id=claim.run_id,
429
+ error_type=error_type,
430
+ message=message,
431
+ traceback=None,
432
+ queue=queue,
433
+ )
434
+ )
435
+ _finish_journal(journal, accepted)
436
+
437
+
438
+ def _returncode_message(returncode: int) -> str:
439
+ if returncode < 0:
440
+ return f"Command terminated by signal {-returncode}."
441
+ return f"Command exited with status {returncode}."
442
+
443
+
444
+ def _write_bytes(destination: Any, value: bytes) -> None:
445
+ buffer = getattr(destination, "buffer", None)
446
+ if buffer is not None:
447
+ buffer.write(value)
448
+ buffer.flush()
449
+ return
450
+ destination.write(value.decode("utf-8", errors="backslashreplace"))
451
+ destination.flush()
452
+
453
+
454
+ def _interactive_terminal() -> bool:
455
+ return sys.stdin.isatty() and sys.stdout.isatty() and sys.stderr.isatty()
456
+
457
+
458
+ def _terminal_size(descriptor: int) -> bytes | None:
459
+ try:
460
+ import fcntl
461
+ import termios
462
+
463
+ return fcntl.ioctl(descriptor, termios.TIOCGWINSZ, b"\0" * 8)
464
+ except OSError:
465
+ return None
466
+
467
+
468
+ def _copy_terminal_size(source: int, destination: int) -> None:
469
+ import termios
470
+
471
+ size = _terminal_size(source)
472
+ if size is None:
473
+ return
474
+ try:
475
+ import fcntl
476
+
477
+ fcntl.ioctl(destination, termios.TIOCSWINSZ, size)
478
+ except OSError:
479
+ pass
480
+
481
+
482
+ @contextmanager
483
+ def _raw_terminal(descriptor: int) -> Iterator[None]:
484
+ import termios
485
+ import tty
486
+
487
+ attributes = termios.tcgetattr(descriptor)
488
+ tty.setraw(descriptor)
489
+ try:
490
+ yield
491
+ finally:
492
+ termios.tcsetattr(descriptor, termios.TCSADRAIN, attributes)