diff --git a/v2transport/handshake_test.go b/v2transport/handshake_test.go new file mode 100644 index 00000000..23c3820f --- /dev/null +++ b/v2transport/handshake_test.go @@ -0,0 +1,313 @@ +package v2transport + +import ( + "bytes" + "errors" + "fmt" + "io" + "testing" + "time" +) + +// splitReadWriter keeps bytes read by the transport separate from bytes the +// transport writes in response. +type splitReadWriter struct { + reader *bytes.Reader + writes bytes.Buffer + beforeWrite func() +} + +// testHandshakeAdmission adapts a function to HandshakeAdmission so each test +// can record when the responder enters and leaves the CPU-bound phase. +type testHandshakeAdmission func() (func(), error) + +// Acquire invokes the test admission function. +func (a testHandshakeAdmission) Acquire() (func(), error) { + return a() +} + +func newSplitReadWriter(input []byte) *splitReadWriter { + return &splitReadWriter{reader: bytes.NewReader(input)} +} + +func (rw *splitReadWriter) Read(p []byte) (int, error) { + return rw.reader.Read(p) +} + +func (rw *splitReadWriter) Write(p []byte) (int, error) { + if rw.beforeWrite != nil { + rw.beforeWrite() + } + + return rw.writes.Write(p) +} + +// bufferedReadWriter is one endpoint of a buffered in-memory duplex pipe. +type bufferedReadWriter struct { + recv <-chan []byte + send chan<- []byte + buf bytes.Buffer +} + +func newBufferedReadWriterPair() (*bufferedReadWriter, *bufferedReadWriter) { + leftToRight := make(chan []byte, 32) + rightToLeft := make(chan []byte, 32) + + return &bufferedReadWriter{ + recv: rightToLeft, + send: leftToRight, + }, &bufferedReadWriter{ + recv: leftToRight, + send: rightToLeft, + } +} + +func (rw *bufferedReadWriter) Read(p []byte) (int, error) { + if rw.buf.Len() == 0 { + rw.buf.Write(<-rw.recv) + } + + return rw.buf.Read(p) +} + +func (rw *bufferedReadWriter) Write(p []byte) (int, error) { + msg := append([]byte(nil), p...) + rw.send <- msg + + return len(p), nil +} + +// TestResponderV1Fallback verifies a full v1 prefix does not consume CPU +// admission or generate responder key material. +func TestResponderV1Fallback(t *testing.T) { + var admissions int + peer := NewPeerWithOptions(WithResponderHandshakeAdmission( + testHandshakeAdmission(func() (func(), error) { + admissions++ + return func() {}, nil + }), + )) + rw := newSplitReadWriter(createV1Prefix(BitcoinNet(0xd9b4bef9))) + peer.UseReadWriter(rw) + + err := peer.RespondV2Handshake(0, BitcoinNet(0xd9b4bef9)) + if !errors.Is(err, ErrUseV1Protocol) { + t.Fatalf("unexpected fallback error: %v", err) + } + if admissions != 0 { + t.Fatalf("v1 fallback consumed %d CPU admissions", admissions) + } + if peer.privkeyOurs != nil { + t.Fatal("v1 fallback generated responder key material") + } + if rw.writes.Len() != 0 { + t.Fatalf("v1 fallback wrote %d bytes", rw.writes.Len()) + } +} + +// TestResponderIncompleteCandidate verifies incomplete v2 candidates do not +// consume CPU admission or generate responder key material. +func TestResponderIncompleteCandidate(t *testing.T) { + for _, length := range []int{1, 15, 16, 63} { + t.Run(fmt.Sprintf("length_%d", length), func(t *testing.T) { + var admissions int + peer := NewPeerWithOptions(WithResponderHandshakeAdmission( + testHandshakeAdmission(func() (func(), error) { + admissions++ + return func() {}, nil + }), + )) + candidate := make([]byte, length) + candidate[0] = 1 + rw := newSplitReadWriter(candidate) + peer.UseReadWriter(rw) + + err := peer.RespondV2Handshake(0, BitcoinNet(0xd9b4bef9)) + if err == nil { + t.Fatal("incomplete candidate unexpectedly succeeded") + } + if admissions != 0 { + t.Fatalf("incomplete candidate consumed %d admissions", + admissions) + } + if peer.privkeyOurs != nil { + t.Fatal("incomplete candidate generated responder key") + } + if rw.writes.Len() != 0 { + t.Fatalf("incomplete candidate wrote %d bytes", + rw.writes.Len()) + } + }) + } +} + +// TestResponderWrongNetworkV1 verifies wrong-network v1 is rejected before +// CPU admission and key generation. +func TestResponderWrongNetworkV1(t *testing.T) { + const expectedNet = BitcoinNet(0xd9b4bef9) + + candidate := make([]byte, 64) + copy(candidate, createV1Prefix(BitcoinNet(0x0709110b))) + + var admissions int + peer := NewPeerWithOptions(WithResponderHandshakeAdmission( + testHandshakeAdmission(func() (func(), error) { + admissions++ + return func() {}, nil + }), + )) + rw := newSplitReadWriter(candidate) + peer.UseReadWriter(rw) + + err := peer.RespondV2Handshake(0, expectedNet) + if !errors.Is(err, errWrongNetV1Peer) { + t.Fatalf("unexpected wrong-network error: %v", err) + } + if admissions != 0 { + t.Fatalf("wrong-network v1 consumed %d admissions", admissions) + } + if peer.privkeyOurs != nil { + t.Fatal("wrong-network v1 generated responder key") + } +} + +// TestResponderAdmissionRejected verifies denial occurs before all responder +// cryptography and network writes. +func TestResponderAdmissionRejected(t *testing.T) { + errRejected := errors.New("rejected") + candidate := make([]byte, 64) + for i := range candidate { + candidate[i] = byte(i) + } + + var admissions int + peer := NewPeerWithOptions(WithResponderHandshakeAdmission( + testHandshakeAdmission(func() (func(), error) { + admissions++ + return nil, errRejected + }), + )) + rw := newSplitReadWriter(candidate) + peer.UseReadWriter(rw) + + err := peer.RespondV2Handshake(0, BitcoinNet(0xd9b4bef9)) + if !errors.Is(err, errRejected) { + t.Fatalf("unexpected admission error: %v", err) + } + if admissions != 1 { + t.Fatalf("admission called %d times, want 1", admissions) + } + if peer.privkeyOurs != nil { + t.Fatal("rejected candidate generated responder key") + } + if rw.writes.Len() != 0 { + t.Fatalf("rejected candidate wrote %d bytes", rw.writes.Len()) + } +} + +// TestResponderAdmissionLease verifies the admission lease covers responder +// cryptography and is released before the response is written. +func TestResponderAdmissionLease(t *testing.T) { + candidate := make([]byte, 64) + for i := range candidate { + candidate[i] = byte(i) + } + + var ( + admissions int + releases int + ) + peer := NewPeerWithOptions(WithResponderHandshakeAdmission( + testHandshakeAdmission(func() (func(), error) { + admissions++ + return func() { releases++ }, nil + }), + )) + rw := newSplitReadWriter(candidate) + rw.beforeWrite = func() { + if releases != 1 { + t.Fatalf("response write began before lease release: got %d", releases) + } + } + peer.UseReadWriter(rw) + + if err := peer.RespondV2Handshake(0, BitcoinNet(0xd9b4bef9)); err != nil { + t.Fatalf("responder handshake failed: %v", err) + } + if admissions != 1 || releases != 1 { + t.Fatalf("unexpected lease counts: admissions=%d releases=%d", + admissions, releases) + } + if peer.privkeyOurs == nil || !peer.responderReady { + t.Fatal("accepted responder did not initialize key material") + } + if rw.writes.Len() != 64 { + t.Fatalf("responder wrote %d bytes, want 64", rw.writes.Len()) + } +} + +// TestResponderHandshakeInteroperability verifies the refactored responder +// ordering preserves the complete v2 wire transcript. +func TestResponderHandshakeInteroperability(t *testing.T) { + const testNet = BitcoinNet(0xd9b4bef9) + + initiatorRW, responderRW := newBufferedReadWriterPair() + initiator := NewPeer() + responder := NewPeer() + initiator.UseReadWriter(initiatorRW) + responder.UseReadWriter(responderRW) + + errs := make(chan error, 2) + go func() { + if err := initiator.InitiateV2Handshake(0); err != nil { + errs <- err + return + } + errs <- initiator.CompleteHandshake(true, []int{0, 8}, testNet) + }() + go func() { + if err := responder.RespondV2Handshake(0, testNet); err != nil { + errs <- err + return + } + errs <- responder.CompleteHandshake(false, []int{3}, testNet) + }() + + for i := 0; i < 2; i++ { + select { + case err := <-errs: + if err != nil { + t.Fatalf("handshake failed: %v", err) + } + case <-time.After(5 * time.Second): + t.Fatal("handshake timed out") + } + } + + payload := []byte("post-handshake payload") + if _, _, err := initiator.V2EncPacket(payload, nil, false); err != nil { + t.Fatalf("packet send failed: %v", err) + } + got, err := responder.V2ReceivePacket(nil) + if err != nil { + t.Fatalf("packet receive failed: %v", err) + } + if !bytes.Equal(got, payload) { + t.Fatalf("packet mismatch: got %x, want %x", got, payload) + } +} + +// TestSendShortWrite verifies short writes are surfaced to callers. +func TestSendShortWrite(t *testing.T) { + peer := NewPeer() + peer.UseReadWriter(shortReadWriter{}) + + if _, err := peer.Send([]byte{1, 2}); !errors.Is(err, io.ErrShortWrite) { + t.Fatalf("unexpected short-write error: %v", err) + } +} + +type shortReadWriter struct{} + +func (shortReadWriter) Read([]byte) (int, error) { return 0, io.EOF } +func (shortReadWriter) Write([]byte) (int, error) { return 1, nil } diff --git a/v2transport/transport.go b/v2transport/transport.go index d9500581..6dfd74e4 100644 --- a/v2transport/transport.go +++ b/v2transport/transport.go @@ -106,6 +106,30 @@ var ( ErrShouldDowngradeToV1 = fmt.Errorf("should downgrade to v1") ) +// HandshakeAdmission controls access to the CPU-bound portion of an inbound +// v2 handshake. +type HandshakeAdmission interface { + // Acquire reserves the resources needed to generate the responder key + // and perform key agreement. The returned function releases the + // reservation once those operations are complete. + Acquire() (release func(), err error) +} + +// PeerOption customizes a v2 transport peer. +type PeerOption func(*Peer) + +// WithResponderHandshakeAdmission installs an admission policy that runs +// after a responder has received a complete v2 key candidate and before it +// performs key generation or key agreement. +func WithResponderHandshakeAdmission( + admission HandshakeAdmission, +) PeerOption { + + return func(p *Peer) { + p.responderAdmission = admission + } +} + // Peer defines the components necessary for sending/receiving data over the v2 // transport. type Peer struct { @@ -120,8 +144,9 @@ type Peer struct { // 4095 bytes. sentGarbage []byte - // receivedPrefix is used to determine which transport protocol we're - // using. + // receivedPrefix contains the bytes consumed while classifying the + // transport. For a v2 responder, it eventually contains the initiator's + // complete ElligatorSwift encoding. receivedPrefix []byte // sendL is the cipher used to send encrypted packet lengths. @@ -166,14 +191,27 @@ type Peer struct { // rw is the underlying object that will be read from / written to in // calls to V2EncPacket and V2ReceivePacket. rw io.ReadWriter + + // responderAdmission optionally reserves resources for the CPU-bound + // portion of an inbound handshake. + responderAdmission HandshakeAdmission + + // responderReady is set after the responder has completed key agreement + // and initialized its packet ciphers. + responderReady bool } // NewPeer returns a new instance of Peer. func NewPeer() *Peer { + return NewPeerWithOptions() +} + +// NewPeerWithOptions returns a new Peer configured with the provided options. +func NewPeerWithOptions(options ...PeerOption) *Peer { // The keys (initiatorL, initiatorP, responderL, responderP) as well as // the sessionID must have space for the hkdf Expand-derived Reader to // work. - return &Peer{ + p := &Peer{ receivedPrefix: make([]byte, 0), initiatorL: make([]byte, 32), initiatorP: make([]byte, 32), @@ -181,6 +219,11 @@ func NewPeer() *Peer { responderP: make([]byte, 32), sessionID: make([]byte, 32), } + for _, option := range options { + option(p) + } + + return p } // createV2Ciphers constructs the packet-length and packet encryption ciphers. @@ -373,15 +416,15 @@ func (p *Peer) InitiateV2Handshake(garbageLen int) error { log.Debugf("Sending ellswift pubkey and garbage (total_len=%d)", len(data)) - p.Send(data) - - return nil + _, err = p.Send(data) + return err } -// RespondV2Handshake responds to the initiator, determines if the initiator -// wants to use the v2 protocol and if so returns our ElligatorSwift-encoded -// public key followed by our garbage data over. If the initiator does not want -// to use the v2 protocol, we'll instead revert to the v1 protocol. +// RespondV2Handshake determines whether the initiator wants to use the v2 +// protocol. For a v2 initiator, it receives the complete key candidate before +// reserving CPU resources, initializes the responder ciphers, and sends the +// responder's ElligatorSwift-encoded public key followed by garbage. For a v1 +// initiator, it returns ErrUseV1Protocol without performing v2 cryptography. func (p *Peer) RespondV2Handshake(garbageLen int, net BitcoinNet) error { v1Prefix := createV1Prefix(net) @@ -402,7 +445,7 @@ func (p *Peer) RespondV2Handshake(garbageLen int, net BitcoinNet) error { var receiveBytes []byte receiveBytes, _, err = p.Receive(1) if err != nil { - log.Errorf("Failed to receive byte for v1 prefix "+ + log.Debugf("Failed to receive byte for v1 prefix "+ "check: %v", err) return err } @@ -418,28 +461,7 @@ func (p *Peer) RespondV2Handshake(garbageLen int, net BitcoinNet) error { "prefix at index %d, assuming v2 peer", p.receivedPrefix[lastIdx], lastIdx) - p.privkeyOurs, p.ellswiftOurs, err = ellswift.EllswiftCreate() - if err != nil { - log.Errorf("Failed to create ellswift "+ - "keypair: %v", err) - return err - } - - log.Tracef("Created ellswift keypair, pubkey=%x", - p.ellswiftOurs) - - data, err := p.generateKeyAndGarbage(garbageLen) - if err != nil { - return err - } - - // Send over our ElligatorSwift-encoded pubkey followed - // by our randomly generated garbage. - log.Debugf("Sending ellswift pubkey and garbage "+ - "(total_len=%d)", len(data)) - p.Send(data) - - return nil + return p.prepareResponderHandshake(garbageLen, net) } } @@ -448,6 +470,90 @@ func (p *Peer) RespondV2Handshake(garbageLen int, net BitcoinNet) error { return ErrUseV1Protocol } +// prepareResponderHandshake receives the rest of the initiator's key and +// performs all CPU-bound responder setup without holding a resource admission +// across a network read or write. +func (p *Peer) prepareResponderHandshake(garbageLen int, net BitcoinNet) error { + remaining := 64 - len(p.receivedPrefix) + recvData, _, err := p.Receive(remaining) + if err != nil { + return err + } + p.receivedPrefix = append(p.receivedPrefix, recvData...) + + if len(p.receivedPrefix) != 64 { + return errInsufficientBytes + } + + v1Prefix := createV1Prefix(net) + if bytes.Equal(p.receivedPrefix[4:16], v1Prefix[4:16]) { + log.Warnf("Peer sent v1 version message for wrong network "+ + "(expected %v)", net) + return errWrongNetV1Peer + } + + release := func() {} + if p.responderAdmission != nil { + release, err = p.responderAdmission.Acquire() + if err != nil { + return err + } + if release == nil { + release = func() {} + } + } + + err = p.initializeResponder(net, release) + if err != nil { + return err + } + + data, err := p.generateKeyAndGarbage(garbageLen) + if err != nil { + return err + } + + log.Debugf("Sending ellswift pubkey and garbage (total_len=%d)", + len(data)) + _, err = p.Send(data) + + return err +} + +// initializeResponder performs the responder's CPU-bound key generation, key +// agreement, and cipher setup. The admission is released before this method +// returns so callers never hold it across a network write. +func (p *Peer) initializeResponder(net BitcoinNet, release func()) error { + defer release() + + var err error + p.privkeyOurs, p.ellswiftOurs, err = ellswift.EllswiftCreate() + if err != nil { + log.Errorf("Failed to create ellswift keypair: %v", err) + return err + } + + log.Tracef("Created ellswift keypair, pubkey=%x", p.ellswiftOurs) + + var ellswiftTheirs [64]byte + copy(ellswiftTheirs[:], p.receivedPrefix) + + ecdhSecret, err := ellswift.V2Ecdh( + p.privkeyOurs, ellswiftTheirs, p.ellswiftOurs, false, + ) + if err != nil { + log.Errorf("Failed to calculate ECDH shared secret: %v", err) + return err + } + + if err := p.createV2Ciphers(ecdhSecret[:], false, net); err != nil { + return err + } + + p.responderReady = true + return nil +} + // generateKeyAndGarbage returns a byte slice containing our ellswift-encoded // public key followed by the garbage we'll send over. func (p *Peer) generateKeyAndGarbage(garbageLen int) ([]byte, error) { @@ -496,80 +602,53 @@ func createV1Prefix(net BitcoinNet) []byte { return v1Prefix } -// CompleteHandshake finishes the v2 protocol negotiation and optionally sends -// decoy packets after sending the garbage terminator. -func (p *Peer) CompleteHandshake(initiating bool, decoyContentLens []int, - btcnet BitcoinNet) error { - - log.Debugf("Completing v2 handshake (initiating=%v, "+ - "num_decoys=%d, net=%v)", initiating, len(decoyContentLens), - btcnet) +// completeKeyExchange receives the remote key and initializes the packet +// ciphers. RespondV2Handshake performs this phase early for responders so its +// CPU admission can be released before any response is written. +func (p *Peer) completeKeyExchange(initiating bool, net BitcoinNet) error { + if p.responderReady { + return nil + } var receivedPrefix []byte if initiating { - log.Trace("Initiator expecting 64 bytes for peer's " + - "ellswift key") - + log.Trace("Initiator expecting 64 bytes for peer's ellswift key") receivedPrefix = make([]byte, 0, 16) } else { - // If we are the responder, we have already received bytes to - // compare against the v1 transport protocol's starting bytes. - // We have to account for these when reading the rest of the 64 - // bytes off the wire to properly parse the remote's - // ellswift-encoded public key. + // A responder might have already received bytes while comparing the + // start of the stream against a v1 version message. receivedPrefix = p.receivedPrefix - log.Tracef("Responder already has prefix_len=%d, expecting %d "+ - "more bytes for peer's ellswift key", - len(receivedPrefix), 64-len(receivedPrefix)) + "more bytes for peer's ellswift key", len(receivedPrefix), + 64-len(receivedPrefix)) } recvData, numRead, err := p.Receive(64 - len(receivedPrefix)) if err != nil { - // If we receive an error when reading off the wire and we read - // zero bytes, then we will reconnect to the peer using v1. - // There are several different errors that Receive can return - // that indicate we should reconnect. Instead of special-casing - // them all, just perform these checks if any error was - // returned. + // An initiator which receives no response most likely connected to a + // v1-only peer. Signal the caller to reconnect using v1. if numRead == 0 && initiating { - // The peer most likely attempted to parse our 64-byte - // elligator-swift key as a version message and failed - // when trying to parse the message header into - // something valid. In this case, return a special - // error that signals to the server that we can - // reconnect with the OG v1 scheme. - log.Debugf("Received transport error during " + - "v2 handshake, retying downgraded v1 " + - "connection.") - + log.Debugf("Received transport error during v2 handshake, " + + "retrying downgraded v1 connection.") p.shouldDowngradeToV1.Store(true) return ErrShouldDowngradeToV1 } - // If we are the recipient, we can fail. - log.Errorf("Failed to receive peer's ellswift key data: %v", - err) + log.Errorf("Failed to receive peer's ellswift key data: %v", err) return err } log.Tracef("Received %d bytes for peer's ellswift key", len(recvData)) var ellswiftTheirs [64]byte - if initiating { - // If we are initiating, read all 64 bytes into ellswiftTheirs. copy(ellswiftTheirs[:], recvData) } else { - // If we are the responder, then we need to account for the - // bytes already received as part of matching against the - // starting v1 transport bytes. We sanity check receivedPrefix - // in case it is too large for some reason. prefixLen := len(receivedPrefix) if prefixLen > 16 { - log.Errorf("Responder's received prefix length %d is "+ - "too large (> 16)", prefixLen) + log.Errorf("Responder's received prefix length %d is too "+ + "large (> 16)", prefixLen) return errPrefixTooLarge } @@ -580,29 +659,14 @@ func (p *Peer) CompleteHandshake(initiating bool, decoyContentLens []int, log.Tracef("Assembled peer's ellswift key: %x", ellswiftTheirs) - // Calculate the v1 protocol's message prefix and see if the bytes read - // read into ellswiftTheirs matches it. - v1Prefix := createV1Prefix(btcnet) - - // ellswiftTheirs should be at least 16 bytes if receive succeeded, but - // just in case, check the size. - if len(ellswiftTheirs) < 16 { - log.Errorf("Received insufficient bytes (%d) for "+ - - "ellswift key", len(ellswiftTheirs)) - return errInsufficientBytes - } - + v1Prefix := createV1Prefix(net) if !initiating && bytes.Equal(ellswiftTheirs[4:16], v1Prefix[4:16]) { log.Warnf("Peer sent v1 version message for wrong network "+ - "(expected %v)", btcnet) + "(expected %v)", net) return errWrongNetV1Peer } log.Debug("Calculating ECDH shared secret") - - // Calculate the shared secret to be used in creating the packet - // ciphers. ecdhSecret, err := ellswift.V2Ecdh( p.privkeyOurs, ellswiftTheirs, p.ellswiftOurs, initiating, ) @@ -613,14 +677,28 @@ func (p *Peer) CompleteHandshake(initiating bool, decoyContentLens []int, log.Tracef("Calculated ECDH shared secret: %x", ecdhSecret) - err = p.createV2Ciphers(ecdhSecret[:], initiating, btcnet) - if err != nil { + return p.createV2Ciphers(ecdhSecret[:], initiating, net) +} + +// CompleteHandshake finishes the v2 protocol negotiation and optionally sends +// decoy packets after sending the garbage terminator. +func (p *Peer) CompleteHandshake(initiating bool, decoyContentLens []int, + btcnet BitcoinNet) error { + + log.Debugf("Completing v2 handshake (initiating=%v, "+ + "num_decoys=%d, net=%v)", initiating, len(decoyContentLens), + btcnet) + + if err := p.completeKeyExchange(initiating, btcnet); err != nil { return err } // Send garbage terminator. log.Debugf("Sending garbage terminator: %x", p.sendGarbageTerm) - p.Send(p.sendGarbageTerm[:]) + _, err := p.Send(p.sendGarbageTerm[:]) + if err != nil { + return err + } // Optionally send decoy packets after garbage terminator. aad := p.sentGarbage @@ -630,15 +708,13 @@ func (p *Peer) CompleteHandshake(initiating bool, decoyContentLens []int, decoyContent := make([]byte, decoyContentLens[i]) - encPacket, _, err := p.V2EncPacket(decoyContent, aad, true) + _, _, err := p.V2EncPacket(decoyContent, aad, true) if err != nil { log.Errorf("Failed to encrypt/send decoy "+ "packet %d: %v", i+1, err) return err } - p.Send(encPacket) - // AAD is only used for the first packet after the handshake. aad = nil } @@ -872,7 +948,7 @@ func (p *Peer) V2ReceivePacket(aad []byte) ([]byte, error) { } } -// ReceivedPrefix returns the partial header bytes we've already received. +// ReceivedPrefix returns the transport-classification bytes already received. func (p *Peer) ReceivedPrefix() []byte { return p.receivedPrefix } @@ -894,6 +970,9 @@ func (p *Peer) UseReadWriter(rw io.ReadWriter) { func (p *Peer) Send(data []byte) (int, error) { log.Tracef("Sending %d bytes", len(data)) n, err := p.rw.Write(data) + if err == nil && n != len(data) { + err = io.ErrShortWrite + } if err != nil { log.Errorf("Send failed after %d bytes: %v", n, err) } else { @@ -930,7 +1009,7 @@ func (p *Peer) Receive(numBytes int) ([]byte, int, error) { n, err := p.rw.Read(b[index:]) if err != nil { - log.Errorf("Receive failed after reading %d bytes "+ + log.Debugf("Receive failed after reading %d bytes "+ "(target %d): %v", total+n, numBytes, err) return nil, total, err }