use mkit_core::hash::Hash;
use mkit_core::protocol::AdvanceOutcome;
use crate::error::Code;
#[derive(Debug, Clone, PartialEq, Eq, Hash, PartialOrd, Ord)]
pub struct ReplayKey(pub Hash);
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ReplayRecord {
pub fingerprint: Hash,
pub expires_at_ms: i64,
pub state: ReplayState,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum ReplayState {
InFlight {
resumable: bool,
},
Committed(StoredResult),
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum StoredResult {
UpdateRef(UpdateRefResult),
AdvanceRefs(AdvanceOutcome),
BeginUpload(BeginUploadResult),
UploadPack,
RepoVisibility,
Rejected(StoredRejection),
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum BeginUploadResult {
AlreadyPresent,
Ticket {
id: Hash,
part_size: u64,
expires_at_ms: u64,
token: Vec<u8>,
},
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum UpdateRefResult {
Committed,
Conflict {
current: Option<Hash>,
},
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct StoredRejection {
code: Code,
message: String,
}
impl StoredRejection {
#[must_use]
pub const fn is_storable(code: Code) -> bool {
matches!(
code,
Code::InvalidArgument
| Code::NotFound
| Code::AlreadyExists
| Code::PermissionDenied
| Code::FailedPrecondition
| Code::OutOfRange
| Code::Unimplemented
)
}
#[must_use]
pub fn new(code: Code, message: impl Into<String>) -> Option<Self> {
Self::is_storable(code).then(|| Self {
code,
message: message.into(),
})
}
#[must_use]
pub fn code(&self) -> Code {
self.code
}
#[must_use]
pub fn message(&self) -> &str {
&self.message
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum ReplayDecision {
New,
Return(StoredResult),
Resume,
RetryLater,
FingerprintMismatch,
}
#[must_use]
pub fn classify(existing: Option<&ReplayRecord>, fingerprint: &Hash) -> ReplayDecision {
match existing {
None => ReplayDecision::New,
Some(record) if record.fingerprint != *fingerprint => ReplayDecision::FingerprintMismatch,
Some(record) => match &record.state {
ReplayState::Committed(result) => ReplayDecision::Return(result.clone()),
ReplayState::InFlight { resumable: true } => ReplayDecision::Resume,
ReplayState::InFlight { resumable: false } => ReplayDecision::RetryLater,
},
}
}
#[cfg(test)]
mod tests {
use super::*;
const FP: Hash = [7; 32];
fn record(state: ReplayState) -> ReplayRecord {
ReplayRecord {
fingerprint: FP,
expires_at_ms: 1_000,
state,
}
}
#[test]
fn classify_none_is_new() {
assert_eq!(classify(None, &FP), ReplayDecision::New);
}
#[test]
fn classify_other_fingerprint_is_mismatch() {
for state in [
ReplayState::InFlight { resumable: true },
ReplayState::Committed(StoredResult::UploadPack),
] {
assert_eq!(
classify(Some(&record(state)), &[8; 32]),
ReplayDecision::FingerprintMismatch
);
}
}
#[test]
fn classify_committed_returns_stored_result() {
let result = StoredResult::UpdateRef(UpdateRefResult::Conflict {
current: Some([1; 32]),
});
let rec = record(ReplayState::Committed(result.clone()));
assert_eq!(classify(Some(&rec), &FP), ReplayDecision::Return(result));
}
#[test]
fn classify_inflight_resumable_is_resume() {
let rec = record(ReplayState::InFlight { resumable: true });
assert_eq!(classify(Some(&rec), &FP), ReplayDecision::Resume);
}
#[test]
fn classify_inflight_not_resumable_is_retry_later() {
let rec = record(ReplayState::InFlight { resumable: false });
assert_eq!(classify(Some(&rec), &FP), ReplayDecision::RetryLater);
}
#[test]
fn no_stored_result_variant_for_challenge() {
fn is_final(result: &StoredResult) -> bool {
match result {
StoredResult::UpdateRef(_)
| StoredResult::AdvanceRefs(_)
| StoredResult::BeginUpload(_)
| StoredResult::UploadPack
| StoredResult::RepoVisibility => true,
StoredResult::Rejected(r) => StoredRejection::is_storable(r.code()),
}
}
let denied = StoredRejection::new(Code::PermissionDenied, "denied").unwrap();
assert_eq!(
(denied.code(), denied.message()),
(Code::PermissionDenied, "denied")
);
assert!(is_final(&StoredResult::Rejected(denied)));
}
#[test]
fn stored_rejection_admits_only_final_codes() {
let final_codes = [
Code::InvalidArgument,
Code::NotFound,
Code::AlreadyExists,
Code::PermissionDenied,
Code::FailedPrecondition,
Code::OutOfRange,
Code::Unimplemented,
];
for code in final_codes {
assert!(StoredRejection::new(code, "m").is_some(), "{code:?}");
}
let refused = [
Code::Unavailable,
Code::Aborted,
Code::ResourceExhausted,
Code::Canceled,
Code::DeadlineExceeded,
Code::Internal,
Code::Unknown,
Code::DataLoss,
Code::Unauthenticated,
];
for code in refused {
assert_eq!(StoredRejection::new(code, "m"), None, "{code:?}");
}
assert_eq!(
final_codes.len() + refused.len(),
16,
"every Code is classified"
);
}
}