From 2e823c0fd579f93087a7386fd26c4ec2b9c33bea Mon Sep 17 00:00:00 2001 From: soypat Date: Sat, 15 Feb 2025 11:15:19 -0300 Subject: [PATCH] begin add tcp.Handler tests --- examples/stack/main.go | 2 +- tcp/control.go | 7 +-- tcp/definitions.go | 3 +- tcp/handler.go | 56 +++++++++++++++++----- tcp/handler_test.go | 105 +++++++++++++++++++++++++++++++++++++---- 5 files changed, 148 insertions(+), 25 deletions(-) diff --git a/examples/stack/main.go b/examples/stack/main.go index c6ce01d..f0151e6 100644 --- a/examples/stack/main.go +++ b/examples/stack/main.go @@ -41,7 +41,7 @@ func main() { if err != nil { log.Fatal(err) } - err = pStack.handler.Open(gen.DstTCP, iss) + err = pStack.handler.OpenListen(gen.DstTCP, iss) if err != nil { log.Fatal(err) } diff --git a/tcp/control.go b/tcp/control.go index b86074f..0077a23 100644 --- a/tcp/control.go +++ b/tcp/control.go @@ -47,13 +47,14 @@ type ControlBlock struct { rcv recvSpace // When FlagRST is set in pending flags rstPtr will contain the sequence number of the RST segment to make it "believable" (See RFC9293) rstPtr Value + logger + // pending is the queue of pending flags to be sent in the next 2 segments. // On a call to Send the queue is advanced and flags set in the segment are unset. // The second position of the queue is used for FIN segments. pending [2]Flags _state State // leading underscore so field not suggested on top of exported State method when developing. challengeAck bool - logger } // State returns the current state of the TCP connection. @@ -144,7 +145,7 @@ type recvSpace struct { func (tcb *ControlBlock) Open(iss Value, wnd Size) (err error) { switch { case tcb._state != StateClosed && tcb._state != StateListen: - err = errTCBNotClosed + err = errNeedClosedTCBToOpen case wnd > math.MaxUint16: err = errWindowTooLarge } @@ -189,7 +190,7 @@ func (tcb *ControlBlock) PendingSegment(payloadLen int) (_ Segment, ok bool) { _ = inFlight maxPayload := tcb.snd.maxSend() if payloadLen > int(maxPayload) { - if maxPayload == 0 && !tcb.pending[0].HasAny(FlagFIN|FlagRST|FlagSYN) { + if maxPayload == 0 && !pending.HasAny(FlagFIN|FlagRST|FlagSYN) { return Segment{}, false } else if maxPayload > tcb.snd.WND { panic("seqs: bad calculation") diff --git a/tcp/definitions.go b/tcp/definitions.go index 9b5f950..0a6f700 100644 --- a/tcp/definitions.go +++ b/tcp/definitions.go @@ -16,7 +16,8 @@ var ( errDropSegment = errors.New("drop segment") errWindowTooLarge = errors.New("invalid window size > 2**16") - errTCBNotClosed = errors.New("TCB not closed") + errBufferTooSmall = errors.New("buffer too small") + errNeedClosedTCBToOpen = errors.New("need closed TCB to call open") errInvalidState = errors.New("invalid state") errConnNotexist = errors.New("connection does not exist") errConnectionClosing = errors.New("connection closing") diff --git a/tcp/handler.go b/tcp/handler.go index 6444bbf..c78bdba 100644 --- a/tcp/handler.go +++ b/tcp/handler.go @@ -19,20 +19,22 @@ var ( // Does NOT implement IP related logic, so no CRC calculation/validation or pseudo header logic. // Does NOT implement connection lifetime handling, so NO deadlines, keepalives, backoffs or anything that requires use of time package. type Handler struct { - scb ControlBlock - bufTx ringTx - bufRx internal.Ring + scb ControlBlock + bufTx ringTx + bufRx internal.Ring + logger + validator lneto2.Validator localPort uint16 remotePort uint16 // connid is a conenction counter that is incremented each time a new // connection is established via Open calls. This disambiguate's whether // Read and Write calls belong to the current connection. - connid uint8 - closing bool - validator lneto2.Validator - logger + connid uint8 + closing bool } +func (h *Handler) State() State { return h.scb.State() } + func (h *Handler) SetBuffers(txbuf, rxbuf []byte, packets int) error { if !h.scb.State().IsClosed() { return errors.New("tcp.Handler must be closed before setting buffers") @@ -43,35 +45,61 @@ func (h *Handler) SetBuffers(txbuf, rxbuf []byte, packets int) error { if len(h.bufRx.Buf) < 1 { return errors.New("short rx buffer") } + h.scb.SetRecvWindow(Size(h.bufRx.Size())) h.bufRx.Reset() return h.bufTx.ResetOrReuse(txbuf, packets, 0) } +// LocalPort returns the local port of the connection. Returns 0 if the connection is closed and uninitialized. func (h *Handler) LocalPort() uint16 { return h.localPort } -// Open prepares a passive connection. Set remote port to non-zero to -func (h *Handler) Open(localPort uint16, iss Value) error { +// RemotePort returns the remote port of the connection if it is set. +// If the connection is passive and has not yet been established it will return 0. +func (h *Handler) RemotePort() uint16 { + return h.remotePort +} + +func (h *Handler) OpenActive(localPort, remotePort uint16, iss Value) error { + if h.bufRx.Size() < minBufferSize || h.bufTx.Size() < minBufferSize { + return errBufferTooSmall + } + if h.scb.State() != StateClosed { + return errNeedClosedTCBToOpen + } + h.reset(localPort, remotePort, iss) + return nil +} + +// OpenListen prepares a passive connection. Set remote port to non-zero to +func (h *Handler) OpenListen(localPort uint16, iss Value) error { + if h.bufRx.Size() < minBufferSize || h.bufTx.Size() < minBufferSize { + return errBufferTooSmall + } // Open will fail unless SCB in closed state. err := h.scb.Open(iss, Size(h.bufRx.Size())) if err != nil { return err } + h.reset(localPort, 0, iss) + return nil +} + +func (h *Handler) reset(localPort, remotePort uint16, iss Value) { *h = Handler{ scb: h.scb, bufTx: h.bufTx, bufRx: h.bufRx, connid: h.connid + 1, localPort: localPort, - remotePort: 0, + remotePort: remotePort, validator: h.validator, logger: h.logger, closing: false, } h.bufTx.ResetOrReuse(nil, 0, iss) h.bufRx.Reset() - return nil } func (h *Handler) Recv(b []byte) error { @@ -132,7 +160,7 @@ func (h *Handler) Recv(b []byte) error { func (h *Handler) Send(b []byte) (int, error) { h.trace("tcp.Handler:start", slog.Uint64("port", uint64(h.localPort))) - if h.isClosed() { + if h.isClosed() && !h.AwaitingSynSend() { return 0, net.ErrClosed } tfrm, err := NewFrame(b) @@ -181,6 +209,10 @@ func (h *Handler) AwaitingSynResponse() bool { return h.remotePort != 0 && h.scb.State() == StateSynSent } +func (h *Handler) AwaitingSynAck() bool { + return h.remotePort == 0 && h.scb.State() == StateListen +} + func (h *Handler) AwaitingSynSend() bool { return h.remotePort != 0 && h.scb.State() == StateClosed } diff --git a/tcp/handler_test.go b/tcp/handler_test.go index aa21417..d55185e 100644 --- a/tcp/handler_test.go +++ b/tcp/handler_test.go @@ -6,15 +6,104 @@ import ( ) func TestHandler(t *testing.T) { - + const mtu = 1500 + const maxpackets = 3 + rng := rand.New(rand.NewSource(0)) + client, server := setupClientServer(t, rng, mtu, mtu, maxpackets, mtu, mtu, maxpackets) + var rawbuf [mtu]byte + establish(t, client, server, rawbuf[:]) } -func setupClientServer(rng *rand.Rand) (client, server Handler) { - - // err := server.Open(StateListen, uint16(rng.Uint32()), 0, 0) - // if err != nil { - // panic(err) - // } - // err = client.Open(State) +func setupClientServer(t *testing.T, rng *rand.Rand, clientTxSize, clientRxSize, clientPackets, serverTxSize, serverRxSize, serverPackets int) (client, server *Handler) { + client = new(Handler) + server = new(Handler) + err := client.SetBuffers(make([]byte, clientTxSize), make([]byte, clientRxSize), clientPackets) + if err != nil { + t.Fatal(err) + } + err = server.SetBuffers(make([]byte, serverTxSize), make([]byte, serverRxSize), serverPackets) + if err != nil { + t.Fatal(err) + } + err = server.OpenListen(uint16(rng.Uint32()), 0) + if err != nil { + t.Fatal(err) + } + err = client.OpenActive(uint16(rng.Uint32()), server.LocalPort(), 0) + if err != nil { + t.Fatal(err) + } + if !client.AwaitingSynSend() { + t.Fatal("client in wrong state") + } + if !server.AwaitingSynAck() { + t.Fatal("server in wrong state") + } return client, server } + +func establish(t *testing.T, client, server *Handler, buf []byte) { + if client.State() != StateClosed { + t.Fatal("client in wrong state") + } else if server.State() != StateListen { + t.Fatal("server in wrong state") + } + clear(buf) + + // Commence 3-way handshake: client sends SYN, server sends SYN-ACK, client sends ACK. + + // Client sends SYN. + n, err := client.Send(buf) + if err != nil { + t.Fatal("client sending:", err) + } else if n < sizeHeaderTCP { + t.Fatal("expected client to send SYN packet") + } else if client.State() != StateSynSent { + t.Fatal("client did not transition to SynSent state:", client.State().String()) + } + err = server.Recv(buf[:n]) // Server receives SYN. + if err != nil { + t.Fatal(err) + } else if server.State() != StateSynRcvd { + t.Fatal("server did not transition to SynReceived state:", server.State().String()) + } + clear(buf) + // Server sends SYNACK response to client's SYN. + n, err = server.Send(buf) + if err != nil { + t.Fatal("server sending:", err) + } else if n < sizeHeaderTCP { + t.Fatal("expected server to send SYNACK packet") + } else if server.State() != StateSynRcvd { + t.Fatal("server should remain in SynReceived state:", server.State().String()) + } + err = client.Recv(buf[:n]) // Client receives SYNACK, is established but must send ACK. + if err != nil { + t.Fatal(err) + } else if client.State() != StateEstablished { + t.Fatal("client did not transition to Established state:", client.State().String()) + } + + clear(buf) + n, err = client.Send(buf) // Client sends ACK. + if err != nil { + t.Fatal("client sending ACK:", err) + } else if n < sizeHeaderTCP { + t.Fatal("expected client to send ACK packet") + } else if client.State() != StateEstablished { + t.Fatal("client should remain in Established state:", client.State().String()) + } + err = server.Recv(buf[:n]) // Server receives ACK. + if err != nil { + t.Fatal(err) + } else if server.State() != StateEstablished { + t.Fatal("server did not transition to Established state on ACK receive:", server.State().String()) + } +} + +func clear[E any, T []E](s T) { + var zero E + for i := range s { + s[i] = zero + } +}