diff --git a/internet/node-tcplistener.go b/internet/node-tcplistener.go index 346f2c4..58e85f9 100644 --- a/internet/node-tcplistener.go +++ b/internet/node-tcplistener.go @@ -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] }