Files
qsc20001102 142e5dc7d6
ci / Go checks (ubuntu-latest) (push) Has been cancelled
ci / Go checks (windows-latest) (push) Has been cancelled
初版功能完成
2026-08-29 13:12:17 +08:00

289 lines
9.2 KiB
Go

package database
import (
"context"
"database/sql"
"encoding/json"
"errors"
"fmt"
"math"
"net/netip"
"path/filepath"
"testing"
"time"
"remlink/internal/model"
)
func TestEmptyListRepositoriesEncodeAsArrays(t *testing.T) {
db := openTestDB(t)
store := NewStore(db)
ctx := context.Background()
nodes, err := store.ListNodes(ctx)
if err != nil {
t.Fatal(err)
}
sessions, err := store.ListSessions(ctx)
if err != nil {
t.Fatal(err)
}
events, err := store.ListEvents(ctx, model.EventLogFilter{})
if err != nil {
t.Fatal(err)
}
for name, value := range map[string]any{"nodes": nodes, "sessions": sessions, "events": events} {
encoded, err := json.Marshal(value)
if err != nil {
t.Fatal(err)
}
if string(encoded) != "[]" {
t.Errorf("empty %s encoded as %s, want []", name, encoded)
}
}
}
func TestOpenMigratesAndReopens(t *testing.T) {
path := filepath.Join(t.TempDir(), "remlink.db")
ctx := context.Background()
db, err := Open(ctx, path)
if err != nil {
t.Fatal(err)
}
wantTables := []string{"settings", "nodes", "sessions", "session_cidrs", "session_stats", "event_logs"}
for _, table := range wantTables {
var count int
if err := db.QueryRowContext(ctx,
`SELECT COUNT(*) FROM sqlite_master WHERE type = 'table' AND name = ?`, table).Scan(&count); err != nil {
t.Fatal(err)
}
if count != 1 {
t.Errorf("table %s count = %d, want 1", table, count)
}
}
var migrationCount int
if err := db.QueryRowContext(ctx, `SELECT COUNT(*) FROM schema_migrations`).Scan(&migrationCount); err != nil {
t.Fatal(err)
}
if migrationCount != 1 {
t.Fatalf("migration count = %d, want 1", migrationCount)
}
if err := db.Close(); err != nil {
t.Fatal(err)
}
db, err = Open(ctx, path)
if err != nil {
t.Fatal(err)
}
defer db.Close()
if err := db.QueryRowContext(ctx, `SELECT COUNT(*) FROM schema_migrations`).Scan(&migrationCount); err != nil {
t.Fatal(err)
}
if migrationCount != 1 {
t.Fatalf("migration count after reopen = %d, want 1", migrationCount)
}
}
func TestSettingsRoundTrip(t *testing.T) {
db := openTestDB(t)
ctx := context.Background()
if value, found, err := GetSetting(ctx, db, "overlay_cidr"); err != nil || found || value != "" {
t.Fatalf("missing GetSetting = %q, %v, %v", value, found, err)
}
if err := SetSetting(ctx, db, "overlay_cidr", "10.88.0.0/16"); err != nil {
t.Fatal(err)
}
if err := SetSetting(ctx, db, "overlay_cidr", "10.99.0.0/16"); err != nil {
t.Fatal(err)
}
value, found, err := GetSetting(ctx, db, "overlay_cidr")
if err != nil || !found || value != "10.99.0.0/16" {
t.Fatalf("GetSetting = %q, %v, %v", value, found, err)
}
}
func TestNodeRepository(t *testing.T) {
db := openTestDB(t)
store := NewStore(db)
ctx := context.Background()
node := model.Node{
ID: "96e5d037-2358-4ff6-b706-1fbf56d196bc",
Type: model.NodeTypeEngineer,
Name: "Engineer-A",
OverlayIP: netip.MustParseAddr("10.88.0.2"),
WGPublicKey: "public-key-a",
NodeTokenHash: []byte("token-hash-a"),
}
if err := store.CreateNode(ctx, node); err != nil {
t.Fatal(err)
}
got, err := store.GetNode(ctx, node.ID)
if err != nil {
t.Fatal(err)
}
if got.ID != node.ID || got.OverlayIP != node.OverlayIP || got.Status != model.NodeOffline {
t.Fatalf("unexpected node: %+v", got)
}
node.Name = "Engineer-Renamed"
node.NodeTokenHash = []byte("token-hash-b")
node.Version = "1.0.0"
if err := store.UpdateNodeRegistration(ctx, node); err != nil {
t.Fatal(err)
}
got, err = store.GetNode(ctx, node.ID)
if err != nil {
t.Fatal(err)
}
if got.Name != node.Name || got.Version != "1.0.0" || string(got.NodeTokenHash) != "token-hash-b" {
t.Fatalf("registration update not persisted: %+v", got)
}
if err := store.UpdateNodeOverlayIP(ctx, node.ID, netip.MustParseAddr("10.88.0.50")); err != nil {
t.Fatal(err)
}
addresses, err := store.ListOverlayIPs(ctx)
if err != nil {
t.Fatal(err)
}
if len(addresses) != 1 || addresses[0].String() != "10.88.0.50" {
t.Fatalf("overlay addresses = %v", addresses)
}
if err := store.DeleteNode(ctx, node.ID); err != nil {
t.Fatal(err)
}
if _, err := store.GetNode(ctx, node.ID); !errors.Is(err, ErrNodeNotFound) {
t.Fatalf("GetNode after delete error = %v", err)
}
}
func TestOverlayMigrationCommitsNodeIPsAndNetworkSettingAtomically(t *testing.T) {
db := openTestDB(t)
store := NewStore(db)
ctx := context.Background()
node := model.Node{
ID: "engineer", Type: model.NodeTypeEngineer, Name: "Engineer",
OverlayIP: netip.MustParseAddr("10.88.0.2"), WGPublicKey: "engineer-key", NodeTokenHash: []byte("hash"),
}
if err := store.CreateNode(ctx, node); err != nil {
t.Fatal(err)
}
if err := SetSetting(ctx, db, "admin.network", "old"); err != nil {
t.Fatal(err)
}
if err := store.ReplaceNodeOverlayIPsAndSetting(ctx,
map[string]netip.Addr{"engineer": netip.MustParseAddr("10.99.0.2")}, "admin.network", "new"); err != nil {
t.Fatal(err)
}
migrated, _ := store.GetNode(ctx, "engineer")
setting, found, err := GetSetting(ctx, db, "admin.network")
if err != nil || !found || migrated.OverlayIP.String() != "10.99.0.2" || setting != "new" {
t.Fatalf("committed migration node=%s setting=%q found=%v error=%v", migrated.OverlayIP, setting, found, err)
}
err = store.ReplaceNodeOverlayIPsAndSetting(ctx, map[string]netip.Addr{
"engineer": netip.MustParseAddr("10.77.0.2"),
"missing": netip.MustParseAddr("10.77.0.3"),
}, "admin.network", "partial")
if err == nil {
t.Fatal("migration with a missing Node unexpectedly succeeded")
}
afterFailure, _ := store.GetNode(ctx, "engineer")
setting, _, _ = GetSetting(ctx, db, "admin.network")
if afterFailure.OverlayIP.String() != "10.99.0.2" || setting != "new" {
t.Fatalf("failed migration was not atomic: node=%s setting=%q", afterFailure.OverlayIP, setting)
}
}
func TestForeignKeysEnabled(t *testing.T) {
db := openTestDB(t)
_, err := db.Exec(`INSERT INTO sessions(session_id, engineer_node_id, site_node_id, status, created_at)
VALUES('1', 'missing-a', 'missing-b', 'CREATING', '2026-08-25T00:00:00Z')`)
if err == nil {
t.Fatal("session with missing node references was accepted")
}
}
func TestEventRepositoryFiltersByTimeAndNormalizesModule(t *testing.T) {
store := NewStore(openTestDB(t))
ctx := context.Background()
base := time.Date(2026, 8, 25, 12, 0, 0, 0, time.UTC)
for index, eventTime := range []time.Time{base.Add(-time.Hour), base, base.Add(time.Hour)} {
if err := store.AppendEvent(ctx, model.EventLog{Time: eventTime, Level: "info", Module: "session", Message: fmt.Sprintf("event-%d", index)}); err != nil {
t.Fatal(err)
}
}
events, err := store.ListEvents(ctx, model.EventLogFilter{Module: "SESSION", From: base.Add(-time.Minute), To: base.Add(time.Minute)})
if err != nil || len(events) != 1 || events[0].Message != "event-1" || events[0].Module != "SESSION" {
t.Fatalf("events = %+v, %v", events, err)
}
if err := store.AppendEvent(ctx, model.EventLog{Level: "INFO", Module: "CUSTOM", Message: "outside taxonomy"}); err == nil {
t.Fatal("AppendEvent accepted a twelfth logging module")
}
}
func TestSessionRepositoryLifecycleAndStartupClose(t *testing.T) {
db := openTestDB(t)
store := NewStore(db)
ctx := context.Background()
createSessionTestNodes(t, store)
session := model.Session{
ID: 18446744073709551614, EngineerNodeID: "engineer", SiteNodeID: "site",
Status: model.SessionCreating, CIDRs: []netip.Prefix{
netip.MustParsePrefix("192.168.13.0/24"), netip.MustParsePrefix("172.20.0.0/16"),
},
}
if err := store.CreateSession(ctx, session); err != nil {
t.Fatal(err)
}
if err := store.UpdateSessionStatus(ctx, session.ID, model.SessionActive, ""); err != nil {
t.Fatal(err)
}
counters := model.SessionCounters{UploadBytes: 100, DownloadBytes: 200, UploadPackets: 3, DownloadPackets: 4}
if err := store.UpdateSessionStats(ctx, session.ID, counters); err != nil {
t.Fatal(err)
}
got, err := store.GetSession(ctx, session.ID)
if err != nil {
t.Fatal(err)
}
if got.Status != model.SessionActive || got.ActiveAt == nil || got.Counters != counters || len(got.CIDRs) != 2 {
t.Fatalf("unexpected persisted Session: %+v", got)
}
if err := store.UpdateSessionStats(ctx, session.ID, model.SessionCounters{UploadBytes: math.MaxUint64}); err == nil {
t.Fatal("SQLite-unsafe counter was accepted")
}
closed, err := store.CloseOpenSessions(ctx)
if err != nil || closed != 1 {
t.Fatalf("CloseOpenSessions = %d, %v", closed, err)
}
got, err = store.GetSession(ctx, session.ID)
if err != nil || got.Status != model.SessionClosed || got.ClosedAt == nil {
t.Fatalf("Session after startup close = %+v, %v", got, err)
}
}
func createSessionTestNodes(t *testing.T, store *Store) {
t.Helper()
for _, node := range []model.Node{
{ID: "engineer", Type: model.NodeTypeEngineer, Name: "Engineer", OverlayIP: netip.MustParseAddr("10.88.0.2"), WGPublicKey: "engineer-key", NodeTokenHash: []byte("a")},
{ID: "site", Type: model.NodeTypeSite, Name: "Site", OverlayIP: netip.MustParseAddr("10.88.0.3"), WGPublicKey: "site-key", NodeTokenHash: []byte("b")},
} {
if err := store.CreateNode(context.Background(), node); err != nil {
t.Fatal(err)
}
}
}
func openTestDB(t *testing.T) *sql.DB {
t.Helper()
db, err := Open(context.Background(), filepath.Join(t.TempDir(), "test.db"))
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = db.Close() })
return db
}