feat: revive wireguard overlay plane

Lee Yunjin committed Apr 8, 2026 at 21:38 UTC b9b1e6e9cc47a9b615c92b631dc0300eff98bbe6
13 files changed +990 -144
cmd/relay-server/main.go
+62 -33
@@ -31,36 +31,42 @@ func main() {
31 }
32
33 type relayServerConfig struct {
34 - PortalURL string
35 - APIPort int
36 - SNIPort int
37 - MinPort int
38 - MaxPort int
39 - UDPEnabled bool
40 - TCPEnabled bool
41 - LandingPageEnabled bool
42 - Bootstraps string
43 - DiscoveryEnabled bool
44 - IdentityPath string
45 - AdminSecretKey string
46 - TrustProxyHeaders bool
47 - TrustedProxyCIDRs string
48 - AdminSettingsPath string
49 - KeylessDir string
50 -
51 - HeadlessShellURL string
52 -
53 - ACMEDNSProvider string
54 - ENSGaslessEnabled bool
55 - CloudflareToken string
56 - GCPProjectID string
57 - GCPManagedZone string
58 - AWSAccessKeyID string
59 - AWSSecretAccessKey string
60 - AWSSessionToken string
61 - AWSRegion string
62 - AWSHostedZoneID string
63 - AWSDNSSECKMSKeyARN string
34 + PortalURL string
35 + APIPort int
36 + SNIPort int
37 + MinPort int
38 + MaxPort int
39 + UDPEnabled bool
40 + TCPEnabled bool
41 + LandingPageEnabled bool
42 + Bootstraps string
43 + DiscoveryEnabled bool
44 + MaxRouting int
45 + OverlayEnabled bool
46 + OverlayMaxHops int
47 + OverlayCongestion float64
48 + WireGuardPrivateKey string
49 + WireGuardEndpoint string
50 + OverlayIPv4 string
51 + OverlayCIDRs string
52 + IdentityPath string
53 + AdminSecretKey string
54 + TrustProxyHeaders bool
55 + TrustedProxyCIDRs string
56 + AdminSettingsPath string
57 + KeylessDir string
58 + HeadlessShellURL string
59 + ACMEDNSProvider string
60 + ENSGaslessEnabled bool
61 + CloudflareToken string
62 + GCPProjectID string
63 + GCPManagedZone string
64 + AWSAccessKeyID string
65 + AWSSecretAccessKey string
66 + AWSSessionToken string
67 + AWSRegion string
68 + AWSHostedZoneID string
69 + AWSDNSSECKMSKeyARN string
70 }
71
72 func runServeCommand(args []string) error {
@@ -77,6 +83,14 @@ func runServeCommand(args []string) error {
83 utils.BoolFlagEnv(fs, &cfg.LandingPageEnabled, "landing-page-enabled", false, "enable landing page by default when no admin setting has been saved yet", "LANDING_PAGE_ENABLED")
84 utils.StringFlagEnv(fs, &cfg.Bootstraps, "bootstraps", "", "additional bootstrap relay API URLs used for discovery expansion", "BOOTSTRAPS")
85 utils.BoolFlagEnv(fs, &cfg.DiscoveryEnabled, "discovery", false, "serve relay discovery endpoints and poll discovery peers", "DISCOVERY")
86 + utils.IntFlagEnv(fs, &cfg.MaxRouting, "max-routing", 1, nil, "maximum number of discovery routing attempts per refresh", "MAX_ROUTING")
87 + utils.BoolFlagEnv(fs, &cfg.OverlayEnabled, "overlay-enabled", false, "enable experimental Pepper overlay route planning", "OVERLAY_ENABLED")
88 + utils.IntFlagEnv(fs, &cfg.OverlayMaxHops, "overlay-max-hops", 0, nil, "Pepper overlay max hops (0 disables overlay route planning)", "OVERLAY_MAX_HOPS")
89 + utils.Float64FlagEnv(fs, &cfg.OverlayCongestion, "overlay-congestion-latency-ms", 120, nil, "latency threshold in ms to trigger reverse-Siamese overlay route selection", "OVERLAY_CONGESTION_LATENCY_MS")
90 + utils.StringFlagEnv(fs, &cfg.WireGuardPrivateKey, "wireguard-private-key", "", "wireguard private key for relay overlay", "WIREGUARD_PRIVATE_KEY")
91 + utils.StringFlagEnv(fs, &cfg.WireGuardEndpoint, "wireguard-endpoint", "", "wireguard endpoint (host:port) for relay overlay", "WIREGUARD_ENDPOINT")
92 + utils.StringFlagEnv(fs, &cfg.OverlayIPv4, "overlay-ipv4", "", "explicit overlay IPv4 override (auto-derived from public key when unset)", "OVERLAY_IPV4")
93 + utils.StringFlagEnv(fs, &cfg.OverlayCIDRs, "overlay-cidrs", "", "comma-separated overlay CIDR allowlist advertised to peers", "OVERLAY_CIDRS")
94 utils.StringFlagEnv(fs, &cfg.IdentityPath, "identity-path", "identity.json", "relay identity json file path", "IDENTITY_PATH")
95 utils.StringFlagEnv(fs, &cfg.AdminSecretKey, "admin-secret-key", "", "admin auth secret", "ADMIN_SECRET_KEY")
96 utils.BoolFlagEnv(fs, &cfg.TrustProxyHeaders, "trust-proxy-headers", false, "trust X-Forwarded-* and X-Real-IP headers from trusted proxies", "TRUST_PROXY_HEADERS")
@@ -108,6 +122,9 @@ func runServeCommand(args []string) error {
122 printRootUsage(os.Stderr)
123 return err
124 }
125 + if err := utils.ValidateMaxRouting(cfg.MaxRouting); err != nil {
126 + return err
127 + }
128
129 log.Info().
130 Str("release_version", types.ReleaseVersion).
@@ -118,6 +135,9 @@ func runServeCommand(args []string) error {
135 Int("max_port", cfg.MaxPort).
136 Bool("landing_page_enabled", cfg.LandingPageEnabled).
137 Bool("discovery_enabled", cfg.DiscoveryEnabled).
138 + Int("max_routing", cfg.MaxRouting).
139 + Bool("overlay_enabled", cfg.OverlayEnabled).
140 + Int("overlay_max_hops", cfg.OverlayMaxHops).
141 Str("acme_dns_provider", cfg.ACMEDNSProvider).
142 Bool("ens_gasless_enabled", cfg.ENSGaslessEnabled).
143 Bool("udp_enabled", cfg.UDPEnabled).
@@ -135,11 +155,16 @@ func runServer(ctx context.Context, cfg relayServerConfig) error {
155 if err != nil {
156 return fmt.Errorf("resolve discovery bootstraps: %w", err)
157 }
158 + overlayCIDRs := utils.SplitCSV(cfg.OverlayCIDRs)
159
160 server, err := portal.NewServer(portal.ServerConfig{
140 - PortalURL: cfg.PortalURL,
141 - IdentityPath: cfg.IdentityPath,
142 - Bootstraps: bootstraps,
161 + PortalURL: cfg.PortalURL,
162 + IdentityPath: cfg.IdentityPath,
163 + Bootstraps: bootstraps,
164 + WireGuardPrivateKey: cfg.WireGuardPrivateKey,
165 + WireGuardEndpoint: cfg.WireGuardEndpoint,
166 + OverlayIPv4: cfg.OverlayIPv4,
167 + OverlayCIDRs: overlayCIDRs,
168 ACME: acme.Config{
169 KeyDir: cfg.KeylessDir,
170 DNSProvider: cfg.ACMEDNSProvider,
@@ -159,6 +184,10 @@ func runServer(ctx context.Context, cfg relayServerConfig) error {
184 TrustedProxyCIDRs: cfg.TrustedProxyCIDRs,
185 TrustProxyHeaders: cfg.TrustProxyHeaders,
186 DiscoveryEnabled: cfg.DiscoveryEnabled,
187 + MaxRouting: cfg.MaxRouting,
188 + OverlayEnabled: cfg.OverlayEnabled,
189 + OverlayMaxHops: cfg.OverlayMaxHops,
190 + OverlayCongestion: cfg.OverlayCongestion,
191 MinPort: cfg.MinPort,
192 MaxPort: cfg.MaxPort,
193 UDPEnabled: cfg.UDPEnabled,
go.mod
+7
@@ -21,6 +21,7 @@ require (
21 golang.org/x/net v0.51.0
22 golang.org/x/oauth2 v0.35.0
23 golang.org/x/sync v0.19.0
24 + golang.zx2c4.com/wireguard v0.0.0-20250521234502-f333402bd9cb
25 google.golang.org/api v0.267.0
26 )
27
@@ -46,6 +47,7 @@ require (
47 github.com/felixge/httpsnoop v1.0.4 // indirect
48 github.com/go-logr/logr v1.4.3 // indirect
49 github.com/go-logr/stdr v1.2.2 // indirect
50 + github.com/google/btree v1.1.2 // indirect
51 github.com/google/s2a-go v0.1.9 // indirect
52 github.com/google/uuid v1.6.0 // indirect
53 github.com/googleapis/enterprise-certificate-proxy v0.3.11 // indirect
@@ -68,8 +70,13 @@ require (
70 golang.org/x/mod v0.32.0 // indirect
71 golang.org/x/sys v0.41.0 // indirect
72 golang.org/x/text v0.34.0 // indirect
73 + golang.org/x/time v0.14.0 // indirect
74 golang.org/x/tools v0.41.0 // indirect
75 + golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2 // indirect
76 google.golang.org/genproto/googleapis/rpc v0.0.0-20260203192932-546029d2fa20 // indirect
77 google.golang.org/grpc v1.79.3 // indirect
78 google.golang.org/protobuf v1.36.11 // indirect
79 + gvisor.dev/gvisor v0.0.0-20250503011706-39ed1f5ac29c // indirect
80 )
81 +
82 +exclude golang.zx2c4.com/wireguard/tun/netstack v0.0.0-20220703234212-c31a7b1ab478
go.sum
+8
@@ -69,6 +69,8 @@ github.com/go-rod/rod v0.116.2/go.mod h1:H+CMO9SCNc2TJ2WfrG+pKhITz57uGNYU43qYHh4
69 github.com/godbus/dbus/v5 v5.0.4/go.mod h1:xhWf0FNVPg57R7Z0UbKHbJfkEywrmjJnf7w5xrFpKfA=
70 github.com/golang/protobuf v1.5.4 h1:i7eJL8qZTpSEXOPTxNKhASYpMn+8e5Q6AdndVa1dWek=
71 github.com/golang/protobuf v1.5.4/go.mod h1:lnTiLA8Wa4RWRcIUkrtSVa5nRhsEGBg48fD6rSs7xps=
72 +github.com/google/btree v1.1.2 h1:xf4v41cLI2Z6FxbKm+8Bu+m8ifhj15JuZ9sa0jZCMUU=
73 +github.com/google/btree v1.1.2/go.mod h1:qOPhT0dTNdNzV6Z/lhRX0YXUafgPLFUh+gZMl761Gm4=
74 github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8=
75 github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU=
76 github.com/google/s2a-go v0.1.9 h1:LGD7gtMgezd8a/Xak7mEWL0PjoTQFvpRudN895yqKW0=
@@ -160,6 +162,10 @@ golang.org/x/time v0.14.0 h1:MRx4UaLrDotUKUdCIqzPC48t1Y9hANFKIRpNx+Te8PI=
162 golang.org/x/time v0.14.0/go.mod h1:eL/Oa2bBBK0TkX57Fyni+NgnyQQN4LitPmob2Hjnqw4=
163 golang.org/x/tools v0.41.0 h1:a9b8iMweWG+S0OBnlU36rzLp20z1Rp10w+IY2czHTQc=
164 golang.org/x/tools v0.41.0/go.mod h1:XSY6eDqxVNiYgezAVqqCeihT4j1U2CCsqvH3WhQpnlg=
165 +golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2 h1:B82qJJgjvYKsXS9jeunTOisW56dUokqW/FOteYJJ/yg=
166 +golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2/go.mod h1:deeaetjYA+DHMHg+sMSMI58GrEteJUUzzw7en6TJQcI=
167 +golang.zx2c4.com/wireguard v0.0.0-20250521234502-f333402bd9cb h1:whnFRlWMcXI9d+ZbWg+4sHnLp52d5yiIPUxMBSt4X9A=
168 +golang.zx2c4.com/wireguard v0.0.0-20250521234502-f333402bd9cb/go.mod h1:rpwXGsirqLqN2L0JDJQlwOboGHmptD5ZD6T2VmcqhTw=
169 gonum.org/v1/gonum v0.16.0 h1:5+ul4Swaf3ESvrOnidPp4GZbzf0mxVQpDCYUQE7OJfk=
170 gonum.org/v1/gonum v0.16.0/go.mod h1:fef3am4MQ93R2HHpKnLk4/Tbh/s0+wqD5nfa6Pnwy4E=
171 google.golang.org/api v0.267.0 h1:w+vfWPMPYeRs8qH1aYYsFX68jMls5acWl/jocfLomwE=
@@ -176,3 +182,5 @@ google.golang.org/protobuf v1.36.11 h1:fV6ZwhNocDyBLK0dj+fg8ektcVegBBuEolpbTQyBN
182 google.golang.org/protobuf v1.36.11/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco=
183 gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA=
184 gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM=
185 +gvisor.dev/gvisor v0.0.0-20250503011706-39ed1f5ac29c h1:m/r7OM+Y2Ty1sgBQ7Qb27VgIMBW8ZZhT4gLnUyDIhzI=
186 +gvisor.dev/gvisor v0.0.0-20250503011706-39ed1f5ac29c/go.mod h1:3r5CMtNQMKIvBlrmM9xWUNamjKBYPOWyXOjmg5Kts3g=
portal/api_server.go
+36 -11
@@ -191,15 +191,24 @@ func (s *Server) handleRelayDiscovery(w http.ResponseWriter, r *http.Request) {
191 }
192
193 self, err := discovery.NormalizeDescriptor(types.RelayDescriptor{
194 - Identity: s.identity.Copy(),
195 - Sequence: uint64(now.UnixMilli()),
196 - Version: 1,
197 - IssuedAt: now,
198 - ExpiresAt: now.Add(2 * types.DiscoveryPollInterval),
199 - APIHTTPSAddr: s.cfg.PortalURL,
200 - IngressTLSAddr: ingressAddr,
201 - SupportsUDP: s.cfg.UDPEnabled && s.quicTunnel != nil,
202 - SupportsTCP: s.cfg.TCPEnabled,
194 + Identity: s.identity.Copy(),
195 + RelayID: s.cfg.PortalURL,
196 + OwnerAddress: s.identity.Address,
197 + SignerPublicKey: s.identity.PublicKey,
198 + Sequence: uint64(now.UnixMilli()),
199 + Version: 1,
200 + IssuedAt: now,
201 + ExpiresAt: now.Add(2 * types.DiscoveryPollInterval),
202 + APIHTTPSAddr: s.cfg.PortalURL,
203 + IngressTLSAddr: ingressAddr,
204 + WireGuardPublicKey: wireGuardField(s.wireGuardPeerPlaneEnabled(), s.cfg.WireGuardPublicKey),
205 + WireGuardEndpoint: wireGuardField(s.wireGuardPeerPlaneEnabled(), s.cfg.WireGuardEndpoint),
206 + OverlayIPv4: wireGuardField(s.wireGuardPeerPlaneEnabled(), s.cfg.OverlayIPv4),
207 + OverlayCIDRs: overlayCIDRsField(s.wireGuardPeerPlaneEnabled(), s.cfg.OverlayCIDRs),
208 + SupportsUDP: s.cfg.UDPEnabled && s.quicTunnel != nil,
209 + SupportsTCP: s.cfg.TCPEnabled,
210 + SupportsOverlayPeer: s.cfg.OverlayEnabled,
211 + Load: float64(s.loadMgr.ActiveConns()),
212 })
213 if err != nil {
214 utils.WriteAPIError(w, http.StatusInternalServerError, types.APIErrorCodeInternal, err.Error())
@@ -212,8 +221,8 @@ func (s *Server) handleRelayDiscovery(w http.ResponseWriter, r *http.Request) {
221 Self: self,
222 Relays: nil,
223 }
215 - if s.relaySet != nil {
216 - resp.Relays = s.relaySet.ActiveRelayDescriptors()
224 + if s.discoveryMgr != nil {
225 + resp.Relays = s.discoveryMgr.ActiveRelayDescriptors()
226 }
227 utils.WriteAPIData(w, http.StatusOK, resp)
228 }
@@ -394,6 +403,22 @@ func (s *Server) handleUnregister(w http.ResponseWriter, r *http.Request) {
403 utils.WriteAPIData(w, http.StatusOK, map[string]any{})
404 }
405
406 +func wireGuardField(enabled bool, value string) string {
407 + if !enabled {
408 + return ""
409 + }
410 + return value
411 +}
412 +
413 +func overlayCIDRsField(enabled bool, cidrs []string) []string {
414 + if !enabled || len(cidrs) == 0 {
415 + return nil
416 + }
417 + out := make([]string, len(cidrs))
418 + copy(out, cidrs)
419 + return out
420 +}
421 +
422 func (s *Server) handleConnect(w http.ResponseWriter, r *http.Request) {
423 if !utils.RequireMethod(w, r, http.MethodGet) {
424 return
portal/discovery/discovery.go
+35
@@ -22,6 +22,14 @@ func NormalizeDescriptor(desc types.RelayDescriptor) (types.RelayDescriptor, err
22 desc.Name = utils.NormalizeHostname(desc.Name)
23 desc.Address = strings.TrimSpace(desc.Address)
24 desc.APIHTTPSAddr = strings.TrimSpace(desc.APIHTTPSAddr)
25 + desc.RelayID = strings.TrimSpace(desc.RelayID)
26 + desc.IngressTLSAddr = strings.TrimSpace(desc.IngressTLSAddr)
27 + desc.WireGuardPublicKey = strings.TrimSpace(desc.WireGuardPublicKey)
28 + desc.WireGuardEndpoint = strings.TrimSpace(desc.WireGuardEndpoint)
29 + desc.OverlayIPv4 = strings.TrimSpace(desc.OverlayIPv4)
30 + desc.OverlayCIDRs = utils.NormalizeIPPrefixes(desc.OverlayCIDRs)
31 + desc.OwnerAddress = strings.TrimSpace(desc.OwnerAddress)
32 + desc.SignerPublicKey = strings.TrimSpace(desc.SignerPublicKey)
33 if !desc.IssuedAt.IsZero() {
34 desc.IssuedAt = desc.IssuedAt.UTC()
35 }
@@ -36,6 +44,16 @@ func NormalizeDescriptor(desc types.RelayDescriptor) (types.RelayDescriptor, err
44 }
45 desc.APIHTTPSAddr = normalized
46 }
47 + if desc.RelayID != "" {
48 + normalized, err := utils.NormalizeRelayURL(desc.RelayID)
49 + if err != nil {
50 + return types.RelayDescriptor{}, fmt.Errorf("normalize relay id: %w", err)
51 + }
52 + desc.RelayID = normalized
53 + }
54 + if desc.RelayID == "" {
55 + desc.RelayID = desc.APIHTTPSAddr
56 + }
57 if desc.Address != "" {
58 normalized, err := utils.NormalizeEVMAddress(desc.Address)
59 if err != nil {
@@ -43,6 +61,19 @@ func NormalizeDescriptor(desc types.RelayDescriptor) (types.RelayDescriptor, err
61 }
62 desc.Address = normalized
63 }
64 + if desc.OwnerAddress == "" {
65 + desc.OwnerAddress = desc.Address
66 + }
67 + if desc.OwnerAddress != "" {
68 + normalized, err := utils.NormalizeEVMAddress(desc.OwnerAddress)
69 + if err != nil {
70 + return types.RelayDescriptor{}, fmt.Errorf("normalize owner address: %w", err)
71 + }
72 + desc.OwnerAddress = normalized
73 + }
74 + if desc.SignerPublicKey == "" {
75 + desc.SignerPublicKey = desc.PublicKey
76 + }
77 return desc, nil
78 }
79
@@ -61,6 +92,10 @@ func ValidateDescriptor(desc types.RelayDescriptor, now time.Time) (types.RelayD
92 return types.RelayDescriptor{}, errors.New("identity.name is required")
93 case normalized.APIHTTPSAddr == "":
94 return types.RelayDescriptor{}, errors.New("api_https_addr is required")
95 + case normalized.RelayID == "":
96 + return types.RelayDescriptor{}, errors.New("relay_id is required")
97 + case normalized.APIHTTPSAddr != "" && normalized.RelayID != normalized.APIHTTPSAddr:
98 + return types.RelayDescriptor{}, errors.New("relay_id must match api_https_addr")
99 case normalized.Sequence == 0:
100 return types.RelayDescriptor{}, errors.New("sequence is required")
101 case normalized.Version == 0:
portal/discovery/relayset.go
+1
@@ -252,6 +252,7 @@ func (s *RelaySet) bootstrapDescriptors() []types.RelayDescriptor {
252 Name: utils.PortalRootHost(relayURL),
253 },
254 APIHTTPSAddr: relayURL,
255 + RelayID: relayURL,
256 Version: 1,
257 })
258 }
portal/server.go
+452 -81
@@ -5,6 +5,7 @@ import (
5 "crypto/tls"
6 "errors"
7 "fmt"
8 + "hash/crc32"
9 "io"
10 "net"
11 "net/http"
@@ -20,8 +21,10 @@ import (
21 "github.com/gosuda/portal-tunnel/v2/portal/acme"
22 "github.com/gosuda/portal-tunnel/v2/portal/discovery"
23 "github.com/gosuda/portal-tunnel/v2/portal/keyless"
24 + "github.com/gosuda/portal-tunnel/v2/portal/overlay"
25 "github.com/gosuda/portal-tunnel/v2/portal/policy"
26 "github.com/gosuda/portal-tunnel/v2/portal/transport"
27 + "github.com/gosuda/portal-tunnel/v2/portal/wireguard"
28 "github.com/gosuda/portal-tunnel/v2/types"
29 "github.com/gosuda/portal-tunnel/v2/utils"
30 "github.com/gosuda/portal-tunnel/v2/utils/thumbnail"
@@ -37,40 +40,57 @@ const (
40 )
41
42 type ServerConfig struct {
40 - PortalURL string
41 - IdentityPath string
42 - Bootstraps []string
43 - ACME acme.Config
44 - APIPort int
45 - SNIPort int
46 - APIListenAddr string
47 - SNIListenAddr string
48 - TrustedProxyCIDRs string
49 - TrustProxyHeaders bool
50 - DiscoveryEnabled bool
51 - MinPort int
52 - MaxPort int
53 - UDPEnabled bool
54 - TCPEnabled bool
55 - HeadlessShellURL string
43 + PortalURL string
44 + IdentityPath string
45 + Bootstraps []string
46 + WireGuardPrivateKey string
47 + WireGuardEndpoint string
48 + WireGuardPublicKey string
49 + OverlayIPv4 string
50 + OverlayCIDRs []string
51 + ACME acme.Config
52 + APIPort int
53 + SNIPort int
54 + APIListenAddr string
55 + SNIListenAddr string
56 + TrustedProxyCIDRs string
57 + TrustProxyHeaders bool
58 + DiscoveryEnabled bool
59 + MaxRouting int
60 + OverlayEnabled bool
61 + OverlayMaxHops int
62 + OverlayCongestion float64
63 + MinPort int
64 + MaxPort int
65 + UDPEnabled bool
66 + TCPEnabled bool
67 + HeadlessShellURL string
68 }
69
70 type Server struct {
71 sniListener net.Listener
72 apiListener net.Listener
73 apiServer *http.Server
74 + wgPeerListener net.Listener
75 + wgPeerServer *http.Server
76 apiTLSClose io.Closer
77 acmeManager *acme.Manager
78 quicTunnel *quic.Listener
79 + wgRuntime *wireguard.Runtime
80 cancel context.CancelFunc
81 group *errgroup.Group
82 registry *leaseRegistry
83 ports *transport.PortAllocator
84 tcpPorts *transport.PortAllocator
85 + loadMgr *policy.LoadManager
86 + weightMgr *policy.WeightManager
87 identity types.Identity
88 cfg ServerConfig
89 trustedProxyCIDRs []*net.IPNet
73 - relaySet *discovery.RelaySet
90 + discoveryMgr *discovery.Manager
91 + overlayPolicy *overlay.RoutePolicy
92 + overlayRoute []uint32
93 + overlayRouteMu sync.RWMutex
94 thumbnails *thumbnail.Service
95 shutdownOnce sync.Once
96 }
@@ -93,8 +113,71 @@ func NewServer(cfg ServerConfig) (*Server, error) {
113 if err != nil {
114 return nil, fmt.Errorf("normalize bootstraps: %w", err)
115 }
116 + selfRelayURL := ""
117 + if trimmedPortalURL := strings.TrimSpace(cfg.PortalURL); trimmedPortalURL != "" {
118 + normalizedPortalURL, err := utils.NormalizeRelayURL(trimmedPortalURL)
119 + if err != nil {
120 + return nil, fmt.Errorf("normalize portal url: %w", err)
121 + }
122 + selfRelayURL = normalizedPortalURL
123 + }
124 + if len(bootstraps) > 0 {
125 + filtered := bootstraps[:0]
126 + for _, relayURL := range bootstraps {
127 + if selfRelayURL != "" && relayURL == selfRelayURL {
128 + continue
129 + }
130 + filtered = append(filtered, relayURL)
131 + }
132 + bootstraps = filtered
133 + }
134 cfg.Bootstraps = bootstraps
97 -
135 + wireGuardConfigured := strings.TrimSpace(cfg.WireGuardPrivateKey) != "" ||
136 + strings.TrimSpace(cfg.WireGuardEndpoint) != "" ||
137 + strings.TrimSpace(cfg.WireGuardPublicKey) != "" ||
138 + strings.TrimSpace(cfg.OverlayIPv4) != "" ||
139 + len(cfg.OverlayCIDRs) > 0
140 + if wireGuardConfigured {
141 + if strings.TrimSpace(cfg.WireGuardPrivateKey) == "" {
142 + return nil, errors.New("wireguard private key is required when overlay is enabled")
143 + }
144 + if strings.TrimSpace(cfg.WireGuardEndpoint) == "" {
145 + return nil, errors.New("wireguard endpoint is required when overlay is enabled")
146 + }
147 + normalizedKey, err := utils.NormalizeWireGuardPrivateKey(cfg.WireGuardPrivateKey)
148 + if err != nil {
149 + return nil, fmt.Errorf("normalize wireguard private key: %w", err)
150 + }
151 + cfg.WireGuardPrivateKey = normalizedKey
152 + if strings.TrimSpace(cfg.WireGuardPublicKey) == "" {
153 + publicKey, err := utils.WireGuardPublicKeyFromPrivate(cfg.WireGuardPrivateKey)
154 + if err != nil {
155 + return nil, fmt.Errorf("derive wireguard public key: %w", err)
156 + }
157 + cfg.WireGuardPublicKey = publicKey
158 + }
159 + if strings.TrimSpace(cfg.OverlayIPv4) == "" {
160 + overlayIP, err := utils.DeriveWireGuardOverlayIPv4(cfg.WireGuardPublicKey)
161 + if err != nil {
162 + return nil, fmt.Errorf("derive overlay ipv4: %w", err)
163 + }
164 + cfg.OverlayIPv4 = overlayIP
165 + }
166 + cfg.OverlayCIDRs = utils.NormalizeIPPrefixes(cfg.OverlayCIDRs)
167 + }
168 + if cfg.OverlayMaxHops < 0 {
169 + return nil, errors.New("overlay max hops must be >= 0")
170 + }
171 + if cfg.OverlayMaxHops > 10 {
172 + return nil, errors.New("overlay max hops must be <= 10")
173 + }
174 + if cfg.OverlayCongestion <= 0 {
175 + cfg.OverlayCongestion = 120
176 + }
177 + cfg.OverlayEnabled = cfg.OverlayEnabled && cfg.OverlayMaxHops > 0
178 + if wireGuardConfigured {
179 + cfg.OverlayEnabled = true
180 + }
181 transportEnabled := cfg.UDPEnabled || cfg.TCPEnabled
182 hasPortRange := cfg.MinPort > 0 && cfg.MaxPort > 0
183 if transportEnabled {
@@ -117,21 +200,21 @@ func NewServer(cfg ServerConfig) (*Server, error) {
200 portMax = cfg.MaxPort
201 }
202
120 - identity, generatedIdentity, err := utils.LoadOrCreateIdentity(cfg.IdentityPath, types.Identity{Name: rootHost})
203 + identity, created, err := utils.LoadOrCreateIdentity(cfg.IdentityPath, types.Identity{Name: rootHost})
204 if err != nil {
205 return nil, fmt.Errorf("load relay identity: %w", err)
206 }
124 - if generatedIdentity {
207 + if created {
208 log.Warn().
209 Str("identity_path", cfg.IdentityPath).
210 Str("address", identity.Address).
211 Msg("generated relay identity and saved it to disk")
212 + } else {
213 + log.Info().
214 + Str("identity_path", cfg.IdentityPath).
215 + Str("address", identity.Address).
216 + Msg("loaded relay identity from disk")
217 }
130 - selfRelayURL, err := utils.NormalizeRelayURL(cfg.PortalURL)
131 - if err != nil {
132 - return nil, fmt.Errorf("normalize portal url: %w", err)
133 - }
134 - cfg.Bootstraps = utils.RemoveRelayURL(cfg.Bootstraps, selfRelayURL)
218
219 tcpPortMin, tcpPortMax := 0, 0
220 if cfg.TCPEnabled {
@@ -139,10 +222,10 @@ func NewServer(cfg ServerConfig) (*Server, error) {
222 tcpPortMax = cfg.MaxPort
223 }
224
142 - policy := policy.NewRuntime()
143 - policy.SetUDPPolicy(cfg.UDPEnabled, 0)
144 - policy.SetTCPPortPolicy(cfg.TCPEnabled, 0)
145 - registry := newLeaseRegistry(policy)
225 + runtimePolicy := policy.NewRuntime()
226 + runtimePolicy.SetUDPPolicy(cfg.UDPEnabled, 0)
227 + runtimePolicy.SetTCPPortPolicy(cfg.TCPEnabled, 0)
228 + registry := newLeaseRegistry(runtimePolicy)
229 ports := transport.NewPortAllocator(portMin, portMax, 5*time.Minute)
230 tcpPorts := transport.NewPortAllocator(tcpPortMin, tcpPortMax, 5*time.Minute)
231
@@ -151,17 +234,26 @@ func NewServer(cfg ServerConfig) (*Server, error) {
234 registry: registry,
235 ports: ports,
236 tcpPorts: tcpPorts,
237 + loadMgr: policy.NewLoadManager(),
238 + weightMgr: policy.NewWeightManager(),
239 identity: identity,
240 trustedProxyCIDRs: trustedProxyCIDRs,
241 thumbnails: thumbnail.NewService(cfg.HeadlessShellURL),
242 }
158 -
243 + if cfg.OverlayEnabled {
244 + s.overlayPolicy = overlay.NewRoutePolicy()
245 + }
246 if cfg.DiscoveryEnabled {
160 - s.relaySet = discovery.NewRelaySet()
161 - if err := s.relaySet.SetSelfRelay(identity, selfRelayURL); err != nil {
162 - return nil, fmt.Errorf("set self relay: %w", err)
247 + manager, err := discovery.NewManager(discovery.ManagerConfig{
248 + Identity: identity,
249 + PortalURL: cfg.PortalURL,
250 + Bootstraps: cfg.Bootstraps,
251 + MaxRouting: cfg.MaxRouting,
252 + })
253 + if err != nil {
254 + return nil, err
255 }
164 - s.relaySet.SetBootstrapRelayURLs(cfg.Bootstraps)
256 + s.discoveryMgr = manager
257 }
258
259 return s, nil
@@ -175,35 +267,30 @@ func (s *Server) Start(ctx context.Context, apiMux *http.ServeMux) error {
267 return err
268 }
269
178 - var cleanups []func()
179 - defer func() {
180 - for i := len(cleanups) - 1; i >= 0; i-- {
181 - cleanups[i]()
182 - }
183 - }()
184 -
185 - cleanups = append(cleanups, acmeManager.Stop)
186 -
270 serverCtx, cancel := context.WithCancel(ctx)
188 - cleanups = append(cleanups, cancel)
189 -
271 var listenConfig net.ListenConfig
272
273 apiListener, err := listenConfig.Listen(serverCtx, "tcp", s.cfg.APIListenAddr)
274 if err != nil {
275 + acmeManager.Stop()
276 + cancel()
277 return fmt.Errorf("listen api: %w", err)
278 }
196 - cleanups = append(cleanups, func() { _ = apiListener.Close() })
197 -
279 sniListener, err := listenConfig.Listen(serverCtx, "tcp", s.cfg.SNIListenAddr)
280 if err != nil {
281 + acmeManager.Stop()
282 + _ = apiListener.Close()
283 + cancel()
284 return fmt.Errorf("listen sni: %w", err)
285 }
202 - cleanups = append(cleanups, func() { _ = sniListener.Close() })
286
287 group, groupCtx := errgroup.WithContext(serverCtx)
288 wrappedAPIListener, apiServer, apiCloser, err := s.newAPIServer(apiListener, apiMux, apiTLS)
289 if err != nil {
290 + acmeManager.Stop()
291 + _ = apiListener.Close()
292 + _ = sniListener.Close()
293 + cancel()
294 return err
295 }
296
@@ -214,13 +301,27 @@ func (s *Server) Start(ctx context.Context, apiMux *http.ServeMux) error {
301 s.acmeManager = acmeManager
302 s.cancel = cancel
303 s.group = group
217 - cleanups = nil
304 +
305 + if s.wireGuardPeerPlaneEnabled() {
306 + if err := s.startWireGuardPeerPlane(); err != nil {
307 + acmeManager.Stop()
308 + _ = apiServer.Close()
309 + _ = apiCloser.Close()
310 + _ = sniListener.Close()
311 + cancel()
312 + return fmt.Errorf("start wireguard peer plane: %w", err)
313 + }
314 + }
315
316 group.Go(s.runAPIServer)
317 + if s.wgPeerServer != nil && s.wgPeerListener != nil {
318 + group.Go(s.runWireGuardPeerAPIServer)
319 + group.Go(func() error { return s.runWireGuardSyncLoop(groupCtx) })
320 + }
321 group.Go(func() error { return s.runSNIListener(groupCtx) })
322 group.Go(func() error { return s.runLeaseJanitor(groupCtx, 5*time.Second) })
323 if s.cfg.DiscoveryEnabled {
223 - group.Go(func() error { return s.relaySet.RunLoop(groupCtx, nil, nil) })
324 + group.Go(func() error { return s.runRelayDiscoveryLoop(groupCtx) })
325 }
326 s.acmeManager.Start(serverCtx)
327
@@ -244,6 +345,9 @@ func (s *Server) Start(ctx context.Context, apiMux *http.ServeMux) error {
345 Int("min_port", s.cfg.MinPort).
346 Int("max_port", s.cfg.MaxPort).
347 Bool("discovery_enabled", s.cfg.DiscoveryEnabled).
348 + Bool("wireguard_enabled", s.wireGuardPeerPlaneEnabled()).
349 + Bool("overlay_enabled", s.cfg.OverlayEnabled).
350 + Int("overlay_max_hops", s.cfg.OverlayMaxHops).
351 Bool("udp_enabled", s.cfg.UDPEnabled).
352 Bool("tcp_enabled", s.cfg.TCPEnabled)
353 if s.quicTunnel != nil {
@@ -258,7 +362,18 @@ func (s *Server) Wait() error {
362 if s.group == nil {
363 return nil
364 }
261 - return s.group.Wait()
365 + err := s.group.Wait()
366 + if errors.Is(err, context.Canceled) {
367 + return nil
368 + }
369 + return err
370 +}
371 +
372 +func (s *Server) Identity() types.Identity {
373 + if s == nil {
374 + return types.Identity{}
375 + }
376 + return s.identity.Copy()
377 }
378
379 func (s *Server) Shutdown(ctx context.Context) error {
@@ -298,6 +413,17 @@ func (s *Server) Shutdown(ctx context.Context) error {
413 shutdownErr = err
414 }
415 }
416 + if s.wgPeerServer != nil {
417 + if err := s.wgPeerServer.Shutdown(ctx); err != nil && shutdownErr == nil && !errors.Is(err, http.ErrServerClosed) {
418 + shutdownErr = err
419 + }
420 + }
421 + if s.wgPeerListener != nil {
422 + _ = s.wgPeerListener.Close()
423 + }
424 + if s.wgRuntime != nil {
425 + _ = s.wgRuntime.Close()
426 + }
427 if s.apiTLSClose != nil {
428 _ = s.apiTLSClose.Close()
429 }
@@ -312,43 +438,83 @@ func (s *Server) Shutdown(ctx context.Context) error {
438 }
439
440 func (s *Server) PolicyRuntime() *policy.Runtime {
441 + if s == nil || s.registry == nil {
442 + return nil
443 + }
444 return s.registry.policy
445 }
446
447 func (s *Server) PortalURL() string {
448 + if s == nil {
449 + return ""
450 + }
451 return s.cfg.PortalURL
452 }
453
454 +func (s *Server) wireGuardPeerPlaneEnabled() bool {
455 + if s == nil {
456 + return false
457 + }
458 + return strings.TrimSpace(s.cfg.WireGuardPrivateKey) != "" &&
459 + strings.TrimSpace(s.cfg.WireGuardEndpoint) != "" &&
460 + strings.TrimSpace(s.cfg.WireGuardPublicKey) != "" &&
461 + strings.TrimSpace(s.cfg.OverlayIPv4) != ""
462 +}
463 +
464 func (s *Server) LeaseSnapshots() []types.Lease {
465 + s.registry.mu.RLock()
466 + defer s.registry.mu.RUnlock()
467 +
468 now := time.Now()
324 - all := s.registry.activeAdminSnapshots()
325 - out := make([]types.Lease, 0, len(all))
326 - for _, snap := range all {
469 + records := make([]*leaseRecord, 0, len(s.registry.leasesByKey))
470 + for _, record := range s.registry.leasesByKey {
471 + records = append(records, record)
472 + }
473 + snapshots := make([]types.Lease, 0, len(records))
474 + for _, record := range records {
475 + if now.After(record.ExpiresAt) {
476 + continue
477 + }
478 + adminSnapshot := s.registry.AdminSnapshot(record)
479 since := time.Duration(0)
328 - if !snap.LastSeenAt.IsZero() {
329 - since = max(now.Sub(snap.LastSeenAt), 0)
480 + if !adminSnapshot.LastSeenAt.IsZero() {
481 + since = max(now.Sub(adminSnapshot.LastSeenAt), 0)
482 }
331 - if snap.IsBanned || snap.IsDenied || !snap.IsApproved || snap.Metadata.Hide {
483 + if adminSnapshot.IsBanned || adminSnapshot.IsDenied || !adminSnapshot.IsApproved || adminSnapshot.Metadata.Hide {
484 continue
485 }
334 - if snap.Ready == 0 && since >= 3*time.Minute {
486 + if adminSnapshot.Ready == 0 && since >= 3*time.Minute {
487 continue
488 }
337 - if snap.Metadata.Thumbnail == "" && s.thumbnails != nil {
338 - if _, _, ok := s.thumbnails.Get(snap.Hostname); ok {
339 - snap.Metadata.Thumbnail = types.PathThumbnailPrefix + snap.Hostname
340 - }
341 - }
342 - out = append(out, snap.Lease)
489 + snapshots = append(snapshots, adminSnapshot.Lease)
490 }
344 - return out
491 + return snapshots
492 }
493
494 func (s *Server) AdminLeaseSnapshots() []types.AdminLease {
348 - return s.registry.activeAdminSnapshots()
495 + s.registry.mu.RLock()
496 + defer s.registry.mu.RUnlock()
497 +
498 + now := time.Now()
499 + records := make([]*leaseRecord, 0, len(s.registry.leasesByKey))
500 + for _, record := range s.registry.leasesByKey {
501 + records = append(records, record)
502 + }
503 + snapshots := make([]types.AdminLease, 0, len(records))
504 + for _, record := range records {
505 + if now.After(record.ExpiresAt) {
506 + continue
507 + }
508 + snapshots = append(snapshots, s.registry.AdminSnapshot(record))
509 + }
510 + return snapshots
511 }
512
513 func (s *Server) LeaseSnapshotByHostname(hostname string) (types.Lease, bool) {
514 + if s == nil || s.registry == nil {
515 + return types.Lease{}, false
516 + }
517 +
518 record, ok := s.registry.Lookup(hostname)
519 if !ok || record == nil || time.Now().After(record.ExpiresAt) {
520 return types.Lease{}, false
@@ -356,6 +522,129 @@ func (s *Server) LeaseSnapshotByHostname(hostname string) (types.Lease, bool) {
522 return s.registry.Snapshot(record), true
523 }
524
525 +func (s *Server) startWireGuardPeerPlane() error {
526 + if s == nil {
527 + return nil
528 + }
529 + runtime, err := wireguard.NewRuntime(wireguard.RuntimeConfig{
530 + PrivateKey: s.cfg.WireGuardPrivateKey,
531 + Endpoint: s.cfg.WireGuardEndpoint,
532 + OverlayIPv4: s.cfg.OverlayIPv4,
533 + })
534 + if err != nil {
535 + return err
536 + }
537 +
538 + listener, err := runtime.ListenTCP(wireguard.DefaultPeerAPIHTTPPort)
539 + if err != nil {
540 + _ = runtime.Close()
541 + return fmt.Errorf("listen wireguard peer api: %w", err)
542 + }
543 +
544 + server := &http.Server{
545 + Handler: s.peerAPIHandler(),
546 + ReadHeaderTimeout: 10 * time.Second,
547 + }
548 +
549 + s.wgRuntime = runtime
550 + s.wgPeerListener = listener
551 + s.wgPeerServer = server
552 + if err := s.syncWireGuardPeers(); err != nil {
553 + _ = server.Close()
554 + _ = runtime.Close()
555 + s.wgRuntime = nil
556 + s.wgPeerListener = nil
557 + s.wgPeerServer = nil
558 + return err
559 + }
560 + return nil
561 +}
562 +
563 +func (s *Server) peerAPIHandler() http.Handler {
564 + mux := http.NewServeMux()
565 + mux.HandleFunc(types.PathRoot, s.handleRoot)
566 + mux.HandleFunc(types.PathHealthz, s.handleHealthz)
567 + if s.cfg.DiscoveryEnabled {
568 + mux.HandleFunc(types.PathDiscovery, s.handleRelayDiscovery)
569 + }
570 + return mux
571 +}
572 +
573 +func (s *Server) runWireGuardPeerAPIServer() error {
574 + if s == nil || s.wgPeerServer == nil || s.wgPeerListener == nil {
575 + return nil
576 + }
577 + err := s.wgPeerServer.Serve(s.wgPeerListener)
578 + if errors.Is(err, http.ErrServerClosed) || errors.Is(err, net.ErrClosed) {
579 + return nil
580 + }
581 + return err
582 +}
583 +
584 +func (s *Server) runWireGuardSyncLoop(ctx context.Context) error {
585 + if s.wgRuntime == nil {
586 + <-ctx.Done()
587 + return nil
588 + }
589 + ticker := time.NewTicker(30 * time.Second)
590 + defer ticker.Stop()
591 + for {
592 + if err := s.syncWireGuardPeers(); err != nil {
593 + log.Warn().Err(err).Msg("sync wireguard peers failed")
594 + }
595 + select {
596 + case <-ctx.Done():
597 + return nil
598 + case <-ticker.C:
599 + }
600 + }
601 +}
602 +
603 +func (s *Server) desiredWireGuardPeers() []types.DesiredPeer {
604 + if s.discoveryMgr == nil {
605 + return nil
606 + }
607 + descs := s.discoveryMgr.ActiveRelayDescriptors()
608 + if len(descs) == 0 {
609 + return nil
610 + }
611 + selfKey := s.identity.Key()
612 + peers := make([]types.DesiredPeer, 0, len(descs))
613 + for _, desc := range descs {
614 + nodeKey := relayNodeKey(desc)
615 + if nodeKey == "" || nodeKey == selfKey {
616 + continue
617 + }
618 + if !desc.SupportsOverlayPeer {
619 + continue
620 + }
621 + if strings.TrimSpace(desc.WireGuardPublicKey) == "" ||
622 + strings.TrimSpace(desc.WireGuardEndpoint) == "" ||
623 + strings.TrimSpace(desc.OverlayIPv4) == "" {
624 + continue
625 + }
626 + allowed := []string{desc.OverlayIPv4 + "/32"}
627 + if len(desc.OverlayCIDRs) > 0 {
628 + allowed = append(allowed, desc.OverlayCIDRs...)
629 + }
630 + peers = append(peers, types.DesiredPeer{
631 + RelayID: nodeKey,
632 + WireGuardPublicKey: desc.WireGuardPublicKey,
633 + WireGuardEndpoint: desc.WireGuardEndpoint,
634 + AllowedIPs: allowed,
635 + })
636 + }
637 + return peers
638 +}
639 +
640 +func (s *Server) syncWireGuardPeers() error {
641 + if s.wgRuntime == nil {
642 + return nil
643 + }
644 + peers := s.desiredWireGuardPeers()
645 + return s.wgRuntime.ApplyPeers(peers)
646 +}
647 +
648 func (s *Server) prepareAPITLS(ctx context.Context) (keyless.TLSMaterialConfig, *acme.Manager, error) {
649 acmeCfg := s.cfg.ACME
650 if baseDomain := utils.NormalizeHostname(acmeCfg.BaseDomain); baseDomain != "" && baseDomain != s.identity.Name {
@@ -426,7 +715,7 @@ func (s *Server) runSNIListener(ctx context.Context) error {
715 _ = wrappedConn.Close()
716 return
717 }
429 - BridgeConns(wrappedConn, upstream)
718 + s.BridgeConns(wrappedConn, upstream)
719 return
720 }
721
@@ -445,12 +734,15 @@ func (s *Server) runSNIListener(ctx context.Context) error {
734 return
735 }
736
448 - BridgeConns(wrappedConn, session)
737 + s.BridgeConns(wrappedConn, session)
738 }(conn)
739 case errors.Is(err, net.ErrClosed):
740 return nil
741 default:
453 - return err
742 + if ctxErr := ctx.Err(); ctxErr != nil {
743 + return ctxErr
744 + }
745 + return fmt.Errorf("accept sni connection: %w", err)
746 }
747 }
748 }
@@ -536,27 +828,106 @@ func (s *Server) runQUICTunnelListener(listener *quic.Listener) error {
828 }
829 }
830
539 -func BridgeConns(left, right net.Conn) {
831 +func (s *Server) runRelayDiscoveryLoop(ctx context.Context) error {
832 + if s.discoveryMgr == nil {
833 + <-ctx.Done()
834 + return nil
835 + }
836 + return s.discoveryMgr.Run(ctx, s.handleDiscoverySnapshot)
837 +}
838 +
839 +func (s *Server) handleDiscoverySnapshot(_ map[string]types.RelayState) {
840 + if s.discoveryMgr == nil || s.overlayPolicy == nil || !s.cfg.OverlayEnabled {
841 + return
842 + }
843 + descs := s.discoveryMgr.ActiveRelayDescriptors()
844 + if len(descs) == 0 {
845 + s.overlayRouteMu.Lock()
846 + s.overlayRoute = nil
847 + s.overlayRouteMu.Unlock()
848 + return
849 + }
850 +
851 + candidates := make([]uint32, 0, len(descs))
852 + for _, d := range descs {
853 + nodeKey := relayNodeKey(d)
854 + if nodeKey == "" {
855 + continue
856 + }
857 + candidates = append(candidates, crc32.ChecksumIEEE([]byte(nodeKey)))
858 + }
859 + if len(candidates) == 0 {
860 + return
861 + }
862 + selfKey := strings.TrimSpace(s.cfg.PortalURL)
863 + if selfKey == "" {
864 + selfKey = s.identity.Key()
865 + }
866 + if selfKey == "" {
867 + return
868 + }
869 + selfID := crc32.ChecksumIEEE([]byte(selfKey))
870 + route, err := s.overlayPolicy.BuildRouteWithLoad(selfID, candidates, s.cfg.OverlayMaxHops, s.weightMgr.Collect(), s.cfg.OverlayCongestion)
871 + if err != nil {
872 + return
873 + }
874 + s.overlayRouteMu.Lock()
875 + s.overlayRoute = route
876 + s.overlayRouteMu.Unlock()
877 +}
878 +
879 +func (s *Server) OverlayRoute() []uint32 {
880 + if s == nil {
881 + return nil
882 + }
883 + s.overlayRouteMu.RLock()
884 + defer s.overlayRouteMu.RUnlock()
885 + if len(s.overlayRoute) == 0 {
886 + return nil
887 + }
888 + out := make([]uint32, len(s.overlayRoute))
889 + copy(out, s.overlayRoute)
890 + return out
891 +}
892 +
893 +func relayNodeKey(desc types.RelayDescriptor) string {
894 + if key := strings.TrimSpace(desc.RelayID); key != "" {
895 + return key
896 + }
897 + if key := strings.TrimSpace(desc.APIHTTPSAddr); key != "" {
898 + return key
899 + }
900 + return desc.Key()
901 +}
902 +
903 +func (s *Server) BridgeConns(left, right net.Conn) {
904 + s.loadMgr.RecordConnStart()
905 + defer s.loadMgr.RecordConnEnd()
906 +
907 defer left.Close()
908 defer right.Close()
909
543 - type closeWriter interface {
544 - CloseWrite() error
545 - }
910 var group errgroup.Group
911 group.Go(func() error {
548 - _, err := io.Copy(right, left)
549 - if cw, ok := right.(closeWriter); ok {
550 - _ = cw.CloseWrite()
551 - }
912 + n, err := io.Copy(right, left)
913 + s.loadMgr.RecordBytesIn(n)
914 + closeWrite(right)
915 return err
916 })
917 group.Go(func() error {
555 - _, err := io.Copy(left, right)
556 - if cw, ok := left.(closeWriter); ok {
557 - _ = cw.CloseWrite()
558 - }
918 + n, err := io.Copy(left, right)
919 + s.loadMgr.RecordBytesOut(n)
920 + closeWrite(left)
921 return err
922 })
923 _ = group.Wait()
924 }
925 +
926 +func closeWrite(conn net.Conn) {
927 + type closeWriter interface {
928 + CloseWrite() error
929 + }
930 + if cw, ok := conn.(closeWriter); ok {
931 + _ = cw.CloseWrite()
932 + }
933 +}
portal/server_test.go
+9 -9
@@ -36,6 +36,7 @@ func mustRelayDescriptor(t *testing.T, relayURL string) types.RelayDescriptor {
36 Identity: types.Identity{
37 Name: utils.PortalRootHost(relayURL),
38 },
39 + RelayID: relayURL,
40 Sequence: uint64(now.UnixMilli()),
41 Version: 1,
42 IssuedAt: now,
@@ -348,17 +349,15 @@ func TestServerStartDiscoveryIncludesIdentityAndOmitsSignerFields(t *testing.T)
349 if err != nil {
350 t.Fatalf("read /discovery response: %v", err)
351 }
351 - for _, key := range []string{"\"address\"", "\"name\"", "signer_public_key", "descriptor_signature"} {
352 - if key == "\"address\"" || key == "\"name\"" {
353 - if !strings.Contains(string(body), key) {
354 - t.Fatalf("/discovery body = %q, want %q present", string(body), key)
355 - }
356 - continue
357 - }
358 - if strings.Contains(string(body), key) {
359 - t.Fatalf("/discovery body = %q, want %q omitted", string(body), key)
352 + bodyText := string(body)
353 + for _, key := range []string{"\"address\"", "\"name\"", "\"relay_id\"", "\"owner_address\"", "\"signer_public_key\""} {
354 + if !strings.Contains(bodyText, key) {
355 + t.Fatalf("/discovery body = %q, want %q present", bodyText, key)
356 }
357 }
358 + if strings.Contains(bodyText, "descriptor_signature") {
359 + t.Fatalf("/discovery body = %q, want descriptor_signature omitted", bodyText)
360 + }
361 }
362
363 func TestServerStartRejectsMismatchedACMEBaseDomain(t *testing.T) {
@@ -532,6 +531,7 @@ func TestServerDiscoverySkipsSelfRelayHint(t *testing.T) {
531 bootstrapDesc := mustRelayDescriptor(t, "https://bootstrap.example.com")
532 selfHint, err := discovery.NormalizeDescriptor(types.RelayDescriptor{
533 Identity: server.identity.Copy(),
534 + RelayID: "https://self-mirror.example.com",
535 Sequence: uint64(now.UnixMilli()),
536 Version: 1,
537 IssuedAt: now,
portal/wireguard/runtime.go new
+245
@@ -0,0 +1,245 @@
1 +package wireguard
2 +
3 +import (
4 + "context"
5 + "encoding/json"
6 + "errors"
7 + "fmt"
8 + "net"
9 + "net/http"
10 + "net/netip"
11 + "net/url"
12 + "strconv"
13 + "strings"
14 + "sync"
15 + "time"
16 +
17 + "golang.zx2c4.com/wireguard/conn"
18 + "golang.zx2c4.com/wireguard/device"
19 + "golang.zx2c4.com/wireguard/tun/netstack"
20 +
21 + "github.com/gosuda/portal-tunnel/v2/types"
22 + "github.com/gosuda/portal-tunnel/v2/utils"
23 +)
24 +
25 +const (
26 + DefaultMTU = 1420
27 + DefaultPeerAPIHTTPPort = 7777
28 + DefaultPersistentKeepalive = 25
29 +
30 + defaultPeerRequestTimeout = 15 * time.Second
31 +)
32 +
33 +type RuntimeConfig struct {
34 + PrivateKey string
35 + Endpoint string
36 + OverlayIPv4 string
37 + MTU int
38 +}
39 +
40 +type Runtime struct {
41 + device *device.Device
42 + net *netstack.Net
43 + overlayIP netip.Addr
44 +
45 + mu sync.Mutex
46 + closed bool
47 +}
48 +
49 +func NewRuntime(cfg RuntimeConfig) (*Runtime, error) {
50 + privateKey, err := utils.NormalizeWireGuardPrivateKey(cfg.PrivateKey)
51 + if err != nil {
52 + return nil, fmt.Errorf("normalize wireguard private key: %w", err)
53 + }
54 + listenPort, err := utils.WireGuardListenPort(cfg.Endpoint)
55 + if err != nil {
56 + return nil, err
57 + }
58 +
59 + ip, err := netip.ParseAddr(strings.TrimSpace(cfg.OverlayIPv4))
60 + if err != nil || !ip.Is4() {
61 + return nil, errors.New("overlay ipv4 must be a valid IPv4 address")
62 + }
63 +
64 + mtu := cfg.MTU
65 + if mtu <= 0 {
66 + mtu = DefaultMTU
67 + }
68 +
69 + tunDev, network, err := netstack.CreateNetTUN([]netip.Addr{ip}, nil, mtu)
70 + if err != nil {
71 + return nil, fmt.Errorf("create netstack tun: %w", err)
72 + }
73 +
74 + wgDevice := device.NewDevice(tunDev, conn.NewDefaultBind(), device.NewLogger(device.LogLevelError, "portal-wg"))
75 + privateKeyHex, err := utils.WireGuardKeyHex(privateKey)
76 + if err != nil {
77 + wgDevice.Close()
78 + <-wgDevice.Wait()
79 + return nil, err
80 + }
81 +
82 + config := fmt.Sprintf("private_key=%s\nlisten_port=%d\n", privateKeyHex, listenPort)
83 + if err := wgDevice.IpcSet(config); err != nil {
84 + wgDevice.Close()
85 + <-wgDevice.Wait()
86 + return nil, fmt.Errorf("configure wireguard device: %w", err)
87 + }
88 + if err := wgDevice.Up(); err != nil {
89 + wgDevice.Close()
90 + <-wgDevice.Wait()
91 + return nil, fmt.Errorf("bring wireguard device up: %w", err)
92 + }
93 +
94 + return &Runtime{
95 + device: wgDevice,
96 + net: network,
97 + overlayIP: ip,
98 + }, nil
99 +}
100 +
101 +func (r *Runtime) ListenTCP(port int) (net.Listener, error) {
102 + if r == nil || r.net == nil {
103 + return nil, errors.New("wireguard runtime not initialized")
104 + }
105 + return r.net.ListenTCP(&net.TCPAddr{
106 + IP: net.ParseIP(r.overlayIP.String()),
107 + Port: port,
108 + })
109 +}
110 +
111 +func (r *Runtime) DialContext(ctx context.Context, network, address string) (net.Conn, error) {
112 + if r == nil || r.net == nil {
113 + return nil, errors.New("wireguard runtime not initialized")
114 + }
115 + switch network {
116 + case "tcp", "tcp4", "tcp6":
117 + default:
118 + return nil, fmt.Errorf("unsupported network %q", network)
119 + }
120 +
121 + host, portText, err := net.SplitHostPort(address)
122 + if err != nil {
123 + return nil, err
124 + }
125 + ip, err := netip.ParseAddr(strings.Trim(host, "[]"))
126 + if err != nil {
127 + return nil, err
128 + }
129 + port, err := strconv.Atoi(portText)
130 + if err != nil || port <= 0 || port > 65535 {
131 + return nil, errors.New("invalid tcp port")
132 + }
133 + return r.net.DialContextTCPAddrPort(ctx, netip.AddrPortFrom(ip, uint16(port)))
134 +}
135 +
136 +func (r *Runtime) Discover(ctx context.Context, overlayIPv4 string, port int) (types.DiscoveryResponse, error) {
137 + var resp types.DiscoveryResponse
138 + if err := r.getPeerJSON(ctx, overlayIPv4, port, types.PathDiscovery, &resp); err != nil {
139 + return types.DiscoveryResponse{}, err
140 + }
141 + return resp, nil
142 +}
143 +
144 +func (r *Runtime) getPeerJSON(ctx context.Context, overlayIPv4 string, port int, path string, out any) error {
145 + if r == nil {
146 + return errors.New("wireguard runtime not initialized")
147 + }
148 + if port == 0 {
149 + port = DefaultPeerAPIHTTPPort
150 + }
151 + ip, err := netip.ParseAddr(strings.TrimSpace(overlayIPv4))
152 + if err != nil || !ip.Is4() {
153 + return errors.New("overlay ipv4 must be a valid IPv4 address")
154 + }
155 +
156 + baseURL := &url.URL{
157 + Scheme: "http",
158 + Host: net.JoinHostPort(ip.String(), strconv.Itoa(port)),
159 + Path: path,
160 + }
161 +
162 + httpReq, err := http.NewRequestWithContext(ctx, http.MethodGet, baseURL.String(), nil)
163 + if err != nil {
164 + return err
165 + }
166 +
167 + client := &http.Client{
168 + Transport: &http.Transport{
169 + DialContext: r.DialContext,
170 + ForceAttemptHTTP2: false,
171 + },
172 + Timeout: defaultPeerRequestTimeout,
173 + }
174 +
175 + resp, err := client.Do(httpReq)
176 + if err != nil {
177 + return err
178 + }
179 + defer resp.Body.Close()
180 +
181 + if resp.StatusCode < http.StatusOK || resp.StatusCode >= http.StatusMultipleChoices {
182 + return utils.DecodeAPIRequestError(resp)
183 + }
184 + if out == nil {
185 + return nil
186 + }
187 + return json.NewDecoder(resp.Body).Decode(out)
188 +}
189 +
190 +func (r *Runtime) ApplyPeers(peers []types.DesiredPeer) error {
191 + if r == nil || r.device == nil {
192 + return errors.New("wireguard runtime not initialized")
193 + }
194 +
195 + var builder strings.Builder
196 + builder.WriteString("replace_peers=true\n")
197 +
198 + for _, peer := range peers {
199 + publicKeyHex, err := utils.WireGuardKeyHex(peer.WireGuardPublicKey)
200 + if err != nil {
201 + return fmt.Errorf("normalize peer %q public key: %w", peer.RelayID, err)
202 + }
203 + builder.WriteString("public_key=")
204 + builder.WriteString(publicKeyHex)
205 + builder.WriteByte('\n')
206 + if endpoint := strings.TrimSpace(peer.WireGuardEndpoint); endpoint != "" {
207 + builder.WriteString("endpoint=")
208 + builder.WriteString(endpoint)
209 + builder.WriteByte('\n')
210 + }
211 +
212 + allowedIPs := utils.NormalizeIPPrefixes(peer.AllowedIPs)
213 + for _, allowedIP := range allowedIPs {
214 + builder.WriteString("allowed_ip=")
215 + builder.WriteString(allowedIP)
216 + builder.WriteByte('\n')
217 + }
218 + if DefaultPersistentKeepalive > 0 {
219 + builder.WriteString("persistent_keepalive_interval=")
220 + builder.WriteString(strconv.Itoa(DefaultPersistentKeepalive))
221 + builder.WriteByte('\n')
222 + }
223 + }
224 +
225 + return r.device.IpcSet(builder.String())
226 +}
227 +
228 +func (r *Runtime) Close() error {
229 + if r == nil || r.device == nil {
230 + return nil
231 + }
232 +
233 + r.mu.Lock()
234 + if r.closed {
235 + r.mu.Unlock()
236 + return nil
237 + }
238 + r.closed = true
239 + device := r.device
240 + r.mu.Unlock()
241 +
242 + device.Close()
243 + <-device.Wait()
244 + return nil
245 +}
sdk/expose_test.go
+1
@@ -17,6 +17,7 @@ func mustRelayDescriptor(t *testing.T, relayName, relayURL string) types.RelayDe
17 Identity: types.Identity{
18 Name: relayName,
19 },
20 + RelayID: relayURL,
21 Sequence: uint64(now.UnixMilli()),
22 Version: 1,
23 IssuedAt: now,
types/identity.go
+19 -10
@@ -79,16 +79,25 @@ type AdminLease struct {
79 type RelayDescriptor struct {
80 Identity
81
82 - Sequence uint64 `json:"sequence"`
83 - Version uint32 `json:"version"`
84 - IssuedAt time.Time `json:"issued_at"`
85 - ExpiresAt time.Time `json:"expires_at"`
86 -
87 - APIHTTPSAddr string `json:"api_https_addr"`
88 - IngressTLSAddr string `json:"ingress_tls_addr,omitempty"`
89 -
90 - SupportsUDP bool `json:"supports_udp,omitempty"`
91 - SupportsTCP bool `json:"supports_tcp,omitempty"`
82 + RelayID string `json:"relay_id,omitempty"`
83 + OwnerAddress string `json:"owner_address,omitempty"`
84 + SignerPublicKey string `json:"signer_public_key,omitempty"`
85 + Sequence uint64 `json:"sequence"`
86 + Version uint32 `json:"version"`
87 + IssuedAt time.Time `json:"issued_at"`
88 + ExpiresAt time.Time `json:"expires_at"`
89 + APIHTTPSAddr string `json:"api_https_addr"`
90 + IngressTLSAddr string `json:"ingress_tls_addr,omitempty"`
91 + WireGuardPublicKey string `json:"wireguard_public_key,omitempty"`
92 + WireGuardEndpoint string `json:"wireguard_endpoint,omitempty"`
93 + OverlayIPv4 string `json:"overlay_ipv4,omitempty"`
94 + OverlayCIDRs []string `json:"overlay_cidrs,omitempty"`
95 + SupportsUDP bool `json:"supports_udp,omitempty"`
96 + SupportsTCP bool `json:"supports_tcp,omitempty"`
97 + SupportsOverlayPeer bool `json:"supports_overlay_peer,omitempty"`
98 + Load float64 `json:"load,omitempty"`
99 + LoadScore float64 `json:"load_score,omitempty"`
100 + LastUpdated int64 `json:"last_updated,omitempty"`
101 }
102
103 const DiscoveryPollInterval = 1 * time.Minute
types/overlay.go new
+10
@@ -0,0 +1,10 @@
1 +package types
2 +
3 +// DesiredPeer describes a WireGuard peer that should be programmed into the
4 +// local runtime.
5 +type DesiredPeer struct {
6 + RelayID string `json:"relay_id"`
7 + WireGuardPublicKey string `json:"wireguard_public_key"`
8 + WireGuardEndpoint string `json:"wireguard_endpoint"`
9 + AllowedIPs []string `json:"allowed_ips,omitempty"`
10 +}
utils/wireguard.go new
+105
@@ -0,0 +1,105 @@
1 +package utils
2 +
3 +import (
4 + "crypto/sha256"
5 + "encoding/base64"
6 + "encoding/hex"
7 + "errors"
8 + "net"
9 + "net/netip"
10 + "strconv"
11 + "strings"
12 +
13 + "golang.org/x/crypto/curve25519"
14 +)
15 +
16 +func NormalizeWireGuardPrivateKey(raw string) (string, error) {
17 + key, err := decodeWireGuardKey(raw)
18 + if err != nil {
19 + return "", err
20 + }
21 + clampWireGuardPrivateKey(&key)
22 + return base64.StdEncoding.EncodeToString(key[:]), nil
23 +}
24 +
25 +func WireGuardPublicKeyFromPrivate(raw string) (string, error) {
26 + privateKey, err := decodeWireGuardKey(raw)
27 + if err != nil {
28 + return "", err
29 + }
30 + clampWireGuardPrivateKey(&privateKey)
31 + var publicKey [32]byte
32 + curve25519.ScalarBaseMult(&publicKey, &privateKey)
33 + return base64.StdEncoding.EncodeToString(publicKey[:]), nil
34 +}
35 +
36 +func WireGuardListenPort(rawEndpoint string) (int, error) {
37 + endpoint := strings.TrimSpace(rawEndpoint)
38 + if endpoint == "" {
39 + return 0, errors.New("wireguard endpoint is required")
40 + }
41 + _, portText, err := net.SplitHostPort(endpoint)
42 + if err != nil {
43 + return 0, errors.New("wireguard endpoint must be host:port")
44 + }
45 + port, err := strconv.Atoi(portText)
46 + if err != nil || port <= 0 || port > 65535 {
47 + return 0, errors.New("wireguard endpoint port is invalid")
48 + }
49 + return port, nil
50 +}
51 +
52 +func DeriveWireGuardOverlayIPv4(publicKey string) (string, error) {
53 + decoded, err := base64.StdEncoding.DecodeString(strings.TrimSpace(publicKey))
54 + if err != nil {
55 + return "", errors.New("wireguard public key must be base64 encoded")
56 + }
57 + if len(decoded) != 32 {
58 + return "", errors.New("wireguard public key must be 32 bytes")
59 + }
60 +
61 + sum := sha256.Sum256(decoded)
62 + return netip.AddrFrom4([4]byte{
63 + 100,
64 + 64 + (sum[0] & 0x3f),
65 + sum[1],
66 + 1 + (sum[2] % 254),
67 + }).String(), nil
68 +}
69 +
70 +func WireGuardKeyHex(raw string) (string, error) {
71 + key, err := decodeWireGuardKey(raw)
72 + if err != nil {
73 + return "", err
74 + }
75 + return hex.EncodeToString(key[:]), nil
76 +}
77 +
78 +func decodeWireGuardKey(raw string) ([32]byte, error) {
79 + var key [32]byte
80 + value := strings.TrimSpace(raw)
81 + if value == "" {
82 + return key, errors.New("wireguard key is required")
83 + }
84 +
85 + var decoded []byte
86 + var err error
87 + if len(value) == 64 && !strings.Contains(value, "=") {
88 + decoded, err = hex.DecodeString(value)
89 + } else {
90 + decoded, err = base64.StdEncoding.DecodeString(value)
91 + }
92 + if err != nil {
93 + return key, errors.New("wireguard key must be base64 or hex encoded")
94 + }
95 + if len(decoded) != len(key) {
96 + return key, errors.New("wireguard key must be 32 bytes")
97 + }
98 + copy(key[:], decoded)
99 + return key, nil
100 +}
101 +
102 +func clampWireGuardPrivateKey(key *[32]byte) {
103 + key[0] &= 248
104 + key[31] = (key[31] & 127) | 64
105 +}