use std::{
any::Any,
collections::HashMap,
time::{Duration, Instant},
};
use num_bigint::BigUint;
use num_traits::{CheckedSub, ToPrimitive};
use serde::{Deserialize, Serialize};
use tycho_common::{
dto::ProtocolStateDelta,
models::token::Token,
simulation::{
errors::{SimulationError, TransitionError},
protocol_sim::{Balances, GetAmountOutResult, ProtocolSim},
},
Bytes,
};
pub const QUOTE_TTL: Duration = super::SLOT;
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct PriceLevelStreamQuote {
pub amount_in: BigUint,
pub amount_out: BigUint,
}
impl PriceLevelStreamQuote {
pub fn new(amount_in: BigUint, amount_out: BigUint) -> Self {
Self { amount_in, amount_out }
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct PriceLevelStreamState {
pub token0: Bytes,
pub token1: Bytes,
pub quotes_0_to_1: Vec<PriceLevelStreamQuote>,
pub quotes_1_to_0: Vec<PriceLevelStreamQuote>,
pub gas_cost: BigUint,
#[serde(skip)]
quotable_until: Option<Instant>,
}
impl PriceLevelStreamState {
pub fn new(
token0: Bytes,
token1: Bytes,
mut quotes_0_to_1: Vec<PriceLevelStreamQuote>,
mut quotes_1_to_0: Vec<PriceLevelStreamQuote>,
gas_cost: BigUint,
) -> Self {
for quotes in [&mut quotes_0_to_1, &mut quotes_1_to_0] {
quotes.sort_by(|a, b| a.amount_in.cmp(&b.amount_in));
quotes.dedup_by(|a, b| a.amount_in == b.amount_in);
}
Self { token0, token1, quotes_0_to_1, quotes_1_to_0, gas_cost, quotable_until: None }
}
pub fn with_quotable_until(mut self, until: Instant) -> Self {
self.quotable_until = Some(until);
self
}
pub fn quotable_until(&self) -> Option<Instant> {
self.quotable_until
}
fn ensure_quotable(&self, now: Instant) -> Result<(), SimulationError> {
match self.quotable_until {
Some(until) if now >= until => Err(SimulationError::RecoverableError(format!(
"price levels expired: the frame that carried them is older than {} s (one slot)",
QUOTE_TTL.as_secs()
))),
Some(_) | None => Ok(()),
}
}
fn quotes(
&self,
token_in: &Bytes,
token_out: &Bytes,
) -> Result<&[PriceLevelStreamQuote], SimulationError> {
if token_in == &self.token0 && token_out == &self.token1 {
Ok(&self.quotes_0_to_1)
} else if token_in == &self.token1 && token_out == &self.token0 {
Ok(&self.quotes_1_to_0)
} else {
Err(SimulationError::RecoverableError(format!(
"Invalid token addresses for pair {}/{}: {token_in}, {token_out}",
self.token0, self.token1
)))
}
}
fn interpolate(
&self,
quotes: &[PriceLevelStreamQuote],
amount_in: &BigUint,
) -> Result<BigUint, SimulationError> {
let idx = quotes.partition_point(|quote| "e.amount_in < amount_in);
let upper = "es[idx];
if &upper.amount_in == amount_in {
return Ok(upper.amount_out.clone());
}
let lower = "es[idx - 1];
let Some(out_span) = upper
.amount_out
.checked_sub(&lower.amount_out)
else {
return Err(SimulationError::RecoverableError(format!(
"Quote ladder {}/{} is not monotonically increasing in amount_out around the \
requested amount {amount_in}: {} -> {}, but {} -> {}",
self.token0,
self.token1,
lower.amount_in,
lower.amount_out,
upper.amount_in,
upper.amount_out,
)));
};
let in_span = &upper.amount_in - &lower.amount_in;
let offset = amount_in - &lower.amount_in;
Ok(&lower.amount_out + out_span * offset / in_span)
}
fn consumed(&self) -> Box<dyn ProtocolSim> {
Box::new(Self {
token0: self.token0.clone(),
token1: self.token1.clone(),
quotes_0_to_1: Vec::new(),
quotes_1_to_0: Vec::new(),
gas_cost: self.gas_cost.clone(),
quotable_until: self.quotable_until,
})
}
}
#[typetag::serde]
impl ProtocolSim for PriceLevelStreamState {
fn fee(&self) -> f64 {
0.0
}
fn spot_price(&self, base: &Token, quote: &Token) -> Result<f64, SimulationError> {
self.ensure_quotable(Instant::now())?;
let quotes = self.quotes(&base.address, "e.address)?;
let best = quotes
.iter()
.find(|q| q.amount_in > BigUint::ZERO && q.amount_out > BigUint::ZERO)
.ok_or_else(|| {
SimulationError::RecoverableError("No liquidity available".to_string())
})?;
let amount_in = best.amount_in.to_f64().ok_or_else(|| {
SimulationError::RecoverableError("Can't convert amount in to f64".to_string())
})?;
let amount_out = best
.amount_out
.to_f64()
.ok_or_else(|| {
SimulationError::RecoverableError("Can't convert amount out to f64".to_string())
})?;
Ok((amount_out / 10f64.powi(quote.decimals as i32)) /
(amount_in / 10f64.powi(base.decimals as i32)))
}
fn get_amount_out(
&self,
amount_in: BigUint,
token_in: &Token,
token_out: &Token,
) -> Result<GetAmountOutResult, SimulationError> {
self.ensure_quotable(Instant::now())?;
let quotes = self.quotes(&token_in.address, &token_out.address)?;
let (Some(first), Some(last)) = (quotes.first(), quotes.last()) else {
return Err(SimulationError::RecoverableError("No liquidity available".to_string()));
};
if amount_in < first.amount_in {
return Err(SimulationError::InvalidInput(
format!(
"Input amount is below the smallest quote. input amount: {amount_in}, minimum quoted amount: {}",
first.amount_in
),
None,
));
}
if amount_in > last.amount_in {
let res = GetAmountOutResult {
amount: last.amount_out.clone(),
gas: self.gas_cost.clone(),
new_state: self.consumed(),
};
return Err(SimulationError::InvalidInput(
format!(
"Not enough liquidity to support complete swap. input amount: {amount_in}, maximum quoted amount: {}",
last.amount_in
),
Some(res),
));
}
Ok(GetAmountOutResult {
amount: self.interpolate(quotes, &amount_in)?,
gas: self.gas_cost.clone(),
new_state: self.consumed(),
})
}
fn get_limits(
&self,
sell_token: Bytes,
buy_token: Bytes,
) -> Result<(BigUint, BigUint), SimulationError> {
self.ensure_quotable(Instant::now())?;
let quotes = self.quotes(&sell_token, &buy_token)?;
match quotes.last() {
Some(largest) => Ok((largest.amount_in.clone(), largest.amount_out.clone())),
None => Ok((BigUint::ZERO, BigUint::ZERO)),
}
}
fn delta_transition(
&mut self,
_delta: ProtocolStateDelta,
_tokens: &HashMap<Bytes, Token>,
_balances: &Balances,
) -> Result<(), TransitionError> {
Err(TransitionError::DecodeError("Not implemented".into()))
}
fn clone_box(&self) -> Box<dyn ProtocolSim> {
Box::new(self.clone())
}
fn as_any(&self) -> &dyn Any {
self
}
fn as_any_mut(&mut self) -> &mut dyn Any {
self
}
fn eq(&self, other: &dyn ProtocolSim) -> bool {
other
.as_any()
.downcast_ref::<PriceLevelStreamState>()
.is_some_and(|other| {
let Self {
token0,
token1,
quotes_0_to_1,
quotes_1_to_0,
gas_cost,
quotable_until: _,
} = other;
&self.token0 == token0 &&
&self.token1 == token1 &&
&self.quotes_0_to_1 == quotes_0_to_1 &&
&self.quotes_1_to_0 == quotes_1_to_0 &&
&self.gas_cost == gas_cost
})
}
}
#[cfg(test)]
mod tests {
use std::time::{Duration, Instant};
use rstest::rstest;
use super::{
super::test_support::{token, USDC, WBTC, WETH},
*,
};
fn wbtc() -> Token {
token(WBTC, "WBTC", 8)
}
fn usdc() -> Token {
token(USDC, "USDC", 6)
}
fn weth() -> Token {
token(WETH, "WETH", 18)
}
fn quote(amount_in: u64, amount_out: u64) -> PriceLevelStreamQuote {
PriceLevelStreamQuote::new(BigUint::from(amount_in), BigUint::from(amount_out))
}
fn state() -> PriceLevelStreamState {
PriceLevelStreamState::new(
wbtc().address,
usdc().address,
vec![quote(100_000_000, 100_000_000_000), quote(200_000_000, 190_000_000_000)],
vec![quote(100_000_000_000, 99_000_000), quote(200_000_000_000, 190_000_000)],
BigUint::from(120_000u64),
)
}
#[test]
fn new_sorts_and_dedups_quotes() {
let state = PriceLevelStreamState::new(
wbtc().address,
usdc().address,
vec![quote(200, 380), quote(100, 200), quote(200, 999)],
vec![],
BigUint::ZERO,
);
assert_eq!(state.quotes_0_to_1, vec![quote(100, 200), quote(200, 380)]);
}
#[test]
fn get_amount_out_exact_level() {
let result = state()
.get_amount_out(BigUint::from(100_000_000u64), &wbtc(), &usdc())
.unwrap();
assert_eq!(result.amount, BigUint::from(100_000_000_000u64));
assert_eq!(result.gas, BigUint::from(120_000u64));
}
#[test]
fn get_amount_out_interpolates_between_levels() {
let result = state()
.get_amount_out(BigUint::from(150_000_000u64), &wbtc(), &usdc())
.unwrap();
assert_eq!(result.amount, BigUint::from(145_000_000_000u64));
}
#[test]
fn get_amount_out_on_glitched_ladder_is_rejected() {
let state = PriceLevelStreamState::new(
wbtc().address,
usdc().address,
vec![quote(100, 200), quote(200, 150)],
vec![],
BigUint::ZERO,
);
let result = state.get_amount_out(BigUint::from(150u64), &wbtc(), &usdc());
assert!(matches!(result, Err(SimulationError::RecoverableError(_))));
let result = state
.get_amount_out(BigUint::from(100u64), &wbtc(), &usdc())
.unwrap();
assert_eq!(result.amount, BigUint::from(200u64));
}
#[test]
fn get_amount_out_below_smallest_level_is_rejected() {
let result = state().get_amount_out(BigUint::from(50_000_000u64), &wbtc(), &usdc());
assert!(matches!(result, Err(SimulationError::InvalidInput(_, None))));
}
#[test]
fn get_amount_out_reverse_direction() {
let result = state()
.get_amount_out(BigUint::from(100_000_000_000u64), &usdc(), &wbtc())
.unwrap();
assert_eq!(result.amount, BigUint::from(99_000_000u64));
}
#[test]
fn get_amount_out_beyond_largest_level_is_partial() {
let result = state().get_amount_out(BigUint::from(300_000_000u64), &wbtc(), &usdc());
match result {
Err(SimulationError::InvalidInput(_, Some(partial))) => {
assert_eq!(partial.amount, BigUint::from(190_000_000_000u64));
}
other => panic!("expected partial InvalidInput, got {other:?}"),
}
}
#[test]
fn get_amount_out_consumes_both_ladders() {
let result = state()
.get_amount_out(BigUint::from(100_000_000u64), &wbtc(), &usdc())
.unwrap();
let new_state = result
.new_state
.as_any()
.downcast_ref::<PriceLevelStreamState>()
.expect("price level state");
assert!(new_state.quotes_0_to_1.is_empty());
assert!(new_state.quotes_1_to_0.is_empty());
}
#[test]
fn get_amount_out_rejects_unknown_tokens() {
let result = state().get_amount_out(BigUint::from(1u64), &weth(), &usdc());
assert!(matches!(result, Err(SimulationError::RecoverableError(_))));
}
#[test]
fn get_amount_out_without_liquidity() {
let state = PriceLevelStreamState::new(
wbtc().address,
usdc().address,
vec![],
vec![],
BigUint::ZERO,
);
let result = state.get_amount_out(BigUint::from(1u64), &wbtc(), &usdc());
assert!(matches!(result, Err(SimulationError::RecoverableError(_))));
}
#[test]
fn spot_price_uses_smallest_quote() {
let price = state()
.spot_price(&wbtc(), &usdc())
.unwrap();
assert!((price - 100_000.0).abs() < 1e-9);
let inverse = state()
.spot_price(&usdc(), &wbtc())
.unwrap();
assert!((inverse - 9.9e-6).abs() < 1e-15);
}
#[test]
fn spot_price_skips_zero_amount_out_quotes() {
let state = PriceLevelStreamState::new(
wbtc().address,
usdc().address,
vec![quote(1, 0), quote(100_000_000, 100_000_000_000)],
vec![],
BigUint::ZERO,
);
let price = state
.spot_price(&wbtc(), &usdc())
.unwrap();
assert!((price - 100_000.0).abs() < 1e-9);
}
#[test]
fn get_limits_returns_largest_quote() {
let (max_in, max_out) = state()
.get_limits(wbtc().address, usdc().address)
.unwrap();
assert_eq!(max_in, BigUint::from(200_000_000u64));
assert_eq!(max_out, BigUint::from(190_000_000_000u64));
}
#[test]
fn get_limits_without_liquidity() {
let state = PriceLevelStreamState::new(
wbtc().address,
usdc().address,
vec![],
vec![],
BigUint::ZERO,
);
let (max_in, max_out) = state
.get_limits(wbtc().address, usdc().address)
.unwrap();
assert_eq!(max_in, BigUint::ZERO);
assert_eq!(max_out, BigUint::ZERO);
}
#[test]
fn eq_compares_quotes() {
let a = state();
let mut b = state();
assert!(a.eq(&b as &dyn ProtocolSim));
b.quotes_0_to_1[0].amount_out += 1u32;
assert!(!a.eq(&b as &dyn ProtocolSim));
}
#[rstest]
#[case::one_nanosecond_before(|until| until - Duration::from_nanos(1), true)]
#[case::at_quotable_until(|until| until, false)]
#[case::one_second_after(|until| until + Duration::from_secs(1), false)]
fn ensure_quotable_around_quotable_until(
#[case] now: fn(Instant) -> Instant,
#[case] quotable: bool,
) {
let until = Instant::now() + Duration::from_secs(60);
let state = state().with_quotable_until(until);
assert_eq!(
state
.ensure_quotable(now(until))
.is_ok(),
quotable
);
}
#[test]
fn state_without_quotable_until_never_expires() {
let far = Instant::now() + Duration::from_secs(1_000_000);
assert!(state().ensure_quotable(far).is_ok());
}
#[test]
fn expired_state_refuses_every_query() {
let state = state().with_quotable_until(Instant::now());
assert!(matches!(
state.get_amount_out(BigUint::from(100_000_000u64), &wbtc(), &usdc()),
Err(SimulationError::RecoverableError(_))
));
assert!(matches!(
state.spot_price(&wbtc(), &usdc()),
Err(SimulationError::RecoverableError(_))
));
assert!(matches!(
state.get_limits(wbtc().address, usdc().address),
Err(SimulationError::RecoverableError(_))
));
}
#[test]
fn fresh_state_answers_every_query() {
let state = state().with_quotable_until(Instant::now() + Duration::from_secs(60));
assert!(state
.get_amount_out(BigUint::from(100_000_000u64), &wbtc(), &usdc())
.is_ok());
assert!(state
.spot_price(&wbtc(), &usdc())
.is_ok());
assert!(state
.get_limits(wbtc().address, usdc().address)
.is_ok());
}
#[test]
fn successor_state_keeps_quotable_until() {
let until = Instant::now() + Duration::from_secs(60);
let state = state().with_quotable_until(until);
let result = state
.get_amount_out(BigUint::from(100_000_000u64), &wbtc(), &usdc())
.expect("fresh state quotes");
let successor = result
.new_state
.as_any()
.downcast_ref::<PriceLevelStreamState>()
.expect("price level state");
assert_eq!(successor.quotable_until, Some(until));
}
#[test]
fn quotable_until_serde_round_trip() {
let live = state().with_quotable_until(Instant::now());
let json = serde_json::to_value(&live).unwrap();
assert!(json
.as_object()
.unwrap()
.get("quotable_until")
.is_none());
let replayed: PriceLevelStreamState = serde_json::from_value(json).unwrap();
assert_eq!(replayed.quotable_until, None);
assert!(replayed
.spot_price(&wbtc(), &usdc())
.is_ok());
}
#[test]
fn eq_ignores_quotable_until() {
let until = Instant::now() + Duration::from_secs(60);
assert!(state().eq(&state().with_quotable_until(until)));
assert!(state()
.with_quotable_until(until)
.eq(&state().with_quotable_until(until + Duration::from_secs(1))));
}
}