Skip to main content

systemprompt_runtime/trace/repository/
audit.rs

1//! Audit-view queries joining messages, tool calls, and linked MCP rows.
2//!
3//! Copyright (c) systemprompt.io — Business Source License 1.1.
4//! See <https://systemprompt.io> for licensing details.
5
6use 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}