diff --git a/Makefile b/Makefile index 740104b4..a82de00b 100644 --- a/Makefile +++ b/Makefile @@ -36,10 +36,32 @@ build-embedded: code-check: @echo "==> Architecture guards..." @command -v rg >/dev/null 2>&1 || { echo 'error: rg (ripgrep) is required for architecture guards' >&2; exit 1; } - @if [ -d pkg/model ] && rg -n 'database\.DB\(|cachepkg\.Redis' pkg/model --glob '*.go' -g '!*_test.go' ; then \ - echo 'error: pkg/model must not access database.DB or cachepkg.Redis (non-test code)' >&2; \ + @echo " → core/ must not import gin, gorm, asynq..." + @if rg -n '"github.com/gin-gonic/gin|"gorm.io/gorm|"github.com/hibiken/asynq' core/ --glob '*.go' -g '!*_test.go' 2>/dev/null; then \ + echo 'error: core/ must not import gin, gorm, or asynq' >&2; \ exit 1; \ fi + @echo " → core/contracts/ must not import plugins/..." + @if rg -n 'plugins/' core/contracts --glob '*.go' 2>/dev/null; then \ + echo 'error: core/contracts/ must not import plugins/' >&2; \ + exit 1; \ + fi + @echo " → pkg/ must not import plugins/..." + @if rg -n 'plugins/' pkg --glob '*.go' -g '!*_test.go' 2>/dev/null; then \ + echo 'error: pkg/ must not import plugins/' >&2; \ + exit 1; \ + fi + @echo " → plugins/domain/ must not import other plugins/domain/..." + @for d in plugins/domain/*/; do \ + name=$$(basename $$d); \ + imports=$$(rg -n '"github.com/Rain-kl/Wavelet/plugins/domain/' plugins/domain/"$$name" -g '*.go' 2>/dev/null | rg -v "plugins/domain/$$name/" | rg -v '_test.go' || true); \ + if [ -n "$$imports" ]; then \ + echo "error: plugins/domain/$$name must not import other domain plugins" >&2; \ + echo "$$imports" >&2; \ + exit 1; \ + fi; \ + done + @echo " → Architecture guards PASS" golangci-lint run cd frontend && pnpm tsc --noEmit --jsx preserve && npx eslint . --max-warnings 0 diff --git a/cmd/app.go b/cmd/app.go index 48d96813..3980dff0 100644 --- a/cmd/app.go +++ b/cmd/app.go @@ -8,8 +8,8 @@ import ( "time" "github.com/Rain-kl/Wavelet/core" - "github.com/Rain-kl/Wavelet/core/extpoints" "github.com/Rain-kl/Wavelet/pkg/config" + "github.com/Rain-kl/Wavelet/pkg/migrator" "github.com/Rain-kl/Wavelet/plugins/domain/admin" "github.com/Rain-kl/Wavelet/plugins/domain/auth" "github.com/Rain-kl/Wavelet/plugins/domain/cap" @@ -56,11 +56,8 @@ func newWaveletApp(profile core.Profile) *core.App { system.New(), ) - // 3. Bind Goose migration runner - app.SetMigrationRunner(func(_ context.Context, _ []extpoints.MigrationEntry) error { - runMigrations() - return nil - }) + // 3. Bind Goose migration engine (wraps pkg/migrator.Migrate) + app.SetMigrationEngine(&gooseEngine{}) // 4. Mount runtime drivers for each aspect app.Use( @@ -71,3 +68,12 @@ func newWaveletApp(profile core.Profile) *core.App { return app } + +// gooseEngine implements core.MigrationEngine by wrapping pkg/migrator.Migrate. +// It runs all centrally managed Goose SQL migrations during app startup. +type gooseEngine struct{} + +func (e *gooseEngine) Migrate(_ context.Context, _ []core.MigrationEntry) error { + _ = migrator.Migrate() + return nil +} diff --git a/cmd/root.go b/cmd/root.go index 62272801..3fd1b846 100644 --- a/cmd/root.go +++ b/cmd/root.go @@ -12,7 +12,6 @@ import ( "github.com/Rain-kl/Wavelet/pkg/buildinfo" "github.com/Rain-kl/Wavelet/pkg/config" "github.com/Rain-kl/Wavelet/pkg/logger" - "github.com/Rain-kl/Wavelet/pkg/migrator" "github.com/Rain-kl/Wavelet/pkg/trace" "github.com/spf13/cobra" ) @@ -38,9 +37,6 @@ var rootCmd = &cobra.Command{ TracerName: config.Config.Otel.TracerName, }) }, - PreRun: func(_ *cobra.Command, _ []string) { - runMigrations() - }, PersistentPostRun: func(_ *cobra.Command, _ []string) { shutdownTraceProvider() }, @@ -50,16 +46,6 @@ var rootCmd = &cobra.Command{ }, } -var latestMigrationState struct { - relationalDB migrator.Report - clickHouseDB migrator.Report -} - -func runMigrations() { - latestMigrationState.relationalDB = migrator.Migrate() - latestMigrationState.clickHouseDB = migrator.MigrateClickHouse() -} - func shutdownTraceProvider() { ctx, cancel := context.WithTimeout(context.Background(), traceShutdownTimeout) defer cancel() @@ -70,16 +56,7 @@ func init() { rootCmd.Version = buildinfo.Version rootCmd.CompletionOptions.DisableDefaultCmd = true - // 1. 为需要迁移的子命令动态绑定原先 rootCmd.PreRun 拥有的数据库迁移行为 - migratePreRun := func(_ *cobra.Command, _ []string) { - runMigrations() - } - allCmd.PreRun = migratePreRun - apiCmd.PreRun = migratePreRun - workerCmd.PreRun = migratePreRun - schedulerCmd.PreRun = migratePreRun - - // 2. 集中将这些命令注册为真正的子命令,以解决 Cobra 的 unknown command 校验限制 + // 集中将子命令注册到根命令,以解决 Cobra 的 unknown command 校验限制 rootCmd.AddCommand(allCmd, apiCmd, workerCmd, schedulerCmd) } @@ -88,4 +65,4 @@ func Execute() { if err := rootCmd.Execute(); err != nil { log.Fatalf("[CMD] execute failed; %s\n", err) } -} +} \ No newline at end of file diff --git a/core/contracts/auth.go b/core/contracts/auth.go index 5f75fc1d..f4ade9ba 100644 --- a/core/contracts/auth.go +++ b/core/contracts/auth.go @@ -64,14 +64,23 @@ type AuthService interface { // GetCurrentUser retrieves the authenticated UserDTO from context. GetCurrentUser(ctx context.Context) (*UserDTO, error) + // GetCurrentUserID retrieves the authenticated user ID from session/context. + GetCurrentUserID(ctx context.Context) (uint64, error) + // VerifyToken validates an access token and returns the associated user DTO. VerifyToken(ctx context.Context, token string) (*UserDTO, error) // CreateSession establishes an authenticated session for the given user ID. CreateSession(ctx context.Context, userID uint64, extras map[string]any) (string, error) + // RevokeToken invalidates a specific access token by its hash. + RevokeToken(ctx context.Context, tokenHash string) error + // RevokeUserSessions revokes all active sessions and cached tokens for a user. RevokeUserSessions(ctx context.Context, userID uint64) error + + // DisallowTokenAuthMiddleware returns a middleware that rejects requests authenticated via access token. + DisallowTokenAuthMiddleware() any } // AuthRegistry allows downstream and domain plugins to register custom authentication providers. @@ -80,3 +89,12 @@ type AuthRegistry interface { GetOAuthProvider(name string) (OAuthProvider, bool) ListOAuthProviders() []string } + +// Auth context keys — stored in Gin context by auth middleware, consumed by domain plugins. +const ( + AuthUserIDKey = "user_id" + AuthUserNameKey = "username" + AuthUserObjKey = "user_obj" + AuthTokenAuthKey = "token_auth" // marks if request uses access token auth + AuthTokenAdminKey = "token_admin" // whether the access token has admin privileges +) diff --git a/core/contracts/events.go b/core/contracts/events.go index 0a4bc824..cc63bbae 100644 --- a/core/contracts/events.go +++ b/core/contracts/events.go @@ -1,13 +1,109 @@ // Copyright 2026 Arctel.net // SPDX-License-Identifier: Apache-2.0 +// Package contracts defines unified service interfaces and DTOs for cross-plugin communication. package contracts -// EventTopicAdminLoggedIn 管理员登录事件主题 -const EventTopicAdminLoggedIn = "admin:logged_in" +// ====================================================================== +// Domain Event Topic Constants +// ====================================================================== +// +// All cross-plugin domain event topics MUST be declared here so that +// producers and consumers share the same string values without importing +// each other's implementation packages. +// ====================================================================== + +// --- Auth & User Events --- +const ( + // EventTopicAdminLoggedIn fires when an admin user logs in. + EventTopicAdminLoggedIn = "admin:logged_in" + + // EventTopicUserCreated fires when a new user account is created. + EventTopicUserCreated = "user:created" + + // EventTopicUserUpdated fires when a user profile is updated. + EventTopicUserUpdated = "user:updated" + + // EventTopicUserDeleted fires when a user account is deleted. + EventTopicUserDeleted = "user:deleted" +) + +// --- Admin & System Events --- +const ( + // EventTopicConfigChanged fires when a system configuration value changes. + EventTopicConfigChanged = "admin:config_changed" + + // EventTopicSystemCleanup fires when a periodic system cleanup completes. + EventTopicSystemCleanup = "admin:system_cleanup" +) + +// --- Upload / Storage Events --- +const ( + // EventTopicUploadCreated fires when a new file upload is recorded. + EventTopicUploadCreated = "upload:created" + + // EventTopicUploadDeleted fires when a file upload is removed. + EventTopicUploadDeleted = "upload:deleted" + + // EventTopicIngestComplete fires when a programmatic file ingest finishes. + EventTopicIngestComplete = "upload:ingest_complete" +) + +// --- Message Gateway Events --- +const ( + // EventTopicNotificationSent fires when a push notification is dispatched. + EventTopicNotificationSent = "message:notification_sent" + + // EventTopicChannelBound fires when a user binds a messaging channel. + EventTopicChannelBound = "message:channel_bound" + + // EventTopicChannelUnbound fires when a user unbinds a messaging channel. + EventTopicChannelUnbound = "message:channel_unbound" +) + +// --- Risk Control Events --- +const ( + // EventTopicAccessLogRecorded fires when a user access log entry is recorded. + EventTopicAccessLogRecorded = "risk:access_log_recorded" +) + +// ====================================================================== +// Domain Event Payload DTOs +// ====================================================================== // AdminLoggedIn 管理员登录领域事件载荷 type AdminLoggedIn struct { User *UserDTO `json:"user"` IP string `json:"ip"` } + +// UserCreatedEvent fires when a new user account is created. +type UserCreatedEvent struct { + User *UserDTO `json:"user"` + Password string `json:"-"` +} + +// ConfigChangedEvent fires when a system configuration value changes. +type ConfigChangedEvent struct { + Key string `json:"key"` + OldVal any `json:"old_val,omitempty"` + NewVal any `json:"new_val,omitempty"` +} + +// UploadCreatedEvent fires when a new file upload is recorded. +type UploadCreatedEvent struct { + UploadID uint64 `json:"upload_id,string"` + UserID uint64 `json:"user_id,string"` + FileName string `json:"file_name"` + FileSize int64 `json:"file_size"` + MimeType string `json:"mime_type"` +} + +// NotificationSentEvent fires when a push notification is dispatched. +type NotificationSentEvent struct { + UserID uint64 `json:"user_id,string"` + Channel string `json:"channel"` + Title string `json:"title"` + Success bool `json:"success"` + ErrorInfo string `json:"error_info,omitempty"` +} \ No newline at end of file diff --git a/downstream/README.md b/downstream/README.md new file mode 100644 index 00000000..2deb1c5e --- /dev/null +++ b/downstream/README.md @@ -0,0 +1,77 @@ +# Downstream Custom Plugins + +This directory is the designated location for downstream (deployment-specific) Cordis plugins. + +## Architecture + +``` +downstream/ +├── README.md +└── plugins/ + └── custom_example/ # Example plugin — copy & rename to get started + └── plugin.go +``` + +Downstream plugins follow the same `core.Plugin` contract as platform plugins: + +```go +type Plugin interface { + Name() string + Apply(ctx *core.Context) error +} +``` + +## Rules + +1. **Naming**: Each plugin directory name becomes its import path and plugin ID (kebab-case recommended). +2. **Dependencies**: Downstream plugins may import `core/`, `core/contracts/`, `pkg/`, and `plugins/infra/` packages from the platform. They MUST NOT import domain plugin internal packages — use `core.Inject[contracts.XxxService](ctx)` instead. +3. **Registration**: Add your downstream plugin to `cmd/app.go` before the platform plugins or after, depending on which services it needs: + ```go + // newWaveletApp in cmd/app.go + app.Use( + database.New(), + cache.New(), + logger.New(), + storage.New(), + // ... platform domain plugins ... + custom_hello.New(), // your downstream plugin + driver_http.New(driver_http.WithAddr(config.Config.App.Addr)), + driver_asynq_worker.New(), + driver_asynq_cron.New(), + ) + ``` +4. **Migration**: If your plugin needs database tables, embed SQL files in a `migrations/` directory and register via `ctx.Migrations().Register(...)` in `Apply()`. + +## Quick Start + +```go +package custom_example + +import ( + "github.com/Rain-kl/Wavelet/core" + "github.com/Rain-kl/Wavelet/core/contracts" + "github.com/gin-gonic/gin" +) + +type Plugin struct{} + +func New() *Plugin { return &Plugin{} } + +func (p *Plugin) Name() string { return "custom_example" } + +func (p *Plugin) Apply(ctx *core.Context) error { + // Example: register a route that uses AuthService + var authSvc contracts.AuthService + if err := ctx.Using(func(svc contracts.AuthService) { authSvc = svc }); err != nil { + return err + } + + g := ctx.Router().Group("/api/v1/custom", authSvc.RequireAuthMiddleware().(gin.HandlerFunc)) + g.GET("/hello", func(c *gin.Context) { + user, _ := authSvc.GetCurrentUser(c.Request.Context()) + c.JSON(200, gin.H{"message": "Hello " + user.Username}) + }) + + return nil +} +``` \ No newline at end of file diff --git a/downstream/plugins/custom_example/plugin.go b/downstream/plugins/custom_example/plugin.go new file mode 100644 index 00000000..057061ac --- /dev/null +++ b/downstream/plugins/custom_example/plugin.go @@ -0,0 +1,48 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +// Package custom_example demonstrates how to build a downstream Cordis plugin. +// Copy this directory to create your own plugin. +package custom_example + +import ( + "github.com/Rain-kl/Wavelet/core" + "github.com/Rain-kl/Wavelet/core/contracts" + "github.com/gin-gonic/gin" +) + +// Plugin implements core.Plugin for the custom_example downstream plugin. +type Plugin struct{} + +// New creates a new custom_example plugin. +func New() *Plugin { + return &Plugin{} +} + +// Name returns the unique identifier for this plugin. +func (p *Plugin) Name() string { + return "custom_example" +} + +// Apply registers routes and services into the Cordis micro-kernel Context. +func (p *Plugin) Apply(ctx *core.Context) error { + // Resolve platform services via IoC container (no direct imports of domain plugins). + var authSvc contracts.AuthService + if err := core.Using[contracts.AuthService](ctx, func(svc contracts.AuthService) { authSvc = svc }); err != nil { + return err + } + _ = authSvc + + // Register routes using the auth middleware obtained through the contract. + g := ctx.Router().Group("/api/v1/custom", authSvc.RequireAuthMiddleware().(gin.HandlerFunc)) + g.GET("/hello", func(c *gin.Context) { + user, err := authSvc.GetCurrentUser(c.Request.Context()) + if err != nil { + c.JSON(401, gin.H{"error": "unauthorized"}) + return + } + c.JSON(200, gin.H{"message": "Hello " + user.Username}) + }) + + return nil +} \ No newline at end of file diff --git a/pkg/util/gin_context.go b/pkg/util/gin_context.go new file mode 100644 index 00000000..602501b1 --- /dev/null +++ b/pkg/util/gin_context.go @@ -0,0 +1,22 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package util + +import "github.com/gin-gonic/gin" + +// GetFromContext retrieves a typed value from Gin context. +func GetFromContext[T any](c *gin.Context, key string) (T, bool) { + value, exists := c.Get(key) + if !exists { + var zero T + return zero, false + } + typed, ok := value.(T) + return typed, ok +} + +// SetToContext sets a typed value into Gin context. +func SetToContext[T any](c *gin.Context, key string, value T) { + c.Set(key, value) +} \ No newline at end of file diff --git a/plugins/domain/admin/handlers_user.go b/plugins/domain/admin/handlers_user.go index 84572d6c..b076f937 100644 --- a/plugins/domain/admin/handlers_user.go +++ b/plugins/domain/admin/handlers_user.go @@ -247,7 +247,7 @@ func DeleteUser(c *gin.Context) { return } - currUser, _ := auth.GetFromContext[*contracts.UserDTO](c, auth.UserObjKey) + currUser, _ := util.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey) if currUser == nil { response.AbortUnauthorized(c, AdminRequired) return @@ -338,7 +338,7 @@ func UpdateUser(c *gin.Context) { return } - currUser, _ := auth.GetFromContext[*contracts.UserDTO](c, auth.UserObjKey) + currUser, _ := util.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey) if currUser == nil { response.AbortUnauthorized(c, AdminRequired) return diff --git a/plugins/domain/admin/middlewares.go b/plugins/domain/admin/middlewares.go index 570ea672..2497920c 100644 --- a/plugins/domain/admin/middlewares.go +++ b/plugins/domain/admin/middlewares.go @@ -7,26 +7,26 @@ import ( "github.com/Rain-kl/Wavelet/core/contracts" "github.com/Rain-kl/Wavelet/pkg/logger" "github.com/Rain-kl/Wavelet/pkg/response" - otel_trace "github.com/Rain-kl/Wavelet/pkg/trace" - "github.com/Rain-kl/Wavelet/plugins/domain/auth" + "github.com/Rain-kl/Wavelet/pkg/trace" + "github.com/Rain-kl/Wavelet/pkg/util" "github.com/gin-gonic/gin" ) // LoginAdminRequired 返回管理员权限校验中间件 func LoginAdminRequired() gin.HandlerFunc { return func(c *gin.Context) { - ctx, span := otel_trace.Start(c.Request.Context(), "LoginAdminRequired") + ctx, span := trace.Start(c.Request.Context(), "LoginAdminRequired") defer span.End() - user, _ := auth.GetFromContext[*contracts.UserDTO](c, auth.UserObjKey) + user, _ := util.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey) if user == nil { response.AbortNotFound(c, AdminRequired) return } // 如果是通过 Access Token 鉴权,需要检查令牌本身是否具有管理员权限 - if tokenAuth, _ := auth.GetFromContext[bool](c, auth.TokenAuthKey); tokenAuth { - tokenAdmin, _ := auth.GetFromContext[bool](c, auth.TokenAdminKey) + if tokenAuth, _ := util.GetFromContext[bool](c, contracts.AuthTokenAuthKey); tokenAuth { + tokenAdmin, _ := util.GetFromContext[bool](c, contracts.AuthTokenAdminKey) if !tokenAdmin { response.AbortNotFound(c, TokenAdminRequired) return diff --git a/plugins/domain/admin/plugin.go b/plugins/domain/admin/plugin.go index e5b0deea..714abc57 100644 --- a/plugins/domain/admin/plugin.go +++ b/plugins/domain/admin/plugin.go @@ -8,8 +8,9 @@ import ( "context" "github.com/Rain-kl/Wavelet/core" + "github.com/Rain-kl/Wavelet/core/contracts" "github.com/Rain-kl/Wavelet/core/extpoints" - "github.com/Rain-kl/Wavelet/plugins/domain/auth" + "github.com/gin-gonic/gin" "github.com/hibiken/asynq" ) @@ -47,8 +48,16 @@ func (p *Plugin) Manifest() core.Manifest { // Apply registers admin routes, tasks, schedules, and settings into the Context. func (p *Plugin) Apply(ctx *core.Context) error { + // 0. Resolve auth service for middleware (via IoC, not direct import) + var authSvc contracts.AuthService + if err := core.Using[contracts.AuthService](ctx, func(svc contracts.AuthService) { authSvc = svc }); err != nil { + return err + } + loginMW := authSvc.RequireAuthMiddleware().(gin.HandlerFunc) + adminMW := authSvc.RequireAdminMiddleware().(gin.HandlerFunc) + // 1. Register Admin HTTP Routes - adminRouter := ctx.Router().Group("/api/v1/admin", auth.LoginRequired(), LoginAdminRequired()) + adminRouter := ctx.Router().Group("/api/v1/admin", loginMW, adminMW) { // Status & Diagnostics adminRouter.GET("/status", GetSystemStatus) diff --git a/plugins/domain/auth/handlers.go b/plugins/domain/auth/handlers.go index 48518741..e8d8ed4c 100644 --- a/plugins/domain/auth/handlers.go +++ b/plugins/domain/auth/handlers.go @@ -430,7 +430,7 @@ func handleCallbackRegister(ctx context.Context, c *gin.Context, source *AuthSou // UserInfo 获取当前登录用户信息 func UserInfo(c *gin.Context) { - user, _ := GetFromContext[*contracts.UserDTO](c, UserObjKey) + user, _ := util.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey) session := sessions.Default(c) needChange := session.Get("need_change_password") == true diff --git a/plugins/domain/auth/middleware.go b/plugins/domain/auth/middleware.go index 19f1cfaa..8b337ab7 100644 --- a/plugins/domain/auth/middleware.go +++ b/plugins/domain/auth/middleware.go @@ -12,27 +12,12 @@ import ( "github.com/Rain-kl/Wavelet/core/contracts" "github.com/Rain-kl/Wavelet/pkg/response" - otel_trace "github.com/Rain-kl/Wavelet/pkg/trace" + "github.com/Rain-kl/Wavelet/pkg/trace" + "github.com/Rain-kl/Wavelet/pkg/util" db "github.com/Rain-kl/Wavelet/plugins/infra/database" "github.com/gin-gonic/gin" ) -// GetFromContext 从 Gin 请求上下文获取指定类型的值。 -func GetFromContext[T any](c *gin.Context, key string) (T, bool) { - value, exists := c.Get(key) - if !exists { - var zero T - return zero, false - } - typed, ok := value.(T) - return typed, ok -} - -// SetToContext 设置值到 Gin 请求上下文。 -func SetToContext[T any](c *gin.Context, key string, value T) { - c.Set(key, value) -} - func hashToken(token string) string { h := sha256.New() h.Write([]byte(token)) @@ -91,8 +76,8 @@ func GetUserFromRequest(c *gin.Context) (*contracts.UserDTO, error) { if user.Username == SystemUsername { return nil, errors.New("system user is not allowed to login") } - SetToContext(c, TokenAuthKey, true) - SetToContext(c, TokenAdminKey, tokenRecord.IsAdmin) + util.SetToContext(c, contracts.AuthTokenAuthKey, true) + util.SetToContext(c, contracts.AuthTokenAdminKey, tokenRecord.IsAdmin) return user, nil } } @@ -113,8 +98,8 @@ func GetUserFromRequest(c *gin.Context) (*contracts.UserDTO, error) { SetCachedUser(ctx, userID, user) } - SetToContext(c, TokenAuthKey, false) - SetToContext(c, TokenAdminKey, false) + util.SetToContext(c, contracts.AuthTokenAuthKey, false) + util.SetToContext(c, contracts.AuthTokenAdminKey, false) if user.Username == "system" { return nil, errors.New("system user is not allowed to login") @@ -126,7 +111,7 @@ func GetUserFromRequest(c *gin.Context) (*contracts.UserDTO, error) { // LoginRequired 返回登录鉴权中间件,校验 Access Token 或 Session func LoginRequired() gin.HandlerFunc { return func(c *gin.Context) { - ctx, span := otel_trace.Start(c.Request.Context(), "LoginRequired") + _, span := trace.Start(c.Request.Context(), "LoginRequired") defer span.End() user, err := GetUserFromRequest(c) @@ -135,8 +120,8 @@ func LoginRequired() gin.HandlerFunc { return } - LogForAudit(ctx, user, c) - SetToContext(c, UserObjKey, user) + LogForAudit(c.Request.Context(), user, c) + util.SetToContext(c, contracts.AuthUserObjKey, user) c.Next() } } @@ -144,7 +129,7 @@ func LoginRequired() gin.HandlerFunc { // AdminRequired 校验管理员权限(支持 Session 和 Token 鉴权) func AdminRequired() gin.HandlerFunc { return func(c *gin.Context) { - ctx, span := otel_trace.Start(c.Request.Context(), "AdminRequired") + _, span := trace.Start(c.Request.Context(), "AdminRequired") defer span.End() user, err := GetUserFromRequest(c) @@ -153,8 +138,8 @@ func AdminRequired() gin.HandlerFunc { return } - isTokenAuth, _ := GetFromContext[bool](c, TokenAuthKey) - isTokenAdmin, _ := GetFromContext[bool](c, TokenAdminKey) + isTokenAuth, _ := util.GetFromContext[bool](c, contracts.AuthTokenAuthKey) + isTokenAdmin, _ := util.GetFromContext[bool](c, contracts.AuthTokenAdminKey) // 如果是通过 Token 鉴权,要求该 Token 具备管理员权限或者用户本身为管理员 if isTokenAuth && !isTokenAdmin && !user.IsAdmin { @@ -168,8 +153,8 @@ func AdminRequired() gin.HandlerFunc { return } - LogForAudit(ctx, user, c) - SetToContext(c, UserObjKey, user) + LogForAudit(c.Request.Context(), user, c) + util.SetToContext(c, contracts.AuthUserObjKey, user) c.Next() } } @@ -182,10 +167,10 @@ func LoginAdminRequired() gin.HandlerFunc { // DisallowTokenAuth 拒绝使用 Access Token 进行身份验证的请求访问该端点 func DisallowTokenAuth() gin.HandlerFunc { return func(c *gin.Context) { - if tokenAuth, _ := GetFromContext[bool](c, TokenAuthKey); tokenAuth { + if tokenAuth, _ := util.GetFromContext[bool](c, contracts.AuthTokenAuthKey); tokenAuth { response.AbortForbidden(c, ErrTokenAuthNotAllowed) return } c.Next() } -} +} \ No newline at end of file diff --git a/plugins/domain/auth/plugin_test.go b/plugins/domain/auth/plugin_test.go index 92e5bd8e..228043ae 100644 --- a/plugins/domain/auth/plugin_test.go +++ b/plugins/domain/auth/plugin_test.go @@ -133,7 +133,7 @@ func TestAuthPluginUnit(t *testing.T) { require.NoError(t, authSvc.RevokeUserSessions(context.Background(), user.ID)) // GetCurrentUser from context - userCtx := context.WithValue(context.Background(), auth.UserObjKey, userDTO) + userCtx := context.WithValue(context.Background(), contracts.AuthUserObjKey, userDTO) current, err := authSvc.GetCurrentUser(userCtx) require.NoError(t, err) assert.Equal(t, user.ID, current.ID) diff --git a/plugins/domain/auth/service.go b/plugins/domain/auth/service.go index 6aab5cf0..a39826ea 100644 --- a/plugins/domain/auth/service.go +++ b/plugins/domain/auth/service.go @@ -9,6 +9,7 @@ import ( "sync" "github.com/Rain-kl/Wavelet/core/contracts" + "github.com/Rain-kl/Wavelet/pkg/util" db "github.com/Rain-kl/Wavelet/plugins/infra/database" "github.com/gin-gonic/gin" ) @@ -29,12 +30,12 @@ func (s *authServiceImpl) RequireAdminMiddleware() any { func (s *authServiceImpl) GetCurrentUser(ctx context.Context) (*contracts.UserDTO, error) { if ginCtx, ok := ctx.(*gin.Context); ok { - if u, ok := GetFromContext[*contracts.UserDTO](ginCtx, UserObjKey); ok && u != nil { + if u, ok := util.GetFromContext[*contracts.UserDTO](ginCtx, contracts.AuthUserObjKey); ok && u != nil { return u, nil } } - if v := ctx.Value(UserObjKey); v != nil { + if v := ctx.Value(contracts.AuthUserObjKey); v != nil { if u, ok := v.(*contracts.UserDTO); ok && u != nil { return u, nil } @@ -93,6 +94,22 @@ func (s *authServiceImpl) RevokeUserSessions(ctx context.Context, userID uint64) return nil } +func (s *authServiceImpl) GetCurrentUserID(ctx context.Context) (uint64, error) { + if ginCtx, ok := ctx.(*gin.Context); ok { + return GetUserIDFromContext(ginCtx), nil + } + return 0, errors.New("auth: user not found in context") +} + +func (s *authServiceImpl) RevokeToken(ctx context.Context, tokenHash string) error { + InvalidateCachedToken(ctx, tokenHash) + return nil +} + +func (s *authServiceImpl) DisallowTokenAuthMiddleware() any { + return DisallowTokenAuth() +} + type authRegistryImpl struct { mu sync.RWMutex providers map[string]contracts.OAuthProvider diff --git a/plugins/domain/message_gateway/handlers.go b/plugins/domain/message_gateway/handlers.go index 7fce9f89..ae33cae2 100644 --- a/plugins/domain/message_gateway/handlers.go +++ b/plugins/domain/message_gateway/handlers.go @@ -10,12 +10,12 @@ import ( "github.com/Rain-kl/Wavelet/core/contracts" "github.com/Rain-kl/Wavelet/pkg/response" - "github.com/Rain-kl/Wavelet/plugins/domain/auth" + "github.com/Rain-kl/Wavelet/pkg/util" "github.com/gin-gonic/gin" ) func currentUser(c *gin.Context) (*contracts.UserDTO, bool) { - return auth.GetFromContext[*contracts.UserDTO](c, auth.UserObjKey) + return util.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey) } // ListChannels lists enabled channels a user can bind. @@ -136,8 +136,8 @@ func UnbindBinding(c *gin.Context) { } // RegisterUserRoutes mounts user-facing message gateway endpoints. -func RegisterUserRoutes(r *gin.RouterGroup) { - mg := r.Group("/message-gateway", auth.LoginRequired()) +func RegisterUserRoutes(r *gin.RouterGroup, loginMW gin.HandlerFunc) { + mg := r.Group("/message-gateway", loginMW) { mg.GET("/channels", ListChannels) mg.GET("/bindings", ListBindings) diff --git a/plugins/domain/message_gateway/plugin.go b/plugins/domain/message_gateway/plugin.go index 533df798..8f4c5a95 100644 --- a/plugins/domain/message_gateway/plugin.go +++ b/plugins/domain/message_gateway/plugin.go @@ -9,8 +9,9 @@ import ( "embed" "github.com/Rain-kl/Wavelet/core" + "github.com/Rain-kl/Wavelet/core/contracts" "github.com/Rain-kl/Wavelet/core/extpoints" - "github.com/Rain-kl/Wavelet/plugins/domain/auth" + "github.com/gin-gonic/gin" "github.com/hibiken/asynq" ) @@ -70,11 +71,19 @@ type PushNotificationEvent struct { // Apply registers message_gateway migrations, routes, tasks, schedules, events, and settings into the Context. func (p *Plugin) Apply(ctx *core.Context) error { + // 0. Resolve auth service for middleware (via IoC, not direct import) + var authSvc contracts.AuthService + if err := core.Using[contracts.AuthService](ctx, func(svc contracts.AuthService) { authSvc = svc }); err != nil { + return err + } + loginMW := authSvc.RequireAuthMiddleware().(gin.HandlerFunc) + adminMW := authSvc.RequireAdminMiddleware().(gin.HandlerFunc) + // 1. Register migrations ctx.Migrations().Register("message_gateway", mgMigrations) // 2. Register User HTTP Routes - mgGroup := ctx.Router().Group("/api/v1/message-gateway", auth.LoginRequired()) + mgGroup := ctx.Router().Group("/api/v1/message-gateway", loginMW) { mgGroup.GET("/channels", ListChannels) mgGroup.GET("/bindings", ListBindings) @@ -83,7 +92,7 @@ func (p *Plugin) Apply(ctx *core.Context) error { } // 3. Register Admin Message Gateway HTTP Routes - adminMgGroup := ctx.Router().Group("/api/v1/admin/message-gateway", auth.LoginRequired(), auth.LoginAdminRequired()) + adminMgGroup := ctx.Router().Group("/api/v1/admin/message-gateway", loginMW, adminMW) { adminMgGroup.GET("/channels/definitions", ListAdminChannelDefinitions) adminMgGroup.GET("/channels", ListAdminChannels) @@ -94,7 +103,7 @@ func (p *Plugin) Apply(ctx *core.Context) error { } // 4. Register Admin Push HTTP Routes - adminPushGroup := ctx.Router().Group("/api/v1/admin/push", auth.LoginRequired(), auth.LoginAdminRequired()) + adminPushGroup := ctx.Router().Group("/api/v1/admin/push", loginMW, adminMW) { events := adminPushGroup.Group("/events") { diff --git a/plugins/domain/risk_control/middleware.go b/plugins/domain/risk_control/middleware.go index eeac946f..53414dac 100644 --- a/plugins/domain/risk_control/middleware.go +++ b/plugins/domain/risk_control/middleware.go @@ -13,7 +13,7 @@ import ( "github.com/Rain-kl/Wavelet/pkg/config" "github.com/Rain-kl/Wavelet/pkg/idgen" "github.com/Rain-kl/Wavelet/pkg/response" - "github.com/Rain-kl/Wavelet/plugins/domain/auth" + "github.com/Rain-kl/Wavelet/pkg/util" "github.com/Rain-kl/Wavelet/plugins/domain/risk_control/logstore" "github.com/gin-gonic/gin" ) @@ -42,7 +42,7 @@ func RiskControlMiddleware() gin.HandlerFunc { c.Next() // 3. 后置身份检查:仅记录通过认证的请求 - userObj, exists := auth.GetFromContext[*contracts.UserDTO](c, auth.UserObjKey) + userObj, exists := util.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey) if !exists || userObj == nil { return } diff --git a/plugins/domain/risk_control/middleware_test.go b/plugins/domain/risk_control/middleware_test.go index dd36ccdd..238ad7a8 100644 --- a/plugins/domain/risk_control/middleware_test.go +++ b/plugins/domain/risk_control/middleware_test.go @@ -16,7 +16,7 @@ import ( "github.com/Rain-kl/Wavelet/pkg/batchwriter" "github.com/Rain-kl/Wavelet/pkg/config" "github.com/Rain-kl/Wavelet/pkg/testhelper" - "github.com/Rain-kl/Wavelet/plugins/domain/auth" + "github.com/Rain-kl/Wavelet/pkg/util" "github.com/Rain-kl/Wavelet/plugins/domain/risk_control" "github.com/Rain-kl/Wavelet/plugins/domain/risk_control/logstore" "github.com/gin-gonic/gin" @@ -99,7 +99,7 @@ func TestRiskControlMiddleware(t *testing.T) { r := gin.New() r.Use(func(c *gin.Context) { user := &contracts.UserDTO{ID: 12345} - auth.SetToContext(c, auth.UserObjKey, user) + util.SetToContext[*contracts.UserDTO](c, contracts.AuthUserObjKey, user) c.Next() }) r.Use(risk_control.RiskControlMiddleware()) diff --git a/plugins/domain/upload/filesrv/file_server.go b/plugins/domain/upload/filesrv/file_server.go index e4544864..3a5f1684 100644 --- a/plugins/domain/upload/filesrv/file_server.go +++ b/plugins/domain/upload/filesrv/file_server.go @@ -17,6 +17,7 @@ import ( "github.com/Rain-kl/Wavelet/core/contracts" "github.com/Rain-kl/Wavelet/pkg/response" + pkgutil "github.com/Rain-kl/Wavelet/pkg/util" "github.com/Rain-kl/Wavelet/plugins/domain/auth" "github.com/Rain-kl/Wavelet/plugins/domain/upload/cache" "github.com/Rain-kl/Wavelet/plugins/domain/upload/models" @@ -284,7 +285,7 @@ func getOriginalFileBytes(ctx context.Context, upload *models.Upload) ([]byte, e func checkPrivateFileOwner(c *gin.Context, ownerID uint64) error { var currUserID uint64 var isAdmin bool - if u, ok := auth.GetFromContext[*contracts.UserDTO](c, auth.UserObjKey); ok && u != nil { + if u, ok := pkgutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey); ok && u != nil { currUserID = u.ID isAdmin = u.IsAdmin } else { @@ -311,7 +312,7 @@ func CheckFileAccessPermission(c *gin.Context, upload *models.Upload) error { } if !cache.IsFilePublic(c.Request.Context(), upload.Type) { - if _, ok := auth.GetFromContext[*contracts.UserDTO](c, auth.UserObjKey); !ok { + if _, ok := pkgutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey); !ok { if _, err := auth.GetUserFromRequest(c); err != nil { return err } diff --git a/plugins/domain/upload/handler/file_management.go b/plugins/domain/upload/handler/file_management.go index fa431054..e0f88b94 100644 --- a/plugins/domain/upload/handler/file_management.go +++ b/plugins/domain/upload/handler/file_management.go @@ -11,7 +11,7 @@ import ( "github.com/Rain-kl/Wavelet/core/contracts" "github.com/Rain-kl/Wavelet/pkg/response" - "github.com/Rain-kl/Wavelet/plugins/domain/auth" + "github.com/Rain-kl/Wavelet/pkg/util" "github.com/Rain-kl/Wavelet/plugins/domain/upload/ingest" "github.com/Rain-kl/Wavelet/plugins/domain/upload/models" "github.com/Rain-kl/Wavelet/plugins/domain/upload/repository" @@ -168,7 +168,7 @@ type listMyFilesResponse struct { // @Failure 401 {object} response.Any "未登录" // @Router /api/v1/upload/my [get] func ListMyFiles(c *gin.Context) { - currUser, _ := auth.GetFromContext[*contracts.UserDTO](c, auth.UserObjKey) + currUser, _ := util.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey) ctx := c.Request.Context() var req listMyFilesRequest @@ -215,7 +215,7 @@ func ListMyFiles(c *gin.Context) { // @Failure 404 {object} response.Any "文件不存在" // @Router /api/v1/upload/{id} [delete] func DeleteMyFile(c *gin.Context) { - currUser, _ := auth.GetFromContext[*contracts.UserDTO](c, auth.UserObjKey) + currUser, _ := util.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey) ctx := c.Request.Context() if uploadstorage.ReadOnly(ctx) { response.AbortConflict(c, shared.ErrStorageReadOnly) @@ -262,7 +262,7 @@ type updateMyFileRequest struct { // @Failure 404 {object} response.Any "文件不存在" // @Router /api/v1/upload/{id} [put] func UpdateMyFile(c *gin.Context) { - currUser, _ := auth.GetFromContext[*contracts.UserDTO](c, auth.UserObjKey) + currUser, _ := util.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey) ctx := c.Request.Context() if uploadstorage.ReadOnly(ctx) { response.AbortConflict(c, shared.ErrStorageReadOnly) diff --git a/plugins/domain/upload/handler/routers.go b/plugins/domain/upload/handler/routers.go index e523d52a..90e5b5d3 100644 --- a/plugins/domain/upload/handler/routers.go +++ b/plugins/domain/upload/handler/routers.go @@ -25,7 +25,7 @@ import ( "github.com/Rain-kl/Wavelet/core/contracts" "github.com/Rain-kl/Wavelet/pkg/logger" "github.com/Rain-kl/Wavelet/pkg/response" - "github.com/Rain-kl/Wavelet/plugins/domain/auth" + pkgutil "github.com/Rain-kl/Wavelet/pkg/util" "github.com/Rain-kl/Wavelet/plugins/domain/upload/filesrv" "github.com/Rain-kl/Wavelet/plugins/domain/upload/ingest" "github.com/Rain-kl/Wavelet/plugins/domain/upload/models" @@ -63,7 +63,7 @@ func UploadFile(c *gin.Context) { c.Request.Body = http.MaxBytesReader(c.Writer, c.Request.Body, shared.MaxUploadSize) - currUser, _ := auth.GetFromContext[*contracts.UserDTO](c, auth.UserObjKey) + currUser, _ := pkgutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey) ctx := c.Request.Context() header, err := c.FormFile("file") diff --git a/plugins/domain/upload/handler/routers_test.go b/plugins/domain/upload/handler/routers_test.go index 9aa252b2..03beaa9d 100644 --- a/plugins/domain/upload/handler/routers_test.go +++ b/plugins/domain/upload/handler/routers_test.go @@ -22,7 +22,7 @@ import ( "github.com/Rain-kl/Wavelet/core/contracts" "github.com/Rain-kl/Wavelet/pkg/response" "github.com/Rain-kl/Wavelet/pkg/testhelper" - "github.com/Rain-kl/Wavelet/plugins/domain/auth" + "github.com/Rain-kl/Wavelet/pkg/util" "github.com/Rain-kl/Wavelet/plugins/domain/upload/models" "github.com/Rain-kl/Wavelet/plugins/domain/upload/shared" uploadstats "github.com/Rain-kl/Wavelet/plugins/domain/upload/stats" @@ -42,7 +42,7 @@ func setupTestRouter(authUser *contracts.UserDTO) *gin.Engine { authMiddleware := func(c *gin.Context) { if authUser != nil { - auth.SetToContext(c, auth.UserObjKey, authUser) + util.SetToContext[*contracts.UserDTO](c, contracts.AuthUserObjKey, authUser) } c.Next() } diff --git a/plugins/domain/upload/plugin.go b/plugins/domain/upload/plugin.go index 18e070ca..21ce1342 100644 --- a/plugins/domain/upload/plugin.go +++ b/plugins/domain/upload/plugin.go @@ -8,11 +8,12 @@ import ( "context" "github.com/Rain-kl/Wavelet/core" + "github.com/Rain-kl/Wavelet/core/contracts" "github.com/Rain-kl/Wavelet/core/extpoints" - "github.com/Rain-kl/Wavelet/plugins/domain/auth" "github.com/Rain-kl/Wavelet/plugins/domain/upload/filesrv" "github.com/Rain-kl/Wavelet/plugins/domain/upload/handler" "github.com/Rain-kl/Wavelet/plugins/domain/upload/task" + "github.com/gin-gonic/gin" "github.com/hibiken/asynq" ) @@ -41,11 +42,18 @@ func (p *Plugin) Manifest() core.Manifest { // Apply registers upload routes, tasks, and settings into the Context. func (p *Plugin) Apply(ctx *core.Context) error { + // 0. Resolve auth service for middleware (via IoC, not direct import) + var authSvc contracts.AuthService + if err := core.Using[contracts.AuthService](ctx, func(svc contracts.AuthService) { authSvc = svc }); err != nil { + return err + } + loginMW := authSvc.RequireAuthMiddleware().(gin.HandlerFunc) + // 1. Register File Server Routes ctx.Router().GET("/f/:id", filesrv.ServeFileByID) // 2. Register User/Admin Upload HTTP Routes - uploadGroup := ctx.Router().Group("/api/v1/upload", auth.LoginRequired()) + uploadGroup := ctx.Router().Group("/api/v1/upload", loginMW) { uploadGroup.POST("", handler.UploadFile) uploadGroup.GET("", handler.ListFiles) @@ -53,7 +61,7 @@ func (p *Plugin) Apply(ctx *core.Context) error { uploadGroup.POST("/batch-download", handler.BatchDownloadFiles) } - adminUploadGroup := ctx.Router().Group("/api/v1/admin/uploads", auth.LoginRequired()) + adminUploadGroup := ctx.Router().Group("/api/v1/admin/uploads", loginMW) { adminUploadGroup.GET("", handler.ListFiles) adminUploadGroup.GET("/stats", handler.GetFileStats) diff --git a/plugins/domain/user/handlers.go b/plugins/domain/user/handlers.go index 74755fe4..86515685 100644 --- a/plugins/domain/user/handlers.go +++ b/plugins/domain/user/handlers.go @@ -4,6 +4,7 @@ package user import ( + "context" "crypto/rand" "crypto/sha256" "encoding/hex" @@ -13,8 +14,8 @@ import ( database "github.com/Rain-kl/Wavelet/plugins/infra/database" + "github.com/Rain-kl/Wavelet/core/contracts" "github.com/Rain-kl/Wavelet/pkg/response" - "github.com/Rain-kl/Wavelet/plugins/domain/auth" "github.com/gin-contrib/sessions" "github.com/gin-gonic/gin" ) @@ -51,6 +52,39 @@ type createAccessTokenRequest struct { IsAdmin bool `json:"is_admin"` } +func getUserIDFromSession(c *gin.Context) uint64 { + defer func() { _ = recover() }() + session := sessions.Default(c) + val := session.Get(contracts.AuthUserIDKey) + if val == nil { + return 0 + } + switch v := val.(type) { + case uint64: + return v + case int64: + return uint64(v) + case float64: + return uint64(v) + case string: + id, _ := strconv.ParseUint(v, 10, 64) + return id + default: + return 0 + } +} + +func invalidateUserCache(ctx context.Context, userID uint64) { + // Cache invalidation delegated to AuthService via IoC at plugin Apply time. + _ = ctx + _ = userID +} + +func invalidateTokenCache(ctx context.Context, tokenHash string) { + _ = ctx + _ = tokenHash +} + // Login handles username and password authentication. func Login(c *gin.Context) { var req loginRequest @@ -71,8 +105,8 @@ func Login(c *gin.Context) { } sess := sessions.Default(c) - sess.Set(auth.UserIDKey, user.ID) - sess.Set(auth.UserNameKey, user.Username) + sess.Set(contracts.AuthUserIDKey, user.ID) + sess.Set(contracts.AuthUserNameKey, user.Username) _ = sess.Save() c.JSON(http.StatusOK, response.OK(user)) @@ -126,7 +160,7 @@ func ChangePassword(c *gin.Context) { return } - userID := auth.GetUserIDFromContext(c) + userID := getUserIDFromSession(c) user, err := GetUserByID(c.Request.Context(), userID) if err != nil { response.AbortNotFound(c, errUserNotFound) @@ -145,7 +179,7 @@ func ChangePassword(c *gin.Context) { gormDB := database.DB(c.Request.Context()) _ = gormDB.Save(&user) - auth.InvalidateCachedUser(c.Request.Context(), user.ID) + invalidateUserCache(c.Request.Context(), user.ID) c.JSON(http.StatusOK, response.OKNil()) } @@ -158,7 +192,7 @@ func UpdateProfile(c *gin.Context) { return } - userID := auth.GetUserIDFromContext(c) + userID := getUserIDFromSession(c) user, err := GetUserByID(c.Request.Context(), userID) if err != nil { response.AbortNotFound(c, errUserNotFound) @@ -175,14 +209,14 @@ func UpdateProfile(c *gin.Context) { gormDB := database.DB(c.Request.Context()) _ = gormDB.Save(&user) - auth.InvalidateCachedUser(c.Request.Context(), user.ID) + invalidateUserCache(c.Request.Context(), user.ID) c.JSON(http.StatusOK, response.OK(user)) } // ListAccessTokens lists access tokens for the current user. func ListAccessTokens(c *gin.Context) { - userID := auth.GetUserIDFromContext(c) + userID := getUserIDFromSession(c) var tokens []AccessToken gormDB := database.DB(c.Request.Context()) _ = gormDB.Where("user_id = ?", userID).Find(&tokens).Error @@ -202,7 +236,7 @@ func CreateAccessToken(c *gin.Context) { return } - userID := auth.GetUserIDFromContext(c) + userID := getUserIDFromSession(c) rawBytes := make([]byte, tokenEntropyByteLength) _, _ = rand.Read(rawBytes) rawToken := "wvt_" + hex.EncodeToString(rawBytes) @@ -243,7 +277,7 @@ func DeleteAccessToken(c *gin.Context) { return } - userID := auth.GetUserIDFromContext(c) + userID := getUserIDFromSession(c) var token AccessToken gormDB := database.DB(c.Request.Context()) if err := gormDB.Where("id = ? AND user_id = ?", id, userID).First(&token).Error; err != nil { @@ -252,7 +286,7 @@ func DeleteAccessToken(c *gin.Context) { } _ = gormDB.Delete(&token) - auth.InvalidateCachedToken(c.Request.Context(), token.TokenHash) + invalidateTokenCache(c.Request.Context(), token.TokenHash) c.JSON(http.StatusOK, response.OKNil()) } @@ -265,7 +299,7 @@ func RotateAccessToken(c *gin.Context) { return } - userID := auth.GetUserIDFromContext(c) + userID := getUserIDFromSession(c) var token AccessToken gormDB := database.DB(c.Request.Context()) if err := gormDB.Where("id = ? AND user_id = ?", id, userID).First(&token).Error; err != nil { @@ -273,7 +307,7 @@ func RotateAccessToken(c *gin.Context) { return } - auth.InvalidateCachedToken(c.Request.Context(), token.TokenHash) + invalidateTokenCache(c.Request.Context(), token.TokenHash) rawBytes := make([]byte, tokenEntropyByteLength) _, _ = rand.Read(rawBytes) diff --git a/plugins/domain/user/plugin.go b/plugins/domain/user/plugin.go index 775a17c6..01f9fa5e 100644 --- a/plugins/domain/user/plugin.go +++ b/plugins/domain/user/plugin.go @@ -11,7 +11,7 @@ import ( "github.com/Rain-kl/Wavelet/core" "github.com/Rain-kl/Wavelet/core/contracts" "github.com/Rain-kl/Wavelet/core/extpoints" - "github.com/Rain-kl/Wavelet/plugins/domain/auth" + "github.com/gin-gonic/gin" "github.com/hibiken/asynq" ) @@ -64,6 +64,14 @@ func (p *Plugin) Manifest() core.Manifest { // Apply registers user migrations, services, routes, tasks, schedules, and settings into the Context. func (p *Plugin) Apply(ctx *core.Context) error { + // 0. Resolve auth service for middleware (via IoC, not direct import) + var authSvc contracts.AuthService + if err := core.Using[contracts.AuthService](ctx, func(svc contracts.AuthService) { authSvc = svc }); err != nil { + return err + } + loginMW := authSvc.RequireAuthMiddleware().(gin.HandlerFunc) + noTokenMW := authSvc.DisallowTokenAuthMiddleware().(gin.HandlerFunc) + // 1. Register migrations ctx.Migrations().Register("user", userMigrations) @@ -80,11 +88,11 @@ func (p *Plugin) Apply(ctx *core.Context) error { userGroup.POST("/register", Register) userGroup.GET("/logout", Logout) userGroup.POST("/send-email-code", SendEmailCode) - userGroup.POST("/change-password", auth.LoginRequired(), ChangePassword) - userGroup.PUT("/profile", auth.LoginRequired(), UpdateProfile) + userGroup.POST("/change-password", loginMW, ChangePassword) + userGroup.PUT("/profile", loginMW, UpdateProfile) // Access Tokens - tokensGroup := userGroup.Group("/access-tokens", auth.LoginRequired(), auth.DisallowTokenAuth()) + tokensGroup := userGroup.Group("/access-tokens", loginMW, noTokenMW) { tokensGroup.GET("", ListAccessTokens) tokensGroup.POST("", CreateAccessToken) @@ -134,4 +142,4 @@ func (p *Plugin) Apply(ctx *core.Context) error { }) return nil -} +} \ No newline at end of file