use crate::arch;
use crate::arch::{Vector, MoveMask};
#[repr(transparent)]
pub struct ErrMsg {
pub msg: &'static str
}
impl ErrMsg {
#[inline] #[must_use]
pub const fn new(msg: &'static str) -> Self {
Self { msg }
}
}
impl From<&'static str> for ErrMsg {
#[inline]
fn from(value: &'static str) -> Self {
Self::new(value)
}
}
impl core::fmt::Display for ErrMsg {
#[inline]
fn fmt(&self, f: &mut core::fmt::Formatter) -> core::fmt::Result {
f.write_str(self.msg)
}
}
impl core::fmt::Debug for ErrMsg {
fn fmt(&self, f: &mut core::fmt::Formatter) -> core::fmt::Result {
write!(f, "Unsatisfied Requirement: {}", self.msg)
}
}
impl core::ops::Deref for ErrMsg {
type Target = str;
#[inline]
fn deref(&self) -> &Self::Target {
self.msg
}
}
#[cfg(feature = "std")]
impl std::error::Error for ErrMsg {}
pub trait Condition {
type Error;
#[must_use]
fn check(&mut self, vector: Vector) -> MoveMask;
fn ok(self) -> Result<(), Self::Error>;
}
pub trait Requirement {
type Error;
fn check(&mut self, vector: Vector);
fn check_partial(&mut self, vector: Vector, len: u32);
fn result(self) -> (bool, Result<(), Self::Error>);
fn results(self) -> (bool, impl Iterator<Item = Result<(), Self::Error>>);
}
pub struct Requires<C, Raise, Err>
where
C: Fn(Vector) -> Vector,
Raise: FnOnce() -> Err
{
pub cond: C,
raise: Raise,
seen: bool
}
impl<C, Raise, Err> Requires<C, Raise, Err>
where
C: Fn(Vector) -> Vector,
Raise: FnOnce() -> Err
{
#[inline] #[must_use]
pub const fn new(cond: C, raise: Raise) -> Self {
Self { cond, raise, seen: false }
}
}
impl<C, Raise, Err> Condition for Requires<C, Raise, Err>
where
C: Fn(Vector) -> Vector,
Raise: FnOnce() -> Err
{
type Error = Err;
#[inline] #[must_use]
fn check(&mut self, vector: Vector) -> MoveMask {
let mask = unsafe { MoveMask::new((self.cond)(vector)) };
self.seen |= mask.any_bit_set();
mask
}
#[inline(always)]
fn ok(self) -> Result<(), Self::Error> {
if self.seen { Ok(()) } else { Err((self.raise)()) }
}
}
#[macro_export]
macro_rules! requirement {
(
$(#[$attr:meta])*
$vis:vis $req_name:ident => $cond:expr =>! $error_message:literal
) => {
$(#[$attr])*
#[must_use]
$vis const fn $req_name () -> impl $crate::require::Condition<Error = $crate::require::ErrMsg> {
let res = $crate::require::Requires::new($cond, || { $crate::require::ErrMsg::new($error_message) });
res
}
};
(
$(#[$attr:meta])*
$vis:vis $req_name:ident => $cond:expr =>! $create_err:expr => $err_ty:ty
) => {
$(#[$attr])*
#[must_use]
$vis const fn $req_name () -> impl $crate::require::Condition<Error = $err_ty> {
let res = $crate::require::Requires::new($cond, || { $create_err });
res
}
};
(
$(#[$attr:meta])*
$vis:vis $req_name:ident => $cond:expr =>! $err:ident ($($args:expr),* $(,)?)
) => {
$(#[$attr])*
#[must_use]
$vis const fn $req_name () -> impl $crate::require::Condition<Error = $err> {
let res = $crate::require::Requires::new($cond, || { $err ($($args),*) });
res
}
};
(
$(#[$attr:meta])*
$vis:vis $req_name:ident => $cond:expr =>! $err:ident :: $func:ident ($($args:expr),* $(,)?)
) => {
$(#[$attr])*
#[must_use]
$vis const fn $req_name () -> impl $crate::require::Condition<Error = $err> {
let res = $crate::require::Requires::new($cond, || { $err :: $func ($($args),*) });
res
}
};
(
$(#[$attr:meta])*
$vis:vis $req_name:ident => $cond:expr =>! $err:ident :: $variant:ident
) => {
$(#[$attr])*
#[must_use]
$vis const fn $req_name () -> impl $crate::require::Condition<Error = $err> {
let res = $crate::require::Requires::new($cond, || { $err :: $variant });
res
}
};
}
#[macro_export]
macro_rules! requirements {
([$($requirement:ident),* $(,)?] $(,)?) => {
$crate::requirements!($crate::require::ErrMsg, [$($requirement),*])
};
($error:ty, [$($requirement:ident),* $(,)?] $(,)?) => {{
#[allow(non_camel_case_types)]
struct Requirements<$($requirement: $crate::require::Condition),*> {
__valid: bool,
$($requirement: $requirement),*
}
#[allow(non_camel_case_types)]
impl<$($requirement),*> $crate::require::Requirement for Requirements<$($requirement),*>
where
$($requirement: $crate::require::Condition,
<$requirement as $crate::require::Condition>::Error: Into<$error>),*
{
type Error = $error;
#[inline]
fn check(&mut self, vector: $crate::arch::Vector) {
#[allow(unused_imports)]
use $crate::require::Condition as _;
self.__valid &= ($(self.$requirement.check(vector) )|*).all_bits_set();
}
#[inline]
fn check_partial(&mut self, vector: $crate::arch::Vector, len: u32) {
#[allow(unused_imports)]
use $crate::require::Condition as _;
self.__valid &= ($(self.$requirement.check(vector) )|*)
.trailing_ones() >= len;
}
#[inline]
fn result(self) -> (bool, Result<(), Self::Error>) {
$(
if let Err(err) = self.$requirement.ok() {
return (self.__valid, Err(err.into()));
};
)*
(self.__valid, Ok(()))
}
#[inline]
fn results(self) -> (bool, impl Iterator<Item=Result<(), Self::Error>>) {
(
self.__valid,
[$(self.$requirement.ok().map_err(|err| -> Self::Error {err.into()})),*]
.into_iter()
)
}
}
Requirements {
__valid: true,
$($requirement: $requirement ()),*
}
}};
}
#[inline]
pub fn check<R: Requirement>(data: &[u8], mut req: R) -> R {
if data.len() >= arch::WIDTH {
unsafe { arch::scan::ensure_requirements(data, req) }
} else {
let len = data.len();
req.check_partial(unsafe { arch::load_partial(data, len) }, len as u32);
req
}
}
#[cfg(test)]
#[test]
fn test() {
use crate::{eq, range};
struct SpecialError(&'static str);
impl From<SpecialError> for ErrMsg {
fn from(value: SpecialError) -> Self {
ErrMsg::new(value.0)
}
}
impl SpecialError {
fn new(msg: &'static str) -> Self {
Self ( msg )
}
}
requirement!(
#[inline] pub uppercase => range!(b'A'..=b'Z') =>! "needs uppercase!"
);
requirement!(
#[inline] pub lowercase => range!(b'a'..=b'z') =>! "needs lowercase!"
);
requirement!(
#[inline] pub numeric => range!(b'0'..=b'9') =>! SpecialError("needs number!")
);
requirement!(
#[inline] pub question_mark => eq(b'?') =>! SpecialError::new("needs a question mark!")
);
let res = check(
b"hello world 12345678910",
requirements!([uppercase, lowercase, numeric, question_mark])
);
for res in res.results().1 {
if let Err(err) = res {
println!("{err:?}");
}
}
let res = check(
b"hello world 12345678910",
requirements!([uppercase, lowercase, numeric, question_mark])
);
println!("{:?}", res.result().1.unwrap_err());
}