Skip to main content

systemprompt_api/routes/agent/
tasks.rs

1//! Agent task listing and lookup endpoints.
2//!
3//! Copyright (c) systemprompt.io — Business Source License 1.1.
4//! See <https://systemprompt.io> for licensing details.
5
6use axum::extract::{Path, Query, State};
7use axum::http::StatusCode;
8use axum::response::IntoResponse;
9use axum::{Extension, Json};
10use serde::Deserialize;
11use systemprompt_identifiers::TaskId;
12
13use systemprompt_agent::models::a2a::TaskState;
14use systemprompt_models::RequestContext;
15use systemprompt_models::api::ApiError;
16use systemprompt_runtime::AppContext;
17
18use crate::error::ApiHttpError;
19
20#[derive(Debug, Deserialize)]
21pub struct TaskFilterParams {
22    pub status: Option<String>,
23    pub limit: Option<u32>,
24}
25
26pub async fn list_tasks_by_context(
27    Extension(req_ctx): Extension<RequestContext>,
28    State(app_context): State<AppContext>,
29    Path(context_id): Path<String>,
30) -> Result<impl IntoResponse, ApiHttpError> {
31    tracing::debug!(context_id = %context_id, "Listing tasks");
32
33    let context_id_typed = super::parse_context_id(&context_id)?;
34
35    let context_repo = app_context.a2a_repositories().contexts.clone();
36    context_repo
37        .validate_context_ownership(&context_id_typed, req_ctx.user_id())
38        .await?;
39
40    let task_repo = app_context.a2a_repositories().tasks.clone();
41    let tasks = task_repo.list_tasks_by_context(&context_id_typed).await?;
42
43    tracing::debug!(context_id = %context_id, count = %tasks.len(), "Tasks listed");
44    Ok((StatusCode::OK, Json(tasks)))
45}
46
47pub async fn get_task(
48    Extension(req_ctx): Extension<RequestContext>,
49    State(app_context): State<AppContext>,
50    Path(task_id): Path<String>,
51) -> Result<impl IntoResponse, ApiHttpError> {
52    tracing::debug!(task_id = %task_id, "Retrieving task");
53
54    let task_repo = app_context.a2a_repositories().tasks.clone();
55
56    let task_id_typed = TaskId::try_new(&task_id).map_err(ApiError::from)?;
57    task_repo
58        .validate_task_ownership(&task_id_typed, req_ctx.user_id())
59        .await?;
60
61    let task = task_repo
62        .find_task(&task_id_typed)
63        .await?
64        .ok_or_else(|| ApiHttpError::not_found(format!("Task '{task_id}' not found")))?;
65
66    tracing::debug!("Task retrieved successfully");
67    Ok((StatusCode::OK, Json(task)))
68}
69
70pub async fn list_tasks_by_user(
71    Extension(req_ctx): Extension<RequestContext>,
72    State(app_context): State<AppContext>,
73    Query(params): Query<TaskFilterParams>,
74) -> Result<impl IntoResponse, ApiHttpError> {
75    let user_id = req_ctx.user_id();
76
77    tracing::debug!(user_id = %user_id, "Listing tasks");
78
79    let task_repo = app_context.a2a_repositories().tasks.clone();
80
81    let task_state = params.status.as_ref().and_then(|s| match s.as_str() {
82        "submitted" => Some(TaskState::Submitted),
83        "working" => Some(TaskState::Working),
84        "input-required" => Some(TaskState::InputRequired),
85        "completed" => Some(TaskState::Completed),
86        "canceled" | "cancelled" => Some(TaskState::Canceled),
87        "failed" => Some(TaskState::Failed),
88        "rejected" => Some(TaskState::Rejected),
89        "auth-required" => Some(TaskState::AuthRequired),
90        _ => None,
91    });
92
93    let limit = params.limit.map(|l| i32::try_from(l).unwrap_or(i32::MAX));
94    let mut tasks = task_repo.get_tasks_by_user_id(user_id, limit, None).await?;
95
96    if let Some(state) = task_state {
97        tasks.retain(|t| t.status.state == state);
98    }
99
100    tracing::debug!(user_id = %user_id, count = %tasks.len(), "Tasks listed");
101    Ok((StatusCode::OK, Json(tasks)))
102}
103
104pub async fn get_messages_by_task(
105    Extension(req_ctx): Extension<RequestContext>,
106    State(app_context): State<AppContext>,
107    Path(task_id): Path<String>,
108) -> Result<impl IntoResponse, ApiHttpError> {
109    tracing::debug!(task_id = %task_id, "Retrieving messages");
110
111    let task_repo = app_context.a2a_repositories().tasks.clone();
112
113    let task_id_typed = TaskId::try_new(&task_id).map_err(ApiError::from)?;
114    task_repo
115        .validate_task_ownership(&task_id_typed, req_ctx.user_id())
116        .await?;
117
118    let messages = task_repo.get_messages_by_task(&task_id_typed).await?;
119
120    tracing::debug!(task_id = %task_id, count = %messages.len(), "Messages retrieved");
121    Ok((StatusCode::OK, Json(messages)))
122}
123
124pub async fn delete_task(
125    Extension(req_ctx): Extension<RequestContext>,
126    State(app_context): State<AppContext>,
127    Path(task_id): Path<String>,
128) -> Result<impl IntoResponse, ApiHttpError> {
129    tracing::debug!(task_id = %task_id, "Deleting task");
130
131    let task_repo = app_context.a2a_repositories().tasks.clone();
132
133    let task_id_typed = TaskId::try_new(&task_id).map_err(ApiError::from)?;
134    task_repo
135        .validate_task_ownership(&task_id_typed, req_ctx.user_id())
136        .await?;
137
138    task_repo.delete_task(&task_id_typed).await?;
139
140    tracing::debug!(task_id = %task_id, "Task deleted");
141    Ok(StatusCode::NO_CONTENT)
142}