labtasker-client 2.4.1__tar.gz → 2.6.0__tar.gz

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.
Files changed (26) hide show
  1. {labtasker_client-2.4.1 → labtasker_client-2.6.0}/PKG-INFO +1 -1
  2. {labtasker_client-2.4.1 → labtasker_client-2.6.0}/pyproject.toml +1 -1
  3. {labtasker_client-2.4.1 → labtasker_client-2.6.0}/src/labtasker/__init__.py +3 -1
  4. {labtasker_client-2.4.1 → labtasker_client-2.6.0}/src/labtasker/cli.py +54 -27
  5. {labtasker_client-2.4.1 → labtasker_client-2.6.0}/src/labtasker/client.py +108 -49
  6. {labtasker_client-2.4.1 → labtasker_client-2.6.0}/src/labtasker/command_worker.py +36 -26
  7. labtasker_client-2.6.0/src/labtasker/config.py +294 -0
  8. {labtasker_client-2.4.1 → labtasker_client-2.6.0}/src/labtasker/execution.py +66 -21
  9. {labtasker_client-2.4.1 → labtasker_client-2.6.0}/src/labtasker/journal.py +33 -18
  10. {labtasker_client-2.4.1 → labtasker_client-2.6.0}/src/labtasker/local.py +40 -48
  11. {labtasker_client-2.4.1 → labtasker_client-2.6.0}/src/labtasker/models.py +22 -4
  12. {labtasker_client-2.4.1 → labtasker_client-2.6.0}/src/labtasker/observations.py +17 -3
  13. {labtasker_client-2.4.1 → labtasker_client-2.6.0}/src/labtasker/worker.py +190 -16
  14. labtasker_client-2.4.1/src/labtasker/config.py +0 -213
  15. {labtasker_client-2.4.1 → labtasker_client-2.6.0}/.gitignore +0 -0
  16. {labtasker_client-2.4.1 → labtasker_client-2.6.0}/LICENSE +0 -0
  17. {labtasker_client-2.4.1 → labtasker_client-2.6.0}/src/labtasker/__main__.py +0 -0
  18. {labtasker_client-2.4.1 → labtasker_client-2.6.0}/src/labtasker/api.py +0 -0
  19. {labtasker_client-2.4.1 → labtasker_client-2.6.0}/src/labtasker/binding.py +0 -0
  20. {labtasker_client-2.4.1 → labtasker_client-2.6.0}/src/labtasker/command_template.py +0 -0
  21. {labtasker_client-2.4.1 → labtasker_client-2.6.0}/src/labtasker/errors.py +0 -0
  22. {labtasker_client-2.4.1 → labtasker_client-2.6.0}/src/labtasker/paths.py +0 -0
  23. {labtasker_client-2.4.1 → labtasker_client-2.6.0}/src/labtasker/py.typed +0 -0
  24. {labtasker_client-2.4.1 → labtasker_client-2.6.0}/src/labtasker/tee.py +0 -0
  25. {labtasker_client-2.4.1 → labtasker_client-2.6.0}/src/labtasker/types.py +0 -0
  26. {labtasker_client-2.4.1 → labtasker_client-2.6.0}/src/labtasker/validation.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.5
2
2
  Name: labtasker-client
3
- Version: 2.4.1
3
+ Version: 2.6.0
4
4
  Summary: A small task queue for parallel model inference and evaluation
5
5
  Project-URL: Homepage, https://github.com/luocfprime/labtasker
6
6
  Project-URL: Repository, https://github.com/luocfprime/labtasker.git
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
4
4
 
5
5
  [project]
6
6
  name = "labtasker-client"
7
- version = "2.4.1"
7
+ version = "2.6.0"
8
8
  description = "A small task queue for parallel model inference and evaluation"
9
9
  requires-python = ">=3.10"
10
10
  license = "Apache-2.0"
@@ -31,6 +31,7 @@ from labtasker.execution import (
31
31
  cancellation_requested,
32
32
  finish,
33
33
  report_progress,
34
+ report_worker_telemetry,
34
35
  set_force_stop_timeout,
35
36
  task_info,
36
37
  )
@@ -49,7 +50,7 @@ from labtasker.models import (
49
50
  from labtasker.types import JSONValue, TaskOrderField, TaskStatus, TaskUpdate
50
51
  from labtasker.worker import loop
51
52
 
52
- __version__ = "2.4.1"
53
+ __version__ = "2.6.0"
53
54
 
54
55
  __all__ = [
55
56
  "APIError",
@@ -89,6 +90,7 @@ __all__ = [
89
90
  "list_workers",
90
91
  "loop",
91
92
  "report_progress",
93
+ "report_worker_telemetry",
92
94
  "requeue_task",
93
95
  "set_force_stop_timeout",
94
96
  "submit_task",
@@ -3,6 +3,9 @@ from __future__ import annotations
3
3
  import json
4
4
  import logging
5
5
  from collections.abc import Callable
6
+ from contextvars import ContextVar
7
+ from dataclasses import dataclass
8
+ from pathlib import Path
6
9
  from typing import Annotated, Any, TypeVar, cast
7
10
 
8
11
  import typer
@@ -15,7 +18,6 @@ from labtasker.command_template import TemplateSyntaxError
15
18
  from labtasker.command_worker import run_command_worker
16
19
  from labtasker.config import resolve_config
17
20
  from labtasker.errors import LabtaskerError
18
- from labtasker.execution import report_progress as report_current_progress
19
21
  from labtasker.types import TaskOrderField, TaskStatus, TaskUpdate
20
22
  from labtasker.validation import RequestValidationError, validate_grouping, validate_json_object
21
23
 
@@ -58,6 +60,15 @@ app.add_typer(config_app, name="config")
58
60
  logger = logging.getLogger("labtasker.cli")
59
61
 
60
62
 
63
+ @dataclass(frozen=True, slots=True)
64
+ class CLIState:
65
+ labtasker_root: Path | None
66
+ auto_start_local_server: bool
67
+
68
+
69
+ _CLI_STATE: ContextVar[CLIState | None] = ContextVar("labtasker_cli_state", default=None)
70
+
71
+
61
72
  def _version_callback(value: bool) -> None:
62
73
  if value:
63
74
  typer.echo(f"labtasker-client {__version__}")
@@ -75,8 +86,20 @@ def main(
75
86
  help="Show the Client package version and exit.",
76
87
  ),
77
88
  ] = False,
89
+ labtasker_root: Annotated[
90
+ Path | None,
91
+ typer.Option(help="Exact configuration, journal, and managed-local root."),
92
+ ] = None,
93
+ auto_start_local_server: Annotated[
94
+ bool,
95
+ typer.Option(
96
+ "--auto-start-local-server",
97
+ help="Allow this invocation to start a managed-local daemon.",
98
+ ),
99
+ ] = False,
78
100
  ) -> None:
79
101
  """Submit, inspect, and execute Labtasker v2 Tasks."""
102
+ _CLI_STATE.set(CLIState(labtasker_root, auto_start_local_server))
80
103
 
81
104
 
82
105
  class _SeparatedCommand(TyperCommand):
@@ -126,6 +149,10 @@ def worker_loop(
126
149
  )
127
150
  ),
128
151
  ] = None,
152
+ metadata: Annotated[
153
+ str,
154
+ typer.Option(help="Static Worker metadata as one strict JSON object."),
155
+ ] = "{}",
129
156
  ) -> None:
130
157
  """Claim matching Tasks and execute one child command for each claim.
131
158
 
@@ -146,6 +173,7 @@ def worker_loop(
146
173
  argv.pop(0)
147
174
  if not argv:
148
175
  raise typer.BadParameter("COMMAND is required after --")
176
+ worker_metadata = _json_object(metadata, option="--metadata")
149
177
  try:
150
178
  run_command_worker(
151
179
  argv,
@@ -154,6 +182,9 @@ def worker_loop(
154
182
  idle_timeout=idle_timeout,
155
183
  max_consecutive_failures=max_consecutive_failures,
156
184
  force_stop_timeout=force_stop_timeout,
185
+ metadata=worker_metadata,
186
+ labtasker_root=_cli_state().labtasker_root,
187
+ auto_start_local_server=_cli_state().auto_start_local_server,
157
188
  )
158
189
  except (TemplateSyntaxError, RequestValidationError) as error:
159
190
  raise typer.BadParameter(str(error)) from error
@@ -170,29 +201,6 @@ def worker_loop(
170
201
  raise typer.Exit(1) from error
171
202
 
172
203
 
173
- @app.command("progress")
174
- def progress_report(
175
- data: Annotated[
176
- str,
177
- typer.Option(help="Latest progress as one strict JSON object."),
178
- ],
179
- ) -> None:
180
- """Replace the current Task run's progress snapshot.
181
-
182
- This command is available inside a command launched by ``labtasker loop``.
183
- It prints whether the best-effort report was accepted; transport failures
184
- and confirmed revocation return ``reported: false`` without failing the
185
- command workload.
186
- """
187
- try:
188
- progress = _json_object(data, option="--data")
189
- reported = _invoke(lambda: report_current_progress(progress))
190
- except RuntimeError as error:
191
- typer.echo(str(error), err=True)
192
- raise typer.Exit(1) from error
193
- _write_json({"reported": reported})
194
-
195
-
196
204
  @task_app.command("submit")
197
205
  def task_submit(
198
206
  args: Annotated[
@@ -599,7 +607,10 @@ def queue_delete(
599
607
  name: Annotated[str, typer.Argument(help="Queue name to delete.")],
600
608
  cascade: Annotated[
601
609
  bool,
602
- typer.Option(help="Also permanently delete every non-running Task in the Queue."),
610
+ typer.Option(
611
+ "--cascade",
612
+ help="Also permanently delete every non-running Task in the Queue.",
613
+ ),
603
614
  ] = False,
604
615
  ) -> None:
605
616
  """Permanently delete one Queue.
@@ -618,14 +629,30 @@ def config_show() -> None:
618
629
  .labtasker/config.toml, then built-in defaults. The token value is never
619
630
  printed.
620
631
  """
621
- _write_json(_invoke(lambda: resolve_config().public_dict()))
632
+ state = _cli_state()
633
+ _write_json(
634
+ _invoke(
635
+ lambda: resolve_config(
636
+ labtasker_root=state.labtasker_root,
637
+ auto_start_local_server=state.auto_start_local_server,
638
+ ).public_dict()
639
+ )
640
+ )
622
641
 
623
642
 
624
643
  def _with_client(operation: Callable[[Client], T]) -> T:
625
- with Client() as client:
644
+ state = _cli_state()
645
+ with Client(
646
+ labtasker_root=state.labtasker_root,
647
+ auto_start_local_server=state.auto_start_local_server,
648
+ ) as client:
626
649
  return operation(client)
627
650
 
628
651
 
652
+ def _cli_state() -> CLIState:
653
+ return _CLI_STATE.get() or CLIState(None, False)
654
+
655
+
629
656
  def _invoke(operation: Callable[[], T]) -> T:
630
657
  try:
631
658
  return operation()
@@ -13,12 +13,7 @@ from pydantic import TypeAdapter, ValidationError
13
13
 
14
14
  from labtasker.config import ResolvedConfig, resolve_config
15
15
  from labtasker.errors import APIError, TransportError
16
- from labtasker.local import (
17
- ensure_local_server,
18
- local_paths,
19
- require_local_capabilities,
20
- socket_transport,
21
- )
16
+ from labtasker.local import ensure_local_server, socket_transport
22
17
  from labtasker.models import (
23
18
  BulkUpdateResult,
24
19
  ClaimResponse,
@@ -53,7 +48,7 @@ from labtasker.validation import (
53
48
 
54
49
  T = TypeVar("T")
55
50
  ModelT = TypeVar("ModelT", bound=ResponseModel)
56
- REQUEST_TIMEOUT_SECONDS = 10.0
51
+ REQUEST_TIMEOUT_SECONDS = 15.0
57
52
  MAX_RETRY_ATTEMPTS = 3
58
53
  RETRY_BACKOFF_SECONDS = (0.05, 0.1)
59
54
  QUEUE_LIST_ADAPTER = TypeAdapter(list[Queue])
@@ -63,21 +58,36 @@ class Client:
63
58
  def __init__(
64
59
  self,
65
60
  url: str | None = None,
61
+ socket: str | Path | None = None,
62
+ labtasker_root: str | Path | None = None,
63
+ auto_start_local_server: bool = False,
66
64
  token: str | None = None,
67
65
  queue: str | None = None,
68
66
  ) -> None:
69
- self._initialize(resolve_config(url=url, token=token, queue=queue))
67
+ self._initialize(
68
+ resolve_config(
69
+ url=url,
70
+ socket=socket,
71
+ labtasker_root=labtasker_root,
72
+ auto_start_local_server=auto_start_local_server,
73
+ token=token,
74
+ queue=queue,
75
+ )
76
+ )
70
77
 
71
78
  @classmethod
72
- def _from_local_directory(cls, directory: Path, *, queue: str) -> Client:
73
- require_local_capabilities()
79
+ def _from_socket(cls, socket: Path, *, queue: str) -> Client:
74
80
  client = cls.__new__(cls)
75
81
  client._initialize(
76
82
  ResolvedConfig(
77
83
  url=None,
84
+ socket=socket,
85
+ managed_local=False,
86
+ labtasker_root=Path("/"),
78
87
  queue=validate_identifier(queue, field="queue"),
79
88
  token=None,
80
- local=local_paths(directory),
89
+ auto_start_local_server=False,
90
+ local=None,
81
91
  )
82
92
  )
83
93
  return client
@@ -87,22 +97,22 @@ class Client:
87
97
  headers = {}
88
98
  if self._config.token is not None:
89
99
  headers["Authorization"] = f"Bearer {self._config.token}"
90
- if self._config.local is None:
91
- assert self._config.url is not None
100
+ if self._config.url is not None:
92
101
  self._http = httpx.Client(
93
102
  base_url=f"{self._config.url}/api/v2/",
94
103
  headers=headers,
95
104
  timeout=REQUEST_TIMEOUT_SECONDS,
96
105
  )
97
106
  else:
107
+ assert self._config.socket is not None
98
108
  self._http = httpx.Client(
99
109
  base_url="http://labtasker/api/v2/",
100
- transport=socket_transport(self._config.local),
110
+ transport=socket_transport(self._config.socket),
101
111
  timeout=REQUEST_TIMEOUT_SECONDS,
102
112
  )
103
113
  self._closed = False
104
114
  self._endpoint_announced = False
105
- self._local_ready = False
115
+ self._local_ensure_result: tuple[int | None, str | None] | None = None
106
116
  self._server_version: str | None = None
107
117
  self._warned_server_versions: set[Version] = set()
108
118
 
@@ -148,7 +158,7 @@ class Client:
148
158
  )
149
159
 
150
160
  @property
151
- def configuration(self) -> ResolvedConfig:
161
+ def _configuration(self) -> ResolvedConfig:
152
162
  return self._config
153
163
 
154
164
  def submit_task(
@@ -538,7 +548,7 @@ class Client:
538
548
  path=f"queues/{queue_name}/tasks/claim",
539
549
  json={"route": normalized_route, "run_id": normalized_run_id},
540
550
  parser=_parse_claim,
541
- retry=True,
551
+ recover_local_connect=False,
542
552
  )
543
553
 
544
554
  def _health(self) -> HealthResponse:
@@ -561,6 +571,7 @@ class Client:
561
571
  task_id: str,
562
572
  run_id: str,
563
573
  queue: str | None = None,
574
+ recover_local_connect: bool = True,
564
575
  ) -> HeartbeatResponse:
565
576
  return self._run_action(
566
577
  "heartbeat",
@@ -569,6 +580,7 @@ class Client:
569
580
  queue=queue,
570
581
  body={},
571
582
  parser=lambda response: _parse_model(response, HeartbeatResponse, {200}),
583
+ recover_local_connect=recover_local_connect,
572
584
  )
573
585
 
574
586
  def _complete(
@@ -607,6 +619,27 @@ class Client:
607
619
  parser=lambda response: _parse_none(response, {204}),
608
620
  )
609
621
 
622
+ def _report_worker_telemetry(
623
+ self,
624
+ *,
625
+ worker_id: str,
626
+ telemetry: dict[str, JSONValue],
627
+ queue: str | None = None,
628
+ ) -> None:
629
+ self._ensure_open()
630
+ queue_name = self._queue(queue)
631
+ normalized_telemetry = validate_json_object(telemetry, field="telemetry")
632
+ from labtasker.validation import validate_worker_id
633
+
634
+ normalized_worker_id = validate_worker_id(worker_id)
635
+ self._call(
636
+ operation="report_worker_telemetry",
637
+ method="POST",
638
+ path=f"queues/{queue_name}/workers/{normalized_worker_id}/telemetry",
639
+ json={"telemetry": normalized_telemetry},
640
+ parser=lambda response: _parse_none(response, {204}),
641
+ )
642
+
610
643
  def _fail(
611
644
  self,
612
645
  *,
@@ -686,6 +719,7 @@ class Client:
686
719
  queue: str | None,
687
720
  body: dict[str, object],
688
721
  parser: Callable[[httpx.Response], T],
722
+ recover_local_connect: bool = True,
689
723
  ) -> T:
690
724
  self._ensure_open()
691
725
  queue_name = self._queue(queue)
@@ -697,6 +731,7 @@ class Client:
697
731
  path=f"queues/{queue_name}/tasks/{normalized_task_id}/{action}",
698
732
  json={"run_id": normalized_run_id, **body},
699
733
  parser=parser,
734
+ recover_local_connect=recover_local_connect,
700
735
  )
701
736
 
702
737
  def _queue(self, queue: str | None) -> str:
@@ -712,9 +747,9 @@ class Client:
712
747
  json: object | None = None,
713
748
  params: dict[str, str | int] | None = None,
714
749
  retry: bool = False,
750
+ recover_local_connect: bool = True,
715
751
  ) -> T:
716
752
  self._ensure_open()
717
- self._prepare_endpoint()
718
753
  attempts = MAX_RETRY_ATTEMPTS if retry else 1
719
754
  last_transport_error: TransportError | None = None
720
755
  local_connect_recovery_used = False
@@ -725,13 +760,14 @@ class Client:
725
760
  except httpx.RequestError as error:
726
761
  last_transport_error = self._connection_error(operation)
727
762
  can_recover_local_connect = (
728
- self._config.local is not None
763
+ recover_local_connect
764
+ and self._config.managed_local
765
+ and self._config.auto_start_local_server
729
766
  and isinstance(error, (httpx.ConnectError, httpx.ConnectTimeout))
730
767
  and not local_connect_recovery_used
731
768
  )
732
769
  if can_recover_local_connect:
733
770
  local_connect_recovery_used = True
734
- self._local_ready = False
735
771
  self._ensure_local_available()
736
772
  if attempt + 1 == attempts:
737
773
  attempts += 1
@@ -739,7 +775,7 @@ class Client:
739
775
  raise last_transport_error from error
740
776
  else:
741
777
  self._observe_server_version(response)
742
- self._announce_http_endpoint()
778
+ self._announce_endpoint()
743
779
  if response.is_error:
744
780
  try:
745
781
  api_error = _parse_api_error(response)
@@ -773,32 +809,55 @@ class Client:
773
809
  raise AssertionError("Request loop ended without a result or error.")
774
810
  raise last_transport_error
775
811
 
812
+ def _repair_local_connection(self, error: TransportError) -> None:
813
+ """Repair an opted-in managed-local connection without replaying a request."""
814
+ cause = error.__cause__
815
+ if (
816
+ self._config.managed_local
817
+ and self._config.auto_start_local_server
818
+ and isinstance(cause, (httpx.ConnectError, httpx.ConnectTimeout))
819
+ ):
820
+ self._ensure_local_available()
821
+
776
822
  def _ensure_open(self) -> None:
777
823
  if self._closed:
778
824
  raise RuntimeError("Client is closed.")
779
825
 
780
826
  @property
781
827
  def _operation_endpoint_details(self) -> dict[str, object]:
782
- if self._config.local is None:
828
+ if self._config.url is not None:
783
829
  return {"url": self._config.url}
784
830
  return {
785
- "directory": str(self._config.local.directory),
786
- "socket": str(self._config.local.socket),
831
+ "labtasker_root": str(self._config.labtasker_root),
832
+ "socket": str(self._config.socket),
787
833
  }
788
834
 
789
- def _prepare_endpoint(self) -> None:
790
- if self._config.local is not None and not self._local_ready:
791
- self._ensure_local_available()
792
-
793
- def _announce_http_endpoint(self) -> None:
794
- if self._endpoint_announced or self._config.local is not None:
835
+ def _announce_endpoint(self) -> None:
836
+ if self._endpoint_announced:
795
837
  return
796
- assert self._config.url is not None
797
- transport = self._config.url.partition(":")[0]
798
- print(
799
- f"[labtasker] connected server=remote transport={transport} url={self._config.url}",
800
- file=sys.stderr,
801
- )
838
+ if self._config.url is not None:
839
+ transport = self._config.url.partition(":")[0]
840
+ message = (
841
+ f"[labtasker] connected server=remote transport={transport} url={self._config.url}"
842
+ )
843
+ elif self._config.managed_local:
844
+ assert self._config.local is not None and self._config.socket is not None
845
+ pid, version = self._local_ensure_result or (None, None)
846
+ message = (
847
+ "[labtasker] connected server=local transport=unix "
848
+ f"labtasker_root={self._config.labtasker_root} "
849
+ f"database={self._config.local.database} socket={self._config.socket}"
850
+ )
851
+ if pid is not None:
852
+ message += f" pid={pid}"
853
+ if version is not None:
854
+ message += f" version={version}"
855
+ else:
856
+ assert self._config.socket is not None
857
+ message = (
858
+ f"[labtasker] connected server=external transport=unix socket={self._config.socket}"
859
+ )
860
+ print(message, file=sys.stderr)
802
861
  self._endpoint_announced = True
803
862
 
804
863
  def _ensure_local_available(self) -> None:
@@ -806,16 +865,7 @@ class Client:
806
865
  if paths is None:
807
866
  return
808
867
  result = ensure_local_server(paths, emit=self._emit_local_transition)
809
- pid = result.pid if result.pid is not None else "unknown"
810
- version = result.server_version if result.server_version is not None else "unknown"
811
- print(
812
- "[labtasker] connected server=local transport=unix "
813
- f"directory={paths.directory} database={paths.database} socket={paths.socket} "
814
- f"pid={pid} version={version}",
815
- file=sys.stderr,
816
- )
817
- self._endpoint_announced = True
818
- self._local_ready = True
868
+ self._local_ensure_result = (result.pid, result.server_version)
819
869
 
820
870
  @staticmethod
821
871
  def _emit_local_transition(message: str) -> None:
@@ -823,16 +873,25 @@ class Client:
823
873
 
824
874
  def _connection_error(self, operation: str) -> TransportError:
825
875
  details: dict[str, object] = {"operation": operation}
826
- if self._config.local is None:
876
+ if self._config.url is not None:
827
877
  details["url"] = self._config.url
878
+ elif not self._config.managed_local:
879
+ details["socket"] = str(self._config.socket)
828
880
  else:
881
+ assert self._config.local is not None
829
882
  details.update(
830
883
  {
831
884
  "state": "unhealthy",
832
- "directory": str(self._config.local.directory),
885
+ "labtasker_root": str(self._config.labtasker_root),
833
886
  "database": str(self._config.local.database),
834
- "socket": str(self._config.local.socket),
887
+ "socket": str(self._config.socket),
835
888
  "log": str(self._config.local.log),
889
+ "remedies": [
890
+ "rerun with --auto-start-local-server",
891
+ "launch labtasker-server serve --connection socket --daemon "
892
+ f"--labtasker-root {self._config.labtasker_root}",
893
+ "configure LABTASKER_URL or LABTASKER_SOCKET",
894
+ ],
836
895
  }
837
896
  )
838
897
  return TransportError("The Labtasker Server could not be reached.", details)
@@ -27,17 +27,16 @@ from labtasker.models import ClaimResponse
27
27
  from labtasker.observations import ObservationReporter
28
28
  from labtasker.tee import configure_worker_logger
29
29
  from labtasker.types import JSONValue
30
- from labtasker.validation import validate_identifier
30
+ from labtasker.validation import validate_identifier, validate_json_object
31
31
  from labtasker.worker import (
32
- POLL_INTERVAL_SECONDS,
33
32
  Heartbeat,
34
33
  _best_effort_unclaim,
35
34
  _ExecutionResult,
36
35
  _FailureGuard,
37
36
  _finish_journal,
38
- _generate_run_id,
39
37
  _guard_worker_topology,
40
38
  _journal_best_effort,
39
+ _next_claim,
41
40
  _preflight,
42
41
  _report_until_resolved,
43
42
  _safe_diagnostic_text,
@@ -57,36 +56,41 @@ def run_command_worker(
57
56
  idle_timeout: float = 300.0,
58
57
  force_stop_timeout: float | None = None,
59
58
  max_consecutive_failures: int = 5,
59
+ metadata: dict[str, JSONValue] | None = None,
60
+ labtasker_root: Path | None = None,
61
+ auto_start_local_server: bool = False,
60
62
  ) -> None:
61
63
  guard = _FailureGuard(max_consecutive_failures)
62
64
  templates = compile_argv(argv)
63
65
  normalized_route = validate_identifier(route, field="route")
64
66
  normalized_idle_timeout = _validate_idle_timeout(idle_timeout)
65
67
  normalized_force_stop_timeout = _validate_force_stop_timeout(force_stop_timeout)
68
+ normalized_metadata = validate_json_object(
69
+ {} if metadata is None else metadata, field="metadata"
70
+ )
66
71
  _guard_command_worker_platform()
67
72
  _guard_worker_topology()
68
73
  configure_worker_logger()
69
- with Client(queue=queue) as client:
70
- queue_name = client.configuration.queue
74
+ with Client(
75
+ queue=queue,
76
+ labtasker_root=labtasker_root,
77
+ auto_start_local_server=auto_start_local_server,
78
+ ) as client:
79
+ queue_name = client._configuration.queue
71
80
  _preflight(client, queue_name)
72
- with ObservationReporter(client.configuration, normalized_route) as observer:
73
- idle_deadline: float | None = None
81
+ with ObservationReporter(
82
+ client._configuration, normalized_route, normalized_metadata
83
+ ) as observer:
74
84
  while True:
75
- claim = client._claim(
85
+ claim = _next_claim(
86
+ client,
76
87
  route=normalized_route,
77
- run_id=_generate_run_id(),
78
88
  queue=queue_name,
89
+ idle_timeout=normalized_idle_timeout,
79
90
  )
80
91
  if claim is None:
81
- now = time.monotonic()
82
- if idle_deadline is None:
83
- idle_deadline = now + normalized_idle_timeout
84
- if now >= idle_deadline:
85
- logger.info("Worker idle timeout reached; stopping normally.")
86
- return
87
- time.sleep(min(POLL_INTERVAL_SECONDS, idle_deadline - now))
88
- continue
89
- idle_deadline = None
92
+ logger.info("Worker idle timeout reached; stopping normally.")
93
+ return
90
94
  observer.activity(claim.task.id)
91
95
  logger.info(
92
96
  "Claimed Task %s as run %s (attempt %d, route %s).",
@@ -102,6 +106,7 @@ def run_command_worker(
102
106
  queue=queue_name,
103
107
  route=normalized_route,
104
108
  force_stop_timeout=normalized_force_stop_timeout,
109
+ worker_id=observer.id,
105
110
  )
106
111
 
107
112
  guard.observe(result, claim.task.id)
@@ -124,13 +129,15 @@ def _run_command_claim(
124
129
  queue: str,
125
130
  route: str,
126
131
  force_stop_timeout: float | None,
132
+ worker_id: str,
127
133
  ) -> _ExecutionResult:
128
134
  try:
129
135
  journal = LocalRunJournal.create(
130
136
  claim=claim,
131
- endpoint=client.configuration.endpoint_dict(),
137
+ endpoint=client._configuration.endpoint_dict(),
132
138
  queue=queue,
133
139
  route=route,
140
+ labtasker_root=client._configuration.labtasker_root,
134
141
  )
135
142
  except Exception:
136
143
  _best_effort_unclaim(client, claim, queue)
@@ -153,7 +160,7 @@ def _run_command_claim(
153
160
  return _report_command_failure(
154
161
  client, journal, claim, queue, "TaskBindingError", str(error), control=control
155
162
  )
156
- environment = _command_environment(client, claim, journal, queue, route)
163
+ environment = _command_environment(client, claim, journal, queue, route, worker_id)
157
164
  try:
158
165
  if _interactive_terminal():
159
166
  process = _run_pty(
@@ -489,6 +496,7 @@ def _command_environment(
489
496
  journal: LocalRunJournal,
490
497
  queue: str,
491
498
  route: str,
499
+ worker_id: str,
492
500
  ) -> dict[str, str]:
493
501
  environment = dict(os.environ)
494
502
  environment.update(
@@ -498,20 +506,22 @@ def _command_environment(
498
506
  "LABTASKER_RUN_ID": claim.run_id,
499
507
  "LABTASKER_ROUTE": route,
500
508
  "LABTASKER_RUN_DIR": str(journal.run_dir),
509
+ "LABTASKER_WORKER_ID": worker_id,
501
510
  }
502
511
  )
503
- configuration = client.configuration
504
- if configuration.local is None:
512
+ configuration = client._configuration
513
+ if configuration.url is not None:
505
514
  assert configuration.url is not None
506
515
  environment["LABTASKER_URL"] = configuration.url
507
516
  environment.pop("LABTASKER_SOCKET", None)
508
- environment.pop("LABTASKER_LOCAL_DIRECTORY", None)
509
517
  else:
510
- environment["LABTASKER_SOCKET"] = str(configuration.local.socket)
511
- environment["LABTASKER_LOCAL_DIRECTORY"] = str(configuration.local.directory)
518
+ assert configuration.socket is not None
519
+ environment["LABTASKER_SOCKET"] = str(configuration.socket)
512
520
  environment.pop("LABTASKER_URL", None)
521
+ environment.pop("LABTASKER_ROOT", None)
522
+ environment.pop("LABTASKER_LOCAL_DIRECTORY", None)
513
523
  token = configuration.token
514
- if token is None or configuration.local is not None:
524
+ if token is None or configuration.url is None:
515
525
  environment.pop("LABTASKER_TOKEN", None)
516
526
  else:
517
527
  environment["LABTASKER_TOKEN"] = token