mirror of
https://github.com/netbirdio/gvisor.git
synced 2026-05-22 17:12:49 -07:00
Make stack.Route thread safe
Currently we rely on the user to take the lock on the endpoint that owns the route, in order to modify it safely. We can instead move `Route.RemoteLinkAddress` under `Route`'s mutex, and allow non-locking and thread-safe access to other fields of `Route`. PiperOrigin-RevId: 345461586
This commit is contained in:
committed by
gVisor bot
parent
6f60a2b0a2
commit
3ff1aef544
@@ -62,7 +62,7 @@ func (e *Endpoint) Capabilities() stack.LinkEndpointCapabilities {
|
||||
|
||||
// WritePacket implements stack.LinkEndpoint.
|
||||
func (e *Endpoint) WritePacket(r *stack.Route, gso *stack.GSO, proto tcpip.NetworkProtocolNumber, pkt *stack.PacketBuffer) *tcpip.Error {
|
||||
e.AddHeader(e.Endpoint.LinkAddress(), r.RemoteLinkAddress, proto, pkt)
|
||||
e.AddHeader(e.Endpoint.LinkAddress(), r.RemoteLinkAddress(), proto, pkt)
|
||||
return e.Endpoint.WritePacket(r, gso, proto, pkt)
|
||||
}
|
||||
|
||||
@@ -71,7 +71,7 @@ func (e *Endpoint) WritePackets(r *stack.Route, gso *stack.GSO, pkts stack.Packe
|
||||
linkAddr := e.Endpoint.LinkAddress()
|
||||
|
||||
for pkt := pkts.Front(); pkt != nil; pkt = pkt.Next() {
|
||||
e.AddHeader(linkAddr, r.RemoteLinkAddress, proto, pkt)
|
||||
e.AddHeader(linkAddr, r.RemoteLinkAddress(), proto, pkt)
|
||||
}
|
||||
|
||||
return e.Endpoint.WritePackets(r, gso, pkts, proto)
|
||||
|
||||
@@ -410,7 +410,7 @@ func (e *endpoint) AddHeader(local, remote tcpip.LinkAddress, protocol tcpip.Net
|
||||
// currently writable, the packet is dropped.
|
||||
func (e *endpoint) WritePacket(r *stack.Route, gso *stack.GSO, protocol tcpip.NetworkProtocolNumber, pkt *stack.PacketBuffer) *tcpip.Error {
|
||||
if e.hdrSize > 0 {
|
||||
e.AddHeader(r.LocalLinkAddress, r.RemoteLinkAddress, protocol, pkt)
|
||||
e.AddHeader(r.LocalLinkAddress, r.RemoteLinkAddress(), protocol, pkt)
|
||||
}
|
||||
|
||||
var builder iovec.Builder
|
||||
@@ -453,7 +453,7 @@ func (e *endpoint) sendBatch(batchFD int, batch []*stack.PacketBuffer) (int, *tc
|
||||
mmsgHdrs := make([]rawfile.MMsgHdr, 0, len(batch))
|
||||
for _, pkt := range batch {
|
||||
if e.hdrSize > 0 {
|
||||
e.AddHeader(pkt.EgressRoute.LocalLinkAddress, pkt.EgressRoute.RemoteLinkAddress, pkt.NetworkProtocolNumber, pkt)
|
||||
e.AddHeader(pkt.EgressRoute.LocalLinkAddress, pkt.EgressRoute.RemoteLinkAddress(), pkt.NetworkProtocolNumber, pkt)
|
||||
}
|
||||
|
||||
var vnetHdrBuf []byte
|
||||
|
||||
@@ -183,9 +183,8 @@ func testWritePacket(t *testing.T, plen int, eth bool, gsoMaxSize uint32, hash u
|
||||
c := newContext(t, &Options{Address: laddr, MTU: mtu, EthernetHeader: eth, GSOMaxSize: gsoMaxSize})
|
||||
defer c.cleanup()
|
||||
|
||||
r := &stack.Route{
|
||||
RemoteLinkAddress: raddr,
|
||||
}
|
||||
var r stack.Route
|
||||
r.ResolveWith(raddr)
|
||||
|
||||
// Build payload.
|
||||
payload := buffer.NewView(plen)
|
||||
@@ -220,7 +219,7 @@ func testWritePacket(t *testing.T, plen int, eth bool, gsoMaxSize uint32, hash u
|
||||
L3HdrLen: header.IPv4MaximumHeaderSize,
|
||||
}
|
||||
}
|
||||
if err := c.ep.WritePacket(r, gso, proto, pkt); err != nil {
|
||||
if err := c.ep.WritePacket(&r, gso, proto, pkt); err != nil {
|
||||
t.Fatalf("WritePacket failed: %v", err)
|
||||
}
|
||||
|
||||
@@ -325,9 +324,9 @@ func TestPreserveSrcAddress(t *testing.T) {
|
||||
|
||||
// Set LocalLinkAddress in route to the value of the bridged address.
|
||||
r := &stack.Route{
|
||||
RemoteLinkAddress: raddr,
|
||||
LocalLinkAddress: baddr,
|
||||
LocalLinkAddress: baddr,
|
||||
}
|
||||
r.ResolveWith(raddr)
|
||||
|
||||
pkt := stack.NewPacketBuffer(stack.PacketBufferOptions{
|
||||
// WritePacket panics given a prependable with anything less than
|
||||
|
||||
@@ -36,14 +36,14 @@ func New(lower stack.LinkEndpoint) stack.LinkEndpoint {
|
||||
|
||||
// WritePacket implements stack.LinkEndpoint.WritePacket.
|
||||
func (e *endpoint) WritePacket(r *stack.Route, gso *stack.GSO, protocol tcpip.NetworkProtocolNumber, pkt *stack.PacketBuffer) *tcpip.Error {
|
||||
e.Endpoint.DeliverOutboundPacket(r.RemoteLinkAddress, r.LocalLinkAddress, protocol, pkt)
|
||||
e.Endpoint.DeliverOutboundPacket(r.RemoteLinkAddress(), r.LocalLinkAddress, protocol, pkt)
|
||||
return e.Endpoint.WritePacket(r, gso, protocol, pkt)
|
||||
}
|
||||
|
||||
// WritePackets implements stack.LinkEndpoint.WritePackets.
|
||||
func (e *endpoint) WritePackets(r *stack.Route, gso *stack.GSO, pkts stack.PacketBufferList, proto tcpip.NetworkProtocolNumber) (int, *tcpip.Error) {
|
||||
for pkt := pkts.Front(); pkt != nil; pkt = pkt.Next() {
|
||||
e.Endpoint.DeliverOutboundPacket(pkt.EgressRoute.RemoteLinkAddress, pkt.EgressRoute.LocalLinkAddress, pkt.NetworkProtocolNumber, pkt)
|
||||
e.Endpoint.DeliverOutboundPacket(pkt.EgressRoute.RemoteLinkAddress(), pkt.EgressRoute.LocalLinkAddress, pkt.NetworkProtocolNumber, pkt)
|
||||
}
|
||||
|
||||
return e.Endpoint.WritePackets(r, gso, pkts, proto)
|
||||
|
||||
@@ -55,7 +55,7 @@ func (e *Endpoint) WritePacket(r *stack.Route, _ *stack.GSO, proto tcpip.Network
|
||||
// remote address from the perspective of the other end of the pipe
|
||||
// (e.linked). Similarly, the remote address from the perspective of this
|
||||
// endpoint is the local address on the other end.
|
||||
e.linked.dispatcher.DeliverNetworkPacket(r.LocalLinkAddress /* remote */, r.RemoteLinkAddress /* local */, proto, stack.NewPacketBuffer(stack.PacketBufferOptions{
|
||||
e.linked.dispatcher.DeliverNetworkPacket(r.LocalLinkAddress /* remote */, r.RemoteLinkAddress() /* local */, proto, stack.NewPacketBuffer(stack.PacketBufferOptions{
|
||||
Data: buffer.NewVectorisedView(pkt.Size(), pkt.Views()),
|
||||
}))
|
||||
|
||||
|
||||
@@ -204,7 +204,7 @@ func (e *endpoint) AddHeader(local, remote tcpip.LinkAddress, protocol tcpip.Net
|
||||
// WritePacket writes outbound packets to the file descriptor. If it is not
|
||||
// currently writable, the packet is dropped.
|
||||
func (e *endpoint) WritePacket(r *stack.Route, _ *stack.GSO, protocol tcpip.NetworkProtocolNumber, pkt *stack.PacketBuffer) *tcpip.Error {
|
||||
e.AddHeader(r.LocalLinkAddress, r.RemoteLinkAddress, protocol, pkt)
|
||||
e.AddHeader(r.LocalLinkAddress, r.RemoteLinkAddress(), protocol, pkt)
|
||||
|
||||
views := pkt.Views()
|
||||
// Transmit the packet.
|
||||
|
||||
@@ -260,9 +260,8 @@ func TestSimpleSend(t *testing.T) {
|
||||
defer c.cleanup()
|
||||
|
||||
// Prepare route.
|
||||
r := stack.Route{
|
||||
RemoteLinkAddress: remoteLinkAddr,
|
||||
}
|
||||
var r stack.Route
|
||||
r.ResolveWith(remoteLinkAddr)
|
||||
|
||||
for iters := 1000; iters > 0; iters-- {
|
||||
func() {
|
||||
@@ -342,9 +341,9 @@ func TestPreserveSrcAddressInSend(t *testing.T) {
|
||||
newLocalLinkAddress := tcpip.LinkAddress(strings.Repeat("0xFE", 6))
|
||||
// Set both remote and local link address in route.
|
||||
r := stack.Route{
|
||||
RemoteLinkAddress: remoteLinkAddr,
|
||||
LocalLinkAddress: newLocalLinkAddress,
|
||||
LocalLinkAddress: newLocalLinkAddress,
|
||||
}
|
||||
r.ResolveWith(remoteLinkAddr)
|
||||
|
||||
pkt := stack.NewPacketBuffer(stack.PacketBufferOptions{
|
||||
// WritePacket panics given a prependable with anything less than
|
||||
@@ -395,9 +394,8 @@ func TestFillTxQueue(t *testing.T) {
|
||||
defer c.cleanup()
|
||||
|
||||
// Prepare to send a packet.
|
||||
r := stack.Route{
|
||||
RemoteLinkAddress: remoteLinkAddr,
|
||||
}
|
||||
var r stack.Route
|
||||
r.ResolveWith(remoteLinkAddr)
|
||||
|
||||
buf := buffer.NewView(100)
|
||||
|
||||
@@ -444,9 +442,8 @@ func TestFillTxQueueAfterBadCompletion(t *testing.T) {
|
||||
c.txq.rx.Flush()
|
||||
|
||||
// Prepare to send a packet.
|
||||
r := stack.Route{
|
||||
RemoteLinkAddress: remoteLinkAddr,
|
||||
}
|
||||
var r stack.Route
|
||||
r.ResolveWith(remoteLinkAddr)
|
||||
|
||||
buf := buffer.NewView(100)
|
||||
|
||||
@@ -509,9 +506,8 @@ func TestFillTxMemory(t *testing.T) {
|
||||
defer c.cleanup()
|
||||
|
||||
// Prepare to send a packet.
|
||||
r := stack.Route{
|
||||
RemoteLinkAddress: remoteLinkAddr,
|
||||
}
|
||||
var r stack.Route
|
||||
r.ResolveWith(remoteLinkAddr)
|
||||
|
||||
buf := buffer.NewView(100)
|
||||
|
||||
@@ -557,9 +553,8 @@ func TestFillTxMemoryWithMultiBuffer(t *testing.T) {
|
||||
defer c.cleanup()
|
||||
|
||||
// Prepare to send a packet.
|
||||
r := stack.Route{
|
||||
RemoteLinkAddress: remoteLinkAddr,
|
||||
}
|
||||
var r stack.Route
|
||||
r.ResolveWith(remoteLinkAddr)
|
||||
|
||||
buf := buffer.NewView(100)
|
||||
|
||||
|
||||
@@ -264,7 +264,7 @@ func (d *Device) encodePkt(info *channel.PacketInfo) (buffer.View, bool) {
|
||||
// If the packet does not already have link layer header, and the route
|
||||
// does not exist, we can't compute it. This is possibly a raw packet, tun
|
||||
// device doesn't support this at the moment.
|
||||
if info.Pkt.LinkHeader().View().IsEmpty() && info.Route.RemoteLinkAddress == "" {
|
||||
if info.Pkt.LinkHeader().View().IsEmpty() && info.Route.RemoteLinkAddress() == "" {
|
||||
return nil, false
|
||||
}
|
||||
|
||||
@@ -272,7 +272,7 @@ func (d *Device) encodePkt(info *channel.PacketInfo) (buffer.View, bool) {
|
||||
if d.hasFlags(linux.IFF_TAP) {
|
||||
// Add ethernet header if not provided.
|
||||
if info.Pkt.LinkHeader().View().IsEmpty() {
|
||||
d.endpoint.AddHeader(info.Route.LocalLinkAddress, info.Route.RemoteLinkAddress, info.Proto, info.Pkt)
|
||||
d.endpoint.AddHeader(info.Route.LocalLinkAddress, info.Route.RemoteLinkAddress(), info.Proto, info.Pkt)
|
||||
}
|
||||
vv.AppendView(info.Pkt.LinkHeader().View())
|
||||
}
|
||||
|
||||
@@ -442,9 +442,9 @@ func (*testInterface) Promiscuous() bool {
|
||||
|
||||
func (t *testInterface) WritePacketToRemote(remoteLinkAddr tcpip.LinkAddress, gso *stack.GSO, protocol tcpip.NetworkProtocolNumber, pkt *stack.PacketBuffer) *tcpip.Error {
|
||||
r := stack.Route{
|
||||
NetProto: protocol,
|
||||
RemoteLinkAddress: remoteLinkAddr,
|
||||
NetProto: protocol,
|
||||
}
|
||||
r.ResolveWith(remoteLinkAddr)
|
||||
return t.LinkEndpoint.WritePacket(&r, gso, protocol, pkt)
|
||||
}
|
||||
|
||||
@@ -557,8 +557,8 @@ func TestLinkAddressRequest(t *testing.T) {
|
||||
t.Fatal("expected to send a link address request")
|
||||
}
|
||||
|
||||
if pkt.Route.RemoteLinkAddress != test.expectedRemoteLinkAddr {
|
||||
t.Errorf("got pkt.Route.RemoteLinkAddress = %s, want = %s", pkt.Route.RemoteLinkAddress, test.expectedRemoteLinkAddr)
|
||||
if got := pkt.Route.RemoteLinkAddress(); got != test.expectedRemoteLinkAddr {
|
||||
t.Errorf("got pkt.Route.RemoteLinkAddress() = %s, want = %s", got, test.expectedRemoteLinkAddr)
|
||||
}
|
||||
|
||||
rep := header.ARP(stack.PayloadSince(pkt.Pkt.NetworkHeader()))
|
||||
|
||||
@@ -2770,8 +2770,8 @@ func TestPacketQueing(t *testing.T) {
|
||||
if p.Proto != header.IPv4ProtocolNumber {
|
||||
t.Errorf("got p.Proto = %d, want = %d", p.Proto, header.IPv4ProtocolNumber)
|
||||
}
|
||||
if p.Route.RemoteLinkAddress != host2NICLinkAddr {
|
||||
t.Errorf("got p.Route.RemoteLinkAddress = %s, want = %s", p.Route.RemoteLinkAddress, host2NICLinkAddr)
|
||||
if got := p.Route.RemoteLinkAddress(); got != host2NICLinkAddr {
|
||||
t.Errorf("got p.Route.RemoteLinkAddress() = %s, want = %s", got, host2NICLinkAddr)
|
||||
}
|
||||
checker.IPv4(t, stack.PayloadSince(p.Pkt.NetworkHeader()),
|
||||
checker.SrcAddr(host1IPv4Addr.AddressWithPrefix.Address),
|
||||
@@ -2813,8 +2813,8 @@ func TestPacketQueing(t *testing.T) {
|
||||
if p.Proto != header.IPv4ProtocolNumber {
|
||||
t.Errorf("got p.Proto = %d, want = %d", p.Proto, header.IPv4ProtocolNumber)
|
||||
}
|
||||
if p.Route.RemoteLinkAddress != host2NICLinkAddr {
|
||||
t.Errorf("got p.Route.RemoteLinkAddress = %s, want = %s", p.Route.RemoteLinkAddress, host2NICLinkAddr)
|
||||
if got := p.Route.RemoteLinkAddress(); got != host2NICLinkAddr {
|
||||
t.Errorf("got p.Route.RemoteLinkAddress() = %s, want = %s", got, host2NICLinkAddr)
|
||||
}
|
||||
checker.IPv4(t, stack.PayloadSince(p.Pkt.NetworkHeader()),
|
||||
checker.SrcAddr(host1IPv4Addr.AddressWithPrefix.Address),
|
||||
@@ -2862,8 +2862,8 @@ func TestPacketQueing(t *testing.T) {
|
||||
if p.Proto != arp.ProtocolNumber {
|
||||
t.Errorf("got p.Proto = %d, want = %d", p.Proto, arp.ProtocolNumber)
|
||||
}
|
||||
if p.Route.RemoteLinkAddress != header.EthernetBroadcastAddress {
|
||||
t.Errorf("got p.Route.RemoteLinkAddress = %s, want = %s", p.Route.RemoteLinkAddress, header.EthernetBroadcastAddress)
|
||||
if got := p.Route.RemoteLinkAddress(); got != header.EthernetBroadcastAddress {
|
||||
t.Errorf("got p.Route.RemoteLinkAddress() = %s, want = %s", got, header.EthernetBroadcastAddress)
|
||||
}
|
||||
rep := header.ARP(p.Pkt.NetworkHeader().View())
|
||||
if got := rep.Op(); got != header.ARPRequest {
|
||||
|
||||
@@ -150,9 +150,9 @@ func (*testInterface) Promiscuous() bool {
|
||||
|
||||
func (t *testInterface) WritePacketToRemote(remoteLinkAddr tcpip.LinkAddress, gso *stack.GSO, protocol tcpip.NetworkProtocolNumber, pkt *stack.PacketBuffer) *tcpip.Error {
|
||||
r := stack.Route{
|
||||
NetProto: protocol,
|
||||
RemoteLinkAddress: remoteLinkAddr,
|
||||
NetProto: protocol,
|
||||
}
|
||||
r.ResolveWith(remoteLinkAddr)
|
||||
return t.LinkEndpoint.WritePacket(&r, gso, protocol, pkt)
|
||||
}
|
||||
|
||||
@@ -600,8 +600,8 @@ func routeICMPv6Packet(t *testing.T, args routeArgs, fn func(*testing.T, header.
|
||||
return
|
||||
}
|
||||
|
||||
if len(args.remoteLinkAddr) != 0 && args.remoteLinkAddr != pi.Route.RemoteLinkAddress {
|
||||
t.Errorf("got remote link address = %s, want = %s", pi.Route.RemoteLinkAddress, args.remoteLinkAddr)
|
||||
if got := pi.Route.RemoteLinkAddress(); len(args.remoteLinkAddr) != 0 && got != args.remoteLinkAddr {
|
||||
t.Errorf("got remote link address = %s, want = %s", got, args.remoteLinkAddr)
|
||||
}
|
||||
|
||||
// Pull the full payload since network header. Needed for header.IPv6 to
|
||||
@@ -1381,8 +1381,8 @@ func TestLinkAddressRequest(t *testing.T) {
|
||||
if !ok {
|
||||
t.Fatal("expected to send a link address request")
|
||||
}
|
||||
if pkt.Route.RemoteLinkAddress != test.expectedRemoteLinkAddr {
|
||||
t.Errorf("got pkt.Route.RemoteLinkAddress = %s, want = %s", pkt.Route.RemoteLinkAddress, test.expectedRemoteLinkAddr)
|
||||
if got := pkt.Route.RemoteLinkAddress(); got != test.expectedRemoteLinkAddr {
|
||||
t.Errorf("got pkt.Route.RemoteLinkAddress() = %s, want = %s", got, test.expectedRemoteLinkAddr)
|
||||
}
|
||||
if pkt.Route.RemoteAddress != test.expectedRemoteAddr {
|
||||
t.Errorf("got pkt.Route.RemoteAddress = %s, want = %s", pkt.Route.RemoteAddress, test.expectedRemoteAddr)
|
||||
@@ -1463,8 +1463,8 @@ func TestPacketQueing(t *testing.T) {
|
||||
if p.Proto != ProtocolNumber {
|
||||
t.Errorf("got p.Proto = %d, want = %d", p.Proto, ProtocolNumber)
|
||||
}
|
||||
if p.Route.RemoteLinkAddress != host2NICLinkAddr {
|
||||
t.Errorf("got p.Route.RemoteLinkAddress = %s, want = %s", p.Route.RemoteLinkAddress, host2NICLinkAddr)
|
||||
if got := p.Route.RemoteLinkAddress(); got != host2NICLinkAddr {
|
||||
t.Errorf("got p.Route.RemoteLinkAddress() = %s, want = %s", got, host2NICLinkAddr)
|
||||
}
|
||||
checker.IPv6(t, stack.PayloadSince(p.Pkt.NetworkHeader()),
|
||||
checker.SrcAddr(host1IPv6Addr.AddressWithPrefix.Address),
|
||||
@@ -1505,8 +1505,8 @@ func TestPacketQueing(t *testing.T) {
|
||||
if p.Proto != ProtocolNumber {
|
||||
t.Errorf("got p.Proto = %d, want = %d", p.Proto, ProtocolNumber)
|
||||
}
|
||||
if p.Route.RemoteLinkAddress != host2NICLinkAddr {
|
||||
t.Errorf("got p.Route.RemoteLinkAddress = %s, want = %s", p.Route.RemoteLinkAddress, host2NICLinkAddr)
|
||||
if got := p.Route.RemoteLinkAddress(); got != host2NICLinkAddr {
|
||||
t.Errorf("got p.Route.RemoteLinkAddress() = %s, want = %s", got, host2NICLinkAddr)
|
||||
}
|
||||
checker.IPv6(t, stack.PayloadSince(p.Pkt.NetworkHeader()),
|
||||
checker.SrcAddr(host1IPv6Addr.AddressWithPrefix.Address),
|
||||
@@ -1556,8 +1556,8 @@ func TestPacketQueing(t *testing.T) {
|
||||
t.Errorf("got Proto = %d, want = %d", p.Proto, ProtocolNumber)
|
||||
}
|
||||
snmc := header.SolicitedNodeAddr(host2IPv6Addr.AddressWithPrefix.Address)
|
||||
if want := header.EthernetAddressFromMulticastIPv6Address(snmc); p.Route.RemoteLinkAddress != want {
|
||||
t.Errorf("got p.Route.RemoteLinkAddress = %s, want = %s", p.Route.RemoteLinkAddress, want)
|
||||
if got, want := p.Route.RemoteLinkAddress(), header.EthernetAddressFromMulticastIPv6Address(snmc); got != want {
|
||||
t.Errorf("got p.Route.RemoteLinkAddress() = %s, want = %s", got, want)
|
||||
}
|
||||
checker.IPv6(t, stack.PayloadSince(p.Pkt.NetworkHeader()),
|
||||
checker.SrcAddr(host1IPv6Addr.AddressWithPrefix.Address),
|
||||
|
||||
@@ -650,8 +650,8 @@ func TestNeighorSolicitationResponse(t *testing.T) {
|
||||
if p.Route.RemoteAddress != respNSDst {
|
||||
t.Errorf("got p.Route.RemoteAddress = %s, want = %s", p.Route.RemoteAddress, respNSDst)
|
||||
}
|
||||
if want := header.EthernetAddressFromMulticastIPv6Address(respNSDst); p.Route.RemoteLinkAddress != want {
|
||||
t.Errorf("got p.Route.RemoteLinkAddress = %s, want = %s", p.Route.RemoteLinkAddress, want)
|
||||
if got, want := p.Route.RemoteLinkAddress(), header.EthernetAddressFromMulticastIPv6Address(respNSDst); got != want {
|
||||
t.Errorf("got p.Route.RemoteLinkAddress() = %s, want = %s", got, want)
|
||||
}
|
||||
|
||||
checker.IPv6(t, stack.PayloadSince(p.Pkt.NetworkHeader()),
|
||||
@@ -706,8 +706,8 @@ func TestNeighorSolicitationResponse(t *testing.T) {
|
||||
if p.Route.RemoteAddress != test.naDst {
|
||||
t.Errorf("got p.Route.RemoteAddress = %s, want = %s", p.Route.RemoteAddress, test.naDst)
|
||||
}
|
||||
if p.Route.RemoteLinkAddress != test.naDstLinkAddr {
|
||||
t.Errorf("got p.Route.RemoteLinkAddress = %s, want = %s", p.Route.RemoteLinkAddress, test.naDstLinkAddr)
|
||||
if got := p.Route.RemoteLinkAddress(); got != test.naDstLinkAddr {
|
||||
t.Errorf("got p.Route.RemoteLinkAddress() = %s, want = %s", got, test.naDstLinkAddr)
|
||||
}
|
||||
|
||||
checker.IPv6(t, stack.PayloadSince(p.Pkt.NetworkHeader()),
|
||||
|
||||
@@ -132,7 +132,6 @@ go_test(
|
||||
"//pkg/tcpip/transport/udp",
|
||||
"//pkg/waiter",
|
||||
"@com_github_google_go_cmp//cmp:go_default_library",
|
||||
"@com_github_google_go_cmp//cmp/cmpopts:go_default_library",
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
@@ -309,7 +309,7 @@ func (e *fwdTestLinkEndpoint) LinkAddress() tcpip.LinkAddress {
|
||||
|
||||
func (e fwdTestLinkEndpoint) WritePacket(r *Route, gso *GSO, protocol tcpip.NetworkProtocolNumber, pkt *PacketBuffer) *tcpip.Error {
|
||||
p := fwdTestPacketInfo{
|
||||
RemoteLinkAddress: r.RemoteLinkAddress,
|
||||
RemoteLinkAddress: r.RemoteLinkAddress(),
|
||||
LocalLinkAddress: r.LocalLinkAddress,
|
||||
Pkt: pkt,
|
||||
}
|
||||
|
||||
@@ -540,8 +540,8 @@ func TestDADResolve(t *testing.T) {
|
||||
|
||||
// Make sure the right remote link address is used.
|
||||
snmc := header.SolicitedNodeAddr(addr1)
|
||||
if want := header.EthernetAddressFromMulticastIPv6Address(snmc); p.Route.RemoteLinkAddress != want {
|
||||
t.Errorf("got remote link address = %s, want = %s", p.Route.RemoteLinkAddress, want)
|
||||
if got, want := p.Route.RemoteLinkAddress(), header.EthernetAddressFromMulticastIPv6Address(snmc); got != want {
|
||||
t.Errorf("got remote link address = %s, want = %s", got, want)
|
||||
}
|
||||
|
||||
// Check NDP NS packet.
|
||||
@@ -5197,8 +5197,8 @@ func TestRouterSolicitation(t *testing.T) {
|
||||
}
|
||||
|
||||
// Make sure the right remote link address is used.
|
||||
if want := header.EthernetAddressFromMulticastIPv6Address(header.IPv6AllRoutersMulticastAddress); p.Route.RemoteLinkAddress != want {
|
||||
t.Errorf("got remote link address = %s, want = %s", p.Route.RemoteLinkAddress, want)
|
||||
if got, want := p.Route.RemoteLinkAddress(), header.EthernetAddressFromMulticastIPv6Address(header.IPv6AllRoutersMulticastAddress); got != want {
|
||||
t.Errorf("got remote link address = %s, want = %s", got, want)
|
||||
}
|
||||
|
||||
checker.IPv6(t, stack.PayloadSince(p.Pkt.NetworkHeader()),
|
||||
|
||||
@@ -279,9 +279,9 @@ func (n *NIC) WritePacket(r *Route, gso *GSO, protocol tcpip.NetworkProtocolNumb
|
||||
// WritePacketToRemote implements NetworkInterface.
|
||||
func (n *NIC) WritePacketToRemote(remoteLinkAddr tcpip.LinkAddress, gso *GSO, protocol tcpip.NetworkProtocolNumber, pkt *PacketBuffer) *tcpip.Error {
|
||||
r := Route{
|
||||
NetProto: protocol,
|
||||
RemoteLinkAddress: remoteLinkAddr,
|
||||
NetProto: protocol,
|
||||
}
|
||||
r.ResolveWith(remoteLinkAddr)
|
||||
return n.writePacket(&r, gso, protocol, pkt)
|
||||
}
|
||||
|
||||
|
||||
+47
-24
@@ -34,10 +34,6 @@ type Route struct {
|
||||
// RemoteAddress is the final destination of the route.
|
||||
RemoteAddress tcpip.Address
|
||||
|
||||
// RemoteLinkAddress is the link-layer (MAC) address of the
|
||||
// final destination of the route.
|
||||
RemoteLinkAddress tcpip.LinkAddress
|
||||
|
||||
// LocalAddress is the local address where the route starts.
|
||||
LocalAddress tcpip.Address
|
||||
|
||||
@@ -64,6 +60,10 @@ type Route struct {
|
||||
|
||||
// localAddressEndpoint is the local address this route is associated with.
|
||||
localAddressEndpoint AssignableAddressEndpoint
|
||||
|
||||
// remoteLinkAddress is the link-layer (MAC) address of the next hop in the
|
||||
// route.
|
||||
remoteLinkAddress tcpip.LinkAddress
|
||||
}
|
||||
|
||||
// outgoingNIC is the interface this route uses to write packets.
|
||||
@@ -113,7 +113,7 @@ func constructAndValidateRoute(netProto tcpip.NetworkProtocolNumber, addressEndp
|
||||
if len(gateway) > 0 {
|
||||
r.NextHop = gateway
|
||||
} else if subnet := addressEndpoint.Subnet(); subnet.IsBroadcast(remoteAddr) {
|
||||
r.RemoteLinkAddress = header.EthernetBroadcastAddress
|
||||
r.ResolveWith(header.EthernetBroadcastAddress)
|
||||
}
|
||||
|
||||
return r
|
||||
@@ -190,6 +190,14 @@ func makeLocalRoute(netProto tcpip.NetworkProtocolNumber, localAddr, remoteAddr
|
||||
return makeRouteInner(netProto, localAddr, remoteAddr, outgoingNIC, localAddressNIC, localAddressEndpoint, loop)
|
||||
}
|
||||
|
||||
// RemoteLinkAddress returns the link-layer (MAC) address of the next hop in
|
||||
// the route.
|
||||
func (r *Route) RemoteLinkAddress() tcpip.LinkAddress {
|
||||
r.mu.RLock()
|
||||
defer r.mu.RUnlock()
|
||||
return r.mu.remoteLinkAddress
|
||||
}
|
||||
|
||||
// NICID returns the id of the NIC from which this route originates.
|
||||
func (r *Route) NICID() tcpip.NICID {
|
||||
return r.outgoingNIC.ID()
|
||||
@@ -251,7 +259,9 @@ func (r *Route) GSOMaxSize() uint32 {
|
||||
// ResolveWith immediately resolves a route with the specified remote link
|
||||
// address.
|
||||
func (r *Route) ResolveWith(addr tcpip.LinkAddress) {
|
||||
r.RemoteLinkAddress = addr
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
r.mu.remoteLinkAddress = addr
|
||||
}
|
||||
|
||||
// Resolve attempts to resolve the link address if necessary. Returns ErrWouldBlock in
|
||||
@@ -264,7 +274,10 @@ func (r *Route) ResolveWith(addr tcpip.LinkAddress) {
|
||||
//
|
||||
// The NIC r uses must not be locked.
|
||||
func (r *Route) Resolve(waker *sleep.Waker) (<-chan struct{}, *tcpip.Error) {
|
||||
if !r.IsResolutionRequired() {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
|
||||
if !r.isResolutionRequiredRLocked() {
|
||||
// Nothing to do if there is no cache (which does the resolution on cache miss) or
|
||||
// link address is already known.
|
||||
return nil, nil
|
||||
@@ -274,7 +287,7 @@ func (r *Route) Resolve(waker *sleep.Waker) (<-chan struct{}, *tcpip.Error) {
|
||||
if nextAddr == "" {
|
||||
// Local link address is already known.
|
||||
if r.RemoteAddress == r.LocalAddress {
|
||||
r.RemoteLinkAddress = r.LocalLinkAddress
|
||||
r.mu.remoteLinkAddress = r.LocalLinkAddress
|
||||
return nil, nil
|
||||
}
|
||||
nextAddr = r.RemoteAddress
|
||||
@@ -292,7 +305,7 @@ func (r *Route) Resolve(waker *sleep.Waker) (<-chan struct{}, *tcpip.Error) {
|
||||
if err != nil {
|
||||
return ch, err
|
||||
}
|
||||
r.RemoteLinkAddress = entry.LinkAddr
|
||||
r.mu.remoteLinkAddress = entry.LinkAddr
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
@@ -300,7 +313,7 @@ func (r *Route) Resolve(waker *sleep.Waker) (<-chan struct{}, *tcpip.Error) {
|
||||
if err != nil {
|
||||
return ch, err
|
||||
}
|
||||
r.RemoteLinkAddress = linkAddr
|
||||
r.mu.remoteLinkAddress = linkAddr
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
@@ -329,7 +342,13 @@ func (r *Route) local() bool {
|
||||
//
|
||||
// The NICs the route is associated with must not be locked.
|
||||
func (r *Route) IsResolutionRequired() bool {
|
||||
if !r.isValidForOutgoing() || r.RemoteLinkAddress != "" || r.local() {
|
||||
r.mu.RLock()
|
||||
defer r.mu.RUnlock()
|
||||
return r.isResolutionRequiredRLocked()
|
||||
}
|
||||
|
||||
func (r *Route) isResolutionRequiredRLocked() bool {
|
||||
if !r.isValidForOutgoingRLocked() || r.mu.remoteLinkAddress != "" || r.local() {
|
||||
return false
|
||||
}
|
||||
|
||||
@@ -337,13 +356,17 @@ func (r *Route) IsResolutionRequired() bool {
|
||||
}
|
||||
|
||||
func (r *Route) isValidForOutgoing() bool {
|
||||
r.mu.RLock()
|
||||
defer r.mu.RUnlock()
|
||||
return r.isValidForOutgoingRLocked()
|
||||
}
|
||||
|
||||
func (r *Route) isValidForOutgoingRLocked() bool {
|
||||
if !r.outgoingNIC.Enabled() {
|
||||
return false
|
||||
}
|
||||
|
||||
r.mu.RLock()
|
||||
localAddressEndpoint := r.mu.localAddressEndpoint
|
||||
r.mu.RUnlock()
|
||||
if localAddressEndpoint == nil || !r.localAddressNIC.isValidForOutgoing(localAddressEndpoint) {
|
||||
return false
|
||||
}
|
||||
@@ -413,17 +436,16 @@ func (r *Route) Clone() *Route {
|
||||
defer r.mu.RUnlock()
|
||||
|
||||
newRoute := &Route{
|
||||
RemoteAddress: r.RemoteAddress,
|
||||
RemoteLinkAddress: r.RemoteLinkAddress,
|
||||
LocalAddress: r.LocalAddress,
|
||||
LocalLinkAddress: r.LocalLinkAddress,
|
||||
NextHop: r.NextHop,
|
||||
NetProto: r.NetProto,
|
||||
Loop: r.Loop,
|
||||
localAddressNIC: r.localAddressNIC,
|
||||
outgoingNIC: r.outgoingNIC,
|
||||
linkCache: r.linkCache,
|
||||
linkRes: r.linkRes,
|
||||
RemoteAddress: r.RemoteAddress,
|
||||
LocalAddress: r.LocalAddress,
|
||||
LocalLinkAddress: r.LocalLinkAddress,
|
||||
NextHop: r.NextHop,
|
||||
NetProto: r.NetProto,
|
||||
Loop: r.Loop,
|
||||
localAddressNIC: r.localAddressNIC,
|
||||
outgoingNIC: r.outgoingNIC,
|
||||
linkCache: r.linkCache,
|
||||
linkRes: r.linkRes,
|
||||
}
|
||||
|
||||
newRoute.mu.Lock()
|
||||
@@ -434,6 +456,7 @@ func (r *Route) Clone() *Route {
|
||||
panic(fmt.Sprintf("failed to increment reference count for local address endpoint = %s", newRoute.LocalAddress))
|
||||
}
|
||||
}
|
||||
newRoute.mu.remoteLinkAddress = r.mu.remoteLinkAddress
|
||||
|
||||
return newRoute
|
||||
}
|
||||
|
||||
@@ -27,7 +27,6 @@ import (
|
||||
"time"
|
||||
|
||||
"github.com/google/go-cmp/cmp"
|
||||
"github.com/google/go-cmp/cmp/cmpopts"
|
||||
"gvisor.dev/gvisor/pkg/rand"
|
||||
"gvisor.dev/gvisor/pkg/sync"
|
||||
"gvisor.dev/gvisor/pkg/tcpip"
|
||||
@@ -1570,8 +1569,8 @@ func verifyRoute(gotRoute, wantRoute *stack.Route) error {
|
||||
if gotRoute.RemoteAddress != wantRoute.RemoteAddress {
|
||||
return fmt.Errorf("bad remote address: got %s, want = %s", gotRoute.RemoteAddress, wantRoute.RemoteAddress)
|
||||
}
|
||||
if gotRoute.RemoteLinkAddress != wantRoute.RemoteLinkAddress {
|
||||
return fmt.Errorf("bad remote link address: got %s, want = %s", gotRoute.RemoteLinkAddress, wantRoute.RemoteLinkAddress)
|
||||
if got, want := gotRoute.RemoteLinkAddress(), wantRoute.RemoteLinkAddress(); got != want {
|
||||
return fmt.Errorf("bad remote link address: got %s, want = %s", got, want)
|
||||
}
|
||||
if gotRoute.NextHop != wantRoute.NextHop {
|
||||
return fmt.Errorf("bad next-hop address: got %s, want = %s", gotRoute.NextHop, wantRoute.NextHop)
|
||||
@@ -3351,11 +3350,16 @@ func TestOutgoingSubnetBroadcast(t *testing.T) {
|
||||
remNetSubnetBcast := remNetSubnet.Broadcast()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
nicAddr tcpip.ProtocolAddress
|
||||
routes []tcpip.Route
|
||||
remoteAddr tcpip.Address
|
||||
expectedRoute *stack.Route
|
||||
name string
|
||||
nicAddr tcpip.ProtocolAddress
|
||||
routes []tcpip.Route
|
||||
remoteAddr tcpip.Address
|
||||
expectedLocalAddress tcpip.Address
|
||||
expectedRemoteAddress tcpip.Address
|
||||
expectedRemoteLinkAddress tcpip.LinkAddress
|
||||
expectedNextHop tcpip.Address
|
||||
expectedNetProto tcpip.NetworkProtocolNumber
|
||||
expectedLoop stack.PacketLooping
|
||||
}{
|
||||
// Broadcast to a locally attached subnet populates the broadcast MAC.
|
||||
{
|
||||
@@ -3370,14 +3374,12 @@ func TestOutgoingSubnetBroadcast(t *testing.T) {
|
||||
NIC: nicID1,
|
||||
},
|
||||
},
|
||||
remoteAddr: ipv4SubnetBcast,
|
||||
expectedRoute: &stack.Route{
|
||||
LocalAddress: ipv4Addr.Address,
|
||||
RemoteAddress: ipv4SubnetBcast,
|
||||
RemoteLinkAddress: header.EthernetBroadcastAddress,
|
||||
NetProto: header.IPv4ProtocolNumber,
|
||||
Loop: stack.PacketOut | stack.PacketLoop,
|
||||
},
|
||||
remoteAddr: ipv4SubnetBcast,
|
||||
expectedLocalAddress: ipv4Addr.Address,
|
||||
expectedRemoteAddress: ipv4SubnetBcast,
|
||||
expectedRemoteLinkAddress: header.EthernetBroadcastAddress,
|
||||
expectedNetProto: header.IPv4ProtocolNumber,
|
||||
expectedLoop: stack.PacketOut | stack.PacketLoop,
|
||||
},
|
||||
// Broadcast to a locally attached /31 subnet does not populate the
|
||||
// broadcast MAC.
|
||||
@@ -3393,13 +3395,11 @@ func TestOutgoingSubnetBroadcast(t *testing.T) {
|
||||
NIC: nicID1,
|
||||
},
|
||||
},
|
||||
remoteAddr: ipv4Subnet31Bcast,
|
||||
expectedRoute: &stack.Route{
|
||||
LocalAddress: ipv4AddrPrefix31.Address,
|
||||
RemoteAddress: ipv4Subnet31Bcast,
|
||||
NetProto: header.IPv4ProtocolNumber,
|
||||
Loop: stack.PacketOut,
|
||||
},
|
||||
remoteAddr: ipv4Subnet31Bcast,
|
||||
expectedLocalAddress: ipv4AddrPrefix31.Address,
|
||||
expectedRemoteAddress: ipv4Subnet31Bcast,
|
||||
expectedNetProto: header.IPv4ProtocolNumber,
|
||||
expectedLoop: stack.PacketOut,
|
||||
},
|
||||
// Broadcast to a locally attached /32 subnet does not populate the
|
||||
// broadcast MAC.
|
||||
@@ -3415,13 +3415,11 @@ func TestOutgoingSubnetBroadcast(t *testing.T) {
|
||||
NIC: nicID1,
|
||||
},
|
||||
},
|
||||
remoteAddr: ipv4Subnet32Bcast,
|
||||
expectedRoute: &stack.Route{
|
||||
LocalAddress: ipv4AddrPrefix32.Address,
|
||||
RemoteAddress: ipv4Subnet32Bcast,
|
||||
NetProto: header.IPv4ProtocolNumber,
|
||||
Loop: stack.PacketOut,
|
||||
},
|
||||
remoteAddr: ipv4Subnet32Bcast,
|
||||
expectedLocalAddress: ipv4AddrPrefix32.Address,
|
||||
expectedRemoteAddress: ipv4Subnet32Bcast,
|
||||
expectedNetProto: header.IPv4ProtocolNumber,
|
||||
expectedLoop: stack.PacketOut,
|
||||
},
|
||||
// IPv6 has no notion of a broadcast.
|
||||
{
|
||||
@@ -3436,13 +3434,11 @@ func TestOutgoingSubnetBroadcast(t *testing.T) {
|
||||
NIC: nicID1,
|
||||
},
|
||||
},
|
||||
remoteAddr: ipv6SubnetBcast,
|
||||
expectedRoute: &stack.Route{
|
||||
LocalAddress: ipv6Addr.Address,
|
||||
RemoteAddress: ipv6SubnetBcast,
|
||||
NetProto: header.IPv6ProtocolNumber,
|
||||
Loop: stack.PacketOut,
|
||||
},
|
||||
remoteAddr: ipv6SubnetBcast,
|
||||
expectedLocalAddress: ipv6Addr.Address,
|
||||
expectedRemoteAddress: ipv6SubnetBcast,
|
||||
expectedNetProto: header.IPv6ProtocolNumber,
|
||||
expectedLoop: stack.PacketOut,
|
||||
},
|
||||
// Broadcast to a remote subnet in the route table is send to the next-hop
|
||||
// gateway.
|
||||
@@ -3459,14 +3455,12 @@ func TestOutgoingSubnetBroadcast(t *testing.T) {
|
||||
NIC: nicID1,
|
||||
},
|
||||
},
|
||||
remoteAddr: remNetSubnetBcast,
|
||||
expectedRoute: &stack.Route{
|
||||
LocalAddress: ipv4Addr.Address,
|
||||
RemoteAddress: remNetSubnetBcast,
|
||||
NextHop: ipv4Gateway,
|
||||
NetProto: header.IPv4ProtocolNumber,
|
||||
Loop: stack.PacketOut,
|
||||
},
|
||||
remoteAddr: remNetSubnetBcast,
|
||||
expectedLocalAddress: ipv4Addr.Address,
|
||||
expectedRemoteAddress: remNetSubnetBcast,
|
||||
expectedNextHop: ipv4Gateway,
|
||||
expectedNetProto: header.IPv4ProtocolNumber,
|
||||
expectedLoop: stack.PacketOut,
|
||||
},
|
||||
// Broadcast to an unknown subnet follows the default route. Note that this
|
||||
// is essentially just routing an unknown destination IP, because w/o any
|
||||
@@ -3484,14 +3478,12 @@ func TestOutgoingSubnetBroadcast(t *testing.T) {
|
||||
NIC: nicID1,
|
||||
},
|
||||
},
|
||||
remoteAddr: remNetSubnetBcast,
|
||||
expectedRoute: &stack.Route{
|
||||
LocalAddress: ipv4Addr.Address,
|
||||
RemoteAddress: remNetSubnetBcast,
|
||||
NextHop: ipv4Gateway,
|
||||
NetProto: header.IPv4ProtocolNumber,
|
||||
Loop: stack.PacketOut,
|
||||
},
|
||||
remoteAddr: remNetSubnetBcast,
|
||||
expectedLocalAddress: ipv4Addr.Address,
|
||||
expectedRemoteAddress: remNetSubnetBcast,
|
||||
expectedNextHop: ipv4Gateway,
|
||||
expectedNetProto: header.IPv4ProtocolNumber,
|
||||
expectedLoop: stack.PacketOut,
|
||||
},
|
||||
}
|
||||
|
||||
@@ -3520,10 +3512,27 @@ func TestOutgoingSubnetBroadcast(t *testing.T) {
|
||||
t.Fatalf("got unexpected address length = %d bytes", l)
|
||||
}
|
||||
|
||||
if r, err := s.FindRoute(unspecifiedNICID, "" /* localAddr */, test.remoteAddr, netProto, false /* multicastLoop */); err != nil {
|
||||
r, err := s.FindRoute(unspecifiedNICID, "" /* localAddr */, test.remoteAddr, netProto, false /* multicastLoop */)
|
||||
if err != nil {
|
||||
t.Fatalf("FindRoute(%d, '', %s, %d): %s", unspecifiedNICID, test.remoteAddr, netProto, err)
|
||||
} else if diff := cmp.Diff(r, test.expectedRoute, cmpopts.IgnoreUnexported(stack.Route{})); diff != "" {
|
||||
t.Errorf("route mismatch (-want +got):\n%s", diff)
|
||||
}
|
||||
if r.LocalAddress != test.expectedLocalAddress {
|
||||
t.Errorf("got r.LocalAddress = %s, want = %s", r.LocalAddress, test.expectedLocalAddress)
|
||||
}
|
||||
if r.RemoteAddress != test.expectedRemoteAddress {
|
||||
t.Errorf("got r.RemoteAddress = %s, want = %s", r.RemoteAddress, test.expectedRemoteAddress)
|
||||
}
|
||||
if got := r.RemoteLinkAddress(); got != test.expectedRemoteLinkAddress {
|
||||
t.Errorf("got r.RemoteLinkAddress() = %s, want = %s", got, test.expectedRemoteLinkAddress)
|
||||
}
|
||||
if r.NextHop != test.expectedNextHop {
|
||||
t.Errorf("got r.NextHop = %s, want = %s", r.NextHop, test.expectedNextHop)
|
||||
}
|
||||
if r.NetProto != test.expectedNetProto {
|
||||
t.Errorf("got r.NetProto = %d, want = %d", r.NetProto, test.expectedNetProto)
|
||||
}
|
||||
if r.Loop != test.expectedLoop {
|
||||
t.Errorf("got r.Loop = %x, want = %x", r.Loop, test.expectedLoop)
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -4195,7 +4204,7 @@ func TestWritePacketToRemote(t *testing.T) {
|
||||
if got, want := pkt.Proto, test.protocol; got != want {
|
||||
t.Fatalf("pkt.Proto = %d, want %d", got, want)
|
||||
}
|
||||
if got, want := pkt.Route.RemoteLinkAddress, linkAddr2; got != want {
|
||||
if got, want := pkt.Route.RemoteLinkAddress(), linkAddr2; got != want {
|
||||
t.Fatalf("pkt.Route.RemoteAddress = %s, want %s", got, want)
|
||||
}
|
||||
if diff := cmp.Diff(pkt.Pkt.Data.ToView(), buffer.View(test.payload)); diff != "" {
|
||||
|
||||
@@ -274,26 +274,8 @@ func (e *endpoint) write(p tcpip.Payloader, opts tcpip.WriteOptions) (int64, <-c
|
||||
}
|
||||
}
|
||||
|
||||
var route *stack.Route
|
||||
if to == nil {
|
||||
route = e.route
|
||||
|
||||
if route.IsResolutionRequired() {
|
||||
// Promote lock to exclusive if using a shared route,
|
||||
// given that it may need to change in Route.Resolve()
|
||||
// call below.
|
||||
e.mu.RUnlock()
|
||||
defer e.mu.RLock()
|
||||
|
||||
e.mu.Lock()
|
||||
defer e.mu.Unlock()
|
||||
|
||||
// Recheck state after lock was re-acquired.
|
||||
if e.state != stateConnected {
|
||||
return 0, nil, tcpip.ErrInvalidEndpointState
|
||||
}
|
||||
}
|
||||
} else {
|
||||
route := e.route
|
||||
if to != nil {
|
||||
// Reject destination address if it goes through a different
|
||||
// NIC than the endpoint was bound to.
|
||||
nicID := to.NIC
|
||||
|
||||
@@ -18,7 +18,6 @@ import (
|
||||
"fmt"
|
||||
"sync/atomic"
|
||||
|
||||
"gvisor.dev/gvisor/pkg/sleep"
|
||||
"gvisor.dev/gvisor/pkg/sync"
|
||||
"gvisor.dev/gvisor/pkg/tcpip"
|
||||
"gvisor.dev/gvisor/pkg/tcpip/buffer"
|
||||
@@ -457,36 +456,9 @@ func (e *endpoint) write(p tcpip.Payloader, opts tcpip.WriteOptions) (int64, <-c
|
||||
}
|
||||
}
|
||||
|
||||
var route *stack.Route
|
||||
var resolve func(waker *sleep.Waker) (ch <-chan struct{}, err *tcpip.Error)
|
||||
var dstPort uint16
|
||||
if to == nil {
|
||||
route = e.route
|
||||
dstPort = e.dstPort
|
||||
resolve = func(waker *sleep.Waker) (ch <-chan struct{}, err *tcpip.Error) {
|
||||
// Promote lock to exclusive if using a shared route, given that it may
|
||||
// need to change in Route.Resolve() call below.
|
||||
e.mu.RUnlock()
|
||||
e.mu.Lock()
|
||||
|
||||
// Recheck state after lock was re-acquired.
|
||||
if e.EndpointState() != StateConnected {
|
||||
err = tcpip.ErrInvalidEndpointState
|
||||
}
|
||||
if err == nil && route.IsResolutionRequired() {
|
||||
ch, err = route.Resolve(waker)
|
||||
}
|
||||
|
||||
e.mu.Unlock()
|
||||
e.mu.RLock()
|
||||
|
||||
// Recheck state after lock was re-acquired.
|
||||
if e.EndpointState() != StateConnected {
|
||||
err = tcpip.ErrInvalidEndpointState
|
||||
}
|
||||
return ch, err
|
||||
}
|
||||
} else {
|
||||
route := e.route
|
||||
dstPort := e.dstPort
|
||||
if to != nil {
|
||||
// Reject destination address if it goes through a different
|
||||
// NIC than the endpoint was bound to.
|
||||
nicID := to.NIC
|
||||
@@ -516,7 +488,6 @@ func (e *endpoint) write(p tcpip.Payloader, opts tcpip.WriteOptions) (int64, <-c
|
||||
|
||||
route = r
|
||||
dstPort = dst.Port
|
||||
resolve = route.Resolve
|
||||
}
|
||||
|
||||
if !e.ops.GetBroadcast() && route.IsOutboundBroadcast() {
|
||||
@@ -524,7 +495,7 @@ func (e *endpoint) write(p tcpip.Payloader, opts tcpip.WriteOptions) (int64, <-c
|
||||
}
|
||||
|
||||
if route.IsResolutionRequired() {
|
||||
if ch, err := resolve(nil); err != nil {
|
||||
if ch, err := route.Resolve(nil); err != nil {
|
||||
if err == tcpip.ErrWouldBlock {
|
||||
return 0, ch, tcpip.ErrNoLinkAddress
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user