systemprompt_runtime/trace/repository/
audit.rs1use systemprompt_identifiers::{AiRequestId, McpServerId, McpToolName, TaskId, TraceId};
7
8use super::{Result, TraceRepository};
9use crate::trace::models::{
10 AuditLookupResult, AuditPage, AuditToolCallRow, ConversationMessage, LinkedMcpCall,
11};
12
13struct AuditRow {
14 id: AiRequestId,
15 provider: Option<String>,
16 model: Option<String>,
17 requested_model: Option<String>,
18 input_tokens: Option<i32>,
19 output_tokens: Option<i32>,
20 cache_read_tokens: Option<i32>,
21 cache_creation_tokens: Option<i32>,
22 reasoning_tokens: Option<i32>,
23 cost_microdollars: i64,
24 latency_ms: Option<i32>,
25 status: String,
26 finish_reason: Option<String>,
27 error_message: Option<String>,
28 task_id: Option<TaskId>,
29 trace_id: Option<TraceId>,
30}
31
32struct MsgRow {
33 role: String,
34 content: String,
35 sequence_number: i32,
36}
37
38fn audit_row_to_result(r: AuditRow) -> AuditLookupResult {
39 AuditLookupResult {
40 id: r.id,
41 provider: r.provider,
42 model: r.model,
43 requested_model: r.requested_model,
44 input_tokens: r.input_tokens,
45 output_tokens: r.output_tokens,
46 cache_read_tokens: r.cache_read_tokens,
47 cache_creation_tokens: r.cache_creation_tokens,
48 reasoning_tokens: r.reasoning_tokens,
49 cost_microdollars: r.cost_microdollars,
50 latency_ms: r.latency_ms,
51 status: r.status,
52 finish_reason: r.finish_reason,
53 error_message: r.error_message,
54 task_id: r.task_id,
55 trace_id: r.trace_id,
56 }
57}
58
59impl TraceRepository {
60 pub async fn find_ai_request_for_audit(&self, id: &str) -> Result<Option<AuditLookupResult>> {
61 let partial = format!("{id}%");
62
63 if let Some(row) = self.find_audit_by_request_id(id, &partial).await? {
64 return Ok(Some(row));
65 }
66 if let Some(row) = self.find_audit_by_task_id(id, &partial).await? {
67 return Ok(Some(row));
68 }
69 self.find_audit_by_trace_id(id, &partial).await
70 }
71
72 async fn find_audit_by_request_id(
73 &self,
74 id: &str,
75 partial: &str,
76 ) -> Result<Option<AuditLookupResult>> {
77 let row = sqlx::query_as!(
78 AuditRow,
79 r#"
80 SELECT id as "id!: AiRequestId", provider, model,
81 requested_model,
82 input_tokens, output_tokens,
83 cache_read_tokens, cache_creation_tokens, reasoning_tokens,
84 cost_microdollars as "cost_microdollars!",
85 latency_ms,
86 status as "status!",
87 finish_reason,
88 error_message,
89 task_id as "task_id: TaskId",
90 trace_id as "trace_id: TraceId"
91 FROM ai_requests WHERE id = $1 OR id LIKE $2 LIMIT 1
92 "#,
93 id,
94 partial
95 )
96 .fetch_optional(&*self.pool)
97 .await?;
98
99 Ok(row.map(audit_row_to_result))
100 }
101
102 async fn find_audit_by_task_id(
103 &self,
104 id: &str,
105 partial: &str,
106 ) -> Result<Option<AuditLookupResult>> {
107 let row = sqlx::query_as!(
108 AuditRow,
109 r#"
110 SELECT id as "id!: AiRequestId", provider, model,
111 requested_model,
112 input_tokens, output_tokens,
113 cache_read_tokens, cache_creation_tokens, reasoning_tokens,
114 cost_microdollars as "cost_microdollars!",
115 latency_ms,
116 status as "status!",
117 finish_reason,
118 error_message,
119 task_id as "task_id: TaskId",
120 trace_id as "trace_id: TraceId"
121 FROM ai_requests WHERE task_id = $1 OR task_id LIKE $2
122 ORDER BY created_at DESC LIMIT 1
123 "#,
124 id,
125 partial
126 )
127 .fetch_optional(&*self.pool)
128 .await?;
129
130 Ok(row.map(audit_row_to_result))
131 }
132
133 async fn find_audit_by_trace_id(
134 &self,
135 id: &str,
136 partial: &str,
137 ) -> Result<Option<AuditLookupResult>> {
138 let row = sqlx::query_as!(
139 AuditRow,
140 r#"
141 SELECT id as "id!: AiRequestId", provider, model,
142 requested_model,
143 input_tokens, output_tokens,
144 cache_read_tokens, cache_creation_tokens, reasoning_tokens,
145 cost_microdollars as "cost_microdollars!",
146 latency_ms,
147 status as "status!",
148 finish_reason,
149 error_message,
150 task_id as "task_id: TaskId",
151 trace_id as "trace_id: TraceId"
152 FROM ai_requests WHERE trace_id = $1 OR trace_id LIKE $2
153 ORDER BY created_at DESC LIMIT 1
154 "#,
155 id,
156 partial
157 )
158 .fetch_optional(&*self.pool)
159 .await?;
160
161 Ok(row.map(audit_row_to_result))
162 }
163
164 pub async fn count_audit_messages(&self, request_id: &AiRequestId) -> Result<i64> {
165 let count = sqlx::query_scalar!(
166 r#"SELECT COUNT(*) as "count!" FROM ai_request_messages WHERE request_id = $1"#,
167 request_id.as_str()
168 )
169 .fetch_one(&*self.pool)
170 .await?;
171 Ok(count)
172 }
173
174 pub async fn count_audit_tool_calls(&self, request_id: &AiRequestId) -> Result<i64> {
175 let count = sqlx::query_scalar!(
176 r#"SELECT COUNT(*) as "count!" FROM ai_request_tool_calls WHERE request_id = $1"#,
177 request_id.as_str()
178 )
179 .fetch_one(&*self.pool)
180 .await?;
181 Ok(count)
182 }
183
184 pub async fn list_audit_messages(
185 &self,
186 request_id: &AiRequestId,
187 page: AuditPage,
188 ) -> Result<Vec<ConversationMessage>> {
189 let rows = sqlx::query_as!(
190 MsgRow,
191 r#"
192 SELECT role as "role!", content as "content!", sequence_number as "sequence_number!"
193 FROM ai_request_messages WHERE request_id = $1 ORDER BY sequence_number
194 OFFSET $2 LIMIT $3
195 "#,
196 request_id.as_str(),
197 page.offset,
198 page.sql_limit()
199 )
200 .fetch_all(&*self.pool)
201 .await?;
202
203 Ok(rows
204 .into_iter()
205 .map(|m| ConversationMessage {
206 role: m.role,
207 content: m.content,
208 sequence_number: m.sequence_number,
209 })
210 .collect())
211 }
212
213 pub async fn list_audit_tool_calls(
214 &self,
215 request_id: &AiRequestId,
216 page: AuditPage,
217 ) -> Result<Vec<AuditToolCallRow>> {
218 let rows = sqlx::query!(
219 r#"
220 SELECT tool_name as "tool_name!", tool_input as "tool_input!",
221 sequence_number as "sequence_number!"
222 FROM ai_request_tool_calls WHERE request_id = $1 ORDER BY sequence_number
223 OFFSET $2 LIMIT $3
224 "#,
225 request_id.as_str(),
226 page.offset,
227 page.sql_limit()
228 )
229 .fetch_all(&*self.pool)
230 .await?;
231
232 Ok(rows
233 .into_iter()
234 .map(|t| AuditToolCallRow {
235 tool_name: McpToolName::new(t.tool_name),
236 tool_input: t.tool_input,
237 sequence_number: t.sequence_number,
238 })
239 .collect())
240 }
241
242 pub async fn list_linked_mcp_calls(
243 &self,
244 request_id: &AiRequestId,
245 ) -> Result<Vec<LinkedMcpCall>> {
246 let rows = sqlx::query!(
247 r#"
248 SELECT
249 mte.tool_name as "tool_name!",
250 mte.server_name as "server_name!",
251 mte.status as "status!",
252 mte.execution_time_ms
253 FROM mcp_tool_executions mte
254 JOIN ai_request_tool_calls artc ON artc.mcp_execution_id = mte.mcp_execution_id
255 WHERE artc.request_id = $1
256 "#,
257 request_id.as_str()
258 )
259 .fetch_all(&*self.pool)
260 .await?;
261
262 Ok(rows
263 .into_iter()
264 .map(|r| LinkedMcpCall {
265 tool_name: McpToolName::new(r.tool_name),
266 server_name: McpServerId::new(r.server_name),
267 status: r.status,
268 execution_time_ms: r.execution_time_ms,
269 })
270 .collect())
271 }
272}