Skip to main content

systemprompt_agent/repository/task/
mutations.rs

1//! Task creation and the `agent_tasks` -> DB status string mapping.
2//!
3//! State transitions (with optimistic-concurrency guards) live in the
4//! sibling `state` submodule.
5
6use 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}