packages feed

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

{-# OPTIONS_HADDOCK hide #-}
{-# LANGUAGE DeriveGeneric #-}

-- |
-- Module: Lightning.Protocol.BOLT4.Process
-- Copyright: (c) 2025 Jared Tobin
-- License: MIT
-- Maintainer: Jared Tobin <jared@ppad.tech>
--
-- Onion packet processing.

module Lightning.Protocol.BOLT4.Process (
    ProcessResult(..)
  , ForwardInfo(..)
  , ReceiveInfo(..)
  , process
  ) where

import Control.DeepSeq (NFData)
import qualified Crypto.Curve.Secp256k1 as Secp256k1
import qualified Data.ByteString as BS
import GHC.Generics (Generic)
import qualified Lightning.Protocol.BOLT1 as BOLT1
import Lightning.Protocol.BOLT4.Blinding (unblind)
import Lightning.Protocol.BOLT4.Codec (decode_hop_payload)
import Lightning.Protocol.BOLT4.Prim
import Lightning.Protocol.BOLT4.Types

-- | The result of processing an onion packet.
data ProcessResult
  = Forward !ForwardInfo
    -- ^ forward the onion to the next hop
  | Receive !ReceiveInfo
    -- ^ this node is the onion's final hop
  deriving (Eq, Show, Generic)

instance NFData ProcessResult

-- | What a forwarding node learns from an onion.
data ForwardInfo = ForwardInfo
  { fwd_payload       :: !HopPayload
    -- ^ this node's payload
  , fwd_blinded       :: !(Maybe BlindedInfo)
    -- ^ the decrypted @encrypted_recipient_data@ and next path key, in
    --   a blinded route
  , fwd_next_packet   :: !OnionPacket
    -- ^ the onion for the next hop
  , fwd_shared_secret :: !SharedSecret
    -- ^ the shared secret, to wrap returned errors with
  } deriving (Eq, Show, Generic)

instance NFData ForwardInfo

-- | What the final node learns from an onion.
data ReceiveInfo = ReceiveInfo
  { rcv_payload       :: !HopPayload
    -- ^ this node's payload
  , rcv_blinded       :: !(Maybe BlindedInfo)
    -- ^ the decrypted @encrypted_recipient_data@, in a blinded route
  , rcv_shared_secret :: !SharedSecret
    -- ^ the shared secret, to construct errors with
  } deriving (Eq, Show, Generic)

instance NFData ReceiveInfo

-- | Process an onion packet, given this node's private key, the packet,
--   its associated data (the payment hash, for a payment), and the
--   @path_key@ received with it (in @update_add_htlc@), if any.
--
--   In a blinded route, the onion is encrypted to this node's blinded
--   node id: given a @path_key@, the matching private key is derived
--   from it. The @encrypted_recipient_data@ in the payload is decrypted
--   with the @path_key@, or, at the introduction node, with the
--   payload's @current_path_key@; a path key without
--   @encrypted_recipient_data@, or both kinds of path key, are
--   rejected.
--
--   The payload is otherwise returned as is. The caller must apply the
--   remaining reader requirements of BOLT #4: the presence of the
--   fields required for forwarding or receiving, the HTLC amount and
--   expiry checks (in a blinded route, against @payment_relay@ and
--   @payment_constraints@), @allowed_features@, the rule that
--   @encrypted_recipient_data@ hold only one of @short_channel_id@ and
--   @next_node_id@, and the rejection of replayed HMACs.
--
--   >>> let Just session = secret_key (BS.replicate 32 0x41)
--   >>> let Just node = secret_key (BS.replicate 32 0x42)
--   >>> let pl = empty_hop_payload { hp_outgoing_cltv_value = Just 144 }
--   >>> let ad = BS.replicate 32 0
--   >>> let Right (pkt, _) = construct session [Hop (public_key node) pl] ad
--   >>> let Right (Receive info) = process node pkt ad Nothing
--   >>> hp_outgoing_cltv_value (rcv_payload info)
--   Just 144
process
  :: SecretKey
  -> OnionPacket
  -> BS.ByteString
  -> Maybe BOLT1.Point
  -> Either ProcessError ProcessResult
process sk pkt ad mpk = do
  let ver = onion_version pkt
  if ver /= 0 then Left (InvalidVersion ver) else Right ()
  eph <- note InvalidPublicKey (from_point (onion_public_key pkt))
  -- with a path_key, the onion is encrypted to the blinded node id
  blinding <- traverse path_secret mpk
  key <- case blinding of
    Nothing -> Right (sk_bytes sk)
    Just (_, bss) ->
      note InvalidPathKey (blind_scalar (sk_bytes sk) (blinded_node_tweak bss))
  ss <- note InvalidPublicKey (ecdh key eph)
  let payloads = un_hop_payloads (onion_hop_payloads pkt)
      mac = hmac (derive_mu ss) (payloads <> ad)
  if   ct_eq mac (un_hmac32 (onion_hmac pkt))
  then Right ()
  else Left HmacMismatch
  let plain = xor_bytes (payloads <> BS.replicate 1300 0)
                        (keystream (derive_rho ss) 2600)
  (pl, next_mac, rest) <- split_payload plain
  hp <- either (Left . InvalidPayload) Right (decode_hop_payload pl)
  binfo <- case (hp_encrypted_data hp, blinding, hp_current_path_key hp) of
    (Nothing, Nothing, Nothing) -> Right Nothing
    (Nothing, _, _)             -> Left UnexpectedPathKey
    (Just _, Just _, Just _)    -> Left UnexpectedPathKey
    (Just _, Nothing, Nothing)  -> Left MissingPathKey
    (Just enc, Just (e, bss), Nothing) -> Just <$> unblind bss e enc
    (Just enc, Nothing, Just cpk) -> do
      (e, bss) <- path_secret cpk
      Just <$> unblind bss e enc
  if BS.all (== 0) next_mac
    then Right (Receive (ReceiveInfo hp binfo ss))
    else do
      eph' <- note InvalidPublicKey
                (blind_pub eph (blinding_factor eph ss) >>= to_point)
      let next = OnionPacket 0 eph' (HopPayloads rest) (Hmac32 next_mac)
      Right (Forward (ForwardInfo hp binfo next ss))
  where
    note :: ProcessError -> Maybe a -> Either ProcessError a
    note e = maybe (Left e) Right

    -- a path key, and its shared secret with this node
    path_secret
      :: BOLT1.Point
      -> Either ProcessError (Secp256k1.Projective, SharedSecret)
    path_secret pk = do
      e <- note InvalidPathKey (from_point pk)
      bss <- note InvalidPathKey (ecdh (sk_bytes sk) e)
      Right (e, bss)

-- Split the decrypted 2600-byte buffer into the payload, the next HMAC
-- and the next hop_payloads. The length prefix, payload and HMAC must
-- fit in the 1300-byte hop_payloads.
split_payload
  :: BS.ByteString
  -> Either ProcessError (BS.ByteString, BS.ByteString, BS.ByteString)
split_payload plain = case BOLT1.decode_bigsize plain of
  Nothing -> Left InvalidPayloadLength
  Just (len, r0)
    | len < 2   -> Left InvalidPayloadLength
    | len > fromIntegral (1300 - prefix - 32) -> Left InvalidPayloadLength
    | otherwise ->
        let (pl, r1)  = BS.splitAt (fromIntegral len) r0
            (mac, r2) = BS.splitAt 32 r1
        in  Right (pl, mac, BS.take 1300 r2)
    where
      prefix = BS.length plain - BS.length r0