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

1"""JWT key rotation helpers. 

2 

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. 

7 

8The :class:`~lexigram.auth.authn.jwt.JWTTokenManager` delegates all key 

9management to this class. 

10 

11Example:: 

12 

13 from lexigram.auth.authn.key_rotation import JWTKeyStore 

14 from lexigram.validation import SecretStr 

15 

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""" 

24 

25from __future__ import annotations 

26 

27from typing import Any, cast 

28 

29from lexigram.logging import get_logger 

30from lexigram.primitives import clock as ambient_clock 

31from lexigram.validation import SecretStr 

32 

33__all__ = ["JWTKeyStore"] 

34 

35logger = get_logger(__name__) 

36 

37 

38class JWTKeyStore: 

39 """Manages JWT signing keys with support for seamless rotation. 

40 

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. 

44 

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 """ 

53 

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 } 

66 

67 # ── Rotation ────────────────────────────────────────────────────────── 

68 

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. 

71 

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() 

89 

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 

94 

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) 

104 

105 for kid in remove: 

106 self.keys.pop(kid, None) 

107 self._key_meta.pop(kid, None) 

108 

109 # ── Key access ──────────────────────────────────────────────────────── 

110 

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()) 

117 

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()) 

129 

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()}