diff --git a/examples/stack/main.go b/examples/stack/main.go index 547f7ac..5943fb7 100644 --- a/examples/stack/main.go +++ b/examples/stack/main.go @@ -65,6 +65,7 @@ func main() { lastHit := time.Now().Add(-standbyDuration) var cap pcap.PacketBreakdown var conn *tcp.Conn + accepted := 0 for { nread, err := tap.Read(buf[:]) if err != nil { @@ -90,6 +91,8 @@ func main() { if err != nil { lg.Error("tryaccept", slog.String("err", err.Error())) } + accepted++ + hdr.Reset(nil) lg.Info("ACCEPT!") } if conn != nil { diff --git a/internet/definitions.go b/internet/definitions.go index 23aa4f8..b7f6b6c 100644 --- a/internet/definitions.go +++ b/internet/definitions.go @@ -51,21 +51,23 @@ var ( _ = net.ErrClosed ) -func handleNodeError(nodesPtr *[]node, nodeIdx int, err error) { +func handleNodeError(nodesPtr *[]node, nodeIdx int, err error) (discarded bool) { if err != nil { if nodeIdx >= len(*nodesPtr) { panic("unreachable") } nodes := *nodesPtr badConnID := nodes[nodeIdx].connID != nil && *nodes[nodeIdx].connID != nodes[nodeIdx].currConnID - if err == net.ErrClosed || nodes[nodeIdx].lastErrs[0] == err || nodes[nodeIdx].lastErrs[1] == err || badConnID { + if err == net.ErrClosed || badConnID { *nodesPtr = slices.Delete(nodes, nodeIdx, nodeIdx+1) + discarded = true } else { // Advance Queue of errors nodes[nodeIdx].lastErrs[1] = nodes[nodeIdx].lastErrs[0] nodes[nodeIdx].lastErrs[0] = err } } + return discarded } func addNode(nodes *[]node, h StackNode, port uint16, protocol uint64) { diff --git a/internet/node-tcplistener.go b/internet/node-tcplistener.go index b8608f1..35ee0ba 100644 --- a/internet/node-tcplistener.go +++ b/internet/node-tcplistener.go @@ -82,6 +82,7 @@ func (listener *NodeTCPListener) TryAccept() (*tcp.Conn, error) { if listener.isClosed() { return nil, net.ErrClosed } + listener.maintainConns() for i, conn := range listener.ready { if conn == nil { continue @@ -99,10 +100,14 @@ func (listener *NodeTCPListener) AcceptRaw() (*tcp.Conn, error) { if listener.isClosed() || connid != listener.connID { return nil, net.ErrClosed } - for i, conn := range listener.ready { + for i, conn := range listener.ready { // Scan ready to see if we can accept. if conn == nil { continue } + state := conn.State() + if state != tcp.StateEstablished { + continue // Do not accept until established. + } listener.accepted = append(listener.accepted, conn) listener.ready[i] = nil // discard from ready. return conn, nil @@ -124,7 +129,7 @@ func (listener *NodeTCPListener) Encapsulate(carrierData []byte, tcpFrameOffset } n, err := conn.Encapsulate(carrierData, tcpFrameOffset) if err != nil { - listener.maintainAccepted(i, err) + err = listener.maintainConn(listener.accepted, i, err) } if n == 0 { continue @@ -151,17 +156,17 @@ func (listener *NodeTCPListener) Demux(carrierData []byte, tcpFrameOffset int) e if dst != listener.port { return errors.New("not our port") } - src := tfrm.DestinationPort() - for i, conn := range listener.accepted { - if conn == nil || conn.RemotePort() != src || !bytes.Equal(conn.RemoteAddr(), addr) { - continue - } - err := conn.Demux(carrierData, tcpFrameOffset) - if err != nil { - listener.maintainAccepted(i, err) - } + src := tfrm.SourcePort() + // Try to demux in accepted: + demuxed, err := listener.tryDemux(listener.accepted, src, addr, carrierData, tcpFrameOffset) + if demuxed { return err } + demuxed, err = listener.tryDemux(listener.ready, src, addr, carrierData, tcpFrameOffset) + if demuxed { + return err + } + // Connection not in ready nor accepted. _, flags := tfrm.OffsetAndFlags() if flags != tcp.FlagSYN { return nil // Not a synchronizing packet, drop it. @@ -186,6 +191,18 @@ func (listener *NodeTCPListener) Demux(carrierData []byte, tcpFrameOffset int) e return nil } +func (listener *NodeTCPListener) tryDemux(conns []*tcp.Conn, remotePort uint16, remoteAddr, carrierData []byte, tcpFrameOffset int) (demuxed bool, err error) { + idx := getConn(conns, remotePort, remoteAddr) + if idx >= 0 { + err := conns[idx].Demux(carrierData, tcpFrameOffset) + if err != nil { + err = listener.maintainConn(conns, idx, err) + } + return true, err + } + return false, nil +} + func (listener *NodeTCPListener) maintainAccepted(connIdx int, err error) { if err == net.ErrClosed { conn := listener.accepted[connIdx] @@ -214,3 +231,28 @@ func removeZeros[S ~[]E, E comparable](s S) S { } return s[:putIdx] } + +func getConn(conns []*tcp.Conn, remotePort uint16, remoteAddr []byte) int { + for i, conn := range conns { + if conn == nil { + continue + } + gotPort := conn.RemotePort() + gotaddr := conn.RemoteAddr() + if remotePort == gotPort && bytes.Equal(remoteAddr, gotaddr) { + return i + } + } + return -1 +} + +func (listener *NodeTCPListener) maintainConn(conns []*tcp.Conn, idx int, err error) error { + if err == net.ErrClosed { + println("CLOSING CONN") + conn := conns[idx] + listener.poolReturn(conn) + conns[idx] = nil + return nil // avoid closing listener entirely. + } + return err +} diff --git a/internet/stack-ports.go b/internet/stack-ports.go index cff9811..65ea9dc 100644 --- a/internet/stack-ports.go +++ b/internet/stack-ports.go @@ -77,5 +77,7 @@ func (ps *StackPorts) Register(h StackNode) error { } func (ps *StackPorts) handleResult(handlerIdx, n int, err error) { - handleNodeError(&ps.handlers, handlerIdx, err) + if handleNodeError(&ps.handlers, handlerIdx, err) { + println("DISCARD", handlerIdx, "witherr", err.Error()) + } }