refactor(core): fix cross-domain auth imports, unify migration, add downstream scaffold

Architecture:
- Move GetFromContext/SetToContext from plugins/domain/auth to pkg/util
- Move auth context key constants to core/contracts (AuthUserObjKey, AuthTokenAuthKey, etc.)
- Add AuthUserIDKey, AuthUserNameKey, GetCurrentUserID, RevokeToken to contracts.AuthService
- All 4 domain plugin Apply() methods now resolve AuthService via core.Using IoC
- Plugin route middleware uses authSvc.RequireAuthMiddleware() cast to gin.HandlerFunc
- DisallowTokenAuth added to AuthService contract

Migration:
- Replace cmd/app.go SetMigrationRunner bridge with gooseEngine implementing core.MigrationEngine
- Remove cmd/root.go PreRun migration hooks and runMigrations() function
- Migrations now run via core.App.Start() → RunMigrations()

Events:
- Add complete domain event topic catalog and payload DTOs to core/contracts/events.go
- 15 event topics across auth, user, admin, upload, message_gateway, risk_control

Downstream:
- Create downstream/ directory with README and custom_example plugin scaffold

CI:
- Update Makefile code-check architecture guards for Cordis layering
- Enforce: core no gin/gorm/asynq, contracts no plugins/, pkg no plugins/, domain no cross-domain
This commit is contained in:
ryan
2026-08-28 11:51:48 +08:00
parent 416603b616
commit c9b702d234
26 changed files with 460 additions and 123 deletions
+2 -2
View File
@@ -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
+6 -6
View File
@@ -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
+11 -2
View File
@@ -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)
+1 -1
View File
@@ -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
+16 -31
View File
@@ -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()
}
}
}
+1 -1
View File
@@ -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)
+19 -2
View File
@@ -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
+4 -4
View File
@@ -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)
+13 -4
View File
@@ -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")
{
+2 -2
View File
@@ -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
}
@@ -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())
+3 -2
View File
@@ -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
}
@@ -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)
+2 -2
View File
@@ -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")
@@ -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()
}
+11 -3
View File
@@ -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)
+47 -13
View File
@@ -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)
+13 -5
View File
@@ -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
}
}