mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-05 15:26: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.
|
// Container manages service registration and resolution using Go reflection and generics.
|
||||||
type Container struct {
|
type Container struct {
|
||||||
mu sync.RWMutex
|
mu sync.RWMutex
|
||||||
parent *Container
|
parent *Container
|
||||||
services map[reflect.Type]any
|
services map[reflect.Type]any
|
||||||
listeners map[reflect.Type][]func(any)
|
interfaceCache map[reflect.Type]any
|
||||||
|
listeners map[reflect.Type][]func(any)
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewContainer creates a new IoC container instance with an optional parent container.
|
// NewContainer creates a new IoC container instance with an optional parent container.
|
||||||
func NewContainer(parent *Container) *Container {
|
func NewContainer(parent *Container) *Container {
|
||||||
return &Container{
|
return &Container{
|
||||||
parent: parent,
|
parent: parent,
|
||||||
services: make(map[reflect.Type]any),
|
services: make(map[reflect.Type]any),
|
||||||
listeners: make(map[reflect.Type][]func(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()
|
c.mu.Lock()
|
||||||
defer c.mu.Unlock()
|
defer c.mu.Unlock()
|
||||||
delete(c.services, targetType)
|
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.
|
// 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) {
|
func (c *Container) provide(targetType reflect.Type, service any) {
|
||||||
c.mu.Lock()
|
c.mu.Lock()
|
||||||
c.services[targetType] = service
|
c.services[targetType] = service
|
||||||
|
c.interfaceCache = make(map[reflect.Type]any)
|
||||||
|
|
||||||
// Collect any matching listeners to invoke outside the lock
|
// Collect any matching listeners to invoke outside the lock
|
||||||
var callbacks []func(any)
|
var callbacks []func(any)
|
||||||
@@ -131,17 +135,14 @@ func (c *Container) resolve(targetType reflect.Type) (any, error) {
|
|||||||
c.mu.RUnlock()
|
c.mu.RUnlock()
|
||||||
return val, nil
|
return val, nil
|
||||||
}
|
}
|
||||||
|
c.mu.RUnlock()
|
||||||
|
|
||||||
// 2. Interface assignment scan
|
// 2. Interface assignment scan & cache
|
||||||
if targetType.Kind() == reflect.Interface {
|
if targetType.Kind() == reflect.Interface {
|
||||||
for _, val := range c.services {
|
if val, found := c.resolveInterface(targetType); found {
|
||||||
if reflect.TypeOf(val).Implements(targetType) {
|
return val, nil
|
||||||
c.mu.RUnlock()
|
|
||||||
return val, nil
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
c.mu.RUnlock()
|
|
||||||
|
|
||||||
// 3. Fallback to parent container
|
// 3. Fallback to parent container
|
||||||
if c.parent != nil {
|
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)
|
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.
|
// MustInject resolves a service of type T or panics if the service is not found.
|
||||||
func MustInject[T any](ctx *Context) T {
|
func MustInject[T any](ctx *Context) T {
|
||||||
s, err := Inject[T](ctx)
|
s, err := Inject[T](ctx)
|
||||||
|
|||||||
@@ -571,10 +571,16 @@ func TestContext_ScopedExtpoints_RevertibleEffects(t *testing.T) {
|
|||||||
root := core.NewContext(context.Background())
|
root := core.NewContext(context.Background())
|
||||||
child := root.Fork()
|
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() {})
|
child.Router().GET("/test-route", func() {})
|
||||||
assert.Equal(t, 1, len(root.Router().Routes()))
|
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() {})
|
child.Tasks().Register("test:task", func() {})
|
||||||
assert.Equal(t, 1, len(root.Tasks().Tasks()))
|
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
|
// 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().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.Tasks().Tasks()))
|
||||||
assert.Equal(t, 0, len(root.Schedules().Schedules()))
|
assert.Equal(t, 0, len(root.Schedules().Schedules()))
|
||||||
assert.Equal(t, 0, len(root.Settings().Schemas()))
|
assert.Equal(t, 0, len(root.Settings().Schemas()))
|
||||||
assert.Equal(t, 0, root.Events().Listeners("test:event"))
|
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.
|
// 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
|
//nolint:contextcheck
|
||||||
func (b *EventBus) Parallel(ctx context.Context, topic string, payload any) error {
|
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)
|
}(l)
|
||||||
}
|
}
|
||||||
|
|
||||||
wg.Wait()
|
done := make(chan struct{})
|
||||||
close(errCh)
|
go func() {
|
||||||
|
wg.Wait()
|
||||||
|
close(done)
|
||||||
|
}()
|
||||||
|
|
||||||
var errs []error
|
select {
|
||||||
for err := range errCh {
|
case <-done:
|
||||||
if err != nil {
|
close(errCh)
|
||||||
errs = append(errs, err)
|
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.
|
// Serial executes subscribers strictly in sequence.
|
||||||
|
|||||||
@@ -11,6 +11,7 @@ import (
|
|||||||
"sync"
|
"sync"
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
"testing"
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
"github.com/stretchr/testify/require"
|
"github.com/stretchr/testify/require"
|
||||||
@@ -368,6 +369,25 @@ func TestEventBusParallel(t *testing.T) {
|
|||||||
assert.Equal(t, int64(20), counter.Load())
|
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) {
|
func TestEventBusSerial(t *testing.T) {
|
||||||
bus := core.NewEventBus()
|
bus := core.NewEventBus()
|
||||||
|
|
||||||
|
|||||||
@@ -379,7 +379,7 @@ func TestContextExtensionPointsIntegration(t *testing.T) {
|
|||||||
func TestExtensionPointsUnregister(t *testing.T) {
|
func TestExtensionPointsUnregister(t *testing.T) {
|
||||||
ctx := core.NewContext(context.Background())
|
ctx := core.NewContext(context.Background())
|
||||||
|
|
||||||
// 1. Router unregister
|
// 1. Router unregister (routes, middlewares, whitelist)
|
||||||
rd := ctx.Router().GET("/temp", "temp_handler")
|
rd := ctx.Router().GET("/temp", "temp_handler")
|
||||||
assert.Greater(t, rd.ID, uint64(0))
|
assert.Greater(t, rd.ID, uint64(0))
|
||||||
assert.Len(t, ctx.Router().Routes(), 1)
|
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.True(t, ctx.Router().UnregisterByID(rd2.ID))
|
||||||
assert.Len(t, ctx.Router().Routes(), 0)
|
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
|
// 2. Task unregister
|
||||||
ctx.Task().Register("temp:task", "handler")
|
ctx.Task().Register("temp:task", "handler")
|
||||||
assert.Len(t, ctx.Task().Tasks(), 1)
|
assert.Len(t, ctx.Task().Tasks(), 1)
|
||||||
|
|||||||
@@ -39,17 +39,26 @@ type RouterExtension interface {
|
|||||||
Middlewares() []any
|
Middlewares() []any
|
||||||
Unregister(method, path string) bool
|
Unregister(method, path string) bool
|
||||||
UnregisterByID(id uint64) bool
|
UnregisterByID(id uint64) bool
|
||||||
|
UnregisterMiddlewareByID(id uint64) bool
|
||||||
RegisterWhitelist(patterns ...string)
|
RegisterWhitelist(patterns ...string)
|
||||||
|
UnregisterWhitelist(patterns ...string)
|
||||||
Whitelist() []string
|
Whitelist() []string
|
||||||
IsWhitelisted(path string) bool
|
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.
|
// RouterRegistry implements RouterExtension as the root route and middleware collector.
|
||||||
type RouterRegistry struct {
|
type RouterRegistry struct {
|
||||||
mu sync.RWMutex
|
mu sync.RWMutex
|
||||||
nextID uint64
|
nextID uint64
|
||||||
|
nextMWID uint64
|
||||||
routes []RouteDefinition
|
routes []RouteDefinition
|
||||||
middlewares []any
|
middlewares []middlewareDefinition
|
||||||
whitelist PathWhitelist
|
whitelist PathWhitelist
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -60,17 +69,48 @@ func NewRouterRegistry() *RouterRegistry {
|
|||||||
|
|
||||||
// Use registers global middlewares to the router.
|
// Use registers global middlewares to the router.
|
||||||
func (r *RouterRegistry) Use(middlewares ...any) {
|
func (r *RouterRegistry) Use(middlewares ...any) {
|
||||||
r.mu.Lock()
|
r.UseWithID(middlewares...)
|
||||||
defer r.mu.Unlock()
|
|
||||||
r.middlewares = append(r.middlewares, 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 {
|
func (r *RouterRegistry) Middlewares() []any {
|
||||||
r.mu.RLock()
|
r.mu.RLock()
|
||||||
defer r.mu.RUnlock()
|
defer r.mu.RUnlock()
|
||||||
res := make([]any, len(r.middlewares))
|
res := make([]any, len(r.middlewares))
|
||||||
copy(res, r.middlewares)
|
for i, mw := range r.middlewares {
|
||||||
|
res[i] = mw.Handler
|
||||||
|
}
|
||||||
return res
|
return res
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -196,6 +236,11 @@ func (r *RouterRegistry) RegisterWhitelist(patterns ...string) {
|
|||||||
r.whitelist.Add(patterns...)
|
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.
|
// Whitelist returns a copy of all registered whitelist path patterns.
|
||||||
func (r *RouterRegistry) Whitelist() []string {
|
func (r *RouterRegistry) Whitelist() []string {
|
||||||
return r.whitelist.Patterns()
|
return r.whitelist.Patterns()
|
||||||
@@ -267,6 +312,11 @@ func (g *RouterGroup) UnregisterByID(id uint64) bool {
|
|||||||
return g.registry.UnregisterByID(id)
|
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.
|
// GET registers a GET route in this group.
|
||||||
func (g *RouterGroup) GET(path string, handlers ...any) RouteDefinition {
|
func (g *RouterGroup) GET(path string, handlers ...any) RouteDefinition {
|
||||||
return g.Handle("GET", path, handlers...)
|
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.
|
// Whitelist returns a copy of all registered whitelist path patterns.
|
||||||
func (g *RouterGroup) Whitelist() []string {
|
func (g *RouterGroup) Whitelist() []string {
|
||||||
return g.registry.Whitelist()
|
return g.registry.Whitelist()
|
||||||
@@ -465,6 +522,27 @@ func (w *PathWhitelist) Replace(patterns ...string) {
|
|||||||
w.patterns = compiled
|
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
|
// Match reports whether path matches any registered pattern. Equivalent to calling
|
||||||
// MatchPathPattern for every pattern, except the path is normalised and split once.
|
// MatchPathPattern for every pattern, except the path is normalised and split once.
|
||||||
func (w *PathWhitelist) Match(path string) bool {
|
func (w *PathWhitelist) Match(path string) bool {
|
||||||
|
|||||||
@@ -22,6 +22,19 @@ func newScopedRouterExtension(ctx *Context, underlying extpoints.RouterExtension
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (s *scopedRouterExtension) Use(middlewares ...any) {
|
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...)
|
s.underlying.Use(middlewares...)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -107,8 +120,22 @@ func (s *scopedRouterExtension) UnregisterByID(id uint64) bool {
|
|||||||
return s.underlying.UnregisterByID(id)
|
return s.underlying.UnregisterByID(id)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (s *scopedRouterExtension) UnregisterMiddlewareByID(id uint64) bool {
|
||||||
|
return s.underlying.UnregisterMiddlewareByID(id)
|
||||||
|
}
|
||||||
|
|
||||||
func (s *scopedRouterExtension) RegisterWhitelist(patterns ...string) {
|
func (s *scopedRouterExtension) RegisterWhitelist(patterns ...string) {
|
||||||
s.underlying.RegisterWhitelist(patterns...)
|
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 {
|
func (s *scopedRouterExtension) Whitelist() []string {
|
||||||
|
|||||||
Reference in New Issue
Block a user