fix(core): optimize ioc interface caching, event parallel timeout and router teardown reversibility

This commit is contained in:
ryan
2026-09-03 09:09:33 +08:00
parent 30ab8810cc
commit 2124bce7ca
7 changed files with 255 additions and 31 deletions
+44 -14
View File
@@ -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)
+47 -1
View File
@@ -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"))
}
+18 -9
View File
@@ -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.
+20
View File
@@ -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()
+15 -1
View File
@@ -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)
+84 -6
View File
@@ -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 {
+27
View File
@@ -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 {