use crate::common::errors::PoolError;
use crate::common::pool_base::PoolBase;
use crate::common::types::*;
use crate::common::utils::{
compute_and_charge_aggregate_swap_fees_raw, copy_to_scaled18_apply_rate_round_up_array,
get_single_input_index, require_unbalanced_liquidity_enabled, to_raw_undo_rate_round_down,
};
use crate::hooks::types::HookState;
use crate::hooks::HookBase;
use crate::vault::base_pool_math::{
compute_proportional_amounts_out, compute_remove_liquidity_single_token_exact_in,
compute_remove_liquidity_single_token_exact_out,
};
use alloy_primitives::U256;
pub fn remove_liquidity(
remove_liquidity_input: &RemoveLiquidityInput,
pool_state: &PoolState,
pool_class: &dyn PoolBase,
hook_class: &dyn HookBase,
hook_state: Option<&HookState>,
) -> Result<RemoveLiquidityResult, PoolError> {
let base_state = pool_state.base();
let min_amounts_out_scaled18 = copy_to_scaled18_apply_rate_round_up_array(
&remove_liquidity_input.min_amounts_out_raw,
&base_state.scaling_factors,
&base_state.token_rates,
)?;
let mut updated_balances_live_scaled18 = base_state.balances_live_scaled_18.clone();
if hook_class.config().should_call_before_remove_liquidity {
let hook_return = hook_class.on_before_remove_liquidity(
remove_liquidity_input.kind.clone(),
&remove_liquidity_input.max_bpt_amount_in_raw,
&remove_liquidity_input.min_amounts_out_raw,
&updated_balances_live_scaled18,
hook_state.unwrap(),
);
if !hook_return.success {
return Err(PoolError::BeforeRemoveLiquidityHookFailed);
}
for (i, adjusted_balance) in hook_return
.hook_adjusted_balances_scaled_18
.iter()
.enumerate()
{
updated_balances_live_scaled18[i] = *adjusted_balance;
}
}
let (bpt_amount_in, amounts_out_scaled18, swap_fee_amounts_scaled18) =
match remove_liquidity_input.kind {
RemoveLiquidityKind::Proportional => {
let bpt_amount_in = remove_liquidity_input.max_bpt_amount_in_raw;
let swap_fee_amounts_scaled18 = vec![U256::ZERO; base_state.tokens.len()];
let amounts_out_scaled18 = compute_proportional_amounts_out(
&updated_balances_live_scaled18,
&base_state.total_supply,
&remove_liquidity_input.max_bpt_amount_in_raw,
)?;
(
bpt_amount_in,
amounts_out_scaled18,
swap_fee_amounts_scaled18,
)
}
RemoveLiquidityKind::SingleTokenExactIn => {
require_unbalanced_liquidity_enabled(pool_state)?;
let bpt_amount_in = remove_liquidity_input.max_bpt_amount_in_raw;
let mut amounts_out_scaled18 = min_amounts_out_scaled18.clone();
let token_out_index =
get_single_input_index(&remove_liquidity_input.min_amounts_out_raw)?;
let computed = compute_remove_liquidity_single_token_exact_in(
&updated_balances_live_scaled18,
token_out_index,
&remove_liquidity_input.max_bpt_amount_in_raw,
&base_state.total_supply,
&base_state.swap_fee,
&pool_class.get_minimum_invariant_ratio(),
&|balances, token_out_index, invariant_ratio| {
pool_class.compute_balance(balances, token_out_index, invariant_ratio)
},
)?;
amounts_out_scaled18[token_out_index] = computed.amount_out_with_fee;
(
bpt_amount_in,
amounts_out_scaled18,
computed.swap_fee_amounts,
)
}
RemoveLiquidityKind::SingleTokenExactOut => {
require_unbalanced_liquidity_enabled(pool_state)?;
let amounts_out_scaled18 = min_amounts_out_scaled18.clone();
let token_out_index =
get_single_input_index(&remove_liquidity_input.min_amounts_out_raw)?;
let computed = compute_remove_liquidity_single_token_exact_out(
&updated_balances_live_scaled18,
token_out_index,
&amounts_out_scaled18[token_out_index],
&base_state.total_supply,
&base_state.swap_fee,
&pool_class.get_minimum_invariant_ratio(),
&|balances, rounding| pool_class.compute_invariant(balances, rounding),
)?;
let bpt_amount_in = computed.bpt_amount_in;
(
bpt_amount_in,
amounts_out_scaled18,
computed.swap_fee_amounts,
)
}
};
let mut amounts_out_raw = vec![U256::ZERO; base_state.tokens.len()];
for i in 0..base_state.tokens.len() {
amounts_out_raw[i] = to_raw_undo_rate_round_down(
&amounts_out_scaled18[i],
&base_state.scaling_factors[i],
&base_state.token_rates[i],
)?;
let aggregate_swap_fee_amount_raw = compute_and_charge_aggregate_swap_fees_raw(
&swap_fee_amounts_scaled18[i],
&base_state.aggregate_swap_fee,
&base_state.scaling_factors,
&base_state.token_rates,
i,
)?;
updated_balances_live_scaled18[i] -=
amounts_out_scaled18[i] + aggregate_swap_fee_amount_raw;
}
if hook_class.config().should_call_after_remove_liquidity {
let hook_return = hook_class.on_after_remove_liquidity(
remove_liquidity_input.kind.clone(),
&bpt_amount_in,
&amounts_out_scaled18,
&amounts_out_raw,
&updated_balances_live_scaled18,
hook_state.unwrap(),
);
if !hook_return.success
|| hook_return.hook_adjusted_amounts_out_raw.len() != amounts_out_raw.len()
{
return Err(PoolError::AfterRemoveLiquidityHookFailed);
}
if hook_class.config().enable_hook_adjusted_amounts {
for (i, adjusted_amount) in hook_return.hook_adjusted_amounts_out_raw.iter().enumerate()
{
amounts_out_raw[i] = *adjusted_amount;
}
}
}
Ok(RemoveLiquidityResult {
bpt_amount_in_raw: bpt_amount_in,
amounts_out_raw,
})
}