auth.py 5.8 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186
  1. """Short-lived signed task tokens bound to one immutable workflow node."""
  2. from __future__ import annotations
  3. import hashlib
  4. import json
  5. import time
  6. import uuid
  7. from collections.abc import Mapping
  8. from dataclasses import dataclass
  9. from typing import Any
  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. deployment_id: str
  36. environment: str
  37. workflow_version: int
  38. correlation_id: str
  39. node_id: str
  40. node_type: str
  41. node_digest: str
  42. purpose: str
  43. write_authorized: bool
  44. issued_at: int
  45. expires_at: int
  46. jti: str
  47. class TaskTokenIssuer:
  48. def __init__(self, secret, *, clock=None, ttl_seconds=60):
  49. self._secret = _secret_bytes(secret)
  50. self._clock = clock or time.time
  51. self._ttl_seconds = int(ttl_seconds)
  52. if self._ttl_seconds < 1 or self._ttl_seconds > 300:
  53. raise ValueError("task token ttl must be between 1 and 300 seconds")
  54. def issue(
  55. self,
  56. *,
  57. task_uid,
  58. dataflow_uid,
  59. deployment_id,
  60. environment,
  61. workflow_version,
  62. correlation_id,
  63. node,
  64. write_authorized=False,
  65. ):
  66. now = int(self._clock())
  67. if environment not in {"development", "test", "production"}:
  68. raise ValueError("task environment is invalid")
  69. purpose = str(node.get("purpose") or "read")
  70. if purpose == "write" and not write_authorized:
  71. raise ValueError("write task requires trusted write authorization")
  72. payload = {
  73. "iss": ISSUER,
  74. "aud": AUDIENCE,
  75. "sub": str(task_uid),
  76. "task_uid": str(task_uid),
  77. "dataflow_uid": str(dataflow_uid),
  78. "deployment_id": str(deployment_id),
  79. "environment": str(environment),
  80. "workflow_version": int(workflow_version),
  81. "correlation_id": str(correlation_id),
  82. "node_id": str(node.get("id") or ""),
  83. "node_type": str(node.get("type") or ""),
  84. "node_digest": node_digest(node),
  85. "purpose": purpose,
  86. "write_authorized": bool(write_authorized),
  87. "iat": now,
  88. "exp": now + self._ttl_seconds,
  89. "jti": str(uuid.uuid4()),
  90. }
  91. return jwt.encode(payload, self._secret, algorithm=ALGORITHM)
  92. class TaskTokenVerifier:
  93. def __init__(self, secret, *, clock=None):
  94. self._secret = _secret_bytes(secret)
  95. self._clock = clock or time.time
  96. def verify(self, token, *, node):
  97. try:
  98. payload = jwt.decode(
  99. token,
  100. self._secret,
  101. algorithms=[ALGORITHM],
  102. audience=AUDIENCE,
  103. issuer=ISSUER,
  104. options={"verify_exp": False, "verify_iat": False},
  105. )
  106. except jwt.PyJWTError as exc:
  107. raise TaskTokenInvalid("task token is invalid") from exc
  108. required = {
  109. "task_uid",
  110. "dataflow_uid",
  111. "deployment_id",
  112. "environment",
  113. "workflow_version",
  114. "correlation_id",
  115. "node_id",
  116. "node_type",
  117. "node_digest",
  118. "purpose",
  119. "iat",
  120. "exp",
  121. "jti",
  122. }
  123. if required - set(payload):
  124. raise TaskTokenInvalid("task token claims are incomplete")
  125. now = int(self._clock())
  126. try:
  127. issued_at = int(payload["iat"])
  128. expires_at = int(payload["exp"])
  129. workflow_version = int(payload["workflow_version"])
  130. except (TypeError, ValueError) as exc:
  131. raise TaskTokenInvalid("task token claims are invalid") from exc
  132. if expires_at < now:
  133. raise TaskTokenExpired("task token has expired")
  134. if issued_at > now + 5 or expires_at - issued_at > 300:
  135. raise TaskTokenInvalid("task token lifetime is invalid")
  136. if node_digest(node) != payload["node_digest"]:
  137. raise TaskTokenInvalid("task token node binding does not match")
  138. if (
  139. payload["node_id"] != node.get("id")
  140. or payload["node_type"] != node.get("type")
  141. or payload["purpose"] != str(node.get("purpose") or "read")
  142. ):
  143. raise TaskTokenInvalid("task token node binding does not match")
  144. if payload["environment"] not in {
  145. "development",
  146. "test",
  147. "production",
  148. }:
  149. raise TaskTokenInvalid("task token environment is invalid")
  150. return TaskClaims(
  151. task_uid=str(payload["task_uid"]),
  152. dataflow_uid=str(payload["dataflow_uid"]),
  153. deployment_id=str(payload["deployment_id"]),
  154. environment=str(payload["environment"]),
  155. workflow_version=workflow_version,
  156. correlation_id=str(payload["correlation_id"]),
  157. node_id=str(payload["node_id"]),
  158. node_type=str(payload["node_type"]),
  159. node_digest=str(payload["node_digest"]),
  160. purpose=str(payload["purpose"]),
  161. write_authorized=bool(payload.get("write_authorized", False)),
  162. issued_at=issued_at,
  163. expires_at=expires_at,
  164. jti=str(payload["jti"]),
  165. )