1use std::collections::{BTreeSet, HashMap};
103use std::fmt::Write as _;
104use std::sync::atomic::{AtomicBool, Ordering};
105use std::sync::{Arc, RwLock};
106use std::time::{Duration, Instant};
107
108use async_trait::async_trait;
109
110use crate::error::JsonRpcError;
111use crate::protocol::{CallToolResult, InputRequests, InputResponses, TaskObject, TaskStatus};
112
113const DEFAULT_TTL_MS: u64 = 300_000;
118
119const DEFAULT_POLL_INTERVAL_MS: u64 = 2_000;
121
122#[derive(Debug)]
124pub struct Task {
125 pub id: String,
127 pub tool_name: String,
129 pub arguments: serde_json::Value,
131 pub status: TaskStatus,
133 pub created_at: Instant,
135 pub created_at_str: String,
137 pub last_updated_at_str: String,
139 pub ttl: u64,
141 pub poll_interval: u64,
143 pub status_message: Option<String>,
145 pub meta: Option<serde_json::Value>,
147 pub result: Option<CallToolResult>,
149 pub error: Option<JsonRpcError>,
155 pub owner: TaskOwner,
160 pub input_requests: InputRequests,
162 pub answered_input_keys: BTreeSet<String>,
164 pub input_responses: InputResponses,
171 pub superseded_input_keys: BTreeSet<String>,
174 pub cancellation_token: CancellationToken,
176 pub completed_at: Option<Instant>,
178 pub completion_notify: Arc<tokio::sync::Notify>,
180}
181
182impl Task {
183 fn new(
185 id: String,
186 tool_name: String,
187 arguments: serde_json::Value,
188 ttl: Option<u64>,
189 owner: TaskOwner,
190 ) -> Self {
191 let cancelled = Arc::new(AtomicBool::new(false));
192 let now_str = chrono_now_iso8601();
193 Self {
194 id,
195 tool_name,
196 arguments,
197 status: TaskStatus::Working,
198 created_at: Instant::now(),
199 created_at_str: now_str.clone(),
200 last_updated_at_str: now_str,
201 ttl: ttl.unwrap_or(DEFAULT_TTL_MS),
202 poll_interval: DEFAULT_POLL_INTERVAL_MS,
203 status_message: Some("Task started".to_string()),
204 meta: None,
205 result: None,
206 error: None,
207 owner,
208 input_requests: InputRequests::new(),
209 answered_input_keys: BTreeSet::new(),
210 input_responses: InputResponses::new(),
211 superseded_input_keys: BTreeSet::new(),
212 cancellation_token: CancellationToken { cancelled },
213 completed_at: None,
214 completion_notify: Arc::new(tokio::sync::Notify::new()),
215 }
216 }
217
218 pub fn to_task_object(&self) -> TaskObject {
220 TaskObject {
221 task_id: self.id.clone(),
222 status: self.status,
223 status_message: self.status_message.clone(),
224 created_at: self.created_at_str.clone(),
225 last_updated_at: self.last_updated_at_str.clone(),
226 ttl: Some(self.ttl),
227 poll_interval: Some(self.poll_interval),
228 result: None,
229 error: None,
230 meta: self.meta.clone(),
231 }
232 }
233
234 pub fn is_expired(&self) -> bool {
241 self.created_at.elapsed() > Duration::from_millis(self.ttl)
242 }
243
244 pub fn outstanding_input_requests(&self) -> &InputRequests {
246 &self.input_requests
247 }
248
249 pub fn is_cancelled(&self) -> bool {
251 self.cancellation_token.is_cancelled()
252 }
253}
254
255pub fn generate_task_id() -> String {
268 let mut bytes = [0u8; 16];
269 getrandom::fill(&mut bytes).expect("system entropy source unavailable for task ID generation");
270 let mut id = String::with_capacity(2 * bytes.len());
271 for byte in bytes {
272 let _ = write!(id, "{byte:02x}");
273 }
274 id
275}
276
277pub type TaskOwner = Option<String>;
286
287pub fn owner_matches(owner: &TaskOwner, principal: Option<&str>) -> bool {
294 owner.as_deref() == principal
295}
296
297#[derive(Debug, Clone, Default, PartialEq, Eq)]
303pub struct AppliedInputResponses {
304 pub accepted: BTreeSet<String>,
306 pub ignored: BTreeSet<String>,
309 pub still_outstanding: BTreeSet<String>,
311}
312
313#[derive(Debug, Clone)]
319#[non_exhaustive]
320pub struct TaskResumeContext {
321 pub tool_name: String,
323 pub arguments: serde_json::Value,
325 pub input_responses: InputResponses,
327}
328
329impl AppliedInputResponses {
330 pub fn is_complete(&self) -> bool {
332 self.still_outstanding.is_empty()
333 }
334}
335
336#[derive(Debug, Clone)]
338pub struct CancellationToken {
339 cancelled: Arc<AtomicBool>,
340}
341
342impl CancellationToken {
343 pub fn is_cancelled(&self) -> bool {
345 self.cancelled.load(Ordering::Relaxed)
346 }
347
348 pub fn cancel(&self) {
350 self.cancelled.store(true, Ordering::Relaxed);
351 }
352}
353
354#[derive(Debug, thiserror::Error)]
362#[non_exhaustive]
363pub enum TaskStoreError {
364 #[error("encode error: {0}")]
366 Encode(String),
367 #[error("decode error: {0}")]
369 Decode(String),
370 #[error("backend error: {0}")]
372 Backend(String),
373}
374
375pub type Result<T> = std::result::Result<T, TaskStoreError>;
377
378pub type TaskSnapshot = (TaskObject, Option<CallToolResult>, Option<JsonRpcError>);
384
385#[async_trait]
407pub trait TaskStore: Send + Sync + 'static {
408 async fn create_task(
414 &self,
415 tool_name: &str,
416 arguments: serde_json::Value,
417 ttl: Option<u64>,
418 owner: TaskOwner,
419 ) -> Result<(String, CancellationToken)>;
420
421 async fn task_owner(&self, task_id: &str) -> Result<Option<TaskOwner>>;
427
428 async fn get_task(&self, task_id: &str) -> Result<Option<TaskObject>>;
430
431 async fn set_task_meta(&self, task_id: &str, meta: serde_json::Value) -> Result<bool> {
436 let _ = (task_id, meta);
437 Ok(false)
438 }
439
440 async fn discard_task(&self, task_id: &str) -> Result<bool> {
445 let _ = task_id;
446 Ok(false)
447 }
448
449 async fn get_task_result(&self, task_id: &str) -> Result<Option<TaskSnapshot>>;
451
452 async fn wait_for_completion(&self, task_id: &str) -> Result<Option<TaskSnapshot>>;
458
459 async fn list_tasks(&self, status_filter: Option<TaskStatus>) -> Result<Vec<TaskObject>>;
461
462 async fn require_input(
471 &self,
472 task_id: &str,
473 requests: InputRequests,
474 message: Option<&str>,
475 ) -> Result<bool>;
476
477 async fn outstanding_input_requests(&self, task_id: &str) -> Result<Option<InputRequests>>;
482
483 async fn apply_input_responses(
491 &self,
492 task_id: &str,
493 responses: InputResponses,
494 ) -> Result<Option<AppliedInputResponses>>;
495
496 async fn resume_context(&self, task_id: &str) -> Result<Option<TaskResumeContext>> {
505 let _ = task_id;
506 Ok(None)
507 }
508
509 async fn set_ttl(&self, task_id: &str, ttl_ms: u64) -> Result<bool>;
514
515 async fn complete_task(&self, task_id: &str, result: CallToolResult) -> Result<bool>;
524
525 async fn fail_task(&self, task_id: &str, error: JsonRpcError) -> Result<bool>;
530
531 async fn cancel_task(&self, task_id: &str, reason: Option<&str>) -> Result<Option<TaskObject>>;
537}
538
539#[derive(Debug, Clone)]
547pub struct MemoryTaskStore {
548 tasks: Arc<RwLock<HashMap<String, Task>>>,
549}
550
551impl Default for MemoryTaskStore {
552 fn default() -> Self {
553 Self::new()
554 }
555}
556
557impl MemoryTaskStore {
558 pub fn new() -> Self {
560 Self {
561 tasks: Arc::new(RwLock::new(HashMap::new())),
562 }
563 }
564
565 pub fn cleanup_expired(&self) -> usize {
573 if let Ok(mut tasks) = self.tasks.write() {
574 let before = tasks.len();
575 tasks.retain(|_, t| !t.is_expired());
576 before - tasks.len()
577 } else {
578 0
579 }
580 }
581
582 #[cfg(test)]
584 pub fn len(&self) -> usize {
585 if let Ok(tasks) = self.tasks.read() {
586 tasks.len()
587 } else {
588 0
589 }
590 }
591
592 #[cfg(test)]
594 pub fn is_empty(&self) -> bool {
595 self.len() == 0
596 }
597}
598
599#[async_trait]
600impl TaskStore for MemoryTaskStore {
601 async fn create_task(
602 &self,
603 tool_name: &str,
604 arguments: serde_json::Value,
605 ttl: Option<u64>,
606 owner: TaskOwner,
607 ) -> Result<(String, CancellationToken)> {
608 let id = generate_task_id();
609 let task = Task::new(id.clone(), tool_name.to_string(), arguments, ttl, owner);
610 let token = task.cancellation_token.clone();
611
612 if let Ok(mut tasks) = self.tasks.write() {
613 tasks.insert(id.clone(), task);
614 }
615
616 Ok((id, token))
617 }
618
619 async fn get_task(&self, task_id: &str) -> Result<Option<TaskObject>> {
620 Ok(if let Ok(tasks) = self.tasks.read() {
621 tasks
622 .get(task_id)
623 .filter(|t| !t.is_expired())
624 .map(|t| t.to_task_object())
625 } else {
626 None
627 })
628 }
629
630 async fn set_task_meta(&self, task_id: &str, meta: serde_json::Value) -> Result<bool> {
631 let Ok(mut tasks) = self.tasks.write() else {
632 return Ok(false);
633 };
634 let Some(task) = tasks.get_mut(task_id).filter(|task| !task.is_expired()) else {
635 return Ok(false);
636 };
637 task.meta = Some(meta);
638 Ok(true)
639 }
640
641 async fn discard_task(&self, task_id: &str) -> Result<bool> {
642 Ok(self
643 .tasks
644 .write()
645 .ok()
646 .and_then(|mut tasks| tasks.remove(task_id))
647 .is_some())
648 }
649
650 async fn task_owner(&self, task_id: &str) -> Result<Option<TaskOwner>> {
651 Ok(if let Ok(tasks) = self.tasks.read() {
652 tasks
653 .get(task_id)
654 .filter(|t| !t.is_expired())
655 .map(|t| t.owner.clone())
656 } else {
657 None
658 })
659 }
660
661 async fn get_task_result(&self, task_id: &str) -> Result<Option<TaskSnapshot>> {
662 Ok(if let Ok(tasks) = self.tasks.read() {
663 tasks
664 .get(task_id)
665 .filter(|t| !t.is_expired())
666 .map(|t| (t.to_task_object(), t.result.clone(), t.error.clone()))
667 } else {
668 None
669 })
670 }
671
672 async fn wait_for_completion(&self, task_id: &str) -> Result<Option<TaskSnapshot>> {
673 let notify = {
675 let Ok(tasks) = self.tasks.read() else {
676 return Ok(None);
677 };
678 let Some(task) = tasks.get(task_id).filter(|t| !t.is_expired()) else {
679 return Ok(None);
680 };
681 if task.status.is_terminal() {
682 return Ok(Some((
683 task.to_task_object(),
684 task.result.clone(),
685 task.error.clone(),
686 )));
687 }
688 task.completion_notify.clone()
689 };
690
691 notify.notified().await;
693
694 self.get_task_result(task_id).await
696 }
697
698 async fn list_tasks(&self, status_filter: Option<TaskStatus>) -> Result<Vec<TaskObject>> {
699 Ok(if let Ok(tasks) = self.tasks.read() {
700 tasks
701 .values()
702 .filter(|t| !t.is_expired())
703 .filter(|t| status_filter.is_none() || status_filter == Some(t.status))
704 .map(|t| t.to_task_object())
705 .collect()
706 } else {
707 vec![]
708 })
709 }
710
711 async fn require_input(
712 &self,
713 task_id: &str,
714 requests: InputRequests,
715 message: Option<&str>,
716 ) -> Result<bool> {
717 let Ok(mut tasks) = self.tasks.write() else {
718 return Ok(false);
719 };
720 let Some(task) = tasks.get_mut(task_id).filter(|t| !t.is_expired()) else {
721 return Ok(false);
722 };
723 if task.status.is_terminal() {
724 return Ok(false);
725 }
726
727 for key in std::mem::take(&mut task.input_requests).into_keys() {
729 if !requests.contains_key(&key) {
730 task.superseded_input_keys.insert(key);
731 }
732 }
733 for key in requests.keys() {
735 task.answered_input_keys.remove(key);
736 task.superseded_input_keys.remove(key);
737 }
738
739 task.input_requests = requests;
740 task.status = TaskStatus::InputRequired;
741 task.status_message = Some(
742 message
743 .map(str::to_string)
744 .unwrap_or_else(|| "Awaiting client input".to_string()),
745 );
746 task.last_updated_at_str = chrono_now_iso8601();
747 Ok(true)
748 }
749
750 async fn outstanding_input_requests(&self, task_id: &str) -> Result<Option<InputRequests>> {
751 Ok(if let Ok(tasks) = self.tasks.read() {
752 tasks
753 .get(task_id)
754 .filter(|t| !t.is_expired())
755 .map(|t| t.input_requests.clone())
756 } else {
757 None
758 })
759 }
760
761 async fn apply_input_responses(
762 &self,
763 task_id: &str,
764 responses: InputResponses,
765 ) -> Result<Option<AppliedInputResponses>> {
766 let Ok(mut tasks) = self.tasks.write() else {
767 return Ok(None);
768 };
769 let Some(task) = tasks.get_mut(task_id).filter(|t| !t.is_expired()) else {
770 return Ok(None);
771 };
772 if task.status.is_terminal() {
773 return Ok(None);
774 }
775
776 let mut applied = AppliedInputResponses::default();
777 for (key, response) in responses {
778 if task.input_requests.remove(&key).is_some() {
779 task.answered_input_keys.insert(key.clone());
780 task.input_responses.insert(key.clone(), response);
781 applied.accepted.insert(key);
782 } else {
783 applied.ignored.insert(key);
787 }
788 }
789 applied.still_outstanding = task.input_requests.keys().cloned().collect();
790
791 if !applied.accepted.is_empty() {
792 task.last_updated_at_str = chrono_now_iso8601();
793 }
794 if applied.is_complete() && task.status == TaskStatus::InputRequired {
795 task.status = TaskStatus::Working;
796 task.status_message = Some("Task resumed".to_string());
797 }
798 Ok(Some(applied))
799 }
800
801 async fn resume_context(&self, task_id: &str) -> Result<Option<TaskResumeContext>> {
802 let Ok(tasks) = self.tasks.read() else {
803 return Ok(None);
804 };
805 Ok(tasks
806 .get(task_id)
807 .filter(|task| !task.is_expired())
808 .map(|task| TaskResumeContext {
809 tool_name: task.tool_name.clone(),
810 arguments: task.arguments.clone(),
811 input_responses: task.input_responses.clone(),
812 }))
813 }
814
815 async fn set_ttl(&self, task_id: &str, ttl_ms: u64) -> Result<bool> {
816 let Ok(mut tasks) = self.tasks.write() else {
817 return Ok(false);
818 };
819 let Some(task) = tasks.get_mut(task_id).filter(|t| !t.is_expired()) else {
820 return Ok(false);
821 };
822 task.ttl = ttl_ms;
823 task.last_updated_at_str = chrono_now_iso8601();
824 Ok(true)
825 }
826
827 async fn complete_task(&self, task_id: &str, result: CallToolResult) -> Result<bool> {
828 let Ok(mut tasks) = self.tasks.write() else {
829 return Ok(false);
830 };
831 let Some(task) = tasks.get_mut(task_id).filter(|t| !t.is_expired()) else {
832 return Ok(false);
833 };
834 if task.status.is_terminal() {
835 return Ok(false);
836 }
837 task.status = TaskStatus::Completed;
838 task.status_message = Some("Task completed".to_string());
839 task.result = Some(result);
840 task.input_requests.clear();
841 task.completed_at = Some(Instant::now());
842 task.last_updated_at_str = chrono_now_iso8601();
843 task.completion_notify.notify_waiters();
844 Ok(true)
845 }
846
847 async fn fail_task(&self, task_id: &str, error: JsonRpcError) -> Result<bool> {
848 let Ok(mut tasks) = self.tasks.write() else {
849 return Ok(false);
850 };
851 let Some(task) = tasks.get_mut(task_id).filter(|t| !t.is_expired()) else {
852 return Ok(false);
853 };
854 if task.status.is_terminal() {
855 return Ok(false);
856 }
857 task.status = TaskStatus::Failed;
858 task.status_message = Some(format!("Task failed: {}", error.message));
859 task.error = Some(error);
860 task.input_requests.clear();
861 task.completed_at = Some(Instant::now());
862 task.last_updated_at_str = chrono_now_iso8601();
863 task.completion_notify.notify_waiters();
864 Ok(true)
865 }
866
867 async fn cancel_task(&self, task_id: &str, reason: Option<&str>) -> Result<Option<TaskObject>> {
868 let Ok(mut tasks) = self.tasks.write() else {
869 return Ok(None);
870 };
871 let Some(task) = tasks.get_mut(task_id).filter(|t| !t.is_expired()) else {
872 return Ok(None);
873 };
874
875 task.cancellation_token.cancel();
877
878 if !task.status.is_terminal() {
880 task.input_requests.clear();
881 task.status = TaskStatus::Cancelled;
882 task.status_message = Some(
883 reason
884 .map(|r| format!("Cancelled: {}", r))
885 .unwrap_or_else(|| "Task cancelled".to_string()),
886 );
887 task.completed_at = Some(Instant::now());
888 task.last_updated_at_str = chrono_now_iso8601();
889 task.completion_notify.notify_waiters();
890 }
891 Ok(Some(task.to_task_object()))
892 }
893}
894
895pub fn tasks_extension() -> crate::ExtensionDeclaration {
900 crate::ExtensionDeclaration::empty(crate::protocol::TASKS_EXTENSION_ID)
901 .expect("the built-in Tasks extension declaration is valid")
902}
903
904impl crate::McpRouter {
905 pub fn with_tasks(self) -> Self {
913 self.with_protocol_extension(tasks_extension())
914 }
915}
916
917impl crate::McpClientBuilder {
918 pub fn with_tasks(self) -> Self {
920 self.with_protocol_extension(tasks_extension())
921 }
922}
923
924impl crate::RequestContext {
925 pub fn supports_tasks(&self) -> bool {
931 self.negotiated_extensions()
932 .is_some_and(|extensions| extensions.contains(crate::protocol::TASKS_EXTENSION_ID))
933 }
934}
935
936fn chrono_now_iso8601() -> String {
938 use std::time::SystemTime;
939
940 let now = SystemTime::now();
941 let duration = now
942 .duration_since(SystemTime::UNIX_EPOCH)
943 .unwrap_or_default();
944
945 let secs = duration.as_secs();
946 let millis = duration.subsec_millis();
947
948 let days = secs / 86400;
951 let remaining = secs % 86400;
952 let hours = remaining / 3600;
953 let remaining = remaining % 3600;
954 let minutes = remaining / 60;
955 let seconds = remaining % 60;
956
957 let mut year = 1970i32;
960 let mut remaining_days = days as i32;
961
962 loop {
963 let days_in_year = if is_leap_year(year) { 366 } else { 365 };
964 if remaining_days < days_in_year {
965 break;
966 }
967 remaining_days -= days_in_year;
968 year += 1;
969 }
970
971 let days_in_months: [i32; 12] = if is_leap_year(year) {
972 [31, 29, 31, 30, 31, 30, 31, 31, 30, 31, 30, 31]
973 } else {
974 [31, 28, 31, 30, 31, 30, 31, 31, 30, 31, 30, 31]
975 };
976
977 let mut month = 1;
978 for days_in_month in days_in_months.iter() {
979 if remaining_days < *days_in_month {
980 break;
981 }
982 remaining_days -= days_in_month;
983 month += 1;
984 }
985
986 let day = remaining_days + 1;
987
988 format!(
989 "{:04}-{:02}-{:02}T{:02}:{:02}:{:02}.{:03}Z",
990 year, month, day, hours, minutes, seconds, millis
991 )
992}
993
994fn is_leap_year(year: i32) -> bool {
995 (year % 4 == 0 && year % 100 != 0) || (year % 400 == 0)
996}
997
998#[cfg(test)]
999mod tests {
1000 use super::*;
1001 use crate::protocol::{
1002 ElicitAction, ElicitResult, InputRequest, InputResponse, ListRootsParams,
1003 };
1004
1005 #[tokio::test]
1006 async fn test_create_task() {
1007 let store = MemoryTaskStore::new();
1008 let (id, token) = store
1009 .create_task("test-tool", serde_json::json!({"a": 1}), None, None)
1010 .await
1011 .unwrap();
1012
1013 assert!(!id.is_empty());
1014 assert!(!token.is_cancelled());
1015
1016 let info = store
1017 .get_task(&id)
1018 .await
1019 .unwrap()
1020 .expect("task should exist");
1021 assert_eq!(info.task_id, id);
1022 assert_eq!(info.status, TaskStatus::Working);
1023 }
1024
1025 #[tokio::test]
1026 async fn test_task_lifecycle() {
1027 let store = MemoryTaskStore::new();
1028 let (id, _) = store
1029 .create_task("test-tool", serde_json::json!({}), None, None)
1030 .await
1031 .unwrap();
1032
1033 assert!(
1035 store
1036 .complete_task(&id, CallToolResult::text("Done"))
1037 .await
1038 .unwrap()
1039 );
1040
1041 let info = store.get_task(&id).await.unwrap().unwrap();
1042 assert_eq!(info.status, TaskStatus::Completed);
1043 }
1044
1045 #[tokio::test]
1046 async fn test_task_cancellation() {
1047 let store = MemoryTaskStore::new();
1048 let (id, token) = store
1049 .create_task("test-tool", serde_json::json!({}), None, None)
1050 .await
1051 .unwrap();
1052
1053 assert!(!token.is_cancelled());
1054
1055 let task_obj = store
1056 .cancel_task(&id, Some("User requested"))
1057 .await
1058 .unwrap();
1059 assert!(task_obj.is_some());
1060 assert_eq!(task_obj.unwrap().status, TaskStatus::Cancelled);
1061 assert!(token.is_cancelled());
1062
1063 let info = store.get_task(&id).await.unwrap().unwrap();
1064 assert_eq!(info.status, TaskStatus::Cancelled);
1065 }
1066
1067 #[tokio::test]
1068 async fn test_task_failure() {
1069 let store = MemoryTaskStore::new();
1070 let (id, _) = store
1071 .create_task("test-tool", serde_json::json!({}), None, None)
1072 .await
1073 .unwrap();
1074
1075 assert!(
1076 store
1077 .fail_task(&id, JsonRpcError::internal_error("Something went wrong"))
1078 .await
1079 .unwrap()
1080 );
1081
1082 let info = store.get_task(&id).await.unwrap().unwrap();
1083 assert_eq!(info.status, TaskStatus::Failed);
1084 assert!(info.status_message.as_ref().unwrap().contains("failed"));
1085 }
1086
1087 #[tokio::test]
1088 async fn test_list_tasks() {
1089 let store = MemoryTaskStore::new();
1090 store
1091 .create_task("tool1", serde_json::json!({}), None, None)
1092 .await
1093 .unwrap();
1094 store
1095 .create_task("tool2", serde_json::json!({}), None, None)
1096 .await
1097 .unwrap();
1098 let (id3, _) = store
1099 .create_task("tool3", serde_json::json!({}), None, None)
1100 .await
1101 .unwrap();
1102
1103 store
1105 .complete_task(&id3, CallToolResult::text("Done"))
1106 .await
1107 .unwrap();
1108
1109 let all = store.list_tasks(None).await.unwrap();
1111 assert_eq!(all.len(), 3);
1112
1113 let working = store.list_tasks(Some(TaskStatus::Working)).await.unwrap();
1115 assert_eq!(working.len(), 2);
1116
1117 let completed = store.list_tasks(Some(TaskStatus::Completed)).await.unwrap();
1119 assert_eq!(completed.len(), 1);
1120 }
1121
1122 #[tokio::test]
1123 async fn test_terminal_state_immutable() {
1124 let store = MemoryTaskStore::new();
1125 let (id, _) = store
1126 .create_task("test-tool", serde_json::json!({}), None, None)
1127 .await
1128 .unwrap();
1129
1130 store
1132 .complete_task(&id, CallToolResult::text("Done"))
1133 .await
1134 .unwrap();
1135
1136 assert!(
1138 !store
1139 .fail_task(&id, JsonRpcError::internal_error("Error"))
1140 .await
1141 .unwrap()
1142 );
1143
1144 let info = store.get_task(&id).await.unwrap().unwrap();
1146 assert_eq!(info.status, TaskStatus::Completed);
1147 }
1148
1149 #[tokio::test]
1150 async fn test_task_ids_unique() {
1151 let store = MemoryTaskStore::new();
1152 let (id1, _) = store
1153 .create_task("tool", serde_json::json!({}), None, None)
1154 .await
1155 .unwrap();
1156 let (id2, _) = store
1157 .create_task("tool", serde_json::json!({}), None, None)
1158 .await
1159 .unwrap();
1160 let (id3, _) = store
1161 .create_task("tool", serde_json::json!({}), None, None)
1162 .await
1163 .unwrap();
1164
1165 assert_ne!(id1, id2);
1166 assert_ne!(id2, id3);
1167 assert_ne!(id1, id3);
1168 }
1169
1170 #[tokio::test]
1171 async fn test_get_task_result() {
1172 let store = MemoryTaskStore::new();
1173 let (id, _) = store
1174 .create_task("test-tool", serde_json::json!({}), None, None)
1175 .await
1176 .unwrap();
1177
1178 let result = CallToolResult::text("The result");
1180 store.complete_task(&id, result).await.unwrap();
1181
1182 let (task_obj, result, error) = store.get_task_result(&id).await.unwrap().unwrap();
1183 assert_eq!(task_obj.status, TaskStatus::Completed);
1184 assert!(result.is_some());
1185 assert!(error.is_none());
1186 }
1187
1188 #[tokio::test]
1189 async fn test_wait_for_completion_returns_terminal_snapshot() {
1190 let store = MemoryTaskStore::new();
1191 let (id, _) = store
1192 .create_task("test-tool", serde_json::json!({}), None, None)
1193 .await
1194 .unwrap();
1195
1196 let waiter_store = store.clone();
1198 let waiter_id = id.clone();
1199 let waiter =
1200 tokio::spawn(async move { waiter_store.wait_for_completion(&waiter_id).await });
1201
1202 tokio::time::sleep(Duration::from_millis(10)).await;
1203 store
1204 .complete_task(&id, CallToolResult::text("Done"))
1205 .await
1206 .unwrap();
1207
1208 let (task_obj, result, error) = waiter.await.unwrap().unwrap().unwrap();
1209 assert_eq!(task_obj.status, TaskStatus::Completed);
1210 assert!(result.is_some());
1211 assert!(error.is_none());
1212 }
1213
1214 #[tokio::test]
1215 async fn dyn_task_store_object_safe() {
1216 let store: Arc<dyn TaskStore> = Arc::new(MemoryTaskStore::new());
1218 let (id, _) = store
1219 .create_task("tool", serde_json::json!({}), None, None)
1220 .await
1221 .unwrap();
1222 assert!(store.get_task(&id).await.unwrap().is_some());
1223 }
1224
1225 #[test]
1226 fn test_iso8601_timestamp() {
1227 let ts = chrono_now_iso8601();
1228 assert!(ts.ends_with('Z'));
1230 assert!(ts.contains('T'));
1231 assert_eq!(ts.len(), 24); }
1233
1234 #[test]
1235 fn test_task_status_display() {
1236 assert_eq!(TaskStatus::Working.to_string(), "working");
1237 assert_eq!(TaskStatus::InputRequired.to_string(), "input_required");
1238 assert_eq!(TaskStatus::Completed.to_string(), "completed");
1239 assert_eq!(TaskStatus::Failed.to_string(), "failed");
1240 assert_eq!(TaskStatus::Cancelled.to_string(), "cancelled");
1241 }
1242
1243 #[test]
1244 fn test_task_status_is_terminal() {
1245 assert!(!TaskStatus::Working.is_terminal());
1246 assert!(!TaskStatus::InputRequired.is_terminal());
1247 assert!(TaskStatus::Completed.is_terminal());
1248 assert!(TaskStatus::Failed.is_terminal());
1249 assert!(TaskStatus::Cancelled.is_terminal());
1250 }
1251
1252 fn requests(keys: &[&str]) -> InputRequests {
1253 keys.iter()
1254 .map(|k| {
1255 (
1256 k.to_string(),
1257 InputRequest::ListRoots(ListRootsParams { meta: None }),
1258 )
1259 })
1260 .collect()
1261 }
1262
1263 fn accept(key: &str) -> (String, InputResponse) {
1264 (
1265 key.to_string(),
1266 InputResponse::Elicit(ElicitResult {
1267 action: ElicitAction::Accept,
1268 content: None,
1269 meta: None,
1270 }),
1271 )
1272 }
1273
1274 async fn working_task(store: &MemoryTaskStore, ttl: Option<u64>) -> String {
1275 store
1276 .create_task("tool", serde_json::json!({}), ttl, None)
1277 .await
1278 .unwrap()
1279 .0
1280 }
1281
1282 #[tokio::test]
1283 async fn task_ids_are_unguessable_not_sequential() {
1284 let store = MemoryTaskStore::new();
1285 let mut ids = BTreeSet::new();
1286 for _ in 0..64 {
1287 ids.insert(working_task(&store, None).await);
1288 }
1289 assert_eq!(ids.len(), 64, "task IDs collided");
1290
1291 for id in &ids {
1292 assert_eq!(id.len(), 32, "expected 128 bits of hex: {id}");
1293 assert!(id.chars().all(|c| c.is_ascii_hexdigit()), "{id}");
1294 assert!(!id.starts_with("task-"), "sequential-looking ID: {id}");
1295 }
1296
1297 let leading: BTreeSet<&str> = ids.iter().map(|id| &id[..2]).collect();
1300 assert!(
1301 leading.len() > 32,
1302 "only {} distinct leading bytes across 64 IDs",
1303 leading.len()
1304 );
1305 }
1306
1307 #[tokio::test]
1308 async fn ttl_runs_from_creation_and_expired_tasks_read_as_absent() {
1309 let store = MemoryTaskStore::new();
1310 let id = working_task(&store, Some(0)).await;
1311
1312 tokio::time::sleep(Duration::from_millis(5)).await;
1315
1316 assert!(store.get_task(&id).await.unwrap().is_none());
1317 assert!(store.get_task_result(&id).await.unwrap().is_none());
1318 assert!(store.list_tasks(None).await.unwrap().is_empty());
1319 assert!(
1320 store
1321 .outstanding_input_requests(&id)
1322 .await
1323 .unwrap()
1324 .is_none()
1325 );
1326 assert!(store.cancel_task(&id, None).await.unwrap().is_none());
1327 assert!(!store.set_ttl(&id, 60_000).await.unwrap());
1328 assert!(
1329 !store
1330 .complete_task(&id, CallToolResult::text("late"))
1331 .await
1332 .unwrap()
1333 );
1334 }
1335
1336 #[tokio::test]
1337 async fn ttl_is_mutable_over_the_task_lifetime() {
1338 let store = MemoryTaskStore::new();
1339 let id = working_task(&store, Some(60_000)).await;
1340
1341 assert!(store.set_ttl(&id, 120_000).await.unwrap());
1342 let task = store.get_task(&id).await.unwrap().unwrap();
1343 assert_eq!(task.ttl, Some(120_000));
1344
1345 assert!(store.set_ttl(&id, 0).await.unwrap());
1347 tokio::time::sleep(Duration::from_millis(5)).await;
1348 assert!(store.get_task(&id).await.unwrap().is_none());
1349 }
1350
1351 #[tokio::test]
1352 async fn require_input_records_requests_and_exposes_them() {
1353 let store = MemoryTaskStore::new();
1354 let id = working_task(&store, None).await;
1355
1356 assert!(
1357 store
1358 .require_input(&id, requests(&["approval", "region"]), Some("need input"))
1359 .await
1360 .unwrap()
1361 );
1362
1363 let task = store.get_task(&id).await.unwrap().unwrap();
1364 assert_eq!(task.status, TaskStatus::InputRequired);
1365 assert_eq!(task.status_message.as_deref(), Some("need input"));
1366
1367 let outstanding = store
1368 .outstanding_input_requests(&id)
1369 .await
1370 .unwrap()
1371 .unwrap();
1372 assert_eq!(
1373 outstanding.keys().collect::<Vec<_>>(),
1374 vec!["approval", "region"],
1375 "every outstanding request must be exposed, not just the newest"
1376 );
1377 }
1378
1379 #[tokio::test]
1380 async fn partial_input_responses_leave_the_rest_outstanding() {
1381 let store = MemoryTaskStore::new();
1382 let id = working_task(&store, None).await;
1383 store
1384 .require_input(&id, requests(&["approval", "region"]), None)
1385 .await
1386 .unwrap();
1387
1388 let applied = store
1389 .apply_input_responses(&id, [accept("approval")].into_iter().collect())
1390 .await
1391 .unwrap()
1392 .unwrap();
1393
1394 assert_eq!(applied.accepted, ["approval".to_string()].into());
1395 assert!(applied.ignored.is_empty());
1396 assert_eq!(applied.still_outstanding, ["region".to_string()].into());
1397 assert!(!applied.is_complete());
1398
1399 let task = store.get_task(&id).await.unwrap().unwrap();
1401 assert_eq!(task.status, TaskStatus::InputRequired);
1402 assert_eq!(
1403 store
1404 .outstanding_input_requests(&id)
1405 .await
1406 .unwrap()
1407 .unwrap()
1408 .keys()
1409 .collect::<Vec<_>>(),
1410 vec!["region"]
1411 );
1412
1413 let applied = store
1415 .apply_input_responses(&id, [accept("region")].into_iter().collect())
1416 .await
1417 .unwrap()
1418 .unwrap();
1419 assert!(applied.is_complete());
1420 assert_eq!(
1421 store.get_task(&id).await.unwrap().unwrap().status,
1422 TaskStatus::Working
1423 );
1424 }
1425
1426 #[tokio::test]
1427 async fn unknown_answered_and_superseded_response_keys_are_ignored() {
1428 let store = MemoryTaskStore::new();
1429 let id = working_task(&store, None).await;
1430 store
1431 .require_input(&id, requests(&["approval", "stale"]), None)
1432 .await
1433 .unwrap();
1434
1435 store
1437 .apply_input_responses(&id, [accept("approval")].into_iter().collect())
1438 .await
1439 .unwrap()
1440 .unwrap();
1441 store
1442 .require_input(&id, requests(&["region"]), None)
1443 .await
1444 .unwrap();
1445
1446 let applied = store
1447 .apply_input_responses(
1448 &id,
1449 [accept("never-issued"), accept("approval"), accept("stale")]
1450 .into_iter()
1451 .collect(),
1452 )
1453 .await
1454 .unwrap()
1455 .unwrap();
1456
1457 assert!(
1458 applied.accepted.is_empty(),
1459 "none of these keys are outstanding"
1460 );
1461 assert_eq!(
1462 applied.ignored,
1463 [
1464 "never-issued".to_string(),
1465 "approval".to_string(),
1466 "stale".to_string()
1467 ]
1468 .into(),
1469 "unknown, already-answered, and superseded keys are all ignored"
1470 );
1471 assert_eq!(applied.still_outstanding, ["region".to_string()].into());
1472 assert_eq!(
1473 store.get_task(&id).await.unwrap().unwrap().status,
1474 TaskStatus::InputRequired,
1475 "ignoring a stale update must not resume or fail the task"
1476 );
1477 }
1478
1479 #[tokio::test]
1480 async fn reissued_key_becomes_a_fresh_question() {
1481 let store = MemoryTaskStore::new();
1482 let id = working_task(&store, None).await;
1483 store
1484 .require_input(&id, requests(&["approval"]), None)
1485 .await
1486 .unwrap();
1487 store
1488 .apply_input_responses(&id, [accept("approval")].into_iter().collect())
1489 .await
1490 .unwrap()
1491 .unwrap();
1492
1493 store
1496 .require_input(&id, requests(&["approval"]), None)
1497 .await
1498 .unwrap();
1499 let applied = store
1500 .apply_input_responses(&id, [accept("approval")].into_iter().collect())
1501 .await
1502 .unwrap()
1503 .unwrap();
1504 assert_eq!(applied.accepted, ["approval".to_string()].into());
1505 assert!(applied.is_complete());
1506 }
1507
1508 #[tokio::test]
1509 async fn failed_tasks_preserve_the_structured_error() {
1510 let store = MemoryTaskStore::new();
1511 let id = working_task(&store, None).await;
1512
1513 let mut error = JsonRpcError::invalid_params("bad region");
1514 error.data = Some(serde_json::json!({"field": "region"}));
1515 assert!(store.fail_task(&id, error).await.unwrap());
1516
1517 let (_, result, error) = store.get_task_result(&id).await.unwrap().unwrap();
1518 assert!(result.is_none());
1519 let error = error.expect("structured error must survive the store");
1520 assert_eq!(
1521 error.code, -32602,
1522 "the original code must not be flattened"
1523 );
1524 assert_eq!(error.message, "bad region");
1525 assert_eq!(error.data.unwrap()["field"], "region");
1526 }
1527
1528 #[tokio::test]
1529 async fn tool_error_results_complete_the_task() {
1530 let store = MemoryTaskStore::new();
1531 let id = working_task(&store, None).await;
1532
1533 let mut result = CallToolResult::text("domain failure");
1534 result.is_error = true;
1535 assert!(store.complete_task(&id, result).await.unwrap());
1536
1537 let (task, result, error) = store.get_task_result(&id).await.unwrap().unwrap();
1538 assert_eq!(
1539 task.status,
1540 TaskStatus::Completed,
1541 "isError is a domain error, not an execution failure"
1542 );
1543 assert!(result.unwrap().is_error);
1544 assert!(error.is_none(), "no JSON-RPC error accompanies isError");
1545 }
1546
1547 #[tokio::test]
1548 async fn tasks_record_their_creating_principal() {
1549 let store = MemoryTaskStore::new();
1550 let (owned, _) = store
1551 .create_task("tool", serde_json::json!({}), None, Some("alice".into()))
1552 .await
1553 .unwrap();
1554 let (unowned, _) = store
1555 .create_task("tool", serde_json::json!({}), None, None)
1556 .await
1557 .unwrap();
1558
1559 assert_eq!(
1560 store.task_owner(&owned).await.unwrap(),
1561 Some(Some("alice".to_string()))
1562 );
1563 assert_eq!(store.task_owner(&unowned).await.unwrap(), Some(None));
1564 assert_eq!(
1565 store.task_owner("does-not-exist").await.unwrap(),
1566 None,
1567 "an unknown task has no owner record at all"
1568 );
1569
1570 let wire = serde_json::to_value(store.get_task(&owned).await.unwrap().unwrap()).unwrap();
1572 assert!(
1573 wire.get("owner").is_none(),
1574 "owner leaked to the wire: {wire}"
1575 );
1576 assert!(!wire.to_string().contains("alice"));
1577 }
1578
1579 #[test]
1580 fn owner_matching_is_equality_not_leniency() {
1581 assert!(owner_matches(&None, None), "no auth configured");
1582 assert!(owner_matches(&Some("alice".into()), Some("alice")));
1583
1584 assert!(
1585 !owner_matches(&Some("alice".into()), Some("bob")),
1586 "a different principal must not inherit the task"
1587 );
1588 assert!(
1589 !owner_matches(&Some("alice".into()), None),
1590 "dropping the token must not grant access"
1591 );
1592 assert!(
1593 !owner_matches(&None, Some("alice")),
1594 "an unowned task belongs to a different security context"
1595 );
1596 }
1597
1598 #[tokio::test]
1599 async fn terminal_states_clear_outstanding_requests() {
1600 for (label, terminate) in [("completed", true), ("cancelled", false)] {
1601 let store = MemoryTaskStore::new();
1602 let id = working_task(&store, None).await;
1603 store
1604 .require_input(&id, requests(&["approval"]), None)
1605 .await
1606 .unwrap();
1607
1608 if terminate {
1609 store
1610 .complete_task(&id, CallToolResult::text("done"))
1611 .await
1612 .unwrap();
1613 } else {
1614 store.cancel_task(&id, None).await.unwrap();
1615 }
1616
1617 assert!(
1618 store
1619 .outstanding_input_requests(&id)
1620 .await
1621 .unwrap()
1622 .unwrap()
1623 .is_empty(),
1624 "{label} task still advertises outstanding input requests"
1625 );
1626 assert!(
1627 store
1628 .apply_input_responses(&id, [accept("approval")].into_iter().collect())
1629 .await
1630 .unwrap()
1631 .is_none(),
1632 "{label} task accepted a late input response"
1633 );
1634 }
1635 }
1636}