use std::time::Duration;
use crate::dao::SaTokenDao;
enum RollbackStep {
DeleteKey { key: String },
RestoreKey {
key: String,
value: String,
ttl: Option<Duration>,
},
ListRemove { key: String, member: String },
}
impl RollbackStep {
fn key(&self) -> &str {
match self {
RollbackStep::DeleteKey { key }
| RollbackStep::RestoreKey { key, .. }
| RollbackStep::ListRemove { key, .. } => key,
}
}
}
#[derive(Debug, Default)]
pub struct RollbackReport {
pub succeeded: usize,
pub failed: Vec<(String, String)>,
}
impl RollbackReport {
pub fn is_clean(&self) -> bool {
self.failed.is_empty()
}
pub fn orphan_keys(&self) -> Vec<&str> {
self.failed.iter().map(|(k, _)| k.as_str()).collect()
}
}
pub struct LoginCompensator {
steps: Vec<RollbackStep>,
}
impl std::fmt::Debug for LoginCompensator {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("LoginCompensator { .. }")
}
}
impl LoginCompensator {
pub fn new() -> Self {
Self { steps: Vec::new() }
}
pub fn on_fail_delete(&mut self, key: impl Into<String>) {
self.steps.push(RollbackStep::DeleteKey { key: key.into() });
}
pub fn on_fail_restore(
&mut self,
key: impl Into<String>,
old_value: impl Into<String>,
ttl: Option<Duration>,
) {
self.steps.push(RollbackStep::RestoreKey {
key: key.into(),
value: old_value.into(),
ttl,
});
}
pub fn on_fail_list_remove(&mut self, key: impl Into<String>, member: impl Into<String>) {
self.steps.push(RollbackStep::ListRemove {
key: key.into(),
member: member.into(),
});
}
pub fn pending(&self) -> usize {
self.steps.len()
}
pub fn commit(self) {
drop(self);
}
pub async fn rollback(&self, dao: &SaTokenDao) -> RollbackReport {
let mut report = RollbackReport::default();
for step in self.steps.iter().rev() {
let outcome = match step {
RollbackStep::DeleteKey { key } => dao.delete(key).await.map(|_| ()),
RollbackStep::RestoreKey { key, value, ttl } => {
dao.set_string(key, value, *ttl).await
}
RollbackStep::ListRemove { key, member } => {
dao.list_remove(key, member).await.map(|_| ())
}
};
match outcome {
Ok(()) => report.succeeded += 1,
Err(e) => report.failed.push((step.key().to_string(), e.to_string())),
}
}
if !report.is_clean() {
tracing::error!(
succeeded = report.succeeded,
failed = report.failed.len(),
orphans = ?report.orphan_keys(),
"login rollback incomplete, orphan keys left in storage"
);
} else if report.succeeded > 0 {
tracing::warn!(
reverted = report.succeeded,
"login failed, all staged writes reverted"
);
}
report
}
}
impl Default for LoginCompensator {
fn default() -> Self {
Self::new()
}
}