diff --git a/token/services/identity/deserializer/deserializer_test.go b/token/services/identity/deserializer/deserializer_test.go index 4d5242697c..6e67bdb809 100644 --- a/token/services/identity/deserializer/deserializer_test.go +++ b/token/services/identity/deserializer/deserializer_test.go @@ -9,6 +9,7 @@ package deserializer import ( "context" "errors" + "sync" "testing" "github.com/LFDT-Panurus/panurus/token/driver" @@ -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() +} diff --git a/token/services/identity/deserializer/eidrh.go b/token/services/identity/deserializer/eidrh.go index 2157a199c0..b243fc6160 100644 --- a/token/services/identity/deserializer/eidrh.go +++ b/token/services/identity/deserializer/eidrh.go @@ -8,6 +8,7 @@ package deserializer import ( "context" + "sync" "github.com/LFDT-Panurus/panurus/token/driver" "github.com/LFDT-Panurus/panurus/token/services/identity" @@ -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 } @@ -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 } @@ -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) } diff --git a/token/services/identity/deserializer/signer.go b/token/services/identity/deserializer/signer.go index 49a7433cb3..1243e719b3 100644 --- a/token/services/identity/deserializer/signer.go +++ b/token/services/identity/deserializer/signer.go @@ -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" @@ -20,6 +22,7 @@ import ( type TypedSignerDeserializer = idriver.TypedSignerDeserializer type TypedSignerDeserializerMultiplex struct { + mutex sync.RWMutex deserializers map[idriver.IdentityType][]TypedSignerDeserializer } @@ -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} @@ -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) diff --git a/token/services/identity/deserializer/verifier.go b/token/services/identity/deserializer/verifier.go index a58e8240da..d6be54c04d 100644 --- a/token/services/identity/deserializer/verifier.go +++ b/token/services/identity/deserializer/verifier.go @@ -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" @@ -23,6 +25,7 @@ var logger = logging.MustGetLogger() type TypedVerifierDeserializer = idriver.TypedVerifierDeserializer type TypedVerifierDeserializerMultiplex struct { + mutex sync.RWMutex deserializers map[idriver.IdentityType][]idriver.TypedVerifierDeserializer } @@ -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} @@ -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) @@ -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) } @@ -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) } @@ -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