初版功能完成
This commit is contained in:
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)})
|
||||
}
|
||||
@@ -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 }
|
||||
Reference in New Issue
Block a user