diff --git a/lifecycle.go b/lifecycle.go index 1fb4320..5616d95 100644 --- a/lifecycle.go +++ b/lifecycle.go @@ -5,8 +5,6 @@ import ( "errors" "fmt" "log/slog" - "sort" - "strings" ) type contextKey string @@ -29,6 +27,7 @@ type Lifecycle struct { modules []*Module opts LifecycleOpts + setupOrder []*Module setupCount int setupTracker map[string]int teardownCount int @@ -38,7 +37,6 @@ type Lifecycle struct { // NewLifecycle creates a new Lifecycle instance with a default logger and the // given modules. It panics if any module has a duplicate name. func NewLifecycle(modules ...*Module) *Lifecycle { - // Ensure modules are unique unique := make(map[string]bool) for _, mod := range modules { if _, exists := unique[mod.name]; exists { @@ -96,17 +94,18 @@ func (app *Lifecycle) Logger() *slog.Logger { return app.opts.Logger } -// Setup initializes all modules in the order they were defined and checks -// for dependencies. It returns an error if any module fails to initialize or -// if dependencies are not satisfied. Lifecycle.Teardown should always be run at the end -// of the application lifecycle to ensure all resources are cleaned up properly. +// Setup initializes all registered modules, resolving dependencies via autoload +// when enabled. Modules run setup in dependency order; among independent modules, +// registration order is preserved. Teardown runs in reverse setup order. func (app *Lifecycle) Setup() error { if app.setupCount > 0 { return fmt.Errorf("lifecycle already set up, cannot set up again") } + setupBefore := len(app.setupOrder) for _, mod := range app.modules { - if err := app.setupSingle(nil, mod); err != nil { + if err := app.setupSingle(nil, mod, nil); err != nil { + app.rollbackFrom(setupBefore) return err } } @@ -117,24 +116,24 @@ func (app *Lifecycle) Setup() error { return nil } -// Teardown runs all module teardown functions in reverse order of setup. -// Teardown should always be run at the end of the application lifecycle -// to ensure all resources are cleaned up properly. All module tear down -// errors are returned as a single error (non-blocking). +// Teardown runs teardown for all set-up modules in reverse setup order. +// All module teardown errors are joined and returned (non-blocking). func (app *Lifecycle) Teardown() error { if app.teardownCount > 0 { return fmt.Errorf("lifecycle already torn down, cannot tear down again") } var err error - for i := len(app.modules) - 1; i >= 0; i-- { - if singleErr := app.teardownSingle(app.modules[i]); singleErr != nil { + var failureCount int + for i := len(app.setupOrder) - 1; i >= 0; i-- { + if singleErr := app.teardownSingle(app.setupOrder[i]); singleErr != nil { err = errors.Join(err, singleErr) + failureCount++ } } if err != nil { - app.Logger().Error("Error tearing down modules", "failures", app.setupCount-app.teardownCount, "error", err) + app.Logger().Error("Error tearing down modules", "failures", failureCount, "error", err) return err } @@ -143,183 +142,3 @@ func (app *Lifecycle) Teardown() error { return nil } - -// Require adds module(s) to the lifecycle and immediately runs any setup -// functions. Relevant when a module is not part of the main application -// but may still be conditionally necessary. Any modules that are already -// set up are ignored. -func (app *Lifecycle) Require(modules ...*Module) error { - return app.require(nil, false, modules...) -} - -// RequireUnique is the same as Require, but it returns an error if any requested -// module is already set up--rather than ignoring it. -func (app *Lifecycle) RequireUnique(modules ...*Module) error { - return app.require(nil, true, modules...) -} - -// RequireL adds module(s) to the lifecycle with a specific logger and -// immediately runs any setup functions. See Require for more details. -// This variation is useful when you need to set up modules with a non- -// default logger. -func (app *Lifecycle) RequireL(logger *slog.Logger, modules ...*Module) error { - if logger == nil { - return fmt.Errorf("logger cannot be nil") - } - return app.require(logger, false, modules...) -} - -// RequireUniqueL is the same as RequireL, but it returns an error if any requested -// module is already set up--rather than ignoring it. -func (app *Lifecycle) RequireUniqueL(logger *slog.Logger, modules ...*Module) error { - if logger == nil { - return fmt.Errorf("logger cannot be nil") - } - return app.require(logger, true, modules...) -} - -// require is a helper function that attempts to add module(s) to the lifecycle -// and immediately run any setup functions. -func (app *Lifecycle) require(logger *slog.Logger, unique bool, modules ...*Module) error { - if len(modules) == 0 { - return fmt.Errorf("no modules to require") - } - - for i, mod := range modules { - if mod == nil { - return fmt.Errorf("module %d is nil", i) - } - - // Check if the module has already been set up - if _, ok := app.setupTracker[mod.name]; ok { - if unique { - return fmt.Errorf("module %s is already set up, cannot require again", mod) - } - - app.Logger().Warn("module already set up, ignoring", "module", mod) - - // Mark duplicate module as loaded - mod.loaded = true - mod.lifecycle = app - mod.logger = logger - - continue - } - - // Add the module to the lifecycle - app.modules = append(app.modules, mod) - - // Run the setup function for the module - if err := app.setupSingle(logger, mod); err != nil { - return fmt.Errorf("error setting up required module %s: %w", mod, err) - } - } - - app.Logger().Info("New modules initialized", "all", mapToString(app.setupTracker)) - return nil -} - -// setupSingle is a helper function to set up a single module. Returns an error -// if the module cannot be set up or if dependencies are not satisfied. -func (app *Lifecycle) setupSingle(logger *slog.Logger, mod *Module) error { - if mod == nil { - return fmt.Errorf("module is nil") - } - - // Set the parent lifecycle and logger override - mod.lifecycle = app - mod.logger = logger - - // Check if all dependencies are satisfied - for _, dep := range mod.depends { - if _, ok := app.setupTracker[dep]; !ok { - if app.opts.DisableAutoload { - return fmt.Errorf("dependency %s not satisfied for '%s'", dep, mod) - } else { - // Attempt to set up the dependency - depmod, err := app.getModuleByName(dep) - if err != nil { - return fmt.Errorf("error getting dependency '%s' for %s: %w", dep, mod, err) - } - - if err := app.setupSingle(logger, depmod); err != nil { - return fmt.Errorf("error setting up dependency %s for %s: %w", depmod, mod, err) - } - } - } - } - - if mod.setup != nil { - // Run the setup function for the module - if err := mod.setup(mod); err != nil { - return fmt.Errorf("error initializing %s: %w", mod, err) - } - } - - // Mark this module as setup - app.setupTracker[mod.name] = app.setupCount - app.setupCount++ - mod.loaded = true - return nil -} - -// teardownSingle is a helper function to tear down a single module. Returns -// an error if anything goes wrong. -func (app *Lifecycle) teardownSingle(mod *Module) error { - if mod == nil { - return fmt.Errorf("module is nil") - } - - // Check if the module was set up - if _, ok := app.setupTracker[mod.name]; !ok { - return fmt.Errorf("module %s is not set up, cannot tear down", mod) - } - - // Run the teardown function for the module - if mod.teardown != nil { - if err := mod.teardown(mod); err != nil { - return fmt.Errorf("error tearing down %s: %w", mod, err) - } - } - - // Mark this module as torn down - app.teardownTracker[mod.name] = app.teardownCount - app.teardownCount++ - mod.loaded = false - return nil -} - -// getModuleByName retrieves a module by its name from the lifecycle. -func (app *Lifecycle) getModuleByName(name string) (*Module, error) { - for _, mod := range app.modules { - if mod.name == name { - return mod, nil - } - } - return nil, ErrModuleNotFound -} - -// mapToString converts a map to an ordered, opinionated string representation. -func mapToString(m map[string]int) string { - if len(m) == 0 { - return "[]" - } - - // Sort the map by value - values := make([]int, 0, len(m)) - reverseMap := make(map[int]string, len(m)) - for k, v := range m { - values = append(values, v) - reverseMap[v] = k - } - sort.Ints(values) - - result := make([]string, 0, len(m)) - for _, v := range values { - if name, ok := reverseMap[v]; ok { - result = append(result, fmt.Sprintf("'%s'", name)) - } - } - - return "[" + strings.Join(result, " ") + "]" -} diff --git a/lifecycle_integration_test.go b/lifecycle_integration_test.go new file mode 100644 index 0000000..5187524 --- /dev/null +++ b/lifecycle_integration_test.go @@ -0,0 +1,81 @@ +package app_test + +import ( + "fmt" + "testing" + + "gitea.auvem.com/go-toolkit/app" + "github.com/stretchr/testify/assert" +) + +func TestSetupTeardownIntegration(t *testing.T) { + var order []string + modA := app.NewModule("a", app.ModuleOpts{ + Setup: func(m *app.Module) error { order = append(order, "setup:a"); return nil }, + Teardown: func(m *app.Module) error { order = append(order, "teardown:a"); return nil }, + }) + modB := app.NewModule("b", app.ModuleOpts{ + Setup: func(m *app.Module) error { order = append(order, "setup:b"); return nil }, + Teardown: func(m *app.Module) error { order = append(order, "teardown:b"); return nil }, + Depends: []string{"a"}, + }) + + lc := app.NewLifecycle(modB, modA) + assert.NoError(t, lc.Setup()) + assert.Equal(t, []string{"setup:a", "setup:b"}, order) + + order = nil + assert.NoError(t, lc.Teardown()) + assert.Equal(t, []string{"teardown:b", "teardown:a"}, order) +} + +func TestSetupNoDoubleSetupWithAutoload(t *testing.T) { + var count int + modA := app.NewModule("a", app.ModuleOpts{ + Setup: func(m *app.Module) error { count++; return nil }, + }) + modB := app.NewModule("b", app.ModuleOpts{ + Setup: func(m *app.Module) error { count++; return nil }, + Depends: []string{"a"}, + }) + + lc := app.NewLifecycle(modB, modA) + assert.NoError(t, lc.Setup()) + assert.Equal(t, 2, count) +} + +func TestSetupPartialFailureRollback(t *testing.T) { + var tornDown bool + modA := app.NewModule("a", app.ModuleOpts{ + Setup: func(m *app.Module) error { return nil }, + Teardown: func(m *app.Module) error { tornDown = true; return nil }, + }) + modB := app.NewModule("b", app.ModuleOpts{ + Setup: func(m *app.Module) error { return fmt.Errorf("fail b") }, + }) + + lc := app.NewLifecycle(modA, modB) + err := lc.Setup() + assert.Error(t, err) + assert.True(t, tornDown) + assert.False(t, modA.Loaded()) +} + +func TestSetupCircularDependency(t *testing.T) { + modA := app.NewModule("a", app.ModuleOpts{Depends: []string{"b"}}) + modB := app.NewModule("b", app.ModuleOpts{Depends: []string{"a"}}) + lc := app.NewLifecycle(modA, modB) + err := lc.Setup() + assert.Error(t, err) + assert.Contains(t, err.Error(), "circular dependency") +} + +func TestGetModule(t *testing.T) { + mod := app.NewModule("db", app.ModuleOpts{}) + lc := app.NewLifecycle(mod) + got, err := lc.GetModule("db") + assert.NoError(t, err) + assert.Equal(t, mod, got) + _, err = lc.GetModule("missing") + assert.Error(t, err) +} diff --git a/lifecycle_internal.go b/lifecycle_internal.go new file mode 100644 index 0000000..8bf2824 --- /dev/null +++ b/lifecycle_internal.go @@ -0,0 +1,170 @@ +package app + +import ( + "fmt" + "log/slog" + "sort" + "strings" +) + +func (app *Lifecycle) require(opts RequireOpts, modules ...*Module) error { + if len(modules) == 0 { + return fmt.Errorf("no modules to require") + } + + if opts.Logger != nil && opts.Unique { + // unique with custom logger is valid + } + if opts.Logger == nil && opts.Unique { + // valid + } + + setupBefore := len(app.setupOrder) + + for i, mod := range modules { + if mod == nil { + app.rollbackFrom(setupBefore) + return fmt.Errorf("module %d is nil", i) + } + + if _, ok := app.setupTracker[mod.name]; ok { + if opts.Unique { + app.rollbackFrom(setupBefore) + return fmt.Errorf("module %s is already set up, cannot require again", mod) + } + + app.Logger().Warn("module already set up, ignoring", "module", mod) + mod.loaded = true + mod.lifecycle = app + mod.logger = opts.Logger + continue + } + + app.modules = append(app.modules, mod) + + if err := app.setupSingle(opts.Logger, mod, nil); err != nil { + app.rollbackFrom(setupBefore) + return fmt.Errorf("error setting up required module %s: %w", mod, err) + } + } + + app.Logger().Info("New modules initialized", "all", mapToString(app.setupTracker)) + return nil +} + +func (app *Lifecycle) setupSingle(logger *slog.Logger, mod *Module, visiting map[string]bool) error { + if mod == nil { + return fmt.Errorf("module is nil") + } + + if _, ok := app.setupTracker[mod.name]; ok { + return nil + } + + if visiting == nil { + visiting = make(map[string]bool) + } + if visiting[mod.name] { + return fmt.Errorf("circular dependency detected involving %s", mod) + } + visiting[mod.name] = true + defer delete(visiting, mod.name) + + mod.lifecycle = app + mod.logger = logger + + for _, dep := range mod.depends { + if _, ok := app.setupTracker[dep]; !ok { + if app.opts.DisableAutoload { + return fmt.Errorf("dependency %s not satisfied for '%s'", dep, mod) + } + + depmod, err := app.getModuleByName(dep) + if err != nil { + return fmt.Errorf("error getting dependency '%s' for %s: %w", dep, mod, err) + } + + if err := app.setupSingle(logger, depmod, visiting); err != nil { + return fmt.Errorf("error setting up dependency %s for %s: %w", depmod, mod, err) + } + } + } + + if mod.setup != nil { + if err := mod.setup(mod); err != nil { + return fmt.Errorf("error initializing %s: %w", mod, err) + } + } + + app.setupTracker[mod.name] = app.setupCount + app.setupOrder = append(app.setupOrder, mod) + app.setupCount++ + mod.loaded = true + return nil +} + +func (app *Lifecycle) teardownSingle(mod *Module) error { + if mod == nil { + return fmt.Errorf("module is nil") + } + + if _, ok := app.setupTracker[mod.name]; !ok { + return fmt.Errorf("module %s is not set up, cannot tear down", mod) + } + + if mod.teardown != nil { + if err := mod.teardown(mod); err != nil { + return fmt.Errorf("error tearing down %s: %w", mod, err) + } + } + + app.teardownTracker[mod.name] = app.teardownCount + app.teardownCount++ + mod.loaded = false + return nil +} + +func (app *Lifecycle) rollbackFrom(startIndex int) { + for i := len(app.setupOrder) - 1; i >= startIndex; i-- { + mod := app.setupOrder[i] + if mod.teardown != nil { + _ = mod.teardown(mod) + } + delete(app.setupTracker, mod.name) + mod.loaded = false + } + app.setupOrder = app.setupOrder[:startIndex] + app.setupCount = startIndex +} + +func (app *Lifecycle) getModuleByName(name string) (*Module, error) { + for _, mod := range app.modules { + if mod.name == name { + return mod, nil + } + } + return nil, ErrModuleNotFound +} + +func mapToString(m map[string]int) string { + if len(m) == 0 { + return "[]" + } + + values := make([]int, 0, len(m)) + reverseMap := make(map[int]string, len(m)) + for k, v := range m { + values = append(values, v) + reverseMap[v] = k + } + sort.Ints(values) + + result := make([]string, 0, len(m)) + for _, v := range values { + if name, ok := reverseMap[v]; ok { + result = append(result, fmt.Sprintf("'%s'", name)) + } + } + + return "[" + strings.Join(result, " ") + "]" +} diff --git a/lifecycle_test.go b/lifecycle_test.go index be2e22a..818060a 100644 --- a/lifecycle_test.go +++ b/lifecycle_test.go @@ -246,10 +246,16 @@ func TestLifecycle_Teardown(t *testing.T) { lc := NewLifecycle(tc.modules...) - // Fake setup for all modules - for _, mod := range tc.modules { - mod.loaded = true // Mark as loaded - lc.setupTracker[mod.name] = 0 // Mark as set up + if len(tc.modules) > 0 { + setupBefore := len(lc.setupOrder) + for i, mod := range tc.modules { + lc.setupOrder = append(lc.setupOrder, mod) + lc.setupTracker[mod.name] = i + mod.loaded = true + mod.lifecycle = lc + } + lc.setupCount = len(lc.setupOrder) + _ = setupBefore } err := lc.Teardown() @@ -283,6 +289,8 @@ func TestLifecycle_Teardown(t *testing.T) { // Fake setup for the module lc.modules[0].loaded = true lc.setupTracker[lc.modules[0].name] = 0 + lc.setupOrder = []*Module{lc.modules[0]} + lc.setupCount = 1 err := lc.Teardown() assert.NoError(err, "expected first Teardown to succeed") @@ -418,7 +426,7 @@ func TestLifecycle_require(t *testing.T) { assert := assert.New(t) lc := NewLifecycle() - err := lc.require(tc.logger, tc.unique, tc.modules...) + err := lc.require(RequireOpts{Logger: tc.logger, Unique: tc.unique}, tc.modules...) if tc.expectedErr == "" { assert.NoError(err, "expected require to succeed") @@ -550,7 +558,7 @@ func TestLifecycle_setupSingle(t *testing.T) { lc = NewLifecycle(tc.modules...) } - err := lc.setupSingle(l, tc.targetModule) + err := lc.setupSingle(l, tc.targetModule, nil) if tc.expectedErr == "" { assert.NoError(err, "expected no error from setupSingle") @@ -616,12 +624,17 @@ func TestLifecycle_teardownSingle(t *testing.T) { lc := NewLifecycle() - // Fake setup for all modules var setupCount int - for _, mod := range tc.modules { - lc.setupTracker[mod] = setupCount + for _, modName := range tc.modules { + lc.setupTracker[modName] = setupCount + for _, m := range lc.modules { + if m.name == modName { + lc.setupOrder = append(lc.setupOrder, m) + } + } setupCount++ } + lc.setupCount = setupCount err := lc.teardownSingle(tc.targetModule)