Coverage for agentos/security/audit_logger.py: 33%

155 statements  

« prev     ^ index     » next       coverage.py v7.14.3, created at 2026-07-05 20:52 +0800

1"""Audit Logger — immutable, append-only audit trail for agent operations. 

2 

3Records every tool call, agent action, and routing decision as JSONL events. 

4Supports session grouping, severity filtering, and stats aggregation. 

5""" 

6 

7from __future__ import annotations 

8 

9import json 

10import time 

11import os 

12from dataclasses import dataclass, field, asdict 

13from enum import Enum 

14from pathlib import Path 

15from typing import Any, Optional 

16 

17 

18__all__ = [ 

19 "AuditSeverity", 

20 "AuditActionCategory", 

21 "AuditEvent", 

22 "AuditLogger", 

23] 

24 

25 

26# ── Enums ───────────────────────────────────────────────────────── 

27 

28class AuditSeverity(Enum): 

29 """Severity level for audit events.""" 

30 DEBUG = "debug" 

31 INFO = "info" 

32 WARNING = "warning" 

33 ERROR = "error" 

34 CRITICAL = "critical" 

35 

36 

37class AuditActionCategory(Enum): 

38 """Category of the audited action.""" 

39 TOOL_CALL = "tool_call" 

40 AGENT_INVOKE = "agent_invoke" 

41 MODEL_ROUTE = "model_route" 

42 CONFIG_CHANGE = "config_change" 

43 SECURITY = "security" 

44 SYSTEM = "system" 

45 USER_ACTION = "user_action" 

46 

47 

48# ── Audit Event ─────────────────────────────────────────────────── 

49 

50@dataclass 

51class AuditEvent: 

52 """A single immutable audit event. 

53 

54 Attributes: 

55 agent: Agent name or identifier. 

56 action: Action description (e.g., "tool:search", "agent_start"). 

57 target: Target of the action (e.g., tool args, task snippet). 

58 result: Outcome — "success", "failure", "pending", "timeout". 

59 severity: Event severity. 

60 category: Action category. 

61 session_id: Session identifier for grouping. 

62 timestamp: Unix timestamp when the event occurred. 

63 duration_ms: Elapsed time in milliseconds. 

64 error_message: Error message if result is "failure". 

65 details: Arbitrary structured metadata. 

66 """ 

67 agent: str 

68 action: str 

69 target: str = "" 

70 result: str = "success" 

71 severity: AuditSeverity = AuditSeverity.INFO 

72 category: AuditActionCategory = AuditActionCategory.SYSTEM 

73 session_id: str = "" 

74 timestamp: float = field(default_factory=time.time) 

75 duration_ms: float = 0.0 

76 error_message: str = "" 

77 details: dict[str, Any] = field(default_factory=dict) 

78 

79 def to_dict(self) -> dict: 

80 d = asdict(self) 

81 d["severity"] = self.severity.value 

82 d["category"] = self.category.value 

83 return d 

84 

85 def to_json(self) -> str: 

86 return json.dumps(self.to_dict(), default=str) 

87 

88 

89# ── Audit Logger ────────────────────────────────────────────────── 

90 

91class AuditLogger: 

92 """Immutable append-only audit logger. 

93 

94 Writes audit events as JSONL (one JSON object per line) to a log file. 

95 Supports in-memory stat tracking and log rotation by date. 

96 

97 Usage: 

98 audit = AuditLogger(log_dir="./audit_logs", auto_flush=True) 

99 audit.log( 

100 agent="production", 

101 action="tool:search", 

102 target="query='latest news'", 

103 result="success", 

104 severity=AuditSeverity.INFO, 

105 category=AuditActionCategory.TOOL_CALL, 

106 session_id="sess-001", 

107 duration_ms=45.2, 

108 details={"arguments": {"query": "latest news"}}, 

109 ) 

110 print(audit.stats_summary()) 

111 """ 

112 

113 def __init__( 

114 self, 

115 log_dir: str = "", 

116 auto_flush: bool = True, 

117 max_events_in_memory: int = 10000, 

118 ): 

119 self._auto_flush = auto_flush 

120 self._max_memory = max_events_in_memory 

121 

122 # Determine log path 

123 if log_dir: 

124 self._log_dir = Path(log_dir) 

125 else: 

126 self._log_dir = Path(os.getcwd()) / "audit_logs" 

127 

128 self._log_dir.mkdir(parents=True, exist_ok=True) 

129 

130 # Rotate daily 

131 date_str = time.strftime("%Y%m%d") 

132 self._log_file = self._log_dir / f"audit-{date_str}.jsonl" 

133 

134 # In-memory event buffer 

135 self._events: list[AuditEvent] = [] 

136 self._stats = _AuditStats() 

137 

138 # Ensure log file exists 

139 if not self._log_file.exists(): 

140 self._log_file.touch() 

141 

142 # ── Log ──────────────────────────────────────────────────────── 

143 

144 def log( 

145 self, 

146 *, 

147 agent: str, 

148 action: str, 

149 target: str = "", 

150 result: str = "success", 

151 severity: AuditSeverity = AuditSeverity.INFO, 

152 category: AuditActionCategory = AuditActionCategory.SYSTEM, 

153 session_id: str = "", 

154 duration_ms: float = 0.0, 

155 error_message: str = "", 

156 details: dict[str, Any] | None = None, 

157 ): 

158 """Append an audit event to the log. 

159 

160 All parameters are keyword-only for clarity at call sites. 

161 """ 

162 event = AuditEvent( 

163 agent=agent, 

164 action=action, 

165 target=target, 

166 result=result, 

167 severity=severity, 

168 category=category, 

169 session_id=session_id, 

170 timestamp=time.time(), 

171 duration_ms=duration_ms, 

172 error_message=error_message, 

173 details=details or {}, 

174 ) 

175 

176 # Write to file immediately 

177 self._write_event(event) 

178 

179 # Track in memory (with eviction if needed) 

180 self._events.append(event) 

181 self._stats.record(event) 

182 

183 # Evict oldest if over memory limit 

184 while len(self._events) > self._max_memory: 

185 self._events.pop(0) 

186 

187 # ── Stats ───────────────────────────────────────────────────── 

188 

189 def stats_summary(self) -> dict: 

190 """Return aggregate stats for all logged events.""" 

191 return self._stats.summary() 

192 

193 # ── Query ───────────────────────────────────────────────────── 

194 

195 def query( 

196 self, 

197 session_id: str = "", 

198 category: Optional[AuditActionCategory] = None, 

199 severity: Optional[AuditSeverity] = None, 

200 agent: str = "", 

201 limit: int = 100, 

202 ) -> list[AuditEvent]: 

203 """Query events by filters. Searches in-memory buffer first, 

204 then falls back to scanning the log file.""" 

205 results = [] 

206 seen_ids = set() 

207 

208 # Search memory buffer 

209 for evt in reversed(self._events): 

210 if self._match(evt, session_id, category, severity, agent): 

211 eid = (evt.agent, evt.action, evt.timestamp) 

212 if eid not in seen_ids: 

213 results.append(evt) 

214 seen_ids.add(eid) 

215 if len(results) >= limit: 

216 return results 

217 

218 # If no file or already have enough, return 

219 if len(results) >= limit: 

220 return results 

221 

222 # Scan log file (backwards) 

223 if self._log_file.exists(): 

224 try: 

225 lines = self._log_file.read_text().strip().split("\n") 

226 for line in reversed(lines): 

227 if not line.strip(): 

228 continue 

229 try: 

230 d = json.loads(line) 

231 evt = self._dict_to_event(d) 

232 if self._match(evt, session_id, category, severity, agent): 

233 eid = (evt.agent, evt.action, evt.timestamp) 

234 if eid not in seen_ids: 

235 results.append(evt) 

236 seen_ids.add(eid) 

237 except (json.JSONDecodeError, KeyError): 

238 continue 

239 if len(results) >= limit: 

240 break 

241 except Exception: 

242 pass 

243 

244 return results[:limit] 

245 

246 # ── Export ──────────────────────────────────────────────────── 

247 

248 def export(self, session_id: str = "", fmt: str = "jsonl") -> str: 

249 """Export audit events for a session or all events.""" 

250 events = self.query(session_id=session_id, limit=999999) if session_id else list(self._events) 

251 if fmt == "json": 

252 return json.dumps([e.to_dict() for e in events], indent=2, default=str) 

253 # jsonl 

254 return "\n".join(e.to_json() for e in events) 

255 

256 # ── Internal ────────────────────────────────────────────────── 

257 

258 def _write_event(self, event: AuditEvent): 

259 """Append a single JSONL line to the log file.""" 

260 try: 

261 with open(self._log_file, "a") as f: 

262 f.write(event.to_json() + "\n") 

263 if self._auto_flush: 

264 f.flush() 

265 except Exception: 

266 pass 

267 

268 def _match( 

269 self, 

270 evt: AuditEvent, 

271 session_id: str, 

272 category: Optional[AuditActionCategory], 

273 severity: Optional[AuditSeverity], 

274 agent: str, 

275 ) -> bool: 

276 if session_id and evt.session_id != session_id: 

277 return False 

278 if category is not None and evt.category != category: 

279 return False 

280 if severity is not None and evt.severity != severity: 

281 return False 

282 if agent and evt.agent != agent: 

283 return False 

284 return True 

285 

286 def _dict_to_event(self, d: dict) -> AuditEvent: 

287 return AuditEvent( 

288 agent=d.get("agent", ""), 

289 action=d.get("action", ""), 

290 target=d.get("target", ""), 

291 result=d.get("result", "success"), 

292 severity=AuditSeverity(d.get("severity", "info")), 

293 category=AuditActionCategory(d.get("category", "system")), 

294 session_id=d.get("session_id", ""), 

295 timestamp=d.get("timestamp", 0.0), 

296 duration_ms=d.get("duration_ms", 0.0), 

297 error_message=d.get("error_message", ""), 

298 details=d.get("details", {}), 

299 ) 

300 

301 

302# ── Internal Stats Tracker ─────────────────────────────────────── 

303 

304class _AuditStats: 

305 """Tracks aggregate statistics for audit events.""" 

306 

307 def __init__(self): 

308 self.total_events = 0 

309 self.success_count = 0 

310 self.failure_count = 0 

311 self.total_duration_ms = 0.0 

312 self.by_severity: dict[str, int] = {} 

313 self.by_category: dict[str, int] = {} 

314 self.by_agent: dict[str, int] = {} 

315 self.first_event_ts: float = 0.0 

316 self.last_event_ts: float = 0.0 

317 

318 def record(self, event: AuditEvent): 

319 self.total_events += 1 

320 if event.result == "success": 

321 self.success_count += 1 

322 else: 

323 self.failure_count += 1 

324 

325 self.total_duration_ms += event.duration_ms 

326 

327 sev = event.severity.value 

328 self.by_severity[sev] = self.by_severity.get(sev, 0) + 1 

329 

330 cat = event.category.value 

331 self.by_category[cat] = self.by_category.get(cat, 0) + 1 

332 

333 agt = event.agent 

334 self.by_agent[agt] = self.by_agent.get(agt, 0) + 1 

335 

336 if self.first_event_ts == 0 or event.timestamp < self.first_event_ts: 

337 self.first_event_ts = event.timestamp 

338 if event.timestamp > self.last_event_ts: 

339 self.last_event_ts = event.timestamp 

340 

341 def summary(self) -> dict: 

342 error_rate = ( 

343 self.failure_count / self.total_events 

344 if self.total_events > 0 

345 else 0.0 

346 ) 

347 return { 

348 "total_events": self.total_events, 

349 "success": self.success_count, 

350 "failure": self.failure_count, 

351 "error_rate": round(error_rate, 4), 

352 "total_duration_ms": round(self.total_duration_ms, 2), 

353 "by_severity": dict(self.by_severity), 

354 "by_category": dict(self.by_category), 

355 "by_agent": dict(self.by_agent), 

356 }