diff --git a/http/httpraw/header.go b/http/httpraw/header.go index 9578267..0274869 100644 --- a/http/httpraw/header.go +++ b/http/httpraw/header.go @@ -111,6 +111,11 @@ func (h *Header) TryParse(asResponse bool) (needMoreData bool, err error) { return err == errNeedMore, err } +// ParsingSuccess returns true if TryParse was successful, that is to say it returned needMoreData==false and err==nil. +func (h *Header) ParsingSuccess() bool { + return h.flags.hasAny(flagDoneParsingHeader) +} + // ReadFromLimited reads at most maxBytesToRead from reader and appends them to underlying buffer. // Used to accumulate HTTP header for later parsing with [Header.TryParse]. // If read is successful (read length>0) and reader returns [io.EOF] then ReadFromLimited will return a nil error. @@ -157,6 +162,14 @@ func (h *Header) ReadFromBytes(b []byte) (int, error) { return len(b), nil } +// BufferReceived returns the amoung of bytes read during calls to Read* methods. +func (h *Header) BufferReceived() int { + if h.flags.hasAny(flagMangledBuffer | flagOOMReached) { + return 0 + } + return len(h.hbuf.buf) +} + // BufferParsed returns the amount of bytes parsed during a call to Parse* methods. // If the Parse* method completed without error then BufferParsed returns the header's length including the final "\r\n\r\n" text. // BufferParsed returns 0 if the buffer is invalid/mangled or if no header data has been parsed succesfully. @@ -321,11 +334,15 @@ func (h *Header) getNonEmptyValue(s headerSlice) []byte { // AppendRequest appends the request header representation to the buffer and returns the result. func (h *Header) AppendRequest(dst []byte) ([]byte, error) { + proto := h.Protocol() if h.flags.hasAny(flagOOMReached) { return dst, errOOM - } else if h.requestURI.len == 0 || h.proto.len == 0 || h.method.len == 0 { - return dst, errors.New("need method/protocol/request URI to create request header") + } else if h.requestURI.len == 0 || h.method.len == 0 { + return dst, errors.New("need method/request URI to create request header") + } else if len(proto) == 0 { + return dst, errNoProto } + method := h.Method() if len(method) == 0 { dst = append(dst, methodGet...) @@ -333,7 +350,6 @@ func (h *Header) AppendRequest(dst []byte) ([]byte, error) { dst = append(dst, method...) } uri := h.RequestURI() - proto := h.Protocol() dst = append(dst, ' ') dst = append(dst, uri...) @@ -348,12 +364,18 @@ func (h *Header) AppendRequest(dst []byte) ([]byte, error) { // AppendResponse appends the response header representation to the buffer and returns the result. func (h *Header) AppendResponse(dst []byte) ([]byte, error) { + proto := h.Protocol() if h.flags.hasAny(flagOOMReached) { return dst, errOOM } else if h.statusCode.len == 0 || h.statusText.len == 0 { return dst, errors.New("invalid status code or text") + } else if len(proto) == 0 { + return dst, errNoProto } code, text := h.Status() + + dst = append(dst, proto...) + dst = append(dst, ' ') dst = append(dst, code...) dst = append(dst, ' ') dst = append(dst, text...) diff --git a/http/httpraw/parse.go b/http/httpraw/parse.go index b8b947e..3229267 100644 --- a/http/httpraw/parse.go +++ b/http/httpraw/parse.go @@ -8,6 +8,7 @@ import ( ) var ( + errNoProto = errors.New("missing protocol, HTTP/0.9 unsupported") errNeedMore = errors.New("need more data: cannot find trailing lf") errUnparsed = errors.New("need to finish parsing") errInvalidName = errors.New("invalid header name") @@ -250,7 +251,7 @@ func (h *Header) appendSlice(value string) headerSlice { h.flags |= flagOOMReached return headerSlice{} } - h.hbuf.buf = slices.Grow(h.hbuf.buf, len(value)) + h.hbuf.buf = slices.Grow(h.hbuf.buf, len(value)+1) // Grow 1 beyond due to slice validity. } h.flags |= flagMangledBuffer return h.hbuf.mustAppendSlice(value) diff --git a/internet/pcap/capture_test.go b/internet/pcap/capture_test.go index 173d7c2..7c6bb6b 100644 --- a/internet/pcap/capture_test.go +++ b/internet/pcap/capture_test.go @@ -30,9 +30,13 @@ func TestCap(t *testing.T) { Flags: tcp.FlagFIN, //tcp.FlagSYN | tcp.FlagACK | tcp.FlagPSH, }) var hdr httpraw.Header + hdr.SetProtocol("HTTP/1.1") hdr.SetStatus("200", "OK") hdr.Set("Cookie", "ABC=123") - pkt, _ = hdr.AppendResponse(pkt) + pkt, err := hdr.AppendResponse(pkt) + if err != nil { + t.Fatal(err) + } pkt = append(pkt, httpBody...) var pbreak PacketBreakdown frames, err := pbreak.CaptureEthernet(nil, pkt, 0) @@ -233,12 +237,12 @@ func TestRightAlignedFields(t *testing.T) { // with right-aligned fields that have trailing bits. func TestAppendFieldRightAligned(t *testing.T) { testCases := []struct { - name string - pkt []byte - fieldBitStart int - bitlen int - rightAligned bool - wantData []byte + name string + pkt []byte + fieldBitStart int + bitlen int + rightAligned bool + wantData []byte }{ { // IPv6 Traffic Class: bits 4-11 (8 bits spanning bytes 0-1) diff --git a/internet/tcplistener_test.go b/internet/tcplistener_test.go index 300bfac..f1ff66f 100644 --- a/internet/tcplistener_test.go +++ b/internet/tcplistener_test.go @@ -47,7 +47,7 @@ func TestListener_SingleConnection(t *testing.T) { if listener.NumberOfReadyToAccept() != 1 { t.Fatalf("after handshake: expected 1 ready, got %d", listener.NumberOfReadyToAccept()) } - acceptedConn, err := listener.TryAccept() + acceptedConn, _, err := listener.TryAccept() if err != nil { t.Fatalf("TryAccept: %v", err) } @@ -91,7 +91,7 @@ func TestListener_AcceptAfterEstablished(t *testing.T) { if listener.NumberOfReadyToAccept() != 1 { t.Fatalf("after client1 handshake: expected 1 ready, got %d", listener.NumberOfReadyToAccept()) } - accepted1, err := listener.TryAccept() + accepted1, _, err := listener.TryAccept() if err != nil { t.Fatalf("TryAccept client1: %v", err) } else if listener.NumberOfReadyToAccept() != 0 { @@ -115,7 +115,7 @@ func TestListener_AcceptAfterEstablished(t *testing.T) { if listener.NumberOfReadyToAccept() != 1 { t.Fatalf("after client2 handshake: expected 1 ready, got %d", listener.NumberOfReadyToAccept()) } - accepted2, err := listener.TryAccept() + accepted2, _, err := listener.TryAccept() if err != nil { t.Fatalf("TryAccept client2: %v", err) } else if listener.NumberOfReadyToAccept() != 0 { @@ -174,7 +174,7 @@ func TestListener_MultiConn(t *testing.T) { // Accept all connections. for i := 0; i < numClients; i++ { var err error - acceptedConns[i], err = listener.TryAccept() + acceptedConns[i], _, err = listener.TryAccept() if err != nil { t.Fatalf("TryAccept client %d: %v", i, err) } @@ -369,16 +369,16 @@ func newMockTCPPool(n, queuesize, bufsize int) *mockTCPPool { return pool } -func (p *mockTCPPool) GetTCP() (*tcp.Conn, tcp.Value) { +func (p *mockTCPPool) GetTCP() (*tcp.Conn, any, tcp.Value) { for i := range p.conns { if !p.acquired[i] { p.acquired[i] = true p.nextISS += 1000 p.naqcuired++ - return &p.conns[i], p.nextISS + return &p.conns[i], nil, p.nextISS } } - return nil, 0 + return nil, nil, 0 } func (p *mockTCPPool) PutTCP(conn *tcp.Conn) { diff --git a/tcp/listener.go b/tcp/listener.go index fd071f5..c74be55 100644 --- a/tcp/listener.go +++ b/tcp/listener.go @@ -13,7 +13,7 @@ import ( // pool is a [sync.Pool] like type pool interface { - GetTCP() (*Conn, Value) + GetTCP() (*Conn, any, Value) PutTCP(*Conn) } @@ -21,15 +21,21 @@ type Listener struct { connID uint64 mu sync.Mutex // incoming stores connections that are potential candidates for acceptance. - incoming []*Conn + incoming []handler // accepted stores all connections that have been accepted and are open. - accepted []*Conn + accepted []handler port uint16 - poolGet func() (*Conn, Value) + poolGet func() (*Conn, any, Value) poolReturn func(*Conn) logger } +type handler struct { + conn *Conn + id uint64 + userData any +} + func (listener *Listener) reset(port uint16, tcppool pool) { listener.accepted = listener.accepted[:0] listener.incoming = listener.incoming[:0] @@ -89,7 +95,8 @@ func (listener *Listener) NumberOfReadyToAccept() (nready int) { if listener.isClosed() { return 0 } - for _, conn := range listener.incoming { + for i := range listener.incoming { + conn := listener.incoming[i].conn if conn == nil || conn.State() != StateEstablished { continue } @@ -99,27 +106,29 @@ func (listener *Listener) NumberOfReadyToAccept() (nready int) { } // TryAccept polls the list of ready connections that have been established -func (listener *Listener) TryAccept() (*Conn, error) { +func (listener *Listener) TryAccept() (*Conn, any, error) { listener.mu.Lock() defer listener.mu.Unlock() if listener.isClosed() { - return nil, net.ErrClosed + return nil, nil, net.ErrClosed } listener.debug("listener:tryaccept", slog.Uint64("port", uint64(listener.port))) listener.maintainConns() - for i, conn := range listener.incoming { + for i := range listener.incoming { + conn := listener.incoming[i].conn if conn == nil || conn.State() != StateEstablished { continue } - listener.accepted = append(listener.accepted, conn) - listener.incoming[i] = nil // discard from ready. - return conn, nil + userData := listener.incoming[i].userData + listener.accepted = append(listener.accepted, listener.incoming[i]) + listener.incoming[i] = handler{} // discard from ready. + return conn, userData, nil } - return nil, errors.New("no conns available") + return nil, nil, errors.New("no conns available") } // Encapsulate implements [StackNode]. -func (listener *Listener) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int) (int, error) { +func (listener *Listener) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int) (n int, err error) { listener.mu.Lock() defer listener.mu.Unlock() if listener.isClosed() { @@ -127,7 +136,8 @@ func (listener *Listener) Encapsulate(carrierData []byte, offsetToIP, offsetToFr } //listener.trace("listener:encaps", slog.Uint64("port", uint64(listener.port))) // First try incoming connections (for handshake SYN-ACK). - for i, conn := range listener.incoming { + for i := range listener.incoming { + conn := listener.incoming[i].conn if conn == nil || conn.State() == StateEstablished { // Nil or already established. continue @@ -143,21 +153,24 @@ func (listener *Listener) Encapsulate(carrierData []byte, offsetToIP, offsetToFr return n, err } // Then try accepted connections. - for i, conn := range listener.accepted { + for i := range listener.accepted { + conn := listener.accepted[i].conn if conn == nil { continue - } - n, err := conn.Encapsulate(carrierData, offsetToIP, offsetToFrame) - if err != nil { - err = listener.maintainConn(listener.accepted, i, err) - } - if n == 0 { + } else if conn.h.connid != listener.accepted[i].id { + listener.returnAccepted(i) continue } - listener.debug("listener:encaps", slog.Uint64("port", uint64(listener.port)), slog.Int("plen", n), slog.String("list", "accepted")) - return n, err + n, err = conn.Encapsulate(carrierData, offsetToIP, offsetToFrame) + if n > 0 { + listener.debug("listener:encaps", slog.Uint64("port", uint64(listener.port)), slog.Int("plen", n), slog.String("list", "accepted")) + break + } else if err == net.ErrClosed { + listener.returnAccepted(i) + err = nil + } } - return 0, nil + return n, err } // Demux implements [StackNode]. @@ -198,7 +211,7 @@ func (listener *Listener) Demux(carrierData []byte, tcpFrameOffset int) error { if flags != FlagSYN { return lneto.ErrPacketDrop // Not a synchronizing packet, drop it. } - conn, iss := listener.poolGet() + conn, userData, iss := listener.poolGet() if conn == nil { slog.Error("tcpListener:no-free-conn") return lneto.ErrPacketDrop @@ -215,15 +228,19 @@ func (listener *Listener) Demux(carrierData []byte, tcpFrameOffset int) error { slog.Error("Listener:demux", slog.String("err", err.Error())) return lneto.ErrPacketDrop } - listener.incoming = append(listener.incoming, conn) + listener.incoming = append(listener.incoming, handler{ + conn: conn, + id: *conn.ConnectionID(), + userData: userData, + }) listener.debug("tcplistener:demux-new", slog.Uint64("lport", uint64(listener.port)), slog.Uint64("rport", uint64(src))) return nil } -func (listener *Listener) tryDemux(conns []*Conn, remotePort uint16, remoteAddr, carrierData []byte, tcpFrameOffset int) (demuxed bool, err error) { +func (listener *Listener) tryDemux(conns []handler, 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) + err := conns[idx].conn.Demux(carrierData, tcpFrameOffset) if err != nil { err = listener.maintainConn(conns, idx, err) } @@ -239,21 +256,22 @@ func (listener *Listener) isClosed() bool { func (listener *Listener) maintainConns() { listener.accepted = internal.DeleteZeroed(listener.accepted) for i := range listener.incoming { - if listener.incoming[i] == nil { + conn := listener.incoming[i].conn + if conn == nil { continue } - state := listener.incoming[i].State() + state := conn.State() if state > StateEstablished || state.IsClosed() { // Something went wrong in handshake or pool aborted/closed the connection. - listener.poolReturn(listener.incoming[i]) - listener.incoming[i] = nil + listener.returnIncoming(i) } } listener.incoming = internal.DeleteZeroed(listener.incoming) } -func getConn(conns []*Conn, remotePort uint16, remoteAddr []byte) int { - for i, conn := range conns { +func getConn(conns []handler, remotePort uint16, remoteAddr []byte) int { + for i := range conns { + conn := conns[i].conn if conn == nil { continue } @@ -266,11 +284,20 @@ func getConn(conns []*Conn, remotePort uint16, remoteAddr []byte) int { return -1 } -func (listener *Listener) maintainConn(conns []*Conn, idx int, err error) error { +func (listener *Listener) maintainConn(conns []handler, idx int, err error) error { if err == net.ErrClosed { - listener.poolReturn(conns[idx]) - conns[idx] = nil + listener.returnAccepted(idx) return nil // avoid closing listener entirely. } return err } + +func (listener *Listener) returnAccepted(idx int) { + listener.poolReturn(listener.accepted[idx].conn) + listener.accepted[idx] = handler{} +} + +func (listener *Listener) returnIncoming(idx int) { + listener.poolReturn(listener.incoming[idx].conn) + listener.incoming[idx] = handler{} +} diff --git a/x/xnet/tcppool.go b/x/xnet/tcppool.go index 279a12d..f881f2e 100644 --- a/x/xnet/tcppool.go +++ b/x/xnet/tcppool.go @@ -15,6 +15,7 @@ type TCPPool struct { mu sync.Mutex naqcuired int conns []tcp.Conn + userData []any acquiredAt []time.Time closingAt []time.Time abortedAt []time.Time @@ -31,9 +32,11 @@ func _() { } type TCPPoolConfig struct { - PoolSize int - QueueSize int - BufferSize int + PoolSize int + QueueSize int + TxBufSize int + RxBufSize int + Logger *slog.Logger ConnLogger *slog.Logger Now func() time.Time @@ -43,6 +46,8 @@ type TCPPoolConfig struct { // ClosingTimeout sets the timeout for a TCP connection to close and be returned to Pool. // If the connection is not closed in this time it will be aborted by the pool. ClosingTimeout time.Duration + // NewUserData is used to create user data used for each individual TCP connection and returned on GetTCP. + NewUserData func() any } func NewTCPPool(cfg TCPPoolConfig) (*TCPPool, error) { @@ -50,29 +55,34 @@ func NewTCPPool(cfg TCPPoolConfig) (*TCPPool, error) { return nil, errors.New("invalid timeout") } n := cfg.PoolSize - bufsize := cfg.BufferSize pool := &TCPPool{ acquiredAt: make([]time.Time, n), closingAt: make([]time.Time, n), abortedAt: make([]time.Time, n), conns: make([]tcp.Conn, n), + userData: make([]any, n), _now: cfg.Now, estbTimeout: cfg.EstablishedTimeout, closingTimeout: cfg.ClosingTimeout, logger: cfg.Logger, } - bufSpace := make([]byte, 2*n*bufsize) + allocPerConn := cfg.TxBufSize + cfg.RxBufSize + bufSpace := make([]byte, n*allocPerConn) for i := range pool.conns { - bufoff := 2 * i * bufsize + bufoff := i * allocPerConn + txOff := bufoff + cfg.RxBufSize err := pool.conns[i].Configure(tcp.ConnConfig{ - RxBuf: bufSpace[bufoff : bufoff+bufsize], - TxBuf: bufSpace[bufoff+bufsize : bufoff+2*bufsize], + RxBuf: bufSpace[bufoff:txOff], + TxBuf: bufSpace[txOff : txOff+cfg.TxBufSize], TxPacketQueueSize: cfg.QueueSize, Logger: cfg.ConnLogger, }) if err != nil { return nil, err } + if cfg.NewUserData != nil { + pool.userData[i] = cfg.NewUserData() + } } return pool, nil } @@ -83,7 +93,7 @@ func (p *TCPPool) NumberOfAcquired() int { return p.naqcuired } -func (p *TCPPool) GetTCP() (*tcp.Conn, tcp.Value) { +func (p *TCPPool) GetTCP() (conn *tcp.Conn, userData any, SuggestedISS tcp.Value) { p.mu.Lock() defer p.mu.Unlock() p.debug("TCPPool:get") @@ -92,10 +102,10 @@ func (p *TCPPool) GetTCP() (*tcp.Conn, tcp.Value) { p.acquiredAt[i] = p.now() p.nextISS += 1000 p.naqcuired++ - return &p.conns[i], p.nextISS + return &p.conns[i], p.userData[i], p.nextISS } } - return nil, 0 + return nil, nil, 0 } func (p *TCPPool) PutTCP(conn *tcp.Conn) { @@ -122,7 +132,8 @@ func (p *TCPPool) CheckTimeouts() { defer p.mu.Unlock() p.debug("TCPPool:checktimeouts", slog.Int("acq", p.naqcuired)) for i := range p.conns { - st := p.conns[i].State() + conn := &p.conns[i] + st := conn.State() if st == tcp.StateEstablished { continue } @@ -134,18 +145,18 @@ func (p *TCPPool) CheckTimeouts() { } else if st.IsPreestablished() && p.since(acq) > p.estbTimeout { // Was acquired and did not reach establishment state so we close. // This is part of a syn-flood defense mechanism. - p.conns[i].Close() + conn.Close() } else if st.IsClosed() || st.IsClosing() { // p.mu.Lock() if p.closingAt[i].IsZero() { p.closingAt[i] = p.now() } else if p.abortedAt[i].IsZero() && p.since(p.closingAt[i]) > p.closingTimeout { p.abortedAt[i] = p.now() - p.conns[i].Abort() - } else if p.since(p.abortedAt[i]) > 10*time.Second { + conn.Abort() + } else if !p.abortedAt[i].IsZero() && p.since(p.abortedAt[i]) > 10*time.Second { println("connection aborted and still not returned to TCPPool") + println("source", conn.LocalPort(), "remote", conn.RemotePort(), "state", conn.State().String()) } - // p.mu.Unlock() } } } diff --git a/x/xnet/xnet_concurrent_test.go b/x/xnet/xnet_concurrent_test.go index 0584dc2..1f5e941 100644 --- a/x/xnet/xnet_concurrent_test.go +++ b/x/xnet/xnet_concurrent_test.go @@ -41,7 +41,8 @@ func TestTCPListener_ConcurrentEcho(t *testing.T) { tcpPool, err := NewTCPPool(TCPPoolConfig{ PoolSize: numClients, QueueSize: 4, - BufferSize: 512, + TxBufSize: 512, + RxBufSize: 512, EstablishedTimeout: 5 * time.Second, ClosingTimeout: 5 * time.Second, }) @@ -198,7 +199,7 @@ func echoServer(ctx context.Context, listener *tcp.Listener) { continue } - conn, err := listener.TryAccept() + conn, _, err := listener.TryAccept() if err != nil || conn == nil { continue } diff --git a/x/xnet/xnet_listener_test.go b/x/xnet/xnet_listener_test.go index d8c9753..a7a1a42 100644 --- a/x/xnet/xnet_listener_test.go +++ b/x/xnet/xnet_listener_test.go @@ -56,7 +56,8 @@ func TestStackAsyncListener_SingleConnection(t *testing.T) { pool, err := NewTCPPool(TCPPoolConfig{ PoolSize: 1, QueueSize: 4, - BufferSize: MTU, + TxBufSize: MTU, + RxBufSize: MTU, EstablishedTimeout: 10e9, ClosingTimeout: 10e9, }) @@ -89,7 +90,7 @@ func TestStackAsyncListener_SingleConnection(t *testing.T) { if listener.NumberOfReadyToAccept() != 1 { t.Fatalf("after handshake: expected 1 ready, got %d", listener.NumberOfReadyToAccept()) } - svConn, err := listener.TryAccept() + svConn, _, err := listener.TryAccept() if err != nil { t.Fatalf("TryAccept: %v", err) } @@ -142,7 +143,8 @@ func TestStackAsyncListener_MultiSequentialConn(t *testing.T) { pool, err := NewTCPPool(TCPPoolConfig{ PoolSize: poolsize, QueueSize: 4, - BufferSize: bufsize, + TxBufSize: bufsize, + RxBufSize: bufsize, EstablishedTimeout: 10e9, ClosingTimeout: 10e9, }) @@ -198,7 +200,7 @@ func TestStackAsyncListener_MultiSequentialConn(t *testing.T) { if listener.NumberOfReadyToAccept() != 1 { t.Fatalf("after handshake: expected 1 ready, got %d", listener.NumberOfReadyToAccept()) } - svconn, err := listener.TryAccept() + svconn, _, err := listener.TryAccept() if err != nil { t.Fatal(err) } else if svconn.RemotePort() != clConn.LocalPort() ||