mirror of
https://github.com/soypat/lneto.git
synced 2026-09-06 06:49:06 +00:00
HTTP and Listener improvements and bugfixes (#17)
* listener revamp * improvements to conn lifetimes
This commit is contained in:
+25
-3
@@ -111,6 +111,11 @@ func (h *Header) TryParse(asResponse bool) (needMoreData bool, err error) {
|
|||||||
return err == errNeedMore, err
|
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.
|
// ReadFromLimited reads at most maxBytesToRead from reader and appends them to underlying buffer.
|
||||||
// Used to accumulate HTTP header for later parsing with [Header.TryParse].
|
// 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.
|
// 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
|
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.
|
// 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.
|
// 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.
|
// 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.
|
// AppendRequest appends the request header representation to the buffer and returns the result.
|
||||||
func (h *Header) AppendRequest(dst []byte) ([]byte, error) {
|
func (h *Header) AppendRequest(dst []byte) ([]byte, error) {
|
||||||
|
proto := h.Protocol()
|
||||||
if h.flags.hasAny(flagOOMReached) {
|
if h.flags.hasAny(flagOOMReached) {
|
||||||
return dst, errOOM
|
return dst, errOOM
|
||||||
} else if h.requestURI.len == 0 || h.proto.len == 0 || h.method.len == 0 {
|
} else if h.requestURI.len == 0 || h.method.len == 0 {
|
||||||
return dst, errors.New("need method/protocol/request URI to create request header")
|
return dst, errors.New("need method/request URI to create request header")
|
||||||
|
} else if len(proto) == 0 {
|
||||||
|
return dst, errNoProto
|
||||||
}
|
}
|
||||||
|
|
||||||
method := h.Method()
|
method := h.Method()
|
||||||
if len(method) == 0 {
|
if len(method) == 0 {
|
||||||
dst = append(dst, methodGet...)
|
dst = append(dst, methodGet...)
|
||||||
@@ -333,7 +350,6 @@ func (h *Header) AppendRequest(dst []byte) ([]byte, error) {
|
|||||||
dst = append(dst, method...)
|
dst = append(dst, method...)
|
||||||
}
|
}
|
||||||
uri := h.RequestURI()
|
uri := h.RequestURI()
|
||||||
proto := h.Protocol()
|
|
||||||
|
|
||||||
dst = append(dst, ' ')
|
dst = append(dst, ' ')
|
||||||
dst = append(dst, uri...)
|
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.
|
// AppendResponse appends the response header representation to the buffer and returns the result.
|
||||||
func (h *Header) AppendResponse(dst []byte) ([]byte, error) {
|
func (h *Header) AppendResponse(dst []byte) ([]byte, error) {
|
||||||
|
proto := h.Protocol()
|
||||||
if h.flags.hasAny(flagOOMReached) {
|
if h.flags.hasAny(flagOOMReached) {
|
||||||
return dst, errOOM
|
return dst, errOOM
|
||||||
} else if h.statusCode.len == 0 || h.statusText.len == 0 {
|
} else if h.statusCode.len == 0 || h.statusText.len == 0 {
|
||||||
return dst, errors.New("invalid status code or text")
|
return dst, errors.New("invalid status code or text")
|
||||||
|
} else if len(proto) == 0 {
|
||||||
|
return dst, errNoProto
|
||||||
}
|
}
|
||||||
code, text := h.Status()
|
code, text := h.Status()
|
||||||
|
|
||||||
|
dst = append(dst, proto...)
|
||||||
|
dst = append(dst, ' ')
|
||||||
dst = append(dst, code...)
|
dst = append(dst, code...)
|
||||||
dst = append(dst, ' ')
|
dst = append(dst, ' ')
|
||||||
dst = append(dst, text...)
|
dst = append(dst, text...)
|
||||||
|
|||||||
@@ -8,6 +8,7 @@ import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
var (
|
var (
|
||||||
|
errNoProto = errors.New("missing protocol, HTTP/0.9 unsupported")
|
||||||
errNeedMore = errors.New("need more data: cannot find trailing lf")
|
errNeedMore = errors.New("need more data: cannot find trailing lf")
|
||||||
errUnparsed = errors.New("need to finish parsing")
|
errUnparsed = errors.New("need to finish parsing")
|
||||||
errInvalidName = errors.New("invalid header name")
|
errInvalidName = errors.New("invalid header name")
|
||||||
@@ -250,7 +251,7 @@ func (h *Header) appendSlice(value string) headerSlice {
|
|||||||
h.flags |= flagOOMReached
|
h.flags |= flagOOMReached
|
||||||
return headerSlice{}
|
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
|
h.flags |= flagMangledBuffer
|
||||||
return h.hbuf.mustAppendSlice(value)
|
return h.hbuf.mustAppendSlice(value)
|
||||||
|
|||||||
@@ -30,9 +30,13 @@ func TestCap(t *testing.T) {
|
|||||||
Flags: tcp.FlagFIN, //tcp.FlagSYN | tcp.FlagACK | tcp.FlagPSH,
|
Flags: tcp.FlagFIN, //tcp.FlagSYN | tcp.FlagACK | tcp.FlagPSH,
|
||||||
})
|
})
|
||||||
var hdr httpraw.Header
|
var hdr httpraw.Header
|
||||||
|
hdr.SetProtocol("HTTP/1.1")
|
||||||
hdr.SetStatus("200", "OK")
|
hdr.SetStatus("200", "OK")
|
||||||
hdr.Set("Cookie", "ABC=123")
|
hdr.Set("Cookie", "ABC=123")
|
||||||
pkt, _ = hdr.AppendResponse(pkt)
|
pkt, err := hdr.AppendResponse(pkt)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
pkt = append(pkt, httpBody...)
|
pkt = append(pkt, httpBody...)
|
||||||
var pbreak PacketBreakdown
|
var pbreak PacketBreakdown
|
||||||
frames, err := pbreak.CaptureEthernet(nil, pkt, 0)
|
frames, err := pbreak.CaptureEthernet(nil, pkt, 0)
|
||||||
@@ -233,12 +237,12 @@ func TestRightAlignedFields(t *testing.T) {
|
|||||||
// with right-aligned fields that have trailing bits.
|
// with right-aligned fields that have trailing bits.
|
||||||
func TestAppendFieldRightAligned(t *testing.T) {
|
func TestAppendFieldRightAligned(t *testing.T) {
|
||||||
testCases := []struct {
|
testCases := []struct {
|
||||||
name string
|
name string
|
||||||
pkt []byte
|
pkt []byte
|
||||||
fieldBitStart int
|
fieldBitStart int
|
||||||
bitlen int
|
bitlen int
|
||||||
rightAligned bool
|
rightAligned bool
|
||||||
wantData []byte
|
wantData []byte
|
||||||
}{
|
}{
|
||||||
{
|
{
|
||||||
// IPv6 Traffic Class: bits 4-11 (8 bits spanning bytes 0-1)
|
// IPv6 Traffic Class: bits 4-11 (8 bits spanning bytes 0-1)
|
||||||
|
|||||||
@@ -47,7 +47,7 @@ func TestListener_SingleConnection(t *testing.T) {
|
|||||||
if listener.NumberOfReadyToAccept() != 1 {
|
if listener.NumberOfReadyToAccept() != 1 {
|
||||||
t.Fatalf("after handshake: expected 1 ready, got %d", listener.NumberOfReadyToAccept())
|
t.Fatalf("after handshake: expected 1 ready, got %d", listener.NumberOfReadyToAccept())
|
||||||
}
|
}
|
||||||
acceptedConn, err := listener.TryAccept()
|
acceptedConn, _, err := listener.TryAccept()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("TryAccept: %v", err)
|
t.Fatalf("TryAccept: %v", err)
|
||||||
}
|
}
|
||||||
@@ -91,7 +91,7 @@ func TestListener_AcceptAfterEstablished(t *testing.T) {
|
|||||||
if listener.NumberOfReadyToAccept() != 1 {
|
if listener.NumberOfReadyToAccept() != 1 {
|
||||||
t.Fatalf("after client1 handshake: expected 1 ready, got %d", listener.NumberOfReadyToAccept())
|
t.Fatalf("after client1 handshake: expected 1 ready, got %d", listener.NumberOfReadyToAccept())
|
||||||
}
|
}
|
||||||
accepted1, err := listener.TryAccept()
|
accepted1, _, err := listener.TryAccept()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("TryAccept client1: %v", err)
|
t.Fatalf("TryAccept client1: %v", err)
|
||||||
} else if listener.NumberOfReadyToAccept() != 0 {
|
} else if listener.NumberOfReadyToAccept() != 0 {
|
||||||
@@ -115,7 +115,7 @@ func TestListener_AcceptAfterEstablished(t *testing.T) {
|
|||||||
if listener.NumberOfReadyToAccept() != 1 {
|
if listener.NumberOfReadyToAccept() != 1 {
|
||||||
t.Fatalf("after client2 handshake: expected 1 ready, got %d", listener.NumberOfReadyToAccept())
|
t.Fatalf("after client2 handshake: expected 1 ready, got %d", listener.NumberOfReadyToAccept())
|
||||||
}
|
}
|
||||||
accepted2, err := listener.TryAccept()
|
accepted2, _, err := listener.TryAccept()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("TryAccept client2: %v", err)
|
t.Fatalf("TryAccept client2: %v", err)
|
||||||
} else if listener.NumberOfReadyToAccept() != 0 {
|
} else if listener.NumberOfReadyToAccept() != 0 {
|
||||||
@@ -174,7 +174,7 @@ func TestListener_MultiConn(t *testing.T) {
|
|||||||
// Accept all connections.
|
// Accept all connections.
|
||||||
for i := 0; i < numClients; i++ {
|
for i := 0; i < numClients; i++ {
|
||||||
var err error
|
var err error
|
||||||
acceptedConns[i], err = listener.TryAccept()
|
acceptedConns[i], _, err = listener.TryAccept()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("TryAccept client %d: %v", i, err)
|
t.Fatalf("TryAccept client %d: %v", i, err)
|
||||||
}
|
}
|
||||||
@@ -369,16 +369,16 @@ func newMockTCPPool(n, queuesize, bufsize int) *mockTCPPool {
|
|||||||
return pool
|
return pool
|
||||||
}
|
}
|
||||||
|
|
||||||
func (p *mockTCPPool) GetTCP() (*tcp.Conn, tcp.Value) {
|
func (p *mockTCPPool) GetTCP() (*tcp.Conn, any, tcp.Value) {
|
||||||
for i := range p.conns {
|
for i := range p.conns {
|
||||||
if !p.acquired[i] {
|
if !p.acquired[i] {
|
||||||
p.acquired[i] = true
|
p.acquired[i] = true
|
||||||
p.nextISS += 1000
|
p.nextISS += 1000
|
||||||
p.naqcuired++
|
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) {
|
func (p *mockTCPPool) PutTCP(conn *tcp.Conn) {
|
||||||
|
|||||||
+64
-37
@@ -13,7 +13,7 @@ import (
|
|||||||
|
|
||||||
// pool is a [sync.Pool] like
|
// pool is a [sync.Pool] like
|
||||||
type pool interface {
|
type pool interface {
|
||||||
GetTCP() (*Conn, Value)
|
GetTCP() (*Conn, any, Value)
|
||||||
PutTCP(*Conn)
|
PutTCP(*Conn)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -21,15 +21,21 @@ type Listener struct {
|
|||||||
connID uint64
|
connID uint64
|
||||||
mu sync.Mutex
|
mu sync.Mutex
|
||||||
// incoming stores connections that are potential candidates for acceptance.
|
// 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 stores all connections that have been accepted and are open.
|
||||||
accepted []*Conn
|
accepted []handler
|
||||||
port uint16
|
port uint16
|
||||||
poolGet func() (*Conn, Value)
|
poolGet func() (*Conn, any, Value)
|
||||||
poolReturn func(*Conn)
|
poolReturn func(*Conn)
|
||||||
logger
|
logger
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type handler struct {
|
||||||
|
conn *Conn
|
||||||
|
id uint64
|
||||||
|
userData any
|
||||||
|
}
|
||||||
|
|
||||||
func (listener *Listener) reset(port uint16, tcppool pool) {
|
func (listener *Listener) reset(port uint16, tcppool pool) {
|
||||||
listener.accepted = listener.accepted[:0]
|
listener.accepted = listener.accepted[:0]
|
||||||
listener.incoming = listener.incoming[:0]
|
listener.incoming = listener.incoming[:0]
|
||||||
@@ -89,7 +95,8 @@ func (listener *Listener) NumberOfReadyToAccept() (nready int) {
|
|||||||
if listener.isClosed() {
|
if listener.isClosed() {
|
||||||
return 0
|
return 0
|
||||||
}
|
}
|
||||||
for _, conn := range listener.incoming {
|
for i := range listener.incoming {
|
||||||
|
conn := listener.incoming[i].conn
|
||||||
if conn == nil || conn.State() != StateEstablished {
|
if conn == nil || conn.State() != StateEstablished {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
@@ -99,27 +106,29 @@ func (listener *Listener) NumberOfReadyToAccept() (nready int) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// TryAccept polls the list of ready connections that have been established
|
// 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()
|
listener.mu.Lock()
|
||||||
defer listener.mu.Unlock()
|
defer listener.mu.Unlock()
|
||||||
if listener.isClosed() {
|
if listener.isClosed() {
|
||||||
return nil, net.ErrClosed
|
return nil, nil, net.ErrClosed
|
||||||
}
|
}
|
||||||
listener.debug("listener:tryaccept", slog.Uint64("port", uint64(listener.port)))
|
listener.debug("listener:tryaccept", slog.Uint64("port", uint64(listener.port)))
|
||||||
listener.maintainConns()
|
listener.maintainConns()
|
||||||
for i, conn := range listener.incoming {
|
for i := range listener.incoming {
|
||||||
|
conn := listener.incoming[i].conn
|
||||||
if conn == nil || conn.State() != StateEstablished {
|
if conn == nil || conn.State() != StateEstablished {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
listener.accepted = append(listener.accepted, conn)
|
userData := listener.incoming[i].userData
|
||||||
listener.incoming[i] = nil // discard from ready.
|
listener.accepted = append(listener.accepted, listener.incoming[i])
|
||||||
return conn, nil
|
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].
|
// 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()
|
listener.mu.Lock()
|
||||||
defer listener.mu.Unlock()
|
defer listener.mu.Unlock()
|
||||||
if listener.isClosed() {
|
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)))
|
//listener.trace("listener:encaps", slog.Uint64("port", uint64(listener.port)))
|
||||||
// First try incoming connections (for handshake SYN-ACK).
|
// 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 {
|
if conn == nil || conn.State() == StateEstablished {
|
||||||
// Nil or already established.
|
// Nil or already established.
|
||||||
continue
|
continue
|
||||||
@@ -143,21 +153,24 @@ func (listener *Listener) Encapsulate(carrierData []byte, offsetToIP, offsetToFr
|
|||||||
return n, err
|
return n, err
|
||||||
}
|
}
|
||||||
// Then try accepted connections.
|
// Then try accepted connections.
|
||||||
for i, conn := range listener.accepted {
|
for i := range listener.accepted {
|
||||||
|
conn := listener.accepted[i].conn
|
||||||
if conn == nil {
|
if conn == nil {
|
||||||
continue
|
continue
|
||||||
}
|
} else if conn.h.connid != listener.accepted[i].id {
|
||||||
n, err := conn.Encapsulate(carrierData, offsetToIP, offsetToFrame)
|
listener.returnAccepted(i)
|
||||||
if err != nil {
|
|
||||||
err = listener.maintainConn(listener.accepted, i, err)
|
|
||||||
}
|
|
||||||
if n == 0 {
|
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
listener.debug("listener:encaps", slog.Uint64("port", uint64(listener.port)), slog.Int("plen", n), slog.String("list", "accepted"))
|
n, err = conn.Encapsulate(carrierData, offsetToIP, offsetToFrame)
|
||||||
return n, err
|
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].
|
// Demux implements [StackNode].
|
||||||
@@ -198,7 +211,7 @@ func (listener *Listener) Demux(carrierData []byte, tcpFrameOffset int) error {
|
|||||||
if flags != FlagSYN {
|
if flags != FlagSYN {
|
||||||
return lneto.ErrPacketDrop // Not a synchronizing packet, drop it.
|
return lneto.ErrPacketDrop // Not a synchronizing packet, drop it.
|
||||||
}
|
}
|
||||||
conn, iss := listener.poolGet()
|
conn, userData, iss := listener.poolGet()
|
||||||
if conn == nil {
|
if conn == nil {
|
||||||
slog.Error("tcpListener:no-free-conn")
|
slog.Error("tcpListener:no-free-conn")
|
||||||
return lneto.ErrPacketDrop
|
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()))
|
slog.Error("Listener:demux", slog.String("err", err.Error()))
|
||||||
return lneto.ErrPacketDrop
|
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)))
|
listener.debug("tcplistener:demux-new", slog.Uint64("lport", uint64(listener.port)), slog.Uint64("rport", uint64(src)))
|
||||||
return nil
|
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)
|
idx := getConn(conns, remotePort, remoteAddr)
|
||||||
if idx >= 0 {
|
if idx >= 0 {
|
||||||
err := conns[idx].Demux(carrierData, tcpFrameOffset)
|
err := conns[idx].conn.Demux(carrierData, tcpFrameOffset)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
err = listener.maintainConn(conns, idx, err)
|
err = listener.maintainConn(conns, idx, err)
|
||||||
}
|
}
|
||||||
@@ -239,21 +256,22 @@ func (listener *Listener) isClosed() bool {
|
|||||||
func (listener *Listener) maintainConns() {
|
func (listener *Listener) maintainConns() {
|
||||||
listener.accepted = internal.DeleteZeroed(listener.accepted)
|
listener.accepted = internal.DeleteZeroed(listener.accepted)
|
||||||
for i := range listener.incoming {
|
for i := range listener.incoming {
|
||||||
if listener.incoming[i] == nil {
|
conn := listener.incoming[i].conn
|
||||||
|
if conn == nil {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
state := listener.incoming[i].State()
|
state := conn.State()
|
||||||
if state > StateEstablished || state.IsClosed() {
|
if state > StateEstablished || state.IsClosed() {
|
||||||
// Something went wrong in handshake or pool aborted/closed the connection.
|
// Something went wrong in handshake or pool aborted/closed the connection.
|
||||||
listener.poolReturn(listener.incoming[i])
|
listener.returnIncoming(i)
|
||||||
listener.incoming[i] = nil
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
listener.incoming = internal.DeleteZeroed(listener.incoming)
|
listener.incoming = internal.DeleteZeroed(listener.incoming)
|
||||||
}
|
}
|
||||||
|
|
||||||
func getConn(conns []*Conn, remotePort uint16, remoteAddr []byte) int {
|
func getConn(conns []handler, remotePort uint16, remoteAddr []byte) int {
|
||||||
for i, conn := range conns {
|
for i := range conns {
|
||||||
|
conn := conns[i].conn
|
||||||
if conn == nil {
|
if conn == nil {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
@@ -266,11 +284,20 @@ func getConn(conns []*Conn, remotePort uint16, remoteAddr []byte) int {
|
|||||||
return -1
|
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 {
|
if err == net.ErrClosed {
|
||||||
listener.poolReturn(conns[idx])
|
listener.returnAccepted(idx)
|
||||||
conns[idx] = nil
|
|
||||||
return nil // avoid closing listener entirely.
|
return nil // avoid closing listener entirely.
|
||||||
}
|
}
|
||||||
return err
|
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{}
|
||||||
|
}
|
||||||
|
|||||||
+27
-16
@@ -15,6 +15,7 @@ type TCPPool struct {
|
|||||||
mu sync.Mutex
|
mu sync.Mutex
|
||||||
naqcuired int
|
naqcuired int
|
||||||
conns []tcp.Conn
|
conns []tcp.Conn
|
||||||
|
userData []any
|
||||||
acquiredAt []time.Time
|
acquiredAt []time.Time
|
||||||
closingAt []time.Time
|
closingAt []time.Time
|
||||||
abortedAt []time.Time
|
abortedAt []time.Time
|
||||||
@@ -31,9 +32,11 @@ func _() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
type TCPPoolConfig struct {
|
type TCPPoolConfig struct {
|
||||||
PoolSize int
|
PoolSize int
|
||||||
QueueSize int
|
QueueSize int
|
||||||
BufferSize int
|
TxBufSize int
|
||||||
|
RxBufSize int
|
||||||
|
|
||||||
Logger *slog.Logger
|
Logger *slog.Logger
|
||||||
ConnLogger *slog.Logger
|
ConnLogger *slog.Logger
|
||||||
Now func() time.Time
|
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.
|
// 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.
|
// If the connection is not closed in this time it will be aborted by the pool.
|
||||||
ClosingTimeout time.Duration
|
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) {
|
func NewTCPPool(cfg TCPPoolConfig) (*TCPPool, error) {
|
||||||
@@ -50,29 +55,34 @@ func NewTCPPool(cfg TCPPoolConfig) (*TCPPool, error) {
|
|||||||
return nil, errors.New("invalid timeout")
|
return nil, errors.New("invalid timeout")
|
||||||
}
|
}
|
||||||
n := cfg.PoolSize
|
n := cfg.PoolSize
|
||||||
bufsize := cfg.BufferSize
|
|
||||||
pool := &TCPPool{
|
pool := &TCPPool{
|
||||||
acquiredAt: make([]time.Time, n),
|
acquiredAt: make([]time.Time, n),
|
||||||
closingAt: make([]time.Time, n),
|
closingAt: make([]time.Time, n),
|
||||||
abortedAt: make([]time.Time, n),
|
abortedAt: make([]time.Time, n),
|
||||||
conns: make([]tcp.Conn, n),
|
conns: make([]tcp.Conn, n),
|
||||||
|
userData: make([]any, n),
|
||||||
_now: cfg.Now,
|
_now: cfg.Now,
|
||||||
estbTimeout: cfg.EstablishedTimeout,
|
estbTimeout: cfg.EstablishedTimeout,
|
||||||
closingTimeout: cfg.ClosingTimeout,
|
closingTimeout: cfg.ClosingTimeout,
|
||||||
logger: cfg.Logger,
|
logger: cfg.Logger,
|
||||||
}
|
}
|
||||||
bufSpace := make([]byte, 2*n*bufsize)
|
allocPerConn := cfg.TxBufSize + cfg.RxBufSize
|
||||||
|
bufSpace := make([]byte, n*allocPerConn)
|
||||||
for i := range pool.conns {
|
for i := range pool.conns {
|
||||||
bufoff := 2 * i * bufsize
|
bufoff := i * allocPerConn
|
||||||
|
txOff := bufoff + cfg.RxBufSize
|
||||||
err := pool.conns[i].Configure(tcp.ConnConfig{
|
err := pool.conns[i].Configure(tcp.ConnConfig{
|
||||||
RxBuf: bufSpace[bufoff : bufoff+bufsize],
|
RxBuf: bufSpace[bufoff:txOff],
|
||||||
TxBuf: bufSpace[bufoff+bufsize : bufoff+2*bufsize],
|
TxBuf: bufSpace[txOff : txOff+cfg.TxBufSize],
|
||||||
TxPacketQueueSize: cfg.QueueSize,
|
TxPacketQueueSize: cfg.QueueSize,
|
||||||
Logger: cfg.ConnLogger,
|
Logger: cfg.ConnLogger,
|
||||||
})
|
})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
if cfg.NewUserData != nil {
|
||||||
|
pool.userData[i] = cfg.NewUserData()
|
||||||
|
}
|
||||||
}
|
}
|
||||||
return pool, nil
|
return pool, nil
|
||||||
}
|
}
|
||||||
@@ -83,7 +93,7 @@ func (p *TCPPool) NumberOfAcquired() int {
|
|||||||
return p.naqcuired
|
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()
|
p.mu.Lock()
|
||||||
defer p.mu.Unlock()
|
defer p.mu.Unlock()
|
||||||
p.debug("TCPPool:get")
|
p.debug("TCPPool:get")
|
||||||
@@ -92,10 +102,10 @@ func (p *TCPPool) GetTCP() (*tcp.Conn, tcp.Value) {
|
|||||||
p.acquiredAt[i] = p.now()
|
p.acquiredAt[i] = p.now()
|
||||||
p.nextISS += 1000
|
p.nextISS += 1000
|
||||||
p.naqcuired++
|
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) {
|
func (p *TCPPool) PutTCP(conn *tcp.Conn) {
|
||||||
@@ -122,7 +132,8 @@ func (p *TCPPool) CheckTimeouts() {
|
|||||||
defer p.mu.Unlock()
|
defer p.mu.Unlock()
|
||||||
p.debug("TCPPool:checktimeouts", slog.Int("acq", p.naqcuired))
|
p.debug("TCPPool:checktimeouts", slog.Int("acq", p.naqcuired))
|
||||||
for i := range p.conns {
|
for i := range p.conns {
|
||||||
st := p.conns[i].State()
|
conn := &p.conns[i]
|
||||||
|
st := conn.State()
|
||||||
if st == tcp.StateEstablished {
|
if st == tcp.StateEstablished {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
@@ -134,18 +145,18 @@ func (p *TCPPool) CheckTimeouts() {
|
|||||||
} else if st.IsPreestablished() && p.since(acq) > p.estbTimeout {
|
} else if st.IsPreestablished() && p.since(acq) > p.estbTimeout {
|
||||||
// Was acquired and did not reach establishment state so we close.
|
// Was acquired and did not reach establishment state so we close.
|
||||||
// This is part of a syn-flood defense mechanism.
|
// This is part of a syn-flood defense mechanism.
|
||||||
p.conns[i].Close()
|
conn.Close()
|
||||||
} else if st.IsClosed() || st.IsClosing() {
|
} else if st.IsClosed() || st.IsClosing() {
|
||||||
// p.mu.Lock()
|
// p.mu.Lock()
|
||||||
if p.closingAt[i].IsZero() {
|
if p.closingAt[i].IsZero() {
|
||||||
p.closingAt[i] = p.now()
|
p.closingAt[i] = p.now()
|
||||||
} else if p.abortedAt[i].IsZero() && p.since(p.closingAt[i]) > p.closingTimeout {
|
} else if p.abortedAt[i].IsZero() && p.since(p.closingAt[i]) > p.closingTimeout {
|
||||||
p.abortedAt[i] = p.now()
|
p.abortedAt[i] = p.now()
|
||||||
p.conns[i].Abort()
|
conn.Abort()
|
||||||
} else if p.since(p.abortedAt[i]) > 10*time.Second {
|
} else if !p.abortedAt[i].IsZero() && p.since(p.abortedAt[i]) > 10*time.Second {
|
||||||
println("connection aborted and still not returned to TCPPool")
|
println("connection aborted and still not returned to TCPPool")
|
||||||
|
println("source", conn.LocalPort(), "remote", conn.RemotePort(), "state", conn.State().String())
|
||||||
}
|
}
|
||||||
// p.mu.Unlock()
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -41,7 +41,8 @@ func TestTCPListener_ConcurrentEcho(t *testing.T) {
|
|||||||
tcpPool, err := NewTCPPool(TCPPoolConfig{
|
tcpPool, err := NewTCPPool(TCPPoolConfig{
|
||||||
PoolSize: numClients,
|
PoolSize: numClients,
|
||||||
QueueSize: 4,
|
QueueSize: 4,
|
||||||
BufferSize: 512,
|
TxBufSize: 512,
|
||||||
|
RxBufSize: 512,
|
||||||
EstablishedTimeout: 5 * time.Second,
|
EstablishedTimeout: 5 * time.Second,
|
||||||
ClosingTimeout: 5 * time.Second,
|
ClosingTimeout: 5 * time.Second,
|
||||||
})
|
})
|
||||||
@@ -198,7 +199,7 @@ func echoServer(ctx context.Context, listener *tcp.Listener) {
|
|||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
conn, err := listener.TryAccept()
|
conn, _, err := listener.TryAccept()
|
||||||
if err != nil || conn == nil {
|
if err != nil || conn == nil {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -56,7 +56,8 @@ func TestStackAsyncListener_SingleConnection(t *testing.T) {
|
|||||||
pool, err := NewTCPPool(TCPPoolConfig{
|
pool, err := NewTCPPool(TCPPoolConfig{
|
||||||
PoolSize: 1,
|
PoolSize: 1,
|
||||||
QueueSize: 4,
|
QueueSize: 4,
|
||||||
BufferSize: MTU,
|
TxBufSize: MTU,
|
||||||
|
RxBufSize: MTU,
|
||||||
EstablishedTimeout: 10e9,
|
EstablishedTimeout: 10e9,
|
||||||
ClosingTimeout: 10e9,
|
ClosingTimeout: 10e9,
|
||||||
})
|
})
|
||||||
@@ -89,7 +90,7 @@ func TestStackAsyncListener_SingleConnection(t *testing.T) {
|
|||||||
if listener.NumberOfReadyToAccept() != 1 {
|
if listener.NumberOfReadyToAccept() != 1 {
|
||||||
t.Fatalf("after handshake: expected 1 ready, got %d", listener.NumberOfReadyToAccept())
|
t.Fatalf("after handshake: expected 1 ready, got %d", listener.NumberOfReadyToAccept())
|
||||||
}
|
}
|
||||||
svConn, err := listener.TryAccept()
|
svConn, _, err := listener.TryAccept()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("TryAccept: %v", err)
|
t.Fatalf("TryAccept: %v", err)
|
||||||
}
|
}
|
||||||
@@ -142,7 +143,8 @@ func TestStackAsyncListener_MultiSequentialConn(t *testing.T) {
|
|||||||
pool, err := NewTCPPool(TCPPoolConfig{
|
pool, err := NewTCPPool(TCPPoolConfig{
|
||||||
PoolSize: poolsize,
|
PoolSize: poolsize,
|
||||||
QueueSize: 4,
|
QueueSize: 4,
|
||||||
BufferSize: bufsize,
|
TxBufSize: bufsize,
|
||||||
|
RxBufSize: bufsize,
|
||||||
EstablishedTimeout: 10e9,
|
EstablishedTimeout: 10e9,
|
||||||
ClosingTimeout: 10e9,
|
ClosingTimeout: 10e9,
|
||||||
})
|
})
|
||||||
@@ -198,7 +200,7 @@ func TestStackAsyncListener_MultiSequentialConn(t *testing.T) {
|
|||||||
if listener.NumberOfReadyToAccept() != 1 {
|
if listener.NumberOfReadyToAccept() != 1 {
|
||||||
t.Fatalf("after handshake: expected 1 ready, got %d", listener.NumberOfReadyToAccept())
|
t.Fatalf("after handshake: expected 1 ready, got %d", listener.NumberOfReadyToAccept())
|
||||||
}
|
}
|
||||||
svconn, err := listener.TryAccept()
|
svconn, _, err := listener.TryAccept()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
} else if svconn.RemotePort() != clConn.LocalPort() ||
|
} else if svconn.RemotePort() != clConn.LocalPort() ||
|
||||||
|
|||||||
Reference in New Issue
Block a user