{-# LANGUAGE BangPatterns #-}

module DataFrame.IO.Parquet.Dictionary (DictVals (..), readDictVals, decodeRLEBitPackedHybrid) where

import Data.Bits
import qualified Data.ByteString as BS
import qualified Data.ByteString.Unsafe as BSU
import Data.Int (Int32, Int64)
import qualified Data.Text as T
import Data.Text.Encoding
import Data.Time (UTCTime)
import qualified Data.Vector as V
import Data.Word
import DataFrame.IO.Parquet.Binary (readUVarInt)
import DataFrame.IO.Parquet.Thrift (ThriftType (..))
import DataFrame.IO.Parquet.Time (int96ToUTCTime)
import DataFrame.Internal.Binary (
    littleEndianInt32,
    littleEndianWord32,
    littleEndianWord64,
 )
import GHC.Float

data DictVals
    = DBool (V.Vector Bool)
    | DInt32 (V.Vector Int32)
    | DInt64 (V.Vector Int64)
    | DInt96 (V.Vector UTCTime)
    | DFloat (V.Vector Float)
    | DDouble (V.Vector Double)
    | DText (V.Vector T.Text)
    deriving (Int -> DictVals -> ShowS
[DictVals] -> ShowS
DictVals -> String
(Int -> DictVals -> ShowS)
-> (DictVals -> String) -> ([DictVals] -> ShowS) -> Show DictVals
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: Int -> DictVals -> ShowS
showsPrec :: Int -> DictVals -> ShowS
$cshow :: DictVals -> String
show :: DictVals -> String
$cshowList :: [DictVals] -> ShowS
showList :: [DictVals] -> ShowS
Show, DictVals -> DictVals -> Bool
(DictVals -> DictVals -> Bool)
-> (DictVals -> DictVals -> Bool) -> Eq DictVals
forall a. (a -> a -> Bool) -> (a -> a -> Bool) -> Eq a
$c== :: DictVals -> DictVals -> Bool
== :: DictVals -> DictVals -> Bool
$c/= :: DictVals -> DictVals -> Bool
/= :: DictVals -> DictVals -> Bool
Eq)

{- | Decode the values from a dictionary page.

The @numVals@ argument is the entry count declared in the dictionary page
header.  It is used to limit BOOLEAN decoding (1-bit-per-value encoding has
no natural delimiter).

The @typeLength@ argument is only meaningful for FIXED_LEN_BYTE_ARRAY: it is
the byte-width of each individual dictionary entry, NOT the total number of
entries.  Passing @numVals@ here (the old behaviour) would cause it to be
misread as an element size, yielding a dictionary that is far too small.
-}
readDictVals :: ThriftType -> BS.ByteString -> Int32 -> Maybe Int32 -> DictVals
readDictVals :: ThriftType -> ByteString -> Int32 -> Maybe Int32 -> DictVals
readDictVals (BOOLEAN Enumeration 0
_) ByteString
bs Int32
count Maybe Int32
_ = Vector Bool -> DictVals
DBool ([Bool] -> Vector Bool
forall a. [a] -> Vector a
V.fromList (Int -> [Bool] -> [Bool]
forall a. Int -> [a] -> [a]
take (Int32 -> Int
forall a b. (Integral a, Num b) => a -> b
fromIntegral Int32
count) ([Bool] -> [Bool]) -> [Bool] -> [Bool]
forall a b. (a -> b) -> a -> b
$ ByteString -> [Bool]
readPageBool ByteString
bs))
readDictVals (INT32 Enumeration 1
_) ByteString
bs Int32
_ Maybe Int32
_ = Vector Int32 -> DictVals
DInt32 ([Int32] -> Vector Int32
forall a. [a] -> Vector a
V.fromList (ByteString -> [Int32]
readPageInt32 ByteString
bs))
readDictVals (INT64 Enumeration 2
_) ByteString
bs Int32
_ Maybe Int32
_ = Vector Int64 -> DictVals
DInt64 ([Int64] -> Vector Int64
forall a. [a] -> Vector a
V.fromList (ByteString -> [Int64]
readPageInt64 ByteString
bs))
readDictVals (INT96 Enumeration 3
_) ByteString
bs Int32
_ Maybe Int32
_ = Vector UTCTime -> DictVals
DInt96 ([UTCTime] -> Vector UTCTime
forall a. [a] -> Vector a
V.fromList (ByteString -> [UTCTime]
readPageInt96Times ByteString
bs))
readDictVals (FLOAT Enumeration 4
_) ByteString
bs Int32
_ Maybe Int32
_ = Vector Float -> DictVals
DFloat ([Float] -> Vector Float
forall a. [a] -> Vector a
V.fromList (ByteString -> [Float]
readPageFloat ByteString
bs))
readDictVals (DOUBLE Enumeration 5
_) ByteString
bs Int32
_ Maybe Int32
_ = Vector Double -> DictVals
DDouble ([Double] -> Vector Double
forall a. [a] -> Vector a
V.fromList (ByteString -> [Double]
readPageWord64 ByteString
bs))
readDictVals (BYTE_ARRAY Enumeration 6
_) ByteString
bs Int32
_ Maybe Int32
_ = Vector Text -> DictVals
DText ([Text] -> Vector Text
forall a. [a] -> Vector a
V.fromList (ByteString -> [Text]
readPageBytes ByteString
bs))
readDictVals (FIXED_LEN_BYTE_ARRAY Enumeration 7
_) ByteString
bs Int32
_ (Just Int32
len) =
    Vector Text -> DictVals
DText ([Text] -> Vector Text
forall a. [a] -> Vector a
V.fromList (ByteString -> Int -> [Text]
readPageFixedBytes ByteString
bs (Int32 -> Int
forall a b. (Integral a, Num b) => a -> b
fromIntegral Int32
len)))
readDictVals ThriftType
t ByteString
_ Int32
_ Maybe Int32
_ = String -> DictVals
forall a. HasCallStack => String -> a
error (String -> DictVals) -> String -> DictVals
forall a b. (a -> b) -> a -> b
$ String
"Unsupported dictionary type: " String -> ShowS
forall a. [a] -> [a] -> [a]
++ ThriftType -> String
forall a. Show a => a -> String
show ThriftType
t

readPageInt32 :: BS.ByteString -> [Int32]
readPageInt32 :: ByteString -> [Int32]
readPageInt32 ByteString
xs
    | ByteString -> Bool
BS.null ByteString
xs = []
    | Bool
otherwise = ByteString -> Int32
littleEndianInt32 (Int -> ByteString -> ByteString
BS.take Int
4 ByteString
xs) Int32 -> [Int32] -> [Int32]
forall a. a -> [a] -> [a]
: ByteString -> [Int32]
readPageInt32 (Int -> ByteString -> ByteString
BS.drop Int
4 ByteString
xs)

readPageWord64 :: BS.ByteString -> [Double]
readPageWord64 :: ByteString -> [Double]
readPageWord64 ByteString
xs
    | ByteString -> Bool
BS.null ByteString
xs = []
    | Bool
otherwise =
        Word64 -> Double
castWord64ToDouble (ByteString -> Word64
littleEndianWord64 (Int -> ByteString -> ByteString
BS.take Int
8 ByteString
xs))
            Double -> [Double] -> [Double]
forall a. a -> [a] -> [a]
: ByteString -> [Double]
readPageWord64 (Int -> ByteString -> ByteString
BS.drop Int
8 ByteString
xs)

readPageBytes :: BS.ByteString -> [T.Text]
readPageBytes :: ByteString -> [Text]
readPageBytes ByteString
xs
    | ByteString -> Bool
BS.null ByteString
xs = []
    | Bool
otherwise =
        let lenBytes :: Int
lenBytes = Int32 -> Int
forall a b. (Integral a, Num b) => a -> b
fromIntegral (ByteString -> Int32
littleEndianInt32 (ByteString -> Int32) -> ByteString -> Int32
forall a b. (a -> b) -> a -> b
$ Int -> ByteString -> ByteString
BS.take Int
4 ByteString
xs)
            totalBytesRead :: Int
totalBytesRead = Int
lenBytes Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
4
         in ByteString -> Text
decodeUtf8Lenient (Int -> ByteString -> ByteString
BS.take Int
lenBytes (Int -> ByteString -> ByteString
BS.drop Int
4 ByteString
xs))
                Text -> [Text] -> [Text]
forall a. a -> [a] -> [a]
: ByteString -> [Text]
readPageBytes (Int -> ByteString -> ByteString
BS.drop Int
totalBytesRead ByteString
xs)

readPageBool :: BS.ByteString -> [Bool]
readPageBool :: ByteString -> [Bool]
readPageBool ByteString
bs =
    (Word8 -> [Bool]) -> [Word8] -> [Bool]
forall (t :: * -> *) a b. Foldable t => (a -> [b]) -> t a -> [b]
concatMap (\Word8
b -> (Int -> Bool) -> [Int] -> [Bool]
forall a b. (a -> b) -> [a] -> [b]
map (\Int
i -> (Word8
b Word8 -> Int -> Word8
forall a. Bits a => a -> Int -> a
`shiftR` Int
i) Word8 -> Word8 -> Word8
forall a. Bits a => a -> a -> a
.&. Word8
1 Word8 -> Word8 -> Bool
forall a. Eq a => a -> a -> Bool
== Word8
1) [Int
0 .. Int
7]) (ByteString -> [Word8]
BS.unpack ByteString
bs)

readPageInt64 :: BS.ByteString -> [Int64]
readPageInt64 :: ByteString -> [Int64]
readPageInt64 ByteString
xs
    | ByteString -> Bool
BS.null ByteString
xs = []
    | Bool
otherwise =
        Word64 -> Int64
forall a b. (Integral a, Num b) => a -> b
fromIntegral (ByteString -> Word64
littleEndianWord64 (Int -> ByteString -> ByteString
BS.take Int
8 ByteString
xs)) Int64 -> [Int64] -> [Int64]
forall a. a -> [a] -> [a]
: ByteString -> [Int64]
readPageInt64 (Int -> ByteString -> ByteString
BS.drop Int
8 ByteString
xs)

readPageFloat :: BS.ByteString -> [Float]
readPageFloat :: ByteString -> [Float]
readPageFloat ByteString
xs
    | ByteString -> Bool
BS.null ByteString
xs = []
    | Bool
otherwise =
        Word32 -> Float
castWord32ToFloat (ByteString -> Word32
littleEndianWord32 (Int -> ByteString -> ByteString
BS.take Int
4 ByteString
xs))
            Float -> [Float] -> [Float]
forall a. a -> [a] -> [a]
: ByteString -> [Float]
readPageFloat (Int -> ByteString -> ByteString
BS.drop Int
4 ByteString
xs)

readNInt96Times :: Int -> BS.ByteString -> ([UTCTime], BS.ByteString)
readNInt96Times :: Int -> ByteString -> ([UTCTime], ByteString)
readNInt96Times Int
0 ByteString
bs = ([], ByteString
bs)
readNInt96Times Int
k ByteString
bs =
    let timestamp96 :: ByteString
timestamp96 = Int -> ByteString -> ByteString
BS.take Int
12 ByteString
bs
        utcTime :: UTCTime
utcTime = ByteString -> UTCTime
int96ToUTCTime ByteString
timestamp96
        bs' :: ByteString
bs' = Int -> ByteString -> ByteString
BS.drop Int
12 ByteString
bs
        ([UTCTime]
times, ByteString
rest) = Int -> ByteString -> ([UTCTime], ByteString)
readNInt96Times (Int
k Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1) ByteString
bs'
     in (UTCTime
utcTime UTCTime -> [UTCTime] -> [UTCTime]
forall a. a -> [a] -> [a]
: [UTCTime]
times, ByteString
rest)

readPageInt96Times :: BS.ByteString -> [UTCTime]
readPageInt96Times :: ByteString -> [UTCTime]
readPageInt96Times ByteString
bs
    | ByteString -> Bool
BS.null ByteString
bs = []
    | Bool
otherwise =
        let ([UTCTime]
times, ByteString
_) = Int -> ByteString -> ([UTCTime], ByteString)
readNInt96Times (ByteString -> Int
BS.length ByteString
bs Int -> Int -> Int
forall a. Integral a => a -> a -> a
`div` Int
12) ByteString
bs
         in [UTCTime]
times

readPageFixedBytes :: BS.ByteString -> Int -> [T.Text]
readPageFixedBytes :: ByteString -> Int -> [Text]
readPageFixedBytes ByteString
xs Int
len
    | ByteString -> Bool
BS.null ByteString
xs = []
    | Bool
otherwise =
        ByteString -> Text
decodeUtf8Lenient (Int -> ByteString -> ByteString
BS.take Int
len ByteString
xs) Text -> [Text] -> [Text]
forall a. a -> [a] -> [a]
: ByteString -> Int -> [Text]
readPageFixedBytes (Int -> ByteString -> ByteString
BS.drop Int
len ByteString
xs) Int
len

unpackBitPacked :: Int -> Int -> BS.ByteString -> ([Word32], BS.ByteString)
unpackBitPacked :: Int -> Int -> ByteString -> ([Word32], ByteString)
unpackBitPacked Int
bw Int
count ByteString
bs
    | Int
count Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
<= Int
0 = ([], ByteString
bs)
    | ByteString -> Bool
BS.null ByteString
bs = ([], ByteString
bs)
    | Bool
otherwise =
        let totalBytes :: Int
totalBytes = (Int
bw Int -> Int -> Int
forall a. Num a => a -> a -> a
* Int
count 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
            chunk :: ByteString
chunk = Int -> ByteString -> ByteString
BS.take Int
totalBytes ByteString
bs
            rest :: ByteString
rest = Int -> ByteString -> ByteString
BS.drop Int
totalBytes ByteString
bs
         in (Int -> Int -> ByteString -> [Word32]
extractBits Int
bw Int
count ByteString
chunk, ByteString
rest)

-- | LSB-first bit accumulator: reads each byte once with no intermediate ByteString allocation.
extractBits :: Int -> Int -> BS.ByteString -> [Word32]
extractBits :: Int -> Int -> ByteString -> [Word32]
extractBits Int
bw Int
count ByteString
bs = Int -> Word64 -> Int -> Int -> [Word32]
forall {t} {a}.
(Ord t, Num t, Num a) =>
Int -> Word64 -> Int -> t -> [a]
go Int
0 (Word64
0 :: Word64) Int
0 Int
count
  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
`shiftL` 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 -> t -> [a]
go !Int
byteIdx !Word64
acc !Int
accBits !t
remaining
        | t
remaining t -> t -> Bool
forall a. Ord a => a -> a -> Bool
<= t
0 = []
        | Int
accBits Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
bw =
            Word64 -> a
forall a b. (Integral a, Num b) => a -> b
fromIntegral (Word64
acc Word64 -> Word64 -> Word64
forall a. Bits a => a -> a -> a
.&. Word64
mask)
                a -> [a] -> [a]
forall a. a -> [a] -> [a]
: Int -> Word64 -> Int -> t -> [a]
go Int
byteIdx (Word64
acc Word64 -> Int -> Word64
forall a. Bits a => a -> Int -> a
`shiftR` Int
bw) (Int
accBits Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
bw) (t
remaining t -> t -> t
forall a. Num a => a -> a -> a
- t
1)
        | Int
byteIdx Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
len = []
        | 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 -> t -> [a]
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
`shiftL` Int
accBits)) (Int
accBits Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
8) t
remaining

decodeRLEBitPackedHybrid :: Int -> BS.ByteString -> ([Word32], BS.ByteString)
decodeRLEBitPackedHybrid :: Int -> ByteString -> ([Word32], ByteString)
decodeRLEBitPackedHybrid Int
bitWidth ByteString
bs
    | Int
bitWidth Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
== Int
0 = ([Word32
0], ByteString
bs)
    | ByteString -> Bool
BS.null ByteString
bs = ([], ByteString
bs)
    | Bool
otherwise =
        -- readUVarInt is evaluated here, inside the guard that has already
        -- confirmed bs is non-empty.  Keeping it in a where clause would cause
        -- it to be forced before the BS.null guard under {-# LANGUAGE Strict #-}.
        let (Word64
hdr64, ByteString
afterHdr) = ByteString -> (Word64, ByteString)
readUVarInt ByteString
bs
            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
                    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
                     in Int -> Int -> ByteString -> ([Word32], ByteString)
unpackBitPacked Int
bitWidth Int
totalVals ByteString
afterHdr
                else
                    let mask :: Word32
mask = if Int
bitWidth 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
bitWidth) Word32 -> Word32 -> Word32
forall a. Num a => a -> a -> a
- Word32
1
                        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
bitWidth 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 :: Int
                        word32 :: Word32
word32 = ByteString -> Word32
littleEndianWord32 (Int -> ByteString -> ByteString
BS.take Int
4 ByteString
afterHdr)
                        value :: Word32
value = Word32
word32 Word32 -> Word32 -> Word32
forall a. Bits a => a -> a -> a
.&. Word32
mask
                     in (Int -> Word32 -> [Word32]
forall a. Int -> a -> [a]
replicate Int
runLen Word32
value, Int -> ByteString -> ByteString
BS.drop Int
nBytes ByteString
afterHdr)