use tycho_common::Bytes;
use crate::encoding::{evm::constants::GROUPABLE_PROTOCOLS, models::Swap};
#[derive(Clone, Debug)]
pub struct SwapGroup {
pub token_in: Bytes,
pub token_out: Bytes,
pub protocol_system: String,
pub swaps: Vec<Swap>,
pub split: f64,
}
impl PartialEq for SwapGroup {
fn eq(&self, other: &Self) -> bool {
self.token_in == other.token_in &&
self.token_out == other.token_out &&
self.protocol_system == other.protocol_system &&
self.swaps == other.swaps &&
self.split == other.split
}
}
pub fn group_swaps(swaps: &[Swap]) -> Vec<SwapGroup> {
let mut grouped_swaps: Vec<SwapGroup> = Vec::new();
let mut current_group: Option<SwapGroup> = None;
let mut last_swap_protocol = "".to_string();
let mut groupable_protocol;
let mut last_swap_out_token = Bytes::default();
for swap in swaps {
let mut current_swap_protocol = swap.component().protocol_system.clone();
if current_swap_protocol == "uniswap_v4_hooks" {
current_swap_protocol = "uniswap_v4".to_string();
};
groupable_protocol = GROUPABLE_PROTOCOLS.contains(¤t_swap_protocol.as_str());
let no_split = swap.split() == 0.0 && *swap.token_in() == last_swap_out_token;
if current_swap_protocol == last_swap_protocol && groupable_protocol && no_split {
if let Some(group) = current_group.as_mut() {
group.swaps.push(swap.clone());
group.token_out = swap.token_out().clone();
}
} else {
if let Some(group) = current_group.as_mut() {
grouped_swaps.push(group.clone());
}
current_group = Some(SwapGroup {
token_in: swap.token_in().clone(),
token_out: swap.token_out().clone(),
protocol_system: current_swap_protocol.clone(),
swaps: vec![swap.clone()],
split: swap.split(),
});
}
last_swap_protocol = current_swap_protocol;
last_swap_out_token = swap.token_out().clone();
}
if let Some(group) = current_group.as_mut() {
grouped_swaps.push(group.clone());
}
grouped_swaps
}
#[cfg(test)]
mod tests {
use std::str::FromStr;
use alloy::primitives::hex;
use tycho_common::{models::protocol::ProtocolComponent, Bytes};
use super::*;
use crate::encoding::models::Swap;
fn weth() -> Bytes {
Bytes::from(hex!("c02aaa39b223fe8d0a0e5c4f27ead9083c756cc2").to_vec())
}
#[test]
fn test_group_swaps_simple() {
let weth = weth();
let wbtc = Bytes::from_str("0x2260fac5e5542a773aa44fbcfedf7c193bc2c599").unwrap();
let usdc = Bytes::from_str("0xa0b86991c6218b36c1d19d4a2e9eb0ce3606eb48").unwrap();
let dai = Bytes::from_str("0x6b175474e89094c44da98b954eedeac495271d0f").unwrap();
let swap_weth_wbtc = Swap::new(
ProtocolComponent { protocol_system: "uniswap_v4".to_string(), ..Default::default() },
weth.clone(),
wbtc.clone(),
);
let swap_wbtc_usdc = Swap::new(
ProtocolComponent { protocol_system: "uniswap_v4".to_string(), ..Default::default() },
wbtc.clone(),
usdc.clone(),
);
let swap_usdc_dai = Swap::new(
ProtocolComponent { protocol_system: "uniswap_v2".to_string(), ..Default::default() },
usdc.clone(),
dai.clone(),
);
let swaps = vec![swap_weth_wbtc.clone(), swap_wbtc_usdc.clone(), swap_usdc_dai.clone()];
let grouped_swaps = group_swaps(&swaps);
assert_eq!(
grouped_swaps,
vec![
SwapGroup {
swaps: vec![swap_weth_wbtc, swap_wbtc_usdc],
token_in: weth,
token_out: usdc.clone(),
protocol_system: "uniswap_v4".to_string(),
split: 0f64,
},
SwapGroup {
swaps: vec![swap_usdc_dai],
token_in: usdc,
token_out: dai,
protocol_system: "uniswap_v2".to_string(),
split: 0f64,
}
]
);
}
#[test]
fn test_group_swaps_complex_split() {
let weth = weth();
let wbtc = Bytes::from_str("0x2260fac5e5542a773aa44fbcfedf7c193bc2c599").unwrap();
let usdc = Bytes::from_str("0xa0b86991c6218b36c1d19d4a2e9eb0ce3606eb48").unwrap();
let dai = Bytes::from_str("0x6b175474e89094c44da98b954eedeac495271d0f").unwrap();
let swap_wbtc_weth = Swap::new(
ProtocolComponent { protocol_system: "uniswap_v4".to_string(), ..Default::default() },
wbtc.clone(),
weth.clone(),
);
let swap_weth_usdc = Swap::new(
ProtocolComponent { protocol_system: "uniswap_v4".to_string(), ..Default::default() },
weth.clone(),
usdc.clone(),
)
.with_split(0.5f64);
let swap_weth_dai = Swap::new(
ProtocolComponent { protocol_system: "uniswap_v4".to_string(), ..Default::default() },
weth.clone(),
dai.clone(),
);
let swap_dai_usdc = Swap::new(
ProtocolComponent { protocol_system: "uniswap_v4".to_string(), ..Default::default() },
dai.clone(),
usdc.clone(),
);
let swaps = vec![
swap_wbtc_weth.clone(),
swap_weth_usdc.clone(),
swap_weth_dai.clone(),
swap_dai_usdc.clone(),
];
let grouped_swaps = group_swaps(&swaps);
assert_eq!(
grouped_swaps,
vec![
SwapGroup {
swaps: vec![swap_wbtc_weth],
token_in: wbtc.clone(),
token_out: weth.clone(),
protocol_system: "uniswap_v4".to_string(),
split: 0f64,
},
SwapGroup {
swaps: vec![swap_weth_usdc],
token_in: weth.clone(),
token_out: usdc.clone(),
protocol_system: "uniswap_v4".to_string(),
split: 0.5f64,
},
SwapGroup {
swaps: vec![swap_weth_dai, swap_dai_usdc],
token_in: weth,
token_out: usdc,
protocol_system: "uniswap_v4".to_string(),
split: 0f64,
}
]
);
}
#[test]
fn test_group_swaps_complex_split_multi_protocol() {
let weth = weth();
let wbtc = Bytes::from_str("0x2260fac5e5542a773aa44fbcfedf7c193bc2c599").unwrap();
let usdc = Bytes::from_str("0xa0b86991c6218b36c1d19d4a2e9eb0ce3606eb48").unwrap();
let dai = Bytes::from_str("0x6b175474e89094c44da98b954eedeac495271d0f").unwrap();
let swap_weth_wbtc = Swap::new(
ProtocolComponent {
protocol_system: "vm:balancer_v3".to_string(),
..Default::default()
},
weth.clone(),
wbtc.clone(),
)
.with_split(0.5f64);
let swap_wbtc_usdc = Swap::new(
ProtocolComponent {
protocol_system: "vm:balancer_v3".to_string(),
..Default::default()
},
wbtc.clone(),
usdc.clone(),
);
let swap_weth_dai = Swap::new(
ProtocolComponent { protocol_system: "uniswap_v4".to_string(), ..Default::default() },
weth.clone(),
dai.clone(),
);
let swap_dai_usdc = Swap::new(
ProtocolComponent { protocol_system: "uniswap_v4".to_string(), ..Default::default() },
dai.clone(),
usdc.clone(),
);
let swaps = vec![
swap_weth_wbtc.clone(),
swap_wbtc_usdc.clone(),
swap_weth_dai.clone(),
swap_dai_usdc.clone(),
];
let grouped_swaps = group_swaps(&swaps);
assert_eq!(
grouped_swaps,
vec![
SwapGroup {
swaps: vec![swap_weth_wbtc, swap_wbtc_usdc],
token_in: weth.clone(),
token_out: usdc.clone(),
protocol_system: "vm:balancer_v3".to_string(),
split: 0.5f64,
},
SwapGroup {
swaps: vec![swap_weth_dai, swap_dai_usdc],
token_in: weth,
token_out: usdc,
protocol_system: "uniswap_v4".to_string(),
split: 0f64,
}
]
);
}
#[test]
fn test_group_swaps_uniswap_v4_with_hooks() {
let weth = weth();
let wbtc = Bytes::from_str("0x2260fac5e5542a773aa44fbcfedf7c193bc2c599").unwrap();
let usdc = Bytes::from_str("0xa0b86991c6218b36c1d19d4a2e9eb0ce3606eb48").unwrap();
let dai = Bytes::from_str("0x6b175474e89094c44da98b954eedeac495271d0f").unwrap();
let swap_weth_wbtc = Swap::new(
ProtocolComponent { protocol_system: "uniswap_v4".to_string(), ..Default::default() },
weth.clone(),
wbtc.clone(),
);
let swap_wbtc_usdc = Swap::new(
ProtocolComponent {
protocol_system: "uniswap_v4_hooks".to_string(),
..Default::default()
},
wbtc.clone(),
usdc.clone(),
);
let swap_usdc_dai = Swap::new(
ProtocolComponent { protocol_system: "uniswap_v2".to_string(), ..Default::default() },
usdc.clone(),
dai.clone(),
);
let swaps = vec![swap_weth_wbtc.clone(), swap_wbtc_usdc.clone(), swap_usdc_dai.clone()];
let grouped_swaps = group_swaps(&swaps);
assert_eq!(grouped_swaps.len(), 2);
assert_eq!(grouped_swaps[0].swaps.len(), 2);
assert_eq!(grouped_swaps[0].token_in, weth);
assert_eq!(grouped_swaps[0].token_out, usdc.clone());
assert_eq!(grouped_swaps[0].protocol_system, "uniswap_v4");
assert_eq!(grouped_swaps[1].swaps.len(), 1);
assert_eq!(grouped_swaps[1].token_in, usdc);
assert_eq!(grouped_swaps[1].token_out, dai);
assert_eq!(grouped_swaps[1].protocol_system, "uniswap_v2");
}
}