From 175be3501193699bed06ab703832863802757650 Mon Sep 17 00:00:00 2001 From: Adin Scannell Date: Thu, 1 Dec 2022 15:36:40 -0800 Subject: [PATCH] Fix lock violations. The locks should not be rewritten. This is especially confusing with locking seemingly respected in the first part of this function. Instead, just update the relevant fields while holding appropriate locks. PiperOrigin-RevId: 492317437 --- .../transport/internal/network/endpoint.go | 42 +++++++++---------- 1 file changed, 19 insertions(+), 23 deletions(-) diff --git a/pkg/tcpip/transport/internal/network/endpoint.go b/pkg/tcpip/transport/internal/network/endpoint.go index 87f45d3df..11c76e920 100644 --- a/pkg/tcpip/transport/internal/network/endpoint.go +++ b/pkg/tcpip/transport/internal/network/endpoint.go @@ -118,10 +118,9 @@ type multicastMembership struct { // Init initializes the endpoint. func (e *Endpoint) Init(s *stack.Stack, netProto tcpip.NetworkProtocolNumber, transProto tcpip.TransportProtocolNumber, ops *tcpip.SocketOptions, waiterQueue *waiter.Queue) { e.mu.Lock() - memberships := e.multicastMemberships - e.mu.Unlock() - if memberships != nil { - panic(fmt.Sprintf("endpoint is already initialized; got e.multicastMemberships = %#v, want = nil", memberships)) + defer e.mu.Unlock() + if e.multicastMemberships != nil { + panic(fmt.Sprintf("endpoint is already initialized; got e.multicastMemberships = %#v, want = nil", e.multicastMemberships)) } switch netProto { @@ -130,27 +129,24 @@ func (e *Endpoint) Init(s *stack.Stack, netProto tcpip.NetworkProtocolNumber, tr panic(fmt.Sprintf("invalid protocol number = %d", netProto)) } - *e = Endpoint{ - stack: s, - ops: ops, - netProto: netProto, - transProto: transProto, - waiterQueue: waiterQueue, - - info: stack.TransportEndpointInfo{ - NetProto: netProto, - TransProto: transProto, - }, - effectiveNetProto: netProto, - ipv4TTL: tcpip.UseDefaultIPv4TTL, - ipv6HopLimit: tcpip.UseDefaultIPv6HopLimit, - // Linux defaults to TTL=1. - multicastTTL: 1, - multicastMemberships: make(map[multicastMembership]struct{}), + e.stack = s + e.ops = ops + e.netProto = netProto + e.transProto = transProto + e.waiterQueue = waiterQueue + e.infoMu.Lock() + e.info = stack.TransportEndpointInfo{ + NetProto: netProto, + TransProto: transProto, } + e.infoMu.Unlock() + e.effectiveNetProto = netProto + e.ipv4TTL = tcpip.UseDefaultIPv4TTL + e.ipv6HopLimit = tcpip.UseDefaultIPv6HopLimit - e.mu.Lock() - defer e.mu.Unlock() + // Linux defaults to TTL=1. + e.multicastTTL = 1 + e.multicastMemberships = make(map[multicastMembership]struct{}) e.setEndpointState(transport.DatagramEndpointStateInitial) }