Co-authored-by: multica-agent <github@multica.ai>
This commit is contained in:
@@ -7,6 +7,7 @@ import (
|
|||||||
|
|
||||||
"github.com/perfect-panel/server/internal/model/group"
|
"github.com/perfect-panel/server/internal/model/group"
|
||||||
"github.com/perfect-panel/server/internal/model/node"
|
"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/svc"
|
||||||
"github.com/perfect-panel/server/internal/types"
|
"github.com/perfect-panel/server/internal/types"
|
||||||
"github.com/perfect-panel/server/pkg/logger"
|
"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)
|
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 删除节点组
|
// 使用 GORM Transaction 删除节点组
|
||||||
return l.svcCtx.DB.Transaction(func(tx *gorm.DB) error {
|
return l.svcCtx.DB.Transaction(func(tx *gorm.DB) error {
|
||||||
// 删除节点组
|
// 删除节点组
|
||||||
|
|||||||
@@ -4,6 +4,7 @@ import (
|
|||||||
"bytes"
|
"bytes"
|
||||||
"context"
|
"context"
|
||||||
"encoding/csv"
|
"encoding/csv"
|
||||||
|
"encoding/json"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
|
||||||
"github.com/perfect-panel/server/internal/model/group"
|
"github.com/perfect-panel/server/internal/model/group"
|
||||||
@@ -52,10 +53,9 @@ func (l *ExportGroupResultLogic) ExportGroupResult(req *types.ExportGroupResultR
|
|||||||
Email string `json:"email"`
|
Email string `json:"email"`
|
||||||
}
|
}
|
||||||
var users []UserInfo
|
var users []UserInfo
|
||||||
if err := l.svcCtx.DB.Raw("SELECT * FROM JSON_ARRAY(?)", detail.UserData).Scan(&users).Error; err != nil {
|
if err := json.Unmarshal([]byte(detail.UserData), &users); err != nil {
|
||||||
// 如果解析失败,尝试用标准 JSON 解析
|
|
||||||
logger.Errorf("failed to parse user data: %v", err)
|
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...)
|
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
|
return result, filename, nil
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -76,16 +76,16 @@ func (l *GetGroupHistoryDetailLogic) GetGroupHistoryDetail(req *types.GetGroupHi
|
|||||||
configSnapshot := make(map[string]interface{})
|
configSnapshot := make(map[string]interface{})
|
||||||
configSnapshot["group_details"] = details
|
configSnapshot["group_details"] = details
|
||||||
|
|
||||||
// 获取配置快照(从 system_config 读取)
|
// 获取配置快照(从 system 读取)
|
||||||
var configValue string
|
var configValue string
|
||||||
if history.GroupMode == "average" {
|
if history.GroupMode == "average" {
|
||||||
l.svcCtx.DB.Table("system_config").
|
l.svcCtx.DB.Table("system").
|
||||||
Where("`key` = ?", "group.average_config").
|
Where("`category` = ? AND `key` = ?", "group", "average_config").
|
||||||
Select("value").
|
Select("value").
|
||||||
Scan(&configValue)
|
Scan(&configValue)
|
||||||
} else if history.GroupMode == "traffic" {
|
} else if history.GroupMode == "traffic" {
|
||||||
l.svcCtx.DB.Table("system_config").
|
l.svcCtx.DB.Table("system").
|
||||||
Where("`key` = ?", "group.traffic_config").
|
Where("`category` = ? AND `key` = ?", "group", "traffic_config").
|
||||||
Select("value").
|
Select("value").
|
||||||
Scan(&configValue)
|
Scan(&configValue)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -44,9 +44,9 @@ func (l *GetGroupHistoryLogic) GetGroupHistory(req *types.GetGroupHistoryRequest
|
|||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
// 分页查询
|
page, size := normalizePagination(req.Page, req.Size)
|
||||||
offset := (req.Page - 1) * req.Size
|
offset := (page - 1) * size
|
||||||
if err := query.Order("id DESC").Offset(offset).Limit(req.Size).Find(&histories).Error; err != nil {
|
if err := query.Order("id DESC").Offset(offset).Limit(size).Find(&histories).Error; err != nil {
|
||||||
logger.Errorf("failed to find group histories: %v", err)
|
logger.Errorf("failed to find group histories: %v", err)
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -38,9 +38,9 @@ func (l *GetNodeGroupListLogic) GetNodeGroupList(req *types.GetNodeGroupListRequ
|
|||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
// 分页查询
|
page, size := normalizePagination(req.Page, req.Size)
|
||||||
offset := (req.Page - 1) * req.Size
|
offset := (page - 1) * size
|
||||||
if err := query.Order("sort ASC").Offset(offset).Limit(req.Size).Find(&nodeGroups).Error; err != nil {
|
if err := query.Order("sort ASC").Offset(offset).Limit(size).Find(&nodeGroups).Error; err != nil {
|
||||||
logger.Errorf("failed to find node groups: %v", err)
|
logger.Errorf("failed to find node groups: %v", err)
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -196,10 +196,10 @@ func (l *PreviewUserNodesLogic) PreviewUserNodes(req *types.PreviewUserNodesRequ
|
|||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
// 6. 过滤出包含至少一个匹配节点组的节点(仅显示用户真正所在分组的节点,不包含公共节点)
|
// 6. 过滤出公共节点和至少一个匹配节点组的节点,与真实订阅下发保持一致
|
||||||
for _, n := range dbNodes {
|
for _, n := range dbNodes {
|
||||||
// 节点未配置节点组(公共节点),预览时不显示
|
|
||||||
if len(n.NodeGroupIds) == 0 {
|
if len(n.NodeGroupIds) == 0 {
|
||||||
|
filteredNodes = append(filteredNodes, n)
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -450,9 +450,13 @@ func (l *PreviewUserNodesLogic) PreviewUserNodes(req *types.PreviewUserNodesRequ
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// 预览模式不显示公共节点(node_group_ids 为空的节点),只展示用户真正所在分组的节点
|
|
||||||
if len(publicNodes) > 0 {
|
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 {
|
} else {
|
||||||
|
|||||||
@@ -679,21 +679,7 @@ func (l *RecalculateGroupLogic) executeTrafficGrouping(tx *gorm.DB, historyId in
|
|||||||
// 将字节转换为 GB
|
// 将字节转换为 GB
|
||||||
usedTrafficGB := float64(us.UsedTraffic) / (1024 * 1024 * 1024)
|
usedTrafficGB := float64(us.UsedTraffic) / (1024 * 1024 * 1024)
|
||||||
|
|
||||||
// 查找匹配的流量范围(使用左闭右开区间 [Min, Max))
|
targetNodeGroupId := matchTrafficNodeGroup(usedTrafficGB, nodeGroups)
|
||||||
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 保持为 0(不分配节点组)
|
// 如果没有匹配到任何范围,targetNodeGroupId 保持为 0(不分配节点组)
|
||||||
|
|
||||||
@@ -734,7 +720,14 @@ func (l *RecalculateGroupLogic) executeTrafficGrouping(tx *gorm.DB, historyId in
|
|||||||
// 4. 创建分组历史详情记录(只统计有用户的节点组)
|
// 4. 创建分组历史详情记录(只统计有用户的节点组)
|
||||||
nodeGroupCount := make(map[int64]int) // node_group_id -> node_count
|
nodeGroupCount := make(map[int64]int) // node_group_id -> node_count
|
||||||
for _, ng := range nodeGroups {
|
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 {
|
for nodeGroupId, userCount := range groupUserCount {
|
||||||
@@ -764,6 +757,20 @@ func (l *RecalculateGroupLogic) executeTrafficGrouping(tx *gorm.DB, historyId in
|
|||||||
return affectedCount, nil
|
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)
|
// containsIgnoreCase checks if a string contains another substring (case-insensitive)
|
||||||
func containsIgnoreCase(s, substr string) bool {
|
func containsIgnoreCase(s, substr string) bool {
|
||||||
if len(substr) == 0 {
|
if len(substr) == 0 {
|
||||||
|
|||||||
@@ -7,8 +7,10 @@ import (
|
|||||||
"github.com/perfect-panel/server/internal/model/node"
|
"github.com/perfect-panel/server/internal/model/node"
|
||||||
"github.com/perfect-panel/server/internal/model/subscribe"
|
"github.com/perfect-panel/server/internal/model/subscribe"
|
||||||
"github.com/perfect-panel/server/internal/model/system"
|
"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/internal/svc"
|
||||||
"github.com/perfect-panel/server/pkg/logger"
|
"github.com/perfect-panel/server/pkg/logger"
|
||||||
|
"gorm.io/gorm"
|
||||||
)
|
)
|
||||||
|
|
||||||
type ResetGroupsLogic struct {
|
type ResetGroupsLogic struct {
|
||||||
@@ -27,56 +29,65 @@ func NewResetGroupsLogic(ctx context.Context, svcCtx *svc.ServiceContext) *Reset
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (l *ResetGroupsLogic) ResetGroups() error {
|
func (l *ResetGroupsLogic) ResetGroups() error {
|
||||||
|
err := l.svcCtx.DB.Transaction(func(tx *gorm.DB) error {
|
||||||
// 1. Delete all node groups
|
// 1. Delete all node groups
|
||||||
err := l.svcCtx.DB.Where("1 = 1").Delete(&group.NodeGroup{}).Error
|
if err := tx.Where("1 = 1").Delete(&group.NodeGroup{}).Error; err != nil {
|
||||||
if err != nil {
|
|
||||||
l.Errorw("Failed to delete all node groups", logger.Field("error", err.Error()))
|
l.Errorw("Failed to delete all node groups", logger.Field("error", err.Error()))
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
l.Infow("Successfully deleted all node groups")
|
l.Infow("Successfully deleted all node groups")
|
||||||
|
|
||||||
// 2. Clear node_group_ids for all subscribes (products)
|
// 2. Clear node_group_id/node_group_ids for all subscribes (products)
|
||||||
err = l.svcCtx.DB.Model(&subscribe.Subscribe{}).Where("1 = 1").Update("node_group_ids", "[]").Error
|
if err := tx.Table((&subscribe.Subscribe{}).TableName()).Where("1 = 1").Updates(map[string]interface{}{
|
||||||
if err != nil {
|
"node_group_id": 0,
|
||||||
l.Errorw("Failed to clear subscribes' node_group_ids", logger.Field("error", err.Error()))
|
"node_group_ids": "[]",
|
||||||
|
}).Error; err != nil {
|
||||||
|
l.Errorw("Failed to clear subscribes' node groups", logger.Field("error", err.Error()))
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
l.Infow("Successfully cleared all subscribes' node_group_ids")
|
l.Infow("Successfully cleared all subscribes' node groups")
|
||||||
|
|
||||||
// 3. Clear node_group_ids for all nodes
|
// 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 := tx.Table((&node.Node{}).TableName()).Where("1 = 1").Update("node_group_ids", "[]").Error; err != nil {
|
||||||
if err != nil {
|
|
||||||
l.Errorw("Failed to clear nodes' node_group_ids", logger.Field("error", err.Error()))
|
l.Errorw("Failed to clear nodes' node_group_ids", logger.Field("error", err.Error()))
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
l.Infow("Successfully cleared all nodes' node_group_ids")
|
l.Infow("Successfully cleared all nodes' node_group_ids")
|
||||||
|
|
||||||
// 4. Clear group history
|
// 4. Clear user_subscribe node_group_id
|
||||||
err = l.svcCtx.DB.Where("1 = 1").Delete(&group.GroupHistory{}).Error
|
if err := tx.Table((&user.Subscribe{}).TableName()).Where("1 = 1").Update("node_group_id", 0).Error; err != nil {
|
||||||
if 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()))
|
l.Errorw("Failed to clear group history", logger.Field("error", err.Error()))
|
||||||
// Non-critical error, continue anyway
|
return err
|
||||||
} else {
|
}
|
||||||
l.Infow("Successfully cleared group history")
|
l.Infow("Successfully cleared group history")
|
||||||
}
|
|
||||||
|
|
||||||
// 7. Clear group history details
|
// 6. Clear group history details
|
||||||
err = l.svcCtx.DB.Where("1 = 1").Delete(&group.GroupHistoryDetail{}).Error
|
if err := tx.Where("1 = 1").Delete(&group.GroupHistoryDetail{}).Error; err != nil {
|
||||||
if err != nil {
|
|
||||||
l.Errorw("Failed to clear group history details", logger.Field("error", err.Error()))
|
l.Errorw("Failed to clear group history details", logger.Field("error", err.Error()))
|
||||||
// Non-critical error, continue anyway
|
return err
|
||||||
} else {
|
|
||||||
l.Infow("Successfully cleared group history details")
|
|
||||||
}
|
}
|
||||||
|
l.Infow("Successfully cleared group history details")
|
||||||
|
|
||||||
// 5. Delete all group config settings
|
// 7. Delete all group config settings
|
||||||
err = l.svcCtx.DB.Where("`category` = ?", "group").Delete(&system.System{}).Error
|
if err := tx.Where("`category` = ?", "group").Delete(&system.System{}).Error; err != nil {
|
||||||
if err != nil {
|
|
||||||
l.Errorw("Failed to delete group config", logger.Field("error", err.Error()))
|
l.Errorw("Failed to delete group config", logger.Field("error", err.Error()))
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
l.Infow("Successfully deleted all group config settings")
|
l.Infow("Successfully deleted all group config settings")
|
||||||
|
|
||||||
|
return nil
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
l.Infow("Group reset completed successfully")
|
l.Infow("Group reset completed successfully")
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user