mirror of
https://github.com/hyprspace/hyprspace.git
synced 2026-09-12 19:51:07 +05:00
Handle packetSize more carefully (#132)
Also move MTU + max packet size to `internal/protocol/network.go`. And add tests.
This commit is contained in:
parent
561d01fc22
commit
a0d0788b59
6
internal/protocol/network.go
Normal file
6
internal/protocol/network.go
Normal file
@ -0,0 +1,6 @@
|
||||
package protocol
|
||||
|
||||
const (
|
||||
TunnelMTU = 1420
|
||||
MaxPacketSize = TunnelMTU
|
||||
)
|
||||
85
node/node.go
85
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")
|
||||
|
||||
123
node/stream_test.go
Normal file
123
node/stream_test.go
Normal 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)
|
||||
}
|
||||
}
|
||||
@ -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 {
|
||||
|
||||
Loading…
Reference in New Issue
Block a user