auth.py 5.7 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182
  1. """Short-lived signed task tokens bound to one immutable workflow node."""
  2. from __future__ import annotations
  3. import hashlib
  4. import base64
  5. import json
  6. import time
  7. import uuid
  8. from dataclasses import dataclass
  9. from typing import Any, Mapping
  10. import jwt
  11. ISSUER = "dataops-platform"
  12. AUDIENCE = "dataops-runner"
  13. ALGORITHM = "HS256"
  14. class TaskTokenInvalid(ValueError):
  15. pass
  16. class TaskTokenExpired(TaskTokenInvalid):
  17. pass
  18. def node_digest(node: Mapping[str, Any]) -> str:
  19. canonical = json.dumps(
  20. node,
  21. sort_keys=True,
  22. separators=(",", ":"),
  23. ensure_ascii=False,
  24. )
  25. return hashlib.sha256(canonical.encode("utf-8")).hexdigest()
  26. def _secret_bytes(secret: str) -> bytes:
  27. encoded = str(secret or "").encode("utf-8")
  28. if len(encoded) < 32:
  29. raise ValueError("task token secret must be at least 32 bytes")
  30. return encoded
  31. @dataclass(frozen=True)
  32. class TaskClaims:
  33. task_uid: str
  34. dataflow_uid: str
  35. workflow_version: int
  36. correlation_id: str
  37. node_id: str
  38. node_type: str
  39. node_digest: str
  40. purpose: str
  41. write_authorized: bool
  42. issued_at: int
  43. expires_at: int
  44. jti: str
  45. class TaskTokenIssuer:
  46. def __init__(self, secret, *, clock=None, ttl_seconds=60):
  47. self._secret = _secret_bytes(secret)
  48. self._clock = clock or time.time
  49. self._ttl_seconds = int(ttl_seconds)
  50. if self._ttl_seconds < 1 or self._ttl_seconds > 300:
  51. raise ValueError("task token ttl must be between 1 and 300 seconds")
  52. def issue(
  53. self,
  54. *,
  55. task_uid,
  56. dataflow_uid,
  57. workflow_version,
  58. correlation_id,
  59. node,
  60. write_authorized=False,
  61. ):
  62. now = int(self._clock())
  63. purpose = str(node.get("purpose") or "read")
  64. if purpose == "write" and not write_authorized:
  65. raise ValueError("write task requires trusted write authorization")
  66. payload = {
  67. "iss": ISSUER,
  68. "aud": AUDIENCE,
  69. "sub": str(task_uid),
  70. "task_uid": str(task_uid),
  71. "dataflow_uid": str(dataflow_uid),
  72. "workflow_version": int(workflow_version),
  73. "correlation_id": str(correlation_id),
  74. "node_id": str(node.get("id") or ""),
  75. "node_type": str(node.get("type") or ""),
  76. "node_digest": node_digest(node),
  77. "purpose": purpose,
  78. "write_authorized": bool(write_authorized),
  79. "iat": now,
  80. "exp": now + self._ttl_seconds,
  81. "jti": str(uuid.uuid4()),
  82. }
  83. return jwt.encode(payload, self._secret, algorithm=ALGORITHM)
  84. class TaskTokenVerifier:
  85. def __init__(self, secret, *, clock=None):
  86. self._secret = _secret_bytes(secret)
  87. self._clock = clock or time.time
  88. def verify(self, token, *, node):
  89. try:
  90. segments = str(token).split(".")
  91. if len(segments) != 3:
  92. raise ValueError("segment count")
  93. for segment in segments:
  94. decoded = base64.urlsafe_b64decode(
  95. segment + "=" * (-len(segment) % 4)
  96. )
  97. canonical = base64.urlsafe_b64encode(decoded).rstrip(b"=").decode("ascii")
  98. if canonical != segment:
  99. raise ValueError("non-canonical base64url")
  100. except (TypeError, ValueError, UnicodeError) as exc:
  101. raise TaskTokenInvalid("task token is invalid") from exc
  102. try:
  103. payload = jwt.decode(
  104. token,
  105. self._secret,
  106. algorithms=[ALGORITHM],
  107. audience=AUDIENCE,
  108. issuer=ISSUER,
  109. options={"verify_exp": False, "verify_iat": False},
  110. )
  111. except jwt.PyJWTError as exc:
  112. raise TaskTokenInvalid("task token is invalid") from exc
  113. required = {
  114. "task_uid",
  115. "dataflow_uid",
  116. "workflow_version",
  117. "correlation_id",
  118. "node_id",
  119. "node_type",
  120. "node_digest",
  121. "purpose",
  122. "iat",
  123. "exp",
  124. "jti",
  125. }
  126. if required - set(payload):
  127. raise TaskTokenInvalid("task token claims are incomplete")
  128. now = int(self._clock())
  129. try:
  130. issued_at = int(payload["iat"])
  131. expires_at = int(payload["exp"])
  132. workflow_version = int(payload["workflow_version"])
  133. except (TypeError, ValueError) as exc:
  134. raise TaskTokenInvalid("task token claims are invalid") from exc
  135. if expires_at < now:
  136. raise TaskTokenExpired("task token has expired")
  137. if issued_at > now + 5 or expires_at - issued_at > 300:
  138. raise TaskTokenInvalid("task token lifetime is invalid")
  139. if node_digest(node) != payload["node_digest"]:
  140. raise TaskTokenInvalid("task token node binding does not match")
  141. if (
  142. payload["node_id"] != node.get("id")
  143. or payload["node_type"] != node.get("type")
  144. or payload["purpose"] != str(node.get("purpose") or "read")
  145. ):
  146. raise TaskTokenInvalid("task token node binding does not match")
  147. return TaskClaims(
  148. task_uid=str(payload["task_uid"]),
  149. dataflow_uid=str(payload["dataflow_uid"]),
  150. workflow_version=workflow_version,
  151. correlation_id=str(payload["correlation_id"]),
  152. node_id=str(payload["node_id"]),
  153. node_type=str(payload["node_type"]),
  154. node_digest=str(payload["node_digest"]),
  155. purpose=str(payload["purpose"]),
  156. write_authorized=bool(payload.get("write_authorized", False)),
  157. issued_at=issued_at,
  158. expires_at=expires_at,
  159. jti=str(payload["jti"]),
  160. )