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