auth.py 5.1 KB

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