-- SPDX-FileCopyrightText: 2026 Alexandra de Wit
--
-- SPDX-License-Identifier: MIT

{- | A growable byte buffer that one read writes front to back, rewinds in place, and copies each
finished value out of. It writes the bytes aeson writes for strings and integers without building
them first.
-}
module Ecluse.Core.Registry.Json.Scratch (
    Scratch,
    newScratch,
    scratchCursor,
    rewindTo,
    scratchBuffer,
    reserve,
    putByte,
    putVarint,
    putEncodedText,
    putPlainBytes,
    putRawBytes,
    putDecimal,
    putAt,
    copyOut,
    decimalLength,
) where

import Control.Monad.ST (ST)
import Data.ByteString qualified as BS
import Data.ByteString.Unsafe qualified as BSU
import Data.Primitive.ByteArray (ByteArray, MutableByteArray, copyMutableByteArray, getSizeofMutableByteArray, newByteArray, unsafeFreezeByteArray, writeByteArray)
import Data.Primitive.MutVar (MutVar, newMutVar, readMutVar, writeMutVar)
import Data.Primitive.PrimVar (PrimVar, newPrimVar, readPrimVar, writePrimVar)

import Ecluse.Core.Registry.Json.Packed (encodedLength, quote, varintSize, writeEncoded, writeVarint)

-- | The buffer and the offset of its next byte.
data Scratch st = Scratch !(MutVar st (MutableByteArray st)) !(PrimVar st Int)

-- | An empty buffer with room for the given bytes.
newScratch :: Int -> ST st (Scratch st)
newScratch :: forall st. Int -> ST st (Scratch st)
newScratch Int
size = MutVar st (MutableByteArray st) -> PrimVar st Int -> Scratch st
forall st.
MutVar st (MutableByteArray st) -> PrimVar st Int -> Scratch st
Scratch (MutVar st (MutableByteArray st) -> PrimVar st Int -> Scratch st)
-> ST st (MutVar st (MutableByteArray st))
-> ST st (PrimVar st Int -> Scratch st)
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
<$> (Int -> ST st (MutableByteArray (PrimState (ST st)))
forall (m :: * -> *).
PrimMonad m =>
Int -> m (MutableByteArray (PrimState m))
newByteArray (Int -> Int -> Int
forall a. Ord a => a -> a -> a
max Int
64 Int
size) ST st (MutableByteArray st)
-> (MutableByteArray st -> ST st (MutVar st (MutableByteArray st)))
-> ST st (MutVar st (MutableByteArray st))
forall a b. ST st a -> (a -> ST st b) -> ST st b
forall (m :: * -> *) a b. Monad m => m a -> (a -> m b) -> m b
>>= MutableByteArray st -> ST st (MutVar st (MutableByteArray st))
MutableByteArray st
-> ST st (MutVar (PrimState (ST st)) (MutableByteArray st))
forall (m :: * -> *) a.
PrimMonad m =>
a -> m (MutVar (PrimState m) a)
newMutVar) ST st (PrimVar st Int -> Scratch st)
-> ST st (PrimVar st Int) -> ST st (Scratch st)
forall a b. ST st (a -> b) -> ST st a -> ST st b
forall (f :: * -> *) a b. Applicative f => f (a -> b) -> f a -> f b
<*> Int -> ST st (PrimVar (PrimState (ST st)) Int)
forall (m :: * -> *) a.
(PrimMonad m, Prim a) =>
a -> m (PrimVar (PrimState m) a)
newPrimVar Int
0

-- | Where the next byte goes.
scratchCursor :: Scratch st -> ST st Int
scratchCursor :: forall st. Scratch st -> ST st Int
scratchCursor (Scratch MutVar st (MutableByteArray st)
_ PrimVar st Int
cursor) = PrimVar (PrimState (ST st)) Int -> ST st Int
forall (m :: * -> *) a.
(PrimMonad m, Prim a) =>
PrimVar (PrimState m) a -> m a
readPrimVar PrimVar st Int
PrimVar (PrimState (ST st)) Int
cursor
{-# INLINE scratchCursor #-}

-- | Forget every byte from the offset on.
rewindTo :: Scratch st -> Int -> ST st ()
rewindTo :: forall st. Scratch st -> Int -> ST st ()
rewindTo (Scratch MutVar st (MutableByteArray st)
_ PrimVar st Int
cursor) = PrimVar (PrimState (ST st)) Int -> Int -> ST st ()
forall (m :: * -> *) a.
(PrimMonad m, Prim a) =>
PrimVar (PrimState m) a -> a -> m ()
writePrimVar PrimVar st Int
PrimVar (PrimState (ST st)) Int
cursor
{-# INLINE rewindTo #-}

-- | The buffer as it stands. Any later write that needs room may replace it.
scratchBuffer :: Scratch st -> ST st (MutableByteArray st)
scratchBuffer :: forall st. Scratch st -> ST st (MutableByteArray st)
scratchBuffer (Scratch MutVar st (MutableByteArray st)
buffer PrimVar st Int
_) = MutVar (PrimState (ST st)) (MutableByteArray st)
-> ST st (MutableByteArray st)
forall (m :: * -> *) a.
PrimMonad m =>
MutVar (PrimState m) a -> m a
readMutVar MutVar st (MutableByteArray st)
MutVar (PrimState (ST st)) (MutableByteArray st)
buffer
{-# INLINE scratchBuffer #-}

-- | The buffer with room for the given bytes past the cursor, and the cursor.
reserve :: Scratch st -> Int -> (MutableByteArray st -> Int -> ST st a) -> ST st a
reserve :: forall st a.
Scratch st
-> Int -> (MutableByteArray st -> Int -> ST st a) -> ST st a
reserve (Scratch MutVar st (MutableByteArray st)
ref PrimVar st Int
cursor) Int
count MutableByteArray st -> Int -> ST st a
use = do
    buffer <- MutVar (PrimState (ST st)) (MutableByteArray st)
-> ST st (MutableByteArray st)
forall (m :: * -> *) a.
PrimMonad m =>
MutVar (PrimState m) a -> m a
readMutVar MutVar st (MutableByteArray st)
MutVar (PrimState (ST st)) (MutableByteArray st)
ref
    at <- readPrimVar cursor
    size <- getSizeofMutableByteArray buffer
    if at + count <= size
        then use buffer at
        else do
            grown <- newByteArray (max (at + count) (2 * size))
            copyMutableByteArray grown 0 buffer 0 at
            writeMutVar ref grown
            use grown at
{-# INLINE reserve #-}

-- | Write bytes at the cursor with the given writer, which returns the offset after them.
putAt :: Scratch st -> Int -> (MutableByteArray st -> Int -> ST st Int) -> ST st ()
putAt :: forall st.
Scratch st
-> Int -> (MutableByteArray st -> Int -> ST st Int) -> ST st ()
putAt scratch :: Scratch st
scratch@(Scratch MutVar st (MutableByteArray st)
_ PrimVar st Int
cursor) Int
count MutableByteArray st -> Int -> ST st Int
write = Scratch st
-> Int -> (MutableByteArray st -> Int -> ST st ()) -> ST st ()
forall st a.
Scratch st
-> Int -> (MutableByteArray st -> Int -> ST st a) -> ST st a
reserve Scratch st
scratch Int
count (\MutableByteArray st
buffer Int
at -> MutableByteArray st -> Int -> ST st Int
write MutableByteArray st
buffer Int
at ST st Int -> (Int -> ST st ()) -> ST st ()
forall a b. ST st a -> (a -> ST st b) -> ST st b
forall (m :: * -> *) a b. Monad m => m a -> (a -> m b) -> m b
>>= PrimVar (PrimState (ST st)) Int -> Int -> ST st ()
forall (m :: * -> *) a.
(PrimMonad m, Prim a) =>
PrimVar (PrimState m) a -> a -> m ()
writePrimVar PrimVar st Int
PrimVar (PrimState (ST st)) Int
cursor)
{-# INLINE putAt #-}

-- | Write one byte.
putByte :: Scratch st -> Word8 -> ST st ()
putByte :: forall st. Scratch st -> Word8 -> ST st ()
putByte Scratch st
scratch Word8
byte = Scratch st
-> Int -> (MutableByteArray st -> Int -> ST st Int) -> ST st ()
forall st.
Scratch st
-> Int -> (MutableByteArray st -> Int -> ST st Int) -> ST st ()
putAt Scratch st
scratch Int
1 (\MutableByteArray st
buffer Int
at -> MutableByteArray (PrimState (ST st)) -> Int -> Word8 -> ST st ()
forall a (m :: * -> *).
(Prim a, PrimMonad m) =>
MutableByteArray (PrimState m) -> Int -> a -> m ()
writeByteArray MutableByteArray st
MutableByteArray (PrimState (ST st))
buffer Int
at Word8
byte ST st () -> ST st Int -> ST st Int
forall a b. ST st a -> ST st b -> ST st b
forall (m :: * -> *) a b. Monad m => m a -> m b -> m b
>> Int -> ST st Int
forall a. a -> ST st a
forall (f :: * -> *) a. Applicative f => a -> f a
pure (Int
at Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1))
{-# INLINE putByte #-}

-- | Write a non-negative integer seven bits at a time, low bits first.
putVarint :: Scratch st -> Int -> ST st ()
putVarint :: forall st. Scratch st -> Int -> ST st ()
putVarint Scratch st
scratch Int
n = Scratch st
-> Int -> (MutableByteArray st -> Int -> ST st Int) -> ST st ()
forall st.
Scratch st
-> Int -> (MutableByteArray st -> Int -> ST st Int) -> ST st ()
putAt Scratch st
scratch (Int -> Int
varintSize Int
n) (Int -> MutableByteArray st -> Int -> ST st Int
forall st. Int -> MutableByteArray st -> Int -> ST st Int
writeVarint Int
n)
{-# INLINE putVarint #-}

{- | Write the varint a header makes of a string's encoded length, then the bytes aeson writes for the
string, quotes included.
-}
putEncodedText :: Scratch st -> (Int -> Int) -> Text -> ST st ()
putEncodedText :: forall st. Scratch st -> (Int -> Int) -> Text -> ST st ()
putEncodedText Scratch st
scratch Int -> Int
header Text
text = Scratch st -> Int -> ST st ()
forall st. Scratch st -> Int -> ST st ()
putVarint Scratch st
scratch (Int -> Int
header Int
len) ST st () -> ST st () -> ST st ()
forall a b. ST st a -> ST st b -> ST st b
forall (m :: * -> *) a b. Monad m => m a -> m b -> m b
>> Scratch st
-> Int -> (MutableByteArray st -> Int -> ST st Int) -> ST st ()
forall st.
Scratch st
-> Int -> (MutableByteArray st -> Int -> ST st Int) -> ST st ()
putAt Scratch st
scratch Int
len (Text -> MutableByteArray st -> Int -> ST st Int
forall st. Text -> MutableByteArray st -> Int -> ST st Int
writeEncoded Text
text)
  where
    len :: Int
len = Text -> Int
encodedLength Text
text

-- | Write a string whose bytes need no escape, between quotes.
putPlainBytes :: Scratch st -> ByteString -> ST st ()
putPlainBytes :: forall st. Scratch st -> ByteString -> ST st ()
putPlainBytes Scratch st
scratch ByteString
bytes = Scratch st
-> Int -> (MutableByteArray st -> Int -> ST st Int) -> ST st ()
forall st.
Scratch st
-> Int -> (MutableByteArray st -> Int -> ST st Int) -> ST st ()
putAt Scratch st
scratch (ByteString -> Int
BS.length ByteString
bytes Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
2) ((MutableByteArray st -> Int -> ST st Int) -> ST st ())
-> (MutableByteArray st -> Int -> ST st Int) -> ST st ()
forall a b. (a -> b) -> a -> b
$ \MutableByteArray st
buffer Int
at -> do
    MutableByteArray (PrimState (ST st)) -> Int -> Word8 -> ST st ()
forall a (m :: * -> *).
(Prim a, PrimMonad m) =>
MutableByteArray (PrimState m) -> Int -> a -> m ()
writeByteArray MutableByteArray st
MutableByteArray (PrimState (ST st))
buffer Int
at Word8
quote
    end <- ByteString -> MutableByteArray st -> Int -> ST st Int
forall st. ByteString -> MutableByteArray st -> Int -> ST st Int
copyBytes ByteString
bytes MutableByteArray st
buffer (Int
at Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
    writeByteArray buffer end quote
    pure (end + 1)

-- | Write bytes as they are.
putRawBytes :: Scratch st -> ByteString -> ST st ()
putRawBytes :: forall st. Scratch st -> ByteString -> ST st ()
putRawBytes Scratch st
scratch ByteString
bytes = Scratch st
-> Int -> (MutableByteArray st -> Int -> ST st Int) -> ST st ()
forall st.
Scratch st
-> Int -> (MutableByteArray st -> Int -> ST st Int) -> ST st ()
putAt Scratch st
scratch (ByteString -> Int
BS.length ByteString
bytes) (ByteString -> MutableByteArray st -> Int -> ST st Int
forall st. ByteString -> MutableByteArray st -> Int -> ST st Int
copyBytes ByteString
bytes)

copyBytes :: ByteString -> MutableByteArray st -> Int -> ST st Int
copyBytes :: forall st. ByteString -> MutableByteArray st -> Int -> ST st Int
copyBytes ByteString
bytes MutableByteArray st
buffer = Int -> Int -> ST st Int
forall {f :: * -> *}.
(PrimState f ~ st, PrimMonad f) =>
Int -> Int -> f Int
go Int
0
  where
    len :: Int
len = ByteString -> Int
BS.length ByteString
bytes
    go :: Int -> Int -> f Int
go !Int
index !Int
at
        | Int
index Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
len = Int -> f Int
forall a. a -> f a
forall (f :: * -> *) a. Applicative f => a -> f a
pure Int
at
        | Bool
otherwise = MutableByteArray (PrimState f) -> Int -> Word8 -> f ()
forall a (m :: * -> *).
(Prim a, PrimMonad m) =>
MutableByteArray (PrimState m) -> Int -> a -> m ()
writeByteArray MutableByteArray st
MutableByteArray (PrimState f)
buffer Int
at (ByteString -> Int -> Word8
BSU.unsafeIndex ByteString
bytes Int
index) f () -> f Int -> f Int
forall a b. f a -> f b -> f b
forall (m :: * -> *) a b. Monad m => m a -> m b -> m b
>> Int -> Int -> f Int
go (Int
index Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1) (Int
at Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)

-- | Write an integer in decimal, as aeson writes one.
putDecimal :: Scratch st -> Int -> ST st ()
putDecimal :: forall st. Scratch st -> Int -> ST st ()
putDecimal Scratch st
scratch Int
n = Scratch st
-> Int -> (MutableByteArray st -> Int -> ST st Int) -> ST st ()
forall st.
Scratch st
-> Int -> (MutableByteArray st -> Int -> ST st Int) -> ST st ()
putAt Scratch st
scratch Int
len ((MutableByteArray st -> Int -> ST st Int) -> ST st ())
-> (MutableByteArray st -> Int -> ST st Int) -> ST st ()
forall a b. (a -> b) -> a -> b
$ \MutableByteArray st
buffer Int
at -> do
    Bool -> ST st () -> ST st ()
forall (f :: * -> *). Applicative f => Bool -> f () -> f ()
when (Int
n Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
< Int
0) (MutableByteArray (PrimState (ST st)) -> Int -> Word8 -> ST st ()
forall a (m :: * -> *).
(Prim a, PrimMonad m) =>
MutableByteArray (PrimState m) -> Int -> a -> m ()
writeByteArray MutableByteArray st
MutableByteArray (PrimState (ST st))
buffer Int
at (Word8
0x2d :: Word8))
    MutableByteArray st -> Int -> Word -> ST st ()
forall st. MutableByteArray st -> Int -> Word -> ST st ()
writeDigits MutableByteArray st
buffer (Int
at Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
len Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1) (Int -> Word
magnitude Int
n)
    Int -> ST st Int
forall a. a -> ST st a
forall (f :: * -> *) a. Applicative f => a -> f a
pure (Int
at Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
len)
  where
    len :: Int
len = Int -> Int
decimalLength Int
n

-- Write the digits of a magnitude backwards from the offset.
writeDigits :: MutableByteArray st -> Int -> Word -> ST st ()
writeDigits :: forall st. MutableByteArray st -> Int -> Word -> ST st ()
writeDigits MutableByteArray st
buffer !Int
at !Word
value = do
    MutableByteArray (PrimState (ST st)) -> Int -> Word8 -> ST st ()
forall a (m :: * -> *).
(Prim a, PrimMonad m) =>
MutableByteArray (PrimState m) -> Int -> a -> m ()
writeByteArray MutableByteArray st
MutableByteArray (PrimState (ST st))
buffer Int
at (Word8
0x30 Word8 -> Word8 -> Word8
forall a. Num a => a -> a -> a
+ Word -> Word8
forall a b. (Integral a, Num b) => a -> b
fromIntegral (Word
value Word -> Word -> Word
forall a. Integral a => a -> a -> a
`rem` Word
10) :: Word8)
    if Word
value Word -> Word -> Bool
forall a. Ord a => a -> a -> Bool
>= Word
10 then MutableByteArray st -> Int -> Word -> ST st ()
forall st. MutableByteArray st -> Int -> Word -> ST st ()
writeDigits MutableByteArray st
buffer (Int
at Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1) (Word
value Word -> Word -> Word
forall a. Integral a => a -> a -> a
`quot` Word
10) else ST st ()
forall (f :: * -> *). Applicative f => f ()
pass

magnitude :: Int -> Word
magnitude :: Int -> Word
magnitude Int
n = if Int
n Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
< Int
0 then Int -> Word
forall a b. (Integral a, Num b) => a -> b
fromIntegral (Int -> Int
forall a. Num a => a -> a
negate Int
n) else Int -> Word
forall a b. (Integral a, Num b) => a -> b
fromIntegral Int
n

-- | The bytes 'putDecimal' writes.
decimalLength :: Int -> Int
decimalLength :: Int -> Int
decimalLength Int
n = (if Int
n Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
< Int
0 then Int
1 else Int
0) Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Word -> Int -> Int
count (Int -> Word
magnitude Int
n) Int
1
  where
    count :: Word -> Int -> Int
    count :: Word -> Int -> Int
count !Word
remaining !Int
total
        | Word
remaining Word -> Word -> Bool
forall a. Ord a => a -> a -> Bool
< Word
10 = Int
total
        | Bool
otherwise = Word -> Int -> Int
count (Word
remaining Word -> Word -> Word
forall a. Integral a => a -> a -> a
`quot` Word
10) (Int
total Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)

-- | A copy of the bytes between two offsets, in an array of their exact size.
copyOut :: Scratch st -> Int -> Int -> ST st ByteArray
copyOut :: forall st. Scratch st -> Int -> Int -> ST st ByteArray
copyOut Scratch st
scratch Int
from Int
to = do
    buffer <- Scratch st -> ST st (MutableByteArray st)
forall st. Scratch st -> ST st (MutableByteArray st)
scratchBuffer Scratch st
scratch
    target <- newByteArray (to - from)
    copyMutableByteArray target 0 buffer from (to - from)
    unsafeFreezeByteArray target