use num_bigint::BigUint;
use crate::error::CombError;
pub const MAX_FACTORIAL_INPUT: u64 = 50_000;
pub const MAX_PARTITION_INPUT: u64 = 10_000;
pub const MAX_BINOMIAL_INPUT: u64 = MAX_FACTORIAL_INPUT;
fn big(n: u64) -> BigUint {
BigUint::from(n)
}
pub fn factorial(n: u64) -> BigUint {
(1..=n).fold(BigUint::from(1u32), |acc, i| acc * big(i))
}
pub fn factorial_checked(n: u64) -> Result<BigUint, CombError> {
if n > MAX_FACTORIAL_INPUT {
return Err(CombError::LimitExceeded(format!(
"factorial input {n} exceeds maximum {MAX_FACTORIAL_INPUT}"
)));
}
Ok(factorial(n))
}
pub fn falling_factorial(n: u64, k: u64) -> BigUint {
if k > n {
return BigUint::from(0u32);
}
(0..k).fold(BigUint::from(1u32), |acc, i| acc * big(n - i))
}
pub fn permutation_count(n: u64, k: u64) -> BigUint {
falling_factorial(n, k)
}
pub fn binomial(n: u64, k: u64) -> BigUint {
if k > n {
return BigUint::from(0u32);
}
let k = k.min(n - k);
let mut result = BigUint::from(1u32);
for i in 1..=k {
result = result * big(n - k + i) / big(i);
}
result
}
pub fn binomial_checked(n: u64, k: u64, max_iter: u64) -> Result<BigUint, CombError> {
let effective = k.min(n.saturating_sub(k));
if effective > max_iter {
return Err(CombError::LimitExceeded(format!(
"binomial k={effective} exceeds maximum {max_iter}"
)));
}
Ok(binomial(n, k))
}
pub fn multinomial(parts: &[u64]) -> BigUint {
let total: u64 = parts.iter().sum();
let mut denom = BigUint::from(1u32);
for &p in parts {
denom *= factorial(p);
}
factorial(total) / denom
}
pub fn stirling2(n: u64, k: u64) -> BigUint {
let (n, k) = (n as usize, k as usize);
let mut dp = vec![BigUint::from(0u32); k + 1];
dp[0] = BigUint::from(1u32); for _ in 1..=n {
let mut next = vec![BigUint::from(0u32); k + 1];
for j in 1..=k {
next[j] = big(j as u64) * &dp[j] + &dp[j - 1];
}
dp = next;
}
if k < dp.len() {
dp[k].clone()
} else {
BigUint::from(0u32)
}
}
pub fn bell_number(n: u64) -> BigUint {
(0..=n).map(|k| stirling2(n, k)).sum()
}
pub fn integer_partition_count(n: u64) -> BigUint {
let n = n as usize;
let mut dp = vec![BigUint::from(0u32); n + 1];
dp[0] = BigUint::from(1u32);
for part in 1..=n {
for j in part..=n {
dp[j] = &dp[j] + &dp[j - part].clone();
}
}
dp[n].clone()
}
pub fn integer_partition_count_checked(n: u64) -> Result<BigUint, CombError> {
if n > MAX_PARTITION_INPUT {
return Err(CombError::LimitExceeded(format!(
"partition-count input {n} exceeds maximum {MAX_PARTITION_INPUT}"
)));
}
Ok(integer_partition_count(n))
}
#[cfg(test)]
mod tests {
use super::*;
fn b(n: u64) -> BigUint {
BigUint::from(n)
}
#[test]
fn factorial_basics() {
assert_eq!(factorial(0), b(1));
assert_eq!(factorial(5), b(120));
}
#[test]
fn binomial_and_permutations() {
assert_eq!(binomial(5, 2), b(10));
assert_eq!(binomial(10, 5), b(252));
assert_eq!(binomial(5, 7), b(0));
assert_eq!(binomial(6, 0), b(1));
assert_eq!(permutation_count(5, 3), b(60));
assert_eq!(falling_factorial(5, 2), b(20));
}
#[test]
fn multinomial_value() {
assert_eq!(multinomial(&[2, 1, 1]), b(12));
}
#[test]
fn stirling_bell_and_partitions() {
assert_eq!(stirling2(4, 2), b(7));
assert_eq!(bell_number(4), b(15));
assert_eq!(integer_partition_count(5), b(7));
}
#[test]
fn checked_counts_accept_small_inputs() {
assert_eq!(factorial_checked(5).unwrap(), b(120));
assert_eq!(integer_partition_count_checked(5).unwrap(), b(7));
assert_eq!(binomial_checked(5, 2, MAX_BINOMIAL_INPUT).unwrap(), b(10));
assert_eq!(binomial_checked(2, 5, MAX_BINOMIAL_INPUT).unwrap(), b(0));
}
#[test]
fn checked_counts_reject_huge_inputs() {
assert!(matches!(
factorial_checked(MAX_FACTORIAL_INPUT + 1),
Err(CombError::LimitExceeded(_))
));
assert!(matches!(
integer_partition_count_checked(u64::MAX),
Err(CombError::LimitExceeded(_))
));
assert!(matches!(
binomial_checked(1_000_000_000_000, 500_000_000_000, MAX_BINOMIAL_INPUT),
Err(CombError::LimitExceeded(_))
));
}
}