use core::fmt;
use crate::isqrt;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) struct SegmentedSieveError;
impl fmt::Display for SegmentedSieveError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "the upper limit was smaller than `N`")
}
}
impl core::error::Error for SegmentedSieveError {}
#[must_use = "the function only returns a new value and does not modify its inputs"]
pub(crate) const fn sieve_segment<const N: usize>(
base_sieve: &[bool; N],
upper_limit: u64,
) -> Result<[bool; N], SegmentedSieveError> {
let mut segment_sieve = [true; N];
let Some(lower_limit) = upper_limit.checked_sub(N as u64) else {
return Err(SegmentedSieveError);
};
if lower_limit == 0 && N > 1 {
return Ok(*base_sieve);
} else if lower_limit == 1 && N > 0 {
segment_sieve[0] = false;
}
let mut i = 0;
while i < N {
if base_sieve[i] {
let prime = i as u64;
let mut composite = (lower_limit / prime) * prime;
if composite < lower_limit {
composite += prime;
}
if composite == prime {
composite += prime;
}
while composite < upper_limit {
segment_sieve[(composite - lower_limit) as usize] = false;
composite += prime;
}
}
i += 1;
}
Ok(segment_sieve)
}
#[must_use = "the function only returns a new value and does not modify its input"]
pub const fn sieve_lt<const N: usize, const MEM: usize>(
upper_limit: u64,
) -> Result<[bool; N], SieveError> {
const { assert!(MEM >= N, "`MEM` must be at least as large as `N`") }
let mem_sqr = const {
let mem64 = MEM as u64;
match mem64.checked_mul(mem64) {
Some(mem_sqr) => mem_sqr,
None => panic!("`MEM`^2 must fit in a `u64`"),
}
};
if upper_limit > mem_sqr {
return Err(SieveError::TooSmallSieveSize);
}
let n64 = N as u64;
if upper_limit < n64 {
return Err(SieveError::TooSmallLimit);
}
if N == 0 {
return Ok([false; N]);
}
if upper_limit == n64 {
return Ok(sieve());
}
let base_sieve: [bool; MEM] = sieve();
let (offset, upper_sieve) = match sieve_segment(&base_sieve, upper_limit) {
Ok(res) => (0, res),
Err(_) => ((MEM as u64 - upper_limit) as usize, base_sieve),
};
let mut i = 0;
let mut ans = [false; N];
while i < N {
ans[N - 1 - i] = upper_sieve[MEM - 1 - i - offset];
i += 1;
}
Ok(ans)
}
#[must_use = "the function only returns a new value"]
pub const fn sieve<const N: usize>() -> [bool; N] {
let mut sieve = [true; N];
if N == 0 {
return sieve;
}
if N > 0 {
sieve[0] = false;
}
if N > 1 {
sieve[1] = false;
}
let mut number: usize = 2;
let bound = isqrt(N as u64);
while (number as u64) <= bound {
if sieve[number] {
let Some(mut composite) = number.checked_mul(number) else {
break;
};
while composite < N {
sieve[composite] = false;
composite = match composite.checked_add(number) {
Some(sum) => sum,
None => break,
};
}
}
number += 1;
}
sieve
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
#[cfg_attr(
feature = "rkyv",
derive(rkyv::Archive, rkyv::Serialize, rkyv::Deserialize)
)]
pub enum SieveError {
TooSmallLimit,
TooSmallSieveSize,
TotalDoesntFitU64,
}
impl fmt::Display for SieveError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::TooSmallLimit => write!(f, "`limit` must be at least `N`"),
Self::TooSmallSieveSize => {
write!(f, "`MEM`^2 was smaller than the largest encountered value")
}
Self::TotalDoesntFitU64 => write!(f, "`MEM + limit` must fit in a `u64`"),
}
}
}
impl core::error::Error for SieveError {}
#[must_use = "the function only returns a new value and does not modify its input"]
pub const fn sieve_geq<const N: usize, const MEM: usize>(
lower_limit: u64,
) -> Result<[bool; N], SieveError> {
const { assert!(MEM >= N, "`MEM` must be at least as large as `N`") }
let (mem64, mem_sqr) = const {
let mem64 = MEM as u64;
match mem64.checked_mul(mem64) {
Some(mem_sqr) => (mem64, mem_sqr),
None => panic!("`MEM`^2 must fit in a `u64`"),
}
};
let Some(upper_limit) = mem64.checked_add(lower_limit) else {
return Err(SieveError::TotalDoesntFitU64);
};
if upper_limit > mem_sqr {
return Err(SieveError::TooSmallSieveSize);
}
if N == 0 {
return Ok([false; N]);
}
if lower_limit == 0 {
return Ok(sieve());
}
let base_sieve: [bool; MEM] = sieve();
let Ok(upper_sieve) = sieve_segment(&base_sieve, upper_limit) else {
panic!("this is already checked above")
};
let mut ans = [false; N];
let mut i = 0;
while i < N {
ans[i] = upper_sieve[i];
i += 1;
}
Ok(ans)
}
#[macro_export]
macro_rules! sieve_segment {
($n:expr; < $lim:expr) => {
$crate::sieve_lt::<
{ $n },
{
let mem: u64 = { $lim };
$crate::isqrt(mem) as ::core::primitive::usize + 1
},
>({ $lim })
};
($n:expr; >= $lim:expr) => {
$crate::sieve_geq::<
{ $n },
{
let mem: u64 = { $lim };
$crate::isqrt(mem) as ::core::primitive::usize + 1 + { $n }
},
>({ $lim })
};
}
#[cfg(test)]
mod test {
use crate::SieveError;
use super::{sieve, sieve_geq, sieve_lt, sieve_segment, SegmentedSieveError};
#[test]
fn test_consistency_of_sieve_segment() {
const P: [bool; 10] = match sieve_segment(&sieve(), 10) {
Ok(s) => s,
Err(_) => panic!(),
};
const PP: [bool; 10] = match sieve_segment(&sieve(), 11) {
Ok(s) => s,
Err(_) => panic!(),
};
assert_eq!(P, sieve());
assert_eq!(PP, sieve::<11>()[1..]);
assert_eq!(
sieve_segment::<5>(&[false, false, true, true, false], 4),
Err(SegmentedSieveError)
);
assert_eq!(sieve_segment(&sieve::<5>(), 5), Ok(sieve()));
}
#[test]
fn test_sieve_lt() {
assert_eq!(sieve_lt::<5, 5>(30), Err(SieveError::TooSmallSieveSize));
assert_eq!(sieve_lt::<5, 5>(4), Err(SieveError::TooSmallLimit));
assert_eq!(sieve_lt::<5, 5>(5), Ok(sieve()));
assert_eq!(sieve_lt::<2, 5>(20), Ok([false, true]));
}
#[test]
fn test_sieve() {
assert_eq!(sieve(), [false; 0]);
}
#[test]
fn test_sieve_geq() {
assert_eq!(
sieve_geq::<5, 5>(u64::MAX),
Err(SieveError::TotalDoesntFitU64)
);
assert_eq!(sieve_geq::<5, 5>(30), Err(SieveError::TooSmallSieveSize));
assert_eq!(sieve_geq::<0, 1>(0), Ok([false; 0]))
}
}