systemprompt_agent/repository/task/
mutations.rs1use sqlx::PgPool;
10use std::sync::Arc;
11use systemprompt_traits::RepositoryError;
12
13use crate::models::a2a::{Task, TaskState};
14
15pub const fn task_state_to_db_string(state: TaskState) -> &'static str {
16 match state {
17 TaskState::Pending => "TASK_STATE_PENDING",
18 TaskState::Submitted => "TASK_STATE_SUBMITTED",
19 TaskState::Working => "TASK_STATE_WORKING",
20 TaskState::InputRequired => "TASK_STATE_INPUT_REQUIRED",
21 TaskState::Completed => "TASK_STATE_COMPLETED",
22 TaskState::Canceled => "TASK_STATE_CANCELED",
23 TaskState::Failed => "TASK_STATE_FAILED",
24 TaskState::Rejected => "TASK_STATE_REJECTED",
25 TaskState::AuthRequired => "TASK_STATE_AUTH_REQUIRED",
26 TaskState::Unknown => "TASK_STATE_UNKNOWN",
27 }
28}
29
30#[expect(
31 missing_debug_implementations,
32 reason = "params struct holds non-Debug references"
33)]
34pub struct CreateTaskParams<'a> {
35 pub pool: &'a Arc<PgPool>,
36 pub task: &'a Task,
37 pub user_id: &'a systemprompt_identifiers::UserId,
38 pub session_id: &'a systemprompt_identifiers::SessionId,
39 pub trace_id: &'a systemprompt_identifiers::TraceId,
40 pub agent_name: &'a str,
41}
42
43pub async fn create_task(params: CreateTaskParams<'_>) -> Result<String, RepositoryError> {
44 let CreateTaskParams {
45 pool,
46 task,
47 user_id,
48 session_id,
49 trace_id,
50 agent_name,
51 } = params;
52 let metadata_json = task.metadata.as_ref().map_or_else(
53 || serde_json::json!({}),
54 |m| {
55 serde_json::to_value(m).unwrap_or_else(|e| {
56 tracing::warn!(error = %e, task_id = %task.id, "Failed to serialize task metadata");
57 serde_json::json!({})
58 })
59 },
60 );
61
62 let status = task_state_to_db_string(task.status.state);
63 let task_id_str = task.id.as_str();
64 let context_id_str = task.context_id.as_str();
65 let user_id_str = user_id.as_ref();
66 let session_id_str = session_id.as_ref();
67 let trace_id_str = trace_id.as_ref();
68
69 sqlx::query!(
70 r#"INSERT INTO agent_tasks (task_id, context_id, status, status_timestamp, user_id, session_id, trace_id, metadata, agent_name)
71 VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9)"#,
72 task_id_str,
73 context_id_str,
74 status,
75 task.status.timestamp,
76 user_id_str,
77 session_id_str,
78 trace_id_str,
79 metadata_json,
80 agent_name
81 )
82 .execute(pool.as_ref())
83 .await
84 .map_err(RepositoryError::database)?;
85
86 Ok(task.id.to_string())
87}
88
89pub async fn track_agent_in_context(
90 pool: &Arc<PgPool>,
91 context_id: &systemprompt_identifiers::ContextId,
92 agent_name: &str,
93) -> Result<(), RepositoryError> {
94 let context_id_str = context_id.as_str();
95 sqlx::query!(
96 r#"INSERT INTO context_agents (context_id, agent_name) VALUES ($1, $2)
97 ON CONFLICT (context_id, agent_name) DO NOTHING"#,
98 context_id_str,
99 agent_name
100 )
101 .execute(pool.as_ref())
102 .await
103 .map_err(RepositoryError::database)?;
104
105 Ok(())
106}