listener refactor: tcpPool+ready/accepted slices

This commit is contained in:
soypat
2025-06-08 23:10:43 -03:00
parent 9a60642283
commit 76be78812f
+100 -60
View File
@@ -14,34 +14,31 @@ import (
var _ StackNode = (*NodeTCPListener)(nil) var _ StackNode = (*NodeTCPListener)(nil)
type NodeTCPListener struct { type tcpPool interface {
connID uint64 GetTCP() (*tcp.Conn, tcp.Value)
conns []tcp.Conn PutTCP(*tcp.Conn)
accepted []bool
port uint16
getISS func() uint32
} }
func (listener *NodeTCPListener) AcceptRaw() (*tcp.Conn, error) { type NodeTCPListener struct {
connid := listener.connID connID uint64
for { // ready have received a
if listener.isClosed() || connid != listener.connID { ready []*tcp.Conn
return nil, net.ErrClosed accepted []*tcp.Conn
}
for i := range listener.conns { port uint16
isAvailable := listener.connReceivedSyn(i) && !listener.connAccepted(i) poolGet func() (*tcp.Conn, tcp.Value)
if !isAvailable { poolReturn func(*tcp.Conn)
continue
}
// Connection received as SYN and is not yet accepted.
listener.accepted[i] = true
return &listener.conns[i], nil
}
time.Sleep(5 * time.Millisecond)
}
panic("unreachable")
} }
// LocalPort implements [StackNode].
func (listener *NodeTCPListener) LocalPort() uint16 { return listener.port }
// ConnectionID implements [StackNode].
func (listener *NodeTCPListener) ConnectionID() *uint64 { return &listener.connID }
// Protocol implements [StackNode].
func (listener *NodeTCPListener) Protocol() uint64 { return uint64(lneto.IPProtoTCP) }
func (listener *NodeTCPListener) Close() error { func (listener *NodeTCPListener) Close() error {
if listener.isClosed() { if listener.isClosed() {
return errors.New("already closed") return errors.New("already closed")
@@ -51,24 +48,56 @@ func (listener *NodeTCPListener) Close() error {
return nil return nil
} }
func (listener *NodeTCPListener) LocalPort() uint16 { return listener.port } func (listener *NodeTCPListener) Reset(port uint16, pool tcpPool) error {
if port == 0 {
return errZeroPort
} else if pool == nil {
return errors.New("nil TCP pool")
}
*listener = NodeTCPListener{
connID: listener.connID + 1,
port: port,
poolGet: pool.GetTCP,
poolReturn: pool.PutTCP,
ready: listener.ready[:0],
accepted: listener.accepted[:0],
}
return nil
}
func (listener *NodeTCPListener) ConnectionID() *uint64 { return &listener.connID } func (listener *NodeTCPListener) AcceptRaw() (*tcp.Conn, error) {
connid := listener.connID
func (listener *NodeTCPListener) Protocol() uint64 { return uint64(lneto.IPProtoTCP) } for {
if listener.isClosed() || connid != listener.connID {
return nil, net.ErrClosed
}
for i, conn := range listener.ready {
if conn == nil {
continue
}
listener.accepted = append(listener.accepted, conn)
listener.ready[i] = nil // discard from ready.
return conn, nil
}
listener.maintainConns()
time.Sleep(5 * time.Millisecond)
}
panic("unreachable")
}
// Encapsulate implements [StackNode].
func (listener *NodeTCPListener) Encapsulate(carrierData []byte, tcpFrameOffset int) (int, error) { func (listener *NodeTCPListener) Encapsulate(carrierData []byte, tcpFrameOffset int) (int, error) {
if listener.isClosed() { if listener.isClosed() {
return 0, net.ErrClosed return 0, net.ErrClosed
} }
for i := range listener.conns { for i, conn := range listener.accepted {
conn := &listener.conns[i] if conn == nil {
if conn.State().IsClosed() {
continue continue
} }
n, err := conn.Encapsulate(carrierData, tcpFrameOffset) n, err := conn.Encapsulate(carrierData, tcpFrameOffset)
if err != nil { if err != nil {
listener.maintainConn(i, err) listener.maintainAccepted(i, err)
} }
if n == 0 { if n == 0 {
continue continue
@@ -78,6 +107,7 @@ func (listener *NodeTCPListener) Encapsulate(carrierData []byte, tcpFrameOffset
return 0, nil return 0, nil
} }
// Demux implements [StackNode].
func (listener *NodeTCPListener) Demux(carrierData []byte, tcpFrameOffset int) error { func (listener *NodeTCPListener) Demux(carrierData []byte, tcpFrameOffset int) error {
if listener.isClosed() { if listener.isClosed() {
return net.ErrClosed return net.ErrClosed
@@ -95,45 +125,45 @@ func (listener *NodeTCPListener) Demux(carrierData []byte, tcpFrameOffset int) e
return errors.New("not our port") return errors.New("not our port")
} }
src := tfrm.DestinationPort() src := tfrm.DestinationPort()
_, flags := tfrm.OffsetAndFlags() for i, conn := range listener.accepted {
for i := range listener.conns { if conn == nil || conn.RemotePort() != src || !bytes.Equal(conn.RemoteAddr(), addr) {
if listener.conns[i].RemotePort() != src || !bytes.Equal(listener.conns[i].RemoteAddr(), addr) {
continue continue
} }
conn := &listener.conns[i]
err := conn.Demux(carrierData, tcpFrameOffset) err := conn.Demux(carrierData, tcpFrameOffset)
if err != nil { if err != nil {
listener.maintainConn(i, err) listener.maintainAccepted(i, err)
} }
return err return err
} }
if !flags.HasAll(tcp.FlagSYN) { _, flags := tfrm.OffsetAndFlags()
if flags != tcp.FlagSYN {
return nil // Not a synchronizing packet, drop it. return nil // Not a synchronizing packet, drop it.
} }
// New connection must be assigned. conn, iss := listener.poolGet()
for i := range listener.conns { if conn == nil {
conn := &listener.conns[i] slog.Error("tcpListener:no-free-conn")
isOpen := !conn.State().IsClosed() return nil
if isOpen {
continue
}
if conn.State() == tcp.StateTimeWait {
conn.Abort()
}
err = conn.OpenListen(dst, tcp.Value(listener.getISS()))
if err != nil {
return err
}
return conn.Demux(carrierData, tcpFrameOffset)
} }
slog.Error("tcpListener:no-free-conn") err = conn.OpenListen(dst, iss)
if err != nil {
slog.Error("NodeTCPListener:open", slog.String("err", err.Error()))
return err // This should not happend
}
err = conn.Demux(carrierData, tcpFrameOffset)
if err != nil {
conn.Abort()
slog.Error("NodeTCPListener:demux", slog.String("err", err.Error()))
return nil
}
listener.ready = append(listener.ready, conn)
return nil return nil
} }
func (listener *NodeTCPListener) maintainConn(connIdx int, err error) { func (listener *NodeTCPListener) maintainAccepted(connIdx int, err error) {
if err == net.ErrClosed { if err == net.ErrClosed {
listener.conns[connIdx].Abort() conn := listener.accepted[connIdx]
listener.poolReturn(conn)
listener.accepted[connIdx] = nil
} }
} }
@@ -141,9 +171,19 @@ func (listener *NodeTCPListener) isClosed() bool {
return listener.port == 0 return listener.port == 0
} }
func (listener *NodeTCPListener) connReceivedSyn(idx int) bool { func (listener *NodeTCPListener) maintainConns() {
return listener.conns[idx].RemotePort() != 0 listener.accepted = removeZeros(listener.accepted)
listener.ready = removeZeros(listener.ready)
} }
func (listener *NodeTCPListener) connAccepted(idx int) bool {
return listener.accepted[idx] func removeZeros[S ~[]E, E comparable](s S) S {
var z E
putIdx := 0
for i := range s {
if s[i] != z {
s[putIdx] = s[i]
putIdx++
}
}
return s[:putIdx]
} }