From 5efe79093726706db5284ee2acea1b6b622ffaeb Mon Sep 17 00:00:00 2001 From: sagit <36596628+Sagit-chu@users.noreply.github.com> Date: Tue, 31 Mar 2026 10:52:49 +0800 Subject: [PATCH] =?UTF-8?q?fix:=20=E5=AE=9E=E7=8E=B0=E7=94=A8=E6=88=B7?= =?UTF-8?q?=E7=AB=AF=E5=8F=A3=E6=95=B0=E9=87=8F=E9=99=90=E5=88=B6=E9=AA=8C?= =?UTF-8?q?=E8=AF=81=20(#399)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 在 ensureUserTunnelForwardAllowed 中添加 User.Num 限制验证 - 在 ensureUserTunnelForwardAllowed 中添加 UserTunnel.Num 限制验证 - 新增 CountActiveForwardsByUser 和 CountActiveForwardsByUserTunnel 函数 - 添加转发数量限制的契约测试 Closes #390 --- .../internal/http/handler/flow_policy.go | 23 +- .../internal/store/repo/repository_flow.go | 18 + .../forward_count_limit_contract_test.go | 346 ++++++++++++++++++ 3 files changed, 386 insertions(+), 1 deletion(-) create mode 100644 go-backend/tests/contract/forward_count_limit_contract_test.go diff --git a/go-backend/internal/http/handler/flow_policy.go b/go-backend/internal/http/handler/flow_policy.go index 9f0321d..77885e1 100644 --- a/go-backend/internal/http/handler/flow_policy.go +++ b/go-backend/internal/http/handler/flow_policy.go @@ -22,6 +22,7 @@ type userTunnelPolicy struct { OutFlow int64 ExpTime int64 Status int + Num int } type gostConfigSnapshot struct { @@ -363,6 +364,16 @@ func (h *Handler) ensureUserTunnelForwardAllowed(userID int64, tunnelID int64, n return err } + if user.Num > 0 { + currentForwardCount, err := h.repo.CountActiveForwardsByUser(userID) + if err != nil { + return err + } + if currentForwardCount >= int64(user.Num) { + return errors.New("转发数量已达上限") + } + } + userTunnelID, _, _, err := h.resolveUserTunnelAndLimiter(userID, tunnelID) if err != nil { return err @@ -392,6 +403,16 @@ func (h *Handler) ensureUserTunnelForwardAllowed(userID int64, tunnelID int64, n return errors.New("该隧道流量已超额,禁止开启转发") } + if policy.Num > 0 { + currentTunnelForwardCount, err := h.repo.CountActiveForwardsByUserTunnel(userID, tunnelID) + if err != nil { + return err + } + if currentTunnelForwardCount >= int64(policy.Num) { + return errors.New("该隧道转发数量已达上限") + } + } + return nil } @@ -442,7 +463,7 @@ func (h *Handler) getUserTunnelPolicy(userTunnelID int64) (*userTunnelPolicy, er return &userTunnelPolicy{ ID: ut.ID, UserID: ut.UserID, TunnelID: ut.TunnelID, Flow: ut.Flow, InFlow: ut.InFlow, OutFlow: ut.OutFlow, - ExpTime: ut.ExpTime, Status: ut.Status, + ExpTime: ut.ExpTime, Status: ut.Status, Num: ut.Num, }, nil } diff --git a/go-backend/internal/store/repo/repository_flow.go b/go-backend/internal/store/repo/repository_flow.go index 603ab46..2248df2 100644 --- a/go-backend/internal/store/repo/repository_flow.go +++ b/go-backend/internal/store/repo/repository_flow.go @@ -262,6 +262,24 @@ func (r *Repository) SpeedLimitExists(id int64) (bool, error) { return count > 0, nil } +func (r *Repository) CountActiveForwardsByUser(userID int64) (int64, error) { + if r == nil || r.db == nil { + return 0, errors.New("repository not initialized") + } + var count int64 + err := r.db.Model(&model.Forward{}).Where("user_id = ? AND status = 1", userID).Count(&count).Error + return count, err +} + +func (r *Repository) CountActiveForwardsByUserTunnel(userID, tunnelID int64) (int64, error) { + if r == nil || r.db == nil { + return 0, errors.New("repository not initialized") + } + var count int64 + err := r.db.Model(&model.Forward{}).Where("user_id = ? AND tunnel_id = ? AND status = 1", userID, tunnelID).Count(&count).Error + return count, err +} + func (r *Repository) GetSpeedLimitSpeed(id int64) (int, error) { if r == nil || r.db == nil { return 0, errors.New("repository not initialized") diff --git a/go-backend/tests/contract/forward_count_limit_contract_test.go b/go-backend/tests/contract/forward_count_limit_contract_test.go new file mode 100644 index 0000000..63c1688 --- /dev/null +++ b/go-backend/tests/contract/forward_count_limit_contract_test.go @@ -0,0 +1,346 @@ +package contract_test + +import ( + "bytes" + "encoding/json" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "go-backend/internal/auth" + "go-backend/internal/http/response" +) + +func TestForwardCreateBlockedWhenUserNumLimitExceeded(t *testing.T) { + secret := "contract-jwt-secret" + router, repo := setupContractRouter(t, secret) + now := time.Now().UnixMilli() + + userID := int64(100) + tunnelID := int64(1) + + if err := repo.DB().Exec(` + INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status) + VALUES(?, 'num_limit_user', 'pwd', 1, 2727251700000, 99999, 0, 0, 1, 2, ?, ?, 1) + `, userID, now, now).Error; err != nil { + t.Fatalf("insert user: %v", err) + } + + if err := repo.DB().Exec(` + INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) + VALUES(?, 'num_limit_tunnel', 1.0, 1, 'tls', 99999, ?, ?, 1, NULL, 0) + `, tunnelID, now, now).Error; err != nil { + t.Fatalf("insert tunnel: %v", err) + } + + if err := repo.DB().Exec(` + INSERT INTO user_tunnel(id, user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status) + VALUES(10, ?, ?, NULL, 99999, 99999, 0, 0, 1, 2727251700000, 1) + `, userID, tunnelID).Error; err != nil { + t.Fatalf("insert user_tunnel: %v", err) + } + + if err := repo.DB().Exec(` + INSERT INTO forward(id, user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx) + VALUES(1, ?, 'num_limit_user', 'existing_forward_1', ?, '8.8.8.8:53', 'fifo', 0, 0, ?, ?, 1, 0) + `, userID, tunnelID, now, now).Error; err != nil { + t.Fatalf("insert existing forward 1: %v", err) + } + + if err := repo.DB().Exec(` + INSERT INTO forward(id, user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx) + VALUES(2, ?, 'num_limit_user', 'existing_forward_2', ?, '8.8.4.4:53', 'fifo', 0, 0, ?, ?, 1, 0) + `, userID, tunnelID, now, now).Error; err != nil { + t.Fatalf("insert existing forward 2: %v", err) + } + + token, err := auth.GenerateToken(userID, "num_limit_user", 1, secret) + if err != nil { + t.Fatalf("generate token: %v", err) + } + + payload := `{"tunnelId":1,"name":"new_forward","remoteAddr":"1.1.1.1:53"}` + req := httptest.NewRequest(http.MethodPost, "/api/v1/forward/create", bytes.NewBufferString(payload)) + req.Header.Set("Authorization", token) + req.Header.Set("Content-Type", "application/json") + res := httptest.NewRecorder() + + router.ServeHTTP(res, req) + + var out response.R + if err := json.NewDecoder(res.Body).Decode(&out); err != nil { + t.Fatalf("decode response: %v", err) + } + if out.Code == 0 { + t.Fatalf("expected non-zero code when num limit exceeded, got code=%d msg=%q", out.Code, out.Msg) + } + if !strings.Contains(out.Msg, "转发数量已达上限") { + t.Fatalf("expected forward count limit message, got %q", out.Msg) + } +} + +func TestForwardResumeBlockedWhenUserNumLimitExceeded(t *testing.T) { + secret := "contract-jwt-secret" + router, repo := setupContractRouter(t, secret) + now := time.Now().UnixMilli() + + userID := int64(101) + tunnelID := int64(1) + pausedForwardID := int64(3) + + if err := repo.DB().Exec(` + INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status) + VALUES(?, 'num_resume_user', 'pwd', 1, 2727251700000, 99999, 0, 0, 1, 2, ?, ?, 1) + `, userID, now, now).Error; err != nil { + t.Fatalf("insert user: %v", err) + } + + if err := repo.DB().Exec(` + INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) + VALUES(?, 'num_resume_tunnel', 1.0, 1, 'tls', 99999, ?, ?, 1, NULL, 0) + `, tunnelID, now, now).Error; err != nil { + t.Fatalf("insert tunnel: %v", err) + } + + if err := repo.DB().Exec(` + INSERT INTO user_tunnel(id, user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status) + VALUES(10, ?, ?, NULL, 99999, 99999, 0, 0, 1, 2727251700000, 1) + `, userID, tunnelID).Error; err != nil { + t.Fatalf("insert user_tunnel: %v", err) + } + + if err := repo.DB().Exec(` + INSERT INTO forward(id, user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx) + VALUES(1, ?, 'num_resume_user', 'existing_forward_1', ?, '8.8.8.8:53', 'fifo', 0, 0, ?, ?, 1, 0) + `, userID, tunnelID, now, now).Error; err != nil { + t.Fatalf("insert existing forward 1: %v", err) + } + + if err := repo.DB().Exec(` + INSERT INTO forward(id, user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx) + VALUES(2, ?, 'num_resume_user', 'existing_forward_2', ?, '8.8.4.4:53', 'fifo', 0, 0, ?, ?, 1, 0) + `, userID, tunnelID, now, now).Error; err != nil { + t.Fatalf("insert existing forward 2: %v", err) + } + + if err := repo.DB().Exec(` + INSERT INTO forward(id, user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx) + VALUES(?, ?, 'num_resume_user', 'paused_forward', ?, '1.1.1.1:53', 'fifo', 0, 0, ?, ?, 0, 0) + `, pausedForwardID, userID, tunnelID, now, now).Error; err != nil { + t.Fatalf("insert paused forward: %v", err) + } + + token, err := auth.GenerateToken(userID, "num_resume_user", 1, secret) + if err != nil { + t.Fatalf("generate token: %v", err) + } + + req := httptest.NewRequest(http.MethodPost, "/api/v1/forward/resume", bytes.NewBufferString(`{"id":3}`)) + req.Header.Set("Authorization", token) + req.Header.Set("Content-Type", "application/json") + res := httptest.NewRecorder() + + router.ServeHTTP(res, req) + + var out response.R + if err := json.NewDecoder(res.Body).Decode(&out); err != nil { + t.Fatalf("decode response: %v", err) + } + if out.Code == 0 { + t.Fatalf("expected non-zero code when num limit exceeded, got code=%d msg=%q", out.Code, out.Msg) + } + if !strings.Contains(out.Msg, "转发数量已达上限") { + t.Fatalf("expected forward count limit message, got %q", out.Msg) + } + + status := mustQueryInt(t, repo, `SELECT status FROM forward WHERE id = ?`, pausedForwardID) + if status != 0 { + t.Fatalf("expected forward status to remain 0, got %d", status) + } +} + +func TestForwardCreateBlockedWhenUserTunnelNumLimitExceeded(t *testing.T) { + secret := "contract-jwt-secret" + router, repo := setupContractRouter(t, secret) + now := time.Now().UnixMilli() + + userID := int64(102) + tunnelID := int64(1) + + if err := repo.DB().Exec(` + INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status) + VALUES(?, 'ut_num_limit_user', 'pwd', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1) + `, userID, now, now).Error; err != nil { + t.Fatalf("insert user: %v", err) + } + + if err := repo.DB().Exec(` + INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) + VALUES(?, 'ut_num_limit_tunnel', 1.0, 1, 'tls', 99999, ?, ?, 1, NULL, 0) + `, tunnelID, now, now).Error; err != nil { + t.Fatalf("insert tunnel: %v", err) + } + + if err := repo.DB().Exec(` + INSERT INTO user_tunnel(id, user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status) + VALUES(10, ?, ?, NULL, 1, 99999, 0, 0, 1, 2727251700000, 1) + `, userID, tunnelID).Error; err != nil { + t.Fatalf("insert user_tunnel with num=1: %v", err) + } + + if err := repo.DB().Exec(` + INSERT INTO forward(id, user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx) + VALUES(1, ?, 'ut_num_limit_user', 'existing_tunnel_forward', ?, '8.8.8.8:53', 'fifo', 0, 0, ?, ?, 1, 0) + `, userID, tunnelID, now, now).Error; err != nil { + t.Fatalf("insert existing forward: %v", err) + } + + token, err := auth.GenerateToken(userID, "ut_num_limit_user", 1, secret) + if err != nil { + t.Fatalf("generate token: %v", err) + } + + payload := `{"tunnelId":1,"name":"new_tunnel_forward","remoteAddr":"1.1.1.1:53"}` + req := httptest.NewRequest(http.MethodPost, "/api/v1/forward/create", bytes.NewBufferString(payload)) + req.Header.Set("Authorization", token) + req.Header.Set("Content-Type", "application/json") + res := httptest.NewRecorder() + + router.ServeHTTP(res, req) + + var out response.R + if err := json.NewDecoder(res.Body).Decode(&out); err != nil { + t.Fatalf("decode response: %v", err) + } + if out.Code == 0 { + t.Fatalf("expected non-zero code when user_tunnel num limit exceeded, got code=%d msg=%q", out.Code, out.Msg) + } + if !strings.Contains(out.Msg, "隧道转发数量已达上限") { + t.Fatalf("expected tunnel forward count limit message, got %q", out.Msg) + } +} + +func TestForwardCreateAllowedWhenBelowUserNumLimit(t *testing.T) { + secret := "contract-jwt-secret" + router, repo := setupContractRouter(t, secret) + now := time.Now().UnixMilli() + + userID := int64(103) + tunnelID := int64(1) + + if err := repo.DB().Exec(` + INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status) + VALUES(?, 'num_ok_user', 'pwd', 1, 2727251700000, 99999, 0, 0, 1, 3, ?, ?, 1) + `, userID, now, now).Error; err != nil { + t.Fatalf("insert user: %v", err) + } + + if err := repo.DB().Exec(` + INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) + VALUES(?, 'num_ok_tunnel', 1.0, 1, 'tls', 99999, ?, ?, 1, NULL, 0) + `, tunnelID, now, now).Error; err != nil { + t.Fatalf("insert tunnel: %v", err) + } + + if err := repo.DB().Exec(` + INSERT INTO user_tunnel(id, user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status) + VALUES(10, ?, ?, NULL, 99999, 99999, 0, 0, 1, 2727251700000, 1) + `, userID, tunnelID).Error; err != nil { + t.Fatalf("insert user_tunnel: %v", err) + } + + if err := repo.DB().Exec(` + INSERT INTO forward(id, user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx) + VALUES(1, ?, 'num_ok_user', 'existing_forward_1', ?, '8.8.8.8:53', 'fifo', 0, 0, ?, ?, 1, 0) + `, userID, tunnelID, now, now).Error; err != nil { + t.Fatalf("insert existing forward 1: %v", err) + } + + token, err := auth.GenerateToken(userID, "num_ok_user", 1, secret) + if err != nil { + t.Fatalf("generate token: %v", err) + } + + payload := `{"tunnelId":1,"name":"new_forward_ok","remoteAddr":"1.1.1.1:53"}` + req := httptest.NewRequest(http.MethodPost, "/api/v1/forward/create", bytes.NewBufferString(payload)) + req.Header.Set("Authorization", token) + req.Header.Set("Content-Type", "application/json") + res := httptest.NewRecorder() + + router.ServeHTTP(res, req) + + var out response.R + if err := json.NewDecoder(res.Body).Decode(&out); err != nil { + t.Fatalf("decode response: %v", err) + } + if out.Code != 0 { + t.Fatalf("expected success (code=0) when below num limit, got code=%d msg=%q", out.Code, out.Msg) + } +} + +func TestForwardCreateAllowedWhenNumZero(t *testing.T) { + secret := "contract-jwt-secret" + router, repo := setupContractRouter(t, secret) + now := time.Now().UnixMilli() + + userID := int64(104) + tunnelID := int64(1) + + if err := repo.DB().Exec(` + INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status) + VALUES(?, 'num_zero_user', 'pwd', 1, 2727251700000, 99999, 0, 0, 1, 0, ?, ?, 1) + `, userID, now, now).Error; err != nil { + t.Fatalf("insert user: %v", err) + } + + if err := repo.DB().Exec(` + INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) + VALUES(?, 'num_zero_tunnel', 1.0, 1, 'tls', 99999, ?, ?, 1, NULL, 0) + `, tunnelID, now, now).Error; err != nil { + t.Fatalf("insert tunnel: %v", err) + } + + if err := repo.DB().Exec(` + INSERT INTO user_tunnel(id, user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status) + VALUES(10, ?, ?, NULL, 0, 99999, 0, 0, 1, 2727251700000, 1) + `, userID, tunnelID).Error; err != nil { + t.Fatalf("insert user_tunnel with num=0: %v", err) + } + + if err := repo.DB().Exec(` + INSERT INTO forward(id, user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx) + VALUES(1, ?, 'num_zero_user', 'existing_forward_1', ?, '8.8.8.8:53', 'fifo', 0, 0, ?, ?, 1, 0) + `, userID, tunnelID, now, now).Error; err != nil { + t.Fatalf("insert existing forward 1: %v", err) + } + + if err := repo.DB().Exec(` + INSERT INTO forward(id, user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx) + VALUES(2, ?, 'num_zero_user', 'existing_forward_2', ?, '8.8.4.4:53', 'fifo', 0, 0, ?, ?, 1, 0) + `, userID, tunnelID, now, now).Error; err != nil { + t.Fatalf("insert existing forward 2: %v", err) + } + + token, err := auth.GenerateToken(userID, "num_zero_user", 1, secret) + if err != nil { + t.Fatalf("generate token: %v", err) + } + + payload := `{"tunnelId":1,"name":"new_forward_zero","remoteAddr":"1.1.1.1:53"}` + req := httptest.NewRequest(http.MethodPost, "/api/v1/forward/create", bytes.NewBufferString(payload)) + req.Header.Set("Authorization", token) + req.Header.Set("Content-Type", "application/json") + res := httptest.NewRecorder() + + router.ServeHTTP(res, req) + + var out response.R + if err := json.NewDecoder(res.Body).Decode(&out); err != nil { + t.Fatalf("decode response: %v", err) + } + if out.Code != 0 { + t.Fatalf("expected success (code=0) when num=0 (unlimited), got code=%d msg=%q", out.Code, out.Msg) + } +} \ No newline at end of file