Skip to main content

tycho_simulation/rfq/protocols/hashflow/
state.rs

1use std::{any::Any, collections::HashMap, fmt};
2
3use async_trait::async_trait;
4use num_bigint::BigUint;
5use num_traits::{FromPrimitive, Pow, ToPrimitive};
6use serde::{Deserialize, Serialize};
7use tycho_common::{
8    dto::ProtocolStateDelta,
9    models::{protocol::GetAmountOutParams, token::Token},
10    simulation::{
11        errors::{SimulationError, TransitionError},
12        indicatively_priced::{IndicativelyPriced, SignedQuote},
13        protocol_sim::{Balances, GetAmountOutResult, ProtocolSim},
14    },
15    Bytes,
16};
17
18use crate::rfq::{
19    client::RFQClient,
20    models::fill_levels,
21    protocols::hashflow::{client::HashflowClient, models::HashflowMarketMakerLevels},
22};
23
24#[derive(Clone, Serialize, Deserialize)]
25pub struct HashflowState {
26    pub base_token: Token,
27    pub quote_token: Token,
28    pub levels: HashflowMarketMakerLevels,
29    pub market_maker: String,
30    pub client: HashflowClient,
31}
32
33impl fmt::Debug for HashflowState {
34    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
35        f.debug_struct("HashflowState")
36            .field("base_token", &self.base_token)
37            .field("quote_token", &self.quote_token)
38            .field("market_maker", &self.market_maker)
39            .finish_non_exhaustive()
40    }
41}
42
43impl HashflowState {
44    pub fn new(
45        base_token: Token,
46        quote_token: Token,
47        levels: HashflowMarketMakerLevels,
48        market_maker: String,
49        client: HashflowClient,
50    ) -> Self {
51        Self { base_token, quote_token, levels, market_maker, client }
52    }
53
54    fn valid_direction_guard(
55        &self,
56        token_address_in: &Bytes,
57        token_address_out: &Bytes,
58    ) -> Result<(), SimulationError> {
59        // The current levels are only valid for the base/quote pair.
60        if !(token_address_in == &self.base_token.address &&
61            token_address_out == &self.quote_token.address)
62        {
63            Err(SimulationError::InvalidInput(
64                format!("Invalid token addresses. Got in={token_address_in}, out={token_address_out}, expected in={}, out={}", self.base_token.address, self.quote_token.address),
65                None,
66            ))
67        } else {
68            Ok(())
69        }
70    }
71
72    fn valid_levels_guard(&self) -> Result<(), SimulationError> {
73        if self.levels.levels.is_empty() {
74            return Err(SimulationError::RecoverableError("No liquidity".into()));
75        }
76        Ok(())
77    }
78}
79
80#[typetag::serde]
81impl ProtocolSim for HashflowState {
82    fn fee(&self) -> f64 {
83        todo!()
84    }
85
86    fn spot_price(&self, base: &Token, quote: &Token) -> Result<f64, SimulationError> {
87        self.valid_direction_guard(&base.address, &quote.address)?;
88
89        // Hashflow's levels are sorted by price, so the first level represents the best price.
90        self.levels
91            .levels
92            .first()
93            .ok_or(SimulationError::RecoverableError("No liquidity".into()))
94            .map(|level| level.price)
95    }
96
97    fn get_amount_out(
98        &self,
99        amount_in: BigUint,
100        token_in: &Token,
101        token_out: &Token,
102    ) -> Result<GetAmountOutResult, SimulationError> {
103        self.valid_direction_guard(&token_in.address, &token_out.address)?;
104        self.valid_levels_guard()?;
105
106        let amount_in = amount_in.to_f64().ok_or_else(|| {
107            SimulationError::RecoverableError("Can't convert amount in to f64".into())
108        })? / 10f64.powi(token_in.decimals as i32);
109
110        // First level represents the minimum amount that can be traded
111        let min_amount = self.levels.levels[0].quantity;
112        if amount_in < min_amount {
113            return Err(SimulationError::RecoverableError(format!(
114                "Amount below minimum. Input amount: {amount_in}, min amount: {min_amount}"
115            )));
116        }
117
118        // Calculate amount out
119        let (amount_out, remaining_amount_in) = fill_levels(&self.levels.levels, amount_in);
120
121        let res = GetAmountOutResult {
122            amount: BigUint::from_f64(amount_out * 10f64.powi(token_out.decimals as i32))
123                .ok_or_else(|| {
124                    SimulationError::RecoverableError("Can't convert amount out to BigUInt".into())
125                })?,
126            gas: BigUint::from(151_000u64), // Rough gas estimation
127            new_state: self.clone_box(),    // The state doesn't change after a swap
128        };
129
130        if remaining_amount_in > 0.0 {
131            return Err(SimulationError::InvalidInput(
132                format!("Pool has not enough liquidity to support complete swap. Input amount: {amount_in}, consumed amount: {}", amount_in-remaining_amount_in),
133                Some(res)));
134        }
135
136        Ok(res)
137    }
138
139    fn get_limits(
140        &self,
141        sell_token: Bytes,
142        buy_token: Bytes,
143    ) -> Result<(BigUint, BigUint), SimulationError> {
144        self.valid_direction_guard(&sell_token, &buy_token)?;
145        self.valid_levels_guard()?;
146
147        let sell_decimals = self.base_token.decimals;
148        let buy_decimals = self.quote_token.decimals;
149        let (total_sell_amount, total_buy_amount) =
150            self.levels
151                .levels
152                .iter()
153                .fold((0.0, 0.0), |(sell_sum, buy_sum), level| {
154                    (sell_sum + level.quantity, buy_sum + level.quantity * level.price)
155                });
156
157        let sell_limit =
158            BigUint::from((total_sell_amount * 10_f64.pow(sell_decimals as f64)) as u128);
159        let buy_limit = BigUint::from((total_buy_amount * 10_f64.pow(buy_decimals as f64)) as u128);
160
161        Ok((sell_limit, buy_limit))
162    }
163
164    fn as_indicatively_priced(&self) -> Result<&dyn IndicativelyPriced, SimulationError> {
165        Ok(self)
166    }
167
168    fn delta_transition(
169        &mut self,
170        _delta: ProtocolStateDelta,
171        _tokens: &HashMap<Bytes, Token>,
172        _balances: &Balances,
173    ) -> Result<(), TransitionError> {
174        todo!()
175    }
176
177    fn clone_box(&self) -> Box<dyn ProtocolSim> {
178        Box::new(self.clone())
179    }
180
181    fn as_any(&self) -> &dyn Any {
182        self
183    }
184
185    fn as_any_mut(&mut self) -> &mut dyn Any {
186        self
187    }
188
189    fn eq(&self, other: &dyn ProtocolSim) -> bool {
190        if let Some(other_state) = other
191            .as_any()
192            .downcast_ref::<HashflowState>()
193        {
194            self.base_token == other_state.base_token &&
195                self.quote_token == other_state.quote_token &&
196                self.levels == other_state.levels
197        } else {
198            false
199        }
200    }
201}
202
203#[async_trait]
204impl IndicativelyPriced for HashflowState {
205    async fn request_signed_quote(
206        &self,
207        params: GetAmountOutParams,
208    ) -> Result<SignedQuote, SimulationError> {
209        Ok(self
210            .client
211            .request_binding_quote(&params)
212            .await?)
213    }
214}
215
216#[cfg(test)]
217mod tests {
218    use std::{collections::HashSet, str::FromStr};
219
220    use tokio::time::Duration;
221    use tycho_common::models::Chain;
222
223    use super::*;
224    use crate::rfq::protocols::hashflow::models::{HashflowPair, HashflowPriceLevel};
225
226    fn wbtc() -> Token {
227        Token::new(
228            &hex::decode("2260fac5e5542a773aa44fbcfedf7c193bc2c599")
229                .unwrap()
230                .into(),
231            "WBTC",
232            8,
233            0,
234            &[Some(10_000)],
235            Chain::Ethereum,
236            100,
237        )
238    }
239
240    fn usdc() -> Token {
241        Token::new(
242            &hex::decode("a0b86991c6218a76c1d19d4a2e9eb0ce3606eb48")
243                .unwrap()
244                .into(),
245            "USDC",
246            6,
247            0,
248            &[Some(10_000)],
249            Chain::Ethereum,
250            100,
251        )
252    }
253
254    fn weth() -> Token {
255        Token::new(
256            &Bytes::from_str("0xc02aaa39b223fe8d0a0e5c4f27ead9083c756cc2").unwrap(),
257            "WETH",
258            18,
259            0,
260            &[],
261            Default::default(),
262            100,
263        )
264    }
265
266    fn empty_hashflow_client() -> HashflowClient {
267        HashflowClient::new(
268            Chain::Ethereum,
269            HashSet::new(),
270            0.0,
271            HashSet::new(),
272            "".to_string(),
273            "".to_string(),
274            Duration::from_secs(0),
275            Duration::from_secs(30),
276        )
277        .unwrap()
278    }
279
280    fn create_test_hashflow_state() -> HashflowState {
281        HashflowState {
282            base_token: weth(),
283            quote_token: usdc(),
284            levels: HashflowMarketMakerLevels {
285                pair: HashflowPair {
286                    base_token: Bytes::from_str("0xc02aaa39b223fe8d0a0e5c4f27ead9083c756cc2")
287                        .unwrap(),
288                    quote_token: Bytes::from_str("0xa0b86991c6218a76c1d19d4a2e9eb0ce3606eb48")
289                        .unwrap(),
290                },
291                levels: vec![
292                    HashflowPriceLevel { quantity: 0.5, price: 3000.0 },
293                    HashflowPriceLevel { quantity: 1.5, price: 3000.0 },
294                    HashflowPriceLevel { quantity: 5.0, price: 2999.0 },
295                ],
296            },
297            market_maker: "test_mm".to_string(),
298            client: empty_hashflow_client(),
299        }
300    }
301
302    mod spot_price {
303        use super::*;
304
305        #[test]
306        fn returns_best_price() {
307            let state = create_test_hashflow_state();
308            let price = state
309                .spot_price(&state.base_token, &state.quote_token)
310                .unwrap();
311            // The best price is the first level's price (3000.0)
312            assert_eq!(price, 3000.0);
313        }
314
315        #[test]
316        fn returns_invalid_input_error() {
317            let state = create_test_hashflow_state();
318            let result = state.spot_price(&wbtc(), &usdc());
319            assert!(result.is_err());
320            if let Err(SimulationError::InvalidInput(msg, _)) = result {
321                assert!(msg.contains("Invalid token addresses"));
322            } else {
323                panic!("Expected InvalidInput");
324            }
325        }
326
327        #[test]
328        fn returns_no_liquidity_error() {
329            let mut state = create_test_hashflow_state();
330            state.levels.levels.clear();
331            let result = state.spot_price(&state.base_token, &state.quote_token);
332            assert!(result.is_err());
333            if let Err(SimulationError::RecoverableError(msg)) = result {
334                assert_eq!(msg, "No liquidity");
335            } else {
336                panic!("Expected RecoverableError");
337            }
338        }
339    }
340
341    mod get_amount_out {
342        use super::*;
343
344        #[test]
345        fn wbtc_to_usdc() {
346            let state = create_test_hashflow_state();
347
348            // Test swapping 1.5 WETH -> USDC
349            // Should consume first level (0.5 WETH at 3000) + partial second level (1.0 WETH at
350            // 3000)
351            let amount_out_result = state
352                .get_amount_out(
353                    BigUint::from_str("1500000000000000000").unwrap(), // 1.5 WETH (18 decimals)
354                    &weth(),
355                    &usdc(),
356                )
357                .unwrap();
358
359            // Expected: (0.5 * 3000) + (1.0 * 3000) = 1500 + 3000 = 4500 USDC
360            assert_eq!(amount_out_result.amount, BigUint::from_str("4500000000").unwrap()); // 6 decimals
361            assert_eq!(amount_out_result.gas, BigUint::from(151_000u64));
362        }
363
364        #[test]
365        fn usdc_to_wbtc() {
366            let state = create_test_hashflow_state();
367
368            // Test swapping 10000 USDC -> WETH
369            // The price levels returned by Hashflow are only valid for the requested pair,
370            // and they can't be inverted to derive the reverse swap.
371            // In that case, we should return an error.
372            let result = state.get_amount_out(
373                BigUint::from_str("10000000000").unwrap(), // 10000 USDC (6 decimals)
374                &usdc(),
375                &weth(),
376            );
377
378            assert!(result.is_err());
379            if let Err(SimulationError::InvalidInput(msg, ..)) = result {
380                assert!(msg.contains("Invalid token addresses"));
381            } else {
382                panic!("Expected InvalidInput");
383            }
384        }
385
386        #[test]
387        fn below_minimum() {
388            let state = create_test_hashflow_state();
389
390            // Test with amount below minimum (first level quantity is 0.5 WETH)
391            let result = state.get_amount_out(
392                BigUint::from_str("250000000000000000").unwrap(), // 0.25 WETH (18 decimals)
393                &weth(),
394                &usdc(),
395            );
396
397            assert!(result.is_err());
398            if let Err(SimulationError::RecoverableError(msg)) = result {
399                assert!(msg.contains("Amount below minimum"));
400            } else {
401                panic!("Expected RecoverableError");
402            }
403        }
404
405        #[test]
406        fn insufficient_liquidity() {
407            let state = create_test_hashflow_state();
408
409            // Test with amount exceeding total liquidity (total is 7.0 WETH)
410            let result = state.get_amount_out(
411                BigUint::from_str("8000000000000000000").unwrap(), // 8.0 WETH (18 decimals)
412                &weth(),
413                &usdc(),
414            );
415
416            assert!(result.is_err());
417            if let Err(SimulationError::InvalidInput(msg, _)) = result {
418                assert!(msg.contains("Pool has not enough liquidity"));
419            } else {
420                panic!("Expected InvalidInput");
421            }
422        }
423
424        #[test]
425        fn invalid_token_pair() {
426            let state = create_test_hashflow_state();
427
428            // Test with invalid token pair (WBTC not in WETH/USDC pool)
429            let result = state.get_amount_out(
430                BigUint::from_str("100000000").unwrap(), // 1 WBTC
431                &wbtc(),
432                &usdc(),
433            );
434
435            assert!(result.is_err());
436            if let Err(SimulationError::InvalidInput(msg, ..)) = result {
437                assert!(msg.contains("Invalid token addresses"));
438            } else {
439                panic!("Expected InvalidInput");
440            }
441        }
442
443        #[test]
444        fn no_liquidity() {
445            let mut state = create_test_hashflow_state();
446            state.levels.levels = vec![]; // Remove all levels
447
448            let result = state.get_amount_out(
449                BigUint::from_str("1000000000000000000").unwrap(), // 1.0 WETH
450                &weth(),
451                &usdc(),
452            );
453
454            assert!(result.is_err());
455            if let Err(SimulationError::RecoverableError(msg)) = result {
456                assert_eq!(msg, "No liquidity");
457            } else {
458                panic!("Expected RecoverableError");
459            }
460        }
461    }
462
463    mod get_limits {
464        use super::*;
465
466        #[test]
467        fn valid_limits() {
468            let state = create_test_hashflow_state();
469            let (sell_limit, buy_limit) = state
470                .get_limits(state.base_token.address.clone(), state.quote_token.address.clone())
471                .unwrap();
472
473            // Total sell: 0.5 + 1.5 + 5.0 = 7.0 WETH (18 decimals)
474            // Total buy: (0.5+1.5)*3000 + 5.0*2999 = 20995 USDC (6 decimals)
475            assert_eq!(sell_limit, BigUint::from((7.0 * 10f64.powi(18)) as u128));
476            assert_eq!(buy_limit, BigUint::from((20995.0 * 10f64.powi(6)) as u128));
477        }
478
479        #[test]
480        fn invalid_token_pair() {
481            let state = create_test_hashflow_state();
482            let result =
483                state.get_limits(wbtc().address.clone(), state.quote_token.address.clone());
484            assert!(result.is_err());
485            if let Err(SimulationError::InvalidInput(msg, _)) = result {
486                assert!(msg.contains("Invalid token addresses"));
487            } else {
488                panic!("Expected InvalidInput");
489            }
490        }
491
492        #[test]
493        fn no_liquidity() {
494            let mut state = create_test_hashflow_state();
495            state.levels.levels = vec![];
496            let result = state
497                .get_limits(state.base_token.address.clone(), state.quote_token.address.clone());
498            assert!(result.is_err());
499            if let Err(SimulationError::RecoverableError(msg)) = result {
500                assert_eq!(msg, "No liquidity");
501            } else {
502                panic!("Expected RecoverableError");
503            }
504        }
505    }
506}