1use std::{
2 any::Any,
3 collections::HashMap,
4 time::{Duration, Instant},
5};
6
7use num_bigint::BigUint;
8use num_traits::{CheckedSub, ToPrimitive};
9use serde::{Deserialize, Serialize};
10use tycho_common::{
11 dto::ProtocolStateDelta,
12 models::token::Token,
13 simulation::{
14 errors::{SimulationError, TransitionError},
15 protocol_sim::{Balances, GetAmountOutResult, ProtocolSim},
16 },
17 Bytes,
18};
19
20pub const QUOTE_TTL: Duration = super::SLOT;
24
25#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
30pub struct PriceLevelStreamQuote {
31 pub amount_in: BigUint,
32 pub amount_out: BigUint,
33}
34
35impl PriceLevelStreamQuote {
36 pub fn new(amount_in: BigUint, amount_out: BigUint) -> Self {
37 Self { amount_in, amount_out }
38 }
39}
40
41#[derive(Debug, Clone, Serialize, Deserialize)]
48pub struct PriceLevelStreamState {
49 pub token0: Bytes,
50 pub token1: Bytes,
51 pub quotes_0_to_1: Vec<PriceLevelStreamQuote>,
52 pub quotes_1_to_0: Vec<PriceLevelStreamQuote>,
53 pub gas_cost: BigUint,
54 #[serde(skip)]
60 quotable_until: Option<Instant>,
61}
62
63impl PriceLevelStreamState {
64 pub fn new(
69 token0: Bytes,
70 token1: Bytes,
71 mut quotes_0_to_1: Vec<PriceLevelStreamQuote>,
72 mut quotes_1_to_0: Vec<PriceLevelStreamQuote>,
73 gas_cost: BigUint,
74 ) -> Self {
75 for quotes in [&mut quotes_0_to_1, &mut quotes_1_to_0] {
76 quotes.sort_by(|a, b| a.amount_in.cmp(&b.amount_in));
77 quotes.dedup_by(|a, b| a.amount_in == b.amount_in);
78 }
79 Self { token0, token1, quotes_0_to_1, quotes_1_to_0, gas_cost, quotable_until: None }
80 }
81
82 pub fn with_quotable_until(mut self, until: Instant) -> Self {
84 self.quotable_until = Some(until);
85 self
86 }
87
88 pub fn quotable_until(&self) -> Option<Instant> {
91 self.quotable_until
92 }
93
94 fn ensure_quotable(&self, now: Instant) -> Result<(), SimulationError> {
97 match self.quotable_until {
98 Some(until) if now >= until => Err(SimulationError::RecoverableError(format!(
99 "price levels expired: the frame that carried them is older than {} s (one slot)",
100 QUOTE_TTL.as_secs()
101 ))),
102 Some(_) | None => Ok(()),
103 }
104 }
105
106 fn quotes(
109 &self,
110 token_in: &Bytes,
111 token_out: &Bytes,
112 ) -> Result<&[PriceLevelStreamQuote], SimulationError> {
113 if token_in == &self.token0 && token_out == &self.token1 {
114 Ok(&self.quotes_0_to_1)
115 } else if token_in == &self.token1 && token_out == &self.token0 {
116 Ok(&self.quotes_1_to_0)
117 } else {
118 Err(SimulationError::RecoverableError(format!(
119 "Invalid token addresses for pair {}/{}: {token_in}, {token_out}",
120 self.token0, self.token1
121 )))
122 }
123 }
124
125 fn interpolate(
133 &self,
134 quotes: &[PriceLevelStreamQuote],
135 amount_in: &BigUint,
136 ) -> Result<BigUint, SimulationError> {
137 let idx = quotes.partition_point(|quote| "e.amount_in < amount_in);
140 let upper = "es[idx];
141 if &upper.amount_in == amount_in {
142 return Ok(upper.amount_out.clone());
143 }
144 let lower = "es[idx - 1];
145 let Some(out_span) = upper
146 .amount_out
147 .checked_sub(&lower.amount_out)
148 else {
149 return Err(SimulationError::RecoverableError(format!(
151 "Quote ladder {}/{} is not monotonically increasing in amount_out around the \
152 requested amount {amount_in}: {} -> {}, but {} -> {}",
153 self.token0,
154 self.token1,
155 lower.amount_in,
156 lower.amount_out,
157 upper.amount_in,
158 upper.amount_out,
159 )));
160 };
161 let in_span = &upper.amount_in - &lower.amount_in;
162 let offset = amount_in - &lower.amount_in;
163 Ok(&lower.amount_out + out_span * offset / in_span)
164 }
165
166 fn consumed(&self) -> Box<dyn ProtocolSim> {
170 Box::new(Self {
171 token0: self.token0.clone(),
172 token1: self.token1.clone(),
173 quotes_0_to_1: Vec::new(),
174 quotes_1_to_0: Vec::new(),
175 gas_cost: self.gas_cost.clone(),
176 quotable_until: self.quotable_until,
177 })
178 }
179}
180
181#[typetag::serde]
182impl ProtocolSim for PriceLevelStreamState {
183 fn fee(&self) -> f64 {
184 0.0
185 }
186
187 fn spot_price(&self, base: &Token, quote: &Token) -> Result<f64, SimulationError> {
188 self.ensure_quotable(Instant::now())?;
189 let quotes = self.quotes(&base.address, "e.address)?;
190 let best = quotes
191 .iter()
192 .find(|q| q.amount_in > BigUint::ZERO && q.amount_out > BigUint::ZERO)
193 .ok_or_else(|| {
194 SimulationError::RecoverableError("No liquidity available".to_string())
195 })?;
196 let amount_in = best.amount_in.to_f64().ok_or_else(|| {
197 SimulationError::RecoverableError("Can't convert amount in to f64".to_string())
198 })?;
199 let amount_out = best
200 .amount_out
201 .to_f64()
202 .ok_or_else(|| {
203 SimulationError::RecoverableError("Can't convert amount out to f64".to_string())
204 })?;
205 Ok((amount_out / 10f64.powi(quote.decimals as i32)) /
206 (amount_in / 10f64.powi(base.decimals as i32)))
207 }
208
209 fn get_amount_out(
210 &self,
211 amount_in: BigUint,
212 token_in: &Token,
213 token_out: &Token,
214 ) -> Result<GetAmountOutResult, SimulationError> {
215 self.ensure_quotable(Instant::now())?;
216 let quotes = self.quotes(&token_in.address, &token_out.address)?;
217 let (Some(first), Some(last)) = (quotes.first(), quotes.last()) else {
218 return Err(SimulationError::RecoverableError("No liquidity available".to_string()));
219 };
220 if amount_in < first.amount_in {
234 return Err(SimulationError::InvalidInput(
235 format!(
236 "Input amount is below the smallest quote. input amount: {amount_in}, minimum quoted amount: {}",
237 first.amount_in
238 ),
239 None,
240 ));
241 }
242 if amount_in > last.amount_in {
245 let res = GetAmountOutResult {
246 amount: last.amount_out.clone(),
247 gas: self.gas_cost.clone(),
248 new_state: self.consumed(),
249 };
250 return Err(SimulationError::InvalidInput(
251 format!(
252 "Not enough liquidity to support complete swap. input amount: {amount_in}, maximum quoted amount: {}",
253 last.amount_in
254 ),
255 Some(res),
256 ));
257 }
258 Ok(GetAmountOutResult {
259 amount: self.interpolate(quotes, &amount_in)?,
260 gas: self.gas_cost.clone(),
261 new_state: self.consumed(),
262 })
263 }
264
265 fn get_limits(
266 &self,
267 sell_token: Bytes,
268 buy_token: Bytes,
269 ) -> Result<(BigUint, BigUint), SimulationError> {
270 self.ensure_quotable(Instant::now())?;
271 let quotes = self.quotes(&sell_token, &buy_token)?;
272 match quotes.last() {
273 Some(largest) => Ok((largest.amount_in.clone(), largest.amount_out.clone())),
274 None => Ok((BigUint::ZERO, BigUint::ZERO)),
275 }
276 }
277
278 fn delta_transition(
279 &mut self,
280 _delta: ProtocolStateDelta,
281 _tokens: &HashMap<Bytes, Token>,
282 _balances: &Balances,
283 ) -> Result<(), TransitionError> {
284 Err(TransitionError::DecodeError("Not implemented".into()))
285 }
286
287 fn clone_box(&self) -> Box<dyn ProtocolSim> {
288 Box::new(self.clone())
289 }
290
291 fn as_any(&self) -> &dyn Any {
292 self
293 }
294
295 fn as_any_mut(&mut self) -> &mut dyn Any {
296 self
297 }
298
299 fn eq(&self, other: &dyn ProtocolSim) -> bool {
303 other
304 .as_any()
305 .downcast_ref::<PriceLevelStreamState>()
306 .is_some_and(|other| {
307 let Self {
308 token0,
309 token1,
310 quotes_0_to_1,
311 quotes_1_to_0,
312 gas_cost,
313 quotable_until: _,
314 } = other;
315 &self.token0 == token0 &&
316 &self.token1 == token1 &&
317 &self.quotes_0_to_1 == quotes_0_to_1 &&
318 &self.quotes_1_to_0 == quotes_1_to_0 &&
319 &self.gas_cost == gas_cost
320 })
321 }
322}
323
324#[cfg(test)]
325mod tests {
326 use std::time::{Duration, Instant};
327
328 use rstest::rstest;
329
330 use super::{
331 super::test_support::{token, USDC, WBTC, WETH},
332 *,
333 };
334
335 fn wbtc() -> Token {
336 token(WBTC, "WBTC", 8)
337 }
338
339 fn usdc() -> Token {
340 token(USDC, "USDC", 6)
341 }
342
343 fn weth() -> Token {
344 token(WETH, "WETH", 18)
345 }
346
347 fn quote(amount_in: u64, amount_out: u64) -> PriceLevelStreamQuote {
348 PriceLevelStreamQuote::new(BigUint::from(amount_in), BigUint::from(amount_out))
349 }
350
351 fn state() -> PriceLevelStreamState {
354 PriceLevelStreamState::new(
355 wbtc().address,
356 usdc().address,
357 vec![quote(100_000_000, 100_000_000_000), quote(200_000_000, 190_000_000_000)],
358 vec![quote(100_000_000_000, 99_000_000), quote(200_000_000_000, 190_000_000)],
359 BigUint::from(120_000u64),
360 )
361 }
362
363 #[test]
364 fn new_sorts_and_dedups_quotes() {
365 let state = PriceLevelStreamState::new(
366 wbtc().address,
367 usdc().address,
368 vec![quote(200, 380), quote(100, 200), quote(200, 999)],
369 vec![],
370 BigUint::ZERO,
371 );
372 assert_eq!(state.quotes_0_to_1, vec![quote(100, 200), quote(200, 380)]);
373 }
374
375 #[test]
376 fn get_amount_out_exact_level() {
377 let result = state()
378 .get_amount_out(BigUint::from(100_000_000u64), &wbtc(), &usdc())
379 .unwrap();
380 assert_eq!(result.amount, BigUint::from(100_000_000_000u64));
381 assert_eq!(result.gas, BigUint::from(120_000u64));
382 }
383
384 #[test]
385 fn get_amount_out_interpolates_between_levels() {
386 let result = state()
388 .get_amount_out(BigUint::from(150_000_000u64), &wbtc(), &usdc())
389 .unwrap();
390 assert_eq!(result.amount, BigUint::from(145_000_000_000u64));
391 }
392
393 #[test]
394 fn get_amount_out_on_glitched_ladder_is_rejected() {
395 let state = PriceLevelStreamState::new(
399 wbtc().address,
400 usdc().address,
401 vec![quote(100, 200), quote(200, 150)],
402 vec![],
403 BigUint::ZERO,
404 );
405 let result = state.get_amount_out(BigUint::from(150u64), &wbtc(), &usdc());
406 assert!(matches!(result, Err(SimulationError::RecoverableError(_))));
407
408 let result = state
410 .get_amount_out(BigUint::from(100u64), &wbtc(), &usdc())
411 .unwrap();
412 assert_eq!(result.amount, BigUint::from(200u64));
413 }
414
415 #[test]
416 fn get_amount_out_below_smallest_level_is_rejected() {
417 let result = state().get_amount_out(BigUint::from(50_000_000u64), &wbtc(), &usdc());
420 assert!(matches!(result, Err(SimulationError::InvalidInput(_, None))));
421 }
422
423 #[test]
424 fn get_amount_out_reverse_direction() {
425 let result = state()
426 .get_amount_out(BigUint::from(100_000_000_000u64), &usdc(), &wbtc())
427 .unwrap();
428 assert_eq!(result.amount, BigUint::from(99_000_000u64));
429 }
430
431 #[test]
432 fn get_amount_out_beyond_largest_level_is_partial() {
433 let result = state().get_amount_out(BigUint::from(300_000_000u64), &wbtc(), &usdc());
434 match result {
435 Err(SimulationError::InvalidInput(_, Some(partial))) => {
436 assert_eq!(partial.amount, BigUint::from(190_000_000_000u64));
437 }
438 other => panic!("expected partial InvalidInput, got {other:?}"),
439 }
440 }
441
442 #[test]
443 fn get_amount_out_consumes_both_ladders() {
444 let result = state()
445 .get_amount_out(BigUint::from(100_000_000u64), &wbtc(), &usdc())
446 .unwrap();
447 let new_state = result
448 .new_state
449 .as_any()
450 .downcast_ref::<PriceLevelStreamState>()
451 .expect("price level state");
452 assert!(new_state.quotes_0_to_1.is_empty());
453 assert!(new_state.quotes_1_to_0.is_empty());
454 }
455
456 #[test]
457 fn get_amount_out_rejects_unknown_tokens() {
458 let result = state().get_amount_out(BigUint::from(1u64), &weth(), &usdc());
459 assert!(matches!(result, Err(SimulationError::RecoverableError(_))));
460 }
461
462 #[test]
463 fn get_amount_out_without_liquidity() {
464 let state = PriceLevelStreamState::new(
465 wbtc().address,
466 usdc().address,
467 vec![],
468 vec![],
469 BigUint::ZERO,
470 );
471 let result = state.get_amount_out(BigUint::from(1u64), &wbtc(), &usdc());
472 assert!(matches!(result, Err(SimulationError::RecoverableError(_))));
473 }
474
475 #[test]
476 fn spot_price_uses_smallest_quote() {
477 let price = state()
479 .spot_price(&wbtc(), &usdc())
480 .unwrap();
481 assert!((price - 100_000.0).abs() < 1e-9);
482
483 let inverse = state()
484 .spot_price(&usdc(), &wbtc())
485 .unwrap();
486 assert!((inverse - 9.9e-6).abs() < 1e-15);
488 }
489
490 #[test]
491 fn spot_price_skips_zero_amount_out_quotes() {
492 let state = PriceLevelStreamState::new(
495 wbtc().address,
496 usdc().address,
497 vec![quote(1, 0), quote(100_000_000, 100_000_000_000)],
498 vec![],
499 BigUint::ZERO,
500 );
501 let price = state
502 .spot_price(&wbtc(), &usdc())
503 .unwrap();
504 assert!((price - 100_000.0).abs() < 1e-9);
505 }
506
507 #[test]
508 fn get_limits_returns_largest_quote() {
509 let (max_in, max_out) = state()
510 .get_limits(wbtc().address, usdc().address)
511 .unwrap();
512 assert_eq!(max_in, BigUint::from(200_000_000u64));
513 assert_eq!(max_out, BigUint::from(190_000_000_000u64));
514 }
515
516 #[test]
517 fn get_limits_without_liquidity() {
518 let state = PriceLevelStreamState::new(
519 wbtc().address,
520 usdc().address,
521 vec![],
522 vec![],
523 BigUint::ZERO,
524 );
525 let (max_in, max_out) = state
526 .get_limits(wbtc().address, usdc().address)
527 .unwrap();
528 assert_eq!(max_in, BigUint::ZERO);
529 assert_eq!(max_out, BigUint::ZERO);
530 }
531
532 #[test]
533 fn eq_compares_quotes() {
534 let a = state();
535 let mut b = state();
536 assert!(a.eq(&b as &dyn ProtocolSim));
537 b.quotes_0_to_1[0].amount_out += 1u32;
538 assert!(!a.eq(&b as &dyn ProtocolSim));
539 }
540
541 #[rstest]
542 #[case::one_nanosecond_before(|until| until - Duration::from_nanos(1), true)]
543 #[case::at_quotable_until(|until| until, false)]
544 #[case::one_second_after(|until| until + Duration::from_secs(1), false)]
545 fn ensure_quotable_around_quotable_until(
546 #[case] now: fn(Instant) -> Instant,
547 #[case] quotable: bool,
548 ) {
549 let until = Instant::now() + Duration::from_secs(60);
550 let state = state().with_quotable_until(until);
551 assert_eq!(
552 state
553 .ensure_quotable(now(until))
554 .is_ok(),
555 quotable
556 );
557 }
558
559 #[test]
560 fn state_without_quotable_until_never_expires() {
561 let far = Instant::now() + Duration::from_secs(1_000_000);
562 assert!(state().ensure_quotable(far).is_ok());
563 }
564
565 #[test]
566 fn expired_state_refuses_every_query() {
567 let state = state().with_quotable_until(Instant::now());
570 assert!(matches!(
571 state.get_amount_out(BigUint::from(100_000_000u64), &wbtc(), &usdc()),
572 Err(SimulationError::RecoverableError(_))
573 ));
574 assert!(matches!(
575 state.spot_price(&wbtc(), &usdc()),
576 Err(SimulationError::RecoverableError(_))
577 ));
578 assert!(matches!(
579 state.get_limits(wbtc().address, usdc().address),
580 Err(SimulationError::RecoverableError(_))
581 ));
582 }
583
584 #[test]
585 fn fresh_state_answers_every_query() {
586 let state = state().with_quotable_until(Instant::now() + Duration::from_secs(60));
587 assert!(state
588 .get_amount_out(BigUint::from(100_000_000u64), &wbtc(), &usdc())
589 .is_ok());
590 assert!(state
591 .spot_price(&wbtc(), &usdc())
592 .is_ok());
593 assert!(state
594 .get_limits(wbtc().address, usdc().address)
595 .is_ok());
596 }
597
598 #[test]
599 fn successor_state_keeps_quotable_until() {
600 let until = Instant::now() + Duration::from_secs(60);
601 let state = state().with_quotable_until(until);
602 let result = state
603 .get_amount_out(BigUint::from(100_000_000u64), &wbtc(), &usdc())
604 .expect("fresh state quotes");
605 let successor = result
606 .new_state
607 .as_any()
608 .downcast_ref::<PriceLevelStreamState>()
609 .expect("price level state");
610 assert_eq!(successor.quotable_until, Some(until));
611 }
612
613 #[test]
614 fn quotable_until_serde_round_trip() {
615 let live = state().with_quotable_until(Instant::now());
616 let json = serde_json::to_value(&live).unwrap();
617 assert!(json
618 .as_object()
619 .unwrap()
620 .get("quotable_until")
621 .is_none());
622 let replayed: PriceLevelStreamState = serde_json::from_value(json).unwrap();
624 assert_eq!(replayed.quotable_until, None);
625 assert!(replayed
626 .spot_price(&wbtc(), &usdc())
627 .is_ok());
628 }
629
630 #[test]
632 fn eq_ignores_quotable_until() {
633 let until = Instant::now() + Duration::from_secs(60);
634 assert!(state().eq(&state().with_quotable_until(until)));
635 assert!(state()
636 .with_quotable_until(until)
637 .eq(&state().with_quotable_until(until + Duration::from_secs(1))));
638 }
639}