diff --git a/docs/changelog/index.md b/docs/changelog/index.md index b52f235e..bfa998bb 100644 --- a/docs/changelog/index.md +++ b/docs/changelog/index.md @@ -27,7 +27,11 @@ sidebar: false ### 修复 +- 修复上传模块(`upload`)的 API 错误响应绕过:将 `c.AbortWithStatus` 和自定义 `c.JSON` 错误响应统一替换为标准的 `response.Abort*` 辅助函数,确保响应格式符合全局信封规范。 +- 修复通用工具包(`pkg/utils`)中的网络和格式化工具 Bug:优化 `isPrivateIPv4` 使其通过 `net.ParseIP` 解析并调用标准库 `ip.IsPrivate()` 检查;修复 `Bytes2Size` 的边界判定,将大小限制变量(`sizeKB`、`sizeMB`、`sizeGB`)改为只读常量以增强不变性。 - 修复 CAP 模块路由错误响应:将 CAP 接口中的所有直接 JSON 错误响应改造为统一的 `response.Abort*` 抛出并挂载到中间件统一写出 JSON,保证全局 `{ "error_msg": "...", "data": null }` 信封规范。 +- 修复 ClickHouse 批量写入(`batchwriter` / `chwriter`)在服务退出时无法安全停机和刷出剩余日志的问题,统一在 Server 优雅停机流程中调用 `bootstrap.Stop()`。 +- 修复 OpenFlare 系统参数并发读取的数据竞争(data race)问题,在读取 OpenResty 配置快照、Agent 和 Relay 配置时引入 `OptionMapRWMutex` 读锁保护。 ### 移除 diff --git a/internal/apps/openflare/agent/helpers.go b/internal/apps/openflare/agent/helpers.go index ac83ae2b..194d2afa 100644 --- a/internal/apps/openflare/agent/helpers.go +++ b/internal/apps/openflare/agent/helpers.go @@ -206,6 +206,9 @@ func isPublicNodeIP(raw string) bool { } func buildAgentSettings(node *model.OpenFlareNode, updateNow bool, updateChannel string, updateTag string, restartOpenrestyNow bool) *Settings { + model.OptionMapRWMutex.RLock() + defer model.OptionMapRWMutex.RUnlock() + autoUpdate := false if node != nil { autoUpdate = node.AutoUpdateEnabled diff --git a/internal/apps/openflare/config_version/snapshot.go b/internal/apps/openflare/config_version/snapshot.go index f0ea71d4..bc1878d7 100644 --- a/internal/apps/openflare/config_version/snapshot.go +++ b/internal/apps/openflare/config_version/snapshot.go @@ -440,6 +440,8 @@ func convertPoWConfig(config *waf.PoWConfig) *openrestyrender.PoWConfig { } func buildOpenRestyConfigSnapshot() openRestyConfigSnapshot { + model.OptionMapRWMutex.RLock() + defer model.OptionMapRWMutex.RUnlock() return openRestyConfigSnapshot{ DefaultServerReturnStatus: model.OpenRestyDefaultServerReturnStatus, WorkerProcesses: model.OpenRestyWorkerProcesses, diff --git a/internal/apps/openflare/relay/helpers.go b/internal/apps/openflare/relay/helpers.go index 49fe80e4..e64abe18 100644 --- a/internal/apps/openflare/relay/helpers.go +++ b/internal/apps/openflare/relay/helpers.go @@ -98,6 +98,9 @@ func buildRelayConfig(node *model.OpenFlareNode) *Config { // BuildSettings returns runtime settings shared by relay and flared clients. func BuildSettings(node *model.OpenFlareNode, updateNow bool, updateChannel, updateTag string) *Settings { + model.OptionMapRWMutex.RLock() + defer model.OptionMapRWMutex.RUnlock() + autoUpdate := false if node != nil { autoUpdate = node.AutoUpdateEnabled diff --git a/internal/apps/risk_control/logics.go b/internal/apps/risk_control/logics.go index b2f969f9..23379421 100644 --- a/internal/apps/risk_control/logics.go +++ b/internal/apps/risk_control/logics.go @@ -62,6 +62,15 @@ func InitLogWriter(ctx context.Context) { logWriter = writer } +// 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() diff --git a/internal/apps/upload/filesrv/file_server.go b/internal/apps/upload/filesrv/file_server.go index 580f18ac..0599318e 100644 --- a/internal/apps/upload/filesrv/file_server.go +++ b/internal/apps/upload/filesrv/file_server.go @@ -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() }() diff --git a/internal/apps/upload/filesrv/file_server_test.go b/internal/apps/upload/filesrv/file_server_test.go index fe213d63..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{ @@ -63,7 +69,7 @@ func TestServeFileByIDAccessControl(t *testing.T) { ID: 8001, UserID: user.ID, FileName: "avatar.png", - FilePath: "uploads/avatar.png", + FilePath: "avatar.png", FileSize: 5, MimeType: "image/png", Extension: "png", @@ -75,7 +81,7 @@ func TestServeFileByIDAccessControl(t *testing.T) { ID: 8002, UserID: user.ID, FileName: "doc.pdf", - FilePath: "uploads/doc.pdf", + FilePath: "doc.pdf", FileSize: 5, MimeType: "application/pdf", Extension: "pdf", @@ -84,9 +90,12 @@ func TestServeFileByIDAccessControl(t *testing.T) { 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,8 +198,7 @@ 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) } @@ -201,7 +208,7 @@ func TestImageCompression(t *testing.T) { ID: 3001, UserID: user.ID, FileName: "test_image.png", - FilePath: filePath, + FilePath: "test_image.png", FileSize: int64(pngBuf.Len()), MimeType: "image/png", Extension: "png", @@ -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 e42defa2..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, "更新文件记录失败") 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/apps/upload/handler/routers_test.go b/internal/apps/upload/handler/routers_test.go index dea5a99a..e8ddce20 100644 --- a/internal/apps/upload/handler/routers_test.go +++ b/internal/apps/upload/handler/routers_test.go @@ -14,6 +14,7 @@ import ( "net/http" "net/http/httptest" "os" + "path/filepath" "strconv" "strings" "testing" @@ -29,6 +30,7 @@ import ( "github.com/Rain-kl/Wavelet/internal/storage" "github.com/Rain-kl/Wavelet/internal/testhelper" "github.com/gin-gonic/gin" + "gorm.io/gorm" ) type testResponse struct { @@ -104,7 +106,9 @@ func createMultipartRequest(t *testing.T, fieldName, fileName string, fileConten func TestUploadFile(t *testing.T) { dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) defer cleanup() - defer func() { _ = os.RemoveAll("uploads") }() // Clean up local files created during tests + + tempDir := t.TempDir() + configureLocalStorageRoot(t, dbConn, tempDir) authUser := &model.User{ID: 1001, Username: "test_user"} router := setupTestRouter(authUser) @@ -322,7 +326,7 @@ func TestUploadFile(t *testing.T) { } // Confirm file was actually written to local disk - fileContent, err := os.ReadFile(localRecord.FilePath) + fileContent, err := os.ReadFile(filepath.Join(tempDir, localRecord.FilePath)) if err != nil { t.Fatalf("failed to read local file: %v", err) } @@ -336,7 +340,9 @@ func TestUploadFile(t *testing.T) { func TestDownloadFile(t *testing.T) { dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) defer cleanup() - defer func() { _ = os.RemoveAll("uploads") }() + + tempDir := t.TempDir() + configureLocalStorageRoot(t, dbConn, tempDir) authUser := &model.User{ID: 1001, Username: "test_user"} router := setupTestRouter(authUser) @@ -346,7 +352,7 @@ func TestDownloadFile(t *testing.T) { ID: 2001, UserID: 1001, FileName: "中文文件名.txt", - FilePath: "uploads/test_download.txt", + FilePath: "test_download.txt", FileSize: 12, MimeType: "text/plain", Extension: "txt", @@ -354,11 +360,7 @@ func TestDownloadFile(t *testing.T) { } // Create local file - err := os.MkdirAll("uploads", 0755) - if err != nil { - t.Fatalf("failed to create directory: %v", err) - } - err = os.WriteFile(localUpload.FilePath, []byte("hello download"), 0644) + err := os.WriteFile(filepath.Join(tempDir, "test_download.txt"), []byte("hello download"), 0644) if err != nil { t.Fatalf("failed to write file: %v", err) } @@ -405,6 +407,9 @@ func TestListFiles(t *testing.T) { dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) defer cleanup() + tempDir := t.TempDir() + configureLocalStorageRoot(t, dbConn, tempDir) + authUser := &model.User{ID: 1001, Username: "test_user"} router := setupTestRouter(authUser) @@ -413,7 +418,7 @@ func TestListFiles(t *testing.T) { ID: 2101, UserID: authUser.ID, FileName: "first-report.txt", - FilePath: "uploads/first-report.txt", + FilePath: "first-report.txt", FileSize: 10, MimeType: "text/plain", Extension: "txt", @@ -423,7 +428,7 @@ func TestListFiles(t *testing.T) { ID: 2102, UserID: authUser.ID, FileName: "Second-Photo.PNG", - FilePath: "uploads/second-photo.png", + FilePath: "second-photo.png", FileSize: 20, MimeType: "image/png", Extension: "png", @@ -433,7 +438,7 @@ func TestListFiles(t *testing.T) { ID: 2103, UserID: authUser.ID, FileName: "third-notes.md", - FilePath: "uploads/third-notes.md", + FilePath: "third-notes.md", FileSize: 30, MimeType: "text/markdown", Extension: "md", @@ -443,7 +448,7 @@ func TestListFiles(t *testing.T) { ID: 2104, UserID: 2002, FileName: "other-user.txt", - FilePath: "uploads/other-user.txt", + FilePath: "other-user.txt", FileSize: 40, MimeType: "text/plain", Extension: "txt", @@ -544,20 +549,22 @@ func TestListFiles(t *testing.T) { func TestBatchDownloadFiles(t *testing.T) { dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) defer cleanup() - defer func() { _ = os.RemoveAll("uploads") }() + + tempDir := t.TempDir() + configureLocalStorageRoot(t, dbConn, tempDir) authUser := &model.User{ID: 1001, Username: "test_user"} router := setupTestRouter(authUser) - // Create and write files locally - err := os.MkdirAll("uploads", 0755) - if err != nil { - t.Fatalf("failed to create local dir: %v", err) + if err := os.WriteFile(filepath.Join(tempDir, "f1.txt"), []byte("file1 content"), 0644); err != nil { + t.Fatalf("failed to write f1.txt: %v", err) + } + if err := os.WriteFile(filepath.Join(tempDir, "f2.txt"), []byte("file2 content"), 0644); err != nil { + t.Fatalf("failed to write f2.txt: %v", err) + } + if err := os.WriteFile(filepath.Join(tempDir, "f3.txt"), []byte("duplicate name file content"), 0644); err != nil { + t.Fatalf("failed to write f3.txt: %v", err) } - - _ = os.WriteFile("uploads/f1.txt", []byte("file1 content"), 0644) - _ = os.WriteFile("uploads/f2.txt", []byte("file2 content"), 0644) - _ = os.WriteFile("uploads/f3.txt", []byte("duplicate name file content"), 0644) // Seed upload records. Note f2 and f3 have the same FileName "file_a.txt" to trigger name collision resolution. uploads := []model.Upload{ @@ -565,7 +572,7 @@ func TestBatchDownloadFiles(t *testing.T) { ID: 3001, UserID: 1001, FileName: "file_a.txt", - FilePath: "uploads/f1.txt", + FilePath: "f1.txt", FileSize: 13, MimeType: "text/plain", Extension: "txt", @@ -575,7 +582,7 @@ func TestBatchDownloadFiles(t *testing.T) { ID: 3002, UserID: 1001, FileName: "file_b.txt", - FilePath: "uploads/f2.txt", + FilePath: "f2.txt", FileSize: 13, MimeType: "text/plain", Extension: "txt", @@ -585,7 +592,7 @@ func TestBatchDownloadFiles(t *testing.T) { ID: 3003, UserID: 1001, FileName: "file_a.txt", // COLLISION with 3001! - FilePath: "uploads/f3.txt", + FilePath: "f3.txt", FileSize: 28, MimeType: "text/plain", Extension: "txt", @@ -656,7 +663,9 @@ func TestBatchDownloadFiles(t *testing.T) { func TestUploadAccessModeAccessControl(t *testing.T) { dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) defer cleanup() - defer func() { _ = os.RemoveAll("uploads") }() + + tempDir := t.TempDir() + configureLocalStorageRoot(t, dbConn, tempDir) user1 := &model.User{ID: 1001, Username: "user1"} user2 := &model.User{ID: 1002, Username: "user2"} @@ -744,6 +753,9 @@ func TestGetFileStats(t *testing.T) { dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) defer cleanup() + tempDir := t.TempDir() + configureLocalStorageRoot(t, dbConn, tempDir) + authUser := &model.User{ID: 1001, Username: "test_user"} router := setupTestRouter(authUser) @@ -753,7 +765,7 @@ func TestGetFileStats(t *testing.T) { ID: 3101, UserID: authUser.ID, FileName: "photo.png", - FilePath: "uploads/photo.png", + FilePath: "photo.png", FileSize: 100, MimeType: "image/png", Extension: "png", @@ -765,7 +777,7 @@ func TestGetFileStats(t *testing.T) { ID: 3102, UserID: authUser.ID, FileName: "video.mp4", - FilePath: "uploads/video.mp4", + FilePath: "video.mp4", FileSize: 500, MimeType: "video/mp4", Extension: "mp4", @@ -777,7 +789,7 @@ func TestGetFileStats(t *testing.T) { ID: 3103, UserID: authUser.ID, FileName: "document.pdf", - FilePath: "uploads/document.pdf", + FilePath: "document.pdf", FileSize: 200, MimeType: "application/pdf", Extension: "pdf", @@ -854,6 +866,9 @@ func TestUserUploadManagement(t *testing.T) { dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) defer cleanup() + tempDir := t.TempDir() + configureLocalStorageRoot(t, dbConn, tempDir) + user1 := &model.User{ID: 1001, Username: "user1"} user2 := &model.User{ID: 1002, Username: "user2"} @@ -868,7 +883,7 @@ func TestUserUploadManagement(t *testing.T) { ID: 4001, UserID: 1001, FileName: "user1-file.txt", - FilePath: "uploads/user1-file.txt", + FilePath: "user1-file.txt", FileSize: 100, MimeType: "text/plain", Extension: "txt", @@ -879,7 +894,7 @@ func TestUserUploadManagement(t *testing.T) { ID: 4002, UserID: 1002, FileName: "user2-file.png", - FilePath: "uploads/user2-file.png", + FilePath: "user2-file.png", FileSize: 200, MimeType: "image/png", Extension: "png", @@ -977,3 +992,26 @@ func TestUserUploadManagement(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/bootstrap/bootstrap.go b/internal/bootstrap/bootstrap.go index b7cf8796..d3e3cbde 100644 --- a/internal/bootstrap/bootstrap.go +++ b/internal/bootstrap/bootstrap.go @@ -89,6 +89,16 @@ func Init(ctx context.Context, opts Options) { }) } +// Stop stops all batch writers and background resources. +func Stop(ctx context.Context) { + if err := risk_control.StopLogWriter(ctx); err != nil { + logger.ErrorF(ctx, "[Bootstrap] stop risk_control log writer failed: %v", err) + } + if err := chwriter.Stop(ctx); err != nil { + logger.ErrorF(ctx, "[Bootstrap] stop chwriter failed: %v", err) + } +} + // ResetInitRuntimeOnceForTest clears initRuntimeOnce so Init can run again in unit tests. func ResetInitRuntimeOnceForTest() { initRuntimeOnce = sync.Once{} diff --git a/internal/router/router.go b/internal/router/router.go index bd133b88..135609b1 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" @@ -101,9 +102,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") diff --git a/pkg/utils/format.go b/pkg/utils/format.go index b21ae3b5..870ef88a 100644 --- a/pkg/utils/format.go +++ b/pkg/utils/format.go @@ -14,22 +14,24 @@ const ( secondsPerMinute = 60 ) -var sizeKB = 1024 -var sizeMB = sizeKB * 1024 -var sizeGB = sizeMB * 1024 +const ( + sizeKB = 1024 + sizeMB = sizeKB * 1024 + sizeGB = sizeMB * 1024 +) // Bytes2Size converts a byte count to a human-readable string with unit (B, KB, MB, GB). func Bytes2Size(num int64) string { numStr := "" unit := "B" switch { - case num/int64(sizeGB) > 1: + case num/int64(sizeGB) >= 1: numStr = fmt.Sprintf("%.2f", float64(num)/float64(sizeGB)) unit = "GB" - case num/int64(sizeMB) > 1: + case num/int64(sizeMB) >= 1: numStr = fmt.Sprintf("%d", int(float64(num)/float64(sizeMB))) unit = "MB" - case num/int64(sizeKB) > 1: + case num/int64(sizeKB) >= 1: numStr = fmt.Sprintf("%d", int(float64(num)/float64(sizeKB))) unit = "KB" default: diff --git a/pkg/utils/format_test.go b/pkg/utils/format_test.go new file mode 100644 index 00000000..dce34809 --- /dev/null +++ b/pkg/utils/format_test.go @@ -0,0 +1,49 @@ +package utils + +import ( + "testing" +) + +func TestBytes2Size(t *testing.T) { + tests := []struct { + input int64 + expected string + }{ + {0, "0 B"}, + {500, "500 B"}, + {1023, "1023 B"}, + {1024, "1 KB"}, + {2048, "2 KB"}, + {1024 * 1024, "1 MB"}, + {1024 * 1024 * 1024, "1.00 GB"}, + {1024 * 1024 * 1024 * 2, "2.00 GB"}, + } + + for _, tt := range tests { + result := Bytes2Size(tt.input) + if result != tt.expected { + t.Errorf("Bytes2Size(%d) = %q, expected %q", tt.input, result, tt.expected) + } + } +} + +func TestSeconds2Time(t *testing.T) { + tests := []struct { + input int + expected string + }{ + {0, "0 秒"}, + {30, "30 秒"}, + {60, "1 分钟 0 秒"}, + {125, "2 分钟 5 秒"}, + {3600, "1 小时 0 秒"}, + {86400, "1 天 0 秒"}, + } + + for _, tt := range tests { + result := Seconds2Time(tt.input) + if result != tt.expected { + t.Errorf("Seconds2Time(%d) = %q, expected %q", tt.input, result, tt.expected) + } + } +} diff --git a/pkg/utils/network.go b/pkg/utils/network.go index c1e7af67..483e20ae 100644 --- a/pkg/utils/network.go +++ b/pkg/utils/network.go @@ -3,7 +3,6 @@ package utils import ( "log/slog" "net" - "strings" ) // GetIP returns the first private IPv4 address found on the local network interfaces. @@ -35,7 +34,9 @@ func privateIPv4FromAddr(addr net.Addr) (string, bool) { } func isPrivateIPv4(ip string) bool { - return strings.HasPrefix(ip, "10") || - strings.HasPrefix(ip, "172") || - strings.HasPrefix(ip, "192.168") + parsedIP := net.ParseIP(ip) + if parsedIP == nil { + return false + } + return parsedIP.IsPrivate() } diff --git a/pkg/utils/network_test.go b/pkg/utils/network_test.go new file mode 100644 index 00000000..879eb610 --- /dev/null +++ b/pkg/utils/network_test.go @@ -0,0 +1,33 @@ +package utils + +import ( + "testing" +) + +func TestIsPrivateIPv4(t *testing.T) { + tests := []struct { + ip string + expected bool + }{ + {"127.0.0.1", false}, // Loopback is not in RFC 1918 private range + {"10.0.0.1", true}, + {"172.16.0.1", true}, + {"192.168.1.1", true}, + {"8.8.8.8", false}, + {"invalid-ip", false}, + } + + for _, tt := range tests { + result := isPrivateIPv4(tt.ip) + if result != tt.expected { + t.Errorf("isPrivateIPv4(%q) = %v, expected %v", tt.ip, result, tt.expected) + } + } +} + +func TestGetIP(t *testing.T) { + ip := GetIP() + // GetIP should return empty if no private IPv4 address is configured, or a valid IP. + // We just ensure it doesn't panic. + t.Logf("GetIP returned: %q", ip) +}