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