From be949d960c6bda8fc3d296a8f68954a1b8c34739 Mon Sep 17 00:00:00 2001 From: Pat Whittingslow Date: Wed, 25 Feb 2026 16:58:55 +0100 Subject: [PATCH] tcp: mss honoring; accept syn with ECE/CWR flags; add RSTQueue type (#41) * tcp: mss honoring; accept syn with ECE/CWR flags; add RSTQueue type * tcp: move option logic to own file * remove prints in pcap * add MSS send threshold inspired by linux/freebsd/lwip thresh --- internet/pcap/capture.go | 2 - internet/stack-ports.go | 23 +++- internet/tcplistener_test.go | 215 +++++++++++++++++++++++++++++++++++ tcp/control.go | 5 + tcp/definitions.go | 141 ----------------------- tcp/handler.go | 33 ++++-- tcp/handler_test.go | 81 +++++++++++++ tcp/listener.go | 50 ++------ tcp/options.go | 148 ++++++++++++++++++++++++ tcp/rst.go | 63 ++++++++++ 10 files changed, 570 insertions(+), 191 deletions(-) create mode 100644 tcp/options.go create mode 100644 tcp/rst.go diff --git a/internet/pcap/capture.go b/internet/pcap/capture.go index 6bce41b..e8d89f8 100644 --- a/internet/pcap/capture.go +++ b/internet/pcap/capture.go @@ -229,7 +229,6 @@ func (pc *PacketBreakdown) CaptureIPv4(dst []Frame, pkt []byte, bitOffset int) ( if err == nil { tfrm.ValidateSize(pc.validator()) if pc.vld.HasError() { - println("BAD TCP") return dst, pc.vld.ErrPop() } ifrm4.CRCWriteTCPPseudo(&crc) @@ -242,7 +241,6 @@ func (pc *PacketBreakdown) CaptureIPv4(dst []Frame, pkt []byte, bitOffset int) ( if err == nil { ufrm.ValidateSize(pc.validator()) if pc.vld.HasError() { - println("BAD UDP") return dst, pc.vld.ErrPop() } if ufrm.CRC() != 0 { diff --git a/internet/stack-ports.go b/internet/stack-ports.go index e70c0e2..fb4f52c 100644 --- a/internet/stack-ports.go +++ b/internet/stack-ports.go @@ -11,6 +11,7 @@ import ( "github.com/soypat/lneto" "github.com/soypat/lneto/ethernet" "github.com/soypat/lneto/internal" + "github.com/soypat/lneto/tcp" ) type StackPorts struct { @@ -18,7 +19,8 @@ type StackPorts struct { handlers handlers dstPortOff uint16 protocol uint16 - // stores last node to demux/encapsulate. + // rstQueue stores pending RST responses for TCP SYNs to unregistered ports. + rstQueue tcp.RSTQueue } func (ps *StackPorts) ResetUDP(maxNodes int) error { @@ -56,6 +58,9 @@ func (ps *StackPorts) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame return 0, io.ErrShortBuffer } _, n, err = ps.handlers.encapsulateAny(carrierData, offsetToIP, offsetToFrame) + if n == 0 { + n, _ = ps.rstQueue.Drain(carrierData, offsetToIP, offsetToFrame) + } return n, err } @@ -65,6 +70,19 @@ func (ps *StackPorts) Demux(b []byte, offset int) (err error) { } port := binary.BigEndian.Uint16(b[int(ps.dstPortOff)+offset:]) _, err = ps.handlers.demuxByPort(b, offset, port) + if err == lneto.ErrPacketDrop && ps.protocol == uint16(lneto.IPProtoTCP) && offset+14 <= len(b) { + // RFC 9293 §3.10.7.1: RST for SYN to port with no listener. + flags := binary.BigEndian.Uint16(b[offset+12:]) & 0x01ff + const flagSYN, flagRST, flagACK = 0x02, 0x04, 0x10 + if flags&flagSYN != 0 && flags&(flagRST|flagACK) == 0 { + srcaddr, _, _, _, iperr := internal.GetIPAddr(b) + if iperr == nil { + remotePort := binary.BigEndian.Uint16(b[offset:]) + segSeq := tcp.Value(binary.BigEndian.Uint32(b[offset+4:])) + ps.rstQueue.Queue(srcaddr, remotePort, port, 0, segSeq+1, tcp.FlagRST|tcp.FlagACK) + } + } + } return err } @@ -149,5 +167,8 @@ func (ps *StackPortsMACFiltered) Encapsulate(carrierData []byte, offsetToIP, off h.error("handlers:encapsulate", slog.String("func", "encapsulateAny"), slog.String("ctx", h.context), slog.String("err", err.Error())) } } + if n, _ := ps.sp.rstQueue.Drain(carrierData, offsetToIP, offsetToFrame); n > 0 { + return n, nil + } return 0, err // Return last written error. } diff --git a/internet/tcplistener_test.go b/internet/tcplistener_test.go index f0a327c..513ccff 100644 --- a/internet/tcplistener_test.go +++ b/internet/tcplistener_test.go @@ -1,6 +1,7 @@ package internet import ( + "encoding/binary" "math/rand" "net/netip" "testing" @@ -390,6 +391,220 @@ func TestListener_RSTOnPoolExhaustion(t *testing.T) { } } +func TestListener_RSTOnStalePacket(t *testing.T) { + // Test Scenario C: stale FIN,ACK to a port with a listener but no matching connection. + // Test at Listener level directly to avoid StackIP CRC validation. + var listener tcp.Listener + pool := newMockTCPPool(1, 3, 2048) + serverPort := uint16(80) + if err := listener.Reset(serverPort, pool); err != nil { + t.Fatal(err) + } + + // Build a raw stale FIN,ACK targeting port 80 from an unknown source. + // IP header (20 bytes) + TCP header (20 bytes). + clientIP := [4]byte{10, 0, 0, 1} + serverIP := [4]byte{10, 0, 0, 2} + rawBuf := make([]byte, 256) + rawBuf[0] = 0x45 // version=4, IHL=5 + rawBuf[9] = 6 // protocol=TCP + copy(rawBuf[12:16], clientIP[:]) + copy(rawBuf[16:20], serverIP[:]) + // TCP header at offset 20. + binary.BigEndian.PutUint16(rawBuf[20:], 1337) // src port + binary.BigEndian.PutUint16(rawBuf[22:], serverPort) // dst port + binary.BigEndian.PutUint32(rawBuf[24:], 500) // SEQ + binary.BigEndian.PutUint32(rawBuf[28:], 200) // ACK + rawBuf[32] = 0x50 // offset=5 + rawBuf[33] = 0x11 // flags = FIN|ACK (0x01|0x10) + + // Demux directly on Listener — no matching connection, should queue RST. + err := listener.Demux(rawBuf[:40], 20) + if err == nil { + t.Fatal("expected error from stale FIN,ACK demux") + } + + // Encapsulate should produce RST. + // Pre-fill IP version (StackIP normally does this before calling children). + var outBuf [256]byte + outBuf[0] = 0x45 + n, err := listener.Encapsulate(outBuf[:], 0, 20) + if err != nil { + t.Fatal("encapsulate RST:", err) + } else if n == 0 { + t.Fatal("no RST produced for stale packet") + } + + tfrm, err := tcp.NewFrame(outBuf[20 : 20+n]) + if err != nil { + t.Fatal("parse RST frame:", err) + } + _, flags := tfrm.OffsetAndFlags() + if flags != tcp.FlagRST { + t.Errorf("RST flags: got %s, want [RST]", flags) + } + if tfrm.Seq() != 200 { + t.Errorf("RST SEQ: got %d, want 200 (stale packet's ACK)", tfrm.Seq()) + } + if tfrm.SourcePort() != serverPort { + t.Errorf("RST source port: got %d, want %d", tfrm.SourcePort(), serverPort) + } + if tfrm.DestinationPort() != 1337 { + t.Errorf("RST dest port: got %d, want 1337", tfrm.DestinationPort()) + } +} + +func TestStackPorts_RSTOnUnknownPort(t *testing.T) { + // Test Scenario A: SYN to a port with no listener (e.g. HTTPS port 443). + // Test at StackPorts level directly to avoid StackIP CRC validation. + var sp StackPorts + var listener tcp.Listener + pool := newMockTCPPool(1, 3, 2048) + if err := sp.ResetTCP(4); err != nil { + t.Fatal(err) + } + if err := listener.Reset(80, pool); err != nil { + t.Fatal(err) + } + if err := sp.Register(&listener); err != nil { + t.Fatal(err) + } + + // Build a SYN to port 443 (no listener). + // StackPorts.Demux expects carrier data starting at IP header with TCP at offset. + clientIP := [4]byte{10, 0, 0, 1} + serverIP := [4]byte{10, 0, 0, 2} + rawBuf := make([]byte, 256) + rawBuf[0] = 0x45 // version=4, IHL=5 + rawBuf[9] = 6 // protocol=TCP + copy(rawBuf[12:16], clientIP[:]) + copy(rawBuf[16:20], serverIP[:]) + // TCP header at offset 20. + binary.BigEndian.PutUint16(rawBuf[20:], 5000) // src port + binary.BigEndian.PutUint16(rawBuf[22:], 443) // dst port (no listener!) + binary.BigEndian.PutUint32(rawBuf[24:], 700) // SEQ = 700 + binary.BigEndian.PutUint32(rawBuf[28:], 0) // ACK = 0 + rawBuf[32] = 0x50 // offset=5 + rawBuf[33] = 0x02 // flags = SYN + + err := sp.Demux(rawBuf[:40], 20) + if err == nil { + t.Fatal("expected error for SYN to unknown port") + } + + // Pre-fill IP version (StackIP normally does this before calling children). + var outBuf [256]byte + outBuf[0] = 0x45 + n, err := sp.Encapsulate(outBuf[:], 0, 20) + if err != nil { + t.Fatal("encapsulate RST:", err) + } else if n == 0 { + t.Fatal("no RST produced for SYN to unknown port") + } + + tfrm, err := tcp.NewFrame(outBuf[20 : 20+n]) + if err != nil { + t.Fatal("parse RST frame:", err) + } + _, flags := tfrm.OffsetAndFlags() + wantFlags := tcp.FlagRST | tcp.FlagACK + if flags != wantFlags { + t.Errorf("RST flags: got %s, want %s", flags, wantFlags) + } + if tfrm.SourcePort() != 443 { + t.Errorf("RST source port: got %d, want 443", tfrm.SourcePort()) + } + if tfrm.DestinationPort() != 5000 { + t.Errorf("RST dest port: got %d, want 5000", tfrm.DestinationPort()) + } + if tfrm.Seq() != 0 { + t.Errorf("RST SEQ: got %d, want 0", tfrm.Seq()) + } + if tfrm.Ack() != 701 { + t.Errorf("RST ACK: got %d, want 701 (SEG.SEQ+1)", tfrm.Ack()) + } +} + +func TestListener_ECN_SYN(t *testing.T) { + // Test that Listener.Demux accepts SYN+ECE+CWR (ECN negotiation per RFC 3168). + // Currently fails: strict flags != FlagSYN check rejects ECN SYNs. + var listener tcp.Listener + pool := newMockTCPPool(1, 3, 2048) + serverPort := uint16(80) + if err := listener.Reset(serverPort, pool); err != nil { + t.Fatal(err) + } + + // Build SYN+ECE+CWR packet targeting port 80. + clientIP := [4]byte{10, 0, 0, 1} + serverIP := [4]byte{10, 0, 0, 2} + rawBuf := make([]byte, 256) + rawBuf[0] = 0x45 // version=4, IHL=5 + rawBuf[9] = 6 // protocol=TCP + copy(rawBuf[12:16], clientIP[:]) + copy(rawBuf[16:20], serverIP[:]) + // TCP header at offset 20. + binary.BigEndian.PutUint16(rawBuf[20:], 5000) // src port + binary.BigEndian.PutUint16(rawBuf[22:], serverPort) + binary.BigEndian.PutUint32(rawBuf[24:], 300) // SEQ + binary.BigEndian.PutUint32(rawBuf[28:], 0) // ACK + rawBuf[32] = 0x50 // offset=5 + rawBuf[33] = byte(tcp.FlagSYN | tcp.FlagECE | tcp.FlagCWR) // SYN+ECE+CWR + + // Should be accepted (create connection), not dropped. + err := listener.Demux(rawBuf[:40], 20) + if err != nil { + t.Errorf("SYN+ECE+CWR was rejected: %v (want accepted as valid SYN)", err) + } +} + +func TestStackPorts_ECN_SYN_RST(t *testing.T) { + // Test that StackPorts queues RST for SYN+ECE+CWR to an unknown port. + // Currently fails: strict flags == flagSYN check doesn't recognize ECN SYN. + var sp StackPorts + var listener tcp.Listener + pool := newMockTCPPool(1, 3, 2048) + if err := sp.ResetTCP(4); err != nil { + t.Fatal(err) + } + if err := listener.Reset(80, pool); err != nil { + t.Fatal(err) + } + if err := sp.Register(&listener); err != nil { + t.Fatal(err) + } + + // Build SYN+ECE+CWR to port 443 (no listener). + clientIP := [4]byte{10, 0, 0, 1} + serverIP := [4]byte{10, 0, 0, 2} + rawBuf := make([]byte, 256) + rawBuf[0] = 0x45 + rawBuf[9] = 6 + copy(rawBuf[12:16], clientIP[:]) + copy(rawBuf[16:20], serverIP[:]) + binary.BigEndian.PutUint16(rawBuf[20:], 5000) // src port + binary.BigEndian.PutUint16(rawBuf[22:], 443) // dst port (no listener) + binary.BigEndian.PutUint32(rawBuf[24:], 700) // SEQ + binary.BigEndian.PutUint32(rawBuf[28:], 0) // ACK + rawBuf[32] = 0x50 + rawBuf[33] = byte(tcp.FlagSYN | tcp.FlagECE | tcp.FlagCWR) + + err := sp.Demux(rawBuf[:40], 20) + if err == nil { + t.Fatal("expected error for SYN to unknown port") + } + + // Should have queued RST (same as bare SYN to unknown port). + var outBuf [256]byte + outBuf[0] = 0x45 + n, err := sp.Encapsulate(outBuf[:], 0, 20) + if err != nil { + t.Fatal("encapsulate RST:", err) + } else if n == 0 { + t.Error("no RST produced for SYN+ECE+CWR to unknown port (want RST,ACK)") + } +} + // tryExchange attempts an exchange but doesn't fail if no data to send. func tryExchange(t *testing.T, from, to *StackIP, buf []byte) { t.Helper() diff --git a/tcp/control.go b/tcp/control.go index c1e9164..72183d4 100644 --- a/tcp/control.go +++ b/tcp/control.go @@ -122,6 +122,7 @@ type sendSpace struct { UNA Value // send unacknowledged. Seqs equal to UNA and above have NOT been acked by remote. Corresponds to local data. NXT Value // send next. This seq and up to UNA+WND-1 are allowed to be sent. Corresponds to local data. WND Size // send window defined by remote. Permitted number of local unacked octets in flight. + MSS Size // maximum segment size advertised by remote peer. 0 means not set. // WL1 Value // segment sequence number used for last window update // WL2 Value // segment acknowledgment number used for last window update } @@ -203,6 +204,10 @@ func (tcb *ControlBlock) PendingSegment(payloadLen int) (_ Segment, ok bool) { } payloadLen = int(maxPayload) } + // Cap by remote MSS. + if tcb.snd.MSS > 0 && payloadLen > int(tcb.snd.MSS) { + payloadLen = int(tcb.snd.MSS) + } if payloadLen > 0 { pending |= FlagPSH // By default ensure all data flushed to destination application immediately on receive. } diff --git a/tcp/definitions.go b/tcp/definitions.go index a17c8f2..84a95fa 100644 --- a/tcp/definitions.go +++ b/tcp/definitions.go @@ -5,7 +5,6 @@ import ( "fmt" "math/bits" "strconv" - "strings" "unsafe" ) @@ -332,143 +331,3 @@ func (s State) isOpen() bool { func (s State) hasIRS() bool { return s.isOpen() && s != StateSynSent && s != StateListen } - -type OptionKind uint8 - -const ( - OptEnd OptionKind = iota // end of option list - OptNop // no-operation - OptMaxSegmentSize // maximum segment size - OptWindowScale // window scale - OptSACKPermitted // SACK permitted - OptSACK // SACK - OptEcho // echo(obsolete) - optEchoReply // echo reply(obsolete) - OptTimestamps // timestamps - optPOCP // partial order connection permitted(obsolete) - optPOSP // partial order service profile(obsolete) - optCC // CC(obsolete) - optCCnew // CC.new(obsolete) - optCCecho // CC.echo(obsolete) - optACR // alternate checksum request(obsolete) - optACD // alternate checksum data(obsolete) - optSkeeter // skeeter - optBubba // bubba - OptTrailerChecksum // trailer checksum - optMD5Signature // MD5 signature(obsolete) - OptSCPSCapabilities // SCPS capabilities - OptSNA // selective negative acks - OptRecordBoundaries // record boundaries - OptCorruptionExperienced // corruption experienced - OptSNAP // SNAP - OptUnassigned // unassigned - OptCompressionFilter // compression filter - OptQuickStartResponse // quick-start response - OptUserTimeout // user timeout or unauthorized use - OptAuthetication // Authentication TCP-AO - OptMultipath // multipath TCP -) - -const ( - OptFastOpenCookie OptionKind = 34 // fast open cookie - OptEncryptionNegotiation OptionKind = 69 // encryption negotiation - OptAccurateECN0 OptionKind = 172 // accurate ECN order 0 - OptAccurateECN1 OptionKind = 174 // accurate ECN order 1 -) - -// IsObsolete returns true if option considered obsolete by newer TCP specifications. -func (kind OptionKind) IsObsolete() bool { - if kind.IsDefined() { - return strings.HasSuffix(kind.String(), "(obsolete)") - } - return false -} - -// IsDefined returns true if the option is a known unreserved option kind. -func (kind OptionKind) IsDefined() bool { - return kind <= 30 || kind == 34 || kind == 69 || kind == 172 || kind == 174 -} - -type OptionCodec struct { - Flags OptionFlags -} - -type OptionFlags uint8 - -const ( - OptFlagSkipSizeValidation OptionFlags = 1 << iota - OptFlagSkipObsolete -) - -func (flags OptionFlags) HasAny(ofTheseFlags OptionFlags) bool { - return flags&ofTheseFlags != 0 -} - -func (op OptionCodec) PutOption16(dst []byte, kind OptionKind, v uint16) (int, error) { - return op.PutOption(dst, kind, byte(v>>8), byte(v)) -} - -func (op OptionCodec) PutOption32(dst []byte, kind OptionKind, v uint32) (int, error) { - return op.PutOption(dst, kind, byte(v>>24), byte(v>>16), byte(v>>7), byte(v)) -} - -func (op OptionCodec) PutOption(dst []byte, kind OptionKind, data ...byte) (int, error) { - putSize := 2 + len(data) - if len(dst) < putSize { - return -1, errBufferTooSmall - } else if putSize > 255 { - return -1, errors.New("option data too large") - } else if kind == OptNop || kind == OptEnd { - return -1, errors.New("cant put Nop or End option type") - } - dst[0] = byte(kind) - dst[1] = byte(putSize) - copy(dst[2:], data) - return putSize, nil -} - -func (op OptionCodec) ForEachOption(opts []byte, fn func(OptionKind, []byte) error) error { - off := 0 - skipSizeValidation := op.Flags.HasAny(OptFlagSkipSizeValidation) - skipObsolete := op.Flags.HasAny(OptFlagSkipObsolete) - for off < len(opts) && opts[off] != 0 { - kind := OptionKind(opts[off]) - off++ - if kind == OptNop { - continue - } - if len(opts[off:]) < 2 { - return errors.New("short TCP options") - } - size := int(opts[off]) - off++ - if len(opts[off:]) < size { - return fmt.Errorf("option %q length %d exceeds buffer size %d", kind.String(), size, len(opts[off:])) - } - - if !skipSizeValidation { - expectSize := -1 - switch kind { - case OptTimestamps: - expectSize = 10 - case OptMaxSegmentSize, OptUserTimeout: - expectSize = 4 - case OptWindowScale: - expectSize = 3 - case OptSACKPermitted: - expectSize = 2 - } - if expectSize != -1 && size != expectSize { - return fmt.Errorf("bad TCP option %q size want %d got %d", kind.String(), expectSize, opts[off]) - } - } - if skipObsolete && kind.IsObsolete() { - err := fn(kind, opts[off:off+size]) - if err != nil { - return err - } - } - off += size - } - return nil -} diff --git a/tcp/handler.go b/tcp/handler.go index 7121c37..30299e9 100644 --- a/tcp/handler.go +++ b/tcp/handler.go @@ -176,7 +176,7 @@ func (h *Handler) Recv(incomingPacket []byte) error { if err != nil { if h.scb.State() == StateClosed { // TODO(soypat): Should return EOF/ErrClosed? - err = err // Connection closed by reset. + err = net.ErrClosed //err // Connection closed by reset. } return err } @@ -198,10 +198,22 @@ func (h *Handler) Recv(incomingPacket []byte) error { // Update TX ring buffer to free up acked data. h.bufTx.RecvACK(segIncoming.ACK) } - if segIncoming.Flags.HasAny(FlagSYN) && h.remotePort == 0 { - // Remote reached out and has given us their port, set it on our side. - h.debug("tcp.Handler:rx-remoteport-set", slog.Uint64("port", uint64(h.localPort)), slog.Uint64("remoteport", uint64(remotePort))) - h.remotePort = remotePort + if segIncoming.Flags.HasAny(FlagSYN) { + // Parse remote MSS from TCP options. + h.optcodec.ForEachOption(tfrm.Options(), func(kind OptionKind, data []byte) error { + if kind == OptMaxSegmentSize && len(data) == 2 { + mss := uint16(data[0])<<8 | uint16(data[1]) + if mss > 0 { + h.scb.snd.MSS = Size(mss) + } + } + return nil + }) + if h.remotePort == 0 { + // Remote reached out and has given us their port, set it on our side. + h.debug("tcp.Handler:rx-remoteport-set", slog.Uint64("port", uint64(h.localPort)), slog.Uint64("remoteport", uint64(remotePort))) + h.remotePort = remotePort + } } if h.logenabled(internal.LevelTrace) { h.trace("tcp.Handler:rx-done", slog.Uint64("port", uint64(h.localPort)), slog.Uint64("remoteport", uint64(remotePort)), slog.String("seg", segIncoming.String())) @@ -343,16 +355,19 @@ func (h *Handler) Read(b []byte) (n int, err error) { // is still 0 after we've Read() data from the buffer. // // Per RFC 9293 §3.8.6.2.2 (SWS avoidance), the window is updated when freed -// space >= min(bufferSize/2, MSS). Since we don't track MSS, we use bufferSize/2. -// Zero-window openings always trigger an update. +// space >= min(bufferSize/2, MSS). This applies uniformly including zero-window +// recovery — the remote uses zero-window probes until enough space opens. func (h *Handler) maybeQueueWindowUpdate() { currentFree := Size(h.bufRx.Free()) lastAdvertised := h.scb.RecvWindow() if currentFree <= lastAdvertised { return // Window hasn't grown. } - bufSize := Size(h.bufRx.Size()) - if lastAdvertised == 0 || currentFree-lastAdvertised >= bufSize/2 { + thresh := Size(h.bufRx.Size()) / 2 + if mss := h.scb.snd.MSS; mss > 0 && mss < thresh { + thresh = mss + } + if currentFree-lastAdvertised >= thresh { h.scb.pending[0] |= FlagACK } } diff --git a/tcp/handler_test.go b/tcp/handler_test.go index 0dc0d35..3910fc6 100644 --- a/tcp/handler_test.go +++ b/tcp/handler_test.go @@ -141,6 +141,87 @@ func establish(t *testing.T, client, server *Handler, packetBuf []byte) { } } +// TestHandler_MSSHonored verifies that the server respects the MSS option from +// the client's SYN when sending data segments. The client uses a small packet +// buffer for its SYN (advertising MSS=100), and the server should not send +// segments with more than 100 bytes of payload. +func TestHandler_MSSHonored(t *testing.T) { + const mtu = 1500 + rng := rand.New(rand.NewSource(0)) + client, server := newHandler(t, mtu, 3), newHandler(t, mtu, 3) + setupClientServer(t, rng, client, server) + + // Use a 120-byte buffer for client SYN so MSS option = 120 - 20 = 100. + var smallBuf [120]byte + var largeBuf [mtu]byte + + // Client sends SYN (MSS=100 in TCP options). + n, err := client.Send(smallBuf[:]) + if err != nil { + t.Fatal("client SYN:", err) + } + err = server.Recv(smallBuf[:n]) + if err != nil { + t.Fatal("server recv SYN:", err) + } + + // Server sends SYN-ACK. + clear(largeBuf[:]) + n, err = server.Send(largeBuf[:]) + if err != nil { + t.Fatal("server SYN-ACK:", err) + } + err = client.Recv(largeBuf[:n]) + if err != nil { + t.Fatal("client recv SYN-ACK:", err) + } + + // Client sends ACK. + clear(largeBuf[:]) + n, err = client.Send(largeBuf[:]) + if err != nil { + t.Fatal("client ACK:", err) + } + err = server.Recv(largeBuf[:n]) + if err != nil { + t.Fatal("server recv ACK:", err) + } + if server.State() != StateEstablished { + t.Fatal("server not established:", server.State()) + } + + // Write 200 bytes to server's TX buffer. + data := make([]byte, 200) + for i := range data { + data[i] = byte(i) + } + nw, err := server.Write(data) + if err != nil { + t.Fatal("server write:", err) + } else if nw != 200 { + t.Fatal("server write short:", nw) + } + + // Server sends data — should be capped at client's MSS (100). + clear(largeBuf[:]) + n, err = server.Send(largeBuf[:]) + if err != nil { + t.Fatal("server send data:", err) + } else if n == 0 { + t.Fatal("server sent nothing") + } + + tfrm, err := NewFrame(largeBuf[:n]) + if err != nil { + t.Fatal("parse server frame:", err) + } + payload := tfrm.Payload() + const clientMSS = 100 + if len(payload) > clientMSS { + t.Errorf("server sent %d bytes payload, want <= %d (client MSS)", len(payload), clientMSS) + } +} + func clear[E any, T []E](s T) { var zero E for i := range s { diff --git a/tcp/listener.go b/tcp/listener.go index 5548613..3e9117d 100644 --- a/tcp/listener.go +++ b/tcp/listener.go @@ -28,17 +28,9 @@ type Listener struct { poolGet func() (*Conn, any, Value) poolReturn func(*Conn) logger - // rstQueue stores pending RST responses for SYNs rejected due to pool exhaustion. - // Per RFC 9293 §3.5.3: RST.SEQ=0, RST.ACK=SEG.SEQ+1, flags=RST|ACK. - rstQueue [4]rstEntry - rstQueueLen uint8 -} - -// rstEntry holds the minimum state needed to construct a stateless RST response. -type rstEntry struct { - remoteAddr [4]byte // IPv4 remote address. - remotePort uint16 - ackNum Value // SEG.SEQ + 1. + // rstQueue stores pending RST responses for rejected segments. + // Per RFC 9293 §3.10.7.1 (CLOSED state processing). + rstQueue RSTQueue } type handler struct { @@ -182,24 +174,8 @@ func (listener *Listener) Encapsulate(carrierData []byte, offsetToIP, offsetToFr } } // Drain one RST entry if no connection data was sent. Lower priority than connection traffic. - if n == 0 && listener.rstQueueLen > 0 && offsetToIP >= 0 { - listener.rstQueueLen-- - entry := &listener.rstQueue[listener.rstQueueLen] - tfrm, err := NewFrame(carrierData[offsetToFrame:]) - if err == nil { - tfrm.SetSourcePort(listener.port) - tfrm.SetDestinationPort(entry.remotePort) - tfrm.SetSegment(Segment{ - SEQ: 0, - ACK: entry.ackNum, - Flags: FlagRST | FlagACK, - }, 5) - tfrm.SetUrgentPtr(0) - err = internal.SetIPAddrs(carrierData[offsetToIP:offsetToFrame], 0, nil, entry.remoteAddr[:]) - if err == nil { - return sizeHeaderTCP, nil - } - } + if n == 0 { + n, _ = listener.rstQueue.Drain(carrierData, offsetToIP, offsetToFrame) } if n == 0 { listener.maintainConns() @@ -242,19 +218,17 @@ func (listener *Listener) Demux(carrierData []byte, tcpFrameOffset int) error { // Connection not in ready nor accepted. _, flags := tfrm.OffsetAndFlags() - if flags != FlagSYN { - return lneto.ErrPacketDrop // Not a synchronizing packet, drop it. + if !flags.HasAll(FlagSYN) || flags.HasAny(FlagACK) { + // RFC 9293 §3.10.7.1: CLOSED state — send RST for non-RST segments. + if !flags.HasAny(FlagRST) && flags.HasAny(FlagACK) { + listener.rstQueue.Queue(srcaddr, src, listener.port, tfrm.Ack(), 0, FlagRST) + } + return lneto.ErrPacketDrop } conn, userData, iss := listener.poolGet() if conn == nil { slog.Error("tcpListener:no-free-conn") - if len(srcaddr) == 4 && listener.rstQueueLen < uint8(len(listener.rstQueue)) { - entry := &listener.rstQueue[listener.rstQueueLen] - entry.remotePort = src - entry.ackNum = tfrm.Seq() + 1 - copy(entry.remoteAddr[:], srcaddr) - listener.rstQueueLen++ - } + listener.rstQueue.Queue(srcaddr, src, listener.port, 0, tfrm.Seq()+1, FlagRST|FlagACK) return lneto.ErrPacketDrop } err = conn.OpenListen(dst, iss) diff --git a/tcp/options.go b/tcp/options.go new file mode 100644 index 0000000..bc6f7dd --- /dev/null +++ b/tcp/options.go @@ -0,0 +1,148 @@ +package tcp + +import ( + "errors" + "fmt" + "strings" +) + +type OptionKind uint8 + +const ( + OptEnd OptionKind = iota // end of option list + OptNop // no-operation + OptMaxSegmentSize // maximum segment size + OptWindowScale // window scale + OptSACKPermitted // SACK permitted + OptSACK // SACK + OptEcho // echo(obsolete) + optEchoReply // echo reply(obsolete) + OptTimestamps // timestamps + optPOCP // partial order connection permitted(obsolete) + optPOSP // partial order service profile(obsolete) + optCC // CC(obsolete) + optCCnew // CC.new(obsolete) + optCCecho // CC.echo(obsolete) + optACR // alternate checksum request(obsolete) + optACD // alternate checksum data(obsolete) + optSkeeter // skeeter + optBubba // bubba + OptTrailerChecksum // trailer checksum + optMD5Signature // MD5 signature(obsolete) + OptSCPSCapabilities // SCPS capabilities + OptSNA // selective negative acks + OptRecordBoundaries // record boundaries + OptCorruptionExperienced // corruption experienced + OptSNAP // SNAP + OptUnassigned // unassigned + OptCompressionFilter // compression filter + OptQuickStartResponse // quick-start response + OptUserTimeout // user timeout or unauthorized use + OptAuthetication // Authentication TCP-AO + OptMultipath // multipath TCP +) + +const ( + OptFastOpenCookie OptionKind = 34 // fast open cookie + OptEncryptionNegotiation OptionKind = 69 // encryption negotiation + OptAccurateECN0 OptionKind = 172 // accurate ECN order 0 + OptAccurateECN1 OptionKind = 174 // accurate ECN order 1 +) + +// IsObsolete returns true if option considered obsolete by newer TCP specifications. +func (kind OptionKind) IsObsolete() bool { + if kind.IsDefined() { + return strings.HasSuffix(kind.String(), "(obsolete)") + } + return false +} + +// IsDefined returns true if the option is a known unreserved option kind. +func (kind OptionKind) IsDefined() bool { + return kind <= 30 || kind == 34 || kind == 69 || kind == 172 || kind == 174 +} + +type OptionCodec struct { + Flags OptionFlags +} + +type OptionFlags uint8 + +const ( + OptFlagSkipSizeValidation OptionFlags = 1 << iota + OptFlagSkipObsolete +) + +func (flags OptionFlags) HasAny(ofTheseFlags OptionFlags) bool { + return flags&ofTheseFlags != 0 +} + +func (op OptionCodec) PutOption16(dst []byte, kind OptionKind, v uint16) (int, error) { + return op.PutOption(dst, kind, byte(v>>8), byte(v)) +} + +func (op OptionCodec) PutOption32(dst []byte, kind OptionKind, v uint32) (int, error) { + return op.PutOption(dst, kind, byte(v>>24), byte(v>>16), byte(v>>7), byte(v)) +} + +func (op OptionCodec) PutOption(dst []byte, kind OptionKind, data ...byte) (int, error) { + putSize := 2 + len(data) + if len(dst) < putSize { + return -1, errBufferTooSmall + } else if putSize > 255 { + return -1, errors.New("option data too large") + } else if kind == OptNop || kind == OptEnd { + return -1, errors.New("cant put Nop or End option type") + } + dst[0] = byte(kind) + dst[1] = byte(putSize) + copy(dst[2:], data) + return putSize, nil +} + +func (op OptionCodec) ForEachOption(opts []byte, fn func(OptionKind, []byte) error) error { + off := 0 + skipSizeValidation := op.Flags.HasAny(OptFlagSkipSizeValidation) + skipObsolete := op.Flags.HasAny(OptFlagSkipObsolete) + for off < len(opts) && opts[off] != 0 { + kind := OptionKind(opts[off]) + off++ + if kind == OptNop { + continue + } + if len(opts[off:]) < 1 { + return errors.New("short TCP options") + } + size := int(opts[off]) // Total option length including kind and length bytes. + off++ + dataLen := size - 2 // Data bytes after kind and length. + if dataLen < 0 || len(opts[off:]) < dataLen { + return fmt.Errorf("option %q length %d exceeds buffer size %d", kind.String(), size, len(opts[off:])) + } + + if !skipSizeValidation { + expectSize := -1 + switch kind { + case OptTimestamps: + expectSize = 10 + case OptMaxSegmentSize, OptUserTimeout: + expectSize = 4 + case OptWindowScale: + expectSize = 3 + case OptSACKPermitted: + expectSize = 2 + } + if expectSize != -1 && size != expectSize { + return fmt.Errorf("bad TCP option %q size want %d got %d", kind.String(), expectSize, size) + } + } + if !(skipObsolete && kind.IsObsolete()) { + err := fn(kind, opts[off:off+dataLen]) + if err != nil { + return err + } + } + off += dataLen + } + return nil +} diff --git a/tcp/rst.go b/tcp/rst.go new file mode 100644 index 0000000..afd327c --- /dev/null +++ b/tcp/rst.go @@ -0,0 +1,63 @@ +package tcp + +import "github.com/soypat/lneto/internal" + +// RSTQueue is a small fixed-size queue of pending stateless RST responses. +// It is not safe for concurrent use; callers must synchronize access. +type RSTQueue struct { + buf [4]rstEntry + len uint8 +} + +type rstEntry struct { + remoteAddr [4]byte + remotePort uint16 + localPort uint16 + seq Value + ack Value + flags Flags +} + +// Queue enqueues a RST response. Silently drops if srcaddr is not IPv4 or queue is full. +func (q *RSTQueue) Queue(srcaddr []byte, remotePort, localPort uint16, seq, ack Value, flags Flags) { + if len(srcaddr) == 4 && q.len < uint8(len(q.buf)) { + entry := &q.buf[q.len] + copy(entry.remoteAddr[:], srcaddr) + entry.remotePort = remotePort + entry.localPort = localPort + entry.seq = seq + entry.ack = ack + entry.flags = flags + q.len++ + } +} + +// Pending returns the number of queued RST entries. +func (q *RSTQueue) Pending() int { return int(q.len) } + +// Drain writes one pending RST to the carrier buffer and returns the TCP frame length written. +// Returns (0, nil) if the queue is empty or offsetToIP < 0. +func (q *RSTQueue) Drain(carrierData []byte, offsetToIP, offsetToFrame int) (int, error) { + if q.len == 0 || offsetToIP < 0 { + return 0, nil + } + q.len-- + entry := &q.buf[q.len] + tfrm, err := NewFrame(carrierData[offsetToFrame:]) + if err != nil { + return 0, nil + } + tfrm.SetSourcePort(entry.localPort) + tfrm.SetDestinationPort(entry.remotePort) + tfrm.SetSegment(Segment{ + SEQ: entry.seq, + ACK: entry.ack, + Flags: entry.flags, + }, 5) + tfrm.SetUrgentPtr(0) + err = internal.SetIPAddrs(carrierData[offsetToIP:offsetToFrame], 0, nil, entry.remoteAddr[:]) + if err != nil { + return 0, nil + } + return sizeHeaderTCP, nil +}