use num_enum::{IntoPrimitive, TryFromPrimitive};
use serde::{Deserialize, Serialize};
use strum::VariantArray;
#[derive(
Clone,
Copy,
Debug,
PartialEq,
Eq,
Hash,
IntoPrimitive,
TryFromPrimitive,
VariantArray,
Serialize,
Deserialize,
)]
#[repr(u8)]
pub enum RegisterClass {
GeneralPurpose = 0,
Segment = 1,
Xmm = 2,
Ymm = 3,
Zmm = 4,
Mask = 5,
St = 6,
Mmx = 7,
Control = 8,
Debug = 9,
Test = 10,
Ip = 11,
Bnd = 12,
}
impl RegisterClass {
#[must_use]
pub const fn name_prefix(self) -> Option<&'static str> {
Some(match self {
Self::Xmm => "xmm",
Self::Ymm => "ymm",
Self::Zmm => "zmm",
Self::Mmx => "mm",
Self::Mask => "k",
Self::Bnd => "bnd",
Self::St => "st",
Self::Control => "cr",
Self::Debug => "dr",
Self::Test => "tr",
Self::GeneralPurpose | Self::Segment | Self::Ip => return None,
})
}
#[must_use]
pub fn from_name(name: &str) -> Option<Self> {
Self::VARIANTS.iter().copied().find(|class| {
class.name_prefix().is_some_and(|prefix| {
name.strip_prefix(prefix)
.is_some_and(|rest| rest.starts_with(|c: char| c.is_ascii_digit()))
})
})
}
}
impl std::fmt::Display for RegisterClass {
#[inline]
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(match self {
Self::GeneralPurpose => "general-purpose",
Self::Segment => "segment",
Self::Xmm => "xmm",
Self::Ymm => "ymm",
Self::Zmm => "zmm",
Self::Mask => "mask",
Self::St => "st",
Self::Mmx => "mmx",
Self::Control => "control",
Self::Debug => "debug",
Self::Test => "test",
Self::Ip => "ip",
Self::Bnd => "bnd",
})
}
}
#[cfg(test)]
mod tests {
use assert2::assert;
use idakit_sys as sys;
use rstest::rstest;
use super::*;
#[test]
fn raw_roundtrips_every_variant() {
for &c in RegisterClass::VARIANTS {
assert!(RegisterClass::try_from(u8::from(c)).ok() == Some(c));
}
}
#[test]
fn try_from_rejects_unknown() {
assert!(RegisterClass::try_from(13).is_err());
assert!(RegisterClass::try_from(14).is_err());
assert!(RegisterClass::try_from(100).is_err());
assert!(RegisterClass::try_from(255).is_err());
}
mod proptests {
use proptest::prelude::*;
use super::*;
proptest! {
#[test]
fn try_from_matches_the_modelled_discriminant_set(byte: u8) {
let modelled = RegisterClass::VARIANTS.iter().any(|&v| u8::from(v) == byte);
prop_assert_eq!(RegisterClass::try_from(byte).is_ok(), modelled);
}
}
}
#[test]
fn name_prefix_present_for_exactly_the_regular_classes() {
for &c in RegisterClass::VARIANTS {
let irregular = matches!(
c,
RegisterClass::GeneralPurpose | RegisterClass::Segment | RegisterClass::Ip
);
assert!(c.name_prefix().is_some() != irregular, "{c:?}");
}
}
#[test]
fn from_name_inverts_name_prefix() {
for &c in RegisterClass::VARIANTS {
if let Some(prefix) = c.name_prefix() {
let name = format!("{prefix}0");
assert!(RegisterClass::from_name(&name) == Some(c), "{name}");
}
}
}
#[test]
fn from_name_rejects_unprefixed_names() {
for n in ["rax", "eax", "al", "es", "cs", "rip", "eip", "r8", "r15"] {
assert!(RegisterClass::from_name(n).is_none(), "{n}");
}
}
#[test]
fn from_name_handles_suffixed_multidigit_and_bare() {
assert!(RegisterClass::from_name("cr8d") == Some(RegisterClass::Control));
assert!(RegisterClass::from_name("zmm31") == Some(RegisterClass::Zmm));
assert!(RegisterClass::from_name("k7") == Some(RegisterClass::Mask));
assert!(RegisterClass::from_name("st").is_none());
}
#[rstest]
#[case::prefix_then_letter("xmmx")]
#[case::prefix_then_letters("crab")]
#[case::bare_prefix_no_index("xmm")]
#[case::empty("")]
fn from_name_rejects_malformed_prefixed_names(#[case] name: &str) {
assert!(RegisterClass::from_name(name).is_none(), "{name}");
}
#[test]
fn reg_class_ids_align_with_the_facade() {
let expected = [
RegisterClass::GeneralPurpose,
RegisterClass::Segment,
RegisterClass::Xmm,
RegisterClass::Ymm,
RegisterClass::Zmm,
RegisterClass::Mask,
RegisterClass::St,
RegisterClass::Mmx,
RegisterClass::Control,
RegisterClass::Debug,
RegisterClass::Test,
RegisterClass::Ip,
RegisterClass::Bnd,
];
assert!(RegisterClass::VARIANTS.len() == expected.len());
let ids = sys::reg_class_ids();
assert!(ids.len() == expected.len());
for (i, cls) in expected.iter().enumerate() {
assert!(
ids[i] == u8::from(*cls),
"reg class {cls:?}: facade {} != discriminant {}",
ids[i],
u8::from(*cls)
);
}
}
#[test]
fn display_renders_every_variant() {
for &c in RegisterClass::VARIANTS {
assert!(!c.to_string().is_empty());
}
assert!(RegisterClass::GeneralPurpose.to_string() == "general-purpose");
assert!(RegisterClass::Xmm.to_string() == "xmm");
}
#[test]
fn serde_roundtrips_every_variant() {
for &c in RegisterClass::VARIANTS {
let json = serde_json::to_string(&c).expect("serialize");
let back: RegisterClass = serde_json::from_str(&json).expect("deserialize");
assert!(back == c);
}
}
#[test]
fn hash_usable_in_set() {
use std::collections::HashSet;
let set: HashSet<RegisterClass> = RegisterClass::VARIANTS.iter().copied().collect();
assert!(set.len() == RegisterClass::VARIANTS.len());
}
#[test]
fn register_serde_roundtrip() {
let reg = Register {
number: 0,
class: RegisterClass::GeneralPurpose,
width: 8,
name: "rax".into(),
};
let json = serde_json::to_string(®).expect("serialize");
let back: Register = serde_json::from_str(&json).expect("deserialize");
assert!(back == reg);
}
}
#[derive(Clone, Debug, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[doc(alias("op_t::reg"))]
pub struct Register {
pub number: u16,
pub class: RegisterClass,
pub width: u8,
pub name: Box<str>,
}