Coverage for src/lexigram/auth/authn/key_rotation.py: 92%
50 statements
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-26 00:58 +0800
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-26 00:58 +0800
1"""JWT key rotation helpers.
3This module provides the :class:`JWTKeyStore` helper that manages the
4lifecycle of JWT signing keys — adding new keys, retiring old ones after a
5grace period, and exposing the correct key material for signing and
6verification.
8The :class:`~lexigram.auth.authn.jwt.JWTTokenManager` delegates all key
9management to this class.
11Example::
13 from lexigram.auth.authn.key_rotation import JWTKeyStore
14 from lexigram.validation import SecretStr
16 store = JWTKeyStore(
17 current_key_id="v1",
18 keys={"v1": SecretStr("super-secret")},
19 grace_period_seconds=3600,
20 )
21 await store.rotate("v2", "new-super-secret")
22 signing_key = store.get_signing_key()
23"""
25from __future__ import annotations
27from typing import Any, cast
29from lexigram.logging import get_logger
30from lexigram.primitives import clock as ambient_clock
31from lexigram.validation import SecretStr
33__all__ = ["JWTKeyStore"]
35logger = get_logger(__name__)
38class JWTKeyStore:
39 """Manages JWT signing keys with support for seamless rotation.
41 Keys are kept for a configurable grace period after rotation so that
42 tokens signed with the old key remain verifiable during the overlap
43 window.
45 Args:
46 current_key_id: The key ID that should be used for signing new tokens.
47 keys: Mapping of key ID → key material. Values may be
48 :class:`~lexigram.validation.SecretStr` instances (symmetric) or
49 dicts with ``"private"``/``"public"`` entries (asymmetric).
50 grace_period_seconds: How long to retain retired keys before deletion.
51 Defaults to 3600 s (1 hour).
52 """
54 def __init__(
55 self,
56 current_key_id: str,
57 keys: dict[str, Any] | None = None,
58 grace_period_seconds: float = 3600.0,
59 ) -> None:
60 self.current_key_id = current_key_id
61 self.keys: dict[str, Any] = keys or {}
62 self.grace_period_seconds = grace_period_seconds
63 self._key_meta: dict[str, dict[str, Any]] = {
64 kid: {"created_at": ambient_clock.now()} for kid in self.keys
65 }
67 # ── Rotation ──────────────────────────────────────────────────────────
69 async def rotate(self, new_key_id: str, new_secret: str | dict) -> None:
70 """Switch to a new signing key, retaining the old one for the grace period.
72 Args:
73 new_key_id: Identifier for the new key.
74 new_secret: Key material — a plain string for HMAC algorithms or a
75 dict with ``"private"`` / ``"public"`` entries for RSA/EC.
76 """
77 if isinstance(new_secret, dict):
78 self.keys[new_key_id] = {
79 sk: SecretStr(sv) if isinstance(sv, str) else sv
80 for sk, sv in new_secret.items()
81 }
82 else:
83 self.keys[new_key_id] = (
84 SecretStr(new_secret) if isinstance(new_secret, str) else new_secret
85 )
86 self._key_meta[new_key_id] = {"created_at": ambient_clock.now()}
87 self.current_key_id = new_key_id
88 await self._cleanup_old_keys()
90 async def _cleanup_old_keys(self) -> None:
91 """Remove keys that have been retired beyond the grace period."""
92 if not self._key_meta:
93 return
95 now = ambient_clock.now()
96 remove: list[str] = []
97 for kid, meta in list(self._key_meta.items()):
98 created = meta.get("created_at")
99 if not created:
100 continue
101 if (now - created).total_seconds() > self.grace_period_seconds:
102 if kid != self.current_key_id:
103 remove.append(kid)
105 for kid in remove:
106 self.keys.pop(kid, None)
107 self._key_meta.pop(kid, None)
109 # ── Key access ────────────────────────────────────────────────────────
111 def get_signing_key(self) -> str:
112 """Return the raw signing key string for the current key ID."""
113 val = self.keys[self.current_key_id]
114 if isinstance(val, dict):
115 return cast("str", val["private"].get_secret_value())
116 return cast("str", val.get_secret_value())
118 def get_verification_key(self, kid: str) -> str | None:
119 """Return the raw verification key string for *kid*, or ``None``."""
120 val = self.keys.get(kid)
121 if val is None:
122 return None
123 if isinstance(val, dict):
124 key_entry = val.get("public") or val.get("private")
125 return cast(
126 "str | None", key_entry.get_secret_value() if key_entry else None
127 )
128 return cast("str", val.get_secret_value())
130 def list_keys(self) -> dict[str, dict[str, Any]]:
131 """Return a copy of the key metadata dict (for inspection/auditing)."""
132 return {k: v.copy() for k, v in self._key_meta.items()}