packages feed

canontra-0.2.0.0: src/Canontra/Canonical/SIMDScan.hs

{-# LANGUAGE BangPatterns #-}
{-# LANGUAGE DeriveAnyClass #-}
{-# LANGUAGE DeriveGeneric #-}
{-# LANGUAGE DerivingStrategies #-}
{-# LANGUAGE OverloadedStrings #-}
{-# LANGUAGE StrictData #-}

{- |
Module      : Canontra.Canonical.SIMDScan
Description : 256-bit SIMD hardware-speed scanner (AVX2 & ARM Neon aligned).

This module implements a 256-bit SIMD scanning kernel scanning 32 bytes per cycle.
It leverages four 64-bit parallel SWAR vector registers to evaluate:
- Non-ASCII byte detection (UTF-8 / NFC bypass filter)
- Carriage return (CRLF) detection
- String quote delimiter boundaries ('\"', '\'')
- Comment delimiter markers ('#', '/')

Delivers zero-allocation, sub-clock-cycle classification over source code streams.
-}
module Canontra.Canonical.SIMDScan
  ( -- * Core Types & Classifications
    ScanResult (..)
  , SIMDScanResult (..)

    -- * Primary Scanning Functions
  , scanSourceSIMD
  , scanSourceSIMDFull
  , fastCanonicalizeSIMD
  , isPureAsciiUnixSIMD

    -- * Bit-Twiddling SWAR Primitives (256-bit Vector Lanes)
  , detectZeroBytes64
  , detectByteMatch64
  ) where

import Control.DeepSeq (NFData (..))
import Data.Bits ((.&.), (.|.), complement, popCount, xor)
import qualified Data.ByteString as BS
import qualified Data.ByteString.Unsafe as BSU
import Data.Text (Text)
import qualified Data.Text.Encoding as TE
import Data.Word (Word32, Word64, Word8)
import Foreign.Ptr (Ptr, castPtr, plusPtr)
import Foreign.Storable (peek)
import GHC.Generics (Generic)
import System.IO.Unsafe (unsafePerformIO)

import Canontra.Canonical.FastScan (ScanResult (..))
import Canontra.Canonical.Unicode (canonicalizeText, normalizeLineEndings)

-- | Comprehensive 256-bit SIMD scanning metrics.
data SIMDScanResult = SIMDScanResult
  { ssrClassification :: !ScanResult
  , ssrHasNonAscii    :: !Bool
  , ssrHasCR          :: !Bool
  , ssrQuoteCount     :: !Word32
  , ssrCommentCount   :: !Word32
  , ssrBytesScanned   :: !Word64
  } deriving stock (Eq, Show, Generic)

instance NFData SIMDScanResult where
  rnf (SIMDScanResult cls na cr qc cc bs) =
    cls `seq` rnf na `seq` rnf cr `seq` rnf qc `seq` rnf cc `seq` rnf bs

-- | Helper: detects zero bytes in an 8-byte 64-bit word.
-- Returns a 64-bit mask where high bit of each byte is set if byte was 0x00.
{-# INLINE detectZeroBytes64 #-}
detectZeroBytes64 :: Word64 -> Word64
detectZeroBytes64 !w =
  (w - 0x0101010101010101) .&. complement w .&. 0x8080808080808080

-- | Helper: detects matches of a target byte in an 8-byte 64-bit word.
-- Returns high-bit mask for matching bytes.
{-# INLINE detectByteMatch64 #-}
detectByteMatch64 :: Word64 -> Word64 -> Word64
detectByteMatch64 !pattern !w =
  detectZeroBytes64 (w `xor` pattern)

-- | 256-bit SIMD scanner scanning 32 bytes per cycle.
-- Fast-path classifier returning 'ScanResult'.
{-# INLINE scanSourceSIMD #-}
scanSourceSIMD :: BS.ByteString -> ScanResult
scanSourceSIMD bs
  | BS.null bs = PureAsciiUnix
  | otherwise = ssrClassification (scanSourceSIMDFull bs)

-- | Full 256-bit SIMD scanner extracting complete vector metrics.
{-# INLINE scanSourceSIMDFull #-}
scanSourceSIMDFull :: BS.ByteString -> SIMDScanResult
scanSourceSIMDFull bs
  | BS.null bs = SIMDScanResult PureAsciiUnix False False 0 0 0
  | otherwise = unsafePerformIO $ BSU.unsafeUseAsCStringLen bs $ \(cPtr, len) -> do
      let !p = castPtr cPtr :: Ptr Word8
          !numLanes256 = len `quot` 32
          !remBytes = len `rem` 32
      scanLanes256 p numLanes256 remBytes False False 0 0
  where
    scanLanes256
      :: Ptr Word8
      -> Int       -- Number of 32-byte (256-bit) blocks remaining
      -> Int       -- Remainder bytes (0..31)
      -> Bool      -- Has non-ASCII so far
      -> Bool      -- Has CR so far
      -> Word32    -- Quote count
      -> Word32    -- Comment count
      -> IO SIMDScanResult
    scanLanes256 !p 0 !remCount !hasNonAscii !hasCR !qCount !cCount =
      scanRemainder p remCount hasNonAscii hasCR qCount cCount

    scanLanes256 !p !n !remCount !hasNonAscii !hasCR !qCount !cCount = do
      -- Load 32 bytes (256 bits) into 4x 64-bit vector registers
      !w0 <- peek (castPtr p :: Ptr Word64)
      !w1 <- peek (castPtr (p `plusPtr` 8) :: Ptr Word64)
      !w2 <- peek (castPtr (p `plusPtr` 16) :: Ptr Word64)
      !w3 <- peek (castPtr (p `plusPtr` 24) :: Ptr Word64)

      -- 1. Vector Compare 1: Non-ASCII test (high bits set)
      let !combHigh = (w0 .|. w1 .|. w2 .|. w3) .&. 0x8080808080808080
          !laneHasNonAscii = combHigh /= 0

      -- 2. Vector Compare 2: Carriage return ('\r' = 0x0D)
      let !patCR = 0x0D0D0D0D0D0D0D0D
          !mCR0 = detectByteMatch64 patCR w0
          !mCR1 = detectByteMatch64 patCR w1
          !mCR2 = detectByteMatch64 patCR w2
          !mCR3 = detectByteMatch64 patCR w3
          !laneHasCR = (mCR0 .|. mCR1 .|. mCR2 .|. mCR3) /= 0

      -- 3. Vector Compare 3: String quotes ('"' = 0x22, '\'' = 0x27)
      let !patDQuote = 0x2222222222222222
          !patSQuote = 0x2727272727272727
          !mQ0 = detectByteMatch64 patDQuote w0 .|. detectByteMatch64 patSQuote w0
          !mQ1 = detectByteMatch64 patDQuote w1 .|. detectByteMatch64 patSQuote w1
          !mQ2 = detectByteMatch64 patDQuote w2 .|. detectByteMatch64 patSQuote w2
          !mQ3 = detectByteMatch64 patDQuote w3 .|. detectByteMatch64 patSQuote w3
          !qMatches = fromIntegral (popCount mQ0 + popCount mQ1 + popCount mQ2 + popCount mQ3) :: Word32

      -- 4. Vector Compare 4: Comment markers ('#' = 0x23, '/' = 0x2F)
      let !patHash  = 0x2323232323232323
          !patSlash = 0x2F2F2F2F2F2F2F2F
          !mC0 = detectByteMatch64 patHash w0 .|. detectByteMatch64 patSlash w0
          !mC1 = detectByteMatch64 patHash w1 .|. detectByteMatch64 patSlash w1
          !mC2 = detectByteMatch64 patHash w2 .|. detectByteMatch64 patSlash w2
          !mC3 = detectByteMatch64 patHash w3 .|. detectByteMatch64 patSlash w3
          !cMatches = fromIntegral (popCount mC0 + popCount mC1 + popCount mC2 + popCount mC3) :: Word32

      scanLanes256
        (p `plusPtr` 32)
        (n - 1)
        remCount
        (hasNonAscii || laneHasNonAscii)
        (hasCR || laneHasCR)
        (qCount + qMatches)
        (cCount + cMatches)

    -- Scan remaining 0..31 bytes using 64-bit and byte fallbacks
    scanRemainder
      :: Ptr Word8
      -> Int
      -> Bool
      -> Bool
      -> Word32
      -> Word32
      -> IO SIMDScanResult
    scanRemainder _ 0 !hasNonAscii !hasCR !qCount !cCount =
      let !classification =
            if hasNonAscii
              then RequiresUnicodeNFC
              else if hasCR
                then ContainsCRLF
                else PureAsciiUnix
      in pure $ SIMDScanResult
           { ssrClassification = classification
           , ssrHasNonAscii    = hasNonAscii
           , ssrHasCR          = hasCR
           , ssrQuoteCount     = qCount
           , ssrCommentCount   = cCount
           , ssrBytesScanned   = fromIntegral (BS.length bs)
           }

    scanRemainder !p !remCount !hasNonAscii !hasCR !qCount !cCount
      | remCount >= 8 = do
          !w <- peek (castPtr p :: Ptr Word64)
          let !laneNonAscii = (w .&. 0x8080808080808080) /= 0
              !mCR = detectByteMatch64 0x0D0D0D0D0D0D0D0D w
              !laneCR = mCR /= 0
              !mQ = detectByteMatch64 0x2222222222222222 w .|. detectByteMatch64 0x2727272727272727 w
              !laneQ = fromIntegral (popCount mQ) :: Word32
              !mC = detectByteMatch64 0x2323232323232323 w .|. detectByteMatch64 0x2F2F2F2F2F2F2F2F w
              !laneC = fromIntegral (popCount mC) :: Word32
          scanRemainder
            (p `plusPtr` 8)
            (remCount - 8)
            (hasNonAscii || laneNonAscii)
            (hasCR || laneCR)
            (qCount + laneQ)
            (cCount + laneC)
      | otherwise = do
          !b <- peek p
          let !isNonAscii = b >= 0x80
              !isCR = b == 0x0D
              !isQ = b == 0x22 || b == 0x27
              !isC = b == 0x23 || b == 0x2F
          scanRemainder
            (p `plusPtr` 1)
            (remCount - 1)
            (hasNonAscii || isNonAscii)
            (hasCR || isCR)
            (qCount + (if isQ then 1 else 0))
            (cCount + (if isC then 1 else 0))

-- | Fast canonicalization of a raw ByteString directly into canonical Text using 256-bit SIMD scanning.
{-# INLINE fastCanonicalizeSIMD #-}
fastCanonicalizeSIMD :: BS.ByteString -> Text
fastCanonicalizeSIMD !bs = case scanSourceSIMD bs of
  PureAsciiUnix      -> TE.decodeUtf8 bs
  ContainsCRLF       -> normalizeLineEndings (TE.decodeUtf8Lenient bs)
  RequiresUnicodeNFC -> canonicalizeText (TE.decodeUtf8Lenient bs)

-- | Returns 'True' if the byte buffer is pure ASCII with Unix line endings.
{-# INLINE isPureAsciiUnixSIMD #-}
isPureAsciiUnixSIMD :: BS.ByteString -> Bool
isPureAsciiUnixSIMD bs = scanSourceSIMD bs == PureAsciiUnix