systemprompt-agent 0.63.0

Agent-to-Agent (A2A) protocol for systemprompt.io AI governance: streaming, JSON-RPC models, task lifecycle, .well-known discovery, and governed agent orchestration.
Documentation
//! Optimistic-concurrency state transitions for `agent_tasks`.
//!
//! Each transition reads the current row `FOR UPDATE`, validates the move
//! against [`TaskState::can_transition_to`], and guards the write with a
//! version check so concurrent updates fail loudly rather than clobber.
//!
//! Copyright (c) systemprompt.io — Business Source License 1.1.
//! See <https://systemprompt.io> for licensing details.

use std::sync::Arc;
use systemprompt_models::errors::ParseEnumError;

use sqlx::PgPool;
use systemprompt_identifiers::TaskId;
use systemprompt_traits::RepositoryError;

use crate::models::a2a::TaskState;

pub async fn update_task_state(
    pool: &Arc<PgPool>,
    task_id: &TaskId,
    state: TaskState,
    timestamp: &chrono::DateTime<chrono::Utc>,
) -> Result<(), RepositoryError> {
    let mut tx = pool.begin().await?;
    transition_in_tx(&mut tx, task_id, state, timestamp).await?;
    tx.commit().await?;
    Ok(())
}

pub(super) async fn transition_in_tx(
    tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
    task_id: &TaskId,
    state: TaskState,
    timestamp: &chrono::DateTime<chrono::Utc>,
) -> Result<(), RepositoryError> {
    let task_id_str = task_id.as_str();
    let (current_state, expected_version) = lock_task_state(tx, task_id_str).await?;

    if current_state == state {
        return Ok(());
    }

    if !current_state.can_transition_to(&state) {
        return Err(RepositoryError::conflict(
            "task",
            task_id_str,
            format!("invalid state transition {current_state:?} -> {state:?}"),
        ));
    }

    let rows_affected =
        execute_state_update(tx, state, timestamp, task_id_str, expected_version).await?;

    if rows_affected == 0 {
        return Err(RepositoryError::conflict(
            "task",
            task_id_str,
            format!("stale update, expected version {expected_version}"),
        ));
    }

    Ok(())
}

async fn lock_task_state(
    tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
    task_id_str: &str,
) -> Result<(TaskState, i64), RepositoryError> {
    let current = sqlx::query!(
        r#"SELECT status, version FROM agent_tasks WHERE task_id = $1 FOR UPDATE"#,
        task_id_str
    )
    .fetch_optional(&mut **tx)
    .await?
    .ok_or_else(|| RepositoryError::not_found("task", task_id_str))?;

    let current_state: TaskState = current
        .status
        .parse()
        .map_err(|e: ParseEnumError| RepositoryError::decode("stored task state", e))?;

    Ok((current_state, current.version))
}

async fn execute_state_update(
    tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
    state: TaskState,
    timestamp: &chrono::DateTime<chrono::Utc>,
    task_id_str: &str,
    expected_version: i64,
) -> Result<u64, RepositoryError> {
    let status = state.as_str();

    let result = if state == TaskState::Completed {
        sqlx::query!(
            r#"UPDATE agent_tasks
               SET status = $1,
                   status_timestamp = $2,
                   updated_at = CURRENT_TIMESTAMP,
                   completed_at = CURRENT_TIMESTAMP,
                   started_at = COALESCE(started_at, CURRENT_TIMESTAMP),
                   execution_time_ms = EXTRACT(EPOCH FROM (CURRENT_TIMESTAMP - COALESCE(started_at, CURRENT_TIMESTAMP))) * 1000,
                   version = version + 1
               WHERE task_id = $3 AND version = $4"#,
            status,
            timestamp,
            task_id_str,
            expected_version
        )
        .execute(&mut **tx)
        .await
    } else if state == TaskState::Working {
        sqlx::query!(
            r#"UPDATE agent_tasks
               SET status = $1,
                   status_timestamp = $2,
                   updated_at = CURRENT_TIMESTAMP,
                   started_at = COALESCE(started_at, CURRENT_TIMESTAMP),
                   version = version + 1
               WHERE task_id = $3 AND version = $4"#,
            status,
            timestamp,
            task_id_str,
            expected_version
        )
        .execute(&mut **tx)
        .await
    } else {
        sqlx::query!(
            r#"UPDATE agent_tasks
               SET status = $1,
                   status_timestamp = $2,
                   updated_at = CURRENT_TIMESTAMP,
                   version = version + 1
               WHERE task_id = $3 AND version = $4"#,
            status,
            timestamp,
            task_id_str,
            expected_version
        )
        .execute(&mut **tx)
        .await
    };

    Ok(result?.rows_affected())
}

pub async fn apply_notification_status(
    pool: &Arc<PgPool>,
    task_id: &TaskId,
    state: &str,
    timestamp: &chrono::DateTime<chrono::Utc>,
) -> Result<(), RepositoryError> {
    let parsed: TaskState = state.parse().map_err(|_unknown: ParseEnumError| {
        RepositoryError::invalid_argument("state", format!("unknown task state {state:?}"))
    })?;
    update_task_state(pool, task_id, parsed, timestamp).await
}

pub async fn update_task_failed_with_error(
    pool: &Arc<PgPool>,
    task_id: &TaskId,
    error_message: &str,
    timestamp: &chrono::DateTime<chrono::Utc>,
) -> Result<(), RepositoryError> {
    let task_id_str = task_id.as_str();

    let mut tx = pool.begin().await?;

    let (current_state, expected_version) = lock_task_state(&mut tx, task_id_str).await?;

    if current_state == TaskState::Failed {
        tx.commit().await?;
        return Ok(());
    }

    if !current_state.can_transition_to(&TaskState::Failed) {
        return Err(RepositoryError::conflict(
            "task",
            task_id_str,
            format!("invalid state transition {current_state:?} -> Failed"),
        ));
    }

    let rows_affected = sqlx::query!(
        r#"UPDATE agent_tasks SET
            status = 'TASK_STATE_FAILED',
            status_timestamp = $1,
            error_message = $2,
            updated_at = CURRENT_TIMESTAMP,
            completed_at = CURRENT_TIMESTAMP,
            started_at = COALESCE(started_at, CURRENT_TIMESTAMP),
            execution_time_ms = EXTRACT(EPOCH FROM (CURRENT_TIMESTAMP - COALESCE(started_at, CURRENT_TIMESTAMP))) * 1000,
            version = version + 1
        WHERE task_id = $3 AND version = $4"#,
        timestamp,
        error_message,
        task_id_str,
        expected_version
    )
    .execute(&mut *tx)
    .await?
    .rows_affected();

    if rows_affected == 0 {
        return Err(RepositoryError::conflict(
            "task",
            task_id_str,
            format!("stale update, expected version {expected_version}"),
        ));
    }

    tx.commit().await?;
    Ok(())
}