use std::fmt;
mod header_repair;
use asupersync::Cx;
use asupersync::http::h1::{ClientError, HttpError};
use fastmcp_core::CanonicalHttpUrl;
use fastmcp_protocol::{CoreRequest, FinalInputResponses, InputRequiredResult, RequestId};
use super::{
InputSelection, ManagedCoreCall, ManagedCoreError, ManagedCoreEvent, ManagedInteraction,
ManagedInteractionError, ManagedInteractionEvent, Step, admit_challenge, admit_fresh_id,
bounded_wait, continuation_request_selected, input_required,
};
use crate::http_auth::managed::OAuthSessionError;
use crate::http_executor::{ModernHttpExecutorError, ModernHttpResponseKind};
pub(crate) const MAX_RECOVERIES: usize = 4;
const MAX_REQUEST_IDS: usize = 1 + 64 * (1 + MAX_RECOVERIES);
#[derive(Clone)]
pub struct ContinuationReplayContract {
resource: CanonicalHttpUrl,
maximum_recoveries: usize,
}
impl ContinuationReplayContract {
pub fn for_configured_endpoint(
resource: CanonicalHttpUrl,
maximum_recoveries: usize,
) -> Result<Self, ContinuationRecoveryError> {
if !resource.as_str().starts_with("https://")
|| !(1..=MAX_RECOVERIES).contains(&maximum_recoveries)
{
return Err(ContinuationRecoveryError::InvalidContract);
}
Ok(Self {
resource,
maximum_recoveries,
})
}
pub(crate) fn admit_endpoint(
&self,
resource: &CanonicalHttpUrl,
) -> Result<usize, ContinuationRecoveryError> {
if resource.as_str() != self.resource.as_str() {
return Err(ContinuationRecoveryError::EndpointMismatch);
}
Ok(self.maximum_recoveries)
}
}
impl fmt::Debug for ContinuationReplayContract {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("ContinuationReplayContract")
.field("maximum_recoveries", &self.maximum_recoveries)
.finish_non_exhaustive()
}
}
#[derive(Debug)]
pub enum ContinuationRecoveryError {
Interrupted,
InvalidContract,
EndpointMismatch,
StateRequired,
WrongPhase,
RecoveryLimit,
JsonReplyRequired,
Interaction(ManagedInteractionError),
}
impl fmt::Display for ContinuationRecoveryError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Interrupted => f.write_str(
"continuation reply interrupted; recovery requires an explicit host decision",
),
Self::InvalidContract => f.write_str("invalid configured continuation replay contract"),
Self::EndpointMismatch => {
f.write_str("continuation replay contract names another endpoint")
}
Self::StateRequired => {
f.write_str("reply recovery requires nonempty server continuation state")
}
Self::WrongPhase => {
f.write_str("continuation recovery operation is not in the required phase")
}
Self::RecoveryLimit => f.write_str("continuation reply recovery budget exhausted"),
Self::JsonReplyRequired => {
f.write_str("continuation reply recovery requires a finite JSON response")
}
Self::Interaction(error) => error.fmt(f),
}
}
}
impl std::error::Error for ContinuationRecoveryError {}
impl From<ManagedInteractionError> for ContinuationRecoveryError {
fn from(error: ManagedInteractionError) -> Self {
Self::Interaction(error)
}
}
impl From<ManagedCoreError> for ContinuationRecoveryError {
fn from(error: ManagedCoreError) -> Self {
Self::Interaction(error.into())
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum Phase {
Prepared,
Reading,
Recoverable,
Delivered,
Closed,
}
#[must_use = "send once, observe the reply, and explicitly decide any recovery"]
pub struct RecoverableManagedContinuation {
interaction: ManagedInteraction,
request: Option<CoreRequest>,
call: Option<Box<ManagedCoreCall>>,
phase: Phase,
maximum_recoveries: usize,
attempts: usize,
answer_count: usize,
}
impl fmt::Debug for RecoverableManagedContinuation {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("RecoverableManagedContinuation")
.field("phase", &self.phase)
.field("attempts", &self.attempts)
.finish_non_exhaustive()
}
}
impl ManagedInteraction {
pub fn prepare_recoverable_continuation(
mut self,
cx: &Cx,
responses: Option<FinalInputResponses>,
contract: ContinuationReplayContract,
) -> Result<RecoverableManagedContinuation, ContinuationRecoveryError> {
self.check(cx)?;
let maximum_recoveries = contract.admit_endpoint(self.session.resource())?;
let input = self
.pending_input()
.ok_or(ManagedInteractionError::NotAwaitingInput)?;
admit_challenge(
&self.original,
input,
self.limits,
self.continuations,
self.input_responses,
)?;
admit_answer_bytes(responses.as_ref(), self.limits.core.request_bytes)?;
let answer_count = responses.as_ref().map_or(0, FinalInputResponses::len);
let request = recovery_request(&self.original, input, responses)?;
let _ = self.prepare_request(request.clone(), RequestId::Number(0))?;
self.step = None;
Ok(RecoverableManagedContinuation {
interaction: self,
request: Some(request),
call: None,
phase: Phase::Prepared,
maximum_recoveries,
attempts: 0,
answer_count,
})
}
}
impl RecoverableManagedContinuation {
pub fn attempts(&self) -> usize {
self.attempts
}
pub fn is_recovery_pending(&self) -> bool {
self.phase == Phase::Recoverable
}
pub async fn send(
&mut self,
cx: &Cx,
request_id: RequestId,
) -> Result<(), ContinuationRecoveryError> {
self.attempt(cx, request_id, false).await
}
pub async fn recover(
&mut self,
cx: &Cx,
request_id: RequestId,
) -> Result<(), ContinuationRecoveryError> {
self.attempt(cx, request_id, true).await
}
pub fn close(&mut self) {
self.call = None;
self.request = None;
self.interaction.close();
self.phase = Phase::Closed;
}
pub async fn next_event(
&mut self,
cx: &Cx,
) -> Result<ManagedInteractionEvent, ContinuationRecoveryError> {
self.check(cx)?;
if self.phase != Phase::Reading {
return Err(ContinuationRecoveryError::WrongPhase);
}
let mut call = self
.call
.take()
.ok_or(ContinuationRecoveryError::WrongPhase)?;
self.phase = Phase::Recoverable;
let outcome = call.next_event(cx).await;
self.check(cx)?;
let result = match outcome {
Ok(Some(ManagedCoreEvent::Result(result))) => result,
Ok(_) => {
self.close();
return Err(ContinuationRecoveryError::JsonReplyRequired);
}
Err(error) => return Err(self.failed(error)),
};
self.interaction.response_bytes = call.decoder.bytes;
self.interaction.notifications = call.decoder.notifications;
let event = if let Some(input) = input_required(&result) {
if self.interaction.response_bytes >= self.interaction.limits.core.total_bytes {
self.close();
return Err(ManagedCoreError::ResponseByteLimit.into());
}
if let Err(error) = admit_challenge(
&self.interaction.original,
input,
self.interaction.limits,
self.interaction.continuations,
self.interaction.input_responses,
) {
self.close();
return Err(error.into());
}
let input = Box::new(input.clone());
self.interaction.step = Some(Step::Awaiting(input.clone()));
ManagedInteractionEvent::InputRequired(input)
} else {
self.interaction.step = Some(Step::Complete);
ManagedInteractionEvent::Complete(result)
};
self.check(cx)?;
self.request = None;
self.phase = Phase::Delivered;
Ok(event)
}
pub fn into_interaction(self) -> Result<ManagedInteraction, ContinuationRecoveryError> {
if self.phase != Phase::Delivered {
return Err(ContinuationRecoveryError::WrongPhase);
}
Ok(self.interaction)
}
async fn attempt(
&mut self,
cx: &Cx,
request_id: RequestId,
recovery: bool,
) -> Result<(), ContinuationRecoveryError> {
self.check(cx)?;
let expected = if recovery {
Phase::Recoverable
} else {
Phase::Prepared
};
if self.phase != expected {
return Err(ContinuationRecoveryError::WrongPhase);
}
if recovery && self.attempts > self.maximum_recoveries {
return Err(ContinuationRecoveryError::RecoveryLimit);
}
if self.interaction.used_ids.len() >= MAX_REQUEST_IDS {
return Err(ContinuationRecoveryError::RecoveryLimit);
}
admit_fresh_id(&self.interaction.used_ids, &request_id)?;
let core = self.interaction.limits.core;
let reserved = reserve_frame(
self.interaction.response_bytes,
core.frame_bytes,
core.total_bytes,
)?;
let request = self
.request
.as_ref()
.ok_or(ContinuationRecoveryError::WrongPhase)?
.clone();
let (wire, mut decoder) = self
.interaction
.prepare_request(request, request_id.clone())?;
decoder.bytes = self.interaction.response_bytes;
decoder.notifications = self.interaction.notifications;
self.check(cx)?;
if !recovery {
self.interaction.continuations += 1;
self.interaction.input_responses += self.answer_count;
}
self.interaction.response_bytes = reserved;
self.interaction.used_ids.push(request_id);
self.attempts += 1;
self.phase = Phase::Recoverable;
let response = bounded_wait(
cx,
&self.interaction.cancellation,
self.interaction.deadline,
async {
self.interaction
.session
.execute_with_cancellation(cx, &self.interaction.cancellation, &wire)
.await
.map_err(ManagedCoreError::from)
},
)
.await;
self.check(cx)?;
let response = match response {
Ok(response) => response,
Err(error) => return Err(self.failed(error)),
};
if response.metadata().kind() != ModernHttpResponseKind::Json {
self.close();
return Err(ContinuationRecoveryError::JsonReplyRequired);
}
let call = ManagedCoreCall::from_response(
response,
decoder,
self.interaction.cancellation.clone(),
self.interaction.deadline,
);
match call {
Ok(call) => {
self.interaction.generation = call.credential_generation();
self.call = Some(Box::new(call));
self.phase = Phase::Reading;
Ok(())
}
Err(error) => Err(self.failed(error)),
}
}
fn check(&mut self, cx: &Cx) -> Result<(), ContinuationRecoveryError> {
if let Err(error) = self.interaction.check(cx) {
self.close();
return Err(error.into());
}
Ok(())
}
fn failed(&mut self, error: ManagedCoreError) -> ContinuationRecoveryError {
if recoverable_transport_failure(&error) {
return ContinuationRecoveryError::Interrupted;
}
self.close();
error.into()
}
}
pub(crate) fn recovery_request(
original: &CoreRequest,
input: &InputRequiredResult,
responses: Option<FinalInputResponses>,
) -> Result<CoreRequest, ContinuationRecoveryError> {
if input.request_state().is_none_or(str::is_empty) {
return Err(ContinuationRecoveryError::StateRequired);
}
let selection = if responses
.as_ref()
.is_some_and(|answers| !answers.is_empty())
{
InputSelection::Partial
} else {
InputSelection::Complete
};
Ok(continuation_request_selected(
original, input, responses, selection,
)?)
}
pub(crate) fn reserve_frame(
used: usize,
frame: usize,
total: usize,
) -> Result<usize, ManagedCoreError> {
used.checked_add(frame)
.filter(|reserved| *reserved <= total)
.ok_or(ManagedCoreError::ResponseByteLimit)
}
pub(crate) fn admit_answer_bytes(
responses: Option<&FinalInputResponses>,
maximum: usize,
) -> Result<(), ManagedCoreError> {
if let Some(responses) = responses {
let mut encoded = super::super::BoundedWriter {
bytes: Vec::new(),
maximum,
};
serde_json::to_writer(&mut encoded, responses)
.map_err(|_| ManagedCoreError::RequestTooLarge)?;
}
Ok(())
}
fn recoverable_transport_failure(error: &ManagedCoreError) -> bool {
matches!(error, ManagedCoreError::MissingTerminal)
|| matches!(error, ManagedCoreError::Session(OAuthSessionError::Http(error))
if recovery_http_interruption(error))
}
pub(crate) fn recovery_http_interruption(error: &ModernHttpExecutorError) -> bool {
matches!(
error,
ModernHttpExecutorError::ResponseBodyReadFailed
| ModernHttpExecutorError::Transport(
ClientError::Io(_) | ClientError::HttpError(HttpError::Io(_))
)
| ModernHttpExecutorError::DispatchUncertain(
ClientError::Io(_) | ClientError::HttpError(HttpError::Io(_))
)
)
}
#[cfg(test)]
mod tests;