diff --git a/pkg/tcpip/stack/nic.go b/pkg/tcpip/stack/nic.go index eb335e121..21fe24591 100644 --- a/pkg/tcpip/stack/nic.go +++ b/pkg/tcpip/stack/nic.go @@ -1100,9 +1100,9 @@ func (n *nic) multicastForwarding(protocol tcpip.NetworkProtocolNumber) (bool, t return ep.MulticastForwarding(), nil } -// ExperimentIPOptionEnabled returns whether the NIC is responsible for +// GetExperimentIPOptionEnabled returns whether the NIC is responsible for // passing the experiment IP option. -func (n *nic) ExperimentIPOptionEnabled() bool { +func (n *nic) GetExperimentIPOptionEnabled() bool { return n.experimentIPOptionEnabled } diff --git a/pkg/tcpip/transport/internal/network/endpoint.go b/pkg/tcpip/transport/internal/network/endpoint.go index a72d57592..6a0523963 100644 --- a/pkg/tcpip/transport/internal/network/endpoint.go +++ b/pkg/tcpip/transport/internal/network/endpoint.go @@ -311,7 +311,7 @@ func (c *WriteContext) newPacketBufferLocked(reserveHdrBytes int, data buffer.Bu // https://github.com/torvalds/linux/blob/38d741cb70b/include/net/sock.h#L2519 // https://github.com/torvalds/linux/blob/38d741cb70b/net/core/sock.c#L2588 var expOptVal uint16 - if nic, err := c.e.stack.GetNICByID(c.route.OutgoingNIC()); err == nil && nic.ExperimentIPOptionEnabled() { + if nic, err := c.e.stack.GetNICByID(c.route.OutgoingNIC()); err == nil && nic.GetExperimentIPOptionEnabled() { expOptVal = c.e.ops.GetExperimentOptionValue() } if c.route.NetProto() == header.IPv6ProtocolNumber && expOptVal != 0 { @@ -352,7 +352,7 @@ func (c *WriteContext) WritePacket(pkt *stack.PacketBuffer, headerIncluded bool) } var expOptVal uint16 - if nic, err := c.e.stack.GetNICByID(c.route.OutgoingNIC()); err == nil && nic.ExperimentIPOptionEnabled() { + if nic, err := c.e.stack.GetNICByID(c.route.OutgoingNIC()); err == nil && nic.GetExperimentIPOptionEnabled() { expOptVal = c.e.ops.GetExperimentOptionValue() } diff --git a/pkg/tcpip/transport/tcp/endpoint.go b/pkg/tcpip/transport/tcp/endpoint.go index fa5ee2762..f92ba1e90 100644 --- a/pkg/tcpip/transport/tcp/endpoint.go +++ b/pkg/tcpip/transport/tcp/endpoint.go @@ -3319,7 +3319,7 @@ func (e *Endpoint) GetAcceptConn() bool { // getExperimentOptionValue returns the experiment option value set on the // endpoint if experiment IP options are enabled on outgoing NIC of the route. func (e *Endpoint) getExperimentOptionValue(route *stack.Route) uint16 { - if nic, err := e.stack.GetNICByID(route.OutgoingNIC()); err == nil && nic.ExperimentIPOptionEnabled() { + if nic, err := e.stack.GetNICByID(route.OutgoingNIC()); err == nil && nic.GetExperimentIPOptionEnabled() { return e.ops.GetExperimentOptionValue() } return 0 diff --git a/pkg/tcpip/transport/tcp/forwarder.go b/pkg/tcpip/transport/tcp/forwarder.go index 39a522156..f6fdf1439 100644 --- a/pkg/tcpip/transport/tcp/forwarder.go +++ b/pkg/tcpip/transport/tcp/forwarder.go @@ -15,6 +15,9 @@ package tcp import ( + "fmt" + + "gvisor.dev/gvisor/pkg/buffer" "gvisor.dev/gvisor/pkg/sync" "gvisor.dev/gvisor/pkg/tcpip" "gvisor.dev/gvisor/pkg/tcpip/header" @@ -170,3 +173,56 @@ func (r *ForwarderRequest) CreateEndpoint(queue *waiter.Queue) (tcpip.Endpoint, return ep, nil } + +// ForwardedPacketExperimentOption returns the experiment option value from the +// forwarded packet and a bool indicating whether an experiment option value was +// found. +func (r *ForwarderRequest) ForwardedPacketExperimentOption() (uint16, bool) { + r.mu.Lock() + defer r.mu.Unlock() + + switch r.segment.pkt.NetworkProtocolNumber { + case header.IPv4ProtocolNumber: + h := header.IPv4(r.segment.pkt.NetworkHeader().Slice()) + opts := h.Options() + iter := opts.MakeIterator() + for { + opt, done, err := iter.Next() + if err != nil { + return 0, false + } + if done { + return 0, false + } + if opt.Type() == header.IPv4OptionExperimentType { + return opt.(*header.IPv4OptionExperiment).Value(), true + } + } + case header.IPv6ProtocolNumber: + h := header.IPv6(r.segment.pkt.NetworkHeader().Slice()) + v := r.segment.pkt.NetworkHeader().View() + if v != nil { + v.TrimFront(header.IPv6MinimumSize) + } + buf := buffer.MakeWithView(v) + buf.Append(r.segment.pkt.TransportHeader().View()) + dataBuf := r.segment.pkt.Data().ToBuffer() + buf.Merge(&dataBuf) + it := header.MakeIPv6PayloadIterator(header.IPv6ExtensionHeaderIdentifier(h.NextHeader()), buf) + + for { + hdr, done, err := it.Next() + if done || err != nil { + break + } + if h, ok := hdr.(header.IPv6ExperimentExtHdr); ok { + hdr.Release() + return h.Value, true + } + hdr.Release() + } + default: + panic(fmt.Sprintf("Unexpected network protocol number %d", r.segment.pkt.NetworkProtocolNumber)) + } + return 0, false +}