From 0fcb600560bd4673b40d48dae76750ab23cf9293 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E6=96=87=E5=BE=90?= Date: Tue, 13 Sep 2022 21:46:22 +0800 Subject: [PATCH 1/6] seperate interface and implement --- rlog/log.go | 130 ----------------------------------------- rlog/logger/default.go | 130 +++++++++++++++++++++++++++++++++++++++++ 2 files changed, 130 insertions(+), 130 deletions(-) create mode 100644 rlog/logger/default.go diff --git a/rlog/log.go b/rlog/log.go index e4070b77..0d26cf68 100644 --- a/rlog/log.go +++ b/rlog/log.go @@ -17,16 +17,6 @@ package rlog -import ( - "os" - "path/filepath" - "strings" - - "gopkg.in/natefinch/lumberjack.v2" - - "github.com/sirupsen/logrus" -) - const ( LogKeyProducerGroup = "producerGroup" LogKeyConsumerGroup = "consumerGroup" @@ -51,128 +41,8 @@ type Logger interface { OutputPath(path string) (err error) } -func init() { - r := &defaultLogger{ - logger: logrus.New(), - } - level := os.Getenv("ROCKETMQ_GO_LOG_LEVEL") - switch strings.ToLower(level) { - case "debug": - r.logger.SetLevel(logrus.DebugLevel) - case "warn": - r.logger.SetLevel(logrus.WarnLevel) - case "error": - r.logger.SetLevel(logrus.ErrorLevel) - case "fatal": - r.logger.SetLevel(logrus.FatalLevel) - default: - r.logger.SetLevel(logrus.InfoLevel) - } - rLog = r -} - var rLog Logger -type defaultLogger struct { - logger *logrus.Logger -} - -func (l *defaultLogger) Debug(msg string, fields map[string]interface{}) { - if msg == "" && len(fields) == 0 { - return - } - l.logger.WithFields(fields).Debug(msg) -} - -func (l *defaultLogger) Info(msg string, fields map[string]interface{}) { - if msg == "" && len(fields) == 0 { - return - } - l.logger.WithFields(fields).Info(msg) -} - -func (l *defaultLogger) Warning(msg string, fields map[string]interface{}) { - if msg == "" && len(fields) == 0 { - return - } - l.logger.WithFields(fields).Warning(msg) -} - -func (l *defaultLogger) Error(msg string, fields map[string]interface{}) { - if msg == "" && len(fields) == 0 { - return - } - l.logger.WithFields(fields).Error(msg) -} - -func (l *defaultLogger) Fatal(msg string, fields map[string]interface{}) { - if msg == "" && len(fields) == 0 { - return - } - l.logger.WithFields(fields).Fatal(msg) -} - -func (l *defaultLogger) Level(level string) { - switch strings.ToLower(level) { - case "debug": - l.logger.SetLevel(logrus.DebugLevel) - case "warn": - l.logger.SetLevel(logrus.WarnLevel) - case "error": - l.logger.SetLevel(logrus.ErrorLevel) - case "fatal": - l.logger.SetLevel(logrus.FatalLevel) - default: - l.logger.SetLevel(logrus.InfoLevel) - } -} - -type Config struct { - OutputPath string - MaxFileSizeMB int - MaxBackups int - MaxAges int - Compress bool - LocalTime bool -} - -func (c *Config) Logger() *lumberjack.Logger { - return &lumberjack.Logger{ - Filename: filepath.ToSlash(c.OutputPath), - MaxSize: c.MaxFileSizeMB, // MB - MaxBackups: c.MaxBackups, - MaxAge: c.MaxAges, // days - Compress: c.Compress, // disabled by default - LocalTime: c.LocalTime, - } -} - -const defaultLogPath = "/tmp/rocketmq-client.log" - -func defaultConfig() Config { - return Config{ - OutputPath: defaultLogPath, - MaxFileSizeMB: 10, - MaxBackups: 5, - MaxAges: 3, - Compress: false, - LocalTime: true, - } -} - -func (l *defaultLogger) Config(conf Config) (err error) { - l.logger.Out = conf.Logger() - return -} - -func (l *defaultLogger) OutputPath(path string) (err error) { - config := defaultConfig() - config.OutputPath = path - - l.logger.Out = config.Logger() - return -} - // SetLogger use specified logger user customized, in general, we suggest user to replace the default logger with specified func SetLogger(logger Logger) { rLog = logger diff --git a/rlog/logger/default.go b/rlog/logger/default.go new file mode 100644 index 00000000..2d476c24 --- /dev/null +++ b/rlog/logger/default.go @@ -0,0 +1,130 @@ +package logger + +import ( + "github.com/apache/rocketmq-client-go/v2/rlog" + "github.com/sirupsen/logrus" + "gopkg.in/natefinch/lumberjack.v2" + "os" + "path/filepath" + "strings" +) + +type defaultLogger struct { + logger *logrus.Logger +} + +func init() { + r := &defaultLogger{ + logger: logrus.New(), + } + level := os.Getenv("ROCKETMQ_GO_LOG_LEVEL") + switch strings.ToLower(level) { + case "debug": + r.logger.SetLevel(logrus.DebugLevel) + case "warn": + r.logger.SetLevel(logrus.WarnLevel) + case "error": + r.logger.SetLevel(logrus.ErrorLevel) + case "fatal": + r.logger.SetLevel(logrus.FatalLevel) + default: + r.logger.SetLevel(logrus.InfoLevel) + } + rlog.SetLogger(r) +} + +func (l *defaultLogger) Debug(msg string, fields map[string]interface{}) { + if msg == "" && len(fields) == 0 { + return + } + l.logger.WithFields(fields).Debug(msg) +} + +func (l *defaultLogger) Info(msg string, fields map[string]interface{}) { + if msg == "" && len(fields) == 0 { + return + } + l.logger.WithFields(fields).Info(msg) +} + +func (l *defaultLogger) Warning(msg string, fields map[string]interface{}) { + if msg == "" && len(fields) == 0 { + return + } + l.logger.WithFields(fields).Warning(msg) +} + +func (l *defaultLogger) Error(msg string, fields map[string]interface{}) { + if msg == "" && len(fields) == 0 { + return + } + l.logger.WithFields(fields).Error(msg) +} + +func (l *defaultLogger) Fatal(msg string, fields map[string]interface{}) { + if msg == "" && len(fields) == 0 { + return + } + l.logger.WithFields(fields).Fatal(msg) +} + +func (l *defaultLogger) Level(level string) { + switch strings.ToLower(level) { + case "debug": + l.logger.SetLevel(logrus.DebugLevel) + case "warn": + l.logger.SetLevel(logrus.WarnLevel) + case "error": + l.logger.SetLevel(logrus.ErrorLevel) + case "fatal": + l.logger.SetLevel(logrus.FatalLevel) + default: + l.logger.SetLevel(logrus.InfoLevel) + } +} + +type Config struct { + OutputPath string + MaxFileSizeMB int + MaxBackups int + MaxAges int + Compress bool + LocalTime bool +} + +func (c *Config) Logger() *lumberjack.Logger { + return &lumberjack.Logger{ + Filename: filepath.ToSlash(c.OutputPath), + MaxSize: c.MaxFileSizeMB, // MB + MaxBackups: c.MaxBackups, + MaxAge: c.MaxAges, // days + Compress: c.Compress, // disabled by default + LocalTime: c.LocalTime, + } +} + +const defaultLogPath = "/tmp/rocketmq-client.log" + +func defaultConfig() Config { + return Config{ + OutputPath: defaultLogPath, + MaxFileSizeMB: 10, + MaxBackups: 5, + MaxAges: 3, + Compress: false, + LocalTime: true, + } +} + +func (l *defaultLogger) Config(conf Config) (err error) { + l.logger.Out = conf.Logger() + return +} + +func (l *defaultLogger) OutputPath(path string) (err error) { + config := defaultConfig() + config.OutputPath = path + + l.logger.Out = config.Logger() + return +} From 4234157ef46c4ab19511ac28bffee38ae798c963 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E6=96=87=E5=BE=90?= Date: Tue, 4 Apr 2023 16:50:10 +0800 Subject: [PATCH 2/6] fix panic when close tracedispatcher --- internal/trace.go | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) diff --git a/internal/trace.go b/internal/trace.go index ad302817..f7cea8de 100644 --- a/internal/trace.go +++ b/internal/trace.go @@ -305,8 +305,11 @@ func (td *traceDispatcher) GetTraceTopicName() string { func (td *traceDispatcher) Start() { td.running = true td.cli.Start() + maxWaitDuration := 5 * time.Millisecond + td.ticker = time.NewTicker(maxWaitDuration) + maxWaitTime := maxWaitDuration.Nanoseconds() go primitive.WithRecover(func() { - td.process() + td.process(maxWaitTime) }) } @@ -334,12 +337,9 @@ func (td *traceDispatcher) Append(ctx TraceContext) bool { } // process -func (td *traceDispatcher) process() { +func (td *traceDispatcher) process(maxWaitTime int64) { var count int var batch []TraceContext - maxWaitDuration := 5 * time.Millisecond - maxWaitTime := maxWaitDuration.Nanoseconds() - td.ticker = time.NewTicker(maxWaitDuration) lastput := time.Now() for { select { From cd6cc59faebd4e7edaf796d32d229c9dea39a861 Mon Sep 17 00:00:00 2001 From: wenxuwan Date: Tue, 4 Apr 2023 16:52:23 +0800 Subject: [PATCH 3/6] Restore rlog/log.go --- rlog/log.go | 130 ++++++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 130 insertions(+) diff --git a/rlog/log.go b/rlog/log.go index 0d26cf68..e4070b77 100644 --- a/rlog/log.go +++ b/rlog/log.go @@ -17,6 +17,16 @@ package rlog +import ( + "os" + "path/filepath" + "strings" + + "gopkg.in/natefinch/lumberjack.v2" + + "github.com/sirupsen/logrus" +) + const ( LogKeyProducerGroup = "producerGroup" LogKeyConsumerGroup = "consumerGroup" @@ -41,8 +51,128 @@ type Logger interface { OutputPath(path string) (err error) } +func init() { + r := &defaultLogger{ + logger: logrus.New(), + } + level := os.Getenv("ROCKETMQ_GO_LOG_LEVEL") + switch strings.ToLower(level) { + case "debug": + r.logger.SetLevel(logrus.DebugLevel) + case "warn": + r.logger.SetLevel(logrus.WarnLevel) + case "error": + r.logger.SetLevel(logrus.ErrorLevel) + case "fatal": + r.logger.SetLevel(logrus.FatalLevel) + default: + r.logger.SetLevel(logrus.InfoLevel) + } + rLog = r +} + var rLog Logger +type defaultLogger struct { + logger *logrus.Logger +} + +func (l *defaultLogger) Debug(msg string, fields map[string]interface{}) { + if msg == "" && len(fields) == 0 { + return + } + l.logger.WithFields(fields).Debug(msg) +} + +func (l *defaultLogger) Info(msg string, fields map[string]interface{}) { + if msg == "" && len(fields) == 0 { + return + } + l.logger.WithFields(fields).Info(msg) +} + +func (l *defaultLogger) Warning(msg string, fields map[string]interface{}) { + if msg == "" && len(fields) == 0 { + return + } + l.logger.WithFields(fields).Warning(msg) +} + +func (l *defaultLogger) Error(msg string, fields map[string]interface{}) { + if msg == "" && len(fields) == 0 { + return + } + l.logger.WithFields(fields).Error(msg) +} + +func (l *defaultLogger) Fatal(msg string, fields map[string]interface{}) { + if msg == "" && len(fields) == 0 { + return + } + l.logger.WithFields(fields).Fatal(msg) +} + +func (l *defaultLogger) Level(level string) { + switch strings.ToLower(level) { + case "debug": + l.logger.SetLevel(logrus.DebugLevel) + case "warn": + l.logger.SetLevel(logrus.WarnLevel) + case "error": + l.logger.SetLevel(logrus.ErrorLevel) + case "fatal": + l.logger.SetLevel(logrus.FatalLevel) + default: + l.logger.SetLevel(logrus.InfoLevel) + } +} + +type Config struct { + OutputPath string + MaxFileSizeMB int + MaxBackups int + MaxAges int + Compress bool + LocalTime bool +} + +func (c *Config) Logger() *lumberjack.Logger { + return &lumberjack.Logger{ + Filename: filepath.ToSlash(c.OutputPath), + MaxSize: c.MaxFileSizeMB, // MB + MaxBackups: c.MaxBackups, + MaxAge: c.MaxAges, // days + Compress: c.Compress, // disabled by default + LocalTime: c.LocalTime, + } +} + +const defaultLogPath = "/tmp/rocketmq-client.log" + +func defaultConfig() Config { + return Config{ + OutputPath: defaultLogPath, + MaxFileSizeMB: 10, + MaxBackups: 5, + MaxAges: 3, + Compress: false, + LocalTime: true, + } +} + +func (l *defaultLogger) Config(conf Config) (err error) { + l.logger.Out = conf.Logger() + return +} + +func (l *defaultLogger) OutputPath(path string) (err error) { + config := defaultConfig() + config.OutputPath = path + + l.logger.Out = config.Logger() + return +} + // SetLogger use specified logger user customized, in general, we suggest user to replace the default logger with specified func SetLogger(logger Logger) { rLog = logger From 5bb300b7dfdbeb66215104938005fdd981ca2de9 Mon Sep 17 00:00:00 2001 From: wenxuwan Date: Tue, 4 Apr 2023 16:53:27 +0800 Subject: [PATCH 4/6] Delete default.go --- rlog/logger/default.go | 130 ----------------------------------------- 1 file changed, 130 deletions(-) delete mode 100644 rlog/logger/default.go diff --git a/rlog/logger/default.go b/rlog/logger/default.go deleted file mode 100644 index 2d476c24..00000000 --- a/rlog/logger/default.go +++ /dev/null @@ -1,130 +0,0 @@ -package logger - -import ( - "github.com/apache/rocketmq-client-go/v2/rlog" - "github.com/sirupsen/logrus" - "gopkg.in/natefinch/lumberjack.v2" - "os" - "path/filepath" - "strings" -) - -type defaultLogger struct { - logger *logrus.Logger -} - -func init() { - r := &defaultLogger{ - logger: logrus.New(), - } - level := os.Getenv("ROCKETMQ_GO_LOG_LEVEL") - switch strings.ToLower(level) { - case "debug": - r.logger.SetLevel(logrus.DebugLevel) - case "warn": - r.logger.SetLevel(logrus.WarnLevel) - case "error": - r.logger.SetLevel(logrus.ErrorLevel) - case "fatal": - r.logger.SetLevel(logrus.FatalLevel) - default: - r.logger.SetLevel(logrus.InfoLevel) - } - rlog.SetLogger(r) -} - -func (l *defaultLogger) Debug(msg string, fields map[string]interface{}) { - if msg == "" && len(fields) == 0 { - return - } - l.logger.WithFields(fields).Debug(msg) -} - -func (l *defaultLogger) Info(msg string, fields map[string]interface{}) { - if msg == "" && len(fields) == 0 { - return - } - l.logger.WithFields(fields).Info(msg) -} - -func (l *defaultLogger) Warning(msg string, fields map[string]interface{}) { - if msg == "" && len(fields) == 0 { - return - } - l.logger.WithFields(fields).Warning(msg) -} - -func (l *defaultLogger) Error(msg string, fields map[string]interface{}) { - if msg == "" && len(fields) == 0 { - return - } - l.logger.WithFields(fields).Error(msg) -} - -func (l *defaultLogger) Fatal(msg string, fields map[string]interface{}) { - if msg == "" && len(fields) == 0 { - return - } - l.logger.WithFields(fields).Fatal(msg) -} - -func (l *defaultLogger) Level(level string) { - switch strings.ToLower(level) { - case "debug": - l.logger.SetLevel(logrus.DebugLevel) - case "warn": - l.logger.SetLevel(logrus.WarnLevel) - case "error": - l.logger.SetLevel(logrus.ErrorLevel) - case "fatal": - l.logger.SetLevel(logrus.FatalLevel) - default: - l.logger.SetLevel(logrus.InfoLevel) - } -} - -type Config struct { - OutputPath string - MaxFileSizeMB int - MaxBackups int - MaxAges int - Compress bool - LocalTime bool -} - -func (c *Config) Logger() *lumberjack.Logger { - return &lumberjack.Logger{ - Filename: filepath.ToSlash(c.OutputPath), - MaxSize: c.MaxFileSizeMB, // MB - MaxBackups: c.MaxBackups, - MaxAge: c.MaxAges, // days - Compress: c.Compress, // disabled by default - LocalTime: c.LocalTime, - } -} - -const defaultLogPath = "/tmp/rocketmq-client.log" - -func defaultConfig() Config { - return Config{ - OutputPath: defaultLogPath, - MaxFileSizeMB: 10, - MaxBackups: 5, - MaxAges: 3, - Compress: false, - LocalTime: true, - } -} - -func (l *defaultLogger) Config(conf Config) (err error) { - l.logger.Out = conf.Logger() - return -} - -func (l *defaultLogger) OutputPath(path string) (err error) { - config := defaultConfig() - config.OutputPath = path - - l.logger.Out = config.Logger() - return -} From 7a87936585ded624f6f533ea323f34904c00840f Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E6=96=87=E5=BE=90?= Date: Mon, 28 Sep 2026 16:45:11 +0800 Subject: [PATCH 5/6] fix dispathcer panic --- consumer/interceptor.go | 39 ++- consumer/pull_consumer.go | 12 + consumer/push_consumer.go | 12 + consumer/trace_lifecycle_test.go | 131 ++++++++++ docs/trace.md | 84 ++++++ internal/namesrv.go | 12 +- internal/route.go | 40 ++- internal/trace.go | 258 ++++++++++--------- internal/trace_client.go | 256 +++++++++++++++++++ internal/trace_client_test.go | 423 +++++++++++++++++++++++++++++++ internal/trace_route_test.go | 126 +++++++++ primitive/trace.go | 14 + producer/interceptor.go | 38 ++- producer/producer.go | 12 + producer/trace_lifecycle_test.go | 108 ++++++++ 15 files changed, 1410 insertions(+), 155 deletions(-) create mode 100644 consumer/trace_lifecycle_test.go create mode 100644 docs/trace.md create mode 100644 internal/trace_client.go create mode 100644 internal/trace_client_test.go create mode 100644 internal/trace_route_test.go create mode 100644 producer/trace_lifecycle_test.go diff --git a/consumer/interceptor.go b/consumer/interceptor.go index 08f6b6ae..e93ed55b 100644 --- a/consumer/interceptor.go +++ b/consumer/interceptor.go @@ -19,8 +19,6 @@ package consumer import ( "context" - "fmt" - "reflect" "time" "github.com/apache/rocketmq-client-go/v2/internal" @@ -31,24 +29,39 @@ import ( // WithTrace support rocketmq trace: https://github.com/apache/rocketmq/wiki/RIP-6-Message-Trace. func WithTrace(traceCfg *primitive.TraceConfig) Option { return func(options *consumerOptions) { - dispatcher := internal.NewTraceDispatcher(traceCfg) - options.TraceDispatcher = dispatcher - ori := options.Interceptors - options.Interceptors = make([]primitive.Interceptor, 0) - options.Interceptors = append(options.Interceptors, newTraceInterceptor(dispatcher)) - options.Interceptors = append(options.Interceptors, ori...) + installTraceInterceptor(options, internal.NewTraceDispatcher(traceCfg)) } } +// WithSharedTrace shares trace connections by logical cluster identity. Discovery +// addresses may change without recreating the client. See SharedTraceClientConfig +// for resolver ownership and key requirements. +func WithSharedTrace(traceCfg *primitive.TraceConfig, shared primitive.SharedTraceClientConfig) Option { + return func(options *consumerOptions) { + installTraceInterceptor(options, internal.NewSharedTraceDispatcher(traceCfg, shared)) + } +} + +func installTraceInterceptor(options *consumerOptions, dispatcher internal.TraceDispatcher) { + if internal.IsNilTraceDispatcher(dispatcher) { + return + } + if !internal.IsNilTraceDispatcher(options.TraceDispatcher) { + options.TraceDispatcher.Close() + } + options.TraceDispatcher = dispatcher + options.Interceptors = append([]primitive.Interceptor{newTraceInterceptor(dispatcher)}, options.Interceptors...) +} + func newTraceInterceptor(dispatcher internal.TraceDispatcher) primitive.Interceptor { - if dispatcher != nil && !reflect.ValueOf(dispatcher).IsNil() { - dispatcher.Start() + if internal.IsNilTraceDispatcher(dispatcher) { + return func(ctx context.Context, req, reply interface{}, next primitive.Invoker) error { + return next(ctx, req, reply) + } } + dispatcher.Start() return func(ctx context.Context, req, reply interface{}, next primitive.Invoker) error { - if dispatcher == nil { - return fmt.Errorf("GetOrNewRocketMQClient faild") - } consumerCtx, exist := primitive.GetConsumerCtx(ctx) if !exist || len(consumerCtx.Msgs) == 0 { return next(ctx, req, reply) diff --git a/consumer/pull_consumer.go b/consumer/pull_consumer.go index 5fc747d9..b0aac921 100644 --- a/consumer/pull_consumer.go +++ b/consumer/pull_consumer.go @@ -95,6 +95,12 @@ type defaultPullConsumer struct { func NewPullConsumer(options ...Option) (*defaultPullConsumer, error) { defaultOpts := defaultPullConsumerOptions() + constructed := false + defer func() { + if !constructed && !internal.IsNilTraceDispatcher(defaultOpts.TraceDispatcher) { + defaultOpts.TraceDispatcher.Close() + } + }() for _, apply := range options { apply(&defaultOpts) } @@ -133,6 +139,7 @@ func NewPullConsumer(options ...Option) (*defaultPullConsumer, error) { dc.mqChanged = c.messageQueueChanged c.submitToConsume = c.consumeMessageConcurrently c.interceptor = primitive.ChainInterceptors(c.option.Interceptors...) + constructed = true return c, nil } @@ -227,6 +234,11 @@ func (pc *defaultPullConsumer) nextPullOffset(mq *primitive.MessageQueue, origin func (pc *defaultPullConsumer) Start() error { var err error + defer func() { + if err != nil && !internal.IsNilTraceDispatcher(pc.option.TraceDispatcher) { + pc.option.TraceDispatcher.Close() + } + }() pc.once.Do(func() { err = pc.validate() if err != nil { diff --git a/consumer/push_consumer.go b/consumer/push_consumer.go index 5d87b51c..f18af466 100644 --- a/consumer/push_consumer.go +++ b/consumer/push_consumer.go @@ -77,6 +77,12 @@ type pushConsumer struct { func NewPushConsumer(opts ...Option) (*pushConsumer, error) { defaultOpts := defaultPushConsumerOptions() + constructed := false + defer func() { + if !constructed && !internal.IsNilTraceDispatcher(defaultOpts.TraceDispatcher) { + defaultOpts.TraceDispatcher.Close() + } + }() for _, apply := range opts { apply(&defaultOpts) } @@ -126,11 +132,17 @@ func NewPushConsumer(opts ...Option) (*pushConsumer, error) { p.interceptor = primitive.ChainInterceptors(p.option.Interceptors...) + constructed = true return p, nil } func (pc *pushConsumer) Start() error { var err error + defer func() { + if err != nil && !internal.IsNilTraceDispatcher(pc.option.TraceDispatcher) { + pc.option.TraceDispatcher.Close() + } + }() pc.once.Do(func() { rlog.Info("the consumer start beginning", map[string]interface{}{ rlog.LogKeyConsumerGroup: pc.consumerGroup, diff --git a/consumer/trace_lifecycle_test.go b/consumer/trace_lifecycle_test.go new file mode 100644 index 00000000..7b140ae1 --- /dev/null +++ b/consumer/trace_lifecycle_test.go @@ -0,0 +1,131 @@ +/* +Licensed to the Apache Software Foundation (ASF) under one or more +contributor license agreements. See the NOTICE file distributed with +this work for additional information regarding copyright ownership. +The ASF licenses this file to You under the Apache License, Version 2.0 +(the "License"); you may not use this file except in compliance with +the License. You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package consumer + +import ( + "context" + "errors" + "sync/atomic" + "testing" + + "github.com/apache/rocketmq-client-go/v2/internal" + "github.com/apache/rocketmq-client-go/v2/primitive" + "github.com/stretchr/testify/require" + atomic2 "go.uber.org/atomic" +) + +type unavailableTrace struct{} + +func (*unavailableTrace) Start() { panic("nil dispatcher started") } +func (*unavailableTrace) Close() { panic("nil dispatcher closed") } +func (*unavailableTrace) Append(internal.TraceContext) bool { panic("nil dispatcher invoked") } +func (*unavailableTrace) GetTraceTopicName() string { panic("nil dispatcher invoked") } + +func TestUnavailableTraceInvokesBusinessCallback(t *testing.T) { + var typedNil *unavailableTrace + for _, dispatcher := range []internal.TraceDispatcher{nil, typedNil} { + interceptor := newTraceInterceptor(dispatcher) + for _, businessError := range []error{nil, errors.New("business error")} { + calls := 0 + ctx := context.WithValue(context.Background(), "test-key", "test-value") + ctx = primitive.WithConsumerCtx(ctx, &primitive.ConsumeMessageContext{ + ConsumerGroup: "test-group", + Msgs: []*primitive.MessageExt{{Message: *primitive.NewMessage("topic", []byte("payload"))}}, + Properties: map[string]string{}, + }) + request, reply := new(int), new(int) + err := interceptor(ctx, request, reply, func(actual context.Context, req, resp interface{}) error { + calls++ + require.Equal(t, ctx, actual) + require.Same(t, request, req) + require.Same(t, reply, resp) + *resp.(*int) = 42 + return businessError + }) + require.Equal(t, 1, calls) + require.Equal(t, 42, *reply) + require.Equal(t, businessError, err) + } + } +} + +func TestFailedTraceInstallationIsOptional(t *testing.T) { + options := defaultPushConsumerOptions() + originalCount := len(options.Interceptors) + WithTrace(&primitive.TraceConfig{})(&options) + require.Nil(t, options.TraceDispatcher) + require.Len(t, options.Interceptors, originalCount) + WithSharedTrace(&primitive.TraceConfig{}, primitive.SharedTraceClientConfig{Key: "missing-factory"})(&options) + require.Nil(t, options.TraceDispatcher) + require.Len(t, options.Interceptors, originalCount) +} + +func TestTraceReleasedOnConstructionFailure(t *testing.T) { + var created, closed int32 + shared := primitive.SharedTraceClientConfig{Key: t.Name(), ResolverFactory: func() (primitive.NsResolver, func(), error) { + atomic.AddInt32(&created, 1) + return primitive.NewPassthroughResolver([]string{"127.0.0.1:9876"}), func() { atomic.AddInt32(&closed, 1) }, nil + }} + for i := 0; i < 3; i++ { + _, err := NewPushConsumer( + WithSharedTrace(&primitive.TraceConfig{}, shared), + WithNsResolver(primitive.NewPassthroughResolver([]string{"invalid"})), + ) + require.Error(t, err) + } + require.Equal(t, int32(3), atomic.LoadInt32(&created)) + require.Equal(t, atomic.LoadInt32(&created), atomic.LoadInt32(&closed)) +} + +func TestPullTraceReleasedOnConstructionFailure(t *testing.T) { + var created, closed int32 + shared := primitive.SharedTraceClientConfig{Key: t.Name(), ResolverFactory: func() (primitive.NsResolver, func(), error) { + atomic.AddInt32(&created, 1) + return primitive.NewPassthroughResolver([]string{"127.0.0.1:9876"}), func() { atomic.AddInt32(&closed, 1) }, nil + }} + for i := 0; i < 3; i++ { + _, err := NewPullConsumer( + WithSharedTrace(&primitive.TraceConfig{}, shared), + WithNsResolver(primitive.NewPassthroughResolver([]string{"invalid"})), + ) + require.Error(t, err) + } + require.Equal(t, int32(3), atomic.LoadInt32(&created)) + require.Equal(t, atomic.LoadInt32(&created), atomic.LoadInt32(&closed)) +} + +func TestTraceReleasedOnStartFailure(t *testing.T) { + for _, pull := range []bool{false, true} { + var closed int32 + opts := defaultPushConsumerOptions() + shared := primitive.SharedTraceClientConfig{Key: t.Name(), ResolverFactory: func() (primitive.NsResolver, func(), error) { + return primitive.NewPassthroughResolver([]string{"127.0.0.1:9876"}), func() { atomic.AddInt32(&closed, 1) }, nil + }} + WithSharedTrace(&primitive.TraceConfig{}, shared)(&opts) + // Invalid group fails validation before any business transport is started. + dc := &defaultConsumer{consumerGroup: "", option: opts, state: atomic2.NewInt32(int32(internal.StateCreateJust))} + var err error + if pull { + err = (&defaultPullConsumer{defaultConsumer: dc}).Start() + } else { + err = (&pushConsumer{defaultConsumer: dc}).Start() + } + require.Error(t, err) + require.Equal(t, int32(1), atomic.LoadInt32(&closed)) + } +} diff --git a/docs/trace.md b/docs/trace.md new file mode 100644 index 00000000..812f6682 --- /dev/null +++ b/docs/trace.md @@ -0,0 +1,84 @@ +# Message trace discovery and lifecycle + +`consumer.WithTrace` and `producer.WithTrace` remain available with the existing +`primitive.TraceConfig`. A trace initialization failure is logged and disables +tracing for that option; it does not prevent the business interceptor from +running or replace its result/error. No public fields were added to `TraceConfig`. + +## Sharing with dynamic discovery + +For a long-running process that creates and closes consumers independently, use +`consumer.WithSharedTrace` (or `producer.WithSharedTrace`) with a stable logical +cluster key and a resolver factory: + +```go +shared := primitive.SharedTraceClientConfig{ + Key: "discovery-endpoint/tenant/cluster", + ResolverFactory: func() (primitive.NsResolver, func(), error) { + // Create an independently owned resolver here. For a dynamic resolver, + // return its stop/close function as the second result. + return primitive.NewPassthroughResolver([]string{"127.0.0.1:9876"}), nil, nil + }, +} + +c, err := rocketmq.NewPushConsumer( + consumer.WithGroupName("orders"), + consumer.WithNameServer([]string{"127.0.0.1:9876"}), + consumer.WithSharedTrace(&primitive.TraceConfig{}, shared), +) +if err != nil { + return err +} +defer c.Shutdown() +// Subscribe, then Start as usual. +``` + +The key must identify the discovery source and destination, including any tenant +or namespace that affects routing. Do not use a consumer group, a resolved IP +list, or a changing address-list hash. `TraceConfig.UnitName`, `Access`, and the +full `Credentials` additionally partition the pool. Changing credentials creates +an isolated resource; users of the old credentials keep their resource until +closed. Different keys never share transport, even if they currently resolve to +the same addresses. + +The factory is called once per shared resource lifetime. All callers using the +same key and partition must supply equivalent factory configuration. The factory +owns an independent resolver: do not return a consumer's resolver if that +consumer will close it. The optional cleanup runs once after the last dispatcher +releases the resource, or if initialization fails. Factory, `Resolve`, and cleanup +must return promptly; `Resolve` must be safe for concurrent calls. The SDK does +not call arbitrary `Close` methods on borrowed resolvers. + +`TraceConfig.NamesrvAddrs` and `Resolver` are ignored by `WithSharedTrace`. The +factory is authoritative. Initialization failure is logged and tracing remains +disabled for that dispatcher; recreating the consumer/producer retries setup. + +## Connections, refresh and shutdown + +Each consumer/producer has its own trace dispatcher and buffer. Dispatchers in the +same partition share one NameServer connection pool, one Broker connection pool, +one resolver, and one discovery worker. Adding consumers therefore does not add +independent trace connection pools for that partition. Connections are opened +lazily; this is not a promise of exactly one TCP connection per cluster. + +NameServer addresses refresh after 10 seconds and then every 2 minutes. Trace +topic routes refresh every 30 seconds, including region-specific cloud trace +topics. NameServer discovery and Broker route discovery are separate operations; +updates are periodic, not immediate. All dispatchers read the same route and +address state. Closing one dispatcher leaves other users active. + +Shutdown rejects new records, drains accepted records, and waits for in-flight +sends. Close waits at most 5 seconds, then cancels outstanding I/O; cleanup finishes +when workers and callbacks exit. Trace delivery remains best effort. After the +last dispatcher finishes, both connection pools are closed and factory cleanup +runs. A concurrent acquisition waits for that cleanup before creating a new +resource. Failed consumer/producer construction or startup also releases trace +ownership. + +Legacy `WithTrace` still uses the default sharing partition and checks resolved +addresses before reuse. Its resolver is borrowed and must remain usable for the +whole shared lifetime. It also benefits from reference-counted cleanup, a common +NameServer object, route refresh, and nil-safe interceptors. Use `WithSharedTrace` +when separate resolvers can return different address snapshots for the same +logical cluster. The general producer/consumer client registry and its +NameServer conflict checks are unchanged. diff --git a/internal/namesrv.go b/internal/namesrv.go index ac7643db..84042cd3 100644 --- a/internal/namesrv.go +++ b/internal/namesrv.go @@ -121,7 +121,7 @@ func NewNamesrv(resolver primitive.NsResolver, config *remote.RemotingClientConf } nameSrvClient := remote.NewRemotingClient(config) return &namesrvs{ - srvs: addr, + srvs: append([]string(nil), addr...), lock: new(sync.Mutex), nameSrvClient: nameSrvClient, brokerVersionMap: make(map[string]map[string]int32, 0), @@ -149,18 +149,22 @@ func (s *namesrvs) getNameServerAddress() string { } func (s *namesrvs) Size() int { + s.lock.Lock() + defer s.lock.Unlock() return len(s.srvs) } func (s *namesrvs) String() string { - return strings.Join(s.srvs, ";") + return strings.Join(s.AddrList(), ";") } func (s *namesrvs) SetCredentials(credentials primitive.Credentials) { s.nameSrvClient.RegisterInterceptor(remote.ACLInterceptor(credentials)) } func (s *namesrvs) AddrList() []string { - return s.srvs + s.lock.Lock() + defer s.lock.Unlock() + return append([]string(nil), s.srvs...) } // UpdateNameServerAddress will update srvs. @@ -179,5 +183,5 @@ func (s *namesrvs) UpdateNameServerAddress() { return } - s.srvs = srvs + s.srvs = append([]string(nil), srvs...) } diff --git a/internal/route.go b/internal/route.go index 4af8619c..d8ca7fdb 100644 --- a/internal/route.go +++ b/internal/route.go @@ -120,6 +120,10 @@ func (s *namesrvs) CheckTopicRouteHasTopic(topic string) bool { } func (s *namesrvs) UpdateTopicRouteInfoWithDefault(topic string, defaultTopic string, defaultQueueNum int) (*TopicRouteData, bool, error) { + return s.updateTopicRouteInfoWithContext(context.Background(), topic, defaultTopic, defaultQueueNum) +} + +func (s *namesrvs) updateTopicRouteInfoWithContext(ctx context.Context, topic string, defaultTopic string, defaultQueueNum int) (*TopicRouteData, bool, error) { s.lockNamesrv.Lock() defer s.lockNamesrv.Unlock() @@ -132,7 +136,7 @@ func (s *namesrvs) UpdateTopicRouteInfoWithDefault(topic string, defaultTopic st if len(defaultTopic) > 0 { t = defaultTopic } - routeData, err = s.queryTopicRouteInfoFromServer(t) + routeData, err = s.queryTopicRouteInfoWithContext(ctx, t) if err != nil { rlog.Warning("query topic route from server error", map[string]interface{}{ @@ -357,6 +361,31 @@ func (s *namesrvs) FetchPublishMessageQueues(topic string) ([]*primitive.Message return publishInfo.MqList, nil } +func (s *namesrvs) fetchPublishMessageQueuesWithContext(ctx context.Context, topic string) ([]*primitive.MessageQueue, error) { + var ( + err error + routeData *TopicRouteData + ) + + v, exist := s.routeDataMap.Load(topic) + if !exist { + routeData, _, err = s.updateTopicRouteInfoWithContext(ctx, topic, "", 0) + if err != nil { + rlog.Error("queryTopicRouteInfoFromServer failed", map[string]interface{}{ + rlog.LogKeyTopic: topic, + }) + return nil, err + } + } else { + // Queue selection sorts in place; keep the shared cache immutable for refreshers. + routeData = v.(*TopicRouteData).clone() + } + + publishInfo := s.routeData2PublishInfo(topic, routeData) + + return publishInfo.MqList, nil +} + func (s *namesrvs) AddBrokerVersion(brokerName, brokerAddr string, version int32) { s.brokerLock.Lock() defer s.brokerLock.Unlock() @@ -383,6 +412,10 @@ func (s *namesrvs) findBrokerVersion(brokerName, brokerAddr string) int32 { } func (s *namesrvs) queryTopicRouteInfoFromServer(topic string) (*TopicRouteData, error) { + return s.queryTopicRouteInfoWithContext(context.Background(), topic) +} + +func (s *namesrvs) queryTopicRouteInfoWithContext(parent context.Context, topic string) (*TopicRouteData, error) { request := &GetRouteInfoRequestHeader{ Topic: topic, } @@ -402,8 +435,11 @@ func (s *namesrvs) queryTopicRouteInfoFromServer(topic string) (*TopicRouteData, } for i := 0; i < s.Size(); i++ { + if err := parent.Err(); err != nil { + return nil, err + } rc := remote.NewRemotingCommand(ReqGetRouteInfoByTopic, request, nil) - ctx, cancel := context.WithTimeout(context.Background(), requestTimeout) + ctx, cancel := context.WithTimeout(parent, requestTimeout) response, err = s.nameSrvClient.InvokeSync(ctx, s.getNameServerAddress(), rc) if err == nil { diff --git a/internal/trace.go b/internal/trace.go index 55bc71d3..56be0291 100644 --- a/internal/trace.go +++ b/internal/trace.go @@ -21,14 +21,12 @@ import ( "bytes" "context" "fmt" - "runtime" "strconv" "strings" + "sync" "sync/atomic" "time" - "github.com/pkg/errors" - "github.com/apache/rocketmq-client-go/v2/internal/remote" "github.com/apache/rocketmq-client-go/v2/primitive" "github.com/apache/rocketmq-client-go/v2/rlog" @@ -221,169 +219,172 @@ type TraceDispatcher interface { } type traceDispatcher struct { - ctx context.Context - cancel context.CancelFunc - running bool + // Keep 64-bit atomic data first for alignment on 32-bit platforms. + discardCount int64 + mu sync.Mutex + started, closed bool + closeOnce sync.Once + ctx context.Context + cancel context.CancelFunc + sendCtx context.Context + cancelSend context.CancelFunc + processDone chan struct{} + closeDone chan struct{} + pending sync.WaitGroup traceTopic string access primitive.AccessChannel - - ticker *time.Ticker - input chan TraceContext - batchCh chan []*TraceContext - - discardCount int64 - - // support deliver trace message to other cluster. - namesrvs *namesrvs - // round robin index - rrindex int32 - cli RMQClient + ticker *time.Ticker + input chan TraceContext + namesrvs *namesrvs + rrindex int32 + cli RMQClient + resource *traceClient } func NewTraceDispatcher(traceCfg *primitive.TraceConfig) *traceDispatcher { - ctx := context.Background() - ctx, cancel := context.WithCancel(ctx) - - t := traceCfg.TraceTopic - if len(t) == 0 { - t = RmqSysTraceTopic - } - - if traceCfg.Access == primitive.Cloud { - t = TraceTopicPrefix + traceCfg.TraceTopic - } - - if len(traceCfg.NamesrvAddrs) == 0 && traceCfg.Resolver == nil { - panic("no NamesrvAddrs or Resolver configured") - } + return newTraceDispatcher(traceCfg, nil) +} - var srvs *namesrvs - var err error - if len(traceCfg.NamesrvAddrs) > 0 { - srvs, err = NewNamesrv(primitive.NewPassthroughResolver(traceCfg.NamesrvAddrs), nil) - } else { - srvs, err = NewNamesrv(traceCfg.Resolver, nil) - } +func NewSharedTraceDispatcher(traceCfg *primitive.TraceConfig, shared primitive.SharedTraceClientConfig) *traceDispatcher { + return newTraceDispatcher(traceCfg, &shared) +} +func newTraceDispatcher(traceCfg *primitive.TraceConfig, shared *primitive.SharedTraceClientConfig) *traceDispatcher { + resource, err := acquireTraceClient(traceCfg, shared) if err != nil { - panic(errors.Wrap(err, "new Namesrv failed.")) + rlog.Error("trace initialization failed; tracing is disabled", map[string]interface{}{ + rlog.LogKeyUnderlayError: err, + }) + return nil } - if !traceCfg.Credentials.IsEmpty() { - srvs.SetCredentials(traceCfg.Credentials) + topic := traceCfg.TraceTopic + if topic == "" { + topic = RmqSysTraceTopic } - - cliOp := DefaultClientOptions() - cliOp.GroupName = traceCfg.GroupName - cliOp.UnitName = traceCfg.UnitName - cliOp.NameServerAddrs = traceCfg.NamesrvAddrs - cliOp.InstanceName = "INNER_TRACE_CLIENT_DEFAULT" - cliOp.RetryTimes = 0 - cliOp.Namesrv = srvs - cliOp.Credentials = traceCfg.Credentials - cli := GetOrNewRocketMQClient(cliOp, nil) - if cli == nil { - return nil + if traceCfg.Access == primitive.Cloud { + topic = TraceTopicPrefix + traceCfg.TraceTopic } - cliOp.Namesrv = cli.GetNameSrv() + ctx, cancel := context.WithCancel(context.Background()) + sendCtx, cancelSend := context.WithCancel(context.Background()) return &traceDispatcher{ - ctx: ctx, - cancel: cancel, - - traceTopic: t, - access: traceCfg.Access, - input: make(chan TraceContext, 1024), - batchCh: make(chan []*TraceContext, 2048), - cli: cli, - namesrvs: srvs, + ctx: ctx, cancel: cancel, sendCtx: sendCtx, cancelSend: cancelSend, + processDone: make(chan struct{}), closeDone: make(chan struct{}), + traceTopic: topic, access: traceCfg.Access, input: make(chan TraceContext, 1024), + cli: resource.cli, namesrvs: resource.namesrvs, resource: resource, } } -func (td *traceDispatcher) GetTraceTopicName() string { - return td.traceTopic -} +func (td *traceDispatcher) GetTraceTopicName() string { return td.traceTopic } func (td *traceDispatcher) Start() { - td.running = true - td.cli.Start() - maxWaitDuration := 5 * time.Millisecond - td.ticker = time.NewTicker(maxWaitDuration) - maxWaitTime := maxWaitDuration.Nanoseconds() + if td == nil { + return + } + td.mu.Lock() + defer td.mu.Unlock() + if td.started || td.closed { + return + } + td.started = true + td.ticker = time.NewTicker(5 * time.Millisecond) go primitive.WithRecover(func() { - td.process(maxWaitTime) + defer close(td.processDone) + td.process() }) } func (td *traceDispatcher) Close() { - td.running = false - td.ticker.Stop() - td.cancel() + if td == nil { + return + } + td.closeOnce.Do(func() { + td.mu.Lock() + td.closed = true + if td.started { + td.ticker.Stop() + } else { + close(td.processDone) + } + td.cancel() + td.mu.Unlock() + // Drain accepted records, then wait for all asynchronous callbacks before + // releasing a transport. If the deadline expires, cancel I/O and finish + // cleanup in the background; never let a late send reopen a closed pool. + go func() { + <-td.processDone + td.pending.Wait() + td.cancelSend() + td.resource.release() + close(td.closeDone) + }() + }) + timer := time.NewTimer(5 * time.Second) + defer timer.Stop() + select { + case <-td.closeDone: + case <-timer.C: + td.cancelSend() + } } func (td *traceDispatcher) Append(ctx TraceContext) bool { - if !td.running { - rlog.Error("traceDispatcher is closed.", nil) + if td == nil { + return false + } + td.mu.Lock() + defer td.mu.Unlock() + if !td.started || td.closed { return false } select { case td.input <- ctx: return true default: - rlog.Warning("buffer full", map[string]interface{}{ + rlog.Warning("trace buffer full", map[string]interface{}{ "discardCount": atomic.AddInt64(&td.discardCount, 1), - "TraceContext": ctx, }) return false } } -// process -func (td *traceDispatcher) process(maxWaitTime int64) { - var count int +func (td *traceDispatcher) submitBatch(batch []TraceContext) { + if len(batch) == 0 { + return + } + td.pending.Add(1) + go primitive.WithRecover(func() { + defer td.pending.Done() + td.batchCommit(batch) + }) +} + +func (td *traceDispatcher) process() { var batch []TraceContext - lastput := time.Now() + flush := func() { td.submitBatch(batch); batch = nil } for { select { case ctx := <-td.input: - count++ - lastput = time.Now() batch = append(batch, ctx) - if count == batchSize { - count = 0 - batchSend := batch - go primitive.WithRecover(func() { - td.batchCommit(batchSend) - }) - batch = make([]TraceContext, 0) + if len(batch) == batchSize { + flush() } case <-td.ticker.C: - delta := time.Since(lastput).Nanoseconds() - if delta > maxWaitTime { - lastput = time.Now() - if len(batch) > 0 { - count = 0 - batchSend := batch - go primitive.WithRecover(func() { - td.batchCommit(batchSend) - }) - batch = make([]TraceContext, 0) - } - } + flush() case <-td.ctx.Done(): - batchSend := batch - go primitive.WithRecover(func() { - td.batchCommit(batchSend) - }) - batch = make([]TraceContext, 0) - - now := time.Now().UnixNano() / int64(time.Millisecond) - end := now + 500 - for now < end { - now = time.Now().UnixNano() / int64(time.Millisecond) - runtime.Gosched() + // Append and Close share a lock, so no more records can enter. + for { + select { + case ctx := <-td.input: + batch = append(batch, ctx) + if len(batch) == batchSize { + flush() + } + default: + flush() + return + } } - rlog.Info(fmt.Sprintf("------end trace send %v %v", td.input, td.batchCh), nil) - return } } } @@ -394,7 +395,7 @@ func (td *traceDispatcher) batchCommit(ctxs []TraceContext) { keyedCtxs := make(map[string][]TraceTransferBean) for _, ctx := range ctxs { if len(ctx.TraceBeans) == 0 { - return + continue } topic := ctx.TraceBeans[0].Topic regionID := ctx.RegionId @@ -469,9 +470,12 @@ func (td *traceDispatcher) sendTraceDataByMQ(keySet Keyset, regionID string, dat } var req = td.buildSendRequest(mq, msg) - ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + ctx, cancel := context.WithTimeout(td.sendCtx, 5*time.Second) + td.pending.Add(1) + var finishOnce sync.Once + finish := func() { finishOnce.Do(func() { cancel(); td.pending.Done() }) } err := td.cli.InvokeAsync(ctx, addr, req, func(command *remote.RemotingCommand, e error) { - cancel() + defer finish() resp := primitive.NewSendResult() if e != nil { rlog.Info("send trace data error.", map[string]interface{}{ @@ -486,7 +490,7 @@ func (td *traceDispatcher) sendTraceDataByMQ(keySet Keyset, regionID string, dat } }) if err != nil { - cancel() + finish() rlog.Info("send trace data error when invoke", map[string]interface{}{ rlog.LogKeyUnderlayError: err, }) @@ -498,7 +502,13 @@ func (td *traceDispatcher) findMq(regionID string) (*primitive.MessageQueue, str if td.access == primitive.Cloud { traceTopic = td.traceTopic + regionID } - mqs, err := td.namesrvs.FetchPublishMessageQueues(traceTopic) + if td.sendCtx.Err() != nil { + return nil, "" + } + td.resource.topics.Store(traceTopic, struct{}{}) + ctx, cancel := context.WithTimeout(td.sendCtx, 5*time.Second) + defer cancel() + mqs, err := td.namesrvs.fetchPublishMessageQueuesWithContext(ctx, traceTopic) if err != nil { rlog.Error("fetch publish message queues failed", map[string]interface{}{ rlog.LogKeyUnderlayError: err, diff --git a/internal/trace_client.go b/internal/trace_client.go new file mode 100644 index 00000000..a9e8dd08 --- /dev/null +++ b/internal/trace_client.go @@ -0,0 +1,256 @@ +/* +Licensed to the Apache Software Foundation (ASF) under one or more +contributor license agreements. See the NOTICE file distributed with +this work for additional information regarding copyright ownership. +The ASF licenses this file to You under the Apache License, Version 2.0 +(the "License"); you may not use this file except in compliance with +the License. You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package internal + +import ( + "context" + "fmt" + "reflect" + "sort" + "sync" + "time" + + "github.com/apache/rocketmq-client-go/v2/internal/remote" + "github.com/apache/rocketmq-client-go/v2/primitive" +) + +// Trace clients have no registered consumers/producers. Their lifetime and +// discovery work belong to this pool, not the general MQ client registry. +// In particular, trace topics must be refreshed even without a producerMap entry. +type traceClientKey struct { + shared bool + key, unit string + access primitive.AccessChannel + credentials primitive.Credentials +} + +type traceClientEntry struct { + ready chan struct{} + client *traceClient + refs int +} + +var traceClients = struct { + sync.Mutex + entries map[traceClientKey]*traceClientEntry +}{entries: make(map[traceClientKey]*traceClientEntry)} + +type traceClient struct { + key traceClientKey + cli *rmqClient + namesrvs *namesrvs + cleanup func() + cancel context.CancelFunc + done chan struct{} + topics sync.Map +} + +func isNilTraceValue(value interface{}) bool { + if value == nil { + return true + } + v := reflect.ValueOf(value) + switch v.Kind() { + case reflect.Chan, reflect.Func, reflect.Interface, reflect.Map, reflect.Ptr, reflect.Slice: + return v.IsNil() + default: + return false + } +} + +// IsNilTraceDispatcher also handles a typed nil returned by a constructor. +func IsNilTraceDispatcher(dispatcher TraceDispatcher) bool { + return isNilTraceValue(dispatcher) +} + +func acquireTraceClient(cfg *primitive.TraceConfig, shared *primitive.SharedTraceClientConfig) (*traceClient, error) { + if cfg == nil { + return nil, fmt.Errorf("trace config is nil") + } + key := traceClientKey{unit: cfg.UnitName, access: cfg.Access, credentials: cfg.Credentials} + var resolver primitive.NsResolver + if shared != nil { + if shared.Key == "" || shared.ResolverFactory == nil { + return nil, fmt.Errorf("shared trace client requires a key and resolver factory") + } + key.shared, key.key = true, shared.Key + } else { + resolver = cfg.Resolver + if len(cfg.NamesrvAddrs) > 0 { + resolver = primitive.NewPassthroughResolver(append([]string(nil), cfg.NamesrvAddrs...)) + } + if isNilTraceValue(resolver) { + return nil, fmt.Errorf("no trace NamesrvAddrs or Resolver configured") + } + } + + for { + traceClients.Lock() + entry := traceClients.entries[key] + if entry != nil { + if entry.ready != nil { + ready := entry.ready + traceClients.Unlock() + <-ready + continue + } + entry.refs++ + client := entry.client + traceClients.Unlock() + // Preserve the legacy guard. Opt-in sharing uses logical identity, + // so a discovery update cannot be mistaken for a different cluster. + if shared == nil && !sameTraceAddresses(resolver.Resolve(), client.namesrvs.resolver.Resolve()) { + client.release() + return nil, fmt.Errorf("different namesrv option in the same trace instance") + } + return client, nil + } + entry = &traceClientEntry{ready: make(chan struct{}), refs: 1} + traceClients.entries[key] = entry + traceClients.Unlock() + + // User discovery code must not run under the pool lock. + client, err := createTraceClient(key, cfg, shared, resolver) + traceClients.Lock() + if err != nil { + delete(traceClients.entries, key) + } else { + entry.client = client + } + close(entry.ready) + entry.ready = nil + traceClients.Unlock() + return client, err + } +} + +func sameTraceAddresses(a, b []string) bool { + if len(a) != len(b) { + return false + } + a, b = append([]string(nil), a...), append([]string(nil), b...) + sort.Strings(a) + sort.Strings(b) + for i := range a { + if a[i] != b[i] { + return false + } + } + return true +} + +func createTraceClient(key traceClientKey, cfg *primitive.TraceConfig, shared *primitive.SharedTraceClientConfig, resolver primitive.NsResolver) (*traceClient, error) { + var cleanup func() + var err error + if shared != nil { + resolver, cleanup, err = shared.ResolverFactory() + } + success := false + defer func() { + if !success && cleanup != nil { + cleanup() + } + }() + if err != nil { + return nil, err + } + if isNilTraceValue(resolver) { + return nil, fmt.Errorf("trace resolver factory returned nil") + } + srvs, err := NewNamesrv(resolver, nil) + if err != nil { + return nil, err + } + transport := remote.NewRemotingClient(nil) + if !cfg.Credentials.IsEmpty() { + srvs.SetCredentials(cfg.Credentials) + transport.RegisterInterceptor(remote.ACLInterceptor(cfg.Credentials)) + } + ctx, cancel := context.WithCancel(context.Background()) + options := DefaultClientOptions() + options.Namesrv = srvs + options.GroupName = cfg.GroupName + options.UnitName = cfg.UnitName + options.Credentials = cfg.Credentials + options.InstanceName = "INNER_TRACE_CLIENT_DEFAULT" + client := &traceClient{ + key: key, namesrvs: srvs, cleanup: cleanup, cancel: cancel, done: make(chan struct{}), + // Only remoting and response decoding are used. Starting the general + // MQ workers here would duplicate discovery and run empty heartbeats. + cli: &rmqClient{option: options, remoteClient: transport}, + } + go primitive.WithRecover(func() { client.refresh(ctx) }) + success = true + return client, nil +} + +func (client *traceClient) refresh(ctx context.Context) { + defer close(client.done) + names := time.NewTimer(10 * time.Second) + routes := time.NewTicker(_PullNameServerInterval) + defer names.Stop() + defer routes.Stop() + for { + select { + case <-ctx.Done(): + return + case <-names.C: + client.namesrvs.UpdateNameServerAddress() + names.Reset(2 * time.Minute) + case <-routes.C: + client.refreshRoutes(ctx) + } + } +} + +func (client *traceClient) refreshRoutes(ctx context.Context) { + client.topics.Range(func(key, _ interface{}) bool { + if ctx.Err() != nil { + return false + } + requestCtx, cancel := context.WithTimeout(ctx, 5*time.Second) + client.namesrvs.updateTopicRouteInfoWithContext(requestCtx, key.(string), "", 0) + cancel() + return true + }) +} + +func (client *traceClient) release() { + traceClients.Lock() + entry := traceClients.entries[client.key] + entry.refs-- + if entry.refs != 0 { + traceClients.Unlock() + return + } + // Keep a closing entry until cleanup completes. A replacement cannot reuse + // a closed resolver or race with disposal of the previous connection pools. + entry.ready = make(chan struct{}) + traceClients.Unlock() + client.cancel() + <-client.done + client.cli.remoteClient.ShutDown() + client.namesrvs.nameSrvClient.ShutDown() + if client.cleanup != nil { + client.cleanup() + } + traceClients.Lock() + delete(traceClients.entries, client.key) + close(entry.ready) + traceClients.Unlock() +} diff --git a/internal/trace_client_test.go b/internal/trace_client_test.go new file mode 100644 index 00000000..15ac2eab --- /dev/null +++ b/internal/trace_client_test.go @@ -0,0 +1,423 @@ +/* +Licensed to the Apache Software Foundation (ASF) under one or more +contributor license agreements. See the NOTICE file distributed with +this work for additional information regarding copyright ownership. +The ASF licenses this file to You under the Apache License, Version 2.0 +(the "License"); you may not use this file except in compliance with +the License. You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package internal + +import ( + "context" + "errors" + "fmt" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/apache/rocketmq-client-go/v2/internal/remote" + "github.com/apache/rocketmq-client-go/v2/primitive" + "github.com/golang/mock/gomock" + "github.com/stretchr/testify/require" +) + +type mutableTraceResolver struct { + mu sync.Mutex + addrs []string +} + +func (r *mutableTraceResolver) Resolve() []string { + r.mu.Lock() + defer r.mu.Unlock() + return append([]string(nil), r.addrs...) +} +func (r *mutableTraceResolver) Description() string { return "test discovery endpoint" } +func (r *mutableTraceResolver) set(addr string) { + r.mu.Lock() + defer r.mu.Unlock() + r.addrs = []string{addr} +} + +func traceTestShared(key string, resolver primitive.NsResolver, created, closed *int32) primitive.SharedTraceClientConfig { + return primitive.SharedTraceClientConfig{Key: key, ResolverFactory: func() (primitive.NsResolver, func(), error) { + atomic.AddInt32(created, 1) + return resolver, func() { atomic.AddInt32(closed, 1) }, nil + }} +} + +func TestSharedTraceDiscoveryAndLifetime(t *testing.T) { + var created, closed int32 + resolver := &mutableTraceResolver{addrs: []string{"127.0.0.1:9876"}} + shared := traceTestShared(t.Name(), resolver, &created, &closed) + first := NewSharedTraceDispatcher(&primitive.TraceConfig{}, shared) + require.NotNil(t, first) + defer first.Close() + first.Start() + first.Start() + resolver.set("127.0.0.2:9876") + second := NewSharedTraceDispatcher(&primitive.TraceConfig{GroupName: "different-consumer"}, shared) + require.NotNil(t, second) + defer second.Close() + require.Same(t, first.resource, second.resource) + require.Same(t, first.namesrvs, second.namesrvs) + require.Same(t, first.cli.GetNameSrv(), second.namesrvs) + require.Equal(t, int32(1), atomic.LoadInt32(&created)) + first.namesrvs.UpdateNameServerAddress() + require.Equal(t, []string{"127.0.0.2:9876"}, second.namesrvs.AddrList()) + first.Close() + first.Close() + require.Equal(t, int32(0), atomic.LoadInt32(&closed)) + second.Start() + require.True(t, second.Append(TraceContext{})) + second.Close() + require.Equal(t, int32(1), atomic.LoadInt32(&closed)) + require.False(t, first.Append(TraceContext{})) + first.Start() + resolver.set("127.0.0.3:9876") + third := NewSharedTraceDispatcher(&primitive.TraceConfig{}, shared) + require.NotNil(t, third) + defer third.Close() + require.NotSame(t, first.resource, third.resource) + require.Equal(t, []string{"127.0.0.3:9876"}, third.namesrvs.AddrList()) + third.Close() // Closing before Start must release ownership too. + require.Equal(t, int32(2), atomic.LoadInt32(&created)) + require.Equal(t, int32(2), atomic.LoadInt32(&closed)) +} + +func TestSharedTraceIdentityIsolation(t *testing.T) { + var created, closed int32 + resolver := primitive.NewPassthroughResolver([]string{"127.0.0.1:9876"}) + shared := traceTestShared(t.Name(), resolver, &created, &closed) + base := NewSharedTraceDispatcher(&primitive.TraceConfig{}, shared) + require.NotNil(t, base) + defer base.Close() + for _, cfg := range []primitive.TraceConfig{ + {UnitName: "other-unit"}, {Access: primitive.Cloud}, + {Credentials: primitive.Credentials{AccessKey: "user", SecretKey: "secret"}}, + {Credentials: primitive.Credentials{AccessKey: "user", SecretKey: "rotated", SecurityToken: "token"}}, + } { + other := NewSharedTraceDispatcher(&cfg, shared) + require.NotNil(t, other) + require.NotSame(t, base.resource, other.resource) + other.Close() + } + shared.Key += "-other-cluster" + other := NewSharedTraceDispatcher(&primitive.TraceConfig{}, shared) + require.NotNil(t, other) + require.NotSame(t, base.resource, other.resource) + other.Close() + base.Close() + require.Equal(t, int32(6), atomic.LoadInt32(&created)) + require.Equal(t, atomic.LoadInt32(&created), atomic.LoadInt32(&closed)) +} + +func TestTraceLegacyMismatchAndRestart(t *testing.T) { + cfg := &primitive.TraceConfig{UnitName: t.Name(), NamesrvAddrs: []string{"127.0.0.1:9876"}} + first := NewTraceDispatcher(cfg) + require.NotNil(t, first) + defer first.Close() + second := NewTraceDispatcher(cfg) + require.NotNil(t, second) + defer second.Close() + require.Same(t, first.namesrvs, second.namesrvs) + cfg.NamesrvAddrs = []string{"127.0.0.2:9876"} + missing := NewTraceDispatcher(cfg) + require.Nil(t, missing) + require.True(t, IsNilTraceDispatcher(missing)) + missing.Start() + missing.Close() + require.False(t, missing.Append(TraceContext{})) + first.Close() + second.Close() + replacement := NewTraceDispatcher(cfg) + require.NotNil(t, replacement) + defer replacement.Close() + require.NotSame(t, first.resource, replacement.resource) + require.Equal(t, cfg.NamesrvAddrs, replacement.namesrvs.AddrList()) +} + +func TestSharedTraceFactoryFailures(t *testing.T) { + var created, closed int32 + require.Nil(t, NewTraceDispatcher(nil)) + require.Nil(t, NewTraceDispatcher(&primitive.TraceConfig{})) + require.Nil(t, NewSharedTraceDispatcher(&primitive.TraceConfig{}, primitive.SharedTraceClientConfig{})) + for _, mode := range []string{"error", "nil", "invalid", "empty"} { + t.Run(mode, func(t *testing.T) { + shared := primitive.SharedTraceClientConfig{Key: t.Name(), ResolverFactory: func() (primitive.NsResolver, func(), error) { + atomic.AddInt32(&created, 1) + cleanup := func() { atomic.AddInt32(&closed, 1) } + switch mode { + case "error": + return nil, cleanup, errors.New("discovery unavailable") + case "nil": + var resolver *mutableTraceResolver + return resolver, cleanup, nil + case "invalid": + return primitive.NewPassthroughResolver([]string{"invalid"}), cleanup, nil + default: + return &mutableTraceResolver{}, cleanup, nil + } + }} + require.Nil(t, NewSharedTraceDispatcher(&primitive.TraceConfig{}, shared)) + require.Nil(t, NewSharedTraceDispatcher(&primitive.TraceConfig{}, shared)) + }) + } + require.Equal(t, int32(8), created) + require.Equal(t, created, closed) +} + +func TestSharedTraceConcurrentAcquireAndClose(t *testing.T) { + var created, closed int32 + resolver := &mutableTraceResolver{addrs: []string{"127.0.0.1:9876"}} + shared := traceTestShared(t.Name(), resolver, &created, &closed) + anchor := NewSharedTraceDispatcher(&primitive.TraceConfig{}, shared) + require.NotNil(t, anchor) + defer anchor.Close() + var wg sync.WaitGroup + for i := 0; i < 40; i++ { + wg.Add(1) + go func() { + defer wg.Done() + for j := 0; j < 10; j++ { + td := NewSharedTraceDispatcher(&primitive.TraceConfig{}, shared) + if td == nil { + t.Error("acquire failed") + return + } + if td.resource != anchor.resource { + t.Error("duplicate transport") + } + var operations sync.WaitGroup + for k := 0; k < 3; k++ { + operations.Add(1) + go func() { defer operations.Done(); td.Start(); td.Append(TraceContext{}); td.Close() }() + } + operations.Wait() + } + }() + } + wg.Wait() + anchor.Close() + require.Equal(t, int32(1), atomic.LoadInt32(&created)) + require.Equal(t, int32(1), atomic.LoadInt32(&closed)) + // Exercise final release racing with a new acquisition, without an anchor. + for i := 0; i < 10; i++ { + wg.Add(1) + go func() { + defer wg.Done() + for j := 0; j < 10; j++ { + td := NewSharedTraceDispatcher(&primitive.TraceConfig{}, shared) + if td == nil { + t.Error("reacquire failed") + return + } + td.Start() + td.Close() + } + }() + } + wg.Wait() + require.Equal(t, atomic.LoadInt32(&created), atomic.LoadInt32(&closed)) + traceClients.Lock() + _, exists := traceClients.entries[anchor.resource.key] + traceClients.Unlock() + require.False(t, exists) +} + +func TestSharedTraceAcquireWaitsForCleanup(t *testing.T) { + cleanupEntered, allowCleanup := make(chan struct{}), make(chan struct{}) + var created int32 + shared := primitive.SharedTraceClientConfig{Key: t.Name(), ResolverFactory: func() (primitive.NsResolver, func(), error) { + generation := atomic.AddInt32(&created, 1) + return primitive.NewPassthroughResolver([]string{"127.0.0.1:9876"}), func() { + if generation == 1 { + close(cleanupEntered) + <-allowCleanup + } + }, nil + }} + first := NewSharedTraceDispatcher(&primitive.TraceConfig{}, shared) + require.NotNil(t, first) + go first.Close() + <-cleanupEntered + acquired := make(chan *traceDispatcher, 1) + go func() { acquired <- NewSharedTraceDispatcher(&primitive.TraceConfig{}, shared) }() + select { + case <-acquired: + t.Fatal("acquired while old resolver was closing") + case <-time.After(20 * time.Millisecond): + } + close(allowCleanup) + second := <-acquired + require.NotNil(t, second) + second.Close() + require.Equal(t, int32(2), atomic.LoadInt32(&created)) +} + +func testTraceRoute(broker string) *remote.RemotingCommand { + return &remote.RemotingCommand{Code: ResSuccess, Body: []byte(fmt.Sprintf(`{"queueDatas":[{"brokerName":"broker","readQueueNums":1,"writeQueueNums":1,"perm":6}],"brokerDatas":[{"cluster":"cluster","brokerName":"broker","brokerAddrs":{"0":%q}}]}`, broker))} +} + +func TestSharedTraceRefreshesBrokerRoute(t *testing.T) { + ctrl := gomock.NewController(t) + defer ctrl.Finish() + var created, closed int32 + resolver := &mutableTraceResolver{addrs: []string{"127.0.0.1:9876"}} + shared := traceTestShared(t.Name(), resolver, &created, &closed) + first := NewSharedTraceDispatcher(&primitive.TraceConfig{}, shared) + require.NotNil(t, first) + defer first.Close() + second := NewSharedTraceDispatcher(&primitive.TraceConfig{}, shared) + require.NotNil(t, second) + defer second.Close() + ns := remote.NewMockRemotingClient(ctrl) + first.namesrvs.nameSrvClient = ns + gomock.InOrder( + ns.EXPECT().InvokeSync(gomock.Any(), "127.0.0.1:9876", gomock.Any()).Return(testTraceRoute("127.0.0.1:10911"), nil), + ns.EXPECT().InvokeSync(gomock.Any(), "127.0.0.2:9876", gomock.Any()).Return(testTraceRoute("127.0.0.2:10911"), nil), + ns.EXPECT().ShutDown(), + ) + mq, addr := first.findMq("") + require.NotNil(t, mq) + require.Equal(t, "127.0.0.1:10911", addr) + resolver.set("127.0.0.2:9876") + first.namesrvs.UpdateNameServerAddress() + first.resource.refreshRoutes(context.Background()) + mq, addr = second.findMq("") + require.NotNil(t, mq) + require.Equal(t, "127.0.0.2:10911", addr) + // The trace-only resource doesn't need a producer/consumer registration. + require.Nil(t, first.namesrvs.bundleClient) +} + +func TestTraceCloseDrainsAndWaitsForSend(t *testing.T) { + ctrl := gomock.NewController(t) + defer ctrl.Finish() + td := NewTraceDispatcher(&primitive.TraceConfig{UnitName: t.Name(), NamesrvAddrs: []string{"127.0.0.1:9876"}}) + require.NotNil(t, td) + defer td.Close() + ns := remote.NewMockRemotingClient(ctrl) + td.namesrvs.nameSrvClient = ns + broker := remote.NewMockRemotingClient(ctrl) + td.resource.cli.remoteClient = broker + ns.EXPECT().InvokeSync(gomock.Any(), gomock.Any(), gomock.Any()).Return(testTraceRoute("127.0.0.1:10911"), nil) + sent := make(chan func(*remote.ResponseFuture), 1) + broker.EXPECT().InvokeAsync(gomock.Any(), "127.0.0.1:10911", gomock.Any(), gomock.Any()).DoAndReturn( + func(ctx context.Context, addr string, req *remote.RemotingCommand, cb func(*remote.ResponseFuture)) error { + require.Contains(t, string(req.Body), "test-message") + sent <- cb + return nil + }) + broker.EXPECT().ShutDown() + ns.EXPECT().ShutDown() + td.Start() + require.True(t, td.Append(TraceContext{TraceType: SubBefore, TraceBeans: []TraceBean{{Topic: "topic", MsgId: "test-message"}}})) + returned := make(chan struct{}) + go func() { td.Close(); close(returned) }() + callback := <-sent + select { + case <-returned: + t.Fatal("closed transport before callback") + case <-time.After(20 * time.Millisecond): + } + callback(&remote.ResponseFuture{ResponseCommand: &remote.RemotingCommand{Code: ResSuccess}}) + select { + case <-returned: + case <-time.After(time.Second): + t.Fatal("close did not finish") + } +} + +func TestTraceRouteCancellation(t *testing.T) { + ctrl := gomock.NewController(t) + defer ctrl.Finish() + ns, err := NewNamesrv(primitive.NewPassthroughResolver([]string{"127.0.0.1:9876", "127.0.0.2:9876"}), nil) + require.NoError(t, err) + remoteClient := remote.NewMockRemotingClient(ctrl) + ns.nameSrvClient = remoteClient + ctx, cancel := context.WithCancel(context.Background()) + remoteClient.EXPECT().InvokeSync(gomock.Any(), gomock.Any(), gomock.Any()).DoAndReturn( + func(ctx context.Context, _ string, _ *remote.RemotingCommand) (*remote.RemotingCommand, error) { + cancel() + return nil, ctx.Err() + }) + _, err = ns.queryTopicRouteInfoWithContext(ctx, "trace-topic") + require.Error(t, err) +} + +func TestTraceDiscoveryConcurrentReaders(t *testing.T) { + resolver := &mutableTraceResolver{addrs: []string{"127.0.0.1:9876"}} + td := NewTraceDispatcher(&primitive.TraceConfig{UnitName: t.Name(), Resolver: resolver}) + require.NotNil(t, td) + defer td.Close() + var wg sync.WaitGroup + wg.Add(2) + go func() { + defer wg.Done() + for i := 0; i < 1000; i++ { + resolver.set(fmt.Sprintf("127.0.0.%d:9876", i%2+1)) + td.namesrvs.UpdateNameServerAddress() + } + }() + go func() { + defer wg.Done() + for i := 0; i < 1000; i++ { + td.namesrvs.Size() + td.namesrvs.String() + td.namesrvs.getNameServerAddress() + snapshot := td.namesrvs.AddrList() + snapshot[0] = "must not mutate shared addresses" + } + }() + wg.Wait() + require.NotEqual(t, "must not mutate shared addresses", td.namesrvs.AddrList()[0]) +} + +func TestTraceSendFailuresReleaseOwnership(t *testing.T) { + for _, async := range []bool{false, true} { + t.Run(fmt.Sprint(async), func(t *testing.T) { + ctrl := gomock.NewController(t) + defer ctrl.Finish() + td := NewTraceDispatcher(&primitive.TraceConfig{UnitName: t.Name(), Access: primitive.Cloud, NamesrvAddrs: []string{"127.0.0.1:9876"}}) + require.NotNil(t, td) + defer td.Close() + ns := remote.NewMockRemotingClient(ctrl) + td.namesrvs.nameSrvClient = ns + broker := remote.NewMockRemotingClient(ctrl) + td.resource.cli.remoteClient = broker + ns.EXPECT().InvokeSync(gomock.Any(), gomock.Any(), gomock.Any()).Return(testTraceRoute("127.0.0.1:10911"), nil) + broker.EXPECT().InvokeAsync(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).DoAndReturn( + func(ctx context.Context, addr string, req *remote.RemotingCommand, cb func(*remote.ResponseFuture)) error { + err := errors.New("trace send failed") + if async { + cb(&remote.ResponseFuture{Err: err}) + return nil + } + return err + }) + broker.EXPECT().ShutDown() + ns.EXPECT().ShutDown() + td.Start() + require.True(t, td.Append(TraceContext{TraceType: SubBefore, RegionId: "region", TraceBeans: []TraceBean{{Topic: "topic", MsgId: "message"}}})) + td.Close() + _, ok := td.resource.topics.Load(td.GetTraceTopicName() + "region") + require.True(t, ok) + select { + case <-td.closeDone: + default: + t.Fatal("failed send leaked the dispatcher") + } + }) + } +} diff --git a/internal/trace_route_test.go b/internal/trace_route_test.go new file mode 100644 index 00000000..4e746e14 --- /dev/null +++ b/internal/trace_route_test.go @@ -0,0 +1,126 @@ +/* +Licensed to the Apache Software Foundation (ASF) under one or more +contributor license agreements. See the NOTICE file distributed with +this work for additional information regarding copyright ownership. +The ASF licenses this file to You under the Apache License, Version 2.0 +(the "License"); you may not use this file except in compliance with +the License. You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package internal + +import ( + "context" + "fmt" + "sync" + "testing" + + "github.com/apache/rocketmq-client-go/v2/internal/remote" + "github.com/apache/rocketmq-client-go/v2/primitive" + "github.com/golang/mock/gomock" + "github.com/stretchr/testify/require" +) + +const traceRouteTestBrokers = 8 + +func multiBrokerTraceRoute(generation int) *TopicRouteData { + route := &TopicRouteData{} + // Deliberately return reverse order so publish queue selection must sort. + for i := traceRouteTestBrokers - 1; i >= 0; i-- { + name := fmt.Sprintf("broker-%02d", i) + route.QueueDataList = append(route.QueueDataList, &QueueData{ + BrokerName: name, ReadQueueNums: 1, WriteQueueNums: 1, Perm: 6, + }) + route.BrokerDataList = append(route.BrokerDataList, &BrokerData{ + Cluster: "cluster", BrokerName: name, + BrokerAddresses: map[int64]string{MasterId: fmt.Sprintf("127.0.0.%d:%d", i+1, 10911+generation%2)}, + }) + } + return route +} + +func TestTraceCachedRoutePreservesSnapshot(t *testing.T) { + ns, err := NewNamesrv(primitive.NewPassthroughResolver([]string{"127.0.0.1:9876"}), nil) + require.NoError(t, err) + defer ns.nameSrvClient.ShutDown() + route := multiBrokerTraceRoute(0) + before := append([]*QueueData(nil), route.QueueDataList...) + ns.routeDataMap.Store("trace-topic", route) + queues, err := ns.fetchPublishMessageQueuesWithContext(context.Background(), "trace-topic") + require.NoError(t, err) + require.Len(t, queues, traceRouteTestBrokers) + for i, queue := range queues { + require.Equal(t, fmt.Sprintf("broker-%02d", i), queue.BrokerName) + } + require.Equal(t, before, route.QueueDataList, "publish queue sorting must not reorder the cached snapshot") +} + +func TestSharedTraceConcurrentMultiBrokerRouteRefresh(t *testing.T) { + ctrl := gomock.NewController(t) + defer ctrl.Finish() + var created, closed int32 + shared := traceTestShared(t.Name(), primitive.NewPassthroughResolver([]string{"127.0.0.1:9876"}), &created, &closed) + first := NewSharedTraceDispatcher(&primitive.TraceConfig{}, shared) + require.NotNil(t, first) + defer first.Close() + second := NewSharedTraceDispatcher(&primitive.TraceConfig{}, shared) + require.NotNil(t, second) + defer second.Close() + require.Same(t, first.namesrvs, second.namesrvs) + topic := first.GetTraceTopicName() + first.namesrvs.routeDataMap.Store(topic, multiBrokerTraceRoute(0)) + first.resource.topics.Store(topic, struct{}{}) + + ns := remote.NewMockRemotingClient(ctrl) + first.namesrvs.nameSrvClient = ns + generation := 0 + const refreshes = 200 + ns.EXPECT().InvokeSync(gomock.Any(), gomock.Any(), gomock.Any()).DoAndReturn( + func(context.Context, string, *remote.RemotingCommand) (*remote.RemotingCommand, error) { + generation++ + return &remote.RemotingCommand{Code: ResSuccess, Body: []byte(multiBrokerTraceRoute(generation).String())}, nil + }).Times(refreshes) + ns.EXPECT().ShutDown() + + start := make(chan struct{}) + var wg sync.WaitGroup + wg.Add(1) + go func() { + defer wg.Done() + <-start + for i := 0; i < refreshes; i++ { + first.resource.refreshRoutes(context.Background()) + } + }() + for i := 0; i < 4; i++ { + dispatcher := []*traceDispatcher{first, second}[i%2] + wg.Add(1) + go func(td *traceDispatcher) { + defer wg.Done() + <-start + for j := 0; j < 1000; j++ { + queues, err := td.namesrvs.fetchPublishMessageQueuesWithContext(context.Background(), topic) + if err != nil || len(queues) != traceRouteTestBrokers { + t.Errorf("incomplete route snapshot: %d queues, error %v", len(queues), err) + return + } + for k, queue := range queues { + if queue.BrokerName != fmt.Sprintf("broker-%02d", k) || queue.QueueId != 0 { + t.Errorf("inconsistent route snapshot at queue %d: %+v", k, queue) + return + } + } + } + }(dispatcher) + } + close(start) + wg.Wait() +} diff --git a/primitive/trace.go b/primitive/trace.go index 4ddb2df9..c2992c69 100644 --- a/primitive/trace.go +++ b/primitive/trace.go @@ -27,3 +27,17 @@ type TraceConfig struct { Resolver NsResolver Credentials // acl config for trace. omit if acl is closed on broker. } + +// SharedTraceClientConfig identifies a trace transport independently of the +// current NameServer addresses. Use a stable key for the logical cluster (for +// example its discovery endpoint and tenant), never a consumer group or an IP +// list. UnitName, Access and Credentials from TraceConfig also partition clients. +type SharedTraceClientConfig struct { + Key string + // ResolverFactory is called once per shared client's lifetime. It must return + // an independently owned resolver, not a consumer's resolver. The optional + // cleanup function is called once after the last dispatcher closes (also on + // initialization failure). Resolve must return promptly and be safe for + // concurrent use. NamesrvAddrs and Resolver in TraceConfig are ignored. + ResolverFactory func() (NsResolver, func(), error) +} diff --git a/producer/interceptor.go b/producer/interceptor.go index b86599b1..834b9bce 100644 --- a/producer/interceptor.go +++ b/producer/interceptor.go @@ -23,7 +23,6 @@ package producer import ( "context" "fmt" - "reflect" "time" "github.com/apache/rocketmq-client-go/v2/internal" @@ -34,24 +33,39 @@ import ( // WithTrace support rocketmq trace: https://github.com/apache/rocketmq/wiki/RIP-6-Message-Trace. func WithTrace(traceCfg *primitive.TraceConfig) Option { return func(options *producerOptions) { - dispatcher := internal.NewTraceDispatcher(traceCfg) - options.TraceDispatcher = dispatcher - ori := options.Interceptors - options.Interceptors = make([]primitive.Interceptor, 0) - options.Interceptors = append(options.Interceptors, newTraceInterceptor(dispatcher)) - options.Interceptors = append(options.Interceptors, ori...) + installTraceInterceptor(options, internal.NewTraceDispatcher(traceCfg)) } } +// WithSharedTrace shares trace connections by logical cluster identity. Discovery +// addresses may change without recreating the client. See SharedTraceClientConfig +// for resolver ownership and key requirements. +func WithSharedTrace(traceCfg *primitive.TraceConfig, shared primitive.SharedTraceClientConfig) Option { + return func(options *producerOptions) { + installTraceInterceptor(options, internal.NewSharedTraceDispatcher(traceCfg, shared)) + } +} + +func installTraceInterceptor(options *producerOptions, dispatcher internal.TraceDispatcher) { + if internal.IsNilTraceDispatcher(dispatcher) { + return + } + if !internal.IsNilTraceDispatcher(options.TraceDispatcher) { + options.TraceDispatcher.Close() + } + options.TraceDispatcher = dispatcher + options.Interceptors = append([]primitive.Interceptor{newTraceInterceptor(dispatcher)}, options.Interceptors...) +} + func newTraceInterceptor(dispatcher internal.TraceDispatcher) primitive.Interceptor { - if dispatcher != nil && !reflect.ValueOf(dispatcher).IsNil() { - dispatcher.Start() + if internal.IsNilTraceDispatcher(dispatcher) { + return func(ctx context.Context, req, reply interface{}, next primitive.Invoker) error { + return next(ctx, req, reply) + } } + dispatcher.Start() return func(ctx context.Context, req, reply interface{}, next primitive.Invoker) error { - if dispatcher == nil { - return fmt.Errorf("GetOrNewRocketMQClient faild") - } beginT := time.Now() producerCtx, ok := primitive.GetProducerCtx(ctx) if !ok { diff --git a/producer/producer.go b/producer/producer.go index 2cbeff68..57dc6af0 100644 --- a/producer/producer.go +++ b/producer/producer.go @@ -53,6 +53,12 @@ type defaultProducer struct { func NewDefaultProducer(opts ...Option) (*defaultProducer, error) { defaultOpts := defaultProducerOptions() + constructed := false + defer func() { + if !constructed && !internal.IsNilTraceDispatcher(defaultOpts.TraceDispatcher) { + defaultOpts.TraceDispatcher.Close() + } + }() for _, apply := range opts { apply(&defaultOpts) } @@ -78,11 +84,17 @@ func NewDefaultProducer(opts ...Option) (*defaultProducer, error) { producer.interceptor = primitive.ChainInterceptors(producer.options.Interceptors...) + constructed = true return producer, nil } func (p *defaultProducer) Start() error { var err error + defer func() { + if err != nil && !internal.IsNilTraceDispatcher(p.options.TraceDispatcher) { + p.options.TraceDispatcher.Close() + } + }() p.startOnce.Do(func() { err = p.client.RegisterProducer(p.group, p) if err != nil { diff --git a/producer/trace_lifecycle_test.go b/producer/trace_lifecycle_test.go new file mode 100644 index 00000000..76e8522e --- /dev/null +++ b/producer/trace_lifecycle_test.go @@ -0,0 +1,108 @@ +/* +Licensed to the Apache Software Foundation (ASF) under one or more +contributor license agreements. See the NOTICE file distributed with +this work for additional information regarding copyright ownership. +The ASF licenses this file to You under the Apache License, Version 2.0 +(the "License"); you may not use this file except in compliance with +the License. You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package producer + +import ( + "context" + "errors" + "sync/atomic" + "testing" + + "github.com/apache/rocketmq-client-go/v2/internal" + "github.com/apache/rocketmq-client-go/v2/primitive" + "github.com/golang/mock/gomock" + "github.com/stretchr/testify/require" +) + +type unavailableTrace struct{} + +func (*unavailableTrace) Start() { panic("nil dispatcher started") } +func (*unavailableTrace) Close() { panic("nil dispatcher closed") } +func (*unavailableTrace) Append(internal.TraceContext) bool { panic("nil dispatcher invoked") } +func (*unavailableTrace) GetTraceTopicName() string { panic("nil dispatcher invoked") } + +func TestUnavailableTraceInvokesBusinessCallback(t *testing.T) { + var typedNil *unavailableTrace + for _, dispatcher := range []internal.TraceDispatcher{nil, typedNil} { + interceptor := newTraceInterceptor(dispatcher) + for _, businessError := range []error{nil, errors.New("business error")} { + calls := 0 + ctx := context.WithValue(context.Background(), "test-key", "test-value") + ctx = primitive.WithProducerCtx(ctx, &primitive.ProducerCtx{ + ProducerGroup: "test-group", + Message: *primitive.NewMessage("topic", []byte("payload")), + }) + request, reply := new(int), new(int) + err := interceptor(ctx, request, reply, func(actual context.Context, req, resp interface{}) error { + calls++ + require.Equal(t, ctx, actual) + require.Same(t, request, req) + require.Same(t, reply, resp) + *resp.(*int) = 42 + return businessError + }) + require.Equal(t, 1, calls) + require.Equal(t, 42, *reply) + require.Equal(t, businessError, err) + } + } +} + +func TestFailedTraceInstallationIsOptional(t *testing.T) { + options := defaultProducerOptions() + originalCount := len(options.Interceptors) + WithTrace(&primitive.TraceConfig{})(&options) + require.Nil(t, options.TraceDispatcher) + require.Len(t, options.Interceptors, originalCount) + WithSharedTrace(&primitive.TraceConfig{}, primitive.SharedTraceClientConfig{Key: "missing-factory"})(&options) + require.Nil(t, options.TraceDispatcher) + require.Len(t, options.Interceptors, originalCount) +} + +func TestTraceReleasedOnConstructionFailure(t *testing.T) { + var created, closed int32 + shared := primitive.SharedTraceClientConfig{Key: t.Name(), ResolverFactory: func() (primitive.NsResolver, func(), error) { + atomic.AddInt32(&created, 1) + return primitive.NewPassthroughResolver([]string{"127.0.0.1:9876"}), func() { atomic.AddInt32(&closed, 1) }, nil + }} + for i := 0; i < 3; i++ { + _, err := NewDefaultProducer( + WithSharedTrace(&primitive.TraceConfig{}, shared), + WithNsResolver(primitive.NewPassthroughResolver([]string{"invalid"})), + ) + require.Error(t, err) + } + require.Equal(t, int32(3), atomic.LoadInt32(&created)) + require.Equal(t, atomic.LoadInt32(&created), atomic.LoadInt32(&closed)) +} + +func TestTraceReleasedOnStartFailure(t *testing.T) { + ctrl := gomock.NewController(t) + defer ctrl.Finish() + var closed int32 + opts := defaultProducerOptions() + shared := primitive.SharedTraceClientConfig{Key: t.Name(), ResolverFactory: func() (primitive.NsResolver, func(), error) { + return primitive.NewPassthroughResolver([]string{"127.0.0.1:9876"}), func() { atomic.AddInt32(&closed, 1) }, nil + }} + WithSharedTrace(&primitive.TraceConfig{}, shared)(&opts) + client := internal.NewMockRMQClient(ctrl) + client.EXPECT().RegisterProducer(gomock.Any(), gomock.Any()).Return(errors.New("duplicate group")) + p := &defaultProducer{client: client, options: opts} + require.Error(t, p.Start()) + require.Equal(t, int32(1), atomic.LoadInt32(&closed)) +} From 8ba014196e2fc17e3ddc035dc410db5a42ae24e4 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E6=96=87=E5=BE=90?= Date: Wed, 30 Sep 2026 16:58:23 +0800 Subject: [PATCH 6/6] fix: preserve trace compatibility and shared client lifecycle Balance client ownership from construction through shutdown, preserve trace batching and NameServer failover, and safely replace trace interceptors. Separate address discovery from route refresh, reuse concurrent route lookups, publish broker addresses before queues, and avoid mutating resolver snapshots. Add lifecycle, migration, concurrency, and timeout regression coverage. --- client_lifecycle_test.go | 97 +++++++++++++++ consumer/consumer.go | 15 +-- consumer/interceptor.go | 3 + consumer/pull_consumer.go | 22 +++- consumer/push_consumer.go | 22 +++- consumer/trace_lifecycle_test.go | 36 ++++++ consumer/trace_replacement_test.go | 89 ++++++++++++++ docs/trace.md | 29 +++-- internal/client.go | 57 ++++++--- internal/client_lifecycle_test.go | 96 +++++++++++++++ internal/route.go | 28 ++++- internal/trace.go | 20 +-- internal/trace_client.go | 40 ++++-- internal/trace_compatibility_test.go | 177 +++++++++++++++++++++++++++ internal/trace_migration_test.go | 150 +++++++++++++++++++++++ internal/trace_route_test.go | 159 ++++++++++++++++++++++++ producer/interceptor.go | 3 + producer/producer.go | 16 ++- producer/trace_lifecycle_test.go | 55 +++++++++ producer/trace_replacement_test.go | 87 +++++++++++++ 20 files changed, 1140 insertions(+), 61 deletions(-) create mode 100644 client_lifecycle_test.go create mode 100644 consumer/trace_replacement_test.go create mode 100644 internal/client_lifecycle_test.go create mode 100644 internal/trace_compatibility_test.go create mode 100644 internal/trace_migration_test.go create mode 100644 producer/trace_replacement_test.go diff --git a/client_lifecycle_test.go b/client_lifecycle_test.go new file mode 100644 index 00000000..1f2dab53 --- /dev/null +++ b/client_lifecycle_test.go @@ -0,0 +1,97 @@ +/* +Licensed to the Apache Software Foundation (ASF) under one or more +contributor license agreements. See the NOTICE file distributed with +this work for additional information regarding copyright ownership. +The ASF licenses this file to You under the Apache License, Version 2.0 +(the "License"); you may not use this file except in compliance with +the License. You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package rocketmq_test + +import ( + "testing" + + rocketmq "github.com/apache/rocketmq-client-go/v2" + "github.com/apache/rocketmq-client-go/v2/consumer" + "github.com/apache/rocketmq-client-go/v2/internal" + "github.com/apache/rocketmq-client-go/v2/producer" + "github.com/stretchr/testify/require" +) + +type clientOwner interface{ Shutdown() error } + +func TestClientOwnershipBeforeStart(t *testing.T) { + constructors := map[string]func(string, string) (clientOwner, error){ + "producer": func(instance, address string) (clientOwner, error) { + return rocketmq.NewProducer(producer.WithInstanceName(instance), producer.WithNameServer([]string{address})) + }, + "push": func(instance, address string) (clientOwner, error) { + return rocketmq.NewPushConsumer(consumer.WithInstance(instance), consumer.WithNameServer([]string{address})) + }, + "pull": func(instance, address string) (clientOwner, error) { + return rocketmq.NewPullConsumer(consumer.WithInstance(instance), consumer.WithNameServer([]string{address})) + }, + } + for name, create := range constructors { + t.Run(name, func(t *testing.T) { + instance := t.Name() + clientID := internal.DefaultClientOptions().ClientIP + "@" + instance + first, err := create(instance, "127.0.0.1:9876") + require.NoError(t, err) + defer first.Shutdown() + second, err := create(instance, "127.0.0.1:9876") + require.NoError(t, err) + defer second.Shutdown() + require.NoError(t, first.Shutdown()) + require.NoError(t, first.Shutdown()) + _, err = internal.GetNamesrv(clientID) + require.NoError(t, err, "another unstarted owner still needs the shared client") + _, err = create(instance, "127.0.0.2:9876") + require.Error(t, err, "a live owner's NameServer compatibility check must remain in force") + require.NoError(t, second.Shutdown()) + _, err = internal.GetNamesrv(clientID) + require.Error(t, err, "the final owner must release the registry entry, including after a rejected acquisition") + replacement, err := create(instance, "127.0.0.2:9876") + require.NoError(t, err, "the same instance must be reusable with new addresses") + require.NoError(t, replacement.Shutdown()) + }) + } +} + +func TestClientOwnershipProtectsRunningProducer(t *testing.T) { + instance := t.Name() + clientID := internal.DefaultClientOptions().ClientIP + "@" + instance + create := func(group string) (rocketmq.Producer, error) { + return rocketmq.NewProducer(producer.WithInstanceName(instance), producer.WithGroupName(group), producer.WithNameServer([]string{"127.0.0.1:9876"})) + } + active, err := create("active-group") + require.NoError(t, err) + defer active.Shutdown() + require.NoError(t, active.Start()) + require.NoError(t, active.Start()) + for i := 0; i < 2; i++ { + duplicate, err := create("active-group") + require.NoError(t, err) + require.Error(t, duplicate.Start(), "failed owners must not unregister the active producer") + require.NoError(t, duplicate.Shutdown()) + } + pending, err := create("pending-group") + require.NoError(t, err) + defer pending.Shutdown() + require.NoError(t, active.Shutdown()) + _, err = internal.GetNamesrv(clientID) + require.NoError(t, err, "the unstarted owner must retain its client after the active owner exits") + require.NoError(t, pending.Start()) + require.NoError(t, pending.Shutdown()) + _, err = internal.GetNamesrv(clientID) + require.Error(t, err) +} diff --git a/consumer/consumer.go b/consumer/consumer.go index 98b44477..34f1cb5c 100644 --- a/consumer/consumer.go +++ b/consumer/consumer.go @@ -247,13 +247,14 @@ type defaultConsumer struct { fromWhere ConsumeFromWhere consumerStartTimestamp int64 - cType ConsumeType - client internal.RMQClient - mqChanged func(topic string, mqAll, mqDivided []*primitive.MessageQueue) - state *atomic.Int32 - pause *atomic.Bool - once sync.Once - option consumerOptions + cType ConsumeType + client internal.RMQClient + mqChanged func(topic string, mqAll, mqDivided []*primitive.MessageQueue) + state *atomic.Int32 + pause *atomic.Bool + lifecycleMu sync.Mutex + once sync.Once + option consumerOptions // key: primitive.MessageQueue // value: *processQueue processQueueTable sync.Map diff --git a/consumer/interceptor.go b/consumer/interceptor.go index e93ed55b..f14b2a63 100644 --- a/consumer/interceptor.go +++ b/consumer/interceptor.go @@ -48,6 +48,9 @@ func installTraceInterceptor(options *consumerOptions, dispatcher internal.Trace } if !internal.IsNilTraceDispatcher(options.TraceDispatcher) { options.TraceDispatcher.Close() + // Trace is always prepended; WithInterceptor appends user interceptors. + // Remove the old wrapper before installing its replacement. + options.Interceptors = options.Interceptors[1:] } options.TraceDispatcher = dispatcher options.Interceptors = append([]primitive.Interceptor{newTraceInterceptor(dispatcher)}, options.Interceptors...) diff --git a/consumer/pull_consumer.go b/consumer/pull_consumer.go index b0aac921..8c383057 100644 --- a/consumer/pull_consumer.go +++ b/consumer/pull_consumer.go @@ -233,6 +233,11 @@ func (pc *defaultPullConsumer) nextPullOffset(mq *primitive.MessageQueue, origin } func (pc *defaultPullConsumer) Start() error { + pc.lifecycleMu.Lock() + defer pc.lifecycleMu.Unlock() + if pc.state.Load() == int32(internal.StateShutdown) { + return internal.ErrServiceState + } var err error defer func() { if err != nil && !internal.IsNilTraceDispatcher(pc.option.TraceDispatcher) { @@ -284,7 +289,7 @@ func (pc *defaultPullConsumer) Start() error { pc.client.UpdateTopicRouteInfo() _, exist := pc.topicSubscribeInfoTable.Load(pc.topic) if !exist { - err = pc.Shutdown() + err = pc.shutdownLocked() if err != nil { rlog.Error("defaultPullConsumer.Shutdown . route info not found, it may not exist", map[string]interface{}{ rlog.LogKeyTopic: pc.topic, @@ -582,12 +587,25 @@ func (pc *defaultPullConsumer) CurrentOffset(queue *primitive.MessageQueue) (int // Shutdown close defaultConsumer, refuse new request. func (pc *defaultPullConsumer) Shutdown() error { + pc.lifecycleMu.Lock() + defer pc.lifecycleMu.Unlock() + return pc.shutdownLocked() +} + +func (pc *defaultPullConsumer) shutdownLocked() error { var err error pc.closeOnce.Do(func() { - if pc.option.TraceDispatcher != nil { + if !internal.IsNilTraceDispatcher(pc.option.TraceDispatcher) { pc.option.TraceDispatcher.Close() } close(pc.done) + // Release the construction reference without unregistering another + // owner's group when this consumer never started successfully. + if pc.state.Load() != int32(internal.StateRunning) { + pc.state.Store(int32(internal.StateShutdown)) + pc.client.Shutdown() + return + } pc.client.UnregisterConsumer(pc.consumerGroup) err = pc.defaultConsumer.shutdown() diff --git a/consumer/push_consumer.go b/consumer/push_consumer.go index f18af466..2d640bc1 100644 --- a/consumer/push_consumer.go +++ b/consumer/push_consumer.go @@ -137,6 +137,11 @@ func NewPushConsumer(opts ...Option) (*pushConsumer, error) { } func (pc *pushConsumer) Start() error { + pc.lifecycleMu.Lock() + defer pc.lifecycleMu.Unlock() + if pc.state.Load() == int32(internal.StateShutdown) { + return internal.ErrServiceState + } var err error defer func() { if err != nil && !internal.IsNilTraceDispatcher(pc.option.TraceDispatcher) { @@ -249,7 +254,7 @@ func (pc *pushConsumer) Start() error { pc.subscribedTopic.Range(func(k, v interface{}) bool { _, exist := pc.topicSubscribeInfoTable.Load(k) if !exist { - pc.Shutdown() + pc.shutdownLocked() err = fmt.Errorf("the topic=%s route info not found, it may not exist", k) return false } @@ -289,12 +294,25 @@ func (pc *pushConsumer) GetOffsetDiffMap() map[string]int64 { } func (pc *pushConsumer) Shutdown() error { + pc.lifecycleMu.Lock() + defer pc.lifecycleMu.Unlock() + return pc.shutdownLocked() +} + +func (pc *pushConsumer) shutdownLocked() error { var err error pc.closeOnce.Do(func() { - if pc.option.TraceDispatcher != nil { + if !internal.IsNilTraceDispatcher(pc.option.TraceDispatcher) { pc.option.TraceDispatcher.Close() } close(pc.done) + // Release the construction reference without unregistering another + // owner's group when this consumer never started successfully. + if pc.state.Load() != int32(internal.StateRunning) { + pc.state.Store(int32(internal.StateShutdown)) + pc.client.Shutdown() + return + } if pc.consumeOrderly && pc.model == Clustering { pc.unlockAll(false) } diff --git a/consumer/trace_lifecycle_test.go b/consumer/trace_lifecycle_test.go index 7b140ae1..2b2d721f 100644 --- a/consumer/trace_lifecycle_test.go +++ b/consumer/trace_lifecycle_test.go @@ -25,6 +25,7 @@ import ( "github.com/apache/rocketmq-client-go/v2/internal" "github.com/apache/rocketmq-client-go/v2/primitive" + "github.com/golang/mock/gomock" "github.com/stretchr/testify/require" atomic2 "go.uber.org/atomic" ) @@ -129,3 +130,38 @@ func TestTraceReleasedOnStartFailure(t *testing.T) { require.Equal(t, int32(1), atomic.LoadInt32(&closed)) } } + +// An unsuccessful Start releases its client reference without unregistering a group. +func TestTraceShutdownBeforeStartOrAfterDuplicateGroup(t *testing.T) { + for _, pull := range []bool{false, true} { + for _, failed := range []bool{false, true} { + ctrl := gomock.NewController(t) + client := internal.NewMockRMQClient(ctrl) + client.EXPECT().Shutdown() + opts := defaultPushConsumerOptions() + var closed int32 + WithSharedTrace(&primitive.TraceConfig{}, primitive.SharedTraceClientConfig{Key: t.Name(), ResolverFactory: func() (primitive.NsResolver, func(), error) { + return primitive.NewPassthroughResolver([]string{"127.0.0.1:9876"}), func() { atomic.AddInt32(&closed, 1) }, nil + }})(&opts) + dc := &defaultConsumer{consumerGroup: "trace-test", option: opts, client: client, state: atomic2.NewInt32(int32(internal.StateCreateJust))} + var c interface { + Start() error + Shutdown() error + } + if pull { + c = &defaultPullConsumer{defaultConsumer: dc, done: make(chan struct{}), SubType: Assign} + } else { + c = &pushConsumer{defaultConsumer: dc, done: make(chan struct{})} + } + if failed { + client.EXPECT().RegisterConsumer(gomock.Any(), gomock.Any()).Return(errors.New("duplicate group")) + require.Error(t, c.Start()) + } + require.NoError(t, c.Shutdown()) + require.NoError(t, c.Shutdown()) + require.Error(t, c.Start()) + require.Equal(t, int32(1), atomic.LoadInt32(&closed)) + ctrl.Finish() + } + } +} diff --git a/consumer/trace_replacement_test.go b/consumer/trace_replacement_test.go new file mode 100644 index 00000000..8c9a4d64 --- /dev/null +++ b/consumer/trace_replacement_test.go @@ -0,0 +1,89 @@ +/* +Licensed to the Apache Software Foundation (ASF) under one or more +contributor license agreements. See the NOTICE file distributed with +this work for additional information regarding copyright ownership. +The ASF licenses this file to You under the Apache License, Version 2.0 +(the "License"); you may not use this file except in compliance with +the License. You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package consumer + +import ( + "context" + "errors" + "testing" + + "github.com/apache/rocketmq-client-go/v2/internal" + "github.com/apache/rocketmq-client-go/v2/primitive" + "github.com/stretchr/testify/require" +) + +type replacementTraceDispatcher struct { + starts, closes int + records []internal.TraceContext +} + +func (d *replacementTraceDispatcher) Start() { d.starts++ } +func (d *replacementTraceDispatcher) Close() { d.closes++ } +func (*replacementTraceDispatcher) GetTraceTopicName() string { return "trace-topic" } +func (d *replacementTraceDispatcher) Append(ctx internal.TraceContext) bool { + d.records = append(d.records, ctx) + return d.closes == 0 +} + +func TestTraceReplacementPreservesInterceptors(t *testing.T) { + options := defaultPushConsumerOptions() + var calls []string + userInterceptor := func(name string) primitive.Interceptor { + return func(ctx context.Context, req, reply interface{}, next primitive.Invoker) error { + calls = append(calls, name) + return next(ctx, req, reply) + } + } + WithInterceptor(userInterceptor("first"))(&options) + old := &replacementTraceDispatcher{} + installTraceInterceptor(&options, old) + WithInterceptor(userInterceptor("second"))(&options) + middle, current := &replacementTraceDispatcher{}, &replacementTraceDispatcher{} + installTraceInterceptor(&options, middle) + installTraceInterceptor(&options, current) + // An unavailable replacement must preserve the working dispatcher and chain. + var typedNil *replacementTraceDispatcher + for _, unavailable := range []internal.TraceDispatcher{nil, typedNil} { + installTraceInterceptor(&options, unavailable) + } + require.Same(t, current, options.TraceDispatcher) + require.Equal(t, 1, old.closes) + require.Equal(t, 1, middle.closes) + require.Zero(t, current.closes) + for _, dispatcher := range []*replacementTraceDispatcher{old, middle, current} { + require.Equal(t, 1, dispatcher.starts) + } + ctx := primitive.WithConsumerCtx(context.Background(), &primitive.ConsumeMessageContext{ + ConsumerGroup: "group", + Msgs: []*primitive.MessageExt{{Message: *primitive.NewMessage("topic", []byte("payload"))}}, + Properties: map[string]string{primitive.PropCtxType: string(primitive.FailedReturn)}, + }) + var reply interface{} + businessError := errors.New("business error") + err := primitive.ChainInterceptors(options.Interceptors...)(ctx, nil, reply, + func(context.Context, interface{}, interface{}) error { + calls = append(calls, "business") + return businessError + }) + require.Equal(t, businessError, err) + require.Equal(t, []string{"first", "second", "business"}, calls) + require.Empty(t, old.records, "replaced dispatcher must not receive trace records") + require.Empty(t, middle.records, "repeated replacement must remove each stale interceptor") + require.Len(t, current.records, 2) + require.Len(t, options.Interceptors, 3) +} diff --git a/docs/trace.md b/docs/trace.md index 812f6682..2fc926b4 100644 --- a/docs/trace.md +++ b/docs/trace.md @@ -33,6 +33,10 @@ defer c.Shutdown() // Subscribe, then Start as usual. ``` +The example uses fixed addresses. To handle NameServer replacement, the factory +must return a resolver that discovers current addresses on each `Resolve` call. +The factory itself is not called again when the addresses change. + The key must identify the discovery source and destination, including any tenant or namespace that affects routing. Do not use a consumer group, a resolved IP list, or a changing address-list hash. `TraceConfig.UnitName`, `Access`, and the @@ -57,15 +61,24 @@ disabled for that dispatcher; recreating the consumer/producer retries setup. Each consumer/producer has its own trace dispatcher and buffer. Dispatchers in the same partition share one NameServer connection pool, one Broker connection pool, -one resolver, and one discovery worker. Adding consumers therefore does not add -independent trace connection pools for that partition. Connections are opened -lazily; this is not a promise of exactly one TCP connection per cluster. +one resolver, and shared address and route refresh workers. Adding consumers +therefore does not add independent trace connection pools for that partition. +Connections are opened lazily; this is not a promise of exactly one TCP +connection per cluster. NameServer addresses refresh after 10 seconds and then every 2 minutes. Trace topic routes refresh every 30 seconds, including region-specific cloud trace topics. NameServer discovery and Broker route discovery are separate operations; -updates are periodic, not immediate. All dispatchers read the same route and -address state. Closing one dispatcher leaves other users active. +updates are periodic, not immediate. Address discovery runs independently so a +slow route query cannot delay discovering replacement NameServers. Concurrent +cache misses reuse the first successful route lookup; scheduled refreshes still +query the servers even when a cached route exists. All dispatchers read the same +route and address state. Closing one dispatcher leaves other users active. + +Route queries retain the existing 6-second timeout per NameServer and try the +next address after a failed attempt. Shutdown can cancel the whole query. +Trace records are sent in batches of 100, or after more than 5 milliseconds +without a new record, preserving the legacy batching policy. Shutdown rejects new records, drains accepted records, and waits for in-flight sends. Close waits at most 5 seconds, then cancels outstanding I/O; cleanup finishes @@ -80,5 +93,7 @@ addresses before reuse. Its resolver is borrowed and must remain usable for the whole shared lifetime. It also benefits from reference-counted cleanup, a common NameServer object, route refresh, and nil-safe interceptors. Use `WithSharedTrace` when separate resolvers can return different address snapshots for the same -logical cluster. The general producer/consumer client registry and its -NameServer conflict checks are unchanged. +logical cluster. The general producer/consumer client registry retains its +NameServer conflict checks. Each successful client acquisition holds a reference +until shutdown, even if the producer or consumer has not started. Closing the +last owner removes the registry entry so that the instance can be created again. diff --git a/internal/client.go b/internal/client.go index 6f27f5e8..c302cf77 100644 --- a/internal/client.go +++ b/internal/client.go @@ -26,7 +26,6 @@ import ( "strconv" "strings" "sync" - "sync/atomic" "time" errors2 "github.com/apache/rocketmq-client-go/v2/errors" @@ -183,7 +182,10 @@ type rmqClient struct { done chan struct{} shutdownOnce sync.Once - instanceCount int32 + // Each successful GetOrNewRocketMQClient owns one reference, including + // clients that have not started yet. refs is protected by clientMapMu. + refs int + ready chan struct{} } func (c *rmqClient) GetNameSrv() Namesrvs { @@ -191,25 +193,37 @@ func (c *rmqClient) GetNameSrv() Namesrvs { } var clientMap sync.Map +var clientMapMu sync.Mutex +// GetOrNewRocketMQClient acquires a reference that must be released by Shutdown, +// whether or not the caller starts the client. func GetOrNewRocketMQClient(option ClientOptions, callbackCh chan interface{}) RMQClient { client := &rmqClient{ option: option, remoteClient: remote.NewRemotingClient(option.RemotingClientConfig), done: make(chan struct{}), + ready: make(chan struct{}), } + clientMapMu.Lock() actual, loaded := clientMap.LoadOrStore(client.ClientID(), client) + client = actual.(*rmqClient) + client.refs++ + clientMapMu.Unlock() if loaded { + // Do not expose a shared client before its request handlers are ready. + <-client.ready // compare namesrv address - client = actual.(*rmqClient) - now := option.Namesrv.(*namesrvs).resolver.Resolve() - old := client.GetNameSrv().(*namesrvs).resolver.Resolve() + // Resolvers may return their own shared slices. Compare private copies + // so sorting cannot mutate discovery state or race with another owner. + now := append([]string(nil), option.Namesrv.(*namesrvs).resolver.Resolve()...) + old := append([]string(nil), client.GetNameSrv().(*namesrvs).resolver.Resolve()...) if len(now) != len(old) { rlog.Error("different namesrv option in the same instance", map[string]interface{}{ "NewNameSrv": now, "BeforeNameSrv": old, }) + client.Shutdown() return nil } sort.Strings(now) @@ -220,6 +234,7 @@ func GetOrNewRocketMQClient(option ClientOptions, callbackCh chan interface{}) R "NewNameSrv": now, "BeforeNameSrv": old, }) + client.Shutdown() return nil } } @@ -404,16 +419,14 @@ func GetOrNewRocketMQClient(option ClientOptions, callbackCh chan interface{}) R } return res }) + // Only the creator binds the NameServer; reuse must not race with readers. + client.GetNameSrv().(*namesrvs).bundleClient = client + close(client.ready) } - // bundle this client to namesrv - client.GetNameSrv().(*namesrvs).bundleClient = client return client } func (c *rmqClient) Start() { - //ctx, cancel := context.WithCancel(context.Background()) - //c.cancel = cancel - atomic.AddInt32(&c.instanceCount, 1) c.once.Do(func() { if !c.option.Credentials.IsEmpty() { c.remoteClient.RegisterInterceptor(remote.ACLInterceptor(c.option.Credentials)) @@ -538,23 +551,29 @@ func (c *rmqClient) Start() { }) } -func (c *rmqClient) removeClient() { - rlog.Info("will remove client from clientMap", map[string]interface{}{ - "clientID": c.ClientID(), - }) - clientMap.Delete(c.ClientID()) -} - func (c *rmqClient) Shutdown() { - if atomic.AddInt32(&c.instanceCount, -1) > 0 { + clientMapMu.Lock() + if c.refs == 0 { + clientMapMu.Unlock() return } + c.refs-- + if c.refs != 0 { + clientMapMu.Unlock() + return + } + // Removal and acquisition must be atomic. A new generation may be created + // after this point; disposing this client's transports cannot remove it. + clientMap.Delete(c.ClientID()) + clientMapMu.Unlock() c.shutdownOnce.Do(func() { + rlog.Info("will remove client from clientMap", map[string]interface{}{ + "clientID": c.ClientID(), + }) close(c.done) c.close = true c.remoteClient.ShutDown() - c.removeClient() }) } diff --git a/internal/client_lifecycle_test.go b/internal/client_lifecycle_test.go new file mode 100644 index 00000000..a2d7502c --- /dev/null +++ b/internal/client_lifecycle_test.go @@ -0,0 +1,96 @@ +/* +Licensed to the Apache Software Foundation (ASF) under one or more +contributor license agreements. See the NOTICE file distributed with +this work for additional information regarding copyright ownership. +The ASF licenses this file to You under the Apache License, Version 2.0 +(the "License"); you may not use this file except in compliance with +the License. You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package internal + +import ( + "sync" + "testing" + + "github.com/apache/rocketmq-client-go/v2/primitive" + "github.com/stretchr/testify/require" +) + +func ownershipTestOptions(t *testing.T) ClientOptions { + options := DefaultClientOptions() + options.InstanceName = t.Name() + ns, err := NewNamesrv(primitive.NewPassthroughResolver([]string{"127.0.0.2:9876", "127.0.0.1:9876"}), nil) + require.NoError(t, err) + options.Namesrv = ns + return options +} + +func TestClientOwnershipConcurrentAcquireAndRelease(t *testing.T) { + const workers = 16 + var wg sync.WaitGroup + for i := 0; i < workers; i++ { + wg.Add(1) + go func() { + defer wg.Done() + for j := 0; j < 20; j++ { + client := GetOrNewRocketMQClient(ownershipTestOptions(t), nil) + if client == nil { + t.Error("acquisition failed") + return + } + actual := client.(*rmqClient) + if actual.GetNameSrv().(*namesrvs).bundleClient != actual { + t.Error("client published before initialization") + } + select { + case <-actual.done: + t.Error("acquired a closed client") + default: + } + client.Shutdown() + } + }() + } + wg.Wait() + options := ownershipTestOptions(t) + _, exists := clientMap.Load((&rmqClient{option: options}).ClientID()) + require.False(t, exists, "all acquisitions must be released") +} + +func TestClientOwnershipOldCleanupKeepsReplacement(t *testing.T) { + options := ownershipTestOptions(t) + old := GetOrNewRocketMQClient(options, nil).(*rmqClient) + old.Shutdown() + replacement := GetOrNewRocketMQClient(options, nil).(*rmqClient) + defer replacement.Shutdown() + require.NotSame(t, old, replacement) + old.Shutdown() + stored, exists := clientMap.Load(replacement.ClientID()) + require.True(t, exists) + require.Same(t, replacement, stored) +} + +func TestClientOwnershipKeepsResolverSnapshot(t *testing.T) { + options := ownershipTestOptions(t) + addresses := []string{"127.0.0.2:9876", "127.0.0.1:9876"} + ns, err := NewNamesrv(primitive.NewPassthroughResolver(addresses), nil) + require.NoError(t, err) + options.Namesrv = ns + first := GetOrNewRocketMQClient(options, nil) + require.NotNil(t, first) + defer first.Shutdown() + second := GetOrNewRocketMQClient(options, nil) + require.NotNil(t, second) + defer second.Shutdown() + require.Equal(t, []string{"127.0.0.2:9876", "127.0.0.1:9876"}, addresses, + "compatibility checks must not sort a resolver-owned snapshot in place") +} diff --git a/internal/route.go b/internal/route.go index d8ca7fdb..3e3da344 100644 --- a/internal/route.go +++ b/internal/route.go @@ -126,7 +126,12 @@ func (s *namesrvs) UpdateTopicRouteInfoWithDefault(topic string, defaultTopic st func (s *namesrvs) updateTopicRouteInfoWithContext(ctx context.Context, topic string, defaultTopic string, defaultQueueNum int) (*TopicRouteData, bool, error) { s.lockNamesrv.Lock() defer s.lockNamesrv.Unlock() + return s.updateTopicRouteInfoLocked(ctx, topic, defaultTopic, defaultQueueNum) +} +// The caller holds lockNamesrv. Explicit refreshes always query the server; +// cache-miss lookups recheck the cache under the same lock before calling here. +func (s *namesrvs) updateTopicRouteInfoLocked(ctx context.Context, topic string, defaultTopic string, defaultQueueNum int) (*TopicRouteData, bool, error) { var ( routeData *TopicRouteData err error @@ -197,15 +202,17 @@ func (s *namesrvs) updateTopicRouteInfoWithContext(ctx context.Context, topic st rlog.Info("change the route for clients", nil) } + // Cache readers may immediately select a queue from the new route. + // Publish its broker addresses before making those queues visible. + for _, brokerData := range routeData.BrokerDataList { + s.brokerAddressesMap.Store(brokerData.BrokerName, brokerData) + } s.routeDataMap.Store(topic, routeData) rlog.Info("the topic route info changed", map[string]interface{}{ rlog.LogKeyTopic: topic, rlog.LogKeyValueChangedFrom: oldRouteData, rlog.LogKeyValueChangedTo: routeData.String(), }) - for _, brokerData := range routeData.BrokerDataList { - s.brokerAddressesMap.Store(brokerData.BrokerName, brokerData) - } } return routeData.clone(), changed, nil @@ -369,7 +376,7 @@ func (s *namesrvs) fetchPublishMessageQueuesWithContext(ctx context.Context, top v, exist := s.routeDataMap.Load(topic) if !exist { - routeData, _, err = s.updateTopicRouteInfoWithContext(ctx, topic, "", 0) + routeData, err = s.loadTopicRouteWithContext(ctx, topic) if err != nil { rlog.Error("queryTopicRouteInfoFromServer failed", map[string]interface{}{ rlog.LogKeyTopic: topic, @@ -386,6 +393,19 @@ func (s *namesrvs) fetchPublishMessageQueuesWithContext(ctx context.Context, top return publishInfo.MqList, nil } +func (s *namesrvs) loadTopicRouteWithContext(ctx context.Context, topic string) (*TopicRouteData, error) { + s.lockNamesrv.Lock() + defer s.lockNamesrv.Unlock() + if err := ctx.Err(); err != nil { + return nil, err + } + if cached, ok := s.routeDataMap.Load(topic); ok { + return cached.(*TopicRouteData).clone(), nil + } + route, _, err := s.updateTopicRouteInfoLocked(ctx, topic, "", 0) + return route, err +} + func (s *namesrvs) AddBrokerVersion(brokerName, brokerAddr string, version int32) { s.brokerLock.Lock() defer s.brokerLock.Unlock() diff --git a/internal/trace.go b/internal/trace.go index 56be0291..533e7108 100644 --- a/internal/trace.go +++ b/internal/trace.go @@ -287,10 +287,11 @@ func (td *traceDispatcher) Start() { return } td.started = true - td.ticker = time.NewTicker(5 * time.Millisecond) + maxWaitDuration := 5 * time.Millisecond + td.ticker = time.NewTicker(maxWaitDuration) go primitive.WithRecover(func() { defer close(td.processDone) - td.process() + td.process(maxWaitDuration) }) } @@ -359,18 +360,23 @@ func (td *traceDispatcher) submitBatch(batch []TraceContext) { }) } -func (td *traceDispatcher) process() { +func (td *traceDispatcher) process(maxWaitDuration time.Duration) { var batch []TraceContext + lastPut := time.Now() flush := func() { td.submitBatch(batch); batch = nil } for { select { case ctx := <-td.input: + lastPut = time.Now() batch = append(batch, ctx) if len(batch) == batchSize { flush() } case <-td.ticker.C: - flush() + if time.Since(lastPut) > maxWaitDuration { + lastPut = time.Now() + flush() + } case <-td.ctx.Done(): // Append and Close share a lock, so no more records can enter. for { @@ -506,9 +512,9 @@ func (td *traceDispatcher) findMq(regionID string) (*primitive.MessageQueue, str return nil, "" } td.resource.topics.Store(traceTopic, struct{}{}) - ctx, cancel := context.WithTimeout(td.sendCtx, 5*time.Second) - defer cancel() - mqs, err := td.namesrvs.fetchPublishMessageQueuesWithContext(ctx, traceTopic) + // Each NameServer attempt has its own timeout; shutdown can still cancel + // the whole lookup through sendCtx without cutting off healthy fallbacks. + mqs, err := td.namesrvs.fetchPublishMessageQueuesWithContext(td.sendCtx, traceTopic) if err != nil { rlog.Error("fetch publish message queues failed", map[string]interface{}{ rlog.LogKeyUnderlayError: err, diff --git a/internal/trace_client.go b/internal/trace_client.go index a9e8dd08..89ec01a1 100644 --- a/internal/trace_client.go +++ b/internal/trace_client.go @@ -200,18 +200,38 @@ func createTraceClient(key traceClientKey, cfg *primitive.TraceConfig, shared *p } func (client *traceClient) refresh(ctx context.Context) { - defer close(client.done) - names := time.NewTimer(10 * time.Second) - routes := time.NewTicker(_PullNameServerInterval) - defer names.Stop() + client.runRefresh(ctx, 10*time.Second, 2*time.Minute, _PullNameServerInterval) +} + +func (client *traceClient) runRefresh(ctx context.Context, initialNamesDelay, namesInterval, routesInterval time.Duration) { + ctx, cancel := context.WithCancel(ctx) + namesDone := make(chan struct{}) + go primitive.WithRecover(func() { + defer close(namesDone) + names := time.NewTimer(initialNamesDelay) + defer names.Stop() + for { + select { + case <-ctx.Done(): + return + case <-names.C: + client.namesrvs.UpdateNameServerAddress() + names.Reset(namesInterval) + } + } + }) + // Wait for both loops before the owner disposes the resolver or transports. + defer func() { + cancel() + <-namesDone + close(client.done) + }() + routes := time.NewTicker(routesInterval) defer routes.Stop() for { select { case <-ctx.Done(): return - case <-names.C: - client.namesrvs.UpdateNameServerAddress() - names.Reset(2 * time.Minute) case <-routes.C: client.refreshRoutes(ctx) } @@ -223,9 +243,9 @@ func (client *traceClient) refreshRoutes(ctx context.Context) { if ctx.Err() != nil { return false } - requestCtx, cancel := context.WithTimeout(ctx, 5*time.Second) - client.namesrvs.updateTopicRouteInfoWithContext(requestCtx, key.(string), "", 0) - cancel() + // Preserve the per-NameServer timeout and fallback attempts. The worker + // context cancels outstanding discovery when the resource is released. + client.namesrvs.updateTopicRouteInfoWithContext(ctx, key.(string), "", 0) return true }) } diff --git a/internal/trace_compatibility_test.go b/internal/trace_compatibility_test.go new file mode 100644 index 00000000..0b50cdca --- /dev/null +++ b/internal/trace_compatibility_test.go @@ -0,0 +1,177 @@ +/* +Licensed to the Apache Software Foundation (ASF) under one or more +contributor license agreements. See the NOTICE file distributed with +this work for additional information regarding copyright ownership. +The ASF licenses this file to You under the Apache License, Version 2.0 +(the "License"); you may not use this file except in compliance with +the License. You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package internal + +import ( + "context" + "fmt" + "strings" + "testing" + "time" + + "github.com/apache/rocketmq-client-go/v2/internal/remote" + "github.com/apache/rocketmq-client-go/v2/primitive" + "github.com/golang/mock/gomock" + "github.com/stretchr/testify/require" +) + +func TestTraceRouteFailoverAfterTimeout(t *testing.T) { + for _, refresh := range []bool{false, true} { + refresh := refresh + t.Run(fmt.Sprintf("refresh=%t", refresh), func(t *testing.T) { + t.Parallel() + ctrl := gomock.NewController(t) + defer ctrl.Finish() + td := NewTraceDispatcher(&primitive.TraceConfig{UnitName: t.Name(), NamesrvAddrs: []string{"127.0.0.1:9876", "127.0.0.2:9876"}}) + require.NotNil(t, td) + defer td.Close() + ns := remote.NewMockRemotingClient(ctrl) + td.namesrvs.nameSrvClient = ns + gomock.InOrder( + ns.EXPECT().InvokeSync(gomock.Any(), "127.0.0.1:9876", gomock.Any()).DoAndReturn( + func(ctx context.Context, _ string, _ *remote.RemotingCommand) (*remote.RemotingCommand, error) { + <-ctx.Done() + return nil, ctx.Err() + }), + ns.EXPECT().InvokeSync(gomock.Any(), "127.0.0.2:9876", gomock.Any()).Return(testTraceRoute("127.0.0.2:10911"), nil), + ) + ns.EXPECT().ShutDown() + if refresh { + td.resource.topics.Store(td.GetTraceTopicName(), struct{}{}) + td.resource.refreshRoutes(context.Background()) + _, exists := td.namesrvs.routeDataMap.Load(td.GetTraceTopicName()) + require.True(t, exists, "refresh must reach the healthy fallback after a timeout") + } + mq, addr := td.findMq("") + require.NotNil(t, mq) + require.Equal(t, "127.0.0.2:10911", addr) + }) + } +} + +func TestTraceRouteLookupCancellation(t *testing.T) { + for _, refresh := range []bool{false, true} { + t.Run(fmt.Sprintf("refresh=%t", refresh), func(t *testing.T) { + ctrl := gomock.NewController(t) + defer ctrl.Finish() + td := NewTraceDispatcher(&primitive.TraceConfig{UnitName: t.Name(), NamesrvAddrs: []string{"127.0.0.1:9876", "127.0.0.2:9876"}}) + require.NotNil(t, td) + defer td.Close() + ns := remote.NewMockRemotingClient(ctrl) + td.namesrvs.nameSrvClient = ns + entered, returned := make(chan struct{}), make(chan struct{}) + ns.EXPECT().InvokeSync(gomock.Any(), "127.0.0.1:9876", gomock.Any()).DoAndReturn( + func(ctx context.Context, _ string, _ *remote.RemotingCommand) (*remote.RemotingCommand, error) { + close(entered) + <-ctx.Done() + return nil, ctx.Err() + }) + ns.EXPECT().ShutDown() + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + go func() { + defer close(returned) + if refresh { + td.resource.topics.Store(td.GetTraceTopicName(), struct{}{}) + td.resource.refreshRoutes(ctx) + } else { + td.findMq("") + } + }() + select { + case <-entered: + case <-time.After(time.Second): + t.Fatal("route lookup did not start") + } + cancel() + td.cancelSend() + select { + case <-returned: + case <-time.After(time.Second): + t.Fatal("route lookup ignored cancellation") + } + _, exists := td.namesrvs.routeDataMap.Load(td.GetTraceTopicName()) + require.False(t, exists) + }) + } +} + +func TestTraceBatchFlushPolicy(t *testing.T) { + for _, tc := range []struct { + name string + records int + idleTimeout time.Duration + closeToFlush bool + }{ + {"full-batch", batchSize, time.Hour, false}, + {"idle", 3, 0, false}, + {"shutdown", 3, time.Hour, true}, + } { + t.Run(tc.name, func(t *testing.T) { + ctrl := gomock.NewController(t) + defer ctrl.Finish() + td := NewTraceDispatcher(&primitive.TraceConfig{UnitName: t.Name(), NamesrvAddrs: []string{"127.0.0.1:9876"}}) + require.NotNil(t, td) + defer td.Close() + ns := remote.NewMockRemotingClient(ctrl) + td.namesrvs.nameSrvClient = ns + broker := remote.NewMockRemotingClient(ctrl) + td.resource.cli.remoteClient = broker + ns.EXPECT().InvokeSync(gomock.Any(), gomock.Any(), gomock.Any()).Return(testTraceRoute("127.0.0.1:10911"), nil) + ns.EXPECT().ShutDown() + broker.EXPECT().ShutDown() + sent := make(chan string, 1) + broker.EXPECT().InvokeAsync(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).DoAndReturn( + func(ctx context.Context, addr string, req *remote.RemotingCommand, callback func(*remote.ResponseFuture)) error { + sent <- string(req.Body) + callback(&remote.ResponseFuture{ResponseCommand: &remote.RemotingCommand{Code: ResSuccess}}) + return nil + }).Times(1) + + // Drive ticker events explicitly. A long idle threshold lets us test active + // traffic without depending on the scheduler meeting a 5ms deadline. + ticks := make(chan time.Time) + td.input = make(chan TraceContext) + td.ticker = time.NewTicker(time.Hour) + td.ticker.C = ticks + td.started = true + go func() { + defer close(td.processDone) + td.process(tc.idleTimeout) + }() + for i := 0; i < tc.records; i++ { + td.input <- TraceContext{TraceType: SubBefore, TraceBeans: []TraceBean{{Topic: "topic", MsgId: "batch-record"}}} + if tc.name == "full-batch" && i < tc.records-1 { + ticks <- time.Now() + } + } + if tc.closeToFlush { + td.Close() + } else if tc.name == "idle" { + ticks <- time.Now() + } + select { + case body := <-sent: + require.Equal(t, tc.records, strings.Count(body, "batch-record")) + case <-time.After(time.Second): + t.Fatal("trace batch was not sent") + } + td.Close() + }) + } +} diff --git a/internal/trace_migration_test.go b/internal/trace_migration_test.go new file mode 100644 index 00000000..3d75b4bf --- /dev/null +++ b/internal/trace_migration_test.go @@ -0,0 +1,150 @@ +/* +Licensed to the Apache Software Foundation (ASF) under one or more +contributor license agreements. See the NOTICE file distributed with +this work for additional information regarding copyright ownership. +The ASF licenses this file to You under the Apache License, Version 2.0 +(the "License"); you may not use this file except in compliance with +the License. You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package internal + +import ( + "context" + "net" + "sync/atomic" + "testing" + "time" + + "github.com/apache/rocketmq-client-go/v2/internal/remote" + "github.com/apache/rocketmq-client-go/v2/primitive" + "github.com/golang/mock/gomock" + "github.com/stretchr/testify/require" +) + +func TestSharedTraceSurvivesNameServerReplacement(t *testing.T) { + ctrl := gomock.NewController(t) + defer ctrl.Finish() + var created, closed, secondFactoryCalls int32 + resolver := &mutableTraceResolver{addrs: []string{"127.0.0.1:9876"}} + shared := traceTestShared(t.Name(), resolver, &created, &closed) + first := NewSharedTraceDispatcher(&primitive.TraceConfig{GroupName: "old-consumer"}, shared) + require.NotNil(t, first) + defer first.Close() + ns := remote.NewMockRemotingClient(ctrl) + first.namesrvs.nameSrvClient = ns + gomock.InOrder( + ns.EXPECT().InvokeSync(gomock.Any(), "127.0.0.1:9876", gomock.Any()).Return(testTraceRoute("127.0.0.1:10911"), nil), + ns.EXPECT().InvokeSync(gomock.Any(), "127.0.0.2:9876", gomock.Any()).Return(testTraceRoute("127.0.0.2:10911"), nil), + ) + ns.EXPECT().ShutDown() + first.Start() + mq, addr := first.findMq("") + require.NotNil(t, mq) + require.Equal(t, "127.0.0.1:10911", addr) + resolver.set("127.0.0.2:9876") + // The new consumer has a different address snapshot but the same logical identity. + second := NewSharedTraceDispatcher(&primitive.TraceConfig{GroupName: "new-consumer"}, primitive.SharedTraceClientConfig{ + Key: shared.Key, + ResolverFactory: func() (primitive.NsResolver, func(), error) { + atomic.AddInt32(&secondFactoryCalls, 1) + return primitive.NewPassthroughResolver([]string{"127.0.0.2:9876"}), nil, nil + }, + }) + require.NotNil(t, second) + defer second.Close() + require.Same(t, first.resource, second.resource) + require.Equal(t, int32(0), atomic.LoadInt32(&secondFactoryCalls)) + first.Close() + require.Equal(t, int32(0), atomic.LoadInt32(&closed)) + second.namesrvs.UpdateNameServerAddress() + second.resource.refreshRoutes(context.Background()) + broker := remote.NewMockRemotingClient(ctrl) + second.resource.cli.remoteClient = broker + broker.EXPECT().InvokeAsync(gomock.Any(), "127.0.0.2:10911", gomock.Any(), gomock.Any()).DoAndReturn( + func(ctx context.Context, addr string, req *remote.RemotingCommand, cb func(*remote.ResponseFuture)) error { + require.Contains(t, string(req.Body), "after-replacement") + cb(&remote.ResponseFuture{ResponseCommand: &remote.RemotingCommand{Code: ResSuccess}}) + return nil + }) + broker.EXPECT().ShutDown() + second.Start() + require.True(t, second.Append(TraceContext{TraceType: SubBefore, TraceBeans: []TraceBean{{Topic: "topic", MsgId: "after-replacement"}}})) + second.Close() + require.Equal(t, int32(1), atomic.LoadInt32(&closed)) + traceClients.Lock() + _, exists := traceClients.entries[first.resource.key] + traceClients.Unlock() + require.False(t, exists) +} +func TestLegacyAndSharedTraceStayIsolated(t *testing.T) { + legacy := NewTraceDispatcher(&primitive.TraceConfig{UnitName: t.Name(), NamesrvAddrs: []string{"127.0.0.1:9876"}}) + require.NotNil(t, legacy) + defer legacy.Close() + shared := NewSharedTraceDispatcher(&primitive.TraceConfig{UnitName: t.Name()}, primitive.SharedTraceClientConfig{ + Key: t.Name(), ResolverFactory: func() (primitive.NsResolver, func(), error) { + return primitive.NewPassthroughResolver([]string{"127.0.0.1:9876"}), nil, nil + }, + }) + require.NotNil(t, shared) + defer shared.Close() + require.NotSame(t, legacy.resource, shared.resource) + require.NotSame(t, legacy.namesrvs, shared.namesrvs) + legacy.Close() + shared.Start() + require.True(t, shared.Append(TraceContext{})) +} + +func TestTraceRealSendTimeoutReleases(t *testing.T) { + listener, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + defer listener.Close() + serverDone := make(chan struct{}) + go func() { + defer close(serverDone) + conn, err := listener.Accept() + if err != nil { + return + } + defer conn.Close() + buffer := make([]byte, 1024) + for { + if _, err := conn.Read(buffer); err != nil { + return + } + } + }() + var created, closed int32 + td := NewSharedTraceDispatcher(&primitive.TraceConfig{}, traceTestShared(t.Name(), + primitive.NewPassthroughResolver([]string{"127.0.0.1:9876"}), &created, &closed)) + require.NotNil(t, td) + defer td.Close() + route := &TopicRouteData{ + QueueDataList: []*QueueData{{BrokerName: "broker", WriteQueueNums: 1, ReadQueueNums: 1, Perm: 6}}, + BrokerDataList: []*BrokerData{{BrokerName: "broker", BrokerAddresses: map[int64]string{MasterId: listener.Addr().String()}}}, + } + td.namesrvs.AddBroker(route) + td.namesrvs.routeDataMap.Store(td.traceTopic, route) + td.Start() + require.True(t, td.Append(TraceContext{TraceType: SubBefore, TraceBeans: []TraceBean{{Topic: "topic", MsgId: "message"}}})) + td.Close() + select { + case <-td.closeDone: + case <-time.After(2 * time.Second): + t.Fatal("send timeout did not release the resource") + } + require.Equal(t, int32(1), atomic.LoadInt32(&closed)) + select { + case <-serverDone: + case <-time.After(time.Second): + t.Fatal("broker connection was not closed") + } +} diff --git a/internal/trace_route_test.go b/internal/trace_route_test.go index 4e746e14..4670ac4c 100644 --- a/internal/trace_route_test.go +++ b/internal/trace_route_test.go @@ -22,6 +22,7 @@ import ( "fmt" "sync" "testing" + "time" "github.com/apache/rocketmq-client-go/v2/internal/remote" "github.com/apache/rocketmq-client-go/v2/primitive" @@ -47,6 +48,164 @@ func multiBrokerTraceRoute(generation int) *TopicRouteData { return route } +func TestSharedTraceConcurrentColdRouteLookup(t *testing.T) { + ctrl := gomock.NewController(t) + defer ctrl.Finish() + var created, closed int32 + shared := traceTestShared(t.Name(), primitive.NewPassthroughResolver([]string{"127.0.0.1:9876"}), &created, &closed) + const callers = 24 + dispatchers := make([]*traceDispatcher, callers) + for i := range dispatchers { + dispatchers[i] = NewSharedTraceDispatcher(&primitive.TraceConfig{}, shared) + require.NotNil(t, dispatchers[i]) + defer dispatchers[i].Close() + } + ns := remote.NewMockRemotingClient(ctrl) + dispatchers[0].namesrvs.nameSrvClient = ns + entered, release := make(chan struct{}), make(chan struct{}) + ns.EXPECT().InvokeSync(gomock.Any(), gomock.Any(), gomock.Any()).DoAndReturn( + func(context.Context, string, *remote.RemotingCommand) (*remote.RemotingCommand, error) { + close(entered) + <-release + return testTraceRoute("127.0.0.1:10911"), nil + }).Times(1) + ns.EXPECT().ShutDown() + var wg sync.WaitGroup + lookup := func(td *traceDispatcher) { + defer wg.Done() + mq, addr := td.findMq("") + if mq == nil || addr != "127.0.0.1:10911" { + t.Errorf("incomplete shared route: queue=%v, address=%q", mq, addr) + } + } + wg.Add(1) + go lookup(dispatchers[0]) + <-entered + started := make(chan struct{}, callers-1) + for _, td := range dispatchers[1:] { + wg.Add(1) + go func(td *traceDispatcher) { started <- struct{}{}; lookup(td) }(td) + } + for i := 1; i < callers; i++ { + <-started + } + // Keep the first RPC pending while the other callers observe the cold cache. + time.Sleep(20 * time.Millisecond) + close(release) + wg.Wait() +} + +func TestTraceFailedRouteLookupCanRetry(t *testing.T) { + ctrl := gomock.NewController(t) + defer ctrl.Finish() + ns, err := NewNamesrv(primitive.NewPassthroughResolver([]string{"127.0.0.1:9876"}), nil) + require.NoError(t, err) + remoteClient := remote.NewMockRemotingClient(ctrl) + ns.nameSrvClient = remoteClient + gomock.InOrder( + remoteClient.EXPECT().InvokeSync(gomock.Any(), gomock.Any(), gomock.Any()).Return(nil, context.DeadlineExceeded), + remoteClient.EXPECT().InvokeSync(gomock.Any(), gomock.Any(), gomock.Any()).Return(testTraceRoute("127.0.0.1:10911"), nil), + remoteClient.EXPECT().InvokeSync(gomock.Any(), gomock.Any(), gomock.Any()).Return(testTraceRoute("127.0.0.2:10911"), nil), + ) + _, err = ns.fetchPublishMessageQueuesWithContext(context.Background(), "trace") + require.Error(t, err) + queues, err := ns.fetchPublishMessageQueuesWithContext(context.Background(), "trace") + require.NoError(t, err) + require.Len(t, queues, 1) + _, _, err = ns.updateTopicRouteInfoWithContext(context.Background(), "trace", "", 0) + require.NoError(t, err) + require.Equal(t, "127.0.0.2:10911", ns.FindBrokerAddrByName("broker"), "explicit refresh must bypass the cached route") +} + +func TestTraceNameServerRefreshWhileRouteQueryBlocked(t *testing.T) { + ctrl := gomock.NewController(t) + defer ctrl.Finish() + resolver := &mutableTraceResolver{addrs: []string{"127.0.0.1:9876"}} + ns, err := NewNamesrv(resolver, nil) + require.NoError(t, err) + remoteClient := remote.NewMockRemotingClient(ctrl) + ns.nameSrvClient = remoteClient + client := &traceClient{namesrvs: ns, done: make(chan struct{})} + client.topics.Store("trace", struct{}{}) + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + entered := make(chan struct{}) + remoteClient.EXPECT().InvokeSync(gomock.Any(), "127.0.0.1:9876", gomock.Any()).DoAndReturn( + func(ctx context.Context, _ string, _ *remote.RemotingCommand) (*remote.RemotingCommand, error) { + close(entered) + <-ctx.Done() + return nil, ctx.Err() + }) + go client.runRefresh(ctx, time.Millisecond, time.Millisecond, time.Millisecond) + select { + case <-entered: + case <-time.After(time.Second): + t.Fatal("route query did not start") + } + resolver.set("127.0.0.2:9876") + require.Eventually(t, func() bool { return ns.AddrList()[0] == "127.0.0.2:9876" }, time.Second, time.Millisecond, + "a blocked route RPC must not delay address discovery") + cancel() + select { + case <-client.done: + case <-time.After(time.Second): + t.Fatal("refresh workers did not exit after cancellation") + } +} + +func TestTraceRoutePublishesBrokerBeforeQueues(t *testing.T) { + ctrl := gomock.NewController(t) + defer ctrl.Finish() + var created, closed int32 + td := NewSharedTraceDispatcher(&primitive.TraceConfig{}, traceTestShared(t.Name(), + primitive.NewPassthroughResolver([]string{"127.0.0.1:9876"}), &created, &closed)) + require.NotNil(t, td) + defer td.Close() + ns := remote.NewMockRemotingClient(ctrl) + td.namesrvs.nameSrvClient = ns + generation := 0 + ns.EXPECT().InvokeSync(gomock.Any(), gomock.Any(), gomock.Any()).DoAndReturn( + func(context.Context, string, *remote.RemotingCommand) (*remote.RemotingCommand, error) { + generation++ + name := fmt.Sprintf("broker-%d", generation) + route := &TopicRouteData{ + QueueDataList: []*QueueData{{BrokerName: name, ReadQueueNums: 1, WriteQueueNums: 1, Perm: 6}}, + BrokerDataList: []*BrokerData{{BrokerName: name, BrokerAddresses: map[int64]string{MasterId: "127.0.0.1:10911"}}}, + } + return &remote.RemotingCommand{Code: ResSuccess, Body: []byte(route.String())}, nil + }).AnyTimes() + ns.EXPECT().ShutDown() + mq, addr := td.findMq("") + require.NotNil(t, mq) + require.NotEmpty(t, addr) + var wg sync.WaitGroup + start := make(chan struct{}) + wg.Add(1) + go func() { + defer wg.Done() + <-start + for i := 0; i < 200; i++ { + td.resource.refreshRoutes(context.Background()) + } + }() + for i := 0; i < 4; i++ { + wg.Add(1) + go func() { + defer wg.Done() + <-start + for j := 0; j < 1000; j++ { + mq, addr := td.findMq("") + if mq == nil || addr == "" { + t.Errorf("visible queue has no broker address: %v", mq) + return + } + } + }() + } + close(start) + wg.Wait() +} + func TestTraceCachedRoutePreservesSnapshot(t *testing.T) { ns, err := NewNamesrv(primitive.NewPassthroughResolver([]string{"127.0.0.1:9876"}), nil) require.NoError(t, err) diff --git a/producer/interceptor.go b/producer/interceptor.go index 834b9bce..1923af6a 100644 --- a/producer/interceptor.go +++ b/producer/interceptor.go @@ -52,6 +52,9 @@ func installTraceInterceptor(options *producerOptions, dispatcher internal.Trace } if !internal.IsNilTraceDispatcher(options.TraceDispatcher) { options.TraceDispatcher.Close() + // Trace is always prepended; WithInterceptor appends user interceptors. + // Remove the old wrapper before installing its replacement. + options.Interceptors = options.Interceptors[1:] } options.TraceDispatcher = dispatcher options.Interceptors = append([]primitive.Interceptor{newTraceInterceptor(dispatcher)}, options.Interceptors...) diff --git a/producer/producer.go b/producer/producer.go index 57dc6af0..59dc1829 100644 --- a/producer/producer.go +++ b/producer/producer.go @@ -47,6 +47,7 @@ type defaultProducer struct { interceptor primitive.Interceptor + lifecycleMu sync.Mutex startOnce sync.Once ShutdownOnce sync.Once } @@ -89,6 +90,11 @@ func NewDefaultProducer(opts ...Option) (*defaultProducer, error) { } func (p *defaultProducer) Start() error { + p.lifecycleMu.Lock() + defer p.lifecycleMu.Unlock() + if atomic.LoadInt32(&p.state) == int32(internal.StateShutdown) { + return internal.ErrServiceState + } var err error defer func() { if err != nil && !internal.IsNilTraceDispatcher(p.options.TraceDispatcher) { @@ -111,12 +117,16 @@ func (p *defaultProducer) Start() error { } func (p *defaultProducer) Shutdown() error { + p.lifecycleMu.Lock() + defer p.lifecycleMu.Unlock() p.ShutdownOnce.Do(func() { - if p.options.TraceDispatcher != nil { + if !internal.IsNilTraceDispatcher(p.options.TraceDispatcher) { p.options.TraceDispatcher.Close() } - atomic.StoreInt32(&p.state, int32(internal.StateShutdown)) - p.client.UnregisterProducer(p.group) + previous := atomic.SwapInt32(&p.state, int32(internal.StateShutdown)) + if previous == int32(internal.StateRunning) { + p.client.UnregisterProducer(p.group) + } p.client.Shutdown() }) return nil diff --git a/producer/trace_lifecycle_test.go b/producer/trace_lifecycle_test.go index 76e8522e..5abdd115 100644 --- a/producer/trace_lifecycle_test.go +++ b/producer/trace_lifecycle_test.go @@ -20,6 +20,7 @@ package producer import ( "context" "errors" + "sync" "sync/atomic" "testing" @@ -106,3 +107,57 @@ func TestTraceReleasedOnStartFailure(t *testing.T) { require.Error(t, p.Start()) require.Equal(t, int32(1), atomic.LoadInt32(&closed)) } + +func TestTraceShutdownBeforeStartOrAfterDuplicateGroup(t *testing.T) { + for _, failed := range []bool{false, true} { + ctrl := gomock.NewController(t) + client := internal.NewMockRMQClient(ctrl) + client.EXPECT().Shutdown() + opts := defaultProducerOptions() + var closed int32 + WithSharedTrace(&primitive.TraceConfig{}, primitive.SharedTraceClientConfig{Key: t.Name(), ResolverFactory: func() (primitive.NsResolver, func(), error) { + return primitive.NewPassthroughResolver([]string{"127.0.0.1:9876"}), func() { atomic.AddInt32(&closed, 1) }, nil + }})(&opts) + p := &defaultProducer{client: client, options: opts} + if failed { + client.EXPECT().RegisterProducer(gomock.Any(), gomock.Any()).Return(errors.New("duplicate group")) + require.Error(t, p.Start()) + } + require.NoError(t, p.Shutdown()) + require.NoError(t, p.Shutdown()) + require.Error(t, p.Start()) + require.Equal(t, int32(1), atomic.LoadInt32(&closed)) + ctrl.Finish() + } +} + +func TestTraceConcurrentStartShutdown(t *testing.T) { + ctrl := gomock.NewController(t) + defer ctrl.Finish() + client := internal.NewMockRMQClient(ctrl) + registered, release := make(chan struct{}), make(chan struct{}) + client.EXPECT().RegisterProducer(gomock.Any(), gomock.Any()).DoAndReturn(func(string, internal.InnerProducer) error { close(registered); <-release; return nil }) + client.EXPECT().Start() + client.EXPECT().UnregisterProducer(gomock.Any()) + client.EXPECT().Shutdown() + p := &defaultProducer{client: client, options: defaultProducerOptions()} + var wg sync.WaitGroup + wg.Add(2) + go func() { + defer wg.Done() + if err := p.Start(); err != nil { + t.Error(err) + } + }() + <-registered + go func() { + defer wg.Done() + if err := p.Shutdown(); err != nil { + t.Error(err) + } + }() + close(release) + wg.Wait() + require.NoError(t, p.Shutdown()) + require.Error(t, p.Start()) +} diff --git a/producer/trace_replacement_test.go b/producer/trace_replacement_test.go new file mode 100644 index 00000000..b26f30ca --- /dev/null +++ b/producer/trace_replacement_test.go @@ -0,0 +1,87 @@ +/* +Licensed to the Apache Software Foundation (ASF) under one or more +contributor license agreements. See the NOTICE file distributed with +this work for additional information regarding copyright ownership. +The ASF licenses this file to You under the Apache License, Version 2.0 +(the "License"); you may not use this file except in compliance with +the License. You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package producer + +import ( + "context" + "errors" + "testing" + + "github.com/apache/rocketmq-client-go/v2/internal" + "github.com/apache/rocketmq-client-go/v2/primitive" + "github.com/stretchr/testify/require" +) + +type replacementTraceDispatcher struct { + starts, closes int + records []internal.TraceContext +} + +func (d *replacementTraceDispatcher) Start() { d.starts++ } +func (d *replacementTraceDispatcher) Close() { d.closes++ } +func (*replacementTraceDispatcher) GetTraceTopicName() string { return "trace-topic" } +func (d *replacementTraceDispatcher) Append(ctx internal.TraceContext) bool { + d.records = append(d.records, ctx) + return d.closes == 0 +} + +func TestTraceReplacementPreservesInterceptors(t *testing.T) { + options := defaultProducerOptions() + var calls []string + userInterceptor := func(name string) primitive.Interceptor { + return func(ctx context.Context, req, reply interface{}, next primitive.Invoker) error { + calls = append(calls, name) + return next(ctx, req, reply) + } + } + WithInterceptor(userInterceptor("first"))(&options) + old := &replacementTraceDispatcher{} + installTraceInterceptor(&options, old) + WithInterceptor(userInterceptor("second"))(&options) + middle, current := &replacementTraceDispatcher{}, &replacementTraceDispatcher{} + installTraceInterceptor(&options, middle) + installTraceInterceptor(&options, current) + // An unavailable replacement must preserve the working dispatcher and chain. + var typedNil *replacementTraceDispatcher + for _, unavailable := range []internal.TraceDispatcher{nil, typedNil} { + installTraceInterceptor(&options, unavailable) + } + require.Same(t, current, options.TraceDispatcher) + require.Equal(t, 1, old.closes) + require.Equal(t, 1, middle.closes) + require.Zero(t, current.closes) + for _, dispatcher := range []*replacementTraceDispatcher{old, middle, current} { + require.Equal(t, 1, dispatcher.starts) + } + ctx := primitive.WithProducerCtx(context.Background(), &primitive.ProducerCtx{ + ProducerGroup: "group", Message: *primitive.NewMessage("topic", []byte("payload")), + }) + reply := &primitive.SendResult{RegionID: "region", TraceOn: true, Status: primitive.SendOK} + businessError := errors.New("business error") + err := primitive.ChainInterceptors(options.Interceptors...)(ctx, nil, reply, + func(context.Context, interface{}, interface{}) error { + calls = append(calls, "business") + return businessError + }) + require.Equal(t, businessError, err) + require.Equal(t, []string{"first", "second", "business"}, calls) + require.Empty(t, old.records, "replaced dispatcher must not receive trace records") + require.Empty(t, middle.records, "repeated replacement must remove each stale interceptor") + require.Len(t, current.records, 1) + require.Len(t, options.Interceptors, 3) +}