| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168 |
- """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"]),
- )
|