Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
121 changes: 121 additions & 0 deletions bolt12/decode_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,121 @@
package bolt12

import (
"bytes"
"testing"

"github.com/lightningnetwork/lnd/lnwire"
"github.com/lightningnetwork/lnd/tlv"
"github.com/stretchr/testify/require"
)

// appendRawRecord writes a single TLV record (type, length, value) to buf.
func appendRawRecord(t *testing.T, buf *bytes.Buffer, typ uint64,
value []byte) {

t.Helper()

var scratch [8]byte
require.NoError(t, tlv.WriteVarInt(buf, typ, &scratch))
require.NoError(t, tlv.WriteVarInt(buf, uint64(len(value)), &scratch))
_, err := buf.Write(value)
require.NoError(t, err)
}

// TestDecodeRejectsNonMinimalFeatures tests that a non-minimally encoded
// feature vector is rejected at decode, so the canonical re-encode of an
// accepted message always reproduces the wire bytes.
func TestDecodeRejectsNonMinimalFeatures(t *testing.T) {
t.Parallel()

// A feature vector holding only bit 0 encodes minimally as 0x01. The
// two-byte form pads it with a leading zero byte.
padded := []byte{0x00, 0x01}

tests := []struct {
name string
typ uint64
decode func([]byte) error
}{
{
name: "offer_features",
typ: 12,
decode: func(b []byte) error {
_, err := decodeOffer(b)
return err
},
},
{
name: "invreq_features",
typ: 84,
decode: func(b []byte) error {
_, err := DecodeInvoiceRequest(b)
return err
},
},
{
name: "invoice_features",
typ: 174,
decode: func(b []byte) error {
_, err := DecodeInvoice(b)
return err
},
},
}

for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()

var buf bytes.Buffer
appendRawRecord(t, &buf, tc.typ, padded)

err := tc.decode(buf.Bytes())
require.ErrorIs(t, err, ErrNonMinimalFeatures)

// The minimal encoding of the same bit set is accepted.
var minimalBuf bytes.Buffer
appendRawRecord(t, &minimalBuf, tc.typ, []byte{0x01})
require.NoError(t, tc.decode(minimalBuf.Bytes()))
})
}
}

// TestDecodeRejectsNonMinimalAmount tests that a non-minimally encoded
// amount is rejected at decode, so the canonical re-encode of an accepted
// message always reproduces the wire bytes.
func TestDecodeRejectsNonMinimalAmount(t *testing.T) {
t.Parallel()

// invreq_amount (type 82) holding the value 1 in two bytes: the
// minimal tu64 encoding of 1 is a single byte.
var buf bytes.Buffer
appendRawRecord(t, &buf, 82, []byte{0x00, 0x01})

_, err := DecodeInvoiceRequest(buf.Bytes())
require.ErrorIs(t, err, tlv.ErrTUintNotMinimal)
}

// TestUnknownOddTLVRoundTripByteExact tests that unknown odd TLV types in the
// signed range are preserved on decode and re-encode, so that the canonical
// re-encode of an accepted message always reproduces the wire bytes.
func TestUnknownOddTLVRoundTripByteExact(t *testing.T) {
t.Parallel()

var buf bytes.Buffer

// invreq_metadata (type 0), then two unknown odd types in the signed
// range: one with a value, one zero-length.
appendRawRecord(t, &buf, 0, []byte("meta"))
appendRawRecord(t, &buf, 93, []byte("xyz"))
appendRawRecord(t, &buf, 95, nil)

wire := buf.Bytes()

ir, err := DecodeInvoiceRequest(wire)
require.NoError(t, err)

var out bytes.Buffer
require.NoError(t, lnwire.EncodePureTLVMessage(ir, &out))
require.Equal(t, wire, out.Bytes())
}
146 changes: 146 additions & 0 deletions bolt12/helpers_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -3,12 +3,15 @@ package bolt12
import (
"bytes"
"encoding/json"
"io"
"os"
"sync"
"testing"
"time"

"github.com/btcsuite/btcd/btcec/v2"
"github.com/lightningnetwork/lnd/lnwire"
"github.com/lightningnetwork/lnd/tlv"
"github.com/stretchr/testify/require"
)

Expand Down Expand Up @@ -115,3 +118,146 @@ func loadOffersVectors(t *testing.T) []offersTestVector {

return vectors
}

// streamToRecords parses an arbitrary TLV byte stream into tlv.Record values
// whose Encode method reproduces the original wire bytes, without going
// through a typed message decoder.
func streamToRecords(t *testing.T, data []byte) []tlv.Record {
t.Helper()

stream, err := tlv.NewStream()
require.NoError(t, err)

typeMap, err := stream.DecodeWithParsedTypesP2P(bytes.NewReader(data))
require.NoError(t, err)

return lnwire.TlvMapToRecords(typeMap)
}

// payerIDFromStream returns the invreq_payer_id public key carried by a raw
// BOLT 12 TLV stream. It reads the field straight from the parsed type map so
// callers stay independent of the typed message decoders.
func payerIDFromStream(t *testing.T, data []byte) *btcec.PublicKey {
t.Helper()

stream, err := tlv.NewStream()
require.NoError(t, err)

typeMap, err := stream.DecodeWithParsedTypesP2P(bytes.NewReader(data))
require.NoError(t, err)

raw, ok := typeMap[invreqPayerIDType]
require.True(t, ok, "stream carries no invreq_payer_id")

pubKey, err := btcec.ParsePubKey(raw)
require.NoError(t, err)

return pubKey
}

// recordFromWireBytes builds a single tlv.Record whose encoding is the
// supplied full TLV byte slice. The slice must be a complete
// type+length+value sequence. Inputs are trusted spec fixtures, so the
// length prefix is allocated without a bound.
func recordFromWireBytes(t *testing.T, full []byte) tlv.Record {
t.Helper()

var buf [8]byte
r := bytes.NewReader(full)

typ, err := tlv.ReadVarInt(r, &buf)
require.NoError(t, err)

length, err := tlv.ReadVarInt(r, &buf)
require.NoError(t, err)

value := make([]byte, length)
_, err = io.ReadFull(r, value)
require.NoError(t, err)

return tlv.MakePrimitiveRecord(tlv.Type(typ), &value)
}

// sigTestVector represents a test case from signature-test.json.
type sigTestVector struct {
Comment string `json:"comment"`
TLV string `json:"tlv"`
Bolt12 string `json:"bolt12"`

//nolint:tagliatelle // BOLT 12 spec vector key.
FirstTLV string `json:"first-tlv"`
Leaves []json.RawMessage `json:"leaves"`
Branches []json.RawMessage `json:"branches"`
Merkle string `json:"merkle"`

SignatureTag string `json:"signature_tag"`
Signature string `json:"signature"`
}

// readSignatureDataOnce reads signature-test.json once so the file is
// parsed only once per test process.
var readSignatureDataOnce = sync.OnceValues(func() ([]byte, error) {
return os.ReadFile("test-vectors/signature-test.json")
})

// loadSignatureVectorsOnce parses signature-test.json into typed
// sigTestVectors. The raw-JSON loader is separate because the JSON
// contains a key ("H(signature_tag,merkle)") that cannot be expressed
// via Go struct tags.
var loadSignatureVectorsOnce = sync.OnceValues(
func() ([]sigTestVector, error) {
data, err := readSignatureDataOnce()
if err != nil {
return nil, err
}

var vectors []sigTestVector
if err := json.Unmarshal(data, &vectors); err != nil {
return nil, err
}

return vectors, nil
},
)

// loadSignatureVectors returns the parsed sigTestVector slice, failing the
// test if signature-test.json is unreadable or malformed.
func loadSignatureVectors(t *testing.T) []sigTestVector {
t.Helper()

vectors, err := loadSignatureVectorsOnce()
require.NoError(t, err)

return vectors
}

// loadSignatureRawOnce parses signature-test.json as a slice of raw
// json.RawMessage so callers can index into keys whose names cannot be
// expressed via struct tags.
var loadSignatureRawOnce = sync.OnceValues(
func() ([]json.RawMessage, error) {
data, err := readSignatureDataOnce()
if err != nil {
return nil, err
}

var raw []json.RawMessage
if err := json.Unmarshal(data, &raw); err != nil {
return nil, err
}

return raw, nil
},
)

// loadSignatureRawVectors returns the raw json.RawMessage view of
// signature-test.json, failing the test if the file is unreadable or
// malformed.
func loadSignatureRawVectors(t *testing.T) []json.RawMessage {
t.Helper()

raw, err := loadSignatureRawOnce()
require.NoError(t, err)

return raw
}
8 changes: 5 additions & 3 deletions bolt12/invoice.go
Original file line number Diff line number Diff line change
Expand Up @@ -338,14 +338,16 @@ func DecodeInvoice(data []byte) (*Invoice, error) {
data,
invreqMetadata.Record(), chains.Record(), offerMeta.Record(),
currency.Record(), offerAmt.Record(), desc.Record(),
offerFeat.Record(), expiry.Record(), offerPaths.Record(),
strictFeaturesRecord(&offerFeat), expiry.Record(),
offerPaths.Record(),
issuer.Record(), qtyMax.Record(), issuerID.Record(),
invreqChain.Record(), invreqAmt.Record(), invreqFeat.Record(),
invreqChain.Record(), invreqAmt.Record(),
strictFeaturesRecord(&invreqFeat),
invreqQty.Record(), payerID.Record(), payerNote.Record(),
invreqPaths.Record(), bip353.Record(), invPaths.Record(),
blindedPay.Record(), createdAt.Record(), relExp.Record(),
payHash.Record(), invAmt.Record(), fallbacks.Record(),
invFeat.Record(), nodeID.Record(), sig.Record(),
strictFeaturesRecord(&invFeat), nodeID.Record(), sig.Record(),
)
if err != nil {
return nil, fmt.Errorf("decode invoice: %w", err)
Expand Down
5 changes: 3 additions & 2 deletions bolt12/invoice_request.go
Original file line number Diff line number Diff line change
Expand Up @@ -198,10 +198,11 @@ func DecodeInvoiceRequest(data []byte) (*InvoiceRequest, error) {
tm, err := decodeStream(
data, invreqMetadata.Record(), chains.Record(),
metadata.Record(), currency.Record(), amount.Record(),
desc.Record(), features.Record(), expiry.Record(),
desc.Record(), strictFeaturesRecord(&features), expiry.Record(),
paths.Record(), issuer.Record(), qtyMax.Record(),
issuerID.Record(), invreqChain.Record(), invreqAmount.Record(),
invreqFeatures.Record(), invreqQty.Record(), payerID.Record(),
strictFeaturesRecord(&invreqFeatures), invreqQty.Record(),
payerID.Record(),
payerNote.Record(), invreqPaths.Record(), bip353.Record(),
sig.Record(),
)
Expand Down
8 changes: 7 additions & 1 deletion bolt12/invoice_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -185,8 +185,14 @@ func TestInvoiceRoundTripPreservesAllTypes(t *testing.T) {
t.Parallel()

inv := validInvoice(t)

// Sign with the fixture's node id (Bob) so the read path's signature
// check accepts the invoice.
priv, _ := bobKey()
sig, err := SignInvoice(inv, priv)
require.NoError(t, err)
inv.Signature = tlv.SomeRecordT(
tlv.NewPrimitiveRecord[tlv.TlvType240, [64]byte]([64]byte{}),
tlv.NewPrimitiveRecord[tlv.TlvType240](sig),
)

encoded, err := inv.Encode()
Expand Down
Loading
Loading