packages feed

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

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

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

module Lightning.Protocol.BOLT4.Construct (
    Hop(..)
  , ConstructError(..)
  , construct
  ) where

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

-- | A hop of a route: the public key the hop's onion layer is encrypted
--   to (its node id, or its blinded node id in a blinded route), and its
--   payload.
data Hop = Hop
  { hop_pubkey  :: !BOLT1.Point
  , hop_payload :: !HopPayload
  } deriving (Eq, Show, Generic)

instance NFData Hop

-- | Why an onion packet could not be constructed.
data ConstructError
  = EmptyRoute
    -- ^ the route has no hops
  | TooManyHops
    -- ^ the route has more than 20 hops
  | InvalidHopPubKey {-# UNPACK #-} !Int
    -- ^ the public key of the hop with this index is not a valid point
  | InvalidHopPayload {-# UNPACK #-} !Int
    -- ^ the payload of the hop with this index can't be encoded, or
    --   encodes to fewer than 2 bytes
  | PayloadsTooLarge {-# UNPACK #-} !Int
    -- ^ the hops' payloads, with their length prefixes and HMACs, take
    --   this many bytes, more than the 1300 available
  deriving (Eq, Show, Generic)

instance NFData ConstructError

-- | Construct an onion packet for a route, given a session key (which
--   must be fresh and random for every onion), the route's hops from
--   the first to the final one, and the associated data (the payment
--   hash, for a payment).
--
--   Returns the packet and the shared secret of each hop, in route
--   order; keep the latter to attribute returned errors with
--   'Lightning.Protocol.BOLT4.unwrap_error'.
--
--   >>> let Just session = secret_key (BS.replicate 32 0x41)
--   >>> let Just node = secret_key (BS.replicate 32 0x42)
--   >>> let Just amt = BOLT1.milli_satoshi 1000
--   >>> :{
--   let payload = empty_hop_payload {
--           hp_amt_to_forward      = Just amt
--         , hp_outgoing_cltv_value = Just 800000
--         }
--       route = [Hop (public_key node) payload]
--   :}
--   >>> let Right (pkt, sss) = construct session route (BS.replicate 32 0)
--   >>> BS.length (encode_onion_packet pkt)
--   1366
--   >>> length sss
--   1
construct
  :: SecretKey
  -> [Hop]
  -> BS.ByteString
  -> Either ConstructError (OnionPacket, [SharedSecret])
construct sk hops ad
  | null hops        = Left EmptyRoute
  | length hops > 20 = Left TooManyHops
  | otherwise = do
      pubs <- traverse parse_pub ihops
      pls <- traverse encode_pl ihops
      let sizes = map shift_size pls
          total = sum sizes
      if total > 1300
        then Left (PayloadsTooLarge total)
        else do
          sss <- shared_secrets (sk_bytes sk) (sk_pub sk) (zip [0 ..] pubs)
          let streams = map (\ss -> keystream (derive_rho ss) 2600) sss
              n = length hops
              fill = filler (take (n - 1) streams) (take (n - 1) sizes)
              start = keystream (derive_pad sk) 1300
              (body, mac) = wrap ad fill start
                              (reverse (zip3 sss streams pls))
              pkt = OnionPacket 0 (public_key sk) (HopPayloads body)
                      (Hmac32 mac)
          pure (pkt, sss)
  where
    ihops = zip [0 ..] hops

    parse_pub (i, h) =
      maybe (Left (InvalidHopPubKey i)) Right (from_point (hop_pubkey h))

    encode_pl (i, h) = case encode_hop_payload (hop_payload h) of
      Right p | BS.length p >= 2 -> Right p
      _ -> Left (InvalidHopPayload i)

-- bytes a payload takes in hop_payloads: length prefix, payload and HMAC
shift_size :: BS.ByteString -> Int
shift_size p =
  let !l = BS.length p
  in  BS.length (BOLT1.encode_bigsize (fromIntegral l)) + l + 32

-- the shared secret of each hop, blinding the ephemeral key pair after
-- each one
shared_secrets
  :: BS.ByteString
  -> Secp256k1.Projective
  -> [(Int, Secp256k1.Projective)]
  -> Either ConstructError [SharedSecret]
shared_secrets _ _ [] = Right []
shared_secrets e epub ((i, pub) : rest) = do
  ss <- note (ecdh e pub)
  case rest of
    [] -> Right [ss]
    _  -> do
      let bf = blinding_factor epub ss
      e' <- note (blind_scalar e bf)
      epub' <- note (blind_pub epub bf)
      (ss :) <$> shared_secrets e' epub' rest
  where
    note = maybe (Left (InvalidHopPubKey i)) Right

-- the filler, from the 2600-byte rho streams and shift sizes of all
-- hops but the final one
filler :: [BS.ByteString] -> [Int] -> BS.ByteString
filler streams sizes = L.foldl' step BS.empty (zip streams sizes)
  where
    step f (s, n) =
      let !ext = f <> BS.replicate n 0
          !off = 1300 - BS.length f
      in  xor_bytes ext (BS.take (BS.length ext) (BS.drop off s))

-- wrap the hops' payloads, given (shared secret, rho stream, payload)
-- from the final hop to the first, returning hop_payloads and the HMAC
wrap
  :: BS.ByteString
  -> BS.ByteString
  -> BS.ByteString
  -> [(SharedSecret, BS.ByteString, BS.ByteString)]
  -> (BS.ByteString, BS.ByteString)
wrap ad fill = go True (BS.replicate 32 0)
  where
    go _ !mac !buf [] = (buf, mac)
    go final !mac !buf ((ss, stream, pl) : rest) =
      let !len = BOLT1.encode_bigsize (fromIntegral (BS.length pl))
          !n = BS.length len + BS.length pl + 32
          !shifted = BS.concat [len, pl, mac, BS.take (1300 - n) buf]
          !obf = xor_bytes shifted stream
          !buf' | final     = BS.take (1300 - BS.length fill) obf <> fill
                | otherwise = obf
          !mac' = hmac (derive_mu ss) (buf' <> ad)
      in  go False mac' buf' rest