packages feed

ppad-bolt4-0.1.0: test/Main.hs

{-# LANGUAGE OverloadedStrings #-}

module Main where

import Control.Monad (forM_)
import qualified Crypto.Cipher.ChaCha20 as ChaCha
import qualified Crypto.Curve.Secp256k1 as Secp256k1
import qualified Crypto.Hash.SHA256 as SHA256
import Data.Bits (xor)
import qualified Data.ByteString as BS
import qualified Data.ByteString.Base16 as B16
import Data.List (nubBy)
import Data.Word (Word8, Word16, Word32, Word64)
import qualified Lightning.Protocol.BOLT1 as BOLT1
import Lightning.Protocol.BOLT4
import qualified Lightning.Protocol.BOLT9 as BOLT9
import Test.Tasty
import Test.Tasty.HUnit
import Test.Tasty.QuickCheck
import qualified Vectors as V

main :: IO ()
main = defaultMain $ testGroup "ppad-bolt4" [
    key_tests
  , codec_tests
  , failure_tests
  , construct_tests
  , process_tests
  , blinding_tests
  , error_tests
  , testGroup "spec vectors" [
        onion_vectors
      , error_vectors
      , trace_vectors
      , route_blinding_vectors
      , blinded_payment_vectors
      ]
  , properties
  ]

-- helpers --------------------------------------------------------------------

demand :: String -> Maybe a -> IO a
demand _ (Just a) = pure a
demand msg Nothing = assertFailure msg

right :: Show e => String -> Either e a -> IO a
right _ (Right a) = pure a
right msg (Left e) = assertFailure (msg ++ ": " ++ show e)

unhex :: BS.ByteString -> IO BS.ByteString
unhex = demand "invalid hex" . B16.decode

hex_key :: BS.ByteString -> IO SecretKey
hex_key h = unhex h >>= demand "invalid secret key" . secret_key

hex_point :: BS.ByteString -> IO BOLT1.Point
hex_point h = unhex h >>= demand "invalid point" . BOLT1.point

hex_secret :: BS.ByteString -> IO SharedSecret
hex_secret h = unhex h >>= demand "invalid shared secret" . shared_secret

-- strip a BigSize length prefix, checking that it is exact
unprefix :: BS.ByteString -> IO BS.ByteString
unprefix bs = do
  (len, body) <- demand "invalid length prefix" (BOLT1.decode_bigsize bs)
  assertEqual "length prefix" (fromIntegral len) (BS.length body)
  pure body

msat :: Word64 -> IO BOLT1.MilliSatoshi
msat = demand "invalid amount" . BOLT1.milli_satoshi

scid :: Word32 -> Word32 -> Word16 -> IO BOLT1.ShortChannelId
scid b t o = demand "invalid scid" (BOLT1.short_channel_id b t o)

tlvs :: [BOLT1.TlvRecord] -> IO BOLT1.TlvStream
tlvs = demand "invalid tlv stream" . BOLT1.tlv_stream

no_tlvs :: BOLT1.TlvStream
no_tlvs = BOLT1.empty_tlv_stream

key :: Word8 -> IO SecretKey
key b = demand "invalid secret key" (secret_key (BS.replicate 32 b))

ad :: BS.ByteString
ad = BS.replicate 32 0x42

-- a point with a valid prefix but not on the curve
off_curve :: IO BOLT1.Point
off_curve =
  case [ p | x <- [0 .. 255 :: Word8]
           , let bs = BS.cons 0x02 (BS.replicate 31 0 <> BS.singleton x)
           , Secp256k1.parse_point bs == Nothing
           , Just p <- [BOLT1.point bs] ] of
    p : _ -> pure p
    []    -> assertFailure "no off-curve point found"

-- a payload with the given outgoing_cltv_value
cltv_payload :: Word32 -> HopPayload
cltv_payload c = empty_hop_payload { hp_outgoing_cltv_value = Just c }

expect_process
  :: ProcessError -> Either ProcessError ProcessResult -> Assertion
expect_process e r = case r of
  Left e' -> e' @?= e
  Right _ -> assertFailure ("expected " ++ show e)

forward :: Either ProcessError ProcessResult -> IO ForwardInfo
forward r = case r of
  Right (Forward f) -> pure f
  Right (Receive _) -> assertFailure "expected Forward, got Receive"
  Left e -> assertFailure ("process: " ++ show e)

receive :: Either ProcessError ProcessResult -> IO ReceiveInfo
receive r = case r of
  Right (Receive i) -> pure i
  Right (Forward _) -> assertFailure "expected Receive, got Forward"
  Left e -> assertFailure ("process: " ++ show e)

-- A single-hop onion, built independently of the library, whose
-- decrypted hop_payloads is the given 1300-byte plaintext.
raw_onion :: SecretKey -> BS.ByteString -> IO OnionPacket
raw_onion node plain = do
  let session = BS.replicate 32 0x77
  e <- demand "session" (Secp256k1.parse_int256 session)
  epub <- demand "session pub" (Secp256k1.derive_pub e)
  n <- demand "node pub"
         (Secp256k1.parse_point (BOLT1.un_point (public_key node)))
  pt <- demand "ecdh" (Secp256k1.mul n e)
  let ss = SHA256.hash (Secp256k1.serialize_point pt)
      derive l = let SHA256.MAC k = SHA256.hmac l ss in k
  stream <- right "chacha"
    (ChaCha.cipher (derive "rho") 0 (BS.replicate 12 0)
       (BS.replicate 1300 0))
  let body = BS.packZipWith xor plain stream
      SHA256.MAC mac = SHA256.hmac (derive "mu") (body <> ad)
  pk <- demand "point" (BOLT1.point (Secp256k1.serialize_point epub))
  OnionPacket 0 pk
    <$> demand "hop_payloads" (hop_payloads body)
    <*> demand "hmac" (hmac32 mac)

-- pad a plaintext region to 1300 bytes
pad1300 :: BS.ByteString -> BS.ByteString
pad1300 bs = bs <> BS.replicate (1300 - BS.length bs) 0

-- keys -----------------------------------------------------------------------

key_tests :: TestTree
key_tests = testGroup "keys" [
    testCase "secret_key accepts 32-byte keys in range" $ do
      _ <- key 0x01
      sk <- demand "n - 1" . secret_key =<< unhex
        "fffffffffffffffffffffffffffffffebaaedce6af48a03bbfd25e8cd0364140"
      BS.length (BOLT1.un_point (public_key sk)) @?= 33
  , testCase "secret_key rejects other lengths" $ do
      assertNone (secret_key (BS.replicate 31 0x01))
      assertNone (secret_key (BS.replicate 33 0x01))
      assertNone (secret_key (BS.cons 0x00 (BS.replicate 32 0x01)))
      assertNone (secret_key BS.empty)
  , testCase "secret_key rejects zero and keys >= n" $ do
      assertNone (secret_key (BS.replicate 32 0x00))
      n <- unhex
        "fffffffffffffffffffffffffffffffebaaedce6af48a03bbfd25e8cd0364141"
      assertNone (secret_key n)
      assertNone (secret_key (BS.replicate 32 0xff))
  , testCase "public_key matches secp256k1" $ do
      sk <- key 0x41
      expected <- unhex V.onionPubKey0
      BOLT1.un_point (public_key sk) @?= expected
  , testCase "shared_secret requires 32 bytes" $ do
      assertNone (shared_secret (BS.replicate 31 0))
      assertNone (shared_secret (BS.replicate 33 0))
      ss <- demand "ss" (shared_secret (BS.replicate 32 7))
      un_shared_secret ss @?= BS.replicate 32 7
  , testCase "shared secrets compare by value" $ do
      a <- demand "a" (shared_secret (BS.replicate 32 7))
      b <- demand "b" (shared_secret (BS.replicate 32 7))
      c <- demand "c" (shared_secret (BS.cons 8 (BS.replicate 31 7)))
      assertBool "equal" (a == b)
      assertBool "not equal" (a /= c)
  , testCase "secrets don't show their bytes" $ do
      sk <- key 0x41
      ss <- demand "ss" (shared_secret (BS.replicate 32 0x41))
      ps <- demand "ps" (payment_secret (BS.replicate 32 0x41))
      show sk @?= "SecretKey <redacted>"
      show ss @?= "SharedSecret <redacted>"
      show ps @?= "PaymentSecret <redacted>"
  ]
  where
    assertNone m = case m of
      Nothing -> pure ()
      Just _  -> assertFailure "expected Nothing"

-- codecs ---------------------------------------------------------------------

codec_tests :: TestTree
codec_tests = testGroup "codecs" [
    testCase "decode_onion_packet rejects other lengths" $ do
      decode_onion_packet (BS.replicate 1365 0) @?= Left InvalidLength
      decode_onion_packet (BS.replicate 1367 0) @?= Left InvalidLength
      decode_onion_packet BS.empty @?= Left InvalidLength
  , testCase "decode_onion_packet rejects a bad point prefix" $ do
      let bs = BS.concat [ BS.singleton 0, BS.singleton 0x04
                         , BS.replicate 1364 0 ]
      decode_onion_packet bs @?= Left InvalidPoint
  , testCase "decode_onion_packet keeps the version byte" $ do
      let bs = BS.concat [ BS.singleton 7, BS.singleton 0x02
                         , BS.replicate 1364 0 ]
      pkt <- right "decode" (decode_onion_packet bs)
      onion_version pkt @?= 7
      encode_onion_packet pkt @?= bs
  , testCase "smart constructors check lengths" $ do
      hop_payloads (BS.replicate 1299 0) @?= Nothing
      hop_payloads (BS.replicate 1301 0) @?= Nothing
      hmac32 (BS.replicate 31 0) @?= Nothing
      onion_hash (BS.replicate 33 0) @?= Nothing
      case payment_secret (BS.replicate 31 0) of
        Nothing -> pure ()
        Just _  -> assertFailure "short payment secret"

  , testCase "hop payload: all typed fields" $ do
      amt <- msat 1000
      tot <- msat 5000
      s <- scid 1 2 3
      ps <- demand "ps" (payment_secret (BS.replicate 32 0xaa))
      pk <- hex_point V.onionPubKey0
      ex <- tlvs [BOLT1.TlvRecord 19 "x"]
      let hp = HopPayload
            { hp_amt_to_forward      = Just amt
            , hp_outgoing_cltv_value = Just 500000
            , hp_short_channel_id    = Just s
            , hp_payment_data        = Just (PaymentData ps tot)
            , hp_encrypted_data      = Just "enc"
            , hp_current_path_key    = Just pk
            , hp_payment_metadata    = Just "meta"
            , hp_total_amount_msat   = Just tot
            , hp_extra               = ex
            }
      bs <- right "encode" (encode_hop_payload hp)
      expected <- unhex $ BS.concat
        [ "020203e8", "040307a120", "06080000010000020003"
        , "0822", BS.concat (replicate 32 "aa"), "1388"
        , "0a03656e63", "0c21", V.onionPubKey0, "10046d657461"
        , "12021388", "130178" ]
      bs @?= expected
      decode_hop_payload bs @?= Right hp
  , testCase "hop payload: zero values encode minimally" $ do
      zero <- msat 0
      let hp = empty_hop_payload { hp_amt_to_forward = Just zero
                                 , hp_outgoing_cltv_value = Just 0 }
      encode_hop_payload hp @?= Right "\STX\NUL\EOT\NUL"
      decode_hop_payload "\STX\NUL\EOT\NUL" @?= Right hp
  , testCase "hop payload: unknown odd types are kept" $ do
      hp <- right "decode" (decode_hop_payload "\STX\SOH\SOH\ETX\NUL")
      ex <- tlvs [BOLT1.TlvRecord 3 ""]
      hp_extra hp @?= ex
      encode_hop_payload hp @?= Right "\STX\SOH\SOH\ETX\NUL"
  , testCase "hop payload: rejects unknown even types" $ do
      decode_hop_payload "\DC4\NUL" @?=
        Left (InvalidTlvStream (BOLT1.TlvUnknownEvenType 20))
      decode_hop_payload "\SO\NUL" @?=
        Left (InvalidTlvStream (BOLT1.TlvUnknownEvenType 14))
  , testCase "hop payload: rejects malformed streams" $ do
      decode_hop_payload "\EOT\SOH\SOH\STX\SOH\SOH" @?=
        Left (InvalidTlvStream BOLT1.TlvNotStrictlyIncreasing)
      decode_hop_payload "\STX\SOH\SOH\STX\SOH\SOH" @?=
        Left (InvalidTlvStream BOLT1.TlvNotStrictlyIncreasing)
      decode_hop_payload "\STX\ENQ\SOH" @?=
        Left (InvalidTlvStream BOLT1.TlvTruncated)
      decode_hop_payload "\STX\255\255\255\255\255\255\255\255\255" @?=
        Left (InvalidTlvStream BOLT1.TlvTruncated)
      decode_hop_payload "\253\NUL\STX\NUL" @?=
        Left (InvalidTlvStream BOLT1.TlvNonMinimalBigSize)
  , testCase "hop payload: rejects malformed values" $ do
      -- non-minimal tu64 and tu32
      decode_hop_payload "\STX\STX\NUL\SOH" @?= Left (InvalidTlvValue 2)
      decode_hop_payload "\EOT\STX\NUL\SOH" @?= Left (InvalidTlvValue 4)
      -- tu32 longer than 4 bytes
      decode_hop_payload "\EOT\ENQ\SOH\SOH\SOH\SOH\SOH" @?=
        Left (InvalidTlvValue 4)
      -- amount above 21M BTC
      decode_hop_payload "\STX\b\255\255\255\255\255\255\255\255" @?=
        Left (InvalidTlvValue 2)
      -- short_channel_id of 7 bytes
      decode_hop_payload (BS.pack [6, 7] <> BS.replicate 7 0) @?=
        Left (InvalidTlvValue 6)
      -- payment_data with a 31-byte secret
      decode_hop_payload (BS.pack [8, 31] <> BS.replicate 31 0) @?=
        Left (InvalidTlvValue 8)
      -- current_path_key of the wrong length or prefix
      decode_hop_payload (BS.pack [12, 32] <> BS.replicate 32 2) @?=
        Left (InvalidTlvValue 12)
      decode_hop_payload (BS.pack [12, 33, 4] <> BS.replicate 32 2) @?=
        Left (InvalidTlvValue 12)
      -- total_amount_msat non-minimal
      decode_hop_payload "\DC2\STX\NUL\SOH" @?= Left (InvalidTlvValue 18)
  , testCase "hop payload: extra records may not have known types" $ do
      ex <- tlvs [BOLT1.TlvRecord 16 "meta"]
      encode_hop_payload empty_hop_payload { hp_extra = ex } @?=
        Left ConflictingTlv
      ex' <- tlvs [BOLT1.TlvRecord 2 "\SOH"]
      encode_hop_payload (cltv_payload 1) { hp_extra = ex' } @?=
        Left ConflictingTlv
  , testCase "hop payload: extra records may have unknown even types" $ do
      ex <- tlvs [BOLT1.TlvRecord 5482373484 (BS.replicate 32 1)]
      bs <- right "encode" (encode_hop_payload (cltv_payload 1)
                              { hp_extra = ex })
      decode_hop_payload bs @?=
        Left (InvalidTlvStream (BOLT1.TlvUnknownEvenType 5482373484))

  , testCase "blinded hop data: all typed fields" $ do
      s <- scid 0 0 1729
      pk <- hex_point V.rbDavePathKey
      m <- msat 1500
      ex <- tlvs [BOLT1.TlvRecord 561 "\DC24V"]
      let d = BlindedHopData
            { bhd_padding                = Just (BS.replicate 3 0)
            , bhd_short_channel_id       = Just s
            , bhd_next_node_id           = Just pk
            , bhd_path_id                = Just "id"
            , bhd_next_path_key_override = Just pk
            , bhd_payment_relay          = Just (PaymentRelay 36 150 10000)
            , bhd_payment_constraints    =
                Just (PaymentConstraints 748005 m)
            , bhd_allowed_features       = Just (BOLT9.parse "\STX")
            , bhd_extra                  = ex
            }
      bs <- right "encode" (encode_blinded_hop_data d)
      expected <- unhex $ BS.concat
        [ "0103000000", "020800000000000006c1", "0421", V.rbDavePathKey
        , "06026964", "0821", V.rbDavePathKey, "0a080024000000962710"
        , "0c06000b69e505dc", "0e0102", "fd023103123456" ]
      bs @?= expected
      decode_blinded_hop_data bs @?= Right d
  , testCase "blinded hop data: unknown types" $ do
      decode_blinded_hop_data "\DLE\NUL" @?=
        Left (InvalidTlvStream (BOLT1.TlvUnknownEvenType 16))
      d <- right "decode" (decode_blinded_hop_data "\ETX\NUL")
      ex <- tlvs [BOLT1.TlvRecord 3 ""]
      bhd_extra d @?= ex
  , testCase "blinded hop data: rejects malformed values" $ do
      -- payment_relay shorter than its fixed fields
      decode_blinded_hop_data "\n\ENQ\NUL\NUL\NUL\NUL\NUL" @?=
        Left (InvalidTlvValue 10)
      -- payment_relay with a non-minimal fee_base_msat
      decode_blinded_hop_data "\n\a\NUL\NUL\NUL\NUL\NUL\NUL\NUL" @?=
        Left (InvalidTlvValue 10)
      -- payment_constraints shorter than its fixed fields
      decode_blinded_hop_data "\f\ETX\NUL\NUL\NUL" @?=
        Left (InvalidTlvValue 12)
      -- next_node_id and next_path_key_override of the wrong length
      decode_blinded_hop_data (BS.pack [4, 2, 2, 2]) @?=
        Left (InvalidTlvValue 4)
      decode_blinded_hop_data (BS.pack [8, 2, 2, 2]) @?=
        Left (InvalidTlvValue 8)
      -- short_channel_id of 9 bytes
      decode_blinded_hop_data (BS.pack [2, 9] <> BS.replicate 9 0) @?=
        Left (InvalidTlvValue 2)
  , testCase "blinded hop data: extra records may not have known types" $ do
      ex <- tlvs [BOLT1.TlvRecord 1 ""]
      encode_blinded_hop_data empty_blinded_hop_data { bhd_extra = ex } @?=
        Left ConflictingTlv
  ]

-- failure messages -----------------------------------------------------------

failure_tests :: TestTree
failure_tests = testGroup "failure messages" [
    testCase "encodings and codes of every failure" $ do
      h <- demand "hash" (onion_hash (BS.replicate 32 0xab))
      m <- msat 100
      let hh = BS.concat (replicate 32 "ab")
          cases =
            [ (TemporaryNodeFailure, "2002")
            , (PermanentNodeFailure, "6002")
            , (RequiredNodeFeatureMissing, "6003")
            , (InvalidOnionVersion h, "c004" <> hh)
            , (InvalidOnionHmac h, "c005" <> hh)
            , (InvalidOnionKey h, "c006" <> hh)
            , (TemporaryChannelFailure "cu", "100700026375")
            , (PermanentChannelFailure, "4008")
            , (RequiredChannelFeatureMissing, "4009")
            , (UnknownNextPeer, "400a")
            , (AmountBelowMinimum m "", "100b00000000000000640000")
            , (FeeInsufficient m "cu", "100c000000000000006400026375")
            , (IncorrectCltvExpiry 800000 "", "100d000c35000000")
            , (ExpiryTooSoon "", "100e0000")
            , ( IncorrectOrUnknownPaymentDetails m 800000
              , "400f0000000000000064000c3500" )
            , (FinalIncorrectCltvExpiry 800000, "0012000c3500")
            , (FinalIncorrectHtlcAmount m, "00130000000000000064")
            , (ChannelDisabled 0 "", "101400000000")
            , (ExpiryTooFar, "0015")
            , (InvalidOnionPayload Nothing, "4016")
            , (InvalidOnionPayload (Just (2, 7)), "401602 0007")
            , (MppTimeout, "0017")
            , (InvalidOnionBlinding h, "c018" <> hh)
            , (UnknownFailure 0x0042 "data", "004264617461")
            ]
      forM_ cases $ \(f, hx) -> do
        bs <- unhex (BS.filter (/= 0x20) hx)
        let fm = FailureMessage f no_tlvs
        encode_failure_message fm @?= Right bs
        decode_failure_message bs @?= Right fm
        Just (failure_code f) @?= fmap fst (BOLT1.decode_u16 bs)
  , testCase "flags" $ do
      h <- demand "hash" (onion_hash (BS.replicate 32 0))
      let flags f = (is_badonion f, is_perm f, is_node f, is_update f)
      flags TemporaryNodeFailure @?= (False, False, True, False)
      flags PermanentNodeFailure @?= (False, True, True, False)
      flags (InvalidOnionBlinding h) @?= (True, True, False, False)
      flags (ExpiryTooSoon "") @?= (False, False, False, True)
      flags MppTimeout @?= (False, False, False, False)
      flags (UnknownFailure 0xf000 "") @?= (True, True, True, True)
  , testCase "spec example: incorrect_or_unknown_payment_details" $ do
      bs <- unhex "400f0000000000000064000c3500"
      m <- msat 100
      decode_failure_message bs @?= Right
        (FailureMessage (IncorrectOrUnknownPaymentDetails m 800000) no_tlvs)
  , testCase "a trailing TLV stream is kept" $ do
      ex <- tlvs [BOLT1.TlvRecord 34001 (BS.replicate 3 0x80)]
      let fm = FailureMessage TemporaryNodeFailure ex
      bs <- right "encode" (encode_failure_message fm)
      bs @?= "\x20\x02\xfd\x84\xd1\x03\x80\x80\x80"
      decode_failure_message bs @?= Right fm
  , testCase "other trailing bytes are ignored" $ do
      let ignored f bs = decode_failure_message bs @?=
            Right (FailureMessage f no_tlvs)
      ignored TemporaryNodeFailure "\x20\x02\xff"
      -- an even type: not a valid stream without known types
      ignored TemporaryNodeFailure "\x20\x02\x02\x00"
      ignored (FinalIncorrectCltvExpiry 1) "\NUL\DC2\NUL\NUL\NUL\SOH\SOH"
  , testCase "truncated data is rejected" $ do
      decode_failure_message "" @?= Left InvalidLength
      decode_failure_message "\x20" @?= Left InvalidLength
      decode_failure_message "\xc0\x05\x00" @?=
        Left (InvalidFailureData 0xc005)
      decode_failure_message "\x10\x07\x00\x02\x00" @?=
        Left (InvalidFailureData 0x1007)
      decode_failure_message "\x40\x0f\x00\x00\x00\x00\x00\x00\x00\x64" @?=
        Left (InvalidFailureData 0x400f)
      decode_failure_message "\x00\x13\x00" @?=
        Left (InvalidFailureData 0x0013)
      decode_failure_message "\x10\x14\x00\x00\x00" @?=
        Left (InvalidFailureData 0x1014)
  , testCase "amounts above 21M BTC are rejected" $
      decode_failure_message "\x00\x13\xff\xff\xff\xff\xff\xff\xff\xff" @?=
        Left (InvalidFailureData 0x0013)
  , testCase "unknown codes keep their data" $
      decode_failure_message "\x12\x34\x01\x00" @?=
        Right (FailureMessage (UnknownFailure 0x1234 "\x01\x00") no_tlvs)
  , testCase "channel_update longer than 65535 bytes" $
      encode_failure_message
        (FailureMessage (ExpiryTooSoon (BS.replicate 65536 0)) no_tlvs)
        @?= Left FieldTooLong
  ]

-- construct ------------------------------------------------------------------

construct_tests :: TestTree
construct_tests = testGroup "construct" [
    testCase "empty route" $ do
      sk <- key 0x41
      assertConstructError EmptyRoute (construct sk [] ad)
  , testCase "more than 20 hops" $ do
      sk <- key 0x41
      node <- key 0x42
      let hop = Hop (public_key node) (cltv_payload 1)
      assertConstructError TooManyHops (construct sk (replicate 21 hop) ad)
      case construct sk (replicate 20 hop) ad of
        Right (_, sss) -> length sss @?= 20
        Left e -> assertFailure (show e)
  , testCase "invalid hop public key" $ do
      sk <- key 0x41
      node <- key 0x42
      bad <- off_curve
      let good = Hop (public_key node) (cltv_payload 1)
      assertConstructError (InvalidHopPubKey 1)
        (construct sk [good, Hop bad (cltv_payload 1), good] ad)
  , testCase "payloads under 2 bytes" $ do
      sk <- key 0x41
      node <- key 0x42
      let good = Hop (public_key node) (cltv_payload 1)
      assertConstructError (InvalidHopPayload 1)
        (construct sk [good, Hop (public_key node) empty_hop_payload] ad)
  , testCase "unencodable payload" $ do
      sk <- key 0x41
      node <- key 0x42
      ex <- tlvs [BOLT1.TlvRecord 4 "\SOH"]
      let bad = Hop (public_key node) empty_hop_payload { hp_extra = ex }
      assertConstructError (InvalidHopPayload 0) (construct sk [bad] ad)
  , testCase "payloads over 1300 bytes" $ do
      sk <- key 0x41
      node <- key 0x42
      -- a 1265-byte payload exactly fills hop_payloads
      ex <- tlvs [BOLT1.TlvRecord 1 (BS.replicate 1261 0)]
      let full = Hop (public_key node) empty_hop_payload { hp_extra = ex }
      case construct sk [full] ad of
        Right _ -> pure ()
        Left e -> assertFailure (show e)
      ex' <- tlvs [BOLT1.TlvRecord 1 (BS.replicate 1262 0)]
      let over = Hop (public_key node) empty_hop_payload { hp_extra = ex' }
      assertConstructError (PayloadsTooLarge 1301) (construct sk [over] ad)
      let hops = replicate 2 (Hop (public_key node) (cltv_payload 1))
      assertConstructError (PayloadsTooLarge 1372)
        (construct sk (full : hops) ad)
  ]
  where
    assertConstructError e r = case r of
      Left e' -> e' @?= e
      Right _ -> assertFailure ("expected " ++ show e)

-- process --------------------------------------------------------------------

process_tests :: TestTree
process_tests = testGroup "process" [
    testCase "rejects versions other than 0" $ do
      (node, pkt) <- simple_onion
      expect_process (InvalidVersion 1)
        (process node pkt { onion_version = 1 } ad Nothing)
  , testCase "rejects a public key not on the curve" $ do
      (node, pkt) <- simple_onion
      bad <- off_curve
      expect_process InvalidPublicKey
        (process node pkt { onion_public_key = bad } ad Nothing)
  , testCase "rejects a wrong HMAC" $ do
      (node, pkt) <- simple_onion
      flipped <- demand "flip" $
        case BS.uncons (un_hop_payloads (onion_hop_payloads pkt)) of
          Just (h, t) -> hop_payloads (BS.cons (h `xor` 1) t)
          Nothing     -> Nothing
      expect_process HmacMismatch
        (process node pkt { onion_hop_payloads = flipped } ad Nothing)
      mac <- demand "mac" (hmac32 (BS.replicate 32 0))
      expect_process HmacMismatch
        (process node pkt { onion_hmac = mac } ad Nothing)
  , testCase "rejects other associated data" $ do
      (node, pkt) <- simple_onion
      expect_process HmacMismatch
        (process node pkt (BS.replicate 32 0) Nothing)
  , testCase "rejects other node keys" $ do
      (_, pkt) <- simple_onion
      other <- key 0x44
      expect_process HmacMismatch (process other pkt ad Nothing)
  , testCase "rejects payload lengths 0 and 1" $ do
      node <- key 0x42
      p0 <- raw_onion node (pad1300 "\NUL")
      expect_process InvalidPayloadLength (process node p0 ad Nothing)
      p1 <- raw_onion node (pad1300 "\SOH\SOH")
      expect_process InvalidPayloadLength (process node p1 ad Nothing)
  , testCase "rejects a malformed length prefix" $ do
      node <- key 0x42
      p <- raw_onion node (pad1300 "\253\NUL\DLE")
      expect_process InvalidPayloadLength (process node p ad Nothing)
  , testCase "rejects payloads past the end of hop_payloads" $ do
      node <- key 0x42
      -- 3-byte prefix, 1265-byte payload and 32-byte HMAC fill it
      let rec n = BS.concat [ "\SOH\253", BOLT1.encode_u16 n
                            , BS.replicate (fromIntegral n) 0 ]
      ok <- raw_onion node (pad1300 ("\253\EOT\241" <> rec 1261))
      _ <- receive (process node ok ad Nothing)
      over <- raw_onion node (pad1300 ("\253\EOT\242" <> rec 1262))
      expect_process InvalidPayloadLength (process node over ad Nothing)
      -- a length reaching into the zero-extended stream
      far <- raw_onion node (pad1300 ("\253\a\208" <> rec 1261))
      expect_process InvalidPayloadLength (process node far ad Nothing)
      huge <- raw_onion node
                (pad1300 ("\255\128\NUL\NUL\NUL\NUL\NUL\NUL\NUL" <> rec 1))
      expect_process InvalidPayloadLength (process node huge ad Nothing)
  , testCase "rejects malformed payloads" $ do
      node <- key 0x42
      let bad pl = raw_onion node
            (pad1300 (BOLT1.encode_bigsize (fromIntegral (BS.length pl))
                        <> pl))
      p0 <- bad "\EOT\SOH\SOH\STX\SOH\SOH"
      expect_process
        (InvalidPayload (InvalidTlvStream BOLT1.TlvNotStrictlyIncreasing))
        (process node p0 ad Nothing)
      p1 <- bad "\DC4\NUL"
      expect_process
        (InvalidPayload (InvalidTlvStream (BOLT1.TlvUnknownEvenType 20)))
        (process node p1 ad Nothing)
      p2 <- bad "\STX\STX\NUL\SOH"
      expect_process (InvalidPayload (InvalidTlvValue 2))
        (process node p2 ad Nothing)
  , testCase "the final hop receives" $ do
      (node, pkt) <- simple_onion
      info <- receive (process node pkt ad Nothing)
      rcv_payload info @?= cltv_payload 144
      rcv_blinded info @?= Nothing
  , testCase "forwarded packets are well formed" $ do
      sk <- key 0x41
      a <- key 0x42
      b <- key 0x43
      (pkt, sss) <- right "construct" $ construct sk
        [Hop (public_key a) (cltv_payload 1), Hop (public_key b)
           (cltv_payload 2)] ad
      fwd <- forward (process a pkt ad Nothing)
      let next = fwd_next_packet fwd
      onion_version next @?= 0
      decode_onion_packet (encode_onion_packet next) @?= Right next
      Just (fwd_shared_secret fwd) @?= safe_head sss
      rcv <- receive (process b next ad Nothing)
      rcv_payload rcv @?= cltv_payload 2
  , testCase "path keys without encrypted_recipient_data" $ do
      sk <- key 0x41
      alice <- key 0x42
      seed <- key 0x01
      path <- right "path" $
        create_blinded_path seed [(public_key alice, empty_blinded_hop_data)]
      bid <- blinded_id path 0
      let pk = bp_first_path_key path
      -- a path_key, onion encrypted to the blinded id
      (p0, _) <- right "construct" $
        construct sk [Hop bid (cltv_payload 1)] ad
      expect_process UnexpectedPathKey (process alice p0 ad (Just pk))
      -- a current_path_key
      let pl = (cltv_payload 1) { hp_current_path_key = Just pk }
      (p1, _) <- right "construct" $
        construct sk [Hop (public_key alice) pl] ad
      expect_process UnexpectedPathKey (process alice p1 ad Nothing)
  , testCase "both a path_key and a current_path_key" $ do
      sk <- key 0x41
      alice <- key 0x42
      seed <- key 0x01
      path <- right "path" $
        create_blinded_path seed [(public_key alice, empty_blinded_hop_data)]
      bid <- blinded_id path 0
      enc <- blinded_data path 0
      let pk = bp_first_path_key path
          pl = empty_hop_payload { hp_encrypted_data = Just enc
                                 , hp_current_path_key = Just pk }
      (pkt, _) <- right "construct" $ construct sk [Hop bid pl] ad
      expect_process UnexpectedPathKey (process alice pkt ad (Just pk))
  , testCase "encrypted_recipient_data without a path key" $ do
      sk <- key 0x41
      alice <- key 0x42
      let pl = empty_hop_payload { hp_encrypted_data = Just "data" }
      (pkt, _) <- right "construct" $
        construct sk [Hop (public_key alice) pl] ad
      expect_process MissingPathKey (process alice pkt ad Nothing)
  , testCase "undecryptable encrypted_recipient_data" $ do
      sk <- key 0x41
      alice <- key 0x42
      pk <- hex_point V.rbBobPathKey
      let pl = empty_hop_payload { hp_encrypted_data = Just
                                     (BS.replicate 40 0)
                                 , hp_current_path_key = Just pk }
      (pkt, _) <- right "construct" $
        construct sk [Hop (public_key alice) pl] ad
      expect_process InvalidRecipientData (process alice pkt ad Nothing)
  , testCase "invalid path keys" $ do
      sk <- key 0x41
      alice <- key 0x42
      bad <- off_curve
      let pl = empty_hop_payload { hp_encrypted_data = Just "data"
                                 , hp_current_path_key = Just bad }
      (pkt, _) <- right "construct" $
        construct sk [Hop (public_key alice) pl] ad
      expect_process InvalidPathKey (process alice pkt ad Nothing)
      (pkt', _) <- right "construct" $
        construct sk [Hop (public_key alice) (cltv_payload 1)] ad
      expect_process InvalidPathKey (process alice pkt' ad (Just bad))
  ]
  where
    simple_onion = do
      sk <- key 0x41
      node <- key 0x42
      (pkt, _) <- right "construct" $
        construct sk [Hop (public_key node) (cltv_payload 144)] ad
      pure (node, pkt)

three :: [a] -> IO (a, a, a)
three xs = case xs of
  [a, b, c] -> pure (a, b, c)
  _ -> assertFailure "expected three elements"

safe_head :: [a] -> Maybe a
safe_head xs = case xs of
  x : _ -> Just x
  []    -> Nothing

safe_index :: [a] -> Int -> Maybe a
safe_index xs i = safe_head (drop i xs)

blinded_id :: BlindedPath -> Int -> IO BOLT1.Point
blinded_id p i =
  demand "blinded hop" (bh_blinded_node_id <$> safe_index (bp_hops p) i)

blinded_data :: BlindedPath -> Int -> IO BS.ByteString
blinded_data p i =
  demand "blinded hop" (bh_encrypted_data <$> safe_index (bp_hops p) i)

-- route blinding -------------------------------------------------------------

blinding_tests :: TestTree
blinding_tests = testGroup "route blinding" [
    testCase "empty path" $ do
      seed <- key 0x01
      create_blinded_path seed [] @?= Left EmptyPath
  , testCase "invalid node id" $ do
      seed <- key 0x01
      alice <- key 0x42
      bad <- off_curve
      create_blinded_path seed
        [ (public_key alice, empty_blinded_hop_data)
        , (bad, empty_blinded_hop_data) ]
        @?= Left (InvalidNodeId 1)
  , testCase "unencodable hop data" $ do
      seed <- key 0x01
      alice <- key 0x42
      ex <- tlvs [BOLT1.TlvRecord 6 "id"]
      create_blinded_path seed
        [(public_key alice, empty_blinded_hop_data { bhd_extra = ex })]
        @?= Left (InvalidHopData 0 ConflictingTlv)
  , testCase "decrypt_recipient_data rejects bad input" $ do
      seed <- key 0x01
      alice <- key 0x42
      bob <- key 0x43
      path <- right "path" $
        create_blinded_path seed [(public_key alice, empty_blinded_hop_data)]
      enc <- blinded_data path 0
      let pk = bp_first_path_key path
      _ <- right "decrypt" (decrypt_recipient_data alice pk enc)
      decrypt_recipient_data bob pk enc @?= Left InvalidRecipientData
      decrypt_recipient_data alice pk (BS.take 15 enc) @?=
        Left InvalidRecipientData
      decrypt_recipient_data alice pk (BS.drop 1 enc) @?=
        Left InvalidRecipientData
      bad <- off_curve
      decrypt_recipient_data alice bad enc @?= Left InvalidPathKey
  , testCase "unknown even types in encrypted data are rejected" $ do
      seed <- key 0x01
      alice <- key 0x42
      ex <- tlvs [BOLT1.TlvRecord 20 ""]
      path <- right "path" $ create_blinded_path seed
        [(public_key alice, empty_blinded_hop_data { bhd_extra = ex })]
      enc <- blinded_data path 0
      decrypt_recipient_data alice (bp_first_path_key path) enc @?=
        Left InvalidRecipientData
  , testCase "a payment through a blinded route" $ do
      -- Carol creates a route Alice -> Bob -> Carol; Dave pays through it,
      -- reaching Alice unblinded with current_path_key.
      seed <- key 0x01
      session <- key 0x02
      alice <- key 0x42
      bob <- key 0x43
      carol <- key 0x44
      s1 <- scid 1 1 1
      s2 <- scid 2 2 2
      m <- msat 1
      let relay = Just (PaymentRelay 40 100 1000)
          cons = Just (PaymentConstraints 900000 m)
          alice_data = empty_blinded_hop_data
            { bhd_short_channel_id = Just s1, bhd_payment_relay = relay
            , bhd_payment_constraints = cons }
          bob_data = empty_blinded_hop_data
            { bhd_short_channel_id = Just s2, bhd_payment_relay = relay
            , bhd_payment_constraints = cons }
          carol_data = empty_blinded_hop_data
            { bhd_path_id = Just "invoice" }
      path <- right "path" $ create_blinded_path seed
        [ (public_key alice, alice_data), (public_key bob, bob_data)
        , (public_key carol, carol_data) ]
      bp_first_node_id path @?= public_key alice
      bp_first_path_key path @?= public_key seed
      (e0, e1, e2) <- three =<< mapM (blinded_data path) [0, 1, 2]
      (_, b1, b2) <- three =<< mapM (blinded_id path) [0, 1, 2]
      amt <- msat 5000
      let intro = empty_hop_payload
            { hp_encrypted_data = Just e0
            , hp_current_path_key = Just (bp_first_path_key path) }
          final = empty_hop_payload
            { hp_encrypted_data = Just e2, hp_amt_to_forward = Just amt
            , hp_outgoing_cltv_value = Just 800000
            , hp_total_amount_msat = Just amt }
          route = [ Hop (public_key alice) intro
                  , Hop b1 empty_hop_payload { hp_encrypted_data = Just e1 }
                  , Hop b2 final ]
      (pkt, sss) <- right "construct" (construct session route ad)
      fa <- forward (process alice pkt ad Nothing)
      ia <- demand "alice blinded" (fwd_blinded fa)
      bi_data ia @?= alice_data
      fb <- forward (process bob (fwd_next_packet fa) ad
                       (Just (bi_next_path_key ia)))
      ib <- demand "bob blinded" (fwd_blinded fb)
      bi_data ib @?= bob_data
      rc <- receive (process carol (fwd_next_packet fb) ad
                       (Just (bi_next_path_key ib)))
      ic <- demand "carol blinded" (rcv_blinded rc)
      bi_data ic @?= carol_data
      rcv_payload rc @?= final
      [fwd_shared_secret fa, fwd_shared_secret fb, rcv_shared_secret rc]
        @?= sss
      -- the wrong path key yields the wrong blinded key
      expect_process HmacMismatch
        (process bob (fwd_next_packet fa) ad
           (Just (bp_first_path_key path)))
  ]

-- returning errors -----------------------------------------------------------

error_tests :: TestTree
error_tests = testGroup "returning errors" [
    testCase "padding brings failure_len + pad_len to 256" $ do
      ss <- demand "ss" (shared_secret (BS.replicate 32 1))
      let fm = FailureMessage TemporaryNodeFailure no_tlvs
      pkt <- right "construct_error" (construct_error ss fm)
      let ErrorPacket raw = wrap_error ss pkt
      BS.length raw @?= 32 + 2 + 2 + 2 + 254
      BS.take 4 (BS.drop 32 raw) @?= "\NUL\STX\x20\x02"
      BS.take 2 (BS.drop 36 raw) @?= "\NUL\254"
  , testCase "long failure messages are not padded" $ do
      ss <- demand "ss" (shared_secret (BS.replicate 32 1))
      ex <- tlvs [BOLT1.TlvRecord 1 (BS.replicate 300 0)]
      let fm = FailureMessage TemporaryNodeFailure ex
      ErrorPacket raw <- wrap_error ss <$>
        right "construct_error" (construct_error ss fm)
      BS.length raw @?= 32 + 2 + 306 + 2
      BS.drop (32 + 2 + 306) raw @?= "\NUL\NUL"
  , testCase "failure messages over 65535 bytes" $ do
      ss <- demand "ss" (shared_secret (BS.replicate 32 1))
      ex <- tlvs [BOLT1.TlvRecord 1 (BS.replicate 65535 0)]
      construct_error ss (FailureMessage TemporaryNodeFailure ex)
        @?= Left FieldTooLong
  , testCase "attributes the failing hop" $ do
      sss <- mapM (\b -> demand "ss" (shared_secret (BS.replicate 32 b)))
               [1 .. 5]
      let fm = FailureMessage PermanentChannelFailure no_tlvs
      forM_ (zip [0 ..] sss) $ \(i, ss) -> do
        pkt <- right "construct_error" (construct_error ss fm)
        let back = foldr wrap_error pkt (take i sss)
        unwrap_error sss back @?= Attributed i fm
  , testCase "reports the hop of a malformed failure" $ do
      sss <- mapM (\b -> demand "ss" (shared_secret (BS.replicate 32 b)))
               [1 .. 3]
      ss2 <- demand "ss2" (safe_index sss 2)
      -- incorrect_cltv_expiry with no data
      let fm = FailureMessage (UnknownFailure 0x100d "") no_tlvs
      pkt <- right "construct_error" (construct_error ss2 fm)
      let back = foldr wrap_error pkt (take 2 sss)
      unwrap_error sss back @?= MalformedFailure 2
  , testCase "reports a malformed return packet" $ do
      ss <- demand "ss" (shared_secret (BS.replicate 32 1))
      -- a valid HMAC over a body whose failure_len overruns it
      let body = "\SOH\NUL\x20\x02"
          SHA256.MAC um = SHA256.hmac "um" (un_shared_secret ss)
          SHA256.MAC mac = SHA256.hmac um body
          pkt = wrap_error ss (ErrorPacket (mac <> body))
      unwrap_error [ss] pkt @?= MalformedFailure 0
  , testCase "unknown origin" $ do
      sss <- mapM (\b -> demand "ss" (shared_secret (BS.replicate 32 b)))
               [1 .. 3]
      other <- demand "ss" (shared_secret (BS.replicate 32 9))
      let fm = FailureMessage TemporaryNodeFailure no_tlvs
      pkt <- right "construct_error" (construct_error other fm)
      unwrap_error sss pkt @?= UnknownOrigin
      unwrap_error [] pkt @?= UnknownOrigin
      unwrap_error sss (ErrorPacket "short") @?= UnknownOrigin
  ]

-- spec vectors: onion-test.json ----------------------------------------------

onion_hops :: IO [(SecretKey, BOLT1.Point, BS.ByteString)]
onion_hops = mapM hop
    [ (V.onionNodeKey0, V.onionPubKey0, V.onionPayload0)
    , (V.onionNodeKey1, V.onionPubKey1, V.onionPayload1)
    , (V.onionNodeKey2, V.onionPubKey2, V.onionPayload2)
    , (V.onionNodeKey3, V.onionPubKey3, V.onionPayload3)
    , (V.onionNodeKey4, V.onionPubKey4, V.onionPayload4)
    ]
  where
    hop (k, p, pl) = (,,) <$> hex_key k <*> hex_point p
                          <*> (unhex pl >>= unprefix)

-- the shared secrets of the onion-test.json route, as listed in
-- onion-error-test.json
route_secrets :: IO [SharedSecret]
route_secrets = mapM hex_secret
  [ V.errSharedSecret0, V.errSharedSecret1, V.errSharedSecret2
  , V.errSharedSecret3, V.errSharedSecret4 ]

onion_vectors :: TestTree
onion_vectors = testGroup "onion-test.json" [
    testCase "construct reproduces the onion" $ do
      hops <- onion_hops
      sk <- hex_key V.onionSessionKey
      assoc <- unhex V.onionAssocData
      expected <- unhex V.onionPacket
      sss <- route_secrets
      route <- mapM (\(_, pub, body) -> do
                 hp <- right "decode_hop_payload" (decode_hop_payload body)
                 encode_hop_payload hp @?= Right body
                 pure (Hop pub hp)) hops
      (pkt, secrets) <- right "construct" (construct sk route assoc)
      encode_onion_packet pkt @?= expected
      secrets @?= sss
  , testCase "process peels the onion at every hop" $ do
      hops <- onion_hops
      assoc <- unhex V.onionAssocData
      sss <- route_secrets
      pkt <- unhex V.onionPacket >>= right "decode" . decode_onion_packet
      let go [] _ = assertFailure "no hops"
          go [((k, _, body), ss)] p = do
            r <- receive (process k p assoc Nothing)
            encode_hop_payload (rcv_payload r) @?= Right body
            rcv_shared_secret r @?= ss
          go (((k, _, body), ss) : rest) p = do
            f <- forward (process k p assoc Nothing)
            encode_hop_payload (fwd_payload f) @?= Right body
            fwd_shared_secret f @?= ss
            fwd_blinded f @?= Nothing
            go rest (fwd_next_packet f)
      go (zip hops sss) pkt
  ]

-- spec vectors: onion-error-test.json ----------------------------------------

-- the failing node's shared secret, and those of the hops back to the
-- origin, from hop 3 to hop 0
error_wrappers :: IO (SharedSecret, [SharedSecret])
error_wrappers = do
  sss <- route_secrets
  case reverse sss of
    ss4 : rest -> pure (ss4, rest)
    []         -> assertFailure "no shared secrets"

error_vectors :: TestTree
error_vectors = testGroup "onion-error-test.json" [
    testCase "failure message" $ do
      msg <- unhex V.errFailureMessage
      let fm = FailureMessage TemporaryNodeFailure no_tlvs
      encode_failure_message fm @?= Right msg
      decode_failure_message msg @?= Right fm
  , testCase "the failing node's return packet" $ do
      (ss4, _) <- error_wrappers
      payload <- unhex V.errPayload4
      let fm = FailureMessage TemporaryNodeFailure no_tlvs
      pkt <- right "construct_error" (construct_error ss4 fm)
      -- removing hop 4's obfuscation leaves the HMAC and payload
      let ErrorPacket raw = wrap_error ss4 pkt
      BS.drop 32 raw @?= payload
  , testCase "wrapping at each hop yields the returned packet" $ do
      (ss4, wrappers) <- error_wrappers
      expected <- unhex V.errPacket
      let fm = FailureMessage TemporaryNodeFailure no_tlvs
      pkt <- right "construct_error" (construct_error ss4 fm)
      foldl (flip wrap_error) pkt wrappers @?= ErrorPacket expected
  , testCase "the origin attributes the packet to hop 4" $ do
      sss <- route_secrets
      pkt <- unhex V.errPacket
      unwrap_error sss (ErrorPacket pkt) @?=
        Attributed 4 (FailureMessage TemporaryNodeFailure no_tlvs)
  ]

-- spec vectors: 'Returning Errors' trace -------------------------------------

trace_failure :: IO FailureMessage
trace_failure = do
  m <- msat 100
  ex <- tlvs [BOLT1.TlvRecord 34001 (BS.replicate 300 0x80)]
  pure (FailureMessage (IncorrectOrUnknownPaymentDetails m 800000) ex)

trace_vectors :: TestTree
trace_vectors = testGroup "returning errors trace" [
    testCase "failure message encoding" $ do
      fm <- trace_failure
      -- the trace lists failuremsg || pad_len || pad
      encoded <- unhex V.traceFailureMessage
      msg <- right "encode" (encode_failure_message fm)
      BS.length msg @?= 320
      BS.take 320 encoded @?= msg
      decode_failure_message msg @?= Right fm
  , testCase "each hop's packet matches the trace" $ do
      (ss4, wrappers) <- error_wrappers
      raw <- unhex V.traceRawPacket
      expected <- mapM unhex
        [ V.tracePacket4, V.tracePacket3, V.tracePacket2
        , V.tracePacket1, V.tracePacket0 ]
      let packets = drop 1 $
            scanl (flip wrap_error) (ErrorPacket raw) (ss4 : wrappers)
      packets @?= map ErrorPacket expected
  , testCase "the origin attributes the packet to hop 4" $ do
      fm <- trace_failure
      sss <- route_secrets
      pkt <- unhex V.tracePacket0
      unwrap_error sss (ErrorPacket pkt) @?= Attributed 4 fm
  ]

-- spec vectors: route-blinding-test.json -------------------------------------

empty_features :: Maybe BOLT9.FeatureVector
empty_features = Just (BOLT9.parse BS.empty)

rb_bob_data :: IO BlindedHopData
rb_bob_data = do
  s <- scid 0 0 1729
  m <- msat 1500
  ex <- tlvs [BOLT1.TlvRecord 561 "\DC24V"]
  pure empty_blinded_hop_data
    { bhd_padding = Just (BS.replicate 26 0)
    , bhd_short_channel_id = Just s
    , bhd_payment_relay = Just (PaymentRelay 36 150 10000)
    , bhd_payment_constraints = Just (PaymentConstraints 748005 m)
    , bhd_allowed_features = empty_features
    , bhd_extra = ex
    }

rb_carol_data :: IO BlindedHopData
rb_carol_data = do
  s <- scid 0 0 1105
  m <- msat 1500
  override <- hex_point V.rbDavePathKey
  pure empty_blinded_hop_data
    { bhd_short_channel_id = Just s
    , bhd_next_path_key_override = Just override
    , bhd_payment_relay = Just (PaymentRelay 48 100 500)
    , bhd_payment_constraints = Just (PaymentConstraints 747969 m)
    , bhd_allowed_features = empty_features
    }

rb_dave_data :: IO BlindedHopData
rb_dave_data = do
  s <- scid 0 0 561
  m <- msat 1500
  pure empty_blinded_hop_data
    { bhd_padding = Just (BS.replicate 35 0)
    , bhd_short_channel_id = Just s
    , bhd_payment_relay = Just (PaymentRelay 144 250 0)
    , bhd_payment_constraints = Just (PaymentConstraints 747921 m)
    , bhd_allowed_features = empty_features
    }

rb_eve_data :: IO BlindedHopData
rb_eve_data = do
  m <- msat 1500
  ex <- tlvs [BOLT1.TlvRecord 65535 "\ACK\193"]
  pure empty_blinded_hop_data
    { bhd_padding = Just (BS.replicate 26 0)
    , bhd_path_id = Just "\222\173\190\239"
    , bhd_payment_constraints = Just (PaymentConstraints 747777 m)
    , bhd_allowed_features =
        Just (BOLT9.parse (BS.cons 0x02 (BS.replicate 14 0)))
    , bhd_extra = ex
    }

-- a hop of route-blinding-test.json
data RbHop = RbHop
  { rb_data          :: !BlindedHopData
  , rb_node_id       :: !BOLT1.Point
  , rb_tlvs          :: !BS.ByteString
  , rb_path_key      :: !BOLT1.Point
  , rb_encrypted     :: !BS.ByteString
  , rb_blinded_id    :: !BOLT1.Point
  , rb_node_key      :: !SecretKey
  , rb_next_path_key :: !BOLT1.Point
  }

rb_hops :: IO [RbHop]
rb_hops = sequence
    [ hop rb_bob_data
        ( V.rbBobNodeId, V.rbBobTlvs, V.rbBobPathKey
        , V.rbBobEncryptedData, V.rbBobBlindedNodeId )
        (V.rbBobNodeKey, V.rbBobNextPathKey)
    , hop rb_carol_data
        ( V.rbCarolNodeId, V.rbCarolTlvs, V.rbCarolPathKey
        , V.rbCarolEncryptedData, V.rbCarolBlindedNodeId )
        (V.rbCarolNodeKey, V.rbCarolNextPathKey)
    , hop rb_dave_data
        ( V.rbDaveNodeId, V.rbDaveTlvs, V.rbDavePathKey
        , V.rbDaveEncryptedData, V.rbDaveBlindedNodeId )
        (V.rbDaveNodeKey, V.rbDaveNextPathKey)
    , hop rb_eve_data
        ( V.rbEveNodeId, V.rbEveTlvs, V.rbEvePathKey
        , V.rbEveEncryptedData, V.rbEveBlindedNodeId )
        (V.rbEveNodeKey, V.rbEveNextPathKey)
    ]
  where
    hop mk (nid, tl, pk, enc, bid) (nk, npk) =
      RbHop <$> mk <*> hex_point nid <*> unhex tl <*> hex_point pk
            <*> unhex enc <*> hex_point bid <*> hex_key nk
            <*> hex_point npk

route_blinding_vectors :: TestTree
route_blinding_vectors = testGroup "route-blinding-test.json" [
    testCase "encrypted_data_tlv encodings" $ do
      hops <- rb_hops
      forM_ hops $ \h -> do
        encode_blinded_hop_data (rb_data h) @?= Right (rb_tlvs h)
        decode_blinded_hop_data (rb_tlvs h) @?= Right (rb_data h)
  , testCase "create_blinded_path reproduces the route" $ do
      hops <- rb_hops
      (bob, carol, dave, eve) <- case hops of
        [b, c, d, e] -> pure (b, c, d, e)
        _ -> assertFailure "expected 4 hops"
      -- Bob creates Bob -> Carol, Eve creates Dave -> Eve
      bob_seed <- hex_key V.rbBobPathPrivKey
      dave_seed <- hex_key V.rbDavePathPrivKey
      let check seed hs = do
            path <- right "create_blinded_path" $ create_blinded_path seed
              [ (rb_node_id h, rb_data h) | h <- hs ]
            case hs of
              h : _ -> do
                bp_first_node_id path @?= rb_node_id h
                bp_first_path_key path @?= rb_path_key h
              [] -> assertFailure "no hops"
            bp_hops path @?=
              [ BlindedHop (rb_blinded_id h) (rb_encrypted h) | h <- hs ]
      check bob_seed [bob, carol]
      check dave_seed [dave, eve]
  , testCase "each hop decrypts its data and next path key" $ do
      hops <- rb_hops
      forM_ hops $ \h -> do
        info <- right "decrypt_recipient_data" $
          decrypt_recipient_data (rb_node_key h) (rb_path_key h)
            (rb_encrypted h)
        bi_data info @?= rb_data h
        encode_blinded_hop_data (bi_data info) @?= Right (rb_tlvs h)
        -- Carol's next path key is her next_path_key_override
        bi_next_path_key info @?= rb_next_path_key h
  , testCase "each hop decrypts an onion to its blinded id" $ do
      -- the onion is encrypted to blinded_node_id, so processing it
      -- requires the blinded private key of the vector
      hops <- rb_hops
      session <- key 0x55
      forM_ hops $ \h -> do
        let pl = empty_hop_payload
                   { hp_encrypted_data = Just (rb_encrypted h) }
        (pkt, sss) <- right "construct" $
          construct session [Hop (rb_blinded_id h) pl] ad
        r <- receive (process (rb_node_key h) pkt ad (Just (rb_path_key h)))
        Just (rcv_shared_secret r) @?= safe_head sss
        fmap bi_next_path_key (rcv_blinded r) @?= Just (rb_next_path_key h)
  ]

-- spec vectors: blinded-payment-onion-test.json ------------------------------

bp_bob_data :: IO BlindedHopData
bp_bob_data = do
  s <- scid 0 0 1
  m <- msat 50
  pure empty_blinded_hop_data
    { bhd_padding = Just (BS.replicate 32 0)
    , bhd_short_channel_id = Just s
    , bhd_payment_relay = Just (PaymentRelay 50 0 10000)
    , bhd_payment_constraints = Just (PaymentConstraints 750150 m)
    , bhd_allowed_features = empty_features
    }

bp_carol_data :: IO BlindedHopData
bp_carol_data = do
  s <- scid 0 0 2
  m <- msat 50
  override <- hex_point V.rbDavePathKey
  pure empty_blinded_hop_data
    { bhd_short_channel_id = Just s
    , bhd_next_path_key_override = Just override
    , bhd_payment_relay = Just (PaymentRelay 75 150 100)
    , bhd_payment_constraints = Just (PaymentConstraints 750100 m)
    , bhd_allowed_features = empty_features
    }

bp_dave_data :: IO BlindedHopData
bp_dave_data = do
  s <- scid 0 0 3
  m <- msat 50
  pure empty_blinded_hop_data
    { bhd_padding = Just (BS.replicate 34 0)
    , bhd_short_channel_id = Just s
    , bhd_payment_relay = Just (PaymentRelay 25 100 0)
    , bhd_payment_constraints = Just (PaymentConstraints 750025 m)
    , bhd_allowed_features = empty_features
    }

bp_eve_data :: IO BlindedHopData
bp_eve_data = do
  m <- msat 50
  path_id <- unhex "c9cf92f45ade68345bc20ae672e2012f4af487ed4415"
  pure empty_blinded_hop_data
    { bhd_padding = Just (BS.replicate 28 0)
    , bhd_path_id = Just path_id
    , bhd_payment_constraints = Just (PaymentConstraints 750000 m)
    , bhd_allowed_features = empty_features
    }

-- per hop: node key, public key the onion is encrypted to, onion
-- received, payload (without its length prefix)
bp_route
  :: IO [(SecretKey, BOLT1.Point, BS.ByteString, BS.ByteString)]
bp_route = mapM hop
    [ (V.onionNodeKey0, V.bpAlicePubKey, V.bpAliceOnion, V.bpAlicePayload)
    , (V.rbBobNodeKey, V.bpBobPubKey, V.bpBobOnion, V.bpBobPayload)
    , (V.rbCarolNodeKey, V.bpCarolPubKey, V.bpCarolOnion, V.bpCarolPayload)
    , (V.rbDaveNodeKey, V.bpDavePubKey, V.bpDaveOnion, V.bpDavePayload)
    , (V.rbEveNodeKey, V.bpEvePubKey, V.bpEveOnion, V.bpEvePayload)
    ]
  where
    hop (k, p, o, pl) = (,,,) <$> hex_key k <*> hex_point p <*> unhex o
                              <*> (unhex pl >>= unprefix)

blinded_payment_vectors :: TestTree
blinded_payment_vectors = testGroup "blinded-payment-onion-test.json" [
    testCase "construct reproduces the onion" $ do
      route <- bp_route
      sk <- hex_key V.bpSessionKey
      assoc <- unhex V.bpAssocData
      expected <- unhex V.bpAliceOnion
      hops <- mapM (\(_, pub, _, body) -> do
                hp <- right "decode_hop_payload" (decode_hop_payload body)
                encode_hop_payload hp @?= Right body
                pure (Hop pub hp)) route
      (pkt, _) <- right "construct" (construct sk hops assoc)
      encode_onion_packet pkt @?= expected
  , testCase "the route's payloads" $ do
      route <- bp_route
      first_pk <- hex_point V.rbBobPathKey
      s <- scid 0 0 10
      (a0, a1, a2) <- three =<< mapM msat [110125, 100000, 150000]
      payloads <- mapM (\(_, _, _, body) ->
        right "decode_hop_payload" (decode_hop_payload body)) route
      case payloads of
        [alice, bob, carol, dave, eve] -> do
          alice @?= empty_hop_payload
            { hp_short_channel_id = Just s, hp_amt_to_forward = Just a0
            , hp_outgoing_cltv_value = Just 749150 }
          hp_current_path_key bob @?= Just first_pk
          forM_ [bob, carol, dave] $ \hp -> do
            assertBool "encrypted data" (hp_encrypted_data hp /= Nothing)
            hp { hp_encrypted_data = Nothing
               , hp_current_path_key = Nothing } @?= empty_hop_payload
          eve { hp_encrypted_data = Nothing } @?= empty_hop_payload
            { hp_amt_to_forward = Just a1, hp_total_amount_msat = Just a2
            , hp_outgoing_cltv_value = Just 749000 }
        _ -> assertFailure "expected 5 payloads"
  , testCase "each hop processes its onion" $ do
      route <- bp_route
      assoc <- unhex V.bpAssocData
      datas <- sequence [bp_bob_data, bp_carol_data, bp_dave_data,
                         bp_eve_data]
      nexts <- mapM hex_point
        [ V.bpBobNextPathKey, V.bpCarolNextPathKey, V.bpDaveNextPathKey
        , V.bpEveNextPathKey ]
      -- (hop, path_key received, expected blinded info)
      let expect = zip3 route
            (Nothing : Nothing : map Just (take 3 nexts))
            (Nothing : map Just (zipWith BlindedInfo datas nexts))
          go [] = assertFailure "no hops"
          go [((k, _, onion, body), mpk, binfo)] = do
            pkt <- right "decode" (decode_onion_packet onion)
            r <- receive (process k pkt assoc mpk)
            encode_hop_payload (rcv_payload r) @?= Right body
            rcv_blinded r @?= binfo
          go (((k, _, onion, body), mpk, binfo) : rest) = do
            pkt <- right "decode" (decode_onion_packet onion)
            f <- forward (process k pkt assoc mpk)
            encode_hop_payload (fwd_payload f) @?= Right body
            fwd_blinded f @?= binfo
            forM_ (take 1 rest) $ \((_, _, next, _), _, _) ->
              encode_onion_packet (fwd_next_packet f) @?= next
            go rest
      go expect
  ]

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

properties :: TestTree
properties = testGroup "properties" [
    testProperty "onion packets round-trip" $
      forAll gen_onion $ \pkt ->
        decode_onion_packet (encode_onion_packet pkt) === Right pkt
  , testProperty "hop payloads round-trip" $
      forAll gen_hop_payload $ \hp ->
        fmap decode_hop_payload (encode_hop_payload hp) === Right (Right hp)
  , testProperty "blinded hop data round-trips" $
      forAll gen_blinded_hop_data $ \d ->
        fmap decode_blinded_hop_data (encode_blinded_hop_data d)
          === Right (Right d)
  , testProperty "failure messages round-trip" $
      forAll gen_failure_message $ \fm ->
        fmap decode_failure_message (encode_failure_message fm)
          === Right (Right fm)
  , testProperty "wrap_error is an involution" $
      forAll gen_secret $ \ss ->
      forAll (gen_bytes 0 600) $ \bs ->
        wrap_error ss (wrap_error ss (ErrorPacket bs)) === ErrorPacket bs
  , testProperty "errors are attributed to the failing hop" $
      forAll (choose (1, 20)) $ \n ->
      forAll (vectorOf n gen_secret) $ \sss ->
      forAll (choose (0, n - 1)) $ \i ->
      forAll gen_failure_message $ \fm ->
        case drop i sss of
          ss : _ | Right pkt <- construct_error ss fm ->
            unwrap_error sss (foldr wrap_error pkt (take i sss))
              === Attributed i fm
          _ -> property False
  , testProperty "routes of 1 to 20 hops construct and process" $
      forAll (choose (1, 20)) $ \n ->
      forAll gen_key $ \sk ->
      forAll (vectorOf n gen_key) $ \keys ->
      forAll (vectorOf n (gen_route_payload n)) $ \pls ->
        route_roundtrip sk (zip keys pls)
  ]

-- construct an onion over a route, then process it at every hop
route_roundtrip :: SecretKey -> [(SecretKey, HopPayload)] -> Property
route_roundtrip sk route =
  case construct sk [Hop (public_key k) hp | (k, hp) <- route] ad of
    Left e -> counterexample (show e) False
    Right (pkt, sss) -> length sss === length route .&&. go route sss pkt
  where
    go [(k, hp)] [ss] pkt = case process k pkt ad Nothing of
      Right (Receive r) ->
        rcv_payload r === hp .&&. property (rcv_shared_secret r == ss)
      other -> counterexample (show other) False
    go ((k, hp) : rest) (ss : sss) pkt = case process k pkt ad Nothing of
      Right (Forward f) ->
             fwd_payload f === hp
        .&&. property (fwd_shared_secret f == ss)
        .&&. go rest sss (fwd_next_packet f)
      other -> counterexample (show other) False
    go _ _ _ = property False

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

gen_bytes :: Int -> Int -> Gen BS.ByteString
gen_bytes lo hi = do
  n <- choose (lo, hi)
  BS.pack <$> vectorOf n arbitrary

gen_key :: Gen SecretKey
gen_key = gen_bytes 32 32 `suchThatMap` secret_key

gen_secret :: Gen SharedSecret
gen_secret = gen_bytes 32 32 `suchThatMap` shared_secret

gen_point :: Gen BOLT1.Point
gen_point = do
  prefix <- elements [0x02, 0x03]
  x <- gen_bytes 32 32
  pure (BS.cons prefix x) `suchThatMap` BOLT1.point

gen_msat :: Gen BOLT1.MilliSatoshi
gen_msat = oneof
  [ choose (0, 0xffff), choose (0, BOLT1.un_milli_satoshi
                                     BOLT1.max_milli_satoshi) ]
  `suchThatMap` BOLT1.milli_satoshi

gen_scid :: Gen BOLT1.ShortChannelId
gen_scid = BOLT1.ShortChannelId <$> arbitrary

gen_maybe :: Gen a -> Gen (Maybe a)
gen_maybe g = oneof [pure Nothing, Just <$> g]

-- extra records with odd types not in the given list
gen_extra :: [Word64] -> Gen BOLT1.TlvStream
gen_extra known = do
  n <- choose (0, 3)
  ts <- vectorOf n (oneof [choose (1, 300), choose (65536, 70000)])
  recs <- mapM (\t -> BOLT1.TlvRecord t <$> gen_bytes 0 20)
            [ t | t <- ts, odd t, t `notElem` known ]
  let uniq = nubBy (\a b -> BOLT1.tlv_type a == BOLT1.tlv_type b) recs
  pure uniq `suchThatMap` BOLT1.tlv_stream

gen_onion :: Gen OnionPacket
gen_onion = OnionPacket
  <$> arbitrary
  <*> gen_point
  <*> (gen_bytes 1300 1300 `suchThatMap` hop_payloads)
  <*> (gen_bytes 32 32 `suchThatMap` hmac32)

gen_hop_payload :: Gen HopPayload
gen_hop_payload = HopPayload
  <$> gen_maybe gen_msat
  <*> gen_maybe arbitrary
  <*> gen_maybe gen_scid
  <*> gen_maybe (PaymentData
                   <$> (gen_bytes 32 32 `suchThatMap` payment_secret)
                   <*> gen_msat)
  <*> gen_maybe (gen_bytes 0 50)
  <*> gen_maybe gen_point
  <*> gen_maybe (gen_bytes 0 50)
  <*> gen_maybe gen_msat
  <*> gen_extra []

gen_blinded_hop_data :: Gen BlindedHopData
gen_blinded_hop_data = BlindedHopData
  <$> gen_maybe (gen_bytes 0 40)
  <*> gen_maybe gen_scid
  <*> gen_maybe gen_point
  <*> gen_maybe (gen_bytes 0 40)
  <*> gen_maybe gen_point
  <*> gen_maybe (PaymentRelay <$> arbitrary <*> arbitrary <*> arbitrary)
  <*> gen_maybe (PaymentConstraints <$> arbitrary <*> gen_msat)
  <*> gen_maybe (BOLT9.parse <$> gen_bytes 0 8)
  <*> gen_extra [1]

gen_failure_message :: Gen FailureMessage
gen_failure_message = do
  f <- gen_failure
  ex <- case f of
    -- these take every byte after the code, or may be ambiguous with a
    -- TLV stream
    UnknownFailure _ _      -> pure no_tlvs
    InvalidOnionPayload Nothing -> pure no_tlvs
    _                       -> gen_extra []
  pure (FailureMessage f ex)

gen_failure :: Gen Failure
gen_failure = oneof
  [ pure TemporaryNodeFailure
  , pure PermanentNodeFailure
  , pure RequiredNodeFeatureMissing
  , InvalidOnionVersion <$> hash
  , InvalidOnionHmac <$> hash
  , InvalidOnionKey <$> hash
  , TemporaryChannelFailure <$> update
  , pure PermanentChannelFailure
  , pure RequiredChannelFeatureMissing
  , pure UnknownNextPeer
  , AmountBelowMinimum <$> gen_msat <*> update
  , FeeInsufficient <$> gen_msat <*> update
  , IncorrectCltvExpiry <$> arbitrary <*> update
  , ExpiryTooSoon <$> update
  , IncorrectOrUnknownPaymentDetails <$> gen_msat <*> arbitrary
  , FinalIncorrectCltvExpiry <$> arbitrary
  , FinalIncorrectHtlcAmount <$> gen_msat
  , ChannelDisabled <$> arbitrary <*> update
  , pure ExpiryTooFar
  , InvalidOnionPayload <$> gen_maybe ((,) <$> arbitrary <*> arbitrary)
  , pure MppTimeout
  , InvalidOnionBlinding <$> hash
  , UnknownFailure <$> (arbitrary `suchThat` unknown) <*> gen_bytes 0 20
  ]
  where
    hash = gen_bytes 32 32 `suchThatMap` onion_hash
    update = gen_bytes 0 100
    unknown c = c `notElem`
      [ 0x2002, 0x6002, 0x6003, 0xc004, 0xc005, 0xc006, 0x1007, 0x4008
      , 0x4009, 0x400a, 0x100b, 0x100c, 0x100d, 0x100e, 0x400f, 0x0012
      , 0x0013, 0x1014, 0x0015, 0x4016, 0x0017, 0xc018 ]

-- a payload whose shift size fits n hops in 1300 bytes
--
-- The typed fields take at most 26 bytes; an odd record of type 65537
-- adds at most 8 bytes of framing plus its value. With a 3-byte length
-- prefix and 32-byte HMAC a hop then shifts at most 69 bytes plus the
-- value's length.
gen_route_payload :: Int -> Gen HopPayload
gen_route_payload n = do
  amt <- gen_msat
  cltv <- arbitrary
  s <- gen_maybe gen_scid
  let budget = 1300 `div` n - 69
  ex <- if budget < 0
    then pure no_tlvs
    else do
      v <- gen_bytes 0 budget
      pure (BOLT1.TlvRecord 65537 v : [])
        `suchThatMap` BOLT1.tlv_stream
  pure empty_hop_payload
    { hp_amt_to_forward = Just amt
    , hp_outgoing_cltv_value = Just cltv
    , hp_short_channel_id = s
    , hp_extra = ex
    }