1from __future__ import annotations
2
3from datetime import datetime
4from typing import TYPE_CHECKING, Any
5
6from lexigram import serialization as json
7from lexigram.ai.governance.audit.models import (
8 AIAuditEvent,
9 AuditEventType,
10 AuditQuery,
11 AuditSummary,
12)
13from lexigram.contracts.audit import AuditEntry, AuditEventSeverity
14
15if TYPE_CHECKING:
16 from lexigram.contracts.audit import AuditLoggerProtocol
17 from lexigram.contracts.data import DatabaseProviderProtocol
18
19_PII_METADATA_KEYS = frozenset(
20 {"prompt", "message", "content", "query", "text", "input", "output", "response"}
21)
22
23_CREATE_TABLE = """
24CREATE TABLE IF NOT EXISTS ai_audit_events (
25 event_id TEXT NOT NULL PRIMARY KEY,
26 event_type TEXT NOT NULL,
27 timestamp TEXT NOT NULL,
28 model TEXT,
29 provider TEXT,
30 user_id TEXT,
31 status TEXT NOT NULL DEFAULT 'success',
32 tokens INTEGER,
33 cost REAL,
34 latency_ms REAL,
35 metadata TEXT NOT NULL DEFAULT '{}'
36)
37"""
38
39_INSERT_EVENT = (
40 "INSERT OR IGNORE INTO ai_audit_events "
41 "(event_id, event_type, timestamp, model, provider, user_id, "
42 "status, tokens, cost, latency_ms, metadata) "
43 "VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)"
44)
45
46
47def _sanitize_metadata(metadata: dict[str, Any]) -> str:
48 """Redact PII-sensitive keys and serialise to JSON for storage."""
49 safe: dict[str, Any] = {}
50 for k, v in metadata.items():
51 safe[k] = "[REDACTED]" if k.lower() in _PII_METADATA_KEYS else v
52 return json.dumps_str(safe)
53
54
55class DatabaseAuditStore:
56 """SQL-backed audit store using :class:`~lexigram.contracts.data.DatabaseProviderProtocol`.
57
58 Writes events to an ``ai_audit_events`` table (created lazily on first
59 use). Metadata fields matching common PII keys (``prompt``, ``content``,
60 ``message``, etc.) are redacted to ``"[REDACTED]"`` before storage.
61
62 ``record()`` is safe to call from hot-paths — it issues a single
63 ``INSERT OR IGNORE`` that will not raise on duplicate ``event_id``.
64
65 Args:
66 db: A connected :class:`~lexigram.contracts.data.DatabaseProviderProtocol`
67 resolved from the DI container.
68 """
69
70 def __init__(
71 self,
72 db: DatabaseProviderProtocol,
73 audit_logger: AuditLoggerProtocol | None = None,
74 ) -> None:
75 self._db = db
76 self._audit_logger = audit_logger
77 self._initialised = False
78
79 async def _ensure_table(self) -> None:
80 if not self._initialised:
81 await self._db.execute(_CREATE_TABLE)
82 self._initialised = True
83
84 async def record(self, event: AIAuditEvent) -> None:
85 """Persist *event* to the database, sanitising PII in metadata.
86
87 Args:
88 event: The audit event to store.
89 """
90 await self._ensure_table()
91 await self._db.execute(
92 _INSERT_EVENT,
93 [
94 event.event_id,
95 event.event_type.value,
96 event.timestamp.isoformat(),
97 event.model,
98 event.provider,
99 event.user_id,
100 event.status,
101 event.tokens,
102 event.cost,
103 event.latency_ms,
104 _sanitize_metadata(event.metadata),
105 ],
106 )
107 if self._audit_logger is not None:
108 meta: dict[str, Any] = {"model": event.model, "provider": event.provider}
109 if event.tokens is not None:
110 meta["tokens"] = event.tokens
111 if event.cost is not None:
112 meta["cost"] = event.cost
113 await self._audit_logger.log(
114 AuditEntry(
115 action=f"ai.{event.event_type.value}",
116 actor_id=event.user_id or "system",
117 resource_type="ai_inference",
118 resource_id=event.event_id,
119 outcome=event.status,
120 severity=AuditEventSeverity.MEDIUM,
121 metadata=meta,
122 source="ai_governance",
123 )
124 )
125
126 async def query(self, query: AuditQuery) -> list[AIAuditEvent]:
127 """Retrieve audit events matching *query*, newest first.
128
129 Args:
130 query: Filter criteria.
131
132 Returns:
133 Matching events ordered by timestamp descending.
134 """
135 await self._ensure_table()
136 sql, params = self._build_where(query)
137 full_sql = (
138 f"SELECT * FROM ai_audit_events{sql} " # noqa: S608 -- WHERE from fixed condition strings; values parameterized
139 f"ORDER BY timestamp DESC "
140 f"LIMIT ? OFFSET ?"
141 )
142 result = await self._db.execute_query(
143 full_sql, [*params, query.limit, query.offset]
144 )
145 return [self._row_to_event(row) for row in result.rows]
146
147 async def aggregate(self, query: AuditQuery) -> AuditSummary:
148 """Compute summary statistics for events matching *query*.
149
150 Executes four targeted SQL queries (totals + three GROUP BY) so that
151 callers do not need to load full event rows into memory.
152
153 Args:
154 query: Filter criteria that scope the aggregation.
155
156 Returns:
157 Summary statistics for the matching events.
158 """
159 await self._ensure_table()
160 where, base_params = self._build_where(query)
161
162 # -- totals ---------------------------------------------------------
163 totals_sql = (
164 "SELECT COUNT(*) AS total_events, " # noqa: S608 -- {where} from fixed condition strings; values parameterized
165 "COALESCE(SUM(cost), 0.0) AS total_spend, "
166 "COALESCE(SUM(tokens), 0) AS total_tokens, "
167 "SUM(CASE WHEN status = 'denied' THEN 1 ELSE 0 END) AS denied_count "
168 f"FROM ai_audit_events{where}"
169 )
170 totals_result = await self._db.execute_query(totals_sql, base_params)
171 row = totals_result.rows[0] if totals_result.rows else {}
172 summary = AuditSummary(
173 total_events=int(row.get("total_events", 0)),
174 total_spend=float(row.get("total_spend", 0.0)),
175 total_tokens=int(row.get("total_tokens", 0)),
176 denied_count=int(row.get("denied_count", 0)),
177 )
178
179 # -- by model -------------------------------------------------------
180 model_result = await self._db.execute_query(
181 f"SELECT COALESCE(model, 'unknown') AS grp, COUNT(*) AS cnt " # noqa: S608 -- {where} from fixed condition strings; values parameterized
182 f"FROM ai_audit_events{where} GROUP BY model",
183 base_params,
184 )
185 summary.by_model = {r["grp"]: int(r["cnt"]) for r in model_result.rows}
186
187 # -- by user --------------------------------------------------------
188 user_result = await self._db.execute_query(
189 f"SELECT COALESCE(user_id, 'anonymous') AS grp, COUNT(*) AS cnt " # noqa: S608 -- {where} from fixed condition strings; values parameterized
190 f"FROM ai_audit_events{where} GROUP BY user_id",
191 base_params,
192 )
193 summary.by_user = {r["grp"]: int(r["cnt"]) for r in user_result.rows}
194
195 # -- by event type --------------------------------------------------
196 type_result = await self._db.execute_query(
197 f"SELECT event_type AS grp, COUNT(*) AS cnt " # noqa: S608 -- {where} from fixed condition strings; values parameterized
198 f"FROM ai_audit_events{where} GROUP BY event_type",
199 base_params,
200 )
201 summary.by_event_type = {r["grp"]: int(r["cnt"]) for r in type_result.rows}
202
203 return summary
204
205 # -----------------------------------------------------------------------
206 # Private helpers
207 # -----------------------------------------------------------------------
208
209 def _build_where(self, query: AuditQuery) -> tuple[str, list[Any]]:
210 """Build a parameterised WHERE clause from *query* filter criteria."""
211 clauses: list[str] = []
212 params: list[Any] = []
213
214 if query.start:
215 clauses.append("timestamp >= ?")
216 params.append(query.start.isoformat())
217 if query.end:
218 clauses.append("timestamp <= ?")
219 params.append(query.end.isoformat())
220 if query.user_id:
221 clauses.append("user_id = ?")
222 params.append(query.user_id)
223 if query.model:
224 clauses.append("model = ?")
225 params.append(query.model)
226 if query.provider:
227 clauses.append("provider = ?")
228 params.append(query.provider)
229 if query.status:
230 clauses.append("status = ?")
231 params.append(query.status)
232 if query.event_types:
233 placeholders = ", ".join("?" * len(query.event_types))
234 clauses.append(f"event_type IN ({placeholders})")
235 params.extend(et.value for et in query.event_types)
236
237 where = (" WHERE " + " AND ".join(clauses)) if clauses else ""
238 return where, params
239
240 @staticmethod
241 def _row_to_event(row: Any) -> AIAuditEvent:
242 """Deserialise a database row into an :class:`AIAuditEvent`."""
243 raw_meta = row.get("metadata", "{}")
244 try:
245 meta = json.loads(raw_meta) if isinstance(raw_meta, str) else raw_meta
246 except (ValueError, TypeError):
247 meta = {}
248
249 return AIAuditEvent(
250 event_id=row["event_id"],
251 event_type=AuditEventType(row["event_type"]),
252 timestamp=datetime.fromisoformat(row["timestamp"]),
253 model=row.get("model"),
254 provider=row.get("provider"),
255 user_id=row.get("user_id"),
256 status=row.get("status", "success"),
257 tokens=row.get("tokens"),
258 cost=row.get("cost"),
259 latency_ms=row.get("latency_ms"),
260 metadata=meta,
261 )