{-# LANGUAGE BangPatterns #-}
{-# LANGUAGE RankNTypes #-}
{-# LANGUAGE ScopedTypeVariables #-}

{- | A flat, unboxed, open-addressing (linear-probe) hash table mapping a row's
key-hash to a dense group id, re-verifying the real key on every hash hit to
reject collisions. Runs in any 'PrimMonad' ('ST' for grouping, 'IO' per worker).
-}
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

{- | An open-addressing linear-probe table. @htMask@ is @capacity - 1@ (capacity
is a power of two) and maps a hash to its home slot.
-}
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
    }

{- | Smallest power of two strictly greater than @n@, at least 2. Sizes the
table so the load factor stays below ~0.5 even when every row is a distinct
group.
-}
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 #-}

{- | Allocate an empty table able to hold up to @n@ distinct groups while
keeping the load factor under ~0.5 (capacity @= nextPow2Above (2*n)@). All
group slots start empty (@-1@).
-}
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 #-}

{- | Look up @row@ (with precomputed @hash@) and return its dense group id: an
empty slot starts a new group via @nextGroup@, a stored-hash match is re-verified
with @eqRow@ before reuse. The 'Bool' is 'True' when a new group was created.
-}
htInsert ::
    (PrimMonad m) =>
    HashTable (PrimState m) ->
    -- | @eqRow a b@: do rows @a@ and @b@ have equal key columns?
    (Int -> Int -> Bool) ->
    -- | Next dense group id to assign if this row starts a new group.
    Int ->
    -- | Row index being inserted.
    Int ->
    -- | Precomputed hash of the row's key.
    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 #-}