use derive_more::Deref;
use itertools::Itertools as _;
use serde::{Deserialize, Deserializer, Serialize};
use thiserror::Error;
const FILL_TOLERANCE: f64 = 4.0 * f64::EPSILON;
#[derive(Clone, Copy, Debug, PartialEq, Serialize, Deserialize)]
pub struct PriceLevel {
pub price: f64,
pub quantity: f64,
}
impl PriceLevel {
fn fillable(self) -> Result<Option<Self>, InvalidLevel> {
if !self.quantity.is_finite() || self.quantity < 0.0 {
return Err(InvalidLevel::Quantity(self.quantity));
}
if self.quantity == 0.0 {
return Ok(None);
}
if !self.price.is_finite() || self.price <= 0.0 {
return Err(InvalidLevel::Price(self.price));
}
Ok(Some(self))
}
}
#[derive(Clone, Copy, Debug, PartialEq, Error)]
pub enum InvalidLevel {
#[error("price level quantity must be a non-negative finite number, got {0}")]
Quantity(f64),
#[error("price level price must be a positive finite number, got {0}")]
Price(f64),
}
#[derive(Clone, Debug, Default, Deref, PartialEq, Serialize)]
#[serde(transparent)]
#[deref(forward)]
pub struct Levels(Vec<PriceLevel>);
impl Levels {
pub fn new(levels: Vec<PriceLevel>) -> Result<Self, InvalidLevel> {
let mut levels: Vec<_> = levels
.into_iter()
.filter_map(|level| level.fillable().transpose())
.try_collect()?;
levels.shrink_to_fit();
Ok(Levels(levels))
}
pub fn fill(&self, amount_in: f64) -> Fill {
if amount_in.is_nan() || amount_in <= 0.0 {
return Fill { amount_out: 0.0, remaining_in: amount_in };
}
let mut consumed = 0.0;
let mut amount_out = 0.0;
for level in &self.0 {
let taken = (amount_in - consumed).min(level.quantity);
amount_out += taken * level.price;
consumed += taken;
if consumed >= amount_in {
break;
}
}
let remaining_in = amount_in - consumed;
let remaining_in = if remaining_in.is_finite() && remaining_in <= amount_in * FILL_TOLERANCE
{
0.0
} else {
remaining_in
};
Fill { amount_out, remaining_in }
}
pub fn invert(&self) -> Levels {
Levels(
self.0
.iter()
.filter_map(|level| {
PriceLevel { price: 1.0 / level.price, quantity: level.quantity * level.price }
.fillable()
.ok()
.flatten()
})
.collect(),
)
}
pub fn notional(&self) -> f64 {
self.totals().1
}
pub fn totals(&self) -> (f64, f64) {
self.0
.iter()
.fold((0.0, 0.0), |(quantity, notional), level| {
(quantity + level.quantity, notional + level.price * level.quantity)
})
}
pub fn average_price(&self, amount_in: f64) -> Option<f64> {
let fill = self.fill(amount_in);
let consumed = amount_in - fill.remaining_in;
(consumed > 0.0).then(|| fill.amount_out / consumed)
}
}
impl<'de> Deserialize<'de> for Levels {
fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
Levels::new(Vec::deserialize(deserializer)?).map_err(serde::de::Error::custom)
}
}
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct Fill {
pub amount_out: f64,
pub remaining_in: f64,
}
impl Fill {
pub fn is_complete(&self) -> bool {
self.remaining_in <= 0.0
}
}
#[cfg(test)]
mod tests {
use rstest::rstest;
use super::*;
fn ladder() -> Levels {
Levels::new(vec![
PriceLevel { price: 3000.0, quantity: 1.0 },
PriceLevel { price: 2999.0, quantity: 2.0 },
])
.unwrap()
}
#[test]
fn new_drops_zero_quantity_placeholders() {
let levels = Levels::new(vec![
PriceLevel { price: 0.0, quantity: 0.0 },
PriceLevel { price: 3000.0, quantity: 1.0 },
PriceLevel { price: -1.0, quantity: 0.0 },
])
.unwrap();
assert_eq!(&*levels, &[PriceLevel { price: 3000.0, quantity: 1.0 }]);
}
#[test]
fn new_rejects_invalid_levels() {
assert_eq!(
Levels::new(vec![PriceLevel { price: 1.0, quantity: -1.0 }]),
Err(InvalidLevel::Quantity(-1.0))
);
assert!(matches!(
Levels::new(vec![PriceLevel { price: 1.0, quantity: f64::NAN }]),
Err(InvalidLevel::Quantity(_))
));
assert_eq!(
Levels::new(vec![PriceLevel { price: 0.0, quantity: 1.0 }]),
Err(InvalidLevel::Price(0.0))
);
assert_eq!(
Levels::new(vec![PriceLevel { price: f64::INFINITY, quantity: 1.0 }]),
Err(InvalidLevel::Price(f64::INFINITY))
);
}
#[test]
fn deserialize_validates_and_serialize_is_transparent() {
let levels: Levels = serde_json::from_str(
r#"[{"price":3000.0,"quantity":1.0},{"price":1.0,"quantity":0.0}]"#,
)
.unwrap();
assert_eq!(levels.len(), 1);
assert_eq!(serde_json::to_string(&levels).unwrap(), r#"[{"price":3000.0,"quantity":1.0}]"#);
assert!(serde_json::from_str::<Levels>(r#"[{"price":0.0,"quantity":1.0}]"#).is_err());
}
#[test]
fn fill_walks_levels_in_order_and_reports_the_unfilled_rest() {
let ladder = ladder();
assert_eq!(ladder.fill(1.0), Fill { amount_out: 3000.0, remaining_in: 0.0 });
assert_eq!(ladder.fill(2.0), Fill { amount_out: 5999.0, remaining_in: 0.0 });
let partial = ladder.fill(5.0);
assert_eq!(partial, Fill { amount_out: 8998.0, remaining_in: 2.0 });
assert!(!partial.is_complete());
}
#[rstest]
#[case::nothing(0.0)]
#[case::less_than_nothing(-1.0)]
#[case::not_a_number(f64::NAN)]
fn fill_of_an_amount_that_is_not_positive_takes_nothing(#[case] amount_in: f64) {
let fill = ladder().fill(amount_in);
assert_eq!(fill.amount_out, 0.0);
assert!(
fill.remaining_in.to_bits() == amount_in.to_bits(),
"the amount comes back untouched"
);
}
#[test]
fn fill_of_an_infinite_amount_takes_the_whole_ladder() {
let fill = ladder().fill(f64::INFINITY);
assert_eq!(fill.amount_out, 8998.0);
assert_eq!(fill.remaining_in, f64::INFINITY);
assert!(!fill.is_complete());
}
#[test]
fn filling_the_whole_ladder_leaves_nothing() {
let quantities = [1.1, 2.2, 3.3, 4.4, 5.5, 6.6, 7.7, 8.8, 9.9, 48.508172];
let ladder = Levels::new(
quantities
.iter()
.map(|&quantity| PriceLevel { price: 2.0, quantity })
.collect(),
)
.unwrap();
let (depth, _) = ladder.totals();
let fill = ladder.fill(depth);
assert_eq!(fill.remaining_in, 0.0);
assert!(fill.is_complete());
}
#[test]
fn filling_past_the_ladder_still_reports_the_shortfall() {
let ladder = ladder();
let (depth, _) = ladder.totals();
let fill = ladder.fill(depth + 0.001);
assert!((fill.remaining_in - 0.001).abs() < 1e-12, "got {}", fill.remaining_in);
assert!(!fill.is_complete());
}
#[test]
fn invert_swaps_units_and_drops_overflowing_levels() {
let inverted = Levels::new(vec![
PriceLevel { price: 0.11, quantity: 3000.0 },
PriceLevel { price: 0.12, quantity: 3000.0 },
PriceLevel { price: 1e308, quantity: 1e308 },
])
.unwrap()
.invert();
assert_eq!(inverted.len(), 2);
assert!((inverted[0].price - 9.090909090909092).abs() < 1e-9);
assert!((inverted[0].quantity - 330.0).abs() < 1e-9);
assert!((inverted[1].price - 8.333333333333334).abs() < 1e-9);
assert!((inverted[1].quantity - 360.0).abs() < 1e-9);
}
#[test]
fn notional_and_totals_sum_the_ladder() {
let ladder = ladder();
assert_eq!(ladder.notional(), 8998.0);
assert_eq!(ladder.totals(), (3.0, 8998.0));
}
#[test]
fn average_price_is_per_consumed_unit() {
let ladder = ladder();
assert_eq!(ladder.average_price(1.0), Some(3000.0));
assert_eq!(ladder.average_price(2.0), Some(2999.5));
assert_eq!(ladder.average_price(5.0), Some(8998.0 / 3.0));
assert_eq!(ladder.average_price(0.0), None);
assert_eq!(Levels::default().average_price(1.0), None);
}
}