diff --git a/internal/logic/admin/group/deleteNodeGroupLogic.go b/internal/logic/admin/group/deleteNodeGroupLogic.go index 16c89d4..8ae2470 100644 --- a/internal/logic/admin/group/deleteNodeGroupLogic.go +++ b/internal/logic/admin/group/deleteNodeGroupLogic.go @@ -7,6 +7,7 @@ import ( "github.com/perfect-panel/server/internal/model/group" "github.com/perfect-panel/server/internal/model/node" + "github.com/perfect-panel/server/internal/model/subscribe" "github.com/perfect-panel/server/internal/svc" "github.com/perfect-panel/server/internal/types" "github.com/perfect-panel/server/pkg/logger" @@ -48,6 +49,33 @@ func (l *DeleteNodeGroupLogic) DeleteNodeGroup(req *types.DeleteNodeGroupRequest return fmt.Errorf("cannot delete group with %d associated nodes, please migrate nodes first", nodeCount) } + var defaultSubscribeCount int64 + if err := l.svcCtx.DB.Model(&subscribe.Subscribe{}).Where("node_group_id = ?", nodeGroup.Id).Count(&defaultSubscribeCount).Error; err != nil { + logger.Errorf("failed to count subscribes with default group: %v", err) + return err + } + if defaultSubscribeCount > 0 { + return fmt.Errorf("cannot delete group referenced by %d subscribes' default node group", defaultSubscribeCount) + } + + var subscribeGroupCount int64 + if err := l.svcCtx.DB.Model(&subscribe.Subscribe{}).Where("JSON_CONTAINS(node_group_ids, ?)", fmt.Sprintf("[%d]", nodeGroup.Id)).Count(&subscribeGroupCount).Error; err != nil { + logger.Errorf("failed to count subscribes in group: %v", err) + return err + } + if subscribeGroupCount > 0 { + return fmt.Errorf("cannot delete group referenced by %d subscribes' node group list", subscribeGroupCount) + } + + var userSubscribeCount int64 + if err := l.svcCtx.DB.Table("user_subscribe").Where("node_group_id = ?", nodeGroup.Id).Count(&userSubscribeCount).Error; err != nil { + logger.Errorf("failed to count user subscribes in group: %v", err) + return err + } + if userSubscribeCount > 0 { + return fmt.Errorf("cannot delete group referenced by %d user subscribes", userSubscribeCount) + } + // 使用 GORM Transaction 删除节点组 return l.svcCtx.DB.Transaction(func(tx *gorm.DB) error { // 删除节点组 diff --git a/internal/logic/admin/group/exportGroupResultLogic.go b/internal/logic/admin/group/exportGroupResultLogic.go index a84befa..4c53564 100644 --- a/internal/logic/admin/group/exportGroupResultLogic.go +++ b/internal/logic/admin/group/exportGroupResultLogic.go @@ -4,6 +4,7 @@ import ( "bytes" "context" "encoding/csv" + "encoding/json" "fmt" "github.com/perfect-panel/server/internal/model/group" @@ -52,10 +53,9 @@ func (l *ExportGroupResultLogic) ExportGroupResult(req *types.ExportGroupResultR Email string `json:"email"` } var users []UserInfo - if err := l.svcCtx.DB.Raw("SELECT * FROM JSON_ARRAY(?)", detail.UserData).Scan(&users).Error; err != nil { - // 如果解析失败,尝试用标准 JSON 解析 + if err := json.Unmarshal([]byte(detail.UserData), &users); err != nil { logger.Errorf("failed to parse user data: %v", err) - continue + return nil, "", fmt.Errorf("parse group history user_data failed: %w", err) } // 查询节点组名称 @@ -123,7 +123,10 @@ func (l *ExportGroupResultLogic) ExportGroupResult(req *types.ExportGroupResultR result = append(result, csvData...) // 生成文件名 - filename := fmt.Sprintf("group_result_%d.csv", req.HistoryId) + filename := "group_result_current.csv" + if req.HistoryId != nil { + filename = fmt.Sprintf("group_result_%d.csv", *req.HistoryId) + } return result, filename, nil } diff --git a/internal/logic/admin/group/getGroupHistoryDetailLogic.go b/internal/logic/admin/group/getGroupHistoryDetailLogic.go index d868d55..2f79d83 100644 --- a/internal/logic/admin/group/getGroupHistoryDetailLogic.go +++ b/internal/logic/admin/group/getGroupHistoryDetailLogic.go @@ -76,16 +76,16 @@ func (l *GetGroupHistoryDetailLogic) GetGroupHistoryDetail(req *types.GetGroupHi configSnapshot := make(map[string]interface{}) configSnapshot["group_details"] = details - // 获取配置快照(从 system_config 读取) + // 获取配置快照(从 system 读取) var configValue string if history.GroupMode == "average" { - l.svcCtx.DB.Table("system_config"). - Where("`key` = ?", "group.average_config"). + l.svcCtx.DB.Table("system"). + Where("`category` = ? AND `key` = ?", "group", "average_config"). Select("value"). Scan(&configValue) } else if history.GroupMode == "traffic" { - l.svcCtx.DB.Table("system_config"). - Where("`key` = ?", "group.traffic_config"). + l.svcCtx.DB.Table("system"). + Where("`category` = ? AND `key` = ?", "group", "traffic_config"). Select("value"). Scan(&configValue) } diff --git a/internal/logic/admin/group/getGroupHistoryLogic.go b/internal/logic/admin/group/getGroupHistoryLogic.go index 6eee9c3..1db7b4c 100644 --- a/internal/logic/admin/group/getGroupHistoryLogic.go +++ b/internal/logic/admin/group/getGroupHistoryLogic.go @@ -44,9 +44,9 @@ func (l *GetGroupHistoryLogic) GetGroupHistory(req *types.GetGroupHistoryRequest return nil, err } - // 分页查询 - offset := (req.Page - 1) * req.Size - if err := query.Order("id DESC").Offset(offset).Limit(req.Size).Find(&histories).Error; err != nil { + page, size := normalizePagination(req.Page, req.Size) + offset := (page - 1) * size + if err := query.Order("id DESC").Offset(offset).Limit(size).Find(&histories).Error; err != nil { logger.Errorf("failed to find group histories: %v", err) return nil, err } diff --git a/internal/logic/admin/group/getNodeGroupListLogic.go b/internal/logic/admin/group/getNodeGroupListLogic.go index 9595393..cad068d 100644 --- a/internal/logic/admin/group/getNodeGroupListLogic.go +++ b/internal/logic/admin/group/getNodeGroupListLogic.go @@ -38,9 +38,9 @@ func (l *GetNodeGroupListLogic) GetNodeGroupList(req *types.GetNodeGroupListRequ return nil, err } - // 分页查询 - offset := (req.Page - 1) * req.Size - if err := query.Order("sort ASC").Offset(offset).Limit(req.Size).Find(&nodeGroups).Error; err != nil { + page, size := normalizePagination(req.Page, req.Size) + offset := (page - 1) * size + if err := query.Order("sort ASC").Offset(offset).Limit(size).Find(&nodeGroups).Error; err != nil { logger.Errorf("failed to find node groups: %v", err) return nil, err } diff --git a/internal/logic/admin/group/group_logic_test.go b/internal/logic/admin/group/group_logic_test.go new file mode 100644 index 0000000..83b3bd8 --- /dev/null +++ b/internal/logic/admin/group/group_logic_test.go @@ -0,0 +1,346 @@ +package group + +import ( + "context" + "encoding/csv" + "fmt" + "strings" + "testing" + "time" + + "github.com/DATA-DOG/go-sqlmock" + modelgroup "github.com/perfect-panel/server/internal/model/group" + "github.com/perfect-panel/server/internal/svc" + "github.com/perfect-panel/server/internal/types" + "github.com/perfect-panel/server/pkg/logger" + "gorm.io/driver/mysql" + "gorm.io/gorm" +) + +func TestGetGroupHistoryDetailReadsConfigFromSystem(t *testing.T) { + db, mock, cleanup := newGroupTestDB(t) + defer cleanup() + + now := time.Unix(1710000000, 0) + mock.ExpectQuery("FROM `group_history`"). + WithArgs(int64(9), 1). + WillReturnRows(sqlmock.NewRows([]string{ + "id", "group_mode", "trigger_type", "state", "total_users", "success_count", "failed_count", "start_time", "end_time", "error_message", "created_at", + }).AddRow(int64(9), "traffic", "manual", "completed", 1, 1, 0, now, now, "", now)) + mock.ExpectQuery("FROM `group_history_detail`"). + WithArgs(int64(9)). + WillReturnRows(sqlmock.NewRows([]string{ + "id", "history_id", "node_group_id", "user_count", "node_count", "user_data", "created_at", + }).AddRow(int64(1), int64(9), int64(3), 1, 2, `[{"id":7,"email":"u@example.com"}]`, now)) + mock.ExpectQuery("FROM `system`"). + WithArgs("group", "traffic_config"). + WillReturnRows(sqlmock.NewRows([]string{"value"}).AddRow(`{"strategy":"closed"}`)) + + logic := newTestGroupHistoryDetailLogic(db) + resp, err := logic.GetGroupHistoryDetail(&types.GetGroupHistoryDetailRequest{Id: 9}) + if err != nil { + t.Fatalf("GetGroupHistoryDetail error: %v", err) + } + if got := resp.ConfigSnapshot["config"].(map[string]interface{})["strategy"]; got != "closed" { + t.Fatalf("config snapshot strategy = %v, want closed", got) + } + assertGroupExpectations(t, mock) +} + +func TestExportGroupResultParsesHistoryUserDataJSON(t *testing.T) { + db, mock, cleanup := newGroupTestDB(t) + defer cleanup() + + historyID := int64(11) + now := time.Unix(1710000000, 0) + mock.ExpectQuery("FROM `group_history_detail`"). + WithArgs(historyID). + WillReturnRows(sqlmock.NewRows([]string{ + "id", "history_id", "node_group_id", "user_count", "node_count", "user_data", "created_at", + }).AddRow(int64(1), historyID, int64(8), 1, 2, `[{"id":42,"email":"u@example.com"}]`, now)) + mock.ExpectQuery("FROM `node_group`"). + WithArgs(int64(8), 1). + WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(int64(8), "VIP")) + + logic := newTestExportGroupResultLogic(db) + data, filename, err := logic.ExportGroupResult(&types.ExportGroupResultRequest{HistoryId: &historyID}) + if err != nil { + t.Fatalf("ExportGroupResult error: %v", err) + } + if filename != "group_result_11.csv" { + t.Fatalf("filename = %q, want group_result_11.csv", filename) + } + + records, err := csv.NewReader(strings.NewReader(strings.TrimPrefix(string(data), "\ufeff"))).ReadAll() + if err != nil { + t.Fatalf("read csv: %v", err) + } + if len(records) != 2 || strings.Join(records[1], ",") != "42,8,VIP" { + t.Fatalf("csv records = %#v, want user/group row", records) + } + assertGroupExpectations(t, mock) +} + +func TestExportGroupResultRejectsInvalidHistoryUserDataJSON(t *testing.T) { + db, mock, cleanup := newGroupTestDB(t) + defer cleanup() + + historyID := int64(12) + now := time.Unix(1710000000, 0) + mock.ExpectQuery("FROM `group_history_detail`"). + WithArgs(historyID). + WillReturnRows(sqlmock.NewRows([]string{ + "id", "history_id", "node_group_id", "user_count", "node_count", "user_data", "created_at", + }).AddRow(int64(1), historyID, int64(8), 1, 2, `{bad-json`, now)) + + logic := newTestExportGroupResultLogic(db) + _, _, err := logic.ExportGroupResult(&types.ExportGroupResultRequest{HistoryId: &historyID}) + if err == nil || !strings.Contains(err.Error(), "parse group history user_data failed") { + t.Fatalf("ExportGroupResult error = %v, want parse error", err) + } + assertGroupExpectations(t, mock) +} + +func TestDeleteNodeGroupRejectsSubscribeReferences(t *testing.T) { + tests := []struct { + name string + defaultSubscribeCount int64 + subscribeGroupCount int64 + userSubscribeCount int64 + wantErr string + }{ + {name: "subscribe default group", defaultSubscribeCount: 1, wantErr: "default node group"}, + {name: "subscribe group list", subscribeGroupCount: 1, wantErr: "node group list"}, + {name: "user subscribe group", userSubscribeCount: 1, wantErr: "user subscribes"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + db, mock, cleanup := newGroupTestDB(t) + defer cleanup() + + expectDeleteNodeGroupBaseChecks(mock, 5) + mock.ExpectQuery("FROM `subscribe`"). + WithArgs(int64(5)). + WillReturnRows(sqlmock.NewRows([]string{"count"}).AddRow(tt.defaultSubscribeCount)) + if tt.defaultSubscribeCount == 0 { + mock.ExpectQuery("FROM `subscribe`"). + WithArgs("[5]"). + WillReturnRows(sqlmock.NewRows([]string{"count"}).AddRow(tt.subscribeGroupCount)) + } + if tt.defaultSubscribeCount == 0 && tt.subscribeGroupCount == 0 { + mock.ExpectQuery("FROM `user_subscribe`"). + WithArgs(int64(5)). + WillReturnRows(sqlmock.NewRows([]string{"count"}).AddRow(tt.userSubscribeCount)) + } + + logic := newTestDeleteNodeGroupLogic(db) + err := logic.DeleteNodeGroup(&types.DeleteNodeGroupRequest{Id: 5}) + if err == nil || !strings.Contains(err.Error(), tt.wantErr) { + t.Fatalf("DeleteNodeGroup error = %v, want contains %q", err, tt.wantErr) + } + assertGroupExpectations(t, mock) + }) + } +} + +func TestResetGroupsRollsBackOnFailure(t *testing.T) { + db, mock, cleanup := newGroupTestDB(t) + defer cleanup() + + mock.ExpectBegin() + mock.ExpectExec("DELETE FROM `node_group`").WillReturnResult(sqlmock.NewResult(0, 1)) + mock.ExpectExec("UPDATE `subscribe`"). + WithArgs(int64(0), "[]"). + WillReturnResult(sqlmock.NewResult(0, 1)) + mock.ExpectExec("UPDATE `nodes`"). + WithArgs("[]"). + WillReturnError(fmt.Errorf("node update failed")) + mock.ExpectRollback() + + logic := newTestResetGroupsLogic(db) + err := logic.ResetGroups() + if err == nil || !strings.Contains(err.Error(), "node update failed") { + t.Fatalf("ResetGroups error = %v, want node update failure", err) + } + assertGroupExpectations(t, mock) +} + +func TestPreviewUserNodesIncludesPublicNodesInGroupMode(t *testing.T) { + db, mock, cleanup := newGroupTestDB(t) + defer cleanup() + + now := time.Unix(1710000000, 0) + mock.ExpectQuery("FROM `user_subscribe`"). + WithArgs(int64(100), int8(0), int8(1)). + WillReturnRows(sqlmock.NewRows([]string{"id", "user_id", "subscribe_id", "node_group_id"}). + AddRow(int64(1), int64(100), int64(10), int64(5))) + mock.ExpectQuery("FROM `subscribe`"). + WithArgs(int64(10)). + WillReturnRows(sqlmock.NewRows([]string{"id", "node_group_id", "node_group_ids", "nodes", "node_tags"}). + AddRow(int64(10), int64(0), "[]", "", "")) + mock.ExpectQuery("FROM `system`"). + WithArgs("group", "enabled"). + WillReturnRows(sqlmock.NewRows([]string{"value"}).AddRow("true")) + mock.ExpectQuery("FROM `nodes`"). + WithArgs(true). + WillReturnRows(sqlmock.NewRows([]string{ + "id", "name", "tags", "port", "address", "server_id", "protocol", "enabled", "sort", "node_group_ids", "created_at", "updated_at", + }). + AddRow(int64(1), "group-node", "", uint16(443), "g.example.com", int64(1), "vless", true, 1, "[5]", now, now). + AddRow(int64(2), "public-node", "", uint16(443), "p.example.com", int64(1), "vless", true, 2, "[]", now, now)) + mock.ExpectQuery("FROM `node_group`"). + WithArgs(int64(5)). + WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(int64(5), "Group A")) + + logic := newTestPreviewUserNodesLogic(db) + resp, err := logic.PreviewUserNodes(&types.PreviewUserNodesRequest{UserId: 100}) + if err != nil { + t.Fatalf("PreviewUserNodes error: %v", err) + } + if !hasNodeGroup(resp.NodeGroups, 0, "public-node") { + t.Fatalf("PreviewUserNodes node groups = %#v, want public node group", resp.NodeGroups) + } + assertGroupExpectations(t, mock) +} + +func TestTrafficRangeMatchingUsesClosedBounds(t *testing.T) { + min0, max10 := int64(0), int64(10) + min10, max20 := int64(10), int64(20) + nodeGroups := []modelgroup.NodeGroup{ + {Id: 1, MinTrafficGB: &min0, MaxTrafficGB: &max10}, + {Id: 2, MinTrafficGB: &min10, MaxTrafficGB: &max20}, + } + + tests := map[float64]int64{ + 0: 1, + 10: 1, + 15: 2, + 21: 0, + } + for used, want := range tests { + if got := matchTrafficNodeGroup(used, nodeGroups); got != want { + t.Fatalf("matchTrafficNodeGroup(%v) = %d, want %d", used, got, want) + } + } +} + +func TestNormalizePaginationDefaultsAndMax(t *testing.T) { + tests := []struct { + page, size int + wantPage int + wantSize int + }{ + {page: 0, size: 0, wantPage: 1, wantSize: defaultPageSize}, + {page: -1, size: -5, wantPage: 1, wantSize: defaultPageSize}, + {page: 2, size: maxPageSize + 1, wantPage: 2, wantSize: maxPageSize}, + } + + for _, tt := range tests { + gotPage, gotSize := normalizePagination(tt.page, tt.size) + if gotPage != tt.wantPage || gotSize != tt.wantSize { + t.Fatalf("normalizePagination(%d, %d) = (%d, %d), want (%d, %d)", tt.page, tt.size, gotPage, gotSize, tt.wantPage, tt.wantSize) + } + } +} + +func newGroupTestDB(t *testing.T) (*gorm.DB, sqlmock.Sqlmock, func()) { + t.Helper() + + sqlDB, mock, err := sqlmock.New(sqlmock.QueryMatcherOption(sqlmock.QueryMatcherFunc(func(expectedSQL, actualSQL string) error { + if strings.Contains(actualSQL, expectedSQL) { + return nil + } + return fmt.Errorf("actual sql %q does not contain %q", actualSQL, expectedSQL) + }))) + if err != nil { + t.Fatalf("create sqlmock: %v", err) + } + + db, err := gorm.Open(mysql.New(mysql.Config{Conn: sqlDB, SkipInitializeWithVersion: true}), &gorm.Config{}) + if err != nil { + _ = sqlDB.Close() + t.Fatalf("open gorm db: %v", err) + } + + return db, mock, func() { + _ = sqlDB.Close() + } +} + +func newTestGroupHistoryDetailLogic(db *gorm.DB) *GetGroupHistoryDetailLogic { + ctx := context.Background() + return &GetGroupHistoryDetailLogic{ + Logger: logger.WithContext(ctx), + ctx: ctx, + svcCtx: &svc.ServiceContext{DB: db}, + } +} + +func newTestExportGroupResultLogic(db *gorm.DB) *ExportGroupResultLogic { + ctx := context.Background() + return &ExportGroupResultLogic{ + Logger: logger.WithContext(ctx), + ctx: ctx, + svcCtx: &svc.ServiceContext{DB: db}, + } +} + +func newTestDeleteNodeGroupLogic(db *gorm.DB) *DeleteNodeGroupLogic { + ctx := context.Background() + return &DeleteNodeGroupLogic{ + Logger: logger.WithContext(ctx), + ctx: ctx, + svcCtx: &svc.ServiceContext{DB: db}, + } +} + +func newTestResetGroupsLogic(db *gorm.DB) *ResetGroupsLogic { + ctx := context.Background() + return &ResetGroupsLogic{ + Logger: logger.WithContext(ctx), + ctx: ctx, + svcCtx: &svc.ServiceContext{DB: db}, + } +} + +func newTestPreviewUserNodesLogic(db *gorm.DB) *PreviewUserNodesLogic { + ctx := context.Background() + return &PreviewUserNodesLogic{ + Logger: logger.WithContext(ctx), + ctx: ctx, + svcCtx: &svc.ServiceContext{DB: db}, + } +} + +func expectDeleteNodeGroupBaseChecks(mock sqlmock.Sqlmock, nodeGroupId int64) { + now := time.Unix(1710000000, 0) + mock.ExpectQuery("FROM `node_group`"). + WithArgs(nodeGroupId, 1). + WillReturnRows(sqlmock.NewRows([]string{"id", "name", "created_at", "updated_at"}). + AddRow(nodeGroupId, "Group", now, now)) + mock.ExpectQuery("FROM `nodes`"). + WithArgs(fmt.Sprintf("[%d]", nodeGroupId)). + WillReturnRows(sqlmock.NewRows([]string{"count"}).AddRow(0)) +} + +func hasNodeGroup(items []types.NodeGroupItem, id int64, nodeName string) bool { + for _, item := range items { + if item.Id != id { + continue + } + for _, n := range item.Nodes { + if n.Name == nodeName { + return true + } + } + } + return false +} + +func assertGroupExpectations(t *testing.T, mock sqlmock.Sqlmock) { + t.Helper() + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatalf("unmet sql expectations: %v", err) + } +} diff --git a/internal/logic/admin/group/helpers.go b/internal/logic/admin/group/helpers.go new file mode 100644 index 0000000..c609bf0 --- /dev/null +++ b/internal/logic/admin/group/helpers.go @@ -0,0 +1,37 @@ +package group + +import ( + "fmt" + + "github.com/perfect-panel/server/internal/model/node" + "gorm.io/gorm" +) + +const ( + defaultPageSize = 20 + maxPageSize = 100 +) + +func normalizePagination(page, size int) (int, int) { + if page <= 0 { + page = 1 + } + if size <= 0 { + size = defaultPageSize + } + if size > maxPageSize { + size = maxPageSize + } + return page, size +} + +func countNodesInGroup(db *gorm.DB, nodeGroupId int64) (int, error) { + var count int64 + err := db.Model(&node.Node{}). + Where("JSON_CONTAINS(node_group_ids, ?)", fmt.Sprintf("[%d]", nodeGroupId)). + Count(&count).Error + if err != nil { + return 0, err + } + return int(count), nil +} diff --git a/internal/logic/admin/group/previewUserNodesLogic.go b/internal/logic/admin/group/previewUserNodesLogic.go index 53b32da..115e1a0 100644 --- a/internal/logic/admin/group/previewUserNodesLogic.go +++ b/internal/logic/admin/group/previewUserNodesLogic.go @@ -196,10 +196,10 @@ func (l *PreviewUserNodesLogic) PreviewUserNodes(req *types.PreviewUserNodesRequ return nil, err } - // 6. 过滤出包含至少一个匹配节点组的节点(仅显示用户真正所在分组的节点,不包含公共节点) + // 6. 过滤出公共节点和至少一个匹配节点组的节点,与真实订阅下发保持一致 for _, n := range dbNodes { - // 节点未配置节点组(公共节点),预览时不显示 if len(n.NodeGroupIds) == 0 { + filteredNodes = append(filteredNodes, n) continue } @@ -450,9 +450,13 @@ func (l *PreviewUserNodesLogic) PreviewUserNodes(req *types.PreviewUserNodesRequ } } - // 预览模式不显示公共节点(node_group_ids 为空的节点),只展示用户真正所在分组的节点 if len(publicNodes) > 0 { - logger.Infof("[PreviewUserNodes] skipping %d public nodes (not in user's assigned group)", len(publicNodes)) + nodeGroupItems = append(nodeGroupItems, types.NodeGroupItem{ + Id: 0, + Name: "公共节点", + Nodes: publicNodes, + }) + logger.Infof("[PreviewUserNodes] adding public nodes group: nodes=%d", len(publicNodes)) } } else { diff --git a/internal/logic/admin/group/recalculateGroupLogic.go b/internal/logic/admin/group/recalculateGroupLogic.go index 9b485b9..2b3105d 100644 --- a/internal/logic/admin/group/recalculateGroupLogic.go +++ b/internal/logic/admin/group/recalculateGroupLogic.go @@ -679,21 +679,7 @@ func (l *RecalculateGroupLogic) executeTrafficGrouping(tx *gorm.DB, historyId in // 将字节转换为 GB usedTrafficGB := float64(us.UsedTraffic) / (1024 * 1024 * 1024) - // 查找匹配的流量范围(使用左闭右开区间 [Min, Max)) - var targetNodeGroupId int64 = 0 - for _, ng := range nodeGroups { - if ng.MinTrafficGB == nil || ng.MaxTrafficGB == nil { - continue - } - minTraffic := float64(*ng.MinTrafficGB) - maxTraffic := float64(*ng.MaxTrafficGB) - - // 检查是否在区间内 [min, max) - if usedTrafficGB >= minTraffic && usedTrafficGB < maxTraffic { - targetNodeGroupId = ng.Id - break - } - } + targetNodeGroupId := matchTrafficNodeGroup(usedTrafficGB, nodeGroups) // 如果没有匹配到任何范围,targetNodeGroupId 保持为 0(不分配节点组) @@ -734,7 +720,14 @@ func (l *RecalculateGroupLogic) executeTrafficGrouping(tx *gorm.DB, historyId in // 4. 创建分组历史详情记录(只统计有用户的节点组) nodeGroupCount := make(map[int64]int) // node_group_id -> node_count for _, ng := range nodeGroups { - nodeGroupCount[ng.Id] = 1 // 每个节点组计为1 + count, err := countNodesInGroup(tx, ng.Id) + if err != nil { + l.Errorw("failed to count nodes in group", + logger.Field("node_group_id", ng.Id), + logger.Field("error", err.Error())) + return affectedCount, err + } + nodeGroupCount[ng.Id] = count } for nodeGroupId, userCount := range groupUserCount { @@ -764,6 +757,20 @@ func (l *RecalculateGroupLogic) executeTrafficGrouping(tx *gorm.DB, historyId in return affectedCount, nil } +func matchTrafficNodeGroup(usedTrafficGB float64, nodeGroups []group.NodeGroup) int64 { + for _, ng := range nodeGroups { + if ng.MinTrafficGB == nil || ng.MaxTrafficGB == nil { + continue + } + minTraffic := float64(*ng.MinTrafficGB) + maxTraffic := float64(*ng.MaxTrafficGB) + if usedTrafficGB >= minTraffic && usedTrafficGB <= maxTraffic { + return ng.Id + } + } + return 0 +} + // containsIgnoreCase checks if a string contains another substring (case-insensitive) func containsIgnoreCase(s, substr string) bool { if len(substr) == 0 { diff --git a/internal/logic/admin/group/resetGroupsLogic.go b/internal/logic/admin/group/resetGroupsLogic.go index eaaa098..bc99cb7 100644 --- a/internal/logic/admin/group/resetGroupsLogic.go +++ b/internal/logic/admin/group/resetGroupsLogic.go @@ -7,8 +7,10 @@ import ( "github.com/perfect-panel/server/internal/model/node" "github.com/perfect-panel/server/internal/model/subscribe" "github.com/perfect-panel/server/internal/model/system" + "github.com/perfect-panel/server/internal/model/user" "github.com/perfect-panel/server/internal/svc" "github.com/perfect-panel/server/pkg/logger" + "gorm.io/gorm" ) type ResetGroupsLogic struct { @@ -27,55 +29,64 @@ func NewResetGroupsLogic(ctx context.Context, svcCtx *svc.ServiceContext) *Reset } func (l *ResetGroupsLogic) ResetGroups() error { - // 1. Delete all node groups - err := l.svcCtx.DB.Where("1 = 1").Delete(&group.NodeGroup{}).Error - if err != nil { - l.Errorw("Failed to delete all node groups", logger.Field("error", err.Error())) - return err - } - l.Infow("Successfully deleted all node groups") + err := l.svcCtx.DB.Transaction(func(tx *gorm.DB) error { + // 1. Delete all node groups + if err := tx.Where("1 = 1").Delete(&group.NodeGroup{}).Error; err != nil { + l.Errorw("Failed to delete all node groups", logger.Field("error", err.Error())) + return err + } + l.Infow("Successfully deleted all node groups") - // 2. Clear node_group_ids for all subscribes (products) - err = l.svcCtx.DB.Model(&subscribe.Subscribe{}).Where("1 = 1").Update("node_group_ids", "[]").Error - if err != nil { - l.Errorw("Failed to clear subscribes' node_group_ids", logger.Field("error", err.Error())) - return err - } - l.Infow("Successfully cleared all subscribes' node_group_ids") + // 2. Clear node_group_id/node_group_ids for all subscribes (products) + if err := tx.Table((&subscribe.Subscribe{}).TableName()).Where("1 = 1").Updates(map[string]interface{}{ + "node_group_id": 0, + "node_group_ids": "[]", + }).Error; err != nil { + l.Errorw("Failed to clear subscribes' node groups", logger.Field("error", err.Error())) + return err + } + l.Infow("Successfully cleared all subscribes' node groups") - // 3. Clear node_group_ids for all nodes - err = l.svcCtx.DB.Model(&node.Node{}).Where("1 = 1").Update("node_group_ids", "[]").Error - if err != nil { - l.Errorw("Failed to clear nodes' node_group_ids", logger.Field("error", err.Error())) - return err - } - l.Infow("Successfully cleared all nodes' node_group_ids") + // 3. Clear node_group_ids for all nodes + if err := tx.Table((&node.Node{}).TableName()).Where("1 = 1").Update("node_group_ids", "[]").Error; err != nil { + l.Errorw("Failed to clear nodes' node_group_ids", logger.Field("error", err.Error())) + return err + } + l.Infow("Successfully cleared all nodes' node_group_ids") - // 4. Clear group history - err = l.svcCtx.DB.Where("1 = 1").Delete(&group.GroupHistory{}).Error - if err != nil { - l.Errorw("Failed to clear group history", logger.Field("error", err.Error())) - // Non-critical error, continue anyway - } else { + // 4. Clear user_subscribe node_group_id + if err := tx.Table((&user.Subscribe{}).TableName()).Where("1 = 1").Update("node_group_id", 0).Error; err != nil { + l.Errorw("Failed to clear user subscribes' node_group_id", logger.Field("error", err.Error())) + return err + } + l.Infow("Successfully cleared all user subscribes' node_group_id") + + // 5. Clear group history + if err := tx.Where("1 = 1").Delete(&group.GroupHistory{}).Error; err != nil { + l.Errorw("Failed to clear group history", logger.Field("error", err.Error())) + return err + } l.Infow("Successfully cleared group history") - } - // 7. Clear group history details - err = l.svcCtx.DB.Where("1 = 1").Delete(&group.GroupHistoryDetail{}).Error - if err != nil { - l.Errorw("Failed to clear group history details", logger.Field("error", err.Error())) - // Non-critical error, continue anyway - } else { + // 6. Clear group history details + if err := tx.Where("1 = 1").Delete(&group.GroupHistoryDetail{}).Error; err != nil { + l.Errorw("Failed to clear group history details", logger.Field("error", err.Error())) + return err + } l.Infow("Successfully cleared group history details") - } - // 5. Delete all group config settings - err = l.svcCtx.DB.Where("`category` = ?", "group").Delete(&system.System{}).Error + // 7. Delete all group config settings + if err := tx.Where("`category` = ?", "group").Delete(&system.System{}).Error; err != nil { + l.Errorw("Failed to delete group config", logger.Field("error", err.Error())) + return err + } + l.Infow("Successfully deleted all group config settings") + + return nil + }) if err != nil { - l.Errorw("Failed to delete group config", logger.Field("error", err.Error())) return err } - l.Infow("Successfully deleted all group config settings") l.Infow("Group reset completed successfully") return nil