1#[cfg(feature = "sqlx-storage")]
7use async_trait::async_trait;
8#[cfg(feature = "sqlx-storage")]
9use serde_json;
10#[cfg(feature = "sqlx-storage")]
11use sqlx::{Row, SqlitePool};
12
13#[cfg(feature = "sqlx-storage")]
14use crate::adapter::business::push_notification::{
15 PushNotificationRegistry, PushNotificationSender,
16};
17
18#[cfg(feature = "sqlx-storage")]
19#[cfg(feature = "http-client")]
20use crate::adapter::business::push_notification::HttpPushNotificationSender;
21#[cfg(feature = "sqlx-storage")]
22#[cfg(not(feature = "http-client"))]
23use crate::adapter::business::push_notification::NoopPushNotificationSender;
24
25#[cfg(feature = "sqlx-storage")]
26use crate::domain::{
27 A2AError, ContextId, Message, Task, TaskId, TaskPushNotificationConfig, TaskState,
28 TaskStateExt, TaskStatus, VersionedTask,
29};
30#[cfg(feature = "sqlx-storage")]
31use crate::port::{
32 AsyncNotificationManager, AsyncPushNotifier, AsyncTaskLifecycle, AsyncTaskQuery,
33 AsyncTaskVersioning,
34};
35
36#[cfg(feature = "sqlx-storage")]
37use std::sync::Arc;
38
39#[cfg(feature = "sqlx-storage")]
40pub struct SqlxTaskStorage {
48 pool: SqlitePool,
50 push_notification_registry: Arc<PushNotificationRegistry>,
52}
53
54#[cfg(feature = "sqlx-storage")]
55use super::database_config::DatabaseType;
56
57#[cfg(feature = "sqlx-storage")]
58impl SqlxTaskStorage {
59 fn validate_url(database_url: &str) -> Result<(), A2AError> {
64 match DatabaseType::from_url(database_url) {
65 Some(DatabaseType::Sqlite) => Ok(()),
66 Some(db_type) => Err(A2AError::DatabaseError(format!(
67 "{db_type} database detected from URL '{database_url}', but SqlxTaskStorage \
68 currently only supports SQLite. For {db_type} support, see the project roadmap."
69 ))),
70 None => Err(A2AError::DatabaseError(format!(
71 "Unrecognized database URL scheme in '{database_url}'. \
72 Expected a URL starting with sqlite:, e.g. 'sqlite::memory:' or 'sqlite:data.db'"
73 ))),
74 }
75 }
76
77 pub async fn new(database_url: &str) -> Result<Self, A2AError> {
82 Self::validate_url(database_url)?;
83
84 let pool = SqlitePool::connect(database_url).await.map_err(|e| {
85 A2AError::DatabaseError(format!("Failed to connect to database: {}", e))
86 })?;
87
88 Self::run_base_migrations(&pool).await?;
90
91 #[cfg(feature = "http-client")]
93 let push_sender = HttpPushNotificationSender::new();
94 #[cfg(not(feature = "http-client"))]
95 let push_sender = NoopPushNotificationSender::default();
96
97 let push_registry = PushNotificationRegistry::new(push_sender);
98
99 Ok(Self {
100 pool,
101 push_notification_registry: Arc::new(push_registry),
102 })
103 }
104
105 pub async fn with_push_sender(
109 database_url: &str,
110 push_sender: impl PushNotificationSender + 'static,
111 ) -> Result<Self, A2AError> {
112 Self::validate_url(database_url)?;
113
114 let pool = SqlitePool::connect(database_url).await.map_err(|e| {
115 A2AError::DatabaseError(format!("Failed to connect to database: {}", e))
116 })?;
117
118 Self::run_base_migrations(&pool).await?;
120
121 let push_registry = PushNotificationRegistry::new(push_sender);
122
123 Ok(Self {
124 pool,
125 push_notification_registry: Arc::new(push_registry),
126 })
127 }
128
129 pub async fn with_migrations(
133 database_url: &str,
134 additional_migrations: &[&str],
135 ) -> Result<Self, A2AError> {
136 Self::validate_url(database_url)?;
137
138 let pool = SqlitePool::connect(database_url).await.map_err(|e| {
139 A2AError::DatabaseError(format!("Failed to connect to database: {}", e))
140 })?;
141
142 Self::run_base_migrations(&pool).await?;
144
145 Self::run_additional_migrations(&pool, additional_migrations).await?;
147
148 #[cfg(feature = "http-client")]
150 let push_sender = HttpPushNotificationSender::new();
151 #[cfg(not(feature = "http-client"))]
152 let push_sender = NoopPushNotificationSender::default();
153
154 let push_registry = PushNotificationRegistry::new(push_sender);
155
156 Ok(Self {
157 pool,
158 push_notification_registry: Arc::new(push_registry),
159 })
160 }
161
162 async fn run_base_migrations(pool: &SqlitePool) -> Result<(), A2AError> {
164 sqlx::query(include_str!("../../../migrations/001_initial_schema.sql"))
165 .execute(pool)
166 .await
167 .map_err(|e| A2AError::DatabaseError(format!("Migration 001 failed: {}", e)))?;
168
169 sqlx::query(include_str!(
170 "../../../migrations/002_v030_push_configs.sql"
171 ))
172 .execute(pool)
173 .await
174 .map_err(|e| A2AError::DatabaseError(format!("Migration 002 failed: {}", e)))?;
175
176 if let Err(e) = sqlx::query(include_str!("../../../migrations/003_task_version.sql"))
180 .execute(pool)
181 .await
182 {
183 let msg = e.to_string();
184 if !msg.contains("duplicate column name") {
185 return Err(A2AError::DatabaseError(format!(
186 "Migration 003 failed: {msg}"
187 )));
188 }
189 }
190
191 Ok(())
192 }
193
194 async fn run_additional_migrations(
196 pool: &SqlitePool,
197 migrations: &[&str],
198 ) -> Result<(), A2AError> {
199 for (i, migration_sql) in migrations.iter().enumerate() {
200 sqlx::query(migration_sql)
201 .execute(pool)
202 .await
203 .map_err(|e| {
204 A2AError::DatabaseError(format!("Additional migration {} failed: {}", i + 1, e))
205 })?;
206 }
207 Ok(())
208 }
209
210 fn row_to_task(row: &sqlx::sqlite::SqliteRow) -> Result<Task, A2AError> {
212 let task_id: String = row
213 .try_get("id")
214 .map_err(|e| A2AError::DatabaseError(format!("Failed to get task_id: {}", e)))?;
215 let context_id: String = row
216 .try_get("context_id")
217 .map_err(|e| A2AError::DatabaseError(format!("Failed to get context_id: {}", e)))?;
218 let status_state: String = row
219 .try_get("status_state")
220 .map_err(|e| A2AError::DatabaseError(format!("Failed to get status_state: {}", e)))?;
221 let status_message_json: Option<String> = row
222 .try_get("status_message")
223 .map_err(|e| A2AError::DatabaseError(format!("Failed to get status_message: {}", e)))?;
224 let metadata_json: Option<String> = row
225 .try_get("metadata")
226 .map_err(|e| A2AError::DatabaseError(format!("Failed to get metadata: {}", e)))?;
227 let artifacts_json: Option<String> = row
228 .try_get("artifacts")
229 .map_err(|e| A2AError::DatabaseError(format!("Failed to get artifacts: {}", e)))?;
230
231 let state = match status_state.as_str() {
233 "submitted" => TaskState::Submitted,
234 "working" => TaskState::Working,
235 "input-required" => TaskState::InputRequired,
236 "completed" => TaskState::Completed,
237 "canceled" => TaskState::Canceled,
238 "failed" => TaskState::Failed,
239 "rejected" => TaskState::Rejected,
240 "auth-required" => TaskState::AuthRequired,
241 "unknown" => TaskState::Unknown,
242 _ => TaskState::Unknown,
243 };
244
245 let status_message = if let Some(msg_str) = status_message_json {
247 Some(serde_json::from_str(&msg_str).map_err(|e| {
248 A2AError::DatabaseError(format!("Failed to parse status message: {}", e))
249 })?)
250 } else {
251 None
252 };
253
254 let metadata =
256 if let Some(meta_str) = metadata_json {
257 Some(serde_json::from_str(&meta_str).map_err(|e| {
258 A2AError::DatabaseError(format!("Failed to parse metadata: {}", e))
259 })?)
260 } else {
261 None
262 };
263
264 let artifacts = if let Some(artifacts_str) = artifacts_json {
266 Some(serde_json::from_str(&artifacts_str).map_err(|e| {
267 A2AError::DatabaseError(format!("Failed to parse artifacts: {}", e))
268 })?)
269 } else {
270 None
271 };
272
273 let now = chrono::Utc::now();
274 let task_status = TaskStatus {
275 state: ::buffa::EnumValue::from(state),
276 message: status_message.into(),
277 timestamp: ::buffa::MessageField::some(::buffa_types::google::protobuf::Timestamp {
278 seconds: now.timestamp(),
279 nanos: now.timestamp_subsec_nanos() as i32,
280 ..Default::default()
281 }),
282 ..Default::default()
283 };
284
285 let task = Task {
286 id: task_id.clone(),
287 context_id,
288 status: ::buffa::MessageField::some(task_status),
289 history: Vec::new(),
290 metadata: metadata.into(),
291 artifacts: artifacts.unwrap_or_default(),
292 ..Default::default()
293 };
294
295 Ok(task)
296 }
297
298 async fn load_task_history(
300 &self,
301 task_id: &str,
302 limit: Option<u32>,
303 ) -> Result<Vec<Message>, A2AError> {
304 let query_str = if let Some(limit) = limit {
305 format!(
306 "SELECT timestamp, status_state, message FROM task_history WHERE task_id = ? ORDER BY timestamp DESC LIMIT {}",
307 limit
308 )
309 } else {
310 "SELECT timestamp, status_state, message FROM task_history WHERE task_id = ? ORDER BY timestamp DESC".to_string()
311 };
312
313 let query = sqlx::query(&query_str);
314
315 let rows = query
316 .bind(task_id)
317 .fetch_all(&self.pool)
318 .await
319 .map_err(|e| A2AError::DatabaseError(format!("Failed to load task history: {}", e)))?;
320
321 let mut history = Vec::new();
322 for row in rows {
323 let message_json: Option<String> = row.try_get("message").map_err(|e| {
324 A2AError::DatabaseError(format!("Failed to get message from history: {}", e))
325 })?;
326
327 if let Some(msg_str) = message_json {
328 let message: Message = serde_json::from_str(&msg_str).map_err(|e| {
329 A2AError::DatabaseError(format!("Failed to parse message from history: {}", e))
330 })?;
331 history.push(message);
332 }
333 }
334
335 history.reverse();
337 Ok(history)
338 }
339
340 async fn add_to_history(
342 &self,
343 task_id: &str,
344 state: TaskState,
345 message: Option<Message>,
346 ) -> Result<(), A2AError> {
347 let state_str = match state {
348 TaskState::Submitted => "submitted",
349 TaskState::Working => "working",
350 TaskState::InputRequired => "input-required",
351 TaskState::Completed => "completed",
352 TaskState::Canceled => "canceled",
353 TaskState::Failed => "failed",
354 TaskState::Rejected => "rejected",
355 TaskState::AuthRequired => "auth-required",
356 TaskState::Unknown => "unknown",
357 };
358
359 let message_json = if let Some(msg) = message {
360 Some(serde_json::to_string(&msg).map_err(|e| {
361 A2AError::DatabaseError(format!("Failed to serialize message: {}", e))
362 })?)
363 } else {
364 None
365 };
366
367 sqlx::query("INSERT INTO task_history (task_id, status_state, message) VALUES (?, ?, ?)")
368 .bind(task_id)
369 .bind(state_str)
370 .bind(message_json)
371 .execute(&self.pool)
372 .await
373 .map_err(|e| A2AError::DatabaseError(format!("Failed to add task history: {}", e)))?;
374
375 Ok(())
376 }
377
378 pub fn push_notifier(&self) -> Arc<dyn AsyncPushNotifier> {
385 self.push_notification_registry.clone()
386 }
387}
388
389#[cfg(feature = "sqlx-storage")]
390#[async_trait]
391impl AsyncTaskLifecycle for SqlxTaskStorage {
392 async fn create(&self, id: &TaskId, context_id: &ContextId) -> Result<Task, A2AError> {
393 let task_id = id.as_str();
394 let context_id = context_id.as_str();
395 let existing = sqlx::query("SELECT id FROM tasks WHERE id = ?")
397 .bind(task_id)
398 .fetch_optional(&self.pool)
399 .await
400 .map_err(|e| {
401 A2AError::DatabaseError(format!("Failed to check existing task: {}", e))
402 })?;
403
404 if existing.is_some() {
405 return Err(A2AError::TaskNotFound(format!(
406 "Task {} already exists",
407 task_id
408 )));
409 }
410
411 let task = Task::new(task_id.to_string(), context_id.to_string());
413
414 let metadata_json = task
416 .metadata
417 .as_option()
418 .map(|m| serde_json::to_string(m).unwrap_or_default());
419 let artifacts_json = serde_json::to_string(&task.artifacts).unwrap_or_default();
420 let status_message_str = task
421 .status
422 .as_option()
423 .and_then(|s| s.message.as_option())
424 .map(|m| serde_json::to_string(m).unwrap_or_default());
425
426 sqlx::query("INSERT INTO tasks (id, context_id, status_state, status_message, metadata, artifacts) VALUES (?, ?, ?, ?, ?, ?)")
428 .bind(&task.id)
429 .bind(&task.context_id)
430 .bind("submitted")
431 .bind(status_message_str)
432 .bind(metadata_json)
433 .bind(artifacts_json)
434 .execute(&self.pool)
435 .await
436 .map_err(|e| A2AError::DatabaseError(format!("Failed to create task: {}", e)))?;
437
438 self.add_to_history(task_id, TaskState::Submitted, None)
440 .await?;
441
442 Ok(task)
443 }
444
445 async fn update_status(
446 &self,
447 id: &TaskId,
448 state: TaskState,
449 message: Option<Message>,
450 ) -> Result<Task, A2AError> {
451 let task_id = id.as_str();
452 let state_str = match state {
454 TaskState::Submitted => "submitted",
455 TaskState::Working => "working",
456 TaskState::InputRequired => "input-required",
457 TaskState::Completed => "completed",
458 TaskState::Canceled => "canceled",
459 TaskState::Failed => "failed",
460 TaskState::Rejected => "rejected",
461 TaskState::AuthRequired => "auth-required",
462 TaskState::Unknown => "unknown",
463 };
464
465 let result =
467 sqlx::query("UPDATE tasks SET status_state = ?, version = version + 1 WHERE id = ?")
468 .bind(state_str)
469 .bind(task_id)
470 .execute(&self.pool)
471 .await
472 .map_err(|e| {
473 A2AError::DatabaseError(format!("Failed to update task status: {}", e))
474 })?;
475
476 if result.rows_affected() == 0 {
477 return Err(A2AError::TaskNotFound(task_id.to_string()));
478 }
479
480 self.add_to_history(task_id, state, message).await?;
482
483 self.get(id, None).await
487 }
488
489 async fn exists(&self, id: &TaskId) -> Result<bool, A2AError> {
490 let task_id = id.as_str();
491 let row = sqlx::query("SELECT id FROM tasks WHERE id = ?")
492 .bind(task_id)
493 .fetch_optional(&self.pool)
494 .await
495 .map_err(|e| {
496 A2AError::DatabaseError(format!("Failed to check task existence: {}", e))
497 })?;
498
499 Ok(row.is_some())
500 }
501
502 async fn get(&self, id: &TaskId, history_length: Option<u32>) -> Result<Task, A2AError> {
503 let task_id = id.as_str();
504 let row = sqlx::query("SELECT * FROM tasks WHERE id = ?")
506 .bind(task_id)
507 .fetch_optional(&self.pool)
508 .await
509 .map_err(|e| A2AError::DatabaseError(format!("Failed to get task: {}", e)))?;
510
511 let Some(row) = row else {
512 return Err(A2AError::TaskNotFound(task_id.to_string()));
513 };
514
515 let mut task = Self::row_to_task(&row)?;
516
517 if history_length.is_some() || history_length.is_none() {
519 let history = self.load_task_history(task_id, history_length).await?;
520 task.history = history;
521 }
522
523 Ok(task)
524 }
525
526 async fn cancel(&self, id: &TaskId) -> Result<Task, A2AError> {
527 let task_id = id.as_str();
528 let task = self.get(id, None).await?;
530
531 if !task.status.state.is_cancelable() {
536 return Err(A2AError::TaskNotCancelable(format!(
537 "Task {} has already finished in state {:?} and cannot be canceled",
538 task_id, task.status.state
539 )));
540 }
541
542 let mut cancel_message = Message::agent_text(
544 format!("Task {} canceled.", task_id),
545 uuid::Uuid::new_v4().to_string(),
546 );
547 cancel_message.task_id = task_id.to_string();
548 cancel_message.context_id = task.context_id.clone();
549
550 sqlx::query("UPDATE tasks SET status_state = ?, version = version + 1 WHERE id = ?")
552 .bind("canceled")
553 .bind(task_id)
554 .execute(&self.pool)
555 .await
556 .map_err(|e| A2AError::DatabaseError(format!("Failed to cancel task: {}", e)))?;
557
558 self.add_to_history(task_id, TaskState::Canceled, Some(cancel_message))
560 .await?;
561
562 self.get(id, None).await
565 }
566}
567
568#[cfg(feature = "sqlx-storage")]
569impl SqlxTaskStorage {
570 async fn current_version(&self, task_id: &str) -> Result<Option<u64>, A2AError> {
572 let row = sqlx::query("SELECT version FROM tasks WHERE id = ?")
573 .bind(task_id)
574 .fetch_optional(&self.pool)
575 .await
576 .map_err(|e| A2AError::DatabaseError(format!("Failed to read task version: {}", e)))?;
577 match row {
578 Some(row) => {
579 let v: i64 = row.try_get("version").map_err(|e| {
580 A2AError::DatabaseError(format!("Failed to get version column: {}", e))
581 })?;
582 Ok(Some(v as u64))
583 }
584 None => Ok(None),
585 }
586 }
587}
588
589#[cfg(feature = "sqlx-storage")]
590#[async_trait]
591impl AsyncTaskVersioning for SqlxTaskStorage {
592 async fn version(&self, id: &TaskId) -> Result<u64, A2AError> {
593 self.current_version(id.as_str())
594 .await?
595 .ok_or_else(|| A2AError::TaskNotFound(id.as_str().to_string()))
596 }
597
598 async fn get_versioned(
599 &self,
600 id: &TaskId,
601 history_length: Option<u32>,
602 ) -> Result<VersionedTask, A2AError> {
603 let task = self.get(id, history_length).await?;
604 let version = self.version(id).await?;
605 Ok(VersionedTask::new(task, version))
606 }
607
608 async fn update_status_checked(
609 &self,
610 id: &TaskId,
611 expected: u64,
612 state: TaskState,
613 message: Option<Message>,
614 ) -> Result<VersionedTask, A2AError> {
615 let task_id = id.as_str();
616 let state_str = match state {
617 TaskState::Submitted => "submitted",
618 TaskState::Working => "working",
619 TaskState::InputRequired => "input-required",
620 TaskState::Completed => "completed",
621 TaskState::Canceled => "canceled",
622 TaskState::Failed => "failed",
623 TaskState::Rejected => "rejected",
624 TaskState::AuthRequired => "auth-required",
625 TaskState::Unknown => "unknown",
626 };
627
628 let result = sqlx::query(
631 "UPDATE tasks SET status_state = ?, version = version + 1 WHERE id = ? AND version = ?",
632 )
633 .bind(state_str)
634 .bind(task_id)
635 .bind(expected as i64)
636 .execute(&self.pool)
637 .await
638 .map_err(|e| A2AError::DatabaseError(format!("Failed to update task status: {}", e)))?;
639
640 if result.rows_affected() == 0 {
641 return match self.current_version(task_id).await? {
643 Some(actual) => Err(A2AError::VersionConflict {
644 id: task_id.to_string(),
645 expected,
646 actual,
647 }),
648 None => Err(A2AError::TaskNotFound(task_id.to_string())),
649 };
650 }
651
652 self.add_to_history(task_id, state, message).await?;
653 let task = self.get(id, None).await?;
654 Ok(VersionedTask::new(task, expected + 1))
655 }
656}
657
658#[cfg(feature = "sqlx-storage")]
659#[async_trait]
660impl AsyncTaskQuery for SqlxTaskStorage {
661 async fn list(
662 &self,
663 params: &crate::domain::ListTasksParams,
664 ) -> Result<crate::domain::ListTasksResult, A2AError> {
665 use crate::domain::ListTasksResult;
666
667 let mut where_conditions = Vec::new();
669
670 if params.context_id.is_some() {
672 where_conditions.push("context_id = ?".to_string());
673 }
674
675 if params.status.is_some() {
677 where_conditions.push("status_state = ?".to_string());
678 }
679
680 let timestamp_str = if let Some(status_timestamp_after) = ¶ms.status_timestamp_after {
682 let timestamp =
684 chrono::DateTime::parse_from_rfc3339(status_timestamp_after).map_err(|e| {
685 A2AError::DatabaseError(format!(
686 "Invalid timestamp value: {} ({})",
687 status_timestamp_after, e
688 ))
689 })?;
690 where_conditions.push("updated_at >= ?".to_string());
691 Some(
692 timestamp
693 .with_timezone(&chrono::Utc)
694 .format("%Y-%m-%d %H:%M:%S")
695 .to_string(),
696 )
697 } else {
698 None
699 };
700
701 let where_clause = if where_conditions.is_empty() {
703 String::new()
704 } else {
705 format!(" WHERE {}", where_conditions.join(" AND "))
706 };
707
708 let count_query = format!("SELECT COUNT(*) as count FROM tasks{}", where_clause);
710 let mut count_q = sqlx::query(&count_query);
711
712 if let Some(ref context_id) = params.context_id {
714 count_q = count_q.bind(context_id);
715 }
716 if let Some(ref status) = params.status {
717 let state_str = match *status {
718 crate::domain::TaskState::Submitted => "submitted",
719 crate::domain::TaskState::Working => "working",
720 crate::domain::TaskState::InputRequired => "input-required",
721 crate::domain::TaskState::Completed => "completed",
722 crate::domain::TaskState::Canceled => "canceled",
723 crate::domain::TaskState::Failed => "failed",
724 crate::domain::TaskState::Rejected => "rejected",
725 crate::domain::TaskState::AuthRequired => "auth-required",
726 crate::domain::TaskState::Unknown => "unknown",
727 };
728 count_q = count_q.bind(state_str);
729 }
730 if let Some(ref ts) = timestamp_str {
731 count_q = count_q.bind(ts);
732 }
733
734 let count_row = count_q
735 .fetch_one(&self.pool)
736 .await
737 .map_err(|e| A2AError::DatabaseError(format!("Failed to count tasks: {}", e)))?;
738
739 let total_size: i32 = count_row
740 .try_get("count")
741 .map_err(|e| A2AError::DatabaseError(format!("Failed to get count: {}", e)))?;
742
743 let page_size = params.page_size.unwrap_or(50).clamp(1, 100);
745 let offset = if let Some(ref token) = params.page_token {
746 token.parse::<i32>().unwrap_or(0)
747 } else {
748 0
749 };
750
751 let main_query = format!(
753 "SELECT * FROM tasks{} ORDER BY updated_at DESC LIMIT ? OFFSET ?",
754 where_clause
755 );
756
757 let mut main_q = sqlx::query(&main_query);
758
759 if let Some(ref context_id) = params.context_id {
761 main_q = main_q.bind(context_id);
762 }
763 if let Some(ref status) = params.status {
764 let state_str = match *status {
765 crate::domain::TaskState::Submitted => "submitted",
766 crate::domain::TaskState::Working => "working",
767 crate::domain::TaskState::InputRequired => "input-required",
768 crate::domain::TaskState::Completed => "completed",
769 crate::domain::TaskState::Canceled => "canceled",
770 crate::domain::TaskState::Failed => "failed",
771 crate::domain::TaskState::Rejected => "rejected",
772 crate::domain::TaskState::AuthRequired => "auth-required",
773 crate::domain::TaskState::Unknown => "unknown",
774 };
775 main_q = main_q.bind(state_str);
776 }
777 if let Some(ref ts) = timestamp_str {
778 main_q = main_q.bind(ts);
779 }
780
781 main_q = main_q.bind(page_size).bind(offset);
783
784 let rows = main_q
785 .fetch_all(&self.pool)
786 .await
787 .map_err(|e| A2AError::DatabaseError(format!("Failed to list tasks: {}", e)))?;
788
789 let mut tasks: Vec<Task> = rows
791 .iter()
792 .filter_map(|row| Self::row_to_task(row).ok())
793 .collect();
794
795 let history_length = params.history_length.unwrap_or(0);
797 for task in &mut tasks {
798 if history_length > 0 {
799 let history = self
800 .load_task_history(&task.id, Some(history_length as u32))
801 .await?;
802 task.history = history;
803 } else {
804 task.history.clear();
805 }
806
807 if !params.include_artifacts.unwrap_or(false) {
809 task.artifacts.clear();
810 }
811 }
812
813 let has_more = offset + page_size < total_size;
815 let next_page_token = if has_more {
816 (offset + page_size).to_string()
817 } else {
818 String::new()
819 };
820
821 Ok(ListTasksResult {
822 tasks,
823 total_size,
824 page_size,
825 next_page_token,
826 })
827 }
828}
829
830#[cfg(feature = "sqlx-storage")]
831#[async_trait]
832impl AsyncNotificationManager for SqlxTaskStorage {
833 async fn get_config(
834 &self,
835 params: &crate::domain::GetTaskPushNotificationConfigParams,
836 ) -> Result<crate::domain::TaskPushNotificationConfig, A2AError> {
837 let row = match params.push_notification_config_id.as_ref() {
842 Some(config_id) => sqlx::query(
843 "SELECT id, task_id, url, token, authentication FROM push_notification_configs WHERE task_id = ? AND id = ?"
844 )
845 .bind(¶ms.id)
846 .bind(config_id),
847 None => sqlx::query(
848 "SELECT id, task_id, url, token, authentication FROM push_notification_configs WHERE task_id = ? ORDER BY id LIMIT 1"
849 )
850 .bind(¶ms.id),
851 }
852 .fetch_optional(&self.pool)
853 .await
854 .map_err(|e| A2AError::DatabaseError(format!("Failed to get push config: {}", e)))?;
855
856 if let Some(row) = row {
857 let id: String = row
858 .try_get("id")
859 .map_err(|e| A2AError::DatabaseError(format!("Failed to get config id: {}", e)))?;
860 let url: String = row
861 .try_get("url")
862 .map_err(|e| A2AError::DatabaseError(format!("Failed to get url: {}", e)))?;
863 let token: Option<String> = row.try_get("token").ok();
864 let auth_json: Option<String> = row.try_get("authentication").ok();
865
866 let auth_info = if let Some(auth_str) = auth_json {
867 serde_json::from_str(&auth_str).ok()
868 } else {
869 None
870 };
871
872 Ok(crate::domain::TaskPushNotificationConfig {
873 task_id: params.id.clone(),
874 id,
875 url,
876 token: token.unwrap_or_default(),
877 authentication: auth_info.into(),
878 tenant: "".to_string(),
879 ..Default::default()
880 })
881 } else {
882 Err(A2AError::TaskNotFound(format!(
883 "Push notification config not found for task {}{}",
884 params.id,
885 params
886 .push_notification_config_id
887 .as_ref()
888 .map(|id| format!(" with id {}", id))
889 .unwrap_or_default()
890 )))
891 }
892 }
893
894 async fn list_configs(
895 &self,
896 params: &crate::domain::ListTaskPushNotificationConfigsParams,
897 ) -> Result<Vec<crate::domain::TaskPushNotificationConfig>, A2AError> {
898 let rows = sqlx::query(
900 "SELECT id, task_id, url, token, authentication FROM push_notification_configs WHERE task_id = ?"
901 )
902 .bind(¶ms.id)
903 .fetch_all(&self.pool)
904 .await
905 .map_err(|e| A2AError::DatabaseError(format!("Failed to list push configs: {}", e)))?;
906
907 let configs: Vec<crate::domain::TaskPushNotificationConfig> = rows
908 .iter()
909 .filter_map(|row| {
910 let id: String = row.try_get("id").ok()?;
911 let url: String = row.try_get("url").ok()?;
912 let token: Option<String> = row.try_get("token").ok().flatten();
913 let auth_json: Option<String> = row.try_get("authentication").ok().flatten();
914
915 let auth_info = if let Some(auth_str) = auth_json {
916 serde_json::from_str(&auth_str).ok()
917 } else {
918 None
919 };
920
921 Some(crate::domain::TaskPushNotificationConfig {
922 task_id: params.id.clone(),
923 id,
924 url,
925 token: token.unwrap_or_default(),
926 authentication: auth_info.into(),
927 tenant: "".to_string(),
928 ..Default::default()
929 })
930 })
931 .collect();
932
933 Ok(configs)
934 }
935
936 async fn delete_config(
937 &self,
938 params: &crate::domain::DeleteTaskPushNotificationConfigParams,
939 ) -> Result<(), A2AError> {
940 let query = if params.push_notification_config_id.is_empty() {
944 sqlx::query("DELETE FROM push_notification_configs WHERE task_id = ?").bind(¶ms.id)
945 } else {
946 sqlx::query("DELETE FROM push_notification_configs WHERE task_id = ? AND id = ?")
947 .bind(¶ms.id)
948 .bind(¶ms.push_notification_config_id)
949 };
950 let _result = query
951 .execute(&self.pool)
952 .await
953 .map_err(|e| A2AError::DatabaseError(format!("Failed to delete push config: {}", e)))?;
954
955 Ok(())
957 }
958
959 async fn set_config(
960 &self,
961 config: &TaskPushNotificationConfig,
962 ) -> Result<TaskPushNotificationConfig, A2AError> {
963 let config_id = if config.id.is_empty() {
965 uuid::Uuid::new_v4().to_string()
966 } else {
967 config.id.clone()
968 };
969
970 let auth_json = config
972 .authentication
973 .as_option()
974 .map(|auth| serde_json::to_string(auth).unwrap_or_default());
975
976 sqlx::query(
978 "INSERT OR REPLACE INTO push_notification_configs (id, task_id, url, token, authentication) VALUES (?, ?, ?, ?, ?)",
979 )
980 .bind(&config_id)
981 .bind(&config.task_id)
982 .bind(&config.url)
983 .bind(&config.token)
984 .bind(auth_json)
985 .execute(&self.pool)
986 .await
987 .map_err(|e| {
988 A2AError::DatabaseError(format!("Failed to set push notification config: {}", e))
989 })?;
990
991 self.push_notification_registry
993 .register(&config.task_id, config.clone())
994 .await?;
995
996 let mut result_config = config.clone();
998 result_config.id = config_id;
999 Ok(result_config)
1000 }
1001}
1002
1003#[cfg(feature = "sqlx-storage")]
1004impl Clone for SqlxTaskStorage {
1005 fn clone(&self) -> Self {
1006 Self {
1007 pool: self.pool.clone(),
1008 push_notification_registry: self.push_notification_registry.clone(),
1009 }
1010 }
1011}