289 lines
9.2 KiB
Go
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
|
|
}
|