use crate::error::{BatchResult, EinError, SsnError};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum TinKind {
Ssn,
Ein,
Itin,
Atin,
Unknown,
}
static EIN_CAMPUS_VALID: [bool; 100] = {
let mut table = [false; 100];
let valid = [
1, 2, 3, 4, 5, 6, 10, 11, 12, 13, 14, 15, 16, 20, 21, 22, 23, 24, 25, 26, 27, 30, 32, 33,
34, 35, 36, 37, 38, 39, 40, 41, 42, 43, 44, 45, 46, 47, 48, 50, 51, 52, 53, 54, 55, 56, 57,
58, 59, 60, 61, 62, 63, 64, 65, 66, 67, 68, 71, 72, 73, 74, 75, 76, 77, 80, 81, 82, 83, 84,
85, 86, 87, 88, 90, 91, 92, 93, 94, 95, 98, 99,
];
let mut i = 0;
while i < valid.len() {
table[valid[i] as usize] = true;
i += 1;
}
table
};
#[derive(Debug, Clone, Copy)]
pub struct TinValidator {
pub allow_formatted: bool,
pub strict_ssn: bool,
}
impl Default for TinValidator {
fn default() -> Self {
Self {
allow_formatted: true,
strict_ssn: true,
}
}
}
impl TinValidator {
pub fn new() -> Self {
Self::default()
}
pub fn strict() -> Self {
Self {
allow_formatted: false,
strict_ssn: true,
}
}
}
#[inline]
fn extract_digits(s: &str) -> Option<[u8; 9]> {
let bytes = s.as_bytes();
let mut digits = [0u8; 9];
let mut count = 0;
if bytes.len() == 9 {
for (i, &b) in bytes.iter().enumerate() {
if b.wrapping_sub(b'0') > 9 {
break;
}
digits[i] = b - b'0';
count += 1;
}
if count == 9 {
return Some(digits);
}
count = 0;
}
for &b in bytes {
let d = b.wrapping_sub(b'0');
if d <= 9 {
if count >= 9 {
return None; }
digits[count] = d;
count += 1;
} else if b != b'-' && b != b' ' {
return None; }
}
if count == 9 {
Some(digits)
} else {
None
}
}
#[inline]
fn digits_to_packed(digits: &[u8; 9]) -> u32 {
let area = (digits[0] as u32) * 100 + (digits[1] as u32) * 10 + (digits[2] as u32);
let group = (digits[3] as u32) * 10 + (digits[4] as u32);
let serial = (digits[5] as u32) * 1000
+ (digits[6] as u32) * 100
+ (digits[7] as u32) * 10
+ (digits[8] as u32);
(area << 21) | (group << 14) | serial
}
#[inline]
fn packed_area(packed: u32) -> u32 {
packed >> 21
}
#[inline]
fn packed_group(packed: u32) -> u32 {
(packed >> 14) & 0x7F
}
#[inline]
fn packed_serial(packed: u32) -> u32 {
packed & 0x3FFF
}
#[inline]
pub fn validate_any(tin: &str) -> bool {
let Some(digits) = extract_digits(tin) else {
return false;
};
if digits.iter().all(|&d| d == 0) {
return false;
}
validate_ssn_digits(&digits).is_ok()
|| validate_ein_digits(&digits).is_ok()
|| validate_itin_digits(&digits)
|| validate_atin_digits(&digits)
}
#[inline]
pub fn validate_ssn(ssn: &str) -> bool {
extract_digits(ssn)
.map(|d| validate_ssn_digits(&d).is_ok())
.unwrap_or(false)
}
#[inline]
pub fn validate_ssn_digits(digits: &[u8; 9]) -> Result<(), SsnError> {
let packed = digits_to_packed(digits);
let area = packed_area(packed);
let group = packed_group(packed);
let serial = packed_serial(packed);
if area == 0 {
return Err(SsnError::AreaZero);
}
if area == 666 {
return Err(SsnError::Area666);
}
if area >= 900 {
return Err(SsnError::AreaReserved);
}
if group == 0 {
return Err(SsnError::GroupZero);
}
if serial == 0 {
return Err(SsnError::SerialZero);
}
let first = digits[0];
if digits.iter().all(|&d| d == first) {
return Err(SsnError::AllSameDigits);
}
Ok(())
}
#[inline]
pub fn validate_ein(ein: &str) -> bool {
extract_digits(ein)
.map(|d| validate_ein_digits(&d).is_ok())
.unwrap_or(false)
}
#[inline]
pub fn validate_ein_digits(digits: &[u8; 9]) -> Result<(), EinError> {
if digits.iter().all(|&d| d == 0) {
return Err(EinError::AllZeros);
}
let campus = (digits[0] as usize) * 10 + (digits[1] as usize);
if !EIN_CAMPUS_VALID[campus] {
return Err(EinError::InvalidCampus);
}
Ok(())
}
#[inline]
pub fn validate_itin(itin: &str) -> bool {
extract_digits(itin)
.map(|d| validate_itin_digits(&d))
.unwrap_or(false)
}
#[inline]
pub fn validate_itin_digits(digits: &[u8; 9]) -> bool {
digits[0] == 9 && (digits[3] == 7 || digits[3] == 8)
}
#[inline]
pub fn validate_atin(atin: &str) -> bool {
extract_digits(atin)
.map(|d| validate_atin_digits(&d))
.unwrap_or(false)
}
#[inline]
pub fn validate_atin_digits(digits: &[u8; 9]) -> bool {
digits[0] == 9 && digits[3] == 9 && digits[4] == 3
}
#[inline]
pub fn detect_tin_kind(digits: &[u8; 9]) -> TinKind {
if digits[0] == 9 {
if digits[3] == 9 && digits[4] == 3 {
return TinKind::Atin;
}
if digits[3] == 7 || digits[3] == 8 {
return TinKind::Itin;
}
}
if validate_ssn_digits(digits).is_ok() {
return TinKind::Ssn;
}
if validate_ein_digits(digits).is_ok() {
return TinKind::Ein;
}
TinKind::Unknown
}
pub fn validate_batch(tins: &[&str]) -> BatchResult {
let mut result = BatchResult::with_capacity(tins.len());
const CHUNK_SIZE: usize = 64;
for (chunk_idx, chunk) in tins.chunks(CHUNK_SIZE).enumerate() {
let base_idx = chunk_idx * CHUNK_SIZE;
let mut digits_buf: [[u8; 9]; CHUNK_SIZE] = [[0; 9]; CHUNK_SIZE];
let mut valid_format: [bool; CHUNK_SIZE] = [false; CHUNK_SIZE];
for (i, tin) in chunk.iter().enumerate() {
if let Some(d) = extract_digits(tin) {
digits_buf[i] = d;
valid_format[i] = true;
}
}
for (i, digits) in digits_buf.iter().enumerate().take(chunk.len()) {
if valid_format[i] {
let is_valid = !digits.iter().all(|&d| d == 0)
&& (validate_ssn_digits(digits).is_ok()
|| validate_ein_digits(digits).is_ok()
|| validate_itin_digits(digits)
|| validate_atin_digits(digits));
result.set(base_idx + i, is_valid);
}
}
}
result
}
pub fn validate_ssn_batch(ssns: &[&str]) -> BatchResult {
let mut result = BatchResult::with_capacity(ssns.len());
for (i, ssn) in ssns.iter().enumerate() {
if let Some(digits) = extract_digits(ssn) {
result.set(i, validate_ssn_digits(&digits).is_ok());
}
}
result
}
pub fn validate_ein_batch(eins: &[&str]) -> BatchResult {
let mut result = BatchResult::with_capacity(eins.len());
for (i, ein) in eins.iter().enumerate() {
if let Some(digits) = extract_digits(ein) {
result.set(i, validate_ein_digits(&digits).is_ok());
}
}
result
}
#[cfg(all(target_arch = "x86_64", feature = "simd"))]
mod simd_x86 {
use super::*;
#[inline]
pub fn has_avx2() -> bool {
is_x86_feature_detected!("avx2")
}
#[cfg(target_feature = "avx2")]
#[inline]
pub unsafe fn validate_ssn_avx2(packed: &[u32; 8]) -> u8 {
use std::arch::x86_64::*;
let tins = _mm256_loadu_si256(packed.as_ptr() as *const __m256i);
let areas = _mm256_srli_epi32(tins, 21);
let zero = _mm256_setzero_si256();
let area_not_zero = _mm256_cmpgt_epi32(areas, zero);
let v666 = _mm256_set1_epi32(666);
let area_not_666 = _mm256_xor_si256(_mm256_cmpeq_epi32(areas, v666), _mm256_set1_epi32(-1));
let v900 = _mm256_set1_epi32(900);
let area_lt_900 = _mm256_cmpgt_epi32(v900, areas);
let groups = _mm256_and_si256(_mm256_srli_epi32(tins, 14), _mm256_set1_epi32(0x7F));
let group_not_zero = _mm256_cmpgt_epi32(groups, zero);
let serials = _mm256_and_si256(tins, _mm256_set1_epi32(0x3FFF));
let serial_not_zero = _mm256_cmpgt_epi32(serials, zero);
let valid = _mm256_and_si256(
_mm256_and_si256(_mm256_and_si256(area_not_zero, area_not_666), area_lt_900),
_mm256_and_si256(group_not_zero, serial_not_zero),
);
_mm256_movemask_ps(_mm256_castsi256_ps(valid)) as u8
}
}
#[cfg(all(target_arch = "aarch64", feature = "simd"))]
#[allow(dead_code, unused_imports)]
mod simd_arm {
use super::*;
#[cfg(target_feature = "neon")]
#[inline]
pub unsafe fn validate_ssn_neon(packed: &[u32; 4]) -> u8 {
use std::arch::aarch64::*;
let tins = vld1q_u32(packed.as_ptr());
let areas = vshrq_n_u32(tins, 21);
let zero = vdupq_n_u32(0);
let area_not_zero = vcgtq_u32(areas, zero);
let v666 = vdupq_n_u32(666);
let area_not_666 = vmvnq_u32(vceqq_u32(areas, v666));
let v900 = vdupq_n_u32(900);
let area_lt_900 = vcltq_u32(areas, v900);
let groups = vandq_u32(vshrq_n_u32(tins, 14), vdupq_n_u32(0x7F));
let group_not_zero = vcgtq_u32(groups, zero);
let serials = vandq_u32(tins, vdupq_n_u32(0x3FFF));
let serial_not_zero = vcgtq_u32(serials, zero);
let valid = vandq_u32(
vandq_u32(vandq_u32(area_not_zero, area_not_666), area_lt_900),
vandq_u32(group_not_zero, serial_not_zero),
);
let narrowed = vmovn_u32(valid);
let bytes = vreinterpret_u8_u16(narrowed);
let mut mask = 0u8;
let arr: [u8; 8] = std::mem::transmute(bytes);
if arr[0] != 0 {
mask |= 1;
}
if arr[2] != 0 {
mask |= 2;
}
if arr[4] != 0 {
mask |= 4;
}
if arr[6] != 0 {
mask |= 8;
}
mask
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_extract_digits() {
assert_eq!(
extract_digits("123456789"),
Some([1, 2, 3, 4, 5, 6, 7, 8, 9])
);
assert_eq!(
extract_digits("123-45-6789"),
Some([1, 2, 3, 4, 5, 6, 7, 8, 9])
);
assert_eq!(
extract_digits("12-3456789"),
Some([1, 2, 3, 4, 5, 6, 7, 8, 9])
);
assert_eq!(extract_digits("12345678"), None); assert_eq!(extract_digits("1234567890"), None); assert_eq!(extract_digits("12345678a"), None); }
#[test]
fn test_ssn_validation() {
assert!(validate_ssn("123-45-6789"));
assert!(validate_ssn("123456789"));
assert!(validate_ssn("078-05-1120"));
assert!(!validate_ssn("000-12-3456")); assert!(!validate_ssn("666-12-3456")); assert!(!validate_ssn("900-12-3456")); assert!(!validate_ssn("123-00-4567")); assert!(!validate_ssn("123-45-0000")); assert!(!validate_ssn("111-11-1111")); assert!(!validate_ssn("000-00-0000")); }
#[test]
fn test_ein_validation() {
assert!(validate_ein("12-3456789"));
assert!(validate_ein("123456789"));
assert!(!validate_ein("00-0000000")); assert!(!validate_ein("07-1234567")); assert!(!validate_ein("08-1234567")); assert!(!validate_ein("09-1234567")); }
#[test]
fn test_itin_validation() {
assert!(validate_itin("912-78-1234"));
assert!(validate_itin("900-70-1234"));
assert!(!validate_itin("123-45-6789")); assert!(!validate_itin("900-60-1234")); }
#[test]
fn test_atin_validation() {
assert!(validate_atin("900-93-1234"));
assert!(!validate_atin("123-93-1234")); assert!(!validate_atin("900-94-1234")); }
#[test]
fn test_detect_tin_kind() {
assert_eq!(detect_tin_kind(&[1, 2, 3, 4, 5, 6, 7, 8, 9]), TinKind::Ssn);
assert_eq!(detect_tin_kind(&[1, 2, 3, 4, 5, 6, 7, 8, 9]), TinKind::Ssn);
assert_eq!(detect_tin_kind(&[9, 0, 0, 7, 0, 1, 2, 3, 4]), TinKind::Itin);
assert_eq!(detect_tin_kind(&[9, 0, 0, 9, 3, 1, 2, 3, 4]), TinKind::Atin);
}
#[test]
fn test_batch_validation() {
let tins = vec![
"123-45-6789", "000-00-0000", "12-3456789", "invalid", "912-78-1234", ];
let result = validate_batch(&tins);
assert!(result.is_valid(0));
assert!(!result.is_valid(1));
assert!(result.is_valid(2));
assert!(!result.is_valid(3));
assert!(result.is_valid(4));
assert_eq!(result.valid_count(), 3);
assert_eq!(result.invalid_count(), 2);
}
#[test]
fn test_packed_representation() {
let digits = [1, 2, 3, 4, 5, 6, 7, 8, 9];
let packed = digits_to_packed(&digits);
assert_eq!(packed_area(packed), 123);
assert_eq!(packed_group(packed), 45);
assert_eq!(packed_serial(packed), 6789);
}
}