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
+}