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

454 lines
20 KiB
Go

package session
import (
"context"
"net/netip"
"path/filepath"
"sync"
"testing"
"time"
"remlink/internal/database"
"remlink/internal/model"
"remlink/internal/protocol"
)
func TestManagerFullLifecycleAndCounters(t *testing.T) {
manager, store, sender, engineer, site := newManagerTest(t, time.Second)
ctx := context.Background()
create := envelope(t, protocol.ControlCreateSession, "request-1", protocol.CreateSessionPayload{
SiteNodeID: site.ID, TargetCIDRs: []string{"192.168.13.0/24", "172.20.0.0/16"},
})
if err := manager.HandleControl(ctx, engineer, create); err != nil {
t.Fatal(err)
}
prepare := sender.last(t, protocol.ControlPrepareSession).payload.(protocol.PrepareSessionPayload)
stored, err := store.GetSession(ctx, prepare.SessionID)
if err != nil || stored.Status != model.SessionPreparingSite {
t.Fatalf("persisted preparing Session = %+v, %v", stored, err)
}
if err := manager.HandleControl(ctx, engineer, create); err != nil {
t.Fatal(err)
}
rejection := sender.lastForNode(t, protocol.ControlStopSession, engineer.ID)
if rejection.requestID != "request-1" || rejection.payload.(protocol.StopSessionPayload).Reason != string(protocol.ErrorEngineerSessionExists) {
t.Fatalf("second Session rejection = %+v", rejection)
}
result := protocol.PrepareResultPayload{
SessionID: prepare.SessionID, OK: true, SubnetGatewayStatus: "netstack", TCPCapacity: 2048, UDPCapacity: 4096,
RouteResults: []protocol.RouteResult{{CIDR: "192.168.13.0/24", Result: "DIRECT"}, {CIDR: "172.20.0.0/16", Result: "ROUTED"}},
}
if err := manager.HandleControl(ctx, site, envelope(t, protocol.ControlPrepareResult, "", result)); err != nil {
t.Fatal(err)
}
configuration := sender.last(t, protocol.ControlSessionConfig)
if configuration.nodeID != engineer.ID || configuration.requestID != "request-1" {
t.Fatalf("SESSION_CONFIG routing = %+v", configuration)
}
if got := configuration.payload.(protocol.SessionConfigPayload); got.PeerOverlayIP != site.OverlayIP.String() || got.UDPPort != 51821 || len(got.CIDRs) != 2 {
t.Fatalf("SESSION_CONFIG = %+v", got)
}
stored, _ = store.GetSession(ctx, prepare.SessionID)
if stored.Status != model.SessionReady {
t.Fatalf("status = %s, want READY", stored.Status)
}
if err := manager.HandleControl(ctx, engineer, envelope(t, protocol.ControlRoutesReady, "", protocol.RoutesReadyPayload{SessionID: prepare.SessionID})); err != nil {
t.Fatal(err)
}
stored, _ = store.GetSession(ctx, prepare.SessionID)
if stored.Status != model.SessionActive || stored.ActiveAt == nil || sender.count(protocol.ControlSessionActive) != 2 {
t.Fatalf("Active Session = %+v, active notifications=%d", stored, sender.count(protocol.ControlSessionActive))
}
counters := model.SessionCounters{UploadBytes: 1234, DownloadBytes: 5678, UploadPackets: 12, DownloadPackets: 34}
stats := envelope(t, protocol.ControlSessionStats, "", protocol.SessionStatsPayload{SessionID: prepare.SessionID, Counters: counters})
if err := manager.HandleControl(ctx, engineer, stats); err != nil {
t.Fatal(err)
}
stored, _ = store.GetSession(ctx, prepare.SessionID)
if stored.Counters != counters {
t.Fatalf("persisted counters = %+v", stored.Counters)
}
if err := manager.HandleControl(ctx, engineer, envelope(t, protocol.ControlSessionStats, "", protocol.SessionStatsPayload{
SessionID: prepare.SessionID, Counters: model.SessionCounters{UploadBytes: 1},
})); err != nil {
t.Fatal(err)
}
stored, _ = store.GetSession(ctx, prepare.SessionID)
if stored.Counters != counters {
t.Fatalf("stale report moved counters backward: %+v", stored.Counters)
}
if err := manager.HandleControl(ctx, site, envelope(t, protocol.ControlStopSession, "", protocol.StopSessionPayload{
SessionID: prepare.SessionID, Reason: "operator",
})); err != nil {
t.Fatal(err)
}
stored, _ = store.GetSession(ctx, prepare.SessionID)
if stored.Status != model.SessionClosed || stored.ClosedAt == nil {
t.Fatalf("closed Session = %+v", stored)
}
}
func TestManagerDisconnectNodeClosesStaleRuntimeAndAllowsReconnect(t *testing.T) {
manager, store, sender, engineer, site := newManagerTest(t, time.Second)
ctx := context.Background()
if err := manager.HandleControl(ctx, engineer, envelope(t, protocol.ControlCreateSession, "before-restart", protocol.CreateSessionPayload{
SiteNodeID: site.ID, TargetCIDRs: []string{"192.168.13.0/24"},
})); err != nil {
t.Fatal(err)
}
prepare := sender.last(t, protocol.ControlPrepareSession).payload.(protocol.PrepareSessionPayload)
if err := manager.HandleControl(ctx, site, envelope(t, protocol.ControlPrepareResult, "", protocol.PrepareResultPayload{
SessionID: prepare.SessionID, OK: true, SubnetGatewayStatus: "netstack", TCPCapacity: 1, UDPCapacity: 1,
RouteResults: []protocol.RouteResult{{CIDR: "192.168.13.0/24", Result: "DIRECT"}},
})); err != nil {
t.Fatal(err)
}
if err := manager.HandleControl(ctx, engineer, envelope(t, protocol.ControlRoutesReady, "", protocol.RoutesReadyPayload{SessionID: prepare.SessionID})); err != nil {
t.Fatal(err)
}
if err := manager.DisconnectNode(ctx, engineer.ID, "NODE_RUNTIME_REBUILT"); err != nil {
t.Fatal(err)
}
stored, err := store.GetSession(ctx, prepare.SessionID)
if err != nil || stored.Status != model.SessionClosed || sender.count(protocol.ControlStopSession) != 2 {
t.Fatalf("reconciled Session=%+v err=%v STOP count=%d", stored, err, sender.count(protocol.ControlStopSession))
}
if err := manager.HandleControl(ctx, engineer, envelope(t, protocol.ControlCreateSession, "after-restart", protocol.CreateSessionPayload{
SiteNodeID: site.ID, TargetCIDRs: []string{"192.168.21.0/24"},
})); err != nil {
t.Fatalf("Node remained locked by stale Session after Bootstrap reconciliation: %v", err)
}
if next := sender.last(t, protocol.ControlPrepareSession).payload.(protocol.PrepareSessionPayload); next.SessionID == prepare.SessionID || sender.count(protocol.ControlPrepareSession) != 2 {
t.Fatalf("new preparation was not accepted after reconciliation: old=%d new=%d count=%d", prepare.SessionID, next.SessionID, sender.count(protocol.ControlPrepareSession))
}
}
func TestManagerClosesActiveSessionWhenSiteBecomesOffline(t *testing.T) {
manager, store, sender, engineer, site := newManagerTest(t, time.Second)
ctx := context.Background()
if err := manager.HandleControl(ctx, engineer, envelope(t, protocol.ControlCreateSession, "site-offline", protocol.CreateSessionPayload{
SiteNodeID: site.ID, TargetCIDRs: []string{"192.168.17.0/24"},
})); err != nil {
t.Fatal(err)
}
prepare := sender.last(t, protocol.ControlPrepareSession).payload.(protocol.PrepareSessionPayload)
if err := manager.HandleControl(ctx, site, envelope(t, protocol.ControlPrepareResult, "", protocol.PrepareResultPayload{
SessionID: prepare.SessionID, OK: true, SubnetGatewayStatus: "netstack", TCPCapacity: 1, UDPCapacity: 1,
RouteResults: []protocol.RouteResult{{CIDR: "192.168.17.0/24", Result: "DIRECT"}},
})); err != nil {
t.Fatal(err)
}
if err := manager.HandleControl(ctx, engineer, envelope(t, protocol.ControlRoutesReady, "", protocol.RoutesReadyPayload{SessionID: prepare.SessionID})); err != nil {
t.Fatal(err)
}
if err := store.UpdateNodeStatus(ctx, site.ID, model.NodeOffline); err != nil {
t.Fatal(err)
}
site.Status = model.NodeOffline
if err := manager.HandleNodeStatusChange(ctx, site, model.NodeOffline); err != nil {
t.Fatal(err)
}
stored, err := store.GetSession(ctx, prepare.SessionID)
if err != nil || stored.Status != model.SessionClosed || stored.ClosedAt == nil {
t.Fatalf("offline Site Session = %+v, %v", stored, err)
}
stop := sender.lastForNode(t, protocol.ControlStopSession, engineer.ID).payload.(protocol.StopSessionPayload)
if stop.SessionID != prepare.SessionID || stop.Reason != string(protocol.ErrorSiteOffline) {
t.Fatalf("Engineer STOP_SESSION = %+v", stop)
}
if err := manager.HandleControl(ctx, engineer, envelope(t, protocol.ControlCreateSession, "after-offline", protocol.CreateSessionPayload{
SiteNodeID: site.ID, TargetCIDRs: []string{"192.168.107.0/24"},
})); err != nil {
t.Fatal(err)
}
if got := sender.lastForNode(t, protocol.ControlStopSession, engineer.ID).payload.(protocol.StopSessionPayload).Reason; got != string(protocol.ErrorSiteOffline) {
t.Fatalf("offline Site reconnect reason = %s", got)
}
}
func TestManagerRejectsInvalidCIDRAndTimesOut(t *testing.T) {
manager, store, sender, engineer, site := newManagerTest(t, 20*time.Millisecond)
ctx := context.Background()
invalid := envelope(t, protocol.ControlCreateSession, "bad", protocol.CreateSessionPayload{
SiteNodeID: site.ID, TargetCIDRs: []string{"10.88.4.0/24"},
})
if err := manager.HandleControl(ctx, engineer, invalid); err != nil {
t.Fatal(err)
}
if got := sender.last(t, protocol.ControlStopSession).payload.(protocol.StopSessionPayload).Reason; got != string(protocol.ErrorCIDROverlayConflict) {
t.Fatalf("invalid CIDR reason = %s", got)
}
valid := envelope(t, protocol.ControlCreateSession, "timeout", protocol.CreateSessionPayload{
SiteNodeID: site.ID, TargetCIDRs: []string{"192.168.13.0/24"},
})
if err := manager.HandleControl(ctx, engineer, valid); err != nil {
t.Fatal(err)
}
prepare := sender.last(t, protocol.ControlPrepareSession).payload.(protocol.PrepareSessionPayload)
deadline := time.Now().Add(time.Second)
for time.Now().Before(deadline) {
stored, err := store.GetSession(ctx, prepare.SessionID)
if err == nil && stored.Status == model.SessionFailed {
if stored.ErrorCode != string(protocol.ErrorSessionTimeout) {
t.Fatalf("timeout error code = %s", stored.ErrorCode)
}
return
}
time.Sleep(5 * time.Millisecond)
}
t.Fatal("Session did not reach FAILED after prepare timeout")
}
func TestManagerRejectsDefaultOnlyPrepareResult(t *testing.T) {
manager, store, sender, engineer, site := newManagerTest(t, time.Second)
ctx := context.Background()
if err := manager.HandleControl(ctx, engineer, envelope(t, protocol.ControlCreateSession, "default-only", protocol.CreateSessionPayload{
SiteNodeID: site.ID, TargetCIDRs: []string{"192.168.13.0/24"},
})); err != nil {
t.Fatal(err)
}
prepare := sender.last(t, protocol.ControlPrepareSession).payload.(protocol.PrepareSessionPayload)
if err := manager.HandleControl(ctx, site, envelope(t, protocol.ControlPrepareResult, "", protocol.PrepareResultPayload{
SessionID: prepare.SessionID, OK: true, SubnetGatewayStatus: "netstack", TCPCapacity: 2048, UDPCapacity: 4096,
RouteResults: []protocol.RouteResult{{CIDR: "192.168.13.0/24", Result: "DEFAULT_ONLY"}},
})); err != nil {
t.Fatal(err)
}
stored, err := store.GetSession(ctx, prepare.SessionID)
if err != nil {
t.Fatal(err)
}
if stored.Status != model.SessionFailed || stored.ErrorCode != string(protocol.ErrorSiteNoRoute) {
t.Fatalf("Session = %+v, want FAILED/%s", stored, protocol.ErrorSiteNoRoute)
}
rejection := sender.lastForNode(t, protocol.ControlStopSession, engineer.ID)
if rejection.nodeID != engineer.ID || rejection.requestID != "default-only" || rejection.payload.(protocol.StopSessionPayload).SessionID != prepare.SessionID {
t.Fatalf("PREPARE rejection lost CREATE correlation: %+v", rejection)
}
}
func TestManagerRunsConcurrentDuplicateCIDRSessions(t *testing.T) {
manager, store, sender, engineerA, siteA := newManagerTest(t, time.Second)
ctx := context.Background()
createOnlineNode := func(id string, nodeType model.NodeType, address string) model.Node {
node := model.Node{ID: id, Type: nodeType, Name: id, OverlayIP: netip.MustParseAddr(address), WGPublicKey: id + "-key", NodeTokenHash: []byte(id)}
if err := store.CreateNode(ctx, node); err != nil {
t.Fatal(err)
}
if err := store.UpdateNodeHeartbeat(ctx, id, model.NodeOnline, time.Now().UTC(), "1.0", "test"); err != nil {
t.Fatal(err)
}
node, _ = store.GetNode(ctx, id)
return node
}
engineerB := createOnlineNode("engineer-b", model.NodeTypeEngineer, "10.88.0.4")
engineerC := createOnlineNode("engineer-c", model.NodeTypeEngineer, "10.88.0.5")
siteB := createOnlineNode("site-b", model.NodeTypeSite, "10.88.0.6")
activate := func(engineer, site model.Node, requestID string) uint64 {
if err := manager.HandleControl(ctx, engineer, envelope(t, protocol.ControlCreateSession, requestID, protocol.CreateSessionPayload{
SiteNodeID: site.ID, TargetCIDRs: []string{"192.168.13.0/24"},
})); err != nil {
t.Fatal(err)
}
prepare := sender.last(t, protocol.ControlPrepareSession).payload.(protocol.PrepareSessionPayload)
if err := manager.HandleControl(ctx, site, envelope(t, protocol.ControlPrepareResult, "", protocol.PrepareResultPayload{
SessionID: prepare.SessionID, OK: true, SubnetGatewayStatus: "netstack", TCPCapacity: 2048, UDPCapacity: 4096,
RouteResults: []protocol.RouteResult{{CIDR: "192.168.13.0/24", Result: "DIRECT"}},
})); err != nil {
t.Fatal(err)
}
if err := manager.HandleControl(ctx, engineer, envelope(t, protocol.ControlRoutesReady, "", protocol.RoutesReadyPayload{SessionID: prepare.SessionID})); err != nil {
t.Fatal(err)
}
return prepare.SessionID
}
idA := activate(engineerA, siteA, "a")
idB := activate(engineerB, siteB, "b")
idC := activate(engineerC, siteA, "c")
if idA == idB || idA == idC || idB == idC {
t.Fatalf("SessionIDs are not unique: %d %d %d", idA, idB, idC)
}
sessions, err := store.ListSessions(ctx)
if err != nil || len(sessions) != 3 {
t.Fatalf("Sessions = %+v, %v", sessions, err)
}
for _, current := range sessions {
if current.Status != model.SessionActive || len(current.CIDRs) != 1 || current.CIDRs[0].String() != "192.168.13.0/24" {
t.Fatalf("unexpected concurrent Session: %+v", current)
}
}
}
func TestManagerUsesMigratedNetworkForNewSessions(t *testing.T) {
manager, store, sender, engineer, site := newManagerTest(t, time.Second)
ctx := context.Background()
newOverlay := netip.MustParsePrefix("10.99.0.0/24")
if err := manager.ReconfigureNetwork(newOverlay, 1400, 6300); err != nil {
t.Fatal(err)
}
if err := store.UpdateNodeOverlayIP(ctx, engineer.ID, netip.MustParseAddr("10.99.0.2")); err != nil {
t.Fatal(err)
}
if err := store.UpdateNodeOverlayIP(ctx, site.ID, netip.MustParseAddr("10.99.0.3")); err != nil {
t.Fatal(err)
}
engineer, _ = store.GetNode(ctx, engineer.ID)
site, _ = store.GetNode(ctx, site.ID)
if err := manager.HandleControl(ctx, engineer, envelope(t, protocol.ControlCreateSession, "new-overlay-conflict", protocol.CreateSessionPayload{
SiteNodeID: site.ID, TargetCIDRs: []string{"10.99.0.0/28"},
})); err != nil {
t.Fatal(err)
}
if got := sender.lastForNode(t, protocol.ControlStopSession, engineer.ID).payload.(protocol.StopSessionPayload).Reason; got != string(protocol.ErrorCIDROverlayConflict) {
t.Fatalf("migrated Overlay conflict reason = %s", got)
}
if err := manager.HandleControl(ctx, engineer, envelope(t, protocol.ControlCreateSession, "after-migration", protocol.CreateSessionPayload{
SiteNodeID: site.ID, TargetCIDRs: []string{"10.88.0.0/24"},
})); err != nil {
t.Fatal(err)
}
prepare := sender.last(t, protocol.ControlPrepareSession).payload.(protocol.PrepareSessionPayload)
if prepare.EngineerOverlayIP != "10.99.0.2" {
t.Fatalf("PREPARE Engineer IP = %s", prepare.EngineerOverlayIP)
}
if err := manager.HandleControl(ctx, site, envelope(t, protocol.ControlPrepareResult, "", protocol.PrepareResultPayload{
SessionID: prepare.SessionID, OK: true, SubnetGatewayStatus: "netstack", TCPCapacity: 2048, UDPCapacity: 4096,
RouteResults: []protocol.RouteResult{{CIDR: "10.88.0.0/24", Result: "DIRECT"}},
})); err != nil {
t.Fatal(err)
}
configured := sender.last(t, protocol.ControlSessionConfig).payload.(protocol.SessionConfigPayload)
if configured.MTU != 1400 || configured.UDPPort != 6300 {
t.Fatalf("SESSION_CONFIG retained stale network values: %+v", configured)
}
}
func TestManagerQuiescesSessionCreationDuringNetworkMigration(t *testing.T) {
manager, _, sender, engineer, site := newManagerTest(t, time.Second)
ctx := context.Background()
if err := manager.BeginNetworkMigration(ctx, "NETWORK_CONFIG_CHANGED"); err != nil {
t.Fatal(err)
}
create := protocol.CreateSessionPayload{SiteNodeID: site.ID, TargetCIDRs: []string{"192.168.13.0/24"}}
if err := manager.HandleControl(ctx, engineer, envelope(t, protocol.ControlCreateSession, "during-migration", create)); err != nil {
t.Fatal(err)
}
rejection := sender.lastForNode(t, protocol.ControlStopSession, engineer.ID)
if rejection.requestID != "during-migration" || rejection.payload.(protocol.StopSessionPayload).Reason != string(protocol.ErrorServerUnreachable) {
t.Fatalf("migration rejection = %+v", rejection)
}
manager.EndNetworkMigration()
if err := manager.HandleControl(ctx, engineer, envelope(t, protocol.ControlCreateSession, "after-migration", create)); err != nil {
t.Fatal(err)
}
if sender.last(t, protocol.ControlPrepareSession).payload.(protocol.PrepareSessionPayload).SessionID == 0 {
t.Fatal("Session creation remained quiesced after migration")
}
}
type sentMessage struct {
nodeID, requestID string
messageType protocol.ControlMessageType
payload any
}
type fakeSender struct {
mu sync.Mutex
messages []sentMessage
}
func (f *fakeSender) Send(_ context.Context, nodeID string, messageType protocol.ControlMessageType, payload any) error {
return f.SendRequest(context.Background(), nodeID, messageType, "", payload)
}
func (f *fakeSender) SendRequest(_ context.Context, nodeID string, messageType protocol.ControlMessageType, requestID string, payload any) error {
f.mu.Lock()
f.messages = append(f.messages, sentMessage{nodeID: nodeID, requestID: requestID, messageType: messageType, payload: payload})
f.mu.Unlock()
return nil
}
func (f *fakeSender) last(t *testing.T, messageType protocol.ControlMessageType) sentMessage {
t.Helper()
f.mu.Lock()
defer f.mu.Unlock()
for index := len(f.messages) - 1; index >= 0; index-- {
if f.messages[index].messageType == messageType {
return f.messages[index]
}
}
t.Fatalf("no %s message", messageType)
return sentMessage{}
}
func (f *fakeSender) lastForNode(t *testing.T, messageType protocol.ControlMessageType, nodeID string) sentMessage {
t.Helper()
f.mu.Lock()
defer f.mu.Unlock()
for index := len(f.messages) - 1; index >= 0; index-- {
if f.messages[index].messageType == messageType && f.messages[index].nodeID == nodeID {
return f.messages[index]
}
}
t.Fatalf("no %s message for Node %s", messageType, nodeID)
return sentMessage{}
}
func (f *fakeSender) count(messageType protocol.ControlMessageType) int {
f.mu.Lock()
defer f.mu.Unlock()
count := 0
for _, message := range f.messages {
if message.messageType == messageType {
count++
}
}
return count
}
func newManagerTest(t *testing.T, timeout time.Duration) (*Manager, *database.Store, *fakeSender, model.Node, model.Node) {
t.Helper()
db, err := database.Open(context.Background(), filepath.Join(t.TempDir(), "test.db"))
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = db.Close() })
store := database.NewStore(db)
engineer := model.Node{ID: "engineer", Type: model.NodeTypeEngineer, Name: "Engineer", OverlayIP: netip.MustParseAddr("10.88.0.2"), WGPublicKey: "engineer-key", NodeTokenHash: []byte("a")}
site := model.Node{ID: "site", Type: model.NodeTypeSite, Name: "Site", OverlayIP: netip.MustParseAddr("10.88.0.3"), WGPublicKey: "site-key", NodeTokenHash: []byte("b")}
for _, node := range []model.Node{engineer, site} {
if err := store.CreateNode(context.Background(), node); err != nil {
t.Fatal(err)
}
if err := store.UpdateNodeHeartbeat(context.Background(), node.ID, model.NodeOnline, time.Now().UTC(), "1.0", "test"); err != nil {
t.Fatal(err)
}
}
engineer, _ = store.GetNode(context.Background(), engineer.ID)
site, _ = store.GetNode(context.Background(), site.ID)
sender := &fakeSender{}
manager, err := NewManager(store, sender, Config{
OverlayCIDR: netip.MustParsePrefix("10.88.0.0/16"), MTU: 1280, UDPPort: 51821, PrepareTimeout: timeout,
})
if err != nil {
t.Fatal(err)
}
return manager, store, sender, engineer, site
}
func envelope(t *testing.T, messageType protocol.ControlMessageType, requestID string, payload any) protocol.ControlEnvelope {
t.Helper()
envelope, err := protocol.NewControlEnvelope(messageType, requestID, payload)
if err != nil {
t.Fatal(err)
}
return envelope
}