use core::slice::from_raw_parts;

/// 基于 SIMD 硬件向量化与多级宽字展开的快速键比对
///
/// 算法优化策略:
/// 1. 长度预筛:不等长即不等;相同指针或零长直接返回 true;
/// 2. >= 16 字节长键:
///    - aarch64:采用 NEON `vld1q_u8` + `vceqq_u8` + `vminvq_u8` 向量比对;
///    - x86_64:采用 SSE2 `_mm_loadu_si128` + `_mm_cmpeq_epi8` + `_mm_movemask_epi8` 向量比对;
///    - 尾部非 16 字节对齐部分通过末尾 16 字节重叠比对,彻底消除余数标量循环分支;
/// 3. < 16 字节短键(标量回退):
///    - 8..16 字节:采用首尾双 64 位无对齐整型比对(覆盖整个区间);
///    - 4..8 字节:采用首尾双 32 位无对齐整型比对;
///    - < 4 字节:直接切片相等判断(编译器优化为内联字节比对);
///    - 非 SIMD 平台的 >= 16 字节长键:8 字节宽字逐块推进 + 尾部重叠收尾,完整覆盖中段。
#[inline]
pub fn fast_key_eq(a: &[u8], b: &[u8]) -> bool {
  let len = a.len();
  if len != b.len() {
    return false;
  }
  if a.as_ptr() == b.as_ptr() || len == 0 {
    return true;
  }

  #[cfg(target_arch = "aarch64")]
  {
    if len >= 16 {
      return unsafe { neon_key_eq(a.as_ptr(), b.as_ptr(), len) };
    }
  }

  #[cfg(target_arch = "x86_64")]
  {
    if len >= 16 {
      return unsafe { sse2_key_eq(a.as_ptr(), b.as_ptr(), len) };
    }
  }

  // 标量多级宽字回退(短键或非 SIMD 平台)
  unsafe { scalar_fallback_key_eq(a.as_ptr(), b.as_ptr(), len) }
}

#[cfg(target_arch = "aarch64")]
#[inline]
unsafe fn neon_key_eq(a: *const u8, b: *const u8, len: usize) -> bool {
  use core::arch::aarch64::{vceqq_u8, vld1q_u8, vminvq_u8};
  unsafe {
    let mut offset = 0;
    while offset + 16 <= len {
      let va = vld1q_u8(a.add(offset));
      let vb = vld1q_u8(b.add(offset));
      let vcmp = vceqq_u8(va, vb);
      if vminvq_u8(vcmp) != 0xFF {
        return false;
      }
      offset += 16;
    }
    if offset < len {
      let va = vld1q_u8(a.add(len - 16));
      let vb = vld1q_u8(b.add(len - 16));
      let vcmp = vceqq_u8(va, vb);
      return vminvq_u8(vcmp) == 0xFF;
    }
    true
  }
}

#[cfg(target_arch = "x86_64")]
#[inline]
unsafe fn sse2_key_eq(a: *const u8, b: *const u8, len: usize) -> bool {
  use core::arch::x86_64::{__m128i, _mm_cmpeq_epi8, _mm_loadu_si128, _mm_movemask_epi8};
  unsafe {
    let mut offset = 0;
    while offset + 16 <= len {
      let va = _mm_loadu_si128(a.add(offset) as *const __m128i);
      let vb = _mm_loadu_si128(b.add(offset) as *const __m128i);
      let vcmp = _mm_cmpeq_epi8(va, vb);
      let mask = _mm_movemask_epi8(vcmp);
      if mask != 0xFFFF {
        return false;
      }
      offset += 16;
    }
    if offset < len {
      let va = _mm_loadu_si128(a.add(len - 16) as *const __m128i);
      let vb = _mm_loadu_si128(b.add(len - 16) as *const __m128i);
      let vcmp = _mm_cmpeq_epi8(va, vb);
      return _mm_movemask_epi8(vcmp) == 0xFFFF;
    }
    true
  }
}

/// 标量宽字回退比对
///
/// `len >= 8` 时按 8 字节宽字逐块推进 + 尾部重叠收尾:SIMD 平台上仅 `len < 16` 的短键
/// 回退至此(循环至多 1 轮,头尾双字并集恰好覆盖整个区间);非 SIMD 平台长键也能完整
/// 覆盖中段字节,杜绝漏比导致的假阳性相等。
#[inline]
unsafe fn scalar_fallback_key_eq(a: *const u8, b: *const u8, len: usize) -> bool {
  unsafe {
    if len >= 8 {
      let mut offset = 0;
      while offset + 8 <= len {
        if (a.add(offset) as *const u64).read_unaligned()
          != (b.add(offset) as *const u64).read_unaligned()
        {
          return false;
        }
        offset += 8;
      }
      let tail = len - 8;
      return (a.add(tail) as *const u64).read_unaligned()
        == (b.add(tail) as *const u64).read_unaligned();
    }
    if len >= 4 {
      let a_head = (a as *const u32).read_unaligned();
      let b_head = (b as *const u32).read_unaligned();
      let a_tail = (a.add(len - 4) as *const u32).read_unaligned();
      let b_tail = (b.add(len - 4) as *const u32).read_unaligned();
      return a_head == b_head && a_tail == b_tail;
    }
    from_raw_parts(a, len) == from_raw_parts(b, len)
  }
}

#[cfg(test)]
mod tests {
  use super::fast_key_eq;

  #[test]
  fn test_fast_key_eq() {
    // 1. 空切片边界测试(同指针与不同指针)
    assert!(fast_key_eq(b"", b""));
    let empty_a: &[u8] = &[];
    let empty_b: &[u8] = &[];
    assert!(fast_key_eq(empty_a, empty_b));

    // 2. 基础切片测试
    assert!(fast_key_eq(b"hello", b"hello"));
    assert!(!fast_key_eq(b"hello", b"world"));
    assert!(!fast_key_eq(b"short", b"shorter"));

    // 3. 重叠内存切片比对(同一缓冲区不同偏移)
    let overlap_buf = b"0123456789abcdef0123456789abcdef";
    assert!(fast_key_eq(&overlap_buf[0..16], &overlap_buf[16..32]));
    assert!(!fast_key_eq(&overlap_buf[0..16], &overlap_buf[1..17]));
    assert!(fast_key_eq(&overlap_buf[0..8], &overlap_buf[16..24]));
    assert!(!fast_key_eq(&overlap_buf[0..8], &overlap_buf[1..9]));

    // 4. 遍历 1..=64 字节所有步长,且穷举变异每一个位置 (0..len)
    for len in 1..=64 {
      let v1 = vec![0x5Au8; len];
      let v2 = vec![0x5Au8; len];
      assert!(fast_key_eq(&v1, &v2));

      // 穷举变异每一个位置
      for pos in 0..len {
        let mut v3 = v1.clone();
        v3[pos] ^= 0xFF;
        assert!(
          !fast_key_eq(&v1, &v3),
          "fast_key_eq 应检测出 len={len} 在 pos={pos} 处的差异"
        );
      }
    }

    // 5. 超长切片(128B, 256B, 1024B)极端与跨步长差异比对
    for &long_len in &[128, 256, 1024] {
      let l1 = vec![0x33u8; long_len];
      let l2 = vec![0x33u8; long_len];
      assert!(fast_key_eq(&l1, &l2));

      // 头部变异
      let mut l_head = l1.clone();
      l_head[0] = 0x44;
      assert!(!fast_key_eq(&l1, &l_head));

      // 尾部变异
      let mut l_tail = l1.clone();
      l_tail[long_len - 1] = 0x44;
      assert!(!fast_key_eq(&l1, &l_tail));

      // 中间变异
      let mut l_mid = l1.clone();
      l_mid[long_len / 2] = 0x44;
      assert!(!fast_key_eq(&l1, &l_mid));

      // 16 字节边界对齐处变异
      for boundary in [15, 16, 31, 32, long_len - 17, long_len - 16] {
        let mut l_bound = l1.clone();
        l_bound[boundary] = 0x44;
        assert!(!fast_key_eq(&l1, &l_bound));
      }
    }
  }
}