1"""
2URL-safe signed JSON objects using HMAC/SHA-256.
3
4Use TimestampSigner for signing with expiration:
5
6 TimestampSigner(salt="my-salt").sign_object({"key": "value"})
7 TimestampSigner(salt="my-salt").unsign_object(token, max_age=3600)
8
9Use Signer for signing without expiration:
10
11 Signer(salt="my-salt").sign_object({"key": "value"})
12 Signer(salt="my-salt").unsign_object(token)
13"""
14
15from __future__ import annotations
16
17import base64
18import datetime
19import hmac
20import json
21import time
22import zlib
23from typing import Any
24
25from plain.runtime import settings
26from plain.utils.crypto import salted_hmac
27from plain.utils.encoding import force_bytes
28from plain.utils.regex_helper import _lazy_re_compile
29
30_SEP_UNSAFE = _lazy_re_compile(r"^[A-z0-9-_=]*$")
31BASE62_ALPHABET = "0123456789ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz"
32
33
34class BadSignature(Exception):
35 """Signature does not match."""
36
37
38class SignatureExpired(BadSignature):
39 """Signature timestamp is older than required max_age."""
40
41
42def b62_encode(s: int) -> str:
43 if s == 0:
44 return "0"
45 sign = "-" if s < 0 else ""
46 s = abs(s)
47 encoded = ""
48 while s > 0:
49 s, remainder = divmod(s, 62)
50 encoded = BASE62_ALPHABET[remainder] + encoded
51 return sign + encoded
52
53
54def b62_decode(s: str) -> int:
55 if s == "0":
56 return 0
57 sign = 1
58 if s[0] == "-":
59 s = s[1:]
60 sign = -1
61 decoded = 0
62 for digit in s:
63 decoded = decoded * 62 + BASE62_ALPHABET.index(digit)
64 return sign * decoded
65
66
67def b64_encode(s: bytes) -> bytes:
68 return base64.urlsafe_b64encode(s).strip(b"=")
69
70
71def b64_decode(s: bytes) -> bytes:
72 pad = b"=" * (-len(s) % 4)
73 return base64.urlsafe_b64decode(s + pad)
74
75
76def base64_hmac(salt: str, value: str, key: str, algorithm: str = "sha1") -> str:
77 return b64_encode(
78 salted_hmac(salt, value, key, algorithm=algorithm).digest()
79 ).decode()
80
81
82class JSONSerializer:
83 """
84 Simple wrapper around json used by Signer.sign_object and
85 Signer.unsign_object.
86 """
87
88 def dumps(self, obj: Any) -> bytes:
89 return json.dumps(obj, separators=(",", ":")).encode("latin-1")
90
91 def loads(self, data: bytes) -> Any:
92 return json.loads(data.decode("latin-1"))
93
94
95class Signer:
96 def __init__(
97 self,
98 *,
99 key: str | None = None,
100 sep: str = ":",
101 salt: str | None = None,
102 algorithm: str = "sha256",
103 fallback_keys: list[str] | None = None,
104 ) -> None:
105 self.key = key or settings.SECRET_KEY
106 self.fallback_keys = (
107 fallback_keys
108 if fallback_keys is not None
109 else settings.SECRET_KEY_FALLBACKS
110 )
111 self.sep = sep
112 self.salt = salt or f"{self.__class__.__module__}.{self.__class__.__name__}"
113 self.algorithm = algorithm
114
115 if _SEP_UNSAFE.match(self.sep):
116 raise ValueError(
117 f"Unsafe Signer separator: {sep!r} (cannot be empty or consist of "
118 "only A-z0-9-_=)",
119 )
120
121 def signature(self, value: str, key: str | None = None) -> str:
122 key = key or self.key
123 return base64_hmac(self.salt + "signer", value, key, algorithm=self.algorithm)
124
125 def sign(self, value: str) -> str:
126 return f"{value}{self.sep}{self.signature(value)}"
127
128 def unsign(self, signed_value: str) -> str:
129 if self.sep not in signed_value:
130 raise BadSignature(f'No "{self.sep}" found in value')
131 value, sig = signed_value.rsplit(self.sep, 1)
132 for key in [self.key, *self.fallback_keys]:
133 if hmac.compare_digest(
134 force_bytes(sig), force_bytes(self.signature(value, key))
135 ):
136 return value
137 raise BadSignature(f'Signature "{sig}" does not match')
138
139 def sign_object(
140 self,
141 obj: Any,
142 serializer: type[JSONSerializer] = JSONSerializer,
143 compress: bool = False,
144 ) -> str:
145 """
146 Return URL-safe, hmac signed base64 compressed JSON string.
147
148 If compress is True (not the default), check if compressing using zlib
149 can save some space. Prepend a '.' to signify compression. This is
150 included in the signature, to protect against zip bombs.
151
152 The serializer is expected to return a bytestring.
153 """
154 data = serializer().dumps(obj)
155 # Flag for if it's been compressed or not.
156 is_compressed = False
157
158 if compress:
159 # Avoid zlib dependency unless compress is being used.
160 compressed = zlib.compress(data)
161 if len(compressed) < (len(data) - 1):
162 data = compressed
163 is_compressed = True
164 base64d = b64_encode(data).decode()
165 if is_compressed:
166 base64d = "." + base64d
167 return self.sign(base64d)
168
169 def unsign_object(
170 self,
171 signed_obj: str,
172 serializer: type[JSONSerializer] = JSONSerializer,
173 **kwargs: Any,
174 ) -> Any:
175 # Signer.unsign() returns str but base64 and zlib compression operate
176 # on bytes.
177 base64d = self.unsign(signed_obj, **kwargs).encode()
178 decompress = base64d[:1] == b"."
179 if decompress:
180 # It's compressed; uncompress it first.
181 base64d = base64d[1:]
182 data = b64_decode(base64d)
183 if decompress:
184 data = zlib.decompress(data)
185 return serializer().loads(data)
186
187
188class TimestampSigner:
189 """A signer that includes a timestamp for max_age validation.
190
191 Uses composition rather than inheritance since the interface
192 intentionally differs from Signer (unsign accepts max_age parameter).
193 """
194
195 def __init__(
196 self,
197 *,
198 key: str | None = None,
199 sep: str = ":",
200 salt: str | None = None,
201 algorithm: str = "sha256",
202 fallback_keys: list[str] | None = None,
203 ) -> None:
204 # Compute default salt here to preserve backwards compatibility.
205 # When TimestampSigner inherited from Signer, the default salt was
206 # "plain.signing.TimestampSigner". Now that we use composition,
207 # we must set it explicitly rather than letting Signer compute its own.
208 if salt is None:
209 salt = f"{self.__class__.__module__}.{self.__class__.__name__}"
210 self._signer = Signer(
211 key=key,
212 sep=sep,
213 salt=salt,
214 algorithm=algorithm,
215 fallback_keys=fallback_keys,
216 )
217
218 @property
219 def sep(self) -> str:
220 return self._signer.sep
221
222 def timestamp(self) -> str:
223 return b62_encode(int(time.time()))
224
225 def sign(self, value: str) -> str:
226 value = f"{value}{self.sep}{self.timestamp()}"
227 return self._signer.sign(value)
228
229 def unsign(
230 self, value: str, max_age: float | datetime.timedelta | None = None
231 ) -> str:
232 """
233 Retrieve original value and check it wasn't signed more
234 than max_age seconds ago.
235 """
236 result = self._signer.unsign(value)
237 value, timestamp = result.rsplit(self.sep, 1)
238 ts = b62_decode(timestamp)
239 if max_age is not None:
240 if isinstance(max_age, datetime.timedelta):
241 max_age = max_age.total_seconds()
242 # Check timestamp is not older than max_age
243 age = time.time() - ts
244 if age > max_age:
245 raise SignatureExpired(f"Signature age {age} > {max_age} seconds")
246 return value
247
248 def sign_object(
249 self,
250 obj: Any,
251 serializer: type[JSONSerializer] = JSONSerializer,
252 compress: bool = False,
253 ) -> str:
254 """
255 Return URL-safe, hmac signed base64 compressed JSON string.
256
257 If compress is True (not the default), check if compressing using zlib
258 can save some space. Prepend a '.' to signify compression. This is
259 included in the signature, to protect against zip bombs.
260
261 The serializer is expected to return a bytestring.
262 """
263 data = serializer().dumps(obj)
264 is_compressed = False
265
266 if compress:
267 compressed = zlib.compress(data)
268 if len(compressed) < (len(data) - 1):
269 data = compressed
270 is_compressed = True
271 base64d = b64_encode(data).decode()
272 if is_compressed:
273 base64d = "." + base64d
274 return self.sign(base64d)
275
276 def unsign_object(
277 self,
278 signed_obj: str,
279 serializer: type[JSONSerializer] = JSONSerializer,
280 max_age: float | datetime.timedelta | None = None,
281 ) -> Any:
282 """Unsign and decode an object, optionally checking max_age."""
283 base64d = self.unsign(signed_obj, max_age=max_age).encode()
284 decompress = base64d[:1] == b"."
285 if decompress:
286 base64d = base64d[1:]
287 data = b64_decode(base64d)
288 if decompress:
289 data = zlib.decompress(data)
290 return serializer().loads(data)