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//!
6//! Copyright (c) systemprompt.io — Business Source License 1.1.
7//! See <https://systemprompt.io> for licensing details.
8
9use 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}