mod artifact_delta;
mod pool;
mod store_impl;
use a2a_protocol_types::error::A2aResult;
use pool::pg_pool;
pub(in crate::store) use pool::to_a2a_error;
use sqlx::postgres::PgPool;
#[derive(Debug, Clone)]
pub struct PostgresTaskStore {
pool: PgPool,
max_page_size: u32,
}
impl PostgresTaskStore {
#[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 = pg_pool(url).await?;
Self::from_pool(pool).await
}
pub async fn with_migrations(url: &str) -> Result<Self, sqlx::Error> {
let pool = pg_pool(url).await?;
let runner = super::pg_migration::PgMigrationRunner::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: PgPool) -> 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 JSONB NOT NULL,
created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
updated_at TIMESTAMPTZ NOT NULL DEFAULT now()
)",
)
.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,
})
}
pub async fn purge_expired(
&self,
policy: &super::retention::RetentionPolicy,
) -> A2aResult<super::retention::PurgeReport> {
super::retention::postgres::purge(&self.pool, "tasks", policy)
.await
.map_err(to_a2a_error)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn to_a2a_error_formats_message() {
let pg_err = sqlx::Error::RowNotFound;
let a2a_err = to_a2a_error(pg_err);
let msg = format!("{a2a_err}");
assert!(
msg.contains("postgres error"),
"error message should contain 'postgres error': {msg}"
);
}
}