Skip to content

Commit 8fd6d90

Browse files
committed
perf(test): reduce 1GB echo allocations and random overhead
- Reuse buffers and use rng.Read to cut per-iteration allocations in random echo test - Separate RNG for payload length to keep data stream deterministic - Optimize UDP source validation by comparing UDPAddr directly to avoid per-packet String allocation
1 parent b27887e commit 8fd6d90

3 files changed

Lines changed: 76 additions & 22 deletions

File tree

readloop.go

Lines changed: 31 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -23,18 +23,34 @@
2323
package kcp
2424

2525
import (
26+
"net"
2627
"sync/atomic"
2728

2829
"github.com/pkg/errors"
2930
)
3031

32+
func sameUDPAddr(a, b *net.UDPAddr) bool {
33+
if a == nil || b == nil {
34+
return false
35+
}
36+
if a.Port != b.Port || a.Zone != b.Zone {
37+
return false
38+
}
39+
return a.IP.Equal(b.IP)
40+
}
41+
3142
// defaultReadLoop is the standard procedure for reading from a connection
3243
func (s *UDPSession) defaultReadLoop() {
3344
buf := make([]byte, mtuLimit)
3445

35-
var src string
46+
var src *net.UDPAddr
47+
var srcStr string
3648
if s.remote != nil {
37-
src = s.remote.String()
49+
if udp, ok := s.remote.(*net.UDPAddr); ok {
50+
src = udp
51+
} else {
52+
srcStr = s.remote.String()
53+
}
3854
}
3955
for {
4056
n, addr, err := s.conn.ReadFrom(buf)
@@ -48,9 +64,19 @@ func (s *UDPSession) defaultReadLoop() {
4864
}
4965

5066
// make sure the packet is from the same source
51-
if src == "" { // set source address if nil
52-
src = addr.String()
53-
} else if addr.String() != src {
67+
if src == nil && srcStr == "" { // set source address if nil
68+
if udp, ok := addr.(*net.UDPAddr); ok {
69+
src = udp
70+
} else {
71+
srcStr = addr.String()
72+
}
73+
} else if src != nil {
74+
udp, ok := addr.(*net.UDPAddr)
75+
if !ok || !sameUDPAddr(src, udp) {
76+
atomic.AddUint64(&DefaultSnmp.InErrs, 1)
77+
continue
78+
}
79+
} else if addr.String() != srcStr {
5480
atomic.AddUint64(&DefaultSnmp.InErrs, 1)
5581
continue
5682
}

readloop_linux.go

Lines changed: 21 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -25,6 +25,7 @@
2525
package kcp
2626

2727
import (
28+
"net"
2829
"sync/atomic"
2930

3031
"github.com/pkg/errors"
@@ -44,9 +45,14 @@ func (s *UDPSession) readLoop() {
4445
}
4546

4647
// x/net version
47-
var src string
48+
var src *net.UDPAddr
49+
var srcStr string
4850
if s.remote != nil {
49-
src = s.remote.String()
51+
if udp, ok := s.remote.(*net.UDPAddr); ok {
52+
src = udp
53+
} else {
54+
srcStr = s.remote.String()
55+
}
5056
}
5157
msgs := make([]ipv4.Message, batchSize)
5258
for k := range msgs {
@@ -68,9 +74,19 @@ func (s *UDPSession) readLoop() {
6874
msg := &msgs[i]
6975

7076
// make sure the packet is from the same source
71-
if src == "" { // set source address if nil
72-
src = msg.Addr.String()
73-
} else if msg.Addr.String() != src {
77+
if src == nil && srcStr == "" { // set source address if nil
78+
if udp, ok := msg.Addr.(*net.UDPAddr); ok {
79+
src = udp
80+
} else {
81+
srcStr = msg.Addr.String()
82+
}
83+
} else if src != nil {
84+
udp, ok := msg.Addr.(*net.UDPAddr)
85+
if !ok || !sameUDPAddr(src, udp) {
86+
atomic.AddUint64(&DefaultSnmp.InErrs, 1)
87+
continue
88+
}
89+
} else if msg.Addr.String() != srcStr {
7490
atomic.AddUint64(&DefaultSnmp.InErrs, 1)
7591
continue
7692
}

sess_test.go

Lines changed: 24 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -400,18 +400,21 @@ func randomEchoTest(t *testing.T, cli *UDPSession, N int64) {
400400
// Writer goroutine
401401
go func() {
402402
r := mrand.New(writerSrc)
403+
lenRand := mrand.New(mrand.NewSource(seed + 1))
403404
lastPrint := int64(0)
405+
sndbuf := make([]byte, 1<<20)
404406
for bytesSent < N {
405-
length := mrand.Intn(1<<20) + 1 // Random length between 1 and 1MB
407+
length := lenRand.Intn(1<<20) + 1 // Random length between 1 and 1MB
406408
if bytesSent+int64(length) > N {
407409
length = int(N - bytesSent)
408410
}
409-
sndbuf := make([]byte, length)
410-
for i := range sndbuf {
411-
sndbuf[i] = byte(r.Int())
411+
payload := sndbuf[:length]
412+
if _, err := r.Read(payload); err != nil {
413+
t.Errorf("Random fill error: %v", err)
414+
return
412415
}
413416

414-
n, err := cli.Write(sndbuf)
417+
n, err := cli.Write(payload)
415418
if err != nil {
416419
t.Errorf("Write error: %v", err)
417420
return
@@ -426,22 +429,31 @@ func randomEchoTest(t *testing.T, cli *UDPSession, N int64) {
426429

427430
// Reader goroutine
428431
r := mrand.New(readerSrc)
432+
lenRand := mrand.New(mrand.NewSource(seed + 2))
429433
lastPrint := int64(0)
434+
rcvbuf := make([]byte, 1<<20)
435+
expbuf := make([]byte, 1<<20)
430436
for bytesReceived < N {
431-
length := mrand.Intn(1<<20) + 1 // Random length between 1 and 1MB
437+
length := lenRand.Intn(1<<20) + 1 // Random length between 1 and 1MB
432438
if bytesReceived+int64(length) > N {
433439
length = int(N - bytesReceived)
434440
}
435-
rcvbuf := make([]byte, length)
436-
n, err := cli.Read(rcvbuf)
441+
buf := rcvbuf[:length]
442+
n, err := cli.Read(buf)
437443
if err != nil && err != io.EOF {
438444
t.Fatalf("Read error: %v", err)
439445
}
440-
for i := range n {
441-
expectedByte := byte(r.Int())
442-
if rcvbuf[i] != expectedByte {
443-
t.Fatalf("Data mismatch at byte %d: got %v, want %v", bytesReceived+int64(i), rcvbuf[i], expectedByte)
446+
expected := expbuf[:n]
447+
if _, err := r.Read(expected); err != nil {
448+
t.Fatalf("Random fill error: %v", err)
449+
}
450+
if !bytes.Equal(buf[:n], expected) {
451+
for i := 0; i < n; i++ {
452+
if buf[i] != expected[i] {
453+
t.Fatalf("Data mismatch at byte %d: got %v, want %v", bytesReceived+int64(i), buf[i], expected[i])
454+
}
444455
}
456+
t.Fatalf("Data mismatch at byte %d", bytesReceived)
445457
}
446458
bytesReceived += int64(n)
447459
if bytesReceived-lastPrint >= 1<<28 { // print every 256MB

0 commit comments

Comments
 (0)