TCP listener seems to be working really well :)

This commit is contained in:
soypat
2025-06-11 19:38:19 -03:00
parent 63cf2e2669
commit 1fcdd9a7b3
4 changed files with 63 additions and 14 deletions
+3
View File
@@ -65,6 +65,7 @@ func main() {
lastHit := time.Now().Add(-standbyDuration) lastHit := time.Now().Add(-standbyDuration)
var cap pcap.PacketBreakdown var cap pcap.PacketBreakdown
var conn *tcp.Conn var conn *tcp.Conn
accepted := 0
for { for {
nread, err := tap.Read(buf[:]) nread, err := tap.Read(buf[:])
if err != nil { if err != nil {
@@ -90,6 +91,8 @@ func main() {
if err != nil { if err != nil {
lg.Error("tryaccept", slog.String("err", err.Error())) lg.Error("tryaccept", slog.String("err", err.Error()))
} }
accepted++
hdr.Reset(nil)
lg.Info("ACCEPT!") lg.Info("ACCEPT!")
} }
if conn != nil { if conn != nil {
+4 -2
View File
@@ -51,21 +51,23 @@ var (
_ = net.ErrClosed _ = net.ErrClosed
) )
func handleNodeError(nodesPtr *[]node, nodeIdx int, err error) { func handleNodeError(nodesPtr *[]node, nodeIdx int, err error) (discarded bool) {
if err != nil { if err != nil {
if nodeIdx >= len(*nodesPtr) { if nodeIdx >= len(*nodesPtr) {
panic("unreachable") panic("unreachable")
} }
nodes := *nodesPtr nodes := *nodesPtr
badConnID := nodes[nodeIdx].connID != nil && *nodes[nodeIdx].connID != nodes[nodeIdx].currConnID 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) *nodesPtr = slices.Delete(nodes, nodeIdx, nodeIdx+1)
discarded = true
} else { } else {
// Advance Queue of errors // Advance Queue of errors
nodes[nodeIdx].lastErrs[1] = nodes[nodeIdx].lastErrs[0] nodes[nodeIdx].lastErrs[1] = nodes[nodeIdx].lastErrs[0]
nodes[nodeIdx].lastErrs[0] = err nodes[nodeIdx].lastErrs[0] = err
} }
} }
return discarded
} }
func addNode(nodes *[]node, h StackNode, port uint16, protocol uint64) { func addNode(nodes *[]node, h StackNode, port uint16, protocol uint64) {
+53 -11
View File
@@ -82,6 +82,7 @@ func (listener *NodeTCPListener) TryAccept() (*tcp.Conn, error) {
if listener.isClosed() { if listener.isClosed() {
return nil, net.ErrClosed return nil, net.ErrClosed
} }
listener.maintainConns()
for i, conn := range listener.ready { for i, conn := range listener.ready {
if conn == nil { if conn == nil {
continue continue
@@ -99,10 +100,14 @@ func (listener *NodeTCPListener) AcceptRaw() (*tcp.Conn, error) {
if listener.isClosed() || connid != listener.connID { if listener.isClosed() || connid != listener.connID {
return nil, net.ErrClosed 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 { if conn == nil {
continue continue
} }
state := conn.State()
if state != tcp.StateEstablished {
continue // Do not accept until established.
}
listener.accepted = append(listener.accepted, conn) listener.accepted = append(listener.accepted, conn)
listener.ready[i] = nil // discard from ready. listener.ready[i] = nil // discard from ready.
return conn, nil return conn, nil
@@ -124,7 +129,7 @@ func (listener *NodeTCPListener) Encapsulate(carrierData []byte, tcpFrameOffset
} }
n, err := conn.Encapsulate(carrierData, tcpFrameOffset) n, err := conn.Encapsulate(carrierData, tcpFrameOffset)
if err != nil { if err != nil {
listener.maintainAccepted(i, err) err = listener.maintainConn(listener.accepted, i, err)
} }
if n == 0 { if n == 0 {
continue continue
@@ -151,17 +156,17 @@ func (listener *NodeTCPListener) Demux(carrierData []byte, tcpFrameOffset int) e
if dst != listener.port { if dst != listener.port {
return errors.New("not our port") return errors.New("not our port")
} }
src := tfrm.DestinationPort() src := tfrm.SourcePort()
for i, conn := range listener.accepted { // Try to demux in accepted:
if conn == nil || conn.RemotePort() != src || !bytes.Equal(conn.RemoteAddr(), addr) { demuxed, err := listener.tryDemux(listener.accepted, src, addr, carrierData, tcpFrameOffset)
continue if demuxed {
}
err := conn.Demux(carrierData, tcpFrameOffset)
if err != nil {
listener.maintainAccepted(i, err)
}
return err 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() _, flags := tfrm.OffsetAndFlags()
if flags != tcp.FlagSYN { if flags != tcp.FlagSYN {
return nil // Not a synchronizing packet, drop it. return nil // Not a synchronizing packet, drop it.
@@ -186,6 +191,18 @@ func (listener *NodeTCPListener) Demux(carrierData []byte, tcpFrameOffset int) e
return nil 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) { func (listener *NodeTCPListener) maintainAccepted(connIdx int, err error) {
if err == net.ErrClosed { if err == net.ErrClosed {
conn := listener.accepted[connIdx] conn := listener.accepted[connIdx]
@@ -214,3 +231,28 @@ func removeZeros[S ~[]E, E comparable](s S) S {
} }
return s[:putIdx] 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
}
+3 -1
View File
@@ -77,5 +77,7 @@ func (ps *StackPorts) Register(h StackNode) error {
} }
func (ps *StackPorts) handleResult(handlerIdx, n int, err 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())
}
} }