Files
RemLink/cmd/server/main.go
T
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

277 lines
9.6 KiB
Go

package main
import (
"context"
"errors"
"flag"
"fmt"
"net"
"net/http"
"net/netip"
"os"
"os/signal"
"path/filepath"
"strconv"
"syscall"
"time"
serverui "remlink/frontend/server"
"remlink/internal/admin"
"remlink/internal/bootstrap"
"remlink/internal/config"
"remlink/internal/control"
"remlink/internal/database"
"remlink/internal/ipam"
"remlink/internal/logging"
"remlink/internal/overlay/serverwg"
sessionmanager "remlink/internal/session"
"remlink/internal/version"
)
const wireGuardEndpointEnvironment = "REMLINK_WG_ENDPOINT"
const adminTokenEnvironment = "REMLINK_ADMIN_TOKEN"
func main() {
if err := run(os.Args[1:]); err != nil {
fmt.Fprintf(os.Stderr, "remlink-server: %v\n", err)
os.Exit(1)
}
}
func run(arguments []string) error {
flags := flag.NewFlagSet("remlink-server", flag.ContinueOnError)
flags.SetOutput(os.Stderr)
configPath := flags.String("config", "config/server.yaml", "Server YAML configuration path")
wgEndpoint := flags.String("wg-endpoint", os.Getenv(wireGuardEndpointEnvironment), "public WireGuard host:port (or REMLINK_WG_ENDPOINT)")
rotateJoinToken := flags.Bool("rotate-join-token", false, "rotate Join Token in SQLite, print it, and exit")
printJoinToken := flags.Bool("print-join-token", false, "print current Join Token and exit")
if err := flags.Parse(arguments); err != nil {
return err
}
serverConfig, err := config.LoadServer(*configPath)
if err != nil {
return err
}
ctx := context.Background()
databasePath := filepath.Join(serverConfig.Data.Directory, "remlink.db")
db, err := database.Open(ctx, databasePath)
if err != nil {
return err
}
defer db.Close()
store := database.NewStore(db)
joins := bootstrap.NewJoinTokens(store)
if *rotateJoinToken || *printJoinToken {
var token string
if *rotateJoinToken {
token, err = joins.Rotate(ctx)
} else {
token, err = joins.Ensure(ctx)
}
if err != nil {
return err
}
fmt.Println(token)
return nil
}
closedSessions, err := store.CloseOpenSessions(ctx)
if err != nil {
return err
}
storedNetwork, err := admin.LoadStoredNetwork(ctx, store, admin.Network{
OverlayCIDR: serverConfig.Network.OverlayCIDR, ServerOverlayIP: serverConfig.Network.ServerOverlayIP,
WireGuardPort: serverConfig.Server.WireGuardPort, SessionUDPPort: serverConfig.Network.SessionUDPPort,
MTU: serverConfig.Network.MTU, ConfigVersion: 1,
})
if err != nil {
return err
}
serverConfig.Network.OverlayCIDR = storedNetwork.OverlayCIDR
serverConfig.Network.ServerOverlayIP = storedNetwork.ServerOverlayIP
serverConfig.Network.SessionUDPPort = storedNetwork.SessionUDPPort
serverConfig.Network.MTU = storedNetwork.MTU
serverConfig.Server.WireGuardPort = storedNetwork.WireGuardPort
_, controlPortText, err := net.SplitHostPort(serverConfig.Server.ControlListen)
if err != nil {
return err
}
serverConfig.Server.ControlListen = net.JoinHostPort(storedNetwork.ServerOverlayIP, controlPortText)
if err := validateEndpoint(*wgEndpoint, serverConfig.Server.WireGuardPort); err != nil {
return err
}
logger, err := logging.New(logging.DefaultConfig(filepath.Join(serverConfig.Data.Directory, "logs", "server.jsonl")))
if err != nil {
return err
}
defer logger.Close()
coreLogger, _ := logger.For(logging.ModuleCore)
wgLogger, _ := logger.For(logging.ModuleWG)
bootstrapLogger, _ := logger.For(logging.ModuleBootstrap)
overlayCIDR, _ := netip.ParsePrefix(serverConfig.Network.OverlayCIDR)
serverIP, _ := netip.ParseAddr(serverConfig.Network.ServerOverlayIP)
serverAddress := netip.PrefixFrom(serverIP, overlayCIDR.Bits())
wireGuard, err := serverwg.New(ctx, serverwg.Config{
InterfaceName: serverwg.DefaultInterfaceName, Address: serverAddress,
ListenPort: serverConfig.Server.WireGuardPort,
PrivateKeyPath: filepath.Join(serverConfig.Data.Directory, "server-wg.key"),
EnableForwarding: true,
})
if err != nil {
return fmt.Errorf("kernel WireGuard preflight/configuration failed: %w", err)
}
defer wireGuard.Close()
wgLogger.Info("内核 WireGuard 中心接口已就绪", "interface", serverwg.DefaultInterfaceName,
"address", serverAddress, "listen_port", serverConfig.Server.WireGuardPort)
nodes, err := store.ListNodes(ctx)
if err != nil {
return err
}
peers := make([]serverwg.Peer, 0, len(nodes))
for _, node := range nodes {
peers = append(peers, serverwg.Peer{PublicKey: node.WGPublicKey, Address: node.OverlayIP})
}
if err := wireGuard.ReconcilePeers(ctx, peers); err != nil {
return err
}
wgLogger.Info("WireGuard 对等节点已完成同步", "count", len(peers))
ipamManager, err := ipam.New(store, overlayCIDR, serverIP)
if err != nil {
return err
}
serverID, err := bootstrap.EnsureServerID(ctx, store)
if err != nil {
return err
}
joinToken, err := joins.Ensure(ctx)
if err != nil {
return err
}
bootstrapLogger.Info("Join Token 已就绪;可使用 Server 命令行查询或轮换", "token_initialized", joinToken != "")
service, err := bootstrap.NewService(store, ipamManager, joins, wireGuard, bootstrap.ServiceConfig{
ServerID: serverID, Version: version.String(), WGPublicKey: wireGuard.PublicKey(),
WGEndpoint: *wgEndpoint, OverlayCIDR: overlayCIDR, ServerOverlayIP: serverIP,
ControlURL: "ws://" + serverConfig.Server.ControlListen + "/control",
SessionUDPPort: serverConfig.Network.SessionUDPPort, MTU: serverConfig.Network.MTU,
ConfigVersion: storedNetwork.ConfigVersion,
})
if err != nil {
return err
}
controlHub, err := control.NewHub(service, store, nil, control.HubConfig{
NetworkConfigVersion: storedNetwork.ConfigVersion, EnforceRemoteIP: true,
})
if err != nil {
return err
}
sessionManager, err := sessionmanager.NewManager(store, controlHub, sessionmanager.Config{
OverlayCIDR: overlayCIDR, MTU: serverConfig.Network.MTU, UDPPort: serverConfig.Network.SessionUDPPort,
})
if err != nil {
return err
}
service.SetNodeBootstrapHandler(func(ctx context.Context, nodeID string) error {
return sessionManager.DisconnectNode(ctx, nodeID, "NODE_RUNTIME_REBUILT")
})
controlHub.SetMessageHandler(sessionManager)
coreLogger.Info("会话恢复完成", "closed_nonterminal_sessions", closedSessions)
controlMux := http.NewServeMux()
controlMux.Handle("/control", controlHub)
_, controlPortText, _ = net.SplitHostPort(serverConfig.Server.ControlListen)
controlPort, _ := strconv.Atoi(controlPortText)
controlSupervisor, err := control.NewSupervisor(controlMux, controlPort)
if err != nil {
return err
}
runtimeContext, cancelRuntime := context.WithCancel(context.Background())
defer cancelRuntime()
if err := controlSupervisor.Start(runtimeContext, serverIP); err != nil {
return err
}
defer controlSupervisor.Close()
networkManager, err := admin.NewNetworkManager(store, ipamManager, wireGuard, service, controlHub, sessionManager,
storedNetwork, func(address netip.Addr) error { return controlSupervisor.Rebind(address) })
if err != nil {
return err
}
adminHandler, err := admin.Handler(admin.HandlerConfig{
Store: store, IPAM: ipamManager, Peers: wireGuard, Control: controlHub,
Sessions: sessionManager, Network: networkManager, JoinTokens: joins,
AdminToken: os.Getenv(adminTokenEnvironment),
})
if err != nil {
return err
}
publicMux := http.NewServeMux()
publicMux.Handle("/api/v1/admin/", adminHandler)
bootstrapHandler := bootstrap.Handler(service)
publicMux.Handle("/api/v1/server/", bootstrapHandler)
publicMux.Handle("/api/v1/bootstrap/", bootstrapHandler)
webUI, err := serverui.Handler()
if err != nil {
return err
}
publicMux.Handle("/", webUI)
bootstrapServer := &http.Server{
Addr: serverConfig.Server.HTTPListen, Handler: publicMux,
ReadHeaderTimeout: 10 * time.Second, ReadTimeout: 15 * time.Second,
WriteTimeout: 30 * time.Second, IdleTimeout: 60 * time.Second,
MaxHeaderBytes: 1 << 20,
}
type runtimeError struct {
component string
err error
}
serverErrors := make(chan runtimeError, 3)
go func() {
coreLogger.Info("公网 Bootstrap API 已开始监听", "address", bootstrapServer.Addr, "version", version.String())
serverErrors <- runtimeError{component: "Bootstrap API", err: bootstrapServer.ListenAndServe()}
}()
go func() {
coreLogger.Info("Overlay Control WebSocket 已开始监听", "address", serverConfig.Server.ControlListen)
serverErrors <- runtimeError{component: "Control WebSocket", err: controlSupervisor.Wait(runtimeContext)}
}()
go func() {
serverErrors <- runtimeError{component: "Control heartbeat monitor", err: controlHub.Run(runtimeContext)}
}()
signalContext, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM)
defer stop()
select {
case <-signalContext.Done():
cancelRuntime()
shutdownContext, cancel := context.WithTimeout(context.Background(), 10*time.Second)
defer cancel()
if err := bootstrapServer.Shutdown(shutdownContext); err != nil {
return fmt.Errorf("shutdown public Bootstrap API: %w", err)
}
if err := controlSupervisor.Close(); err != nil {
return fmt.Errorf("shutdown Control WebSocket: %w", err)
}
coreLogger.Info("Server 已停止")
return nil
case failure := <-serverErrors:
cancelRuntime()
if errors.Is(failure.err, http.ErrServerClosed) || errors.Is(failure.err, context.Canceled) {
return nil
}
return fmt.Errorf("%s failed: %w", failure.component, failure.err)
}
}
func validateEndpoint(endpoint string, listenPort int) error {
host, portText, err := net.SplitHostPort(endpoint)
if err != nil || host == "" {
return fmt.Errorf("--wg-endpoint (or %s) must be a public host:port", wireGuardEndpointEnvironment)
}
port, err := strconv.Atoi(portText)
if err != nil || port != listenPort {
return fmt.Errorf("public WireGuard endpoint port must match configured listen port %d", listenPort)
}
return nil
}