ppad-bolt1-0.1.0: lib/Lightning/Protocol/BOLT1/TLV.hs
{-# OPTIONS_HADDOCK hide #-}
{-# LANGUAGE BangPatterns #-}
{-# LANGUAGE DeriveGeneric #-}
-- |
-- Module: Lightning.Protocol.BOLT1.TLV
-- Copyright: (c) 2025 Jared Tobin
-- License: MIT
-- Maintainer: Jared Tobin <jared@ppad.tech>
--
-- The TLV (type-length-value) format of BOLT #1.
module Lightning.Protocol.BOLT1.TLV (
TlvRecord(..)
, TlvStream
, tlv_stream
, un_tlv_stream
, empty_tlv_stream
, lookup_tlv
, filter_tlv_stream
, TlvError(..)
, encode_tlv_stream
, decode_tlv_stream
) where
import Control.DeepSeq (NFData)
import qualified Data.ByteString as BS
import qualified Data.ByteString.Unsafe as BU
import qualified Data.List as L
import Data.Word (Word64)
import GHC.Generics (Generic)
import Lightning.Protocol.BOLT1.Prim
-- | A single TLV record.
data TlvRecord = TlvRecord
{ tlv_type :: {-# UNPACK #-} !Word64
, tlv_value :: !BS.ByteString
} deriving (Eq, Show, Generic)
instance NFData TlvRecord
-- | A TLV stream: records in strictly increasing type order.
newtype TlvStream = TlvStream [TlvRecord]
deriving (Eq, Show, Generic)
instance NFData TlvStream
-- | Construct a 'TlvStream' from records in any order. Fails if two
-- records share a type.
--
-- >>> let Just s = tlv_stream [TlvRecord 3 "b", TlvRecord 1 "a"]
-- >>> map tlv_type (un_tlv_stream s)
-- [1,3]
-- >>> tlv_stream [TlvRecord 1 "a", TlvRecord 1 "b"]
-- Nothing
tlv_stream :: [TlvRecord] -> Maybe TlvStream
tlv_stream rs
| increasing sorted = Just (TlvStream sorted)
| otherwise = Nothing
where
sorted = L.sortOn tlv_type rs
increasing :: [TlvRecord] -> Bool
increasing (a:rest@(b:_)) = tlv_type a < tlv_type b && increasing rest
increasing _ = True
-- | The records of a 'TlvStream', in increasing type order.
un_tlv_stream :: TlvStream -> [TlvRecord]
un_tlv_stream (TlvStream rs) = rs
{-# INLINE un_tlv_stream #-}
-- | The empty 'TlvStream'.
empty_tlv_stream :: TlvStream
empty_tlv_stream = TlvStream []
-- | The value of the record with the given type, if present.
--
-- >>> lookup_tlv 3 =<< tlv_stream [TlvRecord 1 "a", TlvRecord 3 "b"]
-- Just "b"
lookup_tlv :: Word64 -> TlvStream -> Maybe BS.ByteString
lookup_tlv t (TlvStream rs) = go rs
where
go [] = Nothing
go (TlvRecord u v : more)
| u == t = Just v
| u > t = Nothing
| otherwise = go more
-- | Keep the records whose types satisfy the predicate.
--
-- >>> let Just s = tlv_stream [TlvRecord 1 "", TlvRecord 2 ""]
-- >>> map tlv_type (un_tlv_stream (filter_tlv_stream odd s))
-- [1]
filter_tlv_stream :: (Word64 -> Bool) -> TlvStream -> TlvStream
filter_tlv_stream p (TlvStream rs) = TlvStream (filter (p . tlv_type) rs)
-- | Why a TLV stream failed to decode.
data TlvError
= TlvTruncated -- ^ a type, length or value was cut off
| TlvNonMinimalBigSize -- ^ a type or length wasn't minimal
| TlvNotStrictlyIncreasing -- ^ types out of order or repeated
| TlvUnknownEvenType !Word64 -- ^ an even type the context doesn't know
deriving (Eq, Show, Generic)
instance NFData TlvError
-- | Encode a 'TlvStream'.
--
-- >>> fmap encode_tlv_stream (tlv_stream [TlvRecord 1 "a"])
-- Just "\SOH\SOHa"
encode_tlv_stream :: TlvStream -> BS.ByteString
encode_tlv_stream (TlvStream rs) = mconcat (concatMap enc rs)
where
enc (TlvRecord t v) =
[encode_bigsize t, encode_bigsize (fromIntegral (BS.length v)), v]
-- | Decode a TLV stream occupying the whole input, given a predicate
-- identifying the types known in the stream's context.
--
-- Per BOLT #1, decoding fails on truncation, non-minimal BigSize
-- encodings, types that are not strictly increasing, and unknown
-- even types. Every record is kept, including unknown odd ones, so
-- re-encoding reproduces the input.
--
-- >>> let Right s = decode_tlv_stream (== 1) "\SOH\SOHa\ETX\NUL"
-- >>> map tlv_type (un_tlv_stream s)
-- [1,3]
-- >>> decode_tlv_stream (== 1) "\STX\NUL"
-- Left (TlvUnknownEvenType 2)
decode_tlv_stream
:: (Word64 -> Bool) -- ^ is this type known?
-> BS.ByteString
-> Either TlvError TlvStream
decode_tlv_stream known = go Nothing []
where
go !_ !acc !bs | BS.null bs = Right (TlvStream (reverse acc))
go !prev !acc !bs = do
(t, r0) <- bigsize bs
case prev of
Just p | t <= p -> Left TlvNotStrictlyIncreasing
_ -> pure ()
(l, r1) <- bigsize r0
if l > fromIntegral (BS.length r1)
then Left TlvTruncated
else do
let !n = fromIntegral l
!v = BU.unsafeTake n r1
!r2 = BU.unsafeDrop n r1
if not (known t) && even t
then Left (TlvUnknownEvenType t)
else go (Just t) (TlvRecord t v : acc) r2
bigsize b = case decode_bigsize_detailed b of
Right r -> Right r
Left BigSizeTruncated -> Left TlvTruncated
Left BigSizeNonMinimal -> Left TlvNonMinimalBigSize