use core::future::Future;
use std::collections::HashSet;
use std::sync::Mutex;
use mkit_core::hash::Hash;
use crate::error::ServerError;
use crate::op::{OpKind, Operation};
use crate::rt::{BoxFuture, MaybeSend, MaybeSync};
use crate::store::{Batch, codec};
pub const FAULT_HEADER: &str = "x-mkit-test-fault";
pub const CLOCK_SKEW_HEADER: &str = "x-mkit-test-clock-skew-ms";
pub const TIMER_MS_HEADER: &str = "x-mkit-test-timer-ms";
pub const RUN_TIMERS_HEADER: &str = "x-mkit-test-run-timers";
pub const BUMP_EPOCH_HEADER: &str = "x-mkit-test-bump-epoch";
pub const LEASE_RECOVERED_HEADER: &str = "x-mkit-test-lease-recovered";
pub const RELAY_DELAY_MS_HEADER: &str = "x-mkit-test-relay-delay-ms";
pub(crate) fn relay_delay_key() -> crate::Key {
crate::Key::new(&b"tdr\0"[..])
}
pub(crate) fn delay_relay_batch(
batch: Batch,
directives: &TestDirectives,
op: &Operation,
now_ms: u64,
) -> Batch {
if let Some(delay) = directives.relay_delay_ms
&& matches!(op.kind, OpKind::UpdateRef(_) | OpKind::AdvanceRefs { .. })
&& batch.preconditions.len() + batch.writes.len() < crate::store::MAX_BATCH_OPS
{
batch.put(
relay_delay_key(),
codec::encode_u64(now_ms.saturating_add(delay)),
)
} else {
batch
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum FaultPoint {
AfterAuthenticate,
AfterAuthorize,
AfterReserve,
AfterBlobCommit,
AfterLeaseGrant,
BeforeFinalApply,
}
pub trait FaultHooks: MaybeSend + MaybeSync {
fn at(
&self,
point: FaultPoint,
op: &Operation,
directives: &TestDirectives,
) -> impl Future<Output = Result<(), ServerError>> + MaybeSend;
}
pub(crate) trait DynFaultHooks: MaybeSend + MaybeSync {
fn at_boxed<'a>(
&'a self,
point: FaultPoint,
op: &'a Operation,
directives: &'a TestDirectives,
) -> BoxFuture<'a, Result<(), ServerError>>;
}
impl<F: FaultHooks> DynFaultHooks for F {
fn at_boxed<'a>(
&'a self,
point: FaultPoint,
op: &'a Operation,
directives: &'a TestDirectives,
) -> BoxFuture<'a, Result<(), ServerError>> {
Box::pin(self.at(point, op, directives))
}
}
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct TestDirectives {
pub fault: Option<String>,
pub clock_skew_ms: i64,
pub timer_ms: Option<u64>,
pub bump_epoch: Option<u64>,
pub lease_recovered: bool,
pub run_timers: Option<String>,
pub relay_delay_ms: Option<u64>,
}
impl TestDirectives {
pub fn from_headers(get: impl Fn(&str) -> Option<String>) -> Result<Self, ServerError> {
let clock_skew_ms = match get(CLOCK_SKEW_HEADER) {
Some(v) => v.trim().parse().map_err(|_| {
ServerError::invalid_argument("x-mkit-test-clock-skew-ms is not an integer")
})?,
None => 0,
};
let timer_ms = get(TIMER_MS_HEADER)
.map(|v| {
v.trim().parse::<u64>().map_err(|_| {
ServerError::invalid_argument("x-mkit-test-timer-ms is not an unsigned integer")
})
})
.transpose()?;
let run_timers = get(RUN_TIMERS_HEADER);
let relay_delay_ms = get(RELAY_DELAY_MS_HEADER)
.map(|v| {
v.trim().parse::<u64>().map_err(|_| {
ServerError::invalid_argument(
"x-mkit-test-relay-delay-ms is not an unsigned integer",
)
})
})
.transpose()?;
if run_timers
.as_ref()
.is_some_and(|name| !crate::refs::validate_ref_name(name))
{
return Err(ServerError::invalid_argument(
"x-mkit-test-run-timers is not a ref name",
));
}
let bump_epoch = get(BUMP_EPOCH_HEADER)
.map(|v| {
v.trim().parse::<u64>().map_err(|_| {
ServerError::invalid_argument(
"x-mkit-test-bump-epoch is not an unsigned integer",
)
})
})
.transpose()?;
let lease_recovered = match get(LEASE_RECOVERED_HEADER).as_deref() {
None => false,
Some("1") => true,
Some(_) => {
return Err(ServerError::invalid_argument(
"x-mkit-test-lease-recovered must be 1",
));
}
};
Ok(Self {
bump_epoch,
lease_recovered,
timer_ms,
run_timers,
relay_delay_ms,
fault: get(FAULT_HEADER).filter(|f| !f.is_empty()),
clock_skew_ms,
})
}
}
#[derive(Debug, Default)]
pub struct FailOnce {
failed: Mutex<HashSet<(FaultPoint, Option<Hash>)>>,
}
impl FailOnce {
#[must_use]
pub fn new() -> Self {
Self::default()
}
}
impl FaultHooks for FailOnce {
async fn at(
&self,
point: FaultPoint,
op: &Operation,
directives: &TestDirectives,
) -> Result<(), ServerError> {
let token = match point {
FaultPoint::AfterReserve => "after-reserve",
FaultPoint::AfterBlobCommit => "after-put",
_ => return Ok(()),
};
if directives.fault.as_deref() != Some(token) {
return Ok(());
}
let scope = op.auth.as_ref().map(|a| a.replay_scope);
let first = self
.failed
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.insert((point, scope));
if first {
return Err(ServerError::internal("injected test fault", token));
}
Ok(())
}
}
pub(crate) async fn schedule_timer<S: crate::NamespaceStore>(
directives: &TestDirectives,
store: &S,
p: &crate::Partition,
repo: &crate::RepoName,
name: &str,
business_now: u64,
) -> Result<(), ServerError> {
use crate::store::keys;
use crate::timers::registry::kinds;
use crate::{Batch, BatchOutcome, Value};
let Some(delay) = directives.timer_ms else {
return Ok(());
};
let reference = [repo.as_str().as_bytes(), b"\0", name.as_bytes()].concat();
let batch = Batch::new().put(
keys::timer(
business_now.saturating_add(delay),
kinds::TEST.get(),
&reference,
),
Value::default(),
);
match store.apply(p, batch).await {
Ok(BatchOutcome::Committed) => Ok(()),
outcome => Err(ServerError::internal(
"test timer scheduling failed",
format!("{outcome:?}"),
)),
}
}
#[allow(clippy::too_many_arguments)] pub(crate) async fn run_timers<S: crate::NamespaceStore>(
directives: &TestDirectives,
store: &S,
blobs: &impl crate::MultipartBlobStore,
shards: &dyn super::ShardMap,
repo: &crate::RepoId,
clock: &dyn crate::Clock,
business_now: u64,
gate: Option<&tokio::sync::Mutex<()>>,
) -> Result<(), ServerError> {
use crate::relay::{NoHook, RelayBudget, RelayHandler};
use crate::timers::{TickBudget, TimerRegistry, run_due, test_kind::TestTimer};
if let Some(name) = &directives.run_timers {
let _tick_guard = match gate {
Some(gate) => Some(gate.lock().await),
None => None,
};
let partition = shards.ref_shard(repo, name);
let registry = TimerRegistry::new()
.register(TestTimer)
.register(RelayHandler {
target: crate::store::BorrowedStore(store),
hook: NoHook,
budget: RelayBudget::default(),
})
.register(crate::timers::quota_rollup::QuotaRollup {
coordinator: crate::store::BorrowedStore(store),
metrics: crate::telemetry::NoopMetrics,
})
.register(crate::timers::ticket_expiry::BorrowedTicketExpiry { blobs });
for _ in 0..64 {
let report = run_due(
store,
&partition,
®istry,
clock,
business_now,
&TickBudget::default(),
)
.await
.map_err(|e| ServerError::internal("test timer tick failed", e))?;
if report.fired == 0 {
if report.failed > 0 {
return Err(ServerError::unavailable(format!(
"test timer tick failed: {report:?}"
)));
}
if report.raced > 0 || report.stopped_on_budget || report.deferred > 0 {
continue;
}
return Ok(());
}
}
return Err(ServerError::unavailable(
"test timers did not drain within 64 ticks",
));
}
Ok(())
}
#[cfg(test)]
mod timer_tests {
use super::*;
#[test]
fn timer_directives_validate_values() {
for (header, value) in [
(TIMER_MS_HEADER, "-1"),
(TIMER_MS_HEADER, "18446744073709551616"),
(RELAY_DELAY_MS_HEADER, "-1"),
(RUN_TIMERS_HEADER, "bad ref"),
] {
assert!(TestDirectives::from_headers(|h| (h == header).then(|| value.into())).is_err());
}
let d = TestDirectives::from_headers(|h| match h {
TIMER_MS_HEADER => Some("1000".into()),
RUN_TIMERS_HEADER => Some("refs/heads/x".into()),
RELAY_DELAY_MS_HEADER => Some("10000".into()),
_ => None,
})
.unwrap();
assert_eq!(d.timer_ms, Some(1000));
assert_eq!(d.run_timers.as_deref(), Some("refs/heads/x"));
assert_eq!(d.relay_delay_ms, Some(10_000));
}
}