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/models.py ADDED
@@ -0,0 +1,218 @@
1
+ from __future__ import annotations
2
+
3
+ from datetime import UTC, datetime
4
+ from pathlib import Path
5
+ from typing import Literal
6
+
7
+ from pydantic import BaseModel, ConfigDict, field_validator
8
+
9
+ from labtasker.types import JSONValue, TaskStatus
10
+ from labtasker.validation import (
11
+ validate_identifier,
12
+ validate_int64,
13
+ validate_json_object,
14
+ validate_routes,
15
+ validate_run_id,
16
+ validate_task_id,
17
+ validate_task_name,
18
+ validate_unicode_scalar,
19
+ )
20
+
21
+
22
+ class ResponseModel(BaseModel):
23
+ model_config = ConfigDict(extra="ignore", frozen=True, strict=True)
24
+
25
+
26
+ class LastError(ResponseModel):
27
+ type: str
28
+ message: str
29
+ traceback: str | None
30
+ occurred_at: datetime
31
+ attempt: int
32
+ run_id: str
33
+
34
+ @field_validator("type", "message", "traceback")
35
+ @classmethod
36
+ def validate_strings(cls, value: str | None, info: object) -> str | None:
37
+ if value is None:
38
+ return None
39
+ return validate_unicode_scalar(value, field=getattr(info, "field_name", "last_error"))
40
+
41
+ @field_validator("occurred_at")
42
+ @classmethod
43
+ def validate_time(cls, value: datetime) -> datetime:
44
+ return _utc_datetime(value)
45
+
46
+ @field_validator("attempt")
47
+ @classmethod
48
+ def validate_attempt(cls, value: int) -> int:
49
+ return validate_int64(value, field="attempt")
50
+
51
+ @field_validator("run_id")
52
+ @classmethod
53
+ def validate_run(cls, value: str) -> str:
54
+ return validate_run_id(value)
55
+
56
+
57
+ class Task(ResponseModel):
58
+ id: str
59
+ queue: str
60
+ status: TaskStatus
61
+ name: str | None
62
+ args: dict[str, JSONValue]
63
+ metadata: dict[str, JSONValue]
64
+ priority: int
65
+ attempt: int
66
+ max_attempts: int
67
+ routes: list[str]
68
+ result: dict[str, JSONValue]
69
+ last_error: LastError | None
70
+ last_route: str | None
71
+ created_at: datetime
72
+ updated_at: datetime
73
+ started_at: datetime | None
74
+ finished_at: datetime | None
75
+
76
+ @field_validator("id")
77
+ @classmethod
78
+ def validate_id(cls, value: str) -> str:
79
+ return validate_task_id(value)
80
+
81
+ @field_validator("queue")
82
+ @classmethod
83
+ def validate_queue(cls, value: str) -> str:
84
+ return validate_identifier(value, field="queue")
85
+
86
+ @field_validator("name")
87
+ @classmethod
88
+ def validate_name(cls, value: str | None) -> str | None:
89
+ return validate_task_name(value)
90
+
91
+ @field_validator("args", "metadata", "result")
92
+ @classmethod
93
+ def validate_objects(cls, value: dict[str, JSONValue], info: object) -> dict[str, JSONValue]:
94
+ return validate_json_object(value, field=getattr(info, "field_name", "task"))
95
+
96
+ @field_validator("priority", "attempt")
97
+ @classmethod
98
+ def validate_numbers(cls, value: int, info: object) -> int:
99
+ return validate_int64(value, field=getattr(info, "field_name", "task"))
100
+
101
+ @field_validator("max_attempts")
102
+ @classmethod
103
+ def validate_max_attempts(cls, value: int) -> int:
104
+ return validate_int64(value, field="max_attempts", positive=True)
105
+
106
+ @field_validator("routes")
107
+ @classmethod
108
+ def validate_task_routes(cls, value: list[str]) -> list[str]:
109
+ routes = validate_routes(value)
110
+ if value != routes:
111
+ raise ValueError("routes must be sorted lexicographically")
112
+ return routes
113
+
114
+ @field_validator("last_route")
115
+ @classmethod
116
+ def validate_last_route(cls, value: str | None) -> str | None:
117
+ return None if value is None else validate_identifier(value, field="last_route")
118
+
119
+ @field_validator("created_at", "updated_at", "started_at", "finished_at")
120
+ @classmethod
121
+ def validate_times(cls, value: datetime | None) -> datetime | None:
122
+ return None if value is None else _utc_datetime(value)
123
+
124
+
125
+ class TaskInfo(Task):
126
+ run_id: str
127
+ run_dir: Path
128
+
129
+ @field_validator("run_id")
130
+ @classmethod
131
+ def validate_run(cls, value: str) -> str:
132
+ return validate_run_id(value)
133
+
134
+ @field_validator("run_dir")
135
+ @classmethod
136
+ def validate_run_dir(cls, value: Path) -> Path:
137
+ if not value.is_absolute():
138
+ raise ValueError("run_dir must be absolute")
139
+ return value
140
+
141
+
142
+ class TaskPage(ResponseModel):
143
+ items: list[Task]
144
+ next_cursor: str | None
145
+
146
+
147
+ class ClaimResponse(ResponseModel):
148
+ task: Task
149
+ run_id: str
150
+ lease_expires_at: datetime
151
+
152
+ @field_validator("run_id")
153
+ @classmethod
154
+ def validate_run(cls, value: str) -> str:
155
+ return validate_run_id(value)
156
+
157
+ @field_validator("lease_expires_at")
158
+ @classmethod
159
+ def validate_lease_expires_at(cls, value: datetime) -> datetime:
160
+ return _utc_datetime(value)
161
+
162
+
163
+ class HeartbeatResponse(ResponseModel):
164
+ lease_expires_at: datetime
165
+
166
+ @field_validator("lease_expires_at")
167
+ @classmethod
168
+ def validate_lease_expires_at(cls, value: datetime) -> datetime:
169
+ return _utc_datetime(value)
170
+
171
+
172
+ class HealthResponse(ResponseModel):
173
+ status: Literal["ok"]
174
+ api_version: Literal["2"]
175
+ database: Literal["ok"]
176
+
177
+
178
+ class Queue(ResponseModel):
179
+ name: str
180
+
181
+ @field_validator("name")
182
+ @classmethod
183
+ def validate_name(cls, value: str) -> str:
184
+ return validate_identifier(value, field="queue")
185
+
186
+
187
+ class BulkUpdateResult(ResponseModel):
188
+ matched: int
189
+ updated: int
190
+
191
+ @field_validator("matched", "updated")
192
+ @classmethod
193
+ def validate_count(cls, value: int, info: object) -> int:
194
+ value = validate_int64(value, field=getattr(info, "field_name", "count"))
195
+ if value < 0:
196
+ raise ValueError("count must be non-negative")
197
+ return value
198
+
199
+
200
+ class CountResponse(ResponseModel):
201
+ count: int
202
+
203
+ @field_validator("count")
204
+ @classmethod
205
+ def validate_count(cls, value: int) -> int:
206
+ value = validate_int64(value, field="count")
207
+ if value < 0:
208
+ raise ValueError("count must be non-negative")
209
+ return value
210
+
211
+
212
+ def _utc_datetime(value: datetime) -> datetime:
213
+ offset = value.utcoffset()
214
+ if value.tzinfo is None or offset is None:
215
+ raise ValueError("timestamp must be timezone-aware")
216
+ if offset.total_seconds() != 0:
217
+ raise ValueError("timestamp must use UTC")
218
+ return value.astimezone(UTC)
labtasker/paths.py ADDED
@@ -0,0 +1,34 @@
1
+ """Shared object-only dot paths used by Python and command Workers."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import re
6
+ from typing import Any
7
+
8
+ _PATH_RE = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*(?:\.[A-Za-z_][A-Za-z0-9_]*)*$")
9
+
10
+
11
+ class PathError(ValueError):
12
+ pass
13
+
14
+
15
+ def parse_path(value: str) -> tuple[str, ...]:
16
+ if not isinstance(value, str) or not _PATH_RE.fullmatch(value):
17
+ raise PathError(
18
+ "path must contain dot-separated ASCII identifiers matching [A-Za-z_][A-Za-z0-9_]*"
19
+ )
20
+ return tuple(value.split("."))
21
+
22
+
23
+ def select_path(value: object, path: tuple[str, ...]) -> Any:
24
+ current = value
25
+ traversed: list[str] = []
26
+ for segment in path:
27
+ traversed.append(segment)
28
+ if not isinstance(current, dict):
29
+ parent = ".".join(traversed[:-1]) or "<root>"
30
+ raise PathError(f"path {'.'.join(path)!r} cannot traverse non-object {parent!r}")
31
+ if segment not in current:
32
+ raise PathError(f"path {'.'.join(path)!r} is missing key {segment!r}")
33
+ current = current[segment]
34
+ return current
labtasker/py.typed ADDED
@@ -0,0 +1 @@
1
+
labtasker/tee.py ADDED
@@ -0,0 +1,128 @@
1
+ from __future__ import annotations
2
+
3
+ import logging
4
+ import os
5
+ import sys
6
+ import threading
7
+ import time
8
+ from collections.abc import Callable, Iterator
9
+ from contextlib import contextmanager
10
+ from pathlib import Path
11
+ from time import struct_time
12
+ from typing import TextIO, cast
13
+
14
+ _ACTIVE_TEE: WorkerTee | None = None
15
+ _FORK_HOOK_INSTALLED = False
16
+
17
+
18
+ class _UTCFormatter(logging.Formatter):
19
+ converter: Callable[[float | None], struct_time] = time.gmtime
20
+
21
+
22
+ class _TeeStream:
23
+ def __init__(self, original: TextIO, lock: threading.RLock) -> None:
24
+ self._original = original
25
+ self._lock = lock
26
+ self._destination: TextIO | None = None
27
+
28
+ def write(self, value: str) -> int:
29
+ with self._lock:
30
+ written = self._original.write(value)
31
+ if self._destination is not None:
32
+ self._destination.write(value)
33
+ return written
34
+
35
+ def flush(self) -> None:
36
+ with self._lock:
37
+ self._original.flush()
38
+ if self._destination is not None:
39
+ self._destination.flush()
40
+
41
+ def set_destination(self, destination: TextIO | None) -> None:
42
+ with self._lock:
43
+ self._destination = destination
44
+
45
+ def __getattr__(self, name: str) -> object:
46
+ return getattr(self._original, name)
47
+
48
+
49
+ class WorkerTee:
50
+ def __init__(self) -> None:
51
+ self._lock = threading.RLock()
52
+ self._stdout: _TeeStream | None = None
53
+ self._stderr: _TeeStream | None = None
54
+ self._original_stdout: TextIO | None = None
55
+ self._original_stderr: TextIO | None = None
56
+ self._destination: TextIO | None = None
57
+
58
+ def __enter__(self) -> WorkerTee:
59
+ global _ACTIVE_TEE, _FORK_HOOK_INSTALLED
60
+ self._original_stdout = sys.stdout
61
+ self._original_stderr = sys.stderr
62
+ self._stdout = _TeeStream(sys.stdout, self._lock)
63
+ self._stderr = _TeeStream(sys.stderr, self._lock)
64
+ sys.stdout = cast(TextIO, self._stdout)
65
+ sys.stderr = cast(TextIO, self._stderr)
66
+ if not _FORK_HOOK_INSTALLED and hasattr(os, "register_at_fork"):
67
+ os.register_at_fork(after_in_child=_clear_tee_after_fork)
68
+ _FORK_HOOK_INSTALLED = True
69
+ _ACTIVE_TEE = self
70
+ return self
71
+
72
+ def __exit__(self, *_: object) -> None:
73
+ global _ACTIVE_TEE
74
+ self.clear_destination()
75
+ if self._original_stdout is not None:
76
+ sys.stdout = self._original_stdout
77
+ if self._original_stderr is not None:
78
+ sys.stderr = self._original_stderr
79
+ if _ACTIVE_TEE is self:
80
+ _ACTIVE_TEE = None
81
+
82
+ @contextmanager
83
+ def capture(self, path: Path) -> Iterator[None]:
84
+ if self._destination is not None:
85
+ raise RuntimeError("A Worker log destination is already active.")
86
+ with path.open("a", encoding="utf-8", errors="backslashreplace") as destination:
87
+ self._destination = destination
88
+ if self._stdout is not None:
89
+ self._stdout.set_destination(destination)
90
+ if self._stderr is not None:
91
+ self._stderr.set_destination(destination)
92
+ try:
93
+ yield
94
+ finally:
95
+ self.clear_destination()
96
+
97
+ def clear_destination(self) -> None:
98
+ if self._stdout is not None:
99
+ self._stdout.set_destination(None)
100
+ if self._stderr is not None:
101
+ self._stderr.set_destination(None)
102
+ if self._destination is not None:
103
+ self._destination.flush()
104
+ self._destination = None
105
+
106
+
107
+ def configure_worker_logger() -> logging.Logger:
108
+ logger = logging.getLogger("labtasker")
109
+ if not logger.hasHandlers():
110
+ handler = logging.StreamHandler()
111
+ handler.setLevel(logging.INFO)
112
+ handler.setFormatter(_worker_log_formatter())
113
+ logger.addHandler(handler)
114
+ if logger.level == logging.NOTSET:
115
+ logger.setLevel(logging.INFO)
116
+ return logger
117
+
118
+
119
+ def _worker_log_formatter() -> logging.Formatter:
120
+ return _UTCFormatter(
121
+ "%(asctime)s.%(msecs)03dZ %(levelname)s [labtasker] %(message)s",
122
+ datefmt="%Y-%m-%dT%H:%M:%S",
123
+ )
124
+
125
+
126
+ def _clear_tee_after_fork() -> None:
127
+ if _ACTIVE_TEE is not None:
128
+ _ACTIVE_TEE.clear_destination()
labtasker/types.py ADDED
@@ -0,0 +1,31 @@
1
+ from __future__ import annotations
2
+
3
+ from typing import Literal, TypeAlias, TypedDict
4
+
5
+ from pydantic import JsonValue as PydanticJSONValue
6
+
7
+ JSONValue: TypeAlias = PydanticJSONValue
8
+ TaskStatus: TypeAlias = Literal["pending", "running", "succeeded", "failed", "cancelled"]
9
+ TaskOrderField: TypeAlias = Literal[
10
+ "id",
11
+ "name",
12
+ "status",
13
+ "priority",
14
+ "attempt",
15
+ "max_attempts",
16
+ "last_route",
17
+ "created_at",
18
+ "updated_at",
19
+ "started_at",
20
+ "finished_at",
21
+ ]
22
+
23
+
24
+ class TaskUpdate(TypedDict, total=False):
25
+ name: str | None
26
+ args: dict[str, JSONValue]
27
+ metadata: dict[str, JSONValue]
28
+ priority: int
29
+ max_attempts: int
30
+ routes: list[str]
31
+ result: dict[str, JSONValue]
@@ -0,0 +1,207 @@
1
+ from __future__ import annotations
2
+
3
+ import math
4
+ import re
5
+ import unicodedata
6
+ from collections.abc import Sequence
7
+ from typing import cast
8
+
9
+ from labtasker.errors import ConfigError
10
+ from labtasker.types import JSONValue, TaskOrderField, TaskStatus, TaskUpdate
11
+
12
+ INT64_MIN = -(2**63)
13
+ INT64_MAX = 2**63 - 1
14
+ MAX_JSON_DEPTH = 64
15
+ MAX_FILTER_BYTES = 8192
16
+ IDENTIFIER_RE = re.compile(r"^[A-Za-z0-9][A-Za-z0-9._-]{0,127}$")
17
+ TASK_ID_RE = re.compile(r"^t_[A-Za-z0-9_-]{12}$")
18
+ RUN_ID_RE = re.compile(r"^r_[A-Za-z0-9_-]{12}$")
19
+ TASK_STATUSES = {"pending", "running", "succeeded", "failed", "cancelled"}
20
+ TASK_ORDER_FIELDS = {
21
+ "id",
22
+ "name",
23
+ "status",
24
+ "priority",
25
+ "attempt",
26
+ "max_attempts",
27
+ "last_route",
28
+ "created_at",
29
+ "updated_at",
30
+ "started_at",
31
+ "finished_at",
32
+ }
33
+
34
+
35
+ class RequestValidationError(ValueError):
36
+ pass
37
+
38
+
39
+ def validate_json_object(value: object, *, field: str) -> dict[str, JSONValue]:
40
+ if not isinstance(value, dict):
41
+ raise RequestValidationError(f"{field} must be a JSON object.")
42
+ validate_json_value(value, field=field)
43
+ return cast(dict[str, JSONValue], value)
44
+
45
+
46
+ def validate_json_value(value: object, *, field: str) -> None:
47
+ active_containers: set[int] = set()
48
+
49
+ def walk(current: object, depth: int, path: str) -> None:
50
+ if isinstance(current, str):
51
+ validate_unicode_scalar(current, field=path)
52
+ return
53
+ if current is None or isinstance(current, bool):
54
+ return
55
+ if isinstance(current, int):
56
+ if not INT64_MIN <= current <= INT64_MAX:
57
+ raise RequestValidationError(f"{path} is outside the signed 64-bit range.")
58
+ return
59
+ if isinstance(current, float):
60
+ if not math.isfinite(current):
61
+ raise RequestValidationError(f"{path} must be finite.")
62
+ return
63
+ if isinstance(current, dict):
64
+ if depth >= MAX_JSON_DEPTH:
65
+ raise RequestValidationError(f"{field} exceeds JSON depth {MAX_JSON_DEPTH}.")
66
+ identity = id(current)
67
+ if identity in active_containers:
68
+ raise RequestValidationError(f"{field} contains a cycle.")
69
+ active_containers.add(identity)
70
+ try:
71
+ for key, child in current.items():
72
+ if not isinstance(key, str):
73
+ raise RequestValidationError(f"{path} has a non-string object key.")
74
+ validate_unicode_scalar(key, field=f"{path}.<key>")
75
+ walk(child, depth + 1, f"{path}.{key}")
76
+ finally:
77
+ active_containers.remove(identity)
78
+ return
79
+ if isinstance(current, Sequence) and not isinstance(current, (str, bytes, bytearray)):
80
+ if not isinstance(current, list):
81
+ raise RequestValidationError(f"{path} must use JSON arrays, not Python sequences.")
82
+ if depth >= MAX_JSON_DEPTH:
83
+ raise RequestValidationError(f"{field} exceeds JSON depth {MAX_JSON_DEPTH}.")
84
+ identity = id(current)
85
+ if identity in active_containers:
86
+ raise RequestValidationError(f"{field} contains a cycle.")
87
+ active_containers.add(identity)
88
+ try:
89
+ for index, child in enumerate(current):
90
+ walk(child, depth + 1, f"{path}[{index}]")
91
+ finally:
92
+ active_containers.remove(identity)
93
+ return
94
+ raise RequestValidationError(f"{path} is not representable in strict JSON.")
95
+
96
+ walk(value, 0, field)
97
+
98
+
99
+ def validate_unicode_scalar(value: str, *, field: str) -> str:
100
+ if any(0xD800 <= ord(character) <= 0xDFFF for character in value):
101
+ raise RequestValidationError(f"{field} contains a lone Unicode surrogate.")
102
+ return value
103
+
104
+
105
+ def validate_identifier(value: object, *, field: str) -> str:
106
+ if not isinstance(value, str) or not IDENTIFIER_RE.fullmatch(value):
107
+ raise RequestValidationError(f"{field} must match [A-Za-z0-9][A-Za-z0-9._-]{{0,127}}.")
108
+ return value
109
+
110
+
111
+ def validate_task_id(value: object) -> str:
112
+ if not isinstance(value, str) or not TASK_ID_RE.fullmatch(value):
113
+ raise RequestValidationError("task_id must match t_[A-Za-z0-9_-]{12}.")
114
+ return value
115
+
116
+
117
+ def validate_run_id(value: object) -> str:
118
+ if not isinstance(value, str) or not RUN_ID_RE.fullmatch(value):
119
+ raise RequestValidationError("run_id must match r_[A-Za-z0-9_-]{12}.")
120
+ return value
121
+
122
+
123
+ def validate_task_name(value: object) -> str | None:
124
+ if value is None:
125
+ return None
126
+ if not isinstance(value, str):
127
+ raise RequestValidationError("name must be a string or None.")
128
+ validate_unicode_scalar(value, field="name")
129
+ if len(value) > 256:
130
+ raise RequestValidationError("name exceeds 256 Unicode code points.")
131
+ if any(unicodedata.category(character) == "Cc" for character in value):
132
+ raise RequestValidationError("name contains a control character.")
133
+ return value
134
+
135
+
136
+ def validate_int64(value: object, *, field: str, positive: bool = False) -> int:
137
+ if isinstance(value, bool) or not isinstance(value, int):
138
+ raise RequestValidationError(f"{field} must be an integer.")
139
+ if not INT64_MIN <= value <= INT64_MAX or (positive and value <= 0):
140
+ qualifier = "a positive signed 64-bit integer" if positive else "a signed 64-bit integer"
141
+ raise RequestValidationError(f"{field} must be {qualifier}.")
142
+ return value
143
+
144
+
145
+ def validate_routes(value: object) -> list[str]:
146
+ if not isinstance(value, list) or not value:
147
+ raise RequestValidationError("routes must be a non-empty list of strings.")
148
+ routes = [validate_identifier(route, field="route") for route in value]
149
+ if len(routes) != len(set(routes)):
150
+ raise RequestValidationError("routes must not contain duplicates.")
151
+ return sorted(routes)
152
+
153
+
154
+ def validate_filter(value: object, *, required: bool = False) -> str | None:
155
+ if value is None:
156
+ if required:
157
+ raise RequestValidationError("filter must be a non-empty string.")
158
+ return None
159
+ if not isinstance(value, str) or not value.strip():
160
+ raise RequestValidationError("filter must be a non-empty string.")
161
+ validate_unicode_scalar(value, field="filter")
162
+ if len(value.encode("utf-8")) > MAX_FILTER_BYTES:
163
+ raise RequestValidationError(f"filter exceeds {MAX_FILTER_BYTES} bytes.")
164
+ return value
165
+
166
+
167
+ def validate_status(value: object | None) -> TaskStatus | None:
168
+ if value is None:
169
+ return None
170
+ if not isinstance(value, str) or value not in TASK_STATUSES:
171
+ raise RequestValidationError("status is not a valid Task status.")
172
+ return cast(TaskStatus, value)
173
+
174
+
175
+ def validate_order_field(value: object) -> TaskOrderField:
176
+ if not isinstance(value, str) or value not in TASK_ORDER_FIELDS:
177
+ raise RequestValidationError("order_by is not a supported Task field.")
178
+ return cast(TaskOrderField, value)
179
+
180
+
181
+ def validate_task_update(changes: object) -> TaskUpdate:
182
+ if not isinstance(changes, dict) or not changes:
183
+ raise RequestValidationError("changes must be a non-empty object.")
184
+ allowed = {"name", "args", "metadata", "priority", "max_attempts", "routes", "result"}
185
+ unknown = set(changes) - allowed
186
+ if unknown:
187
+ raise RequestValidationError(f"changes contains unsupported fields: {sorted(unknown)!r}.")
188
+ normalized: dict[str, object] = {}
189
+ for field, value in changes.items():
190
+ if field == "name":
191
+ normalized[field] = validate_task_name(value)
192
+ elif field in {"args", "metadata", "result"}:
193
+ normalized[field] = validate_json_object(value, field=field)
194
+ elif field == "priority":
195
+ normalized[field] = validate_int64(value, field=field)
196
+ elif field == "max_attempts":
197
+ normalized[field] = validate_int64(value, field=field, positive=True)
198
+ elif field == "routes":
199
+ normalized[field] = validate_routes(value)
200
+ return cast(TaskUpdate, normalized)
201
+
202
+
203
+ def invalid_config(message: str, *, source: str, field: str | None = None) -> ConfigError:
204
+ details = {"source": source}
205
+ if field is not None:
206
+ details["field"] = field
207
+ return ConfigError("invalid_config", message, details)