diff --git a/chanbackup/backup_test.go b/chanbackup/backup_test.go index 2c324caa701..05a24090c0a 100644 --- a/chanbackup/backup_test.go +++ b/chanbackup/backup_test.go @@ -160,6 +160,7 @@ func TestFetchStaticChanBackups(t *testing.T) { chanSource.addAddrsForNode(randomChan2.IdentityPub, []net.Addr{addr2}) chanSource.addAddrsForNode(randomChan2.IdentityPub, []net.Addr{addr3}) chanSource.addAddrsForNode(randomChan2.IdentityPub, []net.Addr{addr4}) + chanSource.addAddrsForNode(randomChan2.IdentityPub, []net.Addr{addr5}) // With the channel source populated, we'll now attempt to create a set // of backups for all the channels. This should succeed, as all items diff --git a/chanbackup/multi_test.go b/chanbackup/multi_test.go index 3350c773cf4..d84ca2f111b 100644 --- a/chanbackup/multi_test.go +++ b/chanbackup/multi_test.go @@ -23,7 +23,7 @@ func TestMultiPackUnpack(t *testing.T) { } single := NewSingle( - channel, []net.Addr{addr1, addr2, addr3, addr4}, + channel, []net.Addr{addr1, addr2, addr3, addr4, addr5}, ) originalSingles = append(originalSingles, single) diff --git a/chanbackup/single_test.go b/chanbackup/single_test.go index b3c7471a36b..f1f805c1435 100644 --- a/chanbackup/single_test.go +++ b/chanbackup/single_test.go @@ -41,7 +41,11 @@ var ( OnionService: "3g2upl4pq6kufc4m.onion", Port: 9735, } - addr4 = &lnwire.OpaqueAddrs{ + addr4 = &lnwire.DNSAddress{ + Hostname: "example.com", + Port: 8080, + } + addr5 = &lnwire.OpaqueAddrs{ // The first byte must be an address type we are not yet aware // of for it to be a valid OpaqueAddrs. Payload: []byte{math.MaxUint8, 1, 2, 3, 4}, @@ -320,7 +324,7 @@ func TestSinglePackUnpack(t *testing.T) { require.NoError(t, err, "unable to gen open channel") singleChanBackup := NewSingle( - channel, []net.Addr{addr1, addr2, addr3, addr4}, + channel, []net.Addr{addr1, addr2, addr3, addr4, addr5}, ) keyRing := &lnencrypt.MockKeyRing{} @@ -647,7 +651,7 @@ func TestSingleUnconfirmedChannel(t *testing.T) { channel.FundingBroadcastHeight = fundingBroadcastHeight singleChanBackup := NewSingle( - channel, []net.Addr{addr1, addr2, addr3, addr4}, + channel, []net.Addr{addr1, addr2, addr3, addr4, addr5}, ) keyRing := &lnencrypt.MockKeyRing{} diff --git a/discovery/chan_series.go b/discovery/chan_series.go index ed7140d7d81..cf502694c97 100644 --- a/discovery/chan_series.go +++ b/discovery/chan_series.go @@ -231,6 +231,13 @@ func (c *ChanSeries) UpdatesInHorizon(chain chainhash.Hash, return nil, err } + if err := netann.ValidateNodeAnnFields(nodeUpdate); err != nil { + log.Debugf("Skipping forwarding invalid node "+ + "announcement %x: %v", nodeAnn.PubKeyBytes, err) + + continue + } + updates = append(updates, nodeUpdate) } @@ -282,6 +289,7 @@ func (c *ChanSeries) FilterChannelRange(_ chainhash.Hash, startHeight, // to reply to a QueryShortChanIDs message sent by a remote peer. The response // will contain a unique set of ChannelAnnouncements, the latest ChannelUpdate // for each of the announcements, and a unique set of NodeAnnouncements. +// Invalid node announcements are skipped and logged for debugging purposes. // // NOTE: This is part of the ChannelGraphTimeSeries interface. func (c *ChanSeries) FetchChanAnns(chain chainhash.Hash, @@ -335,8 +343,15 @@ func (c *ChanSeries) FetchChanAnns(chain chainhash.Hash, return nil, err } - chanAnns = append(chanAnns, nodeAnn) - nodePubsSent[nodePub] = struct{}{} + err = netann.ValidateNodeAnnFields(nodeAnn) + if err != nil { + log.Debugf("Skipping forwarding "+ + "invalid node announcement "+ + "%x: %v", nodeAnn.NodeID, err) + } else { + chanAnns = append(chanAnns, nodeAnn) + nodePubsSent[nodePub] = struct{}{} + } } } if edge2 != nil { @@ -354,8 +369,15 @@ func (c *ChanSeries) FetchChanAnns(chain chainhash.Hash, return nil, err } - chanAnns = append(chanAnns, nodeAnn) - nodePubsSent[nodePub] = struct{}{} + err = netann.ValidateNodeAnnFields(nodeAnn) + if err != nil { + log.Debugf("Skipping forwarding "+ + "invalid node announcement "+ + "%x: %v", nodeAnn.NodeID, err) + } else { + chanAnns = append(chanAnns, nodeAnn) + nodePubsSent[nodePub] = struct{}{} + } } } } diff --git a/discovery/gossiper.go b/discovery/gossiper.go index c720a612967..f24286ad928 100644 --- a/discovery/gossiper.go +++ b/discovery/gossiper.go @@ -2229,7 +2229,9 @@ func (d *AuthenticatedGossiper) processZombieUpdate(_ context.Context, } // fetchNodeAnn fetches the latest signed node announcement from our point of -// view for the node with the given public key. +// view for the node with the given public key. It also validates the node +// announcement fields and returns an error if they are invalid to prevent +// forwarding invalid node announcements to our peers. func (d *AuthenticatedGossiper) fetchNodeAnn(ctx context.Context, pubKey [33]byte) (*lnwire.NodeAnnouncement, error) { @@ -2238,7 +2240,12 @@ func (d *AuthenticatedGossiper) fetchNodeAnn(ctx context.Context, return nil, err } - return node.NodeAnnouncement(true) + nodeAnn, err := node.NodeAnnouncement(true) + if err != nil { + return nil, err + } + + return nodeAnn, netann.ValidateNodeAnnFields(nodeAnn) } // isMsgStale determines whether a message retrieved from the backing diff --git a/docs/release-notes/release-notes-0.20.0.md b/docs/release-notes/release-notes-0.20.0.md index be853084fda..d6d077d5dfb 100644 --- a/docs/release-notes/release-notes-0.20.0.md +++ b/docs/release-notes/release-notes-0.20.0.md @@ -227,6 +227,11 @@ reader of a payment request. * [Require invoices to include a payment address or blinded paths](https://github.com/lightningnetwork/lnd/pull/9752) to comply with updated BOLT 11 specifications before sending payments. +* [LND can now recgonize DNS address type in node + announcement msg](https://github.com/lightningnetwork/lnd/pull/9455). This + allows users to forward node announcement with valid DNS address types. The + validity aligns with Bolt 07 DNS constraints. + ## Testing * Previously, automatic peer bootstrapping was disabled for simnet, signet and diff --git a/graph/db/addr.go b/graph/db/addr.go index c68039a2624..836d516b000 100644 --- a/graph/db/addr.go +++ b/graph/db/addr.go @@ -31,8 +31,36 @@ const ( // opaqueAddrs denotes an address (or a set of addresses) that LND was // not able to parse since LND is not yet aware of the address type. opaqueAddrs addressType = 4 + + // dnsAddr denotes a DNS address type. + dnsAddr addressType = 5 ) +// encodeDNSAddr encodes a DNS address. +func encodeDNSAddr(w io.Writer, addr *lnwire.DNSAddress) error { + if _, err := w.Write([]byte{byte(dnsAddr)}); err != nil { + return err + } + + // Write the length of the hostname. + hostLen := len(addr.Hostname) + if _, err := w.Write([]byte{byte(hostLen)}); err != nil { + return err + } + + if _, err := w.Write([]byte(addr.Hostname)); err != nil { + return err + } + + var port [2]byte + byteOrder.PutUint16(port[:], addr.Port) + if _, err := w.Write(port[:]); err != nil { + return err + } + + return nil +} + // encodeTCPAddr serializes a TCP address into its compact raw bytes // representation. func encodeTCPAddr(w io.Writer, addr *net.TCPAddr) error { @@ -230,6 +258,30 @@ func DeserializeAddr(r io.Reader) (net.Addr, error) { Port: port, } + case dnsAddr: + // Read the length of the hostname. + var hostLen [1]byte + if _, err := r.Read(hostLen[:]); err != nil { + return nil, err + } + + // Read the hostname. + hostname := make([]byte, hostLen[0]) + if _, err := r.Read(hostname); err != nil { + return nil, err + } + + // Read the port. + var port [2]byte + if _, err := r.Read(port[:]); err != nil { + return nil, err + } + + address = &lnwire.DNSAddress{ + Hostname: string(hostname), + Port: binary.BigEndian.Uint16(port[:]), + } + case opaqueAddrs: // Read the length of the payload. var l [2]byte @@ -264,6 +316,8 @@ func SerializeAddr(w io.Writer, address net.Addr) error { return encodeOnionAddr(w, addr) case *lnwire.OpaqueAddrs: return encodeOpaqueAddrs(w, addr) + case *lnwire.DNSAddress: + return encodeDNSAddr(w, addr) default: return ErrUnknownAddressType } diff --git a/graph/db/addr_test.go b/graph/db/addr_test.go index b3dbea8cba1..d3c3700a33d 100644 --- a/graph/db/addr_test.go +++ b/graph/db/addr_test.go @@ -6,6 +6,7 @@ import ( "strings" "testing" + "github.com/lightningnetwork/lnd/lnwire" "github.com/lightningnetwork/lnd/tor" "github.com/stretchr/testify/require" ) @@ -38,6 +39,18 @@ var ( OnionService: "vww6ybal4bd7szmgncyruucpgfkqahzddi37ktceo3ah7ngmcopnpyyd.onion", //nolint:ll Port: 80, } + + testOpaqueAddr = &lnwire.OpaqueAddrs{ + // NOTE: the first byte is a protocol level address type. So + // for we set it to 0xff to guarantee that we do not know this + // type yet. + Payload: []byte{0xff, 0x02, 0x03, 0x04, 0x05, 0x06}, + } + + testDNSAddr = &lnwire.DNSAddress{ + Hostname: "example.com", + Port: 8080, + } ) var addrTests = []struct { @@ -57,6 +70,12 @@ var addrTests = []struct { { expAddr: testOnionV3Addr, }, + { + expAddr: testOpaqueAddr, + }, + { + expAddr: testDNSAddr, + }, // Invalid addresses. { diff --git a/graph/db/graph_test.go b/graph/db/graph_test.go index c158ea5d9f0..0d55ef62580 100644 --- a/graph/db/graph_test.go +++ b/graph/db/graph_test.go @@ -37,10 +37,7 @@ var ( Port: 9000} anotherAddr, _ = net.ResolveTCPAddr("tcp", "[2001:db8:85a3:0:0:8a2e:370:7334]:80") - testAddrs = []net.Addr{testAddr, anotherAddr} - testOpaqueAddr = &lnwire.OpaqueAddrs{ - Payload: []byte{0x01, 0x02, 0x03, 0x04, 0x05, 0x06}, - } + testAddrs = []net.Addr{testAddr, anotherAddr} testRBytes, _ = hex.DecodeString("8ce2bc69281ce27da07e6683571319d18" + "e949ddfa2965fb6caa1bf0314f882d7") @@ -208,6 +205,8 @@ func TestNodeInsertionAndDeletion(t *testing.T) { // Add one v2 and one v3 onion address. testOnionV2Addr, testOnionV3Addr, + // Add a DNS host address. + testDNSAddr, // Make sure to also test the opaque address type. testOpaqueAddr, } diff --git a/graph/db/sql_store.go b/graph/db/sql_store.go index 1789c70c4d9..c53a0fe5c32 100644 --- a/graph/db/sql_store.go +++ b/graph/db/sql_store.go @@ -3530,6 +3530,7 @@ const ( addressTypeIPv6 dbAddressType = 2 addressTypeTorV2 dbAddressType = 3 addressTypeTorV3 dbAddressType = 4 + addressTypeDNS dbAddressType = 5 addressTypeOpaque dbAddressType = math.MaxInt8 ) @@ -3545,6 +3546,7 @@ func collectAddressRecords(addresses []net.Addr) (map[dbAddressType][]string, addressTypeIPv6: {}, addressTypeTorV2: {}, addressTypeTorV3: {}, + addressTypeDNS: {}, addressTypeOpaque: {}, } addAddr := func(t dbAddressType, addr net.Addr) { @@ -3574,6 +3576,9 @@ func collectAddressRecords(addresses []net.Addr) (map[dbAddressType][]string, "a tor address") } + case *lnwire.DNSAddress: + addAddr(addressTypeDNS, addr) + case *lnwire.OpaqueAddrs: addAddr(addressTypeOpaque, addr) @@ -4663,6 +4668,23 @@ func parseAddress(addrType dbAddressType, address string) (net.Addr, error) { Port: port, }, nil + case addressTypeDNS: + hostname, portStr, err := net.SplitHostPort(address) + if err != nil { + return nil, fmt.Errorf("unable to split DNS "+ + "address: %v", address) + } + + port, err := strconv.Atoi(portStr) + if err != nil { + return nil, err + } + + return &lnwire.DNSAddress{ + Hostname: hostname, + Port: uint16(port), + }, nil + case addressTypeOpaque: opaque, err := hex.DecodeString(address) if err != nil { diff --git a/lnwire/dns_addr.go b/lnwire/dns_addr.go new file mode 100644 index 00000000000..87ddcd8cd2b --- /dev/null +++ b/lnwire/dns_addr.go @@ -0,0 +1,88 @@ +package lnwire + +import ( + "errors" + "fmt" + "net" + "strconv" +) + +var ( + // ErrEmptyDNSHostname is returned when a DNS hostname is empty. + ErrEmptyDNSHostname = errors.New("hostname cannot be empty") + + // ErrZeroPort is returned when a DNS port is zero. + ErrZeroPort = errors.New("port cannot be zero") + + // ErrHostnameTooLong is returned when a DNS hostname exceeds 255 bytes. + ErrHostnameTooLong = errors.New("DNS hostname length exceeds limit " + + "of 255 bytes") + + // ErrInvalidHostnameCharacter is returned when a DNS hostname contains + // an invalid character. + ErrInvalidHostnameCharacter = errors.New("hostname contains invalid " + + "character") +) + +// DNSAddress is used to represent a DNS address of a node. +type DNSAddress struct { + // Hostname is the DNS hostname of the address. This MUST only contain + // ASCII characters as per Bolt #7. The maximum length that this may + // be is 255 bytes. + Hostname string + + // Port is the port number of the address. + Port uint16 +} + +// A compile-time check to ensure that DNSAddress implements the net.Addr +// interface. +var _ net.Addr = (*DNSAddress)(nil) + +// Network returns the network that this address uses, which is "tcp". +func (d *DNSAddress) Network() string { + return "tcp" +} + +// String returns the address in the form "hostname:port". +func (d *DNSAddress) String() string { + return net.JoinHostPort(d.Hostname, strconv.Itoa(int(d.Port))) +} + +// ValidateDNSAddr validates that the DNS hostname is not empty and contains +// only ASCII characters and of max length 255 characters and port is non zero +// according to BOLT #7. +func ValidateDNSAddr(hostname string, port uint16) error { + if hostname == "" { + return ErrEmptyDNSHostname + } + + // Per BOLT 7, ports must not be zero for type 5 address (DNS address). + if port == 0 { + return ErrZeroPort + } + + if len(hostname) > 255 { + return fmt.Errorf("%w: DNS hostname length %d", + ErrHostnameTooLong, len(hostname)) + } + + // Check if hostname contains only ASCII characters. + for i, r := range hostname { + // Check for valid hostname characters, excluding ASCII control + // characters (0-31), spaces, underscores, delete character + // (127), and the special characters (like /, \, @, #, $, etc.). + if !((r >= 'a' && r <= 'z') || + (r >= 'A' && r <= 'Z') || + (r >= '0' && r <= '9') || + r == '-' || + r == '.') { + + return fmt.Errorf("%w: hostname '%s' contains invalid "+ + "character '%c' at position %d", + ErrInvalidHostnameCharacter, hostname, r, i) + } + } + + return nil +} diff --git a/lnwire/dns_addr_test.go b/lnwire/dns_addr_test.go new file mode 100644 index 00000000000..8cdf2a858c6 --- /dev/null +++ b/lnwire/dns_addr_test.go @@ -0,0 +1,87 @@ +package lnwire + +import ( + "fmt" + "strings" + "testing" + + "github.com/stretchr/testify/require" +) + +// TestValidateDNSAddr tests hostname and port validation per BOLT #7. +func TestValidateDNSAddr(t *testing.T) { + t.Parallel() + + testCases := []struct { + name string + hostname string + port uint16 + err error + }{ + { + name: "empty hostname", + hostname: "", + port: 9735, + err: ErrEmptyDNSHostname, + }, + { + name: "zero port", + hostname: "example.com", + port: 0, + err: ErrZeroPort, + }, + { + name: "hostname too long", + hostname: strings.Repeat("a", 256), + port: 9735, + err: fmt.Errorf("%w: DNS hostname length 256", + ErrHostnameTooLong), + }, + { + name: "hostname with invalid ASCII space", + hostname: "exa mple.com", + port: 9735, + err: fmt.Errorf("%w: hostname 'exa mple.com' contains "+ + "invalid character ' ' at position 3", + ErrInvalidHostnameCharacter), + }, + { + name: "hostname with invalid ASCII underscore", + hostname: "example_node.com", + port: 9735, + err: fmt.Errorf("%w: hostname 'example_node.com' "+ + "contains invalid character '_' at position 7", + ErrInvalidHostnameCharacter), + }, + { + name: "hostname with non-ASCII character", + hostname: "example❄️.com", + port: 9735, + err: fmt.Errorf("%w: hostname 'example❄️.com' "+ + "contains invalid character '❄' at position 7", + ErrInvalidHostnameCharacter), + }, + { + name: "valid hostname", + hostname: "example.com", + port: 9735, + }, + { + name: "valid hostname with numbers", + hostname: "node101.example.com", + port: 9735, + }, + { + name: "valid hostname with hyphens", + hostname: "my-node.example-domain.com", + port: 9735, + }, + } + + for _, tc := range testCases { + t.Run(tc.name, func(t *testing.T) { + err := ValidateDNSAddr(tc.hostname, tc.port) + require.Equal(t, err, tc.err) + }) + } +} diff --git a/lnwire/lnwire.go b/lnwire/lnwire.go index 7d686658297..86f7caa6173 100644 --- a/lnwire/lnwire.go +++ b/lnwire/lnwire.go @@ -28,6 +28,28 @@ const ( MaxMsgBody = 65533 ) +const ( + // tcp4AddrLen is the length of an IPv4 address + // (4 bytes IP + 2 bytes port). + tcp4AddrLen = 6 + + // tcp6AddrLen is the length of an IPv6 address + // (16 bytes IP + 2 bytes port). + tcp6AddrLen = 18 + + // v2OnionAddrLen is the length of a version 2 Tor onion service + // address. + v2OnionAddrLen = 12 + + // v3OnionAddrLen is the length of a version 3 Tor onion service address + // (35 bytes decoded onion + 2 bytes port). + v3OnionAddrLen = 37 + + // dnsAddrOverhead is the fixed overhead for a DNS address: 1 byte for + // the hostname length and 2 bytes for the port. + dnsAddrOverhead = 3 +) + // PkScript is simple type definition which represents a raw serialized public // key script. type PkScript []byte @@ -52,26 +74,10 @@ const ( // v3OnionAddr denotes a version 3 Tor (prop224) onion service address. v3OnionAddr addressType = 4 -) -// AddrLen returns the number of bytes that it takes to encode the target -// address. -func (a addressType) AddrLen() uint16 { - switch a { - case noAddr: - return 0 - case tcp4Addr: - return 6 - case tcp6Addr: - return 18 - case v2OnionAddr: - return 12 - case v3OnionAddr: - return 37 - default: - return 0 - } -} + // dnsAddr denotes a DNS address. + dnsAddr addressType = 5 +) // WriteElement is a one-stop shop to write the big endian representation of // any element which is to be serialized for the wire protocol. @@ -348,6 +354,11 @@ func WriteElement(w *bytes.Buffer, element interface{}) error { return err } + case *DNSAddress: + if err := WriteDNSAddress(w, e); err != nil { + return err + } + case *OpaqueAddrs: if err := WriteOpaqueAddrs(w, e); err != nil { return err @@ -743,7 +754,6 @@ func ReadElement(r io.Reader, element interface{}) error { var address net.Addr switch aType := addressType(descriptor[0]); aType { case noAddr: - addrBytesRead += aType.AddrLen() continue case tcp4Addr: @@ -761,7 +771,7 @@ func ReadElement(r io.Reader, element interface{}) error { IP: net.IP(ip[:]), Port: int(binary.BigEndian.Uint16(port[:])), } - addrBytesRead += aType.AddrLen() + addrBytesRead += tcp4AddrLen case tcp6Addr: var ip [16]byte @@ -778,7 +788,7 @@ func ReadElement(r io.Reader, element interface{}) error { IP: net.IP(ip[:]), Port: int(binary.BigEndian.Uint16(port[:])), } - addrBytesRead += aType.AddrLen() + addrBytesRead += tcp6AddrLen case v2OnionAddr: var h [tor.V2DecodedLen]byte @@ -799,7 +809,7 @@ func ReadElement(r io.Reader, element interface{}) error { OnionService: onionService, Port: port, } - addrBytesRead += aType.AddrLen() + addrBytesRead += v2OnionAddrLen case v3OnionAddr: var h [tor.V3DecodedLen]byte @@ -820,7 +830,35 @@ func ReadElement(r io.Reader, element interface{}) error { OnionService: onionService, Port: port, } - addrBytesRead += aType.AddrLen() + addrBytesRead += v3OnionAddrLen + + case dnsAddr: + var hostnameLen [1]byte + _, err := io.ReadFull(addrBuf, hostnameLen[:]) + if err != nil { + return err + } + + hostname := make([]byte, hostnameLen[0]) + _, err = io.ReadFull(addrBuf, hostname) + if err != nil { + return err + } + + var port [2]byte + _, err = io.ReadFull(addrBuf, port[:]) + if err != nil { + return err + } + + address = &DNSAddress{ + Hostname: string(hostname), + Port: binary.BigEndian.Uint16( + port[:], + ), + } + addrBytesRead += dnsAddrOverhead + + uint16(len(hostname)) default: // If we don't understand this address type, diff --git a/lnwire/message_test.go b/lnwire/message_test.go index d42d5791152..4f95cb0efa1 100644 --- a/lnwire/message_test.go +++ b/lnwire/message_test.go @@ -1000,13 +1000,32 @@ func randV3OnionAddr(t testing.TB, r *rand.Rand) *tor.OnionAddr { return &tor.OnionAddr{OnionService: onionService, Port: addrPort} } +// randDNSAddr generates a random DNS address for testing purposes. +func randDNSAddr(t testing.TB, r *rand.Rand) *lnwire.DNSAddress { + t.Helper() + + var domain [1]byte + _, err := r.Read(domain[:]) + require.NoError(t, err) + + var port [2]byte + _, err = r.Read(port[:]) + require.NoError(t, err, "unable to read port") + + return &lnwire.DNSAddress{ + Hostname: string(domain[:]), + Port: uint16(port[0]), + } +} + func randAddrs(t testing.TB, r *rand.Rand) []net.Addr { tcp4Addr := randTCP4Addr(t, r) tcp6Addr := randTCP6Addr(t, r) v2OnionAddr := randV2OnionAddr(t, r) v3OnionAddr := randV3OnionAddr(t, r) + dnsAddr := randDNSAddr(t, r) - return []net.Addr{tcp4Addr, tcp6Addr, v2OnionAddr, v3OnionAddr} + return []net.Addr{tcp4Addr, tcp6Addr, v2OnionAddr, v3OnionAddr, dnsAddr} } func randAlias(r *rand.Rand) lnwire.NodeAlias { diff --git a/lnwire/writer.go b/lnwire/writer.go index fa6247de0b0..ddf67e8a289 100644 --- a/lnwire/writer.go +++ b/lnwire/writer.go @@ -35,6 +35,9 @@ var ( // ErrNilOpaqueAddrs is returned when the supplied address is nil. ErrNilOpaqueAddrs = errors.New("cannot write nil OpaqueAddrs") + // ErrNilDNSAddress is returned when the supplied address is nil. + ErrNilDNSAddress = errors.New("cannot write nil DNS address") + // ErrNilPublicKey is returned when a nil pubkey is used. ErrNilPublicKey = errors.New("cannot write nil pubkey") @@ -364,6 +367,28 @@ func WriteOnionAddr(buf *bytes.Buffer, addr *tor.OnionAddr) error { return WriteUint16(buf, uint16(addr.Port)) } +// WriteDNSAddress appends the DNS address to the provided buffer. +func WriteDNSAddress(buf *bytes.Buffer, addr *DNSAddress) error { + if addr == nil { + return ErrNilDNSAddress + } + + // Write the descriptor, the hostname length, and the hostname. + if _, err := buf.Write([]byte{byte(dnsAddr)}); err != nil { + return err + } + + if err := WriteUint8(buf, uint8(len(addr.Hostname))); err != nil { + return err + } + + if _, err := buf.WriteString(addr.Hostname); err != nil { + return err + } + + return WriteUint16(buf, addr.Port) +} + // WriteOpaqueAddrs appends the payload of the given OpaqueAddrs to buffer. func WriteOpaqueAddrs(buf *bytes.Buffer, addr *OpaqueAddrs) error { if addr == nil { @@ -397,6 +422,10 @@ func WriteNetAddrs(buf *bytes.Buffer, addresses []net.Addr) error { if err := WriteOpaqueAddrs(addrBuf, a); err != nil { return err } + case *DNSAddress: + if err := WriteDNSAddress(addrBuf, a); err != nil { + return err + } default: return ErrNilNetAddress } diff --git a/lnwire/writer_test.go b/lnwire/writer_test.go index 3e2550443e5..bb2bada06ad 100644 --- a/lnwire/writer_test.go +++ b/lnwire/writer_test.go @@ -559,7 +559,8 @@ func TestWriteOnionAddr(t *testing.T) { } func TestWriteNetAddrs(t *testing.T) { - buf := new(bytes.Buffer) + t.Parallel() + tcpAddr := &net.TCPAddr{ IP: net.IP{127, 0, 0, 1}, Port: 8080, @@ -568,6 +569,10 @@ func TestWriteNetAddrs(t *testing.T) { OnionService: "abcdefghijklmnop.onion", Port: 9065, } + dnsAddr := &DNSAddress{ + Hostname: "example.com", + Port: 8080, + } testCases := []struct { name string @@ -587,24 +592,27 @@ func TestWriteNetAddrs(t *testing.T) { { // Check empty address slice. name: "empty address slice", - addr: []net.Addr{}, expectedErr: nil, // Use two bytes to encode the address size. expectedBytes: []byte{0, 0}, }, { // Check a successful writes of a slice of addresses. - name: "two addresses", - addr: []net.Addr{tcpAddr, onionAddr}, + name: "multiple addresses", + addr: []net.Addr{tcpAddr, onionAddr, dnsAddr}, expectedErr: nil, expectedBytes: []byte{ - // 7 bytes for TCP and 13 bytes for onion. - 0x0, 0x14, + // 7 bytes for TCP and 13 bytes for onion, + // 15 bytes for DNS. + 0x0, 0x23, // TCP address. 0x1, 0x7f, 0x0, 0x0, 0x1, 0x1f, 0x90, // Onion address. 0x3, 0x0, 0x44, 0x32, 0x14, 0xc7, 0x42, 0x54, 0xb6, 0x35, 0xcf, 0x23, 0x69, + // DNS address. + 0x5, 0xb, 0x65, 0x78, 0x61, 0x6d, 0x70, 0x6c, + 0x65, 0x2e, 0x63, 0x6f, 0x6d, 0x1f, 0x90, }, }, } @@ -612,13 +620,25 @@ func TestWriteNetAddrs(t *testing.T) { for _, tc := range testCases { tc := tc t.Run(tc.name, func(t *testing.T) { - oldLen := buf.Len() + buf := new(bytes.Buffer) err := WriteNetAddrs(buf, tc.addr) require.Equal(t, tc.expectedErr, err) - bytesWritten := buf.Bytes()[oldLen:buf.Len()] + bytesWritten := buf.Bytes()[:buf.Len()] require.Equal(t, tc.expectedBytes, bytesWritten) + + if tc.expectedErr != nil { + return + } + + // Read the addresses from the buffer and ensure + // they match the original addresses. + var addrs []net.Addr + err = ReadElement(buf, &addrs) + require.NoError(t, err) + + require.Equal(t, tc.addr, addrs) }) } } diff --git a/netann/node_announcement.go b/netann/node_announcement.go index 71250217057..893e9fc4af7 100644 --- a/netann/node_announcement.go +++ b/netann/node_announcement.go @@ -2,6 +2,7 @@ package netann import ( "bytes" + "errors" "fmt" "image/color" "net" @@ -81,10 +82,45 @@ func SignNodeAnnouncement(signer lnwallet.MessageSigner, return err } -// ValidateNodeAnn validates the node announcement by ensuring that the +// ValidateNodeAnn validates the fields and signature of a node announcement. +func ValidateNodeAnn(a *lnwire.NodeAnnouncement) error { + err := ValidateNodeAnnFields(a) + if err != nil { + return fmt.Errorf("invalid node announcement fields: %w", err) + } + + return ValidateNodeAnnSignature(a) +} + +// ValidateNodeAnnFields validates the fields of a node announcement. +func ValidateNodeAnnFields(a *lnwire.NodeAnnouncement) error { + // Check that it only has at most one DNS address. + hasDNSAddr := false + for _, addr := range a.Addresses { + dnsAddr, ok := addr.(*lnwire.DNSAddress) + if !ok { + continue + } + if hasDNSAddr { + return errors.New("node announcement contains " + + "multiple DNS addresses. Only one is allowed") + } + + hasDNSAddr = true + + err := lnwire.ValidateDNSAddr(dnsAddr.Hostname, dnsAddr.Port) + if err != nil { + return err + } + } + + return nil +} + +// ValidateNodeAnnSignature validates the node announcement by ensuring that the // attached signature is needed a signature of the node announcement under the // specified node public key. -func ValidateNodeAnn(a *lnwire.NodeAnnouncement) error { +func ValidateNodeAnnSignature(a *lnwire.NodeAnnouncement) error { // Reconstruct the data of announcement which should be covered by the // signature so we can verify the signature shortly below data, err := a.DataToSign()