-- | Pure implementation of @bytea@ unescaping. See 'unescapeBytea'.
module Pqi.Native.UnescapeBytea
  ( unescapeBytea,
  )
where

import Data.ByteString (ByteString)
import Data.Either (fromRight)
import Data.Word (Word8)
import Prelude
import PtrPeeker (Variable, fixed, hasMore, runVariableOnByteString, unsignedInt1)
import PtrPoker.Write (Write)
import qualified PtrPoker.Write as Write

-- | Convert the textual representation of a @bytea@ value, as produced by
-- the server, back into raw bytes. Both the modern @\\x@ hex format (lowercase
-- @x@ only) and the legacy escape format are accepted.
--
-- Malformed input is tolerated exactly the way @PQunescapeBytea@ tolerates it:
-- in hex format, characters that are not hex digits (including whitespace) are
-- silently skipped, and a hex digit whose pair character is invalid is
-- dropped; in escape format, an invalid escape simply drops the backslash, and
-- an octal escape must start with @0@..@3@. Input is treated as a C string:
-- the first NUL byte terminates processing.
unescapeBytea :: ByteString -> ByteString
unescapeBytea :: ByteString -> ByteString
unescapeBytea ByteString
input =
  Write -> ByteString
Write.toByteString
    (Write -> ByteString) -> Write -> ByteString
forall a b. (a -> b) -> a -> b
$ Write -> Either Int Write -> Write
forall b a. b -> Either a b -> b
fromRight Write
forall a. Monoid a => a
mempty
    (Either Int Write -> Write) -> Either Int Write -> Write
forall a b. (a -> b) -> a -> b
$ Variable Write -> ByteString -> Either Int Write
forall a. Variable a -> ByteString -> Either Int a
runVariableOnByteString Variable Write
decoder ByteString
input

-- Inline NUL truncation and \x prefix detection so no intermediate ByteStrings
-- are allocated before dispatching to the format-specific decoder.
decoder :: Variable Write
decoder :: Variable Write
decoder = do
  Bool
more <- Variable Bool
hasMore
  if Bool -> Bool
not Bool
more
    then Write -> Variable Write
forall a. a -> Variable a
forall (m :: * -> *) a. Monad m => a -> m a
return Write
forall a. Monoid a => a
mempty
    else do
      Word8
b0 <- Fixed Word8 -> Variable Word8
forall a. Fixed a -> Variable a
fixed Fixed Word8
unsignedInt1
      case Word8
b0 of
        Word8
0x00 -> Write -> Variable Write
forall a. a -> Variable a
forall (m :: * -> *) a. Monad m => a -> m a
return Write
forall a. Monoid a => a
mempty -- NUL: C-string terminator
        Word8
0x5c -> do
          -- backslash: probe for the \x hex-format prefix
          Bool
more2 <- Variable Bool
hasMore
          if Bool -> Bool
not Bool
more2
            then Write -> Variable Write
forall a. a -> Variable a
forall (m :: * -> *) a. Monad m => a -> m a
return Write
forall a. Monoid a => a
mempty -- single trailing backslash
            else do
              Word8
b1 <- Fixed Word8 -> Variable Word8
forall a. Fixed a -> Variable a
fixed Fixed Word8
unsignedInt1
              if Word8
b1 Word8 -> Word8 -> Bool
forall a. Eq a => a -> a -> Bool
== Word8
0x78 -- lowercase 'x': enter hex mode
                then Variable Write
hexDecoder
                else Word8 -> Variable Write
afterBackslash Word8
b1 -- escape mode; b1 follows the consumed '\'
        Word8
_ -> (Word8 -> Write
Write.word8 Word8
b0 Write -> Write -> Write
forall a. Semigroup a => a -> a -> a
<>) (Write -> Write) -> Variable Write -> Variable Write
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
<$> Variable Write
escapeDecoder

-- | Hex-format decoder. Skips non-hex bytes (matching @PQunescapeBytea@),
-- pairs hex nibbles, and stops at a NUL byte (C-string terminator).
hexDecoder :: Variable Write
hexDecoder :: Variable Write
hexDecoder = do
  Bool
more <- Variable Bool
hasMore
  if Bool -> Bool
not Bool
more
    then Write -> Variable Write
forall a. a -> Variable a
forall (m :: * -> *) a. Monad m => a -> m a
return Write
forall a. Monoid a => a
mempty
    else do
      Word8
a <- Fixed Word8 -> Variable Word8
forall a. Fixed a -> Variable a
fixed Fixed Word8
unsignedInt1
      if Word8
a Word8 -> Word8 -> Bool
forall a. Eq a => a -> a -> Bool
== Word8
0x00
        then Write -> Variable Write
forall a. a -> Variable a
forall (m :: * -> *) a. Monad m => a -> m a
return Write
forall a. Monoid a => a
mempty -- NUL: stop
        else case Word8 -> Maybe Word8
hexValue Word8
a of
          Maybe Word8
Nothing -> Variable Write
hexDecoder -- skip non-hex byte
          Just Word8
hi -> do
            Bool
more2 <- Variable Bool
hasMore
            if Bool -> Bool
not Bool
more2
              then Write -> Variable Write
forall a. a -> Variable a
forall (m :: * -> *) a. Monad m => a -> m a
return Write
forall a. Monoid a => a
mempty -- drop unpaired nibble
              else do
                Word8
b <- Fixed Word8 -> Variable Word8
forall a. Fixed a -> Variable a
fixed Fixed Word8
unsignedInt1
                if Word8
b Word8 -> Word8 -> Bool
forall a. Eq a => a -> a -> Bool
== Word8
0x00
                  then Write -> Variable Write
forall a. a -> Variable a
forall (m :: * -> *) a. Monad m => a -> m a
return Write
forall a. Monoid a => a
mempty -- NUL: stop, drop unpaired nibble
                  else case Word8 -> Maybe Word8
hexValue Word8
b of
                    Maybe Word8
Nothing -> Variable Write
hexDecoder -- skip b, look for next pair
                    Just Word8
lo -> (Word8 -> Write
Write.word8 (Word8
hi Word8 -> Word8 -> Word8
forall a. Num a => a -> a -> a
* Word8
16 Word8 -> Word8 -> Word8
forall a. Num a => a -> a -> a
+ Word8
lo) Write -> Write -> Write
forall a. Semigroup a => a -> a -> a
<>) (Write -> Write) -> Variable Write -> Variable Write
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
<$> Variable Write
hexDecoder
  where
    hexValue :: Word8 -> Maybe Word8
    hexValue :: Word8 -> Maybe Word8
hexValue Word8
w
      | Word8
w Word8 -> Word8 -> Bool
forall a. Ord a => a -> a -> Bool
>= Word8
0x30 Bool -> Bool -> Bool
&& Word8
w Word8 -> Word8 -> Bool
forall a. Ord a => a -> a -> Bool
<= Word8
0x39 = Word8 -> Maybe Word8
forall a. a -> Maybe a
Just (Word8
w Word8 -> Word8 -> Word8
forall a. Num a => a -> a -> a
- Word8
0x30)
      | Word8
w Word8 -> Word8 -> Bool
forall a. Ord a => a -> a -> Bool
>= Word8
0x61 Bool -> Bool -> Bool
&& Word8
w Word8 -> Word8 -> Bool
forall a. Ord a => a -> a -> Bool
<= Word8
0x66 = Word8 -> Maybe Word8
forall a. a -> Maybe a
Just (Word8
w Word8 -> Word8 -> Word8
forall a. Num a => a -> a -> a
- Word8
0x57)
      | Word8
w Word8 -> Word8 -> Bool
forall a. Ord a => a -> a -> Bool
>= Word8
0x41 Bool -> Bool -> Bool
&& Word8
w Word8 -> Word8 -> Bool
forall a. Ord a => a -> a -> Bool
<= Word8
0x46 = Word8 -> Maybe Word8
forall a. a -> Maybe a
Just (Word8
w Word8 -> Word8 -> Word8
forall a. Num a => a -> a -> a
- Word8
0x37)
      | Bool
otherwise = Maybe Word8
forall a. Maybe a
Nothing

-- | Escape-format decoder. Processes bytes as escape sequences and stops at
-- a NUL byte (C-string terminator).
escapeDecoder :: Variable Write
escapeDecoder :: Variable Write
escapeDecoder = do
  Bool
more <- Variable Bool
hasMore
  if Bool -> Bool
not Bool
more
    then Write -> Variable Write
forall a. a -> Variable a
forall (m :: * -> *) a. Monad m => a -> m a
return Write
forall a. Monoid a => a
mempty
    else do
      Word8
b <- Fixed Word8 -> Variable Word8
forall a. Fixed a -> Variable a
fixed Fixed Word8
unsignedInt1
      case Word8
b of
        Word8
0x00 -> Write -> Variable Write
forall a. a -> Variable a
forall (m :: * -> *) a. Monad m => a -> m a
return Write
forall a. Monoid a => a
mempty
        Word8
0x5c -> Variable Write
handleEscapeBackslash
        Word8
_ -> (Word8 -> Write
Write.word8 Word8
b Write -> Write -> Write
forall a. Semigroup a => a -> a -> a
<>) (Write -> Write) -> Variable Write -> Variable Write
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
<$> Variable Write
escapeDecoder

-- | Handle the bytes that follow a consumed backslash in escape format.
-- Exported so the top-level dispatcher can reuse it after consuming the
-- @\\x@ prefix check.
afterBackslash :: Word8 -> Variable Write
afterBackslash :: Word8 -> Variable Write
afterBackslash Word8
next = case Word8
next of
  Word8
0x00 -> Write -> Variable Write
forall a. a -> Variable a
forall (m :: * -> *) a. Monad m => a -> m a
return Write
forall a. Monoid a => a
mempty
  Word8
0x5c -> (Word8 -> Write
Write.word8 Word8
0x5c Write -> Write -> Write
forall a. Semigroup a => a -> a -> a
<>) (Write -> Write) -> Variable Write -> Variable Write
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
<$> Variable Write
escapeDecoder
  Word8
_ -> Word8 -> Variable Write
octalOrLiteralDecoder Word8
next

handleEscapeBackslash :: Variable Write
handleEscapeBackslash :: Variable Write
handleEscapeBackslash = do
  Bool
more <- Variable Bool
hasMore
  if Bool -> Bool
not Bool
more
    then Write -> Variable Write
forall a. a -> Variable a
forall (m :: * -> *) a. Monad m => a -> m a
return Write
forall a. Monoid a => a
mempty -- trailing backslash: drop it
    else do
      Word8
next <- Fixed Word8 -> Variable Word8
forall a. Fixed a -> Variable a
fixed Fixed Word8
unsignedInt1
      Word8 -> Variable Write
afterBackslash Word8
next

-- | Try to decode a 3-digit octal starting with @a@ (already consumed).
-- Falls back to emitting @a@ literally and re-routing the consumed lookahead
-- byte(s) through 'afterEscape', reproducing @PQunescapeBytea@'s backtracking.
octalOrLiteralDecoder :: Word8 -> Variable Write
octalOrLiteralDecoder :: Word8 -> Variable Write
octalOrLiteralDecoder Word8
a
  | Word8 -> Bool
isFirstOctal Word8
a = do
      Bool
more <- Variable Bool
hasMore
      if Bool -> Bool
not Bool
more
        then Write -> Variable Write
forall a. a -> Variable a
forall (m :: * -> *) a. Monad m => a -> m a
return (Word8 -> Write
Write.word8 Word8
a)
        else do
          Word8
b <- Fixed Word8 -> Variable Word8
forall a. Fixed a -> Variable a
fixed Fixed Word8
unsignedInt1
          if Bool -> Bool
not (Word8 -> Bool
isOctal Word8
b)
            then (Word8 -> Write
Write.word8 Word8
a Write -> Write -> Write
forall a. Semigroup a => a -> a -> a
<>) (Write -> Write) -> Variable Write -> Variable Write
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
<$> Word8 -> Variable Write
afterEscape Word8
b -- b isn't octal: emit a, re-route b
            else do
              Bool
more2 <- Variable Bool
hasMore
              if Bool -> Bool
not Bool
more2
                then Write -> Variable Write
forall a. a -> Variable a
forall (m :: * -> *) a. Monad m => a -> m a
return (Word8 -> Write
Write.word8 Word8
a Write -> Write -> Write
forall a. Semigroup a => a -> a -> a
<> Word8 -> Write
Write.word8 Word8
b) -- only two digits: both literal
                else do
                  Word8
c <- Fixed Word8 -> Variable Word8
forall a. Fixed a -> Variable a
fixed Fixed Word8
unsignedInt1
                  if Word8 -> Bool
isOctal Word8
c
                    then (Word8 -> Write
Write.word8 (Word8 -> Word8 -> Word8 -> Word8
octal Word8
a Word8
b Word8
c) Write -> Write -> Write
forall a. Semigroup a => a -> a -> a
<>) (Write -> Write) -> Variable Write -> Variable Write
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
<$> Variable Write
escapeDecoder
                    else (\Write
x -> Word8 -> Write
Write.word8 Word8
a Write -> Write -> Write
forall a. Semigroup a => a -> a -> a
<> Word8 -> Write
Write.word8 Word8
b Write -> Write -> Write
forall a. Semigroup a => a -> a -> a
<> Write
x) (Write -> Write) -> Variable Write -> Variable Write
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
<$> Word8 -> Variable Write
afterEscape Word8
c
  | Bool
otherwise = (Word8 -> Write
Write.word8 Word8
a Write -> Write -> Write
forall a. Semigroup a => a -> a -> a
<>) (Write -> Write) -> Variable Write -> Variable Write
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
<$> Variable Write
escapeDecoder
  where
    -- \| Route an already-consumed byte back through the escape-format main loop.
    -- Used when a consumed lookahead byte must be re-processed after a failed
    -- octal-triple attempt.
    afterEscape :: Word8 -> Variable Write
    afterEscape :: Word8 -> Variable Write
afterEscape Word8
b = case Word8
b of
      Word8
0x00 -> Write -> Variable Write
forall a. a -> Variable a
forall (m :: * -> *) a. Monad m => a -> m a
return Write
forall a. Monoid a => a
mempty
      Word8
0x5c -> Variable Write
handleEscapeBackslash
      Word8
_ -> (Word8 -> Write
Write.word8 Word8
b Write -> Write -> Write
forall a. Semigroup a => a -> a -> a
<>) (Write -> Write) -> Variable Write -> Variable Write
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
<$> Variable Write
escapeDecoder

    isFirstOctal :: Word8 -> Bool
    isFirstOctal :: Word8 -> Bool
isFirstOctal Word8
w = Word8
w Word8 -> Word8 -> Bool
forall a. Ord a => a -> a -> Bool
>= Word8
0x30 Bool -> Bool -> Bool
&& Word8
w Word8 -> Word8 -> Bool
forall a. Ord a => a -> a -> Bool
<= Word8
0x33

    isOctal :: Word8 -> Bool
    isOctal :: Word8 -> Bool
isOctal Word8
w = Word8
w Word8 -> Word8 -> Bool
forall a. Ord a => a -> a -> Bool
>= Word8
0x30 Bool -> Bool -> Bool
&& Word8
w Word8 -> Word8 -> Bool
forall a. Ord a => a -> a -> Bool
<= Word8
0x37

    octal :: Word8 -> Word8 -> Word8 -> Word8
    octal :: Word8 -> Word8 -> Word8 -> Word8
octal Word8
a Word8
b Word8
c = (Word8
a Word8 -> Word8 -> Word8
forall a. Num a => a -> a -> a
- Word8
0x30) Word8 -> Word8 -> Word8
forall a. Num a => a -> a -> a
* Word8
64 Word8 -> Word8 -> Word8
forall a. Num a => a -> a -> a
+ (Word8
b Word8 -> Word8 -> Word8
forall a. Num a => a -> a -> a
- Word8
0x30) Word8 -> Word8 -> Word8
forall a. Num a => a -> a -> a
* Word8
8 Word8 -> Word8 -> Word8
forall a. Num a => a -> a -> a
+ (Word8
c Word8 -> Word8 -> Word8
forall a. Num a => a -> a -> a
- Word8
0x30)