"""Short-lived signed task tokens bound to one immutable workflow node.""" from __future__ import annotations import hashlib import json import time import uuid from dataclasses import dataclass from typing import Any, Mapping import jwt ISSUER = "dataops-platform" AUDIENCE = "dataops-runner" ALGORITHM = "HS256" class TaskTokenInvalid(ValueError): pass class TaskTokenExpired(TaskTokenInvalid): pass def node_digest(node: Mapping[str, Any]) -> str: canonical = json.dumps( node, sort_keys=True, separators=(",", ":"), ensure_ascii=False, ) return hashlib.sha256(canonical.encode("utf-8")).hexdigest() def _secret_bytes(secret: str) -> bytes: encoded = str(secret or "").encode("utf-8") if len(encoded) < 32: raise ValueError("task token secret must be at least 32 bytes") return encoded @dataclass(frozen=True) class TaskClaims: task_uid: str dataflow_uid: str workflow_version: int correlation_id: str node_id: str node_type: str node_digest: str purpose: str write_authorized: bool issued_at: int expires_at: int jti: str class TaskTokenIssuer: def __init__(self, secret, *, clock=None, ttl_seconds=60): self._secret = _secret_bytes(secret) self._clock = clock or time.time self._ttl_seconds = int(ttl_seconds) if self._ttl_seconds < 1 or self._ttl_seconds > 300: raise ValueError("task token ttl must be between 1 and 300 seconds") def issue( self, *, task_uid, dataflow_uid, workflow_version, correlation_id, node, write_authorized=False, ): now = int(self._clock()) purpose = str(node.get("purpose") or "read") if purpose == "write" and not write_authorized: raise ValueError("write task requires trusted write authorization") payload = { "iss": ISSUER, "aud": AUDIENCE, "sub": str(task_uid), "task_uid": str(task_uid), "dataflow_uid": str(dataflow_uid), "workflow_version": int(workflow_version), "correlation_id": str(correlation_id), "node_id": str(node.get("id") or ""), "node_type": str(node.get("type") or ""), "node_digest": node_digest(node), "purpose": purpose, "write_authorized": bool(write_authorized), "iat": now, "exp": now + self._ttl_seconds, "jti": str(uuid.uuid4()), } return jwt.encode(payload, self._secret, algorithm=ALGORITHM) class TaskTokenVerifier: def __init__(self, secret, *, clock=None): self._secret = _secret_bytes(secret) self._clock = clock or time.time def verify(self, token, *, node): try: payload = jwt.decode( token, self._secret, algorithms=[ALGORITHM], audience=AUDIENCE, issuer=ISSUER, options={"verify_exp": False, "verify_iat": False}, ) except jwt.PyJWTError as exc: raise TaskTokenInvalid("task token is invalid") from exc required = { "task_uid", "dataflow_uid", "workflow_version", "correlation_id", "node_id", "node_type", "node_digest", "purpose", "iat", "exp", "jti", } if required - set(payload): raise TaskTokenInvalid("task token claims are incomplete") now = int(self._clock()) try: issued_at = int(payload["iat"]) expires_at = int(payload["exp"]) workflow_version = int(payload["workflow_version"]) except (TypeError, ValueError) as exc: raise TaskTokenInvalid("task token claims are invalid") from exc if expires_at < now: raise TaskTokenExpired("task token has expired") if issued_at > now + 5 or expires_at - issued_at > 300: raise TaskTokenInvalid("task token lifetime is invalid") if node_digest(node) != payload["node_digest"]: raise TaskTokenInvalid("task token node binding does not match") if ( payload["node_id"] != node.get("id") or payload["node_type"] != node.get("type") or payload["purpose"] != str(node.get("purpose") or "read") ): raise TaskTokenInvalid("task token node binding does not match") return TaskClaims( task_uid=str(payload["task_uid"]), dataflow_uid=str(payload["dataflow_uid"]), workflow_version=workflow_version, correlation_id=str(payload["correlation_id"]), node_id=str(payload["node_id"]), node_type=str(payload["node_type"]), node_digest=str(payload["node_digest"]), purpose=str(payload["purpose"]), write_authorized=bool(payload.get("write_authorized", False)), issued_at=issued_at, expires_at=expires_at, jti=str(payload["jti"]), )