From c53a5461abedbbb2942ada5da26afe9507969340 Mon Sep 17 00:00:00 2001 From: ryan Date: Thu, 27 Aug 2026 23:47:57 +0800 Subject: [PATCH] feat(core): add typed eventbus and domain extension points --- core/context.go | 84 ++++++++- core/context_test.go | 24 +++ core/events.go | 271 ++++++++++++++++++++++++++ core/events_test.go | 315 +++++++++++++++++++++++++++++++ core/extpoints/extpoints_test.go | 277 +++++++++++++++++++++++++++ core/extpoints/migration.go | 83 ++++++++ core/extpoints/router.go | 266 ++++++++++++++++++++++++++ core/extpoints/schedule.go | 100 ++++++++++ core/extpoints/setting.go | 77 ++++++++ core/extpoints/task.go | 119 ++++++++++++ core/types.go | 18 ++ 11 files changed, 1625 insertions(+), 9 deletions(-) create mode 100644 core/events.go create mode 100644 core/events_test.go create mode 100644 core/extpoints/extpoints_test.go create mode 100644 core/extpoints/migration.go create mode 100644 core/extpoints/router.go create mode 100644 core/extpoints/schedule.go create mode 100644 core/extpoints/setting.go create mode 100644 core/extpoints/task.go diff --git a/core/context.go b/core/context.go index 980521d9..797f2180 100644 --- a/core/context.go +++ b/core/context.go @@ -6,6 +6,8 @@ import ( "fmt" "sync" "time" + + "github.com/Rain-kl/Wavelet/core/extpoints" ) // Context is the central micro-kernel service bus and runtime lifecycle container. @@ -17,6 +19,13 @@ type Context struct { parent *Context container *Container + events *EventBus + router extpoints.RouterExtension + migrations extpoints.MigrationExtension + tasks extpoints.TaskExtension + schedules extpoints.ScheduleExtension + settings extpoints.SettingExtension + mu sync.RWMutex children []*Context disposers []Disposer @@ -34,10 +43,16 @@ func NewContext(base context.Context) *Context { ctx, cancel := context.WithCancel(base) return &Context{ - goCtx: ctx, - cancel: cancel, - container: NewContainer(nil), - values: make(map[any]any), + goCtx: ctx, + cancel: cancel, + container: NewContainer(nil), + events: NewEventBus(), + router: extpoints.NewRouterRegistry(), + migrations: extpoints.NewMigrationRegistry(), + tasks: extpoints.NewTaskRegistry(), + schedules: extpoints.NewScheduleRegistry(), + settings: extpoints.NewSettingRegistry(), + values: make(map[any]any), } } @@ -127,11 +142,17 @@ func (c *Context) ForkWithContext(base context.Context) *Context { ctx, cancel := context.WithCancel(base) child := &Context{ - goCtx: ctx, - cancel: cancel, - parent: c, - container: NewContainer(c.container), - values: make(map[any]any), + goCtx: ctx, + cancel: cancel, + parent: c, + container: NewContainer(c.container), + events: c.Events(), + router: c.Router(), + migrations: c.Migrations(), + tasks: c.Tasks(), + schedules: c.Schedules(), + settings: c.Settings(), + values: make(map[any]any), } c.mu.Lock() @@ -141,6 +162,51 @@ func (c *Context) ForkWithContext(base context.Context) *Context { return child } +// Events returns the domain EventBus associated with this Context hierarchy. +func (c *Context) Events() *EventBus { + return c.events +} + +// Router returns the RouterExtension registry. +func (c *Context) Router() extpoints.RouterExtension { + return c.router +} + +// Migrations returns the MigrationExtension registry. +func (c *Context) Migrations() extpoints.MigrationExtension { + return c.migrations +} + +// Tasks returns the TaskExtension registry. +func (c *Context) Tasks() extpoints.TaskExtension { + return c.tasks +} + +// Task is an alias for Tasks(). +func (c *Context) Task() extpoints.TaskExtension { + return c.tasks +} + +// Schedules returns the ScheduleExtension registry. +func (c *Context) Schedules() extpoints.ScheduleExtension { + return c.schedules +} + +// Schedule is an alias for Schedules(). +func (c *Context) Schedule() extpoints.ScheduleExtension { + return c.schedules +} + +// Settings returns the SettingExtension registry. +func (c *Context) Settings() extpoints.SettingExtension { + return c.settings +} + +// Setting is an alias for Settings(). +func (c *Context) Setting() extpoints.SettingExtension { + return c.settings +} + // OnDispose registers a cleanup callback function to be executed when this Context is disposed. // It accepts func() error, func(), or Disposer. func (c *Context) OnDispose(fn any) { diff --git a/core/context_test.go b/core/context_test.go index dc9d1282..e06edb9e 100644 --- a/core/context_test.go +++ b/core/context_test.go @@ -473,3 +473,27 @@ func TestConcurrentAccess(t *testing.T) { wg.Wait() } + +func TestContextExtensionPointsAccessors(t *testing.T) { + ctx := core.NewContext(nil) + assert.NotNil(t, ctx.Events()) + assert.NotNil(t, ctx.Router()) + assert.NotNil(t, ctx.Migrations()) + assert.NotNil(t, ctx.Tasks()) + assert.NotNil(t, ctx.Task()) + assert.NotNil(t, ctx.Schedules()) + assert.NotNil(t, ctx.Schedule()) + assert.NotNil(t, ctx.Settings()) + assert.NotNil(t, ctx.Setting()) + + child := ctx.Fork() + assert.Equal(t, ctx.Events(), child.Events()) + assert.Equal(t, ctx.Router(), child.Router()) + assert.Equal(t, ctx.Migrations(), child.Migrations()) + assert.Equal(t, ctx.Tasks(), child.Tasks()) + assert.Equal(t, ctx.Task(), child.Task()) + assert.Equal(t, ctx.Schedules(), child.Schedules()) + assert.Equal(t, ctx.Schedule(), child.Schedule()) + assert.Equal(t, ctx.Settings(), child.Settings()) + assert.Equal(t, ctx.Setting(), child.Setting()) +} diff --git a/core/events.go b/core/events.go new file mode 100644 index 00000000..2e9b8534 --- /dev/null +++ b/core/events.go @@ -0,0 +1,271 @@ +package core + +import ( + "context" + "errors" + "fmt" + "reflect" + "sync" + "sync/atomic" +) + +var ctxInterfaceType = reflect.TypeFor[context.Context]() +var errInterfaceType = reflect.TypeFor[error]() + +type eventListener struct { + id uint64 + fnVal reflect.Value + numIn int + hasCtx bool + hasPayload bool + argType reflect.Type + returnsErr bool +} + +// EventBus is a thread-safe, strongly-typed in-process domain event bus. +type EventBus struct { + mu sync.RWMutex + nextID atomic.Uint64 + handlers map[string][]eventListener +} + +// NewEventBus creates a new EventBus instance. +func NewEventBus() *EventBus { + return &EventBus{ + handlers: make(map[string][]eventListener), + } +} + +// On registers an event handler for the given topic. +// +// Supported handler signatures: +// - func(ctx context.Context, event T) error +// - func(ctx context.Context, event T) +// - func(event T) error +// - func(event T) +// - func(ctx context.Context) error +// - func(ctx context.Context) +// - func() error +// - func() +// +// Returns a Disposer function that unregisters the handler when called. +func (b *EventBus) On(topic string, handler any) Disposer { + if handler == nil { + panic("core/events: handler cannot be nil") + } + + fnVal := reflect.ValueOf(handler) + fnType := fnVal.Type() + + if fnType.Kind() != reflect.Func { + panic(fmt.Sprintf("core/events: expected func, got %s", fnType.Kind())) + } + + numIn := fnType.NumIn() + if numIn > 2 { + panic(fmt.Sprintf("core/events: handler has %d parameters, maximum 2 supported (ctx, event)", numIn)) + } + + numOut := fnType.NumOut() + if numOut > 1 { + panic(fmt.Sprintf("core/events: handler has %d return values, maximum 1 supported (error)", numOut)) + } + + returnsErr := false + if numOut == 1 { + outType := fnType.Out(0) + if !outType.Implements(errInterfaceType) { + panic(fmt.Sprintf("core/events: handler return type must be error, got %v", outType)) + } + returnsErr = true + } + + listener := eventListener{ + id: b.nextID.Add(1), + fnVal: fnVal, + numIn: numIn, + returnsErr: returnsErr, + } + + switch numIn { + case 0: + // func() or func() error + case 1: + in0 := fnType.In(0) + if in0.Implements(ctxInterfaceType) { + listener.hasCtx = true + } else { + listener.hasPayload = true + listener.argType = in0 + } + case 2: + in0 := fnType.In(0) + if !in0.Implements(ctxInterfaceType) { + panic(fmt.Sprintf("core/events: first parameter must implement context.Context, got %v", in0)) + } + listener.hasCtx = true + listener.hasPayload = true + listener.argType = fnType.In(1) + } + + b.mu.Lock() + b.handlers[topic] = append(b.handlers[topic], listener) + b.mu.Unlock() + + listenerID := listener.id + var disposed atomic.Bool + + return func() error { + if disposed.Swap(true) { + return nil + } + + b.mu.Lock() + defer b.mu.Unlock() + + list := b.handlers[topic] + for i, l := range list { + if l.id == listenerID { + b.handlers[topic] = append(list[:i], list[i+1:]...) + break + } + } + + if len(b.handlers[topic]) == 0 { + delete(b.handlers, topic) + } + + return nil + } +} + +// Subscribe registers a strongly-typed generic event listener on the given EventBus. +func Subscribe[T any](bus *EventBus, topic string, handler func(ctx context.Context, event T) error) Disposer { + if bus == nil { + panic("core/events: nil EventBus provided to Subscribe") + } + return bus.On(topic, handler) +} + +// Emit publishes an event to all subscribers of the specified topic. +// Handlers are executed synchronously. If any handler panics or returns an error, +// the error is collected and returned via errors.Join. +func (b *EventBus) Emit(ctx context.Context, topic string, payload any) error { + if ctx == nil { + ctx = context.Background() + } + + b.mu.RLock() + rawListeners := b.handlers[topic] + if len(rawListeners) == 0 { + b.mu.RUnlock() + return nil + } + + listeners := make([]eventListener, len(rawListeners)) + copy(listeners, rawListeners) + b.mu.RUnlock() + + var payloadVal reflect.Value + if payload != nil { + payloadVal = reflect.ValueOf(payload) + } + + var errs []error + for _, l := range listeners { + args := b.buildArgs(ctx, l, payloadVal) + + err := func() (resErr error) { + defer func() { + if r := recover(); r != nil { + resErr = fmt.Errorf("core/events: panic in handler for topic %q: %v", topic, r) + } + }() + + results := l.fnVal.Call(args) + if l.returnsErr && len(results) > 0 && !results[0].IsNil() { + resErr = results[0].Interface().(error) + } + return resErr + }() + + if err != nil { + errs = append(errs, err) + } + } + + return errors.Join(errs...) +} + +func (b *EventBus) buildArgs(ctx context.Context, l eventListener, payloadVal reflect.Value) []reflect.Value { + if l.numIn == 0 { + return nil + } + + args := make([]reflect.Value, 0, l.numIn) + if l.hasCtx { + args = append(args, reflect.ValueOf(ctx)) + } + + if l.hasPayload { + arg := b.convertPayload(payloadVal, l.argType) + args = append(args, arg) + } + + return args +} + +func (b *EventBus) convertPayload(payloadVal reflect.Value, targetType reflect.Type) reflect.Value { + if !payloadVal.IsValid() { + return reflect.Zero(targetType) + } + + valType := payloadVal.Type() + + // 1. Direct assignable + if valType.AssignableTo(targetType) { + return payloadVal + } + + // 2. Direct convertible + if valType.ConvertibleTo(targetType) { + return payloadVal.Convert(targetType) + } + + // 3. Payload is pointer *T, target expects T + if valType.Kind() == reflect.Pointer && valType.Elem().AssignableTo(targetType) { + if !payloadVal.IsNil() { + return payloadVal.Elem() + } + return reflect.Zero(targetType) + } + + // 4. Payload is value T, target expects *T + if targetType.Kind() == reflect.Pointer && valType.AssignableTo(targetType.Elem()) { + ptr := reflect.New(valType) + ptr.Elem().Set(payloadVal) + return ptr + } + + // Fallback to zero value of targetType + return reflect.Zero(targetType) +} + +// Listeners returns the number of active listeners for a topic. +func (b *EventBus) Listeners(topic string) int { + b.mu.RLock() + defer b.mu.RUnlock() + return len(b.handlers[topic]) +} + +// Topics returns all topics that have registered listeners. +func (b *EventBus) Topics() []string { + b.mu.RLock() + defer b.mu.RUnlock() + + topics := make([]string, 0, len(b.handlers)) + for t := range b.handlers { + topics = append(topics, t) + } + return topics +} diff --git a/core/events_test.go b/core/events_test.go new file mode 100644 index 00000000..49352758 --- /dev/null +++ b/core/events_test.go @@ -0,0 +1,315 @@ +package core_test + +import ( + "context" + "errors" + "fmt" + "sync" + "sync/atomic" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/Rain-kl/Wavelet/core" +) + +type UserRegisteredEvent struct { + UserID string `json:"user_id"` + Username string `json:"username"` +} + +type OrderCreatedEvent struct { + OrderID string `json:"order_id"` + Amount float64 `json:"amount"` +} + +func TestEventBusPublishSubscribe(t *testing.T) { + bus := core.NewEventBus() + var receivedID string + + disposer := bus.On("user:registered", func(ctx context.Context, e UserRegisteredEvent) error { + receivedID = e.UserID + return nil + }) + require.NotNil(t, disposer) + + assert.Equal(t, []string{"user:registered"}, bus.Topics()) + + err := bus.Emit(context.Background(), "user:registered", UserRegisteredEvent{UserID: "u_999", Username: "alice"}) + assert.NoError(t, err) + assert.Equal(t, "u_999", receivedID) + + // Emit to empty topic returns nil error + err = bus.Emit(nil, "empty:topic", nil) + assert.NoError(t, err) +} + +func TestEventBusGenericSubscribe(t *testing.T) { + bus := core.NewEventBus() + var receivedOrder string + + assert.Panics(t, func() { + core.Subscribe[OrderCreatedEvent](nil, "order:created", func(ctx context.Context, e OrderCreatedEvent) error { + return nil + }) + }) + + disposer := core.Subscribe(bus, "order:created", func(ctx context.Context, e OrderCreatedEvent) error { + receivedOrder = e.OrderID + return nil + }) + require.NotNil(t, disposer) + + err := bus.Emit(context.Background(), "order:created", OrderCreatedEvent{OrderID: "ord_123", Amount: 99.5}) + assert.NoError(t, err) + assert.Equal(t, "ord_123", receivedOrder) + + // Test unsubscribe via disposer + err = disposer() + assert.NoError(t, err) + + receivedOrder = "" + err = bus.Emit(context.Background(), "order:created", OrderCreatedEvent{OrderID: "ord_456", Amount: 100}) + assert.NoError(t, err) + assert.Empty(t, receivedOrder, "handler should not be called after disposal") +} + +func TestEventBusHandlerSignatures(t *testing.T) { + bus := core.NewEventBus() + + var ( + calledWithCtxPayloadErr atomic.Bool + calledWithCtxPayload atomic.Bool + calledWithPayloadErr atomic.Bool + calledWithPayload atomic.Bool + calledWithCtxErr atomic.Bool + calledWithCtx atomic.Bool + calledWithNoArgsErr atomic.Bool + calledWithNoArgs atomic.Bool + ) + + bus.On("test:sig", func(ctx context.Context, e UserRegisteredEvent) error { + calledWithCtxPayloadErr.Store(true) + assert.Equal(t, "u_1", e.UserID) + return nil + }) + + bus.On("test:sig", func(ctx context.Context, e UserRegisteredEvent) { + calledWithCtxPayload.Store(true) + assert.Equal(t, "u_1", e.UserID) + }) + + bus.On("test:sig", func(e UserRegisteredEvent) error { + calledWithPayloadErr.Store(true) + assert.Equal(t, "u_1", e.UserID) + return nil + }) + + bus.On("test:sig", func(e UserRegisteredEvent) { + calledWithPayload.Store(true) + assert.Equal(t, "u_1", e.UserID) + }) + + bus.On("test:sig", func(ctx context.Context) error { + calledWithCtxErr.Store(true) + return nil + }) + + bus.On("test:sig", func(ctx context.Context) { + calledWithCtx.Store(true) + }) + + bus.On("test:sig", func() error { + calledWithNoArgsErr.Store(true) + return nil + }) + + bus.On("test:sig", func() { + calledWithNoArgs.Store(true) + }) + + err := bus.Emit(context.Background(), "test:sig", UserRegisteredEvent{UserID: "u_1", Username: "test"}) + assert.NoError(t, err) + + assert.True(t, calledWithCtxPayloadErr.Load()) + assert.True(t, calledWithCtxPayload.Load()) + assert.True(t, calledWithPayloadErr.Load()) + assert.True(t, calledWithPayload.Load()) + assert.True(t, calledWithCtxErr.Load()) + assert.True(t, calledWithCtx.Load()) + assert.True(t, calledWithNoArgsErr.Load()) + assert.True(t, calledWithNoArgs.Load()) +} + +func TestEventBusPointerAndValueConversion(t *testing.T) { + bus := core.NewEventBus() + + var ( + receivedFromValueToPtr atomic.Bool + receivedFromPtrToValue atomic.Bool + receivedFromNilPtr atomic.Bool + ) + + // Handler expects pointer, payload emitted as value + bus.On("test:ptr", func(ctx context.Context, e *UserRegisteredEvent) error { + if e != nil && e.UserID == "u_ptr" { + receivedFromValueToPtr.Store(true) + } + return nil + }) + + err := bus.Emit(context.Background(), "test:ptr", UserRegisteredEvent{UserID: "u_ptr"}) + assert.NoError(t, err) + assert.True(t, receivedFromValueToPtr.Load()) + + // Handler expects value, payload emitted as pointer + bus.On("test:val", func(ctx context.Context, e UserRegisteredEvent) error { + if e.UserID == "u_val" { + receivedFromPtrToValue.Store(true) + } + return nil + }) + + err = bus.Emit(context.Background(), "test:val", &UserRegisteredEvent{UserID: "u_val"}) + assert.NoError(t, err) + assert.True(t, receivedFromPtrToValue.Load()) + + // Handler expects value, payload is nil pointer + var nilEvent *UserRegisteredEvent + bus.On("test:nil_ptr", func(ctx context.Context, e UserRegisteredEvent) error { + assert.Equal(t, "", e.UserID) + receivedFromNilPtr.Store(true) + return nil + }) + err = bus.Emit(context.Background(), "test:nil_ptr", nilEvent) + assert.NoError(t, err) + assert.True(t, receivedFromNilPtr.Load()) + + // Convertible type test (int to int64) + var receivedConvert int64 + bus.On("test:conv", func(e int64) { + receivedConvert = e + }) + err = bus.Emit(context.Background(), "test:conv", int(42)) + assert.NoError(t, err) + assert.Equal(t, int64(42), receivedConvert) +} + +func TestEventBusErrorCollectionAndPanicRecovery(t *testing.T) { + bus := core.NewEventBus() + + errHandler1 := errors.New("handler 1 failed") + errHandler2 := errors.New("handler 2 failed") + + bus.On("test:err", func() error { + return errHandler1 + }) + + bus.On("test:err", func() { + panic("something went horribly wrong") + }) + + bus.On("test:err", func() error { + return errHandler2 + }) + + err := bus.Emit(context.Background(), "test:err", nil) + require.Error(t, err) + assert.True(t, errors.Is(err, errHandler1) || errors.Is(err, errHandler2)) + assert.Contains(t, err.Error(), "handler 1 failed") + assert.Contains(t, err.Error(), "handler 2 failed") + assert.Contains(t, err.Error(), "panic") +} + +func TestEventBusInvalidHandlerPanics(t *testing.T) { + bus := core.NewEventBus() + + assert.Panics(t, func() { + bus.On("test:invalid", nil) + }) + + assert.Panics(t, func() { + bus.On("test:invalid", "not-a-func") + }) + + assert.Panics(t, func() { + // More than 2 arguments + bus.On("test:invalid", func(a, b, c string) {}) + }) + + assert.Panics(t, func() { + // 2 args, but first is not context + bus.On("test:invalid", func(a string, b int) {}) + }) + + assert.Panics(t, func() { + // More than 1 return value + bus.On("test:invalid", func() (int, error) { return 0, nil }) + }) + + assert.Panics(t, func() { + // Return value is not error + bus.On("test:invalid", func() int { return 0 }) + }) +} + +func TestEventBusListenersCountAndDisposerIdempotence(t *testing.T) { + bus := core.NewEventBus() + + assert.Equal(t, 0, bus.Listeners("topic1")) + + d1 := bus.On("topic1", func() {}) + d2 := bus.On("topic1", func() {}) + assert.Equal(t, 2, bus.Listeners("topic1")) + + _ = d1() + assert.Equal(t, 1, bus.Listeners("topic1")) + + // Calling disposer again should be no-op + _ = d1() + assert.Equal(t, 1, bus.Listeners("topic1")) + + _ = d2() + assert.Equal(t, 0, bus.Listeners("topic1")) +} + +func TestEventBusConcurrentAccess(t *testing.T) { + bus := core.NewEventBus() + var wg sync.WaitGroup + + var receivedCount atomic.Int64 + + // Concurrently subscribe and emit + for i := 0; i < 50; i++ { + wg.Add(1) + go func(idx int) { + defer wg.Done() + topic := fmt.Sprintf("topic:%d", idx%5) + disposer := bus.On(topic, func(ctx context.Context, e UserRegisteredEvent) error { + receivedCount.Add(1) + return nil + }) + + // Emit some events + _ = bus.Emit(context.Background(), topic, UserRegisteredEvent{UserID: fmt.Sprintf("u_%d", idx)}) + + // Randomly dispose + if idx%2 == 0 { + _ = disposer() + } + }(i) + } + + for i := 0; i < 50; i++ { + wg.Add(1) + go func(idx int) { + defer wg.Done() + topic := fmt.Sprintf("topic:%d", idx%5) + _ = bus.Emit(context.Background(), topic, UserRegisteredEvent{UserID: fmt.Sprintf("u_%d", idx)}) + }(i) + } + + wg.Wait() + assert.Greater(t, receivedCount.Load(), int64(0)) +} diff --git a/core/extpoints/extpoints_test.go b/core/extpoints/extpoints_test.go new file mode 100644 index 00000000..dacf4f8d --- /dev/null +++ b/core/extpoints/extpoints_test.go @@ -0,0 +1,277 @@ +package extpoints_test + +import ( + "context" + "testing" + "testing/fstest" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/Rain-kl/Wavelet/core" + "github.com/Rain-kl/Wavelet/core/extpoints" +) + +func TestRouterExtension(t *testing.T) { + r := extpoints.NewRouterRegistry() + require.NotNil(t, r) + + mGlobal := "global_middleware" + r.Use(mGlobal) + assert.Equal(t, []any{mGlobal}, r.Middlewares()) + + // Test root methods + hRoot := "root_handler" + r.GET("/", hRoot) + r.POST("/root_post", hRoot) + r.PUT("/root_put", hRoot) + r.DELETE("/root_del", hRoot) + r.PATCH("/root_patch", hRoot) + r.HEAD("/root_head", hRoot) + r.OPTIONS("/root_opt", hRoot) + anyRootDefs := r.Any("/root_any", hRoot) + assert.Len(t, anyRootDefs, 7) + + // Group and Group.Use + mAPI := "api_middleware" + api := r.Group("/api/v1", mAPI) + api.Use("api_extra_middleware") + assert.Len(t, api.Middlewares(), 2) + + hList := "list_orders_handler" + hCreate := "create_order_handler" + api.GET("/orders", hList) + api.POST("/orders", hCreate) + + mAdmin := "admin_middleware" + admin := api.Group("admin", mAdmin) + + hUserGet := "get_user_handler" + hUserPut := "put_user_handler" + hUserDel := "del_user_handler" + hUserPatch := "patch_user_handler" + hUserHead := "head_user_handler" + hUserOptions := "options_user_handler" + admin.GET("/users/:id", hUserGet) + admin.PUT("/users/:id", hUserPut) + admin.DELETE("/users/:id", hUserDel) + admin.PATCH("/users/:id", hUserPatch) + admin.HEAD("/users/:id", hUserHead) + admin.OPTIONS("/users/:id", hUserOptions) + + hCustom := "custom_handler" + admin.Handle("CUSTOM", "/custom", hCustom) + + hAny := "any_handler" + anyRoutes := admin.Any("/all", hAny) + assert.NotEmpty(t, anyRoutes) + + // Group.Routes() returns root routes + assert.Equal(t, r.Routes(), admin.Routes()) + + routes := r.Routes() + + // Verify route paths and middlewares + var foundOrderGet bool + var foundUserPut bool + for _, route := range routes { + if route.Method == "GET" && route.Path == "/api/v1/orders" { + foundOrderGet = true + assert.Equal(t, []any{mGlobal, mAPI, "api_extra_middleware"}, route.Middlewares) + assert.Equal(t, []any{hList}, route.Handlers) + } + if route.Method == "PUT" && route.Path == "/api/v1/admin/users/:id" { + foundUserPut = true + assert.Equal(t, []any{mGlobal, mAPI, "api_extra_middleware", mAdmin}, route.Middlewares) + assert.Equal(t, []any{hUserPut}, route.Handlers) + } + } + assert.True(t, foundOrderGet) + assert.True(t, foundUserPut) +} + +func TestMigrationExtension(t *testing.T) { + m := extpoints.NewMigrationRegistry() + require.NotNil(t, m) + + fs1 := fstest.MapFS{ + "migrations/001_init.sql": &fstest.MapFile{Data: []byte("CREATE TABLE t1(id int);")}, + } + fs2 := fstest.MapFS{ + "custom/001_order.sql": &fstest.MapFile{Data: []byte("CREATE TABLE t2(id int);")}, + } + + m.Register("auth", fs1) + m.Register("order", fs2, "custom") + + // Update existing entry + fs1Updated := fstest.MapFS{ + "migrations/002_update.sql": &fstest.MapFile{Data: []byte("ALTER TABLE t1 ADD col int;")}, + } + m.Register("auth", fs1Updated, "") + + entries := m.Entries() + require.Len(t, entries, 2) + assert.Equal(t, "auth", entries[0].PluginID) + assert.Equal(t, "migrations", entries[0].Dir) + assert.Equal(t, "order", entries[1].PluginID) + assert.Equal(t, "custom", entries[1].Dir) + + authEntry, ok := m.Get("auth") + assert.True(t, ok) + assert.Equal(t, "auth", authEntry.PluginID) + + _, ok = m.Get("non_existent") + assert.False(t, ok) +} + +func TestTaskExtension(t *testing.T) { + tr := extpoints.NewTaskRegistry() + require.NotNil(t, tr) + + handler := func(ctx context.Context, payload []byte) error { return nil } + + tr.Register("order:cancel_timeout", handler, + extpoints.WithTaskConcurrency(5), + extpoints.WithTaskRetry(3), + extpoints.WithTaskTimeout(10*time.Second), + extpoints.WithTaskMetadata("queue", "critical"), + nil, // test nil option + ) + + // Re-register to test update + tr.Register("order:cancel_timeout", handler, + extpoints.WithTaskConcurrency(10), + extpoints.WithTaskMetadata("queue", "high"), + ) + + tasks := tr.Tasks() + require.Len(t, tasks, 1) + assert.Equal(t, "order:cancel_timeout", tasks[0].Pattern) + assert.Equal(t, 10, tasks[0].Concurrency) + assert.Equal(t, "high", tasks[0].Metadata["queue"]) + + task, ok := tr.Get("order:cancel_timeout") + assert.True(t, ok) + assert.Equal(t, "order:cancel_timeout", task.Pattern) + + _, ok = tr.Get("unknown") + assert.False(t, ok) +} + +func TestScheduleExtension(t *testing.T) { + sr := extpoints.NewScheduleRegistry() + require.NotNil(t, sr) + + type ReportPayload struct { + Type string `json:"type"` + } + + sr.RegisterCron("0 2 * * *", "report:daily_summary", ReportPayload{Type: "daily"}) + sr.Register("@every 1h", "cleanup:expired_sessions", nil, + extpoints.WithScheduleOption("retry", 2), + nil, // test nil option + ) + + // Re-register to test update + sr.RegisterCron("0 3 * * *", "report:daily_summary", ReportPayload{Type: "all"}) + + schedules := sr.Schedules() + require.Len(t, schedules, 2) + + assert.Equal(t, "0 3 * * *", schedules[0].Spec) + assert.Equal(t, "report:daily_summary", schedules[0].TaskType) + assert.Equal(t, ReportPayload{Type: "all"}, schedules[0].Payload) + + assert.Equal(t, "@every 1h", schedules[1].Spec) + assert.Equal(t, "cleanup:expired_sessions", schedules[1].TaskType) + assert.Equal(t, 2, schedules[1].Options["retry"]) + + sched, ok := sr.Get("report:daily_summary") + assert.True(t, ok) + assert.Equal(t, "0 3 * * *", sched.Spec) + + _, ok = sr.Get("unknown") + assert.False(t, ok) +} + +func TestSettingExtension(t *testing.T) { + sr := extpoints.NewSettingRegistry() + require.NotNil(t, sr) + + assert.Panics(t, func() { + sr.Register(extpoints.SettingSchema{}) // empty key panics + }) + + sr.Register(extpoints.SettingSchema{ + Key: "order.auto_cancel_mins", + Default: 15, + Description: "Order auto cancellation timeout in minutes", + Category: "order", + Public: true, + }) + + // Re-register to test update + sr.Register(extpoints.SettingSchema{ + Key: "order.auto_cancel_mins", + Default: 30, + Description: "Updated timeout", + }) + + sr.Register(extpoints.SettingSchema{ + Key: "auth.jwt_secret", + Default: "default-secret", + Description: "JWT secret key", + Category: "auth", + ReadOnly: true, + }) + + schemas := sr.Schemas() + require.Len(t, schemas, 2) + + schema, ok := sr.Get("order.auto_cancel_mins") + assert.True(t, ok) + assert.Equal(t, 30, schema.Default) + + _, ok = sr.Get("unknown") + assert.False(t, ok) +} + +func TestContextExtensionPointsIntegration(t *testing.T) { + ctx := core.NewContext(context.Background()) + + require.NotNil(t, ctx.Events()) + require.NotNil(t, ctx.Router()) + require.NotNil(t, ctx.Migrations()) + require.NotNil(t, ctx.Tasks()) + require.NotNil(t, ctx.Task()) + require.NotNil(t, ctx.Schedules()) + require.NotNil(t, ctx.Schedule()) + require.NotNil(t, ctx.Settings()) + require.NotNil(t, ctx.Setting()) + + // Register from child context and verify shared application registry + child := ctx.Fork() + child.Router().GET("/ping", "pong_handler") + child.Task().Register("sample:task", "handler") + child.Schedule().RegisterCron("@hourly", "sample:cron", nil) + child.Settings().Register(extpoints.SettingSchema{ + Key: "app.name", + Default: "Wavelet", + }) + + assert.Len(t, ctx.Router().Routes(), 1) + assert.Len(t, ctx.Tasks().Tasks(), 1) + assert.Len(t, ctx.Schedules().Schedules(), 1) + assert.Len(t, ctx.Settings().Schemas(), 1) + + // Child and root events + var eventReceived bool + child.Events().On("app:ready", func() { + eventReceived = true + }) + err := ctx.Events().Emit(context.Background(), "app:ready", nil) + assert.NoError(t, err) + assert.True(t, eventReceived) +} diff --git a/core/extpoints/migration.go b/core/extpoints/migration.go new file mode 100644 index 00000000..322a6a34 --- /dev/null +++ b/core/extpoints/migration.go @@ -0,0 +1,83 @@ +package extpoints + +import ( + "io/fs" + "sync" +) + +// MigrationEntry contains the migration filesystem and configuration for a plugin. +type MigrationEntry struct { + PluginID string + FS fs.FS + Dir string +} + +// MigrationExtension defines the interface for registering and querying plugin migrations. +type MigrationExtension interface { + Register(pluginID string, fsys fs.FS, dir ...string) + Entries() []MigrationEntry + Get(pluginID string) (MigrationEntry, bool) +} + +// MigrationRegistry collects and stores migration entries from plugins. +type MigrationRegistry struct { + mu sync.RWMutex + entries []MigrationEntry + lookup map[string]MigrationEntry +} + +// NewMigrationRegistry creates a new migration registry. +func NewMigrationRegistry() *MigrationRegistry { + return &MigrationRegistry{ + lookup: make(map[string]MigrationEntry), + } +} + +// Register adds a migration entry for a plugin. +// If dir is not specified, it defaults to "migrations". +func (m *MigrationRegistry) Register(pluginID string, fsys fs.FS, dir ...string) { + m.mu.Lock() + defer m.mu.Unlock() + + migrationDir := "migrations" + if len(dir) > 0 && dir[0] != "" { + migrationDir = dir[0] + } + + entry := MigrationEntry{ + PluginID: pluginID, + FS: fsys, + Dir: migrationDir, + } + + // If entry already exists, update in-place; otherwise append + if _, exists := m.lookup[pluginID]; exists { + for i, e := range m.entries { + if e.PluginID == pluginID { + m.entries[i] = entry + break + } + } + } else { + m.entries = append(m.entries, entry) + } + + m.lookup[pluginID] = entry +} + +// Entries returns a copy of all registered migration entries in registration order. +func (m *MigrationRegistry) Entries() []MigrationEntry { + m.mu.RLock() + defer m.mu.RUnlock() + res := make([]MigrationEntry, len(m.entries)) + copy(res, m.entries) + return res +} + +// Get retrieves the migration entry for a specific plugin ID. +func (m *MigrationRegistry) Get(pluginID string) (MigrationEntry, bool) { + m.mu.RLock() + defer m.mu.RUnlock() + e, ok := m.lookup[pluginID] + return e, ok +} diff --git a/core/extpoints/router.go b/core/extpoints/router.go new file mode 100644 index 00000000..9905c288 --- /dev/null +++ b/core/extpoints/router.go @@ -0,0 +1,266 @@ +package extpoints + +import ( + "strings" + "sync" +) + +// RouteDefinition holds the metadata and handler list for a single HTTP route. +type RouteDefinition struct { + Method string + Path string + Handlers []any + Middlewares []any +} + +// RouterExtension defines the interface for registering routes and middlewares. +type RouterExtension interface { + Use(middlewares ...any) + Group(prefix string, middlewares ...any) RouterExtension + Handle(method string, path string, handlers ...any) RouteDefinition + GET(path string, handlers ...any) RouteDefinition + POST(path string, handlers ...any) RouteDefinition + PUT(path string, handlers ...any) RouteDefinition + DELETE(path string, handlers ...any) RouteDefinition + PATCH(path string, handlers ...any) RouteDefinition + HEAD(path string, handlers ...any) RouteDefinition + OPTIONS(path string, handlers ...any) RouteDefinition + Any(path string, handlers ...any) []RouteDefinition + Routes() []RouteDefinition + Middlewares() []any +} + +// RouterRegistry implements RouterExtension as the root route and middleware collector. +type RouterRegistry struct { + mu sync.RWMutex + routes []RouteDefinition + middlewares []any +} + +// NewRouterRegistry creates a new root router collector. +func NewRouterRegistry() *RouterRegistry { + return &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...) +} + +// Middlewares returns a copy of registered root middlewares. +func (r *RouterRegistry) Middlewares() []any { + r.mu.RLock() + defer r.mu.RUnlock() + res := make([]any, len(r.middlewares)) + copy(res, r.middlewares) + return res +} + +// Group creates a new RouteGroup under the router. +func (r *RouterRegistry) Group(prefix string, middlewares ...any) RouterExtension { + return &RouterGroup{ + registry: r, + prefix: cleanPath(prefix), + middlewares: middlewares, + } +} + +// Handle registers a route with a custom HTTP method and handlers. +func (r *RouterRegistry) Handle(method, path string, handlers ...any) RouteDefinition { + r.mu.Lock() + defer r.mu.Unlock() + + rd := RouteDefinition{ + Method: strings.ToUpper(method), + Path: cleanPath(path), + Handlers: handlers, + Middlewares: append([]any(nil), r.middlewares...), + } + r.routes = append(r.routes, rd) + return rd +} + +// GET registers a GET route. +func (r *RouterRegistry) GET(path string, handlers ...any) RouteDefinition { + return r.Handle("GET", path, handlers...) +} + +// POST registers a POST route. +func (r *RouterRegistry) POST(path string, handlers ...any) RouteDefinition { + return r.Handle("POST", path, handlers...) +} + +// PUT registers a PUT route. +func (r *RouterRegistry) PUT(path string, handlers ...any) RouteDefinition { + return r.Handle("PUT", path, handlers...) +} + +// DELETE registers a DELETE route. +func (r *RouterRegistry) DELETE(path string, handlers ...any) RouteDefinition { + return r.Handle("DELETE", path, handlers...) +} + +// PATCH registers a PATCH route. +func (r *RouterRegistry) PATCH(path string, handlers ...any) RouteDefinition { + return r.Handle("PATCH", path, handlers...) +} + +// HEAD registers a HEAD route. +func (r *RouterRegistry) HEAD(path string, handlers ...any) RouteDefinition { + return r.Handle("HEAD", path, handlers...) +} + +// OPTIONS registers an OPTIONS route. +func (r *RouterRegistry) OPTIONS(path string, handlers ...any) RouteDefinition { + return r.Handle("OPTIONS", path, handlers...) +} + +// Any registers a route for standard HTTP methods. +func (r *RouterRegistry) Any(path string, handlers ...any) []RouteDefinition { + methods := []string{"GET", "POST", "PUT", "DELETE", "PATCH", "HEAD", "OPTIONS"} + defs := make([]RouteDefinition, 0, len(methods)) + for _, m := range methods { + defs = append(defs, r.Handle(m, path, handlers...)) + } + return defs +} + +// Routes returns a copy of all collected RouteDefinitions. +func (r *RouterRegistry) Routes() []RouteDefinition { + r.mu.RLock() + defer r.mu.RUnlock() + res := make([]RouteDefinition, len(r.routes)) + copy(res, r.routes) + return res +} + +// RouterGroup represents a scoped route group with a path prefix and group-level middlewares. +type RouterGroup struct { + registry *RouterRegistry + prefix string + middlewares []any +} + +// Use adds middlewares to this group. +func (g *RouterGroup) Use(middlewares ...any) { + g.middlewares = append(g.middlewares, middlewares...) +} + +// Group creates a nested RouteGroup. +func (g *RouterGroup) Group(prefix string, middlewares ...any) RouterExtension { + combinedPrefix := joinPaths(g.prefix, prefix) + combinedMiddlewares := make([]any, 0, len(g.middlewares)+len(middlewares)) + combinedMiddlewares = append(combinedMiddlewares, g.middlewares...) + combinedMiddlewares = append(combinedMiddlewares, middlewares...) + + return &RouterGroup{ + registry: g.registry, + prefix: combinedPrefix, + middlewares: combinedMiddlewares, + } +} + +// Handle registers a route under this group. +func (g *RouterGroup) Handle(method, path string, handlers ...any) RouteDefinition { + g.registry.mu.Lock() + defer g.registry.mu.Unlock() + + fullPath := joinPaths(g.prefix, path) + + allMiddlewares := make([]any, 0, len(g.registry.middlewares)+len(g.middlewares)) + allMiddlewares = append(allMiddlewares, g.registry.middlewares...) + allMiddlewares = append(allMiddlewares, g.middlewares...) + + rd := RouteDefinition{ + Method: strings.ToUpper(method), + Path: fullPath, + Handlers: handlers, + Middlewares: allMiddlewares, + } + g.registry.routes = append(g.registry.routes, rd) + return rd +} + +// GET registers a GET route in this group. +func (g *RouterGroup) GET(path string, handlers ...any) RouteDefinition { + return g.Handle("GET", path, handlers...) +} + +// POST registers a POST route in this group. +func (g *RouterGroup) POST(path string, handlers ...any) RouteDefinition { + return g.Handle("POST", path, handlers...) +} + +// PUT registers a PUT route in this group. +func (g *RouterGroup) PUT(path string, handlers ...any) RouteDefinition { + return g.Handle("PUT", path, handlers...) +} + +// DELETE registers a DELETE route in this group. +func (g *RouterGroup) DELETE(path string, handlers ...any) RouteDefinition { + return g.Handle("DELETE", path, handlers...) +} + +// PATCH registers a PATCH route in this group. +func (g *RouterGroup) PATCH(path string, handlers ...any) RouteDefinition { + return g.Handle("PATCH", path, handlers...) +} + +// HEAD registers a HEAD route in this group. +func (g *RouterGroup) HEAD(path string, handlers ...any) RouteDefinition { + return g.Handle("HEAD", path, handlers...) +} + +// OPTIONS registers an OPTIONS route in this group. +func (g *RouterGroup) OPTIONS(path string, handlers ...any) RouteDefinition { + return g.Handle("OPTIONS", path, handlers...) +} + +// Any registers a route in this group for standard HTTP methods. +func (g *RouterGroup) Any(path string, handlers ...any) []RouteDefinition { + methods := []string{"GET", "POST", "PUT", "DELETE", "PATCH", "HEAD", "OPTIONS"} + defs := make([]RouteDefinition, 0, len(methods)) + for _, m := range methods { + defs = append(defs, g.Handle(m, path, handlers...)) + } + return defs +} + +// Routes returns all routes from the parent registry. +func (g *RouterGroup) Routes() []RouteDefinition { + return g.registry.Routes() +} + +// Middlewares returns a copy of the group's middlewares. +func (g *RouterGroup) Middlewares() []any { + res := make([]any, len(g.middlewares)) + copy(res, g.middlewares) + return res +} + +func cleanPath(p string) string { + if p == "" { + return "/" + } + if !strings.HasPrefix(p, "/") { + p = "/" + p + } + if len(p) > 1 && strings.HasSuffix(p, "/") { + p = strings.TrimSuffix(p, "/") + } + return p +} + +func joinPaths(base, relative string) string { + if base == "" || base == "/" { + return cleanPath(relative) + } + if relative == "" || relative == "/" { + return cleanPath(base) + } + base = strings.TrimSuffix(base, "/") + relative = strings.TrimPrefix(relative, "/") + return cleanPath(base + "/" + relative) +} diff --git a/core/extpoints/schedule.go b/core/extpoints/schedule.go new file mode 100644 index 00000000..48aebc14 --- /dev/null +++ b/core/extpoints/schedule.go @@ -0,0 +1,100 @@ +package extpoints + +import "sync" + +// ScheduleDefinition holds the configuration for a scheduled/cron task. +type ScheduleDefinition struct { + Spec string + TaskType string + Payload any + Options map[string]any +} + +// ScheduleOption configures a ScheduleDefinition. +type ScheduleOption func(*ScheduleDefinition) + +// WithScheduleOption adds a custom option to the schedule definition. +func WithScheduleOption(key string, val any) ScheduleOption { + return func(sd *ScheduleDefinition) { + if sd.Options == nil { + sd.Options = make(map[string]any) + } + sd.Options[key] = val + } +} + +// ScheduleExtension defines the interface for registering and querying cron/scheduled tasks. +type ScheduleExtension interface { + Register(spec string, taskType string, payload any, opts ...ScheduleOption) + RegisterCron(spec string, taskType string, payload any, opts ...ScheduleOption) + Schedules() []ScheduleDefinition + Get(taskType string) (ScheduleDefinition, bool) +} + +// ScheduleRegistry collects and manages schedule registrations. +type ScheduleRegistry struct { + mu sync.RWMutex + schedules []ScheduleDefinition + lookup map[string]ScheduleDefinition +} + +// NewScheduleRegistry creates a new schedule registry. +func NewScheduleRegistry() *ScheduleRegistry { + return &ScheduleRegistry{ + lookup: make(map[string]ScheduleDefinition), + } +} + +// Register adds a schedule definition. +func (s *ScheduleRegistry) Register(spec string, taskType string, payload any, opts ...ScheduleOption) { + s.mu.Lock() + defer s.mu.Unlock() + + sd := ScheduleDefinition{ + Spec: spec, + TaskType: taskType, + Payload: payload, + Options: make(map[string]any), + } + + for _, opt := range opts { + if opt != nil { + opt(&sd) + } + } + + if _, exists := s.lookup[taskType]; exists { + for i, item := range s.schedules { + if item.TaskType == taskType { + s.schedules[i] = sd + break + } + } + } else { + s.schedules = append(s.schedules, sd) + } + + s.lookup[taskType] = sd +} + +// RegisterCron is an alias for Register. +func (s *ScheduleRegistry) RegisterCron(spec string, taskType string, payload any, opts ...ScheduleOption) { + s.Register(spec, taskType, payload, opts...) +} + +// Schedules returns a copy of all registered ScheduleDefinitions. +func (s *ScheduleRegistry) Schedules() []ScheduleDefinition { + s.mu.RLock() + defer s.mu.RUnlock() + res := make([]ScheduleDefinition, len(s.schedules)) + copy(res, s.schedules) + return res +} + +// Get retrieves a schedule definition by its task type. +func (s *ScheduleRegistry) Get(taskType string) (ScheduleDefinition, bool) { + s.mu.RLock() + defer s.mu.RUnlock() + sd, ok := s.lookup[taskType] + return sd, ok +} diff --git a/core/extpoints/setting.go b/core/extpoints/setting.go new file mode 100644 index 00000000..4eadd536 --- /dev/null +++ b/core/extpoints/setting.go @@ -0,0 +1,77 @@ +package extpoints + +import "sync" + +// SettingSchema defines the configuration schema and metadata for a system or plugin setting. +type SettingSchema struct { + Key string `json:"key"` + Default any `json:"default"` + Description string `json:"description"` + Type string `json:"type,omitempty"` + ReadOnly bool `json:"read_only,omitempty"` + Public bool `json:"public,omitempty"` + Category string `json:"category,omitempty"` + Validation string `json:"validation,omitempty"` +} + +// SettingExtension defines the interface for registering and querying setting configuration schemas. +type SettingExtension interface { + Register(schema SettingSchema) + Schemas() []SettingSchema + Get(key string) (SettingSchema, bool) +} + +// SettingRegistry collects and manages setting configuration schemas. +type SettingRegistry struct { + mu sync.RWMutex + schemas []SettingSchema + lookup map[string]SettingSchema +} + +// NewSettingRegistry creates a new setting schema registry. +func NewSettingRegistry() *SettingRegistry { + return &SettingRegistry{ + lookup: make(map[string]SettingSchema), + } +} + +// Register registers a SettingSchema into the registry. +// Panics if the schema Key is empty. +func (s *SettingRegistry) Register(schema SettingSchema) { + if schema.Key == "" { + panic("core/extpoints: setting schema key cannot be empty") + } + + s.mu.Lock() + defer s.mu.Unlock() + + if _, exists := s.lookup[schema.Key]; exists { + for i, item := range s.schemas { + if item.Key == schema.Key { + s.schemas[i] = schema + break + } + } + } else { + s.schemas = append(s.schemas, schema) + } + + s.lookup[schema.Key] = schema +} + +// Schemas returns a copy of all registered SettingSchemas. +func (s *SettingRegistry) Schemas() []SettingSchema { + s.mu.RLock() + defer s.mu.RUnlock() + res := make([]SettingSchema, len(s.schemas)) + copy(res, s.schemas) + return res +} + +// Get retrieves a SettingSchema by its key. +func (s *SettingRegistry) Get(key string) (SettingSchema, bool) { + s.mu.RLock() + defer s.mu.RUnlock() + schema, ok := s.lookup[key] + return schema, ok +} diff --git a/core/extpoints/task.go b/core/extpoints/task.go new file mode 100644 index 00000000..f99e56d9 --- /dev/null +++ b/core/extpoints/task.go @@ -0,0 +1,119 @@ +package extpoints + +import ( + "sync" + "time" +) + +// TaskDefinition holds the definition and runtime options for an asynchronous background task. +type TaskDefinition struct { + Pattern string + Handler any + Concurrency int + Retry int + Timeout time.Duration + Metadata map[string]any +} + +// TaskOption configures a TaskDefinition. +type TaskOption func(*TaskDefinition) + +// WithTaskConcurrency sets the concurrency limit for the task. +func WithTaskConcurrency(concurrency int) TaskOption { + return func(td *TaskDefinition) { + td.Concurrency = concurrency + } +} + +// WithTaskRetry sets the maximum retry count for the task. +func WithTaskRetry(retry int) TaskOption { + return func(td *TaskDefinition) { + td.Retry = retry + } +} + +// WithTaskTimeout sets the execution timeout for the task. +func WithTaskTimeout(timeout time.Duration) TaskOption { + return func(td *TaskDefinition) { + td.Timeout = timeout + } +} + +// WithTaskMetadata adds a key-value pair to the task metadata. +func WithTaskMetadata(key string, val any) TaskOption { + return func(td *TaskDefinition) { + if td.Metadata == nil { + td.Metadata = make(map[string]any) + } + td.Metadata[key] = val + } +} + +// TaskExtension defines the interface for registering and querying background task handlers. +type TaskExtension interface { + Register(pattern string, handler any, opts ...TaskOption) + Tasks() []TaskDefinition + Get(pattern string) (TaskDefinition, bool) +} + +// TaskRegistry collects and manages task registrations. +type TaskRegistry struct { + mu sync.RWMutex + tasks []TaskDefinition + lookup map[string]TaskDefinition +} + +// NewTaskRegistry creates a new task registry. +func NewTaskRegistry() *TaskRegistry { + return &TaskRegistry{ + lookup: make(map[string]TaskDefinition), + } +} + +// Register registers a task pattern and its handler with optional configuration. +func (t *TaskRegistry) Register(pattern string, handler any, opts ...TaskOption) { + t.mu.Lock() + defer t.mu.Unlock() + + td := TaskDefinition{ + Pattern: pattern, + Handler: handler, + Metadata: make(map[string]any), + } + + for _, opt := range opts { + if opt != nil { + opt(&td) + } + } + + if _, exists := t.lookup[pattern]; exists { + for i, item := range t.tasks { + if item.Pattern == pattern { + t.tasks[i] = td + break + } + } + } else { + t.tasks = append(t.tasks, td) + } + + t.lookup[pattern] = td +} + +// Tasks returns a copy of all registered TaskDefinitions. +func (t *TaskRegistry) Tasks() []TaskDefinition { + t.mu.RLock() + defer t.mu.RUnlock() + res := make([]TaskDefinition, len(t.tasks)) + copy(res, t.tasks) + return res +} + +// Get retrieves a task definition by its pattern. +func (t *TaskRegistry) Get(pattern string) (TaskDefinition, bool) { + t.mu.RLock() + defer t.mu.RUnlock() + td, ok := t.lookup[pattern] + return td, ok +} diff --git a/core/types.go b/core/types.go index 2014a39b..33beb287 100644 --- a/core/types.go +++ b/core/types.go @@ -3,6 +3,8 @@ package core import ( "context" "errors" + + "github.com/Rain-kl/Wavelet/core/extpoints" ) // Standard sentinel errors returned by core operations. @@ -69,3 +71,19 @@ type Driver interface { // Disposer is a cleanup function executed when a Context is disposed. type Disposer func() error + +// Re-exported extension point types from core/extpoints for convenient usage. +type ( + RouterExtension = extpoints.RouterExtension + RouteDefinition = extpoints.RouteDefinition + MigrationExtension = extpoints.MigrationExtension + MigrationEntry = extpoints.MigrationEntry + TaskExtension = extpoints.TaskExtension + TaskDefinition = extpoints.TaskDefinition + TaskOption = extpoints.TaskOption + ScheduleExtension = extpoints.ScheduleExtension + ScheduleDefinition = extpoints.ScheduleDefinition + ScheduleOption = extpoints.ScheduleOption + SettingExtension = extpoints.SettingExtension + SettingSchema = extpoints.SettingSchema +)