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

{- | Incremental registry extraction within a decompressed body ceiling. With the vendored json-stream
and text 2.1.3, decoded strings and keys own their arrays, including chunk-spanning tokens.
See <https://github.com/ondrap/json-stream/blob/537a43a775e64f50dc63c373193323de98619799/Data/JsonStream/Unescape.hs decoder storage>.
-}
module Ecluse.Core.Registry.JsonStream (
    -- * Bounded reads
    StreamResult (..),
    Steps (..),
    Step,
    readSteps,
    readJsonStream,

    -- * Retained values
    retainedValue,
    withinRetainedDepth,
    Members,
    namedMembers,
    everyMember,
    retainedObjectOr,
    retainedScalar,
    retainedObjectWith,
    retainedArrayWith,
) where

import Data.Aeson (Value (..))
import Data.Aeson.Key qualified as Key
import Data.Aeson.KeyMap qualified as KeyMap
import Data.ByteString qualified as BS
import Data.HashMap.Strict qualified as HashMap
import Data.JsonStream.Parser qualified as J
import Data.Vector qualified as V

import Ecluse.Core.Registry (ParseError (..))
import Ecluse.Core.Security (BodyLimit, LimitError (BodyTooLarge), bodyLimitBytes)

-- | Extracted data and the size of the complete decompressed source, including ignored fields.
data StreamResult a = StreamResult
    { forall a. StreamResult a -> Either ParseError a
streamValue :: Either ParseError a
    , forall a. StreamResult a -> Int
streamBytes :: Int
    }
    deriving stock (StreamResult a -> StreamResult a -> Bool
(StreamResult a -> StreamResult a -> Bool)
-> (StreamResult a -> StreamResult a -> Bool)
-> Eq (StreamResult a)
forall a. Eq a => StreamResult a -> StreamResult a -> Bool
forall a. (a -> a -> Bool) -> (a -> a -> Bool) -> Eq a
$c== :: forall a. Eq a => StreamResult a -> StreamResult a -> Bool
== :: StreamResult a -> StreamResult a -> Bool
$c/= :: forall a. Eq a => StreamResult a -> StreamResult a -> Bool
/= :: StreamResult a -> StreamResult a -> Bool
Eq, Int -> StreamResult a -> ShowS
[StreamResult a] -> ShowS
StreamResult a -> String
(Int -> StreamResult a -> ShowS)
-> (StreamResult a -> String)
-> ([StreamResult a] -> ShowS)
-> Show (StreamResult a)
forall a. Show a => Int -> StreamResult a -> ShowS
forall a. Show a => [StreamResult a] -> ShowS
forall a. Show a => StreamResult a -> String
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: forall a. Show a => Int -> StreamResult a -> ShowS
showsPrec :: Int -> StreamResult a -> ShowS
$cshow :: forall a. Show a => StreamResult a -> String
show :: StreamResult a -> String
$cshowList :: forall a. Show a => [StreamResult a] -> ShowS
showList :: [StreamResult a] -> ShowS
Show)

{- | A read in progress: it needs input, stops on a parse error or a refused value, or has finished.
A read that writes as it goes resumes in its own effect.
-}
data Steps m s
    = NeedData (ByteString -> m (Steps m s))
    | Failed Text
    | Refused LimitError
    | Finished s

-- | A read with no effect of its own.
type Step = Steps Identity

{- | Feed a read in pieces of at most 32 KiB, within the body ceiling, and drain the body after the
read finishes. An empty chunk ends the body. The read resumes in its effect, run in the reader's.
-}
readSteps :: (Monad n) => (forall a. m a -> n a) -> BodyLimit -> Steps m s -> n ByteString -> n (Either LimitError (StreamResult s))
readSteps :: forall (n :: * -> *) (m :: * -> *) s.
Monad n =>
(forall a. m a -> n a)
-> BodyLimit
-> Steps m s
-> n ByteString
-> n (Either LimitError (StreamResult s))
readSteps forall a. m a -> n a
run BodyLimit
bound Steps m s
start n ByteString
readChunk = Int -> Steps m s -> n (Either LimitError (StreamResult s))
forall {a}.
Int -> Steps m a -> n (Either LimitError (StreamResult a))
go Int
0 Steps m s
start
  where
    go :: Int -> Steps m a -> n (Either LimitError (StreamResult a))
go !Int
seen Steps m a
step = case Steps m a
step of
        Refused LimitError
fault -> Either LimitError (StreamResult a)
-> n (Either LimitError (StreamResult a))
forall a. a -> n a
forall (f :: * -> *) a. Applicative f => a -> f a
pure (LimitError -> Either LimitError (StreamResult a)
forall a b. a -> Either a b
Left LimitError
fault)
        Failed Text
err -> Either LimitError (StreamResult a)
-> n (Either LimitError (StreamResult a))
forall a. a -> n a
forall (f :: * -> *) a. Applicative f => a -> f a
pure (StreamResult a -> Either LimitError (StreamResult a)
forall a b. b -> Either a b
Right (Either ParseError a -> Int -> StreamResult a
forall a. Either ParseError a -> Int -> StreamResult a
StreamResult (ParseError -> Either ParseError a
forall a b. a -> Either a b
Left (Text -> ParseError
ParseError Text
err)) Int
seen))
        Steps m a
_ -> do
            chunk <- n ByteString
readChunk
            if BS.null chunk
                then pure . Right $ StreamResult (finish step) seen
                else
                    if BS.length chunk > bodyLimitBytes bound - seen
                        then pure (Left (BodyTooLarge bound))
                        else feed (seen + BS.length chunk) step chunk
    feed :: Int
-> Steps m a
-> ByteString
-> n (Either LimitError (StreamResult a))
feed Int
seen Steps m a
step ByteString
chunk = case Steps m a
step of
        NeedData ByteString -> m (Steps m a)
next
            | Bool -> Bool
not (ByteString -> Bool
BS.null ByteString
chunk) -> do
                let (ByteString
piece, ByteString
remaining) = Int -> ByteString -> (ByteString, ByteString)
BS.splitAt Int
32768 ByteString
chunk
                resumed <- m (Steps m a) -> n (Steps m a)
forall a. m a -> n a
run (ByteString -> m (Steps m a)
next ByteString
piece)
                feed seen resumed remaining
        Steps m a
_ -> Int -> Steps m a -> n (Either LimitError (StreamResult a))
go Int
seen Steps m a
step
    finish :: Steps m b -> Either ParseError b
finish = \case
        Finished b
result -> b -> Either ParseError b
forall a b. b -> Either a b
Right b
result
        Steps m b
_ -> ParseError -> Either ParseError b
forall a b. a -> Either a b
Left (Text -> ParseError
ParseError Text
"incomplete registry JSON")

-- | Fold each value the parser yields through the step, as 'readSteps' feeds it.
readJsonStream :: (Monad m) => BodyLimit -> J.Parser a -> (s -> a -> Either LimitError s) -> s -> m ByteString -> m (Either LimitError (StreamResult s))
readJsonStream :: forall (m :: * -> *) a s.
Monad m =>
BodyLimit
-> Parser a
-> (s -> a -> Either LimitError s)
-> s
-> m ByteString
-> m (Either LimitError (StreamResult s))
readJsonStream BodyLimit
bound Parser a
parser s -> a -> Either LimitError s
step s
initial = (forall a. Identity a -> m a)
-> BodyLimit
-> Steps Identity s
-> m ByteString
-> m (Either LimitError (StreamResult s))
forall (n :: * -> *) (m :: * -> *) s.
Monad n =>
(forall a. m a -> n a)
-> BodyLimit
-> Steps m s
-> n ByteString
-> n (Either LimitError (StreamResult s))
readSteps (a -> m a
forall a. a -> m a
forall (f :: * -> *) a. Applicative f => a -> f a
pure (a -> m a) -> (Identity a -> a) -> Identity a -> m a
forall b c a. (b -> c) -> (a -> b) -> a -> c
. Identity a -> a
forall a. Identity a -> a
runIdentity) BodyLimit
bound (s -> ParseOutput a -> Steps Identity s
parserSteps s
initial (Parser a -> ParseOutput a
forall a. Parser a -> ParseOutput a
J.runParser Parser a
parser))
  where
    parserSteps :: s -> ParseOutput a -> Steps Identity s
parserSteps !s
acc = \case
        J.ParseYield a
value ParseOutput a
next -> (LimitError -> Steps Identity s)
-> (s -> Steps Identity s)
-> Either LimitError s
-> Steps Identity s
forall a c b. (a -> c) -> (b -> c) -> Either a b -> c
either LimitError -> Steps Identity s
forall (m :: * -> *) s. LimitError -> Steps m s
Refused (s -> ParseOutput a -> Steps Identity s
`parserSteps` ParseOutput a
next) (s -> a -> Either LimitError s
step s
acc a
value)
        J.ParseNeedData ByteString -> ParseOutput a
next -> (ByteString -> Identity (Steps Identity s)) -> Steps Identity s
forall (m :: * -> *) s. (ByteString -> m (Steps m s)) -> Steps m s
NeedData (Steps Identity s -> Identity (Steps Identity s)
forall a. a -> Identity a
Identity (Steps Identity s -> Identity (Steps Identity s))
-> (ByteString -> Steps Identity s)
-> ByteString
-> Identity (Steps Identity s)
forall b c a. (b -> c) -> (a -> b) -> a -> c
. s -> ParseOutput a -> Steps Identity s
parserSteps s
acc (ParseOutput a -> Steps Identity s)
-> (ByteString -> ParseOutput a) -> ByteString -> Steps Identity s
forall b c a. (b -> c) -> (a -> b) -> a -> c
. ByteString -> ParseOutput a
next)
        J.ParseFailed String
err -> Text -> Steps Identity s
forall (m :: * -> *) s. Text -> Steps m s
Failed (String -> Text
forall a. ToText a => a -> Text
toText String
err)
        J.ParseDone ByteString
_ -> s -> Steps Identity s
forall (m :: * -> *) s. s -> Steps m s
Finished s
acc

-- | Decode a retained field within a structural budget. Unknown fields never call this parser.
retainedValue :: Int -> J.Parser Value
retainedValue :: Int -> Parser Value
retainedValue Int
depth =
    Int -> Parser Value -> Parser Value
forall a. Int -> Parser a -> Parser a
withinRetainedDepth Int
depth (Parser Value -> Parser Value) -> Parser Value -> Parser Value
forall a b. (a -> b) -> a -> b
$
        Parser Value -> Members -> Parser Value
retainedObjectWith
            (Parser Value -> Parser Value -> Parser Value
retainedArrayWith Parser Value
retainedScalar Parser Value
child)
            (Parser Value -> Members
everyMember Parser Value
child)
  where
    child :: Parser Value
child = Int -> Parser Value
retainedValue (Int
depth Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1)

-- | Charge the parsed value's own level, including empty containers. Children need one less level.
withinRetainedDepth :: Int -> J.Parser a -> J.Parser a
withinRetainedDepth :: forall a. Int -> Parser a -> Parser a
withinRetainedDepth Int
budget Parser a
parser
    | Int
budget Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
<= Int
0 = (() -> Either String a) -> Parser () -> Parser a
forall a b. (a -> Either String b) -> Parser a -> Parser b
J.mapWithFailure (Either String a -> () -> Either String a
forall a b. a -> b -> a
const (String -> Either String a
forall a b. a -> Either a b
Left String
"retained JSON nesting limit")) (() -> Parser ()
forall a. a -> Parser a
forall (f :: * -> *) a. Applicative f => a -> f a
pure ())
    | Bool
otherwise = Parser a
parser

{- | Which members of an object are retained, with what parser, and under which key. Every object
read with one 'Members' value holds each name it knows under one shared key.
-}
data Members
    = NamedMembers (HashMap Text (Key.Key, J.Parser Value))
    | EveryMember (J.Parser Value)

-- | Retain only the named members. The first entry for a name wins.
namedMembers :: [(Text, J.Parser Value)] -> Members
namedMembers :: [(Text, Parser Value)] -> Members
namedMembers [(Text, Parser Value)]
entries = HashMap Text (Key, Parser Value) -> Members
NamedMembers (((Key, Parser Value) -> (Key, Parser Value) -> (Key, Parser Value))
-> [(Text, (Key, Parser Value))]
-> HashMap Text (Key, Parser Value)
forall k v.
(Eq k, Hashable k) =>
(v -> v -> v) -> [(k, v)] -> HashMap k v
HashMap.fromListWith (\(Key, Parser Value)
_ (Key, Parser Value)
earlier -> (Key, Parser Value)
earlier) [(Text
name, (Text -> Key
Key.fromText Text
name, Parser Value
parser)) | (Text
name, Parser Value
parser) <- [(Text, Parser Value)]
entries])

-- | Retain every member with one parser, each under its own key.
everyMember :: J.Parser Value -> Members
everyMember :: Parser Value -> Members
everyMember = Parser Value -> Members
EveryMember

-- | Supply an invalid-shape witness without traversing a valid object through a parallel fallback.
retainedObjectOr :: Value -> Members -> J.Parser Value
retainedObjectOr :: Value -> Members -> Parser Value
retainedObjectOr Value
fallback = (Maybe Value -> Value) -> Parser (Maybe Value) -> Parser Value
forall a b. (a -> b) -> Parser a -> Parser b
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
fmap (Value -> Maybe Value -> Value
forall a. a -> Maybe a -> a
fromMaybe Value
fallback) (Parser (Maybe Value) -> Parser Value)
-> (Members -> Parser (Maybe Value)) -> Members -> Parser Value
forall b c a. (b -> c) -> (a -> b) -> a -> c
. Parser RetainedEvent -> Parser (Maybe Value)
foldRetained (Parser RetainedEvent -> Parser (Maybe Value))
-> (Members -> Parser RetainedEvent)
-> Members
-> Parser (Maybe Value)
forall b c a. (b -> c) -> (a -> b) -> a -> c
. Members -> Parser RetainedEvent
objectEvents

-- | Select object events before folding. The fallback handles scalars and other container shapes.
retainedObjectWith :: J.Parser Value -> Members -> J.Parser Value
retainedObjectWith :: Parser Value -> Members -> Parser Value
retainedObjectWith Parser Value
fallback Members
members = Parser (Maybe Value) -> Parser Value
forall a. Parser (Maybe a) -> Parser a
J.catMaybeI (Parser RetainedEvent -> Parser (Maybe Value)
foldRetained (Members -> Parser RetainedEvent
objectEvents Members
members Parser RetainedEvent
-> Parser RetainedEvent -> Parser RetainedEvent
forall a. Parser a -> Parser a -> Parser a
forall (f :: * -> *) a. Alternative f => f a -> f a -> f a
<|> (Value -> RetainedEvent
OtherValue (Value -> RetainedEvent) -> Parser Value -> Parser RetainedEvent
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
<$> Parser Value
fallback)))

-- | Select array events before folding. A container fallback must yield only its completed value.
retainedArrayWith :: J.Parser Value -> J.Parser Value -> J.Parser Value
retainedArrayWith :: Parser Value -> Parser Value -> Parser Value
retainedArrayWith Parser Value
fallback Parser Value
parser = Parser (Maybe Value) -> Parser Value
forall a. Parser (Maybe a) -> Parser a
J.catMaybeI (Parser RetainedEvent -> Parser (Maybe Value)
foldRetained (Parser Value -> Parser RetainedEvent
arrayEvents Parser Value
parser Parser RetainedEvent
-> Parser RetainedEvent -> Parser RetainedEvent
forall a. Parser a -> Parser a -> Parser a
forall (f :: * -> *) a. Alternative f => f a -> f a -> f a
<|> (Value -> RetainedEvent
OtherValue (Value -> RetainedEvent) -> Parser Value -> Parser RetainedEvent
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
<$> Parser Value
fallback)))

data RetainedEvent = BeginObject | ObjectField Key.Key Value | BeginArray | ArrayItem Value | OtherValue Value | EndContainer

data Retained = Missing | ObjectFields (KeyMap.KeyMap Value) | ArrayItems [Value] | ScalarValue Value

objectEvents :: Members -> J.Parser RetainedEvent
objectEvents :: Members -> Parser RetainedEvent
objectEvents Members
members = RetainedEvent
-> RetainedEvent -> Parser RetainedEvent -> Parser RetainedEvent
forall a. a -> a -> Parser a -> Parser a
J.objectFound RetainedEvent
BeginObject RetainedEvent
EndContainer ((Text -> Parser RetainedEvent) -> Parser RetainedEvent
forall a. (Text -> Parser a) -> Parser a
J.objectKeyValues (Members -> Text -> Parser RetainedEvent
memberEvent Members
members))

memberEvent :: Members -> Text -> J.Parser RetainedEvent
memberEvent :: Members -> Text -> Parser RetainedEvent
memberEvent = \case
    NamedMembers HashMap Text (Key, Parser Value)
named -> \Text
name -> Parser RetainedEvent
-> ((Key, Parser Value) -> Parser RetainedEvent)
-> Maybe (Key, Parser Value)
-> Parser RetainedEvent
forall b a. b -> (a -> b) -> Maybe a -> b
maybe Parser RetainedEvent
forall a. Monoid a => a
mempty (\(Key
key, Parser Value
parser) -> Key -> Value -> RetainedEvent
ObjectField Key
key (Value -> RetainedEvent) -> Parser Value -> Parser RetainedEvent
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
<$> Parser Value
parser) (Text
-> HashMap Text (Key, Parser Value) -> Maybe (Key, Parser Value)
forall k v. (Eq k, Hashable k) => k -> HashMap k v -> Maybe v
HashMap.lookup Text
name HashMap Text (Key, Parser Value)
named)
    EveryMember Parser Value
parser -> \Text
name -> Key -> Value -> RetainedEvent
ObjectField (Text -> Key
Key.fromText Text
name) (Value -> RetainedEvent) -> Parser Value -> Parser RetainedEvent
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
<$> Parser Value
parser

arrayEvents :: J.Parser Value -> J.Parser RetainedEvent
arrayEvents :: Parser Value -> Parser RetainedEvent
arrayEvents Parser Value
parser = RetainedEvent
-> RetainedEvent -> Parser RetainedEvent -> Parser RetainedEvent
forall a. a -> a -> Parser a -> Parser a
J.arrayFound RetainedEvent
BeginArray RetainedEvent
EndContainer (Value -> RetainedEvent
ArrayItem (Value -> RetainedEvent) -> Parser Value -> Parser RetainedEvent
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
<$> Parser Value -> Parser Value
forall a. Parser a -> Parser a
J.arrayOf Parser Value
parser)

-- The fallback contributes one completed value. Unmatched first shapes still skip their input.
foldRetained :: J.Parser RetainedEvent -> J.Parser (Maybe Value)
foldRetained :: Parser RetainedEvent -> Parser (Maybe Value)
foldRetained = (Retained -> Maybe Value)
-> Parser Retained -> Parser (Maybe Value)
forall a b. (a -> b) -> Parser a -> Parser b
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
fmap Retained -> Maybe Value
finish (Parser Retained -> Parser (Maybe Value))
-> (Parser RetainedEvent -> Parser Retained)
-> Parser RetainedEvent
-> Parser (Maybe Value)
forall b c a. (b -> c) -> (a -> b) -> a -> c
. (Retained -> RetainedEvent -> Retained)
-> Retained -> Parser RetainedEvent -> Parser Retained
forall b a. (b -> a -> b) -> b -> Parser a -> Parser b
J.foldI Retained -> RetainedEvent -> Retained
collect Retained
Missing
  where
    collect :: Retained -> RetainedEvent -> Retained
collect Retained
_ RetainedEvent
BeginObject = KeyMap Value -> Retained
ObjectFields KeyMap Value
forall a. Monoid a => a
mempty
    collect Retained
_ RetainedEvent
BeginArray = [Value] -> Retained
ArrayItems []
    collect (ObjectFields KeyMap Value
fields) (ObjectField Key
key Value
value) =
        KeyMap Value -> Retained
ObjectFields (if Key -> KeyMap Value -> Bool
forall a. Key -> KeyMap a -> Bool
KeyMap.member Key
key KeyMap Value
fields then KeyMap Value
fields else Key -> Value -> KeyMap Value -> KeyMap Value
forall v. Key -> v -> KeyMap v -> KeyMap v
KeyMap.insert Key
key Value
value KeyMap Value
fields)
    collect (ArrayItems [Value]
values) (ArrayItem Value
value) = [Value] -> Retained
ArrayItems (Value
value Value -> [Value] -> [Value]
forall a. a -> [a] -> [a]
: [Value]
values)
    collect Retained
_ (OtherValue Value
value) = Value -> Retained
ScalarValue Value
value
    collect Retained
current RetainedEvent
_ = Retained
current
    finish :: Retained -> Maybe Value
finish Retained
Missing = Maybe Value
forall a. Maybe a
Nothing
    finish (ObjectFields KeyMap Value
fields) = Value -> Maybe Value
forall a. a -> Maybe a
Just (KeyMap Value -> Value
Object KeyMap Value
fields)
    finish (ArrayItems [Value]
values) = Value -> Maybe Value
forall a. a -> Maybe a
Just (Array -> Value
Array ([Value] -> Array
forall a. [a] -> Vector a
V.fromList ([Value] -> [Value]
forall a. [a] -> [a]
reverse [Value]
values)))
    finish (ScalarValue Value
value) = Value -> Maybe Value
forall a. a -> Maybe a
Just Value
value

-- | Read a scalar without materialising an object or array when the field has the wrong shape.
retainedScalar :: J.Parser Value
retainedScalar :: Parser Value
retainedScalar = (Text -> Value
String (Text -> Value) -> Parser Text -> Parser Value
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
<$> Parser Text
J.string) Parser Value -> Parser Value -> Parser Value
forall a. Parser a -> Parser a -> Parser a
forall (f :: * -> *) a. Alternative f => f a -> f a -> f a
<|> (Scientific -> Value
Number (Scientific -> Value) -> Parser Scientific -> Parser Value
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
<$> Parser Scientific
J.number) Parser Value -> Parser Value -> Parser Value
forall a. Parser a -> Parser a -> Parser a
forall (f :: * -> *) a. Alternative f => f a -> f a -> f a
<|> (Bool -> Value
Bool (Bool -> Value) -> Parser Bool -> Parser Value
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
<$> Parser Bool
J.bool) Parser Value -> Parser Value -> Parser Value
forall a. Parser a -> Parser a -> Parser a
forall (f :: * -> *) a. Alternative f => f a -> f a -> f a
<|> (Value
Null Value -> Parser () -> Parser Value
forall a b. a -> Parser b -> Parser a
forall (f :: * -> *) a b. Functor f => a -> f b -> f a
<$ Parser ()
J.jNull)