mirror of
https://github.com/netbirdio/gvisor.git
synced 2026-05-22 17:12:49 -07:00
Do not reject TCP SYN w/ ECN flags set.
This change does not add support for ECN flags in Netstack. It just ensures we don't reject valid SYN packets with ECN bits set. For a complete ECN implementation we would need to implement the following section https://datatracker.ietf.org/doc/html/rfc3168#section-6 as well as the IP bits to set the ECN bits accordingly. Fixes #7075 PiperOrigin-RevId: 423345617
This commit is contained in:
committed by
gVisor bot
parent
2a62f43796
commit
0b81b32c95
@@ -60,7 +60,7 @@ func (f TCPFlags) Contains(o TCPFlags) bool {
|
||||
|
||||
// String implements Stringer.String.
|
||||
func (f TCPFlags) String() string {
|
||||
flagsStr := []byte("FSRPAU")
|
||||
flagsStr := []byte("FSRPAUEC")
|
||||
for i := range flagsStr {
|
||||
if f&(1<<uint(i)) == 0 {
|
||||
flagsStr[i] = ' '
|
||||
@@ -77,6 +77,8 @@ const (
|
||||
TCPFlagPsh
|
||||
TCPFlagAck
|
||||
TCPFlagUrg
|
||||
TCPFlagEce
|
||||
TCPFlagCwr
|
||||
)
|
||||
|
||||
// Options that may be present in a TCP segment.
|
||||
@@ -152,6 +154,10 @@ type TCPSynOptions struct {
|
||||
|
||||
// SACKPermitted is true if the SACK option was provided in the SYN/SYN-ACK.
|
||||
SACKPermitted bool
|
||||
|
||||
// Flags if specified are set on the outgoing SYN. The SYN flag is
|
||||
// always set.
|
||||
Flags TCPFlags
|
||||
}
|
||||
|
||||
// SACKBlock represents a single contiguous SACK block.
|
||||
|
||||
@@ -152,14 +152,16 @@ func TestTCPFlags(t *testing.T) {
|
||||
flags header.TCPFlags
|
||||
want string
|
||||
}{
|
||||
{header.TCPFlagFin, "F "},
|
||||
{header.TCPFlagSyn, " S "},
|
||||
{header.TCPFlagRst, " R "},
|
||||
{header.TCPFlagPsh, " P "},
|
||||
{header.TCPFlagAck, " A "},
|
||||
{header.TCPFlagUrg, " U"},
|
||||
{header.TCPFlagSyn | header.TCPFlagAck, " S A "},
|
||||
{header.TCPFlagFin | header.TCPFlagAck, "F A "},
|
||||
{header.TCPFlagFin, "F "},
|
||||
{header.TCPFlagSyn, " S "},
|
||||
{header.TCPFlagRst, " R "},
|
||||
{header.TCPFlagPsh, " P "},
|
||||
{header.TCPFlagAck, " A "},
|
||||
{header.TCPFlagUrg, " U "},
|
||||
{header.TCPFlagEce, " E "},
|
||||
{header.TCPFlagCwr, " C"},
|
||||
{header.TCPFlagSyn | header.TCPFlagAck, " S A "},
|
||||
{header.TCPFlagFin | header.TCPFlagAck, "F A "},
|
||||
} {
|
||||
if got := tt.flags.String(); got != tt.want {
|
||||
t.Errorf("got TCPFlags(%#b).String() = %s, want = %s", tt.flags, got, tt.want)
|
||||
|
||||
@@ -83,6 +83,7 @@ go_test(
|
||||
size = "large",
|
||||
srcs = [
|
||||
"dual_stack_test.go",
|
||||
"forwarder_test.go",
|
||||
"rcv_test.go",
|
||||
"sack_scoreboard_test.go",
|
||||
"tcp_noracedetector_test.go",
|
||||
|
||||
@@ -443,7 +443,7 @@ func (e *endpoint) handleListenSegment(ctx *listenContext, s *segment) tcpip.Err
|
||||
e.stack.Stats().DroppedPackets.Increment()
|
||||
return nil
|
||||
|
||||
case s.flags == header.TCPFlagSyn:
|
||||
case s.flags.Contains(header.TCPFlagSyn):
|
||||
if e.acceptQueueIsFull() {
|
||||
e.stack.Stats().TCP.ListenOverflowSynDrop.Increment()
|
||||
e.stats.ReceiveErrors.ListenOverflowSynDrop.Increment()
|
||||
|
||||
@@ -68,8 +68,8 @@ func (f *Forwarder) HandlePacket(id stack.TransportEndpointID, pkt *stack.Packet
|
||||
s := newIncomingSegment(id, f.stack.Clock(), pkt)
|
||||
defer s.decRef()
|
||||
|
||||
// We only care about well-formed SYN packets.
|
||||
if !s.parse(pkt.RXTransportChecksumValidated) || !s.csumValid || s.flags != header.TCPFlagSyn {
|
||||
// We only care about well-formed SYN packets (not SYN-ACK) packets.
|
||||
if !s.parse(pkt.RXTransportChecksumValidated) || !s.csumValid || !s.flags.Contains(header.TCPFlagSyn) || s.flags.Contains(header.TCPFlagAck) {
|
||||
return false
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,101 @@
|
||||
// Copyright 2018 The gVisor Authors.
|
||||
//
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package tcp_test
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"gvisor.dev/gvisor/pkg/tcpip"
|
||||
"gvisor.dev/gvisor/pkg/tcpip/header"
|
||||
"gvisor.dev/gvisor/pkg/tcpip/transport/tcp"
|
||||
"gvisor.dev/gvisor/pkg/tcpip/transport/tcp/testing/context"
|
||||
)
|
||||
|
||||
func TestForwarderSendMSSLessThanMTU(t *testing.T) {
|
||||
const maxPayload = 100
|
||||
const mtu = 1200
|
||||
c := context.New(t, mtu)
|
||||
defer c.Cleanup()
|
||||
|
||||
s := c.Stack()
|
||||
ch := make(chan tcpip.Error, 1)
|
||||
f := tcp.NewForwarder(s, 65536, 10, func(r *tcp.ForwarderRequest) {
|
||||
var err tcpip.Error
|
||||
c.EP, err = r.CreateEndpoint(&c.WQ)
|
||||
ch <- err
|
||||
close(ch)
|
||||
})
|
||||
s.SetTransportProtocolHandler(tcp.ProtocolNumber, f.HandlePacket)
|
||||
|
||||
// Do 3-way handshake.
|
||||
c.PassiveConnect(maxPayload, -1, header.TCPSynOptions{MSS: mtu - header.IPv4MinimumSize - header.TCPMinimumSize})
|
||||
|
||||
// Wait for connection to be available.
|
||||
select {
|
||||
case err := <-ch:
|
||||
if err != nil {
|
||||
t.Fatalf("Error creating endpoint: %s", err)
|
||||
}
|
||||
case <-time.After(2 * time.Second):
|
||||
t.Fatalf("Timed out waiting for connection")
|
||||
}
|
||||
|
||||
// Check that data gets properly segmented.
|
||||
testBrokenUpWrite(t, c, maxPayload)
|
||||
}
|
||||
|
||||
func TestForwarderDoesNotRejectECNFlags(t *testing.T) {
|
||||
testCases := []struct {
|
||||
name string
|
||||
flags header.TCPFlags
|
||||
}{
|
||||
{name: "non-setup ECN SYN w/ ECE", flags: header.TCPFlagEce},
|
||||
{name: "non-setup ECN SYN w/ CWR", flags: header.TCPFlagCwr},
|
||||
{name: "setup ECN SYN", flags: header.TCPFlagEce | header.TCPFlagCwr},
|
||||
}
|
||||
|
||||
for _, tc := range testCases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
const maxPayload = 100
|
||||
const mtu = 1200
|
||||
c := context.New(t, mtu)
|
||||
defer c.Cleanup()
|
||||
|
||||
s := c.Stack()
|
||||
ch := make(chan tcpip.Error, 1)
|
||||
f := tcp.NewForwarder(s, 65536, 10, func(r *tcp.ForwarderRequest) {
|
||||
var err tcpip.Error
|
||||
c.EP, err = r.CreateEndpoint(&c.WQ)
|
||||
ch <- err
|
||||
close(ch)
|
||||
})
|
||||
s.SetTransportProtocolHandler(tcp.ProtocolNumber, f.HandlePacket)
|
||||
|
||||
// Do 3-way handshake.
|
||||
c.PassiveConnect(maxPayload, -1, header.TCPSynOptions{MSS: mtu - header.IPv4MinimumSize - header.TCPMinimumSize, Flags: tc.flags})
|
||||
|
||||
// Wait for connection to be available.
|
||||
select {
|
||||
case err := <-ch:
|
||||
if err != nil {
|
||||
t.Fatalf("Error creating endpoint: %s", err)
|
||||
}
|
||||
case <-time.After(2 * time.Second):
|
||||
t.Fatalf("Timed out waiting for connection")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -3631,38 +3631,6 @@ func TestSynCookiePassiveSendMSSLessThanMTU(t *testing.T) {
|
||||
testBrokenUpWrite(t, c, maxPayload)
|
||||
}
|
||||
|
||||
func TestForwarderSendMSSLessThanMTU(t *testing.T) {
|
||||
const maxPayload = 100
|
||||
const mtu = 1200
|
||||
c := context.New(t, mtu)
|
||||
defer c.Cleanup()
|
||||
|
||||
s := c.Stack()
|
||||
ch := make(chan tcpip.Error, 1)
|
||||
f := tcp.NewForwarder(s, 65536, 10, func(r *tcp.ForwarderRequest) {
|
||||
var err tcpip.Error
|
||||
c.EP, err = r.CreateEndpoint(&c.WQ)
|
||||
ch <- err
|
||||
})
|
||||
s.SetTransportProtocolHandler(tcp.ProtocolNumber, f.HandlePacket)
|
||||
|
||||
// Do 3-way handshake.
|
||||
c.PassiveConnect(maxPayload, -1, header.TCPSynOptions{MSS: mtu - header.IPv4MinimumSize - header.TCPMinimumSize})
|
||||
|
||||
// Wait for connection to be available.
|
||||
select {
|
||||
case err := <-ch:
|
||||
if err != nil {
|
||||
t.Fatalf("Error creating endpoint: %s", err)
|
||||
}
|
||||
case <-time.After(2 * time.Second):
|
||||
t.Fatalf("Timed out waiting for connection")
|
||||
}
|
||||
|
||||
// Check that data gets properly segmented.
|
||||
testBrokenUpWrite(t, c, maxPayload)
|
||||
}
|
||||
|
||||
func TestSynOptionsOnActiveConnect(t *testing.T) {
|
||||
const mtu = 1400
|
||||
c := context.New(t, mtu)
|
||||
@@ -8671,3 +8639,67 @@ func TestTimestampSynCookies(t *testing.T) {
|
||||
t.Fatalf("got TSVal = %d, want %d", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
// TestECNFlagsAccept tests that an ECN non-setup/setup SYN is accepted
|
||||
// and the connection is correctly completed.
|
||||
func TestECNFlagsAccept(t *testing.T) {
|
||||
testCases := []struct {
|
||||
name string
|
||||
flags header.TCPFlags
|
||||
}{
|
||||
{name: "non-setup ECN SYN w/ ECE", flags: header.TCPFlagEce},
|
||||
{name: "non-setup ECN SYN w/ CWR", flags: header.TCPFlagCwr},
|
||||
{name: "setup ECN SYN", flags: header.TCPFlagEce | header.TCPFlagCwr},
|
||||
}
|
||||
|
||||
for _, tc := range testCases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
|
||||
c := context.New(t, defaultMTU)
|
||||
defer c.Cleanup()
|
||||
|
||||
// Create EP and start listening.
|
||||
wq := &waiter.Queue{}
|
||||
ep, err := c.Stack().NewEndpoint(tcp.ProtocolNumber, ipv4.ProtocolNumber, wq)
|
||||
if err != nil {
|
||||
t.Fatalf("NewEndpoint failed: %s", err)
|
||||
}
|
||||
defer ep.Close()
|
||||
|
||||
if err := ep.Bind(tcpip.FullAddress{Port: context.StackPort}); err != nil {
|
||||
t.Fatalf("Bind failed: %s", err)
|
||||
}
|
||||
|
||||
if err := ep.Listen(10); err != nil {
|
||||
t.Fatalf("Listen failed: %s", err)
|
||||
}
|
||||
|
||||
// Do 3-way handshake.
|
||||
const maxPayload = 100
|
||||
|
||||
c.PassiveConnect(maxPayload, -1 /* wndScale */, header.TCPSynOptions{MSS: defaultIPv4MSS, Flags: tc.flags})
|
||||
|
||||
// Try to accept the connection.
|
||||
we, ch := waiter.NewChannelEntry(waiter.ReadableEvents)
|
||||
wq.EventRegister(&we)
|
||||
defer wq.EventUnregister(&we)
|
||||
|
||||
c.EP, _, err = ep.Accept(nil)
|
||||
if cmp.Equal(&tcpip.ErrWouldBlock{}, err) {
|
||||
// Wait for connection to be established.
|
||||
select {
|
||||
case <-ch:
|
||||
c.EP, _, err = ep.Accept(nil)
|
||||
if err != nil {
|
||||
t.Fatalf("Accept failed: %s", err)
|
||||
}
|
||||
|
||||
case <-time.After(1 * time.Second):
|
||||
t.Fatalf("Timed out waiting for accept")
|
||||
}
|
||||
} else if err != nil {
|
||||
t.Fatalf("Accept failed: %s", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1158,7 +1158,7 @@ func (c *Context) PassiveConnectWithOptions(maxPayload, wndScale int, synOptions
|
||||
c.SendPacket(nil, &Headers{
|
||||
SrcPort: TestPort,
|
||||
DstPort: StackPort,
|
||||
Flags: header.TCPFlagSyn,
|
||||
Flags: header.TCPFlagSyn | synOptions.Flags,
|
||||
SeqNum: iss,
|
||||
RcvWnd: 30000,
|
||||
TCPOpts: opts[:offset],
|
||||
|
||||
@@ -189,7 +189,7 @@ func TestLayerStringFormat(t *testing.T) {
|
||||
"SeqNum:3452155723 " +
|
||||
"AckNum:2596996163 " +
|
||||
"DataOffset:5 " +
|
||||
"Flags: R A " +
|
||||
"Flags: R A " +
|
||||
"WindowSize:64240 " +
|
||||
"Checksum:11819" +
|
||||
"}",
|
||||
|
||||
Reference in New Issue
Block a user