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 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, "e.address)?;
88
89 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 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 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), new_state: self.clone_box(), };
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(¶ms)
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 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 let amount_out_result = state
352 .get_amount_out(
353 BigUint::from_str("1500000000000000000").unwrap(), &weth(),
355 &usdc(),
356 )
357 .unwrap();
358
359 assert_eq!(amount_out_result.amount, BigUint::from_str("4500000000").unwrap()); 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 let result = state.get_amount_out(
373 BigUint::from_str("10000000000").unwrap(), &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 let result = state.get_amount_out(
392 BigUint::from_str("250000000000000000").unwrap(), &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 let result = state.get_amount_out(
411 BigUint::from_str("8000000000000000000").unwrap(), &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 let result = state.get_amount_out(
430 BigUint::from_str("100000000").unwrap(), &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![]; let result = state.get_amount_out(
449 BigUint::from_str("1000000000000000000").unwrap(), &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 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}