use std::borrow::Cow;
use std::future::Future;
use std::pin::Pin;
use a2a_protocol_types::error::{A2aError, A2aResult};
use a2a_protocol_types::params::ListTasksParams;
use a2a_protocol_types::responses::TaskListResponse;
use a2a_protocol_types::task::{Task, TaskId};
use sqlx::sqlite::SqlitePool;
use super::task_store::{ArtifactDelta, TaskStore};
#[derive(Debug, Clone)]
pub struct SqliteTaskStore {
pool: SqlitePool,
max_page_size: u32,
}
impl SqliteTaskStore {
#[must_use]
pub const fn with_max_page_size(mut self, max: u32) -> Self {
self.max_page_size = max;
self
}
pub async fn new(url: &str) -> Result<Self, sqlx::Error> {
let pool = sqlite_pool(url).await?;
Self::from_pool(pool).await
}
pub async fn with_migrations(url: &str) -> Result<Self, sqlx::Error> {
let pool = sqlite_pool(url).await?;
let runner = super::migration::MigrationRunner::new(pool.clone());
runner.run_pending().await?;
Ok(Self {
pool,
max_page_size: crate::store::DEFAULT_MAX_PAGE_SIZE,
})
}
pub async fn from_pool(pool: SqlitePool) -> Result<Self, sqlx::Error> {
sqlx::query(
"CREATE TABLE IF NOT EXISTS tasks (
id TEXT PRIMARY KEY,
context_id TEXT NOT NULL,
state TEXT NOT NULL,
data TEXT NOT NULL,
updated_at TEXT NOT NULL DEFAULT (strftime('%Y-%m-%d %H:%M:%f','now')),
created_at TEXT NOT NULL DEFAULT (datetime('now'))
)",
)
.execute(&pool)
.await?;
sqlx::query(journal::CREATE_TABLE_SQL)
.execute(&pool)
.await?;
sqlx::query("CREATE INDEX IF NOT EXISTS idx_tasks_context_id ON tasks(context_id)")
.execute(&pool)
.await?;
sqlx::query("CREATE INDEX IF NOT EXISTS idx_tasks_state ON tasks(state)")
.execute(&pool)
.await?;
sqlx::query(
"CREATE INDEX IF NOT EXISTS idx_tasks_context_id_state ON tasks(context_id, state)",
)
.execute(&pool)
.await?;
sqlx::query(
"CREATE INDEX IF NOT EXISTS idx_tasks_updated_at ON tasks(updated_at DESC, id DESC)",
)
.execute(&pool)
.await?;
Ok(Self {
pool,
max_page_size: crate::store::DEFAULT_MAX_PAGE_SIZE,
})
}
async fn journal_append(&self, task: &Task, rows: Vec<journal::Row>) -> A2aResult<()> {
let mut query_builder = sqlx::QueryBuilder::new(
"INSERT INTO task_artifact_appends (task_id, artifact, seq, part) ",
);
query_builder.push_values(rows, |mut b, (artifact, seq, part)| {
b.push_bind(task.id.0.clone())
.push_bind(artifact)
.push_bind(seq)
.push_bind(part);
});
query_builder.push(" ON CONFLICT(task_id, artifact, seq) DO NOTHING");
if query_builder.build().execute(&self.pool).await.is_err() {
return self.save(task).await;
}
Ok(())
}
pub async fn purge_expired(
&self,
policy: &super::retention::RetentionPolicy,
) -> A2aResult<super::retention::PurgeReport> {
super::retention::sqlite::purge(&self.pool, "tasks", Some("task_artifact_appends"), policy)
.await
.map_err(to_a2a_error)
}
}
use crate::sqlite_pool::sqlite_pool;
pub(super) mod journal;
#[allow(clippy::needless_pass_by_value)]
fn to_a2a_error(e: sqlx::Error) -> A2aError {
A2aError::internal(format!("sqlite error: {e}"))
}
const MAX_INLINE_APPEND: usize = 8;
fn artifact_delta_sql(task: &Task, delta: ArtifactDelta) -> A2aResult<Option<DeltaStatement>> {
let Some(artifacts) = task.artifacts.as_ref() else {
return Ok(None);
};
match delta {
ArtifactDelta::AppendedParts { index, count } => {
if count == 0 || count > MAX_INLINE_APPEND {
return Ok(None);
}
let Some(artifact) = artifacts.get(index) else {
return Ok(None);
};
if artifact.parts.len() < count {
return Ok(None);
}
let tail = &artifact.parts[artifact.parts.len() - count..];
let payload = serde_json::to_string(tail)
.map_err(|e| A2aError::internal(format!("failed to serialize parts: {e}")))?;
if count == 1 {
return Ok(Some(DeltaStatement {
sql: APPEND_ONE_PART_SQL,
payload,
index: Some(index),
}));
}
let exprs = (0..count)
.map(|i| format!("'$.artifacts[{index}].parts[#]', json_extract(?1, '$[{i}]')"))
.collect::<Vec<_>>()
.join(", ");
Ok(Some(DeltaStatement {
sql: Cow::Owned(format!(
"UPDATE tasks SET data = json_set(data, {exprs}) \
WHERE id = ?2 AND json_type(data, '$.artifacts') = 'array'"
)),
payload,
index: None,
}))
}
ArtifactDelta::Pushed { index } => {
if index + 1 != artifacts.len() {
return Ok(None);
}
let Some(artifact) = artifacts.get(index) else {
return Ok(None);
};
let payload = serde_json::to_string(std::slice::from_ref(artifact))
.map_err(|e| A2aError::internal(format!("failed to serialize artifact: {e}")))?;
Ok(Some(DeltaStatement {
sql: PUSH_ARTIFACT_SQL,
payload,
index: None,
}))
}
}
}
struct DeltaStatement {
sql: Cow<'static, str>,
payload: String,
index: Option<usize>,
}
const APPEND_ONE_PART_SQL: Cow<'static, str> = Cow::Borrowed(
"UPDATE tasks SET data = json_set(data, '$.artifacts[' || ?3 || '].parts[#]', \
json_extract(?1, '$[0]')) \
WHERE id = ?2 AND json_type(data, '$.artifacts') = 'array'",
);
const PUSH_ARTIFACT_SQL: Cow<'static, str> = Cow::Borrowed(
"UPDATE tasks SET data = json_set(data, '$.artifacts[#]', json_extract(?1, '$[0]')) \
WHERE id = ?2 AND json_type(data, '$.artifacts') = 'array'",
);
#[allow(clippy::manual_async_fn)]
impl TaskStore for SqliteTaskStore {
fn save<'a>(
&'a self,
task: &'a Task,
) -> Pin<Box<dyn Future<Output = A2aResult<()>> + Send + 'a>> {
Box::pin(async move {
let id = task.id.0.as_str();
let context_id = task.context_id.0.as_str();
let state = task.status.state.to_string();
let data = serde_json::to_string(task)
.map_err(|e| A2aError::internal(format!("failed to serialize task: {e}")))?;
let status_ts = super::status_timestamp_sqlite(task.status.timestamp.as_deref());
let mut tx = self.pool.begin().await.map_err(to_a2a_error)?;
sqlx::query(
"INSERT INTO tasks (id, context_id, state, data, updated_at)
VALUES (?1, ?2, ?3, ?4, COALESCE(?5, strftime('%Y-%m-%d %H:%M:%f','now')))
ON CONFLICT(id) DO UPDATE SET
context_id = excluded.context_id,
state = excluded.state,
data = excluded.data,
updated_at = excluded.updated_at",
)
.bind(id)
.bind(context_id)
.bind(&state)
.bind(&data)
.bind(&status_ts)
.execute(&mut *tx)
.await
.map_err(to_a2a_error)?;
sqlx::query(journal::DELETE_FOR_TASK_SQL)
.bind(id)
.execute(&mut *tx)
.await
.map_err(to_a2a_error)?;
tx.commit().await.map_err(to_a2a_error)?;
Ok(())
})
}
fn save_artifact_delta<'a>(
&'a self,
task: &'a Task,
delta: ArtifactDelta,
) -> Pin<Box<dyn Future<Output = A2aResult<()>> + Send + 'a>> {
Box::pin(async move {
if let ArtifactDelta::AppendedParts { index, count } = delta {
let Some(rows) = journal::rows_for_append(task, index, count)? else {
return self.save(task).await;
};
return self.journal_append(task, rows).await;
}
let Some(stmt) = artifact_delta_sql(task, delta)? else {
return self.save(task).await;
};
let mut query = sqlx::query(stmt.sql.as_ref())
.bind(&stmt.payload)
.bind(task.id.0.as_str());
if let Some(index) = stmt.index {
query = query.bind(i64::try_from(index).unwrap_or(i64::MAX));
}
let affected = query
.execute(&self.pool)
.await
.map_err(to_a2a_error)?
.rows_affected();
if affected == 0 {
return self.save(task).await;
}
Ok(())
})
}
fn get<'a>(
&'a self,
id: &'a TaskId,
) -> Pin<Box<dyn Future<Output = A2aResult<Option<Task>>> + Send + 'a>> {
Box::pin(async move {
let row: Option<(String,)> = sqlx::query_as("SELECT data FROM tasks WHERE id = ?1")
.bind(id.0.as_str())
.fetch_optional(&self.pool)
.await
.map_err(to_a2a_error)?;
match row {
Some((data,)) => {
let mut task: Task = serde_json::from_str(&data).map_err(|e| {
A2aError::internal(format!("failed to deserialize task: {e}"))
})?;
let rows: Vec<journal::Row> = sqlx::query_as(journal::SELECT_FOR_TASK_SQL)
.bind(id.0.as_str())
.fetch_all(&self.pool)
.await
.map_err(to_a2a_error)?;
journal::splice(&mut task, rows)?;
Ok(Some(task))
}
None => Ok(None),
}
})
}
#[allow(clippy::too_many_lines)]
fn list<'a>(
&'a self,
params: &'a ListTasksParams,
) -> Pin<Box<dyn Future<Output = A2aResult<TaskListResponse>> + Send + 'a>> {
Box::pin(async move {
let mut conditions = Vec::new();
let mut bind_values: Vec<String> = Vec::new();
if let Some(ref ctx) = params.context_id {
conditions.push(format!("context_id = ?{}", bind_values.len() + 1));
bind_values.push(ctx.clone());
}
if let Some(ref status) = params.status {
conditions.push(format!("state = ?{}", bind_values.len() + 1));
bind_values.push(status.to_string());
}
if let Some(ref after) = params.status_timestamp_after {
let Some(after_dt) = super::status_timestamp_sqlite(Some(after)) else {
return Ok(TaskListResponse::new(Vec::new()));
};
conditions.push(format!("updated_at > ?{}", bind_values.len() + 1));
bind_values.push(after_dt);
}
if let Some(ref token) = params.page_token {
let Some((cursor_ua, cursor_id)) = super::cursor::decode(token) else {
return Ok(TaskListResponse::new(Vec::new()));
};
let p = bind_values.len();
conditions.push(format!("(updated_at, id) < (?{}, ?{})", p + 1, p + 2));
bind_values.push(cursor_ua.to_string());
bind_values.push(cursor_id.to_string());
}
let where_clause = if conditions.is_empty() {
String::new()
} else {
format!("WHERE {}", conditions.join(" AND "))
};
let page_size = match params.page_size {
Some(0) | None => 50_u32,
Some(n) => n.min(self.max_page_size),
};
let limit = super::pagination::fetch_limit(page_size);
let limit_param = bind_values.len() + 1;
let sql = format!(
"SELECT updated_at, data FROM tasks {where_clause} \
ORDER BY updated_at DESC, id DESC LIMIT ?{limit_param}"
);
let mut query = sqlx::query_as::<_, (String, String)>(&sql);
for val in &bind_values {
query = query.bind(val);
}
query = query.bind(limit);
let rows: Vec<(String, String)> =
query.fetch_all(&self.pool).await.map_err(to_a2a_error)?;
let mut rows: Vec<(String, Task)> = rows
.into_iter()
.map(|(updated_at, data)| {
serde_json::from_str::<Task>(&data)
.map(|task| (updated_at, task))
.map_err(|e| A2aError::internal(format!("deserialize: {e}")))
})
.collect::<A2aResult<Vec<_>>>()?;
if !rows.is_empty() {
let mut journal_query = sqlx::QueryBuilder::new(
"SELECT task_id, artifact, seq, part FROM task_artifact_appends WHERE task_id IN (",
);
let mut separated = journal_query.separated(", ");
for (_, task) in &rows {
separated.push_bind(task.id.0.clone());
}
journal_query.push(") ORDER BY task_id, artifact, seq");
let journalled: Vec<(String, i64, i64, String)> = journal_query
.build_query_as()
.fetch_all(&self.pool)
.await
.map_err(to_a2a_error)?;
if !journalled.is_empty() {
let mut by_task: std::collections::HashMap<String, Vec<journal::Row>> =
std::collections::HashMap::new();
for (task_id, artifact, seq, part) in journalled {
by_task
.entry(task_id)
.or_default()
.push((artifact, seq, part));
}
for (_, task) in &mut rows {
if let Some(task_rows) = by_task.remove(task.id.0.as_str()) {
journal::splice(task, task_rows)?;
}
}
}
}
let next_page_token =
if super::pagination::has_next_page(rows.len(), page_size as usize) {
rows.truncate(page_size as usize);
rows.last()
.map(|(ua, task)| super::cursor::encode(ua, task.id.0.as_str()))
.unwrap_or_default()
} else {
String::new()
};
#[allow(clippy::cast_possible_truncation)]
let page_len = rows.len() as u32;
let tasks: Vec<Task> = rows.into_iter().map(|(_, task)| task).collect();
let mut response = TaskListResponse::new(tasks);
response.next_page_token = next_page_token;
response.page_size = page_len;
Ok(response)
})
}
fn insert_if_absent<'a>(
&'a self,
task: &'a Task,
) -> Pin<Box<dyn Future<Output = A2aResult<bool>> + Send + 'a>> {
Box::pin(async move {
let id = task.id.0.as_str();
let context_id = task.context_id.0.as_str();
let state = task.status.state.to_string();
let data = serde_json::to_string(task)
.map_err(|e| A2aError::internal(format!("failed to serialize task: {e}")))?;
let status_ts = super::status_timestamp_sqlite(task.status.timestamp.as_deref());
let result = sqlx::query(
"INSERT OR IGNORE INTO tasks (id, context_id, state, data, updated_at)
VALUES (?1, ?2, ?3, ?4, COALESCE(?5, strftime('%Y-%m-%d %H:%M:%f','now')))",
)
.bind(id)
.bind(context_id)
.bind(&state)
.bind(&data)
.bind(&status_ts)
.execute(&self.pool)
.await
.map_err(to_a2a_error)?;
Ok(result.rows_affected() > 0)
})
}
fn delete<'a>(
&'a self,
id: &'a TaskId,
) -> Pin<Box<dyn Future<Output = A2aResult<()>> + Send + 'a>> {
Box::pin(async move {
sqlx::query(journal::DELETE_FOR_TASK_SQL)
.bind(id.0.as_str())
.execute(&self.pool)
.await
.map_err(to_a2a_error)?;
sqlx::query("DELETE FROM tasks WHERE id = ?1")
.bind(id.0.as_str())
.execute(&self.pool)
.await
.map_err(to_a2a_error)?;
Ok(())
})
}
fn count<'a>(&'a self) -> Pin<Box<dyn Future<Output = A2aResult<u64>> + Send + 'a>> {
Box::pin(async move {
let row: (i64,) = sqlx::query_as("SELECT COUNT(*) FROM tasks")
.fetch_one(&self.pool)
.await
.map_err(to_a2a_error)?;
#[allow(clippy::cast_sign_loss)]
Ok(row.0 as u64)
})
}
}
#[cfg(test)]
mod artifact_delta_tests;
#[cfg(test)]
mod retention_tests;
#[cfg(test)]
mod tests;