diff --git a/backend/core/container.go b/backend/core/container.go index 64d8c030..a910a314 100644 --- a/backend/core/container.go +++ b/backend/core/container.go @@ -13,18 +13,20 @@ import ( // Container manages service registration and resolution using Go reflection and generics. type Container struct { - mu sync.RWMutex - parent *Container - services map[reflect.Type]any - listeners map[reflect.Type][]func(any) + mu sync.RWMutex + parent *Container + services map[reflect.Type]any + interfaceCache map[reflect.Type]any + listeners map[reflect.Type][]func(any) } // NewContainer creates a new IoC container instance with an optional parent container. func NewContainer(parent *Container) *Container { return &Container{ - parent: parent, - services: make(map[reflect.Type]any), - listeners: make(map[reflect.Type][]func(any)), + parent: parent, + services: make(map[reflect.Type]any), + interfaceCache: make(map[reflect.Type]any), + listeners: make(map[reflect.Type][]func(any)), } } @@ -45,6 +47,7 @@ func (c *Container) remove(targetType reflect.Type) { c.mu.Lock() defer c.mu.Unlock() delete(c.services, targetType) + c.interfaceCache = make(map[reflect.Type]any) } // Provide registers a typed service implementation into the Context hierarchy's root IoC container. @@ -88,6 +91,7 @@ func ProvideScoped[T any](ctx *Context, service T) { func (c *Container) provide(targetType reflect.Type, service any) { c.mu.Lock() c.services[targetType] = service + c.interfaceCache = make(map[reflect.Type]any) // Collect any matching listeners to invoke outside the lock var callbacks []func(any) @@ -131,17 +135,14 @@ func (c *Container) resolve(targetType reflect.Type) (any, error) { c.mu.RUnlock() return val, nil } + c.mu.RUnlock() - // 2. Interface assignment scan + // 2. Interface assignment scan & cache if targetType.Kind() == reflect.Interface { - for _, val := range c.services { - if reflect.TypeOf(val).Implements(targetType) { - c.mu.RUnlock() - return val, nil - } + if val, found := c.resolveInterface(targetType); found { + return val, nil } } - c.mu.RUnlock() // 3. Fallback to parent container if c.parent != nil { @@ -151,6 +152,35 @@ func (c *Container) resolve(targetType reflect.Type) (any, error) { return nil, fmt.Errorf("%w: %v", ErrServiceNotFound, targetType) } +func (c *Container) resolveInterface(targetType reflect.Type) (any, bool) { + c.mu.RLock() + if val, ok := c.interfaceCache[targetType]; ok { + c.mu.RUnlock() + return val, true + } + + var matched any + for _, val := range c.services { + if reflect.TypeOf(val).Implements(targetType) { + matched = val + break + } + } + c.mu.RUnlock() + + if matched == nil { + return nil, false + } + + c.mu.Lock() + if c.interfaceCache == nil { + c.interfaceCache = make(map[reflect.Type]any) + } + c.interfaceCache[targetType] = matched + c.mu.Unlock() + return matched, true +} + // MustInject resolves a service of type T or panics if the service is not found. func MustInject[T any](ctx *Context) T { s, err := Inject[T](ctx) diff --git a/backend/core/context_test.go b/backend/core/context_test.go index 9a9bce93..f7a0a3e4 100644 --- a/backend/core/context_test.go +++ b/backend/core/context_test.go @@ -571,10 +571,16 @@ func TestContext_ScopedExtpoints_RevertibleEffects(t *testing.T) { root := core.NewContext(context.Background()) child := root.Fork() - // Register route, task, schedule, setting, event on child + // Register route, task, schedule, setting, event, middleware, whitelist on child child.Router().GET("/test-route", func() {}) assert.Equal(t, 1, len(root.Router().Routes())) + child.Router().Use("scoped_middleware") + assert.Equal(t, 1, len(root.Router().Middlewares())) + + child.Router().RegisterWhitelist("/api/v1/scoped/*") + assert.True(t, root.Router().IsWhitelisted("/api/v1/scoped/test")) + child.Tasks().Register("test:task", func() {}) assert.Equal(t, 1, len(root.Tasks().Tasks())) @@ -593,8 +599,48 @@ func TestContext_ScopedExtpoints_RevertibleEffects(t *testing.T) { // All child effects should be cleanly revoked in LIFO order assert.Equal(t, 0, len(root.Router().Routes())) + assert.Equal(t, 0, len(root.Router().Middlewares())) + assert.False(t, root.Router().IsWhitelisted("/api/v1/scoped/test")) assert.Equal(t, 0, len(root.Tasks().Tasks())) assert.Equal(t, 0, len(root.Schedules().Schedules())) assert.Equal(t, 0, len(root.Settings().Schemas())) assert.Equal(t, 0, root.Events().Listeners("test:event")) } + +func TestContainer_InterfaceResolutionCache(t *testing.T) { + ctx := core.NewContext(context.Background()) + svc := &sampleServiceImpl{prefix: "Cached:"} + + core.Provide[SampleService](ctx, svc) + + // 1. Initial resolution populates interfaceCache + res1, err := core.Inject[SampleService](ctx) + require.NoError(t, err) + assert.Equal(t, "Cached: Alice", res1.Greet("Alice")) + + // 2. Subsequent resolutions hit interfaceCache + res2, err := core.Inject[SampleService](ctx) + require.NoError(t, err) + assert.Same(t, res1, res2) + + // 3. Concurrent lookups + var wg sync.WaitGroup + for i := 0; i < 20; i++ { + wg.Add(1) + go func() { + defer wg.Done() + r, e := core.Inject[SampleService](ctx) + assert.NoError(t, e) + assert.Equal(t, "Cached: Bob", r.Greet("Bob")) + }() + } + wg.Wait() + + // 4. Overriding/providing another service invalidates cache + svc2 := &sampleServiceImpl{prefix: "Updated:"} + core.Provide[SampleService](ctx, svc2) + + res3, err := core.Inject[SampleService](ctx) + require.NoError(t, err) + assert.Equal(t, "Updated: Alice", res3.Greet("Alice")) +} diff --git a/backend/core/events.go b/backend/core/events.go index d30be096..e6ee840c 100644 --- a/backend/core/events.go +++ b/backend/core/events.go @@ -283,7 +283,8 @@ func (b *EventBus) Waterfall(ctx context.Context, topic string, initialPayload a } // Parallel executes all subscribers of the topic concurrently in separate goroutines. -// It waits for all handlers to complete and collects any errors via errors.Join. +// It waits for all handlers to complete or returns immediately if ctx is cancelled/timed out, +// collecting any handler errors via errors.Join. // //nolint:contextcheck func (b *EventBus) Parallel(ctx context.Context, topic string, payload any) error { @@ -325,17 +326,25 @@ func (b *EventBus) Parallel(ctx context.Context, topic string, payload any) erro }(l) } - wg.Wait() - close(errCh) + done := make(chan struct{}) + go func() { + wg.Wait() + close(done) + }() - var errs []error - for err := range errCh { - if err != nil { - errs = append(errs, err) + select { + case <-done: + close(errCh) + var errs []error + for err := range errCh { + if err != nil { + errs = append(errs, err) + } } + return errors.Join(errs...) + case <-ctx.Done(): + return ctx.Err() } - - return errors.Join(errs...) } // Serial executes subscribers strictly in sequence. diff --git a/backend/core/events_test.go b/backend/core/events_test.go index 730661fc..d4801632 100644 --- a/backend/core/events_test.go +++ b/backend/core/events_test.go @@ -11,6 +11,7 @@ import ( "sync" "sync/atomic" "testing" + "time" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -368,6 +369,25 @@ func TestEventBusParallel(t *testing.T) { assert.Equal(t, int64(20), counter.Load()) } +func TestEventBusParallelContextTimeout(t *testing.T) { + bus := core.NewEventBus() + + bus.On("test:timeout", func(ctx context.Context) error { + select { + case <-time.After(200 * time.Millisecond): + return nil + case <-ctx.Done(): + return ctx.Err() + } + }) + + ctx, cancel := context.WithTimeout(context.Background(), 20*time.Millisecond) + defer cancel() + + err := bus.Parallel(ctx, "test:timeout", nil) + assert.ErrorIs(t, err, context.DeadlineExceeded) +} + func TestEventBusSerial(t *testing.T) { bus := core.NewEventBus() diff --git a/backend/core/extpoints/extpoints_test.go b/backend/core/extpoints/extpoints_test.go index d5612d7c..1c558db0 100644 --- a/backend/core/extpoints/extpoints_test.go +++ b/backend/core/extpoints/extpoints_test.go @@ -379,7 +379,7 @@ func TestContextExtensionPointsIntegration(t *testing.T) { func TestExtensionPointsUnregister(t *testing.T) { ctx := core.NewContext(context.Background()) - // 1. Router unregister + // 1. Router unregister (routes, middlewares, whitelist) rd := ctx.Router().GET("/temp", "temp_handler") assert.Greater(t, rd.ID, uint64(0)) assert.Len(t, ctx.Router().Routes(), 1) @@ -391,6 +391,20 @@ func TestExtensionPointsUnregister(t *testing.T) { assert.True(t, ctx.Router().UnregisterByID(rd2.ID)) assert.Len(t, ctx.Router().Routes(), 0) + ctx.Router().Use("mw1") + assert.Len(t, ctx.Router().Middlewares(), 1) + if reg, ok := ctx.Router().(*extpoints.RouterRegistry); ok { + ids := reg.UseWithID("mw2") + assert.Len(t, ctx.Router().Middlewares(), 2) + assert.True(t, ctx.Router().UnregisterMiddlewareByID(ids[0])) + assert.Len(t, ctx.Router().Middlewares(), 1) + } + + ctx.Router().RegisterWhitelist("/api/v1/temp/*") + assert.True(t, ctx.Router().IsWhitelisted("/api/v1/temp/item")) + ctx.Router().UnregisterWhitelist("/api/v1/temp/*") + assert.False(t, ctx.Router().IsWhitelisted("/api/v1/temp/item")) + // 2. Task unregister ctx.Task().Register("temp:task", "handler") assert.Len(t, ctx.Task().Tasks(), 1) diff --git a/backend/core/extpoints/router.go b/backend/core/extpoints/router.go index 61e47262..92bcda33 100644 --- a/backend/core/extpoints/router.go +++ b/backend/core/extpoints/router.go @@ -39,17 +39,26 @@ type RouterExtension interface { Middlewares() []any Unregister(method, path string) bool UnregisterByID(id uint64) bool + UnregisterMiddlewareByID(id uint64) bool RegisterWhitelist(patterns ...string) + UnregisterWhitelist(patterns ...string) Whitelist() []string IsWhitelisted(path string) bool } +// middlewareDefinition holds an assigned ID and handler for registered middleware. +type middlewareDefinition struct { + ID uint64 + Handler any +} + // RouterRegistry implements RouterExtension as the root route and middleware collector. type RouterRegistry struct { mu sync.RWMutex nextID uint64 + nextMWID uint64 routes []RouteDefinition - middlewares []any + middlewares []middlewareDefinition whitelist PathWhitelist } @@ -60,17 +69,48 @@ func NewRouterRegistry() *RouterRegistry { // Use registers global middlewares to the router. func (r *RouterRegistry) Use(middlewares ...any) { - r.mu.Lock() - defer r.mu.Unlock() - r.middlewares = append(r.middlewares, middlewares...) + r.UseWithID(middlewares...) } -// Middlewares returns a copy of registered root middlewares. +// UseWithID registers global middlewares to the router and returns their assigned IDs. +func (r *RouterRegistry) UseWithID(middlewares ...any) []uint64 { + r.mu.Lock() + defer r.mu.Unlock() + + ids := make([]uint64, 0, len(middlewares)) + for _, mw := range middlewares { + r.nextMWID++ + r.middlewares = append(r.middlewares, middlewareDefinition{ + ID: r.nextMWID, + Handler: mw, + }) + ids = append(ids, r.nextMWID) + } + return ids +} + +// UnregisterMiddlewareByID removes a registered global middleware by its unique ID. +func (r *RouterRegistry) UnregisterMiddlewareByID(id uint64) bool { + r.mu.Lock() + defer r.mu.Unlock() + + for i, mw := range r.middlewares { + if mw.ID == id { + r.middlewares = append(r.middlewares[:i], r.middlewares[i+1:]...) + return true + } + } + return false +} + +// Middlewares returns a copy of registered root middleware handlers. func (r *RouterRegistry) Middlewares() []any { r.mu.RLock() defer r.mu.RUnlock() res := make([]any, len(r.middlewares)) - copy(res, r.middlewares) + for i, mw := range r.middlewares { + res[i] = mw.Handler + } return res } @@ -196,6 +236,11 @@ func (r *RouterRegistry) RegisterWhitelist(patterns ...string) { r.whitelist.Add(patterns...) } +// UnregisterWhitelist removes path patterns from the whitelist. +func (r *RouterRegistry) UnregisterWhitelist(patterns ...string) { + r.whitelist.Remove(patterns...) +} + // Whitelist returns a copy of all registered whitelist path patterns. func (r *RouterRegistry) Whitelist() []string { return r.whitelist.Patterns() @@ -267,6 +312,11 @@ func (g *RouterGroup) UnregisterByID(id uint64) bool { return g.registry.UnregisterByID(id) } +// UnregisterMiddlewareByID removes a middleware by ID via the root registry. +func (g *RouterGroup) UnregisterMiddlewareByID(id uint64) bool { + return g.registry.UnregisterMiddlewareByID(id) +} + // GET registers a GET route in this group. func (g *RouterGroup) GET(path string, handlers ...any) RouteDefinition { return g.Handle("GET", path, handlers...) @@ -331,6 +381,13 @@ func (g *RouterGroup) RegisterWhitelist(patterns ...string) { } } +// UnregisterWhitelist removes path patterns under this group prefix from the whitelist. +func (g *RouterGroup) UnregisterWhitelist(patterns ...string) { + for _, p := range patterns { + g.registry.UnregisterWhitelist(joinPaths(g.prefix, p)) + } +} + // Whitelist returns a copy of all registered whitelist path patterns. func (g *RouterGroup) Whitelist() []string { return g.registry.Whitelist() @@ -465,6 +522,27 @@ func (w *PathWhitelist) Replace(patterns ...string) { w.patterns = compiled } +// Remove removes matching patterns from the whitelist. +func (w *PathWhitelist) Remove(patterns ...string) { + if len(patterns) == 0 { + return + } + targets := make(map[string]struct{}, len(patterns)) + for _, p := range patterns { + targets[cleanPath(p)] = struct{}{} + } + + w.mu.Lock() + defer w.mu.Unlock() + filtered := w.patterns[:0] + for _, p := range w.patterns { + if _, remove := targets[p.raw]; !remove { + filtered = append(filtered, p) + } + } + w.patterns = filtered +} + // Match reports whether path matches any registered pattern. Equivalent to calling // MatchPathPattern for every pattern, except the path is normalised and split once. func (w *PathWhitelist) Match(path string) bool { diff --git a/backend/core/scoped_extpoints.go b/backend/core/scoped_extpoints.go index f567544a..812709d6 100644 --- a/backend/core/scoped_extpoints.go +++ b/backend/core/scoped_extpoints.go @@ -22,6 +22,19 @@ func newScopedRouterExtension(ctx *Context, underlying extpoints.RouterExtension } func (s *scopedRouterExtension) Use(middlewares ...any) { + if len(middlewares) == 0 { + return + } + if reg, ok := s.underlying.(*extpoints.RouterRegistry); ok { + ids := reg.UseWithID(middlewares...) + s.ctx.OnDispose(func() error { + for _, id := range ids { + reg.UnregisterMiddlewareByID(id) + } + return nil + }) + return + } s.underlying.Use(middlewares...) } @@ -107,8 +120,22 @@ func (s *scopedRouterExtension) UnregisterByID(id uint64) bool { return s.underlying.UnregisterByID(id) } +func (s *scopedRouterExtension) UnregisterMiddlewareByID(id uint64) bool { + return s.underlying.UnregisterMiddlewareByID(id) +} + func (s *scopedRouterExtension) RegisterWhitelist(patterns ...string) { s.underlying.RegisterWhitelist(patterns...) + if len(patterns) > 0 { + s.ctx.OnDispose(func() error { + s.underlying.UnregisterWhitelist(patterns...) + return nil + }) + } +} + +func (s *scopedRouterExtension) UnregisterWhitelist(patterns ...string) { + s.underlying.UnregisterWhitelist(patterns...) } func (s *scopedRouterExtension) Whitelist() []string {