use crate::num::arithmetic::traits::ModIsReduced;
use crate::num::basic::traits::Zero;
use crate::num::basic::unsigneds::PrimitiveUnsigned;
use crate::polynomial::{ModSquareTruncated, ModSquareTruncatedAssign};
use crate::unsigned_polynomial::UnsignedPolynomial;
use crate::unsigned_polynomial::arithmetic::mod_mul::{
MOD_SQUARE_KARATSUBA_THRESHOLD, ModData, mod_add_assign_slice,
};
use crate::unsigned_polynomial::arithmetic::mod_mul_truncated::mod_mul_truncated_karatsuba;
use crate::unsigned_polynomial::arithmetic::mod_power_of_2_mul::from_coefficients_trimmed;
use crate::unsigned_polynomial::arithmetic::mod_power_of_2_mul_truncated::truncated_len;
use crate::unsigned_polynomial::arithmetic::mod_square::{
mod_square_classical_prefix, mod_square_karatsuba,
};
use alloc::vec;
use core::cmp::min;
pub(crate) fn mod_square_truncated_karatsuba<T: PrimitiveUnsigned>(
out: &mut [T],
xs: &[T],
d: &ModData<T>,
) {
let len = out.len();
let xs = &xs[..min(xs.len(), len)];
let n = xs.len();
let full_len = (n << 1) - 1;
if full_len <= len {
mod_square_karatsuba(&mut out[..full_len], xs, d);
out[full_len..].fill(T::ZERO);
return;
}
if n < MOD_SQUARE_KARATSUBA_THRESHOLD {
mod_square_classical_prefix(out, xs, d);
return;
}
let h = len.div_ceil(2);
let x0 = &xs[..min(h, n)];
mod_square_truncated_karatsuba(out, x0, d);
if n > h {
let mut cross = vec![T::ZERO; len - h];
mod_mul_truncated_karatsuba(&mut cross, x0, &xs[h..], d);
mod_add_assign_slice(&mut out[h..], &cross, d.m);
mod_add_assign_slice(&mut out[h..], &cross, d.m);
}
}
fn assert_lengths<T>(out: &[T], xs: &[T]) {
assert!(!out.is_empty());
assert!(!xs.is_empty());
}
crate_test_fn! {
#[allow(dead_code)]
mod_square_truncated_to_out_classical<T: PrimitiveUnsigned>(out: &mut [T], xs: &[T], m: T) {
assert_lengths(out, xs);
mod_square_classical_prefix(out, xs, &ModData::new(m, xs.len().min(out.len())));
}}
crate_test_fn! {
#[allow(dead_code)]
mod_square_truncated_to_out_karatsuba<T: PrimitiveUnsigned>(out: &mut [T], xs: &[T], m: T) {
assert_lengths(out, xs);
mod_square_truncated_karatsuba(out, xs, &ModData::new(m, xs.len().min(out.len())));
}}
#[doc(hidden)]
pub fn mod_square_truncated_to_out<T: PrimitiveUnsigned>(out: &mut [T], xs: &[T], m: T) {
assert_lengths(out, xs);
mod_square_truncated_karatsuba(out, xs, &ModData::new(m, xs.len().min(out.len())));
}
fn assert_reduced<T: PrimitiveUnsigned>(p: &UnsignedPolynomial<T>, m: T) {
assert!(
p.mod_is_reduced(&m),
"self must be reduced mod m, but {p} has a coefficient >= {m}"
);
}
pub(crate) fn mod_square_truncated_helper<T: PrimitiveUnsigned>(
xs: &[T],
len: u64,
m: T,
) -> UnsignedPolynomial<T> {
if len == 0 || xs.is_empty() {
return UnsignedPolynomial::ZERO;
}
let mut out = vec![T::ZERO; truncated_len(xs.len(), xs.len(), len)];
mod_square_truncated_to_out(&mut out, xs, m);
from_coefficients_trimmed(out)
}
impl<T: PrimitiveUnsigned> ModSquareTruncated<T> for UnsignedPolynomial<T> {
type Output = Self;
fn mod_square_truncated(self, len: u64, m: T) -> Self {
assert_reduced(&self, m);
mod_square_truncated_helper(&self.coefficients, len, m)
}
}
impl<T: PrimitiveUnsigned> ModSquareTruncated<T> for &UnsignedPolynomial<T> {
type Output = UnsignedPolynomial<T>;
fn mod_square_truncated(self, len: u64, m: T) -> UnsignedPolynomial<T> {
assert_reduced(self, m);
mod_square_truncated_helper(&self.coefficients, len, m)
}
}
impl<T: PrimitiveUnsigned> ModSquareTruncatedAssign<T> for UnsignedPolynomial<T> {
fn mod_square_truncated_assign(&mut self, len: u64, m: T) {
assert_reduced(self, m);
*self = mod_square_truncated_helper(&self.coefficients, len, m);
}
}