1use std::any::Any;
4
5use alloy::primitives::{Address as AlloyAddress, U256};
6use num_bigint::{BigUint, ToBigUint};
7use serde::{Deserialize, Serialize};
8use tracing::debug;
9use tycho_common::{
10 dto::ProtocolStateDelta,
11 models::token::Token,
12 simulation::{
13 errors::{SimulationError, TransitionError},
14 protocol_sim::{
15 Balances, GetAmountOutResult, PoolSwap, Price, ProtocolSim, QueryPoolSwapParams,
16 SwapConstraint,
17 },
18 },
19 Bytes,
20};
21
22use crate::evm::{
23 engine_db::{create_engine, SHARED_TYCHO_DB},
24 protocol::{
25 curve::{
26 adapter::{build_pool, CurveVariant},
27 math::Pool,
28 swap_to_price::{exchange, swap_to_price, SwapToPriceError},
29 vm,
30 },
31 u256_num::{biguint_to_u256, u256_to_biguint, u256_to_f64},
32 },
33};
34
35const FEE_DENOMINATOR: f64 = 1e10;
37const STABLESWAP_GAS: u64 = 150_000;
39const CRYPTOSWAP_GAS: u64 = 350_000;
41
42#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
52pub struct CurveState {
53 pool_address: Bytes,
55 tokens: Vec<Bytes>,
57 decimals: Vec<u8>,
59 variant: CurveVariant,
61 pool: Pool,
63 #[serde(default)]
65 admin_fee: Option<U256>,
66}
67
68impl CurveState {
69 pub fn new(
72 pool_address: Bytes,
73 tokens: Vec<Bytes>,
74 decimals: Vec<u8>,
75 variant: CurveVariant,
76 pool: Pool,
77 admin_fee: Option<U256>,
78 ) -> Self {
79 Self { pool_address, tokens, decimals, variant, pool, admin_fee }
80 }
81
82 fn admin_fee(&self) -> Result<U256, SimulationError> {
83 if self.variant == CurveVariant::StableSwapNG {
84 return Ok(U256::from(5_000_000_000u64));
85 }
86 self.admin_fee.ok_or_else(|| {
87 SimulationError::RecoverableError(format!(
88 "Missing Curve admin fee for {}; refresh the pool state",
89 self.pool_address
90 ))
91 })
92 }
93
94 fn coin_index(&self, token: &Bytes) -> Result<usize, SimulationError> {
95 self.tokens
96 .iter()
97 .position(|t| t == token)
98 .ok_or_else(|| {
99 SimulationError::InvalidInput(
100 format!("token {token} is not a coin of curve pool {}", self.pool_address),
101 None,
102 )
103 })
104 }
105
106 fn is_crypto(&self) -> bool {
107 matches!(
108 self.variant,
109 CurveVariant::TwoCryptoV1 |
110 CurveVariant::TwoCryptoNG |
111 CurveVariant::TwoCryptoStable |
112 CurveVariant::TriCryptoV1 |
113 CurveVariant::TriCryptoNG
114 )
115 }
116
117 fn gas_estimate(&self) -> u64 {
118 if self.is_crypto() {
119 CRYPTOSWAP_GAS
120 } else {
121 STABLESWAP_GAS
122 }
123 }
124
125 fn swap_to_target_price(
129 &self,
130 token_in: &Token,
131 token_out: &Token,
132 target: &Price,
133 tolerance: f64,
134 ) -> Result<Option<PoolSwap>, SimulationError> {
135 let i = self.coin_index(&token_in.address)?;
136 let j = self.coin_index(&token_out.address)?;
137 if target.numerator.bits() > 256 || target.denominator.bits() > 256 {
138 return Ok(None);
139 }
140 let target_num = biguint_to_u256(&target.numerator);
141 let target_den = biguint_to_u256(&target.denominator);
142
143 match swap_to_price(
144 &self.pool,
145 i,
146 j,
147 target_num,
148 target_den,
149 tolerance,
150 if self.is_crypto() { U256::ZERO } else { self.admin_fee()? },
151 ) {
152 Ok(dx) => {
153 if dx.is_zero() {
154 let swap = PoolSwap::new(BigUint::ZERO, BigUint::ZERO, self.clone_box(), None);
155 return Ok(Some(swap));
156 }
157 let result = self.get_amount_out(u256_to_biguint(dx), token_in, token_out)?;
158 Ok(Some(PoolSwap::new(u256_to_biguint(dx), result.amount, result.new_state, None)))
159 }
160 Err(SwapToPriceError::TargetAboveSpot) => {
161 let spot = self.spot_price(token_in, token_out)?;
162 let decimal_adjustment =
163 10f64.powi(token_in.decimals as i32 - token_out.decimals as i32);
164 let target =
165 u256_to_f64(target_num)? / u256_to_f64(target_den)? * decimal_adjustment;
166 Err(SimulationError::InvalidInput(
167 format!("Target price {target} is above spot price {spot}"),
168 None,
169 ))
170 }
171 Err(SwapToPriceError::TargetBelowLimit) => Err(SimulationError::InvalidInput(
172 format!(
173 "Target price below reachable limit for curve pool {pool}",
174 pool = self.pool_address
175 ),
176 None,
177 )),
178 Err(err @ SwapToPriceError::InvalidInput(_)) => Err(SimulationError::InvalidInput(
179 format!("{err} for curve pool {pool}", pool = self.pool_address),
180 None,
181 )),
182 Err(SwapToPriceError::UnsupportedVariant) => Ok(None),
183 Err(err @ SwapToPriceError::MathFailed) => {
184 debug!(
185 pool = %self.pool_address,
186 %err,
187 "Curve native swap-to-price failed; using the numerical search"
188 );
189 Ok(None)
190 }
191 }
192 }
193}
194
195#[typetag::serde]
196impl ProtocolSim for CurveState {
197 fn fee(&self) -> f64 {
198 let fee = self.pool.fee().or_else(|| {
199 self.pool
200 .crypto_fees()
201 .map(|(mid, _, _)| mid)
202 });
203 fee.and_then(|f| u256_to_f64(f).ok())
204 .map(|f| f / FEE_DENOMINATOR)
205 .unwrap_or(0.0)
206 }
207
208 fn spot_price(&self, base: &Token, quote: &Token) -> Result<f64, SimulationError> {
209 let i = self.coin_index(&base.address)?;
210 let j = self.coin_index("e.address)?;
211 let (numerator, denominator) = self
212 .pool
213 .spot_price(i, j)
214 .ok_or_else(|| {
215 SimulationError::RecoverableError(format!(
216 "curve spot price unavailable for {}",
217 self.pool_address
218 ))
219 })?;
220 let ratio = u256_to_f64(numerator)? / u256_to_f64(denominator)?;
223 let decimal_adjustment = 10f64.powi(base.decimals as i32 - quote.decimals as i32);
224 Ok(ratio * decimal_adjustment)
225 }
226
227 fn get_amount_out(
228 &self,
229 amount_in: BigUint,
230 token_in: &Token,
231 token_out: &Token,
232 ) -> Result<GetAmountOutResult, SimulationError> {
233 let i = self.coin_index(&token_in.address)?;
234 let j = self.coin_index(&token_out.address)?;
235 let dx = biguint_to_u256(&amount_in);
236
237 let mut new_pool = self.pool.clone();
238 let dy = if self.is_crypto() {
239 let dy = self
240 .pool
241 .get_amount_out(i, j, dx)
242 .ok_or_else(|| {
243 SimulationError::RecoverableError(format!(
244 "curve get_amount_out failed for {}",
245 self.pool_address
246 ))
247 })?;
248 let balances = self.pool.balances();
249 new_pool
251 .set_balance(i, balances[i] + dx)
252 .map_err(|e| SimulationError::FatalError(e.to_string()))?;
253 new_pool
254 .set_balance(j, balances[j].saturating_sub(dy))
255 .map_err(|e| SimulationError::FatalError(e.to_string()))?;
256 dy
257 } else {
258 let result = exchange(&self.pool, i, j, dx, self.admin_fee()?).ok_or_else(|| {
259 SimulationError::RecoverableError(format!(
260 "curve exchange accounting failed for {}",
261 self.pool_address
262 ))
263 })?;
264 for (index, balance) in result.balances.into_iter().enumerate() {
265 new_pool
266 .set_balance(index, balance)
267 .map_err(|e| SimulationError::FatalError(e.to_string()))?;
268 }
269 result.amount
270 };
271
272 let new_state = Self { pool: new_pool, ..self.clone() };
273 Ok(GetAmountOutResult::new(
274 u256_to_biguint(dy),
275 self.gas_estimate()
276 .to_biguint()
277 .expect("u64 fits in BigUint"),
278 Box::new(new_state),
279 ))
280 }
281
282 fn get_limits(
283 &self,
284 sell_token: Bytes,
285 buy_token: Bytes,
286 ) -> Result<(BigUint, BigUint), SimulationError> {
287 let i = self.coin_index(&sell_token)?;
288 let j = self.coin_index(&buy_token)?;
289 let (balance_in, balance_out) = {
290 let balances = self.pool.balances();
291 (balances[i], balances[j])
292 };
293 if balance_in.is_zero() || balance_out.is_zero() {
294 return Ok((BigUint::ZERO, BigUint::ZERO));
295 }
296 let max_out_reserve = balance_out.saturating_sub(U256::from(1));
299 let max_out = self
300 .pool
301 .get_amount_out(i, j, balance_in)
302 .ok_or_else(|| {
303 SimulationError::RecoverableError(format!(
304 "curve get_limits: solver failed at max input for {}",
305 self.pool_address
306 ))
307 })?
308 .min(max_out_reserve);
309 Ok((u256_to_biguint(balance_in), u256_to_biguint(max_out)))
310 }
311
312 fn delta_transition(
319 &mut self,
320 delta: ProtocolStateDelta,
321 _tokens: &std::collections::HashMap<Bytes, Token>,
322 _balances: &Balances,
323 ) -> Result<(), TransitionError> {
324 let state = match delta
325 .updated_attributes
326 .get(vm::POOL_STATE_ADJUSTED)
327 {
328 Some(encoded) => {
329 let state = vm::decode_raw_state(encoded)?;
330 if state.variant != self.variant {
331 return Err(SimulationError::FatalError(format!(
332 "Variant mismatch: expected {}, got {}",
333 self.variant, state.variant
334 ))
335 .into())
336 }
337 if state.token_decimals != self.decimals {
338 return Err(SimulationError::FatalError(format!(
339 "Token decimals mismatch: expected {:?}, got {:?}",
340 self.decimals, state.token_decimals
341 ))
342 .into())
343 }
344 state
345 }
346 None => {
347 let engine = create_engine(SHARED_TYCHO_DB.clone(), false).expect("Infallible");
348 let pool_address = AlloyAddress::from_slice(self.pool_address.as_ref());
349 vm::read_raw_pool_state(
350 &engine,
351 &pool_address,
352 self.variant,
353 &self.decimals,
354 &Default::default(),
355 )?
356 }
357 };
358 let pool = build_pool(&state)
359 .map_err(|e| SimulationError::FatalError(format!("curve build_pool failed: {e}")))?;
360 self.pool = pool;
361 self.admin_fee = state.admin_fee;
362 Ok(())
363 }
364
365 fn query_pool_swap(&self, params: &QueryPoolSwapParams) -> Result<PoolSwap, SimulationError> {
369 match params.swap_constraint() {
370 SwapConstraint::TradeLimitPrice { .. } => {
371 crate::evm::query_pool_swap::query_pool_swap(self, params)
372 }
373 SwapConstraint::PoolTargetPrice { target, tolerance, .. } => {
374 let native = self.swap_to_target_price(
375 params.token_in(),
376 params.token_out(),
377 target,
378 *tolerance,
379 )?;
380 match native {
381 Some(swap) => Ok(swap),
382 None => crate::evm::query_pool_swap::query_pool_swap(self, params),
383 }
384 }
385 }
386 }
387
388 fn clone_box(&self) -> Box<dyn ProtocolSim> {
389 Box::new(self.clone())
390 }
391
392 fn as_any(&self) -> &dyn Any {
393 self
394 }
395
396 fn as_any_mut(&mut self) -> &mut dyn Any {
397 self
398 }
399
400 fn eq(&self, other: &dyn ProtocolSim) -> bool {
401 other
402 .as_any()
403 .downcast_ref::<Self>()
404 .is_some_and(|other| self == other)
405 }
406}
407
408#[cfg(test)]
409mod tests {
410 use std::{collections::HashMap, str::FromStr};
411
412 use num_traits::ToPrimitive;
413 use rstest::rstest;
414 use tycho_common::{
415 models::Chain,
416 simulation::protocol_sim::{QueryPoolSwapParams, SwapConstraint},
417 };
418
419 use super::*;
420 use crate::evm::{
421 protocol::curve::{adapter::RawPoolState, vm::encode_raw_state},
422 query_pool_swap::test_helpers::{target_price_params, to_price},
423 };
424
425 const VARIANT: CurveVariant = CurveVariant::TriCryptoNG;
426 const DECIMALS: [u8; 3] = [6, 8, 18];
427
428 fn u(s: &str) -> U256 {
429 s.parse().unwrap()
430 }
431
432 fn raw_pool_state(balances: Vec<U256>) -> RawPoolState {
434 RawPoolState {
435 variant: VARIANT,
436 balances,
437 token_decimals: DECIMALS.to_vec(),
438 amp: u("1707629"),
439 mid_fee: Some(u("3000000")),
440 out_fee: Some(u("30000000")),
441 fee_gamma: Some(u("500000000000000")),
442 d: Some(u("7457948167729606869978625")),
443 gamma: Some(u("11809167828997")),
444 price_scale: Some(vec![u("59372627314351316239076"), u("1565715369034455123313")]),
445 ..Default::default()
446 }
447 }
448
449 fn state(balances: Vec<U256>) -> CurveState {
450 let raw_state = raw_pool_state(balances);
451 let pool = build_pool(&raw_state).expect("build");
452 CurveState::new(
453 Bytes::from([7u8; 20]),
454 vec![Bytes::from([1u8; 20]), Bytes::from([2u8; 20]), Bytes::from([3u8; 20])],
455 DECIMALS.to_vec(),
456 VARIANT,
457 pool,
458 Some(U256::ZERO),
459 )
460 }
461
462 fn delta(attributes: HashMap<String, Bytes>) -> ProtocolStateDelta {
463 ProtocolStateDelta { updated_attributes: attributes, ..Default::default() }
464 }
465
466 #[test]
467 fn test_delta_transition_rebuilds_from_attribute() {
468 let confirmed = vec![u("2466241139205"), u("4200057336"), u("1595469030050811720465")];
469 let pending = vec![u("2470000000000"), u("4190000000"), u("1600000000000000000000")];
470 let mut curve = state(confirmed.clone());
471 let attribute = encode_raw_state(&raw_pool_state(pending.clone())).expect("encode");
472
473 curve
474 .delta_transition(
475 delta(HashMap::from([(vm::POOL_STATE_ADJUSTED.to_string(), attribute)])),
476 &HashMap::new(),
477 &Balances::default(),
478 )
479 .expect("delta transition from attribute failed");
480
481 assert_eq!(
484 curve.pool.balances()[..3],
485 pending[..],
486 "balances must come from the attribute"
487 );
488 assert_ne!(curve.pool.balances()[..3], confirmed[..]);
489 }
490
491 #[test]
492 fn test_delta_transition_errors_on_variant_mismatch() {
493 let mut curve = state(vec![u("1"), u("2"), u("3")]);
494 let mut pending_state = raw_pool_state(curve.pool.balances().to_vec());
495 pending_state.variant = CurveVariant::StableSwapMeta;
496 let encoded = encode_raw_state(&pending_state).expect("encode");
497
498 let result = curve.delta_transition(
499 delta(HashMap::from([(vm::POOL_STATE_ADJUSTED.to_string(), encoded)])),
500 &HashMap::new(),
501 &Balances::default(),
502 );
503
504 assert!(matches!(result, Err(TransitionError::SimulationError(_))), "got {result:?}");
505 }
506
507 #[test]
508 fn test_delta_transition_errors_on_decimals_mismatch() {
509 let mut curve = state(vec![u("1"), u("2"), u("3")]);
510 let mut pending_state = raw_pool_state(curve.pool.balances().to_vec());
511 pending_state.token_decimals = vec![18, 18, 18];
512 let encoded = encode_raw_state(&pending_state).expect("encode");
513
514 let result = curve.delta_transition(
515 delta(HashMap::from([(vm::POOL_STATE_ADJUSTED.to_string(), encoded)])),
516 &HashMap::new(),
517 &Balances::default(),
518 );
519
520 assert!(matches!(result, Err(TransitionError::SimulationError(_))), "got {result:?}");
521 }
522
523 #[test]
524 fn test_delta_transition_rejects_malformed_attribute() {
525 let mut curve = state(vec![u("1"), u("2"), u("3")]);
526
527 let result = curve.delta_transition(
528 delta(HashMap::from([(
529 vm::POOL_STATE_ADJUSTED.to_string(),
530 Bytes::from(b"not json".to_vec()),
531 )])),
532 &HashMap::new(),
533 &Balances::default(),
534 );
535
536 assert!(matches!(result, Err(TransitionError::SimulationError(_))), "got {result:?}");
539 }
540
541 const WAD: u128 = 1_000_000_000_000_000_000;
542 const RATE_6_DEC: u128 = 1_000_000_000_000_000_000_000_000_000_000;
543
544 fn token(index: u8, decimals: u32) -> Token {
545 let address =
546 Bytes::from_str(&format!("0x00000000000000000000000000000000000000{index:02x}"))
547 .expect("valid address");
548 Token::new(
549 &address,
550 &format!("T{index}"),
551 decimals,
552 0,
553 &[Some(10_000)],
554 Chain::Ethereum,
555 100,
556 )
557 }
558
559 fn curve_state(pool: Pool, variant: CurveVariant, decimals: Vec<u8>) -> CurveState {
560 let tokens: Vec<Bytes> = (0..decimals.len())
561 .map(|k| token(k as u8, decimals[k] as u32).address)
562 .collect();
563 CurveState::new(
564 Bytes::from_str("0x00000000000000000000000000000000000000ff").expect("valid address"),
565 tokens,
566 decimals,
567 variant,
568 pool,
569 Some(U256::ZERO),
570 )
571 }
572
573 fn v1_two_coin_state() -> (CurveState, Token, Token) {
574 let pool = Pool::StableSwapV1 {
575 balances: vec![U256::from(50_000_000u128 * WAD), U256::from(48_000_000u128 * WAD)],
576 rates: vec![U256::from(WAD), U256::from(WAD)],
577 amp: U256::from(2000u64),
578 fee: U256::from(1_000_000u64),
579 };
580 (curve_state(pool, CurveVariant::StableSwapV1, vec![18, 18]), token(0, 18), token(1, 18))
581 }
582
583 fn v1_three_coin_mixed_state() -> (CurveState, Token, Token) {
584 let pool = Pool::StableSwapV1 {
586 balances: vec![
587 U256::from(63_975_337_809_806_329_031_583_135u128),
588 U256::from(61_219_263_170_093u128),
589 U256::from(37_832_425_459_809u128),
590 ],
591 rates: vec![U256::from(WAD), U256::from(RATE_6_DEC), U256::from(RATE_6_DEC)],
592 amp: U256::from(4000u64),
593 fee: U256::from(1_500_000u64),
594 };
595 (curve_state(pool, CurveVariant::StableSwapV1, vec![18, 6, 6]), token(0, 18), token(1, 6))
596 }
597
598 fn v1_three_coin_mixed_state_reverse() -> (CurveState, Token, Token) {
599 let (state, token_out, token_in) = v1_three_coin_mixed_state();
600 (state, token_in, token_out)
601 }
602
603 fn ng_dynamic_fee_state() -> (CurveState, Token, Token) {
604 let pool = Pool::StableSwapNG {
605 balances: vec![U256::from(1_500_000u128 * WAD), U256::from(700_000u128 * WAD)],
606 rates: vec![U256::from(WAD), U256::from(WAD)],
607 amp: U256::from(40_000u64),
608 fee: U256::from(4_000_000u64),
609 offpeg_fee_multiplier: U256::from(20_000_000_000u64),
610 };
611 (curve_state(pool, CurveVariant::StableSwapNG, vec![18, 18]), token(0, 18), token(1, 18))
612 }
613
614 fn meta_state() -> (CurveState, Token, Token) {
615 let pool = Pool::StableSwapMeta {
616 balances: vec![U256::from(500_000u128 * WAD), U256::from(480_000u128 * WAD)],
617 rates: vec![U256::from(WAD), U256::from(1_030_000_000_000_000_000u128)],
618 amp: U256::from(50_000u64),
619 fee: U256::from(4_000_000u64),
620 };
621 (curve_state(pool, CurveVariant::StableSwapMeta, vec![18, 18]), token(0, 18), token(1, 18))
622 }
623
624 #[test]
625 fn pending_state_preserves_admin_fee_for_exchange() {
626 let (mut state, token_in, token_out) = v1_two_coin_state();
627 let pool = state.pool.clone();
628 let Pool::StableSwapV1 { balances, rates: _, amp, fee } = pool else { unreachable!() };
629 let raw = RawPoolState {
630 variant: CurveVariant::StableSwapV1,
631 balances,
632 amp,
633 fee: Some(fee),
634 token_decimals: vec![18, 18],
635 admin_fee: Some(U256::from(5_000_000_000u64)),
636 ..Default::default()
637 };
638 state
639 .delta_transition(
640 delta(HashMap::from([(
641 vm::POOL_STATE_ADJUSTED.into(),
642 encode_raw_state(&raw).unwrap(),
643 )])),
644 &HashMap::new(),
645 &Balances::default(),
646 )
647 .unwrap();
648 assert_eq!(state.admin_fee, raw.admin_fee);
649 let mut without_admin = state.clone();
650 without_admin.admin_fee = Some(U256::ZERO);
651 let amount = BigUint::from(WAD);
652 let with_fee = state
653 .get_amount_out(amount.clone(), &token_in, &token_out)
654 .unwrap();
655 let no_fee = without_admin
656 .get_amount_out(amount, &token_in, &token_out)
657 .unwrap();
658 assert_eq!(with_fee.amount, no_fee.amount);
659 let balance = |result: &GetAmountOutResult| {
660 result
661 .new_state
662 .as_any()
663 .downcast_ref::<CurveState>()
664 .unwrap()
665 .pool
666 .balances()[1]
667 };
668 assert!(balance(&with_fee) < balance(&no_fee));
669 }
670
671 fn two_crypto_ng_state() -> (CurveState, Token, Token) {
672 let wad = U256::from(WAD);
673 let pool = Pool::TwoCryptoNG {
674 balances: [U256::from(5000u64) * wad, U256::from(5000u64) * wad],
675 precisions: [U256::from(1u64), U256::from(1u64)],
676 price_scale: wad,
677 d: U256::from(10000u64) * wad,
678 ann: U256::from(540_000u64) * U256::from(10_000u64),
679 gamma: U256::from(11_809_167_828_997u64),
680 mid_fee: U256::from(3_000_000u64),
681 out_fee: U256::from(30_000_000u64),
682 fee_gamma: U256::from(230_000_000_000_000u64),
683 };
684 (curve_state(pool, CurveVariant::TwoCryptoNG, vec![18, 18]), token(0, 18), token(1, 18))
685 }
686
687 const TOLERANCE: f64 = 0.001;
688
689 fn spot_target_params(
692 state: &CurveState,
693 token_in: &Token,
694 token_out: &Token,
695 multiplier: f64,
696 ) -> (QueryPoolSwapParams, f64) {
697 let spot = state
698 .spot_price(token_in, token_out)
699 .expect("spot price");
700 let target_f64 = spot * multiplier;
701 let target = to_price(target_f64, token_in, token_out);
702 (target_price_params(token_in, token_out, target, TOLERANCE), target_f64)
703 }
704
705 #[rstest]
706 #[case::v1_two_coin_shallow(v1_two_coin_state(), 0.9999)]
707 #[case::v1_two_coin_mid(v1_two_coin_state(), 0.999)]
708 #[case::v1_two_coin_deep(v1_two_coin_state(), 0.99)]
709 #[case::v1_mixed_decimals_18_to_6(v1_three_coin_mixed_state(), 0.99)]
710 #[case::v1_mixed_decimals_6_to_18(v1_three_coin_mixed_state_reverse(), 0.99)]
711 #[case::ng_dynamic_fee(ng_dynamic_fee_state(), 0.99)]
712 #[case::ng_dynamic_fee_mid(ng_dynamic_fee_state(), 0.999)]
713 #[case::meta_virtual_price(meta_state(), 0.99)]
714 #[case::meta_virtual_price_mid(meta_state(), 0.999)]
715 fn test_query_pool_swap_native_amounts(
716 #[case] setup: (CurveState, Token, Token),
717 #[case] multiplier: f64,
718 ) {
719 let (state, token_in, token_out) = setup;
720 let (params, target) = spot_target_params(&state, &token_in, &token_out, multiplier);
721
722 let swap = state
723 .query_pool_swap(¶ms)
724 .expect("native query_pool_swap");
725
726 let price = swap
727 .new_state()
728 .spot_price(&token_in, &token_out)
729 .unwrap();
730 assert!(price >= target && price <= target * (1.0 + TOLERANCE));
731 let executed = state
732 .get_amount_out(swap.amount_in().clone(), &token_in, &token_out)
733 .unwrap();
734 assert_eq!(swap.amount_out(), &executed.amount);
735 assert!(ProtocolSim::eq(swap.new_state(), executed.new_state.as_ref()));
736 }
737
738 #[rstest]
741 #[case::v1_two_coin_shallow(v1_two_coin_state(), 0.9999)]
742 #[case::v1_two_coin_mid(v1_two_coin_state(), 0.999)]
743 #[case::v1_two_coin_deep(v1_two_coin_state(), 0.99)]
744 #[case::v1_mixed_decimals_18_to_6(v1_three_coin_mixed_state(), 0.99)]
745 #[case::v1_mixed_decimals_6_to_18(v1_three_coin_mixed_state_reverse(), 0.99)]
746 #[case::ng_dynamic_fee(ng_dynamic_fee_state(), 0.99)]
747 #[case::ng_dynamic_fee_mid(ng_dynamic_fee_state(), 0.999)]
748 #[case::meta_virtual_price(meta_state(), 0.99)]
749 #[case::meta_virtual_price_mid(meta_state(), 0.999)]
750 fn test_query_pool_swap_numerical_comparison(
751 #[case] setup: (CurveState, Token, Token),
752 #[case] multiplier: f64,
753 ) {
754 let (state, token_in, token_out) = setup;
755 let (params, target_f64) = spot_target_params(&state, &token_in, &token_out, multiplier);
756
757 let native = state
758 .query_pool_swap(¶ms)
759 .expect("native query_pool_swap");
760 let numerical = crate::evm::query_pool_swap::query_pool_swap(&state, ¶ms)
761 .expect("numerical query_pool_swap");
762
763 for (label, swap, band) in
764 [("native", &native, TOLERANCE / 2.0), ("numerical", &numerical, 5.0 * TOLERANCE)]
765 {
766 assert!(swap.amount_in() > &BigUint::ZERO, "{label} amount_in should be > 0");
767 let new_spot = swap
768 .new_state()
769 .spot_price(&token_in, &token_out)
770 .expect("post-swap spot");
771 let error = (new_spot - target_f64) / target_f64;
772 assert!(
773 error >= -1e-12,
774 "{label} post-swap spot {new_spot} fell below target {target_f64}"
775 );
776 assert!(
777 error <= band,
778 "{label} post-swap spot {new_spot} outside band of target {target_f64}: {error}"
779 );
780 }
781 }
782
783 #[test]
786 fn test_crypto_variant_delegates_to_numerical() {
787 let (state, token_in, token_out) = two_crypto_ng_state();
788 let (params, _) = spot_target_params(&state, &token_in, &token_out, 0.995);
789
790 let result = state.query_pool_swap(¶ms);
791 let Err(SimulationError::InvalidInput(msg, _)) = result else {
792 panic!("crypto pools must delegate to the numerical search, got {result:?}");
793 };
794 assert!(msg.contains("< limit"), "expected the numerical search's limit error, got: {msg}");
795 }
796
797 #[test]
798 fn test_query_pool_swap_target_wider_than_u256() {
799 let (state, token_in, token_out) = v1_two_coin_state();
800 let (params, target_f64) = spot_target_params(&state, &token_in, &token_out, 0.999);
801 let SwapConstraint::PoolTargetPrice { target, .. } = params.swap_constraint() else {
802 panic!("spot_target_params builds a PoolTargetPrice constraint");
803 };
804 let scale = BigUint::from(1u8) << 256;
805 let wide_target = Price::new(&target.numerator * &scale, &target.denominator * &scale);
806 let params = target_price_params(&token_in, &token_out, wide_target, TOLERANCE);
807
808 let swap = state
809 .query_pool_swap(¶ms)
810 .expect("numerical query_pool_swap");
811
812 assert!(swap.price_points().is_some(), "only the numerical search returns price points");
813 let new_spot = swap
814 .new_state()
815 .spot_price(&token_in, &token_out)
816 .expect("post-swap spot");
817 let error = (new_spot - target_f64) / target_f64;
818 assert!(
819 (-1e-12..=5.0 * TOLERANCE).contains(&error),
820 "post-swap spot {new_spot} missed {target_f64}"
821 );
822 }
823
824 #[test]
825 fn test_trade_limit_price_delegates_to_numerical() {
826 let (state, token_in, token_out) = v1_two_coin_state();
827 let spot = state
828 .spot_price(&token_in, &token_out)
829 .expect("spot price");
830 let limit_f64 = spot * 0.999;
831 let params = QueryPoolSwapParams::new(
832 token_in.clone(),
833 token_out.clone(),
834 SwapConstraint::TradeLimitPrice {
835 limit: to_price(limit_f64, &token_in, &token_out),
836 tolerance: TOLERANCE,
837 min_amount_in: None,
838 max_amount_in: None,
839 },
840 );
841
842 let swap = state
843 .query_pool_swap(¶ms)
844 .expect("trade limit query_pool_swap");
845 assert!(swap.amount_in() > &BigUint::ZERO);
846 assert!(swap.amount_out() > &BigUint::ZERO);
847 let trade_price = swap
848 .amount_out()
849 .to_f64()
850 .expect("failed to convert the output amount to f64") /
851 swap.amount_in()
852 .to_f64()
853 .expect("failed to convert the input amount to f64");
854 assert!(trade_price >= limit_f64, "trade price {trade_price} violates limit {limit_f64}");
855 }
856
857 #[rstest]
858 #[case::above_spot(1.01, "is above spot price")]
859 #[case::below_limit(1e-9, "below reachable limit for curve pool")]
860 fn test_query_pool_swap_unreachable_target(#[case] multiplier: f64, #[case] expected: &str) {
861 let (state, token_in, token_out) = v1_two_coin_state();
862 let (params, _) = spot_target_params(&state, &token_in, &token_out, multiplier);
863
864 let result = state.query_pool_swap(¶ms);
865 let Err(SimulationError::InvalidInput(msg, _)) = result else {
866 panic!("expected InvalidInput, got {result:?}");
867 };
868 assert!(msg.contains(expected), "unexpected message: {msg}");
869 }
870
871 #[test]
872 fn test_query_pool_swap_target_equal_to_spot() {
873 let (state, token_in, token_out) = v1_two_coin_state();
874 let i = state
875 .coin_index(&token_in.address)
876 .expect("token_in index");
877 let j = state
878 .coin_index(&token_out.address)
879 .expect("token_out index");
880 let (num, den) = state
881 .pool
882 .spot_price(i, j)
883 .expect("pool spot price");
884 let target = Price::new(u256_to_biguint(num), u256_to_biguint(den));
885 let params = target_price_params(&token_in, &token_out, target, TOLERANCE);
886
887 let swap = state
888 .query_pool_swap(¶ms)
889 .expect("query_pool_swap at spot");
890 assert_eq!(swap.amount_in(), &BigUint::ZERO);
891 assert_eq!(swap.amount_out(), &BigUint::ZERO);
892 }
893}