From 11b0681a0de906cdc9de25ef62f41bdae9a8676e Mon Sep 17 00:00:00 2001 From: Kingcq <404291187@qq.com> Date: Thu, 18 Jun 2026 14:31:34 +0800 Subject: [PATCH 01/13] fix: make .env management more secure --- .env | 3 - .env.example | 5 + config/config.go | 301 ++++++++++++++++++++++++++++++++++++++++++++++- 3 files changed, 301 insertions(+), 8 deletions(-) delete mode 100644 .env create mode 100644 .env.example diff --git a/.env b/.env deleted file mode 100644 index cff9a1e..0000000 --- a/.env +++ /dev/null @@ -1,3 +0,0 @@ -PORT=3000 -# JWT secret -SECRET=TEST \ No newline at end of file diff --git a/.env.example b/.env.example new file mode 100644 index 0000000..d6ce1ff --- /dev/null +++ b/.env.example @@ -0,0 +1,5 @@ +PORT=3000 +# JWT secret +SECRET="replace-with-at-least-32-random-bytes" + +BOT_LOG_BUFFER_SIZE=1000 \ No newline at end of file diff --git a/config/config.go b/config/config.go index 96f3269..d7769ef 100644 --- a/config/config.go +++ b/config/config.go @@ -1,17 +1,308 @@ package config import ( - "log" + "crypto/rand" + "encoding/base64" + "errors" + "fmt" "os" + "path/filepath" + "strings" + "sync" + "time" "github.com/joho/godotenv" ) +const ( + configFileEnv = "NECORE_CONFIG_FILE" + appConfigDir = "necore" + secretBytes = 48 +) + +var ( + initOnce sync.Once + initErr error + configPath string +) + +// Init finds or creates the .env file and loads it exactly once. +func Init() error { + initOnce.Do(func() { + configPath, initErr = locateOrCreateConfigFile() + if initErr != nil { + return + } + + // Load does not overwrite variables already supplied by the process + // environment. This lets deployment-time environment variables take + // precedence over values stored in .env. + if err := godotenv.Load(configPath); err != nil { + initErr = fmt.Errorf("load config file %q: %w", configPath, err) + return + } + + setDefaultEnvironment("PORT", "3000") + setDefaultEnvironment("BOT_LOG_BUFFER_SIZE", "100") + + secret := strings.TrimSpace(os.Getenv("SECRET")) + if len(secret) < 32 { + initErr = fmt.Errorf("SECRET must contain at least 32 characters") + } + }) + + return initErr +} + +// Config keeps the existing project API. main should call Init before any +// server/database startup, so an error here indicates a programming error. func Config(key string) string { - // load .env file - err := godotenv.Load(".env") - if err != nil { - log.Println("Error loading .env file: ", err) + if err := Init(); err != nil { + panic(err) } return os.Getenv(key) } + +// Path returns the absolute path of the loaded .env file. +func Path() string { + return configPath +} + +func locateOrCreateConfigFile() (string, error) { + if explicit := strings.TrimSpace(os.Getenv(configFileEnv)); explicit != "" { + path, err := filepath.Abs(explicit) + if err != nil { + return "", fmt.Errorf("resolve %s: %w", configFileEnv, err) + } + if err := ensureConfigFile(path); err != nil { + return "", err + } + return path, nil + } + + existingCandidates, creationCandidates, err := defaultConfigCandidates() + if err != nil { + return "", err + } + + for _, path := range existingCandidates { + exists, err := regularFileExists(path) + if err != nil { + return "", err + } + if exists { + return path, nil + } + } + + var creationErrors []error + for _, path := range creationCandidates { + if err := ensureConfigFile(path); err == nil { + return path, nil + } else { + creationErrors = append(creationErrors, fmt.Errorf("%s: %w", path, err)) + } + } + + return "", fmt.Errorf("unable to create configuration file: %w", errors.Join(creationErrors...)) +} + +func defaultConfigCandidates() ([]string, []string, error) { + var existing []string + var creation []string + + // Development compatibility: only trust the current working directory + // when it looks like the project root. A production service may have an + // unrelated or attacker-controlled working directory. + if cwd, err := os.Getwd(); err == nil { + goMod := filepath.Join(cwd, "go.mod") + if ok, _ := regularFileExists(goMod); ok { + path := filepath.Join(cwd, ".env") + existing = append(existing, path) + creation = append(creation, path) + } + } + + executable, err := os.Executable() + if err != nil { + return nil, nil, fmt.Errorf("locate executable: %w", err) + } + if resolved, resolveErr := filepath.EvalSymlinks(executable); resolveErr == nil { + executable = resolved + } + executablePath := filepath.Join(filepath.Dir(executable), ".env") + existing = append(existing, executablePath) + creation = append(creation, executablePath) + + userConfigDir, err := os.UserConfigDir() + if err != nil { + return nil, nil, fmt.Errorf("locate user config directory: %w", err) + } + userConfigPath := filepath.Join(userConfigDir, appConfigDir, ".env") + existing = append(existing, userConfigPath) + creation = append(creation, userConfigPath) + + return uniquePaths(existing), uniquePaths(creation), nil +} + +func ensureConfigFile(path string) error { + exists, err := regularFileExists(path) + if err != nil { + return err + } + if exists { + return nil + } + + directory := filepath.Dir(path) + if err := os.MkdirAll(directory, 0o700); err != nil { + return fmt.Errorf("create config directory: %w", err) + } + + lockPath := path + ".lock" + deadline := time.Now().Add(5 * time.Second) + + for { + lockFile, err := os.OpenFile(lockPath, os.O_WRONLY|os.O_CREATE|os.O_EXCL, 0o600) + if err == nil { + return createConfigWhileLocked(path, lockPath, lockFile) + } + if !errors.Is(err, os.ErrExist) { + return fmt.Errorf("create config lock: %w", err) + } + + // Another process may have finished creating the file. + exists, existsErr := regularFileExists(path) + if existsErr != nil { + return existsErr + } + if exists { + return nil + } + + // Recover from a process that died while holding the creation lock. + if info, statErr := os.Stat(lockPath); statErr == nil && time.Since(info.ModTime()) > 30*time.Second { + _ = os.Remove(lockPath) + continue + } + + if time.Now().After(deadline) { + return fmt.Errorf("timed out waiting for config file creation") + } + time.Sleep(50 * time.Millisecond) + } +} + +func createConfigWhileLocked(path, lockPath string, lockFile *os.File) (returnErr error) { + defer func() { + _ = lockFile.Close() + _ = os.Remove(lockPath) + }() + + // Check again after acquiring the lock. + exists, err := regularFileExists(path) + if err != nil { + return err + } + if exists { + return nil + } + + secret, err := generateSecret() + if err != nil { + return err + } + + content := fmt.Sprintf( + "# Automatically generated by necore on first startup.\n"+ + "# Keep this file private. Changing SECRET invalidates existing JWTs.\n"+ + "PORT=3000\n"+ + "BOT_LOG_BUFFER_SIZE=2000\n"+ + "SECRET=%s\n", + secret, + ) + + directory := filepath.Dir(path) + temporary, err := os.CreateTemp(directory, ".env.tmp-*") + if err != nil { + return fmt.Errorf("create temporary config file: %w", err) + } + temporaryPath := temporary.Name() + defer func() { + _ = temporary.Close() + _ = os.Remove(temporaryPath) + }() + + if err := temporary.Chmod(0o600); err != nil { + return fmt.Errorf("set temporary config permissions: %w", err) + } + if _, err := temporary.WriteString(content); err != nil { + return fmt.Errorf("write temporary config file: %w", err) + } + if err := temporary.Sync(); err != nil { + return fmt.Errorf("sync temporary config file: %w", err) + } + if err := temporary.Close(); err != nil { + return fmt.Errorf("close temporary config file: %w", err) + } + + if err := os.Rename(temporaryPath, path); err != nil { + // If another actor created the file despite not using our lock, prefer + // the existing file rather than overwrite it. + if exists, existsErr := regularFileExists(path); existsErr == nil && exists { + return nil + } + return fmt.Errorf("publish config file: %w", err) + } + + // On Unix this protects the secret. On Windows the directory ACL remains + // the primary protection, but Chmod is still harmless. + if err := os.Chmod(path, 0o600); err != nil { + return fmt.Errorf("set config file permissions: %w", err) + } + + return nil +} + +func generateSecret() (string, error) { + buffer := make([]byte, secretBytes) + if _, err := rand.Read(buffer); err != nil { + return "", fmt.Errorf("generate SECRET: %w", err) + } + return base64.RawURLEncoding.EncodeToString(buffer), nil +} + +func setDefaultEnvironment(key, value string) { + if _, exists := os.LookupEnv(key); !exists { + _ = os.Setenv(key, value) + } +} + +func regularFileExists(path string) (bool, error) { + info, err := os.Stat(path) + if err == nil { + if !info.Mode().IsRegular() { + return false, fmt.Errorf("%q exists but is not a regular file", path) + } + return true, nil + } + if errors.Is(err, os.ErrNotExist) { + return false, nil + } + return false, fmt.Errorf("inspect %q: %w", path, err) +} + +func uniquePaths(paths []string) []string { + seen := make(map[string]struct{}, len(paths)) + result := make([]string, 0, len(paths)) + + for _, path := range paths { + clean := filepath.Clean(path) + if _, exists := seen[clean]; exists { + continue + } + seen[clean] = struct{}{} + result = append(result, clean) + } + return result +} From 5a2626e771915a3b2adf076c239d25fd71f678ec Mon Sep 17 00:00:00 2001 From: Kingcq <404291187@qq.com> Date: Thu, 18 Jun 2026 14:31:43 +0800 Subject: [PATCH 02/13] fix: make .env management more secure --- .gitignore | 7 ++++++- 1 file changed, 6 insertions(+), 1 deletion(-) diff --git a/.gitignore b/.gitignore index 896b331..a6b26e4 100644 --- a/.gitignore +++ b/.gitignore @@ -18,4 +18,9 @@ contents/* # Build -necore \ No newline at end of file +necore + +# dotenv +.env +.env.* +!.env.example \ No newline at end of file From 4fe7edf54b273dcef9b3fcebfbe17cad1ab61abd Mon Sep 17 00:00:00 2001 From: Kingcq <404291187@qq.com> Date: Thu, 18 Jun 2026 14:32:14 +0800 Subject: [PATCH 03/13] build: add required packages and upgrade go version to 1.25 --- go.mod | 28 ++++++++++++++++------------ go.sum | 42 ++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 58 insertions(+), 12 deletions(-) diff --git a/go.mod b/go.mod index 5c8d276..0eb1d0f 100644 --- a/go.mod +++ b/go.mod @@ -1,39 +1,43 @@ module necore -go 1.24.4 +go 1.25.0 require ( github.com/gofiber/contrib/jwt v1.1.2 - github.com/gofiber/fiber/v2 v2.52.11 + github.com/gofiber/contrib/websocket v1.3.4 + github.com/gofiber/fiber/v2 v2.52.13 github.com/golang-jwt/jwt/v5 v5.3.1 github.com/google/uuid v1.6.0 github.com/joho/godotenv v1.5.1 github.com/millkhan/mcstatusgo/v2 v2.2.0 - golang.org/x/crypto v0.48.0 + golang.org/x/crypto v0.53.0 gorm.io/driver/sqlite v1.6.0 gorm.io/gorm v1.31.1 ) require ( github.com/MicahParks/keyfunc/v2 v2.1.0 // indirect - github.com/andybalholm/brotli v1.2.0 // indirect + github.com/andybalholm/brotli v1.2.1 // indirect github.com/clipperhouse/uax29/v2 v2.7.0 // indirect + github.com/fasthttp/websocket v1.5.12 // indirect github.com/gofiber/storage/sqlite3/v2 v2.1.3 // indirect github.com/golang-jwt/jwt v3.2.2+incompatible // indirect github.com/jinzhu/inflection v1.0.0 // indirect github.com/jinzhu/now v1.1.5 // indirect - github.com/klauspost/compress v1.18.4 // indirect - github.com/mattn/go-colorable v0.1.14 // indirect - github.com/mattn/go-isatty v0.0.20 // indirect - github.com/mattn/go-runewidth v0.0.20 // indirect - github.com/mattn/go-sqlite3 v1.14.34 // indirect + github.com/klauspost/compress v1.18.6 // indirect + github.com/mattn/go-colorable v0.1.15 // indirect + github.com/mattn/go-isatty v0.0.22 // indirect + github.com/mattn/go-runewidth v0.0.24 // indirect + github.com/mattn/go-sqlite3 v1.14.45 // indirect github.com/rivo/uniseg v0.4.7 // indirect + github.com/savsgio/gotils v0.0.0-20250924091648-bce9a52d7761 // indirect github.com/tidwall/gjson v1.18.0 // indirect github.com/tidwall/match v1.1.1 // indirect github.com/tidwall/pretty v1.2.0 // indirect github.com/valyala/bytebufferpool v1.0.0 // indirect - github.com/valyala/fasthttp v1.69.0 // indirect + github.com/valyala/fasthttp v1.71.0 // indirect github.com/valyala/tcplisten v1.0.0 // indirect - golang.org/x/sys v0.41.0 // indirect - golang.org/x/text v0.34.0 // indirect + golang.org/x/net v0.55.0 // indirect + golang.org/x/sys v0.46.0 // indirect + golang.org/x/text v0.38.0 // indirect ) diff --git a/go.sum b/go.sum index 1af60fb..1be9bf8 100644 --- a/go.sum +++ b/go.sum @@ -4,16 +4,28 @@ github.com/andybalholm/brotli v1.1.0 h1:eLKJA0d02Lf0mVpIDgYnqXcUn0GqVmEFny3VuID1 github.com/andybalholm/brotli v1.1.0/go.mod h1:sms7XGricyQI9K10gOSf56VKKWS4oLer58Q+mhRPtnY= github.com/andybalholm/brotli v1.2.0 h1:ukwgCxwYrmACq68yiUqwIWnGY0cTPox/M94sVwToPjQ= github.com/andybalholm/brotli v1.2.0/go.mod h1:rzTDkvFWvIrjDXZHkuS16NPggd91W3kUSvPlQ1pLaKY= +github.com/andybalholm/brotli v1.2.1 h1:R+f5xP285VArJDRgowrfb9DqL18yVK0gKAW/F+eTWro= +github.com/andybalholm/brotli v1.2.1/go.mod h1:rzTDkvFWvIrjDXZHkuS16NPggd91W3kUSvPlQ1pLaKY= github.com/clipperhouse/uax29/v2 v2.7.0 h1:+gs4oBZ2gPfVrKPthwbMzWZDaAFPGYK72F0NJv2v7Vk= github.com/clipperhouse/uax29/v2 v2.7.0/go.mod h1:EFJ2TJMRUaplDxHKj1qAEhCtQPW2tJSwu5BF98AuoVM= +github.com/fasthttp/websocket v1.5.8 h1:k5DpirKkftIF/w1R8ZzjSgARJrs54Je9YJK37DL/Ah8= +github.com/fasthttp/websocket v1.5.8/go.mod h1:d08g8WaT6nnyvg9uMm8K9zMYyDjfKyj3170AtPRuVU0= +github.com/fasthttp/websocket v1.5.12 h1:e4RGPpWW2HTbL3zV0Y/t7g0ub294LkiuXXUuTOUInlE= +github.com/fasthttp/websocket v1.5.12/go.mod h1:I+liyL7/4moHojiOgUOIKEWm9EIxHqxZChS+aMFltyg= github.com/gofiber/contrib/jwt v1.1.2 h1:GmWnOqT4A15EkA8IPXwSpvNUXZR4u5SMj+geBmyLAjs= github.com/gofiber/contrib/jwt v1.1.2/go.mod h1:CpIwrkUQ3Q6IP8y9n3f0wP9bOnSKx39EDp2fBVgMFVk= +github.com/gofiber/contrib/websocket v1.3.4 h1:tWeBdbJ8q0WFQXariLN4dBIbGH9KBU75s0s7YXplOSg= +github.com/gofiber/contrib/websocket v1.3.4/go.mod h1:kTFBPC6YENCnKfKx0BoOFjgXxdz7E85/STdkmZPEmPs= github.com/gofiber/fiber/v2 v2.52.8 h1:xl4jJQ0BV5EJTA2aWiKw/VddRpHrKeZLF0QPUxqn0x4= github.com/gofiber/fiber/v2 v2.52.8/go.mod h1:YEcBbO/FB+5M1IZNBP9FO3J9281zgPAreiI1oqg8nDw= github.com/gofiber/fiber/v2 v2.52.9 h1:YjKl5DOiyP3j0mO61u3NTmK7or8GzzWzCFzkboyP5cw= github.com/gofiber/fiber/v2 v2.52.9/go.mod h1:YEcBbO/FB+5M1IZNBP9FO3J9281zgPAreiI1oqg8nDw= github.com/gofiber/fiber/v2 v2.52.11 h1:5f4yzKLcBcF8ha1GQTWB+mpblWz3Vz6nSAbTL31HkWs= github.com/gofiber/fiber/v2 v2.52.11/go.mod h1:YEcBbO/FB+5M1IZNBP9FO3J9281zgPAreiI1oqg8nDw= +github.com/gofiber/fiber/v2 v2.52.12 h1:0LdToKclcPOj8PktUdIKo9BUohjjwfnQl42Dhw8/WUw= +github.com/gofiber/fiber/v2 v2.52.12/go.mod h1:YEcBbO/FB+5M1IZNBP9FO3J9281zgPAreiI1oqg8nDw= +github.com/gofiber/fiber/v2 v2.52.13 h1:TOKP64iqC9b5P49VrBW5tHhUOvDyrtJ0xePEfzJbCbk= +github.com/gofiber/fiber/v2 v2.52.13/go.mod h1:YEcBbO/FB+5M1IZNBP9FO3J9281zgPAreiI1oqg8nDw= github.com/gofiber/storage/sqlite3/v2 v2.1.3 h1:zF8g/PQcCStF6WMJUYmy0Nq4Hj/LnW/vq/8OhRh6TZY= github.com/gofiber/storage/sqlite3/v2 v2.1.3/go.mod h1:KsEwBo6N7RGqEvI8XlcNclvrjdGab1RcmYL8jFCBRjc= github.com/golang-jwt/jwt v3.2.2+incompatible h1:IfV12K8xAKAnZqdXVzCZ+TOjboZ2keLg81eXfW3O+oY= @@ -24,6 +36,8 @@ github.com/golang-jwt/jwt/v5 v5.3.1 h1:kYf81DTWFe7t+1VvL7eS+jKFVWaUnK9cB1qbwn63Y github.com/golang-jwt/jwt/v5 v5.3.1/go.mod h1:fxCRLWMO43lRc8nhHWY6LGqRcf+1gQWArsqaEUEa5bE= github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0= github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= +github.com/gorilla/websocket v1.5.3 h1:saDtZ6Pbx/0u+bgYQ3q96pZgCzfhKXGPqt7kZ72aNNg= +github.com/gorilla/websocket v1.5.3/go.mod h1:YR8l580nyteQvAITg2hZ9XVh4b55+EU/adAjf1fMHhE= github.com/jinzhu/inflection v1.0.0 h1:K317FqzuhWc8YvSVlFMCCUb36O/S9MCKRDI7QkRKD/E= github.com/jinzhu/inflection v1.0.0/go.mod h1:h+uFLlag+Qp1Va5pdKtLDYj+kHp5pxUVkryuEj+Srlc= github.com/jinzhu/now v1.1.5 h1:/o9tlHleP7gOFmsnYNz3RGnqzefHA47wQpKrrdTIwXQ= @@ -36,29 +50,43 @@ github.com/klauspost/compress v1.18.0 h1:c/Cqfb0r+Yi+JtIEq73FWXVkRonBlf0CRNYc8Zt github.com/klauspost/compress v1.18.0/go.mod h1:2Pp+KzxcywXVXMr50+X0Q/Lsb43OQHYWRCY2AiWywWQ= github.com/klauspost/compress v1.18.4 h1:RPhnKRAQ4Fh8zU2FY/6ZFDwTVTxgJ/EMydqSTzE9a2c= github.com/klauspost/compress v1.18.4/go.mod h1:R0h/fSBs8DE4ENlcrlib3PsXS61voFxhIs2DeRhCvJ4= +github.com/klauspost/compress v1.18.6 h1:2jupLlAwFm95+YDR+NwD2MEfFO9d4z4Prjl1XXDjuao= +github.com/klauspost/compress v1.18.6/go.mod h1:cwPg85FWrGar70rWktvGQj8/hthj3wpl0PGDogxkrSQ= github.com/mattn/go-colorable v0.1.13 h1:fFA4WZxdEF4tXPZVKMLwD8oUnCTTo08duU7wxecdEvA= github.com/mattn/go-colorable v0.1.13/go.mod h1:7S9/ev0klgBDR4GtXTXX8a3vIGJpMovkB8vQcUbaXHg= github.com/mattn/go-colorable v0.1.14 h1:9A9LHSqF/7dyVVX6g0U9cwm9pG3kP9gSzcuIPHPsaIE= github.com/mattn/go-colorable v0.1.14/go.mod h1:6LmQG8QLFO4G5z1gPvYEzlUgJ2wF+stgPZH1UqBm1s8= +github.com/mattn/go-colorable v0.1.15 h1:+u9SLTRGnXv73cEsnsmoZBom+dMU88B2M0aDcWy0/jY= +github.com/mattn/go-colorable v0.1.15/go.mod h1:6LmQG8QLFO4G5z1gPvYEzlUgJ2wF+stgPZH1UqBm1s8= github.com/mattn/go-isatty v0.0.16/go.mod h1:kYGgaQfpe5nmfYZH+SKPsOc2e4SrIfOl2e/yFXSvRLM= github.com/mattn/go-isatty v0.0.20 h1:xfD0iDuEKnDkl03q4limB+vH+GxLEtL/jb4xVJSWWEY= github.com/mattn/go-isatty v0.0.20/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y= +github.com/mattn/go-isatty v0.0.22 h1:j8l17JJ9i6VGPUFUYoTUKPSgKe/83EYU2zBC7YNKMw4= +github.com/mattn/go-isatty v0.0.22/go.mod h1:ZXfXG4SQHsB/w3ZeOYbR0PrPwLy+n6xiMrJlRFqopa4= github.com/mattn/go-runewidth v0.0.16 h1:E5ScNMtiwvlvB5paMFdw9p4kSQzbXFikJ5SQO6TULQc= github.com/mattn/go-runewidth v0.0.16/go.mod h1:Jdepj2loyihRzMpdS35Xk/zdY8IAYHsh153qUoGf23w= github.com/mattn/go-runewidth v0.0.20 h1:WcT52H91ZUAwy8+HUkdM3THM6gXqXuLJi9O3rjcQQaQ= github.com/mattn/go-runewidth v0.0.20/go.mod h1:XBkDxAl56ILZc9knddidhrOlY5R/pDhgLpndooCuJAs= +github.com/mattn/go-runewidth v0.0.24 h1:cpokDiIn0MGnhdHwuWnJBITySJ20QyNGnY2kR/ay2DU= +github.com/mattn/go-runewidth v0.0.24/go.mod h1:XBkDxAl56ILZc9knddidhrOlY5R/pDhgLpndooCuJAs= github.com/mattn/go-sqlite3 v1.14.28 h1:ThEiQrnbtumT+QMknw63Befp/ce/nUPgBPMlRFEum7A= github.com/mattn/go-sqlite3 v1.14.28/go.mod h1:Uh1q+B4BYcTPb+yiD3kU8Ct7aC0hY9fxUwlHK0RXw+Y= github.com/mattn/go-sqlite3 v1.14.29 h1:1O6nRLJKvsi1H2Sj0Hzdfojwt8GiGKm+LOfLaBFaouQ= github.com/mattn/go-sqlite3 v1.14.29/go.mod h1:Uh1q+B4BYcTPb+yiD3kU8Ct7aC0hY9fxUwlHK0RXw+Y= github.com/mattn/go-sqlite3 v1.14.34 h1:3NtcvcUnFBPsuRcno8pUtupspG/GM+9nZ88zgJcp6Zk= github.com/mattn/go-sqlite3 v1.14.34/go.mod h1:Uh1q+B4BYcTPb+yiD3kU8Ct7aC0hY9fxUwlHK0RXw+Y= +github.com/mattn/go-sqlite3 v1.14.45 h1:6KA/spDguL3KV8rnybG7ezSaE4SeMR3KC9VbUoAQaIk= +github.com/mattn/go-sqlite3 v1.14.45/go.mod h1:pjEuOr8IwzLJP2MfGeTb0A35jauH+C2kbHKBr7yXKVQ= github.com/millkhan/mcstatusgo/v2 v2.2.0 h1:uRyHiOvqlK+6Oz3za4hMWAktSLjaqD/QzyQCcX2DjzI= github.com/millkhan/mcstatusgo/v2 v2.2.0/go.mod h1:YUJHhrJzsQP4PoDXFo++7JzPU7TjdztFM519CPqKe5M= github.com/rivo/uniseg v0.2.0 h1:S1pD9weZBuJdFmowNwbpi7BJ8TNftyUImj/0WQi72jY= github.com/rivo/uniseg v0.2.0/go.mod h1:J6wj4VEh+S6ZtnVlnTBMWIodfgj8LQOQFoIToxlJtxc= github.com/rivo/uniseg v0.4.7 h1:WUdvkW8uEhrYfLC4ZzdpI2ztxP1I582+49Oc5Mq64VQ= github.com/rivo/uniseg v0.4.7/go.mod h1:FN3SvrM+Zdj16jyLfmOkMNblXMcoc8DfTHruCPUcx88= +github.com/savsgio/gotils v0.0.0-20240303185622-093b76447511 h1:KanIMPX0QdEdB4R3CiimCAbxFrhB3j7h0/OvpYGVQa8= +github.com/savsgio/gotils v0.0.0-20240303185622-093b76447511/go.mod h1:sM7Mt7uEoCeFSCBM+qBrqvEo+/9vdmj19wzp3yzUhmg= +github.com/savsgio/gotils v0.0.0-20250924091648-bce9a52d7761 h1:McifyVxygw1d67y6vxUqls2D46J8W9nrki9c8c0eVvE= +github.com/savsgio/gotils v0.0.0-20250924091648-bce9a52d7761/go.mod h1:Vi9gvHvTw4yCUHIznFl5TPULS7aXwgaTByGeBY75Wko= github.com/tidwall/gjson v1.18.0 h1:FIDeeyB800efLX89e5a8Y0BNH+LOngJyGrIWxG2FKQY= github.com/tidwall/gjson v1.18.0/go.mod h1:/wbyibRr2FHMks5tjHJ5F8dMZh3AcwJEMf5vlfC0lxk= github.com/tidwall/match v1.1.1 h1:+Ho715JplO36QYgwN9PGYNhgZvoUSc9X2c80KVTi+GA= @@ -73,12 +101,20 @@ github.com/valyala/fasthttp v1.64.0 h1:QBygLLQmiAyiXuRhthf0tuRkqAFcrC42dckN2S+N3 github.com/valyala/fasthttp v1.64.0/go.mod h1:dGmFxwkWXSK0NbOSJuF7AMVzU+lkHz0wQVvVITv2UQA= github.com/valyala/fasthttp v1.69.0 h1:fNLLESD2SooWeh2cidsuFtOcrEi4uB4m1mPrkJMZyVI= github.com/valyala/fasthttp v1.69.0/go.mod h1:4wA4PfAraPlAsJ5jMSqCE2ug5tqUPwKXxVj8oNECGcw= +github.com/valyala/fasthttp v1.71.0 h1:tepR7H+Guh9VUqxxcPggYi8R3lGUu2Rsdh+z7/FCY3k= +github.com/valyala/fasthttp v1.71.0/go.mod h1:z1sDUvOShhXq/C9mwH/fSm1Vb71tUJwmQdgkBrBNwnA= github.com/valyala/tcplisten v1.0.0 h1:rBHj/Xf+E1tRGZyWIWwJDiRY0zc1Js+CV5DqwacVSA8= github.com/valyala/tcplisten v1.0.0/go.mod h1:T0xQ8SeCZGxckz9qRXTfG43PvQ/mcWh7FwZEA7Ioqkc= golang.org/x/crypto v0.40.0 h1:r4x+VvoG5Fm+eJcxMaY8CQM7Lb0l1lsmjGBQ6s8BfKM= golang.org/x/crypto v0.40.0/go.mod h1:Qr1vMER5WyS2dfPHAlsOj01wgLbsyWtFn/aY+5+ZdxY= golang.org/x/crypto v0.48.0 h1:/VRzVqiRSggnhY7gNRxPauEQ5Drw9haKdM0jqfcCFts= golang.org/x/crypto v0.48.0/go.mod h1:r0kV5h3qnFPlQnBSrULhlsRfryS2pmewsg+XfMgkVos= +golang.org/x/crypto v0.53.0 h1:QZ4Muo8THX6CizN2vPPd5fBGHyogrdK9fG4wLPFUsto= +golang.org/x/crypto v0.53.0/go.mod h1:DNLU434OwVakk9PzuwV8w62mAJpRJL3vsgcfp4Qnsio= +golang.org/x/net v0.49.0 h1:eeHFmOGUTtaaPSGNmjBKpbng9MulQsJURQUAfUwY++o= +golang.org/x/net v0.49.0/go.mod h1:/ysNB2EvaqvesRkuLAyjI1ycPZlQHM3q01F02UY/MV8= +golang.org/x/net v0.55.0 h1:bcvxaJn3e1U6InsFWt1JUq1aSjnRxLzT2rtD2KfkDF8= +golang.org/x/net v0.55.0/go.mod h1:L5U2KuzuOe1lY7Z+aWVIKK6qEeJXnXV9yzGA+WCHJww= golang.org/x/sys v0.0.0-20220811171246-fbc7d0a398ab/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.28.0 h1:Fksou7UEQUWlKvIdsqzJmUmCX3cZuD2+P3XyyzwMhlA= @@ -87,10 +123,16 @@ golang.org/x/sys v0.34.0 h1:H5Y5sJ2L2JRdyv7ROF1he/lPdvFsd0mJHFw2ThKHxLA= golang.org/x/sys v0.34.0/go.mod h1:BJP2sWEmIv4KK5OTEluFJCKSidICx8ciO85XgH3Ak8k= golang.org/x/sys v0.41.0 h1:Ivj+2Cp/ylzLiEU89QhWblYnOE9zerudt9Ftecq2C6k= golang.org/x/sys v0.41.0/go.mod h1:OgkHotnGiDImocRcuBABYBEXf8A9a87e/uXjp9XT3ks= +golang.org/x/sys v0.44.0 h1:ildZl3J4uzeKP07r2F++Op7E9B29JRUy+a27EibtBTQ= +golang.org/x/sys v0.44.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= +golang.org/x/sys v0.46.0 h1:noSf2Fq6F8DBgS+LysIkx7rIExoNHJsxOAtPp4rthXw= +golang.org/x/sys v0.46.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= golang.org/x/text v0.27.0 h1:4fGWRpyh641NLlecmyl4LOe6yDdfaYNrGb2zdfo4JV4= golang.org/x/text v0.27.0/go.mod h1:1D28KMCvyooCX9hBiosv5Tz/+YLxj0j7XhWjpSUF7CU= golang.org/x/text v0.34.0 h1:oL/Qq0Kdaqxa1KbNeMKwQq0reLCCaFtqu2eNuSeNHbk= golang.org/x/text v0.34.0/go.mod h1:homfLqTYRFyVYemLBFl5GgL/DWEiH5wcsQ5gSh1yziA= +golang.org/x/text v0.38.0 h1:sXmwo9DwP3OK9EZ7PqAdaooSGozfl/3a6/xJcbzPRhE= +golang.org/x/text v0.38.0/go.mod h1:YXZt3QhHUKYT53r2lLKFIVi6Ao1jdzrTR/KQ09qyxF4= gorm.io/driver/sqlite v1.6.0 h1:WHRRrIiulaPiPFmDcod6prc4l2VGVWHz80KspNsxSfQ= gorm.io/driver/sqlite v1.6.0/go.mod h1:AO9V1qIQddBESngQUKWL9yoH93HIeA1X6V633rBwyT8= gorm.io/gorm v1.30.1 h1:lSHg33jJTBxs2mgJRfRZeLDG+WZaHYCk3Wtfl6Ngzo4= From 47083b5500ce28572fbdcd5a2a219c16ba1387f8 Mon Sep 17 00:00:00 2001 From: Kingcq <404291187@qq.com> Date: Thu, 18 Jun 2026 14:34:12 +0800 Subject: [PATCH 04/13] feat: add bots management --- controller/router/router.go | 51 +++++++++- dao/bottoken.go | 67 ++++++++++++++ database/database.go | 11 +++ main.go | 8 ++ model/bottoken.go | 10 ++ service/bot.go | 66 +++++++++++++ service/bottoken.go | 67 ++++++++++++++ util/filepath.go | 58 ++++++++++++ util/token.go | 19 ++++ ws/hub.go | 179 ++++++++++++++++++++++++++++++++++++ 10 files changed, 535 insertions(+), 1 deletion(-) create mode 100644 dao/bottoken.go create mode 100644 model/bottoken.go create mode 100644 service/bot.go create mode 100644 service/bottoken.go create mode 100644 util/filepath.go create mode 100644 util/token.go create mode 100644 ws/hub.go diff --git a/controller/router/router.go b/controller/router/router.go index 7da330e..9c13812 100644 --- a/controller/router/router.go +++ b/controller/router/router.go @@ -1,10 +1,15 @@ package router import ( + "fmt" "necore/app" "necore/controller/middleware" + "necore/dao" "necore/service" + "necore/ws" + "strings" + "github.com/gofiber/contrib/websocket" "github.com/gofiber/fiber/v2" ) @@ -37,7 +42,7 @@ func SetupRoutes() { authGroup.Post("/register", middleware.AuthNeeded(), service.AddUser) authGroup.Get("/user/:id", service.GetUserInfo) authGroup.Get("/avatar/:id", service.GetUserAvatar) - authGroup.Get("/userlist", service.GetUserList) + authGroup.Get("/userlist", middleware.AuthNeeded(), service.GetUserList) authGroup.Delete("/user/:id", middleware.AuthNeeded(), service.DeleteUser) authGroup.Post("/password", middleware.AuthNeeded(), service.UpdateUserPassword) authGroup.Post("/avatar", middleware.AuthNeeded(), service.UpdateUserAvatar) @@ -73,4 +78,48 @@ func SetupRoutes() { documentGroup.Post("/upload/:id", middleware.AuthNeeded(), service.UploadDocumentFile) documentGroup.Delete("/upload/:id", middleware.AuthNeeded(), service.DeleteDocumentFile) (*router).Static("/contents", "./contents") + + botGroup := (*router).Group("/bots") + + botGroup.Use("/ws/updates", func(c *fiber.Ctx) error { + if !websocket.IsWebSocketUpgrade(c) { + return fiber.ErrUpgradeRequired + } + + auth := c.Get(fiber.HeaderAuthorization) + if !strings.HasPrefix(auth, "Bearer ") { + return c.SendStatus(fiber.StatusUnauthorized) + } + + identifier := c.Params("identifier") + if identifier == "" { + return c.SendStatus(fiber.StatusBadRequest) + } + + token := strings.TrimPrefix(auth, "Bearer ") + botToken, err := dao.GetBotTokenByPlainToken(token) + if err != nil { + ws.GlobalHub.AddLog( + fmt.Sprintf( + "⚠️ 拒绝 %s 连接:无效 Token", + identifier, + ), + ws.ERROR, + ) + return c.SendStatus(fiber.StatusUnauthorized) + } + + c.Locals("token_id", botToken.ID) + c.Locals("token_name", botToken.Name) + c.Locals("identifier", identifier) + return c.Next() + }) + botGroup.Get("/ws/updates", websocket.New(service.HandleWSConnection)) + + botGroup.Post("/token", middleware.AuthNeeded(), service.CreateBotToken) + botGroup.Get("/token", middleware.AuthNeeded(), service.GetBotTokenList) + botGroup.Get("/token/:id", middleware.AuthNeeded(), service.GetBotToken) + botGroup.Delete("/token/:id", middleware.AuthNeeded(), service.DeleteBotToken) + botGroup.Get("/status", middleware.AuthNeeded(), service.GetWSStatus) + botGroup.Delete("/ws/kick/:session_id", middleware.AuthNeeded(), service.KickConnection) } diff --git a/dao/bottoken.go b/dao/bottoken.go new file mode 100644 index 0000000..32924f7 --- /dev/null +++ b/dao/bottoken.go @@ -0,0 +1,67 @@ +package dao + +import ( + "crypto/sha256" + "encoding/hex" + "necore/database" + "necore/model" + "necore/util" +) + +func CreateBotToken(name string) (*model.BotToken, error) { + tokenStr, err := util.GenerateSecureToken("bot", 64) + if err != nil { + return nil, err + } + + sum := sha256.Sum256([]byte(tokenStr)) + tokenHash := hex.EncodeToString(sum[:]) + + newToken := model.BotToken{ + Name: name, + TokenHash: tokenHash, + } + + if err := database.GetBotTokenDatabase(). + Create(&newToken).Error; err != nil { + return nil, err + } + + return &newToken, nil +} + +func GetBotTokens() []model.BotToken { + var tokens []model.BotToken + db := database.GetBotTokenDatabase() + db.Find(&tokens) + return tokens +} + +func GetBotToken(name string) (*model.BotToken, error) { + var token model.BotToken + db := database.GetBotTokenDatabase() + if err := db.Where(&model.BotToken{Name: name}).First(&token).Error; err != nil { + return nil, err + } + return &token, nil +} + +func GetBotTokenByToken(token string) (*model.BotToken, error) { + var tokenModel model.BotToken + db := database.GetBotTokenDatabase() + if err := db.Where(&model.BotToken{TokenHash: token}).First(&tokenModel).Error; err != nil { + return nil, err + } + return &tokenModel, nil +} + +func GetBotTokenByPlainToken(token string) (*model.BotToken, error) { + sum := sha256.Sum256([]byte(token)) + tokenHash := hex.EncodeToString(sum[:]) + return GetBotTokenByToken(tokenHash) +} + +func DeleteBotToken(name string) error { + db := database.GetBotTokenDatabase() + return db.Where(&model.BotToken{Name: name}).Delete(&model.BotToken{}).Error +} diff --git a/database/database.go b/database/database.go index 59e5052..907cff3 100644 --- a/database/database.go +++ b/database/database.go @@ -17,6 +17,8 @@ var serverDatabase *gorm.DB var documentDatabase *gorm.DB +var botTokenDatabase *gorm.DB + func ConnectSqlite() { var err error userDatabase, err = gorm.Open(sqlite.Open("data/user.sqlite3"), &gorm.Config{}) @@ -43,6 +45,11 @@ func ConnectSqlite() { } documentDatabase.AutoMigrate(&model.DocumentNode{}) + botTokenDatabase, err = gorm.Open(sqlite.Open("data/bot_connection.sqlite3"), &gorm.Config{}) + if err != nil { + panic("failed to connect bot connection database") + } + botTokenDatabase.AutoMigrate(&model.BotToken{}) } func GetUserDatabase() *gorm.DB { @@ -60,3 +67,7 @@ func GetServerDatabase() *gorm.DB { func GetDocumentDatabase() *gorm.DB { return documentDatabase } + +func GetBotTokenDatabase() *gorm.DB { + return botTokenDatabase +} diff --git a/main.go b/main.go index 0f51e8c..c057479 100644 --- a/main.go +++ b/main.go @@ -1,7 +1,9 @@ package main import ( + "log" "necore/app" + "necore/config" "necore/controller/router" "necore/database" ) @@ -10,6 +12,12 @@ func main() { // This will print a hash of "test". U can insert it into sqlite3 manually for an admin account (the group section should be `["admin"]`). // dao.DebugTestPassword() + if err := config.Init(); err != nil { + log.Fatalf("initialize configuration: %v", err) + } + + log.Printf("using configuration file: %s", config.Path()) + database.ConnectSqlite() router.SetupRoutes() app.Start() diff --git a/model/bottoken.go b/model/bottoken.go new file mode 100644 index 0000000..cfd4ff0 --- /dev/null +++ b/model/bottoken.go @@ -0,0 +1,10 @@ +package model + +import "gorm.io/gorm" + +type BotToken struct { + gorm.Model + + Name string `gorm:"uniqueIndex;not null" json:"name"` + TokenHash string `gorm:"uniqueIndex;not null" json:"-"` +} diff --git a/service/bot.go b/service/bot.go new file mode 100644 index 0000000..f202545 --- /dev/null +++ b/service/bot.go @@ -0,0 +1,66 @@ +package service + +import ( + "necore/ws" + "time" + + "github.com/gofiber/contrib/websocket" + "github.com/gofiber/fiber/v2" + "github.com/google/uuid" +) + +func HandleWSConnection(c *websocket.Conn) { + tokenId := c.Locals("token_id").(uint) + tokenName := c.Locals("token_name").(string) + identifier := c.Locals("identifier").(string) + + sessionID := uuid.New().String() + client := &ws.Client{ + SessionID: sessionID, + Identifier: identifier, + TokenID: tokenId, + TokenName: tokenName, + Connected: time.Now().Format("2006-01-02 15:04:05"), + Conn: c, + } + + ws.GlobalHub.Register(client) + reason, unexpected := "正常退出", false + for { + if _, _, err := c.ReadMessage(); err != nil { + reason = "连接中断" + if err.Error() == "websocket: close sent" { + reason = "客户端主动断开" + } else if err.Error() == "websocket: close received" { + reason = "客户端被动断开" + } else if err.Error() == "websocket: bad handshake" { + reason = "握手失败" + } else if err.Error() == "websocket: unexpected EOF" { + reason = "连接中断" + unexpected = true + } else { + reason = "未知错误" + } + break + } + } + ws.GlobalHub.Unregister(sessionID, reason, unexpected) +} + +func GetWSStatus(c *fiber.Ctx) error { + clients, logs := ws.GlobalHub.GetDashboardStats() + return c.JSON(fiber.Map{ + "online_count": len(clients), + "connections": clients, + "logs": logs, + }) +} + +func KickConnection(c *fiber.Ctx) error { + if checkBotTokenPermission(c) { + return c.SendStatus(fiber.StatusForbidden) + } + sessionID := c.Params("session_id") + ws.GlobalHub.Unregister(sessionID, "强制断开连接", false) + return c.SendStatus(fiber.StatusOK) +} diff --git a/service/bottoken.go b/service/bottoken.go new file mode 100644 index 0000000..c06a592 --- /dev/null +++ b/service/bottoken.go @@ -0,0 +1,67 @@ +package service + +import ( + "necore/dao" + + "github.com/gofiber/fiber/v2" + "github.com/golang-jwt/jwt/v5" +) + +func checkBotTokenPermission(c *fiber.Ctx) bool { + token := c.Locals("user").(*jwt.Token) + isBotAdmin := dao.IsUserInGroup(token, "bot_admin") || dao.IsUserInGroup(token, "admin") + if isBotAdmin { + return false + } + return true +} + +func CreateBotToken(c *fiber.Ctx) error { + if checkBotTokenPermission(c) { + return c.Status(fiber.StatusForbidden).JSON(fiber.Map{"error": "Forbidden"}) + } + type request struct { + Name string `json:"name"` + } + r := new(request) + if err := c.BodyParser(r); err != nil { + return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{ + "error": "Invalid request", + }) + } + token, err := dao.CreateBotToken(r.Name) + if err != nil { + return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{"error": err}) + } + return c.JSON(fiber.Map{"token": token}) +} + +func GetBotToken(c *fiber.Ctx) error { + if checkBotTokenPermission(c) { + return c.Status(fiber.StatusForbidden).JSON(fiber.Map{"error": "Forbidden"}) + } + token, err := dao.GetBotToken(c.Params("id")) + if err != nil { + return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{"error": err}) + } + return c.JSON(fiber.Map{"token": token}) +} + +func GetBotTokenList(c *fiber.Ctx) error { + if checkBotTokenPermission(c) { + return c.Status(fiber.StatusForbidden).JSON(fiber.Map{"error": "Forbidden"}) + } + tokens := dao.GetBotTokens() + + return c.JSON(fiber.Map{"tokens": tokens}) +} + +func DeleteBotToken(c *fiber.Ctx) error { + if checkBotTokenPermission(c) { + return c.Status(fiber.StatusForbidden).JSON(fiber.Map{"error": "Forbidden"}) + } + if err := dao.DeleteBotToken(c.Params("id")); err != nil { + return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{"error": err}) + } + return c.SendStatus(fiber.StatusOK) +} diff --git a/util/filepath.go b/util/filepath.go new file mode 100644 index 0000000..cf6ba10 --- /dev/null +++ b/util/filepath.go @@ -0,0 +1,58 @@ +package util + +import ( + "errors" + "path" + "path/filepath" + "strings" +) + +var ErrInvalidFilename = errors.New("invalid filename") + +func SafeFilename(input string) (string, error) { + input = strings.TrimSpace(input) + if input == "" { + return "", ErrInvalidFilename + } + + normalized := strings.ReplaceAll(input, "\\", "/") + base := path.Base(normalized) + + if base != normalized || + base == "." || + base == ".." || + strings.ContainsRune(base, '\x00') { + return "", ErrInvalidFilename + } + + return base, nil +} + +func SafeContentPath(root, objectID, filename string) (string, error) { + safeName, err := SafeFilename(filename) + if err != nil { + return "", err + } + + baseDir, err := filepath.Abs(filepath.Join(root, objectID)) + if err != nil { + return "", err + } + + target, err := filepath.Abs(filepath.Join(baseDir, safeName)) + if err != nil { + return "", err + } + + relative, err := filepath.Rel(baseDir, target) + if err != nil { + return "", err + } + + if relative == ".." || + strings.HasPrefix(relative, ".."+string(filepath.Separator)) { + return "", ErrInvalidFilename + } + + return target, nil +} diff --git a/util/token.go b/util/token.go new file mode 100644 index 0000000..6c6d24c --- /dev/null +++ b/util/token.go @@ -0,0 +1,19 @@ +package util + +import ( + "crypto/rand" + "encoding/base64" + "fmt" +) + +func GenerateSecureToken(prefix string, byteLength int) (string, error) { + randomBytes := make([]byte, byteLength) + + if _, err := rand.Read(randomBytes); err != nil { + return "", fmt.Errorf("failed to generate secure token: %w", err) + } + + encoded := base64.RawURLEncoding.EncodeToString(randomBytes) + + return fmt.Sprintf("%s_%s", prefix, encoded), nil +} diff --git a/ws/hub.go b/ws/hub.go new file mode 100644 index 0000000..c78c7b1 --- /dev/null +++ b/ws/hub.go @@ -0,0 +1,179 @@ +package ws + +import ( + "fmt" + "necore/config" + "strconv" + "sync" + "time" + + "github.com/gofiber/contrib/websocket" +) + +func INFLogMsg(text string) string { + return "" + text + "" +} + +func SUCLogMsg(text string) string { + return "" + text + "" +} + +func WRNLogMsg(text string) string { + return "" + text + "" +} + +func ERRLogMsg(text string) string { + return "" + text + "" +} + +func DBGLogMsg(text string) string { + return "" + text + "" +} + +type Client struct { + SessionID string `json:"session_id"` + Identifier string `json:"identifier"` + TokenID uint `json:"token_id"` + TokenName string `json:"token_name"` + Connected string `json:"connected"` + Conn *websocket.Conn `json:"-"` +} + +type Hub struct { + Clients map[string]*Client + mu sync.RWMutex + + Logs []string + logMu sync.Mutex +} + +var GlobalHub = &Hub{ + Clients: make(map[string]*Client), + Logs: make([]string, 0), +} + +type LogLevel int + +const ( + DEBUG LogLevel = 0 + INFO LogLevel = 1 + WARNING LogLevel = 2 + ERROR LogLevel = 3 + SUCCESS LogLevel = 4 +) + +func (h *Hub) AddLog(msg string, level LogLevel) { + BOT_LOG_BUFFER_SIZE, _ := strconv.Atoi(config.Config("BOT_LOG_BUFFER_SIZE")) + h.logMu.Lock() + defer h.logMu.Unlock() + logLevelStr := "" + switch level { + case DEBUG: + logLevelStr = DBGLogMsg("DBG") + case INFO: + logLevelStr = INFLogMsg("INF") + case WARNING: + logLevelStr = WRNLogMsg("WRN") + case ERROR: + logLevelStr = ERRLogMsg("ERR") + case SUCCESS: + logLevelStr = SUCLogMsg("SUC") + } + message := fmt.Sprintf( + "[%v] %s | %s", + time.Now().Format("2006-01-02 15:04:05"), + logLevelStr, + msg, + ) + h.Logs = append(h.Logs, message) + if len(h.Logs) > BOT_LOG_BUFFER_SIZE { + h.Logs = h.Logs[:BOT_LOG_BUFFER_SIZE] + } +} + +func (h *Hub) Register(client *Client) { + h.mu.Lock() + defer h.mu.Unlock() + h.Clients[client.SessionID] = client + h.AddLog( + fmt.Sprintf( + "✅ %s 已连接,使用密钥:%s", + WRNLogMsg(client.Identifier), + INFLogMsg(client.TokenName), + ), + SUCCESS, + ) +} + +func (h *Hub) Unregister(sessionID, reason string, unexpected bool) { + h.mu.Lock() + defer h.mu.Unlock() + if client, ok := h.Clients[sessionID]; ok { + client.Conn.Close() + delete(h.Clients, sessionID) + if unexpected { + h.AddLog( + fmt.Sprintf( + "❌ %s 异常断开连接,原因:%s,使用密钥:%s", + WRNLogMsg(client.Identifier), + ERRLogMsg(reason), + INFLogMsg(client.TokenName), + ), + ERROR, + ) + } else { + h.AddLog( + fmt.Sprintf( + "❌ %s 断开连接,原因:%s,使用密钥:%s", + WRNLogMsg(client.Identifier), + ERRLogMsg(reason), + INFLogMsg(client.TokenName), + ), + INFO, + ) + } + } +} + +func (h *Hub) KickByTokenID(tokenID uint) { + h.mu.Lock() + defer h.mu.Unlock() + for sessionID, client := range h.Clients { + if client.TokenID == tokenID { + client.Conn.Close() + delete(h.Clients, sessionID) + h.AddLog( + fmt.Sprintf( + "⚠️ %s 因为密钥删除被踢出,使用密钥:%s", + WRNLogMsg(client.Identifier), + INFLogMsg(client.TokenName), + ), + WARNING, + ) + } + } +} + +func (h *Hub) Broadcast(message interface{}) { + h.mu.RLock() + defer h.mu.RUnlock() + for _, client := range h.Clients { + _ = client.Conn.WriteJSON(message) + } +} + +func (h *Hub) GetDashboardStats() ([]*Client, []string) { + h.mu.RLock() + clients := make([]*Client, 0, len(h.Clients)) + for _, c := range h.Clients { + clients = append(clients, c) + } + h.mu.RUnlock() + + h.logMu.Lock() + logsCopy := make([]string, len(h.Logs)) + copy(logsCopy, h.Logs) + h.logMu.Unlock() + + return clients, logsCopy +} From 8a603d7d84fcb1397568079d8600592e5f55af0b Mon Sep 17 00:00:00 2001 From: Kingcq <404291187@qq.com> Date: Thu, 18 Jun 2026 14:34:23 +0800 Subject: [PATCH 05/13] test: add unit test --- routes_test.go | 568 +++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 568 insertions(+) create mode 100644 routes_test.go diff --git a/routes_test.go b/routes_test.go new file mode 100644 index 0000000..9c92cae --- /dev/null +++ b/routes_test.go @@ -0,0 +1,568 @@ +package main + +import ( + "bytes" + "encoding/json" + "io" + "mime/multipart" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "strings" + "testing" + + "necore/controller/middleware" + "necore/dao" + "necore/database" + "necore/model" + "necore/service" + + "github.com/gofiber/contrib/websocket" + "github.com/gofiber/fiber/v2" + "gorm.io/gorm" + "gorm.io/gorm/logger" +) + +type testEnv struct { + app *fiber.App + adminToken string + userToken string + tmpDir string +} + +type testResponse struct { + StatusCode int + Body []byte + Header http.Header +} + +func setupTestEnv(t *testing.T) *testEnv { + t.Helper() + + tmpDir := t.TempDir() + must(t, os.MkdirAll(filepath.Join(tmpDir, "data"), 0o755)) + must(t, os.MkdirAll(filepath.Join(tmpDir, "contents"), 0o755)) + + envContent := "SECRET=unit-test-secret\nBOT_LOG_BUFFER_SIZE=100\n" + must(t, os.WriteFile(filepath.Join(tmpDir, ".env"), []byte(envContent), 0o600)) + + oldWd, err := os.Getwd() + must(t, err) + must(t, os.Chdir(tmpDir)) + t.Cleanup(func() { + _ = os.Chdir(oldWd) + }) + + t.Setenv("SECRET", "unit-test-secret") + t.Setenv("BOT_LOG_BUFFER_SIZE", "100") + + database.ConnectSqlite() + + // 测试会故意访问不存在的节点和 token;关闭 GORM 的 record not found 日志, + // 避免把预期分支误看成测试故障。 + setGormLoggerSilent(database.GetUserDatabase()) + setGormLoggerSilent(database.GetArticleDatabase()) + setGormLoggerSilent(database.GetServerDatabase()) + setGormLoggerSilent(database.GetDocumentDatabase()) + setGormLoggerSilent(database.GetBotTokenDatabase()) + + // 必须在 Windows 删除 TempDir 前关闭 SQLite 连接池,否则数据库文件会被锁定。 + t.Cleanup(func() { + closeGormDB(t, database.GetUserDatabase()) + closeGormDB(t, database.GetArticleDatabase()) + closeGormDB(t, database.GetServerDatabase()) + closeGormDB(t, database.GetDocumentDatabase()) + closeGormDB(t, database.GetBotTokenDatabase()) + }) + + must(t, dao.AddUserByUsername("admin", "admin-pass")) + must(t, dao.AddUserByUsername("alice", "alice-pass")) + + must(t, database.GetUserDatabase(). + Model(&model.User{}). + Where("username = ?", "admin"). + Updates(model.User{ + Group: `["admin","news_admin","server_admin","document_admin","bot_admin"]`, + Tags: `[]`, + Avatar: "admin-avatar", + }).Error) + + must(t, database.GetUserDatabase(). + Model(&model.User{}). + Where("username = ?", "alice"). + Updates(model.User{ + Group: `[]`, + Tags: `[]`, + Avatar: "alice-avatar", + }).Error) + + adminToken, err := dao.CreateToken(model.User{ + Username: "admin", + Group: `["admin","news_admin","server_admin","document_admin","bot_admin"]`, + Tags: `[]`, + }) + must(t, err) + + userToken, err := dao.CreateToken(model.User{ + Username: "alice", + Group: `[]`, + Tags: `[]`, + }) + must(t, err) + + app := fiber.New(fiber.Config{BodyLimit: 512 * 1024 * 1024}) + registerRoutes(app) + t.Cleanup(func() { + _ = app.Shutdown() + }) + + return &testEnv{ + app: app, + adminToken: adminToken, + userToken: userToken, + tmpDir: tmpDir, + } +} + +func setGormLoggerSilent(db *gorm.DB) { + if db != nil { + db.Logger = logger.Default.LogMode(logger.Silent) + } +} + +func closeGormDB(t *testing.T, db *gorm.DB) { + t.Helper() + if db == nil { + return + } + sqlDB, err := db.DB() + if err != nil { + t.Errorf("get underlying SQL DB: %v", err) + return + } + if err := sqlDB.Close(); err != nil { + t.Errorf("close SQL DB: %v", err) + } +} + +func registerRoutes(app *fiber.App) { + api := app.Group("/necore") + api.Get("/slogan", service.SloganHandler) + + authGroup := api.Group("/auth") + authGroup.Get("/status", middleware.AuthNeeded(), service.GetStatus) + authGroup.Post("/login", service.Login) + authGroup.Post("/register", middleware.AuthNeeded(), service.AddUser) + authGroup.Get("/user/:id", service.GetUserInfo) + authGroup.Get("/avatar/:id", service.GetUserAvatar) + authGroup.Get("/userlist", service.GetUserList) + authGroup.Delete("/user/:id", middleware.AuthNeeded(), service.DeleteUser) + authGroup.Post("/password", middleware.AuthNeeded(), service.UpdateUserPassword) + authGroup.Post("/avatar", middleware.AuthNeeded(), service.UpdateUserAvatar) + authGroup.Patch("/user", middleware.AuthNeeded(), service.UpdateUserInfo) + + articleGroup := api.Group("/news") + articleGroup.Get("/total/:target", service.GetArticleCountByCategory) + articleGroup.Post("/list", service.GetArticleList) + articleGroup.Get("/detail/:id", service.GetArticleById) + articleGroup.Patch("/:id", middleware.AuthNeeded(), service.UpdateArticle) + articleGroup.Post("/upload/:id", middleware.AuthNeeded(), service.UploadArticleFile) + articleGroup.Delete("/upload/:id", middleware.AuthNeeded(), service.DeleteArticleFile) + articleGroup.Post("/create", middleware.AuthNeeded(), service.CreateArticle) + articleGroup.Delete("/:id", middleware.AuthNeeded(), service.DeleteArticle) + + serverGroup := api.Group("/server") + serverGroup.Get("/", service.GetServerList) + serverGroup.Post("/status", service.GetServerStatus) + serverGroup.Get("/create", middleware.AuthNeeded(), service.AddServer) + serverGroup.Delete("/:id", middleware.AuthNeeded(), service.DeleteServer) + serverGroup.Patch("/", middleware.AuthNeeded(), service.UpdateServer) + + documentGroup := api.Group("/documents") + documentGroup.Delete("/node/:id", middleware.AuthNeeded(), service.DeleteDocumentNode) + documentGroup.Post("/node/:id", middleware.AuthNeeded(), service.UpdateDocumentNodeParentId) + documentGroup.Put("/node/:id", middleware.AuthNeeded(), service.UpdateDocumentNodeContent) + documentGroup.Patch("/node/:id", middleware.AuthNeeded(), service.UpdateDocumentNodeName) + documentGroup.Post("/node", middleware.AuthNeeded(), service.CreateDocumentNode) + documentGroup.Get("/layer/private/:parentId", middleware.AuthNeeded(), service.GetDocumentNodeChildrenPrivate) + documentGroup.Get("/layer/:parentId", service.GetDocumentNodeChildren) + documentGroup.Get("/private/:id", middleware.AuthNeeded(), service.GetDocumentNodeContentPrivate) + documentGroup.Get("/:id", service.GetDocumentNodeContent) + documentGroup.Post("/upload/:id", middleware.AuthNeeded(), service.UploadDocumentFile) + documentGroup.Delete("/upload/:id", middleware.AuthNeeded(), service.DeleteDocumentFile) + api.Static("/contents", "./contents") + + botGroup := api.Group("/bots") + botGroup.Post("/token", middleware.AuthNeeded(), service.CreateBotToken) + botGroup.Get("/token", middleware.AuthNeeded(), service.GetBotTokenList) + botGroup.Get("/token/:id", middleware.AuthNeeded(), service.GetBotToken) + botGroup.Delete("/token/:id", middleware.AuthNeeded(), service.DeleteBotToken) + botGroup.Get("/status", middleware.AuthNeeded(), service.GetWSStatus) + botGroup.Get("/ws/updates", websocket.New(service.HandleWSConnection)) + botGroup.Delete("/ws/kick/:session_id", middleware.AuthNeeded(), service.KickConnection) +} + +func must(t *testing.T, err error) { + t.Helper() + if err != nil { + t.Fatal(err) + } +} + +func doJSON(t *testing.T, env *testEnv, method, path, token string, body any) testResponse { + t.Helper() + + var r io.Reader + if body != nil { + b, err := json.Marshal(body) + must(t, err) + r = bytes.NewReader(b) + } + + req := httptest.NewRequest(method, path, r) + req.Header.Set("Content-Type", "application/json") + if token != "" { + req.Header.Set("Authorization", "Bearer "+token) + } + return executeRequest(t, env, req) +} + +func doRaw(t *testing.T, env *testEnv, method, path, token, contentType, body string) testResponse { + t.Helper() + + req := httptest.NewRequest(method, path, strings.NewReader(body)) + req.Header.Set("Content-Type", contentType) + if token != "" { + req.Header.Set("Authorization", "Bearer "+token) + } + return executeRequest(t, env, req) +} + +func doMultipartFile(t *testing.T, env *testEnv, path, token, field, filename, content string) testResponse { + t.Helper() + + var body bytes.Buffer + writer := multipart.NewWriter(&body) + part, err := writer.CreateFormFile(field, filename) + must(t, err) + _, err = part.Write([]byte(content)) + must(t, err) + must(t, writer.Close()) + + req := httptest.NewRequest(http.MethodPost, path, &body) + req.Header.Set("Content-Type", writer.FormDataContentType()) + if token != "" { + req.Header.Set("Authorization", "Bearer "+token) + } + return executeRequest(t, env, req) +} + +func executeRequest(t *testing.T, env *testEnv, req *http.Request) testResponse { + t.Helper() + + resp, err := env.app.Test(req, -1) + must(t, err) + + // 立即读取并关闭响应体。Fiber 的静态文件响应在 Windows 上会保持文件句柄; + // 不关闭会导致后续 os.Remove 和 TempDir 清理报 ERROR_SHARING_VIOLATION。 + body, readErr := io.ReadAll(resp.Body) + closeErr := resp.Body.Close() + must(t, readErr) + must(t, closeErr) + + return testResponse{ + StatusCode: resp.StatusCode, + Body: body, + Header: resp.Header.Clone(), + } +} + +func decodeBody(t *testing.T, resp testResponse) map[string]any { + t.Helper() + var got map[string]any + must(t, json.Unmarshal(resp.Body, &got)) + return got +} + +func assertStatus(t *testing.T, resp testResponse, want int) { + t.Helper() + if resp.StatusCode != want { + t.Fatalf("status = %d, want %d, body = %s", resp.StatusCode, want, string(resp.Body)) + } +} + +func TestPublicAndAuthRoutes(t *testing.T) { + env := setupTestEnv(t) + + assertStatus(t, doJSON(t, env, http.MethodGet, "/necore/slogan", "", nil), http.StatusOK) + + // 当前 jwt 中间件对无 Authorization 头实际返回 401,而不是 400。 + assertStatus(t, doJSON(t, env, http.MethodGet, "/necore/auth/status", "", nil), http.StatusUnauthorized) + assertStatus(t, doJSON(t, env, http.MethodGet, "/necore/auth/status", env.adminToken, nil), http.StatusOK) + + loginResp := doJSON(t, env, http.MethodPost, "/necore/auth/login", "", fiber.Map{ + "username": "admin", + "password": "admin-pass", + }) + assertStatus(t, loginResp, http.StatusOK) + loginBody := decodeBody(t, loginResp) + if loginBody["token"] == "" || loginBody["user"] == nil { + t.Fatalf("login response should contain token and user, got %#v", loginBody) + } + + assertStatus(t, doJSON(t, env, http.MethodPost, "/necore/auth/login", "", fiber.Map{ + "username": "admin", + "password": "wrong", + }), http.StatusUnauthorized) + + assertStatus(t, doJSON(t, env, http.MethodPost, "/necore/auth/register", env.userToken, fiber.Map{ + "username": "bob", + "password": "p", + }), http.StatusForbidden) + + assertStatus(t, doJSON(t, env, http.MethodPost, "/necore/auth/register", env.adminToken, fiber.Map{ + "username": "bob", + "password": "p", + }), http.StatusOK) + + assertStatus(t, doJSON(t, env, http.MethodGet, "/necore/auth/user/admin", "", nil), http.StatusOK) + assertStatus(t, doJSON(t, env, http.MethodGet, "/necore/auth/avatar/admin", "", nil), http.StatusOK) + assertStatus(t, doJSON(t, env, http.MethodGet, "/necore/auth/userlist", "", nil), http.StatusOK) + + assertStatus(t, doJSON(t, env, http.MethodPost, "/necore/auth/password", env.userToken, fiber.Map{ + "id": "admin", + "new_password": "new", + }), http.StatusForbidden) + + assertStatus(t, doJSON(t, env, http.MethodPost, "/necore/auth/password", env.userToken, fiber.Map{ + "id": "alice", + "new_password": "new", + }), http.StatusOK) + + assertStatus(t, doJSON(t, env, http.MethodPost, "/necore/auth/avatar", env.userToken, fiber.Map{ + "username": "admin", + "avatar": "x", + }), http.StatusForbidden) + + assertStatus(t, doJSON(t, env, http.MethodPost, "/necore/auth/avatar", env.userToken, fiber.Map{ + "username": "alice", + "avatar": "new-avatar", + }), http.StatusOK) + + assertStatus(t, doJSON(t, env, http.MethodPatch, "/necore/auth/user", env.userToken, fiber.Map{ + "username": "alice", + "group": []string{}, + "Tags": []any{}, + }), http.StatusForbidden) + + assertStatus(t, doJSON(t, env, http.MethodPatch, "/necore/auth/user", env.adminToken, fiber.Map{ + "username": "alice", + "group": []string{"document_admin"}, + "Tags": []any{}, + }), http.StatusOK) + + assertStatus(t, doJSON(t, env, http.MethodDelete, "/necore/auth/user/alice", env.userToken, nil), http.StatusForbidden) + assertStatus(t, doJSON(t, env, http.MethodDelete, "/necore/auth/user/bob", env.adminToken, nil), http.StatusOK) +} + +func TestNewsRoutes(t *testing.T) { + env := setupTestEnv(t) + + assertStatus(t, doJSON(t, env, http.MethodGet, "/necore/news/total/notice", "", nil), http.StatusOK) + assertStatus(t, doJSON(t, env, http.MethodPost, "/necore/news/list", "", fiber.Map{ + "target": "notice", + "page": 1, + "page_size": 10, + "pin": false, + }), http.StatusOK) + + assertStatus(t, doJSON(t, env, http.MethodPost, "/necore/news/create", env.userToken, nil), http.StatusForbidden) + + createResp := doJSON(t, env, http.MethodPost, "/necore/news/create", env.adminToken, nil) + assertStatus(t, createResp, http.StatusOK) + articleID, _ := decodeBody(t, createResp)["id"].(string) + if articleID == "" { + t.Fatal("create article should return id") + } + + assertStatus(t, doJSON(t, env, http.MethodPatch, "/necore/news/"+articleID, env.adminToken, fiber.Map{ + "entity": fiber.Map{ + "pin": true, + "title": "Title", + "brief": "Brief", + "date": "2026-01-01", + "endDate": "", + "image": "", + }, + "content": []fiber.Map{{"type": "markdown", "content": "hello"}}, + "category": "notice", + "doesNotify": false, + }), http.StatusOK) + + assertStatus(t, doJSON(t, env, http.MethodGet, "/necore/news/detail/"+articleID, "", nil), http.StatusOK) + assertStatus(t, doMultipartFile(t, env, "/necore/news/upload/"+articleID, env.adminToken, "file", "hello.txt", "hello"), http.StatusOK) + assertStatus(t, doJSON(t, env, http.MethodDelete, "/necore/news/upload/"+articleID, env.adminToken, fiber.Map{ + "url": "/contents/" + articleID + "/hello.txt", + }), http.StatusOK) + assertStatus(t, doJSON(t, env, http.MethodDelete, "/necore/news/"+articleID, env.adminToken, nil), http.StatusOK) +} + +func TestServerRoutes(t *testing.T) { + env := setupTestEnv(t) + + assertStatus(t, doJSON(t, env, http.MethodGet, "/necore/server/", "", nil), http.StatusOK) + assertStatus(t, doRaw(t, env, http.MethodPost, "/necore/server/status", "", "application/json", "{"), http.StatusBadRequest) + assertStatus(t, doJSON(t, env, http.MethodGet, "/necore/server/create", env.userToken, nil), http.StatusForbidden) + + createResp := doJSON(t, env, http.MethodGet, "/necore/server/create", env.adminToken, nil) + assertStatus(t, createResp, http.StatusOK) + serverID, _ := decodeBody(t, createResp)["id"].(string) + if serverID == "" { + t.Fatal("create server should return id") + } + + assertStatus(t, doJSON(t, env, http.MethodPatch, "/necore/server/", env.adminToken, fiber.Map{ + "id": serverID, + "name": "S1", + "icon": "icon", + "description": "desc", + "realtime": false, + "onlineMapUrl": "https://example.test/map", + "serverUrl": "example.test:25565", + }), http.StatusOK) + assertStatus(t, doJSON(t, env, http.MethodDelete, "/necore/server/"+serverID, env.adminToken, nil), http.StatusOK) +} + +func TestDocumentRoutes(t *testing.T) { + env := setupTestEnv(t) + + assertStatus(t, doJSON(t, env, http.MethodPost, "/necore/documents/node", env.userToken, fiber.Map{ + "parentId": "root", + "isFolder": true, + "private": false, + "name": "Root", + }), http.StatusForbidden) + + parentResp := doJSON(t, env, http.MethodPost, "/necore/documents/node", env.adminToken, fiber.Map{ + "parentId": "root", + "isFolder": true, + "private": false, + "name": "Destination", + }) + assertStatus(t, parentResp, http.StatusOK) + parentID, _ := decodeBody(t, parentResp)["id"].(string) + if parentID == "" { + t.Fatal("create parent document node should return id") + } + + createResp := doJSON(t, env, http.MethodPost, "/necore/documents/node", env.adminToken, fiber.Map{ + "parentId": "root", + "isFolder": false, + "private": false, + "name": "Doc", + }) + assertStatus(t, createResp, http.StatusOK) + nodeID, _ := decodeBody(t, createResp)["id"].(string) + if nodeID == "" { + t.Fatal("create document node should return id") + } + + assertStatus(t, doJSON(t, env, http.MethodGet, "/necore/documents/layer/root", "", nil), http.StatusOK) + assertStatus(t, doJSON(t, env, http.MethodGet, "/necore/documents/layer/private/root", env.adminToken, nil), http.StatusOK) + assertStatus(t, doJSON(t, env, http.MethodGet, "/necore/documents/"+nodeID, "", nil), http.StatusOK) + assertStatus(t, doJSON(t, env, http.MethodGet, "/necore/documents/private/"+nodeID, env.adminToken, nil), http.StatusOK) + + assertStatus(t, doJSON(t, env, http.MethodPatch, "/necore/documents/node/"+nodeID, env.adminToken, fiber.Map{"name": "Renamed"}), http.StatusOK) + assertStatus(t, doJSON(t, env, http.MethodPut, "/necore/documents/node/"+nodeID, env.adminToken, fiber.Map{ + "private": false, + "content": []fiber.Map{{"type": "markdown", "content": "body"}}, + }), http.StatusOK) + + // 使用真实存在的父节点,避免 GORM 输出无意义的 record not found 日志。 + assertStatus(t, doJSON(t, env, http.MethodPost, "/necore/documents/node/"+nodeID, env.adminToken, fiber.Map{ + "parentId": parentID, + }), http.StatusOK) + + assertStatus(t, doMultipartFile(t, env, "/necore/documents/upload/"+nodeID, env.adminToken, "file", "doc.txt", "file body"), http.StatusOK) + + // 不通过 Fiber Static 读取随后需要删除的同一个文件。 + // Fiber/fasthttp 在 Windows 下可能让 SendFile 的文件句柄存活到请求上下文回收, + // 即使 net/http 响应体已读取并关闭,立即 os.Remove 仍可能得到 ERROR_SHARING_VIOLATION。 + uploadedPath := filepath.Join(env.tmpDir, "contents", nodeID, "doc.txt") + uploadedBody, err := os.ReadFile(uploadedPath) + must(t, err) + if string(uploadedBody) != "file body" { + t.Fatalf("uploaded document body = %q, want %q", string(uploadedBody), "file body") + } + + assertStatus(t, doJSON(t, env, http.MethodDelete, "/necore/documents/upload/"+nodeID, env.adminToken, fiber.Map{ + "url": "/contents/" + nodeID + "/doc.txt", + }), http.StatusOK) + if _, err := os.Stat(uploadedPath); !os.IsNotExist(err) { + t.Fatalf("uploaded file should be deleted, stat err = %v", err) + } + + // 仍覆盖静态文件路由,但请求不存在的文件,避免 Windows 文件句柄影响后续删除。 + assertStatus(t, doJSON(t, env, http.MethodGet, "/necore/contents/not-found.txt", "", nil), http.StatusNotFound) + assertStatus(t, doJSON(t, env, http.MethodDelete, "/necore/documents/node/"+nodeID, env.adminToken, nil), http.StatusOK) + assertStatus(t, doJSON(t, env, http.MethodDelete, "/necore/documents/node/"+parentID, env.adminToken, nil), http.StatusOK) +} + +func TestBotRoutes(t *testing.T) { + env := setupTestEnv(t) + + // Token 管理接口确实要求 bot_admin。 + assertStatus(t, doJSON(t, env, http.MethodPost, "/necore/bots/token", env.userToken, nil), http.StatusForbidden) + + createResp := doJSON(t, env, http.MethodPost, "/necore/bots/token", env.adminToken, nil) + assertStatus(t, createResp, http.StatusOK) + createBody := decodeBody(t, createResp) + tokenObj, ok := createBody["token"].(map[string]any) + if !ok || tokenObj["token"] == "" { + t.Fatalf("create bot token should return token object, got %#v", createBody) + } + + assertStatus(t, doJSON(t, env, http.MethodGet, "/necore/bots/token", env.adminToken, nil), http.StatusOK) + assertStatus(t, doJSON(t, env, http.MethodGet, "/necore/bots/token/missing", env.adminToken, nil), http.StatusInternalServerError) + + // 当前源码只要求“已登录”,没有 bot_admin 权限检查。 + assertStatus(t, doJSON(t, env, http.MethodGet, "/necore/bots/status", env.userToken, nil), http.StatusOK) + assertStatus(t, doJSON(t, env, http.MethodDelete, "/necore/bots/ws/kick/not-exist", env.userToken, nil), http.StatusOK) + assertStatus(t, doJSON(t, env, http.MethodDelete, "/necore/bots/token/missing", env.adminToken, nil), http.StatusOK) +} + +func TestSecurityRegression_FileDeletePathTraversalIsCurrentlyPossible(t *testing.T) { + env := setupTestEnv(t) + + victim := filepath.Join(env.tmpDir, "victim.txt") + must(t, os.WriteFile(victim, []byte("do not delete"), 0o644)) + + resp := doJSON(t, env, http.MethodDelete, "/necore/documents/upload/anything", env.adminToken, fiber.Map{ + "url": "victim.txt", + }) + assertStatus(t, resp, http.StatusOK) + + if _, err := os.Stat(victim); !os.IsNotExist(err) { + t.Fatalf("expected vulnerable handler to delete arbitrary relative file; stat err = %v", err) + } +} + +func TestSecurityRegression_BotDashboardAvailableToAnyAuthenticatedUser(t *testing.T) { + env := setupTestEnv(t) + + // 该测试记录当前安全缺陷:普通登录用户也能查看 bot 状态并调用 kick。 + assertStatus(t, doJSON(t, env, http.MethodGet, "/necore/bots/status", env.userToken, nil), http.StatusOK) + assertStatus(t, doJSON(t, env, http.MethodDelete, "/necore/bots/ws/kick/arbitrary-session", env.userToken, nil), http.StatusOK) +} + +func TestSecurityRegression_PrivateUserDataIsPubliclyEnumerable(t *testing.T) { + env := setupTestEnv(t) + + resp := doJSON(t, env, http.MethodGet, "/necore/auth/userlist", "", nil) + assertStatus(t, resp, http.StatusOK) + + if !strings.Contains(string(resp.Body), "admin") || !strings.Contains(string(resp.Body), "alice") { + t.Fatalf("expected public user list to expose usernames, body=%s", string(resp.Body)) + } +} From 7a79bd3c3e3a50dece04f33357ba40947fea8c00 Mon Sep 17 00:00:00 2001 From: Kingcq <404291187@qq.com> Date: Thu, 18 Jun 2026 14:35:43 +0800 Subject: [PATCH 06/13] feat: secure file upload and delete --- dao/article.go | 1 - service/article.go | 86 +++++++++++++++++++++++++++++++++++++-------- service/document.go | 51 +++++++++++++++++++-------- 3 files changed, 109 insertions(+), 29 deletions(-) diff --git a/dao/article.go b/dao/article.go index 3ad51b4..3a3ab2d 100644 --- a/dao/article.go +++ b/dao/article.go @@ -67,7 +67,6 @@ func GetArticleList(target string, page int, pageSize int, pin bool) ([]model.Ar func DeleteArticle(id string) error { db := database.GetArticleDatabase() - // Delete File os.RemoveAll(fmt.Sprintf("./contents/%s", id)) return db.Where(&model.Article{Id: id}).Delete(&model.Article{}).Error } diff --git a/service/article.go b/service/article.go index 9772bf9..cc8b0b7 100644 --- a/service/article.go +++ b/service/article.go @@ -2,16 +2,45 @@ package service import ( "encoding/json" + "errors" "fmt" "necore/dao" "necore/model" + "necore/util" + "necore/ws" "os" + "path/filepath" + "strings" "github.com/gofiber/fiber/v2" "github.com/golang-jwt/jwt/v5" "github.com/google/uuid" ) +func generateStoredFilename(original string) (string, error) { + safeName, err := util.SafeFilename(original) + if err != nil { + return "", err + } + + extension := strings.ToLower(filepath.Ext(safeName)) + + allowedExtensions := map[string]bool{ + ".png": true, + ".jpg": true, + ".jpeg": true, + ".webp": true, + ".pdf": true, + ".txt": true, + } + + if !allowedExtensions[extension] { + return "", errors.New("unsupported file extension") + } + + return uuid.NewString() + extension, nil +} + func checkNewsPermission(c *fiber.Ctx) bool { // Check if user is admin or news_admin token := c.Locals("user").(*jwt.Token) @@ -63,9 +92,10 @@ func UpdateArticle(c *fiber.Ctx) error { Content string `json:"content"` } type Payload struct { - Entity PayloadEntity `json:"entity"` - Content []PayloadContent `json:"content"` - Category string `json:"category"` + Entity PayloadEntity `json:"entity"` + Content []PayloadContent `json:"content"` + Category string `json:"category"` + DoesNotify bool `json:"doesNotify"` } payload := new(Payload) if err := c.BodyParser(payload); err != nil { @@ -94,6 +124,13 @@ func UpdateArticle(c *fiber.Ctx) error { return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{"error": err}) } + if payload.DoesNotify { + go ws.GlobalHub.Broadcast(fiber.Map{ + "event": "article_updated", + "data": newArticle, + }) + } + return c.SendStatus(fiber.StatusOK) } @@ -200,13 +237,17 @@ func UploadArticleFile(c *fiber.Ctx) error { if err != nil { return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{"error": err}) } - if err := os.MkdirAll(fmt.Sprintf("./contents/%s", id), os.ModePerm); err != nil { + if err := os.MkdirAll(fmt.Sprintf("./contents/%s", id), 0o750); err != nil { return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{"error": err}) } - if err := c.SaveFile(file, fmt.Sprintf("./contents/%s/%s", id, file.Filename)); err != nil { + storedName, err := generateStoredFilename(file.Filename) + if err != nil { return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{"error": err}) } - return c.JSON(fiber.Map{"url": fmt.Sprintf("/contents/%s/%s", id, file.Filename)}) + if err := c.SaveFile(file, fmt.Sprintf("./contents/%s/%s", id, storedName)); err != nil { + return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{"error": err}) + } + return c.JSON(fiber.Map{"url": fmt.Sprintf("/contents/%s/%s", id, storedName)}) } func DeleteArticleFile(c *fiber.Ctx) error { @@ -214,18 +255,35 @@ func DeleteArticleFile(c *fiber.Ctx) error { return c.Status(fiber.StatusForbidden).JSON(fiber.Map{"error": "Forbidden"}) } - // id := c.Params("id") // It is included in the url + id := c.Params("id") + type Payload struct { - Url string `json:"url"` + Filename string `json:"filename"` } - payload := new(Payload) - if err := c.BodyParser(payload); err != nil { - return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{"error": err}) + + var payload Payload + if err := c.BodyParser(&payload); err != nil { + return c.Status(fiber.StatusBadRequest). + JSON(fiber.Map{"error": "Invalid request body"}) } - if err := os.Remove(fmt.Sprintf("./%s", payload.Url)); err != nil { - return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{"error": err}) + + target, err := util.SafeContentPath("./contents", id, payload.Filename) + if err != nil { + return c.Status(fiber.StatusBadRequest). + JSON(fiber.Map{"error": "Invalid filename"}) } - return c.SendStatus(fiber.StatusOK) + + if err := os.Remove(target); err != nil { + if errors.Is(err, os.ErrNotExist) { + return c.Status(fiber.StatusNotFound). + JSON(fiber.Map{"error": "File not found"}) + } + + return c.Status(fiber.StatusInternalServerError). + JSON(fiber.Map{"error": "Internal server error"}) + } + + return c.SendStatus(fiber.StatusNoContent) } func DeleteArticle(c *fiber.Ctx) error { diff --git a/service/document.go b/service/document.go index 6acd53b..228d321 100644 --- a/service/document.go +++ b/service/document.go @@ -2,9 +2,11 @@ package service import ( "encoding/json" + "errors" "fmt" "necore/dao" "necore/model" + "necore/util" "os" "github.com/gofiber/fiber/v2" @@ -303,31 +305,52 @@ func UploadDocumentFile(c *fiber.Ctx) error { "error": err.Error(), }) } - if err := os.MkdirAll(fmt.Sprintf("./contents/%s", id), os.ModePerm); err != nil { + if err := os.MkdirAll(fmt.Sprintf("./contents/%s", id), 0o750); err != nil { return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{"error": err}) } - if err := c.SaveFile(file, fmt.Sprintf("./contents/%s/%s", id, file.Filename)); err != nil { + storedName, err := generateStoredFilename(file.Filename) + if err != nil { + return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{"error": err}) + } + if err := c.SaveFile(file, fmt.Sprintf("./contents/%s/%s", id, storedName)); err != nil { return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{"error": err}) } - return c.JSON(fiber.Map{"url": fmt.Sprintf("/contents/%s/%s", id, file.Filename)}) + return c.JSON(fiber.Map{"url": fmt.Sprintf("/contents/%s/%s", id, storedName)}) } func DeleteDocumentFile(c *fiber.Ctx) error { if !checkDocumentPermission(c) { - return c.Status(fiber.StatusForbidden).JSON(fiber.Map{ - "error": "You don't have permission to update document node name", - }) + return c.Status(fiber.StatusForbidden). + JSON(fiber.Map{"error": "Forbidden"}) } - // id := c.Params("id") // It is included in the url + + id := c.Params("id") + type Payload struct { - Url string `json:"url"` + Filename string `json:"filename"` } - payload := new(Payload) - if err := c.BodyParser(payload); err != nil { - return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{"error": err}) + + var payload Payload + if err := c.BodyParser(&payload); err != nil { + return c.Status(fiber.StatusBadRequest). + JSON(fiber.Map{"error": "Invalid request body"}) } - if err := os.Remove(fmt.Sprintf("./%s", payload.Url)); err != nil { - return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{"error": err}) + + target, err := util.SafeContentPath("./contents", id, payload.Filename) + if err != nil { + return c.Status(fiber.StatusBadRequest). + JSON(fiber.Map{"error": "Invalid filename"}) } - return c.SendStatus(fiber.StatusOK) + + if err := os.Remove(target); err != nil { + if errors.Is(err, os.ErrNotExist) { + return c.Status(fiber.StatusNotFound). + JSON(fiber.Map{"error": "File not found"}) + } + + return c.Status(fiber.StatusInternalServerError). + JSON(fiber.Map{"error": "Internal server error"}) + } + + return c.SendStatus(fiber.StatusNoContent) } From 0459e863d3767d01847cd8362fb55538fa328b73 Mon Sep 17 00:00:00 2001 From: Kingcq <404291187@qq.com> Date: Sat, 20 Jun 2026 02:59:43 +0800 Subject: [PATCH 07/13] fix: authing security issue --- config/config.go | 9 +- controller/middleware/auth.go | 3 + controller/middleware/validator.go | 71 ++++++ controller/router/router.go | 13 +- dao/user.go | 66 ++---- go.mod | 2 + go.sum | 4 + model/user.go | 11 +- routes_test.go | 367 +++++++++++++++++++++++++---- service/article.go | 17 +- service/auth.go | 6 +- service/bottoken.go | 6 +- service/document.go | 21 +- service/server.go | 17 +- service/user.go | 62 +++-- 15 files changed, 535 insertions(+), 140 deletions(-) create mode 100644 controller/middleware/validator.go diff --git a/config/config.go b/config/config.go index d7769ef..f45de9b 100644 --- a/config/config.go +++ b/config/config.go @@ -5,6 +5,7 @@ import ( "encoding/base64" "errors" "fmt" + "necore/util" "os" "path/filepath" "strings" @@ -43,11 +44,13 @@ func Init() error { } setDefaultEnvironment("PORT", "3000") - setDefaultEnvironment("BOT_LOG_BUFFER_SIZE", "100") + defaultSecret, _ := util.GenerateSecureToken("", 32) + setDefaultEnvironment("SECRET", defaultSecret) + setDefaultEnvironment("BOT_LOG_BUFFER_SIZE", "1000") secret := strings.TrimSpace(os.Getenv("SECRET")) - if len(secret) < 32 { - initErr = fmt.Errorf("SECRET must contain at least 32 characters") + if len(secret) < 1 { + initErr = fmt.Errorf("SECRET must contain at least 1 characters, SECRET=%q", secret) } }) diff --git a/controller/middleware/auth.go b/controller/middleware/auth.go index 33dbf1b..eda0eda 100644 --- a/controller/middleware/auth.go +++ b/controller/middleware/auth.go @@ -11,6 +11,9 @@ func AuthNeeded() fiber.Handler { return jwtware.New(jwtware.Config{ SigningKey: jwtware.SigningKey{Key: []byte(config.Config("SECRET"))}, ErrorHandler: jwtError, + SuccessHandler: func(c *fiber.Ctx) error { + return validateTokenVersion(c) + }, }) } diff --git a/controller/middleware/validator.go b/controller/middleware/validator.go new file mode 100644 index 0000000..ccbf3b4 --- /dev/null +++ b/controller/middleware/validator.go @@ -0,0 +1,71 @@ +package middleware + +import ( + "errors" + "necore/database" + "necore/model" + + "github.com/gofiber/fiber/v2" + "github.com/golang-jwt/jwt/v5" + "gorm.io/gorm" +) + +func validateTokenVersion(c *fiber.Ctx) error { + token, ok := c.Locals("user").(*jwt.Token) + if !ok || token == nil { + return c.Status(fiber.StatusUnauthorized).JSON(fiber.Map{ + "error": "Unauthorized", + }) + } + + claims, ok := token.Claims.(jwt.MapClaims) + if !ok { + return c.Status(fiber.StatusUnauthorized).JSON(fiber.Map{ + "error": "Invalid token", + }) + } + + username, ok := claims["name"].(string) + if !ok || username == "" { + return c.Status(fiber.StatusUnauthorized).JSON(fiber.Map{ + "error": "Invalid token", + }) + } + + tokenVersionFloat, ok := claims["ver"].(float64) + if !ok { + return c.Status(fiber.StatusUnauthorized).JSON(fiber.Map{ + "error": "Invalid token", + }) + } + + var user model.User + err := database.GetUserDatabase(). + Select("username", "token_version", "group", "tags"). + Where("username = ?", username). + First(&user).Error + + if errors.Is(err, gorm.ErrRecordNotFound) { + return c.Status(fiber.StatusUnauthorized).JSON(fiber.Map{ + "error": "User no longer exists", + }) + } + + if err != nil { + return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{ + "error": "Internal server error", + }) + } + + if uint(tokenVersionFloat) != user.TokenVersion { + return c.Status(fiber.StatusUnauthorized).JSON(fiber.Map{ + "error": "Token has been revoked", + }) + } + + // 将数据库中的最新用户信息放入 Locals, + // 后续权限中间件直接使用,不再信任 JWT 中的 group。 + c.Locals("currentUser", user) + + return c.Next() +} diff --git a/controller/router/router.go b/controller/router/router.go index 9c13812..cf7e00c 100644 --- a/controller/router/router.go +++ b/controller/router/router.go @@ -8,9 +8,11 @@ import ( "necore/service" "necore/ws" "strings" + "time" "github.com/gofiber/contrib/websocket" "github.com/gofiber/fiber/v2" + "github.com/gofiber/fiber/v2/middleware/limiter" ) type routerInstance struct { @@ -33,12 +35,21 @@ func GetInstance() *routerInstance { } func SetupRoutes() { + loginLimiter := limiter.New(limiter.Config{ + Max: 8, + Expiration: time.Minute, + LimitReached: func(c *fiber.Ctx) error { + return c.Status(fiber.StatusTooManyRequests). + JSON(fiber.Map{"error": "Too many login attempts"}) + }, + }) + router := instance.Router (*router).Get("/slogan", service.SloganHandler) authGroup := (*router).Group("/auth") authGroup.Get("/status", middleware.AuthNeeded(), service.GetStatus) - authGroup.Post("/login", service.Login) + authGroup.Post("/login", loginLimiter, service.Login) authGroup.Post("/register", middleware.AuthNeeded(), service.AddUser) authGroup.Get("/user/:id", service.GetUserInfo) authGroup.Get("/avatar/:id", service.GetUserAvatar) diff --git a/dao/user.go b/dao/user.go index 69bfd90..7197330 100644 --- a/dao/user.go +++ b/dao/user.go @@ -7,6 +7,7 @@ import ( "necore/config" "necore/database" "necore/model" + "slices" "time" "github.com/golang-jwt/jwt/v5" @@ -32,59 +33,33 @@ func DebugTestPassword() { log.Println(`Test Password "test":`, hash) } +func UnitTestPassword() string { + hash, _ := hashPassword("unit-test-password") + return hash +} + // Token func CreateToken(u model.User) (string, error) { token := jwt.New(jwt.SigningMethodHS256) claims := token.Claims.(jwt.MapClaims) - claims["username"] = u.Username - claims["group"] = u.Group - claims["tags"] = u.Tags + + claims["name"] = u.Username + claims["ver"] = u.TokenVersion + claims["iat"] = time.Now().Unix() claims["exp"] = time.Now().Add(time.Hour * 72).Unix() + t, err := token.SignedString([]byte(config.Config("SECRET"))) return t, err } -func GetUsernameFromToken(t *jwt.Token) string { - return t.Claims.(jwt.MapClaims)["username"].(string) -} - -func GetUserGroupsFromToken(t *jwt.Token) []string { - claims := t.Claims.(jwt.MapClaims)["group"] - if claims == nil { - return []string{} - } - +func ContainsGroup(userGroup string, group string) bool { var groups []string - err := json.Unmarshal([]byte(claims.(string)), &groups) - if err != nil { - return []string{} - } - return groups -} - -func GetUserTagsFromToken(t *jwt.Token) []string { - claims := t.Claims.(jwt.MapClaims)["tags"] - if claims == nil { - return []string{} - } - - var tags []string - err := json.Unmarshal([]byte(claims.(string)), &tags) + err := json.Unmarshal([]byte(userGroup), &groups) if err != nil { - return []string{} + groups = []string{} } - return tags -} - -func IsUserInGroup(t *jwt.Token, group string) bool { - groups := GetUserGroupsFromToken(t) - for _, g := range groups { - if g == group { - return true - } - } - return false + return slices.Contains(groups, group) } // Database @@ -147,6 +122,17 @@ func UpdateUserInfo(username string, group string, tags string) error { return db.Model(&user).Updates(model.User{Group: group, Tags: tags}).Error } +func UpdateUserPermissions(username string) error { + db := database.GetUserDatabase() + + return db.Model(&model.User{}). + Where(model.User{Username: username}). + UpdateColumn( + "token_version", + gorm.Expr("token_version + ?", 1), + ).Error +} + func GetUserAvatar(username string) (string, error) { db := database.GetUserDatabase() var user model.User diff --git a/go.mod b/go.mod index 0eb1d0f..5eda76f 100644 --- a/go.mod +++ b/go.mod @@ -29,11 +29,13 @@ require ( github.com/mattn/go-isatty v0.0.22 // indirect github.com/mattn/go-runewidth v0.0.24 // indirect github.com/mattn/go-sqlite3 v1.14.45 // indirect + github.com/philhofer/fwd v1.2.0 // indirect github.com/rivo/uniseg v0.4.7 // indirect github.com/savsgio/gotils v0.0.0-20250924091648-bce9a52d7761 // indirect github.com/tidwall/gjson v1.18.0 // indirect github.com/tidwall/match v1.1.1 // indirect github.com/tidwall/pretty v1.2.0 // indirect + github.com/tinylib/msgp v1.6.4 // indirect github.com/valyala/bytebufferpool v1.0.0 // indirect github.com/valyala/fasthttp v1.71.0 // indirect github.com/valyala/tcplisten v1.0.0 // indirect diff --git a/go.sum b/go.sum index 1be9bf8..ad8e271 100644 --- a/go.sum +++ b/go.sum @@ -79,6 +79,8 @@ github.com/mattn/go-sqlite3 v1.14.45 h1:6KA/spDguL3KV8rnybG7ezSaE4SeMR3KC9VbUoAQ github.com/mattn/go-sqlite3 v1.14.45/go.mod h1:pjEuOr8IwzLJP2MfGeTb0A35jauH+C2kbHKBr7yXKVQ= github.com/millkhan/mcstatusgo/v2 v2.2.0 h1:uRyHiOvqlK+6Oz3za4hMWAktSLjaqD/QzyQCcX2DjzI= github.com/millkhan/mcstatusgo/v2 v2.2.0/go.mod h1:YUJHhrJzsQP4PoDXFo++7JzPU7TjdztFM519CPqKe5M= +github.com/philhofer/fwd v1.2.0 h1:e6DnBTl7vGY+Gz322/ASL4Gyp1FspeMvx1RNDoToZuM= +github.com/philhofer/fwd v1.2.0/go.mod h1:RqIHx9QI14HlwKwm98g9Re5prTQ6LdeRQn+gXJFxsJM= github.com/rivo/uniseg v0.2.0 h1:S1pD9weZBuJdFmowNwbpi7BJ8TNftyUImj/0WQi72jY= github.com/rivo/uniseg v0.2.0/go.mod h1:J6wj4VEh+S6ZtnVlnTBMWIodfgj8LQOQFoIToxlJtxc= github.com/rivo/uniseg v0.4.7 h1:WUdvkW8uEhrYfLC4ZzdpI2ztxP1I582+49Oc5Mq64VQ= @@ -93,6 +95,8 @@ github.com/tidwall/match v1.1.1 h1:+Ho715JplO36QYgwN9PGYNhgZvoUSc9X2c80KVTi+GA= github.com/tidwall/match v1.1.1/go.mod h1:eRSPERbgtNPcGhD8UCthc6PmLEQXEWd3PRB5JTxsfmM= github.com/tidwall/pretty v1.2.0 h1:RWIZEg2iJ8/g6fDDYzMpobmaoGh5OLl4AXtGUGPcqCs= github.com/tidwall/pretty v1.2.0/go.mod h1:ITEVvHYasfjBbM0u2Pg8T2nJnzm8xPwvNhhsoaGGjNU= +github.com/tinylib/msgp v1.6.4 h1:mOwYbyYDLPj35mkA2BjjYejgJk9BuHxDdvRnb6v2ZcQ= +github.com/tinylib/msgp v1.6.4/go.mod h1:RSp0LW9oSxFut3KzESt5Voq4GVWyS+PSulT77roAqEA= github.com/valyala/bytebufferpool v1.0.0 h1:GqA5TC/0021Y/b9FG4Oi9Mr3q7XYx6KllzawFIhcdPw= github.com/valyala/bytebufferpool v1.0.0/go.mod h1:6bBcMArwyJ5K/AmCkWv1jt77kVWyCJ6HpOuEn7z0Csc= github.com/valyala/fasthttp v1.51.0 h1:8b30A5JlZ6C7AS81RsWjYMQmrZG6feChmgAolCl1SqA= diff --git a/model/user.go b/model/user.go index 9ee0c45..33e58f3 100644 --- a/model/user.go +++ b/model/user.go @@ -7,9 +7,10 @@ import ( type User struct { gorm.Model - Username string `gorm:"uniqueIndex;not null" json:"username"` - Password string `gorm:"not null" json:"password"` // sha256 hashed - Group string `json:"group"` // json array: []string - Tags string `json:"tags"` // json array: []string - Avatar string `json:"avatar"` + Username string `gorm:"uniqueIndex;not null" json:"username"` + Password string `gorm:"not null" json:"password"` // sha256 hashed + Group string `json:"group"` // json array: []string + Tags string `json:"tags"` // json array: []string + Avatar string `json:"avatar"` + TokenVersion uint `gorm:"not null;default:1" json:"-"` } diff --git a/routes_test.go b/routes_test.go index 9c92cae..e978b9a 100644 --- a/routes_test.go +++ b/routes_test.go @@ -3,6 +3,7 @@ package main import ( "bytes" "encoding/json" + "fmt" "io" "mime/multipart" "net/http" @@ -11,15 +12,18 @@ import ( "path/filepath" "strings" "testing" + "time" "necore/controller/middleware" "necore/dao" "necore/database" "necore/model" "necore/service" + "necore/ws" "github.com/gofiber/contrib/websocket" "github.com/gofiber/fiber/v2" + "github.com/gofiber/fiber/v2/middleware/limiter" "gorm.io/gorm" "gorm.io/gorm/logger" ) @@ -83,33 +87,27 @@ func setupTestEnv(t *testing.T) *testEnv { Model(&model.User{}). Where("username = ?", "admin"). Updates(model.User{ - Group: `["admin","news_admin","server_admin","document_admin","bot_admin"]`, - Tags: `[]`, - Avatar: "admin-avatar", + Password: dao.UnitTestPassword(), + Group: `["admin","news_admin","server_admin","document_admin","bot_admin"]`, + Tags: `[]`, + Avatar: "admin-avatar", }).Error) must(t, database.GetUserDatabase(). Model(&model.User{}). Where("username = ?", "alice"). Updates(model.User{ - Group: `[]`, - Tags: `[]`, - Avatar: "alice-avatar", + Password: dao.UnitTestPassword(), + Group: `[]`, + Tags: `[]`, + Avatar: "alice-avatar", }).Error) - adminToken, err := dao.CreateToken(model.User{ - Username: "admin", - Group: `["admin","news_admin","server_admin","document_admin","bot_admin"]`, - Tags: `[]`, - }) - must(t, err) - - userToken, err := dao.CreateToken(model.User{ - Username: "alice", - Group: `[]`, - Tags: `[]`, - }) - must(t, err) + // 从数据库重新读取用户后再签发 token。 + // 引入 token_version 后,JWT 中的 ver 必须等于数据库里的 token_version, + // 不能再用手写的 model.User 字面量签发测试 token。 + adminToken := createTokenForUser(t, "admin") + userToken := createTokenForUser(t, "alice") app := fiber.New(fiber.Config{BodyLimit: 512 * 1024 * 1024}) registerRoutes(app) @@ -146,17 +144,98 @@ func closeGormDB(t *testing.T, db *gorm.DB) { } } +func createTokenForUser(t *testing.T, username string) string { + t.Helper() + + user, err := dao.GetUserByUsername(username) + must(t, err) + if user == nil { + t.Fatalf("user %q not found", username) + } + + token, err := dao.CreateToken(*user) + must(t, err) + return token +} + +func loginAndGetToken(t *testing.T, env *testEnv, username, password string) string { + t.Helper() + + resp := doJSON(t, env, http.MethodPost, "/necore/auth/login", "", fiber.Map{ + "username": username, + "password": password, + }) + assertStatus(t, resp, http.StatusOK) + + body := decodeBody(t, resp) + token, ok := body["token"].(string) + if !ok || token == "" { + t.Fatalf("login response should contain token string, got %#v", body) + } + return token +} + +func getUserTokenVersion(t *testing.T, username string) uint { + t.Helper() + + var version uint + result := database.GetUserDatabase(). + Model(&model.User{}). + Select("token_version"). + Where("username = ?", username). + Scan(&version) + + must(t, result.Error) + if result.RowsAffected == 0 { + t.Fatalf("user %q not found while reading token_version", username) + } + return version +} + +func incrementUserTokenVersion(t *testing.T, username string) { + t.Helper() + + result := database.GetUserDatabase(). + Model(&model.User{}). + Where("username = ?", username). + UpdateColumn("token_version", gorm.Expr("token_version + 1")) + + must(t, result.Error) + if result.RowsAffected == 0 { + t.Fatalf("user %q not found while incrementing token_version", username) + } +} + +func assertUserTokenVersion(t *testing.T, username string, want uint) { + t.Helper() + + got := getUserTokenVersion(t, username) + if got != want { + t.Fatalf("token_version for %q = %d, want %d", username, got, want) + } +} + func registerRoutes(app *fiber.App) { api := app.Group("/necore") + + loginLimiter := limiter.New(limiter.Config{ + Max: 8, + Expiration: time.Minute, + LimitReached: func(c *fiber.Ctx) error { + return c.Status(fiber.StatusTooManyRequests). + JSON(fiber.Map{"error": "Too many login attempts"}) + }, + }) + api.Get("/slogan", service.SloganHandler) authGroup := api.Group("/auth") authGroup.Get("/status", middleware.AuthNeeded(), service.GetStatus) - authGroup.Post("/login", service.Login) + authGroup.Post("/login", loginLimiter, service.Login) authGroup.Post("/register", middleware.AuthNeeded(), service.AddUser) authGroup.Get("/user/:id", service.GetUserInfo) authGroup.Get("/avatar/:id", service.GetUserAvatar) - authGroup.Get("/userlist", service.GetUserList) + authGroup.Get("/userlist", middleware.AuthNeeded(), service.GetUserList) authGroup.Delete("/user/:id", middleware.AuthNeeded(), service.DeleteUser) authGroup.Post("/password", middleware.AuthNeeded(), service.UpdateUserPassword) authGroup.Post("/avatar", middleware.AuthNeeded(), service.UpdateUserAvatar) @@ -194,12 +273,47 @@ func registerRoutes(app *fiber.App) { api.Static("/contents", "./contents") botGroup := api.Group("/bots") + + botGroup.Use("/ws/updates", func(c *fiber.Ctx) error { + if !websocket.IsWebSocketUpgrade(c) { + return fiber.ErrUpgradeRequired + } + + auth := c.Get(fiber.HeaderAuthorization) + if !strings.HasPrefix(auth, "Bearer ") { + return c.SendStatus(fiber.StatusUnauthorized) + } + + identifier := c.Params("identifier") + if identifier == "" { + return c.SendStatus(fiber.StatusBadRequest) + } + + token := strings.TrimPrefix(auth, "Bearer ") + botToken, err := dao.GetBotTokenByPlainToken(token) + if err != nil { + ws.GlobalHub.AddLog( + fmt.Sprintf( + "⚠️ 拒绝 %s 连接:无效 Token", + identifier, + ), + ws.ERROR, + ) + return c.SendStatus(fiber.StatusUnauthorized) + } + + c.Locals("token_id", botToken.ID) + c.Locals("token_name", botToken.Name) + c.Locals("identifier", identifier) + return c.Next() + }) + botGroup.Get("/ws/updates", websocket.New(service.HandleWSConnection)) + botGroup.Post("/token", middleware.AuthNeeded(), service.CreateBotToken) botGroup.Get("/token", middleware.AuthNeeded(), service.GetBotTokenList) botGroup.Get("/token/:id", middleware.AuthNeeded(), service.GetBotToken) botGroup.Delete("/token/:id", middleware.AuthNeeded(), service.DeleteBotToken) botGroup.Get("/status", middleware.AuthNeeded(), service.GetWSStatus) - botGroup.Get("/ws/updates", websocket.New(service.HandleWSConnection)) botGroup.Delete("/ws/kick/:session_id", middleware.AuthNeeded(), service.KickConnection) } @@ -303,7 +417,7 @@ func TestPublicAndAuthRoutes(t *testing.T) { loginResp := doJSON(t, env, http.MethodPost, "/necore/auth/login", "", fiber.Map{ "username": "admin", - "password": "admin-pass", + "password": "unit-test-password", }) assertStatus(t, loginResp, http.StatusOK) loginBody := decodeBody(t, loginResp) @@ -328,18 +442,25 @@ func TestPublicAndAuthRoutes(t *testing.T) { assertStatus(t, doJSON(t, env, http.MethodGet, "/necore/auth/user/admin", "", nil), http.StatusOK) assertStatus(t, doJSON(t, env, http.MethodGet, "/necore/auth/avatar/admin", "", nil), http.StatusOK) - assertStatus(t, doJSON(t, env, http.MethodGet, "/necore/auth/userlist", "", nil), http.StatusOK) + assertStatus(t, doJSON(t, env, http.MethodGet, "/necore/auth/userlist", "", nil), http.StatusUnauthorized) + assertStatus(t, doJSON(t, env, http.MethodGet, "/necore/auth/userlist", env.adminToken, nil), http.StatusOK) assertStatus(t, doJSON(t, env, http.MethodPost, "/necore/auth/password", env.userToken, fiber.Map{ - "id": "admin", - "new_password": "new", + "id": "admin", + "self_password": "wrong-password", + "new_password": "new", }), http.StatusForbidden) assertStatus(t, doJSON(t, env, http.MethodPost, "/necore/auth/password", env.userToken, fiber.Map{ - "id": "alice", - "new_password": "new", + "id": "alice", + "self_password": "unit-test-password", + "new_password": "new", }), http.StatusOK) + // 改密码属于安全敏感操作。接入 token_version 后,alice 的旧 token + // 会立即失效,因此后续仍需要 alice 身份的断言必须重新登录获取新 token。 + env.userToken = loginAndGetToken(t, env, "alice", "new") + assertStatus(t, doJSON(t, env, http.MethodPost, "/necore/auth/avatar", env.userToken, fiber.Map{ "username": "admin", "avatar": "x", @@ -362,6 +483,8 @@ func TestPublicAndAuthRoutes(t *testing.T) { "Tags": []any{}, }), http.StatusOK) + env.userToken = createTokenForUser(t, "alice") + assertStatus(t, doJSON(t, env, http.MethodDelete, "/necore/auth/user/alice", env.userToken, nil), http.StatusForbidden) assertStatus(t, doJSON(t, env, http.MethodDelete, "/necore/auth/user/bob", env.adminToken, nil), http.StatusOK) } @@ -401,10 +524,12 @@ func TestNewsRoutes(t *testing.T) { }), http.StatusOK) assertStatus(t, doJSON(t, env, http.MethodGet, "/necore/news/detail/"+articleID, "", nil), http.StatusOK) - assertStatus(t, doMultipartFile(t, env, "/necore/news/upload/"+articleID, env.adminToken, "file", "hello.txt", "hello"), http.StatusOK) + response := doMultipartFile(t, env, "/necore/news/upload/"+articleID, env.adminToken, "file", "hello.txt", "hello") + assertStatus(t, response, http.StatusOK) + filename := strings.Split(decodeBody(t, response)["url"].(string), "/")[3] assertStatus(t, doJSON(t, env, http.MethodDelete, "/necore/news/upload/"+articleID, env.adminToken, fiber.Map{ - "url": "/contents/" + articleID + "/hello.txt", - }), http.StatusOK) + "filename": filename, + }), http.StatusNoContent) assertStatus(t, doJSON(t, env, http.MethodDelete, "/necore/news/"+articleID, env.adminToken, nil), http.StatusOK) } @@ -484,12 +609,14 @@ func TestDocumentRoutes(t *testing.T) { "parentId": parentID, }), http.StatusOK) - assertStatus(t, doMultipartFile(t, env, "/necore/documents/upload/"+nodeID, env.adminToken, "file", "doc.txt", "file body"), http.StatusOK) + response := doMultipartFile(t, env, "/necore/documents/upload/"+nodeID, env.adminToken, "file", "doc.txt", "file body") + assertStatus(t, response, http.StatusOK) + filename := strings.Split(decodeBody(t, response)["url"].(string), "/")[3] // 不通过 Fiber Static 读取随后需要删除的同一个文件。 // Fiber/fasthttp 在 Windows 下可能让 SendFile 的文件句柄存活到请求上下文回收, // 即使 net/http 响应体已读取并关闭,立即 os.Remove 仍可能得到 ERROR_SHARING_VIOLATION。 - uploadedPath := filepath.Join(env.tmpDir, "contents", nodeID, "doc.txt") + uploadedPath := filepath.Join(env.tmpDir, "contents", nodeID, filename) uploadedBody, err := os.ReadFile(uploadedPath) must(t, err) if string(uploadedBody) != "file body" { @@ -497,8 +624,8 @@ func TestDocumentRoutes(t *testing.T) { } assertStatus(t, doJSON(t, env, http.MethodDelete, "/necore/documents/upload/"+nodeID, env.adminToken, fiber.Map{ - "url": "/contents/" + nodeID + "/doc.txt", - }), http.StatusOK) + "filename": filename, + }), http.StatusNoContent) if _, err := os.Stat(uploadedPath); !os.IsNotExist(err) { t.Fatalf("uploaded file should be deleted, stat err = %v", err) } @@ -515,7 +642,9 @@ func TestBotRoutes(t *testing.T) { // Token 管理接口确实要求 bot_admin。 assertStatus(t, doJSON(t, env, http.MethodPost, "/necore/bots/token", env.userToken, nil), http.StatusForbidden) - createResp := doJSON(t, env, http.MethodPost, "/necore/bots/token", env.adminToken, nil) + createResp := doJSON(t, env, http.MethodPost, "/necore/bots/token", env.adminToken, fiber.Map{ + "name": "unit-test", + }) assertStatus(t, createResp, http.StatusOK) createBody := decodeBody(t, createResp) tokenObj, ok := createBody["token"].(map[string]any) @@ -528,7 +657,7 @@ func TestBotRoutes(t *testing.T) { // 当前源码只要求“已登录”,没有 bot_admin 权限检查。 assertStatus(t, doJSON(t, env, http.MethodGet, "/necore/bots/status", env.userToken, nil), http.StatusOK) - assertStatus(t, doJSON(t, env, http.MethodDelete, "/necore/bots/ws/kick/not-exist", env.userToken, nil), http.StatusOK) + assertStatus(t, doJSON(t, env, http.MethodDelete, "/necore/bots/ws/kick/not-exist", env.userToken, nil), http.StatusForbidden) assertStatus(t, doJSON(t, env, http.MethodDelete, "/necore/bots/token/missing", env.adminToken, nil), http.StatusOK) } @@ -539,11 +668,11 @@ func TestSecurityRegression_FileDeletePathTraversalIsCurrentlyPossible(t *testin must(t, os.WriteFile(victim, []byte("do not delete"), 0o644)) resp := doJSON(t, env, http.MethodDelete, "/necore/documents/upload/anything", env.adminToken, fiber.Map{ - "url": "victim.txt", + "filename": "../../victim.txt", }) - assertStatus(t, resp, http.StatusOK) + assertStatus(t, resp, http.StatusBadRequest) - if _, err := os.Stat(victim); !os.IsNotExist(err) { + if _, err := os.Stat(victim); err != nil { t.Fatalf("expected vulnerable handler to delete arbitrary relative file; stat err = %v", err) } } @@ -553,16 +682,170 @@ func TestSecurityRegression_BotDashboardAvailableToAnyAuthenticatedUser(t *testi // 该测试记录当前安全缺陷:普通登录用户也能查看 bot 状态并调用 kick。 assertStatus(t, doJSON(t, env, http.MethodGet, "/necore/bots/status", env.userToken, nil), http.StatusOK) - assertStatus(t, doJSON(t, env, http.MethodDelete, "/necore/bots/ws/kick/arbitrary-session", env.userToken, nil), http.StatusOK) + assertStatus(t, doJSON(t, env, http.MethodDelete, "/necore/bots/ws/kick/arbitrary-session", env.userToken, nil), http.StatusForbidden) } func TestSecurityRegression_PrivateUserDataIsPubliclyEnumerable(t *testing.T) { env := setupTestEnv(t) - resp := doJSON(t, env, http.MethodGet, "/necore/auth/userlist", "", nil) + resp := doJSON(t, env, http.MethodGet, "/necore/auth/userlist", env.userToken, nil) assertStatus(t, resp, http.StatusOK) if !strings.Contains(string(resp.Body), "admin") || !strings.Contains(string(resp.Body), "alice") { t.Fatalf("expected public user list to expose usernames, body=%s", string(resp.Body)) } } + +func TestTokenVersion_StaleTokenRejectedAcrossProtectedRouteGroups(t *testing.T) { + env := setupTestEnv(t) + + assertStatus(t, doJSON(t, env, http.MethodGet, "/necore/auth/status", env.adminToken, nil), http.StatusOK) + + oldVersion := getUserTokenVersion(t, "admin") + incrementUserTokenVersion(t, "admin") + assertUserTokenVersion(t, "admin", oldVersion+1) + + cases := []struct { + name string + method string + path string + body any + }{ + { + name: "auth status", + method: http.MethodGet, + path: "/necore/auth/status", + }, + { + name: "auth register", + method: http.MethodPost, + path: "/necore/auth/register", + body: fiber.Map{ + "username": "stale-created-user", + "password": "password", + }, + }, + { + name: "news create", + method: http.MethodPost, + path: "/necore/news/create", + }, + { + name: "server create", + method: http.MethodGet, + path: "/necore/server/create", + }, + { + name: "documents create node", + method: http.MethodPost, + path: "/necore/documents/node", + body: fiber.Map{ + "parentId": "root", + "isFolder": true, + "private": false, + "name": "Should Not Be Created", + }, + }, + { + name: "bots status", + method: http.MethodGet, + path: "/necore/bots/status", + }, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + resp := doJSON(t, env, tc.method, tc.path, env.adminToken, tc.body) + assertStatus(t, resp, http.StatusUnauthorized) + }) + } + + freshAdminToken := createTokenForUser(t, "admin") + assertStatus(t, doJSON(t, env, http.MethodGet, "/necore/auth/status", freshAdminToken, nil), http.StatusOK) +} + +func TestTokenVersion_PasswordChangeRevokesTargetUserToken(t *testing.T) { + env := setupTestEnv(t) + + assertStatus(t, doJSON(t, env, http.MethodGet, "/necore/auth/status", env.userToken, nil), http.StatusOK) + + oldVersion := getUserTokenVersion(t, "alice") + assertStatus(t, doJSON(t, env, http.MethodPost, "/necore/auth/password", env.userToken, fiber.Map{ + "id": "alice", + "self_password": "unit-test-password", + "new_password": "new-alice-password", + }), http.StatusOK) + assertUserTokenVersion(t, "alice", oldVersion+1) + + // 旧 JWT 已经被撤销,不能再访问任何需要登录的接口。 + assertStatus(t, doJSON(t, env, http.MethodGet, "/necore/auth/status", env.userToken, nil), http.StatusUnauthorized) + + // 旧密码不能登录,新密码登录后拿到的新 JWT 应该有效。 + assertStatus(t, doJSON(t, env, http.MethodPost, "/necore/auth/login", "", fiber.Map{ + "username": "alice", + "password": "alice-pass", + }), http.StatusUnauthorized) + + newToken := loginAndGetToken(t, env, "alice", "new-alice-password") + assertStatus(t, doJSON(t, env, http.MethodGet, "/necore/auth/status", newToken, nil), http.StatusOK) +} + +func TestTokenVersion_UserPermissionChangeRevokesTargetTokenOnly(t *testing.T) { + env := setupTestEnv(t) + + assertStatus(t, doJSON(t, env, http.MethodGet, "/necore/auth/status", env.adminToken, nil), http.StatusOK) + assertStatus(t, doJSON(t, env, http.MethodGet, "/necore/auth/status", env.userToken, nil), http.StatusOK) + + oldAliceVersion := getUserTokenVersion(t, "alice") + oldAdminVersion := getUserTokenVersion(t, "admin") + + assertStatus(t, doJSON(t, env, http.MethodPatch, "/necore/auth/user", env.adminToken, fiber.Map{ + "username": "alice", + "group": []string{"document_admin"}, + "Tags": []any{}, + }), http.StatusOK) + + assertUserTokenVersion(t, "alice", oldAliceVersion+1) + assertUserTokenVersion(t, "admin", oldAdminVersion) + + // 被修改权限的目标用户旧 token 应该失效。 + assertStatus(t, doJSON(t, env, http.MethodGet, "/necore/auth/status", env.userToken, nil), http.StatusUnauthorized) + + // 执行修改的管理员不应该因为修改别人权限而被迫下线。 + assertStatus(t, doJSON(t, env, http.MethodGet, "/necore/auth/status", env.adminToken, nil), http.StatusOK) + + // 重新签发的新 token 应该携带/对应数据库中的最新权限。 + freshAliceToken := createTokenForUser(t, "alice") + assertStatus(t, doJSON(t, env, http.MethodGet, "/necore/documents/layer/private/root", freshAliceToken, nil), http.StatusOK) +} + +func TestTokenVersion_AvatarChangeDoesNotRevokeToken(t *testing.T) { + env := setupTestEnv(t) + + oldVersion := getUserTokenVersion(t, "alice") + assertStatus(t, doJSON(t, env, http.MethodPost, "/necore/auth/avatar", env.userToken, fiber.Map{ + "username": "alice", + "avatar": "avatar-after-change", + }), http.StatusOK) + + assertUserTokenVersion(t, "alice", oldVersion) + assertStatus(t, doJSON(t, env, http.MethodGet, "/necore/auth/status", env.userToken, nil), http.StatusOK) +} + +func TestTokenVersion_DeletedUserTokenIsRejected(t *testing.T) { + env := setupTestEnv(t) + + assertStatus(t, doJSON(t, env, http.MethodPost, "/necore/auth/register", env.adminToken, fiber.Map{ + "username": "charlie", + "password": "charlie-pass", + }), http.StatusOK) + + charlieToken := loginAndGetToken(t, env, "charlie", "charlie-pass") + assertStatus(t, doJSON(t, env, http.MethodGet, "/necore/auth/status", charlieToken, nil), http.StatusOK) + + assertStatus(t, doJSON(t, env, http.MethodDelete, "/necore/auth/user/charlie", env.adminToken, nil), http.StatusOK) + + // 即使 token_version 没有递增,只要鉴权中间件每次查询数据库用户, + // 被删除用户的旧 JWT 也必须被拒绝。 + assertStatus(t, doJSON(t, env, http.MethodGet, "/necore/auth/status", charlieToken, nil), http.StatusUnauthorized) +} diff --git a/service/article.go b/service/article.go index cc8b0b7..1dd65c9 100644 --- a/service/article.go +++ b/service/article.go @@ -13,7 +13,6 @@ import ( "strings" "github.com/gofiber/fiber/v2" - "github.com/golang-jwt/jwt/v5" "github.com/google/uuid" ) @@ -43,9 +42,9 @@ func generateStoredFilename(original string) (string, error) { func checkNewsPermission(c *fiber.Ctx) bool { // Check if user is admin or news_admin - token := c.Locals("user").(*jwt.Token) - isAdmin := dao.IsUserInGroup(token, "admin") - isNewsAdmin := dao.IsUserInGroup(token, "news_admin") + user := c.Locals("currentUser").(model.User) + isAdmin := dao.ContainsGroup(user.Group, "admin") + isNewsAdmin := dao.ContainsGroup(user.Group, "news_admin") if isAdmin || isNewsAdmin { return false } @@ -74,8 +73,8 @@ func UpdateArticle(c *fiber.Ctx) error { if checkNewsPermission(c) { return c.Status(fiber.StatusForbidden).JSON(fiber.Map{"error": "Forbidden"}) } - token := c.Locals("user").(*jwt.Token) - author := dao.GetUsernameFromToken(token) + user := c.Locals("currentUser").(model.User) + author := user.Username id := c.Params("id") // Parse @@ -244,7 +243,11 @@ func UploadArticleFile(c *fiber.Ctx) error { if err != nil { return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{"error": err}) } - if err := c.SaveFile(file, fmt.Sprintf("./contents/%s/%s", id, storedName)); err != nil { + contentPath, err := util.SafeContentPath("./contents", id, storedName) + if err != nil { + return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{"error": err}) + } + if err := c.SaveFile(file, contentPath); err != nil { return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{"error": err}) } return c.JSON(fiber.Map{"url": fmt.Sprintf("/contents/%s/%s", id, storedName)}) diff --git a/service/auth.go b/service/auth.go index cb394c9..adb9d0c 100644 --- a/service/auth.go +++ b/service/auth.go @@ -3,9 +3,9 @@ package service import ( "encoding/json" "necore/dao" + "necore/model" "github.com/gofiber/fiber/v2" - "github.com/golang-jwt/jwt/v5" ) // Handlers @@ -77,8 +77,8 @@ func Login(c *fiber.Ctx) error { // Register by admin func AddUser(c *fiber.Ctx) error { // Check if user is admin - token := c.Locals("user").(*jwt.Token) - if !dao.IsUserInGroup(token, "admin") { + currentUser := c.Locals("currentUser").(model.User) + if !dao.ContainsGroup(currentUser.Group, "admin") { return c.Status(fiber.StatusForbidden).JSON(fiber.Map{"error": "Forbidden"}) } diff --git a/service/bottoken.go b/service/bottoken.go index c06a592..49a67df 100644 --- a/service/bottoken.go +++ b/service/bottoken.go @@ -2,14 +2,14 @@ package service import ( "necore/dao" + "necore/model" "github.com/gofiber/fiber/v2" - "github.com/golang-jwt/jwt/v5" ) func checkBotTokenPermission(c *fiber.Ctx) bool { - token := c.Locals("user").(*jwt.Token) - isBotAdmin := dao.IsUserInGroup(token, "bot_admin") || dao.IsUserInGroup(token, "admin") + user := c.Locals("currentUser").(model.User) + isBotAdmin := dao.ContainsGroup(user.Group, "bot_admin") || dao.ContainsGroup(user.Group, "admin") if isBotAdmin { return false } diff --git a/service/document.go b/service/document.go index 228d321..6a7c8ea 100644 --- a/service/document.go +++ b/service/document.go @@ -10,15 +10,14 @@ import ( "os" "github.com/gofiber/fiber/v2" - "github.com/golang-jwt/jwt/v5" "github.com/google/uuid" ) func checkDocumentPermission(c *fiber.Ctx) bool { // Check if user is admin or document_admin - token := c.Locals("user").(*jwt.Token) - isAdmin := dao.IsUserInGroup(token, "admin") - isDocsAdmin := dao.IsUserInGroup(token, "document_admin") + user := c.Locals("currentUser").(model.User) + isAdmin := dao.ContainsGroup(user.Group, "admin") + isDocsAdmin := dao.ContainsGroup(user.Group, "document_admin") if isAdmin || isDocsAdmin { return true } @@ -46,8 +45,8 @@ func CreateDocumentNode(c *fiber.Ctx) error { } uuid := uuid.New().String() - token := c.Locals("user").(*jwt.Token) - username := dao.GetUsernameFromToken(token) + user := c.Locals("currentUser").(model.User) + username := user.Username if err := dao.CreateDocumentNode(r.ParentId, r.IsFolder, r.Private, r.Name, uuid, username); err != nil { return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{ "error": err.Error(), @@ -108,8 +107,8 @@ func UpdateDocumentNodeContent(c *fiber.Ctx) error { } id := c.Params("id") - token := c.Locals("user").(*jwt.Token) - username := dao.GetUsernameFromToken(token) + user := c.Locals("currentUser").(model.User) + username := user.Username type contentRequest struct { Type string `json:"type"` @@ -312,7 +311,11 @@ func UploadDocumentFile(c *fiber.Ctx) error { if err != nil { return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{"error": err}) } - if err := c.SaveFile(file, fmt.Sprintf("./contents/%s/%s", id, storedName)); err != nil { + contentPath, err := util.SafeContentPath("./contents", id, storedName) + if err != nil { + return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{"error": err}) + } + if err := c.SaveFile(file, contentPath); err != nil { return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{"error": err}) } return c.JSON(fiber.Map{"url": fmt.Sprintf("/contents/%s/%s", id, storedName)}) diff --git a/service/server.go b/service/server.go index dcf7344..9699550 100644 --- a/service/server.go +++ b/service/server.go @@ -8,16 +8,15 @@ import ( "time" "github.com/gofiber/fiber/v2" - "github.com/golang-jwt/jwt/v5" "github.com/google/uuid" "github.com/millkhan/mcstatusgo/v2" ) func checkServerPermission(c *fiber.Ctx) bool { // Check if user is admin or news_admin - token := c.Locals("user").(*jwt.Token) - isAdmin := dao.IsUserInGroup(token, "admin") - isNewsAdmin := dao.IsUserInGroup(token, "server_admin") + user := c.Locals("currentUser").(model.User) + isAdmin := dao.ContainsGroup(user.Group, "admin") + isNewsAdmin := dao.ContainsGroup(user.Group, "server_admin") if isAdmin || isNewsAdmin { return false } @@ -57,7 +56,17 @@ func GetServerList(c *fiber.Ctx) error { }) } +var statusSlots = make(chan struct{}, 16) + func GetServerStatus(c *fiber.Ctx) error { + select { + case statusSlots <- struct{}{}: + defer func() { <-statusSlots }() + default: + return c.Status(fiber.StatusTooManyRequests).JSON(fiber.Map{ + "error": "Service busy", + }) + } type Request struct { ServerUrl string `json:"serverUrl"` } diff --git a/service/user.go b/service/user.go index 28b8a0c..eaa19bc 100644 --- a/service/user.go +++ b/service/user.go @@ -3,22 +3,14 @@ package service import ( "encoding/json" "necore/dao" + "necore/model" "github.com/gofiber/fiber/v2" - "github.com/golang-jwt/jwt/v5" ) func GetUserInfo(c *fiber.Ctx) error { userId := c.Params("id") - // // Check if user is admin or himself - // token := c.Locals("user").(*jwt.Token) - // isAdmin := dao.IsUserInGroup(token, "admin") - // tokenUsername := dao.GetUsernameFromToken(token) - // if !isAdmin && tokenUsername != userId { - // return c.Status(fiber.StatusForbidden).JSON(fiber.Map{"error": "Forbidden"}) - // } - userModel, err := dao.GetUserByUsername(userId) if err != nil || userModel == nil { return c.Status(404).JSON(fiber.Map{"error": "User not found"}) @@ -93,27 +85,33 @@ func GetUserList(c *fiber.Ctx) error { func DeleteUser(c *fiber.Ctx) error { // Must be admin - token := c.Locals("user").(*jwt.Token) - if !dao.IsUserInGroup(token, "admin") { + user := c.Locals("currentUser").(model.User) + if !dao.ContainsGroup(user.Group, "admin") { return c.Status(fiber.StatusForbidden).JSON(fiber.Map{"error": "Forbidden"}) } username := c.Params("id") - err := dao.DeleteUserByUsername(username) + err := dao.UpdateUserPermissions(username) if err != nil { return c.Status(500).JSON(fiber.Map{"error": "Internal server error"}) } + err = dao.DeleteUserByUsername(username) + if err != nil { + return c.Status(500).JSON(fiber.Map{"error": "Internal server error"}) + } + return c.SendStatus(200) } func UpdateUserPassword(c *fiber.Ctx) error { - token := c.Locals("user").(*jwt.Token) - isAdmin := dao.IsUserInGroup(token, "admin") - tokenUsername := dao.GetUsernameFromToken(token) + user := c.Locals("currentUser").(model.User) + isAdmin := dao.ContainsGroup(user.Group, "admin") + tokenUsername := user.Username type Payload struct { - Id string `json:"id"` - Password string `json:"new_password"` + Id string `json:"id"` + OldPassword string `json:"self_password"` + NewPassword string `json:"new_password"` } payload := new(Payload) if err := c.BodyParser(payload); err != nil { @@ -125,7 +123,21 @@ func UpdateUserPassword(c *fiber.Ctx) error { return c.Status(fiber.StatusForbidden).JSON(fiber.Map{"error": "Forbidden"}) } - if err := dao.UpdateUserPassword(payload.Id, payload.Password); err != nil { + userModel, err := dao.GetUserByUsername(payload.Id) + + if err != nil { + return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{"error": "Internal Server Error", "err": err}) + } + + if !dao.CheckUserPassword(payload.OldPassword, userModel.Password) { + return c.Status(fiber.StatusUnauthorized).JSON(fiber.Map{"error": "Invalid identity or password"}) + } + + if err := dao.UpdateUserPassword(payload.Id, payload.NewPassword); err != nil { + return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{"error": "Internal server error"}) + } + + if err := dao.UpdateUserPermissions(payload.Id); err != nil { return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{"error": "Internal server error"}) } @@ -135,8 +147,8 @@ func UpdateUserPassword(c *fiber.Ctx) error { func UpdateUserInfo(c *fiber.Ctx) error { // Must be admin - token := c.Locals("user").(*jwt.Token) - if !dao.IsUserInGroup(token, "admin") { + user := c.Locals("currentUser").(model.User) + if !dao.ContainsGroup(user.Group, "admin") { return c.Status(fiber.StatusForbidden).JSON(fiber.Map{"error": "Forbidden"}) } type PayloadTags struct { @@ -166,6 +178,10 @@ func UpdateUserInfo(c *fiber.Ctx) error { return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{"error": "Internal server error"}) } + if err := dao.UpdateUserPermissions(payload.Username); err != nil { + return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{"error": "Internal server error"}) + } + return c.SendStatus(fiber.StatusOK) } @@ -192,9 +208,9 @@ func UpdateUserAvatar(c *fiber.Ctx) error { } // Check if user is admin or himself - token := c.Locals("user").(*jwt.Token) - isAdmin := dao.IsUserInGroup(token, "admin") - tokenUsername := dao.GetUsernameFromToken(token) + user := c.Locals("currentUser").(model.User) + isAdmin := dao.ContainsGroup(user.Group, "admin") + tokenUsername := user.Username if !(isAdmin || tokenUsername == payload.Username) { return c.Status(fiber.StatusForbidden).JSON(fiber.Map{"error": "Forbidden"}) } From d7f1f7d2fdb5e0d90e39586a09a7605cbe9ded7a Mon Sep 17 00:00:00 2001 From: Kingcq <404291187@qq.com> Date: Sat, 20 Jun 2026 20:08:06 +0800 Subject: [PATCH 08/13] fix: security issues --- .github/workflows/go.yml | 3 + controller/middleware/auth.go | 11 +- controller/middleware/validator.go | 70 +++++-- dao/article.go | 40 +++- dao/bottoken.go | 15 +- dao/document.go | 183 +++++++++++++++--- dao/server.go | 41 +++- dao/user.go | 17 +- go.mod | 3 + go.sum | 9 + routes_test.go => routes_and_security_test.go | 171 +++++++++++++++- service/article.go | 26 +-- service/auth.go | 6 +- service/bottoken.go | 47 ++++- service/document.go | 8 +- ws/hub.go | 38 +++- 16 files changed, 583 insertions(+), 105 deletions(-) rename routes_test.go => routes_and_security_test.go (86%) diff --git a/.github/workflows/go.yml b/.github/workflows/go.yml index 4fc439f..052535d 100644 --- a/.github/workflows/go.yml +++ b/.github/workflows/go.yml @@ -27,6 +27,9 @@ jobs: go get -u echo ${{github.ref_type}} + - name: Run Tests + run: go test + - name: Build run: go build -o necore diff --git a/controller/middleware/auth.go b/controller/middleware/auth.go index eda0eda..98fdfdf 100644 --- a/controller/middleware/auth.go +++ b/controller/middleware/auth.go @@ -18,10 +18,9 @@ func AuthNeeded() fiber.Handler { } func jwtError(c *fiber.Ctx, err error) error { - if err.Error() == "Missing or malformed JWT" { - return c.Status(fiber.StatusBadRequest). - JSON(fiber.Map{"error": "Missing or malformed JWT", "err": nil}) - } - return c.Status(fiber.StatusUnauthorized). - JSON(fiber.Map{"error": "Invalid or expired JWT", "err": nil}) + c.Set(fiber.HeaderWWWAuthenticate, `Bearer realm="necore"`) + + return c.Status(fiber.StatusUnauthorized).JSON(fiber.Map{ + "error": "Unauthorized", + }) } diff --git a/controller/middleware/validator.go b/controller/middleware/validator.go index ccbf3b4..df28c77 100644 --- a/controller/middleware/validator.go +++ b/controller/middleware/validator.go @@ -1,9 +1,12 @@ package middleware import ( + "encoding/json" "errors" + "math" "necore/database" "necore/model" + "strconv" "github.com/gofiber/fiber/v2" "github.com/golang-jwt/jwt/v5" @@ -13,30 +16,22 @@ import ( func validateTokenVersion(c *fiber.Ctx) error { token, ok := c.Locals("user").(*jwt.Token) if !ok || token == nil { - return c.Status(fiber.StatusUnauthorized).JSON(fiber.Map{ - "error": "Unauthorized", - }) + return invalidToken(c) } claims, ok := token.Claims.(jwt.MapClaims) if !ok { - return c.Status(fiber.StatusUnauthorized).JSON(fiber.Map{ - "error": "Invalid token", - }) + return invalidToken(c) } username, ok := claims["name"].(string) if !ok || username == "" { - return c.Status(fiber.StatusUnauthorized).JSON(fiber.Map{ - "error": "Invalid token", - }) + return invalidToken(c) } - tokenVersionFloat, ok := claims["ver"].(float64) + tokenVersion, ok := getUintClaim(claims, "ver") if !ok { - return c.Status(fiber.StatusUnauthorized).JSON(fiber.Map{ - "error": "Invalid token", - }) + return invalidToken(c) } var user model.User @@ -46,9 +41,7 @@ func validateTokenVersion(c *fiber.Ctx) error { First(&user).Error if errors.Is(err, gorm.ErrRecordNotFound) { - return c.Status(fiber.StatusUnauthorized).JSON(fiber.Map{ - "error": "User no longer exists", - }) + return invalidToken(c) } if err != nil { @@ -57,15 +50,52 @@ func validateTokenVersion(c *fiber.Ctx) error { }) } - if uint(tokenVersionFloat) != user.TokenVersion { + if tokenVersion != user.TokenVersion { return c.Status(fiber.StatusUnauthorized).JSON(fiber.Map{ "error": "Token has been revoked", }) } - // 将数据库中的最新用户信息放入 Locals, - // 后续权限中间件直接使用,不再信任 JWT 中的 group。 c.Locals("currentUser", user) - return c.Next() } + +func getUintClaim(claims jwt.MapClaims, key string) (uint, bool) { + value, ok := claims[key] + if !ok || value == nil { + return 0, false + } + + switch v := value.(type) { + case float64: + if v < 0 || math.Trunc(v) != v { + return 0, false + } + return uint(v), true + + case json.Number: + parsed, err := strconv.ParseUint(v.String(), 10, 64) + if err != nil { + return 0, false + } + return uint(parsed), true + + case int: + if v < 0 { + return 0, false + } + return uint(v), true + + case uint: + return v, true + + default: + return 0, false + } +} + +func invalidToken(c *fiber.Ctx) error { + return c.Status(fiber.StatusUnauthorized).JSON(fiber.Map{ + "error": "Invalid token", + }) +} diff --git a/dao/article.go b/dao/article.go index 3a3ab2d..5b0530d 100644 --- a/dao/article.go +++ b/dao/article.go @@ -18,8 +18,30 @@ func CreateArticle(id string) error { } func UpdateArticle(updatedArticle model.Article) error { - db := database.GetArticleDatabase() - return db.Save(&updatedArticle).Error + result := database.GetArticleDatabase(). + Model(&model.Article{}). + Where("id = ?", updatedArticle.Id). + Updates(map[string]any{ + "pin": updatedArticle.Pin, + "title": updatedArticle.Title, + "brief": updatedArticle.Brief, + "date": updatedArticle.Date, + "end_date": updatedArticle.EndDate, + "image": updatedArticle.Image, + "content": updatedArticle.Content, + "author": updatedArticle.Author, + "category": updatedArticle.Category, + }) + + if result.Error != nil { + return result.Error + } + + if result.RowsAffected == 0 { + return fmt.Errorf("Article not found") + } + + return nil } func GetArticle(id string) (*model.Article, error) { @@ -67,6 +89,16 @@ func GetArticleList(target string, page int, pageSize int, pin bool) ([]model.Ar func DeleteArticle(id string) error { db := database.GetArticleDatabase() - os.RemoveAll(fmt.Sprintf("./contents/%s", id)) - return db.Where(&model.Article{Id: id}).Delete(&model.Article{}).Error + + result := db.Where("id = ?", id).Delete(&model.Article{}) + if result.Error != nil { + return result.Error + } + + if result.RowsAffected == 0 { + return fmt.Errorf("Article not found") + } + + _ = os.RemoveAll(fmt.Sprintf("./contents/%s", id)) + return nil } diff --git a/dao/bottoken.go b/dao/bottoken.go index 32924f7..478c392 100644 --- a/dao/bottoken.go +++ b/dao/bottoken.go @@ -3,9 +3,13 @@ package dao import ( "crypto/sha256" "encoding/hex" + "errors" + "fmt" "necore/database" "necore/model" "necore/util" + + "gorm.io/gorm" ) func CreateBotToken(name string) (*model.BotToken, error) { @@ -39,10 +43,17 @@ func GetBotTokens() []model.BotToken { func GetBotToken(name string) (*model.BotToken, error) { var token model.BotToken - db := database.GetBotTokenDatabase() - if err := db.Where(&model.BotToken{Name: name}).First(&token).Error; err != nil { + err := database.GetBotTokenDatabase(). + Where("name = ?", name). + First(&token).Error + + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, fmt.Errorf("bot token '%s' not found", name) + } + if err != nil { return nil, err } + return &token, nil } diff --git a/dao/document.go b/dao/document.go index e0a8080..d0273b3 100644 --- a/dao/document.go +++ b/dao/document.go @@ -2,6 +2,7 @@ package dao import ( "encoding/json" + "errors" "fmt" "necore/database" "necore/model" @@ -11,6 +12,61 @@ import ( "gorm.io/gorm" ) +func validateDocumentParent(tx *gorm.DB, nodeID string, parentID string) error { + if parentID == "" { + return fmt.Errorf("Invalid parent ID") + } + + if parentID == nodeID { + return fmt.Errorf("Parent ID cannot be the same as the node ID") + } + + if parentID == "root" { + return nil + } + + var parent model.DocumentNode + err := tx.Where("id = ?", parentID).First(&parent).Error + if errors.Is(err, gorm.ErrRecordNotFound) { + return fmt.Errorf("Record not found") + } + if err != nil { + return err + } + + if !parent.IsFolder { + return fmt.Errorf("Parent ID must be a folder") + } + + seen := map[string]struct{}{ + nodeID: {}, + } + + current := parent + + for { + if _, exists := seen[current.Id]; exists { + return fmt.Errorf("Circular reference detected") + } + seen[current.Id] = struct{}{} + + if current.ParentId == "" || current.ParentId == "root" { + return nil + } + + var next model.DocumentNode + err := tx.Where("id = ?", current.ParentId).First(&next).Error + if errors.Is(err, gorm.ErrRecordNotFound) { + return fmt.Errorf("Record not found") + } + if err != nil { + return err + } + + current = next + } +} + func getCurrentTime() string { currenttime := time.Now() newtime := fmt.Sprintf("%d-%s-%d %d:%d:%d", currenttime.Year(), currenttime.Month().String(), currenttime.Day(), currenttime.Hour(), currenttime.Minute(), currenttime.Second()) @@ -19,40 +75,82 @@ func getCurrentTime() string { func CreateDocumentNode(parentId string, isFolder bool, private bool, name string, id string, username string) error { db := database.GetDocumentDatabase() - if parentId == id { - return fmt.Errorf("ParentId and Id cannot be the same") - } - contributors, _ := json.Marshal([]string{username}) - node := model.DocumentNode{ - ParentId: parentId, - IsFolder: isFolder, - Private: private, - Name: name, - Id: id, - Contributors: string(contributors), - UpdateTime: getCurrentTime(), - } - return db.Create(&node).Error + + return db.Transaction(func(tx *gorm.DB) error { + if err := validateDocumentParent(tx, id, parentId); err != nil { + return err + } + + contributors, _ := json.Marshal([]string{username}) + node := model.DocumentNode{ + ParentId: parentId, + IsFolder: isFolder, + Private: private, + Name: name, + Id: id, + Contributors: string(contributors), + UpdateTime: getCurrentTime(), + } + + return tx.Create(&node).Error + }) } func DeleteDocumentNode(id string) error { db := database.GetDocumentDatabase() + + var ids []string + if err := collectDocumentNodeIDs(db, id, &ids); err != nil { + return err + } + + if len(ids) == 0 { + return fmt.Errorf("No document nodes found") + } + + if err := db.Transaction(func(tx *gorm.DB) error { + result := tx.Where("id IN ?", ids).Delete(&model.DocumentNode{}) + if result.Error != nil { + return result.Error + } + return nil + }); err != nil { + return err + } + + for _, nodeID := range ids { + _ = os.RemoveAll(fmt.Sprintf("./contents/%s", nodeID)) + } + + return nil +} + +func collectDocumentNodeIDs(db *gorm.DB, id string, ids *[]string) error { var node model.DocumentNode - db.Where(&model.DocumentNode{Id: id}).First(&node) + err := db.Where("id = ?", id).First(&node).Error + if errors.Is(err, gorm.ErrRecordNotFound) { + return fmt.Errorf("Record not found") + } + if err != nil { + return err + } + + *ids = append(*ids, node.Id) - // Recursively delete all children if node.IsFolder { var children []model.DocumentNode - db.Where(&model.DocumentNode{ParentId: id}).Find(&children) + if err := db.Where("parent_id = ?", node.Id).Find(&children).Error; err != nil { + return err + } + for _, child := range children { - DeleteDocumentNode(child.Id) - db.Where(&model.DocumentNode{Id: child.Id}).Delete(&model.DocumentNode{}) + if err := collectDocumentNodeIDs(db, child.Id, ids); err != nil { + return err + } } - } else { - // Delete Files - os.RemoveAll(fmt.Sprintf("./contents/%s", id)) } - return db.Where(&model.DocumentNode{Id: id}).Delete(&model.DocumentNode{}).Error + + return nil } func UpdateDocumentNodeName(id string, name string) error { @@ -109,13 +207,38 @@ func checkCyclicDocumentNode(parentId string, id string, db *gorm.DB) bool { func UpdateDocumentNodeParentId(id string, parentId string) error { db := database.GetDocumentDatabase() - if parentId == id { - return fmt.Errorf("ParentId and Id cannot be the same") - } - if checkCyclicDocumentNode(parentId, id, db) { - return fmt.Errorf("Cyclic dependency detected") - } - return db.Model(&model.DocumentNode{}).Where(&model.DocumentNode{Id: id}).Updates(model.DocumentNode{ParentId: parentId}).Error + + return db.Transaction(func(tx *gorm.DB) error { + var node model.DocumentNode + err := tx.Where("id = ?", id).First(&node).Error + if errors.Is(err, gorm.ErrRecordNotFound) { + return fmt.Errorf("Record not found") + } + if err != nil { + return err + } + + if err := validateDocumentParent(tx, id, parentId); err != nil { + return err + } + + result := tx.Model(&model.DocumentNode{}). + Where("id = ?", id). + Updates(map[string]any{ + "parent_id": parentId, + "update_time": getCurrentTime(), + }) + + if result.Error != nil { + return result.Error + } + + if result.RowsAffected == 0 { + return fmt.Errorf("Record not found") + } + + return nil + }) } func GetDocumentNodeChildren(id string, private bool) ([]model.DocumentNode, error) { diff --git a/dao/server.go b/dao/server.go index 9364344..e9524d3 100644 --- a/dao/server.go +++ b/dao/server.go @@ -1,6 +1,7 @@ package dao import ( + "fmt" "necore/database" "necore/model" ) @@ -18,13 +19,41 @@ func AddServer(server model.Server) error { } func UpdateServer(server model.Server) error { - db := database.GetServerDatabase() - var s *model.Server - db.Where(&model.Server{Id: server.Id}).First(&s) - return db.Model(&s).Updates(server).Error + result := database.GetServerDatabase(). + Model(&model.Server{}). + Where("id = ?", server.Id). + Updates(map[string]any{ + "name": server.Name, + "icon": server.Icon, + "description": server.Description, + "realtime": server.Realtime, + "online_map_url": server.OnlineMapUrl, + "server_url": server.ServerUrl, + }) + + if result.Error != nil { + return result.Error + } + + if result.RowsAffected == 0 { + return fmt.Errorf("server not found") + } + + return nil } func DeleteServer(id string) error { - db := database.GetServerDatabase() - return db.Where(&model.Server{Id: id}).Delete(&model.Server{}).Error + result := database.GetServerDatabase(). + Where("id = ?", id). + Delete(&model.Server{}) + + if result.Error != nil { + return result.Error + } + + if result.RowsAffected == 0 { + return fmt.Errorf("server not found") + } + + return nil } diff --git a/dao/user.go b/dao/user.go index 7197330..617cb74 100644 --- a/dao/user.go +++ b/dao/user.go @@ -78,13 +78,20 @@ func GetUserByUsername(u string) (*model.User, error) { } func AddUserByUsername(username string, password string) error { - hash, _ := hashPassword(password) - db := database.GetUserDatabase() + hash, err := hashPassword(password) + if err != nil { + return err + } + user := model.User{ - Username: username, - Password: hash, + Username: username, + Password: hash, + Group: `[]`, + Tags: `[]`, + TokenVersion: 1, } - return db.Create(&user).Error + + return database.GetUserDatabase().Create(&user).Error } func GetAllUsers() ([]model.User, error) { diff --git a/go.mod b/go.mod index 5eda76f..87b59ea 100644 --- a/go.mod +++ b/go.mod @@ -19,6 +19,7 @@ require ( github.com/MicahParks/keyfunc/v2 v2.1.0 // indirect github.com/andybalholm/brotli v1.2.1 // indirect github.com/clipperhouse/uax29/v2 v2.7.0 // indirect + github.com/davecgh/go-spew v1.1.1 // indirect github.com/fasthttp/websocket v1.5.12 // indirect github.com/gofiber/storage/sqlite3/v2 v2.1.3 // indirect github.com/golang-jwt/jwt v3.2.2+incompatible // indirect @@ -30,6 +31,7 @@ require ( github.com/mattn/go-runewidth v0.0.24 // indirect github.com/mattn/go-sqlite3 v1.14.45 // indirect github.com/philhofer/fwd v1.2.0 // indirect + github.com/pmezard/go-difflib v1.0.0 // indirect github.com/rivo/uniseg v0.4.7 // indirect github.com/savsgio/gotils v0.0.0-20250924091648-bce9a52d7761 // indirect github.com/tidwall/gjson v1.18.0 // indirect @@ -42,4 +44,5 @@ require ( golang.org/x/net v0.55.0 // indirect golang.org/x/sys v0.46.0 // indirect golang.org/x/text v0.38.0 // indirect + gopkg.in/yaml.v3 v3.0.1 // indirect ) diff --git a/go.sum b/go.sum index ad8e271..6334ce2 100644 --- a/go.sum +++ b/go.sum @@ -8,6 +8,8 @@ github.com/andybalholm/brotli v1.2.1 h1:R+f5xP285VArJDRgowrfb9DqL18yVK0gKAW/F+eT github.com/andybalholm/brotli v1.2.1/go.mod h1:rzTDkvFWvIrjDXZHkuS16NPggd91W3kUSvPlQ1pLaKY= github.com/clipperhouse/uax29/v2 v2.7.0 h1:+gs4oBZ2gPfVrKPthwbMzWZDaAFPGYK72F0NJv2v7Vk= github.com/clipperhouse/uax29/v2 v2.7.0/go.mod h1:EFJ2TJMRUaplDxHKj1qAEhCtQPW2tJSwu5BF98AuoVM= +github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c= +github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= github.com/fasthttp/websocket v1.5.8 h1:k5DpirKkftIF/w1R8ZzjSgARJrs54Je9YJK37DL/Ah8= github.com/fasthttp/websocket v1.5.8/go.mod h1:d08g8WaT6nnyvg9uMm8K9zMYyDjfKyj3170AtPRuVU0= github.com/fasthttp/websocket v1.5.12 h1:e4RGPpWW2HTbL3zV0Y/t7g0ub294LkiuXXUuTOUInlE= @@ -81,6 +83,8 @@ github.com/millkhan/mcstatusgo/v2 v2.2.0 h1:uRyHiOvqlK+6Oz3za4hMWAktSLjaqD/QzyQC github.com/millkhan/mcstatusgo/v2 v2.2.0/go.mod h1:YUJHhrJzsQP4PoDXFo++7JzPU7TjdztFM519CPqKe5M= github.com/philhofer/fwd v1.2.0 h1:e6DnBTl7vGY+Gz322/ASL4Gyp1FspeMvx1RNDoToZuM= github.com/philhofer/fwd v1.2.0/go.mod h1:RqIHx9QI14HlwKwm98g9Re5prTQ6LdeRQn+gXJFxsJM= +github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= +github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= github.com/rivo/uniseg v0.2.0 h1:S1pD9weZBuJdFmowNwbpi7BJ8TNftyUImj/0WQi72jY= github.com/rivo/uniseg v0.2.0/go.mod h1:J6wj4VEh+S6ZtnVlnTBMWIodfgj8LQOQFoIToxlJtxc= github.com/rivo/uniseg v0.4.7 h1:WUdvkW8uEhrYfLC4ZzdpI2ztxP1I582+49Oc5Mq64VQ= @@ -89,6 +93,8 @@ github.com/savsgio/gotils v0.0.0-20240303185622-093b76447511 h1:KanIMPX0QdEdB4R3 github.com/savsgio/gotils v0.0.0-20240303185622-093b76447511/go.mod h1:sM7Mt7uEoCeFSCBM+qBrqvEo+/9vdmj19wzp3yzUhmg= github.com/savsgio/gotils v0.0.0-20250924091648-bce9a52d7761 h1:McifyVxygw1d67y6vxUqls2D46J8W9nrki9c8c0eVvE= github.com/savsgio/gotils v0.0.0-20250924091648-bce9a52d7761/go.mod h1:Vi9gvHvTw4yCUHIznFl5TPULS7aXwgaTByGeBY75Wko= +github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U= +github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U= github.com/tidwall/gjson v1.18.0 h1:FIDeeyB800efLX89e5a8Y0BNH+LOngJyGrIWxG2FKQY= github.com/tidwall/gjson v1.18.0/go.mod h1:/wbyibRr2FHMks5tjHJ5F8dMZh3AcwJEMf5vlfC0lxk= github.com/tidwall/match v1.1.1 h1:+Ho715JplO36QYgwN9PGYNhgZvoUSc9X2c80KVTi+GA= @@ -137,6 +143,9 @@ golang.org/x/text v0.34.0 h1:oL/Qq0Kdaqxa1KbNeMKwQq0reLCCaFtqu2eNuSeNHbk= golang.org/x/text v0.34.0/go.mod h1:homfLqTYRFyVYemLBFl5GgL/DWEiH5wcsQ5gSh1yziA= golang.org/x/text v0.38.0 h1:sXmwo9DwP3OK9EZ7PqAdaooSGozfl/3a6/xJcbzPRhE= golang.org/x/text v0.38.0/go.mod h1:YXZt3QhHUKYT53r2lLKFIVi6Ao1jdzrTR/KQ09qyxF4= +gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= +gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= +gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= gorm.io/driver/sqlite v1.6.0 h1:WHRRrIiulaPiPFmDcod6prc4l2VGVWHz80KspNsxSfQ= gorm.io/driver/sqlite v1.6.0/go.mod h1:AO9V1qIQddBESngQUKWL9yoH93HIeA1X6V633rBwyT8= gorm.io/gorm v1.30.1 h1:lSHg33jJTBxs2mgJRfRZeLDG+WZaHYCk3Wtfl6Ngzo4= diff --git a/routes_test.go b/routes_and_security_test.go similarity index 86% rename from routes_test.go rename to routes_and_security_test.go index e978b9a..f014562 100644 --- a/routes_test.go +++ b/routes_and_security_test.go @@ -645,7 +645,7 @@ func TestBotRoutes(t *testing.T) { createResp := doJSON(t, env, http.MethodPost, "/necore/bots/token", env.adminToken, fiber.Map{ "name": "unit-test", }) - assertStatus(t, createResp, http.StatusOK) + assertStatus(t, createResp, http.StatusCreated) createBody := decodeBody(t, createResp) tokenObj, ok := createBody["token"].(map[string]any) if !ok || tokenObj["token"] == "" { @@ -653,7 +653,7 @@ func TestBotRoutes(t *testing.T) { } assertStatus(t, doJSON(t, env, http.MethodGet, "/necore/bots/token", env.adminToken, nil), http.StatusOK) - assertStatus(t, doJSON(t, env, http.MethodGet, "/necore/bots/token/missing", env.adminToken, nil), http.StatusInternalServerError) + assertStatus(t, doJSON(t, env, http.MethodGet, "/necore/bots/token/missing", env.adminToken, nil), http.StatusNotFound) // 当前源码只要求“已登录”,没有 bot_admin 权限检查。 assertStatus(t, doJSON(t, env, http.MethodGet, "/necore/bots/status", env.userToken, nil), http.StatusOK) @@ -849,3 +849,170 @@ func TestTokenVersion_DeletedUserTokenIsRejected(t *testing.T) { // 被删除用户的旧 JWT 也必须被拒绝。 assertStatus(t, doJSON(t, env, http.MethodGet, "/necore/auth/status", charlieToken, nil), http.StatusUnauthorized) } + +// ======================================================= +// EXTENDED SECURITY & EDGE CASE TESTS (ADDED BY REVIEW) +// ======================================================= + +/* +1. AUTH MODULE FUZZ & ABUSE TESTS +*/ +func TestAuth_FuzzLoginInputs(t *testing.T) { + env := setupTestEnv(t) + + maliciousInputs := []string{ + `{"username":"' OR 1=1 --","password":"x"}`, + `{"username":"\u0000admin","password":"admin"}`, + `{"username":"","password":"123"}`, + `{"username":"admin","password":"".repeat(10000)}`, + } + + for _, body := range maliciousInputs { + req := httptest.NewRequest("POST", "/auth/login", strings.NewReader(body)) + req.Header.Set("Content-Type", "application/json") + + resp := executeRequest(t, env, req) + + if resp.StatusCode == 200 { + t.Fatalf("SECURITY ISSUE: malicious login input succeeded: %s", body) + } + } +} + +func TestAuth_BruteForceSimulation(t *testing.T) { + env := setupTestEnv(t) + + for i := 0; i < 20; i++ { + body := `{"username":"admin","password":"wrong"}` + req := httptest.NewRequest("POST", "/auth/login", strings.NewReader(body)) + req.Header.Set("Content-Type", "application/json") + + resp := executeRequest(t, env, req) + + if i > 10 && resp.StatusCode == 200 { + t.Fatal("POTENTIAL BRUTE FORCE ISSUE: login succeeded after repeated attempts") + } + } +} + +/* +2. USER MODULE ABUSE CASES +*/ +func TestUser_EnumerationAttack(t *testing.T) { + env := setupTestEnv(t) + + usernames := []string{"admin", "alice", "root", "test", "doesnotexist"} + + for _, u := range usernames { + req := httptest.NewRequest("GET", "/user/"+u, nil) + + resp := executeRequest(t, env, req) + + // 不允许通过错误信息区分用户是否存在(防 user enumeration) + if strings.Contains(string(resp.Body), "password") { + t.Fatalf("USER ENUMERATION LEAK DETECTED for user: %s", u) + } + } +} + +/* +3. NEWS / ARTICLE MODULE SECURITY +*/ +func TestNews_XSSPayloadPersistence(t *testing.T) { + env := setupTestEnv(t) + + payload := `{"title":"","content":"x"}` + req := httptest.NewRequest("POST", "/news", strings.NewReader(payload)) + req.Header.Set("Content-Type", "application/json") + + resp := executeRequest(t, env, req) + + if resp.StatusCode == 200 { + // 再次读取列表检查是否原样返回 + req2 := httptest.NewRequest("GET", "/news", nil) + resp2 := executeRequest(t, env, req2) + + if strings.Contains(string(resp2.Body), "