mirror of
https://github.com/Ed1s0nZ/CyberStrikeAI.git
synced 2026-08-07 03:18:39 +02:00
Add files via upload
This commit is contained in:
@@ -1,8 +1,11 @@
|
|||||||
package handler
|
package handler
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"errors"
|
||||||
"net/http"
|
"net/http"
|
||||||
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
|
"unicode/utf8"
|
||||||
|
|
||||||
"cyberstrike-ai/internal/database"
|
"cyberstrike-ai/internal/database"
|
||||||
"cyberstrike-ai/internal/security"
|
"cyberstrike-ai/internal/security"
|
||||||
@@ -17,6 +20,11 @@ type GroupHandler struct {
|
|||||||
logger *zap.Logger
|
logger *zap.Logger
|
||||||
}
|
}
|
||||||
|
|
||||||
|
const (
|
||||||
|
maxGroupNameRunes = 64
|
||||||
|
maxGroupIconRunes = 16
|
||||||
|
)
|
||||||
|
|
||||||
// NewGroupHandler 创建新的分组处理器
|
// NewGroupHandler 创建新的分组处理器
|
||||||
func NewGroupHandler(db *database.DB, logger *zap.Logger) *GroupHandler {
|
func NewGroupHandler(db *database.DB, logger *zap.Logger) *GroupHandler {
|
||||||
return &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 创建分组请求
|
// CreateGroupRequest 创建分组请求
|
||||||
type CreateGroupRequest struct {
|
type CreateGroupRequest struct {
|
||||||
Name string `json:"name"`
|
Name string `json:"name"`
|
||||||
@@ -39,13 +82,14 @@ func (h *GroupHandler) CreateGroup(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
if req.Name == "" {
|
name, icon, err := validateGroupFields(req.Name, req.Icon)
|
||||||
c.JSON(http.StatusBadRequest, gin.H{"error": "分组名称不能为空"})
|
if err != nil {
|
||||||
|
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
session, _ := security.CurrentSession(c)
|
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 {
|
if err != nil {
|
||||||
h.logger.Error("创建分组失败", zap.Error(err))
|
h.logger.Error("创建分组失败", zap.Error(err))
|
||||||
// 如果是名称重复错误,返回400状态码
|
// 如果是名称重复错误,返回400状态码
|
||||||
@@ -111,12 +155,13 @@ func (h *GroupHandler) UpdateGroup(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
if req.Name == "" {
|
name, icon, err := validateGroupFields(req.Name, req.Icon)
|
||||||
c.JSON(http.StatusBadRequest, gin.H{"error": "分组名称不能为空"})
|
if err != nil {
|
||||||
|
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||||
return
|
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))
|
h.logger.Error("更新分组失败", zap.Error(err))
|
||||||
// 如果是名称重复错误,返回400状态码
|
// 如果是名称重复错误,返回400状态码
|
||||||
if err.Error() == "分组名称已存在" {
|
if err.Error() == "分组名称已存在" {
|
||||||
|
|||||||
@@ -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: `<img src=x onerror="alert(1)">`, icon: "📁"},
|
||||||
|
{name: "日常安全巡检", icon: `<svg onload=alert(1)>`},
|
||||||
|
{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")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user