use std::sync::OnceLock;
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
pub enum ModelCancellationReason {
#[default]
RequestDropped,
Interrupted,
FrontendDisconnected,
CredentialExpired,
RequestFailed,
}
impl ModelCancellationReason {
#[must_use]
pub const fn as_str(self) -> &'static str {
match self {
Self::RequestDropped => "request_dropped",
Self::Interrupted => "interrupted",
Self::FrontendDisconnected => "frontend_disconnected",
Self::CredentialExpired => "credential_expired",
Self::RequestFailed => "request_failed",
}
}
}
#[derive(Debug, Default)]
pub struct ModelCancellation(OnceLock<ModelCancellationReason>);
impl ModelCancellation {
pub fn record(&self, reason: ModelCancellationReason) {
if reason != ModelCancellationReason::RequestDropped {
let _ = self.0.set(reason);
}
}
#[must_use]
pub fn reason(&self) -> ModelCancellationReason {
self.0.get().copied().unwrap_or_default()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn specific_cause_survives_later_generic_failure() {
let cancellation = ModelCancellation::default();
assert_eq!(
cancellation.reason(),
ModelCancellationReason::RequestDropped
);
cancellation.record(ModelCancellationReason::RequestDropped);
cancellation.record(ModelCancellationReason::CredentialExpired);
cancellation.record(ModelCancellationReason::RequestDropped);
cancellation.record(ModelCancellationReason::RequestFailed);
assert_eq!(
cancellation.reason(),
ModelCancellationReason::CredentialExpired
);
}
}