diff --git a/pkg/tcpip/link/ethernet/ethernet.go b/pkg/tcpip/link/ethernet/ethernet.go index 8211a2031..82b815677 100644 --- a/pkg/tcpip/link/ethernet/ethernet.go +++ b/pkg/tcpip/link/ethernet/ethernet.go @@ -50,6 +50,14 @@ func (e *Endpoint) LinkAddress() tcpip.LinkAddress { return header.UnspecifiedEthernetAddress } +// MTU implements stack.LinkEndpoint. +func (e *Endpoint) MTU() uint32 { + if mtu := e.Endpoint.MTU(); mtu > header.EthernetMinimumSize { + return mtu - header.EthernetMinimumSize + } + return 0 +} + // DeliverNetworkPacket implements stack.NetworkDispatcher. func (e *Endpoint) DeliverNetworkPacket(_, _ tcpip.LinkAddress, _ tcpip.NetworkProtocolNumber, pkt *stack.PacketBuffer) { hdr, ok := pkt.LinkHeader().Consume(header.EthernetMinimumSize) diff --git a/pkg/tcpip/link/ethernet/ethernet_test.go b/pkg/tcpip/link/ethernet/ethernet_test.go index 08a7f1ce1..16b183a61 100644 --- a/pkg/tcpip/link/ethernet/ethernet_test.go +++ b/pkg/tcpip/link/ethernet/ethernet_test.go @@ -15,6 +15,7 @@ package ethernet_test import ( + "fmt" "testing" "gvisor.dev/gvisor/pkg/tcpip" @@ -69,3 +70,52 @@ func TestDeliverNetworkPacket(t *testing.T) { t.Fatalf("got networkDispatcher.networkPackets = %d, want = 1", networkDispatcher.networkPackets) } } + +type testLinkEndpoint struct { + stack.LinkEndpoint + + mtu uint32 +} + +func (t *testLinkEndpoint) MTU() uint32 { + return t.mtu +} + +func TestMTU(t *testing.T) { + const maxFrameSize = 1500 + + tests := []struct { + maxFrameSize uint32 + expectedMTU uint32 + }{ + { + maxFrameSize: 0, + expectedMTU: 0, + }, + { + maxFrameSize: header.EthernetMinimumSize - 1, + expectedMTU: 0, + }, + { + maxFrameSize: header.EthernetMinimumSize, + expectedMTU: 0, + }, + { + maxFrameSize: header.EthernetMinimumSize + 1, + expectedMTU: 1, + }, + { + maxFrameSize: maxFrameSize, + expectedMTU: maxFrameSize - header.EthernetMinimumSize, + }, + } + + for _, test := range tests { + t.Run(fmt.Sprintf("MaxFrameSize=%d", test.maxFrameSize), func(t *testing.T) { + e := ethernet.New(&testLinkEndpoint{mtu: test.maxFrameSize}) + if got := e.MTU(); got != test.expectedMTU { + t.Errorf("got e.MTU() = %d, want = %d", got, test.expectedMTU) + } + }) + } +}