use std::collections::{BTreeSet, HashMap};
use std::fmt::Write as _;
use std::io;
use std::sync::{Arc, Condvar, Mutex, OnceLock, RwLock, Weak};
use std::time::{Duration, Instant};
use async_trait::async_trait;
use crate::error::JsonRpcError;
use crate::protocol::{CallToolResult, InputRequests, InputResponses, TaskObject, TaskStatus};
#[derive(Clone, Debug, Default)]
pub struct CancellationToken {
inner: tokio_util::sync::CancellationToken,
}
impl CancellationToken {
#[must_use]
pub fn new() -> Self {
Self::default()
}
#[must_use]
pub fn is_cancelled(&self) -> bool {
self.inner.is_cancelled()
}
pub fn cancel(&self) {
self.inner.cancel();
}
pub async fn cancelled(&self) {
self.inner.cancelled().await;
}
}
const DEFAULT_TASK_TTL: Duration = Duration::from_secs(5 * 60);
const DEFAULT_TASK_CLEANUP_INTERVAL: Duration = Duration::from_secs(60);
const DEFAULT_MAX_RETAINED_TASKS: usize = 1_024;
const DEFAULT_MAX_TASK_PAYLOAD_BYTES: usize = 4 * 1024 * 1024;
const DEFAULT_MAX_RETAINED_TASK_BYTES: usize = 64 * 1024 * 1024;
const DEFAULT_POLL_INTERVAL_MS: u64 = 2_000;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[non_exhaustive]
pub struct TaskRetentionLimits {
pub max_tasks: usize,
pub max_payload_bytes: usize,
pub max_retained_bytes: usize,
}
impl TaskRetentionLimits {
#[must_use]
pub fn new() -> Self {
Self::default()
}
#[must_use]
pub const fn unbounded() -> Self {
Self {
max_tasks: usize::MAX,
max_payload_bytes: usize::MAX,
max_retained_bytes: usize::MAX,
}
}
#[must_use]
pub const fn max_tasks(mut self, max: usize) -> Self {
self.max_tasks = max;
self
}
#[must_use]
pub const fn max_payload_bytes(mut self, max: usize) -> Self {
self.max_payload_bytes = max;
self
}
#[must_use]
pub const fn max_retained_bytes(mut self, max: usize) -> Self {
self.max_retained_bytes = max;
self
}
}
impl Default for TaskRetentionLimits {
fn default() -> Self {
Self {
max_tasks: DEFAULT_MAX_RETAINED_TASKS,
max_payload_bytes: DEFAULT_MAX_TASK_PAYLOAD_BYTES,
max_retained_bytes: DEFAULT_MAX_RETAINED_TASK_BYTES,
}
}
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
#[non_exhaustive]
pub struct TaskStoreUsage {
pub task_count: usize,
pub retained_bytes: usize,
pub reserved_bytes: usize,
}
impl TaskStoreUsage {
#[must_use]
pub fn charged_bytes(self) -> usize {
self.retained_bytes.saturating_add(self.reserved_bytes)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct MemoryTaskStoreConfig {
pub default_ttl: Duration,
pub cleanup_interval: Duration,
}
impl MemoryTaskStoreConfig {
#[must_use]
pub fn new() -> Self {
Self::default()
}
#[must_use]
pub fn default_ttl(mut self, ttl: Duration) -> Self {
self.default_ttl = ttl;
self
}
#[must_use]
pub fn cleanup_interval(mut self, interval: Duration) -> Self {
self.cleanup_interval = interval;
self
}
}
impl Default for MemoryTaskStoreConfig {
fn default() -> Self {
Self {
default_ttl: DEFAULT_TASK_TTL,
cleanup_interval: DEFAULT_TASK_CLEANUP_INTERVAL,
}
}
}
fn duration_millis_saturated(duration: Duration) -> u64 {
u64::try_from(duration.as_millis()).unwrap_or(u64::MAX)
}
#[derive(Debug, Clone)]
pub struct Task {
pub id: String,
pub tool_name: String,
pub arguments: serde_json::Value,
pub status: TaskStatus,
pub created_at: Instant,
pub created_at_str: String,
pub last_updated_at_str: String,
pub ttl: u64,
pub poll_interval: u64,
pub status_message: Option<String>,
pub meta: Option<serde_json::Value>,
pub result: Option<CallToolResult>,
pub error: Option<JsonRpcError>,
pub owner: TaskOwner,
pub input_requests: InputRequests,
pub answered_input_keys: BTreeSet<String>,
pub input_responses: InputResponses,
pub superseded_input_keys: BTreeSet<String>,
pub cancellation_token: CancellationToken,
pub completed_at: Option<Instant>,
pub completion_notify: Arc<tokio::sync::Notify>,
}
impl Task {
fn new(
id: String,
tool_name: String,
arguments: serde_json::Value,
ttl: u64,
owner: TaskOwner,
) -> Self {
let now_str = chrono_now_iso8601();
Self {
id,
tool_name,
arguments,
status: TaskStatus::Working,
created_at: Instant::now(),
created_at_str: now_str.clone(),
last_updated_at_str: now_str,
ttl,
poll_interval: DEFAULT_POLL_INTERVAL_MS,
status_message: Some("Task started".to_string()),
meta: None,
result: None,
error: None,
owner,
input_requests: InputRequests::new(),
answered_input_keys: BTreeSet::new(),
input_responses: InputResponses::new(),
superseded_input_keys: BTreeSet::new(),
cancellation_token: CancellationToken::new(),
completed_at: None,
completion_notify: Arc::new(tokio::sync::Notify::new()),
}
}
pub fn to_task_object(&self) -> TaskObject {
TaskObject {
task_id: self.id.clone(),
status: self.status,
status_message: self.status_message.clone(),
created_at: self.created_at_str.clone(),
last_updated_at: self.last_updated_at_str.clone(),
ttl: Some(self.ttl),
poll_interval: Some(self.poll_interval),
result: None,
error: None,
meta: self.meta.clone(),
}
}
pub fn is_expired(&self) -> bool {
self.is_expired_at(Instant::now())
}
fn expires_at(&self) -> Option<Instant> {
self.created_at.checked_add(Duration::from_millis(self.ttl))
}
fn is_expired_at(&self, now: Instant) -> bool {
self.expires_at().is_some_and(|deadline| now >= deadline)
}
pub fn outstanding_input_requests(&self) -> &InputRequests {
&self.input_requests
}
pub fn is_cancelled(&self) -> bool {
self.cancellation_token.is_cancelled()
}
fn fail_retention_limit(&mut self) {
self.arguments = serde_json::Value::Null;
self.status = TaskStatus::Failed;
self.status_message = Some(RETENTION_FAILURE_STATUS.to_string());
self.meta = None;
self.result = None;
self.error = Some(JsonRpcError::internal_error(RETENTION_FAILURE_MESSAGE));
self.input_requests = InputRequests::new();
self.answered_input_keys = BTreeSet::new();
self.input_responses = InputResponses::new();
self.superseded_input_keys = BTreeSet::new();
self.completed_at = Some(Instant::now());
self.last_updated_at_str = chrono_now_iso8601();
}
}
#[derive(Debug, Clone)]
struct StoredTask {
task: Task,
expiry_signalled: bool,
retained_bytes: usize,
reserved_bytes: usize,
}
impl StoredTask {
fn new(task: Task, retained_bytes: usize, reserved_bytes: usize) -> Self {
Self {
task,
expiry_signalled: false,
retained_bytes,
reserved_bytes,
}
}
fn signal_expiry(&mut self) -> bool {
if self.expiry_signalled {
return false;
}
self.expiry_signalled = true;
self.task.cancellation_token.cancel();
self.task.completion_notify.notify_waiters();
true
}
fn scrub_expired_payload(&mut self) {
self.task.tool_name = String::new();
self.task.arguments = serde_json::Value::Null;
self.task.status_message = None;
self.task.meta = None;
self.task.result = None;
self.task.error = None;
self.task.input_requests = InputRequests::new();
self.task.answered_input_keys = BTreeSet::new();
self.task.input_responses = InputResponses::new();
self.task.superseded_input_keys = BTreeSet::new();
self.reserved_bytes = 0;
}
}
impl std::ops::Deref for StoredTask {
type Target = Task;
fn deref(&self) -> &Self::Target {
&self.task
}
}
impl std::ops::DerefMut for StoredTask {
fn deref_mut(&mut self) -> &mut Self::Target {
&mut self.task
}
}
fn charged_bytes(retained: usize, reserved: usize, limit: usize) -> Result<usize> {
retained
.checked_add(reserved)
.filter(|charged| *charged <= limit)
.ok_or(TaskStoreError::RetentionLimitExceeded {
kind: TaskRetentionLimitKind::AggregateBytes,
limit,
})
}
fn prepare_stored_task(task: Task, limits: TaskRetentionLimits) -> Result<StoredTask> {
let retained_bytes = retained_payload_size(&task, limits.max_retained_bytes)?;
let reserved_bytes = if task.status.is_terminal() {
0
} else {
let mut fallback = task.clone();
fallback.fail_retention_limit();
let fallback_bytes = retained_payload_size(&fallback, limits.max_retained_bytes)?;
fallback_bytes.saturating_sub(retained_bytes)
};
charged_bytes(retained_bytes, reserved_bytes, limits.max_retained_bytes)?;
Ok(StoredTask::new(task, retained_bytes, reserved_bytes))
}
fn replace_stored_task(
data: &mut MemoryTaskStoreData,
task_id: &str,
replacement: StoredTask,
limit: usize,
) -> Result<bool> {
let Some(current) = data.tasks.get(task_id) else {
return Ok(false);
};
let (retained_bytes, reserved_bytes) = replacement_totals(
data,
current.retained_bytes,
current.reserved_bytes,
replacement.retained_bytes,
replacement.reserved_bytes,
limit,
)?;
data.retained_bytes = retained_bytes;
data.reserved_bytes = reserved_bytes;
data.tasks.insert(task_id.to_string(), replacement);
Ok(true)
}
fn replacement_totals(
data: &MemoryTaskStoreData,
old_retained: usize,
old_reserved: usize,
new_retained: usize,
new_reserved: usize,
limit: usize,
) -> Result<(usize, usize)> {
let retained_bytes = data
.retained_bytes
.checked_sub(old_retained)
.and_then(|bytes| bytes.checked_add(new_retained))
.ok_or(TaskStoreError::RetentionLimitExceeded {
kind: TaskRetentionLimitKind::AggregateBytes,
limit,
})?;
let reserved_bytes = data
.reserved_bytes
.checked_sub(old_reserved)
.and_then(|bytes| bytes.checked_add(new_reserved))
.ok_or(TaskStoreError::RetentionLimitExceeded {
kind: TaskRetentionLimitKind::AggregateBytes,
limit,
})?;
charged_bytes(retained_bytes, reserved_bytes, limit)?;
Ok((retained_bytes, reserved_bytes))
}
fn commit_removed_task(
data: &mut MemoryTaskStoreData,
task_id: &str,
replacement: StoredTask,
old_retained: usize,
old_reserved: usize,
limit: usize,
) -> Result<()> {
let (retained_bytes, reserved_bytes) = replacement_totals(
data,
old_retained,
old_reserved,
replacement.retained_bytes,
replacement.reserved_bytes,
limit,
)?;
data.retained_bytes = retained_bytes;
data.reserved_bytes = reserved_bytes;
data.tasks.insert(task_id.to_string(), replacement);
Ok(())
}
fn commit_retention_failure(
data: &mut MemoryTaskStoreData,
task_id: &str,
mut task: StoredTask,
old_retained: usize,
old_reserved: usize,
limits: TaskRetentionLimits,
) {
task.task.fail_retention_limit();
task.retained_bytes = retained_payload_size(&task.task, limits.max_retained_bytes)
.expect("reserved retention failure must fit the aggregate byte limit");
task.reserved_bytes = 0;
let notify = task.completion_notify.clone();
commit_removed_task(
data,
task_id,
task,
old_retained,
old_reserved,
limits.max_retained_bytes,
)
.expect("reserved retention failure must fit global task-store accounting");
notify.notify_waiters();
}
fn cancel_with_bounded_status(task: &mut Task, status: &'static str, scrub: bool) {
if scrub {
task.arguments = serde_json::Value::Null;
task.meta = None;
task.result = None;
task.error = None;
task.answered_input_keys = BTreeSet::new();
task.input_responses = InputResponses::new();
task.superseded_input_keys = BTreeSet::new();
}
task.input_requests = InputRequests::new();
task.status = TaskStatus::Cancelled;
task.status_message = Some(status.to_string());
task.completed_at = Some(Instant::now());
task.last_updated_at_str = chrono_now_iso8601();
}
pub fn generate_task_id() -> String {
let mut bytes = [0u8; 16];
getrandom::fill(&mut bytes).expect("system entropy source unavailable for task ID generation");
let mut id = String::with_capacity(2 * bytes.len());
for byte in bytes {
let _ = write!(id, "{byte:02x}");
}
id
}
pub type TaskOwner = Option<String>;
pub fn owner_matches(owner: &TaskOwner, principal: Option<&str>) -> bool {
owner.as_deref() == principal
}
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct AppliedInputResponses {
pub accepted: BTreeSet<String>,
pub ignored: BTreeSet<String>,
pub still_outstanding: BTreeSet<String>,
}
#[derive(Debug, Clone)]
#[non_exhaustive]
pub struct TaskResumeContext {
pub tool_name: String,
pub arguments: serde_json::Value,
pub input_responses: InputResponses,
pub cancellation_token: Option<CancellationToken>,
}
impl TaskResumeContext {
#[must_use]
pub fn new(
tool_name: impl Into<String>,
arguments: serde_json::Value,
input_responses: InputResponses,
) -> Self {
Self {
tool_name: tool_name.into(),
arguments,
input_responses,
cancellation_token: None,
}
}
#[must_use]
pub fn with_cancellation_token(mut self, token: CancellationToken) -> Self {
self.cancellation_token = Some(token);
self
}
}
impl AppliedInputResponses {
pub fn is_complete(&self) -> bool {
self.still_outstanding.is_empty()
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
#[non_exhaustive]
pub enum TaskPresence {
Present {
owner: TaskOwner,
},
Expired {
owner: TaskOwner,
},
Missing,
}
impl TaskPresence {
pub fn owner(&self) -> Option<&TaskOwner> {
match self {
Self::Present { owner } | Self::Expired { owner } => Some(owner),
Self::Missing => None,
}
}
pub fn is_known(&self) -> bool {
!matches!(self, Self::Missing)
}
}
#[derive(Debug, thiserror::Error)]
#[non_exhaustive]
pub enum TaskStoreError {
#[error("encode error: {0}")]
Encode(String),
#[error("decode error: {0}")]
Decode(String),
#[error("backend error: {0}")]
Backend(String),
#[error("invalid task transition: {0}")]
InvalidTransition(String),
#[error("task retention {kind} limit exceeded (maximum {limit})")]
RetentionLimitExceeded {
kind: TaskRetentionLimitKind,
limit: usize,
},
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[non_exhaustive]
pub enum TaskRetentionLimitKind {
TaskCount,
PayloadBytes,
AggregateBytes,
}
impl std::fmt::Display for TaskRetentionLimitKind {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(match self {
Self::TaskCount => "task-count",
Self::PayloadBytes => "payload-bytes",
Self::AggregateBytes => "aggregate-bytes",
})
}
}
fn same_input_request(
a: &crate::protocol::InputRequest,
b: &crate::protocol::InputRequest,
) -> bool {
match (serde_json::to_value(a), serde_json::to_value(b)) {
(Ok(a), Ok(b)) => a == b,
_ => false,
}
}
pub type Result<T> = std::result::Result<T, TaskStoreError>;
const RETENTION_FAILURE_MESSAGE: &str = "Task payload exceeded configured retention limits";
const RETENTION_FAILURE_STATUS: &str = "Task failed: retention limit exceeded";
struct CountingWriter {
written: usize,
limit: usize,
exceeded: bool,
}
impl CountingWriter {
fn new(limit: usize) -> Self {
Self {
written: 0,
limit,
exceeded: false,
}
}
}
impl io::Write for CountingWriter {
fn write(&mut self, bytes: &[u8]) -> io::Result<usize> {
let Some(next) = self.written.checked_add(bytes.len()) else {
self.exceeded = true;
return Err(io::Error::other("encoded task payload exceeds byte limit"));
};
if next > self.limit {
self.exceeded = true;
return Err(io::Error::other("encoded task payload exceeds byte limit"));
}
self.written = next;
Ok(bytes.len())
}
fn flush(&mut self) -> io::Result<()> {
Ok(())
}
}
fn encoded_size<T: serde::Serialize + ?Sized>(
value: &T,
limit: usize,
kind: TaskRetentionLimitKind,
) -> Result<usize> {
let mut writer = CountingWriter::new(limit);
match serde_json::to_writer(&mut writer, value) {
Ok(()) => Ok(writer.written),
Err(_) if writer.exceeded => Err(TaskStoreError::RetentionLimitExceeded { kind, limit }),
Err(error) => Err(TaskStoreError::Encode(error.to_string())),
}
}
#[derive(serde::Serialize)]
struct RetainedTaskPayload<'a> {
tool_name: &'a str,
arguments: &'a serde_json::Value,
status_message: &'a Option<String>,
meta: &'a Option<serde_json::Value>,
result: &'a Option<CallToolResult>,
error: &'a Option<JsonRpcError>,
owner: &'a TaskOwner,
input_requests: &'a InputRequests,
answered_input_keys: &'a BTreeSet<String>,
input_responses: &'a InputResponses,
superseded_input_keys: &'a BTreeSet<String>,
}
fn retained_payload_size(task: &Task, limit: usize) -> Result<usize> {
encoded_size(
&RetainedTaskPayload {
tool_name: &task.tool_name,
arguments: &task.arguments,
status_message: &task.status_message,
meta: &task.meta,
result: &task.result,
error: &task.error,
owner: &task.owner,
input_requests: &task.input_requests,
answered_input_keys: &task.answered_input_keys,
input_responses: &task.input_responses,
superseded_input_keys: &task.superseded_input_keys,
},
limit,
TaskRetentionLimitKind::AggregateBytes,
)
}
fn validate_payload<T: serde::Serialize + ?Sized>(
payload: &T,
limits: TaskRetentionLimits,
) -> Result<usize> {
let (limit, kind) = if limits.max_payload_bytes <= limits.max_retained_bytes {
(
limits.max_payload_bytes,
TaskRetentionLimitKind::PayloadBytes,
)
} else {
(
limits.max_retained_bytes,
TaskRetentionLimitKind::AggregateBytes,
)
};
encoded_size(payload, limit, kind)
}
fn validate_prefixed_string(
value: &str,
prefix: &str,
limits: TaskRetentionLimits,
) -> Result<usize> {
fn check(
value: &str,
prefix: &str,
limit: usize,
kind: TaskRetentionLimitKind,
) -> Result<usize> {
let Some(value_limit) = limit.checked_sub(prefix.len()) else {
return Err(TaskStoreError::RetentionLimitExceeded { kind, limit });
};
encoded_size(value, value_limit, kind)?
.checked_add(prefix.len())
.ok_or(TaskStoreError::RetentionLimitExceeded { kind, limit })
}
let (limit, kind) = if limits.max_payload_bytes <= limits.max_retained_bytes {
(
limits.max_payload_bytes,
TaskRetentionLimitKind::PayloadBytes,
)
} else {
(
limits.max_retained_bytes,
TaskRetentionLimitKind::AggregateBytes,
)
};
check(value, prefix, limit, kind)
}
pub type TaskSnapshot = (TaskObject, Option<CallToolResult>, Option<JsonRpcError>);
#[async_trait]
pub trait TaskStore: Send + Sync + 'static {
async fn create_task(
&self,
tool_name: &str,
arguments: serde_json::Value,
ttl: Option<u64>,
owner: TaskOwner,
) -> Result<(String, CancellationToken)>;
async fn task_owner(&self, task_id: &str) -> Result<Option<TaskOwner>>;
async fn get_task(&self, task_id: &str) -> Result<Option<TaskObject>>;
async fn set_task_meta(&self, task_id: &str, meta: serde_json::Value) -> Result<bool> {
let _ = (task_id, meta);
Ok(false)
}
async fn discard_task(&self, task_id: &str) -> Result<bool> {
let _ = task_id;
Ok(false)
}
async fn get_task_result(&self, task_id: &str) -> Result<Option<TaskSnapshot>>;
async fn wait_for_completion(&self, task_id: &str) -> Result<Option<TaskSnapshot>>;
async fn list_tasks(&self, status_filter: Option<TaskStatus>) -> Result<Vec<TaskObject>>;
async fn require_input(
&self,
task_id: &str,
requests: InputRequests,
message: Option<&str>,
) -> Result<bool>;
async fn outstanding_input_requests(&self, task_id: &str) -> Result<Option<InputRequests>>;
async fn apply_input_responses(
&self,
task_id: &str,
responses: InputResponses,
) -> Result<Option<AppliedInputResponses>>;
async fn task_presence(&self, task_id: &str) -> Result<TaskPresence> {
Ok(match self.task_owner(task_id).await? {
Some(owner) => TaskPresence::Present { owner },
None => TaskPresence::Missing,
})
}
async fn input_responses(&self, task_id: &str) -> Result<Option<InputResponses>> {
Ok(self
.resume_context(task_id)
.await?
.map(|resume| resume.input_responses))
}
async fn set_status(
&self,
task_id: &str,
status: TaskStatus,
message: Option<&str>,
) -> Result<bool> {
let _ = (task_id, status, message);
Ok(false)
}
async fn resume_context(&self, task_id: &str) -> Result<Option<TaskResumeContext>> {
let _ = task_id;
Ok(None)
}
async fn set_ttl(&self, task_id: &str, ttl_ms: u64) -> Result<bool>;
async fn complete_task(&self, task_id: &str, result: CallToolResult) -> Result<bool>;
async fn fail_task(&self, task_id: &str, error: JsonRpcError) -> Result<bool>;
async fn cancel_task(&self, task_id: &str, reason: Option<&str>) -> Result<Option<TaskObject>>;
}
#[derive(Debug, Default)]
struct WorkerSignalState {
generation: u64,
shutdown: bool,
}
#[derive(Debug, Default)]
struct WorkerSignal {
state: Mutex<WorkerSignalState>,
changed: Condvar,
}
impl WorkerSignal {
fn generation(&self) -> Option<u64> {
let state = self.state.lock().ok()?;
(!state.shutdown).then_some(state.generation)
}
fn wake(&self) {
if let Ok(mut state) = self.state.lock() {
state.generation = state.generation.wrapping_add(1);
self.changed.notify_one();
}
}
fn shutdown(&self) {
if let Ok(mut state) = self.state.lock() {
state.shutdown = true;
state.generation = state.generation.wrapping_add(1);
self.changed.notify_all();
}
}
fn wait_for_change(&self, generation: u64, timeout: Duration) -> bool {
let Ok(state) = self.state.lock() else {
return true;
};
if state.shutdown {
return true;
}
if state.generation != generation {
return false;
}
match self.changed.wait_timeout(state, timeout) {
Ok((state, _)) => state.shutdown,
Err(_) => true,
}
}
}
#[derive(Debug, Default, Clone, Copy, PartialEq, Eq)]
struct Retirement {
signalled: usize,
removed: usize,
}
#[derive(Debug, Default)]
struct MemoryTaskStoreData {
tasks: HashMap<String, StoredTask>,
retained_bytes: usize,
reserved_bytes: usize,
}
impl MemoryTaskStoreData {
fn usage(&self) -> TaskStoreUsage {
TaskStoreUsage {
task_count: self.tasks.len(),
retained_bytes: self.retained_bytes,
reserved_bytes: self.reserved_bytes,
}
}
}
#[derive(Debug)]
struct MemoryTaskStoreState {
data: RwLock<MemoryTaskStoreData>,
config: MemoryTaskStoreConfig,
retention_limits: TaskRetentionLimits,
worker_signal: Arc<WorkerSignal>,
worker_started: OnceLock<std::result::Result<(), String>>,
}
impl MemoryTaskStoreState {
fn retire_expired(&self, remove: bool) -> Retirement {
let Ok(mut data) = self.data.write() else {
return Retirement::default();
};
let now = Instant::now();
let mut retirement = Retirement::default();
let mut retired_old_retained = 0usize;
let mut retired_new_retained = 0usize;
let mut retired_old_reserved = 0usize;
for task in data.tasks.values_mut() {
if task.is_expired_at(now) && task.signal_expiry() {
retirement.signalled += 1;
let old_retained = task.retained_bytes;
let old_reserved = task.reserved_bytes;
task.scrub_expired_payload();
task.retained_bytes = retained_payload_size(&task.task, usize::MAX)
.expect("built-in task payload serialization is infallible");
retired_old_retained = retired_old_retained
.checked_add(old_retained)
.expect("task-store retained-byte accounting overflowed");
retired_new_retained = retired_new_retained
.checked_add(task.retained_bytes)
.expect("task-store retained-byte accounting overflowed");
retired_old_reserved = retired_old_reserved
.checked_add(old_reserved)
.expect("task-store reserved-byte accounting overflowed");
}
}
data.retained_bytes = data
.retained_bytes
.checked_sub(retired_old_retained)
.and_then(|bytes| bytes.checked_add(retired_new_retained))
.expect("task-store retained-byte accounting invariant violated");
data.reserved_bytes = data
.reserved_bytes
.checked_sub(retired_old_reserved)
.expect("task-store reserved-byte accounting invariant violated");
charged_bytes(
data.retained_bytes,
data.reserved_bytes,
self.retention_limits.max_retained_bytes,
)
.expect("expiry scrubbing exceeded the task-store aggregate-byte invariant");
if remove {
let before = data.tasks.len();
let mut removed_retained = 0usize;
let mut removed_reserved = 0usize;
data.tasks.retain(|_, task| {
let keep = !task.is_expired_at(now);
if !keep {
removed_retained = removed_retained
.checked_add(task.retained_bytes)
.expect("task-store retained-byte accounting overflowed");
removed_reserved = removed_reserved
.checked_add(task.reserved_bytes)
.expect("task-store reserved-byte accounting overflowed");
}
keep
});
data.retained_bytes = data
.retained_bytes
.checked_sub(removed_retained)
.expect("task-store retained-byte accounting invariant violated");
data.reserved_bytes = data
.reserved_bytes
.checked_sub(removed_reserved)
.expect("task-store reserved-byte accounting invariant violated");
retirement.removed = before - data.tasks.len();
}
retirement
}
fn next_expiry(&self) -> Option<Instant> {
self.data
.read()
.ok()?
.tasks
.values()
.filter_map(|task| {
(!task.expiry_signalled)
.then(|| task.expires_at())
.flatten()
})
.min()
}
}
impl Drop for MemoryTaskStoreState {
fn drop(&mut self) {
self.worker_signal.shutdown();
}
}
fn next_deadline(a: Option<Instant>, b: Option<Instant>) -> Option<Instant> {
match (a, b) {
(Some(a), Some(b)) => Some(a.min(b)),
(Some(deadline), None) | (None, Some(deadline)) => Some(deadline),
(None, None) => None,
}
}
fn memory_task_store_worker(state: Weak<MemoryTaskStoreState>, signal: Arc<WorkerSignal>) {
const MAX_SLEEP: Duration = Duration::from_secs(60 * 60);
let Some(initial) = state.upgrade() else {
return;
};
let cleanup_interval = if initial.config.cleanup_interval.is_zero() {
Duration::from_millis(1)
} else {
initial.config.cleanup_interval
};
let mut cleanup_at = Instant::now().checked_add(cleanup_interval);
drop(initial);
loop {
let Some(state) = state.upgrade() else {
break;
};
let Some(generation) = signal.generation() else {
break;
};
let now = Instant::now();
if cleanup_at.is_some_and(|deadline| now >= deadline) {
state.retire_expired(true);
cleanup_at = Instant::now().checked_add(cleanup_interval);
} else {
state.retire_expired(false);
}
let deadline = next_deadline(cleanup_at, state.next_expiry());
let timeout = deadline
.map(|deadline| deadline.saturating_duration_since(Instant::now()))
.unwrap_or(MAX_SLEEP)
.min(MAX_SLEEP);
drop(state);
if signal.wait_for_change(generation, timeout) {
break;
}
}
}
#[derive(Debug, Clone)]
pub struct MemoryTaskStore {
state: Arc<MemoryTaskStoreState>,
}
impl Default for MemoryTaskStore {
fn default() -> Self {
Self::new()
}
}
impl MemoryTaskStore {
pub fn new() -> Self {
Self::with_config(MemoryTaskStoreConfig::default())
}
pub fn with_config(config: MemoryTaskStoreConfig) -> Self {
Self::with_config_and_retention(config, TaskRetentionLimits::default())
}
pub fn with_retention_limits(retention_limits: TaskRetentionLimits) -> Self {
Self::with_config_and_retention(MemoryTaskStoreConfig::default(), retention_limits)
}
pub fn with_config_and_retention(
config: MemoryTaskStoreConfig,
retention_limits: TaskRetentionLimits,
) -> Self {
Self {
state: Arc::new(MemoryTaskStoreState {
data: RwLock::new(MemoryTaskStoreData::default()),
config,
retention_limits,
worker_signal: Arc::new(WorkerSignal::default()),
worker_started: OnceLock::new(),
}),
}
}
fn ensure_worker(&self) -> Result<()> {
let weak = Arc::downgrade(&self.state);
let signal = self.state.worker_signal.clone();
match self.state.worker_started.get_or_init(|| {
std::thread::Builder::new()
.name("tower-mcp-task-expiry".to_string())
.spawn(move || memory_task_store_worker(weak, signal))
.map(drop)
.map_err(|error| format!("failed to start task expiry worker: {error}"))
}) {
Ok(()) => Ok(()),
Err(error) => Err(TaskStoreError::Backend(error.clone())),
}
}
pub fn cleanup_expired(&self) -> usize {
self.state.retire_expired(true).removed
}
#[must_use]
pub fn usage(&self) -> TaskStoreUsage {
match self.state.data.read() {
Ok(data) => data.usage(),
Err(poisoned) => poisoned.into_inner().usage(),
}
}
#[cfg(test)]
pub fn len(&self) -> usize {
if let Ok(data) = self.state.data.read() {
data.tasks.len()
} else {
0
}
}
#[cfg(test)]
pub fn is_empty(&self) -> bool {
self.len() == 0
}
}
#[async_trait]
impl TaskStore for MemoryTaskStore {
async fn create_task(
&self,
tool_name: &str,
arguments: serde_json::Value,
ttl: Option<u64>,
owner: TaskOwner,
) -> Result<(String, CancellationToken)> {
self.ensure_worker()?;
let limits = self.state.retention_limits;
validate_payload(tool_name, limits)?;
validate_payload(&arguments, limits)?;
validate_payload(&owner, limits)?;
let id = generate_task_id();
let ttl = ttl.unwrap_or_else(|| duration_millis_saturated(self.state.config.default_ttl));
let task = Task::new(id.clone(), tool_name.to_string(), arguments, ttl, owner);
let token = task.cancellation_token.clone();
let stored = prepare_stored_task(task, limits)?;
self.state.retire_expired(true);
let mut data = self.state.data.write().map_err(|_| {
TaskStoreError::Backend("in-memory task store lock poisoned".to_string())
})?;
if data.tasks.len() >= limits.max_tasks {
return Err(TaskStoreError::RetentionLimitExceeded {
kind: TaskRetentionLimitKind::TaskCount,
limit: limits.max_tasks,
});
}
let retained_bytes = data
.retained_bytes
.checked_add(stored.retained_bytes)
.ok_or(TaskStoreError::RetentionLimitExceeded {
kind: TaskRetentionLimitKind::AggregateBytes,
limit: limits.max_retained_bytes,
})?;
let reserved_bytes = data
.reserved_bytes
.checked_add(stored.reserved_bytes)
.ok_or(TaskStoreError::RetentionLimitExceeded {
kind: TaskRetentionLimitKind::AggregateBytes,
limit: limits.max_retained_bytes,
})?;
charged_bytes(retained_bytes, reserved_bytes, limits.max_retained_bytes)?;
data.retained_bytes = retained_bytes;
data.reserved_bytes = reserved_bytes;
data.tasks.insert(id.clone(), stored);
drop(data);
self.state.worker_signal.wake();
Ok((id, token))
}
async fn get_task(&self, task_id: &str) -> Result<Option<TaskObject>> {
self.state.retire_expired(false);
Ok(if let Ok(data) = self.state.data.read() {
data.tasks
.get(task_id)
.filter(|t| !t.is_expired())
.map(|t| t.to_task_object())
} else {
None
})
}
async fn set_task_meta(&self, task_id: &str, meta: serde_json::Value) -> Result<bool> {
self.state.retire_expired(false);
let Ok(mut data) = self.state.data.write() else {
return Ok(false);
};
let Some(current) = data.tasks.get(task_id).filter(|task| !task.is_expired()) else {
return Ok(false);
};
let limits = self.state.retention_limits;
validate_payload(&meta, limits)?;
let mut task = current.task.clone();
task.meta = Some(meta);
let replacement = prepare_stored_task(task, limits)?;
replace_stored_task(&mut data, task_id, replacement, limits.max_retained_bytes)
}
async fn discard_task(&self, task_id: &str) -> Result<bool> {
let Ok(mut data) = self.state.data.write() else {
return Ok(false);
};
let Some(removed) = data.tasks.remove(task_id) else {
return Ok(false);
};
data.retained_bytes = data
.retained_bytes
.checked_sub(removed.retained_bytes)
.expect("task-store retained-byte accounting invariant violated");
data.reserved_bytes = data
.reserved_bytes
.checked_sub(removed.reserved_bytes)
.expect("task-store reserved-byte accounting invariant violated");
Ok(true)
}
async fn task_owner(&self, task_id: &str) -> Result<Option<TaskOwner>> {
self.state.retire_expired(false);
Ok(if let Ok(data) = self.state.data.read() {
data.tasks
.get(task_id)
.filter(|t| !t.is_expired())
.map(|t| t.owner.clone())
} else {
None
})
}
async fn task_presence(&self, task_id: &str) -> Result<TaskPresence> {
self.state.retire_expired(false);
let Ok(data) = self.state.data.read() else {
return Ok(TaskPresence::Missing);
};
Ok(match data.tasks.get(task_id) {
Some(task) if task.is_expired() => TaskPresence::Expired {
owner: task.owner.clone(),
},
Some(task) => TaskPresence::Present {
owner: task.owner.clone(),
},
None => TaskPresence::Missing,
})
}
async fn get_task_result(&self, task_id: &str) -> Result<Option<TaskSnapshot>> {
self.state.retire_expired(false);
Ok(if let Ok(data) = self.state.data.read() {
data.tasks
.get(task_id)
.filter(|t| !t.is_expired())
.map(|t| (t.to_task_object(), t.result.clone(), t.error.clone()))
} else {
None
})
}
async fn wait_for_completion(&self, task_id: &str) -> Result<Option<TaskSnapshot>> {
self.state.retire_expired(false);
let (mut notified, cancellation) = {
let Ok(data) = self.state.data.read() else {
return Ok(None);
};
let Some(task) = data.tasks.get(task_id).filter(|t| !t.is_expired()) else {
return Ok(None);
};
if task.status.is_terminal() {
return Ok(Some((
task.to_task_object(),
task.result.clone(),
task.error.clone(),
)));
}
let mut notified = Box::pin(task.completion_notify.clone().notified_owned());
let _ = notified.as_mut().enable();
(notified, task.cancellation_token.clone())
};
let cancellation_fired = tokio::select! {
_ = &mut notified => false,
_ = cancellation.cancelled() => true,
};
if cancellation_fired {
match self.get_task_result(task_id).await? {
Some((task, _, _)) if !task.status.is_terminal() => {
notified.await;
}
snapshot => return Ok(snapshot),
}
}
self.get_task_result(task_id).await
}
async fn list_tasks(&self, status_filter: Option<TaskStatus>) -> Result<Vec<TaskObject>> {
self.state.retire_expired(false);
Ok(if let Ok(data) = self.state.data.read() {
data.tasks
.values()
.filter(|t| !t.is_expired())
.filter(|t| status_filter.is_none() || status_filter == Some(t.status))
.map(|t| t.to_task_object())
.collect()
} else {
vec![]
})
}
async fn require_input(
&self,
task_id: &str,
requests: InputRequests,
message: Option<&str>,
) -> Result<bool> {
self.state.retire_expired(false);
let Ok(mut data) = self.state.data.write() else {
return Ok(false);
};
let Some(current) = data.tasks.get(task_id).filter(|task| !task.is_expired()) else {
return Ok(false);
};
if current.status.is_terminal() {
return Ok(false);
}
let limits = self.state.retention_limits;
validate_payload(&requests, limits)?;
if let Some(message) = message {
validate_payload(message, limits)?;
}
let mut task = current.task.clone();
let reused: Vec<String> = requests
.iter()
.filter(|(key, request)| {
if task.answered_input_keys.contains(*key)
|| task.superseded_input_keys.contains(*key)
{
return true;
}
match task.input_requests.get(*key) {
Some(current) => !same_input_request(current, request),
None => false,
}
})
.map(|(key, _)| key.clone())
.collect();
if !reused.is_empty() {
return Err(TaskStoreError::InvalidTransition(format!(
"input request keys must be unique over a task's lifetime, but {} \
already {} used by this task; use a new key to ask again",
reused.join(", "),
if reused.len() == 1 { "was" } else { "were" },
)));
}
for key in std::mem::take(&mut task.input_requests).into_keys() {
if !requests.contains_key(&key) {
task.superseded_input_keys.insert(key);
}
}
task.input_requests = requests;
task.status = TaskStatus::InputRequired;
task.status_message = Some(
message
.map(str::to_string)
.unwrap_or_else(|| "Awaiting client input".to_string()),
);
task.last_updated_at_str = chrono_now_iso8601();
let replacement = prepare_stored_task(task, limits)?;
replace_stored_task(&mut data, task_id, replacement, limits.max_retained_bytes)
}
async fn outstanding_input_requests(&self, task_id: &str) -> Result<Option<InputRequests>> {
self.state.retire_expired(false);
Ok(if let Ok(data) = self.state.data.read() {
data.tasks
.get(task_id)
.filter(|t| !t.is_expired())
.map(|t| t.input_requests.clone())
} else {
None
})
}
async fn apply_input_responses(
&self,
task_id: &str,
responses: InputResponses,
) -> Result<Option<AppliedInputResponses>> {
self.state.retire_expired(false);
let Ok(mut data) = self.state.data.write() else {
return Ok(None);
};
let Some(current) = data.tasks.get(task_id).filter(|task| !task.is_expired()) else {
return Ok(None);
};
if current.status.is_terminal() {
return Ok(None);
}
let limits = self.state.retention_limits;
let accepted_payload: std::collections::BTreeMap<&str, &crate::protocol::InputResponse> =
responses
.iter()
.filter(|(key, _)| current.input_requests.contains_key(*key))
.map(|(key, response)| (key.as_str(), response))
.collect();
if !accepted_payload.is_empty() {
validate_payload(&accepted_payload, limits)?;
}
let mut task = current.task.clone();
let mut applied = AppliedInputResponses::default();
for (key, response) in responses {
if task.input_requests.remove(&key).is_some() {
task.answered_input_keys.insert(key.clone());
task.input_responses.insert(key.clone(), response);
applied.accepted.insert(key);
} else {
applied.ignored.insert(key);
}
}
applied.still_outstanding = task.input_requests.keys().cloned().collect();
if !applied.accepted.is_empty() {
task.last_updated_at_str = chrono_now_iso8601();
}
if applied.is_complete() && task.status == TaskStatus::InputRequired {
task.status = TaskStatus::Working;
task.status_message = Some("Task resumed".to_string());
}
let replacement = prepare_stored_task(task, limits)?;
if replace_stored_task(&mut data, task_id, replacement, limits.max_retained_bytes)? {
Ok(Some(applied))
} else {
Ok(None)
}
}
async fn set_status(
&self,
task_id: &str,
status: TaskStatus,
message: Option<&str>,
) -> Result<bool> {
if status.is_terminal() {
return Err(TaskStoreError::InvalidTransition(format!(
"set_status is for non-terminal progress; use complete_task, fail_task, or cancel_task to reach {status:?}"
)));
}
self.state.retire_expired(false);
let Ok(mut data) = self.state.data.write() else {
return Ok(false);
};
let Some(current) = data.tasks.get(task_id).filter(|task| !task.is_expired()) else {
return Ok(false);
};
if current.status.is_terminal() {
return Ok(false);
}
let limits = self.state.retention_limits;
if let Some(message) = message {
validate_payload(message, limits)?;
}
let mut task = current.task.clone();
task.status = status;
if let Some(message) = message {
task.status_message = Some(message.to_string());
}
task.last_updated_at_str = chrono_now_iso8601();
let replacement = prepare_stored_task(task, limits)?;
replace_stored_task(&mut data, task_id, replacement, limits.max_retained_bytes)
}
async fn resume_context(&self, task_id: &str) -> Result<Option<TaskResumeContext>> {
self.state.retire_expired(false);
let Ok(data) = self.state.data.read() else {
return Ok(None);
};
Ok(data
.tasks
.get(task_id)
.filter(|task| !task.is_expired())
.map(|task| {
TaskResumeContext::new(
task.tool_name.clone(),
task.arguments.clone(),
task.input_responses.clone(),
)
.with_cancellation_token(task.cancellation_token.clone())
}))
}
async fn set_ttl(&self, task_id: &str, ttl_ms: u64) -> Result<bool> {
let Ok(mut data) = self.state.data.write() else {
return Ok(false);
};
let Some(task) = data.tasks.get_mut(task_id) else {
return Ok(false);
};
if task.is_expired() {
drop(data);
self.state.retire_expired(false);
return Ok(false);
}
task.ttl = ttl_ms;
task.last_updated_at_str = chrono_now_iso8601();
drop(data);
self.state.retire_expired(false);
self.state.worker_signal.wake();
Ok(true)
}
async fn complete_task(&self, task_id: &str, result: CallToolResult) -> Result<bool> {
self.state.retire_expired(false);
let Ok(mut data) = self.state.data.write() else {
return Ok(false);
};
let Some(mut task) = data.tasks.remove(task_id) else {
return Ok(false);
};
if task.status.is_terminal() {
data.tasks.insert(task_id.to_string(), task);
return Ok(false);
}
if task.is_expired() {
data.tasks.insert(task_id.to_string(), task);
drop(data);
self.state.retire_expired(false);
return Ok(false);
}
let limits = self.state.retention_limits;
let old_retained = task.retained_bytes;
let old_reserved = task.reserved_bytes;
if let Err(error) = validate_payload(&result, limits) {
drop(result);
commit_retention_failure(&mut data, task_id, task, old_retained, old_reserved, limits);
return Err(error);
}
task.status = TaskStatus::Completed;
task.status_message = Some("Task completed".to_string());
task.result = Some(result);
task.input_requests = InputRequests::new();
task.completed_at = Some(Instant::now());
task.last_updated_at_str = chrono_now_iso8601();
task.retained_bytes = match retained_payload_size(&task.task, limits.max_retained_bytes) {
Ok(bytes) => bytes,
Err(error) => {
commit_retention_failure(
&mut data,
task_id,
task,
old_retained,
old_reserved,
limits,
);
return Err(error);
}
};
task.reserved_bytes = 0;
let (retained_bytes, reserved_bytes) = match replacement_totals(
&data,
old_retained,
old_reserved,
task.retained_bytes,
0,
limits.max_retained_bytes,
) {
Ok(totals) => totals,
Err(error) => {
commit_retention_failure(
&mut data,
task_id,
task,
old_retained,
old_reserved,
limits,
);
return Err(error);
}
};
let notify = task.completion_notify.clone();
data.retained_bytes = retained_bytes;
data.reserved_bytes = reserved_bytes;
data.tasks.insert(task_id.to_string(), task);
notify.notify_waiters();
Ok(true)
}
async fn fail_task(&self, task_id: &str, error: JsonRpcError) -> Result<bool> {
self.state.retire_expired(false);
let Ok(mut data) = self.state.data.write() else {
return Ok(false);
};
let Some(mut task) = data.tasks.remove(task_id) else {
return Ok(false);
};
if task.status.is_terminal() {
data.tasks.insert(task_id.to_string(), task);
return Ok(false);
}
if task.is_expired() {
data.tasks.insert(task_id.to_string(), task);
drop(data);
self.state.retire_expired(false);
return Ok(false);
}
let limits = self.state.retention_limits;
let old_retained = task.retained_bytes;
let old_reserved = task.reserved_bytes;
let error_validation = validate_payload(&error, limits);
if let Err(retention_error) = error_validation {
drop(error);
commit_retention_failure(&mut data, task_id, task, old_retained, old_reserved, limits);
return Err(retention_error);
}
if let Err(retention_error) =
validate_prefixed_string(&error.message, "Task failed: ", limits)
{
drop(error);
commit_retention_failure(&mut data, task_id, task, old_retained, old_reserved, limits);
return Err(retention_error);
}
let status_message = format!("Task failed: {}", error.message);
task.status = TaskStatus::Failed;
task.status_message = Some(status_message);
task.error = Some(error);
task.input_requests = InputRequests::new();
task.completed_at = Some(Instant::now());
task.last_updated_at_str = chrono_now_iso8601();
task.retained_bytes = match retained_payload_size(&task.task, limits.max_retained_bytes) {
Ok(bytes) => bytes,
Err(error) => {
commit_retention_failure(
&mut data,
task_id,
task,
old_retained,
old_reserved,
limits,
);
return Err(error);
}
};
task.reserved_bytes = 0;
let (retained_bytes, reserved_bytes) = match replacement_totals(
&data,
old_retained,
old_reserved,
task.retained_bytes,
0,
limits.max_retained_bytes,
) {
Ok(totals) => totals,
Err(error) => {
commit_retention_failure(
&mut data,
task_id,
task,
old_retained,
old_reserved,
limits,
);
return Err(error);
}
};
let notify = task.completion_notify.clone();
data.retained_bytes = retained_bytes;
data.reserved_bytes = reserved_bytes;
data.tasks.insert(task_id.to_string(), task);
notify.notify_waiters();
Ok(true)
}
async fn cancel_task(&self, task_id: &str, reason: Option<&str>) -> Result<Option<TaskObject>> {
self.state.retire_expired(false);
let Ok(mut data) = self.state.data.write() else {
return Ok(None);
};
let Some(mut task) = data.tasks.remove(task_id) else {
return Ok(None);
};
if task.is_expired() {
data.tasks.insert(task_id.to_string(), task);
drop(data);
self.state.retire_expired(false);
return Ok(None);
}
task.cancellation_token.cancel();
if !task.status.is_terminal() {
let limits = self.state.retention_limits;
let old_retained = task.retained_bytes;
let old_reserved = task.reserved_bytes;
let bounded_reason = match reason {
Some(reason) => validate_prefixed_string(reason, "Cancelled: ", limits).is_ok(),
None => validate_payload("Task cancelled", limits).is_ok(),
};
if bounded_reason {
let requested_status = reason
.map(|reason| format!("Cancelled: {reason}"))
.unwrap_or_else(|| "Task cancelled".to_string());
task.input_requests = InputRequests::new();
task.status = TaskStatus::Cancelled;
task.status_message = Some(requested_status);
task.completed_at = Some(Instant::now());
task.last_updated_at_str = chrono_now_iso8601();
} else {
cancel_with_bounded_status(
&mut task.task,
"Task cancelled: retention limit exceeded",
true,
);
}
let candidate_bytes = retained_payload_size(&task.task, limits.max_retained_bytes);
let totals = candidate_bytes.and_then(|bytes| {
replacement_totals(
&data,
old_retained,
old_reserved,
bytes,
0,
limits.max_retained_bytes,
)
.map(|totals| (bytes, totals))
});
let (new_retained, (retained_bytes, reserved_bytes)) = match totals {
Ok(totals) => totals,
Err(_) => {
cancel_with_bounded_status(
&mut task.task,
"Task cancelled: retention limit exceeded",
true,
);
let bytes = retained_payload_size(&task.task, limits.max_retained_bytes)
.expect("bounded cancellation must fit the aggregate byte limit");
let totals = replacement_totals(
&data,
old_retained,
old_reserved,
bytes,
0,
limits.max_retained_bytes,
)
.expect("bounded cancellation must fit reserved global accounting");
(bytes, totals)
}
};
task.retained_bytes = new_retained;
task.reserved_bytes = 0;
let notify = task.completion_notify.clone();
data.retained_bytes = retained_bytes;
data.reserved_bytes = reserved_bytes;
let object = task.to_task_object();
data.tasks.insert(task_id.to_string(), task);
notify.notify_waiters();
return Ok(Some(object));
}
let object = task.to_task_object();
data.tasks.insert(task_id.to_string(), task);
Ok(Some(object))
}
}
pub fn tasks_extension() -> crate::ExtensionDeclaration {
crate::ExtensionDeclaration::empty(crate::protocol::TASKS_EXTENSION_ID)
.expect("the built-in Tasks extension declaration is valid")
}
impl crate::McpRouter {
pub fn with_tasks(self) -> Self {
self.with_protocol_extension(tasks_extension())
}
}
impl crate::McpClientBuilder {
pub fn with_tasks(self) -> Self {
self.with_protocol_extension(tasks_extension())
}
}
impl crate::RequestContext {
pub fn supports_tasks(&self) -> bool {
self.negotiated_extensions()
.is_some_and(|extensions| extensions.contains(crate::protocol::TASKS_EXTENSION_ID))
}
}
fn chrono_now_iso8601() -> String {
use std::time::SystemTime;
let now = SystemTime::now();
let duration = now
.duration_since(SystemTime::UNIX_EPOCH)
.unwrap_or_default();
let secs = duration.as_secs();
let millis = duration.subsec_millis();
let days = secs / 86400;
let remaining = secs % 86400;
let hours = remaining / 3600;
let remaining = remaining % 3600;
let minutes = remaining / 60;
let seconds = remaining % 60;
let mut year = 1970i32;
let mut remaining_days = days as i32;
loop {
let days_in_year = if is_leap_year(year) { 366 } else { 365 };
if remaining_days < days_in_year {
break;
}
remaining_days -= days_in_year;
year += 1;
}
let days_in_months: [i32; 12] = if is_leap_year(year) {
[31, 29, 31, 30, 31, 30, 31, 31, 30, 31, 30, 31]
} else {
[31, 28, 31, 30, 31, 30, 31, 31, 30, 31, 30, 31]
};
let mut month = 1;
for days_in_month in days_in_months.iter() {
if remaining_days < *days_in_month {
break;
}
remaining_days -= days_in_month;
month += 1;
}
let day = remaining_days + 1;
format!(
"{:04}-{:02}-{:02}T{:02}:{:02}:{:02}.{:03}Z",
year, month, day, hours, minutes, seconds, millis
)
}
fn is_leap_year(year: i32) -> bool {
(year % 4 == 0 && year % 100 != 0) || (year % 400 == 0)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::protocol::{
ElicitAction, ElicitFieldValue, ElicitFormParams, ElicitFormSchema, ElicitRequestParams,
ElicitResult, InputRequest, InputResponse, ListRootsParams,
};
#[test]
fn task_resume_context_constructor_preserves_every_field() {
let input_responses = InputResponses::from([(
"approval".to_string(),
InputResponse::Elicit(ElicitResult {
action: ElicitAction::Accept,
content: None,
meta: None,
}),
)]);
let context = TaskResumeContext::new(
"build_report",
serde_json::json!({"format": "pdf"}),
input_responses.clone(),
);
assert_eq!(context.tool_name, "build_report");
assert_eq!(context.arguments, serde_json::json!({"format": "pdf"}));
assert!(context.cancellation_token.is_none());
assert_eq!(
serde_json::to_value(&context.input_responses).unwrap(),
serde_json::to_value(&input_responses).unwrap()
);
}
#[tokio::test]
async fn test_create_task() {
let store = MemoryTaskStore::new();
let (id, token) = store
.create_task("test-tool", serde_json::json!({"a": 1}), None, None)
.await
.unwrap();
assert!(!id.is_empty());
assert!(!token.is_cancelled());
let info = store
.get_task(&id)
.await
.unwrap()
.expect("task should exist");
assert_eq!(info.task_id, id);
assert_eq!(info.status, TaskStatus::Working);
assert_eq!(info.ttl, Some(300_000));
}
#[tokio::test]
async fn configured_default_ttl_is_used_when_creation_omits_one() {
let store = MemoryTaskStore::with_config(
MemoryTaskStoreConfig::default().default_ttl(Duration::from_secs(42)),
);
let (id, _) = store
.create_task("test-tool", serde_json::json!({}), None, None)
.await
.unwrap();
assert_eq!(
store.get_task(&id).await.unwrap().unwrap().ttl,
Some(42_000)
);
}
#[tokio::test]
async fn working_task_expiry_cancels_and_wakes_completion_waiter() {
let store = MemoryTaskStore::with_config(
MemoryTaskStoreConfig::default()
.default_ttl(Duration::from_secs(60))
.cleanup_interval(Duration::from_secs(60)),
);
let (id, cancellation) = store
.create_task("test-tool", serde_json::json!({}), None, None)
.await
.unwrap();
let waiting_store = store.clone();
let waiting_id = id.clone();
let waiter = tokio::spawn(async move {
waiting_store
.wait_for_completion(&waiting_id)
.await
.unwrap()
});
let mut waiter = waiter;
assert!(
tokio::time::timeout(Duration::from_millis(20), &mut waiter)
.await
.is_err()
);
assert!(store.set_ttl(&id, 0).await.unwrap());
tokio::time::timeout(Duration::from_secs(2), cancellation.cancelled())
.await
.expect("expiry did not cancel the task token");
let snapshot = tokio::time::timeout(Duration::from_secs(2), waiter)
.await
.expect("expiry did not wake the completion waiter")
.unwrap();
assert!(
snapshot.is_none(),
"an expired task has no visible snapshot"
);
assert!(matches!(
store.task_presence(&id).await.unwrap(),
TaskPresence::Expired { .. }
));
assert!(
!store
.complete_task(&id, CallToolResult::text("late"))
.await
.unwrap(),
"a terminal write must not resurrect an expired task"
);
}
#[tokio::test]
async fn automatic_cleanup_physically_reclaims_expired_tasks() {
let store = MemoryTaskStore::with_config(
MemoryTaskStoreConfig::default()
.default_ttl(Duration::from_secs(60))
.cleanup_interval(Duration::from_millis(25)),
);
let (id, _) = store
.create_task("test-tool", serde_json::json!({}), None, None)
.await
.unwrap();
assert_eq!(store.len(), 1);
assert!(store.set_ttl(&id, 0).await.unwrap());
tokio::time::timeout(Duration::from_secs(2), async {
while !store.is_empty() {
tokio::time::sleep(Duration::from_millis(5)).await;
}
})
.await
.expect("automatic cleanup did not reclaim the expired record");
}
#[tokio::test]
async fn shortening_ttl_reschedules_expiry_from_creation() {
let store = MemoryTaskStore::with_config(
MemoryTaskStoreConfig::default()
.default_ttl(Duration::from_secs(60))
.cleanup_interval(Duration::from_secs(60)),
);
let (id, cancellation) = store
.create_task("test-tool", serde_json::json!({}), None, None)
.await
.unwrap();
assert!(store.set_ttl(&id, 250).await.unwrap());
tokio::time::timeout(Duration::from_secs(3), cancellation.cancelled())
.await
.expect("shorter TTL did not reschedule the expiry wakeup");
assert!(store.get_task(&id).await.unwrap().is_none());
}
#[tokio::test]
async fn manual_and_scheduled_retirement_signal_expiry_only_once() {
let store = MemoryTaskStore::with_config(
MemoryTaskStoreConfig::default()
.default_ttl(Duration::from_secs(60))
.cleanup_interval(Duration::from_secs(60)),
);
let (id, cancellation) = store
.create_task("test-tool", serde_json::json!({}), None, None)
.await
.unwrap();
store
.state
.data
.write()
.unwrap()
.tasks
.get_mut(&id)
.unwrap()
.ttl = 0;
let first = store.state.retire_expired(false);
let second = store.state.retire_expired(false);
assert_eq!(first.signalled + second.signalled, 1);
assert!(cancellation.is_cancelled());
assert_eq!(store.cleanup_expired(), 1);
assert_eq!(store.cleanup_expired(), 0);
}
#[tokio::test]
async fn dropping_the_last_store_clone_releases_worker_state() {
let store = MemoryTaskStore::with_config(
MemoryTaskStoreConfig::default().cleanup_interval(Duration::from_secs(60)),
);
store
.create_task("test-tool", serde_json::json!({}), None, None)
.await
.unwrap();
let weak = Arc::downgrade(&store.state);
drop(store);
tokio::time::timeout(Duration::from_secs(2), async {
while weak.upgrade().is_some() {
tokio::task::yield_now().await;
}
})
.await
.expect("the expiry worker retained the store's task state");
}
#[test]
fn creating_a_task_does_not_require_a_tokio_runtime() {
let store = MemoryTaskStore::with_config(
MemoryTaskStoreConfig::default()
.default_ttl(Duration::from_secs(60))
.cleanup_interval(Duration::from_secs(60)),
);
let (id, token) = futures::executor::block_on(store.create_task(
"test-tool",
serde_json::json!({}),
None,
None,
))
.expect("the standard-thread worker should start without Tokio");
assert!(!token.is_cancelled());
assert!(
futures::executor::block_on(store.get_task(&id))
.unwrap()
.is_some()
);
}
#[tokio::test]
async fn test_task_lifecycle() {
let store = MemoryTaskStore::new();
let (id, _) = store
.create_task("test-tool", serde_json::json!({}), None, None)
.await
.unwrap();
assert!(
store
.complete_task(&id, CallToolResult::text("Done"))
.await
.unwrap()
);
let info = store.get_task(&id).await.unwrap().unwrap();
assert_eq!(info.status, TaskStatus::Completed);
}
#[tokio::test]
async fn test_task_cancellation() {
let store = MemoryTaskStore::new();
let (id, token) = store
.create_task("test-tool", serde_json::json!({}), None, None)
.await
.unwrap();
assert!(!token.is_cancelled());
let task_obj = store
.cancel_task(&id, Some("User requested"))
.await
.unwrap();
assert!(task_obj.is_some());
assert_eq!(task_obj.unwrap().status, TaskStatus::Cancelled);
assert!(token.is_cancelled());
let info = store.get_task(&id).await.unwrap().unwrap();
assert_eq!(info.status, TaskStatus::Cancelled);
}
#[tokio::test]
async fn test_task_failure() {
let store = MemoryTaskStore::new();
let (id, _) = store
.create_task("test-tool", serde_json::json!({}), None, None)
.await
.unwrap();
assert!(
store
.fail_task(&id, JsonRpcError::internal_error("Something went wrong"))
.await
.unwrap()
);
let info = store.get_task(&id).await.unwrap().unwrap();
assert_eq!(info.status, TaskStatus::Failed);
assert_eq!(
info.status_message.as_deref(),
Some("Task failed: Something went wrong")
);
}
#[tokio::test]
async fn test_list_tasks() {
let store = MemoryTaskStore::new();
store
.create_task("tool1", serde_json::json!({}), None, None)
.await
.unwrap();
store
.create_task("tool2", serde_json::json!({}), None, None)
.await
.unwrap();
let (id3, _) = store
.create_task("tool3", serde_json::json!({}), None, None)
.await
.unwrap();
store
.complete_task(&id3, CallToolResult::text("Done"))
.await
.unwrap();
let all = store.list_tasks(None).await.unwrap();
assert_eq!(all.len(), 3);
let working = store.list_tasks(Some(TaskStatus::Working)).await.unwrap();
assert_eq!(working.len(), 2);
let completed = store.list_tasks(Some(TaskStatus::Completed)).await.unwrap();
assert_eq!(completed.len(), 1);
}
#[tokio::test]
async fn test_terminal_state_immutable() {
let store = MemoryTaskStore::new();
let (id, _) = store
.create_task("test-tool", serde_json::json!({}), None, None)
.await
.unwrap();
store
.complete_task(&id, CallToolResult::text("Done"))
.await
.unwrap();
assert!(
!store
.fail_task(&id, JsonRpcError::internal_error("Error"))
.await
.unwrap()
);
let info = store.get_task(&id).await.unwrap().unwrap();
assert_eq!(info.status, TaskStatus::Completed);
}
#[tokio::test]
async fn test_task_ids_unique() {
let store = MemoryTaskStore::new();
let (id1, _) = store
.create_task("tool", serde_json::json!({}), None, None)
.await
.unwrap();
let (id2, _) = store
.create_task("tool", serde_json::json!({}), None, None)
.await
.unwrap();
let (id3, _) = store
.create_task("tool", serde_json::json!({}), None, None)
.await
.unwrap();
assert_ne!(id1, id2);
assert_ne!(id2, id3);
assert_ne!(id1, id3);
}
#[tokio::test]
async fn test_get_task_result() {
let store = MemoryTaskStore::new();
let (id, _) = store
.create_task("test-tool", serde_json::json!({}), None, None)
.await
.unwrap();
let result = CallToolResult::text("The result");
store.complete_task(&id, result).await.unwrap();
let (task_obj, result, error) = store.get_task_result(&id).await.unwrap().unwrap();
assert_eq!(task_obj.status, TaskStatus::Completed);
assert!(result.is_some());
assert!(error.is_none());
}
#[tokio::test]
async fn test_wait_for_completion_returns_terminal_snapshot() {
let store = MemoryTaskStore::new();
let (id, _) = store
.create_task("test-tool", serde_json::json!({}), None, None)
.await
.unwrap();
let waiter_store = store.clone();
let waiter_id = id.clone();
let waiter =
tokio::spawn(async move { waiter_store.wait_for_completion(&waiter_id).await });
tokio::time::sleep(Duration::from_millis(10)).await;
store
.complete_task(&id, CallToolResult::text("Done"))
.await
.unwrap();
let (task_obj, result, error) = waiter.await.unwrap().unwrap().unwrap();
assert_eq!(task_obj.status, TaskStatus::Completed);
assert!(result.is_some());
assert!(error.is_none());
}
#[tokio::test]
async fn wait_for_completion_does_not_treat_live_cancellation_signal_as_terminal() {
let store = MemoryTaskStore::new();
let (id, cancellation) = store
.create_task("test-tool", serde_json::json!({}), None, None)
.await
.unwrap();
let waiter_store = store.clone();
let waiter_id = id.clone();
let mut waiter =
tokio::spawn(
async move { waiter_store.wait_for_completion(&waiter_id).await.unwrap() },
);
cancellation.cancel();
assert!(
tokio::time::timeout(Duration::from_millis(20), &mut waiter)
.await
.is_err(),
"the cooperative signal alone must not return a working snapshot"
);
store
.cancel_task(&id, Some("teardown complete"))
.await
.unwrap();
let (task, _, _) = waiter.await.unwrap().unwrap();
assert_eq!(task.status, TaskStatus::Cancelled);
}
#[tokio::test]
async fn dyn_task_store_object_safe() {
let store: Arc<dyn TaskStore> = Arc::new(MemoryTaskStore::new());
let (id, _) = store
.create_task("tool", serde_json::json!({}), None, None)
.await
.unwrap();
assert!(store.get_task(&id).await.unwrap().is_some());
}
#[test]
fn test_iso8601_timestamp() {
let ts = chrono_now_iso8601();
assert!(ts.ends_with('Z'));
assert!(ts.contains('T'));
assert_eq!(ts.len(), 24); }
#[test]
fn test_task_status_display() {
assert_eq!(TaskStatus::Working.to_string(), "working");
assert_eq!(TaskStatus::InputRequired.to_string(), "input_required");
assert_eq!(TaskStatus::Completed.to_string(), "completed");
assert_eq!(TaskStatus::Failed.to_string(), "failed");
assert_eq!(TaskStatus::Cancelled.to_string(), "cancelled");
}
#[test]
fn test_task_status_is_terminal() {
assert!(!TaskStatus::Working.is_terminal());
assert!(!TaskStatus::InputRequired.is_terminal());
assert!(TaskStatus::Completed.is_terminal());
assert!(TaskStatus::Failed.is_terminal());
assert!(TaskStatus::Cancelled.is_terminal());
}
fn requests(keys: &[&str]) -> InputRequests {
keys.iter()
.map(|k| {
(
k.to_string(),
InputRequest::ListRoots(ListRootsParams { meta: None }),
)
})
.collect()
}
fn accept(key: &str) -> (String, InputResponse) {
(
key.to_string(),
InputResponse::Elicit(ElicitResult {
action: ElicitAction::Accept,
content: None,
meta: None,
}),
)
}
async fn working_task(store: &MemoryTaskStore, ttl: Option<u64>) -> String {
store
.create_task("tool", serde_json::json!({}), ttl, None)
.await
.unwrap()
.0
}
#[tokio::test]
async fn task_ids_are_unguessable_not_sequential() {
let store = MemoryTaskStore::new();
let mut ids = BTreeSet::new();
for _ in 0..64 {
ids.insert(working_task(&store, None).await);
}
assert_eq!(ids.len(), 64, "task IDs collided");
for id in &ids {
assert_eq!(id.len(), 32, "expected 128 bits of hex: {id}");
assert!(id.chars().all(|c| c.is_ascii_hexdigit()), "{id}");
assert!(!id.starts_with("task-"), "sequential-looking ID: {id}");
}
let leading: BTreeSet<&str> = ids.iter().map(|id| &id[..2]).collect();
assert!(
leading.len() > 32,
"only {} distinct leading bytes across 64 IDs",
leading.len()
);
}
#[tokio::test]
async fn ttl_runs_from_creation_and_expired_tasks_read_as_absent() {
let store = MemoryTaskStore::new();
let id = working_task(&store, Some(0)).await;
tokio::time::sleep(Duration::from_millis(5)).await;
assert!(store.get_task(&id).await.unwrap().is_none());
assert!(store.get_task_result(&id).await.unwrap().is_none());
assert!(store.list_tasks(None).await.unwrap().is_empty());
assert!(
store
.outstanding_input_requests(&id)
.await
.unwrap()
.is_none()
);
assert!(store.cancel_task(&id, None).await.unwrap().is_none());
assert!(!store.set_ttl(&id, 60_000).await.unwrap());
assert!(
!store
.complete_task(&id, CallToolResult::text("late"))
.await
.unwrap()
);
}
#[tokio::test]
async fn ttl_is_mutable_over_the_task_lifetime() {
let store = MemoryTaskStore::new();
let id = working_task(&store, Some(60_000)).await;
assert!(store.set_ttl(&id, 120_000).await.unwrap());
let task = store.get_task(&id).await.unwrap().unwrap();
assert_eq!(task.ttl, Some(120_000));
assert!(store.set_ttl(&id, 0).await.unwrap());
tokio::time::sleep(Duration::from_millis(5)).await;
assert!(store.get_task(&id).await.unwrap().is_none());
}
#[tokio::test]
async fn require_input_records_requests_and_exposes_them() {
let store = MemoryTaskStore::new();
let id = working_task(&store, None).await;
assert!(
store
.require_input(&id, requests(&["approval", "region"]), Some("need input"))
.await
.unwrap()
);
let task = store.get_task(&id).await.unwrap().unwrap();
assert_eq!(task.status, TaskStatus::InputRequired);
assert_eq!(task.status_message.as_deref(), Some("need input"));
let outstanding = store
.outstanding_input_requests(&id)
.await
.unwrap()
.unwrap();
assert_eq!(
outstanding.keys().collect::<Vec<_>>(),
vec!["approval", "region"],
"every outstanding request must be exposed, not just the newest"
);
}
#[tokio::test]
async fn partial_input_responses_leave_the_rest_outstanding() {
let store = MemoryTaskStore::new();
let id = working_task(&store, None).await;
store
.require_input(&id, requests(&["approval", "region"]), None)
.await
.unwrap();
let applied = store
.apply_input_responses(&id, [accept("approval")].into_iter().collect())
.await
.unwrap()
.unwrap();
assert_eq!(applied.accepted, ["approval".to_string()].into());
assert!(applied.ignored.is_empty());
assert_eq!(applied.still_outstanding, ["region".to_string()].into());
assert!(!applied.is_complete());
let task = store.get_task(&id).await.unwrap().unwrap();
assert_eq!(task.status, TaskStatus::InputRequired);
assert_eq!(
store
.outstanding_input_requests(&id)
.await
.unwrap()
.unwrap()
.keys()
.collect::<Vec<_>>(),
vec!["region"]
);
let applied = store
.apply_input_responses(&id, [accept("region")].into_iter().collect())
.await
.unwrap()
.unwrap();
assert!(applied.is_complete());
assert_eq!(
store.get_task(&id).await.unwrap().unwrap().status,
TaskStatus::Working
);
}
#[tokio::test]
async fn unknown_answered_and_superseded_response_keys_are_ignored() {
let store = MemoryTaskStore::new();
let id = working_task(&store, None).await;
store
.require_input(&id, requests(&["approval", "stale"]), None)
.await
.unwrap();
store
.apply_input_responses(&id, [accept("approval")].into_iter().collect())
.await
.unwrap()
.unwrap();
store
.require_input(&id, requests(&["region"]), None)
.await
.unwrap();
let applied = store
.apply_input_responses(
&id,
[accept("never-issued"), accept("approval"), accept("stale")]
.into_iter()
.collect(),
)
.await
.unwrap()
.unwrap();
assert!(
applied.accepted.is_empty(),
"none of these keys are outstanding"
);
assert_eq!(
applied.ignored,
[
"never-issued".to_string(),
"approval".to_string(),
"stale".to_string()
]
.into(),
"unknown, already-answered, and superseded keys are all ignored"
);
assert_eq!(applied.still_outstanding, ["region".to_string()].into());
assert_eq!(
store.get_task(&id).await.unwrap().unwrap().status,
TaskStatus::InputRequired,
"ignoring a stale update must not resume or fail the task"
);
}
#[tokio::test]
async fn an_answered_key_cannot_be_reissued() {
let store = MemoryTaskStore::new();
let id = working_task(&store, None).await;
store
.require_input(&id, requests(&["approval"]), None)
.await
.unwrap();
store
.apply_input_responses(&id, [accept("approval")].into_iter().collect())
.await
.unwrap()
.unwrap();
let error = store
.require_input(&id, requests(&["approval"]), None)
.await
.expect_err("an answered key must not be reissued");
assert!(
error.to_string().contains("approval"),
"the message must name the offending key: {error}"
);
}
#[tokio::test]
async fn an_outstanding_request_can_be_carried_forward() {
let store = MemoryTaskStore::new();
let id = working_task(&store, None).await;
store
.require_input(&id, requests(&["approval"]), None)
.await
.unwrap();
assert!(
store
.require_input(&id, requests(&["approval", "region"]), None)
.await
.unwrap()
);
let outstanding = store
.outstanding_input_requests(&id)
.await
.unwrap()
.unwrap();
assert!(outstanding.contains_key("approval"));
assert!(outstanding.contains_key("region"));
let applied = store
.apply_input_responses(
&id,
[accept("approval"), accept("region")].into_iter().collect(),
)
.await
.unwrap()
.unwrap();
assert_eq!(
applied.accepted,
["approval".to_string(), "region".to_string()].into()
);
assert!(applied.is_complete());
}
#[tokio::test]
async fn an_outstanding_key_cannot_change_what_it_asks() {
use crate::protocol::{ElicitFormParams, ElicitFormSchema, ElicitRequestParams};
let store = MemoryTaskStore::new();
let id = working_task(&store, None).await;
store
.require_input(&id, requests(&["approval"]), None)
.await
.unwrap();
let mut changed: InputRequests = Default::default();
changed.insert(
"approval".to_string(),
InputRequest::Elicit(ElicitRequestParams::Form(ElicitFormParams {
mode: None,
message: "a different question".to_string(),
requested_schema: ElicitFormSchema::new(),
meta: None,
})),
);
store
.require_input(&id, changed, None)
.await
.expect_err("a live key must not be repointed at another request");
}
#[tokio::test]
async fn a_superseded_key_cannot_be_reissued() {
let store = MemoryTaskStore::new();
let id = working_task(&store, None).await;
store
.require_input(&id, requests(&["approval"]), None)
.await
.unwrap();
store
.require_input(&id, requests(&["region"]), None)
.await
.unwrap();
store
.require_input(&id, requests(&["approval"]), None)
.await
.expect_err("a superseded key must not be reissued");
}
#[tokio::test]
async fn distinct_keys_across_rounds_are_fine() {
let store = MemoryTaskStore::new();
let id = working_task(&store, None).await;
store
.require_input(&id, requests(&["approval"]), None)
.await
.unwrap();
store
.apply_input_responses(&id, [accept("approval")].into_iter().collect())
.await
.unwrap()
.unwrap();
assert!(
store
.require_input(&id, requests(&["approval_2"]), None)
.await
.unwrap()
);
let applied = store
.apply_input_responses(&id, [accept("approval_2")].into_iter().collect())
.await
.unwrap()
.unwrap();
assert_eq!(applied.accepted, ["approval_2".to_string()].into());
assert!(applied.is_complete());
}
#[tokio::test]
async fn failed_tasks_preserve_the_structured_error() {
let store = MemoryTaskStore::new();
let id = working_task(&store, None).await;
let mut error = JsonRpcError::invalid_params("bad region");
error.data = Some(serde_json::json!({"field": "region"}));
assert!(store.fail_task(&id, error).await.unwrap());
let (_, result, error) = store.get_task_result(&id).await.unwrap().unwrap();
assert!(result.is_none());
let error = error.expect("structured error must survive the store");
assert_eq!(
error.code, -32602,
"the original code must not be flattened"
);
assert_eq!(error.message, "bad region");
assert_eq!(error.data.unwrap()["field"], "region");
}
#[tokio::test]
async fn tool_error_results_complete_the_task() {
let store = MemoryTaskStore::new();
let id = working_task(&store, None).await;
let mut result = CallToolResult::text("domain failure");
result.is_error = true;
assert!(store.complete_task(&id, result).await.unwrap());
let (task, result, error) = store.get_task_result(&id).await.unwrap().unwrap();
assert_eq!(
task.status,
TaskStatus::Completed,
"isError is a domain error, not an execution failure"
);
assert!(result.unwrap().is_error);
assert!(error.is_none(), "no JSON-RPC error accompanies isError");
}
#[tokio::test]
async fn tasks_record_their_creating_principal() {
let store = MemoryTaskStore::new();
let (owned, _) = store
.create_task("tool", serde_json::json!({}), None, Some("alice".into()))
.await
.unwrap();
let (unowned, _) = store
.create_task("tool", serde_json::json!({}), None, None)
.await
.unwrap();
assert_eq!(
store.task_owner(&owned).await.unwrap(),
Some(Some("alice".to_string()))
);
assert_eq!(store.task_owner(&unowned).await.unwrap(), Some(None));
assert_eq!(
store.task_owner("does-not-exist").await.unwrap(),
None,
"an unknown task has no owner record at all"
);
let wire = serde_json::to_value(store.get_task(&owned).await.unwrap().unwrap()).unwrap();
assert!(
wire.get("owner").is_none(),
"owner leaked to the wire: {wire}"
);
assert!(!wire.to_string().contains("alice"));
}
#[test]
fn owner_matching_is_equality_not_leniency() {
assert!(owner_matches(&None, None), "no auth configured");
assert!(owner_matches(&Some("alice".into()), Some("alice")));
assert!(
!owner_matches(&Some("alice".into()), Some("bob")),
"a different principal must not inherit the task"
);
assert!(
!owner_matches(&Some("alice".into()), None),
"dropping the token must not grant access"
);
assert!(
!owner_matches(&None, Some("alice")),
"an unowned task belongs to a different security context"
);
}
#[tokio::test]
async fn terminal_states_clear_outstanding_requests() {
for (label, terminate) in [("completed", true), ("cancelled", false)] {
let store = MemoryTaskStore::new();
let id = working_task(&store, None).await;
store
.require_input(&id, requests(&["approval"]), None)
.await
.unwrap();
if terminate {
store
.complete_task(&id, CallToolResult::text("done"))
.await
.unwrap();
} else {
store.cancel_task(&id, None).await.unwrap();
}
assert!(
store
.outstanding_input_requests(&id)
.await
.unwrap()
.unwrap()
.is_empty(),
"{label} task still advertises outstanding input requests"
);
assert!(
store
.apply_input_responses(&id, [accept("approval")].into_iter().collect())
.await
.unwrap()
.is_none(),
"{label} task accepted a late input response"
);
}
}
fn store_with_limits(limits: TaskRetentionLimits) -> MemoryTaskStore {
MemoryTaskStore::with_retention_limits(limits)
}
fn assert_accounting(store: &MemoryTaskStore) {
let data = store.state.data.read().unwrap();
let retained_bytes = data
.tasks
.values()
.map(|task| {
let measured = retained_payload_size(&task.task, usize::MAX).unwrap();
assert_eq!(task.retained_bytes, measured);
measured
})
.sum::<usize>();
let reserved_bytes = data
.tasks
.values()
.map(|task| task.reserved_bytes)
.sum::<usize>();
assert_eq!(data.retained_bytes, retained_bytes);
assert_eq!(data.reserved_bytes, reserved_bytes);
let expected = data.usage();
drop(data);
assert_eq!(store.usage(), expected);
assert!(store.usage().charged_bytes() <= store.state.retention_limits.max_retained_bytes);
}
fn large_request(key: &str, message: String) -> InputRequests {
[(
key.to_string(),
InputRequest::Elicit(ElicitRequestParams::Form(ElicitFormParams {
mode: None,
message,
requested_schema: ElicitFormSchema::new(),
meta: None,
})),
)]
.into_iter()
.collect()
}
fn accept_text(key: &str, value: String) -> (String, InputResponse) {
(
key.to_string(),
InputResponse::Elicit(ElicitResult::accept(std::collections::HashMap::from([(
"value".to_string(),
ElicitFieldValue::String(value),
)]))),
)
}
#[test]
fn retention_policy_defaults_are_finite_and_unbounded_is_explicit() {
assert_eq!(TaskRetentionLimits::new(), TaskRetentionLimits::default());
let limits = TaskRetentionLimits::default();
assert_eq!(limits.max_tasks, 1_024);
assert_eq!(limits.max_payload_bytes, 4 * 1024 * 1024);
assert_eq!(limits.max_retained_bytes, 64 * 1024 * 1024);
let unbounded = TaskRetentionLimits::unbounded();
assert_eq!(unbounded.max_tasks, usize::MAX);
assert_eq!(unbounded.max_payload_bytes, usize::MAX);
assert_eq!(unbounded.max_retained_bytes, usize::MAX);
let synthetic_overflow = TaskStoreUsage {
retained_bytes: usize::MAX,
reserved_bytes: 1,
..TaskStoreUsage::default()
};
assert_eq!(synthetic_overflow.charged_bytes(), usize::MAX);
let config = MemoryTaskStoreConfig {
default_ttl: Duration::from_secs(30),
cleanup_interval: Duration::from_secs(10),
};
let default_limited = MemoryTaskStore::with_config(config);
assert_eq!(
default_limited.state.retention_limits,
TaskRetentionLimits::default()
);
let explicitly_unbounded =
MemoryTaskStore::with_config_and_retention(config, TaskRetentionLimits::unbounded());
assert_eq!(
explicitly_unbounded.state.retention_limits,
TaskRetentionLimits::unbounded()
);
}
#[test]
fn payload_validation_uses_one_stricter_counting_pass() {
struct Counted<'a>(&'a std::sync::atomic::AtomicUsize);
impl serde::Serialize for Counted<'_> {
fn serialize<S>(&self, serializer: S) -> std::result::Result<S::Ok, S::Error>
where
S: serde::Serializer,
{
self.0.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
serializer.serialize_str("payload")
}
}
let calls = std::sync::atomic::AtomicUsize::new(0);
let limits = TaskRetentionLimits::unbounded().max_retained_bytes(64);
assert_eq!(validate_payload(&Counted(&calls), limits).unwrap(), 9);
assert_eq!(calls.load(std::sync::atomic::Ordering::Relaxed), 1);
let tied = TaskRetentionLimits::unbounded()
.max_payload_bytes(1)
.max_retained_bytes(1);
assert!(matches!(
validate_payload("x", tied),
Err(TaskStoreError::RetentionLimitExceeded {
kind: TaskRetentionLimitKind::PayloadBytes,
limit: 1
})
));
let aggregate_is_stricter = TaskRetentionLimits::unbounded().max_retained_bytes(4);
assert!(matches!(
validate_prefixed_string("x", "prefix: ", aggregate_is_stricter),
Err(TaskStoreError::RetentionLimitExceeded {
kind: TaskRetentionLimitKind::AggregateBytes,
limit: 4
})
));
}
#[tokio::test]
async fn payload_limit_accepts_exact_boundary_and_rejects_without_mutation() {
let arguments = serde_json::json!({"data": "abcd"});
let exact = serde_json::to_vec(&arguments).unwrap().len();
let store = store_with_limits(TaskRetentionLimits::unbounded().max_payload_bytes(exact));
store
.create_task("tool", arguments.clone(), None, None)
.await
.expect("the exact encoded-byte boundary is inclusive");
let before = store.usage();
let error = store
.create_task("tool", serde_json::json!({"data": "abcde"}), None, None)
.await
.unwrap_err();
assert!(matches!(
error,
TaskStoreError::RetentionLimitExceeded {
kind: TaskRetentionLimitKind::PayloadBytes,
limit
} if limit == exact
));
assert_eq!(store.usage(), before);
assert_accounting(&store);
let zero = store_with_limits(TaskRetentionLimits::unbounded().max_payload_bytes(0));
assert!(matches!(
zero.create_task("tool", serde_json::Value::Null, None, None)
.await,
Err(TaskStoreError::RetentionLimitExceeded {
kind: TaskRetentionLimitKind::PayloadBytes,
limit: 0
})
));
}
#[tokio::test]
async fn oversized_owner_rejects_create_atomically() {
let store = store_with_limits(
TaskRetentionLimits::unbounded()
.max_payload_bytes(64)
.max_retained_bytes(4 * 1024),
);
let error = store
.create_task(
"tool",
serde_json::Value::Null,
None,
Some("owner-secret".repeat(128)),
)
.await
.unwrap_err();
assert!(matches!(
error,
TaskStoreError::RetentionLimitExceeded {
kind: TaskRetentionLimitKind::PayloadBytes,
limit: 64
}
));
assert_eq!(store.usage(), TaskStoreUsage::default());
assert_accounting(&store);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn concurrent_task_count_admission_never_overshoots() {
let store = store_with_limits(TaskRetentionLimits::unbounded().max_tasks(1));
let barrier = Arc::new(tokio::sync::Barrier::new(17));
let mut creates = Vec::new();
for _ in 0..16 {
let store = store.clone();
let barrier = barrier.clone();
creates.push(tokio::spawn(async move {
barrier.wait().await;
store
.create_task("tool", serde_json::json!({}), None, None)
.await
}));
}
barrier.wait().await;
let mut accepted = 0;
for create in creates {
match create.await.unwrap() {
Ok(_) => accepted += 1,
Err(TaskStoreError::RetentionLimitExceeded {
kind: TaskRetentionLimitKind::TaskCount,
limit: 1,
}) => {}
Err(other) => panic!("unexpected create error: {other}"),
}
}
assert_eq!(accepted, 1);
assert_eq!(store.usage().task_count, 1);
assert_accounting(&store);
}
#[tokio::test]
async fn replacement_and_input_accumulation_account_exactly() {
let store = store_with_limits(TaskRetentionLimits::unbounded());
let id = working_task(&store, None).await;
let initial = store.usage();
assert!(
store
.set_task_meta(&id, serde_json::json!({"note": "x".repeat(2_000)}))
.await
.unwrap()
);
let large_meta = store.usage();
assert!(large_meta.retained_bytes > initial.retained_bytes);
assert_accounting(&store);
assert!(
store
.set_task_meta(&id, serde_json::json!({"note": "small"}))
.await
.unwrap()
);
let small_meta = store.usage();
assert!(small_meta.retained_bytes < large_meta.retained_bytes);
store
.require_input(&id, large_request("first", "question".repeat(20)), None)
.await
.unwrap();
let requested = store.usage();
assert!(requested.retained_bytes > small_meta.retained_bytes);
store
.apply_input_responses(
&id,
[accept_text("first", "answer".repeat(30))]
.into_iter()
.collect(),
)
.await
.unwrap()
.unwrap();
let first_answer = store.usage();
assert_accounting(&store);
store
.require_input(&id, large_request("second", "next".repeat(20)), None)
.await
.unwrap();
store
.apply_input_responses(
&id,
[accept_text("second", "more".repeat(40))]
.into_iter()
.collect(),
)
.await
.unwrap()
.unwrap();
assert!(store.usage().retained_bytes > first_answer.retained_bytes);
assert_accounting(&store);
}
#[tokio::test]
async fn rejected_payload_mutations_are_atomic_and_ignored_input_is_not_charged() {
let limits = TaskRetentionLimits::unbounded().max_payload_bytes(256);
let store = store_with_limits(limits);
let id = working_task(&store, None).await;
let before_meta = store.usage();
assert!(matches!(
store
.set_task_meta(&id, serde_json::json!({"secret": "m".repeat(1_024)}))
.await,
Err(TaskStoreError::RetentionLimitExceeded {
kind: TaskRetentionLimitKind::PayloadBytes,
..
})
));
assert_eq!(store.usage(), before_meta);
assert!(matches!(
store
.require_input(&id, large_request("large", "q".repeat(1_024)), None)
.await,
Err(TaskStoreError::RetentionLimitExceeded {
kind: TaskRetentionLimitKind::PayloadBytes,
..
})
));
assert_eq!(
store.get_task(&id).await.unwrap().unwrap().status,
TaskStatus::Working
);
store
.require_input(&id, requests(&["approval"]), None)
.await
.unwrap();
let accepted = store
.apply_input_responses(
&id,
[
accept("approval"),
accept_text("ignored", "never-retained".repeat(1_024)),
]
.into_iter()
.collect(),
)
.await
.unwrap()
.unwrap();
assert_eq!(accepted.accepted, ["approval".to_string()].into());
assert_eq!(accepted.ignored, ["ignored".to_string()].into());
assert_accounting(&store);
let other = working_task(&store, None).await;
store
.require_input(&other, requests(&["approval"]), None)
.await
.unwrap();
let before_response = store.usage();
assert!(matches!(
store
.apply_input_responses(
&other,
[accept_text("approval", "secret".repeat(1_024))]
.into_iter()
.collect(),
)
.await,
Err(TaskStoreError::RetentionLimitExceeded {
kind: TaskRetentionLimitKind::PayloadBytes,
..
})
));
assert_eq!(store.usage(), before_response);
assert_eq!(
store
.outstanding_input_requests(&other)
.await
.unwrap()
.unwrap()
.len(),
1
);
assert_accounting(&store);
}
#[tokio::test]
async fn oversized_terminal_result_records_bounded_failure_and_wakes_waiter() {
let store = store_with_limits(
TaskRetentionLimits::unbounded()
.max_payload_bytes(256)
.max_retained_bytes(8 * 1024),
);
let id = working_task(&store, None).await;
let waiting_store = store.clone();
let waiting_id = id.clone();
let mut waiter = tokio::spawn(async move {
waiting_store
.wait_for_completion(&waiting_id)
.await
.unwrap()
.unwrap()
});
assert!(
tokio::time::timeout(Duration::from_millis(20), &mut waiter)
.await
.is_err()
);
let secret = "terminal-secret".repeat(1_024);
let error = store
.complete_task(&id, CallToolResult::text(&secret))
.await
.unwrap_err();
assert!(matches!(
error,
TaskStoreError::RetentionLimitExceeded {
kind: TaskRetentionLimitKind::PayloadBytes,
limit: 256
}
));
let (task, result, error) = tokio::time::timeout(Duration::from_secs(2), waiter)
.await
.expect("bounded failure did not wake completion waiter")
.unwrap();
assert_eq!(task.status, TaskStatus::Failed);
assert_eq!(
task.status_message.as_deref(),
Some(RETENTION_FAILURE_STATUS)
);
assert!(result.is_none());
assert_eq!(error.unwrap().message, RETENTION_FAILURE_MESSAGE);
let data = store.state.data.read().unwrap();
let stored = data.tasks.get(&id).unwrap();
assert_eq!(stored.reserved_bytes, 0);
assert!(!format!("{:?}", stored.task).contains("terminal-secret"));
drop(data);
assert_accounting(&store);
}
#[tokio::test]
async fn oversized_failure_and_cancel_reason_store_only_bounded_terminal_state() {
let store = store_with_limits(
TaskRetentionLimits::unbounded()
.max_payload_bytes(128)
.max_retained_bytes(8 * 1024),
);
let failed = working_task(&store, None).await;
let secret = "diagnostic-secret".repeat(1_024);
assert!(matches!(
store
.fail_task(&failed, JsonRpcError::internal_error(&secret))
.await,
Err(TaskStoreError::RetentionLimitExceeded { .. })
));
let (task, _, error) = store.get_task_result(&failed).await.unwrap().unwrap();
assert_eq!(task.status, TaskStatus::Failed);
assert_eq!(error.unwrap().message, RETENTION_FAILURE_MESSAGE);
let cancelled = working_task(&store, None).await;
let object = store
.cancel_task(&cancelled, Some(&secret))
.await
.unwrap()
.unwrap();
assert_eq!(object.status, TaskStatus::Cancelled);
assert_eq!(
object.status_message.as_deref(),
Some("Task cancelled: retention limit exceeded")
);
let data = store.state.data.read().unwrap();
assert!(
!format!("{:?}", data.tasks.get(&failed).unwrap().task).contains("diagnostic-secret")
);
assert!(
!format!("{:?}", data.tasks.get(&cancelled).unwrap().task)
.contains("diagnostic-secret")
);
drop(data);
assert_accounting(&store);
}
fn working_charge(tool: &str, arguments: serde_json::Value) -> usize {
let limits = TaskRetentionLimits::unbounded();
let task = Task::new(
"prototype".to_string(),
tool.to_string(),
arguments,
60_000,
None,
);
let stored = prepare_stored_task(task, limits).unwrap();
stored.retained_bytes + stored.reserved_bytes
}
#[tokio::test]
async fn aggregate_capacity_recovers_after_deletion_and_expiry() {
let charge = working_charge("tool", serde_json::json!({}));
let limits = TaskRetentionLimits::unbounded()
.max_tasks(2)
.max_retained_bytes(charge);
let store = store_with_limits(limits);
let first = working_task(&store, None).await;
assert_eq!(store.usage().charged_bytes(), charge);
assert!(matches!(
store
.create_task("tool", serde_json::json!({}), None, None)
.await,
Err(TaskStoreError::RetentionLimitExceeded {
kind: TaskRetentionLimitKind::AggregateBytes,
limit
}) if limit == charge
));
assert_accounting(&store);
assert!(store.discard_task(&first).await.unwrap());
assert_eq!(store.usage(), TaskStoreUsage::default());
let expired = working_task(&store, None).await;
let before_expiry = store.usage();
assert!(matches!(
store
.create_task("tool", serde_json::json!({}), None, None)
.await,
Err(TaskStoreError::RetentionLimitExceeded {
kind: TaskRetentionLimitKind::AggregateBytes,
limit
}) if limit == charge
));
assert_accounting(&store);
assert!(store.set_ttl(&expired, 0).await.unwrap());
let tombstone = store.usage();
assert_eq!(tombstone.task_count, 1);
assert_eq!(tombstone.reserved_bytes, 0);
assert!(tombstone.retained_bytes < before_expiry.retained_bytes);
assert!(matches!(
store.task_presence(&expired).await.unwrap(),
TaskPresence::Expired { .. }
));
assert_accounting(&store);
let replacement = working_task(&store, None).await;
assert_ne!(replacement, expired);
assert!(matches!(
store.task_presence(&expired).await.unwrap(),
TaskPresence::Missing
));
assert_eq!(store.usage().task_count, 1);
assert_eq!(store.usage().charged_bytes(), charge);
assert_accounting(&store);
}
#[tokio::test]
async fn expiry_accounting_allows_scrubbed_encoding_to_grow() {
let store = MemoryTaskStore::with_config_and_retention(
MemoryTaskStoreConfig::default().cleanup_interval(Duration::from_secs(60)),
TaskRetentionLimits::unbounded(),
);
let (id, _) = store
.create_task("x", serde_json::Value::Null, None, None)
.await
.unwrap();
assert!(
store
.set_status(&id, TaskStatus::Working, Some(""))
.await
.unwrap()
);
let before_expiry = store.usage();
assert!(store.set_ttl(&id, 0).await.unwrap());
let after_expiry = store.usage();
assert_eq!(after_expiry.task_count, 1);
assert_eq!(after_expiry.reserved_bytes, 0);
assert_eq!(
after_expiry.retained_bytes,
before_expiry.retained_bytes + 1,
"the empty status encodes two bytes smaller than None, while the one-byte tool name is scrubbed"
);
assert!(matches!(
store.task_presence(&id).await.unwrap(),
TaskPresence::Expired { .. }
));
assert_accounting(&store);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn concurrent_completions_atomically_enforce_aggregate_limit() {
let result_text = "r".repeat(2_048);
let unbounded = TaskRetentionLimits::unbounded();
let working = prepare_stored_task(
Task::new(
"prototype".to_string(),
"tool".to_string(),
serde_json::json!({}),
60_000,
None,
),
unbounded,
)
.unwrap();
let mut completed = Task::new(
"prototype".to_string(),
"tool".to_string(),
serde_json::json!({}),
60_000,
None,
);
completed.status = TaskStatus::Completed;
completed.status_message = Some("Task completed".to_string());
completed.result = Some(CallToolResult::text(&result_text));
completed.completed_at = Some(Instant::now());
let completed = prepare_stored_task(completed, unbounded).unwrap();
let working_charge = working.retained_bytes + working.reserved_bytes;
assert!(completed.retained_bytes > working_charge);
let cap = completed.retained_bytes + working_charge;
let store = store_with_limits(
TaskRetentionLimits::unbounded()
.max_tasks(2)
.max_payload_bytes(4 * 1024)
.max_retained_bytes(cap),
);
let first = working_task(&store, None).await;
let second = working_task(&store, None).await;
let barrier = Arc::new(tokio::sync::Barrier::new(3));
let mut completions = Vec::new();
for id in [first.clone(), second.clone()] {
let store = store.clone();
let barrier = barrier.clone();
let result_text = result_text.clone();
completions.push(tokio::spawn(async move {
barrier.wait().await;
store
.complete_task(&id, CallToolResult::text(result_text))
.await
}));
}
barrier.wait().await;
let mut completed_count = 0;
let mut rejected_count = 0;
for completion in completions {
match completion.await.unwrap() {
Ok(true) => completed_count += 1,
Err(TaskStoreError::RetentionLimitExceeded {
kind: TaskRetentionLimitKind::AggregateBytes,
limit,
}) if limit == cap => rejected_count += 1,
other => panic!("unexpected completion outcome: {other:?}"),
}
}
assert_eq!((completed_count, rejected_count), (1, 1));
let states = [
store.get_task(&first).await.unwrap().unwrap().status,
store.get_task(&second).await.unwrap().unwrap().status,
];
assert_eq!(
states
.iter()
.filter(|status| **status == TaskStatus::Completed)
.count(),
1
);
assert_eq!(
states
.iter()
.filter(|status| **status == TaskStatus::Failed)
.count(),
1
);
assert_accounting(&store);
}
}