packages feed

ppad-tx-0.2.0: lib/Bitcoin/Prim/Tx.hs

{-# OPTIONS_HADDOCK prune #-}
{-# LANGUAGE BangPatterns #-}
{-# LANGUAGE OverloadedStrings #-}
{-# LANGUAGE RecordWildCards #-}

-- |
-- Module: Bitcoin.Prim.Tx
-- Copyright: (c) 2025 Jared Tobin
-- License: MIT
-- Maintainer: Jared Tobin <jared@ppad.tech>
--
-- Minimal Bitcoin transaction primitives, including raw transaction
-- types, serialisation to/from bytes, and txid computation.

module Bitcoin.Prim.Tx (
    -- * Transaction Types
    Tx(..)
  , TxIn(..)
  , TxOut(..)
  , OutPoint(..)
  , Witness(..)
  , TxId
  , mk_txid
  , un_txid
  , null_txid

    -- * Serialisation
  , to_bytes
  , from_bytes
  , to_bytes_legacy
  , to_base16
  , from_base16

    -- * TxId
  , txid
  ) where

import Bitcoin.Prim.Tx.Internal
import qualified Crypto.Hash.SHA256 as SHA256
import Data.Bits ((.|.), shiftL)
import qualified Data.ByteString as BS
import qualified Data.ByteString.Base16 as B16
import qualified Data.ByteString.Builder as BSB
import Data.List.NonEmpty (NonEmpty(..))
import qualified Data.List.NonEmpty as NE
import Data.Word (Word32, Word64)

-- | Construct a 'TxId' from 32 bytes in internal byte order (the
--   reverse of the usual hex display).
--
--   Returns 'Nothing' if the input is not exactly 32 bytes.
--
--   >>> fmap (BS.length . un_txid) (mk_txid (BS.replicate 32 0x00))
--   Just 32
--   >>> mk_txid (BS.replicate 31 0x00)
--   Nothing
mk_txid :: BS.ByteString -> Maybe TxId
mk_txid bs
  | BS.length bs == 32 = Just (TxId bs)
  | otherwise          = Nothing

-- | The 32 bytes of a 'TxId', in internal byte order.
un_txid :: TxId -> BS.ByteString
un_txid (TxId bs) = bs
{-# INLINE un_txid #-}

-- | The all-zero 'TxId', referenced by coinbase inputs' outpoints.
null_txid :: TxId
null_txid = TxId (BS.replicate 32 0x00)

-- serialisation --------------------------------------------------------------

-- | Serialise a transaction to bytes.
--
--   Uses the segwit format if any input has a non-empty witness, and
--   the legacy format otherwise (as Bitcoin Core does).
--
--   @
--   -- round-trip
--   from_bytes (to_bytes tx) == Just tx
--   @
to_bytes :: Tx -> BS.ByteString
to_bytes tx@Tx {..}
    | not (any (has_witness . txin_witness) tx_inputs) = to_bytes_legacy tx
    | otherwise = to_strict $
           put_word32_le tx_version
        <> BSB.word8 0x00  -- marker
        <> BSB.word8 0x01  -- flag
        <> put_compact (fromIntegral (NE.length tx_inputs))
        <> foldMap put_txin tx_inputs
        <> put_compact (fromIntegral (NE.length tx_outputs))
        <> foldMap put_txout tx_outputs
        <> foldMap (put_witness . txin_witness) tx_inputs
        <> put_word32_le tx_locktime

-- whether a witness stack is non-empty
has_witness :: Witness -> Bool
has_witness (Witness items) = not (null items)
{-# INLINE has_witness #-}

-- | Serialise a transaction to legacy format (no witness data).
--
--   Used for txid computation. Excludes witness data even if present.
--
--   @
--   -- for a legacy tx (no witnesses), the same as to_bytes
--   to_bytes_legacy legacy_tx == to_bytes legacy_tx
--
--   -- for a segwit tx, strips the witnesses
--   BS.length (to_bytes_legacy segwit_tx) < BS.length (to_bytes segwit_tx)
--   @
to_bytes_legacy :: Tx -> BS.ByteString
to_bytes_legacy Tx {..} = to_strict $
       put_word32_le tx_version
    <> put_compact (fromIntegral (NE.length tx_inputs))
    <> foldMap put_txin tx_inputs
    <> put_compact (fromIntegral (NE.length tx_outputs))
    <> foldMap put_txout tx_outputs
    <> put_word32_le tx_locktime

-- | Serialise a transaction to base16 (hex).
--
--   @
--   to_base16 tx = B16.encode (to_bytes tx)
--   @
to_base16 :: Tx -> BS.ByteString
to_base16 tx = B16.encode (to_bytes tx)

-- | Parse a transaction from base16 (hex).
--
--   @
--   -- round-trip
--   from_base16 (to_base16 tx) == Just tx
--   @
from_base16 :: BS.ByteString -> Maybe Tx
from_base16 b16 = do
  bs <- B16.decode b16
  from_bytes bs

-- decoding -------------------------------------------------------------------

-- | Parse a transaction from bytes.
--
--   Automatically detects segwit vs legacy format by checking for
--   marker byte 0x00 followed by flag 0x01 after the version field.
--
--   Returns 'Nothing' on invalid or truncated input.
--
--   @
--   -- round-trip
--   from_bytes (to_bytes tx) == Just tx
--   @
from_bytes :: BS.ByteString -> Maybe Tx
from_bytes !bs = do
  -- need at least 4 bytes for version
  guard (BS.length bs >= 4)
  let !version = get_word32_le bs 0
      !off0 = 4
  -- check for segwit marker (0x00) and flag (0x01)
  if   BS.length bs > off0 + 1
    && BS.index bs off0 == 0x00
    && BS.index bs (off0 + 1) == 0x01
  then parse_segwit bs version (off0 + 2)
  else parse_legacy bs version off0

-- Parse legacy transaction (no witness data)
parse_legacy :: BS.ByteString -> Word32 -> Int -> Maybe Tx
parse_legacy !bs !version !off0 = do
  -- input count
  (input_count, off1) <- get_compact bs off0
  -- inputs (must have at least one)
  (inputs_list, off2) <- get_many get_txin bs off1 input_count
  inputs <- NE.nonEmpty inputs_list
  -- output count
  (output_count, off3) <- get_compact bs off2
  -- outputs (must have at least one)
  (outputs_list, off4) <- get_many get_txout bs off3 output_count
  outputs <- NE.nonEmpty outputs_list
  -- locktime (4 bytes)
  guard (BS.length bs >= off4 + 4)
  let !locktime = get_word32_le bs off4
      !off5 = off4 + 4
  -- should have consumed all bytes
  guard (off5 == BS.length bs)
  pure $! Tx version inputs outputs locktime

-- Parse segwit transaction (with witness data)
parse_segwit :: BS.ByteString -> Word32 -> Int -> Maybe Tx
parse_segwit !bs !version !off0 = do
  -- input count
  (input_count, off1) <- get_compact bs off0
  -- inputs (must have at least one)
  (inputs_list, off2) <- get_many get_txin bs off1 input_count
  inputs <- NE.nonEmpty inputs_list
  -- output count
  (output_count, off3) <- get_compact bs off2
  -- outputs (must have at least one)
  (outputs_list, off4) <- get_many get_txout bs off3 output_count
  outputs <- NE.nonEmpty outputs_list
  -- witnesses (one per input)
  (witnesses, off5) <- get_many get_witness bs off4 input_count
  -- a marker and flag with no witness data is invalid (Bitcoin Core:
  -- "superfluous witness record")
  guard (any has_witness witnesses)
  -- locktime (4 bytes)
  guard (BS.length bs >= off5 + 4)
  let !locktime = get_word32_le bs off5
      !off6 = off5 + 4
  -- should have consumed all bytes
  guard (off6 == BS.length bs)
  pure $! Tx version (attach_witnesses inputs witnesses) outputs locktime

-- Pair each input with its witness (get_many returns exactly one per
-- input).
attach_witnesses :: NonEmpty TxIn -> [Witness] -> NonEmpty TxIn
attach_witnesses (i :| is) ws = case ws of
  (w : rest) -> set i w :| zipWith set is rest
  []         -> i :| is
  where
    set x w = x { txin_witness = w }

-- internal helpers -----------------------------------------------------------

-- | Guard for Maybe monad.
guard :: Bool -> Maybe ()
guard True  = Just ()
guard False = Nothing
{-# INLINE guard #-}

-- | Decode a 32-bit little-endian word at the given offset.
--   Does not bounds-check; caller must ensure sufficient bytes.
get_word32_le :: BS.ByteString -> Int -> Word32
get_word32_le !bs !off =
  let !b0 = fromIntegral (BS.index bs off) :: Word32
      !b1 = fromIntegral (BS.index bs (off + 1)) :: Word32
      !b2 = fromIntegral (BS.index bs (off + 2)) :: Word32
      !b3 = fromIntegral (BS.index bs (off + 3)) :: Word32
  in  b0 .|. (b1 `shiftL` 8) .|. (b2 `shiftL` 16) .|. (b3 `shiftL` 24)
{-# INLINE get_word32_le #-}

-- | Decode a 64-bit little-endian word at the given offset.
--   Does not bounds-check; caller must ensure sufficient bytes.
get_word64_le :: BS.ByteString -> Int -> Word64
get_word64_le !bs !off =
  let !b0 = fromIntegral (BS.index bs off) :: Word64
      !b1 = fromIntegral (BS.index bs (off + 1)) :: Word64
      !b2 = fromIntegral (BS.index bs (off + 2)) :: Word64
      !b3 = fromIntegral (BS.index bs (off + 3)) :: Word64
      !b4 = fromIntegral (BS.index bs (off + 4)) :: Word64
      !b5 = fromIntegral (BS.index bs (off + 5)) :: Word64
      !b6 = fromIntegral (BS.index bs (off + 6)) :: Word64
      !b7 = fromIntegral (BS.index bs (off + 7)) :: Word64
  in  b0 .|. (b1 `shiftL` 8) .|. (b2 `shiftL` 16) .|. (b3 `shiftL` 24)
          .|. (b4 `shiftL` 32) .|. (b5 `shiftL` 40)
          .|. (b6 `shiftL` 48) .|. (b7 `shiftL` 56)
{-# INLINE get_word64_le #-}

-- | Decode a 16-bit little-endian word at the given offset.
--   Does not bounds-check; caller must ensure sufficient bytes.
get_word16_le :: BS.ByteString -> Int -> Word64
get_word16_le !bs !off =
  let !b0 = fromIntegral (BS.index bs off) :: Word64
      !b1 = fromIntegral (BS.index bs (off + 1)) :: Word64
  in  b0 .|. (b1 `shiftL` 8)
{-# INLINE get_word16_le #-}

-- | Decode compactSize (Bitcoin's variable-length integer).
--   Returns (value, new_offset).
--   Enforces minimal encoding: rejects non-minimal representations.
get_compact :: BS.ByteString -> Int -> Maybe (Word64, Int)
get_compact !bs !off
  | off >= BS.length bs = Nothing
  | otherwise = case BS.index bs off of
      tag | tag <= 0xfc ->
        -- Single byte: value is the tag itself
        Just (fromIntegral tag, off + 1)

      0xfd ->
        -- 2-byte value follows
        if BS.length bs < off + 3
        then Nothing
        else
          let !val = get_word16_le bs (off + 1)
          in  if val < 0xfd
              then Nothing  -- non-minimal encoding
              else Just (val, off + 3)

      0xfe ->
        -- 4-byte value follows
        if BS.length bs < off + 5
        then Nothing
        else
          let !val = fromIntegral (get_word32_le bs (off + 1)) :: Word64
          in  if val <= 0xffff
              then Nothing  -- non-minimal encoding
              else Just (val, off + 5)

      _ -> -- 0xff
        -- 8-byte value follows
        if BS.length bs < off + 9
        then Nothing
        else
          let !val = get_word64_le bs (off + 1)
          in  if val <= 0xffffffff
              then Nothing  -- non-minimal encoding
              else Just (val, off + 9)
{-# INLINE get_compact #-}

-- | Decode an outpoint (txid + vout).
--   Returns (OutPoint, new_offset).
get_outpoint :: BS.ByteString -> Int -> Maybe (OutPoint, Int)
get_outpoint !bs !off
  | BS.length bs < off + 36 = Nothing
  | otherwise =
      let !txid_bytes = BS.take 32 (BS.drop off bs)
          !vout = get_word32_le bs (off + 32)
      in  Just (OutPoint (TxId txid_bytes) vout, off + 36)
{-# INLINE get_outpoint #-}

-- | Decode a transaction input.
--   Returns (TxIn, new_offset).
get_txin :: BS.ByteString -> Int -> Maybe (TxIn, Int)
get_txin !bs !off0 = do
  -- outpoint: 36 bytes
  (outpoint, off1) <- get_outpoint bs off0
  -- scriptSig
  (script_sig, off3) <- get_bytes bs off1
  -- sequence: 4 bytes
  guard (BS.length bs >= off3 + 4)
  let !seqn = get_word32_le bs off3
      !off4 = off3 + 4
  pure (TxIn outpoint script_sig seqn (Witness []), off4)

-- | Decode a transaction output.
--   Returns (TxOut, new_offset).
get_txout :: BS.ByteString -> Int -> Maybe (TxOut, Int)
get_txout !bs !off0 = do
  -- value: 8 bytes
  guard (BS.length bs >= off0 + 8)
  let !value = get_word64_le bs off0
      !off1 = off0 + 8
  -- scriptPubKey
  (script_pk, off2) <- get_bytes bs off1
  pure (TxOut value script_pk, off2)

-- | Decode a witness stack for one input.
--   Returns (Witness, new_offset).
get_witness :: BS.ByteString -> Int -> Maybe (Witness, Int)
get_witness !bs !off0 = do
  -- stack item count
  (item_count, off1) <- get_compact bs off0
  -- each item: length + bytes
  (items, off2) <- get_many get_bytes bs off1 item_count
  pure (Witness items, off2)

-- | Decode compactSize-length-prefixed bytes.
--   Returns (bytes, new_offset).
get_bytes :: BS.ByteString -> Int -> Maybe (BS.ByteString, Int)
get_bytes !bs !off0 = do
  (len, off1) <- get_compact bs off0
  -- compare in Word64, so that huge lengths can't wrap negative
  guard (len <= fromIntegral (BS.length bs - off1))
  let !n = fromIntegral len
  pure (BS.take n (BS.drop off1 bs), off1 + n)

-- | Decode a counted sequence of items using a decoder function.
--   Returns (list of items, new_offset).
get_many :: (BS.ByteString -> Int -> Maybe (a, Int))
         -> BS.ByteString -> Int -> Word64 -> Maybe ([a], Int)
get_many getter !bs !off0 !count
  -- every item occupies at least one byte, so a larger count is
  -- malformed (and might not fit in an Int)
  | count > fromIntegral (BS.length bs - off0) = Nothing
  | otherwise = go [] off0 (fromIntegral count :: Int)
  where
    go !acc !off !n
      | n <= 0    = Just (reverse acc, off)
      | otherwise = do
          (item, off') <- getter bs off
          go (item : acc) off' (n - 1)
{-# INLINE get_many #-}

-- txid -----------------------------------------------------------------------

-- | Compute the transaction ID (double SHA256 of legacy serialisation).
--
--   The txid is computed from the legacy serialisation, so segwit
--   transactions have the same txid regardless of witness data. It is
--   in internal byte order; reverse it for the usual hex display.
--
--   @
--   -- Satoshi->Hal tx (block 170)
--   B16.encode (BS.reverse (un_txid (txid satoshi_hal_tx)))
--     == "f4184fc596403b9d638783cf57adfe4c75c605f6356fbc91338530e9831e9e16"
--   @
txid :: Tx -> TxId
txid tx = TxId (SHA256.hash (SHA256.hash (to_bytes_legacy tx)))