Coverage for agentos/tools/key_rotation.py: 33%

117 statements  

« prev     ^ index     » next       coverage.py v7.14.3, created at 2026-07-06 11:37 +0800

1""" 

2KeyRotation — automated secret key rotation with grace periods and scheduled callbacks. 

3 

4Supports: 

5 - Schedule-based rotation (interval in seconds) 

6 - Manual rotation trigger 

7 - Grace period (old key still valid for verification) 

8 - Current / pending / expired key states 

9 - Rotation hooks (pre_rotate, post_rotate) 

10 - Thread-safe 

11""" 

12 

13from __future__ import annotations 

14 

15import secrets 

16import threading 

17import time 

18from dataclasses import dataclass 

19from enum import Enum 

20from typing import Any, Callable, Dict, List, Optional 

21 

22 

23# ============================================================================ 

24# Key State 

25# ============================================================================ 

26 

27class KeyState(Enum): 

28 CURRENT = "current" 

29 PENDING = "pending" # newly rotated, still in grace 

30 EXPIRED = "expired" 

31 

32 

33@dataclass 

34class KeyEntry: 

35 key: str 

36 state: KeyState 

37 created_at: float 

38 expires_at: Optional[float] = None 

39 

40 

41# ============================================================================ 

42# KeyRotation 

43# ============================================================================ 

44 

45class KeyRotation: 

46 """Automated secret key rotation with grace periods. 

47 

48 Usage: 

49 kr = KeyRotation(rotation_interval=3600, grace_period=300, key_length=32) 

50 kr.start() 

51 

52 current = kr.current_key # Use this for signing/encryption 

53 active_keys = kr.active_keys # All currently valid keys 

54 

55 # When a new key is rotated in: 

56 # - current becomes the new key 

57 # - old key stays in active_keys during grace period (for validation) 

58 # - after grace period expires, old key is removed 

59 """ 

60 

61 def __init__( 

62 self, 

63 rotation_interval: float = 3600.0, 

64 grace_period: float = 300.0, 

65 key_length: int = 32, 

66 ): 

67 if rotation_interval <= 0: 

68 raise ValueError("rotation_interval must be positive") 

69 if grace_period < 0: 

70 raise ValueError("grace_period must be non-negative") 

71 self._interval = rotation_interval 

72 self._grace_period = grace_period 

73 self._key_length = key_length 

74 self._keys: List[KeyEntry] = [] 

75 self._lock = threading.RLock() 

76 self._timer: Optional[threading.Timer] = None 

77 self._running = False 

78 

79 # Hooks 

80 self._pre_rotate: List[Callable[[], None]] = [] 

81 self._post_rotate: List[Callable[[str, str], None]] = [] # (old_key, new_key) 

82 

83 # Seed with initial key 

84 self._rotate_now() 

85 

86 # ---------- Lifecycle ---------- 

87 

88 def start(self) -> None: 

89 with self._lock: 

90 if self._running: 

91 return 

92 self._running = True 

93 self._schedule_next() 

94 

95 def stop(self) -> None: 

96 with self._lock: 

97 self._running = False 

98 if self._timer: 

99 self._timer.cancel() 

100 self._timer = None 

101 

102 def _schedule_next(self) -> None: 

103 with self._lock: 

104 if not self._running: 

105 return 

106 self._timer = threading.Timer(self._interval, self._on_timer) 

107 self._timer.daemon = True 

108 self._timer.start() 

109 

110 def _on_timer(self) -> None: 

111 self._rotate_now() 

112 self._schedule_next() 

113 

114 # ---------- Rotation ---------- 

115 

116 def rotate(self) -> str: 

117 """Manually trigger a rotation. Returns the new key.""" 

118 return self._rotate_now() 

119 

120 def _rotate_now(self) -> str: 

121 new_key = secrets.token_hex(self._key_length) 

122 now = time.time() 

123 

124 with self._lock: 

125 old_key = self._keys[0].key if self._keys else None 

126 

127 # Notify pre-rotation 

128 self._notify_pre_rotate() 

129 

130 # Move current → pending (if grace > 0), otherwise expired 

131 for entry in self._keys: 

132 if entry.state == KeyState.CURRENT: 

133 if self._grace_period > 0: 

134 entry.state = KeyState.PENDING 

135 entry.expires_at = now + self._grace_period 

136 else: 

137 entry.state = KeyState.EXPIRED 

138 

139 # Add new current key 

140 new_entry = KeyEntry(key=new_key, state=KeyState.CURRENT, created_at=now) 

141 self._keys.insert(0, new_entry) 

142 

143 # Clean up expired 

144 self._keys = [e for e in self._keys if e.state != KeyState.EXPIRED] 

145 

146 # Notify post-rotation 

147 self._notify_post_rotate(old_key, new_key) 

148 

149 return new_key 

150 

151 # ---------- Key Access ---------- 

152 

153 @property 

154 def current_key(self) -> Optional[str]: 

155 with self._lock: 

156 for entry in self._keys: 

157 if entry.state == KeyState.CURRENT: 

158 return entry.key 

159 return None 

160 

161 @property 

162 def active_keys(self) -> List[str]: 

163 """All currently valid keys (CURRENT + PENDING).""" 

164 with self._lock: 

165 self._cleanup_expired() 

166 return [e.key for e in self._keys if e.state in (KeyState.CURRENT, KeyState.PENDING)] 

167 

168 @property 

169 def pending_keys(self) -> List[str]: 

170 """Keys in grace period only.""" 

171 with self._lock: 

172 self._cleanup_expired() 

173 return [e.key for e in self._keys if e.state == KeyState.PENDING] 

174 

175 def is_valid(self, key: str) -> bool: 

176 """Check if a key is currently valid (CURRENT or PENDING).""" 

177 return key in self.active_keys 

178 

179 def _cleanup_expired(self) -> None: 

180 now = time.time() 

181 self._keys = [ 

182 e for e in self._keys 

183 if e.state != KeyState.EXPIRED 

184 and (e.expires_at is None or e.expires_at > now) 

185 ] 

186 

187 # ---------- Hooks ---------- 

188 

189 def on_pre_rotate(self, callback: Callable[[], None]) -> None: 

190 self._pre_rotate.append(callback) 

191 

192 def on_post_rotate(self, callback: Callable[[str, str], None]) -> None: 

193 self._post_rotate.append(callback) 

194 

195 def _notify_pre_rotate(self) -> None: 

196 for cb in self._pre_rotate: 

197 try: 

198 cb() 

199 except Exception: 

200 pass 

201 

202 def _notify_post_rotate(self, old_key: Optional[str], new_key: str) -> None: 

203 for cb in self._post_rotate: 

204 try: 

205 cb(old_key, new_key) 

206 except Exception: 

207 pass 

208 

209 # ---------- Info ---------- 

210 

211 @property 

212 def stats(self) -> Dict[str, Any]: 

213 with self._lock: 

214 self._cleanup_expired() 

215 return { 

216 "total_keys": len(self._keys), 

217 "current": 1 if any(e.state == KeyState.CURRENT for e in self._keys) else 0, 

218 "pending": sum(1 for e in self._keys if e.state == KeyState.PENDING), 

219 "rotation_interval": self._interval, 

220 "grace_period": self._grace_period, 

221 "key_length": self._key_length, 

222 "running": self._running, 

223 }