use std::sync::{
Arc,
atomic::{AtomicU64, Ordering},
};
use jiff::Timestamp;
use crate::ids::{AccountId, FencingToken, KeyId, LeaseId, PolicyRevision, RequestId};
use crate::units::CostUnits;
pub trait UsageSlot: Send + 'static {
fn record(self, event: UsageEvent);
}
#[derive(Debug)]
pub struct DiscardedUsage {
count: AtomicU64,
}
impl DiscardedUsage {
#[must_use]
pub fn new() -> Arc<Self> {
Arc::new(Self {
count: AtomicU64::new(0),
})
}
#[must_use]
pub fn slot(self: &Arc<Self>) -> DiscardedUsageSlot {
DiscardedUsageSlot(Arc::clone(self))
}
#[must_use]
pub fn count(&self) -> u64 {
self.count.load(Ordering::Relaxed)
}
}
#[derive(Debug)]
pub struct DiscardedUsageSlot(Arc<DiscardedUsage>);
impl UsageSlot for DiscardedUsageSlot {
fn record(self, _event: UsageEvent) {
self.0.count.fetch_add(1, Ordering::Relaxed);
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub enum UsageSource {
Leased {
lease_id: LeaseId,
fencing_token: FencingToken,
},
Overage,
}
impl UsageSource {
#[must_use]
pub const fn lease_id(self) -> Option<LeaseId> {
match self {
UsageSource::Leased { lease_id, .. } => Some(lease_id),
UsageSource::Overage => None,
}
}
#[must_use]
pub const fn fencing_token(self) -> Option<FencingToken> {
match self {
UsageSource::Leased { fencing_token, .. } => Some(fencing_token),
UsageSource::Overage => None,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
#[non_exhaustive]
pub struct UsageEvent {
pub request_id: RequestId,
pub account_id: AccountId,
pub source: UsageSource,
pub units: CostUnits,
pub occurred_at: Timestamp,
#[cfg_attr(feature = "serde", serde(default))]
pub policy_revision: PolicyRevision,
#[cfg_attr(feature = "serde", serde(default))]
pub key_id: Option<KeyId>,
}
impl UsageEvent {
#[must_use]
pub const fn new(
request_id: RequestId,
account_id: AccountId,
source: UsageSource,
units: CostUnits,
occurred_at: Timestamp,
policy_revision: PolicyRevision,
key_id: Option<KeyId>,
) -> Self {
Self {
request_id,
account_id,
source,
units,
occurred_at,
policy_revision,
key_id,
}
}
}
#[cfg(all(test, feature = "serde"))]
mod revision_wire_tests {
use super::*;
fn event(revision: PolicyRevision) -> UsageEvent {
UsageEvent::new(
RequestId(9),
AccountId(1),
UsageSource::Overage,
CostUnits(70),
Timestamp::UNIX_EPOCH,
revision,
None,
)
}
#[test]
fn an_event_without_a_revision_key_decodes_as_unstated() {
let mut value =
serde_json::to_value(event(PolicyRevision([0xc3; 32]))).expect("an event serializes");
assert_eq!(
value["policy_revision"],
serde_json::Value::String("c3".repeat(32)),
"a stated revision is on the wire in canonical form"
);
assert!(
value
.as_object_mut()
.expect("an event is a JSON object")
.remove("policy_revision")
.is_some()
);
let decoded: UsageEvent =
serde_json::from_value(value).expect("an older payload still decodes");
assert_eq!(decoded.policy_revision, PolicyRevision::UNSTATED);
}
#[test]
fn a_revision_survives_the_event_round_trip_exactly() {
let mut bytes = [0u8; 32];
for (index, byte) in bytes.iter_mut().enumerate() {
*byte = (index as u8).wrapping_mul(7).wrapping_add(1);
}
let original = event(PolicyRevision(bytes));
let encoded = serde_json::to_string(&original).expect("an event serializes");
let decoded: UsageEvent = serde_json::from_str(&encoded).expect("it decodes");
assert_eq!(decoded, original);
assert_eq!(decoded.policy_revision.as_bytes(), &bytes);
}
}
#[cfg(test)]
mod discarded_usage_tests {
use super::*;
fn event(request: u128) -> UsageEvent {
UsageEvent::new(
RequestId(request),
AccountId(1),
UsageSource::Overage,
CostUnits(70),
Timestamp::UNIX_EPOCH,
PolicyRevision::UNSTATED,
None,
)
}
#[test]
fn discarding_is_counted_not_silent() {
let discarded = DiscardedUsage::new();
assert_eq!(discarded.count(), 0);
discarded.slot().record(event(1));
discarded.slot().record(event(2));
assert_eq!(discarded.count(), 2);
}
#[test]
fn slots_are_independent_and_share_one_counter() {
let discarded = DiscardedUsage::new();
let first = discarded.slot();
let second = discarded.slot();
second.record(event(2));
first.record(event(1));
assert_eq!(discarded.count(), 2);
}
#[test]
fn counts_across_threads() {
let discarded = DiscardedUsage::new();
std::thread::scope(|scope| {
for index in 0..8u128 {
let discarded = Arc::clone(&discarded);
scope.spawn(move || discarded.slot().record(event(index)));
}
});
assert_eq!(discarded.count(), 8);
}
}