use fastmcp_core::CanonicalHttpUrl;
use fastmcp_protocol::protocol_policy::{
HttpEndpointBundleKey, HttpModernProbe, HttpProbeBody, ProtocolEra, ProtocolPolicy,
};
use std::sync::Arc;
use crate::negotiation::{
ClientHttpNegotiation, ClientHttpNegotiationDecision, ClientHttpNegotiationError,
};
use crate::session::ClientProtocolPlan;
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum HttpFallbackError {
PolicyForbidsFallback {
policy: ProtocolPolicy,
},
MissingHttpEndpointBundle,
MissingModernPostTarget,
MissingLegacySseTarget,
MissingLegacyMessagePostTarget,
Negotiation(ClientHttpNegotiationError),
InvalidObservation,
CrossBundleObservation,
ModernTargetMismatch,
ObservationAlreadyAdmitted {
admitted_attempt: u64,
},
ReplayedObservation {
attempt_id: u64,
},
IneligibleObservation {
status: u16,
body: HttpProbeBody,
},
LegacyGetNotAuthorized,
LegacyGetAlreadyOpened,
EndpointEventWithoutAuthorization,
EndpointEventMalformed,
EndpointEventTargetMismatch,
DuplicateEndpointEvent,
}
impl std::fmt::Display for HttpFallbackError {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::PolicyForbidsFallback { policy } => {
write!(formatter, "{policy:?} cannot coordinate an HTTP fallback")
}
Self::MissingHttpEndpointBundle => {
formatter.write_str("the plan has no configured HTTP endpoint bundle")
}
Self::MissingModernPostTarget => {
formatter.write_str("the plan has no configured modern POST target")
}
Self::MissingLegacySseTarget => {
formatter.write_str("the plan has no configured legacy SSE GET target")
}
Self::MissingLegacyMessagePostTarget => {
formatter.write_str("the plan has no configured legacy message POST target")
}
Self::Negotiation(error) => write!(formatter, "probe classification refused: {error}"),
Self::InvalidObservation => {
formatter.write_str("an observation needs a nonempty target and nonzero attempt")
}
Self::CrossBundleObservation => {
formatter.write_str("the observation or permit belongs to another coordinator")
}
Self::ModernTargetMismatch => {
formatter.write_str("the observation names another modern POST target")
}
Self::ObservationAlreadyAdmitted { admitted_attempt } => write!(
formatter,
"attempt {admitted_attempt} was already admitted by this coordinator"
),
Self::ReplayedObservation { attempt_id } => {
write!(
formatter,
"attempt {attempt_id} was replayed after settling"
)
}
Self::IneligibleObservation { status, body } => write!(
formatter,
"status {status} with {body:?} cannot authorize a legacy GET"
),
Self::LegacyGetNotAuthorized => {
formatter.write_str("this coordinator has not authorized a legacy GET")
}
Self::LegacyGetAlreadyOpened => {
formatter.write_str("the one authorized legacy GET was already opened")
}
Self::EndpointEventWithoutAuthorization => {
formatter.write_str("an endpoint event requires an opened authorized GET")
}
Self::EndpointEventMalformed => {
formatter.write_str("the endpoint event carried no usable target")
}
Self::EndpointEventTargetMismatch => formatter
.write_str("the advertised endpoint is not the configured message POST target"),
Self::DuplicateEndpointEvent => {
formatter.write_str("the era was already selected by an earlier endpoint event")
}
}
}
}
impl std::error::Error for HttpFallbackError {}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ModernProbeObservation {
bundle_key: HttpEndpointBundleKey,
modern_target: String,
attempt_id: u64,
probe: HttpModernProbe,
}
impl ModernProbeObservation {
pub fn new(
bundle_key: HttpEndpointBundleKey,
modern_target: impl Into<String>,
attempt_id: u64,
probe: HttpModernProbe,
) -> Result<Self, HttpFallbackError> {
let modern_target = modern_target.into();
if modern_target.is_empty() || attempt_id == 0 {
return Err(HttpFallbackError::InvalidObservation);
}
Ok(Self {
bundle_key,
modern_target,
attempt_id,
probe,
})
}
#[must_use]
pub const fn bundle_key(&self) -> &HttpEndpointBundleKey {
&self.bundle_key
}
#[must_use]
pub fn modern_target(&self) -> &str {
&self.modern_target
}
#[must_use]
pub const fn attempt_id(&self) -> u64 {
self.attempt_id
}
#[must_use]
pub const fn probe(&self) -> HttpModernProbe {
self.probe
}
}
#[derive(Debug)]
pub struct LegacyGetPermit {
target: String,
attempt_id: u64,
owner: Arc<()>,
}
impl PartialEq for LegacyGetPermit {
fn eq(&self, other: &Self) -> bool {
Arc::ptr_eq(&self.owner, &other.owner)
&& self.target == other.target
&& self.attempt_id == other.attempt_id
}
}
impl Eq for LegacyGetPermit {}
impl LegacyGetPermit {
#[must_use]
pub fn target(&self) -> &str {
&self.target
}
#[must_use]
pub const fn attempt_id(&self) -> u64 {
self.attempt_id
}
}
#[derive(Debug, PartialEq, Eq)]
pub enum FallbackDecision {
ModernRetained,
LegacyGetAuthorized(LegacyGetPermit),
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub struct FallbackState {
pub observations_admitted: usize,
pub legacy_gets_authorized: usize,
pub legacy_gets_opened: usize,
pub endpoint_events_admitted: usize,
pub selected_era: Option<ProtocolEra>,
pub credential_mutations: usize,
pub era_cache_mutations: usize,
}
pub struct HttpFallbackCoordinator {
plan: ClientProtocolPlan,
bundle_key: HttpEndpointBundleKey,
modern_target: String,
legacy_sse_target: String,
legacy_message_post_target: String,
admitted_message_post_target: Option<CanonicalHttpUrl>,
state: FallbackState,
settled_attempt: Option<u64>,
permit_owner: Arc<()>,
}
impl std::fmt::Debug for HttpFallbackCoordinator {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter
.debug_struct("HttpFallbackCoordinator")
.field("state", &self.state)
.field("settled_attempt", &self.settled_attempt)
.finish_non_exhaustive()
}
}
impl HttpFallbackCoordinator {
pub fn new(plan: ClientProtocolPlan) -> Result<Self, HttpFallbackError> {
let policy = plan.policy();
if !matches!(policy, ProtocolPolicy::Auto) {
return Err(HttpFallbackError::PolicyForbidsFallback { policy });
}
let bundle_key = plan
.http_endpoints()
.ok_or(HttpFallbackError::MissingHttpEndpointBundle)?
.key();
let modern_target = plan
.modern_post_target()
.ok_or(HttpFallbackError::MissingModernPostTarget)?
.to_owned();
let legacy_sse_target = plan
.legacy_sse_target()
.ok_or(HttpFallbackError::MissingLegacySseTarget)?
.to_owned();
let legacy_message_post_target = plan
.legacy_message_post_target()
.ok_or(HttpFallbackError::MissingLegacyMessagePostTarget)?
.to_owned();
ClientHttpNegotiation::from_protocol_plan(&plan).map_err(HttpFallbackError::Negotiation)?;
Ok(Self {
plan,
bundle_key,
modern_target,
legacy_sse_target,
legacy_message_post_target,
admitted_message_post_target: None,
state: FallbackState::default(),
settled_attempt: None,
permit_owner: Arc::new(()),
})
}
#[must_use]
pub const fn state(&self) -> FallbackState {
self.state
}
#[must_use]
pub const fn selected_era(&self) -> Option<ProtocolEra> {
self.state.selected_era
}
#[must_use]
pub fn legacy_sse_target(&self) -> &str {
&self.legacy_sse_target
}
#[must_use]
pub fn legacy_message_post_target(&self) -> &str {
&self.legacy_message_post_target
}
#[must_use]
pub fn advertised_message_post_target(&self) -> Option<&str> {
self.admitted_message_post_target
.as_ref()
.map(CanonicalHttpUrl::as_str)
}
#[must_use]
pub const fn bundle_key(&self) -> &HttpEndpointBundleKey {
&self.bundle_key
}
fn classify(
&self,
probe: HttpModernProbe,
) -> Result<ClientHttpNegotiationDecision, ClientHttpNegotiationError> {
let mut negotiation = ClientHttpNegotiation::from_protocol_plan(&self.plan)?;
negotiation.observe_modern_probe(probe)
}
pub fn observe(
&mut self,
observation: &ModernProbeObservation,
) -> Result<FallbackDecision, HttpFallbackError> {
if observation.bundle_key != self.bundle_key {
return Err(HttpFallbackError::CrossBundleObservation);
}
if observation.modern_target != self.modern_target {
return Err(HttpFallbackError::ModernTargetMismatch);
}
if let Some(settled) = self.settled_attempt {
return Err(if settled == observation.attempt_id {
HttpFallbackError::ReplayedObservation {
attempt_id: settled,
}
} else {
HttpFallbackError::ObservationAlreadyAdmitted {
admitted_attempt: settled,
}
});
}
let probe = observation.probe;
let decision =
self.classify(probe)
.map_err(|_| HttpFallbackError::IneligibleObservation {
status: probe.status,
body: probe.body,
})?;
self.settled_attempt = Some(observation.attempt_id);
self.state.observations_admitted += 1;
match decision {
ClientHttpNegotiationDecision::ModernSelected => Ok(FallbackDecision::ModernRetained),
ClientHttpNegotiationDecision::LegacySseFallbackAuthorized => {
self.state.legacy_gets_authorized += 1;
Ok(FallbackDecision::LegacyGetAuthorized(LegacyGetPermit {
target: self.legacy_sse_target.clone(),
attempt_id: observation.attempt_id,
owner: Arc::clone(&self.permit_owner),
}))
}
}
}
pub fn open_legacy_get(
&mut self,
permit: LegacyGetPermit,
) -> Result<String, HttpFallbackError> {
if self.state.legacy_gets_opened > 0 {
return Err(HttpFallbackError::LegacyGetAlreadyOpened);
}
if self.state.legacy_gets_authorized != 1 {
return Err(HttpFallbackError::LegacyGetNotAuthorized);
}
if !Arc::ptr_eq(&self.permit_owner, &permit.owner)
|| Some(permit.attempt_id) != self.settled_attempt
|| permit.target != self.legacy_sse_target
{
return Err(HttpFallbackError::CrossBundleObservation);
}
self.state.legacy_gets_opened += 1;
Ok(self.legacy_sse_target.clone())
}
pub fn admit_endpoint_event(
&mut self,
advertised: &str,
) -> Result<ProtocolEra, HttpFallbackError> {
if self.state.legacy_gets_opened == 0 {
return Err(HttpFallbackError::EndpointEventWithoutAuthorization);
}
if self.state.endpoint_events_admitted > 0 {
return Err(HttpFallbackError::DuplicateEndpointEvent);
}
let target = CanonicalHttpUrl::parse(advertised)
.map_err(|_| HttpFallbackError::EndpointEventMalformed)?;
if target.as_str() != advertised || target.has_userinfo() || target.fragment().is_some() {
return Err(HttpFallbackError::EndpointEventMalformed);
}
if !advertised_target_is_admissible(&self.legacy_message_post_target, target.as_str()) {
return Err(HttpFallbackError::EndpointEventTargetMismatch);
}
self.admitted_message_post_target = Some(target);
self.state.endpoint_events_admitted += 1;
self.state.selected_era = Some(ProtocolEra::Legacy2024);
Ok(ProtocolEra::Legacy2024)
}
}
fn advertised_target_is_admissible(configured: &str, advertised: &str) -> bool {
if advertised == configured {
return true;
}
match advertised.split_once('?') {
Some((base, query)) => !query.is_empty() && base == configured && !configured.contains('?'),
None => false,
}
}