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
+24 -2
View File
@@ -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
+12 -6
View File
@@ -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
}
+2 -25
View File
@@ -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)
}
}
}
+18
View File
@@ -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
)
+98 -2
View File
@@ -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"`
}
+77
View File
@@ -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
}
```
@@ -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
}
+22
View File
@@ -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)
}
+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
}
}