diff --git a/dataframe-parquet/src/DataFrame/IO/Parquet/Writer.hs b/dataframe-parquet/src/DataFrame/IO/Parquet/Writer.hs index 2e8d4e9f..8dca6ca7 100644 --- a/dataframe-parquet/src/DataFrame/IO/Parquet/Writer.hs +++ b/dataframe-parquet/src/DataFrame/IO/Parquet/Writer.hs @@ -1,4 +1,5 @@ {-# LANGUAGE BangPatterns #-} +{-# LANGUAGE FlexibleContexts #-} {-# LANGUAGE OverloadedRecordDot #-} {-# LANGUAGE OverloadedStrings #-} @@ -13,11 +14,19 @@ module DataFrame.IO.Parquet.Writer ( ) where import Control.Monad (forM_, unless, when) +import Control.Monad.IO.Class (MonadIO) +import Control.Monad.Primitive (PrimBase, PrimMonad, PrimState) import qualified Data.ByteString as BS -import Data.IORef (IORef, modifyIORef', newIORef, readIORef, writeIORef) import Data.Int (Int64) import Data.Maybe (fromJust) import Data.Primitive.ByteArray (getSizeofMutableByteArray) +import Data.Primitive.MutVar ( + MutVar, + modifyMutVar', + newMutVar, + readMutVar, + writeMutVar, + ) import qualified Data.Text as T import qualified Data.Vector as VB import DataFrame.IO.Parquet.Thrift hiding (schema) @@ -71,29 +80,29 @@ import System.Directory (createDirectoryIfMissing) import System.FilePath (takeDirectory) import Text.Printf (printf) -data ParquetWriterState = ParquetWriterState +data ParquetWriterState m = ParquetWriterState { outputFileHandle :: !WritableBinaryHandle - , columnChunks :: !(VB.Vector ColumnChunkState) - , currentFileOffsetRef :: !(IORef Int64) - , scratchBuffer :: !MemoryBuffer - , rowGroupMetadataRef :: !(IORef [RowGroup]) - , rowNumberRef :: !(IORef Int) + , columnChunks :: !(VB.Vector (ColumnChunkState m)) + , currentFileOffsetRef :: !(MutVar (PrimState m) Int64) + , scratchBuffer :: !(MemoryBuffer (PrimState m)) + , rowGroupMetadataRef :: !(MutVar (PrimState m) [RowGroup]) + , rowNumberRef :: !(MutVar (PrimState m) Int) } -data ColumnChunkState = ColumnChunkState +data ColumnChunkState m = ColumnChunkState { columnName :: !T.Text , nullable :: !Bool , schema :: !SchemaElement - , encoder :: !Encoder - , buffer :: !MemoryBuffer - , uncompressedBufferSize :: !(IORef Int64) - , pageState :: !PageState + , encoder :: !(Encoder m) + , buffer :: !(MemoryBuffer (PrimState m)) + , uncompressedBufferSize :: !(MutVar (PrimState m) Int64) + , pageState :: !(PageState m) } -data PageState = PageState - { pageBuffer :: !MemoryBuffer - , definitionLevels :: !DefLevels - , currentRowCount :: !(IORef Int) +data PageState m = PageState + { pageBuffer :: !(MemoryBuffer (PrimState m)) + , definitionLevels :: !(DefLevels (PrimState m)) + , currentRowCount :: !(MutVar (PrimState m) Int) } writeParquet :: FilePath -> DataFrame -> IO () @@ -144,7 +153,12 @@ shardPathFor pattern_ shardIndex = -- | Write rows @[startRow, endRow)@ of the frame to a single Parquet file. writeShard :: - ParquetWriteOptions -> FilePath -> DataFrame -> Int -> Int -> IO () + ParquetWriteOptions -> + FilePath -> + DataFrame -> + Int -> + Int -> + IO () writeShard options path_ df startRow endRow = do let names = columnNames df shardRows = max 0 (endRow - startRow) @@ -161,9 +175,9 @@ writeShard options path_ df startRow endRow = do scratchBuffer_ <- mallocBuffer (max 1 options.pageSize) atomicallyWriteFile path_ $ \path -> withWritableBinaryFile path $ \output -> do writeByteStringToFile output magic - currentFileOffsetRef_ <- newIORef 4 - rowGroupMetadataRef_ <- newIORef [] - rowNumberRef_ <- newIORef 0 + currentFileOffsetRef_ <- newMutVar 4 + rowGroupMetadataRef_ <- newMutVar [] + rowNumberRef_ <- newMutVar 0 let writerState = ParquetWriterState output @@ -174,26 +188,26 @@ writeShard options path_ df startRow endRow = do rowNumberRef_ interval = max 1 options.batchRows subBatch = max 1 options.subBatchRows - writeBatch :: Int -> Int -> IO () writeBatch rowNum batchEnd | rowNum >= batchEnd = pure () | otherwise = do let count = min subBatch (batchEnd - rowNum) VB.forM_ columnChunks_ (writeRows options scratchBuffer_ rowNum count) - modifyIORef' rowNumberRef_ (+ count) + modifyMutVar' rowNumberRef_ (+ count) writeBatch (rowNum + count) batchEnd - loop :: Int -> IO () loop rowNum | rowNum >= endRow = pure () | otherwise = do let batchEnd = rowNum + min interval (endRow - rowNum) writeBatch rowNum batchEnd size <- bufferedSize columnChunks_ - when (size >= options.rowGroupSize) (flushRowGroup options writerState) + when (size >= options.rowGroupSize) $ + flushRowGroup options writerState loop batchEnd loop startRow flushRowGroup options writerState - rowGroupMetadata <- reverse <$> readIORef rowGroupMetadataRef_ + rowGroupMetadata <- + reverse <$> readMutVar rowGroupMetadataRef_ let schemaElements = rootSchemaElement (VB.length columnChunks_) : VB.toList (VB.map schema columnChunks_) @@ -216,7 +230,13 @@ nativeTypeKeyValues names df = ] writeRows :: - ParquetWriteOptions -> MemoryBuffer -> Int -> Int -> ColumnChunkState -> IO () + (PrimBase m, MonadIO m) => + ParquetWriteOptions -> + MemoryBuffer (PrimState m) -> + Int -> + Int -> + ColumnChunkState m -> + m () writeRows options scratch firstRow count ccs = do let page = ccs.pageState buf = page.pageBuffer @@ -224,17 +244,20 @@ writeRows options scratch firstRow count ccs = do dl = page.definitionLevels end = firstRow + count - pos0 <- readIORef buf.positionRef + pos0 <- readMutVar buf.positionRef let margin = options.pageSize arr0 <- ensureCapacity buf (pos0 + max margin (count * 64)) size0 <- getSizeofMutableByteArray arr0 let go !size !pos !row - | row >= end = writeIORef buf.positionRef pos + | row >= end = writeMutVar buf.positionRef pos | pos + margin > size = do -- Rare: buffer nearly full, grow it - writeIORef buf.positionRef pos - arr' <- ensureCapacity buf (pos + max margin ((end - row) * 64)) + writeMutVar buf.positionRef pos + arr' <- + ensureCapacity + buf + (pos + max margin ((end - row) * 64)) size' <- getSizeofMutableByteArray arr' go size' pos row | otherwise = do @@ -246,7 +269,7 @@ writeRows options scratch firstRow count ccs = do go size0 pos0 firstRow -- Batch bookkeeping: once per sub-batch instead of per value - modifyIORef' page.currentRowCount (+ count) + modifyMutVar' page.currentRowCount (+ count) flushDef dl pageRes <- bufferResidency buf defRes <- bufferResidency dl.dlBuf @@ -254,22 +277,35 @@ writeRows options scratch firstRow count ccs = do (pageRes + defRes >= options.pageSize) (flushPage options scratch ccs) -flushPage :: ParquetWriteOptions -> MemoryBuffer -> ColumnChunkState -> IO () +flushPage :: + (PrimBase m, MonadIO m) => + ParquetWriteOptions -> + MemoryBuffer (PrimState m) -> + ColumnChunkState m -> + m () flushPage options scratch columnChunkState = do let page = columnChunkState.pageState - numPageRows <- readIORef page.currentRowCount + numPageRows <- readMutVar page.currentRowCount when (numPageRows > 0) $ do - pos <- readIORef page.pageBuffer.positionRef + pos <- readMutVar page.pageBuffer.positionRef pos' <- columnChunkState.encoder.finishValues page.pageBuffer pos - writeIORef page.pageBuffer.positionRef pos' + writeMutVar page.pageBuffer.positionRef pos' body <- assemblePageBody scratch columnChunkState - writeDataPage options.compressionCodec numPageRows body columnChunkState + writeDataPage + options.compressionCodec + numPageRows + body + columnChunkState resetPosition page.pageBuffer resetPosition page.definitionLevels.dlBuf resetPosition scratch - writeIORef page.currentRowCount 0 + writeMutVar page.currentRowCount 0 -assemblePageBody :: MemoryBuffer -> ColumnChunkState -> IO MemoryBuffer +assemblePageBody :: + (PrimMonad m) => + MemoryBuffer (PrimState m) -> + ColumnChunkState m -> + m (MemoryBuffer (PrimState m)) assemblePageBody scratch columnChunkState | not columnChunkState.nullable = pure columnChunkState.pageState.pageBuffer | otherwise = do @@ -283,12 +319,18 @@ assemblePageBody scratch columnChunkState pure scratch writeDataPage :: - CompressionCodec -> Int -> MemoryBuffer -> ColumnChunkState -> IO () + (PrimBase m, MonadIO m) => + CompressionCodec -> + Int -> + MemoryBuffer (PrimState m) -> + ColumnChunkState m -> + m () writeDataPage codec numPageRows body columnChunkState = do uncompressedPageSize <- bufferResidency body compressedBody <- case codec of UNCOMPRESSED _ -> pure Nothing - SNAPPY _ -> Just . Snappy.compress <$> bufferToByteString body + SNAPPY _ -> + Just . Snappy.compress <$> bufferToByteString body other -> error ("writeParquet: unsupported codec " <> show other) let compressedPageSize = maybe uncompressedPageSize BS.length compressedBody headerBytes = @@ -299,13 +341,17 @@ writeDataPage codec numPageRows body columnChunkState = do case compressedBody of Nothing -> flushBufferToBuffer body columnChunkState.buffer Just bytes -> writeByteString columnChunkState.buffer bytes - modifyIORef' + modifyMutVar' columnChunkState.uncompressedBufferSize (+ fromIntegral (BS.length headerBytes + uncompressedPageSize)) -flushRowGroup :: ParquetWriteOptions -> ParquetWriterState -> IO () +flushRowGroup :: + (PrimBase m, MonadIO m) => + ParquetWriteOptions -> + ParquetWriterState m -> + m () flushRowGroup options writerState = do - rowNumber <- readIORef writerState.rowNumberRef + rowNumber <- readMutVar writerState.rowNumberRef when (rowNumber > 0) $ do VB.forM_ writerState.columnChunks @@ -313,14 +359,22 @@ flushRowGroup options writerState = do (reversedColumnChunks, totalCompressed, totalUncompressed) <- VB.foldM' ( \(acc, totalCompressedSize, totalUncompressedSize) columnChunkState -> do - offset <- readIORef writerState.currentFileOffsetRef - compressedSize <- bufferResidency columnChunkState.buffer - uncompressedSize <- readIORef columnChunkState.uncompressedBufferSize - flushBufferToFile writerState.outputFileHandle columnChunkState.buffer - writeIORef + offset <- + readMutVar writerState.currentFileOffsetRef + compressedSize <- + bufferResidency columnChunkState.buffer + uncompressedSize <- + readMutVar + columnChunkState.uncompressedBufferSize + flushBufferToFile + writerState.outputFileHandle + columnChunkState.buffer + writeMutVar writerState.currentFileOffsetRef (offset + fromIntegral compressedSize) - writeIORef columnChunkState.uncompressedBufferSize 0 + writeMutVar + columnChunkState.uncompressedBufferSize + 0 let columnChunk = mkColumnChunk options.compressionCodec @@ -338,7 +392,7 @@ flushRowGroup options writerState = do ) ([], 0 :: Int64, 0 :: Int64) writerState.columnChunks - modifyIORef' + modifyMutVar' writerState.rowGroupMetadataRef ( mkRowGroup (reverse reversedColumnChunks) @@ -347,22 +401,31 @@ flushRowGroup options writerState = do rowNumber : ) - writeIORef writerState.rowNumberRef 0 + writeMutVar writerState.rowNumberRef 0 -bufferedSize :: VB.Vector ColumnChunkState -> IO Int +bufferedSize :: + (PrimMonad m) => + VB.Vector (ColumnChunkState m) -> + m Int bufferedSize = VB.foldM' ( \total columnChunkState -> do chunkSize <- bufferResidency columnChunkState.buffer - valuesSize <- bufferResidency columnChunkState.pageState.pageBuffer + valuesSize <- + bufferResidency columnChunkState.pageState.pageBuffer defLevelsSize <- - bufferResidency columnChunkState.pageState.definitionLevels.dlBuf + bufferResidency + columnChunkState.pageState.definitionLevels.dlBuf pure (total + chunkSize + valuesSize + defLevelsSize) ) 0 initColumnChunkState :: - ParquetWriteOptions -> T.Text -> Column -> IO ColumnChunkState + (PrimBase m, MonadIO m) => + ParquetWriteOptions -> + T.Text -> + Column -> + m (ColumnChunkState m) initColumnChunkState options columnName_ column = do encoder_ <- buildEncoder column let nullable_ = hasMissing column @@ -387,7 +450,7 @@ initColumnChunkState options columnName_ column = do -- is likely to hit the page limit, the others are liable to be -- much smaller than the limit. buffer_ <- mallocBuffer bufferSize - uncompressedBufferSize_ <- newIORef 0 + uncompressedBufferSize_ <- newMutVar 0 pageState_ <- initPageState bufferSize pure ColumnChunkState @@ -400,11 +463,11 @@ initColumnChunkState options columnName_ column = do , pageState = pageState_ } -initPageState :: Int -> IO PageState +initPageState :: (PrimMonad m, MonadIO m) => Int -> m (PageState m) initPageState bufferSize = do pageBuffer_ <- mallocBuffer bufferSize definitionLevels_ <- newDefLevels - currentRowCount_ <- newIORef 0 + currentRowCount_ <- newMutVar 0 pure PageState { pageBuffer = pageBuffer_ diff --git a/dataframe-parquet/src/DataFrame/IO/Parquet/Writer/DefLevels.hs b/dataframe-parquet/src/DataFrame/IO/Parquet/Writer/DefLevels.hs index 2204c1e2..65efe19b 100644 --- a/dataframe-parquet/src/DataFrame/IO/Parquet/Writer/DefLevels.hs +++ b/dataframe-parquet/src/DataFrame/IO/Parquet/Writer/DefLevels.hs @@ -8,51 +8,53 @@ module DataFrame.IO.Parquet.Writer.DefLevels ( ) where import Control.Monad (when) +import Control.Monad.IO.Class (MonadIO) +import Control.Monad.Primitive (PrimMonad, PrimState) import Data.Bits (shiftL, shiftR, (.&.), (.|.)) -import Data.IORef (IORef, newIORef, readIORef, writeIORef) +import Data.Primitive.MutVar (MutVar, newMutVar, readMutVar, writeMutVar) import Data.Word (Word64) import DataFrame.IO.Utils.RandomAccess (MemoryBuffer, mallocBuffer, writeWord8) -data DefLevels = DefLevels - { dlBuf :: !MemoryBuffer - , dlValue :: !(IORef Int) - , dlCount :: !(IORef Int) +data DefLevels s = DefLevels + { dlBuf :: !(MemoryBuffer s) + , dlValue :: !(MutVar s Int) + , dlCount :: !(MutVar s Int) } -newDefLevels :: IO DefLevels -newDefLevels = DefLevels <$> mallocBuffer 64 <*> newIORef 0 <*> newIORef 0 +newDefLevels :: (PrimMonad m, MonadIO m) => m (DefLevels (PrimState m)) +newDefLevels = DefLevels <$> mallocBuffer 64 <*> newMutVar 0 <*> newMutVar 0 -pushDef :: DefLevels -> Int -> IO () +pushDef :: (PrimMonad m) => DefLevels (PrimState m) -> Int -> m () pushDef dl value = do - count <- readIORef dl.dlCount + count <- readMutVar dl.dlCount if count == 0 - then writeIORef dl.dlValue value >> writeIORef dl.dlCount 1 + then writeMutVar dl.dlValue value >> writeMutVar dl.dlCount 1 else do - current <- readIORef dl.dlValue + current <- readMutVar dl.dlValue if current == value - then writeIORef dl.dlCount (count + 1) + then writeMutVar dl.dlCount (count + 1) else do writeDefRun dl current count - writeIORef dl.dlValue value - writeIORef dl.dlCount 1 + writeMutVar dl.dlValue value + writeMutVar dl.dlCount 1 {-# INLINE pushDef #-} -flushDef :: DefLevels -> IO () +flushDef :: (PrimMonad m) => DefLevels (PrimState m) -> m () flushDef dl = do - count <- readIORef dl.dlCount + count <- readMutVar dl.dlCount when (count > 0) $ do - value <- readIORef dl.dlValue + value <- readMutVar dl.dlValue writeDefRun dl value count - writeIORef dl.dlCount 0 + writeMutVar dl.dlCount 0 {-# INLINE flushDef #-} -writeDefRun :: DefLevels -> Int -> Int -> IO () +writeDefRun :: (PrimMonad m) => DefLevels (PrimState m) -> Int -> Int -> m () writeDefRun dl value count = do writeLeb128 dl.dlBuf (fromIntegral (count `shiftL` 1)) writeWord8 dl.dlBuf (fromIntegral value) {-# INLINE writeDefRun #-} -writeLeb128 :: MemoryBuffer -> Word64 -> IO () +writeLeb128 :: (PrimMonad m) => MemoryBuffer (PrimState m) -> Word64 -> m () writeLeb128 buffer value | value < 0x80 = writeWord8 buffer (fromIntegral value) | otherwise = do diff --git a/dataframe-parquet/src/DataFrame/IO/Parquet/Writer/Encoder.hs b/dataframe-parquet/src/DataFrame/IO/Parquet/Writer/Encoder.hs index 0df7b76d..3b5b6afe 100644 --- a/dataframe-parquet/src/DataFrame/IO/Parquet/Writer/Encoder.hs +++ b/dataframe-parquet/src/DataFrame/IO/Parquet/Writer/Encoder.hs @@ -4,20 +4,20 @@ {-# LANGUAGE OverloadedStrings #-} {-# LANGUAGE ScopedTypeVariables #-} {-# LANGUAGE TypeApplications #-} +{-# OPTIONS_GHC -funfolding-use-threshold=1000 #-} module DataFrame.IO.Parquet.Writer.Encoder ( Encoder (..), buildEncoder, ) where +import Control.Monad.IO.Class (MonadIO, liftIO) +import Control.Monad.Primitive (PrimBase, PrimMonad, PrimState) import Control.Monad.ST (stToIO) import Data.Bits (shiftL, (.|.)) -import Data.IORef (newIORef, readIORef, writeIORef) import Data.Int (Int32, Int64) -import Data.Primitive.ByteArray ( - withMutableByteArrayContents, - writeByteArray, - ) +import Data.Primitive.ByteArray (withMutableByteArrayContents, writeByteArray) +import Data.Primitive.MutVar (newMutVar, readMutVar, writeMutVar) import qualified Data.Text as T import qualified Data.Text.Array as TA import Data.Text.Internal (Text (Text)) @@ -55,15 +55,16 @@ import GHC.Float (castDoubleToWord64, castFloatToWord32) import Pinch (enum, putField) import Type.Reflection (typeRep) -data Encoder = Encoder +data Encoder m = Encoder { encType :: !ThriftType , convertedType :: !(Maybe ConvertedType) , logicalType :: !(Maybe LogicalType) - , encodeValue :: !(MemoryBuffer -> Int -> Int -> IO (Int, Bool)) - , finishValues :: !(MemoryBuffer -> Int -> IO Int) + , encodeValue :: + !(MemoryBuffer (PrimState m) -> Int -> Int -> m (Int, Bool)) + , finishValues :: !(MemoryBuffer (PrimState m) -> Int -> m Int) } -buildEncoder :: Column -> IO Encoder +buildEncoder :: (PrimBase m, MonadIO m) => Column -> m (Encoder m) buildEncoder col | hasElemType @Int32 col = pure $ @@ -71,7 +72,7 @@ buildEncoder col (INT32 enum) Nothing Nothing - (\buffer pos v -> writeWord32At buffer pos (fromIntegral v) >> pure (pos + 4)) + (\buffer pos v -> write32 buffer pos (fromIntegral v)) col | hasElemType @Int64 col = pure $ @@ -79,7 +80,7 @@ buildEncoder col (INT64 enum) Nothing Nothing - (\buffer pos v -> writeWord64At buffer pos (fromIntegral v) >> pure (pos + 8)) + (\buffer pos v -> write64 buffer pos (fromIntegral v)) col -- Ints in GHC can be 32 bit or 64 bit integers depending on the -- underlying computers architecture. So we'll do 64bit integers @@ -90,7 +91,7 @@ buildEncoder col (INT64 enum) Nothing Nothing - (\buffer pos v -> writeWord64At buffer pos (fromIntegral v) >> pure (pos + 8)) + (\buffer pos v -> write64 buffer pos (fromIntegral v)) col | hasElemType @Integer col = pure $ @@ -106,8 +107,7 @@ buildEncoder col (FLOAT enum) Nothing Nothing - ( \buffer pos v -> writeWord32At buffer pos (castFloatToWord32 v) >> pure (pos + 4) - ) + (\buffer pos v -> write32 buffer pos (castFloatToWord32 v)) col | hasElemType @Double col = pure $ @@ -115,85 +115,43 @@ buildEncoder col (DOUBLE enum) Nothing Nothing - ( \buffer pos v -> writeWord64At buffer pos (castDoubleToWord64 v) >> pure (pos + 8) - ) + (\buffer pos v -> write64 buffer pos (castDoubleToWord64 v)) col | hasElemType @Bool col = boolEncoder col | hasElemType @T.Text col = pure (textEncoder col) | hasElemType @UTCTime col = pure (timestampEncoder col) | otherwise = error ("writeParquet: unsupported column type " <> columnTypeString col) + where + write32 buffer pos value = do + writeWord32At buffer pos value + pure (pos + 4) + write64 buffer pos value = do + writeWord64At buffer pos value + pure (pos + 8) + +{-# SPECIALIZE buildEncoder :: Column -> IO (Encoder IO) #-} scalarEncoder :: - forall a. - (Columnable a) => + forall a m. + (Columnable a, Monad m) => ThriftType -> Maybe ConvertedType -> Maybe LogicalType -> - (MemoryBuffer -> Int -> a -> IO Int) -> + (MemoryBuffer (PrimState m) -> Int -> a -> m Int) -> Column -> - Encoder + Encoder m scalarEncoder tt conv logical writePrim col = Encoder tt conv logical (columnWriter @a col writePrim) (\_ pos -> pure pos) -{-# INLINEABLE scalarEncoder #-} -{-# SPECIALIZE scalarEncoder :: - ThriftType -> - Maybe ConvertedType -> - Maybe LogicalType -> - (MemoryBuffer -> Int -> Int32 -> IO Int) -> - Column -> - Encoder - #-} -{-# SPECIALIZE scalarEncoder :: - ThriftType -> - Maybe ConvertedType -> - Maybe LogicalType -> - (MemoryBuffer -> Int -> Int64 -> IO Int) -> - Column -> - Encoder - #-} -{-# SPECIALIZE scalarEncoder :: - ThriftType -> - Maybe ConvertedType -> - Maybe LogicalType -> - (MemoryBuffer -> Int -> Float -> IO Int) -> - Column -> - Encoder - #-} -{-# SPECIALIZE scalarEncoder :: - ThriftType -> - Maybe ConvertedType -> - Maybe LogicalType -> - (MemoryBuffer -> Int -> Double -> IO Int) -> - Column -> - Encoder - #-} -{-# SPECIALIZE scalarEncoder :: - ThriftType -> - Maybe ConvertedType -> - Maybe LogicalType -> - (MemoryBuffer -> Int -> Int -> IO Int) -> - Column -> - Encoder - #-} -{-# SPECIALIZE scalarEncoder :: - ThriftType -> - Maybe ConvertedType -> - Maybe LogicalType -> - (MemoryBuffer -> Int -> Integer -> IO Int) -> - Column -> - Encoder - #-} - columnWriter :: - forall a. - (Columnable a) => + forall a m. + (Columnable a, Monad m) => Column -> - (MemoryBuffer -> Int -> a -> IO Int) -> - MemoryBuffer -> + (MemoryBuffer (PrimState m) -> Int -> a -> m Int) -> + MemoryBuffer (PrimState m) -> Int -> Int -> - IO (Int, Bool) + m (Int, Bool) columnWriter col writePrim = case col of BoxedColumn bitmap (values :: VB.Vector b) -> case testEquality (typeRep @a) (typeRep @b) of @@ -213,114 +171,53 @@ columnWriter col writePrim = case col of mismatch = error ("writeParquet: incompatible column representation for " <> columnTypeString col) -{-# INLINEABLE columnWriter #-} -{-# SPECIALIZE columnWriter :: - Column -> - (MemoryBuffer -> Int -> Int32 -> IO Int) -> - MemoryBuffer -> - Int -> - Int -> - IO (Int, Bool) - #-} -{-# SPECIALIZE columnWriter :: - Column -> - (MemoryBuffer -> Int -> Int64 -> IO Int) -> - MemoryBuffer -> - Int -> - Int -> - IO (Int, Bool) - #-} -{-# SPECIALIZE columnWriter :: - Column -> - (MemoryBuffer -> Int -> Float -> IO Int) -> - MemoryBuffer -> - Int -> - Int -> - IO (Int, Bool) - #-} -{-# SPECIALIZE columnWriter :: - Column -> - (MemoryBuffer -> Int -> Double -> IO Int) -> - MemoryBuffer -> - Int -> - Int -> - IO (Int, Bool) - #-} -{-# SPECIALIZE columnWriter :: - Column -> - (MemoryBuffer -> Int -> Bool -> IO Int) -> - MemoryBuffer -> - Int -> - Int -> - IO (Int, Bool) - #-} -{-# SPECIALIZE columnWriter :: - Column -> - (MemoryBuffer -> Int -> UTCTime -> IO Int) -> - MemoryBuffer -> - Int -> - Int -> - IO (Int, Bool) - #-} -{-# SPECIALIZE columnWriter :: - Column -> - (MemoryBuffer -> Int -> Int -> IO Int) -> - MemoryBuffer -> - Int -> - Int -> - IO (Int, Bool) - #-} -{-# SPECIALIZE columnWriter :: - Column -> - (MemoryBuffer -> Int -> Integer -> IO Int) -> - MemoryBuffer -> - Int -> - Int -> - IO (Int, Bool) - #-} +{-# INLINE columnWriter #-} isPresent :: Maybe Bitmap -> Int -> Bool isPresent Nothing _ = True isPresent (Just bitmap) row = bitmapTestBit bitmap row {-# INLINE isPresent #-} -boolEncoder :: Column -> IO Encoder +boolEncoder :: (PrimMonad m) => Column -> m (Encoder m) boolEncoder col = do - bitsRef <- newIORef (0 :: Word8) - countRef <- newIORef (0 :: Int) + bitsRef <- newMutVar (0 :: Word8) + countRef <- newMutVar (0 :: Int) let addBit buffer pos value = do - bits <- readIORef bitsRef - count <- readIORef countRef + bits <- readMutVar bitsRef + count <- readMutVar countRef let bits' = if value then bits .|. ((1 :: Word8) `shiftL` count) else bits count' = count + 1 if count' == 8 then do - arr <- readIORef buffer.arrayRef + arr <- readMutVar buffer.arrayRef writeByteArray arr pos bits' - writeIORef bitsRef 0 - writeIORef countRef 0 + writeMutVar bitsRef 0 + writeMutVar countRef 0 pure (pos + 1) else do - writeIORef bitsRef bits' - writeIORef countRef count' + writeMutVar bitsRef bits' + writeMutVar countRef count' pure pos finish buffer pos = do - count <- readIORef countRef + count <- readMutVar countRef pos' <- if count > 0 then do - bits <- readIORef bitsRef - arr <- readIORef buffer.arrayRef + bits <- readMutVar bitsRef + arr <- readMutVar buffer.arrayRef writeByteArray arr pos bits pure (pos + 1) else pure pos - writeIORef bitsRef 0 - writeIORef countRef 0 + writeMutVar bitsRef 0 + writeMutVar countRef 0 pure pos' pure (Encoder (BOOLEAN enum) Nothing Nothing (columnWriter @Bool col addBit) finish) -textEncoder :: Column -> Encoder +textEncoder :: + (PrimBase m, MonadIO m) => + Column -> + Encoder m textEncoder col = Encoder (BYTE_ARRAY enum) @@ -351,18 +248,28 @@ textEncoder col = pure (pos', True) | otherwise = pure (pos, False) writeTextSlice buffer pos bytes offset count = do - writeIORef buffer.positionRef pos + writeMutVar buffer.positionRef pos _ <- ensureCapacity buffer (pos + 4 + count) writeWord32At buffer pos (fromIntegral count) - arr <- readIORef buffer.arrayRef + arr <- readMutVar buffer.arrayRef withMutableByteArrayContents arr $ \ptr -> - stToIO (TA.copyToPointer bytes offset (ptr `plusPtr` (pos + 4)) count) + liftIO $ + stToIO + ( TA.copyToPointer + bytes + offset + (ptr `plusPtr` (pos + 4)) + count + ) pure (pos + 4 + count) mismatch = error ("writeParquet: incompatible text representation for " <> columnTypeString col) -timestampEncoder :: Column -> Encoder +timestampEncoder :: + (PrimMonad m) => + Column -> + Encoder m timestampEncoder col = Encoder (INT64 enum) diff --git a/dataframe-parquet/src/DataFrame/IO/Parquet/Writer/Metadata.hs b/dataframe-parquet/src/DataFrame/IO/Parquet/Writer/Metadata.hs index 016b925c..6092b9bc 100644 --- a/dataframe-parquet/src/DataFrame/IO/Parquet/Writer/Metadata.hs +++ b/dataframe-parquet/src/DataFrame/IO/Parquet/Writer/Metadata.hs @@ -1,3 +1,4 @@ +{-# LANGUAGE FlexibleContexts #-} {-# LANGUAGE OverloadedStrings #-} module DataFrame.IO.Parquet.Writer.Metadata ( @@ -10,6 +11,8 @@ module DataFrame.IO.Parquet.Writer.Metadata ( magic, ) where +import Control.Monad.IO.Class (MonadIO) +import Control.Monad.Primitive (PrimBase) import qualified Data.ByteString as BS import Data.Int (Int64) import qualified Data.Text as T @@ -137,12 +140,13 @@ mkRowGroup chunks totalCompressed totalUncompressed rgRows = } writeFooter :: + (PrimBase m, MonadIO m) => WritableBinaryHandle -> [SchemaElement] -> Int -> [RowGroup] -> [(T.Text, T.Text)] -> - IO () + m () writeFooter output schemaElements numRows rowGroupMetadata keyValues = do let metadata = FileMetadata diff --git a/dataframe-parquet/src/DataFrame/IO/Utils/RandomAccess.hs b/dataframe-parquet/src/DataFrame/IO/Utils/RandomAccess.hs index 7acc7e8a..30d86bae 100644 --- a/dataframe-parquet/src/DataFrame/IO/Utils/RandomAccess.hs +++ b/dataframe-parquet/src/DataFrame/IO/Utils/RandomAccess.hs @@ -38,13 +38,12 @@ module DataFrame.IO.Utils.RandomAccess ( import Control.Exception (bracket, bracketOnError, finally) import Control.Monad (when) import Control.Monad.IO.Class (MonadIO (..)) -import Control.Monad.Primitive (RealWorld) +import Control.Monad.Primitive (PrimBase, PrimMonad, PrimState) import Control.Monad.ST (stToIO) import Data.Bits (shiftR) import qualified Data.ByteString as BS import Data.ByteString.Internal (ByteString (PS), create) import qualified Data.ByteString.Unsafe as BU -import Data.IORef (IORef, newIORef, readIORef, writeIORef) import Data.Int (Int64) import Data.Primitive.ByteArray ( MutableByteArray, @@ -54,6 +53,7 @@ import Data.Primitive.ByteArray ( withMutableByteArrayContents, writeByteArray, ) +import Data.Primitive.MutVar (MutVar, newMutVar, readMutVar, writeMutVar) import qualified Data.Text.Array as TA import qualified Data.Vector.Storable as VS import Data.Word (Word32, Word64, Word8) @@ -163,7 +163,10 @@ openWritableBinaryFile filepath = do hSetBuffering h NoBuffering pure . WritableBinaryHandle $ h -atomicallyWriteFile :: FilePath -> (FilePath -> IO a) -> IO a +atomicallyWriteFile :: + FilePath -> + (FilePath -> IO a) -> + IO a atomicallyWriteFile path action = bracketOnError openAction @@ -188,23 +191,27 @@ atomicallyWriteFile path action = pure tmpFile ) -withWritableBinaryFile :: FilePath -> (WritableBinaryHandle -> IO a) -> IO a +withWritableBinaryFile :: + FilePath -> + (WritableBinaryHandle -> IO a) -> + IO a withWritableBinaryFile filepath = bracket (openWritableBinaryFile filepath) (hClose . unHandle) -data MemoryBuffer = MemoryBuffer - { arrayRef :: !(IORef (MutableByteArray RealWorld)) - , positionRef :: !(IORef Int) +data MemoryBuffer s = MemoryBuffer + { arrayRef :: !(MutVar s (MutableByteArray s)) + , positionRef :: !(MutVar s Int) } -mallocBuffer :: Int -> IO MemoryBuffer +mallocBuffer :: + (PrimMonad m, MonadIO m) => Int -> m (MemoryBuffer (PrimState m)) mallocBuffer capacity - | capacity < 0 = ioError $ userError "mallocBuffer: negative capacity" + | capacity < 0 = liftIO $ ioError $ userError "mallocBuffer: negative capacity" | otherwise = do array <- newPinnedByteArray capacity - MemoryBuffer <$> newIORef array <*> newIORef 0 + MemoryBuffer <$> newMutVar array <*> newMutVar 0 -- We're using pinned ByteArrays so we must -- not use the grow function brovided by Data.Primitive @@ -221,53 +228,65 @@ mallocBuffer capacity -- is just a matter of adding a new buffer to the array (which we can -- pre-allocate to three elements to begin with and grow it only on the -- off chance that a buffer required more than three grows). -ensureCapacity :: MemoryBuffer -> Int -> IO (MutableByteArray RealWorld) +ensureCapacity :: + (PrimMonad m) => + MemoryBuffer (PrimState m) -> Int -> m (MutableByteArray (PrimState m)) ensureCapacity buffer needed = do - array <- readIORef buffer.arrayRef + array <- readMutVar buffer.arrayRef maxSize <- getSizeofMutableByteArray array if needed <= maxSize then pure array else do - position <- readIORef buffer.positionRef + position <- readMutVar buffer.positionRef grown <- newPinnedByteArray (needed + (needed `div` 2)) copyMutableByteArray grown 0 array 0 position - writeIORef buffer.arrayRef grown + writeMutVar buffer.arrayRef grown pure grown {-# INLINE ensureCapacity #-} -writeWord8 :: MemoryBuffer -> Word8 -> IO () +writeWord8 :: (PrimMonad m) => MemoryBuffer (PrimState m) -> Word8 -> m () writeWord8 buffer b = do - position <- readIORef buffer.positionRef + position <- readMutVar buffer.positionRef array <- ensureCapacity buffer (position + 1) writeByteArray array position b - writeIORef buffer.positionRef (position + 1) + writeMutVar buffer.positionRef (position + 1) {-# INLINE writeWord8 #-} -writeByteString :: MemoryBuffer -> ByteString -> IO () -writeByteString buffer bs = - BU.unsafeUseAsCStringLen bs $ \(source, len) -> do - position <- readIORef buffer.positionRef - array <- ensureCapacity buffer (position + len) - withMutableByteArrayContents array $ \dst -> - copyBytes (dst `plusPtr` position) (castPtr source) len - writeIORef buffer.positionRef (position + len) +writeByteString :: + (PrimBase m, MonadIO m) => + MemoryBuffer (PrimState m) -> + ByteString -> + m () +writeByteString buffer bs = do + position <- readMutVar buffer.positionRef + let len = BS.length bs + array <- ensureCapacity buffer (position + len) + withMutableByteArrayContents array $ \dst -> + liftIO $ + BU.unsafeUseAsCStringLen bs $ \(source, _) -> do + copyBytes + (dst `plusPtr` position) + (castPtr source) + len + writeMutVar buffer.positionRef (position + len) {-# INLINE writeByteString #-} -writeWord32LE :: MemoryBuffer -> Word32 -> IO () +writeWord32LE :: (PrimMonad m) => MemoryBuffer (PrimState m) -> Word32 -> m () writeWord32LE buffer w = do - position <- readIORef buffer.positionRef + position <- readMutVar buffer.positionRef writeWord32At buffer position w - writeIORef buffer.positionRef (position + 4) + writeMutVar buffer.positionRef (position + 4) {-# INLINE writeWord32LE #-} -writeWord64LE :: MemoryBuffer -> Word64 -> IO () +writeWord64LE :: (PrimMonad m) => MemoryBuffer (PrimState m) -> Word64 -> m () writeWord64LE buffer w = do - position <- readIORef buffer.positionRef + position <- readMutVar buffer.positionRef writeWord64At buffer position w - writeIORef buffer.positionRef (position + 8) + writeMutVar buffer.positionRef (position + 8) {-# INLINE writeWord64LE #-} -writeWord32At :: MemoryBuffer -> Int -> Word32 -> IO () +writeWord32At :: + (PrimMonad m) => MemoryBuffer (PrimState m) -> Int -> Word32 -> m () writeWord32At buffer position w = do array <- ensureCapacity buffer (position + 4) writeByteArray array position (fromIntegral w :: Word8) @@ -276,7 +295,8 @@ writeWord32At buffer position w = do writeByteArray array (position + 3) (fromIntegral (w `shiftR` 24) :: Word8) {-# INLINE writeWord32At #-} -writeWord64At :: MemoryBuffer -> Int -> Word64 -> IO () +writeWord64At :: + (PrimMonad m) => MemoryBuffer (PrimState m) -> Int -> Word64 -> m () writeWord64At buffer position w = do array <- ensureCapacity buffer (position + 8) writeByteArray array position (fromIntegral w :: Word8) @@ -289,14 +309,17 @@ writeWord64At buffer position w = do writeByteArray array (position + 7) (fromIntegral (w `shiftR` 56) :: Word8) {-# INLINE writeWord64At #-} -writeInteger64 :: MemoryBuffer -> Integer -> IO () +writeInteger64 :: + (PrimMonad m, MonadIO m) => MemoryBuffer (PrimState m) -> Integer -> m () writeInteger64 buffer value = do - position <- readIORef buffer.positionRef + position <- readMutVar buffer.positionRef newPosition <- writeInteger64At buffer position value - writeIORef buffer.positionRef newPosition + writeMutVar buffer.positionRef newPosition {-# INLINE writeInteger64 #-} -writeInteger64At :: MemoryBuffer -> Int -> Integer -> IO Int +writeInteger64At :: + (PrimMonad m, MonadIO m) => + MemoryBuffer (PrimState m) -> Int -> Integer -> m Int writeInteger64At buffer position value | value < toInteger (minBound :: Int64) = outOfRange | value > toInteger (maxBound :: Int64) = outOfRange @@ -305,24 +328,27 @@ writeInteger64At buffer position value pure (position + 8) where outOfRange = - ioError (userError "writeParquet: Integer value is outside the INT64 range") + liftIO + (ioError (userError "writeParquet: Integer value is outside the INT64 range")) {-# INLINE writeInteger64At #-} -writeFloatLE :: MemoryBuffer -> Float -> IO () +writeFloatLE :: (PrimMonad m) => MemoryBuffer (PrimState m) -> Float -> m () writeFloatLE buffer = writeWord32LE buffer . castFloatToWord32 {-# INLINE writeFloatLE #-} -writeDoubleLE :: MemoryBuffer -> Double -> IO () +writeDoubleLE :: (PrimMonad m) => MemoryBuffer (PrimState m) -> Double -> m () writeDoubleLE buffer = writeWord64LE buffer . castDoubleToWord64 {-# INLINE writeDoubleLE #-} -flushBufferToBuffer :: MemoryBuffer -> MemoryBuffer -> IO () +flushBufferToBuffer :: + (PrimMonad m) => + MemoryBuffer (PrimState m) -> MemoryBuffer (PrimState m) -> m () flushBufferToBuffer source destination | source.arrayRef == destination.arrayRef = pure () | otherwise = do - sourceArray <- readIORef source.arrayRef - sourcePosition <- readIORef source.positionRef - destinationPosition <- readIORef destination.positionRef + sourceArray <- readMutVar source.arrayRef + sourcePosition <- readMutVar source.positionRef + destinationPosition <- readMutVar destination.positionRef destinationArray <- ensureCapacity destination (destinationPosition + sourcePosition) copyMutableByteArray @@ -331,24 +357,28 @@ flushBufferToBuffer source destination sourceArray 0 sourcePosition - writeIORef destination.positionRef (destinationPosition + sourcePosition) - writeIORef source.positionRef 0 + writeMutVar destination.positionRef (destinationPosition + sourcePosition) + writeMutVar source.positionRef 0 {-# INLINE flushBufferToBuffer #-} -bufferToByteString :: MemoryBuffer -> IO ByteString +bufferToByteString :: + (PrimBase m, MonadIO m) => + MemoryBuffer (PrimState m) -> + m ByteString bufferToByteString buffer = do - array <- readIORef buffer.arrayRef - position <- readIORef buffer.positionRef - create position $ \dst -> - withMutableByteArrayContents array $ \src -> - copyBytes dst (castPtr src) position - -bufferResidency :: MemoryBuffer -> IO Int -bufferResidency buffer = readIORef buffer.positionRef + array <- readMutVar buffer.arrayRef + position <- readMutVar buffer.positionRef + withMutableByteArrayContents array $ \src -> + liftIO $ + create position $ \dst -> + copyBytes dst (castPtr src) position + +bufferResidency :: (PrimMonad m) => MemoryBuffer (PrimState m) -> m Int +bufferResidency buffer = readMutVar buffer.positionRef {-# INLINE bufferResidency #-} -resetPosition :: MemoryBuffer -> IO () -resetPosition buffer = writeIORef buffer.positionRef 0 +resetPosition :: (PrimMonad m) => MemoryBuffer (PrimState m) -> m () +resetPosition buffer = writeMutVar buffer.positionRef 0 {-# INLINE resetPosition #-} -- I tested write speeds by doing (on Apple Silicon) @@ -372,11 +402,13 @@ resetPosition buffer = writeIORef buffer.positionRef 0 -- So when writing to a file to minimize syscall overhead while -- trying not to create dirty pages in the kernel page cache, we'll -- be flushing in 256 KiB chunks. -flushBufferToFile :: WritableBinaryHandle -> MemoryBuffer -> IO () +flushBufferToFile :: + (PrimBase m, MonadIO m) => + WritableBinaryHandle -> MemoryBuffer (PrimState m) -> m () flushBufferToFile (WritableBinaryHandle h) buffer = do - array <- readIORef buffer.arrayRef - position <- readIORef buffer.positionRef - withMutableByteArrayContents array $ \ptr -> do + array <- readMutVar buffer.arrayRef + position <- readMutVar buffer.positionRef + withMutableByteArrayContents array $ \ptr -> liftIO $ do let chunkSize = 262144 go offset | offset >= position = pure () @@ -385,7 +417,7 @@ flushBufferToFile (WritableBinaryHandle h) buffer = do hPutBuf h (ptr `plusPtr` offset) n go (offset + n) go 0 - writeIORef buffer.positionRef 0 + writeMutVar buffer.positionRef 0 writeByteStringToFile :: WritableBinaryHandle -> ByteString -> IO () writeByteStringToFile (WritableBinaryHandle h) bs = @@ -399,13 +431,23 @@ writeByteStringToFile (WritableBinaryHandle h) bs = go (offset + n) go 0 -appendTextArraySlice :: MemoryBuffer -> TA.Array -> Int -> Int -> IO () +appendTextArraySlice :: + (PrimBase m, MonadIO m) => + MemoryBuffer (PrimState m) -> TA.Array -> Int -> Int -> m () appendTextArraySlice buffer source offset count - | count < 0 = ioError $ userError "appendTextArraySlice: negative length" + | count < 0 = + liftIO $ ioError $ userError "appendTextArraySlice: negative length" | otherwise = do - position <- readIORef buffer.positionRef + position <- readMutVar buffer.positionRef array <- ensureCapacity buffer (position + count) withMutableByteArrayContents array $ \destination -> - stToIO (TA.copyToPointer source offset (destination `plusPtr` position) count) - writeIORef buffer.positionRef (position + count) + liftIO $ + stToIO + ( TA.copyToPointer + source + offset + (destination `plusPtr` position) + count + ) + writeMutVar buffer.positionRef (position + count) {-# INLINE appendTextArraySlice #-}