From 7103ddb238dc8ed75d8d24ce1aae5bd36fc61bd8 Mon Sep 17 00:00:00 2001 From: Tony Gong Date: Thu, 7 Jul 2022 18:54:12 -0700 Subject: [PATCH] Add checklocks to addressable_endpoint_state.go Added checklocks annotations to `addressable_endpoint_state.go`. Refactored slightly to appease the analyzer. PiperOrigin-RevId: 459651242 --- pkg/tcpip/stack/addressable_endpoint_state.go | 221 ++++++++++-------- 1 file changed, 129 insertions(+), 92 deletions(-) diff --git a/pkg/tcpip/stack/addressable_endpoint_state.go b/pkg/tcpip/stack/addressable_endpoint_state.go index 4f3ac1640..0107bc074 100644 --- a/pkg/tcpip/stack/addressable_endpoint_state.go +++ b/pkg/tcpip/stack/addressable_endpoint_state.go @@ -31,12 +31,11 @@ type AddressableEndpointState struct { // // AddressableEndpointState.mu // addressState.mu - mu struct { - sync.RWMutex - - endpoints map[tcpip.Address]*addressState - primary []*addressState - } + mu sync.RWMutex + // +checklocks:mu + endpoints map[tcpip.Address]*addressState + // +checklocks:mu + primary []*addressState } // Init initializes the AddressableEndpointState with networkEndpoint. @@ -47,7 +46,7 @@ func (a *AddressableEndpointState) Init(networkEndpoint NetworkEndpoint) { a.mu.Lock() defer a.mu.Unlock() - a.mu.endpoints = make(map[tcpip.Address]*addressState) + a.endpoints = make(map[tcpip.Address]*addressState) } // GetAddress returns the AddressEndpoint for the passed address. @@ -60,7 +59,7 @@ func (a *AddressableEndpointState) GetAddress(addr tcpip.Address) AddressEndpoin a.mu.RLock() defer a.mu.RUnlock() - ep, ok := a.mu.endpoints[addr] + ep, ok := a.endpoints[addr] if !ok { return nil } @@ -74,7 +73,7 @@ func (a *AddressableEndpointState) ForEachEndpoint(f func(AddressEndpoint) bool) a.mu.RLock() defer a.mu.RUnlock() - for _, ep := range a.mu.endpoints { + for _, ep := range a.endpoints { if !f(ep) { return } @@ -88,7 +87,7 @@ func (a *AddressableEndpointState) ForEachPrimaryEndpoint(f func(AddressEndpoint a.mu.RLock() defer a.mu.RUnlock() - for _, ep := range a.mu.primary { + for _, ep := range a.primary { if !f(ep) { return } @@ -101,19 +100,20 @@ func (a *AddressableEndpointState) releaseAddressState(addrState *addressState) a.releaseAddressStateLocked(addrState) } -// releaseAddressState removes addrState from s's address state (primary and endpoints list). +// releaseAddressStateLocked removes addrState from a's address state +// (primary and endpoints list). // -// Preconditions: a.mu must be write locked. +// +checklocks:a.mu func (a *AddressableEndpointState) releaseAddressStateLocked(addrState *addressState) { - oldPrimary := a.mu.primary - for i, s := range a.mu.primary { + oldPrimary := a.primary + for i, s := range a.primary { if s == addrState { - a.mu.primary = append(a.mu.primary[:i], a.mu.primary[i+1:]...) + a.primary = append(a.primary[:i], a.primary[i+1:]...) oldPrimary[len(oldPrimary)-1] = nil break } } - delete(a.mu.endpoints, addrState.addr.Address) + delete(a.endpoints, addrState.addr.Address) } // AddAndAcquirePermanentAddress implements AddressableEndpoint. @@ -179,12 +179,12 @@ func (a *AddressableEndpointState) AddAndAcquireTemporaryAddress(addr tcpip.Addr // address already exists in any other state, then *tcpip.ErrDuplicateAddress is // returned, regardless the kind of address that is being added. // -// Precondition: a.mu must be write locked. +// +checklocks:a.mu func (a *AddressableEndpointState) addAndAcquireAddressLocked(addr tcpip.AddressWithPrefix, properties AddressProperties, permanent bool) (*addressState, tcpip.Error) { // attemptAddToPrimary is false when the address is already in the primary // address list. attemptAddToPrimary := true - addrState, ok := a.mu.endpoints[addr.Address] + addrState, ok := a.endpoints[addr.Address] if ok { if !permanent { // We are adding a non-permanent address but the address exists. No need @@ -193,20 +193,21 @@ func (a *AddressableEndpointState) addAndAcquireAddressLocked(addr tcpip.Address return nil, &tcpip.ErrDuplicateAddress{} } - addrState.mu.Lock() - if addrState.mu.kind.IsPermanent() { - addrState.mu.Unlock() + addrState.mu.RLock() + if addrState.refs == 0 { + panic(fmt.Sprintf("found an address that should have been released (ref count == 0); address = %s", addrState.addr)) + } + isPermanent := addrState.kind.IsPermanent() + addrState.mu.RUnlock() + + if isPermanent { // We are adding a permanent address but a permanent address already // exists. return nil, &tcpip.ErrDuplicateAddress{} } - if addrState.mu.refs == 0 { - panic(fmt.Sprintf("found an address that should have been released (ref count == 0); address = %s", addrState.addr)) - } - // We now promote the address. - for i, s := range a.mu.primary { + for i, s := range a.primary { if s == addrState { switch properties.PEB { case CanBePrimaryEndpoint: @@ -217,19 +218,17 @@ func (a *AddressableEndpointState) addAndAcquireAddressLocked(addr tcpip.Address // The address is already first in the primary address list. attemptAddToPrimary = false } else { - a.mu.primary = append(a.mu.primary[:i], a.mu.primary[i+1:]...) + a.primary = append(a.primary[:i], a.primary[i+1:]...) } case NeverPrimaryEndpoint: - a.mu.primary = append(a.mu.primary[:i], a.mu.primary[i+1:]...) + a.primary = append(a.primary[:i], a.primary[i+1:]...) default: panic(fmt.Sprintf("unrecognized primary endpoint behaviour = %d", properties.PEB)) } break } } - } - - if addrState == nil { + } else { addrState = &addressState{ addressableEndpointState: a, addr: addr, @@ -238,11 +237,10 @@ func (a *AddressableEndpointState) addAndAcquireAddressLocked(addr tcpip.Address // results in allocations on every call. subnet: addr.Subnet(), } - a.mu.endpoints[addr.Address] = addrState - addrState.mu.Lock() + a.endpoints[addr.Address] = addrState // We never promote an address to temporary - it can only be added as such. - // If we are actaully adding a permanent address, it is promoted below. - addrState.mu.kind = Temporary + // If we are actually adding a permanent address, it is promoted below. + addrState.kind = Temporary } // At this point we have an address we are either promoting from an expired or @@ -250,40 +248,41 @@ func (a *AddressableEndpointState) addAndAcquireAddressLocked(addr tcpip.Address // or we are adding a new temporary or permanent address. // // The address MUST be write locked at this point. - defer addrState.mu.Unlock() // +checklocksforce + addrState.mu.Lock() + defer addrState.mu.Unlock() if permanent { - if addrState.mu.kind.IsPermanent() { + if addrState.kind.IsPermanent() { panic(fmt.Sprintf("only non-permanent addresses should be promoted to permanent; address = %s", addrState.addr)) } // Primary addresses are biased by 1. - addrState.mu.refs++ - addrState.mu.kind = Permanent + addrState.refs++ + addrState.kind = Permanent } // Acquire the address before returning it. - addrState.mu.refs++ - addrState.mu.deprecated = properties.Deprecated - addrState.mu.configType = properties.ConfigType + addrState.refs++ + addrState.deprecated = properties.Deprecated + addrState.configType = properties.ConfigType if attemptAddToPrimary { switch properties.PEB { case NeverPrimaryEndpoint: case CanBePrimaryEndpoint: - a.mu.primary = append(a.mu.primary, addrState) + a.primary = append(a.primary, addrState) case FirstPrimaryEndpoint: - if cap(a.mu.primary) == len(a.mu.primary) { - a.mu.primary = append([]*addressState{addrState}, a.mu.primary...) + if cap(a.primary) == len(a.primary) { + a.primary = append([]*addressState{addrState}, a.primary...) } else { // Shift all the endpoints by 1 to make room for the new address at the // front. We could have just created a new slice but this saves // allocations when the slice has capacity for the new address. - primaryCount := len(a.mu.primary) - a.mu.primary = append(a.mu.primary, nil) - if n := copy(a.mu.primary[1:], a.mu.primary); n != primaryCount { + primaryCount := len(a.primary) + a.primary = append(a.primary, nil) + if n := copy(a.primary[1:], a.primary); n != primaryCount { panic(fmt.Sprintf("copied %d elements; expected = %d elements", n, primaryCount)) } - a.mu.primary[0] = addrState + a.primary[0] = addrState } default: panic(fmt.Sprintf("unrecognized primary endpoint behaviour = %d", properties.PEB)) @@ -303,9 +302,9 @@ func (a *AddressableEndpointState) RemovePermanentAddress(addr tcpip.Address) tc // removePermanentAddressLocked is like RemovePermanentAddress but with locking // requirements. // -// Precondition: a.mu must be write locked. +// +checklocks:a.mu func (a *AddressableEndpointState) removePermanentAddressLocked(addr tcpip.Address) tcpip.Error { - addrState, ok := a.mu.endpoints[addr] + addrState, ok := a.endpoints[addr] if !ok { return &tcpip.ErrBadLocalAddress{} } @@ -329,7 +328,7 @@ func (a *AddressableEndpointState) RemovePermanentEndpoint(ep AddressEndpoint) t // removePermanentAddressLocked is like RemovePermanentAddress but with locking // requirements. // -// Precondition: a.mu must be write locked. +// +checklocks:a.mu func (a *AddressableEndpointState) removePermanentEndpointLocked(addrState *addressState) tcpip.Error { if !addrState.GetKind().IsPermanent() { return &tcpip.ErrBadLocalAddress{} @@ -350,25 +349,25 @@ func (a *AddressableEndpointState) decAddressRef(addrState *addressState) { // decAddressRefLocked is like decAddressRef but with locking requirements. // -// Precondition: a.mu must be write locked. +// +checklocks:a.mu func (a *AddressableEndpointState) decAddressRefLocked(addrState *addressState) { addrState.mu.Lock() defer addrState.mu.Unlock() - if addrState.mu.refs == 0 { + if addrState.refs == 0 { panic(fmt.Sprintf("attempted to decrease ref count for AddressEndpoint w/ addr = %s when it is already released", addrState.addr)) } - addrState.mu.refs-- + addrState.refs-- - if addrState.mu.refs != 0 { + if addrState.refs != 0 { return } // A non-expired permanent address must not have its reference count dropped // to 0. - if addrState.mu.kind.IsPermanent() { - panic(fmt.Sprintf("permanent addresses should be removed through the AddressableEndpoint: addr = %s, kind = %d", addrState.addr, addrState.mu.kind)) + if addrState.kind.IsPermanent() { + panic(fmt.Sprintf("permanent addresses should be removed through the AddressableEndpoint: addr = %s, kind = %d", addrState.addr, addrState.kind)) } a.releaseAddressStateLocked(addrState) @@ -376,10 +375,10 @@ func (a *AddressableEndpointState) decAddressRefLocked(addrState *addressState) // SetDeprecated implements stack.AddressableEndpoint. func (a *AddressableEndpointState) SetDeprecated(addr tcpip.Address, deprecated bool) tcpip.Error { - a.mu.Lock() - defer a.mu.Unlock() + a.mu.RLock() + defer a.mu.RUnlock() - addrState, ok := a.mu.endpoints[addr] + addrState, ok := a.endpoints[addr] if !ok { return &tcpip.ErrBadLocalAddress{} } @@ -393,24 +392,36 @@ func (a *AddressableEndpointState) MainAddress() tcpip.AddressWithPrefix { defer a.mu.RUnlock() ep := a.acquirePrimaryAddressRLocked(func(ep *addressState) bool { - return ep.GetKind() == Permanent + switch kind := ep.GetKind(); kind { + case Permanent: + return true + case PermanentTentative, PermanentExpired, Temporary: + return false + default: + panic(fmt.Sprintf("unknown address kind: %d", kind)) + } }) if ep == nil { return tcpip.AddressWithPrefix{} } - addr := ep.AddressWithPrefix() - a.decAddressRefLocked(ep) + // Note that when ep must have a ref count >=2, because its ref count + // must be >=1 in order to be found and the ref count was incremented + // when a reference was acquired. The only way for the ref count to + // drop below 2 is for the endpoint to be removed, which requires a + // write lock; so we're guaranteed to be able to decrement the ref + // count and not need to remove the endpoint from a.primary. + ep.decRefMustNotFree() return addr } // acquirePrimaryAddressRLocked returns an acquired primary address that is // valid according to isValid. // -// Precondition: e.mu must be read locked +// +checklocksread:a.mu func (a *AddressableEndpointState) acquirePrimaryAddressRLocked(isValid func(*addressState) bool) *addressState { var deprecatedEndpoint *addressState - for _, ep := range a.mu.primary { + for _, ep := range a.primary { if !isValid(ep) { continue } @@ -422,8 +433,13 @@ func (a *AddressableEndpointState) acquirePrimaryAddressRLocked(isValid func(*ad // If we kept track of a deprecated endpoint, decrement its reference // count since it was incremented when we decided to keep track of it. if deprecatedEndpoint != nil { - a.decAddressRefLocked(deprecatedEndpoint) - deprecatedEndpoint = nil + // Note that when deprecatedEndpoint was found, its ref count + // must have necessarily been >=1, and after incrementing it + // must be >=2. The only way for the ref count to drop below 2 is + // for the endpoint to be removed, which requires a write lock; + // so we're guaranteed to be able to decrement the ref count + // and not need to remove the endpoint from a.primary. + deprecatedEndpoint.decRefMustNotFree() } return ep @@ -455,7 +471,7 @@ func (a *AddressableEndpointState) acquirePrimaryAddressRLocked(isValid func(*ad // returned. func (a *AddressableEndpointState) AcquireAssignedAddressOrMatching(localAddr tcpip.Address, f func(AddressEndpoint) bool, allowTemp bool, tempPEB PrimaryEndpointBehavior) AddressEndpoint { lookup := func() *addressState { - if addrState, ok := a.mu.endpoints[localAddr]; ok { + if addrState, ok := a.endpoints[localAddr]; ok { if !addrState.IsAssigned(allowTemp) { return nil } @@ -468,7 +484,7 @@ func (a *AddressableEndpointState) AcquireAssignedAddressOrMatching(localAddr tc } if f != nil { - for _, addrState := range a.mu.endpoints { + for _, addrState := range a.endpoints { if addrState.IsAssigned(allowTemp) && f(addrState) && addrState.IncRef() { return addrState } @@ -538,8 +554,8 @@ func (a *AddressableEndpointState) AcquireAssignedAddress(localAddr tcpip.Addres // AcquireOutgoingPrimaryAddress implements AddressableEndpoint. func (a *AddressableEndpointState) AcquireOutgoingPrimaryAddress(remoteAddr tcpip.Address, allowExpired bool) AddressEndpoint { - a.mu.RLock() - defer a.mu.RUnlock() + a.mu.Lock() + defer a.mu.Unlock() ep := a.acquirePrimaryAddressRLocked(func(ep *addressState) bool { return ep.IsAssigned(allowExpired) @@ -557,7 +573,7 @@ func (a *AddressableEndpointState) AcquireOutgoingPrimaryAddress(remoteAddr tcpi // an interface value will therefore be non-nil even when the pointer value V // inside is nil. // - // Since acquirePrimaryAddressRLocked returns a nil value with a non-nil type, + // Since acquirePrimaryAddressLocked returns a nil value with a non-nil type, // we need to explicitly return nil below if ep is (a typed) nil. if ep == nil { return nil @@ -572,13 +588,16 @@ func (a *AddressableEndpointState) PrimaryAddresses() []tcpip.AddressWithPrefix defer a.mu.RUnlock() var addrs []tcpip.AddressWithPrefix - for _, ep := range a.mu.primary { + for _, ep := range a.primary { + switch kind := ep.GetKind(); kind { // Don't include tentative, expired or temporary endpoints // to avoid confusion and prevent the caller from using // those. - switch ep.GetKind() { case PermanentTentative, PermanentExpired, Temporary: continue + case Permanent: + default: + panic(fmt.Sprintf("address %s has unknown kind %d", ep.AddressWithPrefix(), kind)) } addrs = append(addrs, ep.AddressWithPrefix()) @@ -593,7 +612,7 @@ func (a *AddressableEndpointState) PermanentAddresses() []tcpip.AddressWithPrefi defer a.mu.RUnlock() var addrs []tcpip.AddressWithPrefix - for _, ep := range a.mu.endpoints { + for _, ep := range a.endpoints { if !ep.GetKind().IsPermanent() { continue } @@ -609,7 +628,7 @@ func (a *AddressableEndpointState) Cleanup() { a.mu.Lock() defer a.mu.Unlock() - for _, ep := range a.mu.endpoints { + for _, ep := range a.endpoints { // removePermanentEndpointLocked returns *tcpip.ErrBadLocalAddress if ep is // not a permanent address. switch err := a.removePermanentEndpointLocked(ep); err.(type) { @@ -628,18 +647,20 @@ type addressState struct { addr tcpip.AddressWithPrefix subnet tcpip.Subnet temporary bool + // Lock ordering (from outer to inner lock ordering): // // AddressableEndpointState.mu // addressState.mu - mu struct { - sync.RWMutex - - refs uint32 - kind AddressKind - configType AddressConfigType - deprecated bool - } + mu sync.RWMutex + // checklocks:mu + refs uint32 + // checklocks:mu + kind AddressKind + // checklocks:mu + configType AddressConfigType + // checklocks:mu + deprecated bool } // AddressWithPrefix implements AddressEndpoint. @@ -656,14 +677,14 @@ func (a *addressState) Subnet() tcpip.Subnet { func (a *addressState) GetKind() AddressKind { a.mu.RLock() defer a.mu.RUnlock() - return a.mu.kind + return a.kind } // SetKind implements AddressEndpoint. func (a *addressState) SetKind(kind AddressKind) { a.mu.Lock() defer a.mu.Unlock() - a.mu.kind = kind + a.kind = kind } // IsAssigned implements AddressEndpoint. @@ -686,11 +707,11 @@ func (a *addressState) IsAssigned(allowExpired bool) bool { func (a *addressState) IncRef() bool { a.mu.Lock() defer a.mu.Unlock() - if a.mu.refs == 0 { + if a.refs == 0 { return false } - a.mu.refs++ + a.refs++ return true } @@ -699,27 +720,43 @@ func (a *addressState) DecRef() { a.addressableEndpointState.decAddressRef(a) } +// decRefMustNotFree decreases the reference count with the guarantee that the +// reference count will be greater than 0 after the decrement. +// +// Panics if the ref count is less than 2 after acquiring the lock in this +// function. +func (a *addressState) decRefMustNotFree() { + a.mu.Lock() + defer a.mu.Unlock() + + if a.refs < 2 { + panic(fmt.Sprintf("cannot decrease addressState %s ref count %d without freeing the endpoint", a.addr, a.refs)) + } + a.refs-- +} + // ConfigType implements AddressEndpoint. func (a *addressState) ConfigType() AddressConfigType { a.mu.RLock() defer a.mu.RUnlock() - return a.mu.configType + return a.configType } // SetDeprecated implements AddressEndpoint. func (a *addressState) SetDeprecated(d bool) { a.mu.Lock() defer a.mu.Unlock() - a.mu.deprecated = d + a.deprecated = d } // Deprecated implements AddressEndpoint. func (a *addressState) Deprecated() bool { a.mu.RLock() defer a.mu.RUnlock() - return a.mu.deprecated + return a.deprecated } +// Temporary implements AddressEndpoint. func (a *addressState) Temporary() bool { return a.temporary }