use std::fmt;
use std::future::Future;
use asupersync::Cx;
use asupersync::types::Time;
use fastmcp_core::McpRequestCancellation;
use fastmcp_protocol::tasks_extension::Task;
use super::super::{
TaskResumeBinding, TaskResumeError, TaskResumeKey, TaskResumeRecord, checkpoint, wall_now,
};
use super::{admit_record, reconcile_controls};
use crate::http_auth::managed::tasks::ManagedTasksClient;
use crate::http_auth::managed::tasks::watch::cancellation::{
CancellableTaskWatchError, ManagedTaskCancelHandle,
};
use crate::http_auth::managed::tasks::watch::recovery::{
ManagedTaskRecoveryError, ManagedTaskRecoveryPolicy, RecoveringManagedTaskWatch,
};
use crate::http_auth::managed::tasks::watch::{ManagedTaskSnapshot, ManagedTaskWatchPolicy};
use crate::http_auth::managed::{OAuthSessionError, deadline_after};
#[derive(Clone)]
pub struct TaskResumeChange {
previous: TaskResumeRecord,
replacement: Option<TaskResumeRecord>,
}
impl fmt::Debug for TaskResumeChange {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("TaskResumeChange")
.field("removes_record", &self.replacement.is_none())
.finish_non_exhaustive()
}
}
impl TaskResumeChange {
pub fn from_snapshot(
cx: &Cx,
current: &TaskResumeBinding,
previous: &TaskResumeRecord,
task: &Task,
) -> Result<Self, TaskResumeError> {
previous.admit(cx, current)?;
Self::prepare_at(current, previous, task, wall_now())
}
pub fn discard(
cx: &Cx,
current: &TaskResumeBinding,
previous: &TaskResumeRecord,
) -> Result<Self, TaskResumeError> {
checkpoint(cx)?;
if previous.binding != current.digest {
return Err(TaskResumeError::Unavailable);
}
previous.validate()?;
Ok(Self {
previous: previous.clone(),
replacement: None,
})
}
fn prepare_at(
current: &TaskResumeBinding,
previous: &TaskResumeRecord,
task: &Task,
now: i128,
) -> Result<Self, TaskResumeError> {
let replacement = reconcile_controls(previous, current, task, now)?;
Ok(Self {
previous: previous.clone(),
replacement,
})
}
pub fn key(&self) -> TaskResumeKey {
self.previous.key()
}
pub fn previous(&self) -> &TaskResumeRecord {
&self.previous
}
pub fn replacement(&self) -> Option<&TaskResumeRecord> {
self.replacement.as_ref()
}
#[cfg(any(target_os = "linux", test))]
fn admit_expected(&self, actual: Option<&TaskResumeRecord>) -> Result<(), TaskResumeError> {
if actual != Some(&self.previous) {
return Err(TaskResumeError::ConflictingSnapshot);
}
Ok(())
}
#[cfg(target_os = "linux")]
pub fn apply<P: super::super::store::TaskResumeProtector>(
&self,
cx: &Cx,
current: &TaskResumeBinding,
store: &mut super::super::store::TaskResumeStore<P>,
) -> Result<(), super::super::store::TaskResumeStoreError> {
let Some(record) = &self.replacement else {
return store.remove_expected(cx, current, &self.previous);
};
self.previous.admit(cx, current)?;
let actual = store.get(cx, current, self.key())?;
self.admit_expected(actual.as_ref())?;
if record != &self.previous {
store.put(cx, current, record.clone())?;
}
Ok(())
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum TaskResumePersistenceState {
NotAttempted,
Unconfirmed,
Acknowledged,
}
pub struct PendingTaskResumeSnapshot {
snapshot: ManagedTaskSnapshot,
change: TaskResumeChange,
persistence: TaskResumePersistenceState,
}
impl PendingTaskResumeSnapshot {
pub fn snapshot(&self) -> &ManagedTaskSnapshot {
&self.snapshot
}
pub fn change(&self) -> &TaskResumeChange {
&self.change
}
pub fn persistence(&self) -> TaskResumePersistenceState {
self.persistence
}
pub fn into_parts(
self,
) -> (
ManagedTaskSnapshot,
TaskResumeChange,
TaskResumePersistenceState,
) {
(self.snapshot, self.change, self.persistence)
}
}
impl fmt::Debug for PendingTaskResumeSnapshot {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("PendingTaskResumeSnapshot")
.field("persistence", &self.persistence)
.finish_non_exhaustive()
}
}
pub enum PersistedTaskWatchError<E> {
Resume(TaskResumeError),
Recovery(ManagedTaskRecoveryError),
Session(OAuthSessionError),
Persistence(E),
CancellationRequested,
TerminalAcknowledgementRequired,
NoTerminal,
Closed,
}
impl<E> fmt::Debug for PersistedTaskWatchError<E> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Resume(error) => f.debug_tuple("Resume").field(error).finish(),
Self::Recovery(error) => f.debug_tuple("Recovery").field(error).finish(),
Self::Session(error) => f.debug_tuple("Session").field(error).finish(),
Self::Persistence(_) => f.write_str("Persistence(<host error>)"),
Self::CancellationRequested => f.write_str("CancellationRequested"),
Self::TerminalAcknowledgementRequired => f.write_str("TerminalAcknowledgementRequired"),
Self::NoTerminal => f.write_str("NoTerminal"),
Self::Closed => f.write_str("Closed"),
}
}
}
impl<E> fmt::Display for PersistedTaskWatchError<E> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Resume(error) => error.fmt(f),
Self::Recovery(error) => error.fmt(f),
Self::Session(error) => error.fmt(f),
Self::Persistence(_) => {
f.write_str("Task observation persistence was not acknowledged")
}
Self::CancellationRequested => {
f.write_str("Task cancellation acknowledged; persisted observation stopped")
}
Self::TerminalAcknowledgementRequired => {
f.write_str("acknowledge the delivered terminal before completing observation")
}
Self::NoTerminal => f.write_str("no terminal Task snapshot has been delivered"),
Self::Closed => f.write_str("persisted Task watch is closed"),
}
}
}
impl<E: std::error::Error + 'static> std::error::Error for PersistedTaskWatchError<E> {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
match self {
Self::Resume(error) => Some(error),
Self::Recovery(error) => Some(error),
Self::Session(error) => Some(error),
Self::Persistence(error) => Some(error),
Self::CancellationRequested
| Self::TerminalAcknowledgementRequired
| Self::NoTerminal
| Self::Closed => None,
}
}
}
impl<E> From<TaskResumeError> for PersistedTaskWatchError<E> {
fn from(error: TaskResumeError) -> Self {
Self::Resume(error)
}
}
impl<E> From<ManagedTaskRecoveryError> for PersistedTaskWatchError<E> {
fn from(error: ManagedTaskRecoveryError) -> Self {
Self::Recovery(error)
}
}
impl<E> From<OAuthSessionError> for PersistedTaskWatchError<E> {
fn from(error: OAuthSessionError) -> Self {
Self::Session(error)
}
}
impl<E> From<CancellableTaskWatchError> for PersistedTaskWatchError<E> {
fn from(error: CancellableTaskWatchError) -> Self {
match error {
CancellableTaskWatchError::CancellationRequested => Self::CancellationRequested,
CancellableTaskWatchError::Closed => Self::Closed,
CancellableTaskWatchError::Watch(error) => {
Self::Recovery(ManagedTaskRecoveryError::Watch(error))
}
CancellableTaskWatchError::Recovery(error) => Self::Recovery(error),
CancellableTaskWatchError::Session(error) => Self::Session(error),
}
}
}
impl ManagedTasksClient {
#[allow(clippy::too_many_arguments)]
pub async fn resume_task_watch_persisted<P, F, E>(
&self,
cx: &Cx,
current: TaskResumeBinding,
record: TaskResumeRecord,
id_prefix: String,
policy: ManagedTaskWatchPolicy,
recovery: ManagedTaskRecoveryPolicy,
persist: P,
) -> Result<PersistedManagedTaskWatch<P>, PersistedTaskWatchError<E>>
where
P: FnMut(TaskResumeChange) -> F,
F: Future<Output = Result<(), E>>,
{
self.resume_task_watch_persisted_with_cancellation(
cx,
&McpRequestCancellation::new(),
current,
record,
id_prefix,
policy,
recovery,
persist,
)
.await
}
#[allow(clippy::too_many_arguments)]
pub async fn resume_task_watch_persisted_with_cancellation<P, F, E>(
&self,
cx: &Cx,
cancellation: &McpRequestCancellation,
current: TaskResumeBinding,
record: TaskResumeRecord,
id_prefix: String,
policy: ManagedTaskWatchPolicy,
recovery: ManagedTaskRecoveryPolicy,
persist: P,
) -> Result<PersistedManagedTaskWatch<P>, PersistedTaskWatchError<E>>
where
P: FnMut(TaskResumeChange) -> F,
F: Future<Output = Result<(), E>>,
{
self.session.check(cx, cancellation)?;
let anchor = cx.now();
let now = wall_now();
admit_record(&record, ¤t, self.session.resource().as_str(), now)?;
let remaining = record
.retain_until
.checked_sub(now)
.ok_or(TaskResumeError::Unavailable)?;
let retention_deadline =
anchor.saturating_add_nanos(u64::try_from(remaining).unwrap_or(u64::MAX));
let deadline = retention_deadline.min(deadline_after(cx, policy.timeout)?);
let remote_cancel = ManagedTaskCancelHandle::for_observation(
self,
record.task_id().clone(),
&id_prefix,
cancellation,
deadline,
)?;
let watch = Box::pin(
self.session
.await_active(cx, cancellation, deadline, None, async {
Ok(self
.watch_tasks_recovering_with_cancellation(
cx,
cancellation,
vec![record.task_id().clone()],
id_prefix,
policy,
recovery,
)
.await)
}),
)
.await??;
let result = PersistedManagedTaskWatch {
client: self.clone(),
current,
record,
cancellation: cancellation.clone(),
deadline,
watch: Some(watch),
remote_cancel,
persist,
pending: None,
terminal_cleanup: None,
cleanup_state: TaskResumePersistenceState::NotAttempted,
closed: false,
finished: false,
};
result.check::<E>(cx)?;
Ok(result)
}
}
#[must_use = "poll snapshots, close explicitly, or drop the observation owner"]
pub struct PersistedManagedTaskWatch<P> {
client: ManagedTasksClient,
current: TaskResumeBinding,
record: TaskResumeRecord,
cancellation: McpRequestCancellation,
deadline: Time,
watch: Option<RecoveringManagedTaskWatch>,
remote_cancel: ManagedTaskCancelHandle,
persist: P,
pending: Option<PendingTaskResumeSnapshot>,
terminal_cleanup: Option<TaskResumeChange>,
cleanup_state: TaskResumePersistenceState,
closed: bool,
finished: bool,
}
impl<P> PersistedManagedTaskWatch<P> {
pub fn last_published_record(&self) -> &TaskResumeRecord {
&self.record
}
pub fn pending(&self) -> Option<&PendingTaskResumeSnapshot> {
self.pending.as_ref()
}
pub fn take_pending(&mut self) -> Option<PendingTaskResumeSnapshot> {
self.pending.take()
}
pub fn terminal_cleanup(&self) -> Option<&TaskResumeChange> {
self.terminal_cleanup.as_ref()
}
pub fn cleanup_state(&self) -> TaskResumePersistenceState {
self.cleanup_state
}
pub fn cancel_handle(&self) -> ManagedTaskCancelHandle {
self.remote_cancel.clone()
}
pub fn close(&mut self) {
self.remote_cancel.close_observation();
self.watch = None;
self.closed = true;
}
pub async fn acknowledge_terminal<F, E>(
&mut self,
cx: &Cx,
) -> Result<(), PersistedTaskWatchError<E>>
where
P: FnMut(TaskResumeChange) -> F,
F: Future<Output = Result<(), E>>,
{
if self.finished {
return Ok(());
}
if self.closed {
return Err(PersistedTaskWatchError::Closed);
}
let change = self
.terminal_cleanup
.clone()
.ok_or(PersistedTaskWatchError::NoTerminal)?;
self.closed = true;
self.check::<E>(cx)?;
let client = self.client.clone();
let cancellation = self.cancellation.clone();
let persist = &mut self.persist;
let state = &mut self.cleanup_state;
Box::pin(
client
.session
.await_active(cx, &cancellation, self.deadline, None, async {
Ok(persist_change(state, persist, change).await)
}),
)
.await?
.map_err(PersistedTaskWatchError::Persistence)?;
self.finished = true;
Ok(())
}
fn check<E>(&self, cx: &Cx) -> Result<(), PersistedTaskWatchError<E>> {
if self.remote_cancel.cancellation_requested() {
return Err(PersistedTaskWatchError::CancellationRequested);
}
self.client.session.check(cx, &self.cancellation)?;
self.record.admit(cx, &self.current)?;
if cx.now() >= self.deadline {
return Err(OAuthSessionError::TimedOut.into());
}
Ok(())
}
pub async fn next_snapshot<F, E>(
&mut self,
cx: &Cx,
) -> Result<Option<ManagedTaskSnapshot>, PersistedTaskWatchError<E>>
where
P: FnMut(TaskResumeChange) -> F,
F: Future<Output = Result<(), E>>,
{
if self.finished {
return Ok(None);
}
if self.remote_cancel.cancellation_requested() {
self.close();
return Err(PersistedTaskWatchError::CancellationRequested);
}
if self.closed {
return Err(PersistedTaskWatchError::Closed);
}
if self.terminal_cleanup.is_some() {
return Err(PersistedTaskWatchError::TerminalAcknowledgementRequired);
}
let mut watch = self.watch.take().ok_or(PersistedTaskWatchError::Closed)?;
let remote_cancel = self.remote_cancel.clone();
let mut lease = remote_cancel.read_lease();
self.check::<E>(cx)?;
let client = self.client.clone();
let cancellation = self.cancellation.clone();
let deadline = self.deadline;
let read = Box::pin(client.session.await_active(
cx,
&cancellation,
deadline,
None,
async { Ok(watch.next_snapshot(cx).await) },
));
let snapshot = remote_cancel
.until_acknowledged(read)
.await???
.ok_or(TaskResumeError::InvalidRecord)?;
let change =
TaskResumeChange::from_snapshot(cx, &self.current, &self.record, &snapshot.task)?;
if change.replacement.is_none() {
self.check::<E>(cx)?;
remote_cancel.select_terminal()?;
self.terminal_cleanup = Some(change);
return Ok(Some(snapshot));
}
self.pending = Some(PendingTaskResumeSnapshot {
snapshot,
change,
persistence: TaskResumePersistenceState::NotAttempted,
});
self.check::<E>(cx)?;
let pending = self
.pending
.as_mut()
.ok_or(TaskResumeError::InvalidRecord)?;
let persist = &mut self.persist;
let writing = Box::pin(client.session.await_active(
cx,
&cancellation,
deadline,
None,
async {
Ok(persist_change(&mut pending.persistence, persist, pending.change.clone()).await)
},
));
remote_cancel
.until_acknowledged(writing)
.await??
.map_err(PersistedTaskWatchError::Persistence)?;
self.check::<E>(cx)?;
let pending = self.pending.take().ok_or(TaskResumeError::InvalidRecord)?;
self.record = pending
.change
.replacement
.ok_or(TaskResumeError::InvalidRecord)?;
self.watch = Some(watch);
lease.disarm();
Ok(Some(pending.snapshot))
}
}
impl<P> Drop for PersistedManagedTaskWatch<P> {
fn drop(&mut self) {
self.remote_cancel.close_observation();
}
}
async fn persist_change<P, F, E>(
state: &mut TaskResumePersistenceState,
persist: &mut P,
change: TaskResumeChange,
) -> Result<(), E>
where
P: FnMut(TaskResumeChange) -> F,
F: Future<Output = Result<(), E>>,
{
*state = TaskResumePersistenceState::Unconfirmed;
let result = persist(change).await;
if result.is_ok() {
*state = TaskResumePersistenceState::Acknowledged;
}
result
}
#[cfg(test)]
mod tests;