Handle packetSize more carefully (#132)

Also move MTU + max packet size to `internal/protocol/network.go`.
 And add tests.
This commit is contained in:
2tefan 2026-08-27 00:47:40 +02:00 committed by GitHub
parent 561d01fc22
commit a0d0788b59
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
4 changed files with 197 additions and 22 deletions

View File

@ -0,0 +1,6 @@
package protocol
const (
TunnelMTU = 1420
MaxPacketSize = TunnelMTU
)

View File

@ -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")

123
node/stream_test.go Normal file
View File

@ -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)
}
}

View File

@ -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 {