diff --git a/cmd/splitcli/main.go b/cmd/splitcli/main.go index a55d520..ecbed11 100644 --- a/cmd/splitcli/main.go +++ b/cmd/splitcli/main.go @@ -3,6 +3,7 @@ package main import ( "fmt" "os" + "strings" "time" "github.com/splitio/go-toolkit/v5/logging" @@ -67,7 +68,20 @@ func executeCall(c types.ClientInterface, a *conf.CliArgs) (string, error) { case "treatment": res, err := c.Treatment(a.Key, a.BucketingKey, a.Feature, a.Attributes) return res.Treatment, err - case "treatments", "treatmentWithConfig", "treatmentsWithConfig", "track": + case "treatments": + res, err := c.Treatments(a.Key, a.BucketingKey, a.Features, a.Attributes) + var sb strings.Builder + for _, result := range res { + if sb.Len() == 0 { // first item doesn't require a leading ',' + sb.WriteString(result.Treatment) + } else { + sb.WriteString("," + result.Treatment) + } + } + return sb.String(), err + case "track": + return "", c.Track(a.Key, a.TrafficType, a.EventType, a.EventVal, nil) + case "treatmentWithConfig", "treatmentsWithConfig": return "", fmt.Errorf("method '%s' is not yet implemented", a.Method) default: return "", fmt.Errorf("unknwon method '%s'", a.Method) diff --git a/cmd/splitd/main.go b/cmd/splitd/main.go index da34b7f..9ff0625 100644 --- a/cmd/splitd/main.go +++ b/cmd/splitd/main.go @@ -15,25 +15,27 @@ import ( func main() { - printHeader() + printHeader() cfg, err := conf.ReadConfig() if err != nil { fmt.Println("error reading config: ", err.Error()) os.Exit(1) } - handleFlags(cfg) + handleFlags(cfg) - logger := logging.NewLogger(cfg.Logger.ToLoggerOptions()) + loggerCfg, err := cfg.Logger.ToLoggerOptions() + exitOnErr("logging setup", err) + logger := logging.NewLogger(loggerCfg) splitSDK, err := sdk.New(logger, cfg.SDK.Apikey, cfg.SDK.ToSDKConf()) - exitOnErr("sdk initialization", err) + exitOnErr("sdk initialization", err) - linkCFG, err := cfg.Link.ToListenerOpts() - exitOnErr("link config", err) + linkCFG, err := cfg.Link.ToListenerOpts() + exitOnErr("link config", err) errc, lShutdown, err := link.Listen(logger, splitSDK, linkCFG) - exitOnErr("rpc listener setup", err) + exitOnErr("rpc listener setup", err) shutdown := util.NewShutdownHandler() shutdown.RegisterHook(func() { @@ -46,26 +48,26 @@ func main() { // Wait for connection to end (either gracefully of because of an error) err = <-errc - exitOnErr("shutdown: ", err) + exitOnErr("shutdown: ", err) } func printHeader() { - fmt.Println(splitio.ASCILogo) - fmt.Printf("Splitd Agent - Version %s. (2023)\n\n", splitio.Version) + fmt.Println(splitio.ASCILogo) + fmt.Printf("Splitd Agent - Version %s. (2023)\n\n", splitio.Version) } func handleFlags(cfg *conf.Config) { - printConf := flag.Bool("outputConfig", false, "print config (with partially obfuscated apikey)") - flag.Parse() - if *printConf { - fmt.Printf("\nConfig: %s\n", cfg) - os.Exit(0) - } + printConf := flag.Bool("outputConfig", false, "print config (with partially obfuscated apikey)") + flag.Parse() + if *printConf { + fmt.Printf("\nConfig: %s\n", cfg) + os.Exit(0) + } } func exitOnErr(ctxStr string, err error) { if err != nil { - fmt.Printf("%s: startup error: %s\n", ctxStr, err.Error()) + fmt.Printf("%s: startup error: %s\n", ctxStr, err.Error()) os.Exit(1) } diff --git a/splitio/conf/splitcli.go b/splitio/conf/splitcli.go index 9555cdd..5a50fcf 100644 --- a/splitio/conf/splitcli.go +++ b/splitio/conf/splitcli.go @@ -14,6 +14,7 @@ import ( ) type CliArgs struct { + ID string LogLevel string Protocol string Serialization string @@ -31,13 +32,15 @@ type CliArgs struct { Features []string TrafficType string EventType string - EventVal float64 + EventVal *float64 Attributes map[string]interface{} } func (a *CliArgs) LinkOpts() (*link.ConsumerOptions, error) { opts := link.DefaultConsumerOptions() + cc.SetIfNotEmpty(&opts.Consumer.ID, &a.ID) + var err error if a.Protocol != "" { @@ -45,13 +48,11 @@ func (a *CliArgs) LinkOpts() (*link.ConsumerOptions, error) { return nil, fmt.Errorf("invalid protocol version %s", a.Protocol) } } - if a.ConnType != "" { if opts.Transfer.ConnType, err = cc.ParseConnType(a.ConnType); err != nil { return nil, fmt.Errorf("invalid connection type %s", a.ConnType) } } - if a.Serialization != "" { if opts.Serialization, err = cc.ParseSerializer(a.Serialization); err != nil { return nil, fmt.Errorf("invalid serialization %s", a.Serialization) @@ -69,6 +70,7 @@ func (a *CliArgs) LinkOpts() (*link.ConsumerOptions, error) { func ParseCliArgs() (*CliArgs, error) { cliFlags := flag.NewFlagSet(os.Args[0], flag.ContinueOnError) + id := cliFlags.String("id", "", "ID to use for internal event queuing separation. Defaults to app's pid") p := cliFlags.String("protocol", "", "Protocol version [v1]") s := cliFlags.String("serialization", "", "Client-Daemon communication serialization mechanism [msgpack]") ll := cliFlags.String("log-level", "INFO", "log level [ERROR,WARNING,INFO,DEBUG]") @@ -89,9 +91,13 @@ func ParseCliArgs() (*CliArgs, error) { return nil, fmt.Errorf("error parsing arguments: %w", err) } - val, err := strconv.ParseFloat(*ev, 64) - if *ev != "" && err != nil { - return nil, fmt.Errorf("error parsing event value") + var eventVal *float64 + if *ev != "" { + val, err := strconv.ParseFloat(*ev, 64) + if err != nil { + return nil, fmt.Errorf("error parsing event value") + } + eventVal = &val } if *at == "" { @@ -103,20 +109,21 @@ func ParseCliArgs() (*CliArgs, error) { } return &CliArgs{ - Serialization: *s, - Protocol: *p, - LogLevel: *ll, - ConnType: *ct, - ConnAddr: *ca, - BufSize: *bs, - Method: *m, - Key: *k, - BucketingKey: *bk, - Feature: *f, - Features: strings.Split(*fs, ","), - TrafficType: *tt, - EventType: *et, - EventVal: val, - Attributes: attrs, + ID: *id, + Serialization: *s, + Protocol: *p, + LogLevel: *ll, + ConnType: *ct, + ConnAddr: *ca, + BufSize: *bs, + Method: *m, + Key: *k, + BucketingKey: *bk, + Feature: *f, + Features: strings.Split(*fs, ","), + TrafficType: *tt, + EventType: *et, + EventVal: eventVal, + Attributes: attrs, }, nil } diff --git a/splitio/conf/splitcli_test.go b/splitio/conf/splitcli_test.go index e935a7e..31b9290 100644 --- a/splitio/conf/splitcli_test.go +++ b/splitio/conf/splitcli_test.go @@ -39,7 +39,7 @@ func TestCliConfig(t *testing.T) { assert.Equal(t, []string{"someFeature1", "someFeature2"}, parsed.Features) assert.Equal(t, "someTrafficType", parsed.TrafficType) assert.Equal(t, "someEventType", parsed.EventType) - assert.Equal(t, 0.123, parsed.EventVal) + assert.Equal(t, ref(float64(0.123)), parsed.EventVal) assert.Equal(t, map[string]interface{}{"some": "attribute"}, parsed.Attributes) // test bad buffer size @@ -63,32 +63,32 @@ func TestCliConfig(t *testing.T) { } func TestLinkOptions(t *testing.T) { - // test defaults - os.Args = os.Args[:1] + // test defaults + os.Args = os.Args[:1] parsed, err := ParseCliArgs() assert.Nil(t, err) lo, err := parsed.LinkOpts() assert.Nil(t, err) assert.Equal(t, link.DefaultConsumerOptions(), *lo) - // test bad protocol - os.Args = []string{os.Args[0], "-protocol=sarasa"} + // test bad protocol + os.Args = []string{os.Args[0], "-protocol=sarasa"} parsed, err = ParseCliArgs() assert.Nil(t, err) lo, err = parsed.LinkOpts() assert.NotNil(t, err) assert.ErrorContains(t, err, "protocol") - // test bad conn type - os.Args = []string{os.Args[0], "-conn-type=sarasa"} + // test bad conn type + os.Args = []string{os.Args[0], "-conn-type=sarasa"} parsed, err = ParseCliArgs() assert.Nil(t, err) lo, err = parsed.LinkOpts() assert.NotNil(t, err) assert.ErrorContains(t, err, "connection type") - // test bad serialization - os.Args = []string{os.Args[0], "-serialization=pinpin"} + // test bad serialization + os.Args = []string{os.Args[0], "-serialization=pinpin"} parsed, err = ParseCliArgs() assert.Nil(t, err) lo, err = parsed.LinkOpts() diff --git a/splitio/conf/splitd.go b/splitio/conf/splitd.go index 615edc0..e1fc39e 100644 --- a/splitio/conf/splitd.go +++ b/splitio/conf/splitd.go @@ -11,6 +11,7 @@ import ( "github.com/splitio/go-toolkit/v5/logging" "github.com/splitio/splitd/splitio/link" + sdlogging "github.com/splitio/splitd/splitio/logging" sdkConf "github.com/splitio/splitd/splitio/sdk/conf" cc "github.com/splitio/splitd/splitio/util/conf" "gopkg.in/yaml.v3" @@ -94,20 +95,50 @@ func (l *Link) ToListenerOpts() (*link.ListenerOptions, error) { } type SDK struct { - Apikey string `yaml:"apikey"` - LabelsEnabled *bool `yaml:"labelsEnabled"` - StreamingEnabled *bool `yaml:"streamingEnabled"` - URLs URLs `yaml:"urls"` + Apikey string `yaml:"apikey"` + LabelsEnabled *bool `yaml:"labelsEnabled"` + StreamingEnabled *bool `yaml:"streamingEnabled"` + URLs URLs `yaml:"urls"` + FeatureFlags FeatureFlags `yaml:"featureFlags"` + Impressions Impressions `yaml:"impressions"` } -func (s *SDK) ToSDKConf() *sdkConf.Config { +type FeatureFlags struct { + SplitNotificationQueueSize *int `yaml:"splitNotificationQueueSize"` + SplitRefreshRateSeconds *int `yaml:"splitRefreshSeconds"` + SegmentNotificationQueueSize *int `yaml:"segmentNotificationQueueSize"` + SegmentRefreshRateSeconds *int `yaml:"segmentRefreshSeconds"` + SegmentUpdateWorkers *int `yaml:"segmentUpdateWorkers"` + SegmentUpdateQueueSize *int `yaml:"segmentUpdateQueueSize"` +} +type Impressions struct { + Mode *string `yaml:"mode"` + RefreshRateSeconds *int `yaml:"refreshRateSeconds"` + CountRefreshRateSeconds *int `yaml:"countRefreshRateSeconds"` + QueueSize *int `yaml:"queueSize"` + ObserverSize *int `yaml:"observerSize"` + Watermark *int `yaml:"watermark"` +} + +func (s *SDK) ToSDKConf() *sdkConf.Config { cfg := sdkConf.DefaultConfig() + durationFromSeconds := func(seconds int) time.Duration { return time.Duration(seconds) * time.Second } cc.SetIfNotNil(&cfg.LabelsEnabled, s.LabelsEnabled) cc.SetIfNotNil(&cfg.StreamingEnabled, s.StreamingEnabled) + cc.SetIfNotEmpty(&cfg.Splits.UpdateBufferSize, s.FeatureFlags.SplitNotificationQueueSize) + cc.MapIfNotNil(&cfg.Splits.SyncPeriod, s.FeatureFlags.SplitRefreshRateSeconds, durationFromSeconds) + cc.SetIfNotEmpty(&cfg.Segments.UpdateBufferSize, s.FeatureFlags.SegmentNotificationQueueSize) + cc.SetIfNotEmpty(&cfg.Segments.QueueSize, s.FeatureFlags.SegmentUpdateQueueSize) + cc.SetIfNotEmpty(&cfg.Segments.WorkerCount, s.FeatureFlags.SegmentUpdateWorkers) + cc.MapIfNotNil(&cfg.Segments.SyncPeriod, s.FeatureFlags.SegmentRefreshRateSeconds, durationFromSeconds) + cc.SetIfNotEmpty(&cfg.Impressions.Mode, s.Impressions.Mode) + cc.SetIfNotEmpty(&cfg.Impressions.ObserverSize, s.Impressions.ObserverSize) + cc.SetIfNotEmpty(&cfg.Impressions.QueueSize, s.Impressions.QueueSize) + cc.MapIfNotNil(&cfg.Impressions.SyncPeriod, s.Impressions.RefreshRateSeconds, durationFromSeconds) + cc.MapIfNotNil(&cfg.Impressions.CountSyncPeriod, s.Impressions.CountRefreshRateSeconds, durationFromSeconds) s.URLs.updateSDKConfURLs(&cfg.URLs) return cfg - } type URLs struct { @@ -127,21 +158,34 @@ func (u *URLs) updateSDKConfURLs(dst *sdkConf.URLs) { } type Logger struct { - Level *string `yaml:"level"` + Level *string `yaml:"level"` + Output *string `yaml:"file"` + RotationMaxFiles *int `yaml:"rotationMaxFiles"` + RotationMaxBytesPerFile *int `yaml:"rotationMaxBytesPerFile"` } -func (l *Logger) ToLoggerOptions() *logging.LoggerOptions { +func (l *Logger) ToLoggerOptions() (*logging.LoggerOptions, error) { + + writer, err := sdlogging.GetWriter(l.Output, l.RotationMaxFiles, l.RotationMaxBytesPerFile) + if err != nil { + return nil, fmt.Errorf("error parsing logger options: %w", err) + } opts := &logging.LoggerOptions{ LogLevel: logging.LevelError, StandardLoggerFlags: log.Ltime | log.Lshortfile, + ErrorWriter: writer, + WarningWriter: writer, + InfoWriter: writer, + DebugWriter: writer, + VerboseWriter: writer, } if l.Level != nil { opts.LogLevel = logging.Level(strings.ToUpper(*l.Level)) } - return opts + return opts, nil } func ReadConfig() (*Config, error) { diff --git a/splitio/conf/splitd_test.go b/splitio/conf/splitd_test.go index 09acfcc..9683d05 100644 --- a/splitio/conf/splitd_test.go +++ b/splitio/conf/splitd_test.go @@ -117,6 +117,22 @@ func TestSDK(t *testing.T) { Streaming: ref("streamingURL"), Telemetry: ref("telemetryURL"), }, + FeatureFlags: FeatureFlags{ + SplitNotificationQueueSize: ref(1), + SplitRefreshRateSeconds: ref(2), + SegmentNotificationQueueSize: ref(3), + SegmentRefreshRateSeconds: ref(4), + SegmentUpdateWorkers: ref(5), + SegmentUpdateQueueSize: ref(6), + }, + Impressions: Impressions{ + Mode: ref("optimized"), + RefreshRateSeconds: ref(1), + CountRefreshRateSeconds: ref(2), + QueueSize: ref(3), + ObserverSize: ref(4), + Watermark: ref(5), + }, } expected := conf.DefaultConfig() @@ -127,7 +143,17 @@ func TestSDK(t *testing.T) { expected.URLs.Events = "eventsURL" expected.URLs.Streaming = "streamingURL" expected.URLs.Telemetry = "telemetryURL" - + expected.Splits.UpdateBufferSize = 1 + expected.Splits.SyncPeriod = 2 * time.Second + expected.Segments.UpdateBufferSize = 3 + expected.Segments.SyncPeriod = 4 * time.Second + expected.Segments.WorkerCount = 5 + expected.Segments.QueueSize = 6 + expected.Impressions.Mode = "optimized" + expected.Impressions.SyncPeriod = 1 * time.Second + expected.Impressions.CountSyncPeriod = 2 * time.Second + expected.Impressions.QueueSize = 3 + expected.Impressions.ObserverSize = 4 assert.Equal(t, expected, sdkCFG.ToSDKConf()) } diff --git a/splitio/link/client/client.go b/splitio/link/client/client.go index 0f9d45f..9102372 100644 --- a/splitio/link/client/client.go +++ b/splitio/link/client/client.go @@ -2,6 +2,8 @@ package client import ( "fmt" + "os" + "strconv" "github.com/splitio/go-toolkit/v5/logging" "github.com/splitio/splitd/splitio/link/client/types" @@ -14,18 +16,20 @@ import ( func New(logger logging.LoggerInterface, conn transfer.RawConn, serial serializer.Interface, opts Options) (types.ClientInterface, error) { switch opts.Protocol { case protocol.V1: - return clientv1.New(logger, conn, serial, opts.ImpressionsFeedback) + return clientv1.New(opts.ID, logger, conn, serial, opts.ImpressionsFeedback) } return nil, fmt.Errorf("unknown protocol version: '%d'", opts.Protocol) } type Options struct { + ID string Protocol protocol.Version ImpressionsFeedback bool } func DefaultOptions() Options { return Options{ + ID: strconv.Itoa(os.Getpid()), Protocol: protocol.V1, ImpressionsFeedback: false, } diff --git a/splitio/link/client/types/interfaces.go b/splitio/link/client/types/interfaces.go index 638e9b5..67f3487 100644 --- a/splitio/link/client/types/interfaces.go +++ b/splitio/link/client/types/interfaces.go @@ -5,6 +5,7 @@ import "github.com/splitio/go-split-commons/v4/dtos" type ClientInterface interface { Treatment(key string, bucketingKey string, feature string, attrs map[string]interface{}) (*Result, error) Treatments(key string, bucketingKey string, features []string, attrs map[string]interface{}) (Results, error) + Track(key string, trafficType string, eventType string, value *float64, properties map[string]interface{}) error Shutdown() error } diff --git a/splitio/link/client/v1/impl.go b/splitio/link/client/v1/impl.go index f4b23d9..1d9a701 100644 --- a/splitio/link/client/v1/impl.go +++ b/splitio/link/client/v1/impl.go @@ -2,8 +2,6 @@ package v1 import ( "fmt" - "os" - "strconv" "github.com/splitio/go-split-commons/v4/dtos" "github.com/splitio/go-toolkit/v5/logging" @@ -26,7 +24,7 @@ type Impl struct { listenerFeedback bool } -func New(logger logging.LoggerInterface, conn transfer.RawConn, serializer serializer.Interface, listenerFeedback bool) (*Impl, error) { +func New(id string, logger logging.LoggerInterface, conn transfer.RawConn, serializer serializer.Interface, listenerFeedback bool) (*Impl, error) { i := &Impl{ logger: logger, conn: conn, @@ -34,7 +32,7 @@ func New(logger logging.LoggerInterface, conn transfer.RawConn, serializer seria listenerFeedback: listenerFeedback, } - if err := i.register(listenerFeedback); err != nil { + if err := i.register(id, listenerFeedback); err != nil { i.conn.Shutdown() return nil, fmt.Errorf("error during client registration: %w", err) } @@ -120,7 +118,28 @@ func (c *Impl) Treatments(key string, bucketingKey string, features []string, at return results, nil } -func (c *Impl) register(impressionsFeedback bool) error { +// Track implements types.ClientInterface +func (c *Impl) Track(key string, trafficType string, eventType string, value *float64, properties map[string]interface{}) error { + + rpc := protov1.RPC{ + RPCBase: protocol.RPCBase{Version: protocol.V1}, + OpCode: protov1.OCTrack, + Args: protov1.TrackArgs{Key: key, TrafficType: trafficType, EventType: eventType, Value: value, Properties: properties}.Encode(), + } + + resp, err := doRPC[protov1.ResponseWrapper[protov1.TrackPayload]](c, &rpc) + if err != nil { + return fmt.Errorf("error executing treatment rpc: %w", err) + } + + if resp.Status != protov1.ResultOk { + return fmt.Errorf("server responded treatment rpc with error %d", resp.Status) + } + + return nil +} + +func (c *Impl) register(id string, impressionsFeedback bool) error { var flags protov1.RegisterFlags if impressionsFeedback { flags |= 1 << protov1.RegisterFlagReturnImpressionData @@ -128,7 +147,7 @@ func (c *Impl) register(impressionsFeedback bool) error { rpc := protov1.RPC{ RPCBase: protocol.RPCBase{Version: protocol.V1}, OpCode: protov1.OCRegister, - Args: protov1.RegisterArgs{ID: strconv.Itoa(os.Getpid()), SDKVersion: fmt.Sprintf("splitd-%s", splitio.Version), Flags: flags}.Encode(), + Args: protov1.RegisterArgs{ID: id, SDKVersion: fmt.Sprintf("splitd-%s", splitio.Version), Flags: flags}.Encode(), } resp, err := doRPC[protov1.ResponseWrapper[protov1.RegisterPayload]](c, &rpc) diff --git a/splitio/link/client/v1/impl_test.go b/splitio/link/client/v1/impl_test.go index e365ab7..743fcce 100644 --- a/splitio/link/client/v1/impl_test.go +++ b/splitio/link/client/v1/impl_test.go @@ -24,7 +24,7 @@ func TestClientGetTreatmentNoImpression(t *testing.T) { rawConnMock.On("ReceiveMessage").Return([]byte("treatmentResult"), nil).Once() serializerMock := &serializerMocks.SerializerMock{} - serializerMock.On("Serialize", proto1Mocks.NewRegisterRPC(false)).Return([]byte("registrationMessage"), nil).Once() + serializerMock.On("Serialize", proto1Mocks.NewRegisterRPC("some", false)).Return([]byte("registrationMessage"), nil).Once() serializerMock.On("Parse", []byte("registrationSuccess"), mock.Anything).Return(nil).Run(func(args mock.Arguments) { *args.Get(1).(*v1.ResponseWrapper[v1.RegisterPayload]) = v1.ResponseWrapper[v1.RegisterPayload]{Status: v1.ResultOk} }).Once() @@ -37,7 +37,7 @@ func TestClientGetTreatmentNoImpression(t *testing.T) { Payload: v1.TreatmentPayload{Treatment: "on"}, } }).Once() - client, err := New(logger, rawConnMock, serializerMock, false) + client, err := New("some", logger, rawConnMock, serializerMock, false) assert.NotNil(t, client) assert.Nil(t, err) @@ -47,6 +47,35 @@ func TestClientGetTreatmentNoImpression(t *testing.T) { assert.Nil(t, res.Impression) } +func TestTrack(t *testing.T) { + + logger := logging.NewLogger(nil) + + rawConnMock := &transferMocks.RawConnMock{} + rawConnMock.On("SendMessage", []byte("registrationMessage")).Return(nil).Once() + rawConnMock.On("ReceiveMessage").Return([]byte("registrationSuccess"), nil).Once() + rawConnMock.On("SendMessage", []byte("trackMessage")).Return(nil).Once() + rawConnMock.On("ReceiveMessage").Return([]byte("trackResult"), nil).Once() + + serializerMock := &serializerMocks.SerializerMock{} + serializerMock.On("Serialize", proto1Mocks.NewRegisterRPC("some", false)).Return([]byte("registrationMessage"), nil).Once() + serializerMock.On("Parse", []byte("registrationSuccess"), mock.Anything).Return(nil).Run(func(args mock.Arguments) { + *args.Get(1).(*v1.ResponseWrapper[v1.RegisterPayload]) = v1.ResponseWrapper[v1.RegisterPayload]{Status: v1.ResultOk} + }).Once() + + serializerMock.On("Serialize", proto1Mocks.NewTrackRPC("key1", "user", "checkin", ref(2.74), map[string]interface{}{"p1": 123})). + Return([]byte("trackMessage"), nil).Once() + serializerMock.On("Parse", []byte("trackResult"), mock.Anything).Return(nil).Run(func(args mock.Arguments) { + *args.Get(1).(*v1.ResponseWrapper[v1.TrackPayload]) = *proto1Mocks.NewTrackResp(true) + }).Once() + client, err := New("some", logger, rawConnMock, serializerMock, false) + assert.NotNil(t, client) + assert.Nil(t, err) + + err = client.Track("key1", "user", "checkin", ref(2.74), map[string]interface{}{"p1": 123}) + assert.Nil(t, err) +} + func TestClientGetTreatmentWithImpression(t *testing.T) { logger := logging.NewLogger(nil) @@ -58,7 +87,7 @@ func TestClientGetTreatmentWithImpression(t *testing.T) { rawConnMock.On("ReceiveMessage").Return([]byte("treatmentResult"), nil).Once() serializerMock := &serializerMocks.SerializerMock{} - serializerMock.On("Serialize", proto1Mocks.NewRegisterRPC(true)).Return([]byte("registrationMessage"), nil).Once() + serializerMock.On("Serialize", proto1Mocks.NewRegisterRPC("some", true)).Return([]byte("registrationMessage"), nil).Once() serializerMock.On("Parse", []byte("registrationSuccess"), mock.Anything).Return(nil).Run(func(args mock.Arguments) { *args.Get(1).(*v1.ResponseWrapper[v1.RegisterPayload]) = v1.ResponseWrapper[v1.RegisterPayload]{Status: v1.ResultOk} }).Once() @@ -74,7 +103,7 @@ func TestClientGetTreatmentWithImpression(t *testing.T) { }, } }).Once() - client, err := New(logger, rawConnMock, serializerMock, true) + client, err := New("some", logger, rawConnMock, serializerMock, true) assert.NotNil(t, client) assert.Nil(t, err) @@ -104,7 +133,7 @@ func TestClientGetTreatmentsNoImpression(t *testing.T) { rawConnMock.On("ReceiveMessage").Return([]byte("treatmentsResult"), nil).Once() serializerMock := &serializerMocks.SerializerMock{} - serializerMock.On("Serialize", proto1Mocks.NewRegisterRPC(false)).Return([]byte("registrationMessage"), nil).Once() + serializerMock.On("Serialize", proto1Mocks.NewRegisterRPC("some", false)).Return([]byte("registrationMessage"), nil).Once() serializerMock.On("Parse", []byte("registrationSuccess"), mock.Anything).Return(nil).Run(func(args mock.Arguments) { *args.Get(1).(*v1.ResponseWrapper[v1.RegisterPayload]) = v1.ResponseWrapper[v1.RegisterPayload]{Status: v1.ResultOk} }).Once() @@ -116,7 +145,7 @@ func TestClientGetTreatmentsNoImpression(t *testing.T) { Status: v1.ResultOk, Payload: v1.TreatmentsPayload{Results: []v1.TreatmentPayload{{Treatment: "on"}, {Treatment: "off"}, {Treatment: "na"}}}} }).Once() - client, err := New(logger, rawConnMock, serializerMock, false) + client, err := New("some", logger, rawConnMock, serializerMock, false) assert.NotNil(t, client) assert.Nil(t, err) @@ -142,7 +171,7 @@ func TestClientGetTreatmentsWithImpression(t *testing.T) { rawConnMock.On("ReceiveMessage").Return([]byte("treatmentsResult"), nil).Once() serializerMock := &serializerMocks.SerializerMock{} - serializerMock.On("Serialize", proto1Mocks.NewRegisterRPC(true)).Return([]byte("registrationMessage"), nil).Once() + serializerMock.On("Serialize", proto1Mocks.NewRegisterRPC("some", true)).Return([]byte("registrationMessage"), nil).Once() serializerMock.On("Parse", []byte("registrationSuccess"), mock.Anything).Return(nil).Run(func(args mock.Arguments) { *args.Get(1).(*v1.ResponseWrapper[v1.RegisterPayload]) = v1.ResponseWrapper[v1.RegisterPayload]{Status: v1.ResultOk} }).Once() @@ -158,7 +187,7 @@ func TestClientGetTreatmentsWithImpression(t *testing.T) { {Treatment: "na", ListenerData: &v1.ListenerExtraData{Label: "l3", Timestamp: 3, ChangeNumber: 7}}, }}} }).Once() - client, err := New(logger, rawConnMock, serializerMock, true) + client, err := New("some", logger, rawConnMock, serializerMock, true) assert.NotNil(t, client) assert.Nil(t, err) diff --git a/splitio/link/link.go b/splitio/link/link.go index 8e35efb..6b45eeb 100644 --- a/splitio/link/link.go +++ b/splitio/link/link.go @@ -61,24 +61,24 @@ type ListenerOptions struct { } func DefaultListenerOptions() ListenerOptions { - return ListenerOptions{ - Transfer: transfer.DefaultOpts(), - Acceptor: transfer.DefaultAcceptorConfig(), - Serialization: serializer.MsgPack, - Protocol: protocol.V1, - } + return ListenerOptions{ + Transfer: transfer.DefaultOpts(), + Acceptor: transfer.DefaultAcceptorConfig(), + Serialization: serializer.MsgPack, + Protocol: protocol.V1, + } } type ConsumerOptions struct { - Transfer transfer.Options - Consumer client.Options - Serialization serializer.Mechanism + Transfer transfer.Options + Consumer client.Options + Serialization serializer.Mechanism } func DefaultConsumerOptions() ConsumerOptions { - return ConsumerOptions{ - Transfer: transfer.DefaultOpts(), - Consumer: client.DefaultOptions(), - Serialization: serializer.MsgPack, - } + return ConsumerOptions{ + Transfer: transfer.DefaultOpts(), + Consumer: client.DefaultOptions(), + Serialization: serializer.MsgPack, + } } diff --git a/splitio/link/protocol/v1/mocks/mocks.go b/splitio/link/protocol/v1/mocks/mocks.go index 0787802..abf4034 100644 --- a/splitio/link/protocol/v1/mocks/mocks.go +++ b/splitio/link/protocol/v1/mocks/mocks.go @@ -2,8 +2,6 @@ package mocks import ( "fmt" - "os" - "strconv" "github.com/splitio/splitd/splitio" "github.com/splitio/splitd/splitio/link/protocol" @@ -11,7 +9,7 @@ import ( "github.com/splitio/splitd/splitio/sdk" ) -func NewRegisterRPC(listener bool) *v1.RPC { +func NewRegisterRPC(id string, listener bool) *v1.RPC { var flags v1.RegisterFlags if listener { flags = 1 << v1.RegisterFlagReturnImpressionData @@ -19,7 +17,7 @@ func NewRegisterRPC(listener bool) *v1.RPC { return &v1.RPC{ RPCBase: protocol.RPCBase{Version: protocol.V1}, OpCode: v1.OCRegister, - Args: []interface{}{strconv.Itoa(os.Getpid()), fmt.Sprintf("splitd-%s", splitio.Version), flags}, + Args: []interface{}{id, fmt.Sprintf("splitd-%s", splitio.Version), flags}, } } @@ -39,6 +37,14 @@ func NewTreatmentsRPC(key string, bucketing string, features []string, attrs map } } +func NewTrackRPC(key string, trafficType string, eventType string, eventVal *float64, props map[string]interface{}) *v1.RPC { + return &v1.RPC{ + RPCBase: protocol.RPCBase{Version: protocol.V1}, + OpCode: v1.OCTrack, + Args: []interface{}{key, trafficType, eventType, nilOrVal(eventVal), props}, + } +} + func NewRegisterResp(ok bool) *v1.ResponseWrapper[v1.RegisterPayload] { res := v1.ResultOk if !ok { @@ -64,7 +70,7 @@ func NewTreatmentResp(ok bool, treatment string, ilData *v1.ListenerExtraData) * } } -func NewTreatmentsResp(ok bool, data []sdk.Result) *v1.ResponseWrapper[v1.TreatmentsPayload] { +func NewTreatmentsResp(ok bool, data []sdk.EvaluationResult) *v1.ResponseWrapper[v1.TreatmentsPayload] { res := v1.ResultOk if !ok { res = v1.ResultInternalError @@ -88,3 +94,21 @@ func NewTreatmentsResp(ok bool, data []sdk.Result) *v1.ResponseWrapper[v1.Treatm Payload: v1.TreatmentsPayload{Results: payload}, } } + +func NewTrackResp(ok bool) *v1.ResponseWrapper[v1.TrackPayload] { + res := v1.ResultOk + if !ok { + res = v1.ResultInternalError + } + return &v1.ResponseWrapper[v1.TrackPayload]{ + Status: res, + Payload: v1.TrackPayload{Success: ok}, + } +} + +func nilOrVal(v *float64) interface{} { + if v == nil { + return nil + } + return *v +} diff --git a/splitio/link/protocol/v1/responses.go b/splitio/link/protocol/v1/responses.go index 968be85..88aee19 100644 --- a/splitio/link/protocol/v1/responses.go +++ b/splitio/link/protocol/v1/responses.go @@ -33,6 +33,10 @@ type TreatmentsWithConfigPayload struct { Results []TreatmentWithConfigPayload `msgpack:"r"` } +type TrackPayload struct { + Success bool `msgpack:"s"` +} + type ListenerExtraData struct { Label string `msgpack:"l"` Timestamp int64 `msgpack:"m"` @@ -44,5 +48,6 @@ type validPayloadsConstraint interface { TreatmentsPayload | TreatmentWithConfigPayload | TreatmentsWithConfigPayload | + TrackPayload | RegisterPayload } diff --git a/splitio/link/protocol/v1/rpcs.go b/splitio/link/protocol/v1/rpcs.go index d0bed80..dba2d60 100644 --- a/splitio/link/protocol/v1/rpcs.go +++ b/splitio/link/protocol/v1/rpcs.go @@ -206,7 +206,6 @@ const ( TrackArgEventTypeIdx int = 2 TrackArgValueIdx int = 3 TrackArgPropertiesIdx int = 4 - TrackArgTimestampIdx int = 5 ) type TrackArgs struct { @@ -215,18 +214,24 @@ type TrackArgs struct { EventType string `msgpack:"e"` Value *float64 `msgpack:"v"` Properties map[string]interface{} `msgpack:"p"` - Timestamp int64 `msgpack:"m"` } func (r TrackArgs) Encode() []interface{} { - return []interface{}{r.Key, r.TrafficType, r.EventType, r.Value, r.Properties, r.Timestamp} + asInterface := make([]interface{}, 0, 5) + asInterface = append(asInterface, r.Key, r.TrafficType, r.EventType) + if r.Value == nil { + asInterface = append(asInterface, nil) + } + asInterface = append(asInterface, *r.Value) + asInterface = append(asInterface, r.Properties) + return asInterface } func (t *TrackArgs) PopulateFromRPC(rpc *RPC) error { if rpc.OpCode != OCTrack { return RPCParseError{Code: PECOpCodeMismatch} } - if len(rpc.Args) != 6 { + if len(rpc.Args) != 5 { return RPCParseError{Code: PECWrongArgCount} } @@ -257,11 +262,6 @@ func (t *TrackArgs) PopulateFromRPC(rpc *RPC) error { return RPCParseError{Code: PECInvalidArgType, Data: int64(TrackArgPropertiesIdx)} } - if t.Timestamp, ok = rpc.Args[TrackArgTimestampIdx].(int64); !ok { - return RPCParseError{Code: PECInvalidArgType, Data: int64(TrackArgTimestampIdx)} - - } - return nil } diff --git a/splitio/link/protocol/v1/rpcs_test.go b/splitio/link/protocol/v1/rpcs_test.go index d7194cf..d27fedf 100644 --- a/splitio/link/protocol/v1/rpcs_test.go +++ b/splitio/link/protocol/v1/rpcs_test.go @@ -174,37 +174,28 @@ func TestTrackRPCParsing(t *testing.T) { ) assert.Equal(t, RPCParseError{Code: PECInvalidArgType, Data: int64(TrackArgKeyIdx)}, - r.PopulateFromRPC(&RPC{RPCBase: protocol.RPCBase{Version: protocol.V1}, OpCode: OCTrack, Args: []interface{}{nil, nil, nil, 123, 123, nil}}), + r.PopulateFromRPC(&RPC{RPCBase: protocol.RPCBase{Version: protocol.V1}, OpCode: OCTrack, Args: []interface{}{nil, nil, nil, 123, 123}}), ) assert.Equal(t, RPCParseError{Code: PECInvalidArgType, Data: int64(TrackArgTrafficTypeIdx)}, - r.PopulateFromRPC(&RPC{RPCBase: protocol.RPCBase{Version: protocol.V1}, OpCode: OCTrack, Args: []interface{}{"key", nil, nil, "asd", 123, nil}}), + r.PopulateFromRPC(&RPC{RPCBase: protocol.RPCBase{Version: protocol.V1}, OpCode: OCTrack, Args: []interface{}{"key", nil, nil, "asd", 123}}), ) assert.Equal(t, RPCParseError{Code: PECInvalidArgType, Data: int64(TrackArgEventTypeIdx)}, - r.PopulateFromRPC(&RPC{RPCBase: protocol.RPCBase{Version: protocol.V1}, OpCode: OCTrack, Args: []interface{}{"key", "tt", nil, "asd", 123, nil}}), + r.PopulateFromRPC(&RPC{RPCBase: protocol.RPCBase{Version: protocol.V1}, OpCode: OCTrack, Args: []interface{}{"key", "tt", nil, "asd", 123}}), ) assert.Equal(t, RPCParseError{Code: PECInvalidArgType, Data: int64(TrackArgValueIdx)}, - r.PopulateFromRPC(&RPC{RPCBase: protocol.RPCBase{Version: protocol.V1}, OpCode: OCTrack, Args: []interface{}{"key", "tt", "et", "asd", 123, nil}})) + r.PopulateFromRPC(&RPC{RPCBase: protocol.RPCBase{Version: protocol.V1}, OpCode: OCTrack, Args: []interface{}{"key", "tt", "et", "asd", 123}})) assert.Equal(t, RPCParseError{Code: PECInvalidArgType, Data: int64(TrackArgPropertiesIdx)}, - r.PopulateFromRPC(&RPC{RPCBase: protocol.RPCBase{Version: protocol.V1}, OpCode: OCTrack, Args: []interface{}{"key", "tt", "et", 2.8, 123, nil}})) + r.PopulateFromRPC(&RPC{RPCBase: protocol.RPCBase{Version: protocol.V1}, OpCode: OCTrack, Args: []interface{}{"key", "tt", "et", 2.8, 123}})) - assert.Equal(t, - RPCParseError{Code: PECInvalidArgType, Data: int64(TrackArgTimestampIdx)}, - r.PopulateFromRPC(&RPC{ - RPCBase: protocol.RPCBase{Version: protocol.V1}, - OpCode: OCTrack, - Args: []interface{}{"key", "tt", "et", 2.8, map[string]interface{}{"a": 1}, nil}, - })) - - now := time.Now() err := r.PopulateFromRPC(&RPC{ RPCBase: protocol.RPCBase{Version: protocol.V1}, OpCode: OCTrack, - Args: []interface{}{"key", "tt", "et", 2.8, map[string]interface{}{"a": int64(1)}, now.UnixMilli()}, + Args: []interface{}{"key", "tt", "et", 2.8, map[string]interface{}{"a": int64(1)}}, }) assert.Nil(t, err) assert.Equal(t, "key", r.Key) @@ -212,13 +203,12 @@ func TestTrackRPCParsing(t *testing.T) { assert.Equal(t, "et", r.EventType) assert.Equal(t, ref(float64(2.8)), r.Value) assert.Equal(t, map[string]interface{}{"a": int64(1)}, r.Properties) - assert.Equal(t, now.UnixMilli(), r.Timestamp) // nil properties err = r.PopulateFromRPC(&RPC{ RPCBase: protocol.RPCBase{Version: protocol.V1}, OpCode: OCTrack, - Args: []interface{}{"key", "tt", "et", 2.8, nil, now.UnixMilli()}, + Args: []interface{}{"key", "tt", "et", 2.8, nil}, }) assert.Nil(t, err) assert.Equal(t, "key", r.Key) @@ -226,14 +216,13 @@ func TestTrackRPCParsing(t *testing.T) { assert.Equal(t, "et", r.EventType) assert.Equal(t, ref(float64(2.8)), r.Value) assert.Nil(t, r.Properties) - assert.Equal(t, now.UnixMilli(), r.Timestamp) // nil value r = TrackArgs{} err = r.PopulateFromRPC(&RPC{ RPCBase: protocol.RPCBase{Version: protocol.V1}, OpCode: OCTrack, - Args: []interface{}{"key", "tt", "et", nil, map[string]interface{}{"a": int64(1)}, now.UnixMilli()}, + Args: []interface{}{"key", "tt", "et", nil, map[string]interface{}{"a": int64(1)}}, }) assert.Nil(t, err) assert.Equal(t, "key", r.Key) @@ -241,7 +230,6 @@ func TestTrackRPCParsing(t *testing.T) { assert.Equal(t, "et", r.EventType) assert.Nil(t, r.Value) assert.Equal(t, map[string]interface{}{"a": int64(1)}, r.Properties) - assert.Equal(t, now.UnixMilli(), r.Timestamp) } @@ -320,16 +308,13 @@ func TestRPCEncoding(t *testing.T) { EventType: "someEventType", Value: ref(123.), Properties: map[string]interface{}{"a": 1}, - Timestamp: 123456, } encodedTrA := tra.Encode() assert.Equal(t, tra.Key, encodedTrA[TrackArgKeyIdx].(string)) assert.Equal(t, tra.TrafficType, encodedTrA[TrackArgTrafficTypeIdx].(string)) assert.Equal(t, tra.EventType, encodedTrA[TrackArgEventTypeIdx].(string)) - assert.Equal(t, *tra.Value, *encodedTrA[TrackArgValueIdx].(*float64)) + assert.Equal(t, *tra.Value, encodedTrA[TrackArgValueIdx].(float64)) assert.Equal(t, tra.Properties, encodedTrA[TrackArgPropertiesIdx].(map[string]interface{})) - assert.Equal(t, tra.Timestamp, encodedTrA[TrackArgTimestampIdx].(int64)) - } func ref[T any](t T) *T { diff --git a/splitio/link/service/v1/clientmgr.go b/splitio/link/service/v1/clientmgr.go index b74f679..f19cd67 100644 --- a/splitio/link/service/v1/clientmgr.go +++ b/splitio/link/service/v1/clientmgr.go @@ -5,6 +5,7 @@ import ( "fmt" "io" "os" + "runtime/debug" "github.com/splitio/go-toolkit/v5/logging" @@ -41,6 +42,7 @@ func (m *ClientManager) Manage() { defer func() { if r := recover(); r != nil { m.logger.Error("CRITICAL - connection handler is panicking: ", r) + m.logger.Error(string(debug.Stack())) } }() err := m.handleClientInteractions() @@ -131,6 +133,13 @@ func (m *ClientManager) handleRPC(rpc *protov1.RPC) (interface{}, error) { return nil, fmt.Errorf("error parsing treatments arguments: %w", err) } return m.handleGetTreatments(&args) + case protov1.OCTrack: + var args protov1.TrackArgs + if err := args.PopulateFromRPC(rpc); err != nil { + return nil, fmt.Errorf("error parsing track argumentts: %w", err) + } + return m.handleTrack(&args) + } return nil, fmt.Errorf("RPC not implemented") } @@ -199,3 +208,17 @@ func (m *ClientManager) handleGetTreatments(args *protov1.TreatmentsArgs) (inter return response, nil } + +func (m *ClientManager) handleTrack(args *protov1.TrackArgs) (interface{}, error) { + err := m.splitSDK.Track(m.clientConfig, args.Key, args.TrafficType, args.EventType, args.Value, args.Properties) + if err != nil && !errors.Is(err, sdk.ErrEventsQueueFull) { + return &protov1.ResponseWrapper[protov1.TreatmentPayload]{Status: protov1.ResultInternalError}, err + } + + response := &protov1.ResponseWrapper[protov1.TrackPayload]{ + Status: protov1.ResultOk, + Payload: protov1.TrackPayload{Success: err == nil}, // if err != nil it can only be ErrEventsQueueFull at this point + } + + return response, nil +} diff --git a/splitio/link/service/v1/clientmgr_test.go b/splitio/link/service/v1/clientmgr_test.go index 30970c7..b22a892 100644 --- a/splitio/link/service/v1/clientmgr_test.go +++ b/splitio/link/service/v1/clientmgr_test.go @@ -49,7 +49,7 @@ func TestRegisterAndTreatmentHappyPath(t *testing.T) { sdkMock := &sdkMocks.SDKMock{} sdkMock. On("Treatment", &types.ClientConfig{Metadata: types.ClientMetadata{ID: "someID", SdkVersion: "some_sdk-1.2.3"}}, "key", (*string)(nil), "someFeature", map[string]interface{}(nil)). - Return(&sdk.Result{Treatment: "on"}, nil).Once() + Return(&sdk.EvaluationResult{Treatment: "on"}, nil).Once() logger := logging.NewLogger(nil) cm := NewClientManager(rawConnMock, logger, sdkMock, serializerMock) @@ -83,7 +83,7 @@ func TestRegisterAndTreatmentsHappyPath(t *testing.T) { Args: []interface{}{"key", nil, []interface{}{"feat1", "feat2", "feat3"}, map[string]interface{}(nil)}, } }).Once() - serializerMock.On("Serialize", proto1Mocks.NewTreatmentsResp(true, []sdk.Result{ + serializerMock.On("Serialize", proto1Mocks.NewTreatmentsResp(true, []sdk.EvaluationResult{ {Treatment: "on"}, {Treatment: "off"}, {Treatment: "control"}, })).Return([]byte("successPayload"), nil).Once() @@ -96,7 +96,7 @@ func TestRegisterAndTreatmentsHappyPath(t *testing.T) { (*string)(nil), []string{"feat1", "feat2", "feat3"}, map[string]interface{}(nil), - ).Return(map[string]sdk.Result{ + ).Return(map[string]sdk.EvaluationResult{ "feat1": {Treatment: "on"}, "feat2": {Treatment: "off"}, "feat3": {Treatment: "control"}, @@ -142,7 +142,44 @@ func TestRegisterWithImpsAndTreatmentHappyPath(t *testing.T) { On("Treatment", &types.ClientConfig{Metadata: types.ClientMetadata{ID: "someID", SdkVersion: "some_sdk-1.2.3"}, ReturnImpressionData: true}, "key", (*string)(nil), "someFeature", map[string]interface{}(nil)). - Return(&sdk.Result{Treatment: "on", Impression: &dtos.Impression{Label: "l1", Time: 1234556}}, nil).Once() + Return(&sdk.EvaluationResult{Treatment: "on", Impression: &dtos.Impression{Label: "l1", Time: 1234556}}, nil).Once() + + logger := logging.NewLogger(nil) + cm := NewClientManager(rawConnMock, logger, sdkMock, serializerMock) + err := cm.handleClientInteractions() + assert.Nil(t, err) + rawConnMock.AssertNumberOfCalls(t, "Shutdown", 1) +} + +func TestTrack(t *testing.T) { + rawConnMock := &transferMocks.RawConnMock{} + rawConnMock.On("ReceiveMessage").Return([]byte("registrationMessage"), nil).Once() + rawConnMock.On("SendMessage", []byte("successRegistration")).Return(nil).Once() + rawConnMock.On("ReceiveMessage").Return([]byte("trackMessage"), nil).Once() + rawConnMock.On("SendMessage", []byte("successPayload")).Return(nil).Once() + rawConnMock.On("ReceiveMessage").Return([]byte(nil), io.EOF).Once() + rawConnMock.On("Shutdown").Return(nil).Once() + + serializerMock := &serializerMocks.SerializerMock{} + serializerMock.On("Parse", []byte("registrationMessage"), mock.Anything).Return(nil).Run(func(args mock.Arguments) { + *args.Get(1).(*v1.RPC) = v1.RPC{ + RPCBase: protocol.RPCBase{Version: protocol.V1}, + OpCode: v1.OCRegister, + Args: []interface{}{"someID", "some_sdk-1.2.3", uint64(0)}, + } + }).Once() + serializerMock.On("Serialize", proto1Mocks.NewRegisterResp(true)).Return([]byte("successRegistration"), nil).Once() + serializerMock.On("Parse", []byte("trackMessage"), mock.Anything).Return(nil).Run(func(args mock.Arguments) { + *args.Get(1).(*v1.RPC) = *proto1Mocks.NewTrackRPC("key1", "user", "checkin", ref(2.75), map[string]interface{}{"a": 1}) + }).Once() + serializerMock.On("Serialize", proto1Mocks.NewTrackResp(true)).Return([]byte("successPayload"), nil).Once() + + sdkMock := &sdkMocks.SDKMock{} + sdkMock. + On("Track", + &types.ClientConfig{Metadata: types.ClientMetadata{ID: "someID", SdkVersion: "some_sdk-1.2.3"}}, + "key1", "user", "checkin", ref(float64(2.75)), map[string]interface{}{"a": 1}). + Return((error)(nil)).Once() logger := logging.NewLogger(nil) cm := NewClientManager(rawConnMock, logger, sdkMock, serializerMock) @@ -194,6 +231,7 @@ func TestManagePanicRecovers(t *testing.T) { logger := &loggerMock{} logger.On("Error", "CRITICAL - connection handler is panicking: ", "some panic").Once() + logger.On("Error", mock.AnythingOfType("string")).Once() serializerMock := &serializerMocks.SerializerMock{} sdkMock := &sdkMocks.SDKMock{} @@ -267,8 +305,8 @@ func TestHandleRPCErrors(t *testing.T) { assert.Nil(t, res) assert.ErrorContains(t, err, "error parsing register arguments") - // set the config to allow other rpcs to be handled - cm.clientConfig = &types.ClientConfig{ReturnImpressionData: true} + // set the config to allow other rpcs to be handled + cm.clientConfig = &types.ClientConfig{ReturnImpressionData: true} // treatment wrong args res, err = cm.handleRPC(&v1.RPC{RPCBase: protocol.RPCBase{Version: protocol.V1}, OpCode: v1.OCTreatment, Args: []interface{}{1, "hola"}}) @@ -289,4 +327,8 @@ func (m *loggerMock) Info(msg ...interface{}) { m.Called(msg...) } func (m *loggerMock) Verbose(msg ...interface{}) { m.Called(msg...) } func (m *loggerMock) Warning(msg ...interface{}) { m.Called(msg...) } +func ref[T any](t T) *T { + return &t +} + var _ logging.LoggerInterface = (*loggerMock)(nil) diff --git a/splitio/logging/helpers.go b/splitio/logging/helpers.go new file mode 100644 index 0000000..b872edc --- /dev/null +++ b/splitio/logging/helpers.go @@ -0,0 +1,53 @@ +package logging + +import ( + "fmt" + "io" + "os" + + "github.com/splitio/go-toolkit/v5/logging" +) + +const ( + defaultMaxFiles = 10 + defaultMaxFileSize = 1024 * 1024 // 1M +) + +var defaultWriter = os.Stdout + +func GetWriter(source *string, maxFiles *int, maxFileSize *int) (io.Writer, error) { + if source == nil { + return defaultWriter, nil + } + + switch *source { + case "stdout", "/dev/stdout": + return os.Stdout, nil + case "stderr", "/dev/stderr": + return os.Stderr, nil + default: + // assume it's a regular file + if maxFiles == nil && maxFileSize == nil { + fileWriter, err := os.OpenFile(*source, os.O_CREATE|os.O_WRONLY|os.O_APPEND, 0644) + if err != nil { + return nil, fmt.Errorf("error creating log-output file: %w", err) + } + return fileWriter, nil + } + + mf := valueOr(maxFiles, defaultMaxFiles) + mfs := valueOr(maxFileSize, defaultMaxFileSize) + return logging.NewFileRotate(&logging.FileRotateOptions{ + MaxBytes: int64(mfs), + BackupCount: mf, + Path: *source, + }) + } +} + +func valueOr[T any](t *T, fallback T) T { + if t == nil { + return fallback + } + return *t +} diff --git a/splitio/logging/helpers_test.go b/splitio/logging/helpers_test.go new file mode 100644 index 0000000..64016fb --- /dev/null +++ b/splitio/logging/helpers_test.go @@ -0,0 +1,95 @@ +package logging + +import ( + "io" + "io/ioutil" + "os" + "strings" + "testing" + "time" + + "github.com/stretchr/testify/assert" +) + +func TestGetWriterStd(t *testing.T) { + // stdout/stderr + writer, err := GetWriter(ref("stdout"), nil, nil) + assert.Nil(t, err) + assert.Equal(t, os.Stdout, writer) + writer, err = GetWriter(ref("/dev/stdout"), nil, nil) + assert.Nil(t, err) + assert.Equal(t, os.Stdout, writer) + writer, err = GetWriter(ref("stderr"), nil, nil) + assert.Nil(t, err) + assert.Equal(t, os.Stderr, writer) + writer, err = GetWriter(ref("/dev/stderr"), nil, nil) + assert.Nil(t, err) + assert.Equal(t, os.Stderr, writer) +} + +func TestGetWriterFileNoRotation(t *testing.T) { + // regular file, no rotation + fn := strings.Join([]string{os.TempDir(), "someLogFile2"}, string(os.PathSeparator)) + defer func() { + if err := os.Remove(fn); err != nil { + t.Error("error deleting file: ", err) + } + }() + writer, err := GetWriter(&fn, nil, nil) + assert.Nil(t, err) + + writer.Write([]byte("hola que tal")) + writer.(io.Closer).Close() + + assertFileContents(t, "hola que tal", fn) +} + +func TestGetWriterFileRotation(t *testing.T) { + + // remove any old file + fl, err := os.ReadDir(os.TempDir()) + assert.Nil(t, err) + for _, fe := range fl { + if strings.HasPrefix(fe.Name(), "someLogFile") { + os.Remove(fe.Name()) + } + } + + + fn := strings.Join([]string{os.TempDir(), "someLogFile"}, string(os.PathSeparator)) + writer, err := GetWriter(&fn, ref(2), ref(5)) + assert.Nil(t, err) + + writer.Write([]byte("12345")) + writer.Write([]byte("67890")) + writer.Write([]byte("qwert")) + writer.Write([]byte("asdfg")) + writer.Write([]byte("zxcvb")) + writer.Write([]byte("hjaiu")) + + + time.Sleep(1*time.Second) // file rotate writer is async.. give it a second before looking at the fs + fl, err = os.ReadDir(os.TempDir()) + assert.Nil(t, err) + names := make([]string, 0, 3) + for _, fe := range fl { + if strings.HasPrefix(fe.Name(), "someLogFile") { + names = append(names, fe.Name()) + } + } + + assert.Contains(t, names, "someLogFile") + assert.Contains(t, names, "someLogFile.1") + assert.Contains(t, names, "someLogFile.2") + assert.NotContains(t, names, "someLogFile.3") +} + +func assertFileContents(t *testing.T, expected string, fn string) { + t.Helper() + contents, err := ioutil.ReadFile(fn) + assert.Nil(t, err) + assert.Equal(t, expected, string(contents)) +} +func ref[T any](t T) *T { + return &t +} diff --git a/splitio/sdk/conf/conf.go b/splitio/sdk/conf/conf.go index 2f1bbe8..411fc30 100644 --- a/splitio/sdk/conf/conf.go +++ b/splitio/sdk/conf/conf.go @@ -6,12 +6,18 @@ import ( "github.com/splitio/go-split-commons/v4/conf" ) +const ( + defaultImpressionsMode = "optimized" + minimumImpressionsRefreshRate = 30 * time.Minute +) + type Config struct { LabelsEnabled bool StreamingEnabled bool Splits Splits Segments Segments Impressions Impressions + Events Events URLs URLs } @@ -36,6 +42,12 @@ type Impressions struct { PostConcurrency int } +type Events struct { + QueueSize int + SyncPeriod time.Duration + PostConcurrency int +} + type URLs struct { Auth string SDK string @@ -46,14 +58,23 @@ type URLs struct { func (c *Config) ToAdvancedConfig() *conf.AdvancedConfig { d := conf.GetDefaultAdvancedConfig() + d.SplitsRefreshRate = int(c.Splits.SyncPeriod.Seconds()) + d.SplitUpdateQueueSize = int64(c.Splits.UpdateBufferSize) d.SegmentsRefreshRate = int(c.Segments.SyncPeriod.Seconds()) + d.SegmentQueueSize = c.Segments.QueueSize + d.SegmentUpdateQueueSize = int64(c.Segments.UpdateBufferSize) + d.SegmentWorkers = c.Segments.WorkerCount d.StreamingEnabled = c.StreamingEnabled + d.AuthServiceURL = c.URLs.Auth d.SdkURL = c.URLs.SDK d.EventsURL = c.URLs.Events d.StreamingServiceURL = c.URLs.Streaming d.TelemetryServiceURL = c.URLs.Telemetry + + d.ImpressionsQueueSize = c.Impressions.QueueSize + return &d } @@ -75,10 +96,15 @@ func DefaultConfig() *Config { Mode: "optimized", ObserverSize: 500000, QueueSize: 8192, - SyncPeriod: 5 * time.Second, - CountSyncPeriod: 5 * time.Second, + SyncPeriod: 30 * time.Minute, + CountSyncPeriod: 60 * time.Minute, PostConcurrency: 1, }, + Events: Events{ + QueueSize: 8192, + SyncPeriod: 1*time.Minute, + PostConcurrency: 1, + }, URLs: URLs{ Auth: "https://auth.split.io", SDK: "https://sdk.split.io/api", @@ -88,3 +114,13 @@ func DefaultConfig() *Config { }, } } + +func (c *Config) Normalize() []string { + var warnings []string + if c.Impressions.Mode == "optimized" && c.Impressions.SyncPeriod < minimumImpressionsRefreshRate { + warnings = append(warnings, "minimum impressions refresh rate is 30 min. ignoring user config") + c.Impressions.SyncPeriod = minimumImpressionsRefreshRate + } + + return warnings +} diff --git a/splitio/sdk/conf/conf_test.go b/splitio/sdk/conf/conf_test.go new file mode 100644 index 0000000..f3e87ea --- /dev/null +++ b/splitio/sdk/conf/conf_test.go @@ -0,0 +1,35 @@ +package conf + +import ( + "testing" + "time" + + "github.com/stretchr/testify/assert" +) + +func TestSDKConf(t *testing.T) { + dc := DefaultConfig() + dc.Impressions.SyncPeriod = 1 * time.Minute + warns := dc.Normalize() + assert.Equal(t, warns, []string{"minimum impressions refresh rate is 30 min. ignoring user config"}) + assert.Equal(t, 30*time.Minute, dc.Impressions.SyncPeriod) + + adv := dc.ToAdvancedConfig() + assert.Equal(t, 30, adv.HTTPTimeout) + assert.Equal(t, dc.Segments.QueueSize, adv.SegmentQueueSize) + assert.Equal(t, dc.Segments.WorkerCount, adv.SegmentWorkers) + assert.Equal(t, dc.URLs.SDK, adv.SdkURL) + assert.Equal(t, dc.URLs.Events, adv.EventsURL) + assert.Equal(t, dc.URLs.Telemetry, adv.TelemetryServiceURL) + // assert.Equal(t, TODO, adv.EventsBulkSize) + // assert.Equal(t, TODO, adv.EventsQueueSize) + assert.Equal(t, dc.Impressions.QueueSize, adv.ImpressionsQueueSize) + // assert.Equal(t, TODO, adv.ImpressionsBulkSize) + assert.Equal(t, dc.StreamingEnabled, adv.StreamingEnabled) + assert.Equal(t, dc.URLs.Auth, adv.AuthServiceURL) + assert.Equal(t, dc.URLs.Streaming, adv.StreamingServiceURL) + assert.Equal(t, int64(dc.Splits.UpdateBufferSize), adv.SplitUpdateQueueSize) + assert.Equal(t, int64(dc.Segments.UpdateBufferSize), adv.SegmentUpdateQueueSize) + assert.Equal(t, int(dc.Splits.SyncPeriod.Seconds()), adv.SplitsRefreshRate) + assert.Equal(t, int(dc.Segments.SyncPeriod.Seconds()), adv.SegmentsRefreshRate) +} diff --git a/splitio/sdk/helpers.go b/splitio/sdk/helpers.go index 2f4de7d..17b2332 100644 --- a/splitio/sdk/helpers.go +++ b/splitio/sdk/helpers.go @@ -31,13 +31,14 @@ func setupWorkers( api *api.SplitAPI, str *storages, hc application.MonitorProducerInterface, - cfg *sdkConf.Impressions, + cfg *sdkConf.Config, ) *synchronizer.Workers { return &synchronizer.Workers{ SplitFetcher: split.NewSplitFetcher(str.splits, api.SplitFetcher, logger, str.telemetry, hc), SegmentFetcher: segment.NewSegmentFetcher(str.splits, str.segments, api.SegmentFetcher, logger, str.telemetry, hc), - ImpressionRecorder: workers.NewImpressionsWorker(logger, str.telemetry, api.ImpressionRecorder, str.impressions, cfg), + ImpressionRecorder: workers.NewImpressionsWorker(logger, str.telemetry, api.ImpressionRecorder, str.impressions, &cfg.Impressions), + EventRecorder: workers.NewEventsWorker(logger, str.telemetry, api.EventRecorder, str.events, &cfg.Events), } } @@ -51,22 +52,33 @@ func setupTasks( api *api.SplitAPI, ) *synchronizer.SplitTasks { impCfg := cfg.Impressions - return &synchronizer.SplitTasks{ - SplitSyncTask: tasks.NewFetchSplitsTask(workers.SplitFetcher, int(cfg.Splits.SyncPeriod.Seconds()), logger), - SegmentSyncTask: tasks.NewFetchSegmentsTask(workers.SegmentFetcher, int(cfg.Segments.SyncPeriod.Seconds()), cfg.Segments.WorkerCount, cfg.Segments.QueueSize, logger), - ImpressionSyncTask: tasks.NewRecordImpressionsTask(workers.ImpressionRecorder, int(impCfg.SyncPeriod.Seconds()), logger, 5000), - //ImpressionSyncTask: tss.NewImpressionSyncTask(workers.ImpressionRecorder, logger, cfg.Impressions), - ImpressionsCountSyncTask: tasks.NewRecordImpressionsCountTask( - impressionscount.NewRecorderSingle(impComponents.counter, api.ImpressionRecorder, md, logger, str.telemetry), + evCfg := cfg.Events + tg := &synchronizer.SplitTasks{ + SplitSyncTask: tasks.NewFetchSplitsTask(workers.SplitFetcher, int(cfg.Splits.SyncPeriod.Seconds()), logger), + SegmentSyncTask: tasks.NewFetchSegmentsTask( + workers.SegmentFetcher, + int(cfg.Segments.SyncPeriod.Seconds()), + cfg.Segments.WorkerCount, + cfg.Segments.QueueSize, logger, - int(impCfg.CountSyncPeriod.Seconds()), ), + ImpressionSyncTask: tasks.NewRecordImpressionsTask(workers.ImpressionRecorder, int(impCfg.SyncPeriod.Seconds()), logger, 5000), + EventSyncTask: tasks.NewRecordEventsTask(workers.EventRecorder, 5000, int(evCfg.SyncPeriod.Seconds()), logger), TelemetrySyncTask: &NoOpTask{}, - EventSyncTask: &NoOpTask{}, UniqueKeysTask: &NoOpTask{}, CleanFilterTask: &NoOpTask{}, ImpsCountConsumerTask: &NoOpTask{}, } + + if impCfg.Mode == "optimized" { + tg.ImpressionsCountSyncTask = tasks.NewRecordImpressionsCountTask( + impressionscount.NewRecorderSingle(impComponents.counter, api.ImpressionRecorder, md, logger, str.telemetry), + logger, + int(impCfg.CountSyncPeriod.Seconds()), + ) + } + + return tg } type impComponents struct { @@ -103,16 +115,19 @@ type storages struct { segments storage.SegmentStorage telemetry storage.TelemetryStorage impressions *sss.ImpressionsStorage + events *sss.EventsStorage } func setupStorages(cfg *sdkConf.Config) *storages { ts, _ := inmemory.NewTelemetryStorage() iq, _ := sss.NewImpressionsQueue(cfg.Impressions.QueueSize) + eq, _ := sss.NewEventsQueue(cfg.Events.QueueSize) return &storages{ splits: mutexmap.NewMMSplitStorage(), segments: mutexmap.NewMMSegmentStorage(), impressions: iq, + events: eq, telemetry: ts, } } diff --git a/splitio/sdk/mocks/sdk.go b/splitio/sdk/mocks/sdk.go index 8d3059e..e53a180 100644 --- a/splitio/sdk/mocks/sdk.go +++ b/splitio/sdk/mocks/sdk.go @@ -17,9 +17,9 @@ func (m *SDKMock) Treatment( bucketingKey *string, feature string, attributes map[string]interface{}, -) (*sdk.Result, error) { +) (*sdk.EvaluationResult, error) { args := m.Called(md, key, bucketingKey, feature, attributes) - return args.Get(0).(*sdk.Result), nil + return args.Get(0).(*sdk.EvaluationResult), args.Error(1) } // Treatments implements sdk.Interface @@ -29,10 +29,15 @@ func (m *SDKMock) Treatments( bucketingKey *string, features []string, attributes map[string]interface{}, -) (map[string]sdk.Result, error) { +) (map[string]sdk.EvaluationResult, error) { args := m.Called(md, key, bucketingKey, features, attributes) - return args.Get(0).(map[string]sdk.Result), nil + return args.Get(0).(map[string]sdk.EvaluationResult), args.Error(1) } +// Track implements sdk.Interface +func (m *SDKMock) Track(cfg *types.ClientConfig, key string, trafficType string, eventType string, value *float64, properties map[string]interface{}) error { + args := m.Called(cfg, key, trafficType, eventType, value, properties) + return args.Error(0) +} var _ sdk.Interface = (*SDKMock)(nil) diff --git a/splitio/sdk/results.go b/splitio/sdk/results.go index c697566..546a454 100644 --- a/splitio/sdk/results.go +++ b/splitio/sdk/results.go @@ -2,7 +2,7 @@ package sdk import "github.com/splitio/go-split-commons/v4/dtos" -type Result struct { +type EvaluationResult struct { Treatment string Impression *dtos.Impression Config *string diff --git a/splitio/sdk/sdk.go b/splitio/sdk/sdk.go index 056b633..e9fd470 100644 --- a/splitio/sdk/sdk.go +++ b/splitio/sdk/sdk.go @@ -1,6 +1,7 @@ package sdk import ( + "errors" "fmt" "time" @@ -24,13 +25,19 @@ import ( const ( impressionsFullNotif = "IMPRESSIONS_FULL" + eventsFullNotif = "EVENTS_FULL" +) + +var ( + ErrEventsQueueFull = errors.New("events queue full") ) type Attributes = map[string]interface{} type Interface interface { - Treatment(cfg *types.ClientConfig, Key string, BucketingKey *string, Feature string, attributes map[string]interface{}) (*Result, error) - Treatments(cfg *types.ClientConfig, Key string, BucketingKey *string, Features []string, attributes map[string]interface{}) (map[string]Result, error) + Treatment(cfg *types.ClientConfig, key string, bucketingKey *string, feature string, attributes map[string]interface{}) (*EvaluationResult, error) + Treatments(cfg *types.ClientConfig, key string, bucketingKey *string, features []string, attributes map[string]interface{}) (map[string]EvaluationResult, error) + Track(cfg *types.ClientConfig, key string, trafficType string, eventType string, value *float64, properties map[string]interface{}) error } type Impl struct { @@ -39,14 +46,22 @@ type Impl struct { sm synchronizer.Manager ss synchronizer.Synchronizer is *storage.ImpressionsStorage + es *storage.EventsStorage iq provisional.ImpressionManager cfg sdkConf.Config status chan int queueFullChan chan string + validator Validator } func New(logger logging.LoggerInterface, apikey string, c *conf.Config) (*Impl, error) { + if warnings := c.Normalize(); len(warnings) > 0 { + for _, w := range warnings { + logger.Warning(w) + } + } + md := dtos.Metadata{SDKVersion: fmt.Sprintf("splitd-%s", splitio.Version)} advCfg := c.ToAdvancedConfig() @@ -58,9 +73,9 @@ func New(logger logging.LoggerInterface, apikey string, c *conf.Config) (*Impl, hc := &application.Dummy{} - queueFullChan := make(chan string, 1) // Only one item so that we don't queue N flushes (which makes no sense) if we're getting hit too hard + queueFullChan := make(chan string, 2) splitApi := api.NewSplitAPI(apikey, *advCfg, logger, md) - workers := setupWorkers(logger, splitApi, stores, hc, &c.Impressions) + workers := setupWorkers(logger, splitApi, stores, hc, c) tasks := setupTasks(c, stores, logger, workers, impc, md, splitApi) sync := synchronizer.NewSynchronizer(*advCfg, *tasks, *workers, logger, queueFullChan, nil) @@ -83,21 +98,23 @@ func New(logger logging.LoggerInterface, apikey string, c *conf.Config) (*Impl, ss: sync, ev: evaluator.NewEvaluator(stores.splits, stores.segments, engine.NewEngine(logger), logger), is: stores.impressions, + es: stores.events, iq: impc.manager, cfg: *c, queueFullChan: queueFullChan, + validator: Validator{logger: logger, splits: stores.splits}, }, nil } // Treatment implements Interface -func (i *Impl) Treatment(cfg *types.ClientConfig, key string, bk *string, feature string, attributes Attributes) (*Result, error) { +func (i *Impl) Treatment(cfg *types.ClientConfig, key string, bk *string, feature string, attributes Attributes) (*EvaluationResult, error) { res := i.ev.EvaluateFeature(key, bk, feature, attributes) if res == nil { return nil, fmt.Errorf("nil result") } imp := i.handleImpression(key, bk, feature, res, cfg.Metadata) - return &Result{ + return &EvaluationResult{ Treatment: res.Treatment, Impression: imp, Config: res.Config, @@ -105,19 +122,19 @@ func (i *Impl) Treatment(cfg *types.ClientConfig, key string, bk *string, featur } // Treatment implements Interface -func (i *Impl) Treatments(cfg *types.ClientConfig, key string, bk *string, features []string, attributes Attributes) (map[string]Result, error) { +func (i *Impl) Treatments(cfg *types.ClientConfig, key string, bk *string, features []string, attributes Attributes) (map[string]EvaluationResult, error) { res := i.ev.EvaluateFeatures(key, bk, features, attributes) - toRet := make(map[string]Result, len(res.Evaluations)) + toRet := make(map[string]EvaluationResult, len(res.Evaluations)) for _, feature := range features { curr, ok := res.Evaluations[feature] if !ok { - toRet[feature] = Result{Treatment: "control"} + toRet[feature] = EvaluationResult{Treatment: "control"} continue } - var eres Result + var eres EvaluationResult eres.Treatment = curr.Treatment eres.Impression = i.handleImpression(key, bk, feature, &curr, cfg.Metadata) eres.Config = curr.Config @@ -127,6 +144,44 @@ func (i *Impl) Treatments(cfg *types.ClientConfig, key string, bk *string, featu return toRet, nil } +func (i *Impl) Track(cfg *types.ClientConfig, key string, trafficType string, eventType string, value *float64, properties map[string]interface{}) error { + + // TODO(mredolatti): validate traffic type & truncate properties if needed + trafficType, err := i.validator.validateTrafficType(trafficType) + if err != nil { + return err + } + + properties, _, err = i.validator.validateTrackProperties(properties) + if err != nil { + return err + } + + event := &dtos.EventDTO{ + Key: key, + TrafficTypeName: trafficType, + EventTypeID: eventType, + Value: value, + Timestamp: timeMillis(), + Properties: properties, + } + + _, err = i.es.Push(cfg.Metadata, *event) + if err != nil { + if err == storage.ErrQueueFull { + select { + case i.queueFullChan <- eventsFullNotif: + default: + i.logger.Warning("events queue has filled up and is currently performing a flush. Current event will be dropped") + } + return ErrEventsQueueFull + } + i.logger.Error("error handling event: ", err) + return err + } + return nil +} + func (i *Impl) handleImpression(key string, bk *string, f string, r *evaluator.Result, cm types.ClientMetadata) *dtos.Impression { var label string if i.cfg.LabelsEnabled { diff --git a/splitio/sdk/sdk_test.go b/splitio/sdk/sdk_test.go index 1001c77..6ef20f6 100644 --- a/splitio/sdk/sdk_test.go +++ b/splitio/sdk/sdk_test.go @@ -2,6 +2,7 @@ package sdk import ( "fmt" + "strings" "testing" "time" @@ -9,8 +10,10 @@ import ( "github.com/splitio/go-split-commons/v4/dtos" "github.com/splitio/go-split-commons/v4/provisional" "github.com/splitio/go-split-commons/v4/service" + scstorage "github.com/splitio/go-split-commons/v4/storage" "github.com/splitio/go-split-commons/v4/storage/inmemory" "github.com/splitio/go-split-commons/v4/synchronizer" + "github.com/splitio/go-toolkit/v5/datastructures/set" "github.com/splitio/go-toolkit/v5/logging" "github.com/splitio/splitd/splitio/sdk/conf" "github.com/splitio/splitd/splitio/sdk/storage" @@ -279,6 +282,116 @@ func TestImpressionsQueueFull(t *testing.T) { assert.Equal(t, 1, totalSize) // assert no more impressions in queue } +func TestTrack(t *testing.T) { + + es, _ := storage.NewEventsQueue(1000) + logger := logging.NewLogger(nil) + + ss := &SplitStorageMock{} + ss.On("TrafficTypeExists", "user").Return(true) + + client := &Impl{ + logger: logging.NewLogger(nil), + es: es, + cfg: conf.Config{LabelsEnabled: false}, + validator: Validator{logger, ss}, + } + + md := types.ClientConfig{Metadata: types.ClientMetadata{ID: "some", SdkVersion: "go-1.2.3"}} + + err := client.Track(&md, "key1", "user", "checkin", ref(123.4), map[string]interface{}{"a": 123}) + assert.Nil(t, err) + + err = es.RangeAndClear(func(md types.ClientMetadata, st *storage.LockingQueue[dtos.EventDTO]) { + assert.Equal(t, types.ClientMetadata{ID: "some", SdkVersion: "go-1.2.3"}, md) + assert.Equal(t, 1, st.Len()) + + var evs []dtos.EventDTO + n, err := st.Pop(1, &evs) + assert.Nil(t, nil) + assert.Equal(t, 1, n) + assert.Equal(t, 1, len(evs)) + assertEventEq(t, &dtos.EventDTO{ + Key: "key1", + TrafficTypeName: "user", + EventTypeID: "checkin", + Value: ref(123.4), + Properties: map[string]interface{}{"a": 123}, + }, &evs[0]) + n, err = st.Pop(1, &evs) + assert.ErrorIs(t, err, storage.ErrQueueEmpty) + + }) + assert.Nil(t, err) + + err = client.Track(&md, "key1", "", "checkin", ref(123.4), map[string]interface{}{"a": 123}) + assert.ErrorIs(t, err, ErrEmtpyTrafficType) + + err = client.Track(&md, "key1", "user", "checkin", ref(123.4), map[string]interface{}{"a": strings.Repeat("qwertyui", 100000)}) + assert.ErrorIs(t, err, ErrEventTooBig) + +} + +func TestTrackEventsFlush(t *testing.T) { + + es, _ := storage.NewEventsQueue(4) + logger := logging.NewLogger(nil) + + ss := &SplitStorageMock{} + ss.On("TrafficTypeExists", "user").Return(true) + + client := &Impl{ + logger: logging.NewLogger(nil), + queueFullChan: make(chan string, 2), + es: es, + cfg: conf.Config{LabelsEnabled: false}, + validator: Validator{logger, ss}, + } + + md := types.ClientConfig{Metadata: types.ClientMetadata{ID: "some", SdkVersion: "go-1.2.3"}} + + err := client.Track(&md, "key1", "user", "checkin", ref(123.4), map[string]interface{}{"a": 123}) + assert.Nil(t, err) + err = client.Track(&md, "key2", "user", "checkin", ref(123.4), map[string]interface{}{"a": 123}) + assert.Nil(t, err) + err = client.Track(&md, "key3", "user", "checkin", ref(123.4), map[string]interface{}{"a": 123}) + assert.Nil(t, err) + err = client.Track(&md, "key4", "user", "checkin", ref(123.4), map[string]interface{}{"a": 123}) + assert.ErrorIs(t, err, ErrEventsQueueFull) + + assert.Equal(t, "EVENTS_FULL", <-client.queueFullChan) + /* + err = es.RangeAndClear(func(md types.ClientMetadata, st *storage.LockingQueue[dtos.EventDTO]) { + assert.Equal(t, types.ClientMetadata{ID: "some", SdkVersion: "go-1.2.3"}, md) + assert.Equal(t, 1, st.Len()) + + var evs []dtos.EventDTO + n, err := st.Pop(1, &evs) + assert.Nil(t, nil) + assert.Equal(t, 1, n) + assert.Equal(t, 1, len(evs)) + assertEventEq(t, &dtos.EventDTO{ + Key: "key1", + TrafficTypeName: "user", + EventTypeID: "checkin", + Value: ref(123.4), + Properties: map[string]interface{}{"a": 123}, + }, &evs[0]) + n, err = st.Pop(1, &evs) + assert.ErrorIs(t, err, storage.ErrQueueEmpty) + + }) + assert.Nil(t, err) + + err = client.Track(&md, "key1", "", "checkin", ref(123.4), map[string]interface{}{"a": 123}) + assert.ErrorIs(t, err, ErrEmtpyTrafficType) + + err = client.Track(&md, "key1", "user", "checkin", ref(123.4), map[string]interface{}{"a": strings.Repeat("qwertyui", 100000)}) + assert.ErrorIs(t, err, ErrEventTooBig) + */ + +} + func assertImpEq(t *testing.T, i1, i2 *dtos.Impression) { t.Helper() assert.Equal(t, i1.KeyName, i2.KeyName) @@ -289,6 +402,15 @@ func assertImpEq(t *testing.T, i1, i2 *dtos.Impression) { assert.Equal(t, i1.ChangeNumber, i2.ChangeNumber) } +func assertEventEq(t *testing.T, e1, e2 *dtos.EventDTO) { + t.Helper() + assert.Equal(t, e1.Key, e2.Key) + assert.Equal(t, e1.TrafficTypeName, e2.TrafficTypeName) + assert.Equal(t, e1.EventTypeID, e2.EventTypeID) + assert.Equal(t, e1.Value, e2.Value) + assert.Equal(t, e1.Properties, e2.Properties) +} + // mocks type EvaluatorMock struct { @@ -339,6 +461,28 @@ func (m *ImpressionRecorderMock) RecordImpressionsCount(pf dtos.ImpressionsCount return args.Error(0) } +type SplitStorageMock struct{ mock.Mock } + +func (m *SplitStorageMock) All() []dtos.SplitDTO { panic("unimplemented") } +func (m *SplitStorageMock) ChangeNumber() (int64, error) { panic("unimplemented") } +func (m *SplitStorageMock) FetchMany([]string) map[string]*dtos.SplitDTO { panic("unimplemented") } +func (m *SplitStorageMock) KillLocally(string, string, int64) { panic("unimplemented") } +func (m *SplitStorageMock) SegmentNames() *set.ThreadUnsafeSet { panic("unimplemented") } +func (m *SplitStorageMock) SetChangeNumber(changeNumber int64) error { panic("unimplemented") } +func (m *SplitStorageMock) Split(splitName string) *dtos.SplitDTO { panic("unimplemented") } +func (m *SplitStorageMock) SplitNames() []string { panic("unimplemented") } +func (m *SplitStorageMock) Update([]dtos.SplitDTO, []dtos.SplitDTO, int64) { panic("unimplemented") } + +func (m *SplitStorageMock) TrafficTypeExists(trafficType string) bool { + args := m.Called(trafficType) + return args.Bool(0) +} + +func ref[T any](t T) *T { + return &t +} + +var _ scstorage.SplitStorage = (*SplitStorageMock)(nil) var _ evaluator.Interface = (*EvaluatorMock)(nil) var _ provisional.ImpressionManager = (*ImpressionManagerMock)(nil) var _ service.ImpressionsRecorder = (*ImpressionRecorderMock)(nil) diff --git a/splitio/sdk/tasks/events.go b/splitio/sdk/tasks/events.go new file mode 100644 index 0000000..53e5278 --- /dev/null +++ b/splitio/sdk/tasks/events.go @@ -0,0 +1,29 @@ +package tasks + +import ( + "github.com/splitio/go-toolkit/v5/asynctask" + "github.com/splitio/go-toolkit/v5/logging" + sdkconf "github.com/splitio/splitd/splitio/sdk/conf" + "github.com/splitio/splitd/splitio/sdk/workers" +) + +const ( + defaultEventsBulkSize = 5000 +) + +func NewEventsSyncTask( + worker *workers.MultiMetaEventsWorker, + logger logging.LoggerInterface, + cfg *sdkconf.Impressions, +) *asynctask.AsyncTask { + + // TODO(mredolatti): pass a proper bulk size (currently ignored, everything is flushed) + return asynctask.NewAsyncTask( + "events-sender", + func(logging.LoggerInterface) error { worker.SynchronizeEvents(defaultEventsBulkSize); return nil }, + int(cfg.SyncPeriod.Seconds()), + nil, + func(logging.LoggerInterface) { worker.SynchronizeEvents(defaultEventsBulkSize) }, + logger, + ) +} diff --git a/splitio/sdk/tasks/impressions.go b/splitio/sdk/tasks/impressions.go index 71bb17a..e37f876 100644 --- a/splitio/sdk/tasks/impressions.go +++ b/splitio/sdk/tasks/impressions.go @@ -8,7 +8,7 @@ import ( ) const ( - defaultBulkSize = 5000 + defaultImpressionsBulkSize = 5000 ) func NewImpressionSyncTask( @@ -20,10 +20,10 @@ func NewImpressionSyncTask( // TODO(mredolatti): pass a proper bulk size (currently ignored, everything is flushed) return asynctask.NewAsyncTask( "impressions-sender", - func(logging.LoggerInterface) error { worker.SynchronizeImpressions(defaultBulkSize); return nil }, + func(logging.LoggerInterface) error { worker.SynchronizeImpressions(defaultImpressionsBulkSize); return nil }, int(cfg.SyncPeriod.Seconds()), nil, - func(logging.LoggerInterface) { worker.SynchronizeImpressions(defaultBulkSize) }, + func(logging.LoggerInterface) { worker.SynchronizeImpressions(defaultImpressionsBulkSize) }, logger, ) } diff --git a/splitio/sdk/validators.go b/splitio/sdk/validators.go new file mode 100644 index 0000000..529343f --- /dev/null +++ b/splitio/sdk/validators.go @@ -0,0 +1,71 @@ +package sdk + +import ( + "errors" + "strings" + + "github.com/splitio/go-split-commons/v4/storage" + "github.com/splitio/go-toolkit/v5/logging" +) + +// MaxEventLength constant to limit the event size +const MaxEventLength = 32768 + +var ErrEventTooBig = errors.New("The maximum size allowed for the properties is 32kb. Event not queued") +var ErrEmtpyTrafficType = errors.New("Traffic type cannot be empty") + +type Validator struct { + logger logging.LoggerInterface + splits storage.SplitStorage +} + +func (i *Validator) validateTrafficType(trafficType string) (string, error) { + if len(trafficType) == 0 { + return "", ErrEmtpyTrafficType + } + + toLower := strings.ToLower(trafficType) + if toLower != trafficType { + i.logger.Warning("Track: traffic type should be all lowercase - converting string to lowercase") + } + + if !i.splits.TrafficTypeExists(toLower) { + i.logger.Warning("Track: traffic type " + toLower + " does not have any corresponding feature flags in this environment, " + + "make sure you’re tracking your events to a valid traffic type defined in the Split user interface") + } + + return toLower, nil +} + +func (i *Validator) validateTrackProperties(properties map[string]interface{}) (map[string]interface{}, int, error) { + if len(properties) == 0 { + return nil, 0, nil + } + + if len(properties) > 300 { + i.logger.Warning("Track: Event has more than 300 properties. Some of them will be trimmed when processed") + } + + processed := make(map[string]interface{}) + size := 1024 // Average event size is ~750 bytes. Using 1kbyte as a starting point. + for name, value := range properties { + size += len(name) + switch value.(type) { + case int, int8, int16, int32, int64, uint, uint8, uint16, uint32, uint64, float32, float64, bool, nil: + processed[name] = value + case string: + asStr := value.(string) + size += len(asStr) + processed[name] = value + default: + i.logger.Warning("Property %s is of invalid type. Setting value to nil") + processed[name] = nil + } + + if size > MaxEventLength { + i.logger.Error("The maximum size allowed for the properties is 32kb. Event not queued") + return nil, size, ErrEventTooBig + } + } + return processed, size, nil +} diff --git a/splitio/sdk/workers/events.go b/splitio/sdk/workers/events.go new file mode 100644 index 0000000..d5c95b0 --- /dev/null +++ b/splitio/sdk/workers/events.go @@ -0,0 +1,94 @@ +package workers + +import ( + "errors" + "sync" + + sdkconf "github.com/splitio/splitd/splitio/sdk/conf" + sss "github.com/splitio/splitd/splitio/sdk/storage" + "github.com/splitio/splitd/splitio/sdk/types" + serrors "github.com/splitio/splitd/splitio/util/errors" + + "github.com/splitio/go-split-commons/v4/dtos" + "github.com/splitio/go-split-commons/v4/service" + "github.com/splitio/go-split-commons/v4/storage" + "github.com/splitio/go-split-commons/v4/synchronizer/worker/event" + "github.com/splitio/go-toolkit/v5/logging" + gtsync "github.com/splitio/go-toolkit/v5/sync" +) + +type MultiMetaEventsWorker struct { + logger logging.LoggerInterface + telemetry storage.TelemetryRuntimeProducer + llrec service.EventsRecorder + iq *sss.EventsStorage + cfg *sdkconf.Events + runnning gtsync.AtomicBool +} + +func NewEventsWorker( + logger logging.LoggerInterface, + telemetry storage.TelemetryRuntimeProducer, + llrec service.EventsRecorder, + iq *sss.EventsStorage, + cfg *sdkconf.Events, +) *MultiMetaEventsWorker { + return &MultiMetaEventsWorker{ + logger: logger, + telemetry: telemetry, + llrec: llrec, + iq: iq, + cfg: cfg, + } +} + +// FlushImpressions implements impression.ImpressionRecorder +// TODO(mredolatti): take `bulkSize` into account +func (m *MultiMetaEventsWorker) FlushEvents(bulkSize int64) error { + + // prevent 2 evictions from happening at the same time. we don't want a sync.Mutex since that would only cause 43928729 + // function calls to pile up and get called after each mutex release. + if !m.runnning.TestAndSet() { + m.logger.Warning("flush/sync requested while another one is in progress. ignoring") + return nil + } + defer m.runnning.Unset() + + var errs serrors.ConcurrentErrorCollector + var wg sync.WaitGroup + + // same logic as impressions workers, without the need for formatting. check impressions.go for a better + // description of what's being done + if err := m.iq.RangeAndClear(func(md types.ClientMetadata, q *sss.LockingQueue[dtos.EventDTO]) { + extracted := make([]dtos.EventDTO, 0, q.Len()) + n, err := q.Pop(q.Len(), &extracted) + if err != nil && !errors.Is(err, sss.ErrQueueEmpty) { + m.logger.Error("error fetching items from queue: ", err) + return // continue with queue + } + + if n == 0 { + return // nothing to do here + } + + wg.Add(1) + go func(events []dtos.EventDTO, cc types.ClientMetadata) { + defer wg.Done() + if err := m.llrec.Record(events, dtos.Metadata{SDKVersion: cc.SdkVersion}); err != nil { + errs.Append(err) + } + }(extracted, md) + }); err != nil { + m.logger.Error("error traversing event queues: ", err) + } + + wg.Wait() + return errs.Join() +} + +// SynchronizeImpressions implements impression.ImpressionRecorder +func (m *MultiMetaEventsWorker) SynchronizeEvents(bulkSize int64) error { + return m.FlushEvents(bulkSize) +} + +var _ event.EventRecorder = (*MultiMetaEventsWorker)(nil) diff --git a/splitio/sdk/workers/events_test.go b/splitio/sdk/workers/events_test.go new file mode 100644 index 0000000..3bc83ea --- /dev/null +++ b/splitio/sdk/workers/events_test.go @@ -0,0 +1,115 @@ +package workers + +import ( + "testing" + "time" + + "github.com/splitio/go-split-commons/v4/dtos" + "github.com/splitio/go-split-commons/v4/service" + "github.com/splitio/go-split-commons/v4/storage/inmemory" + "github.com/splitio/go-toolkit/v5/logging" + "github.com/splitio/splitd/splitio/sdk/conf" + sss "github.com/splitio/splitd/splitio/sdk/storage" + "github.com/splitio/splitd/splitio/sdk/types" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/mock" +) + +func TestEventsTask(t *testing.T) { + is, _ := sss.NewEventsQueue(100) + ts, _ := inmemory.NewTelemetryStorage() + logger := logging.NewLogger(nil) + rec := &EventsRecorderMock{} + + worker := NewEventsWorker(logger, ts, rec, is, &conf.Events{}) + + rec.On("Record", []dtos.EventDTO{ + {Key: "key1", TrafficTypeName: "user", EventTypeID: "checkin", Value: nil, Timestamp: 123, Properties: map[string]interface{}{"a": 2}}, + }, dtos.Metadata{SDKVersion: "php-1.2.3", MachineIP: "", MachineName: ""}). + Return(nil). + Once() + + rec.On("Record", []dtos.EventDTO{ + {Key: "key2", TrafficTypeName: "user", EventTypeID: "checkin", Value: nil, Timestamp: 123, Properties: map[string]interface{}{"a": 2}}, + }, dtos.Metadata{SDKVersion: "go-1.2.3", MachineIP: "", MachineName: ""}). + Return(nil). + Once() + + rec.On("Record", []dtos.EventDTO{ + {Key: "key3", TrafficTypeName: "user", EventTypeID: "checkout", Value: nil, Timestamp: 123, Properties: map[string]interface{}{"a": 2}}, + {Key: "key4", TrafficTypeName: "user", EventTypeID: "checkout", Value: nil, Timestamp: 123, Properties: map[string]interface{}{"a": 2}}, + }, dtos.Metadata{SDKVersion: "python-1.2.3", MachineIP: "", MachineName: ""}).Return(nil).Once() + + is.Push(types.ClientMetadata{ID: "i1", SdkVersion: "php-1.2.3"}, + dtos.EventDTO{Key: "key1", TrafficTypeName: "user", EventTypeID: "checkin", Value: nil, Timestamp: 123, Properties: map[string]interface{}{"a": 2}}, + ) + is.Push(types.ClientMetadata{ID: "i2", SdkVersion: "go-1.2.3"}, + dtos.EventDTO{Key: "key2", TrafficTypeName: "user", EventTypeID: "checkin", Value: nil, Timestamp: 123, Properties: map[string]interface{}{"a": 2}}, + ) + + worker.SynchronizeEvents(5000) + is.Push(types.ClientMetadata{ID: "i3", SdkVersion: "python-1.2.3"}, + dtos.EventDTO{Key: "key3", TrafficTypeName: "user", EventTypeID: "checkout", Value: nil, Timestamp: 123, Properties: map[string]interface{}{"a": 2}}, + dtos.EventDTO{Key: "key4", TrafficTypeName: "user", EventTypeID: "checkout", Value: nil, Timestamp: 123, Properties: map[string]interface{}{"a": 2}}, + ) + + worker.SynchronizeEvents(5000) + + rec.AssertExpectations(t) +} + +func TestEventsTaskNoParallelism(t *testing.T) { + + // to test this, we set up a Recorder that sleeps for 1 second and returns (no err). + // we one call to `SyncrhonizeImpressions()` wait for 500ms, and fire another one. + // the second one should finish immediately, (becase it does nothing). The second one + // should finish after 2 seconds + + es, _ := sss.NewEventsQueue(100) + ts, _ := inmemory.NewTelemetryStorage() + logger := logging.NewLogger(nil) + rec := &EventsRecorderMock{} + + worker := NewEventsWorker(logger, ts, rec, es, &conf.Events{}) + + rec.On("Record", mock.Anything, mock.Anything).Run(func(mock.Arguments) { time.Sleep(1 * time.Second) }).Return(nil).Twice() + + es.Push(types.ClientMetadata{ID: "i1", SdkVersion: "php-1.2.3"}, + dtos.EventDTO{Key: "key1", TrafficTypeName: "user", EventTypeID: "checkout", Value: nil, Timestamp: 123, Properties: map[string]interface{}{"a": 2}}, + ) + es.Push(types.ClientMetadata{ID: "i2", SdkVersion: "go-1.2.3"}, + dtos.EventDTO{Key: "key2", TrafficTypeName: "user", EventTypeID: "checkout", Value: nil, Timestamp: 123, Properties: map[string]interface{}{"a": 2}}, + ) + + done := make(chan struct{}) + + go func() { + worker.SynchronizeEvents(5000) + done <- struct{}{} + }() + + time.Sleep(500 * time.Millisecond) + assert.Nil(t, worker.SynchronizeEvents(5000)) + + // 2nd call has finished, assert that the first one hasn't: + select { + case <-done: // first call has finished, fail the test + assert.Fail(t, "first call shouldn't have finished yet") + default: + } + + <-done // blocking wait for 1st to finish + +} + +type EventsRecorderMock struct { + mock.Mock +} + +// Record implements service.EventsRecorder +func (m *EventsRecorderMock) Record(events []dtos.EventDTO, metadata dtos.Metadata) error { + args := m.Called(events, metadata) + return args.Error(0) +} + +var _ service.EventsRecorder = (*EventsRecorderMock)(nil) diff --git a/splitio/sdk/workers/impressions.go b/splitio/sdk/workers/impressions.go index e600f01..6b97f73 100644 --- a/splitio/sdk/workers/impressions.go +++ b/splitio/sdk/workers/impressions.go @@ -2,15 +2,19 @@ package workers import ( "errors" + "sync" + + sdkconf "github.com/splitio/splitd/splitio/sdk/conf" + sss "github.com/splitio/splitd/splitio/sdk/storage" + "github.com/splitio/splitd/splitio/sdk/types" + serrors "github.com/splitio/splitd/splitio/util/errors" "github.com/splitio/go-split-commons/v4/dtos" "github.com/splitio/go-split-commons/v4/service" "github.com/splitio/go-split-commons/v4/storage" "github.com/splitio/go-split-commons/v4/synchronizer/worker/impression" "github.com/splitio/go-toolkit/v5/logging" - sdkconf "github.com/splitio/splitd/splitio/sdk/conf" - sss "github.com/splitio/splitd/splitio/sdk/storage" - "github.com/splitio/splitd/splitio/sdk/types" + gtsync "github.com/splitio/go-toolkit/v5/sync" ) type MultiMetaImpressionWorker struct { @@ -19,6 +23,7 @@ type MultiMetaImpressionWorker struct { llrec service.ImpressionsRecorder iq *sss.ImpressionsStorage cfg *sdkconf.Impressions + runnning gtsync.AtomicBool } func NewImpressionsWorker( @@ -38,10 +43,24 @@ func NewImpressionsWorker( } // FlushImpressions implements impression.ImpressionRecorder +// TODO(mredolatti): take `bulkSize` into account func (m *MultiMetaImpressionWorker) FlushImpressions(bulkSize int64) error { - // TODO(mredolatti): take `bulkSize` into account - var errs []error + // prevent 2 evictions from happening at the same time. we don't want a sync.Mutex since that would only cause 43928729 + // function calls to pile up and get called after each mutex release. + if !m.runnning.TestAndSet() { + m.logger.Warning("flush/sync requested while another one is in progress. ignoring") + return nil + } + defer m.runnning.Unset() + + var errs serrors.ConcurrentErrorCollector + var wg sync.WaitGroup + + // iterate all internal queues (one per thin-client associate-data) + // for each [metadata, impressions] tuple, format impressions accordingly, and create a goroutine to post them in BG. + // after all impressions posting-goroutines have been created, wait for all of them to complete, collect errors, + // and unset the `running` flag so that this func can be called again if err := m.iq.RangeAndClear(func(md types.ClientMetadata, q *sss.LockingQueue[dtos.Impression]) { extracted := make([]dtos.Impression, 0, q.Len()) n, err := q.Pop(q.Len(), &extracted) @@ -55,14 +74,20 @@ func (m *MultiMetaImpressionWorker) FlushImpressions(bulkSize int64) error { } formatted := formatImpressions(extracted) - if err := m.llrec.Record(formatted, dtos.Metadata{SDKVersion: md.SdkVersion}, nil); err != nil { - errs = append(errs, err) - } + + wg.Add(1) + go func(imps []dtos.ImpressionsDTO, md dtos.Metadata) { + defer wg.Done() + if err := m.llrec.Record(imps, md, nil); err != nil { + errs.Append(err) + } + }(formatted, dtos.Metadata{SDKVersion: md.SdkVersion}) }); err != nil { m.logger.Error("error traversing impression queues: ", err) } - return errors.Join(errs...) + wg.Wait() + return errs.Join() } // SynchronizeImpressions implements impression.ImpressionRecorder diff --git a/splitio/sdk/workers/impressions_test.go b/splitio/sdk/workers/impressions_test.go index fb718b0..1285ddf 100644 --- a/splitio/sdk/workers/impressions_test.go +++ b/splitio/sdk/workers/impressions_test.go @@ -4,6 +4,7 @@ import ( "reflect" "sort" "testing" + "time" "github.com/splitio/go-split-commons/v4/dtos" "github.com/splitio/go-split-commons/v4/service" @@ -12,6 +13,7 @@ import ( "github.com/splitio/splitd/splitio/sdk/conf" sss "github.com/splitio/splitd/splitio/sdk/storage" "github.com/splitio/splitd/splitio/sdk/types" + "github.com/stretchr/testify/assert" "github.com/stretchr/testify/mock" ) @@ -73,6 +75,48 @@ func TestImpressionsTask(t *testing.T) { rec.AssertExpectations(t) } +func TestImpressionsTaskNoParallelism(t *testing.T) { + + // to test this, we set up a Recorder that sleeps for 1 second and returns (no err). + // we one call to `SyncrhonizeImpressions()` wait for 500ms, and fire another one. + // the second one should finish immediately, (becase it does nothing). The second one + // should finish after 2 seconds + + is, _ := sss.NewImpressionsQueue(100) + ts, _ := inmemory.NewTelemetryStorage() + logger := logging.NewLogger(nil) + rec := &RecorderMock{} + + worker := NewImpressionsWorker(logger, ts, rec, is, &conf.Impressions{}) + + rec.On("Record", mock.Anything, mock.Anything, mock.Anything).Run(func(mock.Arguments) { time.Sleep(1 * time.Second) }).Return(nil).Twice() + + is.Push(types.ClientMetadata{ID: "i1", SdkVersion: "php-1.2.3"}, + dtos.Impression{KeyName: "k1", FeatureName: "f1", Treatment: "on", Label: "l1", ChangeNumber: 123, Time: 123456}) + is.Push(types.ClientMetadata{ID: "i2", SdkVersion: "go-1.2.3"}, + dtos.Impression{KeyName: "k2", FeatureName: "f2", Treatment: "off", Label: "l2", ChangeNumber: 456, Time: 123457}) + + done := make(chan struct{}) + + go func() { + worker.SynchronizeImpressions(5000) + done <- struct{}{} + }() + + time.Sleep(500 * time.Millisecond) + assert.Nil(t, worker.SynchronizeImpressions(5000)) + + // 2nd call has finished, assert that the first one hasn't: + select { + case <-done: // first call has finished, fail the test + assert.Fail(t, "first call shouldn't have finished yet") + default: + } + + <-done // blocking wait for 1st to finish + +} + type RecorderMock struct { mock.Mock } diff --git a/splitio/util/errors/concurrent.go b/splitio/util/errors/concurrent.go new file mode 100644 index 0000000..02dcca8 --- /dev/null +++ b/splitio/util/errors/concurrent.go @@ -0,0 +1,23 @@ +package errors + +import ( + "errors" + "sync" +) + +type ConcurrentErrorCollector struct { + errors []error + mutex sync.Mutex +} + +func (c *ConcurrentErrorCollector) Append(err error) { + c.mutex.Lock() + c.errors = append(c.errors, err) + c.mutex.Unlock() +} + +func (c *ConcurrentErrorCollector) Join() error { + c.mutex.Lock() + defer c.mutex.Unlock() + return errors.Join(c.errors...) +} diff --git a/splitio/util/errors/concurrent_test.go b/splitio/util/errors/concurrent_test.go new file mode 100644 index 0000000..874efb9 --- /dev/null +++ b/splitio/util/errors/concurrent_test.go @@ -0,0 +1,40 @@ +package errors + +import ( + "errors" + "fmt" + "sync" + "testing" + + "github.com/stretchr/testify/assert" +) + +func TestConcurrentErrors(t *testing.T) { + + // test setup: + // errors are interfaces containing pointers, in order for them to be compared with errors.Is, it MUST be the same instance. + // to do the test we create a slice with many vectors, use them and then compare against the original + original := make([]error, 100) + for idx := range original { + original[idx] = errors.New(fmt.Sprintf("err_%d", idx)) + } + + var c ConcurrentErrorCollector + + var wg sync.WaitGroup + wg.Add(100) + for idx := 0; idx < 100; idx++ { + go func(i int) { + c.Append(original[i]) + wg.Done() + }(idx) + } + + wg.Wait() + + je := c.Join() + for idx := 0; idx < 100; idx++ { + assert.ErrorIs(t, je, original[idx]) + } + +}