diff --git a/go.mod b/go.mod index 09a14d211d..c6428f0d4d 100644 --- a/go.mod +++ b/go.mod @@ -157,6 +157,7 @@ require ( go.opentelemetry.io/otel/exporters/otlp/otlptrace/otlptracehttp v1.44.0 go.opentelemetry.io/otel/sdk v1.44.0 go.opentelemetry.io/otel/trace v1.44.0 + go.uber.org/goleak v1.3.0 go.uber.org/multierr v1.11.0 golang.org/x/crypto v0.52.0 golang.org/x/net v0.55.0 diff --git a/internal/impl/protobuf/multimodule_watcher.go b/internal/impl/protobuf/multimodule_watcher.go index 463ea7a269..5cc6a025c7 100644 --- a/internal/impl/protobuf/multimodule_watcher.go +++ b/internal/impl/protobuf/multimodule_watcher.go @@ -10,6 +10,7 @@ import ( "buf.build/gen/go/bufbuild/reflect/connectrpc/go/buf/reflect/v1beta1/reflectv1beta1connect" connectrpc "connectrpc.com/connect" + "github.com/Jeffail/shutdown" "github.com/bufbuild/prototransform" "google.golang.org/protobuf/reflect/protoreflect" "google.golang.org/protobuf/reflect/protoregistry" @@ -19,6 +20,8 @@ import ( type MultiModuleWatcher struct { bsrClients map[string]*prototransform.SchemaWatcher + httpClient *http.Client + shutSig *shutdown.Signaller } var _ prototransform.Resolver = &MultiModuleWatcher{} @@ -27,10 +30,19 @@ func newMultiModuleWatcher(bsrModules []*service.ParsedConfig) (*MultiModuleWatc if len(bsrModules) == 0 { return nil, errors.New("no modules provided") } - multiModuleWatcher := &MultiModuleWatcher{} + + httpClient := &http.Client{ + Transport: &http.Transport{}, + } + multiModuleWatcher := &MultiModuleWatcher{ + bsrClients: make(map[string]*prototransform.SchemaWatcher), + httpClient: httpClient, + shutSig: shutdown.NewSignaller(), + } + + ctx, _ := multiModuleWatcher.shutSig.SoftStopCtx(context.Background()) // Initialise one client for each module - multiModuleWatcher.bsrClients = make(map[string]*prototransform.SchemaWatcher) for _, bsrModule := range bsrModules { var bsrURL string bsrURL, err := bsrModule.FieldString(fieldBSRUrl) @@ -53,7 +65,7 @@ func newMultiModuleWatcher(bsrModules []*service.ParsedConfig) (*MultiModuleWatc return nil, err } - watcher, err := newSchemaWatcher(context.Background(), bsrURL, bsrAPIKey, module, version) + watcher, err := newSchemaWatcher(ctx, httpClient, bsrURL, bsrAPIKey, module, version) if err != nil { return nil, err } @@ -63,7 +75,7 @@ func newMultiModuleWatcher(bsrModules []*service.ParsedConfig) (*MultiModuleWatc return multiModuleWatcher, nil } -func newSchemaWatcher(ctx context.Context, bsrURL string, bsrAPIKey string, module string, version string) (*prototransform.SchemaWatcher, error) { +func newSchemaWatcher(ctx context.Context, httpClient *http.Client, bsrURL string, bsrAPIKey string, module string, version string) (*prototransform.SchemaWatcher, error) { // If no BSR url provided, extract from module if bsrURL == "" { segments := strings.Split(module, "/") @@ -80,7 +92,7 @@ func newSchemaWatcher(ctx context.Context, bsrURL string, bsrAPIKey string, modu if bsrAPIKey != "" { opts = append(opts, connectrpc.WithInterceptors(prototransform.NewAuthInterceptor(bsrAPIKey))) } - client := reflectv1beta1connect.NewFileDescriptorSetServiceClient(http.DefaultClient, bsrURL, opts...) + client := reflectv1beta1connect.NewFileDescriptorSetServiceClient(httpClient, bsrURL, opts...) cfg := &prototransform.SchemaWatcherConfig{ SchemaPoller: prototransform.NewSchemaPoller( @@ -173,3 +185,14 @@ func (w *MultiModuleWatcher) FindEnumByName(enum protoreflect.FullName) (protore } return nil, fmt.Errorf("could not find %s in any loaded modules", enum) } + +func (m *MultiModuleWatcher) Close() { + m.shutSig.TriggerHardStop() + for _, v := range m.bsrClients { + v.Stop() + } + m.bsrClients = nil + if t, ok := m.httpClient.Transport.(*http.Transport); ok { + t.CloseIdleConnections() + } +} diff --git a/internal/impl/protobuf/processor_protobuf.go b/internal/impl/protobuf/processor_protobuf.go index f6f028a526..eddc02e291 100644 --- a/internal/impl/protobuf/processor_protobuf.go +++ b/internal/impl/protobuf/processor_protobuf.go @@ -549,6 +549,10 @@ func (p *protobufProc) Process(ctx context.Context, msg *service.Message) (servi return service.MessageBatch{msg}, nil } -func (p *protobufProc) Close(context.Context) error { +func (p *protobufProc) Close(ctx context.Context) error { + if p.multiModuleWatcher != nil { + p.multiModuleWatcher.Close() + p.multiModuleWatcher = nil + } return nil } diff --git a/internal/impl/protobuf/processor_protobuf_test.go b/internal/impl/protobuf/processor_protobuf_test.go index 7b2f0cc4ce..e6da6a387a 100644 --- a/internal/impl/protobuf/processor_protobuf_test.go +++ b/internal/impl/protobuf/processor_protobuf_test.go @@ -14,6 +14,7 @@ import ( "connectrpc.com/connect" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + "go.uber.org/goleak" "google.golang.org/protobuf/reflect/protodesc" "google.golang.org/protobuf/reflect/protoreflect" "google.golang.org/protobuf/reflect/protoregistry" @@ -22,6 +23,10 @@ import ( "github.com/warpstreamlabs/bento/public/service" ) +func TestMain(m *testing.M) { + goleak.VerifyTestMain(m) +} + const protosPath = "../../../config/test/protobuf/schema" func TestProtobufFromJSON(t *testing.T) { @@ -103,6 +108,7 @@ discard_unknown: %t assert.Contains(t, string(mBytes), exp) } require.NoError(t, msgs[0].GetError()) + }) t.Run(test.name+" bsr", func(t *testing.T) { @@ -125,6 +131,9 @@ discard_unknown: %t require.NoError(t, res) require.Len(t, msgs, 1) + err = proc.Close(context.Background()) + require.NoError(t, err) + mBytes, err := msgs[0].AsBytes() require.NoError(t, err) @@ -238,6 +247,9 @@ emit_unpopulated: %t require.NoError(t, res) require.Len(t, msgs, 1) + err = proc.Close(context.Background()) + require.NoError(t, err) + mBytes, err := msgs[0].AsBytes() require.NoError(t, err) @@ -266,6 +278,9 @@ emit_unpopulated: %t require.NoError(t, res) require.Len(t, msgs, 1) + err = proc.Close(context.Background()) + require.NoError(t, err) + mBytes, err := msgs[0].AsBytes() require.NoError(t, err) @@ -324,6 +339,9 @@ import_paths: [ %v ] _, err = proc.Process(context.Background(), service.NewMessage([]byte(test.input))) require.Error(t, err) require.Contains(t, err.Error(), test.output) + + err = proc.Close(context.Background()) + require.NoError(t, err) }) t.Run(test.name+" bsr", func(tt *testing.T) { @@ -344,6 +362,9 @@ bsr: _, err = proc.Process(context.Background(), service.NewMessage([]byte(test.input))) require.Error(t, err) require.Contains(t, err.Error(), test.output) + + err = proc.Close(context.Background()) + require.NoError(t, err) }) } } @@ -456,16 +477,21 @@ func runMockBSRServer(t *testing.T) string { mux := http.NewServeMux() fileDescriptorSetServer := &fileDescriptorSetServer{fileDescriptorSet: fileDescriptorSet} mux.Handle(reflectv1beta1connect.NewFileDescriptorSetServiceHandler(fileDescriptorSetServer)) - go func() { - srv := &http.Server{Handler: mux} - srv.Protocols = new(http.Protocols) - srv.Protocols.SetHTTP1(true) - srv.Protocols.SetUnencryptedHTTP2(true) - if err := http.Serve(listener, srv.Handler); err != nil && !errors.Is(err, http.ErrServerClosed) { - require.NoError(t, err) + srv := &http.Server{Handler: mux} + srv.Protocols = new(http.Protocols) + srv.Protocols.SetHTTP1(true) + srv.Protocols.SetUnencryptedHTTP2(true) + + go func() { + if err := srv.Serve(listener); err != nil && !errors.Is(err, http.ErrServerClosed) { + t.Errorf("mock BSR server error: %v", err) } }() + t.Cleanup(func() { + srv.Close() + }) + return listener.Addr().String() }