use std::fmt;
use std::future::{Future, poll_fn};
use std::sync::Arc;
use std::sync::atomic::{AtomicU8, Ordering};
use std::task::Poll;
use asupersync::Cx;
use asupersync::types::Time;
use fastmcp_core::McpRequestCancellation;
use fastmcp_protocol::RequestId;
use fastmcp_protocol::tasks_extension::{Task, TaskId};
use super::recovery::{
ClientCredentialsTaskRecoveryError, ClientCredentialsTaskRecoveryPolicy, RecoveryState,
};
use super::{
ClientCredentialsError, ClientCredentialsSnapshot, ClientCredentialsTaskWatch,
ClientCredentialsTaskWatchError, ClientCredentialsTaskWatchPolicy,
ClientCredentialsTasksClient, ManagedTaskEvent, ManagedTaskRequest, ManagedTaskSnapshot,
WatchIds, WatchState, active, check_watch, copy_binding, discovery_deadline, prepare_pinned,
request_pinned,
};
pub use crate::http_auth::managed::tasks::watch::cancellation::TaskCancellationState;
const LIVE: u8 = 0;
const CANCEL_REQUESTED: u8 = 1;
const TERMINAL_DELIVERED: u8 = 2;
const CLOSED: u8 = 3;
const READY: u8 = 0;
const UNCONFIRMED: u8 = 1;
const ACKNOWLEDGED: u8 = 2;
#[derive(Debug)]
pub enum ClientCredentialsTaskCancellationError {
Closed,
AlreadyAttempted,
NotAttempted(ClientCredentialsTaskWatchError),
Unconfirmed(ClientCredentialsTaskWatchError),
}
impl fmt::Display for ClientCredentialsTaskCancellationError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Closed => f.write_str("machine Task cancellation owner is closed"),
Self::AlreadyAttempted => f.write_str(
"machine Task cancellation was already attempted; it will not be replayed",
),
Self::NotAttempted(error) => {
write!(f, "machine Task cancellation did not start: {error}")
}
Self::Unconfirmed(error) => write!(
f,
"machine Task cancellation acknowledgement is unknown: {error}"
),
}
}
}
impl std::error::Error for ClientCredentialsTaskCancellationError {}
#[derive(Debug)]
pub enum CancellableClientCredentialsTaskWatchError {
CancellationRequested,
Closed,
Watch(ClientCredentialsTaskWatchError),
Recovery(ClientCredentialsTaskRecoveryError),
}
impl fmt::Display for CancellableClientCredentialsTaskWatchError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::CancellationRequested => {
f.write_str("machine Task cancellation acknowledged; observation stopped")
}
Self::Closed => f.write_str("cancellable machine Task watch is closed"),
Self::Watch(error) => error.fmt(f),
Self::Recovery(error) => error.fmt(f),
}
}
}
impl std::error::Error for CancellableClientCredentialsTaskWatchError {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
match self {
Self::Watch(error) => Some(error),
Self::Recovery(error) => Some(error),
Self::CancellationRequested | Self::Closed => None,
}
}
}
impl From<ClientCredentialsTaskWatchError> for CancellableClientCredentialsTaskWatchError {
fn from(error: ClientCredentialsTaskWatchError) -> Self {
Self::Watch(error)
}
}
impl From<ClientCredentialsError> for CancellableClientCredentialsTaskWatchError {
fn from(error: ClientCredentialsError) -> Self {
Self::Watch(error.into())
}
}
impl From<ClientCredentialsTaskRecoveryError> for CancellableClientCredentialsTaskWatchError {
fn from(error: ClientCredentialsTaskRecoveryError) -> Self {
match error {
ClientCredentialsTaskRecoveryError::Watch(error) => Self::Watch(error),
error => Self::Recovery(error),
}
}
}
struct Control {
completion: AtomicU8,
attempt: AtomicU8,
acknowledged: McpRequestCancellation,
closed: McpRequestCancellation,
}
impl Control {
fn new() -> Self {
Self {
completion: AtomicU8::new(LIVE),
attempt: AtomicU8::new(READY),
acknowledged: McpRequestCancellation::new(),
closed: McpRequestCancellation::new(),
}
}
fn claim(&self) -> Result<(), ClientCredentialsTaskCancellationError> {
if self.completion.load(Ordering::Acquire) != LIVE {
return Err(ClientCredentialsTaskCancellationError::Closed);
}
self.attempt
.compare_exchange(READY, UNCONFIRMED, Ordering::AcqRel, Ordering::Acquire)
.map_err(|_| ClientCredentialsTaskCancellationError::AlreadyAttempted)?;
Ok(())
}
fn state(&self) -> TaskCancellationState {
match self.attempt.load(Ordering::Acquire) {
READY => TaskCancellationState::Ready,
UNCONFIRMED => TaskCancellationState::Unconfirmed,
ACKNOWLEDGED => TaskCancellationState::Acknowledged,
_ => unreachable!("private cancellation state"),
}
}
fn acknowledge(&self) {
self.attempt.store(ACKNOWLEDGED, Ordering::Release);
if self
.completion
.compare_exchange(LIVE, CANCEL_REQUESTED, Ordering::AcqRel, Ordering::Acquire)
.is_ok()
{
self.acknowledged.cancel();
}
}
fn terminal(&self) -> Result<(), CancellableClientCredentialsTaskWatchError> {
self.completion
.compare_exchange(
LIVE,
TERMINAL_DELIVERED,
Ordering::AcqRel,
Ordering::Acquire,
)
.map_err(|state| {
if state == CANCEL_REQUESTED {
CancellableClientCredentialsTaskWatchError::CancellationRequested
} else {
CancellableClientCredentialsTaskWatchError::Closed
}
})?;
self.closed.cancel();
Ok(())
}
fn close(&self) {
let _ = self
.completion
.compare_exchange(LIVE, CLOSED, Ordering::AcqRel, Ordering::Acquire);
self.closed.cancel();
}
}
struct CancelScope {
control: Arc<Control>,
client: ClientCredentialsTasksClient,
task_id: TaskId,
ids: (RequestId, RequestId),
cancellation: McpRequestCancellation,
deadline: Time,
binding: Option<Arc<ClientCredentialsSnapshot>>,
}
#[derive(Clone)]
pub struct ClientCredentialsTaskCancelHandle {
scope: Arc<CancelScope>,
}
impl fmt::Debug for ClientCredentialsTaskCancelHandle {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("ClientCredentialsTaskCancelHandle")
.field("state", &self.state())
.finish_non_exhaustive()
}
}
impl ClientCredentialsTaskCancelHandle {
pub fn state(&self) -> TaskCancellationState {
self.scope.control.state()
}
pub(super) fn for_observation(
client: &ClientCredentialsTasksClient,
task_id: TaskId,
prefix: &str,
cancellation: &McpRequestCancellation,
deadline: Time,
) -> Result<Self, CancellableClientCredentialsTaskWatchError> {
let ids = cancellation_ids(prefix)?;
let _ = prepare_pinned(client, &ids, ManagedTaskRequest::Cancel(task_id.clone()))?;
Ok(Self {
scope: Arc::new(CancelScope {
control: Arc::new(Control::new()),
client: client.clone(),
task_id,
ids,
cancellation: cancellation.clone(),
deadline,
binding: None,
}),
})
}
pub(super) fn pin_binding(
&mut self,
binding: ClientCredentialsSnapshot,
deadline: Time,
) -> Result<(), CancellableClientCredentialsTaskWatchError> {
let scope = Arc::get_mut(&mut self.scope)
.ok_or(CancellableClientCredentialsTaskWatchError::Closed)?;
if scope.binding.is_some() {
return Err(CancellableClientCredentialsTaskWatchError::Closed);
}
scope.binding = Some(Arc::new(binding));
scope.deadline = scope.deadline.min(deadline);
Ok(())
}
pub(super) fn cancellation_requested(&self) -> bool {
self.scope.control.completion.load(Ordering::Acquire) == CANCEL_REQUESTED
}
pub(super) fn close_observation(&self) {
self.scope.control.close();
}
pub(super) fn select_terminal(&self) -> Result<(), CancellableClientCredentialsTaskWatchError> {
self.scope.control.terminal()
}
pub(super) fn read_lease(&self) -> ReadLease {
ReadLease {
control: Arc::clone(&self.scope.control),
armed: true,
}
}
pub(super) async fn until_acknowledged<T>(
&self,
work: impl Future<Output = T>,
) -> Result<T, CancellableClientCredentialsTaskWatchError> {
until_signal(&self.scope.control.acknowledged, work)
.await
.map_err(|()| CancellableClientCredentialsTaskWatchError::CancellationRequested)
}
pub async fn request_cancel(
&self,
cx: &Cx,
) -> Result<(), ClientCredentialsTaskCancellationError> {
self.request_cancel_with_cancellation(cx, &McpRequestCancellation::new())
.await
}
pub async fn request_cancel_with_cancellation(
&self,
cx: &Cx,
cancellation: &McpRequestCancellation,
) -> Result<(), ClientCredentialsTaskCancellationError> {
let scope = &self.scope;
if scope.control.state() != TaskCancellationState::Ready {
return Err(ClientCredentialsTaskCancellationError::AlreadyAttempted);
}
if scope.control.completion.load(Ordering::Acquire) != LIVE {
return Err(ClientCredentialsTaskCancellationError::Closed);
}
let binding = scope
.binding
.as_deref()
.ok_or(ClientCredentialsTaskCancellationError::Closed)?;
let owner = &scope.client.client.inner.closed;
let deadline = (|| {
check_watch(cx, scope.deadline, owner, &scope.cancellation, binding)?;
check_watch(cx, scope.deadline, owner, cancellation, binding)?;
let call_deadline = discovery_deadline(
cx,
scope
.client
.limits
.timeout
.min(scope.client.client.inner.timeout),
)
.map_err(ClientCredentialsError::from)?;
Ok::<_, ClientCredentialsTaskWatchError>(scope.deadline.min(call_deadline))
})()
.map_err(ClientCredentialsTaskCancellationError::NotAttempted)?;
scope.control.claim()?;
let attempted = Box::pin(active(
cx,
deadline,
owner,
&scope.cancellation,
Some(binding),
async {
Ok(async {
let mut call = request_pinned(
&scope.client,
cx,
cancellation,
binding,
deadline,
scope.ids.clone(),
ManagedTaskRequest::Cancel(scope.task_id.clone()),
)
.await?;
if !matches!(
call.next_event(cx).await?,
Some(ManagedTaskEvent::Cancelled(_))
) {
return Err(ClientCredentialsTaskWatchError::UnexpectedEvent);
}
scope.control.acknowledge();
Ok::<(), ClientCredentialsTaskWatchError>(())
}
.await)
},
));
let result = until_signal(&scope.control.closed, attempted).await;
if scope.control.state() == TaskCancellationState::Acknowledged {
return Ok(());
}
let error = match result {
Err(()) => ClientCredentialsTaskWatchError::Closed,
Ok(Err(error)) => error.into(),
Ok(Ok(Err(error))) => error,
Ok(Ok(Ok(()))) => ClientCredentialsTaskWatchError::UnexpectedEvent,
};
Err(ClientCredentialsTaskCancellationError::Unconfirmed(error))
}
}
#[must_use = "retain observation custody and poll snapshots or explicitly close"]
pub struct CancellableClientCredentialsTaskWatch {
remote_cancel: ClientCredentialsTaskCancelHandle,
observation: Option<ClientCredentialsTaskWatch>,
recovery: Option<RecoveryState>,
}
impl CancellableClientCredentialsTaskWatch {
pub fn cancel_handle(&self) -> ClientCredentialsTaskCancelHandle {
self.remote_cancel.clone()
}
pub fn reconnection_attempts(&self) -> usize {
self.recovery
.as_ref()
.map_or(0, RecoveryState::reconnection_attempts)
}
pub fn close(&mut self) {
self.remote_cancel.close_observation();
self.observation = None;
}
pub async fn next_snapshot(
&mut self,
cx: &Cx,
) -> Result<Option<ManagedTaskSnapshot>, CancellableClientCredentialsTaskWatchError> {
match self
.remote_cancel
.scope
.control
.completion
.load(Ordering::Acquire)
{
CANCEL_REQUESTED => {
self.close();
return Err(CancellableClientCredentialsTaskWatchError::CancellationRequested);
}
TERMINAL_DELIVERED => return Ok(None),
CLOSED => return Err(CancellableClientCredentialsTaskWatchError::Closed),
_ => {}
}
let mut observation = self
.observation
.take()
.ok_or(CancellableClientCredentialsTaskWatchError::Closed)?;
let handle = self.remote_cancel.clone();
let mut lease = handle.read_lease();
let scope = Arc::clone(&handle.scope);
let binding = scope
.binding
.as_deref()
.ok_or(CancellableClientCredentialsTaskWatchError::Closed)?;
let read = Box::pin(active(
cx,
scope.deadline,
&scope.client.client.inner.closed,
&scope.cancellation,
Some(binding),
async {
Ok(match self.recovery.as_mut() {
Some(recovery) => {
Box::pin(recovery.next_snapshot(cx, &mut observation, Some(binding)))
.await
.map_err(CancellableClientCredentialsTaskWatchError::from)
}
None => Box::pin(observation.next_snapshot(cx))
.await
.map_err(CancellableClientCredentialsTaskWatchError::from),
})
},
));
let snapshot = handle.until_acknowledged(read).await???;
if handle.cancellation_requested() {
return Err(CancellableClientCredentialsTaskWatchError::CancellationRequested);
}
let snapshot = snapshot.ok_or(ClientCredentialsTaskWatchError::UnexpectedEvent)?;
if matches!(
&*snapshot.task,
Task::Completed { .. } | Task::Failed { .. } | Task::Cancelled(_)
) {
handle.select_terminal()?;
} else {
self.observation = Some(observation);
lease.disarm();
}
Ok(Some(snapshot))
}
}
impl Drop for CancellableClientCredentialsTaskWatch {
fn drop(&mut self) {
self.remote_cancel.close_observation();
}
}
impl fmt::Debug for CancellableClientCredentialsTaskWatch {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("CancellableClientCredentialsTaskWatch")
.field("cancellation", &self.remote_cancel.state())
.field("reconnection_attempts", &self.reconnection_attempts())
.finish_non_exhaustive()
}
}
pub(super) struct ReadLease {
control: Arc<Control>,
armed: bool,
}
impl ReadLease {
pub(super) fn disarm(&mut self) {
self.armed = false;
}
}
impl Drop for ReadLease {
fn drop(&mut self) {
if self.armed {
self.control.close();
}
}
}
async fn until_signal<T>(
signal: &McpRequestCancellation,
work: impl Future<Output = T>,
) -> Result<T, ()> {
let mut stopped = std::pin::pin!(signal.cancelled());
let mut work = std::pin::pin!(work);
poll_fn(|task| {
if stopped.as_mut().poll(task).is_ready() {
return Poll::Ready(Err(()));
}
let result = work.as_mut().poll(task);
if signal.is_cancel_requested() {
return Poll::Ready(Err(()));
}
result.map(Ok)
})
.await
}
fn cancellation_ids(
prefix: &str,
) -> Result<(RequestId, RequestId), ClientCredentialsTaskWatchError> {
let _ = WatchIds::new(prefix.to_owned())?;
Ok((
RequestId::String(format!("{prefix}:cancel:discovery")),
RequestId::String(format!("{prefix}:cancel:operation")),
))
}
impl ClientCredentialsTasksClient {
pub async fn watch_task_cancellable(
&self,
cx: &Cx,
task_id: TaskId,
id_prefix: String,
policy: ClientCredentialsTaskWatchPolicy,
) -> Result<CancellableClientCredentialsTaskWatch, CancellableClientCredentialsTaskWatchError>
{
self.watch_task_cancellable_with_cancellation(
cx,
&McpRequestCancellation::new(),
task_id,
id_prefix,
policy,
)
.await
}
pub async fn watch_task_cancellable_with_cancellation(
&self,
cx: &Cx,
cancellation: &McpRequestCancellation,
task_id: TaskId,
id_prefix: String,
policy: ClientCredentialsTaskWatchPolicy,
) -> Result<CancellableClientCredentialsTaskWatch, CancellableClientCredentialsTaskWatchError>
{
self.open_cancellable_watch(cx, cancellation, task_id, id_prefix, policy, None)
.await
}
pub async fn watch_task_cancellable_recovering(
&self,
cx: &Cx,
task_id: TaskId,
id_prefix: String,
policy: ClientCredentialsTaskWatchPolicy,
recovery: ClientCredentialsTaskRecoveryPolicy,
) -> Result<CancellableClientCredentialsTaskWatch, CancellableClientCredentialsTaskWatchError>
{
self.watch_task_cancellable_recovering_with_cancellation(
cx,
&McpRequestCancellation::new(),
task_id,
id_prefix,
policy,
recovery,
)
.await
}
#[allow(clippy::too_many_arguments)]
pub async fn watch_task_cancellable_recovering_with_cancellation(
&self,
cx: &Cx,
cancellation: &McpRequestCancellation,
task_id: TaskId,
id_prefix: String,
policy: ClientCredentialsTaskWatchPolicy,
recovery: ClientCredentialsTaskRecoveryPolicy,
) -> Result<CancellableClientCredentialsTaskWatch, CancellableClientCredentialsTaskWatchError>
{
self.open_cancellable_watch(cx, cancellation, task_id, id_prefix, policy, Some(recovery))
.await
}
#[allow(clippy::too_many_arguments)]
async fn open_cancellable_watch(
&self,
cx: &Cx,
cancellation: &McpRequestCancellation,
task_id: TaskId,
id_prefix: String,
policy: ClientCredentialsTaskWatchPolicy,
recovery: Option<ClientCredentialsTaskRecoveryPolicy>,
) -> Result<CancellableClientCredentialsTaskWatch, CancellableClientCredentialsTaskWatchError>
{
let connection_policy = match recovery {
Some(recovery) => recovery.connection_policy(policy)?,
None => policy,
};
let _ = WatchState::new(vec![task_id.clone()], policy.maximum_snapshots)?;
let deadline =
discovery_deadline(cx, policy.timeout).map_err(ClientCredentialsError::from)?;
let mut handle = ClientCredentialsTaskCancelHandle::for_observation(
self,
task_id.clone(),
&id_prefix,
cancellation,
deadline,
)?;
let owner = &self.client.inner.closed;
let mut observation = Box::pin(active(cx, deadline, owner, cancellation, None, async {
Ok(self
.watch_tasks_with_cancellation(
cx,
cancellation,
vec![task_id],
id_prefix,
connection_policy,
)
.await)
}))
.await??;
observation.deadline = observation.deadline.min(deadline);
check_watch(
cx,
observation.deadline,
owner,
cancellation,
&observation.binding,
)?;
handle.pin_binding(copy_binding(&observation.binding), observation.deadline)?;
let recovery =
recovery.map(|policy| RecoveryState::new(&observation, connection_policy, policy));
Ok(CancellableClientCredentialsTaskWatch {
remote_cancel: handle,
observation: Some(observation),
recovery,
})
}
}
#[cfg(test)]
mod tests;
#[cfg(test)]
mod recovery_tests {
use super::super::tests::{consumer, runtime};
use super::super::{ClientCredentialsTasksError, OAuthDiscoveryError};
use super::*;
use std::time::Duration;
#[test]
fn cancellable_recovery_reserves_records_before_credentials_or_listen() {
runtime().block_on(async {
let cx = Cx::current().unwrap();
let client = consumer();
let policy =
ClientCredentialsTaskWatchPolicy::new(Duration::from_secs(1), 1, 2).unwrap();
let result = Box::pin(client.watch_task_cancellable_recovering(
&cx,
TaskId::parse("one").unwrap(),
"bounded".to_owned(),
policy,
ClientCredentialsTaskRecoveryPolicy::default(),
))
.await;
assert!(matches!(
result,
Err(CancellableClientCredentialsTaskWatchError::Recovery(
ClientCredentialsTaskRecoveryError::InvalidPolicy
))
));
assert!(
client
.client
.inner
.state
.try_lock_owned()
.unwrap()
.current
.is_none()
);
});
}
#[test]
fn cancellable_recovery_local_cancellation_and_closed_owner_never_acquire() {
runtime().block_on(async {
let cx = Cx::current().unwrap();
for closed in [false, true] {
let client = consumer();
let cancellation = McpRequestCancellation::new();
if closed {
client.client.close();
} else {
cancellation.cancel();
}
let result = Box::pin(client.watch_task_cancellable_recovering_with_cancellation(
&cx,
&cancellation,
TaskId::parse("one").unwrap(),
"bounded".to_owned(),
ClientCredentialsTaskWatchPolicy::default(),
ClientCredentialsTaskRecoveryPolicy::default(),
))
.await;
let Err(CancellableClientCredentialsTaskWatchError::Watch(
ClientCredentialsTaskWatchError::Task(
ClientCredentialsTasksError::Authentication(error),
),
)) = result
else {
panic!("cancelled or closed recovery must stop before admission");
};
if closed {
assert!(matches!(error, ClientCredentialsError::Closed));
} else {
assert!(matches!(
error,
ClientCredentialsError::Discovery(OAuthDiscoveryError::Cancelled)
));
}
assert!(
client
.client
.inner
.state
.try_lock_owned()
.unwrap()
.current
.is_none()
);
}
});
}
#[test]
fn cancellable_recovery_preserves_exhaustion_causes_across_input_error_conversion() {
let error = CancellableClientCredentialsTaskWatchError::from(
ClientCredentialsTaskRecoveryError::RecoveryLimit {
last_error: ClientCredentialsTaskWatchError::Interrupted,
},
);
assert!(std::error::Error::source(&error).is_some());
let error = super::super::drive::ClientCredentialsTaskWatchDriveError::from(error);
assert!(matches!(
error,
super::super::drive::ClientCredentialsTaskWatchDriveError::Recovery(
ClientCredentialsTaskRecoveryError::RecoveryLimit {
last_error: ClientCredentialsTaskWatchError::Interrupted
}
)
));
let error = CancellableClientCredentialsTaskWatchError::from(
ClientCredentialsTaskRecoveryError::Watch(
ClientCredentialsTaskWatchError::SnapshotLimit,
),
);
assert!(matches!(
error,
CancellableClientCredentialsTaskWatchError::Watch(
ClientCredentialsTaskWatchError::SnapshotLimit
)
));
}
}