systemprompt_runtime/trace/repository/
step.rs1use 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}