#[cfg(feature = "sqlx-storage")]
use async_trait::async_trait;
#[cfg(feature = "sqlx-storage")]
use serde_json;
#[cfg(feature = "sqlx-storage")]
use sqlx::{
AnyPool, ConnectOptions, Row,
any::{AnyConnectOptions, AnyPoolOptions},
};
#[cfg(feature = "sqlx-storage")]
use std::{
collections::HashMap,
str::FromStr,
time::{Duration, Instant},
};
#[cfg(feature = "sqlx-storage")]
use crate::adapter::business::push_notification::{
PushNotificationRegistry, PushNotificationSender,
};
#[cfg(feature = "sqlx-storage")]
#[cfg(feature = "http-client")]
use crate::adapter::business::push_notification::HttpPushNotificationSender;
#[cfg(feature = "sqlx-storage")]
#[cfg(not(feature = "http-client"))]
use crate::adapter::business::push_notification::NoopPushNotificationSender;
#[cfg(feature = "sqlx-storage")]
use crate::domain::{
A2AError, ContextId, ContextState, Conversation, Digest, Message, ReadRefresh, Remembered,
RetentionPolicy, Seq, SequencedMessage, StateKey, StateScope, Swept, Task, TaskId,
TaskPushNotificationConfig, TaskState, TaskStateExt, TaskStatus, VersionedTask,
};
#[cfg(feature = "sqlx-storage")]
use crate::port::{
AsyncContextStateStore, AsyncConversationStore, AsyncEventLog, AsyncNotificationManager,
AsyncPushNotifier, AsyncRetention, AsyncTaskLifecycle, AsyncTaskQuery, AsyncTaskVersioning,
Replay, SeqEvent, UpdateEvent, context_state::scope_key,
};
#[cfg(feature = "sqlx-storage")]
use std::sync::Arc;
#[cfg(feature = "sqlx-storage")]
pub struct SqlxTaskStorage {
pool: AnyPool,
dialect: Dialect,
push_notification_registry: Arc<PushNotificationRegistry>,
event_log_capacity: Option<u64>,
read_refresh: ReadRefresh,
claim_cache: Option<Arc<ClaimCache>>,
}
#[cfg(feature = "sqlx-storage")]
use super::database_config::DatabaseType;
#[cfg(feature = "sqlx-storage")]
use super::dialect::Dialect;
#[cfg(feature = "sqlx-storage")]
#[cfg(feature = "sqlx-storage")]
fn state_str(state: TaskState) -> &'static str {
match state {
TaskState::Submitted => "submitted",
TaskState::Working => "working",
TaskState::InputRequired => "input-required",
TaskState::Completed => "completed",
TaskState::Canceled => "canceled",
TaskState::Failed => "failed",
TaskState::Rejected => "rejected",
TaskState::AuthRequired => "auth-required",
TaskState::Unknown => "unknown",
}
}
#[cfg(feature = "sqlx-storage")]
fn status_message_json(message: Option<&Message>) -> Option<String> {
message.map(|m| serde_json::to_string(m).unwrap_or_default())
}
const TASK_COLUMNS: &str = "id, context_id, status_state, status_message, metadata, artifacts";
#[cfg(feature = "sqlx-storage")]
pub const DEFAULT_EVENT_LOG_CAPACITY: u64 = 1024;
#[cfg(feature = "sqlx-storage")]
pub const DEFAULT_CLAIM_CACHE_TTL: Duration = Duration::from_secs(5);
#[cfg(feature = "sqlx-storage")]
const CLAIM_CACHE_CAPACITY: usize = 1024;
#[cfg(feature = "sqlx-storage")]
struct ClaimCache {
entries: std::sync::Mutex<HashMap<String, (ContextClaim, Instant)>>,
ttl: Duration,
}
#[cfg(feature = "sqlx-storage")]
impl ClaimCache {
fn new(ttl: Duration) -> Self {
Self {
entries: std::sync::Mutex::new(HashMap::new()),
ttl,
}
}
fn get(&self, context_id: &str) -> Option<ContextClaim> {
let entries = self.entries.lock().ok()?;
let (claim, at) = entries.get(context_id)?;
(at.elapsed() < self.ttl).then(|| claim.clone())
}
fn insert(&self, context_id: &str, claim: &ContextClaim) {
let Ok(mut entries) = self.entries.lock() else {
return;
};
if entries.len() >= CLAIM_CACHE_CAPACITY {
entries.clear();
}
entries.insert(context_id.to_string(), (claim.clone(), Instant::now()));
}
fn forget(&self, context_id: &str) {
if let Ok(mut entries) = self.entries.lock() {
entries.remove(context_id);
}
}
}
#[cfg(feature = "sqlx-storage")]
struct PoolSettings {
max_connections: u32,
acquire_timeout: Duration,
log_statements: bool,
}
#[cfg(feature = "sqlx-storage")]
impl Default for PoolSettings {
fn default() -> Self {
Self {
max_connections: 10,
acquire_timeout: Duration::from_secs(30),
log_statements: false,
}
}
}
#[cfg(feature = "sqlx-storage")]
impl PoolSettings {
fn connect_options(&self, url: &str) -> Result<AnyConnectOptions, A2AError> {
let options = AnyConnectOptions::from_str(url)
.map_err(|e| A2AError::DatabaseError(format!("Invalid database URL '{url}': {e}")))?;
Ok(if self.log_statements {
options
} else {
options.disable_statement_logging()
})
}
}
#[cfg(feature = "sqlx-storage")]
pub struct SqlxStorageBuilder {
url: String,
pool: PoolSettings,
push_sender: Option<Arc<dyn PushNotificationSender>>,
additional_migrations: Vec<String>,
event_log_capacity: Option<u64>,
read_refresh: ReadRefresh,
claim_cache_ttl: Option<Duration>,
}
#[cfg(feature = "sqlx-storage")]
impl SqlxStorageBuilder {
pub fn from_config(config: &super::database_config::DatabaseConfig) -> Self {
SqlxTaskStorage::builder(&config.url)
.max_connections(config.max_connections)
.acquire_timeout(Duration::from_secs(config.timeout_seconds))
.log_statements(config.enable_logging)
}
pub fn max_connections(mut self, max: u32) -> Self {
self.pool.max_connections = max;
self
}
pub fn acquire_timeout(mut self, timeout: Duration) -> Self {
self.pool.acquire_timeout = timeout;
self
}
pub fn log_statements(mut self, log: bool) -> Self {
self.pool.log_statements = log;
self
}
pub fn push_sender(mut self, sender: impl PushNotificationSender + 'static) -> Self {
self.push_sender = Some(Arc::new(sender));
self
}
pub fn event_log_capacity(mut self, capacity: Option<u64>) -> Self {
self.event_log_capacity = capacity.map(|capacity| capacity.max(1));
self
}
pub fn claim_cache(mut self, ttl: Option<Duration>) -> Self {
self.claim_cache_ttl = ttl;
self
}
pub fn read_refresh(mut self, read_refresh: ReadRefresh) -> Self {
self.read_refresh = read_refresh;
self
}
pub fn migrations<S: AsRef<str>>(mut self, migrations: impl IntoIterator<Item = S>) -> Self {
self.additional_migrations
.extend(migrations.into_iter().map(|s| s.as_ref().to_string()));
self
}
pub async fn connect(self) -> Result<SqlxTaskStorage, A2AError> {
if self.pool.max_connections == 0 {
return Err(A2AError::DatabaseError(
"max_connections must be greater than 0; a pool that hands out no connections \
fails every query"
.to_string(),
));
}
let (pool, dialect) = SqlxTaskStorage::connect(&self.url, &self.pool).await?;
SqlxTaskStorage::run_additional_migrations(&pool, &self.additional_migrations).await?;
let push_registry = match self.push_sender {
Some(sender) => PushNotificationRegistry::from_shared(sender),
None => {
#[cfg(feature = "http-client")]
let sender = HttpPushNotificationSender::new();
#[cfg(not(feature = "http-client"))]
let sender = NoopPushNotificationSender::default();
PushNotificationRegistry::new(sender)
}
};
Ok(SqlxTaskStorage {
pool,
dialect,
push_notification_registry: Arc::new(push_registry),
event_log_capacity: self.event_log_capacity,
read_refresh: self.read_refresh,
claim_cache: self
.claim_cache_ttl
.map(|ttl| Arc::new(ClaimCache::new(ttl))),
})
}
}
#[cfg(feature = "sqlx-storage")]
#[derive(Clone)]
enum ContextClaim {
Open,
Owner(String),
}
#[cfg(feature = "sqlx-storage")]
impl ContextClaim {
fn verdict(&self, context_id: &str, caller: Option<&str>) -> Result<(), A2AError> {
match self {
Self::Open => Ok(()),
Self::Owner(owner) if Some(owner.as_str()) == caller => Ok(()),
Self::Owner(_) => Err(A2AError::ContextAccessDenied {
context_id: context_id.to_string(),
}),
}
}
}
#[cfg(feature = "sqlx-storage")]
impl SqlxTaskStorage {
fn dialect_for(database_url: &str) -> Result<Dialect, A2AError> {
let Some(database_type) = DatabaseType::from_url(database_url) else {
return Err(A2AError::DatabaseError(format!(
"Unrecognized database URL scheme in '{database_url}'. Expected sqlite: or \
postgres:, e.g. 'sqlite::memory:' or 'postgres://user:pass@localhost/a2a'"
)));
};
let Some(dialect) = Dialect::of(database_type) else {
return Err(A2AError::DatabaseError(format!(
"{database_type} is not supported by SqlxTaskStorage. It stores tasks in SQLite \
or PostgreSQL; there is no {database_type} schema."
)));
};
if !database_type.is_feature_enabled() {
return Err(A2AError::DatabaseError(format!(
"{database_type} detected from URL '{database_url}', but the '{}' feature is not \
enabled. Add `features = [\"{}\"]` to your a2a-rs dependency.",
database_type.feature_name(),
database_type.feature_name(),
)));
}
Ok(dialect)
}
fn pooled_url(database_url: &str) -> std::borrow::Cow<'_, str> {
static NEXT: std::sync::atomic::AtomicUsize = std::sync::atomic::AtomicUsize::new(0);
let anonymous_memory = (database_url.contains(":memory:")
|| database_url.contains("mode=memory"))
&& !database_url.contains("cache=shared");
if !anonymous_memory {
return std::borrow::Cow::Borrowed(database_url);
}
let n = NEXT.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
std::borrow::Cow::Owned(format!(
"sqlite:file:a2a-in-memory-{n}?mode=memory&cache=shared"
))
}
fn pool_options(dialect: Dialect) -> AnyPoolOptions {
match dialect {
Dialect::Sqlite => AnyPoolOptions::new().after_connect(|conn, _meta| {
Box::pin(async move {
sqlx::query("PRAGMA foreign_keys = ON")
.execute(conn)
.await?;
Ok(())
})
}),
Dialect::Postgres => AnyPoolOptions::new(),
}
}
async fn connect(
database_url: &str,
settings: &PoolSettings,
) -> Result<(AnyPool, Dialect), A2AError> {
let dialect = Self::dialect_for(database_url)?;
sqlx::any::install_default_drivers();
let url = match dialect {
Dialect::Sqlite => Self::pooled_url(database_url),
Dialect::Postgres => std::borrow::Cow::Borrowed(database_url),
};
let pool = Self::pool_options(dialect)
.max_connections(settings.max_connections)
.acquire_timeout(settings.acquire_timeout)
.connect_with(settings.connect_options(&url)?)
.await
.map_err(|e| A2AError::DatabaseError(format!("Failed to connect to database: {e}")))?;
let migrations = Self::pool_options(dialect)
.max_connections(1)
.acquire_timeout(settings.acquire_timeout)
.connect_with(settings.connect_options(&url)?)
.await
.map_err(|e| {
A2AError::DatabaseError(format!("Failed to open the migration connection: {e}"))
})?;
let migrated = Self::run_base_migrations(migrations.clone(), dialect).await;
migrations.close().await;
migrated?;
Ok((pool, dialect))
}
pub async fn new(database_url: &str) -> Result<Self, A2AError> {
Self::builder(database_url).connect().await
}
pub fn builder(database_url: impl Into<String>) -> SqlxStorageBuilder {
SqlxStorageBuilder {
url: database_url.into(),
pool: PoolSettings::default(),
push_sender: None,
additional_migrations: Vec::new(),
event_log_capacity: Some(DEFAULT_EVENT_LOG_CAPACITY),
read_refresh: ReadRefresh::never(),
claim_cache_ttl: Some(DEFAULT_CLAIM_CACHE_TTL),
}
}
async fn run_base_migrations(pool: AnyPool, dialect: Dialect) -> Result<(), A2AError> {
if let Some(lock) = dialect.migration_lock() {
sqlx::raw_sql(lock).execute(&pool).await.map_err(|e| {
A2AError::DatabaseError(format!("Failed to take the migration lock: {e}"))
})?;
}
let [initial, push_configs, rest @ ..] = dialect.migrations();
Self::run_migration(pool.clone(), initial).await?;
Self::drop_legacy_push_configs(pool.clone(), dialect).await?;
Self::run_migration(pool.clone(), push_configs).await?;
for migration in rest {
Self::run_migration(pool.clone(), migration).await?;
}
sqlx::raw_sql(
"UPDATE task_history SET context_id = \
(SELECT context_id FROM tasks WHERE tasks.id = task_history.task_id) \
WHERE context_id IS NULL",
)
.execute(&pool)
.await
.map_err(|e| A2AError::DatabaseError(format!("Migration 004 backfill failed: {e}")))?;
Self::drop_dead_context_state_column(pool.clone(), dialect).await;
Ok(())
}
async fn drop_dead_context_state_column(pool: AnyPool, dialect: Dialect) {
let probe = sqlx::query(dialect.dead_context_state_column_probe())
.fetch_optional(&pool)
.await;
if !matches!(probe, Ok(Some(_))) {
return;
}
if let Err(e) = sqlx::raw_sql("ALTER TABLE contexts DROP COLUMN state")
.execute(&pool)
.await
{
#[cfg(feature = "tracing")]
tracing::debug!("left the unused contexts.state column in place: {e}");
#[cfg(not(feature = "tracing"))]
let _ = e;
}
}
async fn run_migration(
pool: AnyPool,
migration: super::dialect::Migration,
) -> Result<(), A2AError> {
let mut attempt = sqlx::raw_sql(migration.sql).execute(&pool).await;
if attempt
.as_ref()
.err()
.is_some_and(super::dialect::is_concurrent_ddl_conflict)
{
attempt = sqlx::raw_sql(migration.sql).execute(&pool).await;
}
match attempt {
Ok(_) => Ok(()),
Err(e)
if migration.tolerates_existing_column
&& e.to_string().contains("duplicate column name") =>
{
Ok(())
}
Err(e) => Err(A2AError::DatabaseError(format!(
"Migration {} failed: {e}",
migration.name
))),
}
}
async fn drop_legacy_push_configs(pool: AnyPool, dialect: Dialect) -> Result<(), A2AError> {
let legacy = sqlx::query(dialect.legacy_push_config_probe())
.fetch_optional(&pool)
.await
.map_err(|e| {
A2AError::DatabaseError(format!("Failed to inspect push config table: {e}"))
})?;
if legacy.is_some() {
sqlx::raw_sql("DROP TABLE IF EXISTS push_notification_configs")
.execute(&pool)
.await
.map_err(|e| {
A2AError::DatabaseError(format!(
"Failed to drop the pre-v0.3 push config table: {e}"
))
})?;
}
Ok(())
}
async fn run_additional_migrations(
pool: &AnyPool,
migrations: &[String],
) -> Result<(), A2AError> {
for (i, migration_sql) in migrations.iter().enumerate() {
sqlx::raw_sql(migration_sql)
.execute(pool)
.await
.map_err(|e| {
A2AError::DatabaseError(format!("Additional migration {} failed: {}", i + 1, e))
})?;
}
Ok(())
}
fn sql<'a>(&self, sql: &'a str) -> std::borrow::Cow<'a, str> {
self.dialect.bind_params(sql)
}
fn row_to_task(row: &sqlx::any::AnyRow) -> Result<Task, A2AError> {
let task_id: String = row
.try_get("id")
.map_err(|e| A2AError::DatabaseError(format!("Failed to get task_id: {}", e)))?;
let context_id: String = row
.try_get("context_id")
.map_err(|e| A2AError::DatabaseError(format!("Failed to get context_id: {}", e)))?;
let status_state: String = row
.try_get("status_state")
.map_err(|e| A2AError::DatabaseError(format!("Failed to get status_state: {}", e)))?;
let status_message_json: Option<String> = row
.try_get("status_message")
.map_err(|e| A2AError::DatabaseError(format!("Failed to get status_message: {}", e)))?;
let metadata_json: Option<String> = row
.try_get("metadata")
.map_err(|e| A2AError::DatabaseError(format!("Failed to get metadata: {}", e)))?;
let artifacts_json: Option<String> = row
.try_get("artifacts")
.map_err(|e| A2AError::DatabaseError(format!("Failed to get artifacts: {}", e)))?;
let state = match status_state.as_str() {
"submitted" => TaskState::Submitted,
"working" => TaskState::Working,
"input-required" => TaskState::InputRequired,
"completed" => TaskState::Completed,
"canceled" => TaskState::Canceled,
"failed" => TaskState::Failed,
"rejected" => TaskState::Rejected,
"auth-required" => TaskState::AuthRequired,
"unknown" => TaskState::Unknown,
_ => TaskState::Unknown,
};
let status_message = if let Some(msg_str) = status_message_json {
Some(serde_json::from_str(&msg_str).map_err(|e| {
A2AError::DatabaseError(format!("Failed to parse status message: {}", e))
})?)
} else {
None
};
let metadata =
if let Some(meta_str) = metadata_json {
Some(serde_json::from_str(&meta_str).map_err(|e| {
A2AError::DatabaseError(format!("Failed to parse metadata: {}", e))
})?)
} else {
None
};
let artifacts = if let Some(artifacts_str) = artifacts_json {
Some(serde_json::from_str(&artifacts_str).map_err(|e| {
A2AError::DatabaseError(format!("Failed to parse artifacts: {}", e))
})?)
} else {
None
};
let now = chrono::Utc::now();
let task_status = TaskStatus {
state: ::buffa::EnumValue::from(state),
message: status_message.into(),
timestamp: ::buffa::MessageField::some(::buffa_types::google::protobuf::Timestamp {
seconds: now.timestamp(),
nanos: now.timestamp_subsec_nanos() as i32,
..Default::default()
}),
..Default::default()
};
let task = Task {
id: task_id.clone(),
context_id,
status: ::buffa::MessageField::some(task_status),
history: Vec::new(),
metadata: metadata.into(),
artifacts: artifacts.unwrap_or_default(),
..Default::default()
};
Ok(task)
}
async fn load_task_history(
&self,
task_id: &str,
limit: Option<u32>,
) -> Result<Vec<Message>, A2AError> {
let query_str = if let Some(limit) = limit {
format!(
"SELECT id, status_state, message FROM task_history \
WHERE task_id = ? AND message IS NOT NULL ORDER BY id DESC LIMIT {}",
limit
)
} else {
"SELECT id, status_state, message FROM task_history \
WHERE task_id = ? AND message IS NOT NULL ORDER BY id DESC"
.to_string()
};
let query_str = self.sql(&query_str);
let rows = sqlx::query(&query_str)
.bind(task_id)
.fetch_all(&self.pool)
.await
.map_err(|e| A2AError::DatabaseError(format!("Failed to load task history: {}", e)))?;
let mut history = Vec::new();
for row in rows {
let message_json: Option<String> = row.try_get("message").map_err(|e| {
A2AError::DatabaseError(format!("Failed to get message from history: {}", e))
})?;
if let Some(msg_str) = message_json {
let message: Message = serde_json::from_str(&msg_str).map_err(|e| {
A2AError::DatabaseError(format!("Failed to parse message from history: {}", e))
})?;
history.push(message);
}
}
history.reverse();
Ok(history)
}
async fn add_to_history(
&self,
task_id: &str,
state: TaskState,
message: Option<Message>,
) -> Result<(), A2AError> {
let state_str = state_str(state);
let message_json = if let Some(msg) = message {
Some(serde_json::to_string(&msg).map_err(|e| {
A2AError::DatabaseError(format!("Failed to serialize message: {}", e))
})?)
} else {
None
};
let sql = self.sql(
"INSERT INTO task_history (task_id, context_id, status_state, message) \
VALUES (?, (SELECT context_id FROM tasks WHERE id = ?), ?, ?)",
);
sqlx::query(&sql)
.bind(task_id)
.bind(task_id)
.bind(state_str)
.bind(message_json)
.execute(&self.pool)
.await
.map_err(|e| A2AError::DatabaseError(format!("Failed to add task history: {}", e)))?;
Ok(())
}
async fn claim_or_check_context(
&self,
context_id: &str,
caller: Option<&str>,
) -> Result<(), A2AError> {
if let Some(claim) = self.read_claim(context_id).await? {
return claim.verdict(context_id, caller);
}
sqlx::query(self.dialect.insert_context_if_absent())
.bind(context_id)
.bind(caller)
.execute(&self.pool)
.await
.map_err(|e| A2AError::DatabaseError(format!("Failed to register context: {}", e)))?;
match self.read_claim(context_id).await? {
Some(claim) => claim.verdict(context_id, caller),
None => Ok(()),
}
}
async fn read_claim(&self, context_id: &str) -> Result<Option<ContextClaim>, A2AError> {
if let Some(claim) = self.claim_cache.as_ref().and_then(|c| c.get(context_id)) {
return Ok(Some(claim));
}
let sql = self.sql("SELECT owner FROM contexts WHERE id = ?");
let row = sqlx::query(&sql)
.bind(context_id)
.fetch_optional(&self.pool)
.await
.map_err(|e| A2AError::DatabaseError(format!("Failed to read context owner: {}", e)))?;
let Some(row) = row else {
return Ok(None);
};
let owner: Option<String> = row
.try_get("owner")
.map_err(|e| A2AError::DatabaseError(format!("Failed to get context owner: {}", e)))?;
let claim = match owner {
Some(owner) => ContextClaim::Owner(owner),
None => ContextClaim::Open,
};
if let Some(cache) = self.claim_cache.as_ref() {
cache.insert(context_id, &claim);
}
Ok(Some(claim))
}
async fn refresh_user_state(&self, principal: &str) -> Result<(), A2AError> {
let Some(cutoff) = self.read_refresh.cutoff(chrono::Utc::now()) else {
return Ok(());
};
sqlx::query(self.dialect.refresh_user_state())
.bind(principal)
.bind(self.dialect.format_timestamp(cutoff))
.execute(&self.pool)
.await
.map_err(|e| A2AError::DatabaseError(format!("Failed to refresh user state: {}", e)))?;
Ok(())
}
pub fn push_notifier(&self) -> Arc<dyn AsyncPushNotifier> {
self.push_notification_registry.clone()
}
pub fn max_connections(&self) -> u32 {
self.pool.options().get_max_connections()
}
}
#[cfg(feature = "sqlx-storage")]
#[async_trait]
impl AsyncTaskLifecycle for SqlxTaskStorage {
async fn create(&self, id: &TaskId, context_id: &ContextId) -> Result<Task, A2AError> {
let task_id = id.as_str();
let context_id = context_id.as_str();
let exists_sql = self.sql("SELECT id FROM tasks WHERE id = ?");
let existing = sqlx::query(&exists_sql)
.bind(task_id)
.fetch_optional(&self.pool)
.await
.map_err(|e| {
A2AError::DatabaseError(format!("Failed to check existing task: {}", e))
})?;
if existing.is_some() {
return Err(A2AError::TaskNotFound(format!(
"Task {} already exists",
task_id
)));
}
let task = Task::new(task_id.to_string(), context_id.to_string());
let metadata_json = task
.metadata
.as_option()
.map(|m| serde_json::to_string(m).unwrap_or_default());
let artifacts_json = serde_json::to_string(&task.artifacts).unwrap_or_default();
let status_message_str = task
.status
.as_option()
.and_then(|s| s.message.as_option())
.map(|m| serde_json::to_string(m).unwrap_or_default());
let insert_sql = self.sql(
"INSERT INTO tasks (id, context_id, status_state, status_message, metadata, artifacts) \
VALUES (?, ?, ?, ?, ?, ?)",
);
sqlx::query(&insert_sql)
.bind(&task.id)
.bind(&task.context_id)
.bind("submitted")
.bind(status_message_str)
.bind(metadata_json)
.bind(artifacts_json)
.execute(&self.pool)
.await
.map_err(|e| A2AError::DatabaseError(format!("Failed to create task: {}", e)))?;
self.add_to_history(task_id, TaskState::Submitted, None)
.await?;
Ok(task)
}
async fn update_status(
&self,
id: &TaskId,
state: TaskState,
message: Option<Message>,
) -> Result<Task, A2AError> {
let task_id = id.as_str();
let state_str = state_str(state);
let sql = self.sql(
"UPDATE tasks SET status_state = ?, status_message = ?, version = version + 1 \
WHERE id = ?",
);
let result = sqlx::query(&sql)
.bind(state_str)
.bind(status_message_json(message.as_ref()))
.bind(task_id)
.execute(&self.pool)
.await
.map_err(|e| A2AError::DatabaseError(format!("Failed to update task status: {}", e)))?;
if result.rows_affected() == 0 {
return Err(A2AError::TaskNotFound(task_id.to_string()));
}
self.add_to_history(task_id, state, message).await?;
self.get(id, None).await
}
async fn exists(&self, id: &TaskId) -> Result<bool, A2AError> {
let task_id = id.as_str();
let sql = self.sql("SELECT id FROM tasks WHERE id = ?");
let row = sqlx::query(&sql)
.bind(task_id)
.fetch_optional(&self.pool)
.await
.map_err(|e| {
A2AError::DatabaseError(format!("Failed to check task existence: {}", e))
})?;
Ok(row.is_some())
}
async fn get(&self, id: &TaskId, history_length: Option<u32>) -> Result<Task, A2AError> {
let task_id = id.as_str();
let query_str = format!("SELECT {TASK_COLUMNS} FROM tasks WHERE id = ?");
let sql = self.sql(&query_str);
let row = sqlx::query(&sql)
.bind(task_id)
.fetch_optional(&self.pool)
.await
.map_err(|e| A2AError::DatabaseError(format!("Failed to get task: {}", e)))?;
let Some(row) = row else {
return Err(A2AError::TaskNotFound(task_id.to_string()));
};
let mut task = Self::row_to_task(&row)?;
if history_length.is_some() || history_length.is_none() {
let history = self.load_task_history(task_id, history_length).await?;
task.history = history;
}
Ok(task)
}
async fn cancel(&self, id: &TaskId) -> Result<Task, A2AError> {
let task_id = id.as_str();
let task = self.get(id, None).await?;
if !task.status.state.is_cancelable() {
return Err(A2AError::TaskNotCancelable(format!(
"Task {} has already finished in state {:?} and cannot be canceled",
task_id, task.status.state
)));
}
let mut cancel_message = Message::agent_text(
format!("Task {} canceled.", task_id),
uuid::Uuid::new_v4().to_string(),
);
cancel_message.task_id = task_id.to_string();
cancel_message.context_id = task.context_id.clone();
let sql = self.sql(
"UPDATE tasks SET status_state = ?, status_message = ?, version = version + 1 \
WHERE id = ?",
);
sqlx::query(&sql)
.bind(state_str(TaskState::Canceled))
.bind(status_message_json(Some(&cancel_message)))
.bind(task_id)
.execute(&self.pool)
.await
.map_err(|e| A2AError::DatabaseError(format!("Failed to cancel task: {}", e)))?;
self.add_to_history(task_id, TaskState::Canceled, Some(cancel_message))
.await?;
self.get(id, None).await
}
}
#[cfg(feature = "sqlx-storage")]
impl SqlxTaskStorage {
async fn current_version(&self, task_id: &str) -> Result<Option<u64>, A2AError> {
let sql = self.sql("SELECT version FROM tasks WHERE id = ?");
let row = sqlx::query(&sql)
.bind(task_id)
.fetch_optional(&self.pool)
.await
.map_err(|e| A2AError::DatabaseError(format!("Failed to read task version: {}", e)))?;
match row {
Some(row) => {
let v: i64 = row.try_get("version").map_err(|e| {
A2AError::DatabaseError(format!("Failed to get version column: {}", e))
})?;
Ok(Some(v as u64))
}
None => Ok(None),
}
}
}
#[cfg(feature = "sqlx-storage")]
#[async_trait]
impl AsyncTaskVersioning for SqlxTaskStorage {
async fn version(&self, id: &TaskId) -> Result<u64, A2AError> {
self.current_version(id.as_str())
.await?
.ok_or_else(|| A2AError::TaskNotFound(id.as_str().to_string()))
}
async fn get_versioned(
&self,
id: &TaskId,
history_length: Option<u32>,
) -> Result<VersionedTask, A2AError> {
let task = self.get(id, history_length).await?;
let version = self.version(id).await?;
Ok(VersionedTask::new(task, version))
}
async fn update_status_checked(
&self,
id: &TaskId,
expected: u64,
state: TaskState,
message: Option<Message>,
) -> Result<VersionedTask, A2AError> {
let task_id = id.as_str();
let state_str = state_str(state);
let sql = self.sql(
"UPDATE tasks SET status_state = ?, status_message = ?, version = version + 1 \
WHERE id = ? AND version = ?",
);
let result = sqlx::query(&sql)
.bind(state_str)
.bind(status_message_json(message.as_ref()))
.bind(task_id)
.bind(expected as i64)
.execute(&self.pool)
.await
.map_err(|e| A2AError::DatabaseError(format!("Failed to update task status: {}", e)))?;
if result.rows_affected() == 0 {
return match self.current_version(task_id).await? {
Some(actual) => Err(A2AError::VersionConflict {
id: task_id.to_string(),
expected,
actual,
}),
None => Err(A2AError::TaskNotFound(task_id.to_string())),
};
}
self.add_to_history(task_id, state, message).await?;
let task = self.get(id, None).await?;
Ok(VersionedTask::new(task, expected + 1))
}
}
#[cfg(feature = "sqlx-storage")]
#[async_trait]
impl AsyncTaskQuery for SqlxTaskStorage {
async fn list(
&self,
params: &crate::domain::ListTasksParams,
) -> Result<crate::domain::ListTasksResult, A2AError> {
use crate::domain::ListTasksResult;
let mut where_conditions = Vec::new();
if params.context_id.is_some() {
where_conditions.push("context_id = ?".to_string());
}
if params.status.is_some() {
where_conditions.push("status_state = ?".to_string());
}
let timestamp_str = if let Some(status_timestamp_after) = ¶ms.status_timestamp_after {
let timestamp =
chrono::DateTime::parse_from_rfc3339(status_timestamp_after).map_err(|e| {
A2AError::DatabaseError(format!(
"Invalid timestamp value: {} ({})",
status_timestamp_after, e
))
})?;
where_conditions.push(self.dialect.updated_since_predicate().to_string());
Some(
self.dialect
.format_timestamp(timestamp.with_timezone(&chrono::Utc)),
)
} else {
None
};
let where_clause = if where_conditions.is_empty() {
String::new()
} else {
format!(" WHERE {}", where_conditions.join(" AND "))
};
let count_sql = format!("SELECT COUNT(*) as count FROM tasks{}", where_clause);
let count_query = self.sql(&count_sql);
let mut count_q = sqlx::query(&count_query);
if let Some(ref context_id) = params.context_id {
count_q = count_q.bind(context_id);
}
if let Some(ref status) = params.status {
let state_str = state_str(*status);
count_q = count_q.bind(state_str);
}
if let Some(ref ts) = timestamp_str {
count_q = count_q.bind(ts);
}
let count_row = count_q
.fetch_one(&self.pool)
.await
.map_err(|e| A2AError::DatabaseError(format!("Failed to count tasks: {}", e)))?;
let total_size: i32 = count_row
.try_get::<i64, _>("count")
.map_err(|e| A2AError::DatabaseError(format!("Failed to get count: {}", e)))?
.try_into()
.unwrap_or(i32::MAX);
let page_size = params.page_size.unwrap_or(50).clamp(1, 100);
let offset = if let Some(ref token) = params.page_token {
token.parse::<i32>().unwrap_or(0)
} else {
0
};
let main_sql = format!(
"SELECT {TASK_COLUMNS} FROM tasks{} ORDER BY updated_at DESC LIMIT ? OFFSET ?",
where_clause
);
let main_query = self.sql(&main_sql);
let mut main_q = sqlx::query(&main_query);
if let Some(ref context_id) = params.context_id {
main_q = main_q.bind(context_id);
}
if let Some(ref status) = params.status {
let state_str = state_str(*status);
main_q = main_q.bind(state_str);
}
if let Some(ref ts) = timestamp_str {
main_q = main_q.bind(ts);
}
main_q = main_q.bind(page_size).bind(offset);
let rows = main_q
.fetch_all(&self.pool)
.await
.map_err(|e| A2AError::DatabaseError(format!("Failed to list tasks: {}", e)))?;
let mut tasks: Vec<Task> = rows
.iter()
.filter_map(|row| Self::row_to_task(row).ok())
.collect();
let history_length = params.history_length.unwrap_or(0);
for task in &mut tasks {
if history_length > 0 {
let history = self
.load_task_history(&task.id, Some(history_length as u32))
.await?;
task.history = history;
} else {
task.history.clear();
}
if !params.include_artifacts.unwrap_or(false) {
task.artifacts.clear();
}
}
let has_more = offset + page_size < total_size;
let next_page_token = if has_more {
(offset + page_size).to_string()
} else {
String::new()
};
Ok(ListTasksResult {
tasks,
total_size,
page_size,
next_page_token,
})
}
}
#[cfg(feature = "sqlx-storage")]
#[async_trait]
impl AsyncNotificationManager for SqlxTaskStorage {
async fn get_config(
&self,
params: &crate::domain::GetTaskPushNotificationConfigParams,
) -> Result<crate::domain::TaskPushNotificationConfig, A2AError> {
let by_id = self.sql(
"SELECT id, task_id, url, token, authentication FROM push_notification_configs \
WHERE task_id = ? AND id = ?",
);
let by_task = self.sql(
"SELECT id, task_id, url, token, authentication FROM push_notification_configs \
WHERE task_id = ? ORDER BY id LIMIT 1",
);
let row = match params.push_notification_config_id.as_ref() {
Some(config_id) => sqlx::query(&by_id).bind(¶ms.id).bind(config_id),
None => sqlx::query(&by_task).bind(¶ms.id),
}
.fetch_optional(&self.pool)
.await
.map_err(|e| A2AError::DatabaseError(format!("Failed to get push config: {}", e)))?;
if let Some(row) = row {
let id: String = row
.try_get("id")
.map_err(|e| A2AError::DatabaseError(format!("Failed to get config id: {}", e)))?;
let url: String = row
.try_get("url")
.map_err(|e| A2AError::DatabaseError(format!("Failed to get url: {}", e)))?;
let token: Option<String> = row.try_get("token").ok();
let auth_json: Option<String> = row.try_get("authentication").ok();
let auth_info = if let Some(auth_str) = auth_json {
serde_json::from_str(&auth_str).ok()
} else {
None
};
Ok(crate::domain::TaskPushNotificationConfig {
task_id: params.id.clone(),
id,
url,
token: token.unwrap_or_default(),
authentication: auth_info.into(),
tenant: "".to_string(),
..Default::default()
})
} else {
Err(A2AError::TaskNotFound(format!(
"Push notification config not found for task {}{}",
params.id,
params
.push_notification_config_id
.as_ref()
.map(|id| format!(" with id {}", id))
.unwrap_or_default()
)))
}
}
async fn list_configs(
&self,
params: &crate::domain::ListTaskPushNotificationConfigsParams,
) -> Result<Vec<crate::domain::TaskPushNotificationConfig>, A2AError> {
let sql = self.sql(
"SELECT id, task_id, url, token, authentication FROM push_notification_configs \
WHERE task_id = ?",
);
let rows = sqlx::query(&sql)
.bind(¶ms.id)
.fetch_all(&self.pool)
.await
.map_err(|e| A2AError::DatabaseError(format!("Failed to list push configs: {}", e)))?;
let configs: Vec<crate::domain::TaskPushNotificationConfig> = rows
.iter()
.filter_map(|row| {
let id: String = row.try_get("id").ok()?;
let url: String = row.try_get("url").ok()?;
let token: Option<String> = row.try_get("token").ok().flatten();
let auth_json: Option<String> = row.try_get("authentication").ok().flatten();
let auth_info = if let Some(auth_str) = auth_json {
serde_json::from_str(&auth_str).ok()
} else {
None
};
Some(crate::domain::TaskPushNotificationConfig {
task_id: params.id.clone(),
id,
url,
token: token.unwrap_or_default(),
authentication: auth_info.into(),
tenant: "".to_string(),
..Default::default()
})
})
.collect();
Ok(configs)
}
async fn delete_config(
&self,
params: &crate::domain::DeleteTaskPushNotificationConfigParams,
) -> Result<(), A2AError> {
let all_for_task = self.sql("DELETE FROM push_notification_configs WHERE task_id = ?");
let one = self.sql("DELETE FROM push_notification_configs WHERE task_id = ? AND id = ?");
let query = if params.push_notification_config_id.is_empty() {
sqlx::query(&all_for_task).bind(¶ms.id)
} else {
sqlx::query(&one)
.bind(¶ms.id)
.bind(¶ms.push_notification_config_id)
};
let _result = query
.execute(&self.pool)
.await
.map_err(|e| A2AError::DatabaseError(format!("Failed to delete push config: {}", e)))?;
Ok(())
}
async fn set_config(
&self,
config: &TaskPushNotificationConfig,
) -> Result<TaskPushNotificationConfig, A2AError> {
let config_id = if config.id.is_empty() {
uuid::Uuid::new_v4().to_string()
} else {
config.id.clone()
};
let auth_json = config
.authentication
.as_option()
.map(|auth| serde_json::to_string(auth).unwrap_or_default());
sqlx::query(self.dialect.upsert_push_config())
.bind(&config_id)
.bind(&config.task_id)
.bind(&config.url)
.bind(&config.token)
.bind(auth_json)
.execute(&self.pool)
.await
.map_err(|e| {
A2AError::DatabaseError(format!("Failed to set push notification config: {}", e))
})?;
self.push_notification_registry
.register(&config.task_id, config.clone())
.await?;
let mut result_config = config.clone();
result_config.id = config_id;
Ok(result_config)
}
}
#[cfg(feature = "sqlx-storage")]
impl Clone for SqlxTaskStorage {
fn clone(&self) -> Self {
Self {
pool: self.pool.clone(),
dialect: self.dialect,
push_notification_registry: self.push_notification_registry.clone(),
event_log_capacity: self.event_log_capacity,
read_refresh: self.read_refresh,
claim_cache: self.claim_cache.clone(),
}
}
}
#[cfg(feature = "sqlx-storage")]
#[async_trait]
impl AsyncConversationStore for SqlxTaskStorage {
async fn load(
&self,
context_id: &ContextId,
caller: Option<&str>,
limit: Option<u32>,
) -> Result<Conversation, A2AError> {
let context_id = context_id.as_str();
self.claim_or_check_context(context_id, caller).await?;
let digest_sql = self.sql(
"SELECT covers_through_seq, summary, replaced_messages, model \
FROM context_digests WHERE context_id = ? \
ORDER BY covers_through_seq DESC LIMIT 1",
);
let digest_row = sqlx::query(&digest_sql)
.bind(context_id)
.fetch_optional(&self.pool)
.await
.map_err(|e| {
A2AError::DatabaseError(format!("Failed to load context digest: {}", e))
})?;
let digest = match digest_row {
Some(row) => {
let covers_through: i64 = row.try_get("covers_through_seq").map_err(|e| {
A2AError::DatabaseError(format!("Failed to get digest watermark: {}", e))
})?;
let summary: String = row.try_get("summary").map_err(|e| {
A2AError::DatabaseError(format!("Failed to get digest summary: {}", e))
})?;
let replaced_messages: i64 = row.try_get("replaced_messages").map_err(|e| {
A2AError::DatabaseError(format!("Failed to get digest message count: {}", e))
})?;
let model: String = row.try_get("model").map_err(|e| {
A2AError::DatabaseError(format!("Failed to get digest model: {}", e))
})?;
Some(Digest {
covers_through: Seq::new(covers_through.max(0) as u64),
summary,
replaced_messages: replaced_messages.max(0) as u32,
model,
})
}
None => None,
};
let watermark = digest
.as_ref()
.map(|digest| digest.covers_through.get())
.unwrap_or(0) as i64;
let query = match limit {
Some(limit) => format!(
"SELECT id, message FROM task_history \
WHERE context_id = ? AND id > ? AND message IS NOT NULL \
ORDER BY id DESC LIMIT {}",
limit
),
None => "SELECT id, message FROM task_history \
WHERE context_id = ? AND id > ? AND message IS NOT NULL \
ORDER BY id DESC"
.to_string(),
};
let query = self.sql(&query);
let rows = sqlx::query(&query)
.bind(context_id)
.bind(watermark)
.fetch_all(&self.pool)
.await
.map_err(|e| A2AError::DatabaseError(format!("Failed to load conversation: {}", e)))?;
let mut tail = Vec::with_capacity(rows.len());
for row in rows {
let seq: i64 = row.try_get("id").map_err(|e| {
A2AError::DatabaseError(format!("Failed to get history sequence: {}", e))
})?;
let message_json: String = row.try_get("message").map_err(|e| {
A2AError::DatabaseError(format!("Failed to get history message: {}", e))
})?;
let message: Message = serde_json::from_str(&message_json).map_err(|e| {
A2AError::DatabaseError(format!("Failed to parse history message: {}", e))
})?;
tail.push(SequencedMessage {
seq: Seq::new(seq.max(0) as u64),
message,
});
}
tail.reverse();
Ok(Conversation { digest, tail })
}
async fn compact(
&self,
context_id: &ContextId,
caller: Option<&str>,
digest: Digest,
) -> Result<(), A2AError> {
let context_id = context_id.as_str();
self.claim_or_check_context(context_id, caller).await?;
let sql = self.sql(
"INSERT INTO context_digests \
(context_id, covers_through_seq, summary, replaced_messages, model) \
VALUES (?, ?, ?, ?, ?)",
);
sqlx::query(&sql)
.bind(context_id)
.bind(digest.covers_through.get() as i64)
.bind(&digest.summary)
.bind(digest.replaced_messages as i64)
.bind(&digest.model)
.execute(&self.pool)
.await
.map_err(|e| {
A2AError::DatabaseError(format!("Failed to append context digest: {}", e))
})?;
Ok(())
}
}
#[cfg(feature = "sqlx-storage")]
fn scope_column(scope: StateScope) -> Option<&'static str> {
match scope {
StateScope::User => Some("user"),
StateScope::Context => Some("context"),
StateScope::Temp => None,
}
}
#[cfg(feature = "sqlx-storage")]
#[async_trait]
impl AsyncContextStateStore for SqlxTaskStorage {
async fn load_state(
&self,
context_id: &ContextId,
caller: Option<&str>,
) -> Result<ContextState, A2AError> {
let context_id = context_id.as_str();
self.claim_or_check_context(context_id, caller).await?;
let sql = self.sql(
"SELECT scope, name, value FROM context_state \
WHERE (scope = 'context' AND scope_key = ?) \
OR (scope = 'user' AND scope_key = ?)",
);
let rows = sqlx::query(&sql)
.bind(context_id)
.bind(caller)
.fetch_all(&self.pool)
.await
.map_err(|e| A2AError::DatabaseError(format!("Failed to load context state: {}", e)))?;
let mut state = ContextState::new();
for row in rows {
let scope: String = row.try_get("scope").map_err(|e| {
A2AError::DatabaseError(format!("Failed to get state scope: {}", e))
})?;
let name: String = row
.try_get("name")
.map_err(|e| A2AError::DatabaseError(format!("Failed to get state key: {}", e)))?;
let value: String = row.try_get("value").map_err(|e| {
A2AError::DatabaseError(format!("Failed to get state value: {}", e))
})?;
let scope = match scope.as_str() {
"user" => StateScope::User,
"context" => StateScope::Context,
_other => {
#[cfg(feature = "tracing")]
tracing::warn!("ignoring state row with unknown scope '{_other}'");
continue;
}
};
match StateKey::scoped(scope, &name) {
Ok(key) => state.insert(key, value),
Err(_e) => {
#[cfg(feature = "tracing")]
tracing::warn!("ignoring unusable state key '{name}': {_e}");
}
}
}
if let Some(caller) = caller {
self.refresh_user_state(caller).await?;
}
Ok(state)
}
async fn remember(
&self,
context_id: &ContextId,
caller: Option<&str>,
key: &StateKey,
value: &str,
) -> Result<Remembered, A2AError> {
let context_id = context_id.as_str();
self.claim_or_check_context(context_id, caller).await?;
let (Some(scope_key), Some(scope)) = (
scope_key(key.scope(), context_id, caller, key)?,
scope_column(key.scope()),
) else {
return Ok(Remembered::NotStored);
};
let mut tx = self.pool.begin().await.map_err(|e| {
A2AError::DatabaseError(format!("Failed to open a state transaction: {}", e))
})?;
let read = self
.sql("SELECT value FROM context_state WHERE scope = ? AND scope_key = ? AND name = ?");
let previous: Option<String> = sqlx::query(&read)
.bind(scope)
.bind(scope_key)
.bind(key.name())
.fetch_optional(&mut *tx)
.await
.map_err(|e| A2AError::DatabaseError(format!("Failed to read context state: {}", e)))?
.map(|row| row.try_get("value"))
.transpose()
.map_err(|e| A2AError::DatabaseError(format!("Failed to read context state: {}", e)))?;
sqlx::query(self.dialect.upsert_context_state())
.bind(scope)
.bind(scope_key)
.bind(key.name())
.bind(value)
.execute(&mut *tx)
.await
.map_err(|e| {
A2AError::DatabaseError(format!("Failed to write context state: {}", e))
})?;
tx.commit().await.map_err(|e| {
A2AError::DatabaseError(format!("Failed to commit context state: {}", e))
})?;
Ok(match previous {
None => Remembered::Stored,
Some(previous) if previous == value => Remembered::Unchanged,
Some(previous) => Remembered::Replaced { previous },
})
}
async fn forget(
&self,
context_id: &ContextId,
caller: Option<&str>,
key: &StateKey,
) -> Result<bool, A2AError> {
let context_id = context_id.as_str();
self.claim_or_check_context(context_id, caller).await?;
let (Some(scope_key), Some(scope)) = (
scope_key(key.scope(), context_id, caller, key)?,
scope_column(key.scope()),
) else {
return Ok(false);
};
let sql =
self.sql("DELETE FROM context_state WHERE scope = ? AND scope_key = ? AND name = ?");
let deleted = sqlx::query(&sql)
.bind(scope)
.bind(scope_key)
.bind(key.name())
.execute(&self.pool)
.await
.map_err(|e| A2AError::DatabaseError(format!("Failed to drop context state: {}", e)))?;
Ok(deleted.rows_affected() > 0)
}
}
#[cfg(feature = "sqlx-storage")]
const KIND_STATUS: &str = "status-update";
#[cfg(feature = "sqlx-storage")]
const KIND_ARTIFACT: &str = "artifact-update";
#[cfg(feature = "sqlx-storage")]
#[async_trait]
impl AsyncEventLog for SqlxTaskStorage {
async fn append(&self, task_id: &str, event: UpdateEvent) -> Result<SeqEvent, A2AError> {
let (kind, payload) = match &event {
UpdateEvent::StatusUpdate(update) => (KIND_STATUS, serde_json::to_string(update)),
UpdateEvent::ArtifactUpdate(update) => (KIND_ARTIFACT, serde_json::to_string(update)),
};
let payload = payload.map_err(|e| {
A2AError::DatabaseError(format!("Failed to serialize a stream event: {e}"))
})?;
let sql = self.dialect.insert_task_event();
let id: i64 = sqlx::query(sql)
.bind(task_id)
.bind(kind)
.bind(&payload)
.bind(task_id)
.fetch_one(&self.pool)
.await
.and_then(|row| row.try_get("id"))
.map_err(|e| {
A2AError::DatabaseError(format!("Failed to log a stream event for {task_id}: {e}"))
})?;
let id = id as u64;
if let Some(capacity) = self.event_log_capacity
&& let Some(cutoff) = id.checked_sub(capacity)
{
let sql = self.sql("DELETE FROM task_events WHERE task_id = ? AND id <= ?");
sqlx::query(&sql)
.bind(task_id)
.bind(cutoff as i64)
.execute(&self.pool)
.await
.map_err(|e| {
A2AError::DatabaseError(format!(
"Failed to trim the stream log for {task_id}: {e}"
))
})?;
}
Ok(SeqEvent::new(id, event))
}
async fn replay(&self, task_id: &str, from: u64) -> Result<Replay, A2AError> {
let sql = self.sql("SELECT MIN(id) AS oldest FROM task_events WHERE task_id = ?");
let oldest: Option<i64> = sqlx::query(&sql)
.bind(task_id)
.fetch_one(&self.pool)
.await
.and_then(|row| row.try_get("oldest"))
.map_err(|e| {
A2AError::DatabaseError(format!("Failed to read the stream log for {task_id}: {e}"))
})?;
let sql = self.sql(
"SELECT id, kind, payload FROM task_events \
WHERE task_id = ? AND id > ? ORDER BY id",
);
let rows = sqlx::query(&sql)
.bind(task_id)
.bind(from as i64)
.fetch_all(&self.pool)
.await
.map_err(|e| {
A2AError::DatabaseError(format!("Failed to replay the log for {task_id}: {e}"))
})?;
let events = rows
.iter()
.map(Self::row_to_seq_event)
.collect::<Result<Vec<_>, _>>()?;
Ok(Replay::bounded_by(
oldest.map(|oldest| oldest as u64),
from,
events,
))
}
async fn discard(&self, task_id: &str) -> Result<(), A2AError> {
let sql = self.sql("DELETE FROM task_events WHERE task_id = ?");
sqlx::query(&sql)
.bind(task_id)
.execute(&self.pool)
.await
.map_err(|e| {
A2AError::DatabaseError(format!(
"Failed to discard the stream log for {task_id}: {e}"
))
})?;
Ok(())
}
}
#[cfg(feature = "sqlx-storage")]
impl SqlxTaskStorage {
fn row_to_seq_event(row: &sqlx::any::AnyRow) -> Result<SeqEvent, A2AError> {
let read = |column: &str| -> Result<String, A2AError> {
row.try_get(column).map_err(|e| {
A2AError::DatabaseError(format!("Failed to read stream event {column}: {e}"))
})
};
let id: i64 = row
.try_get("id")
.map_err(|e| A2AError::DatabaseError(format!("Failed to read stream event id: {e}")))?;
let kind = read("kind")?;
let payload = read("payload")?;
fn parse(what: &'static str) -> impl Fn(serde_json::Error) -> A2AError {
move |e| A2AError::DatabaseError(format!("Failed to parse a logged {what}: {e}"))
}
let event = match kind.as_str() {
KIND_STATUS => UpdateEvent::StatusUpdate(
serde_json::from_str(&payload).map_err(parse("status update"))?,
),
KIND_ARTIFACT => UpdateEvent::ArtifactUpdate(
serde_json::from_str(&payload).map_err(parse("artifact update"))?,
),
other => {
return Err(A2AError::DatabaseError(format!(
"Unknown stream event kind {other:?} in task_events"
)));
}
};
Ok(SeqEvent::new(id as u64, event))
}
}
#[cfg(feature = "sqlx-storage")]
#[async_trait]
impl AsyncRetention for SqlxTaskStorage {
async fn sweep(
&self,
policy: &RetentionPolicy,
now: chrono::DateTime<chrono::Utc>,
) -> Result<Swept, A2AError> {
let mut swept = Swept::default();
if let Some(cutoff) = policy.context_cutoff(now) {
for context_id in self.idle_contexts(cutoff).await? {
swept += self.delete_context(&context_id).await?;
}
}
if let Some(cutoff) = policy.user_state_cutoff(now) {
for principal in self.idle_principals(cutoff).await? {
swept.state_keys += self.delete_user_state(&principal).await?;
}
}
Ok(swept)
}
}
#[cfg(feature = "sqlx-storage")]
impl SqlxTaskStorage {
async fn idle_contexts(
&self,
cutoff: chrono::DateTime<chrono::Utc>,
) -> Result<Vec<String>, A2AError> {
let rows = sqlx::query(self.dialect.idle_contexts())
.bind(self.dialect.format_timestamp(cutoff))
.fetch_all(&self.pool)
.await
.map_err(|e| A2AError::DatabaseError(format!("Failed to find idle contexts: {e}")))?;
rows.iter()
.map(|row| {
row.try_get("ctx").map_err(|e| {
A2AError::DatabaseError(format!("Failed to read idle context id: {e}"))
})
})
.collect()
}
async fn idle_principals(
&self,
cutoff: chrono::DateTime<chrono::Utc>,
) -> Result<Vec<String>, A2AError> {
let rows = sqlx::query(self.dialect.idle_principals())
.bind(self.dialect.format_timestamp(cutoff))
.fetch_all(&self.pool)
.await
.map_err(|e| A2AError::DatabaseError(format!("Failed to find idle principals: {e}")))?;
rows.iter()
.map(|row| {
row.try_get("scope_key").map_err(|e| {
A2AError::DatabaseError(format!("Failed to read idle principal: {e}"))
})
})
.collect()
}
async fn delete_context(&self, context_id: &str) -> Result<Swept, A2AError> {
let mut tx = self.pool.begin().await.map_err(|e| {
A2AError::DatabaseError(format!("Failed to open a sweep transaction: {e}"))
})?;
let fail = |table: &str, e: sqlx::Error| {
A2AError::DatabaseError(format!("Failed to sweep {table} for {context_id}: {e}"))
};
let sql = self.sql(
"DELETE FROM push_notification_configs WHERE task_id IN (SELECT id FROM tasks WHERE context_id = ?)",
);
sqlx::query(&sql)
.bind(context_id)
.execute(&mut *tx)
.await
.map_err(|e| fail("push configs", e))?;
let sql = self.sql(
"SELECT COUNT(*) AS count FROM task_history WHERE message IS NOT NULL AND (context_id = ? OR task_id IN (SELECT id FROM tasks WHERE context_id = ?))",
);
let messages: i64 = sqlx::query(&sql)
.bind(context_id)
.bind(context_id)
.fetch_one(&mut *tx)
.await
.and_then(|row| row.try_get("count"))
.map_err(|e| fail("history", e))?;
let sql = self.sql(
"DELETE FROM task_history WHERE context_id = ? OR task_id IN (SELECT id FROM tasks WHERE context_id = ?)",
);
sqlx::query(&sql)
.bind(context_id)
.bind(context_id)
.execute(&mut *tx)
.await
.map_err(|e| fail("history", e))?;
let sql = self.sql(
"DELETE FROM task_events WHERE task_id IN (SELECT id FROM tasks WHERE context_id = ?)",
);
sqlx::query(&sql)
.bind(context_id)
.execute(&mut *tx)
.await
.map_err(|e| fail("stream events", e))?;
let sql = self.sql("DELETE FROM tasks WHERE context_id = ?");
let tasks = sqlx::query(&sql)
.bind(context_id)
.execute(&mut *tx)
.await
.map_err(|e| fail("tasks", e))?;
let sql = self.sql("DELETE FROM context_digests WHERE context_id = ?");
let digests = sqlx::query(&sql)
.bind(context_id)
.execute(&mut *tx)
.await
.map_err(|e| fail("digests", e))?;
let sql = self.sql("DELETE FROM context_state WHERE scope = 'context' AND scope_key = ?");
let state_keys = sqlx::query(&sql)
.bind(context_id)
.execute(&mut *tx)
.await
.map_err(|e| fail("state", e))?;
let sql = self.sql("DELETE FROM contexts WHERE id = ?");
sqlx::query(&sql)
.bind(context_id)
.execute(&mut *tx)
.await
.map_err(|e| fail("the context row", e))?;
tx.commit().await.map_err(|e| {
A2AError::DatabaseError(format!("Failed to commit the sweep of {context_id}: {e}"))
})?;
if let Some(cache) = self.claim_cache.as_ref() {
cache.forget(context_id);
}
Ok(Swept {
contexts: 1,
tasks: tasks.rows_affected(),
messages: messages.max(0) as u64,
digests: digests.rows_affected(),
state_keys: state_keys.rows_affected(),
})
}
async fn delete_user_state(&self, principal: &str) -> Result<u64, A2AError> {
let sql = self.sql("DELETE FROM context_state WHERE scope = 'user' AND scope_key = ?");
let deleted = sqlx::query(&sql)
.bind(principal)
.execute(&self.pool)
.await
.map_err(|e| {
A2AError::DatabaseError(format!("Failed to sweep user state for {principal}: {e}"))
})?;
Ok(deleted.rows_affected())
}
}
#[cfg(all(test, feature = "sqlx-storage"))]
mod tests {
use super::*;
#[tokio::test]
async fn the_dead_state_column_is_dropped_on_the_next_start() {
let dir = tempfile::tempdir().unwrap();
let url = format!("sqlite:{}?mode=rwc", dir.path().join("a2a.db").display());
let storage = SqlxTaskStorage::new(&url).await.unwrap();
sqlx::raw_sql("ALTER TABLE contexts ADD COLUMN state TEXT NOT NULL DEFAULT '{}'")
.execute(&storage.pool)
.await
.expect("put the pre-006 column back");
assert!(has_state_column(&storage).await);
drop(storage);
let restarted = SqlxTaskStorage::new(&url).await.unwrap();
assert!(
!has_state_column(&restarted).await,
"the unused column should be gone after the migration runs"
);
}
#[tokio::test]
async fn sqlite_connections_enforce_foreign_keys() {
let storage = SqlxTaskStorage::new("sqlite::memory:").await.unwrap();
let on: i64 = sqlx::query("PRAGMA foreign_keys")
.fetch_one(&storage.pool)
.await
.unwrap()
.try_get(0)
.unwrap();
assert_eq!(on, 1);
}
#[tokio::test]
async fn the_url_cannot_speak_for_foreign_keys() {
let dir = tempfile::tempdir().unwrap();
let url = format!(
"sqlite:{}?mode=rwc&foreign_keys=off",
dir.path().join("a2a.db").display()
);
let refused = SqlxTaskStorage::new(&url).await;
assert!(
refused.is_err(),
"a SQLite-specific URL parameter should not reach the driver"
);
}
#[tokio::test]
async fn deleting_a_task_cascades_to_its_history() {
let storage = SqlxTaskStorage::new("sqlite::memory:").await.unwrap();
let (task, context) = (tid("task-cascade"), cid("ctx-cascade"));
storage.create(&task, &context).await.unwrap();
storage
.update_status(&task, TaskState::Working, None)
.await
.unwrap();
assert!(history_rows(&storage, "task-cascade").await > 0);
sqlx::query("DELETE FROM tasks WHERE id = ?")
.bind("task-cascade")
.execute(&storage.pool)
.await
.unwrap();
assert_eq!(
history_rows(&storage, "task-cascade").await,
0,
"history should go with the task it references"
);
}
async fn history_rows(storage: &SqlxTaskStorage, task_id: &str) -> i64 {
sqlx::query("SELECT COUNT(*) AS count FROM task_history WHERE task_id = ?")
.bind(task_id)
.fetch_one(&storage.pool)
.await
.unwrap()
.try_get("count")
.unwrap()
}
fn tid(s: &str) -> TaskId {
s.parse().unwrap()
}
fn cid(s: &str) -> ContextId {
s.parse().unwrap()
}
async fn has_state_column(storage: &SqlxTaskStorage) -> bool {
sqlx::query(storage.dialect.dead_context_state_column_probe())
.fetch_optional(&storage.pool)
.await
.unwrap()
.is_some()
}
#[test]
fn a_cached_claim_is_only_returned_inside_its_window() {
let cache = ClaimCache::new(Duration::from_secs(60));
cache.insert("ctx-1", &ContextClaim::Owner("alice".to_string()));
assert!(matches!(
cache.get("ctx-1"),
Some(ContextClaim::Owner(owner)) if owner == "alice"
));
let expired = ClaimCache::new(Duration::ZERO);
expired.insert("ctx-1", &ContextClaim::Open);
assert!(expired.get("ctx-1").is_none());
}
#[test]
fn forgetting_a_context_drops_what_was_cached_for_it() {
let cache = ClaimCache::new(Duration::from_secs(60));
cache.insert("ctx-1", &ContextClaim::Open);
cache.forget("ctx-1");
assert!(cache.get("ctx-1").is_none());
}
#[test]
fn a_full_cache_starts_over_rather_than_growing() {
let cache = ClaimCache::new(Duration::from_secs(60));
for i in 0..=CLAIM_CACHE_CAPACITY {
cache.insert(&format!("ctx-{i}"), &ContextClaim::Open);
}
assert!(cache.entries.lock().unwrap().len() <= CLAIM_CACHE_CAPACITY);
}
#[tokio::test]
async fn a_cached_owner_outlives_a_row_changed_behind_the_store() {
let dir = tempfile::tempdir().unwrap();
let url = format!("sqlite:{}?mode=rwc", dir.path().join("a2a.db").display());
let storage = SqlxTaskStorage::builder(&url)
.max_connections(1)
.claim_cache(Some(Duration::from_secs(60)))
.connect()
.await
.unwrap();
let context = cid("ctx-cached");
storage.load_state(&context, Some("alice")).await.unwrap();
sqlx::raw_sql("UPDATE contexts SET owner = 'bob' WHERE id = 'ctx-cached'")
.execute(&storage.pool)
.await
.unwrap();
storage
.load_state(&context, Some("alice"))
.await
.expect("the claim alice opened is still the cached one");
drop(storage);
let uncached = SqlxTaskStorage::builder(&url)
.max_connections(1)
.claim_cache(None)
.connect()
.await
.unwrap();
assert!(
matches!(
uncached.load_state(&context, Some("alice")).await,
Err(A2AError::ContextAccessDenied { .. })
),
"a store that reads every time sees the row as it now is"
);
}
#[tokio::test]
async fn a_read_refreshes_a_bag_that_is_old_enough() {
let day = Duration::from_secs(24 * 60 * 60);
let storage = SqlxTaskStorage::builder("sqlite::memory:")
.max_connections(1)
.read_refresh(ReadRefresh::after(day))
.connect()
.await
.unwrap();
let context = cid("ctx-refresh");
let name = StateKey::scoped(StateScope::User, "name").unwrap();
storage
.remember(&context, Some("alice"), &name, "Emil")
.await
.unwrap();
age_the_bag(&storage, "alice").await;
storage.load_state(&context, Some("alice")).await.unwrap();
let policy = RetentionPolicy::keep_everything().delete_user_state_idle_for(day);
let swept = storage.sweep(&policy, chrono::Utc::now()).await.unwrap();
assert_eq!(
swept.state_keys, 0,
"the read moved the bag out of the sweep's reach"
);
}
#[tokio::test]
async fn without_a_refresh_the_same_read_leaves_the_bag_sweepable() {
let day = Duration::from_secs(24 * 60 * 60);
let storage = SqlxTaskStorage::builder("sqlite::memory:")
.max_connections(1)
.connect()
.await
.unwrap();
let context = cid("ctx-no-refresh");
let name = StateKey::scoped(StateScope::User, "name").unwrap();
storage
.remember(&context, Some("alice"), &name, "Emil")
.await
.unwrap();
age_the_bag(&storage, "alice").await;
storage.load_state(&context, Some("alice")).await.unwrap();
let policy = RetentionPolicy::keep_everything().delete_user_state_idle_for(day);
let swept = storage.sweep(&policy, chrono::Utc::now()).await.unwrap();
assert_eq!(swept.state_keys, 1);
}
async fn age_the_bag(storage: &SqlxTaskStorage, principal: &str) {
let sql = storage
.sql("UPDATE context_state SET updated_at = ? WHERE scope = 'user' AND scope_key = ?");
let long_ago = storage
.dialect
.format_timestamp(chrono::Utc::now() - chrono::TimeDelta::days(7));
sqlx::query(&sql)
.bind(long_ago)
.bind(principal)
.execute(&storage.pool)
.await
.unwrap();
}
#[tokio::test]
async fn a_swept_context_leaves_nothing_cached() {
let storage = SqlxTaskStorage::builder("sqlite::memory:")
.max_connections(1)
.claim_cache(Some(Duration::from_secs(60)))
.connect()
.await
.unwrap();
let context = cid("ctx-swept");
storage.load_state(&context, Some("alice")).await.unwrap();
let policy =
RetentionPolicy::keep_everything().delete_contexts_idle_for(Duration::from_secs(60));
let swept = storage
.sweep(&policy, chrono::Utc::now() + chrono::TimeDelta::days(30))
.await
.unwrap();
assert_eq!(swept.contexts, 1);
storage
.load_state(&context, Some("bob"))
.await
.expect("nobody holds a context that was swept");
}
}