mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-09-28 05:46:36 +08:00
fix(core): optimize ioc interface caching, event parallel timeout and router teardown reversibility
This commit is contained in:
+44
-14
@@ -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)
|
||||
|
||||
@@ -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
@@ -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.
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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 {
|
||||
|
||||
Reference in New Issue
Block a user