use crate::common::errors::PoolError;
use crate::common::pool_base::PoolBase;
use crate::common::utils::{
compute_and_charge_aggregate_swap_fees_raw, copy_to_scaled18_apply_rate_round_down_array,
get_single_input_index, require_unbalanced_liquidity_enabled, to_raw_undo_rate_round_up,
};
use crate::common::{to_scaled_18_apply_rate_round_down, types::*};
use crate::hooks::types::HookState;
use crate::hooks::HookBase;
use alloy_primitives::U256;
pub fn add_liquidity(
add_liquidity_input: &AddLiquidityInput,
pool_state: &PoolState,
pool_class: &dyn PoolBase,
hook_class: &dyn HookBase,
hook_state: Option<&HookState>,
) -> Result<AddLiquidityResult, PoolError> {
let base_state = pool_state.base();
let max_amounts_in_scaled18 = copy_to_scaled18_apply_rate_round_down_array(
&add_liquidity_input.max_amounts_in_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_add_liquidity {
let hook_return = hook_class.on_before_add_liquidity(
add_liquidity_input.kind.clone(),
&add_liquidity_input.max_amounts_in_raw, &add_liquidity_input.min_bpt_amount_out_raw,
&updated_balances_live_scaled18,
hook_state.unwrap(),
);
if !hook_return.success {
return Err(PoolError::BeforeAddLiquidityHookFailed);
}
for (i, adjusted_balance) in hook_return
.hook_adjusted_balances_scaled_18
.iter()
.enumerate()
{
updated_balances_live_scaled18[i] = *adjusted_balance;
}
}
let mut amounts_in_scaled18 = vec![U256::ZERO; base_state.tokens.len()];
let (bpt_amount_out, swap_fee_amounts_scaled18) = match add_liquidity_input.kind {
AddLiquidityKind::Unbalanced => {
require_unbalanced_liquidity_enabled(pool_state)?;
amounts_in_scaled18 = max_amounts_in_scaled18.clone();
let computed = crate::vault::base_pool_math::compute_add_liquidity_unbalanced(
&updated_balances_live_scaled18,
&max_amounts_in_scaled18,
&base_state.total_supply,
&base_state.swap_fee,
&pool_class.get_maximum_invariant_ratio(),
&|balances, rounding| pool_class.compute_invariant(balances, rounding),
)?;
(computed.bpt_amount_out, computed.swap_fee_amounts)
}
AddLiquidityKind::SingleTokenExactOut => {
require_unbalanced_liquidity_enabled(pool_state)?;
let token_index = get_single_input_index(&max_amounts_in_scaled18)?;
let bpt_amount_out = add_liquidity_input.min_bpt_amount_out_raw;
let computed =
crate::vault::base_pool_math::compute_add_liquidity_single_token_exact_out(
&updated_balances_live_scaled18,
token_index,
&bpt_amount_out,
&base_state.total_supply,
&base_state.swap_fee,
&pool_class.get_maximum_invariant_ratio(),
&|balances, token_in_index, invariant_ratio| {
pool_class.compute_balance(balances, token_in_index, invariant_ratio)
},
)?;
amounts_in_scaled18[token_index] = computed.amount_in_with_fee;
(bpt_amount_out, computed.swap_fee_amounts)
}
};
let mut amounts_in_raw = vec![U256::ZERO; base_state.tokens.len()];
for i in 0..base_state.tokens.len() {
amounts_in_raw[i] = to_raw_undo_rate_round_up(
&amounts_in_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,
)?;
let aggregate_swap_fee_amount_scaled_18 = to_scaled_18_apply_rate_round_down(
&aggregate_swap_fee_amount_raw,
&base_state.scaling_factors[i],
&base_state.token_rates[i],
)?;
updated_balances_live_scaled18[i] = updated_balances_live_scaled18[i]
+ amounts_in_scaled18[i]
- aggregate_swap_fee_amount_scaled_18;
}
if hook_class.config().should_call_after_add_liquidity {
let hook_return = hook_class.on_after_add_liquidity(
add_liquidity_input.kind.clone(),
&amounts_in_scaled18,
&amounts_in_raw,
&bpt_amount_out,
&updated_balances_live_scaled18,
hook_state.unwrap(),
);
if !hook_return.success
|| hook_return.hook_adjusted_amounts_in_raw.len() != amounts_in_raw.len()
{
return Err(PoolError::AfterAddLiquidityHookFailed);
}
if hook_class.config().enable_hook_adjusted_amounts {
for (i, adjusted_amount) in hook_return.hook_adjusted_amounts_in_raw.iter().enumerate()
{
amounts_in_raw[i] = *adjusted_amount;
}
}
}
Ok(AddLiquidityResult {
bpt_amount_out_raw: bpt_amount_out,
amounts_in_raw,
})
}