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)
}