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

155 statements  

« prev     ^ index     » next       coverage.py v7.14.3, created at 2026-07-09 10:28 +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 os 

11import time 

12from dataclasses import asdict, dataclass, field 

13from enum import Enum 

14from pathlib import Path 

15from typing import Any 

16 

17__all__ = [ 

18 "AuditSeverity", 

19 "AuditActionCategory", 

20 "AuditEvent", 

21 "AuditLogger", 

22] 

23 

24 

25# ── Enums ───────────────────────────────────────────────────────── 

26 

27 

28class AuditSeverity(Enum): 

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

30 

31 DEBUG = "debug" 

32 INFO = "info" 

33 WARNING = "warning" 

34 ERROR = "error" 

35 CRITICAL = "critical" 

36 

37 

38class AuditActionCategory(Enum): 

39 """Category of the audited action.""" 

40 

41 TOOL_CALL = "tool_call" 

42 AGENT_INVOKE = "agent_invoke" 

43 MODEL_ROUTE = "model_route" 

44 CONFIG_CHANGE = "config_change" 

45 SECURITY = "security" 

46 SYSTEM = "system" 

47 USER_ACTION = "user_action" 

48 

49 

50# ── Audit Event ─────────────────────────────────────────────────── 

51 

52 

53@dataclass 

54class AuditEvent: 

55 """A single immutable audit event. 

56 

57 Attributes: 

58 agent: Agent name or identifier. 

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

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

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

62 severity: Event severity. 

63 category: Action category. 

64 session_id: Session identifier for grouping. 

65 timestamp: Unix timestamp when the event occurred. 

66 duration_ms: Elapsed time in milliseconds. 

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

68 details: Arbitrary structured metadata. 

69 """ 

70 

71 agent: str 

72 action: str 

73 target: str = "" 

74 result: str = "success" 

75 severity: AuditSeverity = AuditSeverity.INFO 

76 category: AuditActionCategory = AuditActionCategory.SYSTEM 

77 session_id: str = "" 

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

79 duration_ms: float = 0.0 

80 error_message: str = "" 

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

82 

83 def to_dict(self) -> dict: 

84 d = asdict(self) 

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

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

87 return d 

88 

89 def to_json(self) -> str: 

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

91 

92 

93# ── Audit Logger ────────────────────────────────────────────────── 

94 

95 

96class AuditLogger: 

97 """Immutable append-only audit logger. 

98 

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

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

101 

102 Usage: 

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

104 audit.log( 

105 agent="production", 

106 action="tool:search", 

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

108 result="success", 

109 severity=AuditSeverity.INFO, 

110 category=AuditActionCategory.TOOL_CALL, 

111 session_id="sess-001", 

112 duration_ms=45.2, 

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

114 ) 

115 print(audit.stats_summary()) 

116 """ 

117 

118 def __init__( 

119 self, 

120 log_dir: str = "", 

121 auto_flush: bool = True, 

122 max_events_in_memory: int = 10000, 

123 ): 

124 self._auto_flush = auto_flush 

125 self._max_memory = max_events_in_memory 

126 

127 # Determine log path 

128 if log_dir: 

129 self._log_dir = Path(log_dir) 

130 else: 

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

132 

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

134 

135 # Rotate daily 

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

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

138 

139 # In-memory event buffer 

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

141 self._stats = _AuditStats() 

142 

143 # Ensure log file exists 

144 if not self._log_file.exists(): 

145 self._log_file.touch() 

146 

147 # ── Log ──────────────────────────────────────────────────────── 

148 

149 def log( 

150 self, 

151 *, 

152 agent: str, 

153 action: str, 

154 target: str = "", 

155 result: str = "success", 

156 severity: AuditSeverity = AuditSeverity.INFO, 

157 category: AuditActionCategory = AuditActionCategory.SYSTEM, 

158 session_id: str = "", 

159 duration_ms: float = 0.0, 

160 error_message: str = "", 

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

162 ): 

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

164 

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

166 """ 

167 event = AuditEvent( 

168 agent=agent, 

169 action=action, 

170 target=target, 

171 result=result, 

172 severity=severity, 

173 category=category, 

174 session_id=session_id, 

175 timestamp=time.time(), 

176 duration_ms=duration_ms, 

177 error_message=error_message, 

178 details=details or {}, 

179 ) 

180 

181 # Write to file immediately 

182 self._write_event(event) 

183 

184 # Track in memory (with eviction if needed) 

185 self._events.append(event) 

186 self._stats.record(event) 

187 

188 # Evict oldest if over memory limit 

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

190 self._events.pop(0) 

191 

192 # ── Stats ───────────────────────────────────────────────────── 

193 

194 def stats_summary(self) -> dict: 

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

196 return self._stats.summary() 

197 

198 # ── Query ───────────────────────────────────────────────────── 

199 

200 def query( 

201 self, 

202 session_id: str = "", 

203 category: AuditActionCategory | None = None, 

204 severity: AuditSeverity | None = None, 

205 agent: str = "", 

206 limit: int = 100, 

207 ) -> list[AuditEvent]: 

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

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

210 results = [] 

211 seen_ids = set() 

212 

213 # Search memory buffer 

214 for evt in reversed(self._events): 

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

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

217 if eid not in seen_ids: 

218 results.append(evt) 

219 seen_ids.add(eid) 

220 if len(results) >= limit: 

221 return results 

222 

223 # If no file or already have enough, return 

224 if len(results) >= limit: 

225 return results 

226 

227 # Scan log file (backwards) 

228 if self._log_file.exists(): 

229 try: 

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

231 for line in reversed(lines): 

232 if not line.strip(): 

233 continue 

234 try: 

235 d = json.loads(line) 

236 evt = self._dict_to_event(d) 

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

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

239 if eid not in seen_ids: 

240 results.append(evt) 

241 seen_ids.add(eid) 

242 except (json.JSONDecodeError, KeyError): 

243 continue 

244 if len(results) >= limit: 

245 break 

246 except Exception: 

247 pass 

248 

249 return results[:limit] 

250 

251 # ── Export ──────────────────────────────────────────────────── 

252 

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

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

255 events = ( 

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

257 ) 

258 if fmt == "json": 

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

260 # jsonl 

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

262 

263 # ── Internal ────────────────────────────────────────────────── 

264 

265 def _write_event(self, event: AuditEvent): 

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

267 try: 

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

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

270 if self._auto_flush: 

271 f.flush() 

272 except Exception: 

273 pass 

274 

275 def _match( 

276 self, 

277 evt: AuditEvent, 

278 session_id: str, 

279 category: AuditActionCategory | None, 

280 severity: AuditSeverity | None, 

281 agent: str, 

282 ) -> bool: 

283 if session_id and evt.session_id != session_id: 

284 return False 

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

286 return False 

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

288 return False 

289 if agent and evt.agent != agent: 

290 return False 

291 return True 

292 

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

294 return AuditEvent( 

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

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

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

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

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

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

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

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

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

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

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

306 ) 

307 

308 

309# ── Internal Stats Tracker ─────────────────────────────────────── 

310 

311 

312class _AuditStats: 

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

314 

315 def __init__(self): 

316 self.total_events = 0 

317 self.success_count = 0 

318 self.failure_count = 0 

319 self.total_duration_ms = 0.0 

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

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

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

323 self.first_event_ts: float = 0.0 

324 self.last_event_ts: float = 0.0 

325 

326 def record(self, event: AuditEvent): 

327 self.total_events += 1 

328 if event.result == "success": 

329 self.success_count += 1 

330 else: 

331 self.failure_count += 1 

332 

333 self.total_duration_ms += event.duration_ms 

334 

335 sev = event.severity.value 

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

337 

338 cat = event.category.value 

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

340 

341 agt = event.agent 

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

343 

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

345 self.first_event_ts = event.timestamp 

346 if event.timestamp > self.last_event_ts: 

347 self.last_event_ts = event.timestamp 

348 

349 def summary(self) -> dict: 

350 error_rate = self.failure_count / self.total_events if self.total_events > 0 else 0.0 

351 return { 

352 "total_events": self.total_events, 

353 "success": self.success_count, 

354 "failure": self.failure_count, 

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

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

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

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

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

360 }