mirror of
https://github.com/soypat/lneto.git
synced 2026-08-12 10:53:44 +00:00
listener refactor: tcpPool+ready/accepted slices
This commit is contained in:
+100
-60
@@ -14,34 +14,31 @@ import (
|
||||
|
||||
var _ StackNode = (*NodeTCPListener)(nil)
|
||||
|
||||
type NodeTCPListener struct {
|
||||
connID uint64
|
||||
conns []tcp.Conn
|
||||
accepted []bool
|
||||
port uint16
|
||||
getISS func() uint32
|
||||
type tcpPool interface {
|
||||
GetTCP() (*tcp.Conn, tcp.Value)
|
||||
PutTCP(*tcp.Conn)
|
||||
}
|
||||
|
||||
func (listener *NodeTCPListener) AcceptRaw() (*tcp.Conn, error) {
|
||||
connid := listener.connID
|
||||
for {
|
||||
if listener.isClosed() || connid != listener.connID {
|
||||
return nil, net.ErrClosed
|
||||
}
|
||||
for i := range listener.conns {
|
||||
isAvailable := listener.connReceivedSyn(i) && !listener.connAccepted(i)
|
||||
if !isAvailable {
|
||||
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")
|
||||
type NodeTCPListener struct {
|
||||
connID uint64
|
||||
// ready have received a
|
||||
ready []*tcp.Conn
|
||||
accepted []*tcp.Conn
|
||||
|
||||
port uint16
|
||||
poolGet func() (*tcp.Conn, tcp.Value)
|
||||
poolReturn func(*tcp.Conn)
|
||||
}
|
||||
|
||||
// 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 {
|
||||
if listener.isClosed() {
|
||||
return errors.New("already closed")
|
||||
@@ -51,24 +48,56 @@ func (listener *NodeTCPListener) Close() error {
|
||||
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) {
|
||||
if listener.isClosed() {
|
||||
return 0, net.ErrClosed
|
||||
}
|
||||
for i := range listener.conns {
|
||||
conn := &listener.conns[i]
|
||||
if conn.State().IsClosed() {
|
||||
for i, conn := range listener.accepted {
|
||||
if conn == nil {
|
||||
continue
|
||||
}
|
||||
n, err := conn.Encapsulate(carrierData, tcpFrameOffset)
|
||||
if err != nil {
|
||||
listener.maintainConn(i, err)
|
||||
listener.maintainAccepted(i, err)
|
||||
}
|
||||
if n == 0 {
|
||||
continue
|
||||
@@ -78,6 +107,7 @@ func (listener *NodeTCPListener) Encapsulate(carrierData []byte, tcpFrameOffset
|
||||
return 0, nil
|
||||
}
|
||||
|
||||
// Demux implements [StackNode].
|
||||
func (listener *NodeTCPListener) Demux(carrierData []byte, tcpFrameOffset int) error {
|
||||
if listener.isClosed() {
|
||||
return net.ErrClosed
|
||||
@@ -95,45 +125,45 @@ func (listener *NodeTCPListener) Demux(carrierData []byte, tcpFrameOffset int) e
|
||||
return errors.New("not our port")
|
||||
}
|
||||
src := tfrm.DestinationPort()
|
||||
_, flags := tfrm.OffsetAndFlags()
|
||||
for i := range listener.conns {
|
||||
if listener.conns[i].RemotePort() != src || !bytes.Equal(listener.conns[i].RemoteAddr(), addr) {
|
||||
for i, conn := range listener.accepted {
|
||||
if conn == nil || conn.RemotePort() != src || !bytes.Equal(conn.RemoteAddr(), addr) {
|
||||
continue
|
||||
}
|
||||
conn := &listener.conns[i]
|
||||
err := conn.Demux(carrierData, tcpFrameOffset)
|
||||
if err != nil {
|
||||
listener.maintainConn(i, err)
|
||||
listener.maintainAccepted(i, err)
|
||||
}
|
||||
return err
|
||||
}
|
||||
if !flags.HasAll(tcp.FlagSYN) {
|
||||
_, flags := tfrm.OffsetAndFlags()
|
||||
if flags != tcp.FlagSYN {
|
||||
return nil // Not a synchronizing packet, drop it.
|
||||
}
|
||||
// New connection must be assigned.
|
||||
for i := range listener.conns {
|
||||
conn := &listener.conns[i]
|
||||
isOpen := !conn.State().IsClosed()
|
||||
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)
|
||||
conn, iss := listener.poolGet()
|
||||
if conn == nil {
|
||||
slog.Error("tcpListener:no-free-conn")
|
||||
return nil
|
||||
}
|
||||
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
|
||||
}
|
||||
|
||||
func (listener *NodeTCPListener) maintainConn(connIdx int, err error) {
|
||||
func (listener *NodeTCPListener) maintainAccepted(connIdx int, err error) {
|
||||
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
|
||||
}
|
||||
|
||||
func (listener *NodeTCPListener) connReceivedSyn(idx int) bool {
|
||||
return listener.conns[idx].RemotePort() != 0
|
||||
func (listener *NodeTCPListener) maintainConns() {
|
||||
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]
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user