初版功能完成
ci / Go checks (ubuntu-latest) (push) Has been cancelled
ci / Go checks (windows-latest) (push) Has been cancelled

This commit is contained in:
qsc
2026-08-29 13:12:17 +08:00
commit 142e5dc7d6
217 changed files with 21313 additions and 0 deletions
+328
View File
@@ -0,0 +1,328 @@
package admin
import (
"context"
"crypto/subtle"
"encoding/json"
"errors"
"fmt"
"io"
"net/http"
"net/netip"
"strconv"
"strings"
"time"
"remlink/internal/database"
"remlink/internal/logging"
"remlink/internal/model"
"remlink/internal/protocol"
)
const maxAdminBody = 1 << 20
type JoinTokenRotator interface {
Rotate(context.Context) (string, error)
}
type HandlerConfig struct {
Store *database.Store
IPAM IPAM
Peers PeerManager
Control ControlNetwork
Sessions SessionControl
Network *NetworkManager
JoinTokens JoinTokenRotator
AdminToken string
}
// Handler exposes only the exact /api/v1/admin surface from R11.
func Handler(config HandlerConfig) (http.Handler, error) {
if config.Store == nil || config.IPAM == nil || config.Peers == nil || config.Control == nil ||
config.Sessions == nil || config.Network == nil || config.JoinTokens == nil {
return nil, errors.New("Admin handler dependencies are required")
}
mux := http.NewServeMux()
mux.HandleFunc("GET /api/v1/admin/nodes", func(writer http.ResponseWriter, request *http.Request) {
nodes, err := config.Store.ListNodes(request.Context())
if err != nil {
writeResult(writer, nil, err)
return
}
views := make([]nodeView, 0, len(nodes))
for _, node := range nodes {
handshake, handshakeErr := config.Peers.LastHandshake(request.Context(), node.WGPublicKey)
if handshakeErr != nil {
writeResult(writer, nil, handshakeErr)
return
}
views = append(views, nodeView{Node: node, WGHandshake: handshake})
}
writeResult(writer, views, nil)
})
mux.HandleFunc("PATCH /api/v1/admin/nodes/{id}", func(writer http.ResponseWriter, request *http.Request) {
var input struct {
Name *string `json:"name"`
OverlayIP *string `json:"overlay_ip"`
}
if err := decodeAdminJSON(writer, request, &input); err != nil {
writeAdminError(writer, http.StatusBadRequest, "INVALID_REQUEST", err)
return
}
if input.Name == nil && input.OverlayIP == nil {
writeAdminError(writer, http.StatusBadRequest, "INVALID_REQUEST", errors.New("name or overlay_ip is required"))
return
}
var desiredName *string
if input.Name != nil {
name := strings.TrimSpace(*input.Name)
if name == "" || len(name) > 128 {
writeAdminError(writer, http.StatusBadRequest, "INVALID_NODE_NAME", errors.New("name must contain 1 to 128 bytes"))
return
}
desiredName = &name
}
var desiredAddress netip.Addr
if input.OverlayIP != nil {
var parseErr error
desiredAddress, parseErr = netip.ParseAddr(*input.OverlayIP)
if parseErr != nil || !desiredAddress.Is4() {
writeAdminError(writer, http.StatusBadRequest, "INVALID_OVERLAY_IP", errors.New("overlay_ip must be IPv4"))
return
}
}
nodeID := request.PathValue("id")
before, err := config.Store.GetNode(request.Context(), nodeID)
if err != nil {
writeAdminError(writer, http.StatusNotFound, "NODE_NOT_FOUND", err)
return
}
addressChanged := input.OverlayIP != nil && desiredAddress != before.OverlayIP
if addressChanged {
if err := config.Sessions.DisconnectNode(request.Context(), nodeID, "NODE_OVERLAY_IP_CHANGED"); err != nil {
writeAdminError(writer, http.StatusConflict, "SESSION_DISCONNECT_FAILED", err)
return
}
if err := config.IPAM.ChangeNodeAddress(request.Context(), nodeID, desiredAddress); err != nil {
writeAdminError(writer, http.StatusConflict, "OVERLAY_IP_UNAVAILABLE", err)
return
}
// The updated address is now visible through the public Bootstrap API.
// Notify over the still-usable old Peer before replacing its /32;
// changing the Peer first would cut the very Control path used by T17.
current := config.Network.Current()
_ = config.Control.Send(request.Context(), nodeID, protocol.ControlRebootstrapRequired,
protocol.RebootstrapRequiredPayload{ConfigVersion: current.ConfigVersion, Reason: "NODE_OVERLAY_IP_CHANGED"})
if err := config.Peers.EnsurePeer(request.Context(), before.WGPublicKey, desiredAddress); err != nil {
_ = config.IPAM.ChangeNodeAddress(request.Context(), nodeID, before.OverlayIP)
_ = config.Peers.EnsurePeer(request.Context(), before.WGPublicKey, before.OverlayIP)
config.Control.ResetNodeConnection(nodeID, "Node Overlay IP update rolled back")
writeAdminError(writer, http.StatusInternalServerError, "PEER_UPDATE_FAILED", err)
return
}
}
if desiredName != nil {
if err := config.Store.UpdateNodeName(request.Context(), nodeID, *desiredName); err != nil {
if addressChanged {
_ = config.IPAM.ChangeNodeAddress(request.Context(), nodeID, before.OverlayIP)
_ = config.Peers.EnsurePeer(request.Context(), before.WGPublicKey, before.OverlayIP)
config.Control.ResetNodeConnection(nodeID, "Node update rolled back")
}
writeAdminError(writer, http.StatusBadRequest, "NODE_UPDATE_FAILED", err)
return
}
}
if addressChanged {
config.Control.ResetNodeConnection(nodeID, "Node Overlay IP changed")
}
updated, err := config.Store.GetNode(request.Context(), nodeID)
if err == nil {
recordAdminEvent(request.Context(), config.Store, nodeID, 0, "节点配置已更新", map[string]any{"overlay_ip": updated.OverlayIP.String(), "name": updated.Name})
}
writeResult(writer, updated, err)
})
mux.HandleFunc("DELETE /api/v1/admin/nodes/{id}", func(writer http.ResponseWriter, request *http.Request) {
nodeID := request.PathValue("id")
node, err := config.Store.GetNode(request.Context(), nodeID)
if err != nil {
writeAdminError(writer, http.StatusNotFound, "NODE_NOT_FOUND", err)
return
}
if err := config.Sessions.DisconnectNode(request.Context(), nodeID, "NODE_REVOKED"); err != nil {
writeAdminError(writer, http.StatusConflict, "SESSION_DISCONNECT_FAILED", err)
return
}
if err := config.Peers.RemovePeer(request.Context(), node.WGPublicKey); err != nil {
writeAdminError(writer, http.StatusInternalServerError, "PEER_REVOKE_FAILED", err)
return
}
if err := config.IPAM.ReleaseNode(request.Context(), nodeID); err != nil {
_ = config.Peers.EnsurePeer(request.Context(), node.WGPublicKey, node.OverlayIP)
writeAdminError(writer, http.StatusInternalServerError, "NODE_DELETE_FAILED", err)
return
}
config.Control.ResetNodeConnection(nodeID, "Node revoked")
recordAdminEvent(request.Context(), config.Store, nodeID, 0, "节点已撤销", nil)
writer.WriteHeader(http.StatusNoContent)
})
mux.HandleFunc("GET /api/v1/admin/sessions", func(writer http.ResponseWriter, request *http.Request) {
sessions, err := config.Store.ListSessions(request.Context())
writeResult(writer, sessions, err)
})
mux.HandleFunc("POST /api/v1/admin/sessions/{id}/disconnect", func(writer http.ResponseWriter, request *http.Request) {
id, err := strconv.ParseUint(request.PathValue("id"), 10, 64)
if err != nil || id == 0 {
writeAdminError(writer, http.StatusBadRequest, "INVALID_SESSION_ID", errors.New("SessionID must be uint64"))
return
}
if err := config.Sessions.Disconnect(request.Context(), id, "ADMIN_DISCONNECT"); err != nil {
writeAdminError(writer, http.StatusConflict, "SESSION_DISCONNECT_FAILED", err)
return
}
recordAdminEvent(request.Context(), config.Store, "", id, "管理员已强制断开会话", nil)
writeAdminJSON(writer, http.StatusOK, map[string]any{"session_id": id, "status": model.SessionClosed})
})
mux.HandleFunc("GET /api/v1/admin/network", func(writer http.ResponseWriter, _ *http.Request) {
writeAdminJSON(writer, http.StatusOK, config.Network.View())
})
mux.HandleFunc("PUT /api/v1/admin/network", func(writer http.ResponseWriter, request *http.Request) {
var input NetworkUpdate
if err := decodeAdminJSON(writer, request, &input); err != nil {
writeAdminError(writer, http.StatusBadRequest, "INVALID_REQUEST", err)
return
}
current := config.Network.Current()
updated := current
var err error
if !networkInputMatches(input, current) {
updated, err = config.Network.Update(request.Context(), input)
if err != nil {
writeAdminError(writer, http.StatusConflict, "NETWORK_UPDATE_FAILED", err)
return
}
}
response := struct {
Network
JoinToken string `json:"join_token,omitempty"`
}{Network: updated}
if input.RotateJoinToken {
response.JoinToken, err = config.JoinTokens.Rotate(request.Context())
if err != nil {
writeAdminError(writer, http.StatusInternalServerError, "JOIN_TOKEN_ROTATE_FAILED", err)
return
}
}
recordAdminEvent(request.Context(), config.Store, "", 0, "网络配置已更新", map[string]any{"config_version": updated.ConfigVersion, "join_token_rotated": input.RotateJoinToken})
writeAdminJSON(writer, http.StatusOK, response)
})
mux.HandleFunc("GET /api/v1/admin/logs", func(writer http.ResponseWriter, request *http.Request) {
filter := model.EventLogFilter{
Level: request.URL.Query().Get("level"), Module: request.URL.Query().Get("module"),
NodeID: request.URL.Query().Get("node_id"),
}
if raw := request.URL.Query().Get("session_id"); raw != "" {
var parseErr error
filter.SessionID, parseErr = strconv.ParseUint(raw, 10, 64)
if parseErr != nil || filter.SessionID == 0 {
writeAdminError(writer, http.StatusBadRequest, "INVALID_SESSION_ID", errors.New("session_id must be uint64"))
return
}
}
if raw := request.URL.Query().Get("limit"); raw != "" {
var parseErr error
filter.Limit, parseErr = strconv.Atoi(raw)
if parseErr != nil || filter.Limit < 1 || filter.Limit > 1000 {
writeAdminError(writer, http.StatusBadRequest, "INVALID_LIMIT", errors.New("limit must be between 1 and 1000"))
return
}
}
for name, destination := range map[string]*time.Time{"from": &filter.From, "to": &filter.To} {
if raw := request.URL.Query().Get(name); raw != "" {
parsed, parseErr := time.Parse(time.RFC3339, raw)
if parseErr != nil {
writeAdminError(writer, http.StatusBadRequest, "INVALID_TIME", fmt.Errorf("%s must be RFC3339", name))
return
}
*destination = parsed
}
}
if !filter.From.IsZero() && !filter.To.IsZero() && filter.From.After(filter.To) {
writeAdminError(writer, http.StatusBadRequest, "INVALID_TIME_RANGE", errors.New("from must not be after to"))
return
}
events, err := config.Store.ListEvents(request.Context(), filter)
writeResult(writer, events, err)
})
return adminSecurity(config.AdminToken, mux), nil
}
type nodeView struct {
model.Node
WGHandshake *time.Time `json:"wg_handshake,omitempty"`
}
func networkInputMatches(input NetworkUpdate, current Network) bool {
return input.OverlayCIDR == current.OverlayCIDR && input.ServerOverlayIP == current.ServerOverlayIP &&
input.WireGuardPort == current.WireGuardPort && input.SessionUDPPort == current.SessionUDPPort && input.MTU == current.MTU
}
func recordAdminEvent(ctx context.Context, store *database.Store, nodeID string, sessionID uint64, message string, fields map[string]any) {
if fields == nil {
fields = map[string]any{}
}
raw, _ := json.Marshal(fields)
_ = store.AppendEvent(ctx, model.EventLog{
Level: "INFO", Module: string(logging.ModuleCore), NodeID: nodeID, SessionID: sessionID,
Message: message, FieldsJSON: raw,
})
}
func adminSecurity(token string, next http.Handler) http.Handler {
return http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
writer.Header().Set("Cache-Control", "no-store")
writer.Header().Set("X-Content-Type-Options", "nosniff")
writer.Header().Set("X-Frame-Options", "DENY")
if token != "" {
provided := strings.TrimPrefix(request.Header.Get("Authorization"), "Bearer ")
if len(provided) != len(token) || subtle.ConstantTimeCompare([]byte(provided), []byte(token)) != 1 {
writeAdminError(writer, http.StatusUnauthorized, "ADMIN_AUTH_FAILED", errors.New("valid Bearer Admin Token required"))
return
}
}
next.ServeHTTP(writer, request)
})
}
func decodeAdminJSON(writer http.ResponseWriter, request *http.Request, destination any) error {
if contentType := request.Header.Get("Content-Type"); contentType != "" && !strings.HasPrefix(strings.ToLower(contentType), "application/json") {
return errors.New("Content-Type must be application/json")
}
request.Body = http.MaxBytesReader(writer, request.Body, maxAdminBody)
decoder := json.NewDecoder(request.Body)
decoder.DisallowUnknownFields()
if err := decoder.Decode(destination); err != nil {
return err
}
if err := decoder.Decode(&struct{}{}); !errors.Is(err, io.EOF) {
return errors.New("request body must contain one JSON object")
}
return nil
}
func writeResult(writer http.ResponseWriter, value any, err error) {
if err != nil {
writeAdminError(writer, http.StatusInternalServerError, "ADMIN_OPERATION_FAILED", err)
return
}
writeAdminJSON(writer, http.StatusOK, value)
}
func writeAdminError(writer http.ResponseWriter, status int, code string, err error) {
writeAdminJSON(writer, status, map[string]any{"error": map[string]string{"code": code, "message": err.Error()}})
}
func writeAdminJSON(writer http.ResponseWriter, status int, value any) {
writer.Header().Set("Content-Type", "application/json; charset=utf-8")
writer.WriteHeader(status)
_ = json.NewEncoder(writer).Encode(value)
}
+183
View File
@@ -0,0 +1,183 @@
package admin
import (
"bytes"
"context"
"encoding/json"
"net/http"
"net/http/httptest"
"net/netip"
"path/filepath"
"testing"
"time"
"remlink/internal/bootstrap"
"remlink/internal/database"
"remlink/internal/ipam"
"remlink/internal/model"
)
type fakeJoinTokenRotator struct{ count int }
func (f *fakeJoinTokenRotator) Rotate(context.Context) (string, error) {
f.count++
return "rotated-token", nil
}
func TestAdminHandlerAuthNodeUpdateLogsAndRotateOnly(t *testing.T) {
ctx := context.Background()
db, err := database.Open(ctx, filepath.Join(t.TempDir(), "admin.db"))
if err != nil {
t.Fatal(err)
}
defer db.Close()
store := database.NewStore(db)
if err := store.CreateNode(ctx, model.Node{
ID: "engineer", Type: model.NodeTypeEngineer, Name: "Engineer",
OverlayIP: netip.MustParseAddr("10.88.0.2"), WGPublicKey: "wg-key", NodeTokenHash: []byte("hash"),
}); err != nil {
t.Fatal(err)
}
ipamManager, err := ipam.New(store, netip.MustParsePrefix("10.88.0.0/24"), netip.MustParseAddr("10.88.0.1"))
if err != nil {
t.Fatal(err)
}
peers := &fakeAdminPeers{}
bootstrapNetwork := &fakeBootstrapNetwork{config: bootstrap.ServiceConfig{
WGEndpoint: "203.0.113.4:51820", ControlURL: "ws://10.88.0.1:7001/control",
OverlayCIDR: netip.MustParsePrefix("10.88.0.0/24"), ServerOverlayIP: netip.MustParseAddr("10.88.0.1"),
SessionUDPPort: 6200, MTU: 1280, ConfigVersion: 1,
}}
controlNetwork := &fakeAdminControl{}
sessions := &fakeSessionControl{}
network, err := NewNetworkManager(store, ipamManager, peers, bootstrapNetwork, controlNetwork, sessions, Network{
OverlayCIDR: "10.88.0.0/24", ServerOverlayIP: "10.88.0.1", WireGuardPort: 51820,
SessionUDPPort: 6200, MTU: 1280, ConfigVersion: 1,
}, nil)
if err != nil {
t.Fatal(err)
}
rotator := &fakeJoinTokenRotator{}
handler, err := Handler(HandlerConfig{
Store: store, IPAM: ipamManager, Peers: peers, Control: controlNetwork,
Sessions: sessions, Network: network, JoinTokens: rotator, AdminToken: "secret",
})
if err != nil {
t.Fatal(err)
}
unauthorized := httptest.NewRecorder()
handler.ServeHTTP(unauthorized, httptest.NewRequest(http.MethodGet, "/api/v1/admin/nodes", nil))
if unauthorized.Code != http.StatusUnauthorized {
t.Fatalf("unauthorized status = %d", unauthorized.Code)
}
nodesResponse := adminRequest(t, handler, http.MethodGet, "/api/v1/admin/nodes", nil)
if nodesResponse.Code != http.StatusOK || !bytes.Contains(nodesResponse.Body.Bytes(), []byte(`"wg_handshake"`)) {
t.Fatalf("nodes response status=%d body=%s", nodesResponse.Code, nodesResponse.Body.String())
}
patch := adminRequest(t, handler, http.MethodPatch, "/api/v1/admin/nodes/engineer", map[string]any{"name": "Field Engineer"})
if patch.Code != http.StatusOK || bytes.Contains(patch.Body.Bytes(), []byte("node_token")) {
t.Fatalf("PATCH response status=%d body=%s", patch.Code, patch.Body.String())
}
updated, err := store.GetNode(ctx, "engineer")
if err != nil || updated.Name != "Field Engineer" {
t.Fatalf("updated node = %+v, %v", updated, err)
}
invalidCombined := adminRequest(t, handler, http.MethodPatch, "/api/v1/admin/nodes/engineer", map[string]any{
"name": "Must Not Persist", "overlay_ip": "not-an-ip",
})
unchanged, _ := store.GetNode(ctx, "engineer")
if invalidCombined.Code != http.StatusBadRequest || unchanged.Name != "Field Engineer" {
t.Fatalf("invalid combined PATCH status=%d node=%+v", invalidCombined.Code, unchanged)
}
steps := make([]string, 0, 2)
peers.steps = &steps
controlNetwork.steps = &steps
addressPatch := adminRequest(t, handler, http.MethodPatch, "/api/v1/admin/nodes/engineer", map[string]any{
"name": "Moved Engineer", "overlay_ip": "10.88.0.9",
})
moved, moveErr := store.GetNode(ctx, "engineer")
if addressPatch.Code != http.StatusOK || moveErr != nil || moved.Name != "Moved Engineer" || moved.OverlayIP.String() != "10.88.0.9" {
t.Fatalf("address PATCH status=%d node=%+v error=%v body=%s", addressPatch.Code, moved, moveErr, addressPatch.Body.String())
}
wantSteps := []string{"notify", "peers"}
if len(steps) != len(wantSteps) || steps[0] != wantSteps[0] || steps[1] != wantSteps[1] {
t.Fatalf("Node address change steps=%v, want=%v", steps, wantSteps)
}
if peers.ensured != moved.OverlayIP || !controlNetwork.resetNode || sessions.nodeDisconnects != 1 || sessions.nodeReason != "NODE_OVERLAY_IP_CHANGED" {
t.Fatalf("Node address orchestration peer=%s reset=%v disconnects=%d reason=%s", peers.ensured, controlNetwork.resetNode, sessions.nodeDisconnects, sessions.nodeReason)
}
peers.steps = nil
controlNetwork.steps = nil
bad := adminRequest(t, handler, http.MethodPatch, "/api/v1/admin/nodes/engineer", map[string]any{"unknown": true})
if bad.Code != http.StatusBadRequest {
t.Fatalf("unknown JSON field status = %d, want 400", bad.Code)
}
rotate := adminRequest(t, handler, http.MethodPut, "/api/v1/admin/network", map[string]any{
"overlay_cidr": "10.88.0.0/24", "server_overlay_ip": "10.88.0.1",
"wireguard_port": 51820, "session_udp_port": 6200, "mtu": 1280, "rotate_join_token": true,
})
if rotate.Code != http.StatusOK || rotator.count != 1 || sessions.all || controlNetwork.resetAll {
t.Fatalf("rotate-only response=%s count=%d migrated=%v reset=%v", rotate.Body.String(), rotator.count, sessions.all, controlNetwork.resetAll)
}
var rotateBody map[string]any
_ = json.Unmarshal(rotate.Body.Bytes(), &rotateBody)
if rotateBody["join_token"] != "rotated-token" || rotateBody["config_version"] != float64(1) {
t.Fatalf("rotate response = %v", rotateBody)
}
from := time.Now().Add(-time.Hour).UTC().Format(time.RFC3339)
to := time.Now().Add(time.Hour).UTC().Format(time.RFC3339)
logs := adminRequest(t, handler, http.MethodGet, "/api/v1/admin/logs?module=CORE&limit=10&from="+from+"&to="+to, nil)
if logs.Code != http.StatusOK {
t.Fatalf("logs status=%d body=%s", logs.Code, logs.Body.String())
}
var events []model.EventLog
if err := json.Unmarshal(logs.Body.Bytes(), &events); err != nil || len(events) < 2 {
t.Fatalf("Admin events = %+v, %v", events, err)
}
badTime := adminRequest(t, handler, http.MethodGet, "/api/v1/admin/logs?from=not-a-time", nil)
if badTime.Code != http.StatusBadRequest {
t.Fatalf("invalid log time status = %d", badTime.Code)
}
view := adminRequest(t, handler, http.MethodGet, "/api/v1/admin/network", nil)
var networkView map[string]any
_ = json.Unmarshal(view.Body.Bytes(), &networkView)
if _, exists := networkView["uptime_seconds"]; !exists {
t.Fatalf("network response has no uptime_seconds: %s", view.Body.String())
}
deleted := adminRequest(t, handler, http.MethodDelete, "/api/v1/admin/nodes/engineer", nil)
if deleted.Code != http.StatusNoContent || peers.removed != 1 || sessions.nodeDisconnects != 2 || controlNetwork.resetNodeCount != 2 || controlNetwork.resetReason != "Node revoked" {
t.Fatalf("DELETE status=%d peer removals=%d disconnects=%d resets=%d reason=%s",
deleted.Code, peers.removed, sessions.nodeDisconnects, controlNetwork.resetNodeCount, controlNetwork.resetReason)
}
if _, err := store.GetNode(ctx, "engineer"); err == nil {
t.Fatal("deleted Node remains in the registry")
}
}
func adminRequest(t *testing.T, handler http.Handler, method, path string, body any) *httptest.ResponseRecorder {
t.Helper()
var raw []byte
if body != nil {
var err error
raw, err = json.Marshal(body)
if err != nil {
t.Fatal(err)
}
}
request := httptest.NewRequest(method, path, bytes.NewReader(raw))
request.Header.Set("Authorization", "Bearer secret")
if body != nil {
request.Header.Set("Content-Type", "application/json")
}
recorder := httptest.NewRecorder()
handler.ServeHTTP(recorder, request)
return recorder
}
+301
View File
@@ -0,0 +1,301 @@
// Package admin implements the v1 Server Admin API and network migration.
package admin
import (
"context"
"encoding/json"
"errors"
"fmt"
"net"
"net/netip"
"net/url"
"strconv"
"sync"
"time"
"remlink/internal/bootstrap"
"remlink/internal/database"
"remlink/internal/ipam"
"remlink/internal/overlay/serverwg"
"remlink/internal/protocol"
)
const networkSettingKey = "admin.network"
type Network struct {
OverlayCIDR string `json:"overlay_cidr"`
ServerOverlayIP string `json:"server_overlay_ip"`
WireGuardPort int `json:"wireguard_port"`
SessionUDPPort int `json:"session_udp_port"`
MTU int `json:"mtu"`
ConfigVersion uint64 `json:"config_version"`
}
type NetworkUpdate struct {
OverlayCIDR string `json:"overlay_cidr"`
ServerOverlayIP string `json:"server_overlay_ip"`
WireGuardPort int `json:"wireguard_port"`
SessionUDPPort int `json:"session_udp_port"`
MTU int `json:"mtu"`
RotateJoinToken bool `json:"rotate_join_token,omitempty"`
}
type NetworkView struct {
Network
UptimeSeconds int64 `json:"uptime_seconds"`
}
type IPAM interface {
ChangeNodeAddress(context.Context, string, netip.Addr) error
ReleaseNode(context.Context, string) error
Reconfigure(netip.Prefix, netip.Addr) error
}
type PeerManager interface {
EnsurePeer(context.Context, string, netip.Addr) error
RemovePeer(context.Context, string) error
Reconfigure(context.Context, netip.Prefix, int, []serverwg.Peer) error
LastHandshake(context.Context, string) (*time.Time, error)
}
type BootstrapNetwork interface {
NetworkSnapshot() bootstrap.ServiceConfig
UpdateNetwork(bootstrap.ServiceConfig) error
}
type ControlNetwork interface {
Send(context.Context, string, protocol.ControlMessageType, any) error
SetNetworkConfigVersion(uint64) error
ResetNodeConnection(string, string)
ResetConnections(string)
}
type SessionControl interface {
Disconnect(context.Context, uint64, string) error
DisconnectAll(context.Context, string) error
DisconnectNode(context.Context, string, string) error
BeginNetworkMigration(context.Context, string) error
EndNetworkMigration()
ReconfigureNetwork(netip.Prefix, int, int) error
}
type NetworkManager struct {
mu sync.Mutex
store *database.Store
ipam IPAM
peers PeerManager
bootstrap BootstrapNetwork
control ControlNetwork
sessions SessionControl
current Network
startedAt time.Time
rebindControl func(netip.Addr) error
}
func LoadStoredNetwork(ctx context.Context, store *database.Store, fallback Network) (Network, error) {
value, found, err := database.GetSetting(ctx, store.DB(), networkSettingKey)
if err != nil || !found {
return fallback, err
}
var network Network
if err := json.Unmarshal([]byte(value), &network); err != nil {
return Network{}, fmt.Errorf("decode stored Admin network: %w", err)
}
if _, _, err := validateNetwork(network); err != nil {
return Network{}, fmt.Errorf("validate stored Admin network: %w", err)
}
return network, nil
}
func NewNetworkManager(store *database.Store, ipamManager IPAM, peers PeerManager, bootstrapService BootstrapNetwork,
control ControlNetwork, sessions SessionControl, initial Network, rebindControl func(netip.Addr) error) (*NetworkManager, error) {
if store == nil || ipamManager == nil || peers == nil || bootstrapService == nil || control == nil || sessions == nil {
return nil, errors.New("Admin NetworkManager dependencies are required")
}
if _, _, err := validateNetwork(initial); err != nil {
return nil, err
}
if rebindControl == nil {
rebindControl = func(netip.Addr) error { return nil }
}
return &NetworkManager{
store: store, ipam: ipamManager, peers: peers, bootstrap: bootstrapService,
control: control, sessions: sessions, current: initial, startedAt: time.Now(), rebindControl: rebindControl,
}, nil
}
func (m *NetworkManager) View() NetworkView {
m.mu.Lock()
defer m.mu.Unlock()
return NetworkView{Network: m.current, UptimeSeconds: int64(time.Since(m.startedAt).Seconds())}
}
func (m *NetworkManager) Current() Network {
m.mu.Lock()
defer m.mu.Unlock()
return m.current
}
func (m *NetworkManager) Update(ctx context.Context, input NetworkUpdate) (Network, error) {
m.mu.Lock()
defer m.mu.Unlock()
desired := Network{
OverlayCIDR: input.OverlayCIDR, ServerOverlayIP: input.ServerOverlayIP,
WireGuardPort: input.WireGuardPort, SessionUDPPort: input.SessionUDPPort, MTU: input.MTU,
ConfigVersion: m.current.ConfigVersion + 1,
}
prefix, serverIP, err := validateNetwork(desired)
if err != nil {
return Network{}, err
}
nodes, err := m.store.ListNodes(ctx)
if err != nil {
return Network{}, err
}
assignments, err := ipam.PlanNodeAddresses(nodes, prefix, serverIP)
if err != nil {
return Network{}, err
}
if err := m.sessions.BeginNetworkMigration(ctx, "NETWORK_CONFIG_CHANGED"); err != nil {
return Network{}, err
}
defer m.sessions.EndNetworkMigration()
old := m.current
oldPrefix, oldServerIP, _ := validateNetwork(old)
oldAssignments := make(map[string]netip.Addr, len(nodes))
oldPeers := make([]serverwg.Peer, 0, len(nodes))
newPeers := make([]serverwg.Peer, 0, len(nodes))
for _, node := range nodes {
oldAssignments[node.ID] = node.OverlayIP
oldPeers = append(oldPeers, serverwg.Peer{PublicKey: node.WGPublicKey, Address: node.OverlayIP})
newPeers = append(newPeers, serverwg.Peer{PublicKey: node.WGPublicKey, Address: assignments[node.ID]})
}
serviceConfig := m.bootstrap.NetworkSnapshot()
oldServiceConfig := serviceConfig
serviceConfig.OverlayCIDR = prefix
serviceConfig.ServerOverlayIP = serverIP
serviceConfig.WGEndpoint, err = replaceEndpointPort(serviceConfig.WGEndpoint, desired.WireGuardPort)
if err != nil {
return Network{}, err
}
serviceConfig.ControlURL, err = replaceControlHost(serviceConfig.ControlURL, serverIP)
if err != nil {
return Network{}, err
}
serviceConfig.SessionUDPPort = desired.SessionUDPPort
serviceConfig.MTU = desired.MTU
serviceConfig.ConfigVersion = desired.ConfigVersion
oldEncoded, _ := json.Marshal(old)
desiredEncoded, _ := json.Marshal(desired)
if err := m.sessions.ReconfigureNetwork(prefix, desired.MTU, desired.SessionUDPPort); err != nil {
return Network{}, err
}
rollbackSessions := func() { _ = m.sessions.ReconfigureNetwork(oldPrefix, old.MTU, old.SessionUDPPort) }
if err := m.store.ReplaceNodeOverlayIPsAndSetting(ctx, assignments, networkSettingKey, string(desiredEncoded)); err != nil {
rollbackSessions()
return Network{}, err
}
rollbackDatabase := func() {
_ = m.store.ReplaceNodeOverlayIPsAndSetting(context.Background(), oldAssignments, networkSettingKey, string(oldEncoded))
}
if err := m.ipam.Reconfigure(prefix, serverIP); err != nil {
rollbackDatabase()
rollbackSessions()
return Network{}, err
}
if err := m.bootstrap.UpdateNetwork(serviceConfig); err != nil {
_ = m.ipam.Reconfigure(oldPrefix, oldServerIP)
rollbackDatabase()
rollbackSessions()
return Network{}, err
}
if err := m.control.SetNetworkConfigVersion(desired.ConfigVersion); err != nil {
_ = m.bootstrap.UpdateNetwork(oldServiceConfig)
_ = m.ipam.Reconfigure(oldPrefix, oldServerIP)
rollbackDatabase()
rollbackSessions()
return Network{}, err
}
rollbackPublishedConfig := func() {
_ = m.control.SetNetworkConfigVersion(old.ConfigVersion)
_ = m.bootstrap.UpdateNetwork(oldServiceConfig)
_ = m.ipam.Reconfigure(oldPrefix, oldServerIP)
rollbackDatabase()
rollbackSessions()
}
// Publish the authoritative Bootstrap snapshot before notifying over the
// still-live old Control path. Switching wg0 peers or the listener first
// would make REBOOTSTRAP_REQUIRED physically undeliverable.
for _, node := range nodes {
_ = m.control.Send(ctx, node.ID, protocol.ControlRebootstrapRequired, protocol.RebootstrapRequiredPayload{
ConfigVersion: desired.ConfigVersion, Reason: "NETWORK_CONFIG_CHANGED",
})
}
if err := m.peers.Reconfigure(ctx, netip.PrefixFrom(serverIP, prefix.Bits()), desired.WireGuardPort, newPeers); err != nil {
rollbackPublishedConfig()
return Network{}, err
}
rollbackKernel := func() {
_ = m.peers.Reconfigure(context.Background(), netip.PrefixFrom(oldServerIP, oldPrefix.Bits()), old.WireGuardPort, oldPeers)
}
if serverIP != oldServerIP {
if err := m.rebindControl(serverIP); err != nil {
rollbackKernel()
rollbackPublishedConfig()
return Network{}, fmt.Errorf("rebind Control listener after migration notification: %w", err)
}
}
m.control.ResetConnections("Network configuration changed")
m.current = desired
return desired, nil
}
func validateNetwork(network Network) (netip.Prefix, netip.Addr, error) {
prefix, err := netip.ParsePrefix(network.OverlayCIDR)
if err != nil || !prefix.Addr().Is4() || prefix != prefix.Masked() || prefix.Bits() == 0 || prefix.Bits() > 30 {
return netip.Prefix{}, netip.Addr{}, errors.New("Overlay CIDR must be canonical IPv4 with usable hosts")
}
serverIP, err := netip.ParseAddr(network.ServerOverlayIP)
if err != nil || !serverIP.Is4() || !prefix.Contains(serverIP) || serverIP == prefix.Addr() || serverIP == lastIPv4(prefix) {
return netip.Prefix{}, netip.Addr{}, errors.New("Server Overlay IP must be a usable address in Overlay CIDR")
}
for name, port := range map[string]int{"WireGuard port": network.WireGuardPort, "Session UDP port": network.SessionUDPPort} {
if port < 1 || port > 65535 {
return netip.Prefix{}, netip.Addr{}, fmt.Errorf("%s must be between 1 and 65535", name)
}
}
if network.MTU < 576 || network.MTU > 65535 || network.ConfigVersion == 0 {
return netip.Prefix{}, netip.Addr{}, errors.New("MTU or config version is invalid")
}
return prefix, serverIP, nil
}
func replaceEndpointPort(endpoint string, port int) (string, error) {
host, _, err := net.SplitHostPort(endpoint)
if err != nil {
return "", err
}
return net.JoinHostPort(host, strconv.Itoa(port)), nil
}
func replaceControlHost(raw string, host netip.Addr) (string, error) {
parsed, err := url.Parse(raw)
if err != nil {
return "", err
}
_, port, err := net.SplitHostPort(parsed.Host)
if err != nil {
return "", err
}
parsed.Host = net.JoinHostPort(host.String(), port)
return parsed.String(), nil
}
func lastIPv4(prefix netip.Prefix) netip.Addr {
bytes := prefix.Masked().Addr().As4()
value := uint32(bytes[0])<<24 | uint32(bytes[1])<<16 | uint32(bytes[2])<<8 | uint32(bytes[3])
value |= ^uint32(0) >> prefix.Bits()
return netip.AddrFrom4([4]byte{byte(value >> 24), byte(value >> 16), byte(value >> 8), byte(value)})
}
+243
View File
@@ -0,0 +1,243 @@
package admin
import (
"context"
"errors"
"net/netip"
"path/filepath"
"testing"
"time"
"remlink/internal/bootstrap"
"remlink/internal/database"
"remlink/internal/ipam"
"remlink/internal/model"
"remlink/internal/overlay/serverwg"
"remlink/internal/protocol"
)
func TestNetworkManagerRunsSevenStepMigration(t *testing.T) {
ctx := context.Background()
db, err := database.Open(ctx, filepath.Join(t.TempDir(), "test.db"))
if err != nil {
t.Fatal(err)
}
defer db.Close()
store := database.NewStore(db)
for _, node := range []model.Node{
{ID: "engineer", Type: model.NodeTypeEngineer, Name: "Engineer", OverlayIP: netip.MustParseAddr("10.88.0.2"), WGPublicKey: "key-a", NodeTokenHash: []byte("a")},
{ID: "site", Type: model.NodeTypeSite, Name: "Site", OverlayIP: netip.MustParseAddr("10.88.0.3"), WGPublicKey: "key-b", NodeTokenHash: []byte("b")},
} {
if err := store.CreateNode(ctx, node); err != nil {
t.Fatal(err)
}
}
ipamManager, _ := ipam.New(store, netip.MustParsePrefix("10.88.0.0/24"), netip.MustParseAddr("10.88.0.1"))
steps := make([]string, 0, 4)
peers := &fakeAdminPeers{steps: &steps}
bootstrapService := &fakeBootstrapNetwork{config: bootstrap.ServiceConfig{
WGEndpoint: "203.0.113.4:51820", ControlURL: "ws://10.88.0.1:7001/control",
OverlayCIDR: netip.MustParsePrefix("10.88.0.0/24"), ServerOverlayIP: netip.MustParseAddr("10.88.0.1"), ConfigVersion: 1,
}}
control := &fakeAdminControl{steps: &steps}
sessions := &fakeSessionControl{}
var rebound netip.Addr
manager, err := NewNetworkManager(store, ipamManager, peers, bootstrapService, control, sessions, Network{
OverlayCIDR: "10.88.0.0/24", ServerOverlayIP: "10.88.0.1", WireGuardPort: 51820,
SessionUDPPort: 6200, MTU: 1280, ConfigVersion: 1,
}, func(address netip.Addr) error { steps = append(steps, "rebind"); rebound = address; return nil })
if err != nil {
t.Fatal(err)
}
updated, err := manager.Update(ctx, NetworkUpdate{
OverlayCIDR: "10.99.0.0/24", ServerOverlayIP: "10.99.0.1", WireGuardPort: 51830,
SessionUDPPort: 6300, MTU: 1400,
})
if err != nil {
t.Fatal(err)
}
if updated.ConfigVersion != 2 || !sessions.began || !sessions.ended || peers.address.String() != "10.99.0.1/24" || peers.port != 51830 {
t.Fatalf("migration state updated=%+v sessions=%+v peers=%s:%d", updated, sessions, peers.address, peers.port)
}
if sessions.prefix.String() != "10.99.0.0/24" || sessions.mtu != 1400 || sessions.udpPort != 6300 {
t.Fatalf("Session Manager network remained stale: prefix=%s mtu=%d udp=%d", sessions.prefix, sessions.mtu, sessions.udpPort)
}
engineer, _ := store.GetNode(ctx, "engineer")
site, _ := store.GetNode(ctx, "site")
if engineer.OverlayIP.String() != "10.99.0.2" || site.OverlayIP.String() != "10.99.0.3" {
t.Fatalf("migrated addresses Engineer=%s Site=%s", engineer.OverlayIP, site.OverlayIP)
}
if rebound.String() != "10.99.0.1" || control.version != 2 || control.notifications != 2 || !control.resetAll {
t.Fatalf("Control migration rebound=%s version=%d notifications=%d reset=%v", rebound, control.version, control.notifications, control.resetAll)
}
if bootstrapService.config.ControlURL != "ws://10.99.0.1:7001/control" || bootstrapService.config.WGEndpoint != "203.0.113.4:51830" {
t.Fatalf("Bootstrap config = %+v", bootstrapService.config)
}
wantSteps := []string{"notify", "notify", "peers", "rebind"}
if len(steps) != len(wantSteps) {
t.Fatalf("migration steps = %v, want %v", steps, wantSteps)
}
for index := range wantSteps {
if steps[index] != wantSteps[index] {
t.Fatalf("migration steps = %v, want %v", steps, wantSteps)
}
}
loaded, err := LoadStoredNetwork(ctx, store, Network{})
if err != nil || loaded != updated {
t.Fatalf("stored network = %+v, %v", loaded, err)
}
}
func TestNetworkManagerRollsBackKernelWhenControlRebindFails(t *testing.T) {
ctx := context.Background()
db, err := database.Open(ctx, filepath.Join(t.TempDir(), "test.db"))
if err != nil {
t.Fatal(err)
}
defer db.Close()
store := database.NewStore(db)
ipamManager, _ := ipam.New(store, netip.MustParsePrefix("10.88.0.0/24"), netip.MustParseAddr("10.88.0.1"))
peers := &fakeAdminPeers{}
bootstrapService := &fakeBootstrapNetwork{config: bootstrap.ServiceConfig{
WGEndpoint: "203.0.113.4:51820", ControlURL: "ws://10.88.0.1:7001/control",
OverlayCIDR: netip.MustParsePrefix("10.88.0.0/24"), ServerOverlayIP: netip.MustParseAddr("10.88.0.1"), ConfigVersion: 1,
}}
sessions := &fakeSessionControl{}
manager, err := NewNetworkManager(store, ipamManager, peers, bootstrapService, &fakeAdminControl{}, sessions, Network{
OverlayCIDR: "10.88.0.0/24", ServerOverlayIP: "10.88.0.1", WireGuardPort: 51820,
SessionUDPPort: 6200, MTU: 1280, ConfigVersion: 1,
}, func(netip.Addr) error { return errors.New("address unavailable") })
if err != nil {
t.Fatal(err)
}
if _, err := manager.Update(ctx, NetworkUpdate{
OverlayCIDR: "10.99.0.0/24", ServerOverlayIP: "10.99.0.1", WireGuardPort: 51830,
SessionUDPPort: 6300, MTU: 1400,
}); err == nil {
t.Fatal("network migration unexpectedly succeeded despite Control rebind failure")
}
if current := manager.Current(); current.ConfigVersion != 1 || current.OverlayCIDR != "10.88.0.0/24" {
t.Fatalf("failed migration changed current config: %+v", current)
}
if peers.address.String() != "10.88.0.1/24" || peers.port != 51820 {
t.Fatalf("failed migration left kernel config at %s:%d", peers.address, peers.port)
}
if sessions.prefix.String() != "10.88.0.0/24" || sessions.mtu != 1280 || sessions.udpPort != 6200 {
t.Fatalf("failed migration did not restore Session Manager: %s mtu=%d udp=%d", sessions.prefix, sessions.mtu, sessions.udpPort)
}
if bootstrapService.config.OverlayCIDR.String() != "10.88.0.0/24" || bootstrapService.config.ConfigVersion != 1 {
t.Fatalf("failed migration did not restore Bootstrap config: %+v", bootstrapService.config)
}
stored, err := LoadStoredNetwork(ctx, store, Network{})
if err != nil || stored != manager.Current() {
t.Fatalf("failed migration did not restore atomic database setting: stored=%+v current=%+v err=%v", stored, manager.Current(), err)
}
}
func TestAdminNetworkRejectsExitNodeOverlay(t *testing.T) {
_, _, err := validateNetwork(Network{
OverlayCIDR: "0.0.0.0/0", ServerOverlayIP: "10.88.0.1",
WireGuardPort: 51820, SessionUDPPort: 6200, MTU: 1280, ConfigVersion: 1,
})
if err == nil {
t.Fatal("Admin network accepted 0.0.0.0/0 Exit Node Overlay")
}
}
type fakeAdminPeers struct {
address netip.Prefix
port int
peers []serverwg.Peer
steps *[]string
ensured netip.Addr
removed int
}
func (f *fakeAdminPeers) EnsurePeer(_ context.Context, _ string, address netip.Addr) error {
if f.steps != nil {
*f.steps = append(*f.steps, "peers")
}
f.ensured = address
return nil
}
func (f *fakeAdminPeers) RemovePeer(context.Context, string) error {
f.removed++
return nil
}
func (*fakeAdminPeers) LastHandshake(context.Context, string) (*time.Time, error) {
value := time.Date(2026, 8, 25, 1, 2, 3, 0, time.UTC)
return &value, nil
}
func (f *fakeAdminPeers) Reconfigure(_ context.Context, address netip.Prefix, port int, peers []serverwg.Peer) error {
if f.steps != nil {
*f.steps = append(*f.steps, "peers")
}
f.address, f.port, f.peers = address, port, peers
return nil
}
type fakeBootstrapNetwork struct{ config bootstrap.ServiceConfig }
func (f *fakeBootstrapNetwork) NetworkSnapshot() bootstrap.ServiceConfig { return f.config }
func (f *fakeBootstrapNetwork) UpdateNetwork(config bootstrap.ServiceConfig) error {
f.config = config
return nil
}
type fakeAdminControl struct {
version uint64
notifications int
resetAll bool
steps *[]string
resetNode bool
resetNodeCount int
resetReason string
}
func (f *fakeAdminControl) Send(_ context.Context, _ string, messageType protocol.ControlMessageType, _ any) error {
if messageType == protocol.ControlRebootstrapRequired {
f.notifications++
if f.steps != nil {
*f.steps = append(*f.steps, "notify")
}
}
return nil
}
func (f *fakeAdminControl) SetNetworkConfigVersion(version uint64) error {
f.version = version
return nil
}
func (f *fakeAdminControl) ResetNodeConnection(_, reason string) {
f.resetNode, f.resetReason = true, reason
f.resetNodeCount++
}
func (f *fakeAdminControl) ResetConnections(string) { f.resetAll = true }
type fakeSessionControl struct {
all bool
began bool
ended bool
prefix netip.Prefix
mtu int
udpPort int
nodeDisconnects int
nodeReason string
}
func (*fakeSessionControl) Disconnect(context.Context, uint64, string) error { return nil }
func (f *fakeSessionControl) DisconnectAll(context.Context, string) error { f.all = true; return nil }
func (f *fakeSessionControl) DisconnectNode(_ context.Context, _, reason string) error {
f.nodeDisconnects++
f.nodeReason = reason
return nil
}
func (f *fakeSessionControl) ReconfigureNetwork(prefix netip.Prefix, mtu, udpPort int) error {
f.prefix, f.mtu, f.udpPort = prefix, mtu, udpPort
return nil
}
func (f *fakeSessionControl) BeginNetworkMigration(context.Context, string) error {
f.began = true
return nil
}
func (f *fakeSessionControl) EndNetworkMigration() { f.ended = true }