mirror of
https://github.com/soypat/lneto.git
synced 2026-08-20 14:39:02 +00:00
do not rely on unspecified Go behaviour; fix TCP ring buffer bug (#16)
* do not rely on unspecified Go behaviour; fix TCP ring buffer bug * even more stricter error handling on tcp connections * stricter lock copying in tcp.Conn, even though likely not a problem * add listener and tcp pool logging * forgot reqAddr unused * add logging to TCPConn and friends
This commit is contained in:
@@ -43,7 +43,7 @@ func (s StackRetrying) DoNTP(ntpHost netip.Addr, timeout time.Duration, retries
|
||||
expectEnd := time.Now().Add(timeout * time.Duration(retries))
|
||||
for i := 0; i < retries; i++ {
|
||||
if i > 0 {
|
||||
println("Retrying DHCP")
|
||||
println("Retrying NTP")
|
||||
}
|
||||
offset, err = s.block.DoNTP(ntpHost, timeout)
|
||||
if err == nil {
|
||||
|
||||
+33
-8
@@ -1,6 +1,7 @@
|
||||
package xnet
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"log/slog"
|
||||
"sync"
|
||||
@@ -21,6 +22,7 @@ type TCPPool struct {
|
||||
_now func() time.Time
|
||||
estbTimeout time.Duration
|
||||
closingTimeout time.Duration
|
||||
logger *slog.Logger
|
||||
}
|
||||
|
||||
func _() {
|
||||
@@ -32,6 +34,7 @@ type TCPPoolConfig struct {
|
||||
PoolSize int
|
||||
QueueSize int
|
||||
BufferSize int
|
||||
Logger *slog.Logger
|
||||
ConnLogger *slog.Logger
|
||||
Now func() time.Time
|
||||
// EstablishedTimeout sets the timeout for a TCP connection since it is acquired until it is established.
|
||||
@@ -56,6 +59,7 @@ func NewTCPPool(cfg TCPPoolConfig) (*TCPPool, error) {
|
||||
_now: cfg.Now,
|
||||
estbTimeout: cfg.EstablishedTimeout,
|
||||
closingTimeout: cfg.ClosingTimeout,
|
||||
logger: cfg.Logger,
|
||||
}
|
||||
bufSpace := make([]byte, 2*n*bufsize)
|
||||
for i := range pool.conns {
|
||||
@@ -73,9 +77,16 @@ func NewTCPPool(cfg TCPPoolConfig) (*TCPPool, error) {
|
||||
return pool, nil
|
||||
}
|
||||
|
||||
func (p *TCPPool) NumberOfAcquired() int {
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
return p.naqcuired
|
||||
}
|
||||
|
||||
func (p *TCPPool) GetTCP() (*tcp.Conn, tcp.Value) {
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
p.debug("TCPPool:get")
|
||||
for i := range p.conns {
|
||||
if p.acquiredAt[i].IsZero() {
|
||||
p.acquiredAt[i] = p.now()
|
||||
@@ -88,15 +99,18 @@ func (p *TCPPool) GetTCP() (*tcp.Conn, tcp.Value) {
|
||||
}
|
||||
|
||||
func (p *TCPPool) PutTCP(conn *tcp.Conn) {
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
p.debug("TCPPool:put", slog.Uint64("lport", uint64(conn.LocalPort())))
|
||||
for i := range p.conns {
|
||||
if &p.conns[i] == conn {
|
||||
p.mu.Lock()
|
||||
// p.mu.Lock()
|
||||
p.conns[i].Abort()
|
||||
p.acquiredAt[i] = time.Time{}
|
||||
p.abortedAt[i] = time.Time{}
|
||||
p.closingAt[i] = time.Time{}
|
||||
p.naqcuired--
|
||||
p.mu.Unlock()
|
||||
// p.mu.Unlock()
|
||||
return
|
||||
}
|
||||
}
|
||||
@@ -104,14 +118,17 @@ func (p *TCPPool) PutTCP(conn *tcp.Conn) {
|
||||
}
|
||||
|
||||
func (p *TCPPool) CheckTimeouts() {
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
p.debug("TCPPool:checktimeouts", slog.Int("acq", p.naqcuired))
|
||||
for i := range p.conns {
|
||||
st := p.conns[i].State()
|
||||
if st == tcp.StateEstablished {
|
||||
continue
|
||||
}
|
||||
p.mu.Lock()
|
||||
// p.mu.Lock()
|
||||
acq := p.acquiredAt[i]
|
||||
p.mu.Unlock()
|
||||
// p.mu.Unlock()
|
||||
if acq.IsZero() {
|
||||
continue
|
||||
} else if st.IsPreestablished() && p.since(acq) > p.estbTimeout {
|
||||
@@ -119,7 +136,7 @@ func (p *TCPPool) CheckTimeouts() {
|
||||
// This is part of a syn-flood defense mechanism.
|
||||
p.conns[i].Close()
|
||||
} else if st.IsClosed() || st.IsClosing() {
|
||||
p.mu.Lock()
|
||||
// p.mu.Lock()
|
||||
if p.closingAt[i].IsZero() {
|
||||
p.closingAt[i] = p.now()
|
||||
} else if p.abortedAt[i].IsZero() && p.since(p.closingAt[i]) > p.closingTimeout {
|
||||
@@ -128,7 +145,7 @@ func (p *TCPPool) CheckTimeouts() {
|
||||
} else if p.since(p.abortedAt[i]) > 10*time.Second {
|
||||
println("connection aborted and still not returned to TCPPool")
|
||||
}
|
||||
p.mu.Unlock()
|
||||
// p.mu.Unlock()
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -147,6 +164,14 @@ func (p *TCPPool) now() time.Time {
|
||||
return p._now()
|
||||
}
|
||||
|
||||
func (p *TCPPool) NumberOfAcquired() int {
|
||||
return p.naqcuired
|
||||
func (p *TCPPool) trace(msg string, attrs ...slog.Attr) {
|
||||
p.log(slog.LevelDebug-2, msg, attrs...)
|
||||
}
|
||||
func (p *TCPPool) debug(msg string, attrs ...slog.Attr) {
|
||||
p.log(slog.LevelDebug, msg, attrs...)
|
||||
}
|
||||
func (p *TCPPool) log(lvl slog.Level, msg string, attrs ...slog.Attr) {
|
||||
if p.logger != nil {
|
||||
p.logger.LogAttrs(context.Background(), lvl, msg, attrs...)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,282 @@
|
||||
package xnet
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"fmt"
|
||||
"math/rand"
|
||||
"net/netip"
|
||||
"runtime"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/soypat/lneto/tcp"
|
||||
)
|
||||
|
||||
func TestTCPListener_ConcurrentEcho(t *testing.T) {
|
||||
const (
|
||||
numClients = 10
|
||||
serverPort = 8080
|
||||
MTU = 1500
|
||||
seed = 1
|
||||
)
|
||||
|
||||
// 1. Setup server stack with tcp.Listener.
|
||||
var serverStack StackAsync
|
||||
serverMAC := [6]byte{0xaa, 0xbb, 0xcc, 0x00, 0x00, 0x01}
|
||||
serverIP := netip.AddrFrom4([4]byte{10, 0, 0, 1})
|
||||
err := serverStack.Reset(StackConfig{
|
||||
Hostname: "Server",
|
||||
RandSeed: seed,
|
||||
StaticAddress: serverIP,
|
||||
MaxTCPConns: numClients,
|
||||
HardwareAddress: serverMAC,
|
||||
MTU: MTU,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
tcpPool, err := NewTCPPool(TCPPoolConfig{
|
||||
PoolSize: numClients,
|
||||
QueueSize: 4,
|
||||
BufferSize: 512,
|
||||
EstablishedTimeout: 5 * time.Second,
|
||||
ClosingTimeout: 5 * time.Second,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
var listener tcp.Listener
|
||||
err = listener.Reset(serverPort, tcpPool)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
err = serverStack.RegisterListener(&listener)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
// 2. Setup client stacks (one per client).
|
||||
clientStacks := make([]StackAsync, numClients)
|
||||
clientConns := make([]tcp.Conn, numClients)
|
||||
connBufs := make([]byte, numClients*MTU*2) // RX+TX buffer space for all clients
|
||||
|
||||
for i := range clientStacks {
|
||||
clientMAC := [6]byte{0xaa, 0xbb, 0xcc, 0x00, 0x01, byte(i + 1)}
|
||||
clientIP := netip.AddrFrom4([4]byte{10, 0, 0, byte(i + 10)})
|
||||
err := clientStacks[i].Reset(StackConfig{
|
||||
Hostname: fmt.Sprintf("Client%d", i),
|
||||
RandSeed: int64(seed + i + 1),
|
||||
StaticAddress: clientIP,
|
||||
MaxTCPConns: 1,
|
||||
HardwareAddress: clientMAC,
|
||||
MTU: MTU,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("client %d reset: %v", i, err)
|
||||
}
|
||||
// Client gateway points to server.
|
||||
clientStacks[i].SetGateway6(serverMAC)
|
||||
|
||||
// Configure client connection buffers.
|
||||
bufOff := i * MTU * 2
|
||||
err = clientConns[i].Configure(tcp.ConnConfig{
|
||||
RxBuf: connBufs[bufOff : bufOff+MTU],
|
||||
TxBuf: connBufs[bufOff+MTU : bufOff+2*MTU],
|
||||
TxPacketQueueSize: 4,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("client %d conn configure: %v", i, err)
|
||||
}
|
||||
}
|
||||
|
||||
// 3. Start "kernel" goroutine - routes packets between stacks.
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||
defer cancel()
|
||||
|
||||
go kernelLoop(ctx, &serverStack, clientStacks)
|
||||
|
||||
// 4. Start server goroutine - accepts and echoes.
|
||||
go echoServer(ctx, &listener)
|
||||
|
||||
// 5. Start client goroutines.
|
||||
var wg sync.WaitGroup
|
||||
clientSuccess := make([]bool, numClients)
|
||||
for i := range numClients {
|
||||
wg.Add(1)
|
||||
go func(clientID int) {
|
||||
defer wg.Done()
|
||||
if runClient(t, clientID, &clientStacks[clientID], &clientConns[clientID],
|
||||
serverIP, serverPort) {
|
||||
clientSuccess[clientID] = true
|
||||
}
|
||||
}(i)
|
||||
}
|
||||
|
||||
// 6. Wait for all clients to complete.
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
wg.Wait()
|
||||
close(done)
|
||||
}()
|
||||
|
||||
select {
|
||||
case <-done:
|
||||
// Check all clients succeeded.
|
||||
for i, ok := range clientSuccess {
|
||||
if !ok {
|
||||
t.Errorf("client %d did not complete successfully", i)
|
||||
}
|
||||
}
|
||||
case <-ctx.Done():
|
||||
t.Fatal("test timed out")
|
||||
}
|
||||
cancel()
|
||||
}
|
||||
|
||||
func kernelLoop(ctx context.Context, server *StackAsync, clients []StackAsync) {
|
||||
const MTU = 1500
|
||||
buf := make([]byte, MTU)
|
||||
rng := rand.New(rand.NewSource(1)) // Seed 1 for deterministic but randomized order
|
||||
order := make([]int, len(clients))
|
||||
for i := range order {
|
||||
order[i] = i
|
||||
}
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
default:
|
||||
}
|
||||
|
||||
// Process server outgoing -> route to appropriate client based on dest IP.
|
||||
if n, _ := server.Encapsulate(buf, -1, 0); n > 0 {
|
||||
routePacketToClient(buf[:n], clients)
|
||||
}
|
||||
|
||||
// Process each client outgoing in randomized order.
|
||||
rng.Shuffle(len(order), func(i, j int) { order[i], order[j] = order[j], order[i] })
|
||||
for _, idx := range order {
|
||||
if n, _ := clients[idx].Encapsulate(buf, -1, 0); n > 0 {
|
||||
server.Demux(buf[:n], 0) // All clients talk to server.
|
||||
}
|
||||
}
|
||||
|
||||
runtime.Gosched() // Yield to other goroutines.
|
||||
}
|
||||
}
|
||||
|
||||
func routePacketToClient(pkt []byte, clients []StackAsync) {
|
||||
// Extract destination IP from IPv4 header (offset 16-19 in IP header, after 14 byte Ethernet header).
|
||||
if len(pkt) < 34 { // 14 ethernet + 20 min IP header
|
||||
return
|
||||
}
|
||||
dstIP := netip.AddrFrom4([4]byte{pkt[30], pkt[31], pkt[32], pkt[33]})
|
||||
|
||||
for i := range clients {
|
||||
if clients[i].Addr() == dstIP {
|
||||
clients[i].Demux(pkt, 0)
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func echoServer(ctx context.Context, listener *tcp.Listener) {
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
default:
|
||||
}
|
||||
|
||||
if listener.NumberOfReadyToAccept() == 0 {
|
||||
time.Sleep(time.Millisecond)
|
||||
continue
|
||||
}
|
||||
|
||||
conn, err := listener.TryAccept()
|
||||
if err != nil || conn == nil {
|
||||
continue
|
||||
}
|
||||
|
||||
// Handle connection in separate goroutine (like real example).
|
||||
go func(c *tcp.Conn) {
|
||||
var buf [512]byte
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
default:
|
||||
}
|
||||
|
||||
n, err := c.Read(buf[:])
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
if n > 0 {
|
||||
_, err = c.Write(buf[:n])
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
}(conn)
|
||||
}
|
||||
}
|
||||
|
||||
func runClient(t *testing.T, id int, stack *StackAsync, conn *tcp.Conn,
|
||||
serverAddr netip.Addr, serverPort uint16) bool {
|
||||
// Dial server.
|
||||
clientPort := uint16(10000 + id)
|
||||
err := stack.DialTCP(conn, clientPort, netip.AddrPortFrom(serverAddr, serverPort))
|
||||
if err != nil {
|
||||
t.Errorf("client %d dial failed: %v", id, err)
|
||||
return false
|
||||
}
|
||||
|
||||
// Wait for connection established (handshake via kernel loop).
|
||||
deadline := time.Now().Add(5 * time.Second)
|
||||
for conn.State() != tcp.StateEstablished {
|
||||
if time.Now().After(deadline) {
|
||||
t.Errorf("client %d: timeout waiting for established state, got %s", id, conn.State())
|
||||
return false
|
||||
}
|
||||
time.Sleep(time.Millisecond)
|
||||
}
|
||||
|
||||
// Send test data.
|
||||
testData := []byte(fmt.Sprintf("hello from client %d", id))
|
||||
_, err = conn.Write(testData)
|
||||
if err != nil {
|
||||
t.Errorf("client %d write failed: %v", id, err)
|
||||
return false
|
||||
}
|
||||
|
||||
// Read echo response.
|
||||
var buf [64]byte
|
||||
deadline = time.Now().Add(5 * time.Second)
|
||||
var totalRead int
|
||||
for totalRead < len(testData) {
|
||||
if time.Now().After(deadline) {
|
||||
t.Errorf("client %d: timeout waiting for echo response, got %d/%d bytes", id, totalRead, len(testData))
|
||||
return false
|
||||
}
|
||||
n, err := conn.Read(buf[totalRead:])
|
||||
if err != nil {
|
||||
t.Errorf("client %d read failed: %v", id, err)
|
||||
return false
|
||||
}
|
||||
totalRead += n
|
||||
}
|
||||
|
||||
// Verify echo.
|
||||
if !bytes.Equal(buf[:totalRead], testData) {
|
||||
t.Errorf("client %d: expected %q, got %q", id, testData, buf[:totalRead])
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
@@ -7,6 +7,7 @@ import (
|
||||
"net/netip"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/soypat/lneto/arp"
|
||||
"github.com/soypat/lneto/ethernet"
|
||||
@@ -22,6 +23,83 @@ const (
|
||||
finack = tcp.FlagFIN | tcp.FlagACK
|
||||
)
|
||||
|
||||
func TestTCPConn_ReadBlocksUntilDataAvailable(t *testing.T) {
|
||||
const seed = 5678
|
||||
const MTU = 1500
|
||||
const svPort = 8080
|
||||
client, sv, clconn, svconn := newTCPStacks(t, seed, MTU)
|
||||
tst := testerFrom(t, MTU)
|
||||
|
||||
tst.TestTCPSetupAndEstablish(sv, client, svconn, clconn, svPort, 1337)
|
||||
|
||||
// Verify no data buffered initially.
|
||||
if svconn.BufferedInput() != 0 {
|
||||
t.Fatal("expected no buffered input on server conn")
|
||||
}
|
||||
|
||||
sendData := []byte("blocking test data")
|
||||
readDone := make(chan struct{})
|
||||
var readN int
|
||||
var readErr error
|
||||
var readBuf [64]byte
|
||||
|
||||
// Start a goroutine to read from svconn - this should block since no data available.
|
||||
go func() {
|
||||
readN, readErr = svconn.Read(readBuf[:])
|
||||
close(readDone)
|
||||
}()
|
||||
|
||||
// Give Read time to enter blocking state.
|
||||
select {
|
||||
case <-readDone:
|
||||
t.Fatal("Read returned immediately without data - expected blocking")
|
||||
case <-time.After(50 * time.Millisecond):
|
||||
// Good - Read is blocking as expected.
|
||||
}
|
||||
|
||||
// Write data on client side.
|
||||
_, err := clconn.Write(sendData)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
// Perform packet exchange to deliver data.
|
||||
tst.bufmu.Lock()
|
||||
buf := tst.buf[:cap(tst.buf)]
|
||||
n, err := client.Encapsulate(buf, -1, 0)
|
||||
if err != nil {
|
||||
tst.bufmu.Unlock()
|
||||
t.Fatal(err)
|
||||
}
|
||||
if n == 0 {
|
||||
tst.bufmu.Unlock()
|
||||
t.Fatal("expected data packet from client")
|
||||
}
|
||||
err = sv.Demux(buf[:n], 0)
|
||||
tst.bufmu.Unlock()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
// Now Read should unblock and return data.
|
||||
select {
|
||||
case <-readDone:
|
||||
// Good - Read unblocked.
|
||||
case <-time.After(500 * time.Millisecond):
|
||||
t.Fatal("Read did not unblock after data became available")
|
||||
}
|
||||
|
||||
if readErr != nil {
|
||||
t.Fatalf("Read returned error: %v", readErr)
|
||||
}
|
||||
if readN != len(sendData) {
|
||||
t.Fatalf("expected to read %d bytes, got %d", len(sendData), readN)
|
||||
}
|
||||
if !bytes.Equal(readBuf[:readN], sendData) {
|
||||
t.Fatalf("read data mismatch: got %q, want %q", readBuf[:readN], sendData)
|
||||
}
|
||||
}
|
||||
|
||||
func TestStackAsyncTCP_multipacket(t *testing.T) {
|
||||
const seed = 1234
|
||||
const MTU = 512
|
||||
|
||||
Reference in New Issue
Block a user