auth.py 7.2 KB

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