use crate::error::{Error, Result};
use core::fmt;
pub const BACKEND_ENV_VAR: &str = "HYALITE_BACKEND";
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
#[non_exhaustive]
pub enum Backend {
Scalar,
Sse41,
Avx2,
Neon,
}
impl Backend {
#[must_use]
pub const fn name(self) -> &'static str {
match self {
Backend::Scalar => "scalar",
Backend::Sse41 => "sse4.1",
Backend::Avx2 => "avx2",
Backend::Neon => "neon",
}
}
#[must_use]
pub(crate) fn simd_lanes(self) -> Option<usize> {
match self {
Backend::Scalar => None,
Backend::Sse41 | Backend::Neon => Some(16),
Backend::Avx2 => Some(32),
}
}
#[must_use]
pub fn is_available(self) -> bool {
match self {
Backend::Scalar => true,
Backend::Sse41 => sse41_detected(),
Backend::Avx2 => avx2_detected(),
Backend::Neon => neon_available(),
}
}
}
#[cfg(target_arch = "x86_64")]
fn sse41_detected() -> bool {
std::is_x86_feature_detected!("sse4.1")
}
#[cfg(not(target_arch = "x86_64"))]
fn sse41_detected() -> bool {
false
}
#[cfg(target_arch = "x86_64")]
fn avx2_detected() -> bool {
std::is_x86_feature_detected!("avx2")
}
#[cfg(not(target_arch = "x86_64"))]
fn avx2_detected() -> bool {
false
}
#[cfg(target_arch = "aarch64")]
fn neon_available() -> bool {
true
}
#[cfg(not(target_arch = "aarch64"))]
fn neon_available() -> bool {
false
}
impl fmt::Display for Backend {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(self.name())
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Default)]
pub enum BackendChoice {
#[default]
Auto,
Force(Backend),
}
impl BackendChoice {
pub fn parse(s: &str) -> Result<Self> {
match s.trim().to_ascii_lowercase().as_str() {
"auto" => Ok(BackendChoice::Auto),
"scalar" => Ok(BackendChoice::Force(Backend::Scalar)),
"sse4.1" | "sse41" => Ok(BackendChoice::Force(Backend::Sse41)),
"avx2" => Ok(BackendChoice::Force(Backend::Avx2)),
"neon" => Ok(BackendChoice::Force(Backend::Neon)),
_ => Err(Error::InvalidBackendName {
name: s.to_string(),
}),
}
}
}
fn detect_best() -> Backend {
if Backend::Avx2.is_available() {
Backend::Avx2
} else if Backend::Sse41.is_available() {
Backend::Sse41
} else if Backend::Neon.is_available() {
Backend::Neon
} else {
Backend::Scalar
}
}
pub(crate) fn resolve(choice: BackendChoice) -> Result<Backend> {
match choice {
BackendChoice::Auto => Ok(detect_best()),
BackendChoice::Force(backend) => {
if backend.is_available() {
Ok(backend)
} else {
Err(Error::BackendUnavailable { backend })
}
}
}
}
pub(crate) fn choice_from_env() -> Result<Option<BackendChoice>> {
match std::env::var(BACKEND_ENV_VAR) {
Ok(s) if s.trim().is_empty() => Ok(None),
Ok(s) => BackendChoice::parse(&s).map(Some),
Err(std::env::VarError::NotPresent) => Ok(None),
Err(std::env::VarError::NotUnicode(_)) => Err(Error::InvalidBackendName {
name: "<non-unicode>".to_string(),
}),
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn availability_matches_arch_and_cpu() {
assert!(Backend::Scalar.is_available());
assert_eq!(Backend::Sse41.is_available(), sse41_detected());
assert_eq!(Backend::Avx2.is_available(), avx2_detected());
assert_eq!(Backend::Neon.is_available(), neon_available());
#[cfg(not(target_arch = "x86_64"))]
{
assert!(!Backend::Sse41.is_available());
assert!(!Backend::Avx2.is_available());
}
#[cfg(not(target_arch = "aarch64"))]
assert!(!Backend::Neon.is_available());
}
#[test]
fn name_round_trips_through_parse_for_every_backend() {
for b in [
Backend::Scalar,
Backend::Sse41,
Backend::Avx2,
Backend::Neon,
] {
assert_eq!(
BackendChoice::parse(b.name()).unwrap(),
BackendChoice::Force(b)
);
assert_eq!(b.to_string(), b.name());
}
}
#[test]
fn parse_accepts_aliases_and_is_case_insensitive() {
assert_eq!(BackendChoice::parse("auto").unwrap(), BackendChoice::Auto);
assert_eq!(BackendChoice::parse("AUTO").unwrap(), BackendChoice::Auto);
assert_eq!(
BackendChoice::parse(" SSE41 ").unwrap(),
BackendChoice::Force(Backend::Sse41)
);
assert_eq!(
BackendChoice::parse("sse4.1").unwrap(),
BackendChoice::Force(Backend::Sse41)
);
assert_eq!(
BackendChoice::parse("Avx2").unwrap(),
BackendChoice::Force(Backend::Avx2)
);
}
#[test]
fn parse_rejects_unknown_names() {
for bad in ["", "sse2", "ssse3", "avx512", "gpu", "x"] {
let err = BackendChoice::parse(bad).unwrap_err();
assert_eq!(
err,
Error::InvalidBackendName {
name: bad.to_string()
}
);
}
}
#[test]
fn resolve_auto_picks_an_available_backend() {
let b = resolve(BackendChoice::Auto).unwrap();
assert!(b.is_available());
assert_eq!(b, detect_best());
}
#[test]
fn resolve_forcing_scalar_always_succeeds() {
assert_eq!(
resolve(BackendChoice::Force(Backend::Scalar)).unwrap(),
Backend::Scalar
);
}
#[test]
fn resolve_forcing_a_backend_tracks_its_availability() {
for b in [Backend::Sse41, Backend::Avx2, Backend::Neon] {
let got = resolve(BackendChoice::Force(b));
if b.is_available() {
assert_eq!(got.unwrap(), b);
} else {
assert_eq!(got.unwrap_err(), Error::BackendUnavailable { backend: b });
}
}
}
}