use core::fmt;
use crate::{snapshot::PermissionBits, units::CostUnits};
pub trait OpIndex {
fn index(&self) -> usize;
}
impl<O: OpIndex + ?Sized> OpIndex for &O {
#[inline]
fn index(&self) -> usize {
(**self).index()
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct CostQuote {
pub total: CostUnits,
pub fixed: CostUnits,
pub variable: CostUnits,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub enum QuoteError {
EmptyWorkload,
UnknownOperation { index: usize },
Overflow,
}
impl fmt::Display for QuoteError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
QuoteError::EmptyWorkload => f.write_str("workload is empty"),
QuoteError::UnknownOperation { index } => {
write!(f, "operation index {index} is not in the cost table")
}
QuoteError::Overflow => f.write_str("cost quote overflowed"),
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct CostTable {
fixed_request: CostUnits,
minimum_charge: CostUnits,
weights: Box<[Option<CostUnits>]>,
#[cfg_attr(
feature = "serde",
serde(default, skip_serializing_if = "<[PermissionBits]>::is_empty")
)]
permissions: Box<[PermissionBits]>,
}
impl CostTable {
#[must_use]
pub fn builder(fixed_request: CostUnits, minimum_charge: CostUnits) -> CostTableBuilder {
CostTableBuilder {
fixed_request,
minimum_charge,
weights: Vec::new(),
permissions: Vec::new(),
}
}
#[inline]
pub fn quote(&self, op: &impl OpIndex, items: u64) -> Result<CostQuote, QuoteError> {
if items == 0 {
return Err(QuoteError::EmptyWorkload);
}
self.quote_weight(self.weight_at(op.index())?, items)
}
#[inline]
pub fn quote_workload<O: OpIndex>(
&self,
workload: &[(O, u64)],
) -> Result<(CostQuote, u64, PermissionBits), QuoteError> {
let mut items = 0_u64;
let mut variable = CostUnits::ZERO;
let mut required = PermissionBits::NONE;
for (op, count) in workload {
if *count == 0 {
continue;
}
items = items.checked_add(*count).ok_or(QuoteError::Overflow)?;
let index = op.index();
let per_item = self.weight_at(index)?;
let entry = per_item.checked_mul(*count).ok_or(QuoteError::Overflow)?;
variable = variable.checked_add(entry).ok_or(QuoteError::Overflow)?;
required = required.union(self.required_at(index));
}
if items == 0 {
return Err(QuoteError::EmptyWorkload);
}
let subtotal = self
.fixed_request
.checked_add(variable)
.ok_or(QuoteError::Overflow)?;
Ok((
CostQuote {
total: subtotal.max(self.minimum_charge),
fixed: self.fixed_request,
variable,
},
items,
required,
))
}
#[inline]
fn weight_at(&self, index: usize) -> Result<CostUnits, QuoteError> {
match self.weights.get(index) {
Some(Some(weight)) => Ok(*weight),
_ => Err(QuoteError::UnknownOperation { index }),
}
}
#[inline]
fn required_at(&self, index: usize) -> PermissionBits {
self.permissions
.get(index)
.copied()
.unwrap_or(PermissionBits::NONE)
}
#[inline]
pub(crate) fn quote_weight(
&self,
per_item: CostUnits,
items: u64,
) -> Result<CostQuote, QuoteError> {
let variable = per_item.checked_mul(items).ok_or(QuoteError::Overflow)?;
let subtotal = self
.fixed_request
.checked_add(variable)
.ok_or(QuoteError::Overflow)?;
Ok(CostQuote {
total: subtotal.max(self.minimum_charge),
fixed: self.fixed_request,
variable,
})
}
pub(crate) fn maximum_weight(&self) -> Option<(usize, CostUnits)> {
let mut maximum = None;
for (index, weight) in self.weights.iter().enumerate() {
let Some(weight) = *weight else {
continue;
};
if maximum.is_none_or(|(_, current)| weight > current) {
maximum = Some((index, weight));
}
}
maximum
}
#[must_use]
pub fn fixed_request(&self) -> CostUnits {
self.fixed_request
}
#[must_use]
pub fn minimum_charge(&self) -> CostUnits {
self.minimum_charge
}
}
#[derive(Debug, Clone)]
pub struct CostTableBuilder {
fixed_request: CostUnits,
minimum_charge: CostUnits,
weights: Vec<Option<CostUnits>>,
permissions: Vec<PermissionBits>,
}
impl CostTableBuilder {
#[must_use]
pub fn weight(self, op: &impl OpIndex, per_item: CostUnits) -> Self {
self.class(op, per_item, PermissionBits::NONE)
}
#[must_use]
pub fn class(
mut self,
op: &impl OpIndex,
per_item: CostUnits,
required: PermissionBits,
) -> Self {
let index = op.index();
if index >= self.weights.len() {
self.weights.resize(index + 1, None);
}
if index >= self.permissions.len() {
self.permissions.resize(index + 1, PermissionBits::NONE);
}
self.weights[index] = Some(per_item);
self.permissions[index] = required;
self
}
#[must_use]
pub fn build(mut self) -> CostTable {
while self.permissions.last() == Some(&PermissionBits::NONE) {
self.permissions.pop();
}
CostTable {
fixed_request: self.fixed_request,
minimum_charge: self.minimum_charge,
weights: self.weights.into_boxed_slice(),
permissions: self.permissions.into_boxed_slice(),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[derive(Clone, Copy)]
enum Op {
Price,
Greeks,
Unpriced,
}
impl OpIndex for Op {
fn index(&self) -> usize {
*self as usize
}
}
fn table() -> CostTable {
CostTable::builder(CostUnits(50), CostUnits(50))
.weight(&Op::Price, CostUnits(1))
.weight(&Op::Greeks, CostUnits(5))
.build()
}
proptest::proptest! {
#[test]
fn quote_agrees_with_the_one_entry_workload(
class in 0_usize..3,
items in 0_u64..=u64::MAX,
) {
let op = [Op::Price, Op::Greeks, Op::Unpriced][class];
let table = table();
let folded = table
.quote_workload(&[(op, items)])
.map(|(quote, _, _)| quote);
proptest::prop_assert_eq!(table.quote(&op, items), folded);
}
}
#[test]
fn quote_is_fixed_plus_weighted_items() {
let q = table().quote(&Op::Greeks, 10).unwrap();
assert_eq!(q.total, CostUnits(100));
assert_eq!(q.fixed, CostUnits(50));
assert_eq!(q.variable, CostUnits(50));
}
#[test]
fn minimum_charge_applies() {
let t = CostTable::builder(CostUnits(0), CostUnits(25))
.weight(&Op::Price, CostUnits(1))
.build();
assert_eq!(t.quote(&Op::Price, 3).unwrap().total, CostUnits(25));
}
#[test]
fn workload_applies_fixed_once_and_sums_repeated_classes() {
let quote = table()
.quote_workload(&[(Op::Price, 2), (Op::Greeks, 3), (Op::Price, 4)])
.unwrap();
assert_eq!(quote.1, 9);
assert_eq!(quote.0.fixed, CostUnits(50));
assert_eq!(quote.0.variable, CostUnits(21));
assert_eq!(quote.0.total, CostUnits(71));
}
#[test]
fn empty_or_all_zero_workload_is_refused() {
assert_eq!(
table().quote_workload::<Op>(&[]),
Err(QuoteError::EmptyWorkload)
);
assert_eq!(
table().quote_workload(&[(Op::Price, 0), (Op::Greeks, 0)]),
Err(QuoteError::EmptyWorkload)
);
assert_eq!(table().quote(&Op::Price, 0), Err(QuoteError::EmptyWorkload));
}
#[test]
fn workload_checks_the_item_sum_and_variable_sum() {
assert_eq!(
table().quote_workload(&[(Op::Price, u64::MAX), (Op::Price, 1)]),
Err(QuoteError::Overflow)
);
let overflowing = CostTable::builder(CostUnits::ZERO, CostUnits::ZERO)
.weight(&Op::Price, CostUnits(u64::MAX))
.weight(&Op::Greeks, CostUnits(1))
.build();
assert_eq!(
overflowing.quote_workload(&[(Op::Price, 1), (Op::Greeks, 1)]),
Err(QuoteError::Overflow)
);
}
#[test]
fn unregistered_operation_denies() {
assert_eq!(
table().quote(&Op::Unpriced, 1),
Err(QuoteError::UnknownOperation { index: 2 })
);
}
#[test]
fn overflow_denies_instead_of_wrapping() {
let t = CostTable::builder(CostUnits(1), CostUnits(0))
.weight(&Op::Price, CostUnits(u64::MAX))
.build();
assert_eq!(t.quote(&Op::Price, 2), Err(QuoteError::Overflow));
assert_eq!(t.quote(&Op::Price, u64::MAX), Err(QuoteError::Overflow));
}
#[test]
fn the_accessors_report_the_schedule_a_quote_applies() {
let table = CostTable::builder(CostUnits(50), CostUnits(80))
.weight(&Op::Price, CostUnits(1))
.build();
assert_eq!(table.fixed_request(), CostUnits(50));
assert_eq!(table.minimum_charge(), CostUnits(80));
let priced = table.quote(&Op::Price, 100).unwrap();
assert_eq!(priced.fixed, table.fixed_request());
assert_eq!(priced.total, CostUnits(150));
let floored = table.quote(&Op::Price, 1).unwrap();
assert_eq!(floored.total, table.minimum_charge());
}
#[test]
fn a_repeated_class_is_summed_not_quoted_twice() {
let table = table();
let (split, split_items, _) = table
.quote_workload(&[(Op::Price, 2), (Op::Price, 3)])
.unwrap();
let (grouped, grouped_items, _) = table.quote_workload(&[(Op::Price, 5)]).unwrap();
assert_eq!(split, grouped, "a repeated class changed the quote");
assert_eq!(split_items, grouped_items);
assert_eq!(split.fixed, table.fixed_request());
}
#[test]
fn work_permissions_are_the_union_of_the_classes_quoted() {
let price = PermissionBits::bit(1);
let greeks = PermissionBits::bit(2);
let table = CostTable::builder(CostUnits(50), CostUnits(50))
.class(&Op::Price, CostUnits(1), price)
.class(&Op::Greeks, CostUnits(5), greeks)
.build();
let (_, _, one) = table.quote_workload(&[(Op::Price, 1)]).unwrap();
assert_eq!(one, price);
let (_, _, both) = table
.quote_workload(&[(Op::Price, 1), (Op::Greeks, 1)])
.unwrap();
assert_eq!(both, price.union(greeks));
let (_, _, skipped) = table
.quote_workload(&[(Op::Price, 1), (Op::Greeks, 0)])
.unwrap();
assert_eq!(skipped, price);
}
#[test]
fn weight_registers_a_class_that_requires_nothing() {
let table = table();
let (_, _, required) = table.quote_workload(&[(Op::Price, 1)]).unwrap();
assert_eq!(required, PermissionBits::NONE);
}
#[cfg(feature = "serde")]
#[test]
fn legacy_cost_table_round_trips_canonically() {
let legacy = r#"{"fixed_request":50,"minimum_charge":50,"weights":[1,5]}"#;
let decoded: CostTable = serde_json::from_str(legacy).expect("legacy table decodes");
assert_eq!(decoded, table());
let reserialized = serde_json::to_string(&decoded).expect("table serializes");
assert_eq!(reserialized, legacy);
}
#[cfg(feature = "serde")]
#[test]
fn an_all_none_permission_array_is_not_serialized() {
let built = table();
let rendered = serde_json::to_string(&built).expect("table serializes");
assert!(
!rendered.contains("permissions"),
"an all-NONE array was serialized: {rendered}"
);
}
#[cfg(feature = "serde")]
#[test]
fn a_table_with_permissions_round_trips() {
let table = CostTable::builder(CostUnits(50), CostUnits(50))
.class(&Op::Price, CostUnits(1), PermissionBits::bit(1))
.class(&Op::Greeks, CostUnits(5), PermissionBits::bit(2))
.build();
let rendered = serde_json::to_string(&table).expect("table serializes");
let decoded: CostTable = serde_json::from_str(&rendered).expect("table decodes");
assert_eq!(decoded, table);
}
}