diff --git a/internal/apps/risk_control/logics.go b/internal/apps/risk_control/logics.go index 067ad23c..db3f1079 100644 --- a/internal/apps/risk_control/logics.go +++ b/internal/apps/risk_control/logics.go @@ -5,88 +5,107 @@ package risk_control import ( "context" - "time" + "sync" "github.com/Rain-kl/Wavelet/internal/config" + "github.com/Rain-kl/Wavelet/internal/db/batchwriter" + "github.com/Rain-kl/Wavelet/internal/lifecycle" "github.com/Rain-kl/Wavelet/internal/model/analytics" analyticsrepo "github.com/Rain-kl/Wavelet/internal/repository/analytics" "github.com/Rain-kl/Wavelet/pkg/logger" ) -var logChan chan *analytics.UserAccessLog - -const ( - defaultQueueSize = 10000 - maxBatchSize = 1000 - flushInterval = 1 * time.Second +var ( + logWriterMu sync.RWMutex + logWriter *batchwriter.Writer[*analytics.UserAccessLog] ) -// InitLogWriter 初始化日志写入通道和后台写入协程 +// InitLogWriter initializes the ClickHouse access-log batch writer. func InitLogWriter(ctx context.Context) { if !config.Config.ClickHouse.Enabled { return } - logChan = make(chan *analytics.UserAccessLog, defaultQueueSize) - go startBatchWorker(context.WithoutCancel(ctx)) -} - -// IsBufferFull 检查当前本地缓冲队列是否已满 -// 如果没有启用 ClickHouse,默认返回 false,不触发限流 -func IsBufferFull() bool { - if !config.Config.ClickHouse.Enabled || logChan == nil { - return false - } - return len(logChan) >= cap(logChan) -} - -// QueueAccessLog 异步非阻塞地将日志推入缓冲队列 -func QueueAccessLog(logItem *analytics.UserAccessLog) { - if !config.Config.ClickHouse.Enabled || logChan == nil { + logWriterMu.Lock() + defer logWriterMu.Unlock() + if logWriter != nil { return } - select { - case logChan <- logItem: - default: - logger.WarnF(context.Background(), "[RiskControl] Log queue full, dropping log item for path: %s", logItem.Path) + cfg := batchwriter.DefaultConfig() + writer, err := batchwriter.New[*analytics.UserAccessLog](cfg, func(ctx context.Context, items []*analytics.UserAccessLog) error { + rows := make([]analytics.UserAccessLog, 0, len(items)) + for _, item := range items { + if item == nil { + continue + } + rows = append(rows, *item) + } + return analyticsrepo.BatchInsert(ctx, rows) + }, + batchwriter.WithDropHandler[*analytics.UserAccessLog](func(item *analytics.UserAccessLog) { + path := "" + if item != nil { + path = item.Path + } + logger.WarnF(context.Background(), "[RiskControl] Log queue full, dropping log item for path: %s", path) + }), + batchwriter.WithFlushErrorHandler[*analytics.UserAccessLog](func(ctx context.Context, batchSize int, err error) { + logger.ErrorF(ctx, "[RiskControl] Send ClickHouse batch failed (batch=%d): %v", batchSize, err) + }), + ) + if err != nil { + logger.ErrorF(ctx, "[RiskControl] init log writer failed: %v", err) + return + } + + writer.Start(ctx) + logWriter = writer + lifecycle.OnShutdown("risk_control_log_writer", StopLogWriter) +} + +// StopLogWriter stops the ClickHouse access-log batch writer and drains pending logs. +func StopLogWriter(ctx context.Context) error { + writer := currentLogWriter() + if writer == nil { + return nil + } + return writer.Stop(ctx) +} + +// IsBufferFull reports whether the access-log queue has no remaining capacity. +func IsBufferFull() bool { + writer := currentLogWriter() + if writer == nil { + return false + } + return writer.IsFull() +} + +// QueueAccessLog enqueues an access log without blocking. +func QueueAccessLog(logItem *analytics.UserAccessLog) { + writer := currentLogWriter() + if writer == nil || logItem == nil { + return + } + writer.TryEnqueue(logItem) +} + +// SetLogWriterForTest swaps the access-log writer for unit tests. +func SetLogWriterForTest(writer *batchwriter.Writer[*analytics.UserAccessLog]) func() { + logWriterMu.Lock() + previous := logWriter + logWriter = writer + logWriterMu.Unlock() + return func() { + logWriterMu.Lock() + logWriter = previous + logWriterMu.Unlock() } } -func startBatchWorker(ctx context.Context) { - ticker := time.NewTicker(flushInterval) - defer ticker.Stop() - - var batch []*analytics.UserAccessLog - - flush := func() { - if len(batch) == 0 { - return - } - - items := make([]analytics.UserAccessLog, len(batch)) - for i, item := range batch { - items[i] = *item - } - if err := analyticsrepo.BatchInsert(ctx, items); err != nil { - logger.ErrorF(ctx, "[RiskControl] Send ClickHouse batch failed: %v", err) - } - batch = nil - } - - for { - select { - case item, ok := <-logChan: - if !ok { - flush() - return - } - batch = append(batch, item) - if len(batch) >= maxBatchSize { - flush() - } - case <-ticker.C: - flush() - } - } -} \ No newline at end of file +func currentLogWriter() *batchwriter.Writer[*analytics.UserAccessLog] { + logWriterMu.RLock() + defer logWriterMu.RUnlock() + return logWriter +} diff --git a/internal/apps/upload/filesrv/file_server.go b/internal/apps/upload/filesrv/file_server.go index 330f85a3..0599318e 100644 --- a/internal/apps/upload/filesrv/file_server.go +++ b/internal/apps/upload/filesrv/file_server.go @@ -25,7 +25,7 @@ import ( "github.com/Rain-kl/Wavelet/internal/db" "github.com/Rain-kl/Wavelet/internal/diskcache" "github.com/Rain-kl/Wavelet/internal/model" - + "github.com/Rain-kl/Wavelet/pkg/logger" "github.com/gin-gonic/gin" "golang.org/x/sync/singleflight" @@ -66,14 +66,14 @@ func ServeFileByID(c *gin.Context) { upload, err := GetUploadRecordByID(c) if err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { - c.AbortWithStatus(http.StatusNotFound) + response.AbortNotFound(c, "文件记录未找到") return } if _, ok := err.(*strconv.NumError); ok { - c.JSON(http.StatusBadRequest, gin.H{"error": "Invalid upload ID"}) + response.AbortBadRequest(c, "无效的上传ID") return } - c.AbortWithStatus(http.StatusInternalServerError) + response.AbortInternal(c, "服务器内部错误") return } @@ -263,7 +263,7 @@ func ImageCompressionCacheKey(upload *model.Upload, quality string) string { func serveOriginal(c *gin.Context, upload *model.Upload) { obj, err := uploadstorage.OpenStoredObject(c.Request.Context(), upload) if err != nil { - c.AbortWithStatus(http.StatusNotFound) + response.AbortNotFound(c, "文件未找到") return } defer func() { _ = obj.Body.Close() }() @@ -313,4 +313,4 @@ func CheckFileAccessPermission(c *gin.Context, upload *model.Upload) error { } } return nil -} \ No newline at end of file +} diff --git a/internal/apps/upload/filesrv/file_server_test.go b/internal/apps/upload/filesrv/file_server_test.go index aa8b7027..8c33561d 100644 --- a/internal/apps/upload/filesrv/file_server_test.go +++ b/internal/apps/upload/filesrv/file_server_test.go @@ -6,6 +6,7 @@ package filesrv import ( "bytes" + "context" "encoding/json" "image" "image/color" @@ -13,6 +14,7 @@ import ( "net/http" "net/http/httptest" "os" + "path/filepath" "testing" "github.com/Rain-kl/Wavelet/internal/apps/upload/cache" @@ -20,12 +22,16 @@ import ( "github.com/Rain-kl/Wavelet/internal/apps/upload/util" "github.com/Rain-kl/Wavelet/internal/common" "github.com/Rain-kl/Wavelet/internal/common/response" + "github.com/Rain-kl/Wavelet/internal/db" "github.com/Rain-kl/Wavelet/internal/diskcache" "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/internal/storage" "github.com/Rain-kl/Wavelet/internal/testhelper" "github.com/gin-contrib/sessions" "github.com/gin-contrib/sessions/cookie" "github.com/gin-gonic/gin" + "gorm.io/gorm" ) func TestServeFileByIDAccessControl(t *testing.T) { @@ -33,8 +39,8 @@ func TestServeFileByIDAccessControl(t *testing.T) { defer cleanup() cache.ResetAccessCaches() - // Ensure uploads dir is cleaned up - defer func() { _ = os.RemoveAll("uploads") }() + tempDir := t.TempDir() + configureLocalStorageRoot(t, dbConn, tempDir) // Create a user in DB user := model.User{ @@ -60,33 +66,36 @@ func TestServeFileByIDAccessControl(t *testing.T) { // Create two files: one in whitelist (avatar), one not in whitelist (attachment) avatarFile := model.Upload{ - ID: 8001, - UserID: user.ID, - FileName: "avatar.png", - FilePath: "uploads/avatar.png", - FileSize: 5, - MimeType: "image/png", - Extension: "png", - Type: "avatar", - Status: model.UploadStatusUsed, - AccessMode: 1, + ID: 8001, + UserID: user.ID, + FileName: "avatar.png", + FilePath: "avatar.png", + FileSize: 5, + MimeType: "image/png", + Extension: "png", + Type: "avatar", + Status: model.UploadStatusUsed, + AccessMode: 1, } attachmentFile := model.Upload{ - ID: 8002, - UserID: user.ID, - FileName: "doc.pdf", - FilePath: "uploads/doc.pdf", - FileSize: 5, - MimeType: "application/pdf", - Extension: "pdf", - Type: "attachment", - Status: model.UploadStatusUsed, - AccessMode: 1, + ID: 8002, + UserID: user.ID, + FileName: "doc.pdf", + FilePath: "doc.pdf", + FileSize: 5, + MimeType: "application/pdf", + Extension: "pdf", + Type: "attachment", + Status: model.UploadStatusUsed, + AccessMode: 1, } - _ = os.MkdirAll("uploads", 0755) - _ = os.WriteFile(avatarFile.FilePath, []byte("image"), 0644) - _ = os.WriteFile(attachmentFile.FilePath, []byte("bytes"), 0644) + if err := os.WriteFile(filepath.Join(tempDir, "avatar.png"), []byte("image"), 0644); err != nil { + t.Fatalf("failed to write avatar file: %v", err) + } + if err := os.WriteFile(filepath.Join(tempDir, "doc.pdf"), []byte("bytes"), 0644); err != nil { + t.Fatalf("failed to write attachment file: %v", err) + } dbConn.Create(&avatarFile) dbConn.Create(&attachmentFile) @@ -159,15 +168,14 @@ func TestImageCompression(t *testing.T) { dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) defer cleanup() + tempDir := t.TempDir() + configureLocalStorageRoot(t, dbConn, tempDir) + cache := diskcache.GetGlobalCache() if err := cache.Clear(); err != nil { t.Fatalf("failed to clear disk cache before test: %v", err) } - // Ensure uploads dir is cleaned up - defer func() { - _ = os.RemoveAll("uploads") - }() defer func() { if err := cache.Clear(); err != nil { t.Errorf("failed to clear disk cache after test: %v", err) @@ -190,24 +198,23 @@ func TestImageCompression(t *testing.T) { t.Fatalf("failed to encode test png: %v", err) } - _ = os.MkdirAll("uploads", 0755) - filePath := "uploads/test_image.png" + filePath := filepath.Join(tempDir, "test_image.png") if err := os.WriteFile(filePath, pngBuf.Bytes(), 0644); err != nil { t.Fatalf("failed to write test png: %v", err) } // Save upload record to DB uploadRecord := model.Upload{ - ID: 3001, - UserID: user.ID, - FileName: "test_image.png", - FilePath: filePath, - FileSize: int64(pngBuf.Len()), - MimeType: "image/png", - Extension: "png", - Type: "avatar", // Whitelisted by default - Status: model.UploadStatusUsed, - AccessMode: 1, + ID: 3001, + UserID: user.ID, + FileName: "test_image.png", + FilePath: "test_image.png", + FileSize: int64(pngBuf.Len()), + MimeType: "image/png", + Extension: "png", + Type: "avatar", // Whitelisted by default + Status: model.UploadStatusUsed, + AccessMode: 1, } dbConn.Create(&uploadRecord) @@ -344,3 +351,26 @@ func TestNormalizeImageQuality(t *testing.T) { }) } } + +func configureLocalStorageRoot(t *testing.T, dbConn *gorm.DB, tempDir string) { + var sc model.SystemConfig + if err := dbConn.Where("key = ?", model.ConfigKeyStorageConfig).First(&sc).Error; err != nil { + t.Fatalf("failed to find storage config: %v", err) + } + var cfg storage.Config + if err := json.Unmarshal([]byte(sc.Value), &cfg); err != nil { + t.Fatalf("failed to unmarshal storage config: %v", err) + } + cfg.Local.Root = tempDir + newVal, err := json.Marshal(cfg) + if err != nil { + t.Fatalf("failed to marshal storage config: %v", err) + } + sc.Value = string(newVal) + if err := dbConn.Save(&sc).Error; err != nil { + t.Fatalf("failed to save storage config: %v", err) + } + _ = db.HSetJSON(context.Background(), repository.SystemConfigRedisHashKey, sc.Key, &sc) + repository.ResetSystemConfigRAMCacheForTest() + storage.ResetCache() +} diff --git a/internal/apps/upload/handler/file_management.go b/internal/apps/upload/handler/file_management.go index aa446184..075d72e8 100644 --- a/internal/apps/upload/handler/file_management.go +++ b/internal/apps/upload/handler/file_management.go @@ -112,7 +112,7 @@ func DeleteFile(c *gin.Context) { if _, err := softDeleteUpload(ctx, uploadID); err != nil { if isRecordNotFound(err) { - c.AbortWithStatus(http.StatusNotFound) + response.AbortNotFound(c, "文件记录未找到") return } response.AbortBadRequest(c, shared.ErrDeleteFileFailed) @@ -233,11 +233,11 @@ func DeleteMyFile(c *gin.Context) { if _, err := softDeleteOwnedUpload(ctx, currUser.ID, uploadID); err != nil { if isRecordNotFound(err) { - c.AbortWithStatus(http.StatusNotFound) + response.AbortNotFound(c, "文件记录未找到") return } if err == ingest.ErrForbidden { - c.AbortWithStatus(http.StatusForbidden) + response.AbortForbidden(c, "无权操作") return } response.AbortBadRequest(c, shared.ErrDeleteFileFailed) @@ -287,11 +287,11 @@ func UpdateMyFile(c *gin.Context) { upload, err := updateOwnedUpload(ctx, currUser.ID, uploadID, updateMyUploadInput(req)) if err != nil { if isRecordNotFound(err) { - c.AbortWithStatus(http.StatusNotFound) + response.AbortNotFound(c, "文件记录未找到") return } if err == ingest.ErrForbidden { - c.AbortWithStatus(http.StatusForbidden) + response.AbortForbidden(c, "无权操作") return } response.AbortBadRequest(c, "更新文件记录失败") @@ -299,4 +299,4 @@ func UpdateMyFile(c *gin.Context) { } c.JSON(http.StatusOK, response.OK(upload)) -} \ No newline at end of file +} diff --git a/internal/apps/upload/handler/routers.go b/internal/apps/upload/handler/routers.go index d85998f8..025052c4 100644 --- a/internal/apps/upload/handler/routers.go +++ b/internal/apps/upload/handler/routers.go @@ -170,7 +170,7 @@ func DownloadFile(c *gin.Context) { upload, err := filesrv.GetUploadRecordByID(c) if err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { - c.AbortWithStatus(http.StatusNotFound) + response.AbortNotFound(c, "文件记录未找到") return } if _, ok := err.(*strconv.NumError); ok { diff --git a/internal/bootstrap/bootstrap.go b/internal/bootstrap/bootstrap.go index b21af1fd..75b4de49 100644 --- a/internal/bootstrap/bootstrap.go +++ b/internal/bootstrap/bootstrap.go @@ -12,6 +12,7 @@ import ( admin_push "github.com/Rain-kl/Wavelet/internal/apps/admin/push" "github.com/Rain-kl/Wavelet/internal/apps/admin/push/custom_events" "github.com/Rain-kl/Wavelet/internal/apps/risk_control" + "github.com/Rain-kl/Wavelet/internal/lifecycle" taskhandlers "github.com/Rain-kl/Wavelet/internal/task/handlers" "github.com/Rain-kl/Wavelet/pkg/logger" ) @@ -87,7 +88,12 @@ func Init(ctx context.Context, opts Options) { }) } +// Stop stops all batch writers and background resources. +func Stop(ctx context.Context) { + lifecycle.Stop(ctx) +} + // ResetInitRuntimeOnceForTest clears initRuntimeOnce so Init can run again in unit tests. func ResetInitRuntimeOnceForTest() { initRuntimeOnce = sync.Once{} -} \ No newline at end of file +} diff --git a/internal/lifecycle/lifecycle.go b/internal/lifecycle/lifecycle.go new file mode 100644 index 00000000..edfe2d60 --- /dev/null +++ b/internal/lifecycle/lifecycle.go @@ -0,0 +1,66 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +// Package lifecycle manages global application and business component shutdown hooks. +package lifecycle + +import ( + "context" + "log" + "sync" +) + +// ShutdownFunc defines the signature for a graceful shutdown callback. +type ShutdownFunc func(ctx context.Context) error + +type hook struct { + name string + fn ShutdownFunc +} + +var ( + hooks []hook + mu sync.Mutex +) + +// OnShutdown registers a callback to be run during graceful shutdown. +func OnShutdown(name string, fn ShutdownFunc) { + mu.Lock() + defer mu.Unlock() + hooks = append(hooks, hook{name: name, fn: fn}) +} + +// Stop executes all registered shutdown hooks concurrently and waits for completion or context timeout. +func Stop(ctx context.Context) { + mu.Lock() + localHooks := make([]hook, len(hooks)) + copy(localHooks, hooks) + mu.Unlock() + + var wg sync.WaitGroup + for _, h := range localHooks { + wg.Add(1) + go func(name string, fn ShutdownFunc) { + defer wg.Done() + log.Printf("[Lifecycle] stopping %s...\n", name) + if err := fn(ctx); err != nil { + log.Printf("[Lifecycle] stop %s failed: %v\n", name, err) + } else { + log.Printf("[Lifecycle] %s stopped successfully\n", name) + } + }(h.name, h.fn) + } + + done := make(chan struct{}) + go func() { + wg.Wait() + close(done) + }() + + select { + case <-done: + log.Println("[Lifecycle] all services stopped gracefully") + case <-ctx.Done(): + log.Printf("[Lifecycle] shutdown timed out: %v\n", ctx.Err()) + } +} diff --git a/internal/router/router.go b/internal/router/router.go index c21fecac..74ba75e3 100644 --- a/internal/router/router.go +++ b/internal/router/router.go @@ -16,6 +16,7 @@ import ( "time" "github.com/Rain-kl/Wavelet/internal/apps/risk_control" + "github.com/Rain-kl/Wavelet/internal/bootstrap" router_root "github.com/Rain-kl/Wavelet/internal/router/root" v1 "github.com/Rain-kl/Wavelet/internal/router/v1" @@ -99,9 +100,11 @@ func Serve() { if err := srv.Shutdown(shutdownCtx); err != nil { log.Printf("[API] server forced to shutdown: %v\n", err) + bootstrap.Stop(shutdownCtx) cancel() os.Exit(1) } + bootstrap.Stop(shutdownCtx) cancel() log.Println("[API] server exited")