duckdb-simple-0.3.0.0: src/Database/DuckDB/Simple/VariantCodec.hs
{-# LANGUAGE BlockArguments #-}
{-# LANGUAGE LambdaCase #-}
{-# LANGUAGE OverloadedStrings #-}
{-# LANGUAGE ScopedTypeVariables #-}
{- | Checked access to DuckDB 1.5's private VARIANT payload.
Only public C handles and vector accessors cross the native boundary.
-}
module Database.DuckDB.Simple.VariantCodec (
decodeVariant,
prepareVariantDecoder,
decodeVariantPayload,
) where
import Control.Exception (bracket, throwIO)
import Control.Monad (forM, forM_, unless, when)
import Control.Monad.Trans.Class (lift)
import Control.Monad.Trans.State.Strict (StateT, evalStateT, gets, modify')
import Data.Array (Array, bounds, listArray, (!))
import Data.Bits (complement, finiteBitSize, shiftL, (.&.), (.|.))
import Data.ByteString (ByteString)
import qualified Data.ByteString as BS
import qualified Data.IntMap.Strict as IntMap
import qualified Data.IntSet as IntSet
import qualified Data.Set as Set
import Data.Text (Text)
import qualified Data.Text as Text
import qualified Data.Text.Encoding as Text
import Data.Word (Word32, Word8)
import Database.DuckDB.FFI
import Database.DuckDB.Simple.Element (bitStringFromBytes, chunkDecodeBlob, chunkIsRowValid, decodeElement)
import Database.DuckDB.Simple.FromField (
BigNum (..),
DecimalValue (..),
FieldValue (..),
RawGeometry (..),
fromBigNumBytes,
)
import Database.DuckDB.Simple.Internal (destroyLogicalType)
import Database.DuckDB.Simple.Variant (variantObject)
import Foreign.C.String (peekCString)
import Foreign.Marshal.Alloc (alloca, allocaBytesAligned)
import Foreign.Marshal.Utils (copyBytes)
import Foreign.Ptr (Ptr, castPtr, nullPtr)
import Foreign.Storable (Storable, peek, peekElemOff, poke, sizeOf)
import Text.Read (readMaybe)
-- | Raise a codec error before an invalid native operation.
codecError :: String -> IO a
codecError message = throwIO (userError ("duckdb-simple: VARIANT: " <> message))
-- | Convert a checked pure result to an IO result.
checked :: Either String a -> IO a
checked = either codecError pure
-- | Reject native versions outside the supported private-format range.
checkVersion :: IO ()
checkVersion = do
version <- c_duckdb_library_version >>= peekCString
let parts = Text.splitOn "." (Text.pack version)
patch = case parts of
["v1", "5", p] -> readMaybe (Text.unpack p) :: Maybe Int
_ -> Nothing
unless (maybe False (>= 3) patch) (codecError ("unsupported native version " <> version))
-- | Check the byte order and pointer width required by native scalar loads.
checkPlatform :: IO ()
checkPlatform = alloca \ptr -> do
poke ptr (1 :: Word32)
first <- peek (castPtr ptr :: Ptr Word8)
unless (first == 1 && sizeOf (nullPtr :: Ptr ()) == 8) $
codecError "the private codec requires a 64-bit little-endian host"
-- | Acquire a non-NULL native handle.
nonNull :: String -> IO (Ptr a) -> IO (Ptr a)
nonNull label action = do
ptr <- action
when (ptr == nullPtr) (codecError (label <> " returned NULL"))
pure ptr
-- | Check a logical type's tag.
expectType :: DuckDBType -> DuckDBLogicalType -> IO ()
expectType expected logical = do
actual <- c_duckdb_get_type_id logical
unless (actual == expected) (codecError "unexpected physical child type")
-- | Check a STRUCT's names and owned child descriptors.
checkStruct :: DuckDBLogicalType -> [(Text, DuckDBLogicalType -> IO ())] -> IO ()
checkStruct logical fields = do
count <- c_duckdb_struct_type_child_count logical
unless (count == fromIntegral (length fields)) (codecError "unexpected physical child count")
sequence_
[ do
bracket (nonNull "child name" (c_duckdb_struct_type_child_name logical index)) (c_duckdb_free . castPtr) \name -> do
bytes <- BS.packCString name
unless (bytes == Text.encodeUtf8 expectedName) (codecError "unexpected physical child name")
bracket (nonNull "child type" (c_duckdb_struct_type_child_type logical index)) destroyLogicalType checkChild
| (index, (expectedName, checkChild)) <- zip [0 ..] fields
]
-- | Check a LIST and its owned element descriptor.
checkList :: (DuckDBLogicalType -> IO ()) -> DuckDBLogicalType -> IO ()
checkList checkChild logical = do
expectType DuckDBTypeList logical
bracket (nonNull "list child type" (c_duckdb_list_type_child_type logical)) destroyLogicalType checkChild
-- | Check the complete, unshredded four-child VARIANT schema.
checkSchema :: DuckDBLogicalType -> IO ()
checkSchema logical = do
expectType DuckDBTypeVariant logical
checkStruct
logical
[ ("keys", checkList (expectType DuckDBTypeVarchar))
,
( "children"
, checkList \child -> do
expectType DuckDBTypeStruct child
checkStruct child [("keys_index", expectType DuckDBTypeUInteger), ("values_index", expectType DuckDBTypeUInteger)]
)
,
( "values"
, checkList \child -> do
expectType DuckDBTypeStruct child
checkStruct child [("type_id", expectType DuckDBTypeUTinyInt), ("byte_offset", expectType DuckDBTypeUInteger)]
)
, ("data", expectType DuckDBTypeBlob)
]
-- | Check an element index before native pointer arithmetic.
checkIndex :: Int -> Int -> IO ()
checkIndex width index =
when (index < 0 || index > maxBound `div` width) (codecError "native element index exceeds Int range")
-- | Prepare validity access for a chunk. The caller checks row bounds.
prepareValidity :: DuckDBVector -> IO (Int -> IO Bool)
prepareValidity vector = do
validity <- c_duckdb_vector_get_validity vector
pure (chunkIsRowValid validity . fromIntegral)
-- | Borrow a fixed-width buffer until its chunk is destroyed.
prepareElementReader :: (Storable a) => DuckDBVector -> IO (Int -> IO a)
prepareElementReader vector = do
ptr <- c_duckdb_vector_get_data vector
valid <- prepareValidity vector
pure (readAt valid (castPtr ptr))
where
readAt :: (Storable a) => (Int -> IO Bool) -> Ptr a -> Int -> IO a
readAt valid ptr index = do
checkIndex (sizeOfElement ptr) index
when (ptr == nullPtr) (codecError "NULL vector data")
present <- valid index
unless present (codecError "NULL physical payload element")
peekElemOff ptr index
sizeOfElement :: (Storable a) => Ptr a -> Int
sizeOfElement ptr = sizeOf (undefined `asTypeOf` element ptr)
element :: Ptr a -> a
element _ = undefined
-- | Borrow a string buffer and copy each requested value into Haskell memory.
prepareBytesReader :: DuckDBVector -> IO (Int -> IO ByteString)
prepareBytesReader vector = do
base <- c_duckdb_vector_get_data vector
valid <- prepareValidity vector
pure \index -> do
checkIndex 16 index
present <- valid index
unless present (codecError "NULL physical string element")
when (base == nullPtr) (codecError "NULL string vector data")
chunkDecodeBlob base (fromIntegral index)
-- | Convert a nonnegative bounded integer to Int.
checkedInt :: Integer -> IO Int
checkedInt n
| n < 0 || n > toInteger (maxBound :: Int) = codecError "payload size exceeds Int range"
| otherwise = pure (fromInteger n)
-- | Prepare LIST bounds checks against the chunk's fixed child size.
prepareListBounds :: DuckDBVector -> IO (Int -> IO (Int, Int))
prepareListBounds vector = do
readEntry <- prepareElementReader vector
size <- c_duckdb_list_vector_get_size vector
pure \row -> do
DuckDBListEntry offset count <- readEntry row
unless (offset <= size && count <= size - offset) (codecError "LIST bounds exceed child size")
start <- checkedInt (toInteger offset)
len <- checkedInt (toInteger count)
_ <- checkedInt (toInteger offset + toInteger count)
pure (start, len)
-- | Decode one row while its flattened result chunk remains alive.
decodeVariant :: DuckDBVector -> Int -> IO FieldValue
decodeVariant vector row = prepareVariantDecoder vector >>= ($ row)
{- | Check the format once and borrow buffers for a flattened result chunk.
The returned reader must not outlive the chunk. The caller supplies row indices
within that chunk. DuckDB result Fetch flattens nested vectors in 1.5.
Each read copies its payload into Haskell memory, including referenced keys.
-}
prepareVariantDecoder :: DuckDBVector -> IO (Int -> IO FieldValue)
prepareVariantDecoder vector = do
checkVersion
checkPlatform
when (vector == nullPtr) (codecError "NULL vector")
bracket (nonNull "vector type" (c_duckdb_vector_get_column_type vector)) destroyLogicalType checkSchema
valid <- prepareValidity vector
keys <- child vector 0
children <- child vector 1
values <- child vector 2
blob <- child vector 3
keyBounds <- prepareListBounds keys
childBounds <- prepareListBounds children
valueBounds <- prepareListBounds values
keyVector <- nonNull "keys vector" (c_duckdb_list_vector_get_child keys)
childVector <- nonNull "children vector" (c_duckdb_list_vector_get_child children)
valueVector <- nonNull "values vector" (c_duckdb_list_vector_get_child values)
keyIndices <- child childVector 0
valueIndices <- child childVector 1
tags <- child valueVector 0
offsets <- child valueVector 1
readTag <- prepareElementReader tags
readOffset <- prepareElementReader offsets
hasKey <- prepareValidity keyIndices
readKey <- prepareElementReader keyIndices
readValue <- prepareElementReader valueIndices
readKeyBytes <- prepareBytesReader keyVector
readBlob <- prepareBytesReader blob
pure \row -> do
checkIndex 16 row
present <- valid row
if not present
then pure FieldNull
else do
(keyStart, keyCount) <- keyBounds row
(childStart, childCount) <- childBounds row
(valueStart, valueCount) <- valueBounds row
valueRows <- forM [valueStart .. valueStart + valueCount - 1] \index ->
(,) <$> readTag index <*> readOffset index
childRows <- forM [childStart .. childStart + childCount - 1] \index -> do
keyed <- hasKey index
key <- if keyed then Just <$> readKey index else pure Nothing
value <- readValue index
when (toInteger value >= toInteger valueCount) (codecError "child value index out of bounds")
case key of
Just k | toInteger k >= toInteger keyCount -> codecError "child key index out of bounds"
_ -> pure ()
pure (key, value)
let usedKeys = Set.toList (Set.fromList [k | (Just k, _) <- childRows])
keyRows <- forM usedKeys \index -> do
bytes <- readKeyBytes (keyStart + fromIntegral index)
text <- checked (either (Left . show) Right (Text.decodeUtf8' bytes))
pure (index, text)
bytes <- readBlob row
decodeVariantPayload valueRows childRows keyRows bytes
where
child parent index = nonNull "STRUCT vector child" (c_duckdb_struct_vector_get_child parent index)
{- | Decode copied 1.5 payload data for one non-NULL row.
Values are (tag, byte offset). Children are (optional key index, value index).
Keys pair an index with its text. Indices are relative to this row's LISTs.
The root value has index zero. This helper checks all metadata before decoding.
It raises an error for cycles. Shared values use a memo table.
-}
decodeVariantPayload :: [(Word8, Word32)] -> [(Maybe Word32, Word32)] -> [(Word32, Text)] -> ByteString -> IO FieldValue
decodeVariantPayload valueRows childRows keyRows bytes = do
when (finiteBitSize (0 :: Int) < 64) (codecError "the private codec requires a 64-bit host")
let values = listArray (0, length valueRows - 1) valueRows
children = listArray (0, length childRows - 1) childRows
keys = IntMap.fromList [(fromIntegral k, t) | (k, t) <- keyRows]
valueCount = length valueRows
childCount = length childRows
when (null valueRows) (codecError "missing root value")
unless (IntMap.size keys == length keyRows) (codecError "duplicate key dictionary index")
forM_ valueRows \(tag, offset) -> do
when (tag > 33) (codecError "unknown payload tag")
when (toInteger offset > toInteger (BS.length bytes)) (codecError "byte offset exceeds data size")
forM_ childRows \(key, value) -> do
when (toInteger value >= toInteger valueCount) (codecError "child value index out of bounds")
case key of
Just k -> unless (IntMap.member (fromIntegral k) keys) (codecError "missing child key")
Nothing -> pure ()
evalStateT (visit values children keys childCount IntSet.empty 0) IntMap.empty
where
visit :: Array Int (Word8, Word32) -> Array Int (Maybe Word32, Word32) -> IntMap.IntMap Text -> Int -> IntSet.IntSet -> Int -> StateT (IntMap.IntMap FieldValue) IO FieldValue
visit values children keys childCount ancestors index = do
when (IntSet.member index ancestors) (lift (codecError "cyclic child reference"))
cached <- gets (IntMap.lookup index)
case cached of
Just result -> pure result
Nothing -> do
(tag, offset) <- lift (checked (arrayElement values index))
let payload = BS.drop (fromIntegral offset) bytes
result <- case tag of
29 -> nested True payload
30 -> nested False payload
_ -> lift (decodeScalar tag payload)
modify' (IntMap.insert index result)
pure result
where
nested object payload = do
(count, rest) <- lift (checked (readVarint payload))
start <- if count == 0 then pure 0 else fst <$> lift (checked (readVarint rest))
unless (toInteger start + toInteger count <= toInteger childCount) $
lift (codecError "container child range out of bounds")
entries <- forM [fromIntegral start .. fromIntegral start + fromIntegral count - 1] \childIndex -> do
(key, childIndexValue) <- lift (checked (arrayElement children childIndex))
name <- case (object, key) of
(True, Just k) -> case IntMap.lookup (fromIntegral k) keys of
Just text -> pure text
Nothing -> lift (codecError "missing object key")
(False, Nothing) -> pure Text.empty
_ -> lift (codecError "container key validity does not match its tag")
item <- visit values children keys childCount (IntSet.insert index ancestors) (fromIntegral childIndexValue)
pure (name, item)
if object
then do
unless (Set.size (Set.fromList (map fst entries)) == length entries) $
lift (codecError "duplicate object key")
pure (variantObject entries)
else pure (FieldList (map snd entries))
-- | Read an array element after checking both bounds.
arrayElement :: Array Int a -> Int -> Either String a
arrayElement array index
| index < lower || index > upper = Left "value index out of bounds"
| otherwise = Right (array ! index)
where
(lower, upper) = bounds array
-- | Read an unsigned base-128 uint32 without overflow or truncation.
readVarint :: ByteString -> Either String (Word32, ByteString)
readVarint = go 0 0
where
go shift value input = case BS.uncons input of
Nothing -> Left "truncated varint"
Just (byte, rest)
| shift == 28 && byte > 15 -> Left "varint exceeds uint32"
| otherwise ->
let result = value .|. (fromIntegral (byte .&. 127) `shiftL` shift)
in if byte .&. 128 == 0 then Right (result, rest) else go (shift + 7) result rest
-- | Read a bounded fixed-size payload in little-endian order.
readUnsigned :: Int -> ByteString -> Either String (Integer, ByteString)
readUnsigned count bytes
| BS.length bytes < count = Left "truncated scalar payload"
| otherwise =
let (part, rest) = BS.splitAt count bytes
in Right (BS.foldr (\byte value -> value `shiftL` 8 .|. toInteger byte) 0 part, rest)
-- | Read a two's-complement integer with a checked payload length.
readSigned :: Int -> ByteString -> Either String Integer
readSigned count bytes = do
(value, _) <- readUnsigned count bytes
pure (if value >= 2 ^ (count * 8 - 1) then value - 2 ^ (count * 8) else value)
-- | Read a length-prefixed string or blob.
readString :: ByteString -> Either String ByteString
readString bytes = do
(count, rest) <- readVarint bytes
if toInteger count > toInteger (BS.length rest)
then Left "string length exceeds payload size"
else Right (BS.take (fromIntegral count) rest)
-- | Decode the scalar tags from VariantLogicalType in DuckDB 1.5.
decodeScalar :: Word8 -> ByteString -> IO FieldValue
decodeScalar tag bytes = case tag of
0 -> pure FieldNull
1 -> pure (FieldBool True)
2 -> pure (FieldBool False)
15 -> checked do
(precision, rest) <- readVarint bytes
(scale, digits) <- readVarint rest
when (precision < 1 || precision > 38 || scale > precision) (Left "invalid decimal metadata")
let count = if precision <= 4 then 2 else if precision <= 9 then 4 else if precision <= 18 then 8 else 16
value <- readSigned count digits
when (abs value >= 10 ^ precision) (Left "unscaled decimal exceeds its precision")
pure (FieldDecimal (DecimalValue (fromIntegral precision) (fromIntegral scale) value))
16 -> checked do
string <- readString bytes
FieldText <$> either (Left . show) Right (Text.decodeUtf8' string)
17 -> FieldBlob <$> checked (readString bytes)
31 -> FieldBigNum . BigNum <$> checked (readString bytes >>= decodeBigNum)
32 -> checked do
bitBytes <- readString bytes
case BS.unpack (BS.take 2 bitBytes) of
[padding, first] -> do
when (padding > 7) (Left "BIT padding exceeds seven")
let maskBits = paddingMask padding
unless (first .&. maskBits == maskBits) (Left "invalid native BIT padding")
pure (FieldBit (bitStringFromBytes bitBytes))
_ -> Left "BIT requires a padding byte and nonempty data"
33 -> FieldGeometry . (`RawGeometry` Nothing) <$> checked (readString bytes)
_ -> case fixedWidthTag tag of
Just (dtype, size) -> decodeFixedWidth dtype size bytes
Nothing -> codecError "unknown scalar tag"
{- | The type and the payload size of each fixed-width scalar tag. These
payloads have the memory layout of one vector element of the type.
-}
fixedWidthTag :: Word8 -> Maybe (DuckDBType, Int)
fixedWidthTag = \case
3 -> Just (DuckDBTypeTinyInt, 1)
4 -> Just (DuckDBTypeSmallInt, 2)
5 -> Just (DuckDBTypeInteger, 4)
6 -> Just (DuckDBTypeBigInt, 8)
7 -> Just (DuckDBTypeHugeInt, 16)
8 -> Just (DuckDBTypeUTinyInt, 1)
9 -> Just (DuckDBTypeUSmallInt, 2)
10 -> Just (DuckDBTypeUInteger, 4)
11 -> Just (DuckDBTypeUBigInt, 8)
12 -> Just (DuckDBTypeUHugeInt, 16)
13 -> Just (DuckDBTypeFloat, 4)
14 -> Just (DuckDBTypeDouble, 8)
18 -> Just (DuckDBTypeUUID, 16)
19 -> Just (DuckDBTypeDate, 4)
20 -> Just (DuckDBTypeTime, 8)
21 -> Just (DuckDBTypeTimeNs, 8)
22 -> Just (DuckDBTypeTimestampS, 8)
23 -> Just (DuckDBTypeTimestampMs, 8)
24 -> Just (DuckDBTypeTimestamp, 8)
25 -> Just (DuckDBTypeTimestampNs, 8)
26 -> Just (DuckDBTypeTimeTz, 8)
27 -> Just (DuckDBTypeTimestampTz, 8)
28 -> Just (DuckDBTypeInterval, 16)
_ -> Nothing
-- | Copy a fixed-width payload to aligned memory and decode it as one element.
decodeFixedWidth :: DuckDBType -> Int -> ByteString -> IO FieldValue
decodeFixedWidth dtype size bytes = do
when (BS.length bytes < size) (codecError "truncated scalar payload")
allocaBytesAligned size 16 \buffer -> do
BS.useAsCStringLen bytes \(source, _) -> copyBytes buffer (castPtr source) size
decodeElement dtype (castPtr buffer) 0
-- | Get the high-bit mask used for native BIT padding.
paddingMask :: Word8 -> Word8
paddingMask padding = complement ((1 `shiftL` (8 - fromIntegral padding)) - 1)
-- | Check BIGNUM's sign, three-byte length header, and magnitude, then decode it.
decodeBigNum :: ByteString -> Either String Integer
decodeBigNum bytes = do
when (BS.length bytes < 4) (Left "truncated BIGNUM header or magnitude")
let (header, raw) = BS.splitAt 3 bytes
encoded = BS.foldl' (\n byte -> n `shiftL` 8 .|. fromIntegral byte) (0 :: Word32) header
negative = encoded .&. 0x800000 == 0
decoded = if negative then complement encoded .&. 0xffffff else encoded
count = decoded .&. 0x7fffff
magnitude = if negative then BS.map complement raw else raw
unless (toInteger count == toInteger (BS.length magnitude)) (Left "BIGNUM header length mismatch")
case BS.uncons magnitude of
Just (0, rest) | not (BS.null rest) || negative -> Left "noncanonical BIGNUM magnitude"
_ -> Right ()
pure (fromBigNumBytes (BS.unpack bytes))