diff --git a/internal/protocol/network.go b/internal/protocol/network.go new file mode 100644 index 0000000..b6e517d --- /dev/null +++ b/internal/protocol/network.go @@ -0,0 +1,6 @@ +package protocol + +const ( + TunnelMTU = 1420 + MaxPacketSize = TunnelMTU +) diff --git a/node/node.go b/node/node.go index 852594c..804cdfc 100644 --- a/node/node.go +++ b/node/node.go @@ -5,6 +5,7 @@ import ( "encoding/binary" "errors" "fmt" + "io" "io/fs" "net" "net/http" @@ -15,6 +16,7 @@ import ( "github.com/hyprspace/hyprspace/config" hsdns "github.com/hyprspace/hyprspace/dns" + "github.com/hyprspace/hyprspace/internal/protocol" "github.com/hyprspace/hyprspace/p2p" hsrpc "github.com/hyprspace/hyprspace/rpc" "github.com/hyprspace/hyprspace/svc" @@ -54,6 +56,8 @@ type Node struct { wg *sync.WaitGroup } +var errInvalidPacketSize = errors.New("invalid packet size") + func New(ctx context.Context, configPath string, ifName string) Node { innerCtx, ctxCancel := context.WithCancel(ctx) @@ -86,7 +90,7 @@ func (node *Node) Run() error { node.cfg.Interface, tun.Address(node.cfg.BuiltinAddr4.String()+"/32"), tun.Address(node.cfg.BuiltinAddr6.String()+"/128"), - tun.MTU(1420), + tun.MTU(protocol.TunnelMTU), ) if err != nil { logger.With(err).Error("Failed to create TUN Device") @@ -265,7 +269,7 @@ func (node *Node) Run() error { node.activeStreams = make(map[peer.ID]SharedStream) go func() { for { - var packet = make([]byte, 1420) + var packet = make([]byte, protocol.MaxPacketSize) // Read in a packet from the tun device. plen, err := node.tunDev.Iface.Read(packet) if errors.Is(err, fs.ErrClosed) { @@ -343,8 +347,44 @@ func (node *Node) expireActiveStream(pid peer.ID) { delete(node.activeStreams, pid) } +func readStreamPacket(stream io.Reader, packet []byte) (uint16, error) { + var sizeBytes [2]byte + + // Read the incoming packet's size as a binary value. + if _, err := io.ReadFull(stream, sizeBytes[:]); err != nil { + return 0, fmt.Errorf("read packet size: %w", err) + } + + // Decode the incoming packet's size from binary. + size := binary.LittleEndian.Uint16(sizeBytes[:]) + + // Check that size makes sense. + if size == 0 { + return 0, fmt.Errorf("%w: size is zero", errInvalidPacketSize) + } + + if int(size) > len(packet) { + return 0, fmt.Errorf( + "%w: size %d exceeds maximum %d", + errInvalidPacketSize, + size, + len(packet), + ) + } + + // Read in the packet until completion. + if _, err := io.ReadFull(stream, packet[:int(size)]); err != nil { + return 0, fmt.Errorf("read packet body: %w", err) + } + + return size, nil +} + func (node *Node) streamHandler(stream network.Stream) { remotePeerID := stream.Conn().RemotePeer() + log := logger.With( + zap.String("peer", remotePeerID.String()), + ) // If the remote node ID isn't in the list of known nodes don't respond. if _, ok := config.FindPeer(node.cfg.Peers, remotePeerID); !ok { @@ -365,29 +405,34 @@ func (node *Node) streamHandler(stream network.Stream) { } } - var packet = make([]byte, 1420) - var packetSize = make([]byte, 2) + var packet = make([]byte, protocol.MaxPacketSize) for { - // Read the incoming packet's size as a binary value. - _, err := stream.Read(packetSize) + size, err := readStreamPacket(stream, packet) if err != nil { - stream.Close() + errLog := log.With(zap.Error(err)) + + switch { + // Packet size outside of expected range. + case errors.Is(err, errInvalidPacketSize): + errLog.Warn("received invalid packet size") + _ = stream.Reset() + + // Normal shutdown or peer disconnect. + case errors.Is(err, io.EOF), + errors.Is(err, net.ErrClosed), + errors.Is(err, context.Canceled), + errors.Is(err, network.ErrReset): + errLog.Debug("stream read ended") + _ = stream.Close() + + default: + errLog.Warn("failed to read packet from stream") + _ = stream.Close() + } + return } - // Decode the incoming packet's size from binary. - size := binary.LittleEndian.Uint16(packetSize) - - // Read in the packet until completion. - var plen uint16 = 0 - for plen < size { - tmp, err := stream.Read(packet[plen:size]) - plen += uint16(tmp) - if err != nil { - stream.Close() - return - } - } err = stream.SetWriteDeadline(time.Now().Add(25 * time.Second)) if err != nil { logger.With(err).Error("Failed to set write deadline") diff --git a/node/stream_test.go b/node/stream_test.go new file mode 100644 index 0000000..43b2905 --- /dev/null +++ b/node/stream_test.go @@ -0,0 +1,123 @@ +package node + +import ( + "bytes" + "encoding/binary" + "errors" + "io" + "testing" + + "github.com/hyprspace/hyprspace/internal/protocol" +) + +func TestReadStreamPacketReadsValidPacket(t *testing.T) { + var packet [protocol.MaxPacketSize]byte + payload := []byte{0xde, 0xad, 0xbe, 0xef} + var input bytes.Buffer + + err := binary.Write(&input, binary.LittleEndian, uint16(len(payload))) + if err != nil { + t.Fatalf("failed to write packet size: %v", err) + } + input.Write(payload) + + size, err := readStreamPacket(&input, packet[:]) + if err != nil { + t.Fatalf("expected valid packet, got error: %v", err) + } + if int(size) != len(payload) { + t.Fatalf("expected size %d, got %d", len(payload), size) + } + if !bytes.Equal(packet[:size], payload) { + t.Fatalf("expected payload %x, got %x", payload, packet[:size]) + } +} + +func TestReadStreamPacketRejectsOversizedPacket(t *testing.T) { + var packet [protocol.MaxPacketSize]byte + var input bytes.Buffer + + err := binary.Write(&input, binary.LittleEndian, uint16(len(packet)+1)) + if err != nil { + t.Fatalf("failed to write packet size: %v", err) + } + + _, err = readStreamPacket(&input, packet[:]) + if !errors.Is(err, errInvalidPacketSize) { + t.Fatalf("expected errInvalidPacketSize, got %v", err) + } +} + +func TestReadStreamPacketRequiresFullSizeHeader(t *testing.T) { + var packet [protocol.MaxPacketSize]byte + input := bytes.NewBuffer([]byte{0x01}) + + _, err := readStreamPacket(input, packet[:]) + if !errors.Is(err, io.ErrUnexpectedEOF) { + t.Fatalf("expected io.ErrUnexpectedEOF, got %v", err) + } +} + +func TestReadStreamPacketRejectsZeroSize(t *testing.T) { + packet := make([]byte, protocol.MaxPacketSize) + input := bytes.NewReader([]byte{0x00, 0x00}) + + _, err := readStreamPacket(input, packet) + if !errors.Is(err, errInvalidPacketSize) { + t.Fatalf("expected errInvalidPacketSize, got %v", err) + } +} + +func TestReadStreamPacketAcceptsMaximumSize(t *testing.T) { + packet := make([]byte, protocol.MaxPacketSize) + payload := make([]byte, protocol.MaxPacketSize) + + var input bytes.Buffer + if err := binary.Write( + &input, + binary.LittleEndian, + uint16(len(payload)), + ); err != nil { + t.Fatalf("write packet size: %v", err) + } + + input.Write(payload) + + size, err := readStreamPacket(&input, packet) + if err != nil { + t.Fatalf("read maximum-size packet: %v", err) + } + + if int(size) != len(payload) { + t.Fatalf("expected size %d, got %d", len(payload), size) + } +} + +func TestReadStreamPacketRequiresFullPayload(t *testing.T) { + packet := make([]byte, protocol.MaxPacketSize) + + var input bytes.Buffer + if err := binary.Write( + &input, + binary.LittleEndian, + uint16(4), + ); err != nil { + t.Fatalf("write packet size: %v", err) + } + + input.Write([]byte{0xde, 0xad}) + + _, err := readStreamPacket(&input, packet) + if !errors.Is(err, io.ErrUnexpectedEOF) { + t.Fatalf("expected io.ErrUnexpectedEOF, got %v", err) + } +} + +func TestReadStreamPacketReturnsEOFForEmptyStream(t *testing.T) { + packet := make([]byte, protocol.MaxPacketSize) + + _, err := readStreamPacket(bytes.NewReader(nil), packet) + if !errors.Is(err, io.EOF) { + t.Fatalf("expected io.EOF, got %v", err) + } +} diff --git a/svc/network.go b/svc/network.go index 5fbee1a..ff0d7ce 100644 --- a/svc/network.go +++ b/svc/network.go @@ -6,6 +6,7 @@ import ( "net/netip" "github.com/hyprspace/hyprspace/config" + "github.com/hyprspace/hyprspace/internal/protocol" "github.com/hyprspace/hyprspace/netstack" hstun "github.com/hyprspace/hyprspace/tun" "github.com/ipfs/go-log/v2" @@ -93,7 +94,7 @@ func NewServiceNetwork(host host.Host, cfg *config.Config, tunDev *hstun.TUN) Se netip.AddrFrom16([16]byte([]byte("\xfd\x00hyprspinternal"))), }, []netip.Addr{}, - 1420, + protocol.TunnelMTU, ) if err != nil { logger.With(err).Fatal("Failed to Create service-network tunnel device") @@ -101,7 +102,7 @@ func NewServiceNetwork(host host.Host, cfg *config.Config, tunDev *hstun.TUN) Se go func() { sizes := make([]int, 1) - buffer := make([]byte, 1420) + buffer := make([]byte, protocol.MaxPacketSize) buffers := make([][]byte, 1) buffers[0] = buffer for {