Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
74 changes: 74 additions & 0 deletions token/services/identity/deserializer/deserializer_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@ package deserializer
import (
"context"
"errors"
"sync"
"testing"

"github.com/LFDT-Panurus/panurus/token/driver"
Expand Down Expand Up @@ -880,3 +881,76 @@ func TestTypedIdentityVerifierDeserializer(t *testing.T) {
assert.Equal(t, 1, mockMatcherDeserializer.GetAuditInfoMatcherCallCount())
})
}

// TestDeserializerMultiplexConcurrency exercises the three deserializer
// multiplexes under concurrent registration (map writes) and deserialization
// (map reads). Without the internal mutex this reliably trips the race
// detector or the runtime's fatal "concurrent map read and map write" check;
// with it, the test runs clean. Run under `go test -race` for full coverage.
func TestDeserializerMultiplexConcurrency(t *testing.T) {
const (
goroutines = 16
iterations = 200
)

signerMultiplex := NewTypedSignerDeserializerMultiplex()
verifierMultiplex := NewTypedVerifierDeserializerMultiplex()
eidrhDeserializer := NewEIDRHDeserializer()

// Pre-register a deserializer for the type the readers look up so the read
// paths traverse the stored slice rather than bailing out on a miss.
mockSigner := &drivermock.Signer{}
signerDes := &identitydrivermock.TypedSignerDeserializer{}
signerDes.DeserializeSignerReturns(mockSigner, nil)
signerMultiplex.AddTypedSignerDeserializer(identity.Type(99), signerDes)

mockVerifier := &drivermock.Verifier{}
verifierDes := &identitydrivermock.TypedVerifierDeserializer{}
verifierDes.DeserializeVerifierReturns(mockVerifier, nil)
verifierDes.RecipientsReturns([]driver.Identity{[]byte("recipient")}, nil)
verifierDes.GetAuditInfoMatcherReturns(&drivermock.Matcher{}, nil)
verifierDes.GetAuditInfoReturns([]byte("audit-info"), nil)
verifierMultiplex.AddTypedVerifierDeserializer(identity.Type(99), verifierDes)

mockAuditInfo := &identitydrivermock.AuditInfo{}
auditDes := &identitydrivermock.AuditInfoDeserializer{}
auditDes.DeserializeAuditInfoReturns(mockAuditInfo, nil)
eidrhDeserializer.AddDeserializer(identity.Type(99), auditDes)

typedID := createTypedIdentity(t, identity.Type(99), []byte("raw-identity"))
provider := &drivermock.AuditInfoProvider{}

var wg sync.WaitGroup

// Writers: keep registering new deserializers concurrently.
for range goroutines {
wg.Go(func() {
for range iterations {
// Register under a distinct type so writers keep hammering the
// same map concurrently without clobbering the entry the
// readers rely on below.
signerMultiplex.AddTypedSignerDeserializer(identity.Type(100), &identitydrivermock.TypedSignerDeserializer{})
verifierMultiplex.AddTypedVerifierDeserializer(identity.Type(100), &identitydrivermock.TypedVerifierDeserializer{})
eidrhDeserializer.AddDeserializer(identity.Type(100), &identitydrivermock.AuditInfoDeserializer{})
}
})
}

// Readers: keep deserializing concurrently.
for range goroutines {
wg.Go(func() {
ctx := context.Background()
for range iterations {
_, _ = signerMultiplex.DeserializeSigner(ctx, typedID)
_, _ = verifierMultiplex.DeserializeVerifier(ctx, typedID)
_, _ = verifierMultiplex.Recipients(typedID)
_, _ = verifierMultiplex.GetAuditInfoMatcher(ctx, typedID, []byte("audit-info"))
_ = verifierMultiplex.MatchIdentity(ctx, typedID, []byte("audit-info"))
_, _ = verifierMultiplex.GetAuditInfo(ctx, typedID, provider)
_, _, _ = eidrhDeserializer.GetEIDAndRH(ctx, typedID, []byte("audit-info"))
}
})
}

wg.Wait()
}
6 changes: 6 additions & 0 deletions token/services/identity/deserializer/eidrh.go
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@ package deserializer

import (
"context"
"sync"

"github.com/LFDT-Panurus/panurus/token/driver"
"github.com/LFDT-Panurus/panurus/token/services/identity"
Expand All @@ -17,6 +18,7 @@ import (

// EIDRHDeserializer returns enrollment IDs behind the owners of token
type EIDRHDeserializer struct {
mutex sync.RWMutex
deserializers map[identity.Type]driver2.AuditInfoDeserializer
}

Expand All @@ -28,6 +30,8 @@ func NewEIDRHDeserializer() *EIDRHDeserializer {
}

func (e *EIDRHDeserializer) AddDeserializer(typ identity.Type, d driver2.AuditInfoDeserializer) {
e.mutex.Lock()
defer e.mutex.Unlock()
e.deserializers[typ] = d
}

Expand Down Expand Up @@ -69,7 +73,9 @@ func (e *EIDRHDeserializer) DeserializeAuditInfo(ctx context.Context, id driver.
if err != nil {
return nil, errors.Wrap(err, "failed to unmarshal to TypedIdentity")
}
e.mutex.RLock()
d, ok := e.deserializers[si.Type]
e.mutex.RUnlock()
if !ok {
return nil, errors.Errorf("no deserializer found for [%v]", si.Type)
}
Expand Down
11 changes: 9 additions & 2 deletions token/services/identity/deserializer/signer.go
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,8 @@ package deserializer
import (
"context"
errors2 "errors"
"slices"
"sync"

"github.com/LFDT-Panurus/panurus/token/driver"
"github.com/LFDT-Panurus/panurus/token/services/identity"
Expand All @@ -20,6 +22,7 @@ import (
type TypedSignerDeserializer = idriver.TypedSignerDeserializer

type TypedSignerDeserializerMultiplex struct {
mutex sync.RWMutex
deserializers map[idriver.IdentityType][]TypedSignerDeserializer
}
Comment thread
Effi-S marked this conversation as resolved.

Expand All @@ -28,6 +31,8 @@ func NewTypedSignerDeserializerMultiplex() *TypedSignerDeserializerMultiplex {
}

func (v *TypedSignerDeserializerMultiplex) AddTypedSignerDeserializer(typ idriver.IdentityType, d idriver.TypedSignerDeserializer) {
v.mutex.Lock()
defer v.mutex.Unlock()
_, ok := v.deserializers[typ]
if !ok {
v.deserializers[typ] = []TypedSignerDeserializer{d}
Expand All @@ -42,8 +47,10 @@ func (v *TypedSignerDeserializerMultiplex) DeserializeSigner(ctx context.Context
if err != nil {
return nil, errors.Wrap(err, "failed to unmarshal to TypedIdentity")
}
dess, ok := v.deserializers[si.Type]
if !ok {
v.mutex.RLock()
dess := slices.Clone(v.deserializers[si.Type])
v.mutex.RUnlock()
if dess == nil {
return nil, errors.Errorf("no deserializer found for [%v]", si.Type)
}
logger.DebugfContext(ctx, "deserializing [%s] with type [%v]", logging.Base64(id), si.Type)
Expand Down
29 changes: 21 additions & 8 deletions token/services/identity/deserializer/verifier.go
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,8 @@ package deserializer
import (
"context"
errors2 "errors"
"slices"
"sync"

"github.com/LFDT-Panurus/panurus/token/driver"
"github.com/LFDT-Panurus/panurus/token/services/identity"
Expand All @@ -23,6 +25,7 @@ var logger = logging.MustGetLogger()
type TypedVerifierDeserializer = idriver.TypedVerifierDeserializer

type TypedVerifierDeserializerMultiplex struct {
mutex sync.RWMutex
deserializers map[idriver.IdentityType][]idriver.TypedVerifierDeserializer
}

Expand All @@ -31,6 +34,8 @@ func NewTypedVerifierDeserializerMultiplex() *TypedVerifierDeserializerMultiplex
}

func (v *TypedVerifierDeserializerMultiplex) AddTypedVerifierDeserializer(typ idriver.IdentityType, d TypedVerifierDeserializer) {
v.mutex.Lock()
defer v.mutex.Unlock()
_, ok := v.deserializers[typ]
if !ok {
v.deserializers[typ] = []TypedVerifierDeserializer{d}
Expand All @@ -45,8 +50,10 @@ func (v *TypedVerifierDeserializerMultiplex) DeserializeVerifier(ctx context.Con
if err != nil {
return nil, errors.Wrap(err, "failed to unmarshal to TypedIdentity")
}
dess, ok := v.deserializers[si.Type]
if !ok {
v.mutex.RLock()
dess := slices.Clone(v.deserializers[si.Type])
v.mutex.RUnlock()
if dess == nil {
return nil, errors.Errorf("no deserializer found for [%v]", si.Type)
}
logger.DebugfContext(ctx, "deserializing [%s] with type [%v]", id, si.Type)
Expand All @@ -73,8 +80,10 @@ func (v *TypedVerifierDeserializerMultiplex) Recipients(id driver.Identity) ([]d
if err != nil {
return nil, errors.Wrap(err, "failed to unmarshal to TypedIdentity")
}
dess, ok := v.deserializers[si.Type]
if !ok {
v.mutex.RLock()
dess := slices.Clone(v.deserializers[si.Type])
v.mutex.RUnlock()
if dess == nil {
return nil, errors.Errorf("no deserializer found for [%v]", si.Type)
}

Expand Down Expand Up @@ -110,8 +119,10 @@ func (v *TypedVerifierDeserializerMultiplex) GetAuditInfoMatcher(ctx context.Con
}

func (v *TypedVerifierDeserializerMultiplex) getMatcher(ctx context.Context, idType idriver.IdentityType, id driver.Identity, auditInfo []byte) (driver.Matcher, error) {
dess, ok := v.deserializers[idType]
if !ok {
v.mutex.RLock()
dess := slices.Clone(v.deserializers[idType])
v.mutex.RUnlock()
if dess == nil {
return nil, errors.Errorf("no deserializer found for [%v]", idType)
}

Expand Down Expand Up @@ -153,8 +164,10 @@ func (v *TypedVerifierDeserializerMultiplex) GetAuditInfo(ctx context.Context, i
if err != nil {
return nil, errors.Wrap(err, "failed to unmarshal to TypedIdentity")
}
dess, ok := v.deserializers[si.Type]
if !ok {
v.mutex.RLock()
dess := slices.Clone(v.deserializers[si.Type])
v.mutex.RUnlock()
if dess == nil {
return nil, errors.Errorf("no deserializer found for [%v]", si.Type)
}
var errs []error
Expand Down
Loading