feat: self-service worker enrollment bound to a user account
Let a signed-in user turn their own machine into a worker without the shared token. The coordinator already binds a JWT-authenticated registration to owner_id as untrusted; this adds the missing pieces. userservice: long-lived worker keys (scimesh_wk_live_*, hash-at-rest) with create/list/revoke and a public /worker-tokens/exchange that trades a key for a short-lived JWT carrying the owner current role/verified. python worker: SCIMESH_WORKER_KEY + SCIMESH_USERSERVICE_URL; a token provider exchanges the key and refreshes the JWT proactively and on 401, so a long-running worker survives token expiry. Static bearer token path is unchanged. coordinator UI: an "add your machine" page that mints a key and shows a ready-to-run command, proxying key management to the userservice; the dashboard gains an owner-scoped "my machines" section. docs: how to run a worker from your account, plus the untrusted/quorum/ verified trust model.
This commit is contained in:
@@ -7,9 +7,11 @@ import http.client
|
||||
import json
|
||||
from pathlib import Path
|
||||
from typing import Protocol
|
||||
from urllib.error import HTTPError
|
||||
from urllib.parse import quote, urljoin, urlsplit
|
||||
from urllib.request import Request, build_opener
|
||||
|
||||
from .auth import StaticTokenProvider, TokenProvider
|
||||
from .coordinator import CoordinatorConflictError
|
||||
from .models import ClaimedTask, ProducedArtifact, UploadedArtifact
|
||||
from .transport import SameOriginAuthRedirectHandler, origin
|
||||
@@ -29,23 +31,51 @@ class ArtifactClient(Protocol):
|
||||
class HttpArtifactClient:
|
||||
"""Transfers artifacts through the coordinator without leaking credentials."""
|
||||
|
||||
def __init__(self, coordinator_url: str, timeout: float, bearer_token: str | None = None) -> None:
|
||||
def __init__(
|
||||
self,
|
||||
coordinator_url: str,
|
||||
timeout: float,
|
||||
bearer_token: str | None = None,
|
||||
*,
|
||||
token_provider: TokenProvider | None = None,
|
||||
) -> None:
|
||||
self.coordinator_url = coordinator_url.rstrip("/")
|
||||
self.timeout = timeout
|
||||
self.bearer_token = bearer_token
|
||||
self._tokens: TokenProvider = token_provider or StaticTokenProvider(bearer_token)
|
||||
self.coordinator_origin = origin(coordinator_url)
|
||||
self._opener = build_opener(SameOriginAuthRedirectHandler(self.coordinator_origin))
|
||||
|
||||
@property
|
||||
def bearer_token(self) -> str | None:
|
||||
return self._tokens.token()
|
||||
|
||||
def download(self, uri: str, destination: Path) -> None:
|
||||
destination.parent.mkdir(parents=True, exist_ok=True)
|
||||
resolved_uri = urljoin(f"{self.coordinator_url}/", uri)
|
||||
self._download_once(resolved_uri, destination, allow_refresh=True)
|
||||
|
||||
def _download_once(self, resolved_uri: str, destination: Path, *, allow_refresh: bool) -> None:
|
||||
request = Request(resolved_uri, headers=self._auth_headers_for(resolved_uri))
|
||||
with self._opener.open(request, timeout=self.timeout) as response, destination.open("wb") as target:
|
||||
while chunk := response.read(1024 * 1024):
|
||||
target.write(chunk)
|
||||
try:
|
||||
with self._opener.open(request, timeout=self.timeout) as response, destination.open("wb") as target:
|
||||
while chunk := response.read(1024 * 1024):
|
||||
target.write(chunk)
|
||||
except HTTPError as error:
|
||||
# Refresh an expired token and retry once, mirroring the coordinator
|
||||
# client, so a token that lapses mid-task does not fail the download.
|
||||
if error.code == 401 and allow_refresh:
|
||||
self._tokens.refresh()
|
||||
self._download_once(resolved_uri, destination, allow_refresh=False)
|
||||
return
|
||||
raise
|
||||
|
||||
def upload(
|
||||
self, task: ClaimedTask, worker_id: str, artifact: ProducedArtifact
|
||||
) -> UploadedArtifact:
|
||||
return self._upload_once(task, worker_id, artifact, allow_refresh=True)
|
||||
|
||||
def _upload_once(
|
||||
self, task: ClaimedTask, worker_id: str, artifact: ProducedArtifact, *, allow_refresh: bool
|
||||
) -> UploadedArtifact:
|
||||
"""Stream an artifact and require durable coordinator-owned metadata."""
|
||||
url = (
|
||||
@@ -76,6 +106,11 @@ class HttpArtifactClient:
|
||||
connection.send(chunk)
|
||||
response = connection.getresponse()
|
||||
body = response.read()
|
||||
if response.status == 401 and allow_refresh:
|
||||
# Token lapsed mid-task: refresh and retry the upload once.
|
||||
self._tokens.refresh()
|
||||
connection.close()
|
||||
return self._upload_once(task, worker_id, artifact, allow_refresh=False)
|
||||
if response.status == 409:
|
||||
raise CoordinatorConflictError("artifact upload rejected because the task lease was lost")
|
||||
if response.status != 200:
|
||||
@@ -93,8 +128,9 @@ class HttpArtifactClient:
|
||||
|
||||
def _auth_headers_for(self, uri: str) -> dict[str, str]:
|
||||
"""Only coordinator-owned URLs receive the coordinator bearer token."""
|
||||
if self.bearer_token and origin(uri) == self.coordinator_origin:
|
||||
return {"Authorization": f"Bearer {self.bearer_token}"}
|
||||
token = self._tokens.token()
|
||||
if token and origin(uri) == self.coordinator_origin:
|
||||
return {"Authorization": f"Bearer {token}"}
|
||||
return {}
|
||||
|
||||
def sha256_file(path: Path) -> str:
|
||||
|
||||
@@ -0,0 +1,128 @@
|
||||
"""Bearer-token strategies for the worker's coordinator calls.
|
||||
|
||||
A worker authenticates in one of two ways:
|
||||
|
||||
* a *static* token — the shared service token or a directly supplied JWT, fixed
|
||||
for the life of the process; or
|
||||
* a *worker key* — a long-lived per-user credential the worker trades for a
|
||||
short-lived JWT at the userservice, refreshing before that JWT expires.
|
||||
|
||||
Both are exposed through the small ``TokenProvider`` protocol so the HTTP
|
||||
clients neither know nor care which one is in play.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import time
|
||||
from typing import Callable, Protocol
|
||||
from urllib.error import HTTPError, URLError
|
||||
from urllib.request import Request, build_opener
|
||||
|
||||
from .transport import NoRedirectHandler
|
||||
|
||||
|
||||
class TokenExchangeError(RuntimeError):
|
||||
"""The userservice refused or failed to exchange a worker key."""
|
||||
|
||||
|
||||
class TokenProvider(Protocol):
|
||||
def token(self) -> str | None:
|
||||
"""Return the current bearer token, refreshing it if necessary."""
|
||||
|
||||
def refresh(self) -> None:
|
||||
"""Force the next token to be re-fetched (e.g. after a 401)."""
|
||||
|
||||
|
||||
class StaticTokenProvider:
|
||||
"""Serves a fixed token forever. ``None`` means "send no Authorization"."""
|
||||
|
||||
def __init__(self, token: str | None) -> None:
|
||||
self._token = token
|
||||
|
||||
def token(self) -> str | None:
|
||||
return self._token
|
||||
|
||||
def refresh(self) -> None: # noqa: D401 - nothing to refresh
|
||||
return None
|
||||
|
||||
|
||||
class WorkerKeyTokenProvider:
|
||||
"""Exchanges a long-lived worker key for short-lived JWTs and refreshes them.
|
||||
|
||||
The token is cached until roughly ``1 - refresh_leeway`` of its lifetime has
|
||||
elapsed, so the worker renews ahead of expiry instead of waiting for a 401.
|
||||
A monotonic clock is injectable to keep tests deterministic.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
userservice_url: str,
|
||||
worker_key: str,
|
||||
timeout: float,
|
||||
*,
|
||||
refresh_leeway: float = 0.2,
|
||||
now: Callable[[], float] = time.monotonic,
|
||||
) -> None:
|
||||
self._url = userservice_url.rstrip("/")
|
||||
self._key = worker_key
|
||||
self._timeout = timeout
|
||||
self._leeway = refresh_leeway
|
||||
self._now = now
|
||||
self._token: str | None = None
|
||||
self._refresh_at: float = 0.0
|
||||
self._opener = build_opener(NoRedirectHandler())
|
||||
|
||||
def token(self) -> str:
|
||||
if self._token is None or self._now() >= self._refresh_at:
|
||||
self._exchange()
|
||||
assert self._token is not None # _exchange sets it or raises
|
||||
return self._token
|
||||
|
||||
def refresh(self) -> None:
|
||||
self._exchange()
|
||||
|
||||
def _exchange(self) -> None:
|
||||
request = Request(
|
||||
f"{self._url}/worker-tokens/exchange",
|
||||
data=json.dumps({"key": self._key}).encode(),
|
||||
method="POST",
|
||||
headers={"Content-Type": "application/json"},
|
||||
)
|
||||
try:
|
||||
with self._opener.open(request, timeout=self._timeout) as response:
|
||||
raw = response.read()
|
||||
data = json.loads(raw) if raw else {}
|
||||
except HTTPError as error:
|
||||
# A revoked or unknown key is a permanent 401; there is nothing the
|
||||
# worker can do but stop, so surface it rather than retry forever.
|
||||
raise TokenExchangeError(
|
||||
f"worker key exchange rejected with status {error.code}"
|
||||
) from error
|
||||
except (URLError, TimeoutError, json.JSONDecodeError) as error:
|
||||
raise TokenExchangeError("worker key exchange request failed") from error
|
||||
|
||||
token = data.get("token")
|
||||
if not isinstance(token, str) or not token:
|
||||
raise TokenExchangeError("worker key exchange response is missing a token")
|
||||
|
||||
expires_in = data.get("expires_in")
|
||||
ttl = float(expires_in) if isinstance(expires_in, (int, float)) and expires_in > 0 else 0.0
|
||||
self._token = token
|
||||
# Renew once ~(1 - leeway) of the lifetime is gone. An unknown TTL falls
|
||||
# back to re-exchanging on the next call — correct, just chattier.
|
||||
self._refresh_at = self._now() + ttl * (1.0 - self._leeway)
|
||||
|
||||
|
||||
def provider_from_config(
|
||||
*,
|
||||
worker_key: str | None,
|
||||
userservice_url: str | None,
|
||||
bearer_token: str | None,
|
||||
request_timeout: float,
|
||||
) -> TokenProvider:
|
||||
"""Pick the token strategy: a worker key (exchange mode) wins over a static
|
||||
bearer token, which in turn wins over no credential at all."""
|
||||
if worker_key and userservice_url:
|
||||
return WorkerKeyTokenProvider(userservice_url, worker_key, request_timeout)
|
||||
return StaticTokenProvider(bearer_token)
|
||||
+26
-4
@@ -7,6 +7,7 @@ import logging
|
||||
from pathlib import Path
|
||||
|
||||
from .artifacts import HttpArtifactClient
|
||||
from .auth import provider_from_config
|
||||
from .config import WorkerConfig
|
||||
from .coordinator import HttpCoordinatorClient
|
||||
from .daemon import WorkerDaemon
|
||||
@@ -22,14 +23,23 @@ def build_parser() -> argparse.ArgumentParser:
|
||||
"SCIMESH_WORKER_NAME, SCIMESH_CPU_COUNT, SCIMESH_MEMORY_MB, "
|
||||
"SCIMESH_POLL_INTERVAL, SCIMESH_REQUEST_TIMEOUT, "
|
||||
"SCIMESH_HEARTBEAT_INTERVAL, SCIMESH_CLEANUP_AFTER_SECONDS, "
|
||||
"SCIMESH_MAX_TASKS, and SCIMESH_BEARER_TOKEN. "
|
||||
"SCIMESH_WORKER_ID is a legacy/test override."
|
||||
"SCIMESH_MAX_TASKS, SCIMESH_BEARER_TOKEN, SCIMESH_WORKER_KEY, and "
|
||||
"SCIMESH_USERSERVICE_URL. SCIMESH_WORKER_ID is a legacy/test override."
|
||||
),
|
||||
)
|
||||
parser.add_argument("--coordinator-url")
|
||||
parser.add_argument("--worker-id")
|
||||
parser.add_argument("--work-dir")
|
||||
parser.add_argument("--worker-name")
|
||||
parser.add_argument(
|
||||
"--worker-key",
|
||||
help="Long-lived worker key from the web UI; the worker exchanges it for "
|
||||
"short-lived tokens, binding it to your account. Requires --userservice-url.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--userservice-url",
|
||||
help="Base URL of the userservice that issues tokens for --worker-key",
|
||||
)
|
||||
parser.add_argument("--cpu-count", type=int)
|
||||
parser.add_argument("--memory-mb", type=int)
|
||||
parser.add_argument("--poll-interval", type=float)
|
||||
@@ -68,11 +78,23 @@ def main(argv: list[str] | None = None) -> int:
|
||||
except (TypeError, ValueError) as error:
|
||||
parser.error(str(error))
|
||||
logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s")
|
||||
client = HttpCoordinatorClient(config.coordinator_url, config.request_timeout, config.bearer_token)
|
||||
# One shared token strategy backs both clients: a worker key (exchanged and
|
||||
# refreshed) or a static bearer token, decided by what the config carries.
|
||||
tokens = provider_from_config(
|
||||
worker_key=config.worker_key,
|
||||
userservice_url=config.userservice_url,
|
||||
bearer_token=config.bearer_token,
|
||||
request_timeout=config.request_timeout,
|
||||
)
|
||||
client = HttpCoordinatorClient(
|
||||
config.coordinator_url, config.request_timeout, token_provider=tokens
|
||||
)
|
||||
completed_without_interruption = WorkerDaemon(
|
||||
config,
|
||||
client,
|
||||
HttpArtifactClient(config.coordinator_url, config.request_timeout, config.bearer_token),
|
||||
HttpArtifactClient(
|
||||
config.coordinator_url, config.request_timeout, token_provider=tokens
|
||||
),
|
||||
SciMeshRunner(),
|
||||
).run_forever()
|
||||
return 0 if completed_without_interruption else 130
|
||||
|
||||
@@ -11,6 +11,14 @@ from typing import Mapping
|
||||
from urllib.parse import urlsplit
|
||||
|
||||
|
||||
def _clean_url(value: object | None) -> str | None:
|
||||
"""Normalise an optional URL: drop a blank one, strip a trailing slash."""
|
||||
if value is None:
|
||||
return None
|
||||
text = str(value).strip()
|
||||
return text.rstrip("/") or None
|
||||
|
||||
|
||||
def _positive_number(value: object, name: str, *, allow_zero: bool = False) -> None:
|
||||
if (
|
||||
isinstance(value, bool)
|
||||
@@ -35,6 +43,11 @@ class WorkerConfig:
|
||||
request_timeout: float = 30.0
|
||||
heartbeat_interval: float = 15.0
|
||||
bearer_token: str | None = None
|
||||
# A long-lived per-user credential. When set (with userservice_url), the
|
||||
# worker exchanges it for short-lived JWTs instead of using bearer_token,
|
||||
# binding the worker to that user's account.
|
||||
worker_key: str | None = None
|
||||
userservice_url: str | None = None
|
||||
cleanup_after_seconds: float | None = None
|
||||
max_tasks: int | None = None
|
||||
exit_when_idle: bool = False
|
||||
@@ -54,6 +67,12 @@ class WorkerConfig:
|
||||
raise ValueError("coordinator_url must be an absolute HTTP(S) URL")
|
||||
if not isinstance(self.worker_name, str) or not self.worker_name.strip():
|
||||
raise ValueError("worker_name must be non-empty")
|
||||
if self.userservice_url is not None:
|
||||
us = urlsplit(self.userservice_url)
|
||||
if us.scheme not in {"http", "https"} or not us.hostname:
|
||||
raise ValueError("userservice_url must be an absolute HTTP(S) URL")
|
||||
if self.worker_key is not None and not self.userservice_url:
|
||||
raise ValueError("worker_key requires userservice_url (SCIMESH_USERSERVICE_URL)")
|
||||
if isinstance(self.cpu_count, bool) or not isinstance(self.cpu_count, int) or self.cpu_count < 1:
|
||||
raise ValueError("cpu_count must be positive")
|
||||
if self.worker_id is not None and not isinstance(self.worker_id, str):
|
||||
@@ -114,6 +133,8 @@ class WorkerConfig:
|
||||
request_timeout=float(value("request_timeout", "SCIMESH_REQUEST_TIMEOUT", "30")),
|
||||
heartbeat_interval=float(value("heartbeat_interval", "SCIMESH_HEARTBEAT_INTERVAL", "15")),
|
||||
bearer_token=value("bearer_token", "SCIMESH_BEARER_TOKEN"),
|
||||
worker_key=value("worker_key", "SCIMESH_WORKER_KEY"),
|
||||
userservice_url=_clean_url(value("userservice_url", "SCIMESH_USERSERVICE_URL")),
|
||||
cleanup_after_seconds=float(cleanup) if cleanup else None,
|
||||
max_tasks=int(max_tasks) if max_tasks is not None else None,
|
||||
exit_when_idle=bool(values.get("exit_when_idle", False)),
|
||||
|
||||
@@ -7,6 +7,7 @@ from typing import Any, Protocol
|
||||
from urllib.error import HTTPError, URLError
|
||||
from urllib.request import Request, build_opener
|
||||
|
||||
from .auth import StaticTokenProvider, TokenProvider
|
||||
from .models import ClaimedTask, RegisteredWorker
|
||||
from .transport import NoRedirectHandler
|
||||
|
||||
@@ -38,12 +39,25 @@ class CoordinatorClient(Protocol):
|
||||
|
||||
|
||||
class HttpCoordinatorClient:
|
||||
def __init__(self, base_url: str, timeout: float, bearer_token: str | None = None) -> None:
|
||||
def __init__(
|
||||
self,
|
||||
base_url: str,
|
||||
timeout: float,
|
||||
bearer_token: str | None = None,
|
||||
*,
|
||||
token_provider: TokenProvider | None = None,
|
||||
) -> None:
|
||||
self.base_url = base_url.rstrip("/")
|
||||
self.timeout = timeout
|
||||
self.bearer_token = bearer_token
|
||||
# A bearer_token argument keeps older call sites working; internally
|
||||
# everything goes through a provider so refresh is uniform.
|
||||
self._tokens: TokenProvider = token_provider or StaticTokenProvider(bearer_token)
|
||||
self._opener = build_opener(NoRedirectHandler())
|
||||
|
||||
@property
|
||||
def bearer_token(self) -> str | None:
|
||||
return self._tokens.token()
|
||||
|
||||
def register(
|
||||
self, name: str, capabilities: tuple[str, ...], cpu_count: int, memory_mb: int | None
|
||||
) -> RegisteredWorker:
|
||||
@@ -102,6 +116,11 @@ class HttpCoordinatorClient:
|
||||
return lease_expires_at
|
||||
|
||||
def _request(self, method: str, path: str, payload: dict[str, Any]) -> tuple[int, dict[str, Any]]:
|
||||
return self._request_once(method, path, payload, allow_refresh=True)
|
||||
|
||||
def _request_once(
|
||||
self, method: str, path: str, payload: dict[str, Any], *, allow_refresh: bool
|
||||
) -> tuple[int, dict[str, Any]]:
|
||||
request = Request(
|
||||
f"{self.base_url}{path}", data=json.dumps(payload).encode(), method=method,
|
||||
headers={"Content-Type": "application/json", **self._auth_header()},
|
||||
@@ -114,6 +133,11 @@ class HttpCoordinatorClient:
|
||||
except json.JSONDecodeError as error:
|
||||
raise CoordinatorError("coordinator returned invalid JSON") from error
|
||||
except HTTPError as error:
|
||||
# A 401 usually means the short-lived JWT expired; mint a fresh one
|
||||
# and retry exactly once so an in-flight worker rides over the gap.
|
||||
if error.code == 401 and allow_refresh:
|
||||
self._tokens.refresh()
|
||||
return self._request_once(method, path, payload, allow_refresh=False)
|
||||
if error.code >= 500:
|
||||
raise CoordinatorTransientError(f"coordinator returned {error.code}") from error
|
||||
return error.code, {}
|
||||
@@ -121,4 +145,5 @@ class HttpCoordinatorClient:
|
||||
raise CoordinatorTransientError("coordinator request failed") from error
|
||||
|
||||
def _auth_header(self) -> dict[str, str]:
|
||||
return {"Authorization": f"Bearer {self.bearer_token}"} if self.bearer_token else {}
|
||||
token = self._tokens.token()
|
||||
return {"Authorization": f"Bearer {token}"} if token else {}
|
||||
|
||||
Reference in New Issue
Block a user