Skip to main content

tycho_simulation/evm/protocol/curve/
state.rs

1//! [`CurveState`] — a hybrid Curve pool: pure-Rust quote math (`curve_math::Pool`) over state read
2//! from the locally indexed VM storage.
3use 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
35/// Curve fee denominator (`10^10`); both StableSwap `fee` and CryptoSwap `mid_fee` use it.
36const FEE_DENOMINATOR: f64 = 1e10;
37/// Representative gas cost of a StableSwap exchange.
38const STABLESWAP_GAS: u64 = 150_000;
39/// Representative gas cost of a CryptoSwap exchange (heavier math + price oracle update).
40const CRYPTOSWAP_GAS: u64 = 350_000;
41
42/// A single Curve pool quoted via `curve_math`.
43///
44/// `tokens` and `decimals` are ordered to match the pool's coin indices, so a token address maps
45/// directly to a `curve_math` coin index. State (`pool`) is rebuilt from the VM on every
46/// `delta_transition`.
47///
48/// StableSwap exchanges update pricing balances net of admin fees and recompute `D` on the
49/// next quote. CryptoSwap updates balances only, holding `D` and `price_scale` fixed: re-quoting
50/// the same CryptoSwap pool is approximate because execution updates these via `tweak_price`.
51#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
52pub struct CurveState {
53    /// Pool contract address (the Tycho component id).
54    pool_address: Bytes,
55    /// Coin addresses in pool index order.
56    tokens: Vec<Bytes>,
57    /// Coin decimals in pool index order.
58    decimals: Vec<u8>,
59    /// Resolved math variant.
60    variant: CurveVariant,
61    /// Constructed math pool used for quoting.
62    pool: Pool,
63    /// Admin share scaled by 1e10. Absent in old snapshots; legacy StableSwap must refresh.
64    #[serde(default)]
65    admin_fee: Option<U256>,
66}
67
68impl CurveState {
69    /// Construct a pool with its admin fee share (scaled by 1e10). Legacy StableSwap requires
70    /// `Some(admin_fee)` to quote; NG uses its fixed 50% share and CryptoSwap ignores this field.
71    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    /// Finds the swap to `target` with the native solver. Returns `None` when the target does not
126    /// fit in U256 or the variant has no native solver. Also returns `None`, and logs the reason,
127    /// when the solver's math fails.
128    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(&quote.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        // curve_math returns dy/dx (quote per base) in native token units and fee-inclusive;
221        // rescale to human units of quote per 1 base.
222        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            // CryptoSwap retains its existing balance-only approximation.
250            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        // Soft limit: cap the input at the pool's own balance of the sell token. Beyond this the
297        // solver math becomes unreliable and output approaches the available reserve.
298        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    /// When `updated_attributes` carries [`vm::POOL_STATE_ADJUSTED`], the pool is rebuilt from
313    /// those readings. Otherwise the view getters are read from the indexed VM storage.
314    ///
315    /// The attribute exists for pending blocks, whose state never reaches that storage: an
316    /// indexer that has already read the pool under the pending block's overrides passes the
317    /// readings through instead.
318    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    /// Answers [`SwapConstraint::PoolTargetPrice`] on StableSwap pools with the native solver,
366    /// which ignores `min_amount_in` and `max_amount_in` and returns no `price_points`. All other
367    /// cases use the numerical search.
368    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    /// The TriCryptoNG USDC/WBTC/WETH pool (0x7f86bf…), balances aside.
433    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        // The readings must come from the attribute. A VM read would fail here anyway: the
482        // shared engine has no block set in this test.
483        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        // Falling back to the indexed VM state would silently price a pending block against
537        // confirmed state, so a malformed attribute must fail instead.
538        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        // 3pool state at block 24669924: DAI (18 dec) in, USDC (6 dec) out.
585        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    /// Builds `PoolTargetPrice` params for a target of `spot * multiplier`, and returns the
690    /// target as f64.
691    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(&params)
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    /// The native result must land in the lower half of the tolerance band, and the numerical
739    /// result within five times the band.
740    #[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(&params)
759            .expect("native query_pool_swap");
760        let numerical = crate::evm::query_pool_swap::query_pool_swap(&state, &params)
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    /// CryptoSwap pools delegate to the numerical search. That search rejects every target,
784    /// because `get_amount_out` keeps the stored `D` (no `tweak_price` port).
785    #[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(&params);
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(&params)
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(&params)
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(&params);
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(&params)
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}