fix: enforce per-IP speed update permissions

This commit is contained in:
sagitchu
2026-04-27 22:26:34 +08:00
parent 9e8d27d98e
commit dec337d46b
2 changed files with 126 additions and 14 deletions
+17 -14
View File
@@ -1921,21 +1921,24 @@ func (h *Handler) forwardUpdate(w http.ResponseWriter, r *http.Request) {
}
rawIPSpeedID, hasIPSpeedID := req["ipSpeedId"]
requestedIPSpeedID := asAnyToInt64Ptr(rawIPSpeedID)
if actorRole != 0 && hasIPSpeedID && requestedIPSpeedID != nil && !sameSpeedLimitSelection(forward.IPSpeedID, requestedIPSpeedID) {
response.WriteJSON(w, response.Err(-1, "普通用户无法修改每 IP 限速规则"))
return
}
ipSpeedID := requestedIPSpeedID
ipSpeedID, err = h.normalizeSpeedLimitReference(ipSpeedID)
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
newIPSpeedID := forward.IPSpeedID
if ipSpeedID != nil {
newIPSpeedID = sql.NullInt64{Int64: *ipSpeedID, Valid: true}
} else if _, ok := req["ipSpeedId"]; ok {
newIPSpeedID = sql.NullInt64{Valid: false}
if actorRole != 0 {
if hasIPSpeedID && !sameSpeedLimitSelection(forward.IPSpeedID, requestedIPSpeedID) {
response.WriteJSON(w, response.Err(-1, "普通用户无法修改每 IP 限速规则"))
return
}
} else {
ipSpeedID := requestedIPSpeedID
ipSpeedID, err = h.normalizeSpeedLimitReference(ipSpeedID)
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
if ipSpeedID != nil {
newIPSpeedID = sql.NullInt64{Int64: *ipSpeedID, Valid: true}
} else if hasIPSpeedID {
newIPSpeedID = sql.NullInt64{Valid: false}
}
}
port := asInt(req["inPort"], 0)
@@ -1149,6 +1149,115 @@ func TestForwardIPSpeedLimitPermission(t *testing.T) {
assertCodeMsg(t, res, -1, "普通用户无法设置每 IP 限速规则")
}
func TestForwardIPSpeedLimitUpdatePermission(t *testing.T) {
secret := "contract-jwt-secret"
router, repo := setupContractRouter(t, secret)
server := httptest.NewServer(router)
defer server.Close()
now := time.Now().UnixMilli()
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(2, 'normal_user_ip_update', 'pwd', 1, ?, 99999, 0, 0, 1, 10, ?, ?, 1)
`, now+86400000, 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(13, 'ip-speed-update-permission-tunnel', 1.0, 1, 'tls', 99999, ?, ?, 1, NULL, 0)
`, now, now).Error; err != nil {
t.Fatalf("insert tunnel: %v", err)
}
if err := repo.DB().Exec(`
INSERT INTO node(id, name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx)
VALUES(21, 'ip-speed-update-permission-node', 'ip-speed-update-permission-secret', '10.22.0.2', '10.22.0.2', '', '32300-32310', '', 'v1', 1, 1, 1, ?, ?, 1, '[::]', '[::]', 0)
`, now, now).Error; err != nil {
t.Fatalf("insert node: %v", err)
}
if err := repo.DB().Exec(`
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
VALUES(13, 1, 21, 32301, 'round', 1, 'tls')
`).Error; err != nil {
t.Fatalf("insert chain_tunnel: %v", err)
}
if err := repo.DB().Exec(`
INSERT INTO speed_limit(id, name, speed, created_time, status)
VALUES(10, 'per-ip-10m-update', 10, ?, 1), (11, 'per-ip-20m-update', 20, ?, 1)
`, now, now).Error; err != nil {
t.Fatalf("insert speed limits: %v", err)
}
if err := repo.DB().Exec(`
INSERT INTO user_tunnel(user_id, tunnel_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status)
VALUES(2, 13, 10, 99999, 0, 0, 1, ?, 1)
`, now+86400000).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, ip_speed_id, in_flow, out_flow, created_time, updated_time, status, inx)
VALUES(30, 2, 'normal_user_ip_update', 'ip-speed-update-forward', 13, '1.1.1.1:443', 'fifo', 10, 0, 0, ?, ?, 1, 0)
`, now, now).Error; err != nil {
t.Fatalf("insert forward: %v", err)
}
userToken, err := auth.GenerateToken(2, "normal_user_ip_update", 1, secret)
if err != nil {
t.Fatalf("generate user token: %v", err)
}
stopNode := startMockNodeSession(t, server.URL, "ip-speed-update-permission-secret")
defer stopNode()
updateForward := func(t *testing.T, ipSpeedID interface{}) *httptest.ResponseRecorder {
t.Helper()
if err := repo.DB().Exec(`UPDATE forward SET ip_speed_id = 10 WHERE id = 30`).Error; err != nil {
t.Fatalf("reset forward ip speed limit: %v", err)
}
body, err := json.Marshal(map[string]interface{}{
"id": 30,
"name": "ip-speed-update-forward",
"tunnelId": 13,
"remoteAddr": "1.1.1.1:443",
"ipSpeedId": ipSpeedID,
})
if err != nil {
t.Fatalf("marshal update payload: %v", err)
}
req := httptest.NewRequest(http.MethodPost, "/api/v1/forward/update", bytes.NewReader(body))
req.Header.Set("Authorization", userToken)
req.Header.Set("Content-Type", "application/json")
res := httptest.NewRecorder()
router.ServeHTTP(res, req)
return res
}
assertStoredIPSpeedID := func(t *testing.T, want int64) {
t.Helper()
var got sql.NullInt64
if err := repo.DB().Raw(`SELECT ip_speed_id FROM forward WHERE id = 30`).Scan(&got).Error; err != nil {
t.Fatalf("read forward ip_speed_id: %v", err)
}
if !got.Valid || got.Int64 != want {
t.Fatalf("expected ip_speed_id %d, got valid=%v value=%d", want, got.Valid, got.Int64)
}
}
t.Run("non-admin cannot change existing ipSpeedId", func(t *testing.T) {
res := updateForward(t, 11)
assertCodeMsg(t, res, -1, "普通用户无法修改每 IP 限速规则")
assertStoredIPSpeedID(t, 10)
})
t.Run("non-admin cannot clear existing ipSpeedId", func(t *testing.T) {
res := updateForward(t, nil)
assertCodeMsg(t, res, -1, "普通用户无法修改每 IP 限速规则")
assertStoredIPSpeedID(t, 10)
})
t.Run("non-admin can keep existing ipSpeedId", func(t *testing.T) {
res := updateForward(t, 10)
assertCode(t, res, 0)
assertStoredIPSpeedID(t, 10)
})
}
func TestNonAdminCannotSetSpeedIdOrPort(t *testing.T) {
secret := "contract-jwt-secret-perm"
router, repo := setupContractRouter(t, secret)