{-# LANGUAGE BangPatterns #-}
{-# LANGUAGE RankNTypes #-}
{-# LANGUAGE ScopedTypeVariables #-}
module DataFrame.Internal.HashTable (
HashTable (..),
newHashTable,
htInsert,
nextPow2Above,
) where
import Control.Monad.Primitive (PrimMonad, PrimState)
import Data.Bits ((.&.))
import qualified Data.Vector.Unboxed.Mutable as VUM
data HashTable s = HashTable
{ forall s. HashTable s -> MVector s Int
htHash :: !(VUM.MVector s Int)
, forall s. HashTable s -> MVector s Int
htGroup :: !(VUM.MVector s Int)
, forall s. HashTable s -> MVector s Int
htRep :: !(VUM.MVector s Int)
, forall s. HashTable s -> Int
htMask :: !Int
}
nextPow2Above :: Int -> Int
nextPow2Above :: Int -> Int
nextPow2Above Int
n = Int -> Int
go Int
2
where
go :: Int -> Int
go !Int
p
| Int
p Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
> Int
n = Int
p
| Bool
otherwise = Int -> Int
go (Int
p Int -> Int -> Int
forall a. Num a => a -> a -> a
* Int
2)
{-# INLINE nextPow2Above #-}
newHashTable :: (PrimMonad m) => Int -> m (HashTable (PrimState m))
newHashTable :: forall (m :: * -> *).
PrimMonad m =>
Int -> m (HashTable (PrimState m))
newHashTable Int
n = do
let !cap :: Int
cap = Int -> Int
nextPow2Above (Int
2 Int -> Int -> Int
forall a. Num a => a -> a -> a
* Int -> Int -> Int
forall a. Ord a => a -> a -> a
max Int
1 Int
n)
MVector (PrimState m) Int
h <- Int -> m (MVector (PrimState m) Int)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> m (MVector (PrimState m) a)
VUM.unsafeNew Int
cap
MVector (PrimState m) Int
g <- Int -> Int -> m (MVector (PrimState m) Int)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> a -> m (MVector (PrimState m) a)
VUM.replicate Int
cap (-Int
1)
MVector (PrimState m) Int
r <- Int -> m (MVector (PrimState m) Int)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> m (MVector (PrimState m) a)
VUM.unsafeNew Int
cap
HashTable (PrimState m) -> m (HashTable (PrimState m))
forall a. a -> m a
forall (f :: * -> *) a. Applicative f => a -> f a
pure (MVector (PrimState m) Int
-> MVector (PrimState m) Int
-> MVector (PrimState m) Int
-> Int
-> HashTable (PrimState m)
forall s.
MVector s Int
-> MVector s Int -> MVector s Int -> Int -> HashTable s
HashTable MVector (PrimState m) Int
h MVector (PrimState m) Int
g MVector (PrimState m) Int
r (Int
cap Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1))
{-# INLINE newHashTable #-}
htInsert ::
(PrimMonad m) =>
HashTable (PrimState m) ->
(Int -> Int -> Bool) ->
Int ->
Int ->
Int ->
m (Int, Bool)
htInsert :: forall (m :: * -> *).
PrimMonad m =>
HashTable (PrimState m)
-> (Int -> Int -> Bool) -> Int -> Int -> Int -> m (Int, Bool)
htInsert HashTable (PrimState m)
ht Int -> Int -> Bool
eqRow Int
nextGroup Int
row Int
hash = Int -> m (Int, Bool)
forall {m :: * -> *}.
(PrimState m ~ PrimState m, PrimMonad m) =>
Int -> m (Int, Bool)
go (Int
hash Int -> Int -> Int
forall a. Bits a => a -> a -> a
.&. Int
mask)
where
!mask :: Int
mask = HashTable (PrimState m) -> Int
forall s. HashTable s -> Int
htMask HashTable (PrimState m)
ht
!hs :: MVector (PrimState m) Int
hs = HashTable (PrimState m) -> MVector (PrimState m) Int
forall s. HashTable s -> MVector s Int
htHash HashTable (PrimState m)
ht
!gs :: MVector (PrimState m) Int
gs = HashTable (PrimState m) -> MVector (PrimState m) Int
forall s. HashTable s -> MVector s Int
htGroup HashTable (PrimState m)
ht
!rs :: MVector (PrimState m) Int
rs = HashTable (PrimState m) -> MVector (PrimState m) Int
forall s. HashTable s -> MVector s Int
htRep HashTable (PrimState m)
ht
go :: Int -> m (Int, Bool)
go !Int
slot = do
Int
g <- MVector (PrimState m) Int -> Int -> m Int
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.unsafeRead MVector (PrimState m) Int
MVector (PrimState m) Int
gs Int
slot
if Int
g Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
< Int
0
then do
MVector (PrimState m) Int -> Int -> Int -> m ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite MVector (PrimState m) Int
MVector (PrimState m) Int
hs Int
slot Int
hash
MVector (PrimState m) Int -> Int -> Int -> m ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite MVector (PrimState m) Int
MVector (PrimState m) Int
gs Int
slot Int
nextGroup
MVector (PrimState m) Int -> Int -> Int -> m ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite MVector (PrimState m) Int
MVector (PrimState m) Int
rs Int
slot Int
row
(Int, Bool) -> m (Int, Bool)
forall a. a -> m a
forall (f :: * -> *) a. Applicative f => a -> f a
pure (Int
nextGroup, Bool
True)
else do
Int
h <- MVector (PrimState m) Int -> Int -> m Int
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.unsafeRead MVector (PrimState m) Int
MVector (PrimState m) Int
hs Int
slot
if Int
h Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
== Int
hash
then do
Int
rep <- MVector (PrimState m) Int -> Int -> m Int
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.unsafeRead MVector (PrimState m) Int
MVector (PrimState m) Int
rs Int
slot
if Int -> Int -> Bool
eqRow Int
rep Int
row
then (Int, Bool) -> m (Int, Bool)
forall a. a -> m a
forall (f :: * -> *) a. Applicative f => a -> f a
pure (Int
g, Bool
False)
else Int -> m (Int, Bool)
go ((Int
slot Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1) Int -> Int -> Int
forall a. Bits a => a -> a -> a
.&. Int
mask)
else Int -> m (Int, Bool)
go ((Int
slot Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1) Int -> Int -> Int
forall a. Bits a => a -> a -> a
.&. Int
mask)
{-# INLINE htInsert #-}