{-# LANGUAGE BangPatterns #-}
{-# LANGUAGE CPP #-}

module DataFrame.IO.Parquet.Encoding (
    -- Kept from the original Encoding module (used by Levels)
    ceilLog2,
    bitWidthForMaxLevel,
    -- Vector-based RLE/bit-packed decoder (from new parser)
    decodeRLEBitPackedHybrid,
    extractBitsInto,
    fillRun,
    decodeDictIndices,
) where

import Control.Monad.ST (ST, runST)
import Data.Bits
import qualified Data.ByteString as BS
import qualified Data.ByteString.Unsafe as BSU
import qualified Data.Vector.Unboxed as VU
import qualified Data.Vector.Unboxed.Mutable as VUM
import Data.Word
import DataFrame.IO.Parquet.Binary (readUVarInt)
import DataFrame.Internal.Binary (littleEndianWord32)

-- ---------------------------------------------------------------------------
-- Level-width helpers (used by Levels.hs)
-- ---------------------------------------------------------------------------

ceilLog2 :: Int -> Int
ceilLog2 :: Int -> Int
ceilLog2 Int
x
    | Int
x Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
<= Int
1 = Int
0
    | Bool
otherwise = Int
1 Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int -> Int
ceilLog2 ((Int
x Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1) Int -> Int -> Int
forall a. Integral a => a -> a -> a
`div` Int
2)

bitWidthForMaxLevel :: Int -> Int
bitWidthForMaxLevel :: Int -> Int
bitWidthForMaxLevel Int
maxLevel = Int -> Int
ceilLog2 (Int
maxLevel Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)

-- ---------------------------------------------------------------------------
-- Vector-based RLE / bit-packed hybrid decoder
-- ---------------------------------------------------------------------------

decodeRLEBitPackedHybrid ::
    -- | Bit width per value (0 = all zeros, use 'VU.replicate')
    Int ->
    -- | Exact number of values to decode
    Int ->
    BS.ByteString ->
    (VU.Vector Word32, BS.ByteString)
decodeRLEBitPackedHybrid :: Int -> Int -> ByteString -> (Vector Word32, ByteString)
decodeRLEBitPackedHybrid Int
bw Int
need ByteString
bs
    | Int
bw Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
== Int
0 = (Int -> Word32 -> Vector Word32
forall a. Unbox a => Int -> a -> Vector a
VU.replicate Int
need Word32
0, ByteString
bs)
    | Bool
otherwise = (forall s. ST s (Vector Word32, ByteString))
-> (Vector Word32, ByteString)
forall a. (forall s. ST s a) -> a
runST ((forall s. ST s (Vector Word32, ByteString))
 -> (Vector Word32, ByteString))
-> (forall s. ST s (Vector Word32, ByteString))
-> (Vector Word32, ByteString)
forall a b. (a -> b) -> a -> b
$ do
        STVector s Word32
mv <- Int -> ST s (MVector (PrimState (ST s)) Word32)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> m (MVector (PrimState m) a)
VUM.new Int
need
        ByteString
rest <- STVector s Word32 -> Int -> ByteString -> ST s ByteString
forall s. STVector s Word32 -> Int -> ByteString -> ST s ByteString
go STVector s Word32
mv Int
0 ByteString
bs
        Vector Word32
dat <- MVector (PrimState (ST s)) Word32 -> ST s (Vector Word32)
forall a (m :: * -> *).
(Unbox a, PrimMonad m) =>
MVector (PrimState m) a -> m (Vector a)
VU.unsafeFreeze STVector s Word32
MVector (PrimState (ST s)) Word32
mv
        (Vector Word32, ByteString) -> ST s (Vector Word32, ByteString)
forall a. a -> ST s a
forall (m :: * -> *) a. Monad m => a -> m a
return (Vector Word32
dat, ByteString
rest)
  where
    !mask :: Word32
mask = if Int
bw Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
== Int
32 then Word32
forall a. Bounded a => a
maxBound else (Word32
1 Word32 -> Int -> Word32
forall a. Bits a => a -> Int -> a
`shiftL` Int
bw) Word32 -> Word32 -> Word32
forall a. Num a => a -> a -> a
- Word32
1 :: Word32
    go :: VUM.STVector s Word32 -> Int -> BS.ByteString -> ST s BS.ByteString
    go :: forall s. STVector s Word32 -> Int -> ByteString -> ST s ByteString
go STVector s Word32
mv !Int
filled !ByteString
buf
        | Int
filled Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
need = ByteString -> ST s ByteString
forall a. a -> ST s a
forall (m :: * -> *) a. Monad m => a -> m a
return ByteString
buf
        | ByteString -> Bool
BS.null ByteString
buf = ByteString -> ST s ByteString
forall a. a -> ST s a
forall (m :: * -> *) a. Monad m => a -> m a
return ByteString
buf
        | Bool
otherwise =
            let (Word64
hdr64, ByteString
afterHdr) = ByteString -> (Word64, ByteString)
readUVarInt ByteString
buf
                isPacked :: Bool
isPacked = (Word64
hdr64 Word64 -> Word64 -> Word64
forall a. Bits a => a -> a -> a
.&. Word64
1) Word64 -> Word64 -> Bool
forall a. Eq a => a -> a -> Bool
== Word64
1
             in if Bool
isPacked
                    then do
                        let groups :: Int
groups = Word64 -> Int
forall a b. (Integral a, Num b) => a -> b
fromIntegral (Word64
hdr64 Word64 -> Int -> Word64
forall a. Bits a => a -> Int -> a
`shiftR` Int
1) :: Int
                            totalVals :: Int
totalVals = Int
groups Int -> Int -> Int
forall a. Num a => a -> a -> a
* Int
8
                            takeN :: Int
takeN = Int -> Int -> Int
forall a. Ord a => a -> a -> a
min (Int
need Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
filled) Int
totalVals
                            -- Consume all the bytes for this group even if we
                            -- only need a subset of the values.
                            bytesN :: Int
bytesN = (Int
bw Int -> Int -> Int
forall a. Num a => a -> a -> a
* Int
totalVals Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
7) Int -> Int -> Int
forall a. Integral a => a -> a -> a
`div` Int
8
                            (ByteString
chunk, ByteString
rest) = Int -> ByteString -> (ByteString, ByteString)
BS.splitAt Int
bytesN ByteString
afterHdr
                        Int -> Int -> ByteString -> STVector s Word32 -> Int -> ST s ()
forall s.
Int -> Int -> ByteString -> STVector s Word32 -> Int -> ST s ()
extractBitsInto Int
bw Int
takeN ByteString
chunk STVector s Word32
mv Int
filled
                        STVector s Word32 -> Int -> ByteString -> ST s ByteString
forall s. STVector s Word32 -> Int -> ByteString -> ST s ByteString
go STVector s Word32
mv (Int
filled Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
takeN) ByteString
rest
                    else do
                        let runLen :: Int
runLen = Word64 -> Int
forall a b. (Integral a, Num b) => a -> b
fromIntegral (Word64
hdr64 Word64 -> Int -> Word64
forall a. Bits a => a -> Int -> a
`shiftR` Int
1) :: Int
                            nbytes :: Int
nbytes = (Int
bw Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
7) Int -> Int -> Int
forall a. Integral a => a -> a -> a
`div` Int
8
                            val :: Word32
val = ByteString -> Word32
littleEndianWord32 (Int -> ByteString -> ByteString
BS.take Int
4 ByteString
afterHdr) Word32 -> Word32 -> Word32
forall a. Bits a => a -> a -> a
.&. Word32
mask
                            takeN :: Int
takeN = Int -> Int -> Int
forall a. Ord a => a -> a -> a
min (Int
need Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
filled) Int
runLen
                        -- Fill the run directly — no list, no reverse.
                        STVector s Word32 -> Int -> Int -> Word32 -> ST s ()
forall s. STVector s Word32 -> Int -> Int -> Word32 -> ST s ()
fillRun STVector s Word32
mv Int
filled (Int
filled Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
takeN) Word32
val
                        STVector s Word32 -> Int -> ByteString -> ST s ByteString
forall s. STVector s Word32 -> Int -> ByteString -> ST s ByteString
go STVector s Word32
mv (Int
filled Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
takeN) (Int -> ByteString -> ByteString
BS.drop Int
nbytes ByteString
afterHdr)
{-# INLINE decodeRLEBitPackedHybrid #-}

-- | Fill @mv[start..end-1]@ with @val@ using a bulk @memset@-style write.
fillRun :: VUM.STVector s Word32 -> Int -> Int -> Word32 -> ST s ()
fillRun :: forall s. STVector s Word32 -> Int -> Int -> Word32 -> ST s ()
fillRun STVector s Word32
mv Int
i Int
end = MVector (PrimState (ST s)) Word32 -> Word32 -> ST s ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> a -> m ()
VUM.set (Int -> Int -> STVector s Word32 -> STVector s Word32
forall a s. Unbox a => Int -> Int -> MVector s a -> MVector s a
VUM.unsafeSlice Int
i (Int
end Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
i) STVector s Word32
mv)
{-# INLINE fillRun #-}

{- | Write @count@ bit-width-@bw@ values from @bs@ into @mv@ starting at
@offset@, reading the byte buffer with a single-pass LSB-first accumulator.
No intermediate list or ByteString allocation.
-}
extractBitsInto ::
    -- | Bit width
    Int ->
    -- | Number of values to extract
    Int ->
    BS.ByteString ->
    VUM.STVector s Word32 ->
    -- | Write offset into @mv@
    Int ->
    ST s ()
extractBitsInto :: forall s.
Int -> Int -> ByteString -> STVector s Word32 -> Int -> ST s ()
extractBitsInto Int
bw Int
count ByteString
bs STVector s Word32
mv Int
off = Int -> Word64 -> Int -> Int -> ST s ()
forall {m :: * -> *}.
(PrimState m ~ s, PrimMonad m) =>
Int -> Word64 -> Int -> Int -> m ()
go Int
0 (Word64
0 :: Word64) Int
0 Int
0
  where
    !mask :: Word64
mask = if Int
bw Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
== Int
32 then Word64
forall a. Bounded a => a
maxBound else (Word64
1 Word64 -> Int -> Word64
forall a. Bits a => a -> Int -> a
`unsafeShiftL` Int
bw) Word64 -> Word64 -> Word64
forall a. Num a => a -> a -> a
- Word64
1 :: Word64
    !len :: Int
len = ByteString -> Int
BS.length ByteString
bs
    go :: Int -> Word64 -> Int -> Int -> m ()
go !Int
byteIdx !Word64
acc !Int
accBits !Int
done
        | Int
done Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
count = () -> m ()
forall a. a -> m a
forall (m :: * -> *) a. Monad m => a -> m a
return ()
        | Int
accBits Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
bw = do
            MVector (PrimState m) Word32 -> Int -> Word32 -> m ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite STVector s Word32
MVector (PrimState m) Word32
mv (Int
off Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
done) (Word64 -> Word32
forall a b. (Integral a, Num b) => a -> b
fromIntegral (Word64
acc Word64 -> Word64 -> Word64
forall a. Bits a => a -> a -> a
.&. Word64
mask))
            Int -> Word64 -> Int -> Int -> m ()
go Int
byteIdx (Word64
acc Word64 -> Int -> Word64
forall a. Bits a => a -> Int -> a
`unsafeShiftR` Int
bw) (Int
accBits Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
bw) (Int
done Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
        | Int
byteIdx Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
len = () -> m ()
forall a. a -> m a
forall (m :: * -> *) a. Monad m => a -> m a
return ()
        | Bool
otherwise =
            let b :: Word64
b = Word8 -> Word64
forall a b. (Integral a, Num b) => a -> b
fromIntegral (ByteString -> Int -> Word8
BSU.unsafeIndex ByteString
bs Int
byteIdx) :: Word64
             in Int -> Word64 -> Int -> Int -> m ()
go (Int
byteIdx Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1) (Word64
acc Word64 -> Word64 -> Word64
forall a. Bits a => a -> a -> a
.|. (Word64
b Word64 -> Int -> Word64
forall a. Bits a => a -> Int -> a
`unsafeShiftL` Int
accBits)) (Int
accBits Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
8) Int
done
{-# INLINE extractBitsInto #-}

{- | Decode @need@ dictionary indices from a DATA_PAGE bit-width-prefixed
stream (the first byte encodes the bit-width of all subsequent RLE\/bitpacked
values).

Returns the index vector (as 'Int') and the unconsumed bytes.
-}
decodeDictIndices :: Int -> BS.ByteString -> (VU.Vector Int, BS.ByteString)
decodeDictIndices :: Int -> ByteString -> (Vector Int, ByteString)
decodeDictIndices Int
need ByteString
bs = case ByteString -> Maybe (Word8, ByteString)
BS.uncons ByteString
bs of
    Maybe (Word8, ByteString)
Nothing -> [Char] -> (Vector Int, ByteString)
forall a. HasCallStack => [Char] -> a
error [Char]
"decodeDictIndices: empty stream"
    Just (Word8
w0, ByteString
rest0) ->
        let bw :: Int
bw = Word8 -> Int
forall a b. (Integral a, Num b) => a -> b
fromIntegral Word8
w0 :: Int
            (Vector Word32
raw, ByteString
rest1) = Int -> Int -> ByteString -> (Vector Word32, ByteString)
decodeRLEBitPackedHybrid Int
bw Int
need ByteString
rest0
         in ((Word32 -> Int) -> Vector Word32 -> Vector Int
forall a b. (Unbox a, Unbox b) => (a -> b) -> Vector a -> Vector b
VU.map Word32 -> Int
forall a b. (Integral a, Num b) => a -> b
fromIntegral Vector Word32
raw, ByteString
rest1)
{-# INLINE decodeDictIndices #-}