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: code-check:
@echo "==> Architecture guards..." @echo "==> Architecture guards..."
@command -v rg >/dev/null 2>&1 || { echo 'error: rg (ripgrep) is required for architecture guards' >&2; exit 1; } @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 " → core/ must not import gin, gorm, asynq..."
echo 'error: pkg/model must not access database.DB or cachepkg.Redis (non-test code)' >&2; \ @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; \ exit 1; \
fi 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 golangci-lint run
cd frontend && pnpm tsc --noEmit --jsx preserve && npx eslint . --max-warnings 0 cd frontend && pnpm tsc --noEmit --jsx preserve && npx eslint . --max-warnings 0
+12 -6
View File
@@ -8,8 +8,8 @@ import (
"time" "time"
"github.com/Rain-kl/Wavelet/core" "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/config"
"github.com/Rain-kl/Wavelet/pkg/migrator"
"github.com/Rain-kl/Wavelet/plugins/domain/admin" "github.com/Rain-kl/Wavelet/plugins/domain/admin"
"github.com/Rain-kl/Wavelet/plugins/domain/auth" "github.com/Rain-kl/Wavelet/plugins/domain/auth"
"github.com/Rain-kl/Wavelet/plugins/domain/cap" "github.com/Rain-kl/Wavelet/plugins/domain/cap"
@@ -56,11 +56,8 @@ func newWaveletApp(profile core.Profile) *core.App {
system.New(), system.New(),
) )
// 3. Bind Goose migration runner // 3. Bind Goose migration engine (wraps pkg/migrator.Migrate)
app.SetMigrationRunner(func(_ context.Context, _ []extpoints.MigrationEntry) error { app.SetMigrationEngine(&gooseEngine{})
runMigrations()
return nil
})
// 4. Mount runtime drivers for each aspect // 4. Mount runtime drivers for each aspect
app.Use( app.Use(
@@ -71,3 +68,12 @@ func newWaveletApp(profile core.Profile) *core.App {
return 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
}
+1 -24
View File
@@ -12,7 +12,6 @@ import (
"github.com/Rain-kl/Wavelet/pkg/buildinfo" "github.com/Rain-kl/Wavelet/pkg/buildinfo"
"github.com/Rain-kl/Wavelet/pkg/config" "github.com/Rain-kl/Wavelet/pkg/config"
"github.com/Rain-kl/Wavelet/pkg/logger" "github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/Rain-kl/Wavelet/pkg/migrator"
"github.com/Rain-kl/Wavelet/pkg/trace" "github.com/Rain-kl/Wavelet/pkg/trace"
"github.com/spf13/cobra" "github.com/spf13/cobra"
) )
@@ -38,9 +37,6 @@ var rootCmd = &cobra.Command{
TracerName: config.Config.Otel.TracerName, TracerName: config.Config.Otel.TracerName,
}) })
}, },
PreRun: func(_ *cobra.Command, _ []string) {
runMigrations()
},
PersistentPostRun: func(_ *cobra.Command, _ []string) { PersistentPostRun: func(_ *cobra.Command, _ []string) {
shutdownTraceProvider() 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() { func shutdownTraceProvider() {
ctx, cancel := context.WithTimeout(context.Background(), traceShutdownTimeout) ctx, cancel := context.WithTimeout(context.Background(), traceShutdownTimeout)
defer cancel() defer cancel()
@@ -70,16 +56,7 @@ func init() {
rootCmd.Version = buildinfo.Version rootCmd.Version = buildinfo.Version
rootCmd.CompletionOptions.DisableDefaultCmd = true rootCmd.CompletionOptions.DisableDefaultCmd = true
// 1. 为需要迁移的子命令动态绑定原先 rootCmd.PreRun 拥有的数据库迁移行为 // 集中将子命令注册到根命令,以解决 Cobra 的 unknown command 校验限制
migratePreRun := func(_ *cobra.Command, _ []string) {
runMigrations()
}
allCmd.PreRun = migratePreRun
apiCmd.PreRun = migratePreRun
workerCmd.PreRun = migratePreRun
schedulerCmd.PreRun = migratePreRun
// 2. 集中将这些命令注册为真正的子命令,以解决 Cobra 的 unknown command 校验限制
rootCmd.AddCommand(allCmd, apiCmd, workerCmd, schedulerCmd) rootCmd.AddCommand(allCmd, apiCmd, workerCmd, schedulerCmd)
} }
+18
View File
@@ -64,14 +64,23 @@ type AuthService interface {
// GetCurrentUser retrieves the authenticated UserDTO from context. // GetCurrentUser retrieves the authenticated UserDTO from context.
GetCurrentUser(ctx context.Context) (*UserDTO, error) 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 validates an access token and returns the associated user DTO.
VerifyToken(ctx context.Context, token string) (*UserDTO, error) VerifyToken(ctx context.Context, token string) (*UserDTO, error)
// CreateSession establishes an authenticated session for the given user ID. // CreateSession establishes an authenticated session for the given user ID.
CreateSession(ctx context.Context, userID uint64, extras map[string]any) (string, error) 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 revokes all active sessions and cached tokens for a user.
RevokeUserSessions(ctx context.Context, userID uint64) error 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. // AuthRegistry allows downstream and domain plugins to register custom authentication providers.
@@ -80,3 +89,12 @@ type AuthRegistry interface {
GetOAuthProvider(name string) (OAuthProvider, bool) GetOAuthProvider(name string) (OAuthProvider, bool)
ListOAuthProviders() []string 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 // Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0 // SPDX-License-Identifier: Apache-2.0
// Package contracts defines unified service interfaces and DTOs for cross-plugin communication.
package contracts 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 管理员登录领域事件载荷 // AdminLoggedIn 管理员登录领域事件载荷
type AdminLoggedIn struct { type AdminLoggedIn struct {
User *UserDTO `json:"user"` User *UserDTO `json:"user"`
IP string `json:"ip"` 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 return
} }
currUser, _ := auth.GetFromContext[*contracts.UserDTO](c, auth.UserObjKey) currUser, _ := util.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
if currUser == nil { if currUser == nil {
response.AbortUnauthorized(c, AdminRequired) response.AbortUnauthorized(c, AdminRequired)
return return
@@ -338,7 +338,7 @@ func UpdateUser(c *gin.Context) {
return return
} }
currUser, _ := auth.GetFromContext[*contracts.UserDTO](c, auth.UserObjKey) currUser, _ := util.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
if currUser == nil { if currUser == nil {
response.AbortUnauthorized(c, AdminRequired) response.AbortUnauthorized(c, AdminRequired)
return return
+6 -6
View File
@@ -7,26 +7,26 @@ import (
"github.com/Rain-kl/Wavelet/core/contracts" "github.com/Rain-kl/Wavelet/core/contracts"
"github.com/Rain-kl/Wavelet/pkg/logger" "github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/Rain-kl/Wavelet/pkg/response" "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/plugins/domain/auth" "github.com/Rain-kl/Wavelet/pkg/util"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
) )
// LoginAdminRequired 返回管理员权限校验中间件 // LoginAdminRequired 返回管理员权限校验中间件
func LoginAdminRequired() gin.HandlerFunc { func LoginAdminRequired() gin.HandlerFunc {
return func(c *gin.Context) { 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() defer span.End()
user, _ := auth.GetFromContext[*contracts.UserDTO](c, auth.UserObjKey) user, _ := util.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
if user == nil { if user == nil {
response.AbortNotFound(c, AdminRequired) response.AbortNotFound(c, AdminRequired)
return return
} }
// 如果是通过 Access Token 鉴权,需要检查令牌本身是否具有管理员权限 // 如果是通过 Access Token 鉴权,需要检查令牌本身是否具有管理员权限
if tokenAuth, _ := auth.GetFromContext[bool](c, auth.TokenAuthKey); tokenAuth { if tokenAuth, _ := util.GetFromContext[bool](c, contracts.AuthTokenAuthKey); tokenAuth {
tokenAdmin, _ := auth.GetFromContext[bool](c, auth.TokenAdminKey) tokenAdmin, _ := util.GetFromContext[bool](c, contracts.AuthTokenAdminKey)
if !tokenAdmin { if !tokenAdmin {
response.AbortNotFound(c, TokenAdminRequired) response.AbortNotFound(c, TokenAdminRequired)
return return
+11 -2
View File
@@ -8,8 +8,9 @@ import (
"context" "context"
"github.com/Rain-kl/Wavelet/core" "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/core/extpoints"
"github.com/Rain-kl/Wavelet/plugins/domain/auth" "github.com/gin-gonic/gin"
"github.com/hibiken/asynq" "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. // Apply registers admin routes, tasks, schedules, and settings into the Context.
func (p *Plugin) Apply(ctx *core.Context) error { 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 // 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 // Status & Diagnostics
adminRouter.GET("/status", GetSystemStatus) adminRouter.GET("/status", GetSystemStatus)
+1 -1
View File
@@ -430,7 +430,7 @@ func handleCallbackRegister(ctx context.Context, c *gin.Context, source *AuthSou
// UserInfo 获取当前登录用户信息 // UserInfo 获取当前登录用户信息
func UserInfo(c *gin.Context) { func UserInfo(c *gin.Context) {
user, _ := GetFromContext[*contracts.UserDTO](c, UserObjKey) user, _ := util.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
session := sessions.Default(c) session := sessions.Default(c)
needChange := session.Get("need_change_password") == true needChange := session.Get("need_change_password") == true
+15 -30
View File
@@ -12,27 +12,12 @@ import (
"github.com/Rain-kl/Wavelet/core/contracts" "github.com/Rain-kl/Wavelet/core/contracts"
"github.com/Rain-kl/Wavelet/pkg/response" "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" db "github.com/Rain-kl/Wavelet/plugins/infra/database"
"github.com/gin-gonic/gin" "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 { func hashToken(token string) string {
h := sha256.New() h := sha256.New()
h.Write([]byte(token)) h.Write([]byte(token))
@@ -91,8 +76,8 @@ func GetUserFromRequest(c *gin.Context) (*contracts.UserDTO, error) {
if user.Username == SystemUsername { if user.Username == SystemUsername {
return nil, errors.New("system user is not allowed to login") return nil, errors.New("system user is not allowed to login")
} }
SetToContext(c, TokenAuthKey, true) util.SetToContext(c, contracts.AuthTokenAuthKey, true)
SetToContext(c, TokenAdminKey, tokenRecord.IsAdmin) util.SetToContext(c, contracts.AuthTokenAdminKey, tokenRecord.IsAdmin)
return user, nil return user, nil
} }
} }
@@ -113,8 +98,8 @@ func GetUserFromRequest(c *gin.Context) (*contracts.UserDTO, error) {
SetCachedUser(ctx, userID, user) SetCachedUser(ctx, userID, user)
} }
SetToContext(c, TokenAuthKey, false) util.SetToContext(c, contracts.AuthTokenAuthKey, false)
SetToContext(c, TokenAdminKey, false) util.SetToContext(c, contracts.AuthTokenAdminKey, false)
if user.Username == "system" { if user.Username == "system" {
return nil, errors.New("system user is not allowed to login") 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 // LoginRequired 返回登录鉴权中间件,校验 Access Token 或 Session
func LoginRequired() gin.HandlerFunc { func LoginRequired() gin.HandlerFunc {
return func(c *gin.Context) { return func(c *gin.Context) {
ctx, span := otel_trace.Start(c.Request.Context(), "LoginRequired") _, span := trace.Start(c.Request.Context(), "LoginRequired")
defer span.End() defer span.End()
user, err := GetUserFromRequest(c) user, err := GetUserFromRequest(c)
@@ -135,8 +120,8 @@ func LoginRequired() gin.HandlerFunc {
return return
} }
LogForAudit(ctx, user, c) LogForAudit(c.Request.Context(), user, c)
SetToContext(c, UserObjKey, user) util.SetToContext(c, contracts.AuthUserObjKey, user)
c.Next() c.Next()
} }
} }
@@ -144,7 +129,7 @@ func LoginRequired() gin.HandlerFunc {
// AdminRequired 校验管理员权限(支持 Session 和 Token 鉴权) // AdminRequired 校验管理员权限(支持 Session 和 Token 鉴权)
func AdminRequired() gin.HandlerFunc { func AdminRequired() gin.HandlerFunc {
return func(c *gin.Context) { return func(c *gin.Context) {
ctx, span := otel_trace.Start(c.Request.Context(), "AdminRequired") _, span := trace.Start(c.Request.Context(), "AdminRequired")
defer span.End() defer span.End()
user, err := GetUserFromRequest(c) user, err := GetUserFromRequest(c)
@@ -153,8 +138,8 @@ func AdminRequired() gin.HandlerFunc {
return return
} }
isTokenAuth, _ := GetFromContext[bool](c, TokenAuthKey) isTokenAuth, _ := util.GetFromContext[bool](c, contracts.AuthTokenAuthKey)
isTokenAdmin, _ := GetFromContext[bool](c, TokenAdminKey) isTokenAdmin, _ := util.GetFromContext[bool](c, contracts.AuthTokenAdminKey)
// 如果是通过 Token 鉴权,要求该 Token 具备管理员权限或者用户本身为管理员 // 如果是通过 Token 鉴权,要求该 Token 具备管理员权限或者用户本身为管理员
if isTokenAuth && !isTokenAdmin && !user.IsAdmin { if isTokenAuth && !isTokenAdmin && !user.IsAdmin {
@@ -168,8 +153,8 @@ func AdminRequired() gin.HandlerFunc {
return return
} }
LogForAudit(ctx, user, c) LogForAudit(c.Request.Context(), user, c)
SetToContext(c, UserObjKey, user) util.SetToContext(c, contracts.AuthUserObjKey, user)
c.Next() c.Next()
} }
} }
@@ -182,7 +167,7 @@ func LoginAdminRequired() gin.HandlerFunc {
// DisallowTokenAuth 拒绝使用 Access Token 进行身份验证的请求访问该端点 // DisallowTokenAuth 拒绝使用 Access Token 进行身份验证的请求访问该端点
func DisallowTokenAuth() gin.HandlerFunc { func DisallowTokenAuth() gin.HandlerFunc {
return func(c *gin.Context) { 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) response.AbortForbidden(c, ErrTokenAuthNotAllowed)
return return
} }
+1 -1
View File
@@ -133,7 +133,7 @@ func TestAuthPluginUnit(t *testing.T) {
require.NoError(t, authSvc.RevokeUserSessions(context.Background(), user.ID)) require.NoError(t, authSvc.RevokeUserSessions(context.Background(), user.ID))
// GetCurrentUser from context // GetCurrentUser from context
userCtx := context.WithValue(context.Background(), auth.UserObjKey, userDTO) userCtx := context.WithValue(context.Background(), contracts.AuthUserObjKey, userDTO)
current, err := authSvc.GetCurrentUser(userCtx) current, err := authSvc.GetCurrentUser(userCtx)
require.NoError(t, err) require.NoError(t, err)
assert.Equal(t, user.ID, current.ID) assert.Equal(t, user.ID, current.ID)
+19 -2
View File
@@ -9,6 +9,7 @@ import (
"sync" "sync"
"github.com/Rain-kl/Wavelet/core/contracts" "github.com/Rain-kl/Wavelet/core/contracts"
"github.com/Rain-kl/Wavelet/pkg/util"
db "github.com/Rain-kl/Wavelet/plugins/infra/database" db "github.com/Rain-kl/Wavelet/plugins/infra/database"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
) )
@@ -29,12 +30,12 @@ func (s *authServiceImpl) RequireAdminMiddleware() any {
func (s *authServiceImpl) GetCurrentUser(ctx context.Context) (*contracts.UserDTO, error) { func (s *authServiceImpl) GetCurrentUser(ctx context.Context) (*contracts.UserDTO, error) {
if ginCtx, ok := ctx.(*gin.Context); ok { 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 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 { if u, ok := v.(*contracts.UserDTO); ok && u != nil {
return u, nil return u, nil
} }
@@ -93,6 +94,22 @@ func (s *authServiceImpl) RevokeUserSessions(ctx context.Context, userID uint64)
return nil 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 { type authRegistryImpl struct {
mu sync.RWMutex mu sync.RWMutex
providers map[string]contracts.OAuthProvider 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/core/contracts"
"github.com/Rain-kl/Wavelet/pkg/response" "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" "github.com/gin-gonic/gin"
) )
func currentUser(c *gin.Context) (*contracts.UserDTO, bool) { 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. // ListChannels lists enabled channels a user can bind.
@@ -136,8 +136,8 @@ func UnbindBinding(c *gin.Context) {
} }
// RegisterUserRoutes mounts user-facing message gateway endpoints. // RegisterUserRoutes mounts user-facing message gateway endpoints.
func RegisterUserRoutes(r *gin.RouterGroup) { func RegisterUserRoutes(r *gin.RouterGroup, loginMW gin.HandlerFunc) {
mg := r.Group("/message-gateway", auth.LoginRequired()) mg := r.Group("/message-gateway", loginMW)
{ {
mg.GET("/channels", ListChannels) mg.GET("/channels", ListChannels)
mg.GET("/bindings", ListBindings) mg.GET("/bindings", ListBindings)
+13 -4
View File
@@ -9,8 +9,9 @@ import (
"embed" "embed"
"github.com/Rain-kl/Wavelet/core" "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/core/extpoints"
"github.com/Rain-kl/Wavelet/plugins/domain/auth" "github.com/gin-gonic/gin"
"github.com/hibiken/asynq" "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. // Apply registers message_gateway migrations, routes, tasks, schedules, events, and settings into the Context.
func (p *Plugin) Apply(ctx *core.Context) error { 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 // 1. Register migrations
ctx.Migrations().Register("message_gateway", mgMigrations) ctx.Migrations().Register("message_gateway", mgMigrations)
// 2. Register User HTTP Routes // 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("/channels", ListChannels)
mgGroup.GET("/bindings", ListBindings) mgGroup.GET("/bindings", ListBindings)
@@ -83,7 +92,7 @@ func (p *Plugin) Apply(ctx *core.Context) error {
} }
// 3. Register Admin Message Gateway HTTP Routes // 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/definitions", ListAdminChannelDefinitions)
adminMgGroup.GET("/channels", ListAdminChannels) adminMgGroup.GET("/channels", ListAdminChannels)
@@ -94,7 +103,7 @@ func (p *Plugin) Apply(ctx *core.Context) error {
} }
// 4. Register Admin Push HTTP Routes // 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") 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/config"
"github.com/Rain-kl/Wavelet/pkg/idgen" "github.com/Rain-kl/Wavelet/pkg/idgen"
"github.com/Rain-kl/Wavelet/pkg/response" "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/Rain-kl/Wavelet/plugins/domain/risk_control/logstore"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
) )
@@ -42,7 +42,7 @@ func RiskControlMiddleware() gin.HandlerFunc {
c.Next() c.Next()
// 3. 后置身份检查:仅记录通过认证的请求 // 3. 后置身份检查:仅记录通过认证的请求
userObj, exists := auth.GetFromContext[*contracts.UserDTO](c, auth.UserObjKey) userObj, exists := util.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
if !exists || userObj == nil { if !exists || userObj == nil {
return return
} }
@@ -16,7 +16,7 @@ import (
"github.com/Rain-kl/Wavelet/pkg/batchwriter" "github.com/Rain-kl/Wavelet/pkg/batchwriter"
"github.com/Rain-kl/Wavelet/pkg/config" "github.com/Rain-kl/Wavelet/pkg/config"
"github.com/Rain-kl/Wavelet/pkg/testhelper" "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"
"github.com/Rain-kl/Wavelet/plugins/domain/risk_control/logstore" "github.com/Rain-kl/Wavelet/plugins/domain/risk_control/logstore"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
@@ -99,7 +99,7 @@ func TestRiskControlMiddleware(t *testing.T) {
r := gin.New() r := gin.New()
r.Use(func(c *gin.Context) { r.Use(func(c *gin.Context) {
user := &contracts.UserDTO{ID: 12345} user := &contracts.UserDTO{ID: 12345}
auth.SetToContext(c, auth.UserObjKey, user) util.SetToContext[*contracts.UserDTO](c, contracts.AuthUserObjKey, user)
c.Next() c.Next()
}) })
r.Use(risk_control.RiskControlMiddleware()) 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/core/contracts"
"github.com/Rain-kl/Wavelet/pkg/response" "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/auth"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/cache" "github.com/Rain-kl/Wavelet/plugins/domain/upload/cache"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/models" "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 { func checkPrivateFileOwner(c *gin.Context, ownerID uint64) error {
var currUserID uint64 var currUserID uint64
var isAdmin bool 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 currUserID = u.ID
isAdmin = u.IsAdmin isAdmin = u.IsAdmin
} else { } else {
@@ -311,7 +312,7 @@ func CheckFileAccessPermission(c *gin.Context, upload *models.Upload) error {
} }
if !cache.IsFilePublic(c.Request.Context(), upload.Type) { 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 { if _, err := auth.GetUserFromRequest(c); err != nil {
return err return err
} }
@@ -11,7 +11,7 @@ import (
"github.com/Rain-kl/Wavelet/core/contracts" "github.com/Rain-kl/Wavelet/core/contracts"
"github.com/Rain-kl/Wavelet/pkg/response" "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/ingest"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/models" "github.com/Rain-kl/Wavelet/plugins/domain/upload/models"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/repository" "github.com/Rain-kl/Wavelet/plugins/domain/upload/repository"
@@ -168,7 +168,7 @@ type listMyFilesResponse struct {
// @Failure 401 {object} response.Any "未登录" // @Failure 401 {object} response.Any "未登录"
// @Router /api/v1/upload/my [get] // @Router /api/v1/upload/my [get]
func ListMyFiles(c *gin.Context) { 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() ctx := c.Request.Context()
var req listMyFilesRequest var req listMyFilesRequest
@@ -215,7 +215,7 @@ func ListMyFiles(c *gin.Context) {
// @Failure 404 {object} response.Any "文件不存在" // @Failure 404 {object} response.Any "文件不存在"
// @Router /api/v1/upload/{id} [delete] // @Router /api/v1/upload/{id} [delete]
func DeleteMyFile(c *gin.Context) { 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() ctx := c.Request.Context()
if uploadstorage.ReadOnly(ctx) { if uploadstorage.ReadOnly(ctx) {
response.AbortConflict(c, shared.ErrStorageReadOnly) response.AbortConflict(c, shared.ErrStorageReadOnly)
@@ -262,7 +262,7 @@ type updateMyFileRequest struct {
// @Failure 404 {object} response.Any "文件不存在" // @Failure 404 {object} response.Any "文件不存在"
// @Router /api/v1/upload/{id} [put] // @Router /api/v1/upload/{id} [put]
func UpdateMyFile(c *gin.Context) { 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() ctx := c.Request.Context()
if uploadstorage.ReadOnly(ctx) { if uploadstorage.ReadOnly(ctx) {
response.AbortConflict(c, shared.ErrStorageReadOnly) 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/core/contracts"
"github.com/Rain-kl/Wavelet/pkg/logger" "github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/Rain-kl/Wavelet/pkg/response" "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/filesrv"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/ingest" "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/models"
@@ -63,7 +63,7 @@ func UploadFile(c *gin.Context) {
c.Request.Body = http.MaxBytesReader(c.Writer, c.Request.Body, shared.MaxUploadSize) 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() ctx := c.Request.Context()
header, err := c.FormFile("file") header, err := c.FormFile("file")
@@ -22,7 +22,7 @@ import (
"github.com/Rain-kl/Wavelet/core/contracts" "github.com/Rain-kl/Wavelet/core/contracts"
"github.com/Rain-kl/Wavelet/pkg/response" "github.com/Rain-kl/Wavelet/pkg/response"
"github.com/Rain-kl/Wavelet/pkg/testhelper" "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/models"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/shared" "github.com/Rain-kl/Wavelet/plugins/domain/upload/shared"
uploadstats "github.com/Rain-kl/Wavelet/plugins/domain/upload/stats" 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) { authMiddleware := func(c *gin.Context) {
if authUser != nil { if authUser != nil {
auth.SetToContext(c, auth.UserObjKey, authUser) util.SetToContext[*contracts.UserDTO](c, contracts.AuthUserObjKey, authUser)
} }
c.Next() c.Next()
} }
+11 -3
View File
@@ -8,11 +8,12 @@ import (
"context" "context"
"github.com/Rain-kl/Wavelet/core" "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/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/filesrv"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/handler" "github.com/Rain-kl/Wavelet/plugins/domain/upload/handler"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/task" "github.com/Rain-kl/Wavelet/plugins/domain/upload/task"
"github.com/gin-gonic/gin"
"github.com/hibiken/asynq" "github.com/hibiken/asynq"
) )
@@ -41,11 +42,18 @@ func (p *Plugin) Manifest() core.Manifest {
// Apply registers upload routes, tasks, and settings into the Context. // Apply registers upload routes, tasks, and settings into the Context.
func (p *Plugin) Apply(ctx *core.Context) error { 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 // 1. Register File Server Routes
ctx.Router().GET("/f/:id", filesrv.ServeFileByID) ctx.Router().GET("/f/:id", filesrv.ServeFileByID)
// 2. Register User/Admin Upload HTTP Routes // 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.POST("", handler.UploadFile)
uploadGroup.GET("", handler.ListFiles) uploadGroup.GET("", handler.ListFiles)
@@ -53,7 +61,7 @@ func (p *Plugin) Apply(ctx *core.Context) error {
uploadGroup.POST("/batch-download", handler.BatchDownloadFiles) 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("", handler.ListFiles)
adminUploadGroup.GET("/stats", handler.GetFileStats) adminUploadGroup.GET("/stats", handler.GetFileStats)
+47 -13
View File
@@ -4,6 +4,7 @@
package user package user
import ( import (
"context"
"crypto/rand" "crypto/rand"
"crypto/sha256" "crypto/sha256"
"encoding/hex" "encoding/hex"
@@ -13,8 +14,8 @@ import (
database "github.com/Rain-kl/Wavelet/plugins/infra/database" 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/pkg/response"
"github.com/Rain-kl/Wavelet/plugins/domain/auth"
"github.com/gin-contrib/sessions" "github.com/gin-contrib/sessions"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
) )
@@ -51,6 +52,39 @@ type createAccessTokenRequest struct {
IsAdmin bool `json:"is_admin"` 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. // Login handles username and password authentication.
func Login(c *gin.Context) { func Login(c *gin.Context) {
var req loginRequest var req loginRequest
@@ -71,8 +105,8 @@ func Login(c *gin.Context) {
} }
sess := sessions.Default(c) sess := sessions.Default(c)
sess.Set(auth.UserIDKey, user.ID) sess.Set(contracts.AuthUserIDKey, user.ID)
sess.Set(auth.UserNameKey, user.Username) sess.Set(contracts.AuthUserNameKey, user.Username)
_ = sess.Save() _ = sess.Save()
c.JSON(http.StatusOK, response.OK(user)) c.JSON(http.StatusOK, response.OK(user))
@@ -126,7 +160,7 @@ func ChangePassword(c *gin.Context) {
return return
} }
userID := auth.GetUserIDFromContext(c) userID := getUserIDFromSession(c)
user, err := GetUserByID(c.Request.Context(), userID) user, err := GetUserByID(c.Request.Context(), userID)
if err != nil { if err != nil {
response.AbortNotFound(c, errUserNotFound) response.AbortNotFound(c, errUserNotFound)
@@ -145,7 +179,7 @@ func ChangePassword(c *gin.Context) {
gormDB := database.DB(c.Request.Context()) gormDB := database.DB(c.Request.Context())
_ = gormDB.Save(&user) _ = gormDB.Save(&user)
auth.InvalidateCachedUser(c.Request.Context(), user.ID) invalidateUserCache(c.Request.Context(), user.ID)
c.JSON(http.StatusOK, response.OKNil()) c.JSON(http.StatusOK, response.OKNil())
} }
@@ -158,7 +192,7 @@ func UpdateProfile(c *gin.Context) {
return return
} }
userID := auth.GetUserIDFromContext(c) userID := getUserIDFromSession(c)
user, err := GetUserByID(c.Request.Context(), userID) user, err := GetUserByID(c.Request.Context(), userID)
if err != nil { if err != nil {
response.AbortNotFound(c, errUserNotFound) response.AbortNotFound(c, errUserNotFound)
@@ -175,14 +209,14 @@ func UpdateProfile(c *gin.Context) {
gormDB := database.DB(c.Request.Context()) gormDB := database.DB(c.Request.Context())
_ = gormDB.Save(&user) _ = gormDB.Save(&user)
auth.InvalidateCachedUser(c.Request.Context(), user.ID) invalidateUserCache(c.Request.Context(), user.ID)
c.JSON(http.StatusOK, response.OK(user)) c.JSON(http.StatusOK, response.OK(user))
} }
// ListAccessTokens lists access tokens for the current user. // ListAccessTokens lists access tokens for the current user.
func ListAccessTokens(c *gin.Context) { func ListAccessTokens(c *gin.Context) {
userID := auth.GetUserIDFromContext(c) userID := getUserIDFromSession(c)
var tokens []AccessToken var tokens []AccessToken
gormDB := database.DB(c.Request.Context()) gormDB := database.DB(c.Request.Context())
_ = gormDB.Where("user_id = ?", userID).Find(&tokens).Error _ = gormDB.Where("user_id = ?", userID).Find(&tokens).Error
@@ -202,7 +236,7 @@ func CreateAccessToken(c *gin.Context) {
return return
} }
userID := auth.GetUserIDFromContext(c) userID := getUserIDFromSession(c)
rawBytes := make([]byte, tokenEntropyByteLength) rawBytes := make([]byte, tokenEntropyByteLength)
_, _ = rand.Read(rawBytes) _, _ = rand.Read(rawBytes)
rawToken := "wvt_" + hex.EncodeToString(rawBytes) rawToken := "wvt_" + hex.EncodeToString(rawBytes)
@@ -243,7 +277,7 @@ func DeleteAccessToken(c *gin.Context) {
return return
} }
userID := auth.GetUserIDFromContext(c) userID := getUserIDFromSession(c)
var token AccessToken var token AccessToken
gormDB := database.DB(c.Request.Context()) gormDB := database.DB(c.Request.Context())
if err := gormDB.Where("id = ? AND user_id = ?", id, userID).First(&token).Error; err != nil { 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) _ = gormDB.Delete(&token)
auth.InvalidateCachedToken(c.Request.Context(), token.TokenHash) invalidateTokenCache(c.Request.Context(), token.TokenHash)
c.JSON(http.StatusOK, response.OKNil()) c.JSON(http.StatusOK, response.OKNil())
} }
@@ -265,7 +299,7 @@ func RotateAccessToken(c *gin.Context) {
return return
} }
userID := auth.GetUserIDFromContext(c) userID := getUserIDFromSession(c)
var token AccessToken var token AccessToken
gormDB := database.DB(c.Request.Context()) gormDB := database.DB(c.Request.Context())
if err := gormDB.Where("id = ? AND user_id = ?", id, userID).First(&token).Error; err != nil { 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 return
} }
auth.InvalidateCachedToken(c.Request.Context(), token.TokenHash) invalidateTokenCache(c.Request.Context(), token.TokenHash)
rawBytes := make([]byte, tokenEntropyByteLength) rawBytes := make([]byte, tokenEntropyByteLength)
_, _ = rand.Read(rawBytes) _, _ = rand.Read(rawBytes)
+12 -4
View File
@@ -11,7 +11,7 @@ import (
"github.com/Rain-kl/Wavelet/core" "github.com/Rain-kl/Wavelet/core"
"github.com/Rain-kl/Wavelet/core/contracts" "github.com/Rain-kl/Wavelet/core/contracts"
"github.com/Rain-kl/Wavelet/core/extpoints" "github.com/Rain-kl/Wavelet/core/extpoints"
"github.com/Rain-kl/Wavelet/plugins/domain/auth" "github.com/gin-gonic/gin"
"github.com/hibiken/asynq" "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. // Apply registers user migrations, services, routes, tasks, schedules, and settings into the Context.
func (p *Plugin) Apply(ctx *core.Context) error { 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 // 1. Register migrations
ctx.Migrations().Register("user", userMigrations) ctx.Migrations().Register("user", userMigrations)
@@ -80,11 +88,11 @@ func (p *Plugin) Apply(ctx *core.Context) error {
userGroup.POST("/register", Register) userGroup.POST("/register", Register)
userGroup.GET("/logout", Logout) userGroup.GET("/logout", Logout)
userGroup.POST("/send-email-code", SendEmailCode) userGroup.POST("/send-email-code", SendEmailCode)
userGroup.POST("/change-password", auth.LoginRequired(), ChangePassword) userGroup.POST("/change-password", loginMW, ChangePassword)
userGroup.PUT("/profile", auth.LoginRequired(), UpdateProfile) userGroup.PUT("/profile", loginMW, UpdateProfile)
// Access Tokens // Access Tokens
tokensGroup := userGroup.Group("/access-tokens", auth.LoginRequired(), auth.DisallowTokenAuth()) tokensGroup := userGroup.Group("/access-tokens", loginMW, noTokenMW)
{ {
tokensGroup.GET("", ListAccessTokens) tokensGroup.GET("", ListAccessTokens)
tokensGroup.POST("", CreateAccessToken) tokensGroup.POST("", CreateAccessToken)