packages feed

ppad-bolt4-0.1.0: lib/Lightning/Protocol/BOLT4/Codec.hs

{-# OPTIONS_HADDOCK hide #-}

-- |
-- Module: Lightning.Protocol.BOLT4.Codec
-- Copyright: (c) 2025 Jared Tobin
-- License: MIT
-- Maintainer: Jared Tobin <jared@ppad.tech>
--
-- Encodings of the BOLT #4 types.

module Lightning.Protocol.BOLT4.Codec (
    encode_onion_packet
  , decode_onion_packet
  , encode_hop_payload
  , decode_hop_payload
  , encode_blinded_hop_data
  , decode_blinded_hop_data
  , encode_failure_message
  , decode_failure_message
  ) where

import Data.Bifunctor (first)
import qualified Data.ByteString as BS
import Data.Maybe (catMaybes)
import Data.Word (Word16, Word64)
import qualified Lightning.Protocol.BOLT1 as BOLT1
import Lightning.Protocol.BOLT4.Types
import qualified Lightning.Protocol.BOLT9 as BOLT9

-- onion packets --------------------------------------------------------------

-- | Encode an t'OnionPacket' (1366 bytes).
--
--   >>> let Just pk = BOLT1.point (BS.cons 0x02 (BS.replicate 32 0x01))
--   >>> let Just hp = hop_payloads (BS.replicate 1300 0)
--   >>> let Just mac = hmac32 (BS.replicate 32 0)
--   >>> BS.length (encode_onion_packet (OnionPacket 0 pk hp mac))
--   1366
encode_onion_packet :: OnionPacket -> BS.ByteString
encode_onion_packet (OnionPacket v k hp mac) = BS.concat
  [BS.singleton v, BOLT1.un_point k, un_hop_payloads hp, un_hmac32 mac]

-- | Decode an t'OnionPacket' from exactly 1366 bytes. The public key must
--   carry a compressed-encoding prefix; the version byte is not checked
--   (see 'Lightning.Protocol.BOLT4.process').
--
--   >>> decode_onion_packet (BS.replicate 1365 0)
--   Left InvalidLength
decode_onion_packet :: BS.ByteString -> Either DecodeError OnionPacket
decode_onion_packet bs = case BS.uncons bs of
  Just (v, rest) | BS.length bs == 1366 -> do
    let (k, r0)   = BS.splitAt 33 rest
        (hp, mac) = BS.splitAt 1300 r0
    pk <- maybe (Left InvalidPoint) Right (BOLT1.point k)
    pure (OnionPacket v pk (HopPayloads hp) (Hmac32 mac))
  _ -> Left InvalidLength

-- TLV helpers ----------------------------------------------------------------

-- encode typed records with extra ones, none of which may have a known
-- type
encode_tlvs
  :: [Word64]
  -> [BOLT1.TlvRecord]
  -> BOLT1.TlvStream
  -> Either EncodeError BS.ByteString
encode_tlvs known typed extra
  | any ((`elem` known) . BOLT1.tlv_type) ex = Left ConflictingTlv
  | otherwise = case BOLT1.tlv_stream (typed ++ ex) of
      Nothing -> Left ConflictingTlv
      Just s  -> Right (BOLT1.encode_tlv_stream s)
  where
    ex = BOLT1.un_tlv_stream extra

record :: Word64 -> (a -> BS.ByteString) -> Maybe a -> Maybe BOLT1.TlvRecord
record t f = fmap (BOLT1.TlvRecord t . f)

-- decode a TLV stream with the given known types
decode_tlvs :: [Word64] -> BS.ByteString -> Either DecodeError BOLT1.TlvStream
decode_tlvs known =
  first InvalidTlvStream . BOLT1.decode_tlv_stream (`elem` known)

-- parse the value of a record, if present
field
  :: Word64
  -> (BS.ByteString -> Maybe a)
  -> BOLT1.TlvStream
  -> Either DecodeError (Maybe a)
field t p s = case BOLT1.lookup_tlv t s of
  Nothing -> Right Nothing
  Just v  -> maybe (Left (InvalidTlvValue t)) (Right . Just) (p v)

-- the unknown records of a stream
extra_tlvs :: [Word64] -> BOLT1.TlvStream -> BOLT1.TlvStream
extra_tlvs known = BOLT1.filter_tlv_stream (`notElem` known)

-- a decoder that must consume its whole input
whole
  :: (BS.ByteString -> Maybe (a, BS.ByteString))
  -> BS.ByteString
  -> Maybe a
whole d bs = case d bs of
  Just (a, r) | BS.null r -> Just a
  _ -> Nothing

encode_tu_msat :: BOLT1.MilliSatoshi -> BS.ByteString
encode_tu_msat = BOLT1.encode_tu64 . BOLT1.un_milli_satoshi

decode_tu_msat :: BS.ByteString -> Maybe BOLT1.MilliSatoshi
decode_tu_msat v = BOLT1.decode_tu64 v >>= BOLT1.milli_satoshi

-- hop payloads ---------------------------------------------------------------

hop_payload_types :: [Word64]
hop_payload_types = [2, 4, 6, 8, 10, 12, 16, 18]

-- | Encode a t'HopPayload' as a @payload@ TLV stream (without its length
--   prefix). Fails if 'hp_extra' has a record of a known type.
--
--   >>> let hp = empty_hop_payload { hp_outgoing_cltv_value = Just 144 }
--   >>> encode_hop_payload hp
--   Right "\EOT\SOH\144"
encode_hop_payload :: HopPayload -> Either EncodeError BS.ByteString
encode_hop_payload hp = encode_tlvs hop_payload_types typed (hp_extra hp)
  where
    typed = catMaybes
      [ record 2 encode_tu_msat (hp_amt_to_forward hp)
      , record 4 BOLT1.encode_tu32 (hp_outgoing_cltv_value hp)
      , record 6 BOLT1.encode_short_channel_id (hp_short_channel_id hp)
      , record 8 encode_payment_data (hp_payment_data hp)
      , record 10 id (hp_encrypted_data hp)
      , record 12 BOLT1.un_point (hp_current_path_key hp)
      , record 16 id (hp_payment_metadata hp)
      , record 18 encode_tu_msat (hp_total_amount_msat hp)
      ]

-- | Decode a @payload@ TLV stream (without its length prefix). Fails on
--   a malformed stream, an unknown even type, or a malformed known
--   record.
--
--   >>> fmap hp_outgoing_cltv_value (decode_hop_payload "\EOT\SOH\144")
--   Right (Just 144)
--   >>> decode_hop_payload "\DC4\NUL"
--   Left (InvalidTlvStream (TlvUnknownEvenType 20))
decode_hop_payload :: BS.ByteString -> Either DecodeError HopPayload
decode_hop_payload bs = do
  s <- decode_tlvs hop_payload_types bs
  HopPayload
    <$> field 2 decode_tu_msat s
    <*> field 4 BOLT1.decode_tu32 s
    <*> field 6 (whole BOLT1.decode_short_channel_id) s
    <*> field 8 decode_payment_data s
    <*> field 10 Just s
    <*> field 12 BOLT1.point s
    <*> field 16 Just s
    <*> field 18 decode_tu_msat s
    <*> pure (extra_tlvs hop_payload_types s)

encode_payment_data :: PaymentData -> BS.ByteString
encode_payment_data (PaymentData s t) = un_payment_secret s <> encode_tu_msat t

decode_payment_data :: BS.ByteString -> Maybe PaymentData
decode_payment_data v = do
  let (s, t) = BS.splitAt 32 v
  PaymentData <$> payment_secret s <*> decode_tu_msat t

-- blinded hop data -----------------------------------------------------------

blinded_hop_data_types :: [Word64]
blinded_hop_data_types = [1, 2, 4, 6, 8, 10, 12, 14]

-- | Encode t'BlindedHopData' as an @encrypted_data_tlv@ stream. Fails if
--   'bhd_extra' has a record of a known type.
encode_blinded_hop_data :: BlindedHopData -> Either EncodeError BS.ByteString
encode_blinded_hop_data d =
    encode_tlvs blinded_hop_data_types typed (bhd_extra d)
  where
    typed = catMaybes
      [ record 1 id (bhd_padding d)
      , record 2 BOLT1.encode_short_channel_id (bhd_short_channel_id d)
      , record 4 BOLT1.un_point (bhd_next_node_id d)
      , record 6 id (bhd_path_id d)
      , record 8 BOLT1.un_point (bhd_next_path_key_override d)
      , record 10 encode_payment_relay (bhd_payment_relay d)
      , record 12 encode_payment_constraints (bhd_payment_constraints d)
      , record 14 BOLT9.render (bhd_allowed_features d)
      ]

-- | Decode an @encrypted_data_tlv@ stream. Fails on a malformed stream,
--   an unknown even type, or a malformed known record.
--
--   >>> fmap bhd_path_id (decode_blinded_hop_data "\ACK\apath id")
--   Right (Just "path id")
decode_blinded_hop_data :: BS.ByteString -> Either DecodeError BlindedHopData
decode_blinded_hop_data bs = do
  s <- decode_tlvs blinded_hop_data_types bs
  BlindedHopData
    <$> field 1 Just s
    <*> field 2 (whole BOLT1.decode_short_channel_id) s
    <*> field 4 BOLT1.point s
    <*> field 6 Just s
    <*> field 8 BOLT1.point s
    <*> field 10 decode_payment_relay s
    <*> field 12 decode_payment_constraints s
    <*> field 14 (Just . BOLT9.parse) s
    <*> pure (extra_tlvs blinded_hop_data_types s)

encode_payment_relay :: PaymentRelay -> BS.ByteString
encode_payment_relay (PaymentRelay c p b) =
  BOLT1.encode_u16 c <> BOLT1.encode_u32 p <> BOLT1.encode_tu32 b

decode_payment_relay :: BS.ByteString -> Maybe PaymentRelay
decode_payment_relay v = do
  (c, r0) <- BOLT1.decode_u16 v
  (p, r1) <- BOLT1.decode_u32 r0
  PaymentRelay c p <$> BOLT1.decode_tu32 r1

encode_payment_constraints :: PaymentConstraints -> BS.ByteString
encode_payment_constraints (PaymentConstraints c m) =
  BOLT1.encode_u32 c <> encode_tu_msat m

decode_payment_constraints :: BS.ByteString -> Maybe PaymentConstraints
decode_payment_constraints v = do
  (c, r) <- BOLT1.decode_u32 v
  PaymentConstraints c <$> decode_tu_msat r

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

-- | Encode a t'FailureMessage': the failure code, its data and the TLV
--   stream. Fails if a @channel_update@ exceeds 65535 bytes.
--
--   >>> let tlvs = BOLT1.empty_tlv_stream
--   >>> encode_failure_message (FailureMessage TemporaryNodeFailure tlvs)
--   Right " \STX"
encode_failure_message :: FailureMessage -> Either EncodeError BS.ByteString
encode_failure_message (FailureMessage f tlvs) = do
  body <- encode_failure f
  pure (BOLT1.encode_u16 (failure_code f) <> body
          <> BOLT1.encode_tlv_stream tlvs)

encode_failure :: Failure -> Either EncodeError BS.ByteString
encode_failure f = case f of
  TemporaryNodeFailure                 -> Right mempty
  PermanentNodeFailure                 -> Right mempty
  RequiredNodeFeatureMissing           -> Right mempty
  InvalidOnionVersion h                -> Right (un_onion_hash h)
  InvalidOnionHmac h                   -> Right (un_onion_hash h)
  InvalidOnionKey h                    -> Right (un_onion_hash h)
  TemporaryChannelFailure u            -> update u
  PermanentChannelFailure              -> Right mempty
  RequiredChannelFeatureMissing        -> Right mempty
  UnknownNextPeer                      -> Right mempty
  AmountBelowMinimum m u               -> (msat m <>) <$> update u
  FeeInsufficient m u                  -> (msat m <>) <$> update u
  IncorrectCltvExpiry c u              -> (BOLT1.encode_u32 c <>) <$> update u
  ExpiryTooSoon u                      -> update u
  IncorrectOrUnknownPaymentDetails m h -> Right (msat m <> BOLT1.encode_u32 h)
  FinalIncorrectCltvExpiry c           -> Right (BOLT1.encode_u32 c)
  FinalIncorrectHtlcAmount m           -> Right (msat m)
  ChannelDisabled d u                  -> (BOLT1.encode_u16 d <>) <$> update u
  ExpiryTooFar                         -> Right mempty
  InvalidOnionPayload Nothing          -> Right mempty
  InvalidOnionPayload (Just (t, o))    ->
    Right (BOLT1.encode_bigsize t <> BOLT1.encode_u16 o)
  MppTimeout                           -> Right mempty
  InvalidOnionBlinding h               -> Right (un_onion_hash h)
  UnknownFailure _ d                   -> Right d
  where
    msat = BOLT1.encode_milli_satoshi
    update = maybe (Left FieldTooLong) Right . BOLT1.encode_u16_prefixed

-- | Decode a failure message.
--
--   Per BOLT #4, bytes following a failure's data are ignored unless
--   they form a valid TLV stream (which, as no failure TLV types are
--   defined, may hold only odd types). The data of an unknown failure
--   code is every byte following it.
--
--   >>> fmap fm_failure (decode_failure_message " \STX")
--   Right TemporaryNodeFailure
--   >>> decode_failure_message "\DLE\r"
--   Left (InvalidFailureData 4109)
decode_failure_message :: BS.ByteString -> Either DecodeError FailureMessage
decode_failure_message bs = case BOLT1.decode_u16 bs of
  Nothing -> Left InvalidLength
  Just (c, rest) -> case decode_failure c rest of
    Nothing -> Left (InvalidFailureData c)
    Just (f, extra) ->
      let tlvs = either (const BOLT1.empty_tlv_stream) id
                   (BOLT1.decode_tlv_stream (const False) extra)
      in  Right (FailureMessage f tlvs)

decode_failure :: Word16 -> BS.ByteString -> Maybe (Failure, BS.ByteString)
decode_failure c bs = case c of
  0x2002 -> plain TemporaryNodeFailure
  0x6002 -> plain PermanentNodeFailure
  0x6003 -> plain RequiredNodeFeatureMissing
  0xc004 -> hashed InvalidOnionVersion
  0xc005 -> hashed InvalidOnionHmac
  0xc006 -> hashed InvalidOnionKey
  0x1007 -> do
    (u, r) <- update bs
    pure (TemporaryChannelFailure u, r)
  0x4008 -> plain PermanentChannelFailure
  0x4009 -> plain RequiredChannelFeatureMissing
  0x400a -> plain UnknownNextPeer
  0x100b -> do
    (m, r0) <- BOLT1.decode_milli_satoshi bs
    (u, r1) <- update r0
    pure (AmountBelowMinimum m u, r1)
  0x100c -> do
    (m, r0) <- BOLT1.decode_milli_satoshi bs
    (u, r1) <- update r0
    pure (FeeInsufficient m u, r1)
  0x100d -> do
    (e, r0) <- BOLT1.decode_u32 bs
    (u, r1) <- update r0
    pure (IncorrectCltvExpiry e u, r1)
  0x100e -> do
    (u, r) <- update bs
    pure (ExpiryTooSoon u, r)
  0x400f -> do
    (m, r0) <- BOLT1.decode_milli_satoshi bs
    (h, r1) <- BOLT1.decode_u32 r0
    pure (IncorrectOrUnknownPaymentDetails m h, r1)
  0x0012 -> do
    (e, r) <- BOLT1.decode_u32 bs
    pure (FinalIncorrectCltvExpiry e, r)
  0x0013 -> do
    (m, r) <- BOLT1.decode_milli_satoshi bs
    pure (FinalIncorrectHtlcAmount m, r)
  0x1014 -> do
    (d, r0) <- BOLT1.decode_u16 bs
    (u, r1) <- update r0
    pure (ChannelDisabled d u, r1)
  0x0015 -> plain ExpiryTooFar
  0x4016 -> Just $ case BOLT1.decode_bigsize bs of
    Just (t, r0) | Just (o, r1) <- BOLT1.decode_u16 r0 ->
      (InvalidOnionPayload (Just (t, o)), r1)
    _ -> (InvalidOnionPayload Nothing, bs)
  0x0017 -> plain MppTimeout
  0xc018 -> hashed InvalidOnionBlinding
  _      -> Just (UnknownFailure c bs, BS.empty)
  where
    plain f = Just (f, bs)
    hashed g
      | BS.length bs < 32 = Nothing
      | otherwise =
          let (h, r) = BS.splitAt 32 bs
          in  Just (g (OnionHash h), r)
    update = BOLT1.decode_u16_prefixed