packages feed

ppad-base64-0.1.1: test/Main.hs

{-# OPTIONS_GHC -fno-warn-missing-signatures #-}
{-# LANGUAGE OverloadedStrings #-}
{-# LANGUAGE PackageImports #-}

module Main where

import Control.Monad (forM_, when)
import qualified Data.ByteString as BS
import qualified Data.ByteString.Char8 as B8
import qualified "ppad-base64" Data.ByteString.Base64 as B64
import qualified "ppad-base64" Data.ByteString.Base64.Pure as Pure
import qualified "base64-bytestring" Data.ByteString.Base64 as R0
import Data.Word (Word8)
import Test.Tasty
import qualified Test.Tasty.QuickCheck as Q
import qualified Test.Tasty.HUnit as H

-- generators -----------------------------------------------------------------

-- random bytes, as a slice at a random offset into a larger buffer, so
-- that inputs with a non-zero 'ByteString' offset are exercised
bytes :: Int -> Q.Gen BS.ByteString
bytes k = do
  l <- Q.chooseInt (0, k)
  o <- Q.chooseInt (0, 32)
  v <- Q.vectorOf (o + l) Q.arbitrary
  pure (BS.drop o (BS.pack v))

newtype BS = BS BS.ByteString
  deriving (Eq, Show)

instance Q.Arbitrary BS where
  arbitrary = do
    b <- bytes 1024
    pure (BS b)

alphabet :: BS.ByteString
alphabet =
  "ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789+/"

-- valid encodings subjected to a few random mutations: substitution
-- (by an alphabet char, '=', or an arbitrary byte), deletion, or
-- insertion, with positions biased toward the final quartet
newtype Base64ish = Base64ish BS.ByteString
  deriving (Eq, Show)

instance Q.Arbitrary Base64ish where
  arbitrary = do
    b <- Q.frequency [(3, bytes 64), (1, bytes 384)]
    k <- Q.chooseInt (0, 3)
    s <- go k (B64.encode b)
    o <- Q.chooseInt (0, 32)
    p <- Q.vectorOf o (Q.elements (BS.unpack alphabet))
    pure (Base64ish (BS.drop o (BS.pack p <> s)))
    where
      go :: Int -> BS.ByteString -> Q.Gen BS.ByteString
      go 0 s = pure s
      go j s = mutate s >>= go (j - 1)

      mutate s = do
        let l = BS.length s
        i <- Q.oneof [
            Q.chooseInt (0, l)
          , Q.chooseInt (max 0 (l - 4), l)
          ]
        c <- Q.frequency [
            (6, Q.elements (BS.unpack alphabet))
          , (2, pure 0x3d)
          , (1, Q.arbitrary)
          ]
        Q.elements [
            BS.take i s <> BS.singleton c <> BS.drop (i + 1) s
          , BS.take i s <> BS.drop (i + 1) s
          , BS.take i s <> BS.singleton c <> BS.drop i s
          ]

-- decoders under test --------------------------------------------------------

decoders :: [(String, BS.ByteString -> Maybe BS.ByteString)]
decoders = [("public", B64.decode), ("pure", Pure.decode)]

ref_decode :: BS.ByteString -> Maybe BS.ByteString
ref_decode bs = case R0.decode bs of
  Left _  -> Nothing
  Right d -> Just d

-- properties -----------------------------------------------------------------

encode_matches_reference :: BS -> Bool
encode_matches_reference (BS bs) =
  let r0 = R0.encode bs
  in  B64.encode bs == r0 && Pure.encode bs == r0

decode_inverts_encode :: BS -> Bool
decode_inverts_encode (BS bs) =
  let enc = B64.encode bs
  in  all (\(_, dec) -> dec enc == Just bs) decoders

-- base64-bytestring's 'decode' is likewise strict (RFC 4648): it
-- rejects missing or misplaced padding and non-canonical final quartets
decode_matches_reference :: Base64ish -> Bool
decode_matches_reference (Base64ish bs) =
  let r0 = ref_decode bs
  in  B64.decode bs == r0 && Pure.decode bs == r0

-- unit tests -----------------------------------------------------------------

case_rfc_vectors :: TestTree
case_rfc_vectors = H.testCase "RFC 4648 \167 10 vectors" $ do
  let vectors = [
          ("",       "")
        , ("f",      "Zg==")
        , ("fo",     "Zm8=")
        , ("foo",    "Zm9v")
        , ("foob",   "Zm9vYg==")
        , ("fooba",  "Zm9vYmE=")
        , ("foobar", "Zm9vYmFy")
        ]
      check (input, expected) = do
        H.assertEqual ("encode " <> show input)
          expected (B64.encode input)
        H.assertEqual ("pure encode " <> show input)
          expected (Pure.encode input)
        forM_ decoders $ \(nam, dec) ->
          H.assertEqual (nam <> " decode " <> show expected)
            (Just input) (dec expected)
  mapM_ check vectors

-- malformed final quartets, alone and after a body long enough to
-- take the NEON path
malformed :: TestTree
malformed = H.testCase "malformed padding and non-canonical tails" $ do
  let bad = [
          "Zh=="     -- non-canonical: nonzero trailing bits
        , "Zm9="     -- non-canonical: nonzero trailing bits
        , "Zm=v"     -- '=' before a data char
        , "===="
        , "Z==="
        , "=Zg="
        , "Zg=a"
        , "Zg=", "Zg", "Z"
        , "Zm9vY"
        , "Zm=vYmFy" -- '=' in the body
        , "Zg==Zg==" -- padding mid-stream
        , "Zm9v\n"
        ]
      good = [("Zg==", "f"), ("Zm8=", "fo"), ("Zm9v", "foo")]
      body = B64.encode (BS.replicate 72 0xa5)
  forM_ decoders $ \(nam, dec) -> do
    forM_ bad $ \s -> do
      H.assertEqual (nam <> " " <> show s) Nothing (dec s)
      H.assertEqual (nam <> " body <> " <> show s) Nothing (dec (body <> s))
    forM_ good $ \(s, d) ->
      H.assertEqual (nam <> " body <> " <> show s)
        (Just (BS.replicate 72 0xa5 <> d)) (dec (body <> s))

-- every byte value, at every position of encodings that span the NEON
-- body loop, the scalar body tail, and each kind of final quartet,
-- decodes exactly as the reference does; in the body, a byte is
-- accepted iff it is an alphabet char
every_byte_every_position :: TestTree
every_byte_every_position =
  H.testCase "every byte at every position" $
    forM_ [28, 29, 30 :: Int] $ \n -> do
      let base = B64.encode (BS.replicate n 0x00)
          l    = BS.length base
      forM_ [0 .. 255 :: Int] $ \c ->
        forM_ [0 .. l - 1] $ \i -> do
          let w   = fromIntegral c :: Word8
              inp = BS.take i base <> BS.singleton w <> BS.drop (i + 1) base
              pec = ref_decode inp
              msg = "n = " <> show n <> ", byte " <> show c
                    <> ", position " <> show i
          when (i < l - 4) $
            H.assertEqual ("reference, " <> msg)
              (BS.elem w alphabet) (pec /= Nothing)
          forM_ decoders $ \(nam, dec) ->
            H.assertEqual (nam <> ", " <> msg) pec (dec inp)

-- for every length spanning several NEON iterations and each final
-- quartet shape, a single invalid byte at any position (or '=' anywhere
-- in the body) causes decoding to fail
single_corruption :: TestTree
single_corruption = H.testCase "single invalid byte anywhere" $ do
  let bad = BS.pack [0x00, 0x20, 0x7f, 0x80, 0xff] <> "-_.:@[`{"
  forM_ [1 .. 100 :: Int] $ \n -> do
    let raw = BS.pack (fmap fromIntegral [7 * j + 3 | j <- [1 .. n]])
        enc = B64.encode raw
        l   = BS.length enc
        put i w = BS.take i enc <> BS.singleton w <> BS.drop (i + 1) enc
    forM_ decoders $ \(nam, dec) ->
      H.assertEqual (nam <> ", intact, n = " <> show n) (Just raw) (dec enc)
    forM_ [0 .. l - 1] $ \i -> do
      let ws = BS.unpack bad <> (if i < l - 4 then [0x3d] else [])
      forM_ ws $ \w -> do
        let msg = "n = " <> show n <> ", position " <> show i
                  <> ", byte " <> show w
        forM_ decoders $ \(nam, dec) ->
          H.assertEqual (nam <> ", " <> msg) Nothing (dec (put i w))

bad_lengths :: TestTree
bad_lengths = H.testCase "lengths not a multiple of 4" $
  forM_ (filter (\l -> l `rem` 4 /= 0) [1 .. 99 :: Int]) $ \l ->
    forM_ decoders $ \(nam, dec) ->
      H.assertEqual (nam <> ", length " <> show l) Nothing
        (dec (B8.replicate l 'A'))

main :: IO ()
main = defaultMain $
  testGroup "ppad-base64" [
    testGroup "property tests" [
      Q.testProperty "encode matches reference" $
        Q.withMaxSuccess 5000 encode_matches_reference
    , Q.testProperty "decode . encode ~ id" $
        Q.withMaxSuccess 5000 decode_inverts_encode
    , Q.testProperty "decode matches reference" $
        Q.withMaxSuccess 10000 decode_matches_reference
    ]
  , testGroup "unit tests" [
      case_rfc_vectors
    , malformed
    , every_byte_every_position
    , single_corruption
    , bad_lengths
    ]
  ]