diff --git a/go-backend/internal/http/handler/mutations.go b/go-backend/internal/http/handler/mutations.go index d48bb13..2f38113 100644 --- a/go-backend/internal/http/handler/mutations.go +++ b/go-backend/internal/http/handler/mutations.go @@ -2489,38 +2489,100 @@ func (h *Handler) upsertUserTunnel(req map[string]interface{}) error { } db := h.repo.DB() var existingID int64 - err := db.QueryRow(`SELECT id FROM user_tunnel WHERE user_id = ? AND tunnel_id = ? LIMIT 1`, userID, tunnelID).Scan(&existingID) - flow := asInt64(req["flow"], -1) - num := asInt(req["num"], -1) - expTime := asInt64(req["expTime"], -1) - flowReset := asInt64(req["flowResetTime"], -1) - status := asInt(req["status"], 1) + var currentFlow, currentNum, currentExpTime, currentFlowReset int64 + var currentSpeedID sql.NullInt64 + var currentStatus int + + err := db.QueryRow(` + SELECT id, flow, num, exp_time, flow_reset_time, speed_id, status + FROM user_tunnel + WHERE user_id = ? AND tunnel_id = ? + LIMIT 1 + `, userID, tunnelID).Scan(&existingID, ¤tFlow, ¤tNum, ¤tExpTime, ¤tFlowReset, ¤tSpeedID, ¤tStatus) + speedID := asAnyToInt64Ptr(req["speedId"]) + reqFlow := asInt64(req["flow"], -1) + reqNum := asInt(req["num"], -1) + reqExpTime := asInt64(req["expTime"], -1) + reqFlowReset := asInt64(req["flowResetTime"], -1) + reqStatus := asInt(req["status"], -1) + if err == sql.ErrNoRows { - if flow < 0 || num < 0 || expTime < 0 || flowReset < 0 { - _ = db.QueryRow(`SELECT flow, num, exp_time, flow_reset_time FROM user WHERE id = ?`, userID).Scan(&flow, &num, &expTime, &flowReset) + if reqFlow < 0 || reqNum < 0 || reqExpTime < 0 || reqFlowReset < 0 { + var uFlow, uNum, uExp, uReset int64 + if uErr := db.QueryRow(`SELECT flow, num, exp_time, flow_reset_time FROM user WHERE id = ?`, userID).Scan(&uFlow, &uNum, &uExp, &uReset); uErr == nil { + if reqFlow < 0 { + reqFlow = uFlow + } + if reqNum < 0 { + reqNum = int(uNum) + } + if reqExpTime < 0 { + reqExpTime = uExp + } + if reqFlowReset < 0 { + reqFlowReset = uReset + } + } } + if reqFlow < 0 { + reqFlow = 0 + } + if reqNum < 0 { + reqNum = 0 + } + if reqExpTime < 0 { + reqExpTime = time.Now().Add(365 * 24 * time.Hour).UnixMilli() + } + if reqFlowReset < 0 { + reqFlowReset = 1 + } + if reqStatus < 0 { + reqStatus = 1 + } + _, err = db.Exec(`INSERT INTO user_tunnel(user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status) VALUES(?, ?, ?, ?, ?, 0, 0, ?, ?, ?)`, - userID, tunnelID, nullableInt(speedID), num, flow, flowReset, expTime, status) + userID, tunnelID, nullableInt(speedID), reqNum, reqFlow, reqFlowReset, reqExpTime, reqStatus) return err } if err != nil { return err } - if flow < 0 { - flow = 0 + + newFlow := currentFlow + if reqFlow >= 0 { + newFlow = reqFlow } - if num < 0 { - num = 0 + + newNum := int(currentNum) + if reqNum >= 0 { + newNum = reqNum } - if expTime < 0 { - expTime = time.Now().Add(365 * 24 * time.Hour).UnixMilli() + + newExpTime := currentExpTime + if reqExpTime >= 0 { + newExpTime = reqExpTime } - if flowReset < 0 { - flowReset = 1 + + newFlowReset := currentFlowReset + if reqFlowReset >= 0 { + newFlowReset = reqFlowReset } + + newStatus := currentStatus + if reqStatus >= 0 { + newStatus = reqStatus + } + + newSpeedID := currentSpeedID + if speedID != nil { + newSpeedID = sql.NullInt64{Int64: *speedID, Valid: true} + } else if _, ok := req["speedId"]; ok { + newSpeedID = sql.NullInt64{Valid: false} + } + _, err = db.Exec(`UPDATE user_tunnel SET speed_id = ?, flow = ?, num = ?, exp_time = ?, flow_reset_time = ?, status = ? WHERE id = ?`, - nullableInt(speedID), flow, num, expTime, flowReset, status, existingID) + newSpeedID, newFlow, newNum, newExpTime, newFlowReset, newStatus, existingID) return err } diff --git a/go-backend/internal/store/sqlite/sql/schema.sql b/go-backend/internal/store/sqlite/sql/schema.sql index 9330f04..efeca39 100644 --- a/go-backend/internal/store/sqlite/sql/schema.sql +++ b/go-backend/internal/store/sqlite/sql/schema.sql @@ -173,6 +173,7 @@ CREATE UNIQUE INDEX IF NOT EXISTS idx_tunnel_group_tunnel_unique ON tunnel_group CREATE UNIQUE INDEX IF NOT EXISTS idx_user_group_user_unique ON user_group_user(user_group_id, user_id); CREATE UNIQUE INDEX IF NOT EXISTS idx_group_permission_unique ON group_permission(user_group_id, tunnel_group_id); CREATE UNIQUE INDEX IF NOT EXISTS idx_group_permission_grant_unique ON group_permission_grant(user_group_id, tunnel_group_id, user_tunnel_id); +CREATE UNIQUE INDEX IF NOT EXISTS idx_user_tunnel_unique ON user_tunnel(user_id, tunnel_id); CREATE TABLE IF NOT EXISTS vite_config ( id INTEGER PRIMARY KEY AUTOINCREMENT, diff --git a/go-backend/tests/contract/forward_contract_test.go b/go-backend/tests/contract/forward_contract_test.go index 88f2f7d..b438761 100644 --- a/go-backend/tests/contract/forward_contract_test.go +++ b/go-backend/tests/contract/forward_contract_test.go @@ -447,6 +447,89 @@ func TestForwardBatchChangeTunnelRollbackOnSyncFailure(t *testing.T) { } } +func TestUserTunnelReassignmentKeepsStableID(t *testing.T) { + secret := "contract-jwt-secret" + router, repo := setupContractRouter(t, secret) + now := time.Now().UnixMilli() + + adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret) + if err != nil { + t.Fatalf("generate admin token: %v", err) + } + + 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(100, 'stable_user', 'pwd', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1) + `, now, now); err != nil { + t.Fatalf("insert user: %v", err) + } + + tunnelRes, err := repo.DB().Exec(` + INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) + VALUES('stable-tunnel', 1.0, 1, 'tls', 99999, ?, ?, 1, NULL, 0) + `, now, now) + if err != nil { + t.Fatalf("insert tunnel: %v", err) + } + tunnelID, _ := tunnelRes.LastInsertId() + + // 1. Assign permission (creates new user_tunnel) + // userTunnelBatchAssign expects structure: {userId: 123, tunnels: [{tunnelId: 456, ...}]} + assignPayload := `{"userId":100,"tunnels":[{"tunnelId":` + jsonNumber(tunnelID) + `}]}` + req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/user/batch-assign", bytes.NewBufferString(assignPayload)) + req.Header.Set("Authorization", adminToken) + 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 code 0, got %d msg=%q", out.Code, out.Msg) + } + + var initialID int64 + if err := repo.DB().QueryRow(`SELECT id FROM user_tunnel WHERE user_id = 100 AND tunnel_id = ?`, tunnelID).Scan(&initialID); err != nil { + t.Fatalf("query initial user_tunnel id: %v", err) + } + + // 2. Re-assign permission (should UPDATE, not INSERT) + reassignPayload := `{"userId":100,"tunnels":[{"tunnelId":` + jsonNumber(tunnelID) + `}]}` + req2 := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/user/batch-assign", bytes.NewBufferString(reassignPayload)) + req2.Header.Set("Authorization", adminToken) + req2.Header.Set("Content-Type", "application/json") + res2 := httptest.NewRecorder() + router.ServeHTTP(res2, req2) + + var out2 response.R + if err := json.NewDecoder(res2.Body).Decode(&out2); err != nil { + t.Fatalf("decode response 2: %v", err) + } + if out2.Code != 0 { + t.Fatalf("expected code 0, got %d msg=%q", out2.Code, out2.Msg) + } + + // 3. Verify stable ID and no duplicates + var count int + if err := repo.DB().QueryRow(`SELECT COUNT(1) FROM user_tunnel WHERE user_id = 100 AND tunnel_id = ?`, tunnelID).Scan(&count); err != nil { + t.Fatalf("query count: %v", err) + } + if count != 1 { + t.Fatalf("expected exactly 1 user_tunnel record, got %d", count) + } + + var currentID int64 + if err := repo.DB().QueryRow(`SELECT id FROM user_tunnel WHERE user_id = 100 AND tunnel_id = ?`, tunnelID).Scan(¤tID); err != nil { + t.Fatalf("query current user_tunnel: %v", err) + } + + if currentID != initialID { + t.Fatalf("user_tunnel ID changed from %d to %d (unstable ID!)", initialID, currentID) + } +} + func jsonNumber(v int64) string { return strconv.FormatInt(v, 10) }