refactor(tests): centralize DB query assertions with helpers

Reduce repetitive raw SQL in test bodies by routing scalar and multi-column checks through shared helpers, keeping test intent clearer without changing behavior.
This commit is contained in:
Antigravity
2026-02-17 06:02:46 +00:00
parent d82c099c7f
commit 3d1a8c8963
15 changed files with 288 additions and 419 deletions
@@ -0,0 +1,50 @@
package handler
import (
"testing"
"go-backend/internal/store/repo"
)
func mustLastInsertID(t *testing.T, r *repo.Repository, label string) int64 {
t.Helper()
var id int64
if err := r.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&id); err != nil {
t.Fatalf("read last_insert_rowid for %s: %v", label, err)
}
if id <= 0 {
t.Fatalf("invalid last_insert_rowid for %s: %d", label, id)
}
return id
}
func mustQueryInt(t *testing.T, r *repo.Repository, query string, args ...interface{}) int {
t.Helper()
var v int
if err := r.DB().Raw(query, args...).Row().Scan(&v); err != nil {
t.Fatalf("query int failed: %v (query=%q)", err, query)
}
return v
}
func mustQueryInt64Int64String(t *testing.T, r *repo.Repository, query string, args ...interface{}) (int64, int64, string) {
t.Helper()
var a int64
var b int64
var c string
if err := r.DB().Raw(query, args...).Row().Scan(&a, &b, &c); err != nil {
t.Fatalf("query int64+int64+string failed: %v (query=%q)", err, query)
}
return a, b, c
}
func mustQueryInt64Int64Int(t *testing.T, r *repo.Repository, query string, args ...interface{}) (int64, int64, int) {
t.Helper()
var a int64
var b int64
var c int
if err := r.DB().Raw(query, args...).Row().Scan(&a, &b, &c); err != nil {
t.Fatalf("query int64+int64+int failed: %v (query=%q)", err, query)
}
return a, b, c
}
@@ -128,11 +128,7 @@ func TestPrepareTunnelCreateStateRemoteAutoPortDefersToFederation(t *testing.T)
`, name, name+"-secret", "10.0.0.1", "10.0.0.1", "", portRange, "", "v1", 1, 1, 1, now, now, status, "[::]", "[::]", 0, isRemote, "http://peer", "peer-token", `{"shareId":1}`).Error; execErr != nil {
t.Fatalf("insert node %s: %v", name, execErr)
}
var id int64
if idErr := r.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&id); idErr != nil {
t.Fatalf("node id %s: %v", name, idErr)
}
return id
return mustLastInsertID(t, r, name)
}
entryID := insertNode("entry", 1, "31000-31010", 0)
@@ -188,11 +184,7 @@ func TestPrepareTunnelCreateStateAllowsOfflineRemoteMiddleNode(t *testing.T) {
`, name, name+"-secret", "10.0.0.1", "10.0.0.1", "", portRange, "", "v1", 1, 1, 1, now, now, status, "[::]", "[::]", 0, isRemote, "http://peer", "peer-token", `{"shareId":2}`).Error; execErr != nil {
t.Fatalf("insert node %s: %v", name, execErr)
}
var id int64
if idErr := r.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&id); idErr != nil {
t.Fatalf("node id %s: %v", name, idErr)
}
return id
return mustLastInsertID(t, r, name)
}
entryID := insertNode("entry-local", 1, "32000-32010", 0)
@@ -31,10 +31,7 @@ func TestFederationShareCreateRejectsRemoteNode(t *testing.T) {
`, "remote-share-node", "remote-share-secret", "10.10.10.1", "10.10.10.1", "", "20000-20010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0, 1, "http://peer.example", "peer-token", `{"shareId":1}`).Error; err != nil {
t.Fatalf("insert remote node: %v", err)
}
var remoteNodeID int64
if err := r.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&remoteNodeID); err != nil {
t.Fatalf("get remote node id: %v", err)
}
remoteNodeID := mustLastInsertID(t, r, "remote-share-node")
body, err := json.Marshal(createPeerShareRequest{
Name: "remote-node-share",
@@ -69,10 +66,7 @@ func TestFederationShareCreateRejectsRemoteNode(t *testing.T) {
t.Fatalf("expected rejection message %q, got %q", "Only local nodes can be shared", payload.Msg)
}
var shareCount int
if err := r.DB().Raw(`SELECT COUNT(1) FROM peer_share WHERE node_id = ?`, remoteNodeID).Row().Scan(&shareCount); err != nil {
t.Fatalf("query peer_share count: %v", err)
}
shareCount := mustQueryInt(t, r, `SELECT COUNT(1) FROM peer_share WHERE node_id = ?`, remoteNodeID)
if shareCount != 0 {
t.Fatalf("expected no share rows for remote node, got %d", shareCount)
}
@@ -94,10 +88,7 @@ func TestFederationShareCreateRejectsInvalidAllowedIPs(t *testing.T) {
`, "local-share-node", "local-share-secret", "10.20.30.40", "10.20.30.40", "", "21000-21010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0, 0, "", "", "").Error; err != nil {
t.Fatalf("insert local node: %v", err)
}
var localNodeID int64
if err := r.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&localNodeID); err != nil {
t.Fatalf("get local node id: %v", err)
}
localNodeID := mustLastInsertID(t, r, "local-share-node")
body, err := json.Marshal(createPeerShareRequest{
Name: "local-node-share",
@@ -133,10 +124,7 @@ func TestFederationShareCreateRejectsInvalidAllowedIPs(t *testing.T) {
t.Fatalf("expected invalid IP message, got %q", payload.Msg)
}
var shareCount int
if err := r.DB().Raw(`SELECT COUNT(1) FROM peer_share WHERE node_id = ?`, localNodeID).Row().Scan(&shareCount); err != nil {
t.Fatalf("query peer_share count: %v", err)
}
shareCount := mustQueryInt(t, r, `SELECT COUNT(1) FROM peer_share WHERE node_id = ?`, localNodeID)
if shareCount != 0 {
t.Fatalf("expected no share rows for node, got %d", shareCount)
}
@@ -275,10 +263,7 @@ func TestFederationShareDeleteCleansUpRuntimes(t *testing.T) {
t.Fatalf("insert peer_share_runtime rows: %v", err)
}
var runtimeCount int
if err := r.DB().Raw(`SELECT COUNT(1) FROM peer_share_runtime WHERE share_id = ? AND status = 1`, share.ID).Row().Scan(&runtimeCount); err != nil {
t.Fatalf("count active runtimes before: %v", err)
}
runtimeCount := mustQueryInt(t, r, `SELECT COUNT(1) FROM peer_share_runtime WHERE share_id = ? AND status = 1`, share.ID)
if runtimeCount != 2 {
t.Fatalf("expected 2 active runtimes before delete, got %d", runtimeCount)
}
@@ -304,18 +289,12 @@ func TestFederationShareDeleteCleansUpRuntimes(t *testing.T) {
t.Fatalf("expected response code 0, got %d (%s)", payload.Code, payload.Msg)
}
var shareCount int
if err := r.DB().Raw(`SELECT COUNT(1) FROM peer_share WHERE id = ?`, share.ID).Row().Scan(&shareCount); err != nil {
t.Fatalf("count peer_share after: %v", err)
}
shareCount := mustQueryInt(t, r, `SELECT COUNT(1) FROM peer_share WHERE id = ?`, share.ID)
if shareCount != 0 {
t.Fatalf("expected peer_share deleted, got %d rows", shareCount)
}
var runtimeCountAfter int
if err := r.DB().Raw(`SELECT COUNT(1) FROM peer_share_runtime WHERE share_id = ?`, share.ID).Row().Scan(&runtimeCountAfter); err != nil {
t.Fatalf("count peer_share_runtime after: %v", err)
}
runtimeCountAfter := mustQueryInt(t, r, `SELECT COUNT(1) FROM peer_share_runtime WHERE share_id = ?`, share.ID)
if runtimeCountAfter != 0 {
t.Fatalf("expected all peer_share_runtime rows deleted, got %d", runtimeCountAfter)
}
@@ -451,26 +430,17 @@ func TestFederationRemoteUsageList(t *testing.T) {
`, "remote-consumer-node", "remote-consumer-secret", "10.30.40.50", "10.30.40.50", "", "31000-31010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0, 1, "http://peer.example", "peer-token", `{"shareId":88,"maxBandwidth":2147483648,"currentFlow":1073741824,"portRangeStart":31000,"portRangeEnd":31010}`).Error; err != nil {
t.Fatalf("insert remote node: %v", err)
}
var nodeID int64
if err := r.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&nodeID); err != nil {
t.Fatalf("remote node id: %v", err)
}
nodeID := mustLastInsertID(t, r, "remote-consumer-node")
if err := r.DB().Exec(`INSERT INTO tunnel(name, type, protocol, flow, created_time, updated_time, status, in_ip, inx) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?)`, "consumer-tunnel-a", 2, "tls", 1, now, now, 1, "", 0).Error; err != nil {
t.Fatalf("insert tunnel a: %v", err)
}
var tunnelAID int64
if err := r.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&tunnelAID); err != nil {
t.Fatal(err)
}
tunnelAID := mustLastInsertID(t, r, "consumer-tunnel-a")
if err := r.DB().Exec(`INSERT INTO tunnel(name, type, protocol, flow, created_time, updated_time, status, in_ip, inx) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?)`, "consumer-tunnel-b", 2, "tls", 1, now, now, 1, "", 0).Error; err != nil {
t.Fatalf("insert tunnel b: %v", err)
}
var tunnelBID int64
if err := r.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&tunnelBID); err != nil {
t.Fatal(err)
}
tunnelBID := mustLastInsertID(t, r, "consumer-tunnel-b")
if err := r.DB().Exec(`
INSERT INTO federation_tunnel_binding(tunnel_id, node_id, chain_type, hop_inx, remote_url, resource_key, remote_binding_id, allocated_port, status, created_time, updated_time)
+5 -24
View File
@@ -33,20 +33,12 @@ func TestRunStatisticsFlowJobTracksIncrementAndPrunes(t *testing.T) {
h.runStatisticsFlowJob(now)
var staleCount int
if err := r.DB().Raw(`SELECT COUNT(1) FROM statistics_flow WHERE created_time < ?`, nowMs-int64((48*time.Hour)/time.Millisecond)).Row().Scan(&staleCount); err != nil {
t.Fatalf("query stale statistics rows: %v", err)
}
staleCount := mustQueryInt(t, r, `SELECT COUNT(1) FROM statistics_flow WHERE created_time < ?`, nowMs-int64((48*time.Hour)/time.Millisecond))
if staleCount != 0 {
t.Fatalf("expected stale statistics rows to be pruned, got %d", staleCount)
}
var flow int64
var total int64
var hour string
if err := r.DB().Raw(`SELECT flow, total_flow, time FROM statistics_flow WHERE user_id = 1 ORDER BY id DESC LIMIT 1`).Row().Scan(&flow, &total, &hour); err != nil {
t.Fatalf("query latest statistics row: %v", err)
}
flow, total, hour := mustQueryInt64Int64String(t, r, `SELECT flow, total_flow, time FROM statistics_flow WHERE user_id = 1 ORDER BY id DESC LIMIT 1`)
if flow != 50 {
t.Fatalf("expected increment flow 50, got %d", flow)
}
@@ -100,28 +92,17 @@ func TestRunResetAndExpiryJobResetsFlowAndDisablesExpiredRecords(t *testing.T) {
h.runResetAndExpiryJob(now)
var userIn, userOut int64
var userStatus int
if err := r.DB().Raw(`SELECT in_flow, out_flow, status FROM user WHERE id = 2`).Row().Scan(&userIn, &userOut, &userStatus); err != nil {
t.Fatalf("query user after maintenance: %v", err)
}
userIn, userOut, userStatus := mustQueryInt64Int64Int(t, r, `SELECT in_flow, out_flow, status FROM user WHERE id = 2`)
if userIn != 0 || userOut != 0 || userStatus != 0 {
t.Fatalf("expected user reset+disabled, got in=%d out=%d status=%d", userIn, userOut, userStatus)
}
var utIn, utOut int64
var utStatus int
if err := r.DB().Raw(`SELECT in_flow, out_flow, status FROM user_tunnel WHERE id = 10`).Row().Scan(&utIn, &utOut, &utStatus); err != nil {
t.Fatalf("query user_tunnel after maintenance: %v", err)
}
utIn, utOut, utStatus := mustQueryInt64Int64Int(t, r, `SELECT in_flow, out_flow, status FROM user_tunnel WHERE id = 10`)
if utIn != 0 || utOut != 0 || utStatus != 0 {
t.Fatalf("expected user_tunnel reset+disabled, got in=%d out=%d status=%d", utIn, utOut, utStatus)
}
var forwardStatus int
if err := r.DB().Raw(`SELECT status FROM forward WHERE id = 20`).Row().Scan(&forwardStatus); err != nil {
t.Fatalf("query forward after maintenance: %v", err)
}
forwardStatus := mustQueryInt(t, r, `SELECT status FROM forward WHERE id = 20`)
if forwardStatus != 0 {
t.Fatalf("expected forward status=0 after expiry handling, got %d", forwardStatus)
}
@@ -0,0 +1,19 @@
package contract
import (
"testing"
"go-backend/internal/store/repo"
)
func mustLastInsertID(t *testing.T, r *repo.Repository, label string) int64 {
t.Helper()
var id int64
if err := r.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&id); err != nil {
t.Fatalf("read last_insert_rowid for %s: %v", label, err)
}
if id <= 0 {
t.Fatalf("invalid last_insert_rowid for %s: %d", label, id)
}
return id
}
@@ -0,0 +1,119 @@
package contract_test
import (
"database/sql"
"testing"
"go-backend/internal/store/repo"
)
func mustLastInsertID(t *testing.T, r *repo.Repository, label string) int64 {
t.Helper()
var id int64
if err := r.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&id); err != nil {
t.Fatalf("read last_insert_rowid for %s: %v", label, err)
}
if id <= 0 {
t.Fatalf("invalid last_insert_rowid for %s: %d", label, id)
}
return id
}
func mustQueryInt(t *testing.T, r *repo.Repository, query string, args ...interface{}) int {
t.Helper()
var v int
if err := r.DB().Raw(query, args...).Row().Scan(&v); err != nil {
t.Fatalf("query int failed: %v (query=%q)", err, query)
}
return v
}
func mustQueryInt64(t *testing.T, r *repo.Repository, query string, args ...interface{}) int64 {
t.Helper()
var v int64
if err := r.DB().Raw(query, args...).Row().Scan(&v); err != nil {
t.Fatalf("query int64 failed: %v (query=%q)", err, query)
}
return v
}
func mustQueryString(t *testing.T, r *repo.Repository, query string, args ...interface{}) string {
t.Helper()
var v string
if err := r.DB().Raw(query, args...).Row().Scan(&v); err != nil {
t.Fatalf("query string failed: %v (query=%q)", err, query)
}
return v
}
func mustQueryInt64Int(t *testing.T, r *repo.Repository, query string, args ...interface{}) (int64, int) {
t.Helper()
var a int64
var b int
if err := r.DB().Raw(query, args...).Row().Scan(&a, &b); err != nil {
t.Fatalf("query int64+int failed: %v (query=%q)", err, query)
}
return a, b
}
func tryQueryString(t *testing.T, r *repo.Repository, query string, args ...interface{}) (string, error) {
t.Helper()
var v string
err := r.DB().Raw(query, args...).Row().Scan(&v)
if err != nil {
return "", err
}
return v, nil
}
func mustQueryNullString(t *testing.T, r *repo.Repository, query string, args ...interface{}) sql.NullString {
t.Helper()
var v sql.NullString
if err := r.DB().Raw(query, args...).Row().Scan(&v); err != nil {
t.Fatalf("query null string failed: %v (query=%q)", err, query)
}
return v
}
func mustQueryTwoNullStrings(t *testing.T, r *repo.Repository, query string, args ...interface{}) (sql.NullString, sql.NullString) {
t.Helper()
var a sql.NullString
var b sql.NullString
if err := r.DB().Raw(query, args...).Row().Scan(&a, &b); err != nil {
t.Fatalf("query two null strings failed: %v (query=%q)", err, query)
}
return a, b
}
func mustQueryNodePorts(t *testing.T, r *repo.Repository, query string, args ...interface{}) map[int64]int {
t.Helper()
rows, err := r.DB().Raw(query, args...).Rows()
if err != nil {
t.Fatalf("query node ports failed: %v (query=%q)", err, query)
}
defer rows.Close()
out := make(map[int64]int)
for rows.Next() {
var nodeID int64
var port int
if err := rows.Scan(&nodeID, &port); err != nil {
t.Fatalf("scan node ports row failed: %v (query=%q)", err, query)
}
out[nodeID] = port
}
if err := rows.Err(); err != nil {
t.Fatalf("iterate node ports rows failed: %v (query=%q)", err, query)
}
return out
}
func tryQueryInt(t *testing.T, r *repo.Repository, query string, args ...interface{}) (int, error) {
t.Helper()
var v int
err := r.DB().Raw(query, args...).Row().Scan(&v)
if err != nil {
return 0, err
}
return v, nil
}
@@ -37,10 +37,7 @@ func TestDiagnosisChainCoverageContracts(t *testing.T) {
`, "diagnose-chain-tunnel", 1.0, 2, "tls", 99999, now, now, 1, nil, 0).Error; err != nil {
t.Fatalf("insert tunnel: %v", err)
}
var tunnelID int64
if err := r.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&tunnelID); err != nil {
t.Fatalf("get tunnel id: %v", err)
}
tunnelID := mustLastInsertID(t, r, "diagnose-chain-tunnel")
insertNode := func(name, ip string) int64 {
if err := r.DB().Exec(`
@@ -49,11 +46,7 @@ func TestDiagnosisChainCoverageContracts(t *testing.T) {
`, name, name+"-secret", ip, ip, "", "30000-30010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0).Error; err != nil {
t.Fatalf("insert node %s: %v", name, err)
}
var id int64
if err := r.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&id); err != nil {
t.Fatalf("get node id %s: %v", name, err)
}
return id
return mustLastInsertID(t, r, name)
}
entryNodeID := insertNode("entry-node", "10.0.1.10")
@@ -85,10 +78,7 @@ func TestDiagnosisChainCoverageContracts(t *testing.T) {
`, 2, "normal_user", "chain-forward", tunnelID, "8.8.8.8:53", "fifo", now, now, 0).Error; err != nil {
t.Fatalf("insert forward: %v", err)
}
var forwardID int64
if err := r.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&forwardID); err != nil {
t.Fatalf("get forward id: %v", err)
}
forwardID := mustLastInsertID(t, r, "chain-forward")
userToken, err := auth.GenerateToken(2, "normal_user", 1, secret)
if err != nil {
@@ -259,11 +249,7 @@ func TestDiagnosisUsesFederationRuntimeForRemoteNodes(t *testing.T) {
`, name, name+"-secret", ip, ip, "", "30000-30010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0).Error; err != nil {
t.Fatalf("insert local node %s: %v", name, err)
}
var id int64
if err := r.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&id); err != nil {
t.Fatalf("get local node id %s: %v", name, err)
}
return id
return mustLastInsertID(t, r, name)
}
insertRemoteNode := func(name, ip string) int64 {
@@ -273,11 +259,7 @@ func TestDiagnosisUsesFederationRuntimeForRemoteNodes(t *testing.T) {
`, name, name+"-secret", ip, "", "", "31000-31010", "", "", now, now, "[::]", "[::]", 1, remoteServer.URL, remoteToken, `{"shareId": 123}`).Error; err != nil {
t.Fatalf("insert remote node %s: %v", name, err)
}
var id int64
if err := r.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&id); err != nil {
t.Fatalf("get remote node id %s: %v", name, err)
}
return id
return mustLastInsertID(t, r, name)
}
entryNodeID := insertLocalNode("entry-local", "10.50.0.10")
@@ -290,10 +272,7 @@ func TestDiagnosisUsesFederationRuntimeForRemoteNodes(t *testing.T) {
`, "diagnose-remote-tunnel", 1.0, 2, "tls", 99999, now, now, 1, nil, 0).Error; err != nil {
t.Fatalf("insert tunnel: %v", err)
}
var tunnelID int64
if err := r.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&tunnelID); err != nil {
t.Fatalf("get tunnel id: %v", err)
}
tunnelID := mustLastInsertID(t, r, "diagnose-remote-tunnel")
if err := r.DB().Exec(`
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
@@ -112,10 +112,7 @@ func TestFederationDualPanelMiddleExitAutoPortContract(t *testing.T) {
consumerRouter.ServeHTTP(res, req)
assertCode(t, res, 0)
var tunnelID int64
if err := consumerRepo.DB().Raw(`SELECT id FROM tunnel WHERE name = ? ORDER BY id DESC LIMIT 1`, name).Row().Scan(&tunnelID); err != nil {
t.Fatalf("query tunnel id (%s): %v", name, err)
}
tunnelID := mustQueryInt64(t, consumerRepo, `SELECT id FROM tunnel WHERE name = ? ORDER BY id DESC LIMIT 1`, name)
if tunnelID <= 0 {
t.Fatalf("invalid tunnel id for %s", name)
}
@@ -261,10 +258,7 @@ func TestFederationDualPanelRemoteDiagnosisContract(t *testing.T) {
consumerRouter.ServeHTTP(createRes, createReq)
assertCode(t, createRes, 0)
var tunnelID int64
if err := consumerRepo.DB().Raw(`SELECT id FROM tunnel WHERE name = ? ORDER BY id DESC LIMIT 1`, "dual-panel-diagnose-remote").Row().Scan(&tunnelID); err != nil {
t.Fatalf("query tunnel id: %v", err)
}
tunnelID := mustQueryInt64(t, consumerRepo, `SELECT id FROM tunnel WHERE name = ? ORDER BY id DESC LIMIT 1`, "dual-panel-diagnose-remote")
if tunnelID <= 0 {
t.Fatalf("invalid tunnel id")
}
@@ -414,10 +408,7 @@ func TestFederationDualPanelRemoteEntryRuntimeContract(t *testing.T) {
consumerRouter.ServeHTTP(res, req)
assertCode(t, res, 0)
var tunnelID int64
if err := consumerRepo.DB().Raw(`SELECT id FROM tunnel WHERE name = ? ORDER BY id DESC LIMIT 1`, name).Row().Scan(&tunnelID); err != nil {
t.Fatalf("query tunnel id (%s): %v", name, err)
}
tunnelID := mustQueryInt64(t, consumerRepo, `SELECT id FROM tunnel WHERE name = ? ORDER BY id DESC LIMIT 1`, name)
if tunnelID <= 0 {
t.Fatalf("invalid tunnel id for %s", name)
}
@@ -455,11 +446,7 @@ func insertContractNode(t *testing.T, r *repo.Repository, name, ip, portRange, s
`, name, secret, ip, ip, "", portRange, "", "v1", 1, 1, 1, now, now, status, "[::]", "[::]", 0).Error; err != nil {
t.Fatalf("insert node %s: %v", name, err)
}
var id int64
if err := r.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&id); err != nil {
t.Fatalf("node id %s: %v", name, err)
}
return id
return mustLastInsertID(t, r, name)
}
func insertPeerShare(t *testing.T, r *repo.Repository, share *repo.PeerShare) int64 {
@@ -499,10 +486,7 @@ func importRemoteNodeForContract(t *testing.T, router http.Handler, adminToken,
func queryRemoteNodeIDByToken(t *testing.T, r *repo.Repository, token string) int64 {
t.Helper()
var id int64
if err := r.DB().Raw(`SELECT id FROM node WHERE is_remote = 1 AND remote_token = ? ORDER BY id DESC LIMIT 1`, token).Row().Scan(&id); err != nil {
t.Fatalf("query remote node by token %s: %v", token, err)
}
id := mustQueryInt64(t, r, `SELECT id FROM node WHERE is_remote = 1 AND remote_token = ? ORDER BY id DESC LIMIT 1`, token)
if id <= 0 {
t.Fatalf("invalid remote node id for token %s", token)
}
@@ -511,16 +495,7 @@ func queryRemoteNodeIDByToken(t *testing.T, r *repo.Repository, token string) in
func assertTunnelPortInRange(t *testing.T, r *repo.Repository, tunnelID int64, chainType int, nodeID int64, minPort int, maxPort int) {
t.Helper()
var port int
err := r.DB().Raw(`
SELECT port
FROM chain_tunnel
WHERE tunnel_id = ? AND chain_type = ? AND node_id = ?
LIMIT 1
`, tunnelID, chainType, nodeID).Row().Scan(&port)
if err != nil {
t.Fatalf("query tunnel=%d chainType=%d node=%d port: %v", tunnelID, chainType, nodeID, err)
}
port := mustQueryInt(t, r, `SELECT port FROM chain_tunnel WHERE tunnel_id = ? AND chain_type = ? AND node_id = ? LIMIT 1`, tunnelID, chainType, nodeID)
if port < minPort || port > maxPort {
t.Fatalf("expected port in range [%d,%d], got %d", minPort, maxPort, port)
}
@@ -528,10 +503,7 @@ func assertTunnelPortInRange(t *testing.T, r *repo.Repository, tunnelID int64, c
func assertCount(t *testing.T, r *repo.Repository, query string, arg interface{}, expected int) {
t.Helper()
var got int
if err := r.DB().Raw(query, arg).Row().Scan(&got); err != nil {
t.Fatalf("count query failed: %v", err)
}
got := mustQueryInt(t, r, query, arg)
if got != expected {
t.Fatalf("expected count %d, got %d (query: %s, arg: %v)", expected, got, query, arg)
}
@@ -641,8 +613,8 @@ func waitNodeStatus(t *testing.T, r *repo.Repository, nodeID int64, expectedStat
t.Helper()
deadline := time.Now().Add(2 * time.Second)
for {
var status int
if err := r.DB().Raw(`SELECT status FROM node WHERE id = ?`, nodeID).Row().Scan(&status); err == nil && status == expectedStatus {
status, err := tryQueryInt(t, r, `SELECT status FROM node WHERE id = ?`, nodeID)
if err == nil && status == expectedStatus {
return
}
if time.Now().After(deadline) {
@@ -31,10 +31,7 @@ func TestForwardOwnershipAndScopeContracts(t *testing.T) {
`, "contract-tunnel", 1.0, 1, "tls", 99999, now, now, 1, nil, 0).Error; err != nil {
t.Fatalf("insert tunnel: %v", err)
}
var tunnelID int64
if err := repo.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&tunnelID); err != nil {
t.Fatalf("get tunnel id: %v", err)
}
tunnelID := mustLastInsertID(t, repo, "contract-tunnel")
if err := repo.DB().Exec(`
INSERT INTO node(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)
@@ -42,10 +39,7 @@ func TestForwardOwnershipAndScopeContracts(t *testing.T) {
`, "entry-node", "entry-secret", "10.0.0.10", "10.0.0.10", "", "20000-20010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0).Error; err != nil {
t.Fatalf("insert node: %v", err)
}
var entryNodeID int64
if err := repo.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&entryNodeID); err != nil {
t.Fatalf("get node id: %v", err)
}
entryNodeID := mustLastInsertID(t, repo, "entry-node")
if err := repo.DB().Exec(`
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
@@ -60,10 +54,7 @@ func TestForwardOwnershipAndScopeContracts(t *testing.T) {
`, 1, "admin_user", "admin-forward", tunnelID, "1.1.1.1:443", "fifo", now, now, 0).Error; err != nil {
t.Fatalf("insert admin forward: %v", err)
}
var adminForwardID int64
if err := repo.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&adminForwardID); err != nil {
t.Fatalf("get admin forward id: %v", err)
}
adminForwardID := mustLastInsertID(t, repo, "admin-forward")
if err := repo.DB().Exec(`
INSERT INTO forward(user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx)
@@ -71,10 +62,7 @@ func TestForwardOwnershipAndScopeContracts(t *testing.T) {
`, 2, "normal_user", "user-forward", tunnelID, "8.8.8.8:53", "fifo", now, now, 1).Error; err != nil {
t.Fatalf("insert user forward: %v", err)
}
var userForwardID int64
if err := repo.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&userForwardID); err != nil {
t.Fatalf("get user forward id: %v", err)
}
userForwardID := mustLastInsertID(t, repo, "user-forward")
userToken, err := auth.GenerateToken(2, "normal_user", 1, secret)
if err != nil {
@@ -217,11 +205,7 @@ func TestForwardSwitchTunnelRollbackOnSyncFailure(t *testing.T) {
`, name, 1.0, 1, "tls", 99999, now, now, 1, nil, inx).Error; err != nil {
t.Fatalf("insert tunnel %s: %v", name, err)
}
var id int64
if err := repo.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&id); err != nil {
t.Fatalf("get tunnel id %s: %v", name, err)
}
return id
return mustLastInsertID(t, repo, name)
}
insertNode := func(name, ip, portRange string, inx int) int64 {
@@ -231,11 +215,7 @@ func TestForwardSwitchTunnelRollbackOnSyncFailure(t *testing.T) {
`, name, name+"-secret", ip, ip, "", portRange, "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", inx).Error; err != nil {
t.Fatalf("insert node %s: %v", name, err)
}
var id int64
if err := repo.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&id); err != nil {
t.Fatalf("get node id %s: %v", name, err)
}
return id
return mustLastInsertID(t, repo, name)
}
tunnelA := insertTunnel("switch-tunnel-a", 0)
@@ -275,10 +255,7 @@ func TestForwardSwitchTunnelRollbackOnSyncFailure(t *testing.T) {
`, tunnelA, now, now).Error; err != nil {
t.Fatalf("insert forward: %v", err)
}
var forwardID int64
if err := repo.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&forwardID); err != nil {
t.Fatalf("get forward id: %v", err)
}
forwardID := mustLastInsertID(t, repo, "switch-forward")
if err := repo.DB().Exec(`INSERT INTO forward_port(forward_id, node_id, port) VALUES(?, ?, ?)`, forwardID, nodeA, 21001).Error; err != nil {
t.Fatalf("insert forward_port: %v", err)
@@ -300,19 +277,12 @@ func TestForwardSwitchTunnelRollbackOnSyncFailure(t *testing.T) {
t.Fatalf("expected update failure when node is offline")
}
var tunnelAfter int64
if err := repo.DB().Raw(`SELECT tunnel_id FROM forward WHERE id = ?`, forwardID).Row().Scan(&tunnelAfter); err != nil {
t.Fatalf("query forward tunnel_id: %v", err)
}
tunnelAfter := mustQueryInt64(t, repo, `SELECT tunnel_id FROM forward WHERE id = ?`, forwardID)
if tunnelAfter != tunnelA {
t.Fatalf("expected tunnel rollback to %d, got %d", tunnelA, tunnelAfter)
}
var nodeAfter int64
var portAfter int
if err := repo.DB().Raw(`SELECT node_id, port FROM forward_port WHERE forward_id = ? LIMIT 1`, forwardID).Row().Scan(&nodeAfter, &portAfter); err != nil {
t.Fatalf("query forward_port: %v", err)
}
nodeAfter, portAfter := mustQueryInt64Int(t, repo, `SELECT node_id, port FROM forward_port WHERE forward_id = ? LIMIT 1`, forwardID)
if nodeAfter != nodeA || portAfter != 21001 {
t.Fatalf("expected forward_port rollback to node=%d port=21001, got node=%d port=%d", nodeA, nodeAfter, portAfter)
}
@@ -341,10 +311,7 @@ func TestForwardBatchChangeTunnelRollbackOnSyncFailure(t *testing.T) {
`, now, now).Error; err != nil {
t.Fatalf("insert tunnel A: %v", err)
}
var tunnelA int64
if err := repo.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&tunnelA); err != nil {
t.Fatal(err)
}
tunnelA := mustLastInsertID(t, repo, "batch-switch-tunnel-a")
if err := repo.DB().Exec(`
INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
@@ -352,10 +319,7 @@ func TestForwardBatchChangeTunnelRollbackOnSyncFailure(t *testing.T) {
`, now, now).Error; err != nil {
t.Fatalf("insert tunnel B: %v", err)
}
var tunnelB int64
if err := repo.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&tunnelB); err != nil {
t.Fatal(err)
}
tunnelB := mustLastInsertID(t, repo, "batch-switch-tunnel-b")
if err := repo.DB().Exec(`
INSERT INTO node(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)
@@ -363,10 +327,7 @@ func TestForwardBatchChangeTunnelRollbackOnSyncFailure(t *testing.T) {
`, now, now).Error; err != nil {
t.Fatalf("insert node A: %v", err)
}
var nodeA int64
if err := repo.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&nodeA); err != nil {
t.Fatal(err)
}
nodeA := mustLastInsertID(t, repo, "batch-switch-node-a")
if err := repo.DB().Exec(`
INSERT INTO node(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)
@@ -374,10 +335,7 @@ func TestForwardBatchChangeTunnelRollbackOnSyncFailure(t *testing.T) {
`, now, now).Error; err != nil {
t.Fatalf("insert node B: %v", err)
}
var nodeB int64
if err := repo.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&nodeB); err != nil {
t.Fatal(err)
}
nodeB := mustLastInsertID(t, repo, "batch-switch-node-b")
if err := repo.DB().Exec(`INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, 1, ?, 23001, 'round', 1, 'tls')`, tunnelA, nodeA).Error; err != nil {
t.Fatalf("insert chain_tunnel A: %v", err)
@@ -399,10 +357,7 @@ func TestForwardBatchChangeTunnelRollbackOnSyncFailure(t *testing.T) {
`, tunnelA, now, now).Error; err != nil {
t.Fatalf("insert forward: %v", err)
}
var forwardID int64
if err := repo.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&forwardID); err != nil {
t.Fatal(err)
}
forwardID := mustLastInsertID(t, repo, "batch-switch-forward")
if err := repo.DB().Exec(`INSERT INTO forward_port(forward_id, node_id, port) VALUES(?, ?, ?)`, forwardID, nodeA, 23001).Error; err != nil {
t.Fatalf("insert forward_port: %v", err)
@@ -432,19 +387,12 @@ func TestForwardBatchChangeTunnelRollbackOnSyncFailure(t *testing.T) {
t.Fatalf("expected failCount=1, got %v", result["failCount"])
}
var tunnelAfter int64
if err := repo.DB().Raw(`SELECT tunnel_id FROM forward WHERE id = ?`, forwardID).Row().Scan(&tunnelAfter); err != nil {
t.Fatalf("query forward tunnel_id: %v", err)
}
tunnelAfter := mustQueryInt64(t, repo, `SELECT tunnel_id FROM forward WHERE id = ?`, forwardID)
if tunnelAfter != tunnelA {
t.Fatalf("expected tunnel rollback to %d, got %d", tunnelA, tunnelAfter)
}
var nodeAfter int64
var portAfter int
if err := repo.DB().Raw(`SELECT node_id, port FROM forward_port WHERE forward_id = ? LIMIT 1`, forwardID).Row().Scan(&nodeAfter, &portAfter); err != nil {
t.Fatalf("query forward_port: %v", err)
}
nodeAfter, portAfter := mustQueryInt64Int(t, repo, `SELECT node_id, port FROM forward_port WHERE forward_id = ? LIMIT 1`, forwardID)
if nodeAfter != nodeA || portAfter != 23001 {
t.Fatalf("expected forward_port rollback to node=%d port=23001, got node=%d port=%d", nodeA, nodeAfter, portAfter)
}
@@ -473,10 +421,7 @@ func TestUserTunnelReassignmentKeepsStableID(t *testing.T) {
`, now, now).Error; err != nil {
t.Fatalf("insert tunnel: %v", err)
}
var tunnelID int64
if err := repo.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&tunnelID); err != nil {
t.Fatal(err)
}
tunnelID := mustLastInsertID(t, repo, "stable-tunnel")
// 1. Assign permission (creates new user_tunnel)
// userTunnelBatchAssign expects structure: {userId: 123, tunnels: [{tunnelId: 456, ...}]}
@@ -495,10 +440,7 @@ func TestUserTunnelReassignmentKeepsStableID(t *testing.T) {
t.Fatalf("expected code 0, got %d msg=%q", out.Code, out.Msg)
}
var initialID int64
if err := repo.DB().Raw(`SELECT id FROM user_tunnel WHERE user_id = 100 AND tunnel_id = ?`, tunnelID).Row().Scan(&initialID); err != nil {
t.Fatalf("query initial user_tunnel id: %v", err)
}
initialID := mustQueryInt64(t, repo, `SELECT id FROM user_tunnel WHERE user_id = 100 AND tunnel_id = ?`, tunnelID)
// 2. Re-assign permission (should UPDATE, not INSERT)
reassignPayload := `{"userId":100,"tunnels":[{"tunnelId":` + jsonNumber(tunnelID) + `}]}`
@@ -517,18 +459,12 @@ func TestUserTunnelReassignmentKeepsStableID(t *testing.T) {
}
// 3. Verify stable ID and no duplicates
var count int
if err := repo.DB().Raw(`SELECT COUNT(1) FROM user_tunnel WHERE user_id = 100 AND tunnel_id = ?`, tunnelID).Row().Scan(&count); err != nil {
t.Fatalf("query count: %v", err)
}
count := mustQueryInt(t, repo, `SELECT COUNT(1) FROM user_tunnel WHERE user_id = 100 AND tunnel_id = ?`, tunnelID)
if count != 1 {
t.Fatalf("expected exactly 1 user_tunnel record, got %d", count)
}
var currentID int64
if err := repo.DB().Raw(`SELECT id FROM user_tunnel WHERE user_id = 100 AND tunnel_id = ?`, tunnelID).Row().Scan(&currentID); err != nil {
t.Fatalf("query current user_tunnel: %v", err)
}
currentID := mustQueryInt64(t, repo, `SELECT id FROM user_tunnel WHERE user_id = 100 AND tunnel_id = ?`, tunnelID)
if currentID != initialID {
t.Fatalf("user_tunnel ID changed from %d to %d (unstable ID!)", initialID, currentID)
@@ -28,26 +28,17 @@ func TestGroupUserUnbindRevokesInheritedTunnelPermission(t *testing.T) {
`, now, now).Error; err != nil {
t.Fatalf("insert tunnel: %v", err)
}
var tunnelID int64
if err := repo.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&tunnelID); err != nil {
t.Fatalf("read tunnel id: %v", err)
}
tunnelID := mustLastInsertID(t, repo, "group-contract-tunnel")
if err := repo.DB().Exec(`INSERT INTO user_group(name, created_time, updated_time, status) VALUES('ug-contract', ?, ?, 1)`, now, now).Error; err != nil {
t.Fatalf("insert user_group: %v", err)
}
var userGroupID int64
if err := repo.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&userGroupID); err != nil {
t.Fatalf("read user_group id: %v", err)
}
userGroupID := mustLastInsertID(t, repo, "ug-contract")
if err := repo.DB().Exec(`INSERT INTO tunnel_group(name, created_time, updated_time, status) VALUES('tg-contract', ?, ?, 1)`, now, now).Error; err != nil {
t.Fatalf("insert tunnel_group: %v", err)
}
var tunnelGroupID int64
if err := repo.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&tunnelGroupID); err != nil {
t.Fatalf("read tunnel_group id: %v", err)
}
tunnelGroupID := mustLastInsertID(t, repo, "tg-contract")
if err := repo.DB().Exec(`INSERT INTO tunnel_group_tunnel(tunnel_group_id, tunnel_id, created_time) VALUES(?, ?, ?)`, tunnelGroupID, tunnelID, now).Error; err != nil {
t.Fatalf("insert tunnel_group_tunnel: %v", err)
@@ -67,15 +58,9 @@ func TestGroupUserUnbindRevokesInheritedTunnelPermission(t *testing.T) {
router.ServeHTTP(bindRes, bindReq)
assertCode(t, bindRes, 0)
var userTunnelID int64
if err := repo.DB().Raw(`SELECT id FROM user_tunnel WHERE user_id = 200 AND tunnel_id = ?`, tunnelID).Row().Scan(&userTunnelID); err != nil {
t.Fatalf("query user_tunnel after bind: %v", err)
}
userTunnelID := mustQueryInt64(t, repo, `SELECT id FROM user_tunnel WHERE user_id = 200 AND tunnel_id = ?`, tunnelID)
var grantCount int
if err := repo.DB().Raw(`SELECT COUNT(1) FROM group_permission_grant WHERE user_tunnel_id = ?`, userTunnelID).Row().Scan(&grantCount); err != nil {
t.Fatalf("query group_permission_grant after bind: %v", err)
}
grantCount := mustQueryInt(t, repo, `SELECT COUNT(1) FROM group_permission_grant WHERE user_tunnel_id = ?`, userTunnelID)
if grantCount == 0 {
t.Fatalf("expected non-zero grants after bind")
}
@@ -86,17 +71,12 @@ func TestGroupUserUnbindRevokesInheritedTunnelPermission(t *testing.T) {
router.ServeHTTP(unbindRes, unbindReq)
assertCode(t, unbindRes, 0)
if err := repo.DB().Raw(`SELECT COUNT(1) FROM group_permission_grant WHERE user_tunnel_id = ?`, userTunnelID).Row().Scan(&grantCount); err != nil {
t.Fatalf("query group_permission_grant after unbind: %v", err)
}
grantCount = mustQueryInt(t, repo, `SELECT COUNT(1) FROM group_permission_grant WHERE user_tunnel_id = ?`, userTunnelID)
if grantCount != 0 {
t.Fatalf("expected grants revoked after unbind, got %d", grantCount)
}
var userTunnelCount int
if err := repo.DB().Raw(`SELECT COUNT(1) FROM user_tunnel WHERE id = ?`, userTunnelID).Row().Scan(&userTunnelCount); err != nil {
t.Fatalf("query user_tunnel after unbind: %v", err)
}
userTunnelCount := mustQueryInt(t, repo, `SELECT COUNT(1) FROM user_tunnel WHERE id = ?`, userTunnelID)
if userTunnelCount != 0 {
t.Fatalf("expected user_tunnel revoked after unbind, got %d", userTunnelCount)
}
@@ -120,26 +100,17 @@ func TestGroupPermissionRemoveRevokesInheritedTunnelPermission(t *testing.T) {
`, now, now).Error; err != nil {
t.Fatalf("insert tunnel: %v", err)
}
var tunnelID int64
if err := repo.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&tunnelID); err != nil {
t.Fatalf("read tunnel id: %v", err)
}
tunnelID := mustLastInsertID(t, repo, "group-remove-tunnel")
if err := repo.DB().Exec(`INSERT INTO user_group(name, created_time, updated_time, status) VALUES('ug-remove-contract', ?, ?, 1)`, now, now).Error; err != nil {
t.Fatalf("insert user_group: %v", err)
}
var userGroupID int64
if err := repo.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&userGroupID); err != nil {
t.Fatalf("read user_group id: %v", err)
}
userGroupID := mustLastInsertID(t, repo, "ug-remove-contract")
if err := repo.DB().Exec(`INSERT INTO tunnel_group(name, created_time, updated_time, status) VALUES('tg-remove-contract', ?, ?, 1)`, now, now).Error; err != nil {
t.Fatalf("insert tunnel_group: %v", err)
}
var tunnelGroupID int64
if err := repo.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&tunnelGroupID); err != nil {
t.Fatalf("read tunnel_group id: %v", err)
}
tunnelGroupID := mustLastInsertID(t, repo, "tg-remove-contract")
adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret)
if err != nil {
@@ -164,20 +135,11 @@ func TestGroupPermissionRemoveRevokesInheritedTunnelPermission(t *testing.T) {
router.ServeHTTP(assignPermissionRes, assignPermissionReq)
assertCode(t, assignPermissionRes, 0)
var permissionID int64
if err := repo.DB().Raw(`SELECT id FROM group_permission WHERE user_group_id = ? AND tunnel_group_id = ?`, userGroupID, tunnelGroupID).Row().Scan(&permissionID); err != nil {
t.Fatalf("query group_permission id: %v", err)
}
permissionID := mustQueryInt64(t, repo, `SELECT id FROM group_permission WHERE user_group_id = ? AND tunnel_group_id = ?`, userGroupID, tunnelGroupID)
var userTunnelID int64
if err := repo.DB().Raw(`SELECT id FROM user_tunnel WHERE user_id = 201 AND tunnel_id = ?`, tunnelID).Row().Scan(&userTunnelID); err != nil {
t.Fatalf("query user_tunnel after assign: %v", err)
}
userTunnelID := mustQueryInt64(t, repo, `SELECT id FROM user_tunnel WHERE user_id = 201 AND tunnel_id = ?`, tunnelID)
var grantCount int
if err := repo.DB().Raw(`SELECT COUNT(1) FROM group_permission_grant WHERE user_tunnel_id = ?`, userTunnelID).Row().Scan(&grantCount); err != nil {
t.Fatalf("query group_permission_grant after assign: %v", err)
}
grantCount := mustQueryInt(t, repo, `SELECT COUNT(1) FROM group_permission_grant WHERE user_tunnel_id = ?`, userTunnelID)
if grantCount == 0 {
t.Fatalf("expected non-zero grants after permission assign")
}
@@ -188,25 +150,17 @@ func TestGroupPermissionRemoveRevokesInheritedTunnelPermission(t *testing.T) {
router.ServeHTTP(removeRes, removeReq)
assertCode(t, removeRes, 0)
var permissionCount int
if err := repo.DB().Raw(`SELECT COUNT(1) FROM group_permission WHERE id = ?`, permissionID).Row().Scan(&permissionCount); err != nil {
t.Fatalf("query group_permission after remove: %v", err)
}
permissionCount := mustQueryInt(t, repo, `SELECT COUNT(1) FROM group_permission WHERE id = ?`, permissionID)
if permissionCount != 0 {
t.Fatalf("expected group_permission removed, got %d", permissionCount)
}
if err := repo.DB().Raw(`SELECT COUNT(1) FROM group_permission_grant WHERE user_tunnel_id = ?`, userTunnelID).Row().Scan(&grantCount); err != nil {
t.Fatalf("query group_permission_grant after remove: %v", err)
}
grantCount = mustQueryInt(t, repo, `SELECT COUNT(1) FROM group_permission_grant WHERE user_tunnel_id = ?`, userTunnelID)
if grantCount != 0 {
t.Fatalf("expected grants removed after permission remove, got %d", grantCount)
}
var userTunnelCount int
if err := repo.DB().Raw(`SELECT COUNT(1) FROM user_tunnel WHERE id = ?`, userTunnelID).Row().Scan(&userTunnelCount); err != nil {
t.Fatalf("query user_tunnel after permission remove: %v", err)
}
userTunnelCount := mustQueryInt(t, repo, `SELECT COUNT(1) FROM user_tunnel WHERE id = ?`, userTunnelID)
if userTunnelCount != 0 {
t.Fatalf("expected user_tunnel revoked after permission remove, got %d", userTunnelCount)
}
@@ -94,10 +94,7 @@ func TestOpenAPISubStoreContracts(t *testing.T) {
"contract-tunnel", 1.0, 1, "tls", 1, now, now, 1, nil, 0).Error; err != nil {
t.Fatalf("insert tunnel: %v", err)
}
var tunnelID int64
if err := r.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&tunnelID); err != nil {
t.Fatalf("last insert id: %v", err)
}
tunnelID := mustLastInsertID(t, r, "contract-tunnel")
if err := r.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(?, ?, NULL, ?, ?, ?, ?, ?, ?, ?)`,
1, tunnelID, 99999, tunnelFlowGB, tunnelInFlow, tunnelOutFlow, 1, tunnelExpTimeMs, 1).Error; err != nil {
t.Fatalf("insert user_tunnel: %v", err)
@@ -317,10 +314,7 @@ func TestBackupExportImportRestoreContracts(t *testing.T) {
`, "backup-forward-tunnel", 1.0, 1, "tls", 0, now, now, 1, "", 88).Error; err != nil {
t.Fatalf("seed tunnel for forward backup: %v", err)
}
var tunnelID int64
if err := r.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&tunnelID); err != nil {
t.Fatalf("read tunnel id for forward backup: %v", err)
}
tunnelID := mustLastInsertID(t, r, "backup-forward-tunnel")
if err := r.DB().Exec(`
INSERT INTO forward(user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx)
@@ -328,10 +322,7 @@ func TestBackupExportImportRestoreContracts(t *testing.T) {
`, 1, "admin_user", "backup-forward", tunnelID, "127.0.0.1:9000", "fifo", 0, 0, now, now, 1, 88).Error; err != nil {
t.Fatalf("seed forward for backup: %v", err)
}
var forwardID int64
if err := r.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&forwardID); err != nil {
t.Fatalf("read forward id for backup: %v", err)
}
forwardID := mustLastInsertID(t, r, "backup-forward")
expected := map[int64]int{
2001: 21001,
@@ -442,24 +433,7 @@ func TestBackupExportImportRestoreContracts(t *testing.T) {
t.Fatalf("expected forwards import code 0, got %d (%s)", out.Code, out.Msg)
}
rows, err := r.DB().Raw(`SELECT node_id, port FROM forward_port WHERE forward_id = ? ORDER BY id ASC`, forwardID).Rows()
if err != nil {
t.Fatalf("query forward ports after import: %v", err)
}
defer rows.Close()
after := make(map[int64]int)
for rows.Next() {
var nodeID int64
var port int
if err := rows.Scan(&nodeID, &port); err != nil {
t.Fatalf("scan forward_port row: %v", err)
}
after[nodeID] = port
}
if err := rows.Err(); err != nil {
t.Fatalf("iterate forward_port rows: %v", err)
}
after := mustQueryNodePorts(t, r, `SELECT node_id, port FROM forward_port WHERE forward_id = ? ORDER BY id ASC`, forwardID)
if len(after) != len(expected) {
t.Fatalf("expected %d forward ports after import, got %d (%v)", len(expected), len(after), after)
@@ -479,10 +453,7 @@ func TestBackupExportImportRestoreContracts(t *testing.T) {
`, "legacy-null-chain", 1.0, 1, "tls", 1000, now, now, 1, nil, 1).Error; err != nil {
t.Fatalf("seed tunnel for nullable chain export: %v", err)
}
var tunnelID int64
if err := r.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&tunnelID); err != nil {
t.Fatalf("read tunnel id for nullable chain export: %v", err)
}
tunnelID := mustLastInsertID(t, r, "legacy-null-chain")
if err := r.DB().Exec(`
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
@@ -696,25 +667,19 @@ func TestOpenMigratesLegacyNodeDualStackColumns(t *testing.T) {
func readTableColumns(t *testing.T, db *gorm.DB, table string) map[string]bool {
t.Helper()
rows, err := db.Raw("PRAGMA table_info(" + table + ")").Rows()
columnTypes, err := db.Migrator().ColumnTypes(table)
if err != nil {
t.Fatalf("inspect %s columns: %v", table, err)
}
defer rows.Close()
columns := map[string]bool{}
for rows.Next() {
var cid, notNull, pk int
var name, typ string
var defaultValue sql.NullString
if err := rows.Scan(&cid, &name, &typ, &notNull, &defaultValue, &pk); err != nil {
t.Fatalf("scan %s pragma row: %v", table, err)
for _, col := range columnTypes {
name := strings.TrimSpace(col.Name())
if name == "" {
continue
}
columns[name] = true
}
if err := rows.Err(); err != nil {
t.Fatalf("iterate %s pragma rows: %v", table, err)
}
return columns
}
@@ -64,17 +64,14 @@ func TestPostgresNodeCreateRepairsMissingIDDefaultContract(t *testing.T) {
_ = r.Close()
})
var columnDefault sql.NullString
if err := r.DB().Raw(`
columnDefault := mustQueryNullString(t, r, `
SELECT column_default
FROM information_schema.columns
WHERE table_schema = current_schema()
AND table_name = 'node'
AND column_name = 'id'
LIMIT 1
`).Row().Scan(&columnDefault); err != nil {
t.Fatalf("query node.id default: %v", err)
}
`)
if !columnDefault.Valid || !strings.Contains(strings.ToLower(columnDefault.String), "nextval(") {
t.Fatalf("expected node.id default to be nextval(...), got %q", columnDefault.String)
}
@@ -94,10 +91,7 @@ func TestPostgresNodeCreateRepairsMissingIDDefaultContract(t *testing.T) {
router.ServeHTTP(resp, req)
assertCode(t, resp, 0)
var nodeID int64
if err := r.DB().Raw(`SELECT id FROM node WHERE name = ? ORDER BY id DESC LIMIT 1`, "pg-repair-node").Row().Scan(&nodeID); err != nil {
t.Fatalf("query created node: %v", err)
}
nodeID := mustQueryInt64(t, r, `SELECT id FROM node WHERE name = ? ORDER BY id DESC LIMIT 1`, "pg-repair-node")
if nodeID <= 0 {
t.Fatalf("expected positive node id, got %d", nodeID)
}
@@ -2,7 +2,6 @@ package contract_test
import (
"bytes"
"database/sql"
"encoding/json"
"net/http"
"net/http/httptest"
@@ -32,11 +31,7 @@ func TestTunnelCreateRuntimeRollbackContract(t *testing.T) {
`, name, name+"-secret", ip, ip, "", portRange, "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0).Error; err != nil {
t.Fatalf("insert node %s: %v", name, err)
}
var id int64
if err := repo.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&id); err != nil {
t.Fatalf("get node id %s: %v", name, err)
}
return id
return mustLastInsertID(t, repo, name)
}
entryID := insertNode("create-entry", "10.20.0.1", "30000-30010")
@@ -62,18 +57,12 @@ func TestTunnelCreateRuntimeRollbackContract(t *testing.T) {
t.Fatalf("expected node-related error, got %q", out.Msg)
}
var tunnelCount int
if err := repo.DB().Raw(`SELECT COUNT(1) FROM tunnel WHERE name = ?`, "runtime-rollback-tunnel").Row().Scan(&tunnelCount); err != nil {
t.Fatalf("count tunnel: %v", err)
}
tunnelCount := mustQueryInt(t, repo, `SELECT COUNT(1) FROM tunnel WHERE name = ?`, "runtime-rollback-tunnel")
if tunnelCount != 0 {
t.Fatalf("expected tunnel rollback, found %d records", tunnelCount)
}
var chainCount int
if err := repo.DB().Raw(`SELECT COUNT(1) FROM chain_tunnel`).Row().Scan(&chainCount); err != nil {
t.Fatalf("count chain_tunnel: %v", err)
}
chainCount := mustQueryInt(t, repo, `SELECT COUNT(1) FROM chain_tunnel`)
if chainCount != 0 {
t.Fatalf("expected chain_tunnel rollback, found %d records", chainCount)
}
@@ -96,11 +85,7 @@ func TestTunnelUpdateAssignsChainPortsContract(t *testing.T) {
`, name, name+"-secret", ip, ip, "", portRange, "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0).Error; err != nil {
t.Fatalf("insert node %s: %v", name, err)
}
var id int64
if err := repo.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&id); err != nil {
t.Fatalf("get node id %s: %v", name, err)
}
return id
return mustLastInsertID(t, repo, name)
}
entryID := insertNode("update-entry", "10.30.0.1", "40000-40010")
@@ -113,10 +98,7 @@ func TestTunnelUpdateAssignsChainPortsContract(t *testing.T) {
`, "update-port-tunnel", 1.0, 1, "tls", 99999, now, now, 1, nil, 0).Error; err != nil {
t.Fatalf("insert tunnel: %v", err)
}
var tunnelID int64
if err := repo.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&tunnelID); err != nil {
t.Fatalf("get tunnel id: %v", err)
}
tunnelID := mustLastInsertID(t, repo, "update-port-tunnel")
payload := `{"id":` + jsonInt(tunnelID) + `,"name":"update-port-tunnel","type":2,"flow":99999,"trafficRatio":1.0,"status":1,"inNodeId":[{"nodeId":` + jsonInt(entryID) + `,"protocol":"tls"}],"chainNodes":[[{"nodeId":` + jsonInt(chainID) + `,"protocol":"tls","strategy":"round"}]],"outNodeId":[{"nodeId":` + jsonInt(exitID) + `,"protocol":"tls"}]}`
req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/update", bytes.NewBufferString(payload))
@@ -127,26 +109,17 @@ func TestTunnelUpdateAssignsChainPortsContract(t *testing.T) {
router.ServeHTTP(res, req)
assertCode(t, res, 0)
var chainPort int
if err := repo.DB().Raw(`SELECT port FROM chain_tunnel WHERE tunnel_id = ? AND chain_type = 2 LIMIT 1`, tunnelID).Row().Scan(&chainPort); err != nil {
t.Fatalf("query chain port: %v", err)
}
chainPort := mustQueryInt(t, repo, `SELECT port FROM chain_tunnel WHERE tunnel_id = ? AND chain_type = 2 LIMIT 1`, tunnelID)
if chainPort <= 0 {
t.Fatalf("expected chain node port to be assigned, got %d", chainPort)
}
var outPort int
if err := repo.DB().Raw(`SELECT port FROM chain_tunnel WHERE tunnel_id = ? AND chain_type = 3 LIMIT 1`, tunnelID).Row().Scan(&outPort); err != nil {
t.Fatalf("query out port: %v", err)
}
outPort := mustQueryInt(t, repo, `SELECT port FROM chain_tunnel WHERE tunnel_id = ? AND chain_type = 3 LIMIT 1`, tunnelID)
if outPort <= 0 {
t.Fatalf("expected out node port to be assigned, got %d", outPort)
}
var entryStrategy sql.NullString
if err := repo.DB().Raw(`SELECT strategy FROM chain_tunnel WHERE tunnel_id = ? AND chain_type = 1 LIMIT 1`, tunnelID).Row().Scan(&entryStrategy); err != nil {
t.Fatalf("query entry strategy: %v", err)
}
entryStrategy := mustQueryNullString(t, repo, `SELECT strategy FROM chain_tunnel WHERE tunnel_id = ? AND chain_type = 1 LIMIT 1`, tunnelID)
if !entryStrategy.Valid || strings.TrimSpace(entryStrategy.String) == "" {
t.Fatalf("expected entry strategy to be non-null and non-empty")
}
@@ -30,11 +30,7 @@ func TestTunnelCreateWithIPPreferenceContract(t *testing.T) {
`, name, name+"-secret", v4, v4, v6, portRange, "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0).Error; err != nil {
t.Fatalf("insert node %s: %v", name, err)
}
var id int64
if err := repo.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&id); err != nil {
t.Fatalf("get node id %s: %v", name, err)
}
return id
return mustLastInsertID(t, repo, name)
}
entryID := insertDualStackNode("ip-pref-entry", "10.50.0.1", "2001:db8::1", "50000-50010")
@@ -61,8 +57,7 @@ func TestTunnelCreateWithIPPreferenceContract(t *testing.T) {
t.Fatalf("decode response: %v", err)
}
var stored string
err := repo.DB().Raw(`SELECT COALESCE(ip_preference, '') FROM tunnel WHERE name = ?`, "tunnel-"+tc.name).Row().Scan(&stored)
stored, err := tryQueryString(t, repo, `SELECT COALESCE(ip_preference, '') FROM tunnel WHERE name = ?`, "tunnel-"+tc.name)
if err != nil {
if err == sql.ErrNoRows {
t.Skipf("tunnel not created (nodes offline), skipping DB verification")
@@ -93,11 +88,7 @@ func TestTunnelUpdateIPPreferenceContract(t *testing.T) {
`, name, name+"-secret", v4, v4, v6, portRange, "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0).Error; err != nil {
t.Fatalf("insert node %s: %v", name, err)
}
var id int64
if err := repo.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&id); err != nil {
t.Fatalf("get node id %s: %v", name, err)
}
return id
return mustLastInsertID(t, repo, name)
}
entryID := insertDualStackNode("upd-entry", "10.60.0.1", "2001:db8:1::1", "60000-60010")
@@ -109,10 +100,7 @@ func TestTunnelUpdateIPPreferenceContract(t *testing.T) {
`, "update-ip-pref-tunnel", 1.0, 1, "tls", 99999, now, now, 1, nil, 0, "").Error; err != nil {
t.Fatalf("insert tunnel: %v", err)
}
var tunnelID int64
if err := repo.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&tunnelID); err != nil {
t.Fatalf("get tunnel id: %v", err)
}
tunnelID := mustLastInsertID(t, repo, "update-ip-pref-tunnel")
payload := `{"id":` + jsonInt(tunnelID) + `,"name":"update-ip-pref-tunnel","type":2,"flow":99999,"trafficRatio":1.0,"status":1,"ipPreference":"v6","inNodeId":[{"nodeId":` + jsonInt(entryID) + `,"protocol":"tls"}],"chainNodes":[],"outNodeId":[{"nodeId":` + jsonInt(exitID) + `,"protocol":"tls"}]}`
req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/update", bytes.NewBufferString(payload))
@@ -126,10 +114,7 @@ func TestTunnelUpdateIPPreferenceContract(t *testing.T) {
t.Fatalf("decode response: %v", err)
}
var stored string
if err := repo.DB().Raw(`SELECT COALESCE(ip_preference, '') FROM tunnel WHERE id = ?`, tunnelID).Row().Scan(&stored); err != nil {
t.Fatalf("query ip_preference: %v", err)
}
stored := mustQueryString(t, repo, `SELECT COALESCE(ip_preference, '') FROM tunnel WHERE id = ?`, tunnelID)
if stored != "v6" {
t.Fatalf("expected ip_preference='v6' after update, got %q", stored)
}
@@ -202,10 +187,7 @@ func TestIPPreferenceColumnDefaultContract(t *testing.T) {
t.Fatalf("insert tunnel without ip_preference: %v", err)
}
var stored string
if err := repo.DB().Raw(`SELECT COALESCE(ip_preference, '') FROM tunnel WHERE name = ?`, "no-pref-tunnel").Row().Scan(&stored); err != nil {
t.Fatalf("query ip_preference: %v", err)
}
stored := mustQueryString(t, repo, `SELECT COALESCE(ip_preference, '') FROM tunnel WHERE name = ?`, "no-pref-tunnel")
if stored != "" {
t.Fatalf("expected default ip_preference='', got %q", stored)
}
@@ -214,11 +196,7 @@ func TestIPPreferenceColumnDefaultContract(t *testing.T) {
func TestIPPreferenceColumnMigrationContract(t *testing.T) {
_, repo := setupContractRouter(t, "contract-jwt-secret")
var colCount int
err := repo.DB().Raw(`SELECT COUNT(*) FROM pragma_table_info('tunnel') WHERE name = 'ip_preference'`).Row().Scan(&colCount)
if err != nil {
t.Fatalf("check column existence: %v", err)
}
colCount := mustQueryInt(t, repo, `SELECT COUNT(*) FROM pragma_table_info('tunnel') WHERE name = 'ip_preference'`)
if colCount != 1 {
t.Fatalf("expected ip_preference column to exist in tunnel table, found %d", colCount)
}
@@ -235,10 +213,7 @@ func TestIPPreferenceCoalesceNullSafety(t *testing.T) {
t.Skipf("DB does not allow NULL ip_preference (NOT NULL constraint): %v", err)
}
var stored string
if err := repo.DB().Raw(`SELECT COALESCE(ip_preference, '') FROM tunnel WHERE name = ?`, "null-pref-tunnel").Row().Scan(&stored); err != nil {
t.Fatalf("query ip_preference: %v", err)
}
stored := mustQueryString(t, repo, `SELECT COALESCE(ip_preference, '') FROM tunnel WHERE name = ?`, "null-pref-tunnel")
if stored != "" {
t.Fatalf("COALESCE should convert NULL to empty string, got %q", stored)
}
@@ -255,10 +230,7 @@ func TestDualStackNodeIPFieldsStoredContract(t *testing.T) {
t.Fatalf("insert dual-stack node: %v", err)
}
var v4, v6 sql.NullString
if err := repo.DB().Raw(`SELECT server_ip_v4, server_ip_v6 FROM node WHERE name = ?`, "ds-verify-node").Row().Scan(&v4, &v6); err != nil {
t.Fatalf("query node IPs: %v", err)
}
v4, v6 := mustQueryTwoNullStrings(t, repo, `SELECT server_ip_v4, server_ip_v6 FROM node WHERE name = ?`, "ds-verify-node")
if !v4.Valid || v4.String != "10.70.0.1" {
t.Fatalf("expected server_ip_v4='10.70.0.1', got %v", v4)
}
@@ -283,10 +255,7 @@ func TestIPPreferenceValidValuesContract(t *testing.T) {
t.Fatalf("insert tunnel with ip_preference=%q: %v", pref, err)
}
var stored string
if err := repo.DB().Raw(`SELECT COALESCE(ip_preference, '') FROM tunnel WHERE name = ?`, name).Row().Scan(&stored); err != nil {
t.Fatalf("query ip_preference for %s: %v", name, err)
}
stored := mustQueryString(t, repo, `SELECT COALESCE(ip_preference, '') FROM tunnel WHERE name = ?`, name)
if stored != pref {
t.Fatalf("expected ip_preference=%q, got %q for %s", pref, stored, name)
}
@@ -30,11 +30,7 @@ func TestUserTunnelVisibleListContracts(t *testing.T) {
`, name, 1.0, 1, "tls", 99999, now, now, status, nil, inx).Error; err != nil {
t.Fatalf("insert tunnel %s: %v", name, err)
}
var id int64
if err := repo.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&id); err != nil {
t.Fatalf("get tunnel id %s: %v", name, err)
}
return id
return mustLastInsertID(t, repo, name)
}
enabledA := insertTunnel("enabled-A", 1, 1)