1use std::{any::Any, collections::HashMap, fmt, sync::Arc};
2
3use async_trait::async_trait;
4use num_bigint::BigUint;
5use serde::{Deserialize, Serialize};
6use tycho_common::{
7 dto::ProtocolStateDelta,
8 models::{protocol::GetAmountOutParams, token::Token},
9 simulation::{
10 errors::{SimulationError, TransitionError},
11 indicatively_priced::{IndicativelyPriced, SignedQuote},
12 protocol_sim::{Balances, GetAmountOutResult, ProtocolSim},
13 },
14 Bytes,
15};
16
17use crate::rfq::protocols::native::{
18 client::NativeClient, models::NativePriceData, state::NativeState,
19};
20
21#[derive(Clone, Serialize, Deserialize)]
26pub struct NativeAllPairsState {
27 states: Arc<Vec<NativeState>>,
28 directions: Arc<Vec<((Bytes, Bytes), usize)>>,
31 used: bool,
33}
34
35impl fmt::Debug for NativeAllPairsState {
36 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
37 f.debug_struct("NativeAllPairsState")
38 .field("books", &self.states.len())
39 .field("used", &self.used)
40 .finish_non_exhaustive()
41 }
42}
43
44impl NativeAllPairsState {
45 pub fn new(
48 books: Vec<NativePriceData>,
49 tokens: HashMap<Bytes, Token>,
50 client: NativeClient,
51 ) -> Result<Self, SimulationError> {
52 let mut states = Vec::with_capacity(books.len());
53 for book in books {
54 let (Some(base_token), Some(quote_token)) =
56 (tokens.get(&book.base_address), tokens.get(&book.quote_address))
57 else {
58 return Err(SimulationError::FatalError(
59 "Native book token addresses do not match state tokens".to_string(),
60 ));
61 };
62 states.push(NativeState::new(
63 base_token.clone(),
64 quote_token.clone(),
65 book,
66 client.clone(),
67 )?);
68 }
69 let mut directions = HashMap::new();
70 for (index, state) in states.iter().enumerate() {
71 directions
72 .entry((state.book.base_address.clone(), state.book.quote_address.clone()))
73 .or_insert(index);
74 }
75 for (index, state) in states.iter().enumerate() {
76 directions
77 .entry((state.book.quote_address.clone(), state.book.base_address.clone()))
78 .or_insert(index);
79 }
80 let mut directions: Vec<_> = directions.into_iter().collect();
81 directions.sort();
82 Ok(Self { states: Arc::new(states), directions: Arc::new(directions), used: false })
83 }
84
85 fn pair_state(
87 &self,
88 token_in: &Bytes,
89 token_out: &Bytes,
90 ) -> Result<&NativeState, SimulationError> {
91 let index = self
92 .directions
93 .binary_search_by(|((a, b), _)| (a, b).cmp(&(token_in, token_out)))
94 .map_err(|_| {
95 SimulationError::InvalidInput(
96 format!("Invalid token addresses. Got in={token_in}, out={token_out}"),
97 None,
98 )
99 })?;
100 Ok(&self.states[self.directions[index].1])
101 }
102
103 fn quotable_pair_state(
105 &self,
106 token_in: &Bytes,
107 token_out: &Bytes,
108 ) -> Result<&NativeState, SimulationError> {
109 let state = self.pair_state(token_in, token_out)?;
110 if self.used {
111 return Err(SimulationError::RecoverableError(
112 "Native already quoted in this route".to_string(),
113 ));
114 }
115 Ok(state)
116 }
117
118 fn used_state(&self) -> Box<dyn ProtocolSim> {
119 Box::new(Self {
120 states: self.states.clone(),
121 directions: self.directions.clone(),
122 used: true,
123 })
124 }
125}
126
127#[typetag::serde]
128impl ProtocolSim for NativeAllPairsState {
129 fn fee(&self) -> f64 {
130 0.0
131 }
132
133 fn spot_price(&self, base: &Token, quote: &Token) -> Result<f64, SimulationError> {
134 self.quotable_pair_state(&base.address, "e.address)?
135 .spot_price(base, quote)
136 }
137
138 fn get_amount_out(
139 &self,
140 amount_in: BigUint,
141 token_in: &Token,
142 token_out: &Token,
143 ) -> Result<GetAmountOutResult, SimulationError> {
144 let state = self.quotable_pair_state(&token_in.address, &token_out.address)?;
145 match state.get_amount_out(amount_in, token_in, token_out) {
146 Ok(mut res) => {
147 res.new_state = self.used_state();
148 Ok(res)
149 }
150 Err(SimulationError::InvalidInput(message, Some(mut res))) => {
151 res.new_state = self.used_state();
152 Err(SimulationError::InvalidInput(message, Some(res)))
153 }
154 Err(e) => Err(e),
155 }
156 }
157
158 fn get_limits(
159 &self,
160 sell_token: Bytes,
161 buy_token: Bytes,
162 ) -> Result<(BigUint, BigUint), SimulationError> {
163 self.quotable_pair_state(&sell_token, &buy_token)?
164 .get_limits(sell_token, buy_token)
165 }
166
167 fn delta_transition(
168 &mut self,
169 _delta: ProtocolStateDelta,
170 _tokens: &HashMap<Bytes, Token>,
171 _balances: &Balances,
172 ) -> Result<(), TransitionError> {
173 Err(TransitionError::DecodeError("Not implemented".into()))
174 }
175
176 fn clone_box(&self) -> Box<dyn ProtocolSim> {
177 Box::new(self.clone())
178 }
179
180 fn as_any(&self) -> &dyn Any {
181 self
182 }
183
184 fn as_any_mut(&mut self) -> &mut dyn Any {
185 self
186 }
187
188 fn eq(&self, other: &dyn ProtocolSim) -> bool {
189 let Some(other) = other
190 .as_any()
191 .downcast_ref::<NativeAllPairsState>()
192 else {
193 return false;
194 };
195 self.used == other.used &&
196 self.states.len() == other.states.len() &&
197 self.states
198 .iter()
199 .zip(other.states.iter())
200 .all(|(a, b)| a.book == b.book)
201 }
202
203 fn as_indicatively_priced(&self) -> Result<&dyn IndicativelyPriced, SimulationError> {
204 Ok(self)
205 }
206}
207
208#[async_trait]
209impl IndicativelyPriced for NativeAllPairsState {
210 async fn request_signed_quote(
211 &self,
212 params: GetAmountOutParams,
213 ) -> Result<SignedQuote, SimulationError> {
214 self.pair_state(¶ms.token_in, ¶ms.token_out)?
215 .request_signed_quote(params)
216 .await
217 }
218}
219
220#[cfg(test)]
221mod tests {
222 use std::collections::HashSet;
223
224 use tokio::time::Duration;
225 use tycho_common::models::Chain;
226
227 use super::*;
228 use crate::rfq::{
229 models::ComponentLayout,
230 protocols::{
231 native::models::NativePriceLevel,
232 test_utils::{token, usdc, weth},
233 },
234 };
235
236 fn book() -> NativePriceData {
237 NativePriceData {
238 base_address: weth().address,
239 quote_address: usdc().address,
240 minimum_in_base: 100_000_000_000.0,
241 minimum_in_quote: 100.0,
242 minimum_out_base: 0.0,
243 minimum_out_quote: 0.0,
244 bids: vec![NativePriceLevel { quantity: 1.0, price: 2_000.0 }],
245 asks: vec![NativePriceLevel { quantity: 1.0, price: 2_000.0 }],
246 }
247 }
248
249 fn state_with(books: Vec<NativePriceData>) -> Result<NativeAllPairsState, SimulationError> {
250 let client = NativeClient::new(
251 Chain::Ethereum,
252 String::new(),
253 HashSet::new(),
254 0.0,
255 HashSet::new(),
256 Duration::from_secs(5),
257 Duration::from_secs(5),
258 )
259 .unwrap()
260 .with_component_layout(ComponentLayout::AllPairs);
261 NativeAllPairsState::new(
262 books,
263 HashMap::from([(weth().address, weth()), (usdc().address, usdc())]),
264 client,
265 )
266 }
267
268 fn state() -> NativeAllPairsState {
269 state_with(vec![book()]).unwrap()
270 }
271
272 #[test]
273 fn once_per_venue() {
274 let state = state();
275 let first = state
276 .get_amount_out(BigUint::from(500_000_000_000_000_000u64), &weth(), &usdc())
277 .unwrap();
278 let after_first = first
279 .new_state
280 .as_any()
281 .downcast_ref::<NativeAllPairsState>()
282 .unwrap();
283 assert!(after_first.used);
284 assert!(matches!(
285 after_first.get_amount_out(BigUint::from(1_000_000_000u64), &usdc(), &weth()),
286 Err(SimulationError::RecoverableError(message)) if message == "Native already quoted in this route"
287 ));
288 assert!(matches!(
289 after_first.spot_price(&weth(), &usdc()),
290 Err(SimulationError::RecoverableError(message)) if message == "Native already quoted in this route"
291 ));
292 assert!(matches!(
293 after_first.get_limits(weth().address, usdc().address),
294 Err(SimulationError::RecoverableError(message)) if message == "Native already quoted in this route"
295 ));
296 }
297
298 #[test]
299 fn book_quoting_the_direction_beats_the_inverted_one() {
300 let mut forward = book();
303 forward.minimum_out_base = 2_000_000_000_000_000_000.0;
304 let mut reverse = book();
305 reverse.base_address = usdc().address;
306 reverse.quote_address = weth().address;
307 reverse.minimum_in_base = 0.0;
308 reverse.bids = vec![NativePriceLevel { quantity: 2_000.0, price: 0.0005 }];
309 reverse.asks = vec![];
310 let state = state_with(vec![forward, reverse]).unwrap();
311 let result = state
312 .get_amount_out(BigUint::from(2_000_000_000u64), &usdc(), &weth())
313 .unwrap();
314 assert_eq!(result.amount, BigUint::from(1_000_000_000_000_000_000u64));
315 }
316
317 #[test]
318 fn returns_partial_result_when_amount_exceeds_depth() {
319 let state = state();
320 let result =
321 state.get_amount_out(BigUint::from(2_000_000_000_000_000_000u64), &weth(), &usdc());
322 let Err(SimulationError::InvalidInput(_, Some(partial))) = result else {
323 panic!("Expected insufficient-liquidity result, got {result:?}");
324 };
325 assert_eq!(partial.amount, BigUint::from(2_000_000_000u64));
326 let new_state = partial
327 .new_state
328 .as_any()
329 .downcast_ref::<NativeAllPairsState>()
330 .unwrap();
331 assert!(new_state.used);
332 }
333
334 #[test]
335 fn rejects_invalid_pair() {
336 let other = token("0x1111111111111111111111111111111111111111", "OTHER", 18);
337 let state = state();
338 assert!(matches!(
339 state.get_amount_out(BigUint::from(1u64), &other, &usdc()),
340 Err(SimulationError::InvalidInput(_, None))
341 ));
342 assert!(matches!(
343 state.get_limits(other.address.clone(), usdc().address),
344 Err(SimulationError::InvalidInput(_, None))
345 ));
346 let mut empty = book();
348 empty.bids.clear();
349 empty.asks.clear();
350 let state = state_with(vec![empty]).unwrap();
351 assert!(matches!(
352 state.spot_price(&other, &usdc()),
353 Err(SimulationError::InvalidInput(message, None))
354 if message.contains("Invalid token addresses")
355 ));
356 }
357
358 #[test]
359 fn reports_no_liquidity_for_empty_direction() {
360 let mut book = book();
361 book.bids.clear();
362 let state = state_with(vec![book]).unwrap();
363 assert!(matches!(
364 state.get_amount_out(BigUint::from(500_000_000_000_000_000u64), &weth(), &usdc()),
365 Err(SimulationError::RecoverableError(_))
366 ));
367 assert!(matches!(
368 state.get_limits(weth().address, usdc().address),
369 Err(SimulationError::RecoverableError(_))
370 ));
371 }
372}