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: ``}, + {name: "日常安全巡检`onmouseover=alert(1)", icon: "📁"}, + {name: "日常安全巡检\x00", icon: "📁"}, + {name: strings.Repeat("分", maxGroupNameRunes+1), icon: "📁"}, + {name: "日常安全巡检", icon: strings.Repeat("📁", maxGroupIconRunes+1)}, + } + + for _, tt := range tests { + t.Run(tt.name+"/"+tt.icon, func(t *testing.T) { + if _, _, err := validateGroupFields(tt.name, tt.icon); err == nil { + t.Fatal("validateGroupFields returned nil error for unsafe input") + } + }) + } +}