devora-python 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
devora_sdk/guard.py ADDED
@@ -0,0 +1,485 @@
1
+ from __future__ import annotations
2
+
3
+ import inspect
4
+ import math
5
+ import threading
6
+ import time
7
+ from collections import OrderedDict
8
+ from dataclasses import dataclass, field
9
+ from typing import Any, Callable, Mapping, Optional, Sequence
10
+ from urllib.parse import unquote
11
+
12
+ from .policy import ScopeConfig
13
+ from .utils import is_ambiguous_request_path, match_endpoint_pattern
14
+
15
+
16
+ @dataclass(frozen=True)
17
+ class ScopeEndpoint:
18
+ method: str
19
+ pattern: str
20
+
21
+
22
+ @dataclass(frozen=True)
23
+ class ImpersonationContext:
24
+ is_impersonation: bool
25
+ scope: str
26
+ session_id: Optional[str] = None
27
+ expires_at: Optional[int] = None
28
+ impersonator: Optional[dict[str, str]] = None
29
+ actor: Optional[dict[str, str]] = None
30
+ subject: Optional[dict[str, str]] = None
31
+ auth_method: Optional[str] = None
32
+ authorization_source: Optional[str] = None
33
+ recording_allowed: Optional[bool] = None
34
+
35
+
36
+ @dataclass(frozen=True)
37
+ class GuardDecision:
38
+ allowed: bool
39
+ status_code: int = 200
40
+ body: Optional[dict[str, object]] = None
41
+
42
+
43
+ @dataclass
44
+ class _LivenessLookup:
45
+ event: threading.Event = field(default_factory=threading.Event)
46
+ live: Optional[bool] = None
47
+ checked_at_ms: float = 0
48
+
49
+
50
+ class SessionLivenessChecker:
51
+ """Caches server-side session liveness results for the impersonation guard.
52
+
53
+ ``is_live`` returns True (live), False (ended), or None (could not determine — Devora
54
+ unreachable). When unreachable, behavior follows ``on_unavailable``: "allow" (fail open,
55
+ returns None) or "deny" (fail closed, returns False).
56
+ """
57
+
58
+ def __init__(self, sdk: Any, cache_ttl_ms: int = 5_000, on_unavailable: str = "deny") -> None:
59
+ if on_unavailable not in ("allow", "deny"):
60
+ raise ValueError("Invalid liveness unavailable policy")
61
+ if not isinstance(cache_ttl_ms, (int, float)) or not math.isfinite(cache_ttl_ms) or cache_ttl_ms < 0:
62
+ raise ValueError("Invalid liveness cache TTL")
63
+ self._sdk = sdk
64
+ self._cache_ttl_ms = cache_ttl_ms
65
+ self._unavailable_ttl_ms = min(1_000, cache_ttl_ms)
66
+ self._on_unavailable = on_unavailable
67
+ # value: (live | None for "unavailable", checked_at_ms)
68
+ self._cache: OrderedDict[str, tuple[Optional[bool], float]] = OrderedDict()
69
+ self._pending: dict[str, _LivenessLookup] = {}
70
+ self._lock = threading.Lock()
71
+
72
+ @property
73
+ def fails_closed(self) -> bool:
74
+ return self._on_unavailable == "deny"
75
+
76
+ def check(self, session_id: str) -> tuple[Optional[bool], bool]:
77
+ """Return ``(live, unavailable)``.
78
+
79
+ ``live`` is True/False when Devora answered and None when it could not be
80
+ reached; ``unavailable`` is True only when the verdict is unknown *and* the
81
+ checker is configured to fail closed. Concurrent requests for one session
82
+ share a single lookup, and an unavailable verdict is cached for one second
83
+ so an outage costs one lookup per session per second, not per request.
84
+ """
85
+ if not session_id or len(session_id) > 128:
86
+ return None, self.fails_closed
87
+ now = time.monotonic() * 1000
88
+ with self._lock:
89
+ cached = self._cache.get(session_id)
90
+ if cached:
91
+ ttl = self._unavailable_ttl_ms if cached[0] is None else self._cache_ttl_ms
92
+ if now - cached[1] < ttl:
93
+ self._cache.move_to_end(session_id)
94
+ return cached[0], cached[0] is None and self.fails_closed
95
+ del self._cache[session_id]
96
+ pending = self._pending.get(session_id)
97
+ if pending is None:
98
+ # Reject excess distinct lookups without allocating another pending
99
+ # key or starting network I/O. Same-session callers still coalesce.
100
+ if len(self._pending) >= 64:
101
+ return None, self.fails_closed
102
+ pending = _LivenessLookup()
103
+ self._pending[session_id] = pending
104
+ owner = True
105
+ else:
106
+ owner = False
107
+ if not owner:
108
+ if not pending.event.wait(timeout=6):
109
+ return None, self.fails_closed
110
+ # Read this lookup's result, not a potentially evicted/replaced cache
111
+ # entry. A delayed waiter must not resurrect an expired live verdict.
112
+ live = pending.live
113
+ ttl = self._unavailable_ttl_ms if live is None else self._cache_ttl_ms
114
+ if time.monotonic() * 1000 - pending.checked_at_ms >= ttl:
115
+ live = None
116
+ return live, live is None and self.fails_closed
117
+ live: Optional[bool] = None
118
+ try:
119
+ status = self._sdk.get_session_status(session_id)
120
+ live = None if status is None else status.get("valid") is True
121
+ except Exception:
122
+ # A custom transport may raise instead of returning unavailable.
123
+ live = None
124
+ finally:
125
+ with self._lock:
126
+ pending.live = live
127
+ pending.checked_at_ms = time.monotonic() * 1000
128
+ self._cache[session_id] = (live, pending.checked_at_ms)
129
+ self._cache.move_to_end(session_id)
130
+ # Hard LRU bound even when every entry is fresh. No dictionary
131
+ # copying or full-cache scans while holding the shared lock.
132
+ if len(self._cache) > 1024:
133
+ self._cache.popitem(last=False)
134
+ self._pending.pop(session_id, None)
135
+ pending.event.set()
136
+ return live, live is None and self.fails_closed
137
+
138
+ def is_live(self, session_id: str) -> Optional[bool]:
139
+ """Backwards-compatible verdict: False when unavailable and failing closed."""
140
+ live, unavailable = self.check(session_id)
141
+ if unavailable:
142
+ return False
143
+ return live
144
+
145
+
146
+ READ_METHODS = ("GET", "HEAD", "OPTIONS")
147
+
148
+ #: Headers through which common middleware lets a client change the effective method.
149
+ METHOD_OVERRIDE_HEADERS = ("x-http-method-override", "x-http-method", "x-method-override")
150
+
151
+ _AUTHORIZATION_SOURCES = ("standard", "self_approved", "self_approved_read", "break_glass")
152
+
153
+
154
+ class InvalidImpersonationContext(TypeError):
155
+ """The context extractor returned something that is not a recognised context."""
156
+
157
+
158
+ METHOD_OVERRIDE_QUERY_PARAM = "_method"
159
+
160
+
161
+ def policy_methods(
162
+ method: str, headers: Optional[Mapping[str, Any]] = None, query: Optional[str] = None
163
+ ) -> list[str]:
164
+ """Every method a request might execute as, including method-override headers
165
+ and ``_method`` query parameters (``query`` is the raw query string).
166
+
167
+ The guard judges all of them, so override middleware ordered after the guard
168
+ cannot turn an allowed POST into a blocked DELETE. A ``_method`` form field
169
+ in the body is not visible here; resolve it before the guard.
170
+ """
171
+ methods = [method.upper()]
172
+ if query:
173
+ from urllib.parse import parse_qsl
174
+
175
+ for key, value in parse_qsl(query, keep_blank_values=False, errors="replace"):
176
+ value = value.strip().upper()
177
+ if key == METHOD_OVERRIDE_QUERY_PARAM and value and value not in methods:
178
+ methods.append(value)
179
+ if headers:
180
+ # Every value of every override header counts. Starlette's Headers.items()
181
+ # yields each raw pair, and a repeated header must not be collapsed to one
182
+ # value: an override middleware may read the first while a dict keeps the
183
+ # last. Raw ASGI pairs (bytes) are accepted too.
184
+ pairs = headers.items() if hasattr(headers, "items") else headers
185
+ for key, value in pairs:
186
+ name = key.decode("latin-1") if isinstance(key, (bytes, bytearray)) else str(key)
187
+ if name.lower() not in METHOD_OVERRIDE_HEADERS or value is None:
188
+ continue
189
+ for item in value if isinstance(value, (list, tuple)) else [value]:
190
+ text = item.decode("latin-1") if isinstance(item, (bytes, bytearray)) else str(item)
191
+ for part in text.split(","):
192
+ part = part.strip().upper()
193
+ if part and part not in methods:
194
+ methods.append(part)
195
+ return methods
196
+
197
+
198
+ def evaluate_impersonation_guard(
199
+ method: str,
200
+ path: str,
201
+ context: Optional[ImpersonationContext],
202
+ policy: Optional[ScopeConfig],
203
+ on_blocked: Optional[Callable[[ImpersonationContext], None]] = None,
204
+ session_live: Optional[bool] = None,
205
+ is_impersonation_allowed: Optional[Callable[[ImpersonationContext], bool]] = None,
206
+ liveness_unavailable: bool = False,
207
+ bridge_path: Optional[str] = None,
208
+ *,
209
+ raw_path: Optional[str] = None,
210
+ alias_paths: Sequence[str] = (),
211
+ methods: Optional[Sequence[str]] = None,
212
+ ) -> GuardDecision:
213
+ """Decide whether an impersonated request may reach its handler.
214
+
215
+ ``path`` is the decoded, full client-visible path the framework dispatches on
216
+ (mount prefix included, no query). It is never re-parsed as a URL. Policy
217
+ patterns match this convention.
218
+
219
+ ``raw_path`` is the percent-encoded wire path when the adapter has it.
220
+ ``alias_paths`` are further decoded views of the same request (mount-relative,
221
+ locale prefix removed); they can only cause a block. ``methods`` defaults to
222
+ ``[method]``; adapters pass :func:`policy_methods` to include overrides.
223
+
224
+ ``bridge_path``, when set, allows a read-scope ``POST`` to exactly that path
225
+ (the browser-session bridge). It is off by default.
226
+ """
227
+ if context is None:
228
+ return GuardDecision(allowed=True)
229
+ if not isinstance(context, ImpersonationContext) or not isinstance(context.is_impersonation, bool):
230
+ return GuardDecision(allowed=False, status_code=500, body=_guard_error_body(
231
+ "Invalid impersonation context", "INVALID_IMPERSONATION_CONTEXT"
232
+ ))
233
+ if context.is_impersonation is False:
234
+ return GuardDecision(allowed=True)
235
+
236
+ judged = [path, *alias_paths]
237
+ if any(is_ambiguous_request_path(candidate) for candidate in judged) or (
238
+ raw_path is not None and is_ambiguous_request_path(raw_path)
239
+ ):
240
+ if on_blocked:
241
+ on_blocked(context)
242
+ return GuardDecision(
243
+ allowed=False,
244
+ status_code=403,
245
+ body=_guard_error_body(
246
+ "This endpoint is blocked during impersonation", "IMPERSONATION_ENDPOINT_BLOCKED"
247
+ ),
248
+ )
249
+ # Deny rules also see one further decoding of every view (a literal "%75" in a
250
+ # decoded path must not dodge "/users"), and the encoded wire form.
251
+ judged += [unquote(candidate) for candidate in judged]
252
+ if raw_path is not None:
253
+ judged += [raw_path, unquote(raw_path)]
254
+ judged = list(dict.fromkeys(judged))
255
+
256
+ valid, expired = _validate_context(context)
257
+ if not valid:
258
+ return GuardDecision(
259
+ allowed=False,
260
+ status_code=401,
261
+ body=_guard_error_body(
262
+ "Impersonation session expired" if expired else "Invalid impersonation context",
263
+ "IMPERSONATION_EXPIRED" if expired else "INVALID_IMPERSONATION_CONTEXT",
264
+ ),
265
+ )
266
+
267
+ # Liveness could not be verified and the guard fails closed: this is a
268
+ # distinct, retryable condition, not a terminated session.
269
+ if liveness_unavailable:
270
+ return GuardDecision(
271
+ allowed=False,
272
+ status_code=503,
273
+ body=_guard_error_body(
274
+ "Impersonation session liveness could not be verified",
275
+ "IMPERSONATION_LIVENESS_UNAVAILABLE",
276
+ ),
277
+ )
278
+
279
+ # Optional server-side liveness: block sessions terminated/revoked before their expiry.
280
+ if session_live is False:
281
+ if on_blocked:
282
+ on_blocked(context)
283
+ return GuardDecision(
284
+ allowed=False,
285
+ status_code=401,
286
+ body=_guard_error_body(
287
+ "Impersonation session is no longer active", "IMPERSONATION_SESSION_ENDED"
288
+ ),
289
+ )
290
+
291
+ if policy is None:
292
+ return GuardDecision(
293
+ allowed=False,
294
+ status_code=503,
295
+ body=_guard_error_body(
296
+ "Impersonation policy is temporarily unavailable", "IMPERSONATION_POLICY_UNAVAILABLE"
297
+ ),
298
+ )
299
+
300
+ request_methods = [m.upper() for m in (methods or [method])]
301
+ if method.upper() not in request_methods:
302
+ request_methods.insert(0, method.upper())
303
+ # Frameworks answer HEAD with the GET handler, so GET deny rules cover HEAD.
304
+ deny_methods = request_methods + (["GET"] if "HEAD" in request_methods and "GET" not in request_methods else [])
305
+
306
+ # 1. Deny rules FIRST - blocked regardless of scope, on any path view.
307
+ if _find_endpoint(deny_methods, judged, policy.blocked_endpoints, deny=True):
308
+ if on_blocked:
309
+ on_blocked(context)
310
+ return GuardDecision(
311
+ allowed=False,
312
+ status_code=403,
313
+ body=_guard_error_body(
314
+ "This endpoint is blocked during impersonation", "IMPERSONATION_ENDPOINT_BLOCKED"
315
+ ),
316
+ )
317
+
318
+ # 2. Application-specific semantic authorization.
319
+ semantic_allowed = True
320
+ if is_impersonation_allowed:
321
+ semantic_allowed = is_impersonation_allowed(context)
322
+ # A coroutine is truthy even when its eventual decision is False. Sync
323
+ # guards must fail closed for async hooks; async adapters await first.
324
+ if inspect.iscoroutine(semantic_allowed):
325
+ semantic_allowed.close()
326
+ if semantic_allowed is not True:
327
+ if on_blocked:
328
+ on_blocked(context)
329
+ return GuardDecision(
330
+ allowed=False,
331
+ status_code=403,
332
+ body=_guard_error_body(
333
+ "This action is not permitted during impersonation", "IMPERSONATION_ENDPOINT_BLOCKED"
334
+ ),
335
+ )
336
+
337
+ # 3. Write scope allows everything that is not denied.
338
+ if context.scope == "write":
339
+ return GuardDecision(allowed=True)
340
+
341
+ write_methods = [m for m in request_methods if m not in READ_METHODS]
342
+ if not write_methods:
343
+ return GuardDecision(allowed=True)
344
+
345
+ # 4. The browser-session bridge mints a Devora resume code for this
346
+ # already-authenticated, live session. It is not a customer write.
347
+ if bridge_path and request_methods == ["POST"] and path == bridge_path and (
348
+ raw_path is None or unquote(raw_path) in (path, *alias_paths)
349
+ ):
350
+ return GuardDecision(allowed=True)
351
+
352
+ # 5. Read scope: every write method must be allowlisted on the routed path.
353
+ if all(_find_endpoint([m], [path], policy.safe_read_endpoints, deny=False) for m in write_methods):
354
+ return GuardDecision(allowed=True)
355
+
356
+ if on_blocked:
357
+ on_blocked(context)
358
+ return GuardDecision(
359
+ allowed=False,
360
+ status_code=403,
361
+ body=_guard_error_body(
362
+ "This action is blocked during impersonation", "IMPERSONATION_SCOPE_VIOLATION"
363
+ ),
364
+ )
365
+
366
+
367
+ def _guard_error_body(error: str, error_code: str) -> dict[str, object]:
368
+ """Standard blocked response matching the Node guard shape."""
369
+ return {
370
+ "success": False,
371
+ "error": error,
372
+ "errorCode": error_code,
373
+ }
374
+
375
+
376
+ def _strict_int(value: object) -> Optional[int]:
377
+ """An integer, or an integral finite float; never a bool, string, NaN or infinity."""
378
+ if isinstance(value, bool):
379
+ return None
380
+ if isinstance(value, int):
381
+ return value
382
+ if isinstance(value, float) and math.isfinite(value) and value.is_integer():
383
+ return int(value)
384
+ return None
385
+
386
+
387
+ def coerce_impersonation_context(value: Any) -> Optional[ImpersonationContext]:
388
+ """Convert what a context extractor returned into an :class:`ImpersonationContext`.
389
+
390
+ Accepts ``None`` (ordinary traffic), an ``ImpersonationContext`` or a mapping
391
+ using either camelCase or snake_case keys. Anything else, including awaitables,
392
+ dataclasses and other objects, raises :class:`InvalidImpersonationContext`
393
+ so the guard fails closed instead of treating it as ordinary traffic.
394
+ """
395
+ if value is None:
396
+ return value
397
+ if isinstance(value, ImpersonationContext):
398
+ if not isinstance(value.is_impersonation, bool):
399
+ raise InvalidImpersonationContext("is_impersonation must be a boolean")
400
+ return value
401
+ if inspect.isawaitable(value):
402
+ if inspect.iscoroutine(value):
403
+ value.close()
404
+ raise InvalidImpersonationContext("Context extractor returned an awaitable in a sync guard")
405
+ if not isinstance(value, Mapping):
406
+ raise InvalidImpersonationContext(
407
+ f"Context extractor returned {type(value).__name__}; return None, a dict or ImpersonationContext"
408
+ )
409
+ is_impersonation = value.get("isImpersonation", value.get("is_impersonation"))
410
+ if not isinstance(is_impersonation, bool):
411
+ raise InvalidImpersonationContext("isImpersonation must be a boolean")
412
+ return ImpersonationContext(
413
+ is_impersonation=is_impersonation,
414
+ scope=value.get("scope", ""),
415
+ session_id=value.get("sessionId", value.get("session_id")),
416
+ expires_at=value.get("expiresAt", value.get("expires_at")),
417
+ impersonator=value.get("impersonator"),
418
+ actor=value.get("actor", value.get("impersonator")),
419
+ subject=value.get("subject", value.get("targetUser", value.get("target_user"))),
420
+ auth_method=value.get("authMethod", value.get("auth_method")),
421
+ authorization_source=value.get("authorizationSource", value.get("authorization_source")),
422
+ recording_allowed=value.get("recordingAllowed", value.get("recording_allowed")),
423
+ )
424
+
425
+
426
+ def _identity(value: object) -> Optional[str]:
427
+ if isinstance(value, Mapping):
428
+ identifier = value.get("id")
429
+ if isinstance(identifier, str) and identifier:
430
+ return identifier
431
+ return None
432
+
433
+
434
+ def validate_impersonation_context(
435
+ context: Optional[ImpersonationContext],
436
+ ) -> dict[str, Any]:
437
+ """Strictly validate a context: types are checked, nothing is coerced."""
438
+ if not isinstance(context, ImpersonationContext):
439
+ return {"valid": False, "error": "No impersonation context"}
440
+ if context.is_impersonation is not True:
441
+ return {"valid": False, "error": "Not an impersonation session"}
442
+ if context.scope not in ("read", "write"):
443
+ return {"valid": False, "error": "Invalid scope"}
444
+ if not _identity(context.actor or context.impersonator) or not _identity(context.subject):
445
+ return {"valid": False, "error": "Missing impersonation actor or subject"}
446
+ if context.auth_method != "devora_impersonation":
447
+ return {"valid": False, "error": "Invalid impersonation authentication method"}
448
+ if context.authorization_source not in _AUTHORIZATION_SOURCES:
449
+ return {"valid": False, "error": "Invalid authorization source"}
450
+ if not isinstance(context.recording_allowed, bool):
451
+ return {"valid": False, "error": "Missing recording authorization"}
452
+ if not isinstance(context.session_id, str) or not context.session_id:
453
+ return {"valid": False, "error": "Missing session ID"}
454
+ expires_at = _strict_int(context.expires_at)
455
+ if expires_at is None:
456
+ return {"valid": False, "error": "Missing expiration"}
457
+ if expires_at < 946684800000:
458
+ return {"valid": False, "error": "Expiration must be a Unix timestamp in milliseconds"}
459
+ if int(time.time() * 1000) > expires_at:
460
+ return {"valid": False, "expired": True, "error": "Impersonation session expired"}
461
+ return {"valid": True}
462
+
463
+
464
+ def _validate_context(context: ImpersonationContext) -> tuple[bool, bool]:
465
+ result = validate_impersonation_context(context)
466
+ return bool(result["valid"]), bool(result.get("expired"))
467
+
468
+
469
+ def _find_endpoint(
470
+ methods: Sequence[str], paths: Sequence[str], endpoints: list[dict[str, str]], deny: bool
471
+ ) -> Optional[dict[str, str]]:
472
+ """Deny rules: any method on any path, ignoring case and trailing slashes.
473
+ Allow rules: the exact method on every path, strictly."""
474
+ for endpoint in endpoints:
475
+ endpoint_method = str(endpoint.get("method", "")).upper()
476
+ pattern = endpoint.get("pattern", "")
477
+ if not any((endpoint_method == "*" and deny) or endpoint_method == m for m in methods):
478
+ continue
479
+ matches = [
480
+ match_endpoint_pattern(pattern, candidate, case_sensitive=not deny, ignore_trailing_slash=deny)
481
+ for candidate in paths
482
+ ]
483
+ if (any(matches) if deny else all(matches)):
484
+ return endpoint
485
+ return None