Skip to main content

systemprompt_runtime/trace/repository/
step.rs

1//! Execution-step queries for trace views.
2//!
3//! Copyright (c) systemprompt.io — Business Source License 1.1.
4//! See <https://systemprompt.io> for licensing details.
5
6use serde_json::json;
7
8use systemprompt_identifiers::{ContextId, SessionId, TaskId, TraceId, UserId};
9
10use super::{Result, TraceRepository};
11use crate::trace::models::{ExecutionStepSummary, McpExecutionSummary, TraceEvent};
12
13impl TraceRepository {
14    pub async fn fetch_mcp_execution_summary(
15        &self,
16        trace_id: &TraceId,
17    ) -> Result<McpExecutionSummary> {
18        let row = sqlx::query!(
19            r#"
20        SELECT
21            COUNT(*)::bigint as execution_count,
22            COALESCE(SUM(execution_time_ms), 0)::bigint as total_execution_time_ms
23        FROM mcp_tool_executions
24        WHERE trace_id = $1
25        "#,
26            trace_id.as_str()
27        )
28        .fetch_one(&*self.pool)
29        .await?;
30
31        Ok(McpExecutionSummary {
32            execution_count: row.execution_count.unwrap_or(0),
33            total_execution_time_ms: row.total_execution_time_ms.unwrap_or(0),
34        })
35    }
36
37    pub async fn fetch_mcp_execution_events(&self, trace_id: &TraceId) -> Result<Vec<TraceEvent>> {
38        let rows = sqlx::query!(
39            r#"
40        SELECT
41            started_at as timestamp,
42            tool_name,
43            server_name,
44            execution_time_ms,
45            status,
46            error_message,
47            user_id,
48            session_id,
49            task_id,
50            context_id
51        FROM mcp_tool_executions
52        WHERE trace_id = $1
53        ORDER BY started_at ASC
54        "#,
55            trace_id.as_str()
56        )
57        .fetch_all(&*self.pool)
58        .await?;
59
60        Ok(rows
61        .into_iter()
62        .map(|row| {
63            let base_details = format!(
64                "{}/{}: {} ({}ms)",
65                row.server_name,
66                row.tool_name,
67                row.status,
68                row.execution_time_ms.unwrap_or(0)
69            );
70
71            let details = if row.status == "failed" {
72                if let Some(ref error) = row.error_message {
73                    let truncated_error = if error.len() > 80 {
74                        format!("{}...", &error[..80])
75                    } else {
76                        error.clone()
77                    };
78                    format!("{} | {}", base_details, truncated_error)
79                } else {
80                    base_details
81                }
82            } else {
83                base_details
84            };
85
86            let metadata = json!({
87                "execution_time_ms": row.execution_time_ms,
88                "tool_name": row.tool_name,
89                "server_name": row.server_name,
90                "error_message": row.error_message
91            });
92
93            TraceEvent {
94                event_type: "MCP".to_owned(),
95                timestamp: row.timestamp,
96                details,
97                user_id: Some(UserId::new(row.user_id)),
98                session_id: row.session_id.map(SessionId::new),
99                task_id: row.task_id.map(TaskId::new),
100                context_id: row.context_id.and_then(|s| {
101                    ContextId::try_new(&s)
102                        .map_err(|e| {
103                            tracing::warn!(error = %e, raw = %s, "Skipping non-UUID context_id from mcp_tool_executions row");
104                            e
105                        })
106                        .ok()
107                }),
108                metadata: Some(metadata.to_string()),
109            }
110        })
111        .collect())
112    }
113
114    pub async fn fetch_task_id_for_trace(&self, trace_id: &TraceId) -> Result<Option<String>> {
115        let row = sqlx::query!(
116            "SELECT task_id FROM agent_tasks WHERE trace_id = $1 LIMIT 1",
117            trace_id.as_str()
118        )
119        .fetch_optional(&*self.pool)
120        .await?;
121
122        Ok(row.map(|r| r.task_id))
123    }
124
125    pub async fn fetch_execution_step_summary(
126        &self,
127        trace_id: &TraceId,
128    ) -> Result<ExecutionStepSummary> {
129        let row = sqlx::query!(
130        r#"
131        SELECT
132            COUNT(*)::bigint as step_count,
133            COUNT(*) FILTER (WHERE s.status = 'completed')::bigint as completed_count,
134            COUNT(*) FILTER (WHERE s.status = 'failed')::bigint as failed_count,
135            COUNT(*) FILTER (WHERE s.status = 'pending' OR s.status = 'in_progress')::bigint as pending_count
136        FROM task_execution_steps s
137        JOIN agent_tasks t ON s.task_id = t.task_id
138        WHERE t.trace_id = $1
139        "#,
140        trace_id.as_str()
141    )
142    .fetch_one(&*self.pool)
143    .await?;
144
145        Ok(ExecutionStepSummary {
146            total: row.step_count.unwrap_or(0),
147            completed: row.completed_count.unwrap_or(0),
148            failed: row.failed_count.unwrap_or(0),
149            pending: row.pending_count.unwrap_or(0),
150        })
151    }
152
153    pub async fn fetch_execution_step_events(&self, trace_id: &TraceId) -> Result<Vec<TraceEvent>> {
154        let rows = sqlx::query!(
155            r#"
156        SELECT
157            s.started_at as timestamp,
158            s.content,
159            s.status,
160            s.duration_ms,
161            t.user_id,
162            t.session_id,
163            t.task_id,
164            t.context_id
165        FROM task_execution_steps s
166        JOIN agent_tasks t ON s.task_id = t.task_id
167        WHERE t.trace_id = $1
168        ORDER BY s.started_at ASC
169        "#,
170            trace_id.as_str()
171        )
172        .fetch_all(&*self.pool)
173        .await?;
174
175        Ok(rows
176        .into_iter()
177        .map(|row| {
178            let content = row.content.clone();
179            let step_type = content
180                .get("type")
181                .and_then(|v| v.as_str())
182                .unwrap_or("unknown");
183            let tool_name = content.get("tool_name").and_then(|v| v.as_str());
184            let skill_name = content.get("skill_name").and_then(|v| v.as_str());
185
186            let details = match step_type {
187                "understanding" => format!("[{}] Analyzing request... - {}", step_type, row.status),
188                "planning" => format!("[{}] Planning response... - {}", step_type, row.status),
189                "skill_usage" => {
190                    let name = skill_name.unwrap_or("unknown");
191                    format!("[{}] Using {} skill... - {}", step_type, name, row.status)
192                },
193                "tool_execution" => {
194                    let name = tool_name.unwrap_or("unknown");
195                    format!("[{}] Running {}... - {}", step_type, name, row.status)
196                },
197                "completion" => format!("[{}] Complete - {}", step_type, row.status),
198                _ => format!("[{}] - {}", step_type, row.status),
199            };
200
201            let metadata = json!({
202                "step_type": step_type,
203                "status": row.status,
204                "duration_ms": row.duration_ms,
205                "tool_name": tool_name,
206                "skill_name": skill_name
207            });
208
209            TraceEvent {
210                event_type: "STEP".to_owned(),
211                timestamp: row.timestamp,
212                details,
213                user_id: row.user_id.map(UserId::new),
214                session_id: row.session_id.map(SessionId::new),
215                task_id: Some(TaskId::new(row.task_id)),
216                context_id: ContextId::try_new(&row.context_id)
217                    .map_err(|e| {
218                        tracing::warn!(error = %e, raw = %row.context_id, "Skipping non-UUID context_id from task_execution_steps join row");
219                        e
220                    })
221                    .ok(),
222                metadata: Some(metadata.to_string()),
223            }
224        })
225        .collect())
226    }
227}