From 95c8e3b1a27893a8ae7c95dfd1dfe29476b39e88 Mon Sep 17 00:00:00 2001
From: =?UTF-8?q?=E5=85=AC=E6=98=8E?=
<83812544+Ed1s0nZ@users.noreply.github.com>
Date: Thu, 6 Aug 2026 17:16:50 +0800
Subject: [PATCH] Add files via upload
---
internal/handler/group.go | 57 ++++++++++++++++++++++++++++++----
internal/handler/group_test.go | 41 ++++++++++++++++++++++++
2 files changed, 92 insertions(+), 6 deletions(-)
create mode 100644 internal/handler/group_test.go
diff --git a/internal/handler/group.go b/internal/handler/group.go
index 23045e32..b3df88ca 100644
--- a/internal/handler/group.go
+++ b/internal/handler/group.go
@@ -1,8 +1,11 @@
package handler
import (
+ "errors"
"net/http"
+ "strings"
"time"
+ "unicode/utf8"
"cyberstrike-ai/internal/database"
"cyberstrike-ai/internal/security"
@@ -17,6 +20,11 @@ type GroupHandler struct {
logger *zap.Logger
}
+const (
+ maxGroupNameRunes = 64
+ maxGroupIconRunes = 16
+)
+
// NewGroupHandler 创建新的分组处理器
func NewGroupHandler(db *database.DB, logger *zap.Logger) *GroupHandler {
return &GroupHandler{
@@ -25,6 +33,41 @@ func NewGroupHandler(db *database.DB, logger *zap.Logger) *GroupHandler {
}
}
+func validateGroupTextField(field, value string, maxRunes int, required bool) (string, error) {
+ value = strings.TrimSpace(value)
+ if value == "" {
+ if required {
+ return "", errors.New(field + "不能为空")
+ }
+ return "", nil
+ }
+ if utf8.RuneCountInString(value) > maxRunes {
+ return "", errors.New(field + "过长")
+ }
+ for _, r := range value {
+ switch r {
+ case '<', '>', '"', '\'', '`':
+ return "", errors.New(field + "包含非法字符")
+ }
+ if r < 0x20 || r == 0x7f {
+ return "", errors.New(field + "包含非法控制字符")
+ }
+ }
+ return value, nil
+}
+
+func validateGroupFields(name, icon string) (string, string, error) {
+ validName, err := validateGroupTextField("分组名称", name, maxGroupNameRunes, true)
+ if err != nil {
+ return "", "", err
+ }
+ validIcon, err := validateGroupTextField("分组图标", icon, maxGroupIconRunes, false)
+ if err != nil {
+ return "", "", err
+ }
+ return validName, validIcon, nil
+}
+
// CreateGroupRequest 创建分组请求
type CreateGroupRequest struct {
Name string `json:"name"`
@@ -39,13 +82,14 @@ func (h *GroupHandler) CreateGroup(c *gin.Context) {
return
}
- if req.Name == "" {
- c.JSON(http.StatusBadRequest, gin.H{"error": "分组名称不能为空"})
+ name, icon, err := validateGroupFields(req.Name, req.Icon)
+ if err != nil {
+ c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
session, _ := security.CurrentSession(c)
- group, err := h.db.CreateGroup(req.Name, req.Icon, session.UserID)
+ group, err := h.db.CreateGroup(name, icon, session.UserID)
if err != nil {
h.logger.Error("创建分组失败", zap.Error(err))
// 如果是名称重复错误,返回400状态码
@@ -111,12 +155,13 @@ func (h *GroupHandler) UpdateGroup(c *gin.Context) {
return
}
- if req.Name == "" {
- c.JSON(http.StatusBadRequest, gin.H{"error": "分组名称不能为空"})
+ name, icon, err := validateGroupFields(req.Name, req.Icon)
+ if err != nil {
+ c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
- if err := h.db.UpdateGroup(id, req.Name, req.Icon); err != nil {
+ if err := h.db.UpdateGroup(id, name, icon); err != nil {
h.logger.Error("更新分组失败", zap.Error(err))
// 如果是名称重复错误,返回400状态码
if err.Error() == "分组名称已存在" {
diff --git a/internal/handler/group_test.go b/internal/handler/group_test.go
new file mode 100644
index 00000000..1e7cbde8
--- /dev/null
+++ b/internal/handler/group_test.go
@@ -0,0 +1,41 @@
+package handler
+
+import (
+ "strings"
+ "testing"
+)
+
+func TestValidateGroupFieldsAllowsNormalNamesAndIcons(t *testing.T) {
+ name, icon, err := validateGroupFields(" 日常安全巡检 ", " 📁 ")
+ if err != nil {
+ t.Fatalf("validateGroupFields returned error: %v", err)
+ }
+ if name != "日常安全巡检" {
+ t.Fatalf("name = %q, want trimmed normal name", name)
+ }
+ if icon != "📁" {
+ t.Fatalf("icon = %q, want trimmed icon", icon)
+ }
+}
+
+func TestValidateGroupFieldsRejectsStoredXSSPayloads(t *testing.T) {
+ tests := []struct {
+ name string
+ icon string
+ }{
+ {name: `
`, icon: "📁"},
+ {name: "日常安全巡检", icon: `