From 9e979aa82a6ee898976a4222478be743a4434fb2 Mon Sep 17 00:00:00 2001 From: sagitchu Date: Sat, 28 Feb 2026 12:13:28 +0800 Subject: [PATCH] refactor(tests): consolidate contract test helpers --- .../contract/db_test_helper_internal_test.go | 19 ------ .../tests/contract/db_test_helper_test.go | 42 +++++++++++++ .../tests/contract/diagnosis_contract_test.go | 63 ++----------------- .../federation_dual_panel_contract_test.go | 36 ----------- .../tests/contract/forward_contract_test.go | 12 +++- .../tunnel_visibility_contract_test.go | 11 ++-- 6 files changed, 63 insertions(+), 120 deletions(-) delete mode 100644 go-backend/tests/contract/db_test_helper_internal_test.go diff --git a/go-backend/tests/contract/db_test_helper_internal_test.go b/go-backend/tests/contract/db_test_helper_internal_test.go deleted file mode 100644 index c61fa9b..0000000 --- a/go-backend/tests/contract/db_test_helper_internal_test.go +++ /dev/null @@ -1,19 +0,0 @@ -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 -} diff --git a/go-backend/tests/contract/db_test_helper_test.go b/go-backend/tests/contract/db_test_helper_test.go index 1d4b2eb..322ade9 100644 --- a/go-backend/tests/contract/db_test_helper_test.go +++ b/go-backend/tests/contract/db_test_helper_test.go @@ -2,6 +2,8 @@ package contract_test import ( "database/sql" + "strconv" + "strings" "testing" "go-backend/internal/store/repo" @@ -117,3 +119,43 @@ func tryQueryInt(t *testing.T, r *repo.Repository, query string, args ...interfa } return v, nil } + +func valueAsInt(v interface{}) int { + switch n := v.(type) { + case float64: + return int(n) + case int: + return n + case int64: + return int(n) + default: + return 0 + } +} + +func valueAsString(v interface{}) string { + s, _ := v.(string) + return s +} + +func valueAsBool(v interface{}) bool { + switch b := v.(type) { + case bool: + return b + case float64: + return b != 0 + case int: + return b != 0 + case int64: + return b != 0 + case string: + s := strings.TrimSpace(strings.ToLower(b)) + return s == "1" || s == "t" || s == "true" || s == "yes" || s == "y" + default: + return false + } +} + +func jsonInt64(v int64) string { + return strconv.FormatInt(v, 10) +} diff --git a/go-backend/tests/contract/diagnosis_contract_test.go b/go-backend/tests/contract/diagnosis_contract_test.go index 8d414cc..e8983e9 100644 --- a/go-backend/tests/contract/diagnosis_contract_test.go +++ b/go-backend/tests/contract/diagnosis_contract_test.go @@ -1,11 +1,10 @@ -package contract +package contract_test import ( "bytes" "encoding/json" "net/http" "net/http/httptest" - "path/filepath" "strconv" "strings" "sync/atomic" @@ -13,15 +12,12 @@ import ( "time" "go-backend/internal/auth" - httpserver "go-backend/internal/http" - "go-backend/internal/http/handler" "go-backend/internal/http/response" - "go-backend/internal/store/repo" ) func TestDiagnosisChainCoverageContracts(t *testing.T) { secret := "contract-jwt-secret" - router, r := setupDiagnosisContractRouter(t, secret) + router, r := setupContractRouter(t, secret) now := time.Now().UnixMilli() if err := r.DB().Exec(` @@ -195,7 +191,7 @@ func TestDiagnosisChainCoverageContracts(t *testing.T) { func TestForwardDiagnosisRespectsTunnelIPPreferenceContract(t *testing.T) { secret := "contract-jwt-secret" - router, r := setupDiagnosisContractRouter(t, secret) + router, r := setupContractRouter(t, secret) now := time.Now().UnixMilli() if err := r.DB().Exec(` @@ -315,7 +311,7 @@ func TestForwardDiagnosisRespectsTunnelIPPreferenceContract(t *testing.T) { func TestDiagnosisUsesFederationRuntimeForRemoteNodes(t *testing.T) { secret := "contract-jwt-secret" - router, r := setupDiagnosisContractRouter(t, secret) + router, r := setupContractRouter(t, secret) now := time.Now().UnixMilli() remoteToken := "remote-diagnose-token" @@ -465,54 +461,3 @@ func TestDiagnosisUsesFederationRuntimeForRemoteNodes(t *testing.T) { t.Fatalf("expected federation runtime diagnose endpoint to be called") } } - -func valueAsInt(v interface{}) int { - switch n := v.(type) { - case float64: - return int(n) - case int: - return n - case int64: - return int(n) - default: - return 0 - } -} - -func valueAsString(v interface{}) string { - s, _ := v.(string) - return s -} - -func valueAsBool(v interface{}) bool { - switch b := v.(type) { - case bool: - return b - case float64: - return b != 0 - case int: - return b != 0 - case int64: - return b != 0 - case string: - s := strings.TrimSpace(strings.ToLower(b)) - return s == "1" || s == "t" || s == "true" || s == "yes" || s == "y" - default: - return false - } -} - -func setupDiagnosisContractRouter(t *testing.T, jwtSecret string) (http.Handler, *repo.Repository) { - t.Helper() - dbPath := filepath.Join(t.TempDir(), "diagnosis-contract.db") - r, err := repo.Open(dbPath) - if err != nil { - t.Fatalf("open sqlite: %v", err) - } - t.Cleanup(func() { - _ = r.Close() - }) - - h := handler.New(r, jwtSecret) - return httpserver.NewRouter(h, jwtSecret), r -} diff --git a/go-backend/tests/contract/federation_dual_panel_contract_test.go b/go-backend/tests/contract/federation_dual_panel_contract_test.go index b73c277..7c1e40c 100644 --- a/go-backend/tests/contract/federation_dual_panel_contract_test.go +++ b/go-backend/tests/contract/federation_dual_panel_contract_test.go @@ -624,42 +624,6 @@ func waitNodeStatus(t *testing.T, r *repo.Repository, nodeID int64, expectedStat } } -func valueAsInt(v interface{}) int { - switch n := v.(type) { - case float64: - return int(n) - case int: - return n - case int64: - return int(n) - default: - return 0 - } -} - -func valueAsString(v interface{}) string { - s, _ := v.(string) - return s -} - -func valueAsBool(v interface{}) bool { - switch b := v.(type) { - case bool: - return b - case float64: - return b != 0 - case int: - return b != 0 - case int64: - return b != 0 - case string: - s := strings.TrimSpace(strings.ToLower(b)) - return s == "1" || s == "t" || s == "true" || s == "yes" || s == "y" - default: - return false - } -} - func TestFederationRuntimeCommandPortRangeEnforcement(t *testing.T) { providerSecret := "provider-portrange-jwt" providerRouter, providerRepo := setupContractRouter(t, providerSecret) diff --git a/go-backend/tests/contract/forward_contract_test.go b/go-backend/tests/contract/forward_contract_test.go index b4d1c4a..eaaa53f 100644 --- a/go-backend/tests/contract/forward_contract_test.go +++ b/go-backend/tests/contract/forward_contract_test.go @@ -109,7 +109,11 @@ func TestForwardOwnershipAndScopeContracts(t *testing.T) { if !ok { t.Fatalf("expected object item, got %T", arr[0]) } - if got := int64(item["id"].(float64)); got != userForwardID { + idFloat, ok := item["id"].(float64) + if !ok { + t.Fatalf("expected id to be float64, got %T", item["id"]) + } + if got := int64(idFloat); got != userForwardID { t.Fatalf("expected forward id %d, got %d", userForwardID, got) } }) @@ -144,7 +148,11 @@ func TestForwardOwnershipAndScopeContracts(t *testing.T) { if _, ok := first["message"]; !ok { t.Fatalf("expected message field in diagnosis result") } - if got := int(first["fromChainType"].(float64)); got != 1 { + fromChainTypeFloat, ok := first["fromChainType"].(float64) + if !ok { + t.Fatalf("expected fromChainType to be float64, got %T", first["fromChainType"]) + } + if got := int(fromChainTypeFloat); got != 1 { t.Fatalf("expected fromChainType=1, got %d", got) } }) diff --git a/go-backend/tests/contract/tunnel_visibility_contract_test.go b/go-backend/tests/contract/tunnel_visibility_contract_test.go index 56c33e9..30553a2 100644 --- a/go-backend/tests/contract/tunnel_visibility_contract_test.go +++ b/go-backend/tests/contract/tunnel_visibility_contract_test.go @@ -1,4 +1,4 @@ -package contract +package contract_test import ( "encoding/json" @@ -13,7 +13,7 @@ import ( func TestUserTunnelVisibleListContracts(t *testing.T) { secret := "contract-jwt-secret" - router, repo := setupDiagnosisContractRouter(t, secret) + router, repo := setupContractRouter(t, secret) now := time.Now().UnixMilli() if err := repo.DB().Exec(` @@ -126,8 +126,11 @@ func collectTunnelIDs(t *testing.T, data interface{}) map[int64]bool { if !ok { t.Fatalf("expected object item, got %T", item) } - id := int64(obj["id"].(float64)) - ids[id] = true + idFloat, ok := obj["id"].(float64) + if !ok { + t.Fatalf("expected id to be float64, got %T", obj["id"]) + } + ids[int64(idFloat)] = true } return ids }