packages feed

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

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

-- |
-- Module: Lightning.Protocol.BOLT4.Error
-- Copyright: (c) 2025 Jared Tobin
-- License: MIT
-- Maintainer: Jared Tobin <jared@ppad.tech>
--
-- Returning errors: construction, wrapping and attribution of return
-- packets.

module Lightning.Protocol.BOLT4.Error (
    ErrorPacket(..)
  , Attribution(..)
  , construct_error
  , wrap_error
  , unwrap_error
  ) where

import Control.DeepSeq (NFData)
import qualified Data.ByteString as BS
import GHC.Generics (Generic)
import qualified Lightning.Protocol.BOLT1 as BOLT1
import Lightning.Protocol.BOLT4.Codec
import Lightning.Protocol.BOLT4.Prim
import Lightning.Protocol.BOLT4.Types

-- | An obfuscated return packet (the @reason@ of @update_fail_htlc@).
newtype ErrorPacket = ErrorPacket BS.ByteString
  deriving (Eq, Show, Generic)

instance NFData ErrorPacket

-- | The origin of a return packet, as found by 'unwrap_error'.
data Attribution
  = Attributed {-# UNPACK #-} !Int !FailureMessage
    -- ^ the hop with this index (from 0, in route order) returned the
    --   failure
  | MalformedFailure {-# UNPACK #-} !Int
    -- ^ the hop with this index returned the packet, but its failure
    --   message can't be decoded
  | UnknownOrigin
    -- ^ no hop's HMAC matches
  deriving (Eq, Show, Generic)

instance NFData Attribution

-- | Construct a return packet at the failing node, given the shared
--   secret from processing the onion and the failure.
--
--   The failure message is padded so that its length plus that of the
--   padding is at least 256 bytes. Fails if the failure message can't
--   be encoded, or exceeds 65535 bytes.
--
--   >>> let Just ss = shared_secret (BS.replicate 32 0x01)
--   >>> let fm = FailureMessage TemporaryNodeFailure BOLT1.empty_tlv_stream
--   >>> let Right (ErrorPacket pkt) = construct_error ss fm
--   >>> BS.length pkt
--   292
construct_error
  :: SharedSecret
  -> FailureMessage
  -> Either EncodeError ErrorPacket
construct_error ss fm = do
  msg <- encode_failure_message fm
  let len = BS.length msg
  if len > 65535
    then Left FieldTooLong
    else do
      let pad  = max 0 (256 - len)
          body = BS.concat
            [ BOLT1.encode_u16 (fromIntegral len), msg
            , BOLT1.encode_u16 (fromIntegral pad), BS.replicate pad 0 ]
          mac  = hmac (derive_um ss) body
      pure (wrap_error ss (ErrorPacket (mac <> body)))

-- | Wrap a return packet from downstream at an intermediate node, given
--   the shared secret from processing the onion.
--
--   >>> let Just ss = shared_secret (BS.replicate 32 0x01)
--   >>> wrap_error ss (wrap_error ss (ErrorPacket "packet"))
--   ErrorPacket "packet"
wrap_error :: SharedSecret -> ErrorPacket -> ErrorPacket
wrap_error ss (ErrorPacket p) =
  ErrorPacket (xor_bytes p (keystream (derive_ammag ss) (BS.length p)))

-- | Find the origin of a return packet at the origin node, given the
--   shared secrets from 'Lightning.Protocol.BOLT4.construct', in route
--   order.
--
--   >>> let Just ss0 = shared_secret (BS.replicate 32 0x01)
--   >>> let Just ss1 = shared_secret (BS.replicate 32 0x02)
--   >>> let fm = FailureMessage TemporaryNodeFailure BOLT1.empty_tlv_stream
--   >>> let Right pkt = construct_error ss1 fm
--   >>> unwrap_error [ss0, ss1] (wrap_error ss0 pkt) == Attributed 1 fm
--   True
--   >>> unwrap_error [ss0] (wrap_error ss0 pkt)
--   UnknownOrigin
unwrap_error :: [SharedSecret] -> ErrorPacket -> Attribution
unwrap_error = go 0
  where
    go _ [] _ = UnknownOrigin
    go i (ss : rest) p =
      let ErrorPacket q = wrap_error ss p
          (mac, body) = BS.splitAt 32 q
      in  if   ct_eq mac (hmac (derive_um ss) body)
          then either (const (MalformedFailure i)) (Attributed i)
                 (failure_message body)
          else go (i + 1) rest (ErrorPacket q)

-- The failure message of a return packet (after its HMAC).
failure_message :: BS.ByteString -> Either DecodeError FailureMessage
failure_message body = case BOLT1.decode_u16 body of
  Just (len, r0)
    | BS.length r0 >= fromIntegral len
    , let (msg, r1) = BS.splitAt (fromIntegral len) r0
    , Just (pad, r2) <- BOLT1.decode_u16 r1
    , BS.length r2 >= fromIntegral pad
    -> decode_failure_message msg
  _ -> Left InvalidLength