diff --git a/.github/workflows/go.yml b/.github/workflows/go.yml index 794821a..d9ab68b 100644 --- a/.github/workflows/go.yml +++ b/.github/workflows/go.yml @@ -31,7 +31,7 @@ jobs: # go-package: ./... - name: Test - run: go test -v -coverprofile=coverage.txt -covermode=atomic ./... + run: go test -v -coverprofile=coverage.txt -covermode=atomic -race ./... - name: Upload coverage reports to Codecov uses: codecov/codecov-action@v5 diff --git a/dns/dns.go b/dns/dns.go index cb88297..c3abd1f 100644 --- a/dns/dns.go +++ b/dns/dns.go @@ -7,6 +7,8 @@ import ( "slices" "strconv" "strings" + + "github.com/soypat/lneto/internal" ) // Global parameters. @@ -288,11 +290,27 @@ func (m *Message) AddAdditionals(rsc []Resource) { } } +// LimitResourceDecoding sets the maximum number of resources that can be decoded +// by a subsequent call to [Message.Decode]. This is useful for limiting memory +// usage when decoding untrusted DNS messages. +// +// After calling LimitResourceDecoding, a call to Decode will: +// - Decode at most maxQ questions +// - Decode at most maxAns answers +// - Decode at most maxAuth authority records +// - Decode at most maxAdd additional records +// +// If the message contains more resources than the limits, Decode returns +// incompleteButOK=true along with an error indicating which resource type +// exceeded the limit. The message is still usable with the decoded resources. +// +// Call this method before Decode to set up the limits. The limits are based on +// slice capacity, which is set exactly to the specified values. func (m *Message) LimitResourceDecoding(maxQ, maxAns, maxAuth, maxAdd uint16) { - m.Questions = slices.Grow(m.Questions, int(maxQ)) - m.Answers = slices.Grow(m.Answers, int(maxQ)) - m.Authorities = slices.Grow(m.Authorities, int(maxQ)) - m.Additionals = slices.Grow(m.Additionals, int(maxQ)) + internal.SliceReuse(&m.Questions, int(maxQ)) + internal.SliceReuse(&m.Answers, int(maxAns)) + internal.SliceReuse(&m.Authorities, int(maxAuth)) + internal.SliceReuse(&m.Additionals, int(maxAdd)) } func (m *Message) Reset() { @@ -616,10 +634,10 @@ LOOP: } func (dst *Message) CopyFrom(m Message) { - reuseGrowSlice(&dst.Questions, len(m.Questions)) - reuseGrowSlice(&dst.Answers, len(m.Answers)) - reuseGrowSlice(&dst.Authorities, len(m.Authorities)) - reuseGrowSlice(&dst.Additionals, len(m.Additionals)) + internal.SliceReuse(&dst.Questions, len(m.Questions)) + internal.SliceReuse(&dst.Answers, len(m.Answers)) + internal.SliceReuse(&dst.Authorities, len(m.Authorities)) + internal.SliceReuse(&dst.Additionals, len(m.Additionals)) for i := range dst.Questions { dst.Questions[i].CopyFrom(m.Questions[i]) } @@ -652,10 +670,3 @@ func (dst *ResourceHeader) CopyFrom(rh ResourceHeader) { dst.TTL = rh.TTL dst.Length = rh.Length } - -func reuseGrowSlice[T any](dst *[]T, n int) { - if n == 0 { - return - } - *dst = slices.Grow(*dst, n)[:n] -} diff --git a/dns/dns_test.go b/dns/dns_test.go index 5914922..5380a7c 100644 --- a/dns/dns_test.go +++ b/dns/dns_test.go @@ -178,14 +178,15 @@ func TestMessageAppendEncodeIncompleteOK(t *testing.T) { } var msg Message - msg.LimitResourceDecoding(uint16(len(tt.Message.Questions)), uint16(len(tt.Message.Answers)), uint16(len(tt.Message.Authorities)), uint16(len(tt.Message.Additionals))) + // Limit answers to 1 to test incomplete parsing (message has 2 answers). + msg.LimitResourceDecoding(uint16(len(tt.Message.Questions)), 1, uint16(len(tt.Message.Authorities)), uint16(len(tt.Message.Additionals))) _, incomplete, err := msg.Decode(b) if err != nil && !incomplete { t.Fatal(err) } else if !incomplete { t.Fatal("expected incomplete parse") } - tt.Message.Answers = tt.Message.Answers[:1] // Trim off the last answer that was not parsed. + tt.Message.Answers = tt.Message.Answers[:1] // Trim to match the limited decode. if msg.String() != tt.Message.String() { t.Errorf("mismatch message strings after append/decode:\n%s\n%s", tt.Message.String(), msg.String()) } diff --git a/internal/ip.go b/internal/ip.go index 362d07c..ac86109 100644 --- a/internal/ip.go +++ b/internal/ip.go @@ -56,34 +56,3 @@ func SetIPAddrs(buf []byte, id uint16, src, dst []byte) (err error) { copy(dstaddr, dst) return nil } - -// IsZeroed returns true if all arguments are set to their zero value. -func IsZeroed[T comparable](a ...T) bool { - var z T - for i := range a { - if a[i] != z { - return false - } - } - return true -} - -// DeleteZeroed deletes zero values in-place contained within the -// slice and returns the modified slice without zero values. -// Does not modify capacity. -func DeleteZeroed[T comparable](a []T) []T { - var z T - off := 0 - deleted := false - for i := 0; i < len(a); i++ { - if a[i] != z { - if deleted { - a[off] = a[i] - } - off++ - } else if !deleted { - deleted = true - } - } - return a[:off] -} diff --git a/internal/ring.go b/internal/ring.go index f7ec7cc..5300cd5 100644 --- a/internal/ring.go +++ b/internal/ring.go @@ -131,7 +131,7 @@ func (r *Ring) ReadPeek(b []byte) (int, error) { // Read reads up to len(b) bytes from the ring buffer and advances the read pointer. [io.EOF] returned when no data available. func (r *Ring) Read(b []byte) (int, error) { n, err := r.read(b) - if err != nil { + if err != nil || len(b) == 0 { return n, err } r.onReadEnd(n) @@ -139,7 +139,9 @@ func (r *Ring) Read(b []byte) (int, error) { } func (r *Ring) read(b []byte) (n int, err error) { - if r.IsEmpty() { + if len(b) == 0 { + return 0, nil + } else if r.IsEmpty() { return 0, io.EOF } if r.End > r.Off { diff --git a/internal/slices.go b/internal/slices.go new file mode 100644 index 0000000..f1b11fa --- /dev/null +++ b/internal/slices.go @@ -0,0 +1,48 @@ +package internal + +// IsZeroed returns true if all arguments are set to their zero value. +func IsZeroed[T comparable](a ...T) bool { + var z T + for i := range a { + if a[i] != z { + return false + } + } + return true +} + +// DeleteZeroed deletes zero values in-place contained within the +// slice and returns the modified slice without zero values. +// Does not modify capacity. +func DeleteZeroed[T comparable](a []T) []T { + var z T + off := 0 + deleted := false + for i := 0; i < len(a); i++ { + if a[i] != z { + if deleted { + a[off] = a[i] + } + off++ + } else if !deleted { + deleted = true + } + } + return a[:off] +} + +// SliceReuse prepares a slice for reuse with capacity at least n. +// After calling SliceReuse, the slice will have: +// - length = 0 +// - capacity >= n (exactly n if a new allocation was needed) +// +// This function provides specified behavior unlike [slices.Grow] which +// has unspecified capacity growth behavior that differs between Go and TinyGo. +// Use this when the exact capacity matters for subsequent logic. +func SliceReuse[T any](buf *[]T, n int) { + if cap(*buf) < n { + *buf = make([]T, 0, n) + } else { + *buf = (*buf)[:0] + } +} diff --git a/tcp/conn.go b/tcp/conn.go index e83dea5..661793e 100644 --- a/tcp/conn.go +++ b/tcp/conn.go @@ -40,6 +40,19 @@ type Conn struct { ipID uint16 } +// reset must be called while holding [Conn.mu]. +func (conn *Conn) reset(h Handler) { + // Reset fields individually - DO NOT copy the mutex (undefined behavior in Go). + // "A Mutex must not be copied after first use." - sync package docs. + // Copying a locked mutex causes corruption on multi-core systems. + conn.h = h + conn.remoteAddr = conn.remoteAddr[:0] + conn.rdead = time.Time{} + conn.wdead = time.Time{} + conn.abortErr = nil + conn.ipID = 0 +} + type ConnConfig struct { RxBuf []byte TxBuf []byte @@ -123,7 +136,8 @@ func (conn *Conn) OpenActive(localPort uint16, remote netip.AddrPort, iss Value) if !remote.IsValid() { return errInvalidIP } - err := conn.h.OpenActive(localPort, remote.Port(), iss) + rport := remote.Port() + err := conn.h.OpenActive(localPort, rport, iss) if err != nil { return err } @@ -136,6 +150,7 @@ func (conn *Conn) OpenActive(localPort uint16, remote netip.AddrPort, iss Value) addr6 := raddr.As16() conn.remoteAddr = append(conn.remoteAddr[:0], addr6[:]...) } + conn.debug("conn:dial", slog.Uint64("lport", uint64(localPort)), slog.Uint64("rport", uint64(rport))) return nil } @@ -149,13 +164,14 @@ func (conn *Conn) OpenListen(localPort uint16, iss Value) error { return err } conn.reset(conn.h) + conn.debug("conn:listen", slog.Uint64("lport", uint64(localPort))) return nil } func (conn *Conn) Close() error { conn.mu.Lock() defer conn.mu.Unlock() - conn.trace("TCPConn.Close") + conn.trace("TCPConn.Close", slog.Uint64("lport", uint64(conn.h.localPort))) return conn.h.Close() } @@ -163,14 +179,9 @@ func (conn *Conn) Close() error { func (conn *Conn) Abort() { conn.mu.Lock() defer conn.mu.Unlock() + conn.trace("TCPConn.Abort", slog.Uint64("lport", uint64(conn.h.localPort))) conn.h.Abort() - *conn = Conn{ - mu: conn.mu, - h: conn.h, - remoteAddr: conn.remoteAddr[:0], - logger: conn.logger, - ipID: conn.ipID, - } + conn.reset(conn.h) } // InternalHandler returns the internal [Handler] instance. The Handler contains lower level implementation logic for a TCP connection. @@ -186,8 +197,10 @@ func (conn *Conn) Write(b []byte) (int, error) { if err != nil { return 0, err } + rport := conn.RemotePort() plen := len(b) - conn.trace("TCPConn.Write:start") + lport := conn.LocalPort() + conn.trace("TCPConn.Write:start", slog.Uint64("lport", uint64(lport)), slog.Uint64("rport", uint64(rport))) if conn.deadlineExceeded(&conn.wdead) { return 0, errDeadlineExceeded } else if plen == 0 { @@ -200,11 +213,12 @@ func (conn *Conn) Write(b []byte) (int, error) { return 0, err } conn.mu.Lock() - ngot, _ := conn.h.Write(b) + var ngot int + ngot, err = conn.h.Write(b) conn.mu.Unlock() n += ngot b = b[ngot:] - if n == plen { + if err != nil || n == plen { break } else if ngot > 0 { backoff.Hit() @@ -212,12 +226,12 @@ func (conn *Conn) Write(b []byte) (int, error) { } else { backoff.Miss() } - conn.trace("TCPConn.Write:insuf-buf", slog.Int("missing", plen-n)) + conn.trace("TCPConn.Write:insuf-buf", slog.Int("missing", plen-n), slog.Uint64("lport", uint64(lport)), slog.Uint64("rport", uint64(rport))) if conn.deadlineExceeded(&conn.wdead) { return n, errDeadlineExceeded } } - return n, nil + return n, err } func (conn *Conn) Flush() error { @@ -242,15 +256,22 @@ func (conn *Conn) Flush() error { // Read reads data from the socket's input buffer. If the buffer is empty, // Read will block until data is available or connection closes. +// Returns io.EOF when the remote has closed the connection and all buffered data has been read. func (conn *Conn) Read(b []byte) (int, error) { connid, err := conn.lockPipeConnID() if err != nil { return 0, err } - conn.trace("TCPConn.Read:start") + lport := conn.LocalPort() + rport := conn.RemotePort() + conn.trace("TCPConn.Read:start", slog.Uint64("lport", uint64(lport)), slog.Uint64("rport", uint64(rport))) backoff := internal.NewBackoff(internal.BackoffTCPConn) - for conn.BufferedInput() == 0 && conn.State() == StateEstablished { - if err := conn.checkPipe(connid, &conn.rdead); err != nil { + for conn.BufferedInput() == 0 { + state := conn.State() + if !state.RxDataOpen() { + // No use waiting for data, jump to read and return corresponding error from there. + break + } else if err := conn.checkPipe(connid, &conn.rdead); err != nil { return 0, err } backoff.Miss() @@ -281,7 +302,7 @@ func (conn *Conn) checkPipe(connID uint64, deadline *time.Time) (err error) { } else if !deadline.IsZero() && time.Since(*deadline) > 0 { err = errDeadlineExceeded } - return nil + return err } func (conn *Conn) checkPipeOpen() error { @@ -298,7 +319,6 @@ func (conn *Conn) checkPipeOpen() error { func (conn *Conn) Demux(buf []byte, off int) (err error) { conn.mu.Lock() defer conn.mu.Unlock() - conn.trace("tcpconn.Recv:start") if off >= len(buf) { return errors.New("bad offset in TCPConn.Recv") } @@ -309,6 +329,7 @@ func (conn *Conn) Demux(buf []byte, off int) (err error) { if conn.isRaddrSet() && !bytes.Equal(conn.remoteAddr, raddr) { return errors.New("IP addr mismatch on TCPConn") } + conn.trace("tcpconn.Recv", slog.Uint64("lport", uint64(conn.h.LocalPort())), slog.Uint64("rport", uint64(conn.h.remotePort))) err = conn.h.Recv(buf[off:]) if err != nil { return err @@ -336,6 +357,7 @@ func (conn *Conn) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int) } else if len(raddr) != len(conn.remoteAddr) { return 0, errMismatchedIPVersion } + conn.trace("TCPConn.encaps", slog.Uint64("lport", uint64(conn.h.LocalPort())), slog.Uint64("rport", uint64(conn.h.remotePort))) n, err = conn.h.Send(carrierData[offsetToFrame:]) if err != nil { return 0, err @@ -356,18 +378,6 @@ func (conn *Conn) isRaddrSet() bool { return len(conn.remoteAddr) != 0 } -func (conn *Conn) reset(h Handler) { - if conn.mu.TryLock() { - panic("reset must be called from within locked conn") - } - *conn = Conn{ - h: h, - mu: conn.mu, - remoteAddr: conn.remoteAddr[:0], - logger: conn.logger, - } -} - // SetDeadline sets the read and write deadlines associated // with the connection. It is equivalent to calling both // SetReadDeadline and SetWriteDeadline. Implements [net.Conn]. diff --git a/tcp/handler.go b/tcp/handler.go index f50611e..efa4989 100644 --- a/tcp/handler.go +++ b/tcp/handler.go @@ -304,20 +304,29 @@ func (h *Handler) SizeRx() int { // Write implements [io.Writer] by copying b to a internal buffer to be sent over the network on the next // [Handler.Send] call that can send data to remote peer. Use [Handler.Free] to know the maximum length the argument slice can be before erroring. func (h *Handler) Write(b []byte) (int, error) { + state := h.State() if h.closing { return 0, errConnectionClosing - } else if h.State().IsClosed() { // Reject write call if data cannot be sent. + } else if !state.TxDataOpen() { // Reject write call if data cannot be sent. return 0, net.ErrClosed } return h.bufTx.Write(b) } // Read implements [io.Reader] by reading received data from remote peer in internal buffer. -func (h *Handler) Read(b []byte) (int, error) { - if h.State().IsClosed() { // Reject read call if state is at StateClosed. Note this is less strict than Write call condition. - return 0, net.ErrClosed +func (h *Handler) Read(b []byte) (n int, err error) { + if h.bufRx.Buffered() > 0 { + n, err = h.bufRx.Read(b) } - return h.bufRx.Read(b) + if n == 0 && err == nil { + state := h.State() + if state.IsClosed() { + err = net.ErrClosed + } else if !state.RxDataOpen() { + err = io.EOF + } + } + return n, err } // BufferedInput returns amount of bytes buffered in receive(input) buffer and ready to read diff --git a/tcp/listener.go b/tcp/listener.go index a7f50c7..fd071f5 100644 --- a/tcp/listener.go +++ b/tcp/listener.go @@ -27,6 +27,22 @@ type Listener struct { port uint16 poolGet func() (*Conn, Value) poolReturn func(*Conn) + logger +} + +func (listener *Listener) reset(port uint16, tcppool pool) { + listener.accepted = listener.accepted[:0] + listener.incoming = listener.incoming[:0] + listener.connID++ + listener.port = port + listener.poolGet = tcppool.GetTCP + listener.poolReturn = tcppool.PutTCP +} + +func (listener *Listener) SetLogger(logger *slog.Logger) { + listener.mu.Lock() + defer listener.mu.Unlock() + listener.logger.log = logger } // LocalPort implements [StackNode]. @@ -48,6 +64,7 @@ func (listener *Listener) Close() error { if listener.isClosed() { return errors.New("already closed") } + listener.debug("listener:reset", slog.Uint64("port", uint64(listener.port))) listener.connID++ listener.port = 0 return nil @@ -61,15 +78,8 @@ func (listener *Listener) Reset(port uint16, pool pool) error { } listener.mu.Lock() defer listener.mu.Unlock() - *listener = Listener{ - mu: listener.mu, - connID: listener.connID + 1, - port: port, - poolGet: pool.GetTCP, - poolReturn: pool.PutTCP, - incoming: listener.incoming[:0], - accepted: listener.accepted[:0], - } + listener.debug("listener:reset", slog.Uint64("port", uint64(port))) + listener.reset(port, pool) return nil } @@ -95,6 +105,7 @@ func (listener *Listener) TryAccept() (*Conn, error) { if listener.isClosed() { return nil, net.ErrClosed } + listener.debug("listener:tryaccept", slog.Uint64("port", uint64(listener.port))) listener.maintainConns() for i, conn := range listener.incoming { if conn == nil || conn.State() != StateEstablished { @@ -114,6 +125,7 @@ func (listener *Listener) Encapsulate(carrierData []byte, offsetToIP, offsetToFr if listener.isClosed() { return 0, net.ErrClosed } + //listener.trace("listener:encaps", slog.Uint64("port", uint64(listener.port))) // First try incoming connections (for handshake SYN-ACK). for i, conn := range listener.incoming { if conn == nil || conn.State() == StateEstablished { @@ -127,6 +139,7 @@ func (listener *Listener) Encapsulate(carrierData []byte, offsetToIP, offsetToFr if n == 0 { continue } + listener.debug("listener:encaps", slog.Uint64("port", uint64(listener.port)), slog.Int("plen", n), slog.String("list", "incoming")) return n, err } // Then try accepted connections. @@ -141,6 +154,7 @@ func (listener *Listener) Encapsulate(carrierData []byte, offsetToIP, offsetToFr if n == 0 { continue } + listener.debug("listener:encaps", slog.Uint64("port", uint64(listener.port)), slog.Int("plen", n), slog.String("list", "accepted")) return n, err } return 0, nil @@ -166,15 +180,19 @@ func (listener *Listener) Demux(carrierData []byte, tcpFrameOffset int) error { return errors.New("not our port") } src := tfrm.SourcePort() + // Try to demux in accepted: + accepted := true demuxed, err := listener.tryDemux(listener.accepted, src, srcaddr, carrierData, tcpFrameOffset) + if !demuxed { + accepted = false + demuxed, err = listener.tryDemux(listener.incoming, src, srcaddr, carrierData, tcpFrameOffset) + } if demuxed { + listener.debug("tcplistener:demux", slog.Uint64("lport", uint64(listener.port)), slog.Uint64("rport", uint64(src)), slog.Bool("accepted", accepted)) return err } - demuxed, err = listener.tryDemux(listener.incoming, src, srcaddr, carrierData, tcpFrameOffset) - if demuxed { - return err - } + // Connection not in ready nor accepted. _, flags := tfrm.OffsetAndFlags() if flags != FlagSYN { @@ -198,6 +216,7 @@ func (listener *Listener) Demux(carrierData []byte, tcpFrameOffset int) error { return lneto.ErrPacketDrop } listener.incoming = append(listener.incoming, conn) + listener.debug("tcplistener:demux-new", slog.Uint64("lport", uint64(listener.port)), slog.Uint64("rport", uint64(src))) return nil } @@ -223,7 +242,8 @@ func (listener *Listener) maintainConns() { if listener.incoming[i] == nil { continue } - if listener.incoming[i].State() > StateEstablished || listener.incoming[i].State().IsClosed() { + state := listener.incoming[i].State() + if state > StateEstablished || state.IsClosed() { // Something went wrong in handshake or pool aborted/closed the connection. listener.poolReturn(listener.incoming[i]) listener.incoming[i] = nil diff --git a/tcp/txqueue.go b/tcp/txqueue.go index f3e6904..c1ec4c0 100644 --- a/tcp/txqueue.go +++ b/tcp/txqueue.go @@ -2,7 +2,6 @@ package tcp import ( "errors" - "slices" "github.com/soypat/lneto/internal" ) @@ -63,6 +62,7 @@ func (rtx *ringTx) Reset(buf []byte, maxqueuedPackets int, iss Value) error { *rtx = ringTx{ rawbuf: buf, + slist: rtx.slist, } rtx.slist.Reset(maxqueuedPackets, iss) rtx.iss = iss @@ -260,8 +260,11 @@ type sentlist struct { pkts []ringidx } +// Reset clears the sent packet list and prepares it for reuse. +// The packet queue capacity is set to exactly pktQueueSize. +// The initial sequence number is set to iss. func (sl *sentlist) Reset(pktQueueSize int, iss Value) { - sl.pkts = slices.Grow(sl.pkts[:0], pktQueueSize) + internal.SliceReuse(&sl.pkts, pktQueueSize) sl.ssn = iss } diff --git a/x/xnet/stack-retrying.go b/x/xnet/stack-retrying.go index 226a6ab..e11764f 100644 --- a/x/xnet/stack-retrying.go +++ b/x/xnet/stack-retrying.go @@ -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 { diff --git a/x/xnet/tcppool.go b/x/xnet/tcppool.go index c3a9081..279a12d 100644 --- a/x/xnet/tcppool.go +++ b/x/xnet/tcppool.go @@ -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...) + } } diff --git a/x/xnet/xnet_concurrent_test.go b/x/xnet/xnet_concurrent_test.go new file mode 100644 index 0000000..0584dc2 --- /dev/null +++ b/x/xnet/xnet_concurrent_test.go @@ -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 +} diff --git a/x/xnet/xnet_test.go b/x/xnet/xnet_test.go index aaa99ea..8d46c43 100644 --- a/x/xnet/xnet_test.go +++ b/x/xnet/xnet_test.go @@ -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