refact: remove wireguard and simplify discovery

Kim committed Apr 3, 2026 at 17:30 UTC 3fe29d44a13a33b08e8ee614b0d4c5e146c4adcc
26 files changed +856 -2076
.env.example
-2
@@ -3,8 +3,6 @@ PORTAL_URL=https://localhost:4017
3 BOOTSTRAPS=https://localhost:4017
4 DISCOVERY=true
5 IDENTITY_PATH=/portal-certs/identity.json
6 -WIREGUARD_ENDPOINT=
7 -WIREGUARD_PRIVATE_KEY=
6
7 # Listener ports
8 API_PORT=4017
cmd/relay-server/main.go
+31 -40
@@ -31,36 +31,34 @@ 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 - WireGuardPrivateKey string
46 - DiscoveryPort int
47 - WireGuardEndpoint string
48 - AdminSecretKey string
49 - TrustProxyHeaders bool
50 - TrustedProxyCIDRs string
51 - AdminSettingsPath string
52 - KeylessDir string
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 + IdentityPath string
45 + AdminSecretKey string
46 + TrustProxyHeaders bool
47 + TrustedProxyCIDRs string
48 + AdminSettingsPath string
49 + KeylessDir string
50 +
51 + ACMEDNSProvider string
52 + ENSGaslessEnabled bool
53 + CloudflareToken string
54 + GCPProjectID string
55 + GCPManagedZone string
56 + AWSAccessKeyID string
57 + AWSSecretAccessKey string
58 + AWSSessionToken string
59 + AWSRegion string
60 + AWSHostedZoneID string
61 + AWSDNSSECKMSKeyARN string
62 }
63
64 func runServeCommand(args []string) error {
@@ -78,9 +76,6 @@ func runServeCommand(args []string) error {
76 utils.StringFlagEnv(fs, &cfg.Bootstraps, "bootstraps", "", "additional bootstrap relay API URLs used for discovery expansion", "BOOTSTRAPS")
77 utils.BoolFlagEnv(fs, &cfg.DiscoveryEnabled, "discovery", false, "serve relay discovery endpoints and poll discovery peers", "DISCOVERY")
78 utils.StringFlagEnv(fs, &cfg.IdentityPath, "identity-path", "identity.json", "relay identity json file path", "IDENTITY_PATH")
81 - utils.StringFlagEnv(fs, &cfg.WireGuardPrivateKey, "wireguard-private-key", "", "wireguard private key for relay peer overlay", "WIREGUARD_PRIVATE_KEY")
82 - utils.IntFlagEnv(fs, &cfg.DiscoveryPort, "discovery-port", 0, utils.ParsePortNumber, "public UDP listen port advertised for relay-peer discovery overlay (defaults to 51820 when wireguard is enabled)", "DISCOVERY_PORT")
83 - utils.StringFlagEnv(fs, &cfg.WireGuardEndpoint, "wireguard-endpoint", "", "explicit public WireGuard endpoint advertised for relay peer overlay (host:port or ip:port); defaults to PORTAL_URL host + DISCOVERY_PORT when empty", "WIREGUARD_ENDPOINT")
79 utils.StringFlagEnv(fs, &cfg.AdminSecretKey, "admin-secret-key", "", "admin auth secret", "ADMIN_SECRET_KEY")
80 utils.BoolFlagEnv(fs, &cfg.TrustProxyHeaders, "trust-proxy-headers", false, "trust X-Forwarded-* and X-Real-IP headers from trusted proxies", "TRUST_PROXY_HEADERS")
81 utils.StringFlagEnv(fs, &cfg.TrustedProxyCIDRs, "trusted-proxy-cidrs", "", "trusted proxy CIDR allowlist for forwarded headers, comma-separated; defaults to private/loopback proxy ranges when trust-proxy-headers is enabled", "TRUSTED_PROXY_CIDRS")
@@ -121,7 +116,6 @@ func runServeCommand(args []string) error {
116 Bool("discovery_enabled", cfg.DiscoveryEnabled).
117 Str("acme_dns_provider", cfg.ACMEDNSProvider).
118 Bool("ens_gasless_enabled", cfg.ENSGaslessEnabled).
124 - Bool("wireguard_enabled", strings.TrimSpace(cfg.WireGuardPrivateKey) != "").
119 Bool("udp_enabled", cfg.UDPEnabled).
120 Bool("tcp_enabled", cfg.TCPEnabled).
121 Msg("configured relay server")
@@ -139,12 +133,9 @@ func runServer(ctx context.Context, cfg relayServerConfig) error {
133 }
134
135 server, err := portal.NewServer(portal.ServerConfig{
142 - PortalURL: cfg.PortalURL,
143 - IdentityPath: cfg.IdentityPath,
144 - Bootstraps: bootstraps,
145 - WireGuardPrivateKey: cfg.WireGuardPrivateKey,
146 - DiscoveryPort: cfg.DiscoveryPort,
147 - WireGuardEndpoint: cfg.WireGuardEndpoint,
136 + PortalURL: cfg.PortalURL,
137 + IdentityPath: cfg.IdentityPath,
138 + Bootstraps: bootstraps,
139 ACME: acme.Config{
140 KeyDir: cfg.KeylessDir,
141 DNSProvider: cfg.ACMEDNSProvider,
docker-compose.yml
-4
@@ -9,7 +9,6 @@ services:
9 - "${API_PORT:-4017}:${API_PORT:-4017}"
10 - "${SNI_PORT:-443}:${SNI_PORT:-443}"
11 # Uncomment for UDP backhaul, public UDP lease ports, and raw TCP lease ports as needed.
12 - # - "${DISCOVERY_PORT:-51820}:${DISCOVERY_PORT:-51820}/udp"
12 # - "${SNI_PORT:-443}:${SNI_PORT:-443}/udp"
13 # - "${MIN_PORT:-40000}-${MAX_PORT:-40009}:${MIN_PORT:-40000}-${MAX_PORT:-40009}/udp"
14 # - "${MIN_PORT:-40000}-${MAX_PORT:-40009}:${MIN_PORT:-40000}-${MAX_PORT:-40009}"
@@ -19,9 +18,6 @@ services:
18 BOOTSTRAPS: ${BOOTSTRAPS:-}
19 DISCOVERY: ${DISCOVERY:-true}
20 IDENTITY_PATH: ${IDENTITY_PATH:-/portal-certs/identity.json}
22 - WIREGUARD_PRIVATE_KEY: ${WIREGUARD_PRIVATE_KEY:-}
23 - DISCOVERY_PORT: ${DISCOVERY_PORT:-51820}
24 - WIREGUARD_ENDPOINT: ${WIREGUARD_ENDPOINT:-}
21
22 # Listener ports (published to the host below)
23 API_PORT: ${API_PORT:-4017}
docs/deployment.md
-5
@@ -204,7 +204,6 @@ PORTAL_URL=https://example.com
204 BOOTSTRAPS=
205 DISCOVERY=true
206 IDENTITY_PATH=/portal-certs/identity.json
207 -WIREGUARD_ENDPOINT=
207 SNI_PORT=443
208 ADMIN_SECRET_KEY=your-admin-secret
209 KEYLESS_DIR=/portal-certs
@@ -226,7 +225,6 @@ PORTAL_URL=https://example.com
225 BOOTSTRAPS=
226 DISCOVERY=true
227 IDENTITY_PATH=/portal-certs/identity.json
229 -WIREGUARD_ENDPOINT=
228 SNI_PORT=443
229 ADMIN_SECRET_KEY=your-admin-secret
230 KEYLESS_DIR=/portal-certs
@@ -244,7 +242,6 @@ PORTAL_URL=https://example.com
242 BOOTSTRAPS=
243 DISCOVERY=true
244 IDENTITY_PATH=/portal-certs/identity.json
247 -WIREGUARD_ENDPOINT=
245 SNI_PORT=443
246 ADMIN_SECRET_KEY=your-admin-secret
247 KEYLESS_DIR=/portal-certs
@@ -289,8 +286,6 @@ Notes:
286
287 - For non-apex deployments, set `PORTAL_URL` to the non-apex host value, for example `https://portal.example.com:8443`
288 - Portal uses the `PORTAL_URL` host for public lease hostnames
292 -- `WIREGUARD_ENDPOINT` is optional. When empty, Portal advertises `PORTAL_URL` host with `DISCOVERY_PORT`
293 -- Set `WIREGUARD_ENDPOINT` explicitly only when relay-peer discovery UDP is exposed on a different address than `PORTAL_URL`
289 - `IDENTITY_PATH` stores the relay identity JSON inside the container
290 - `KEYLESS_DIR` stores relay certificate material inside the container
291 - The Docker Compose stack stores relay identity JSON and certificate state under `./.portal-certs` on the host
docs/examples/nginx-proxy-multi-service/.env.example
-2
@@ -6,8 +6,6 @@ PORTAL_URL=https://portal.example.com
6 BOOTSTRAPS=
7 DISCOVERY=true
8 IDENTITY_PATH=/portal-certs/identity.json
9 -WIREGUARD_ENDPOINT=
10 -WIREGUARD_PRIVATE_KEY=
9
10 # Listener ports
11 API_PORT=4017
docs/examples/nginx-proxy-multi-service/docker-compose.yaml
-5
@@ -54,7 +54,6 @@ services:
54 # NAT-traversal relay server.
55 # TCP (4017, 4443) is reached by nginx via host.docker.internal.
56 # SNI_PORT is 4443 to avoid conflicting with nginx on 443.
57 - # If you enable relay-peer discovery overlay, expose DISCOVERY_PORT/udp as well.
57 # If you enable UDP, expose SNI_PORT/udp and MIN_PORT-MAX_PORT/udp as well.
58 # If you enable raw TCP transport, expose MIN_PORT-MAX_PORT/tcp as well.
59 portal:
@@ -63,7 +62,6 @@ services:
62 ports:
63 - "${API_PORT:-4017}:${API_PORT:-4017}/tcp"
64 - "${SNI_PORT:-4443}:${SNI_PORT:-4443}/tcp"
66 - - "${DISCOVERY_PORT:-51820}:${DISCOVERY_PORT:-51820}/udp"
65 # Uncomment below when enabling UDP transport (host SNI_PORT/udp is free — nginx only uses 443/tcp):
66 # - "${SNI_PORT:-4443}:${SNI_PORT:-4443}/udp"
67 # - "${MIN_PORT:-40000}-${MAX_PORT:-40009}:${MIN_PORT:-40000}-${MAX_PORT:-40009}/udp"
@@ -77,9 +75,6 @@ services:
75 API_PORT: ${API_PORT:-4017}
76 SNI_PORT: ${SNI_PORT:-4443}
77 IDENTITY_PATH: ${IDENTITY_PATH:-/portal-certs/identity.json}
80 - WIREGUARD_PRIVATE_KEY: ${WIREGUARD_PRIVATE_KEY:-}
81 - DISCOVERY_PORT: ${DISCOVERY_PORT:-51820}
82 - WIREGUARD_ENDPOINT: ${WIREGUARD_ENDPOINT:-}
78 MIN_PORT: ${MIN_PORT:-0}
79 MAX_PORT: ${MAX_PORT:-0}
80 UDP_ENABLED: ${UDP_ENABLED:-false}
docs/examples/nginx-proxy/.env.example
-2
@@ -6,8 +6,6 @@ PORTAL_URL=https://portal.example.com
6 BOOTSTRAPS=
7 DISCOVERY=true
8 IDENTITY_PATH=/portal-certs/identity.json
9 -WIREGUARD_ENDPOINT=
10 -WIREGUARD_PRIVATE_KEY=
9
10 # Listener ports
11 API_PORT=4017
docs/examples/nginx-proxy/docker-compose.yaml
-5
@@ -47,7 +47,6 @@ services:
47 # NAT-traversal relay server.
48 # TCP (4017, 4443) is reached by nginx via 127.0.0.1.
49 # SNI_PORT is set to 4443 to avoid conflicting with nginx on port 443.
50 - # If you enable relay-peer discovery overlay, expose DISCOVERY_PORT/udp as well.
50 # If you enable UDP, expose SNI_PORT/udp and MIN_PORT-MAX_PORT/udp as well.
51 # If you enable raw TCP transport, expose MIN_PORT-MAX_PORT/tcp as well.
52 portal:
@@ -56,7 +55,6 @@ services:
55 ports:
56 - "${API_PORT:-4017}:${API_PORT:-4017}/tcp"
57 - "${SNI_PORT:-4443}:${SNI_PORT:-4443}/tcp"
59 - - "${DISCOVERY_PORT:-51820}:${DISCOVERY_PORT:-51820}/udp"
58 # Uncomment below when enabling UDP transport (host SNI_PORT/udp is free — nginx only uses 443/tcp):
59 # - "${SNI_PORT:-4443}:${SNI_PORT:-4443}/udp"
60 # - "${MIN_PORT:-40000}-${MAX_PORT:-40009}:${MIN_PORT:-40000}-${MAX_PORT:-40009}/udp"
@@ -73,9 +71,6 @@ services:
71 # Use a non-443 port to avoid conflict with nginx on the host.
72 SNI_PORT: ${SNI_PORT:-4443}
73 IDENTITY_PATH: ${IDENTITY_PATH:-/portal-certs/identity.json}
76 - WIREGUARD_PRIVATE_KEY: ${WIREGUARD_PRIVATE_KEY:-}
77 - DISCOVERY_PORT: ${DISCOVERY_PORT:-51820}
78 - WIREGUARD_ENDPOINT: ${WIREGUARD_ENDPOINT:-}
74
75 MIN_PORT: ${MIN_PORT:-0}
76 MAX_PORT: ${MAX_PORT:-0}
go.mod
+2 -7
@@ -10,7 +10,7 @@ require (
10 github.com/aws/aws-sdk-go-v2/service/route53 v1.62.1
11 github.com/decred/dcrd/dcrec/secp256k1/v4 v4.1.0
12 github.com/go-acme/lego/v4 v4.32.0
13 - github.com/go-jose/go-jose/v4 v4.1.3
13 + github.com/go-jose/go-jose/v4 v4.1.4
14 github.com/gosuda/keyless_tls v0.0.1-0.20260304212324-7733f8366abc
15 github.com/quic-go/quic-go v0.59.0
16 github.com/rs/zerolog v1.34.0
@@ -19,7 +19,6 @@ require (
19 golang.org/x/net v0.51.0
20 golang.org/x/oauth2 v0.35.0
21 golang.org/x/sync v0.19.0
22 - golang.zx2c4.com/wireguard v0.0.0-20250521234502-f333402bd9cb
22 google.golang.org/api v0.267.0
23 )
24
@@ -45,7 +44,6 @@ require (
44 github.com/felixge/httpsnoop v1.0.4 // indirect
45 github.com/go-logr/logr v1.4.3 // indirect
46 github.com/go-logr/stdr v1.2.2 // indirect
48 - github.com/google/btree v1.1.2 // indirect
47 github.com/google/s2a-go v0.1.9 // indirect
48 github.com/google/uuid v1.6.0 // indirect
49 github.com/googleapis/enterprise-certificate-proxy v0.3.11 // indirect
@@ -63,11 +61,8 @@ require (
61 golang.org/x/mod v0.32.0 // indirect
62 golang.org/x/sys v0.41.0 // indirect
63 golang.org/x/text v0.34.0 // indirect
66 - golang.org/x/time v0.14.0 // indirect
64 golang.org/x/tools v0.41.0 // indirect
68 - golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2 // indirect
65 google.golang.org/genproto/googleapis/rpc v0.0.0-20260203192932-546029d2fa20 // indirect
70 - google.golang.org/grpc v1.78.0 // indirect
66 + google.golang.org/grpc v1.79.3 // indirect
67 google.golang.org/protobuf v1.36.11 // indirect
72 - gvisor.dev/gvisor v0.0.0-20250503011706-39ed1f5ac29c // indirect
68 )
go.sum
+5 -8
@@ -57,6 +57,8 @@ github.com/go-acme/lego/v4 v4.32.0 h1:z7Ss7aa1noabhKj+DBzhNCO2SM96xhE3b0ucVW3x8T
57 github.com/go-acme/lego/v4 v4.32.0/go.mod h1:lI2fZNdgeM/ymf9xQ9YKbgZm6MeDuf91UrohMQE4DhI=
58 github.com/go-jose/go-jose/v4 v4.1.3 h1:CVLmWDhDVRa6Mi/IgCgaopNosCaHz7zrMeF9MlZRkrs=
59 github.com/go-jose/go-jose/v4 v4.1.3/go.mod h1:x4oUasVrzR7071A4TnHLGSPpNOm2a21K9Kf04k1rs08=
60 +github.com/go-jose/go-jose/v4 v4.1.4 h1:moDMcTHmvE6Groj34emNPLs/qtYXRVcd6S7NHbHz3kA=
61 +github.com/go-jose/go-jose/v4 v4.1.4/go.mod h1:x4oUasVrzR7071A4TnHLGSPpNOm2a21K9Kf04k1rs08=
62 github.com/go-logr/logr v1.2.2/go.mod h1:jdQByPbusPIv2/zmleS9BjJVeZ6kBagPoEUsqbVz/1A=
63 github.com/go-logr/logr v1.4.3 h1:CjnDlHq8ikf6E492q6eKboGOC0T8CDaOvkHCIg8idEI=
64 github.com/go-logr/logr v1.4.3/go.mod h1:9T104GzyrTigFIr8wt5mBrctHMim0Nb2HLGrmQ40KvY=
@@ -65,8 +67,6 @@ github.com/go-logr/stdr v1.2.2/go.mod h1:mMo/vtBO5dYbehREoey6XUKy/eSumjCCveDpRre
67 github.com/godbus/dbus/v5 v5.0.4/go.mod h1:xhWf0FNVPg57R7Z0UbKHbJfkEywrmjJnf7w5xrFpKfA=
68 github.com/golang/protobuf v1.5.4 h1:i7eJL8qZTpSEXOPTxNKhASYpMn+8e5Q6AdndVa1dWek=
69 github.com/golang/protobuf v1.5.4/go.mod h1:lnTiLA8Wa4RWRcIUkrtSVa5nRhsEGBg48fD6rSs7xps=
68 -github.com/google/btree v1.1.2 h1:xf4v41cLI2Z6FxbKm+8Bu+m8ifhj15JuZ9sa0jZCMUU=
69 -github.com/google/btree v1.1.2/go.mod h1:qOPhT0dTNdNzV6Z/lhRX0YXUafgPLFUh+gZMl761Gm4=
70 github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8=
71 github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU=
72 github.com/google/s2a-go v0.1.9 h1:LGD7gtMgezd8a/Xak7mEWL0PjoTQFvpRudN895yqKW0=
@@ -117,6 +117,7 @@ go.opentelemetry.io/otel/sdk v1.39.0 h1:nMLYcjVsvdui1B/4FRkwjzoRVsMK8uL/cj0OyhKz
117 go.opentelemetry.io/otel/sdk v1.39.0/go.mod h1:vDojkC4/jsTJsE+kh+LXYQlbL8CgrEcwmt1ENZszdJE=
118 go.opentelemetry.io/otel/sdk/metric v1.38.0 h1:aSH66iL0aZqo//xXzQLYozmWrXxyFkBJ6qT5wthqPoM=
119 go.opentelemetry.io/otel/sdk/metric v1.38.0/go.mod h1:dg9PBnW9XdQ1Hd6ZnRz689CbtrUp0wMMs9iPcgT9EZA=
120 +go.opentelemetry.io/otel/sdk/metric v1.39.0 h1:cXMVVFVgsIf2YL6QkRF4Urbr/aMInf+2WKg+sEJTtB8=
121 go.opentelemetry.io/otel/trace v1.39.0 h1:2d2vfpEDmCJ5zVYz7ijaJdOF59xLomrvj7bjt6/qCJI=
122 go.opentelemetry.io/otel/trace v1.39.0/go.mod h1:88w4/PnZSazkGzz/w84VHpQafiU4EtqqlVdxWy+rNOA=
123 go.uber.org/mock v0.5.2 h1:LbtPTcP8A5k9WPXj54PPPbjcI4Y6lhyOZXn+VS7wNko=
@@ -142,10 +143,6 @@ golang.org/x/time v0.14.0 h1:MRx4UaLrDotUKUdCIqzPC48t1Y9hANFKIRpNx+Te8PI=
143 golang.org/x/time v0.14.0/go.mod h1:eL/Oa2bBBK0TkX57Fyni+NgnyQQN4LitPmob2Hjnqw4=
144 golang.org/x/tools v0.41.0 h1:a9b8iMweWG+S0OBnlU36rzLp20z1Rp10w+IY2czHTQc=
145 golang.org/x/tools v0.41.0/go.mod h1:XSY6eDqxVNiYgezAVqqCeihT4j1U2CCsqvH3WhQpnlg=
145 -golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2 h1:B82qJJgjvYKsXS9jeunTOisW56dUokqW/FOteYJJ/yg=
146 -golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2/go.mod h1:deeaetjYA+DHMHg+sMSMI58GrEteJUUzzw7en6TJQcI=
147 -golang.zx2c4.com/wireguard v0.0.0-20250521234502-f333402bd9cb h1:whnFRlWMcXI9d+ZbWg+4sHnLp52d5yiIPUxMBSt4X9A=
148 -golang.zx2c4.com/wireguard v0.0.0-20250521234502-f333402bd9cb/go.mod h1:rpwXGsirqLqN2L0JDJQlwOboGHmptD5ZD6T2VmcqhTw=
146 gonum.org/v1/gonum v0.16.0 h1:5+ul4Swaf3ESvrOnidPp4GZbzf0mxVQpDCYUQE7OJfk=
147 gonum.org/v1/gonum v0.16.0/go.mod h1:fef3am4MQ93R2HHpKnLk4/Tbh/s0+wqD5nfa6Pnwy4E=
148 google.golang.org/api v0.267.0 h1:w+vfWPMPYeRs8qH1aYYsFX68jMls5acWl/jocfLomwE=
@@ -158,9 +155,9 @@ google.golang.org/genproto/googleapis/rpc v0.0.0-20260203192932-546029d2fa20 h1:
155 google.golang.org/genproto/googleapis/rpc v0.0.0-20260203192932-546029d2fa20/go.mod h1:j9x/tPzZkyxcgEFkiKEEGxfvyumM01BEtsW8xzOahRQ=
156 google.golang.org/grpc v1.78.0 h1:K1XZG/yGDJnzMdd/uZHAkVqJE+xIDOcmdSFZkBUicNc=
157 google.golang.org/grpc v1.78.0/go.mod h1:I47qjTo4OKbMkjA/aOOwxDIiPSBofUtQUI5EfpWvW7U=
158 +google.golang.org/grpc v1.79.3 h1:sybAEdRIEtvcD68Gx7dmnwjZKlyfuc61Dyo9pGXXkKE=
159 +google.golang.org/grpc v1.79.3/go.mod h1:KmT0Kjez+0dde/v2j9vzwoAScgEPx/Bw1CYChhHLrHQ=
160 google.golang.org/protobuf v1.36.11 h1:fV6ZwhNocDyBLK0dj+fg8ektcVegBBuEolpbTQyBNVE=
161 google.golang.org/protobuf v1.36.11/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco=
162 gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA=
163 gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM=
165 -gvisor.dev/gvisor v0.0.0-20250503011706-39ed1f5ac29c h1:m/r7OM+Y2Ty1sgBQ7Qb27VgIMBW8ZZhT4gLnUyDIhzI=
166 -gvisor.dev/gvisor v0.0.0-20250503011706-39ed1f5ac29c/go.mod h1:3r5CMtNQMKIvBlrmM9xWUNamjKBYPOWyXOjmg5Kts3g=
portal/api_server.go
+10 -19
@@ -146,25 +146,16 @@ func (s *Server) handleRelayDiscovery(w http.ResponseWriter, r *http.Request) {
146 ingressAddr = fmt.Sprintf("%s:%d", ingressAddr, s.cfg.SNIPort)
147 }
148
149 - supportsOverlayPeer := s.wgConfig.PublicKey != "" &&
150 - s.wgConfig.Endpoint != "" &&
151 - s.wgConfig.OverlayIPv4 != ""
152 -
149 self, err := discovery.NormalizeDescriptor(types.RelayDescriptor{
154 - Identity: s.identity.Copy(),
155 - Sequence: uint64(now.UnixMilli()),
156 - Version: 1,
157 - IssuedAt: now,
158 - ExpiresAt: now.Add(2 * types.DiscoveryPollInterval),
159 - APIHTTPSAddr: s.cfg.PortalURL,
160 - IngressTLSAddr: ingressAddr,
161 - SupportsUDP: s.cfg.UDPEnabled && s.quicTunnel != nil,
162 - SupportsTCP: s.cfg.TCPEnabled,
163 - SupportsOverlayPeer: supportsOverlayPeer,
164 - WireGuardPublicKey: s.wgConfig.PublicKey,
165 - WireGuardEndpoint: s.wgConfig.Endpoint,
166 - OverlayIPv4: s.wgConfig.OverlayIPv4,
167 - OverlayCIDRs: append([]string(nil), s.wgConfig.OverlayCIDRs...),
150 + Identity: s.identity.Copy(),
151 + Sequence: uint64(now.UnixMilli()),
152 + Version: 1,
153 + IssuedAt: now,
154 + ExpiresAt: now.Add(2 * types.DiscoveryPollInterval),
155 + APIHTTPSAddr: s.cfg.PortalURL,
156 + IngressTLSAddr: ingressAddr,
157 + SupportsUDP: s.cfg.UDPEnabled && s.quicTunnel != nil,
158 + SupportsTCP: s.cfg.TCPEnabled,
159 })
160 if err != nil {
161 utils.WriteAPIError(w, http.StatusInternalServerError, types.APIErrorCodeInternal, err.Error())
@@ -178,7 +169,7 @@ func (s *Server) handleRelayDiscovery(w http.ResponseWriter, r *http.Request) {
169 Relays: nil,
170 }
171 if s.relaySet != nil {
181 - resp.Relays = s.relaySet.AdvertisedDescriptors()
172 + resp.Relays = s.relaySet.ActiveRelayDescriptors()
173 }
174 utils.WriteAPIData(w, http.StatusOK, resp)
175 }
portal/auth/auth.go
+3 -3
@@ -8,7 +8,7 @@ import (
8 "time"
9
10 "github.com/decred/dcrd/dcrec/secp256k1/v4"
11 - secp256k1ecdsa "github.com/decred/dcrd/dcrec/secp256k1/v4/ecdsa"
11 + "github.com/decred/dcrd/dcrec/secp256k1/v4/ecdsa"
12 jose "github.com/go-jose/go-jose/v4"
13 "github.com/go-jose/go-jose/v4/jwt"
14 "github.com/spruceid/siwe-go"
@@ -58,7 +58,7 @@ func (s *es256kOpaqueSigner) SignPayload(payload []byte, alg jose.SignatureAlgor
58 }
59
60 hash := sha256.Sum256(payload)
61 - compact := secp256k1ecdsa.SignCompact(s.privateKey, hash[:], false)
61 + compact := ecdsa.SignCompact(s.privateKey, hash[:], false)
62 if len(compact) != 65 {
63 return nil, errors.New("invalid compact signature length")
64 }
@@ -97,7 +97,7 @@ func (v *es256kOpaqueVerifier) VerifyPayload(payload []byte, signature []byte, a
97 }
98
99 func verifyRawSignature(hash []byte, r, s *secp256k1.ModNScalar, publicKey *secp256k1.PublicKey) error {
100 - signature := secp256k1ecdsa.NewSignature(r, s)
100 + signature := ecdsa.NewSignature(r, s)
101 if !signature.Verify(hash, publicKey) {
102 return errors.New("token signature is invalid")
103 }
portal/discovery/discovery.go
+12 -62
@@ -9,6 +9,8 @@ import (
9 "strings"
10 "time"
11
12 + "github.com/rs/zerolog/log"
13 +
14 "github.com/gosuda/portal/v2/portal/keyless"
15 "github.com/gosuda/portal/v2/types"
16 "github.com/gosuda/portal/v2/utils"
@@ -20,9 +22,6 @@ 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)
23 - desc.WireGuardPublicKey = strings.TrimSpace(desc.WireGuardPublicKey)
24 - desc.WireGuardEndpoint = strings.TrimSpace(desc.WireGuardEndpoint)
25 - desc.OverlayIPv4 = strings.TrimSpace(desc.OverlayIPv4)
25 if !desc.IssuedAt.IsZero() {
26 desc.IssuedAt = desc.IssuedAt.UTC()
27 }
@@ -44,19 +43,6 @@ func NormalizeDescriptor(desc types.RelayDescriptor) (types.RelayDescriptor, err
43 }
44 desc.Address = normalized
45 }
47 - if len(desc.OverlayCIDRs) > 0 {
48 - normalized, err := utils.NormalizeOverlayCIDRs(desc.OverlayCIDRs)
49 - if err != nil {
50 - return types.RelayDescriptor{}, err
51 - }
52 - desc.OverlayCIDRs = normalized
53 - }
54 - if !desc.SupportsOverlayPeer {
55 - desc.WireGuardPublicKey = ""
56 - desc.WireGuardEndpoint = ""
57 - desc.OverlayIPv4 = ""
58 - desc.OverlayCIDRs = nil
59 - }
46 return desc, nil
47 }
48
@@ -88,17 +74,6 @@ func ValidateDescriptor(desc types.RelayDescriptor, now time.Time) (types.RelayD
74 case normalized.IssuedAt.After(normalized.ExpiresAt):
75 return types.RelayDescriptor{}, errors.New("issued_at must be before expires_at")
76 }
91 - if normalized.SupportsOverlayPeer {
92 - if err := utils.ValidateWireGuardPublicKey(normalized.WireGuardPublicKey); err != nil {
93 - return types.RelayDescriptor{}, err
94 - }
95 - if err := utils.ValidateWireGuardEndpoint(normalized.WireGuardEndpoint); err != nil {
96 - return types.RelayDescriptor{}, err
97 - }
98 - if err := utils.ValidateOverlayIPv4(normalized.OverlayIPv4); err != nil {
99 - return types.RelayDescriptor{}, err
100 - }
101 - }
77 return normalized, nil
78 }
79
@@ -115,23 +90,28 @@ func ValidateRelayDiscoveryResponse(resp types.DiscoveryResponse, now time.Time)
90
91 seen := map[string]struct{}{self.Key(): {}}
92 relays := make([]types.RelayDescriptor, 0, len(resp.Relays))
118 - var validateErr error
93 for _, descriptor := range resp.Relays {
94 verified, err := ValidateDescriptor(descriptor, now)
95 if err != nil {
122 - if validateErr == nil {
123 - validateErr = err
124 - }
96 + log.Warn().
97 + Err(err).
98 + Str("relay", strings.TrimSpace(descriptor.APIHTTPSAddr)).
99 + Str("name", strings.TrimSpace(descriptor.Name)).
100 + Msg("skipping invalid discovery relay hint")
101 continue
102 }
103 identityKey := verified.Key()
104 if _, ok := seen[identityKey]; ok {
105 + log.Debug().
106 + Str("relay", verified.APIHTTPSAddr).
107 + Str("identity_key", identityKey).
108 + Msg("skipping duplicate discovery relay hint")
109 continue
110 }
111 seen[identityKey] = struct{}{}
112 relays = append(relays, verified)
113 }
134 - return self, relays, validateErr
114 + return self, relays, nil
115 }
116
117 // ValidateDescriptorTarget checks if a descriptor matches expected target identity.
@@ -208,33 +188,3 @@ func DiscoveryUnavailableStatus(err error) (statusCode int, code string, unavail
188 }
189 return 0, "", false
190 }
211 -
212 -func SeedDescriptor(apiURL string) (types.RelayDescriptor, error) {
213 - normalized, err := utils.NormalizeRelayURL(apiURL)
214 - if err != nil {
215 - return types.RelayDescriptor{}, err
216 - }
217 - return types.RelayDescriptor{
218 - Identity: types.Identity{
219 - Name: utils.PortalRootHost(normalized),
220 - },
221 - APIHTTPSAddr: normalized,
222 - Version: 1,
223 - }, nil
224 -}
225 -
226 -func RequireOverlayRelayDescriptor(desc types.RelayDescriptor) error {
227 - if !desc.SupportsOverlayPeer {
228 - return errors.New("descriptor does not support overlay peer")
229 - }
230 - if desc.WireGuardPublicKey == "" {
231 - return errors.New("descriptor wireguard public key is required")
232 - }
233 - if desc.WireGuardEndpoint == "" {
234 - return errors.New("descriptor wireguard endpoint is required")
235 - }
236 - if desc.OverlayIPv4 == "" {
237 - return errors.New("descriptor overlay ipv4 is required")
238 - }
239 - return nil
240 -}
portal/discovery/relayset.go
+397 -524
@@ -1,9 +1,9 @@
1 package discovery
2
3 import (
4 + "context"
5 "errors"
6 "net/http"
6 - "reflect"
7 "sort"
8 "strings"
9 "sync"
@@ -15,232 +15,228 @@ import (
15 "github.com/gosuda/portal/v2/utils"
16 )
17
18 -type RelayView struct {
19 - Descriptor types.RelayDescriptor
20 - FirstSeenAt time.Time
21 - LastSeenAt time.Time
22 -}
23 -
18 type RelayLocalState struct {
19 Banned bool
26 - BanReason string
27 - Bootstrap bool
20 Advertised bool
21 Expired bool
30 - Reachable bool
22 ConsecutiveFailures int
32 - LastSuccessAt time.Time
33 - LastFailureAt time.Time
23 }
24
36 -type RelaySummary struct {
37 - Known int
38 - Banned int
39 - Bootstrap int
40 - Advertised int
41 - Expired int
42 - Syncable int
43 - Reachable int
44 - Unreachable int
45 -}
25 +const discoveryRecoveryFailures = 3
26
47 -// RelaySet owns the shared relay discovery view: known relay URLs, stable relay
48 -// id/url mappings, the latest validated descriptor seen for each relay, and common
49 -// process-local relay state such as ban/reachability/failure tracking.
50 -//
51 -// Runtime-specific policy such as bootstrap classification, relay lifecycle, or
52 -// listener ownership belongs in the caller's projection.
27 +// RelaySet owns the current discovery state:
28 +// explicit bootstrap relay URLs, the latest descriptor seen for each relay,
29 +// and the minimal local state required to ban or expire a relay.
30 type RelaySet struct {
54 - mu sync.RWMutex
55 - knownRelayURLs []string
56 - relayKeysByURL map[string]string
57 - relays map[string]RelayView
58 - localByURL map[string]RelayLocalState
59 - lastStatusReachable map[string]bool
60 - lastStatusSummary RelaySummary
61 - haveLastStatus bool
31 + mu sync.RWMutex
32 + bootstrapRelayURLs []string
33 + relayKeysByURL map[string]string
34 + relays map[string]types.RelayDescriptor
35 + localByURL map[string]RelayLocalState
36 + activeRelayURLs []string
37 + activeRelays []types.RelayDescriptor
38 + selfRelayKey string
39 + selfRelayURL string
40 }
41
42 func NewRelaySet() *RelaySet {
43 return &RelaySet{
44 relayKeysByURL: make(map[string]string),
67 - relays: make(map[string]RelayView),
45 + relays: make(map[string]types.RelayDescriptor),
46 localByURL: make(map[string]RelayLocalState),
47 }
48 }
49
72 -func (s *RelaySet) trackedRelayURLs() []string {
73 - if s == nil {
50 +func relayExpiredAt(desc types.RelayDescriptor, state RelayLocalState, now time.Time) bool {
51 + if state.Expired {
52 + return true
53 + }
54 + if desc.ExpiresAt.IsZero() {
55 + return false
56 + }
57 + if now.IsZero() {
58 + now = time.Now().UTC()
59 + }
60 + return !desc.ExpiresAt.After(now)
61 +}
62 +
63 +func (s *RelaySet) bootstrapRelayURLSetLocked() map[string]struct{} {
64 + if len(s.bootstrapRelayURLs) == 0 {
65 return nil
66 }
67 + out := make(map[string]struct{}, len(s.bootstrapRelayURLs))
68 + for _, relayURL := range s.bootstrapRelayURLs {
69 + out[relayURL] = struct{}{}
70 + }
71 + return out
72 +}
73
77 - urls := make([]string, 0, len(s.knownRelayURLs)+len(s.relays))
78 - seen := make(map[string]struct{}, len(s.knownRelayURLs)+len(s.relays))
79 - for _, relayURL := range s.knownRelayURLs {
80 - relayURL = strings.TrimSpace(relayURL)
81 - if relayURL == "" {
82 - continue
83 - }
84 - if _, ok := seen[relayURL]; ok {
85 - continue
86 - }
87 - seen[relayURL] = struct{}{}
88 - urls = append(urls, relayURL)
74 +func (s *RelaySet) descriptorByURLLocked(relayURL string) (types.RelayDescriptor, bool) {
75 + relayKey, ok := s.relayKeysByURL[relayURL]
76 + if !ok {
77 + return types.RelayDescriptor{}, false
78 }
90 - for _, view := range s.relays {
91 - relayURL := view.Descriptor.APIHTTPSAddr
92 - if relayURL == "" {
93 - continue
94 - }
95 - if _, ok := seen[relayURL]; ok {
96 - continue
97 - }
98 - seen[relayURL] = struct{}{}
99 - urls = append(urls, relayURL)
79 + desc, ok := s.relays[relayKey]
80 + return desc, ok
81 +}
82 +
83 +func (s *RelaySet) isSelfRelayURLLocked(relayURL string) bool {
84 + if s == nil {
85 + return false
86 }
101 - return urls
87 + relayURL = strings.TrimSpace(relayURL)
88 + return relayURL != "" && s.selfRelayURL != "" && relayURL == s.selfRelayURL
89 }
90
104 -func (s *RelaySet) ActiveRelayURLs() []string {
91 +func (s *RelaySet) isSelfRelayDescriptorLocked(desc types.RelayDescriptor) bool {
92 if s == nil {
106 - return nil
93 + return false
94 }
108 - s.mu.RLock()
109 - defer s.mu.RUnlock()
110 - if len(s.knownRelayURLs) == 0 {
111 - return nil
95 + if s.selfRelayKey != "" {
96 + if relayKey := desc.Key(); relayKey != "" && relayKey == s.selfRelayKey {
97 + return true
98 + }
99 }
100 + return s.isSelfRelayURLLocked(desc.APIHTTPSAddr)
101 +}
102
114 - out := make([]string, 0, len(s.knownRelayURLs))
115 - for _, relayURL := range s.knownRelayURLs {
116 - if state, ok := s.localByURL[relayURL]; ok && state.Banned {
117 - continue
103 +func (s *RelaySet) removeSelfRelayLocked() {
104 + if s == nil {
105 + return
106 + }
107 + if s.selfRelayURL != "" {
108 + filtered := s.bootstrapRelayURLs[:0]
109 + for _, relayURL := range s.bootstrapRelayURLs {
110 + if s.isSelfRelayURLLocked(relayURL) {
111 + continue
112 + }
113 + filtered = append(filtered, relayURL)
114 }
119 - out = append(out, relayURL)
115 + s.bootstrapRelayURLs = filtered
116 +
117 + delete(s.localByURL, s.selfRelayURL)
118 + delete(s.relayKeysByURL, s.selfRelayURL)
119 }
121 - if len(out) == 0 {
122 - return nil
120 + if s.selfRelayKey != "" {
121 + if desc, ok := s.relays[s.selfRelayKey]; ok {
122 + delete(s.localByURL, desc.APIHTTPSAddr)
123 + delete(s.relayKeysByURL, desc.APIHTTPSAddr)
124 + }
125 + delete(s.relays, s.selfRelayKey)
126 }
124 - return out
127 }
128
127 -func relayExpiredAt(view RelayView, state RelayLocalState, now time.Time) bool {
128 - if state.Expired {
129 - return true
129 +func (s *RelaySet) SetSelfRelay(identity types.Identity, relayURL string) error {
130 + if s == nil {
131 + return nil
132 }
131 - if view.Descriptor.ExpiresAt.IsZero() {
132 - return false
133 + relayURL = strings.TrimSpace(relayURL)
134 + if relayURL != "" {
135 + normalized, err := utils.NormalizeRelayURL(relayURL)
136 + if err != nil {
137 + return err
138 + }
139 + relayURL = normalized
140 }
141 +
142 + s.mu.Lock()
143 + defer s.mu.Unlock()
144 + s.selfRelayKey = identity.Key()
145 + s.selfRelayURL = relayURL
146 + s.removeSelfRelayLocked()
147 + s.syncActiveLocked(time.Now().UTC())
148 + return nil
149 +}
150 +
151 +func (s *RelaySet) syncActiveLocked(now time.Time) {
152 if now.IsZero() {
153 now = time.Now().UTC()
154 }
137 - return !view.Descriptor.ExpiresAt.After(now)
138 -}
155
140 -func (s *RelaySet) logStatusChange() {
141 - now := time.Now().UTC()
142 - var currentReachable map[string]bool
143 - trackedRelayURLs := s.trackedRelayURLs()
144 - if len(trackedRelayURLs) > 0 {
145 - currentReachable = make(map[string]bool, len(trackedRelayURLs))
146 - for _, relayURL := range trackedRelayURLs {
147 - state := s.localByURL[relayURL]
148 - currentReachable[relayURL] = !state.Banned && state.Reachable
149 - }
150 - }
151 - summary := RelaySummary{}
152 - for _, relayURL := range trackedRelayURLs {
153 - summary.Known++
156 + activeRelayURLs := make([]string, 0, len(s.bootstrapRelayURLs)+len(s.relays))
157 + activeRelays := make([]types.RelayDescriptor, 0, len(s.relays))
158 + seen := make(map[string]struct{}, len(s.bootstrapRelayURLs)+len(s.relays))
159 + for _, relayURL := range s.bootstrapRelayURLs {
160 + if s.isSelfRelayURLLocked(relayURL) {
161 + continue
162 + }
163 state := s.localByURL[relayURL]
155 - relayKey := s.relayKeysByURL[relayURL]
156 - view, ok := s.relays[relayKey]
157 - expired := ok && relayExpiredAt(view, state, now) || !ok && state.Expired
164 if state.Banned {
159 - summary.Banned++
165 continue
166 }
162 - if state.Bootstrap {
163 - summary.Bootstrap++
164 - }
165 - if state.Advertised && !expired {
166 - summary.Advertised++
167 - }
168 - if expired {
169 - summary.Expired++
170 - }
171 - if state.Reachable {
172 - summary.Reachable++
173 - } else {
174 - summary.Unreachable++
175 - }
176 - if ok && !state.Bootstrap && !expired && view.Descriptor.SupportsOverlayPeer {
177 - summary.Syncable++
167 + seen[relayURL] = struct{}{}
168 + activeRelayURLs = append(activeRelayURLs, relayURL)
169 + desc, ok := s.descriptorByURLLocked(relayURL)
170 + if !ok || desc.APIHTTPSAddr == "" || !state.Advertised || relayExpiredAt(desc, state, now) {
171 + continue
172 }
179 - }
180 - if s.haveLastStatus && summary == s.lastStatusSummary && reflect.DeepEqual(currentReachable, s.lastStatusReachable) {
181 - return
173 + activeRelays = append(activeRelays, desc)
174 }
175
184 - activated := make([]string, 0)
185 - deactivated := make([]string, 0)
186 - for relayURL, reachable := range currentReachable {
187 - if s.lastStatusReachable == nil || s.lastStatusReachable[relayURL] == reachable {
176 + discovered := make([]types.RelayDescriptor, 0, len(s.relays))
177 + for _, desc := range s.relays {
178 + if s.isSelfRelayDescriptorLocked(desc) {
179 continue
180 }
190 - if reachable {
191 - activated = append(activated, relayURL)
192 - } else {
193 - deactivated = append(deactivated, relayURL)
181 + relayURL := strings.TrimSpace(desc.APIHTTPSAddr)
182 + if relayURL == "" {
183 + continue
184 }
195 - }
196 - for relayURL, reachable := range s.lastStatusReachable {
197 - if _, ok := currentReachable[relayURL]; ok || !reachable {
185 + if _, ok := seen[relayURL]; ok {
186 + continue
187 + }
188 + state := s.localByURL[relayURL]
189 + if state.Banned || !state.Advertised || relayExpiredAt(desc, state, now) {
190 continue
191 }
200 - deactivated = append(deactivated, relayURL)
192 + seen[relayURL] = struct{}{}
193 + discovered = append(discovered, desc)
194 + }
195 + sort.Slice(discovered, func(i, j int) bool {
196 + return discovered[i].APIHTTPSAddr < discovered[j].APIHTTPSAddr
197 + })
198 + for _, desc := range discovered {
199 + activeRelayURLs = append(activeRelayURLs, desc.APIHTTPSAddr)
200 + activeRelays = append(activeRelays, desc)
201 }
202
203 - event := log.Info().
204 - Int("banned", summary.Banned).
205 - Int("bootstrap", summary.Bootstrap).
206 - Int("advertised", summary.Advertised).
207 - Int("expired", summary.Expired).
208 - Int("syncable", summary.Syncable).
209 - Int("reachable", summary.Reachable).
210 - Int("unreachable", summary.Unreachable)
211 - if len(activated) > 0 {
212 - event = event.Strs("activated", activated)
203 + s.activeRelayURLs = activeRelayURLs
204 + s.activeRelays = activeRelays
205 +}
206 +
207 +func (s *RelaySet) ActiveRelayURLs() []string {
208 + if s == nil {
209 + return nil
210 }
214 - if len(deactivated) > 0 {
215 - event = event.Strs("deactivated", deactivated)
211 + s.mu.RLock()
212 + defer s.mu.RUnlock()
213 + if len(s.activeRelayURLs) == 0 {
214 + return nil
215 }
217 - event.Msg("relay status")
218 - s.lastStatusReachable = currentReachable
219 - s.lastStatusSummary = summary
220 - s.haveLastStatus = true
216 + return append([]string(nil), s.activeRelayURLs...)
217 }
218
223 -func (s *RelaySet) BootstrapDescriptors() []types.RelayDescriptor {
219 +func (s *RelaySet) bootstrapDescriptors() []types.RelayDescriptor {
220 if s == nil {
221 return nil
222 }
223 s.mu.RLock()
224 defer s.mu.RUnlock()
229 - if len(s.knownRelayURLs) == 0 {
225 + if len(s.bootstrapRelayURLs) == 0 {
226 return nil
227 }
228
233 - out := make([]types.RelayDescriptor, 0, len(s.knownRelayURLs))
234 - for _, relayURL := range s.knownRelayURLs {
235 - state, ok := s.localByURL[relayURL]
236 - if !ok || !state.Bootstrap {
229 + out := make([]types.RelayDescriptor, 0, len(s.bootstrapRelayURLs))
230 + for _, relayURL := range s.bootstrapRelayURLs {
231 + if s.isSelfRelayURLLocked(relayURL) {
232 continue
233 }
239 - if relayKey, ok := s.relayKeysByURL[relayURL]; ok {
240 - if view, ok := s.relays[relayKey]; ok && view.Descriptor.APIHTTPSAddr != "" {
241 - out = append(out, view.Descriptor)
242 - continue
243 - }
234 + if s.localByURL[relayURL].Banned {
235 + continue
236 + }
237 + if desc, ok := s.descriptorByURLLocked(relayURL); ok && desc.APIHTTPSAddr != "" {
238 + out = append(out, desc)
239 + continue
240 }
241 out = append(out, types.RelayDescriptor{
242 Identity: types.Identity{
@@ -256,155 +252,19 @@ func (s *RelaySet) BootstrapDescriptors() []types.RelayDescriptor {
252 return out
253 }
254
259 -func (s *RelaySet) BanRelayURL(relayURL, reason string) bool {
260 - if s == nil {
261 - return false
262 - }
263 - s.mu.Lock()
264 - defer s.mu.Unlock()
265 - relayURL = strings.TrimSpace(relayURL)
266 - if relayURL == "" {
267 - return false
268 - }
269 -
270 - state := s.localByURL[relayURL]
271 - reason = strings.TrimSpace(reason)
272 - changed := !state.Banned || strings.TrimSpace(state.BanReason) != reason
273 - state.Banned = true
274 - state.BanReason = reason
275 - state.Reachable = false
276 - s.localByURL[relayURL] = state
277 - if changed {
278 - s.logStatusChange()
279 - }
280 - return changed
281 -}
282 -
283 -func (s *RelaySet) MarkRelayUnreachable(relayURL string) bool {
284 - if s == nil {
285 - return false
286 - }
287 - s.mu.Lock()
288 - defer s.mu.Unlock()
289 - relayURL = strings.TrimSpace(relayURL)
290 - if relayURL == "" {
291 - return false
292 - }
293 -
294 - state := s.localByURL[relayURL]
295 - if state.Banned {
296 - return false
297 - }
298 - if !state.Reachable {
299 - return false
300 - }
301 - state.Reachable = false
302 - s.localByURL[relayURL] = state
303 - s.logStatusChange()
304 - return true
305 -}
306 -
307 -func (s *RelaySet) MarkRelayReachable(relayURL string, now time.Time) bool {
308 - if s == nil {
309 - return false
310 - }
311 - s.mu.Lock()
312 - defer s.mu.Unlock()
313 - relayURL = strings.TrimSpace(relayURL)
314 - if relayURL == "" {
315 - return false
316 - }
317 - if now.IsZero() {
318 - now = time.Now().UTC()
319 - }
320 -
321 - state := s.localByURL[relayURL]
322 - changed := !state.Reachable || state.ConsecutiveFailures != 0 || state.LastSuccessAt != now
323 - state.Reachable = true
324 - state.ConsecutiveFailures = 0
325 - state.LastSuccessAt = now
326 - s.localByURL[relayURL] = state
327 - if changed {
328 - s.logStatusChange()
329 - }
330 - return changed
331 -}
332 -
333 -func (s *RelaySet) MarkRelayFailure(relayURL string, now time.Time) RelayLocalState {
334 - if s == nil {
335 - return RelayLocalState{}
336 - }
337 - s.mu.Lock()
338 - defer s.mu.Unlock()
339 - relayURL = strings.TrimSpace(relayURL)
340 - if relayURL == "" {
341 - return RelayLocalState{}
342 - }
343 - if now.IsZero() {
344 - now = time.Now().UTC()
345 - }
346 -
347 - state := s.localByURL[relayURL]
348 - state.Reachable = false
349 - state.ConsecutiveFailures++
350 - state.LastFailureAt = now
351 - s.localByURL[relayURL] = state
352 - s.logStatusChange()
353 - return state
354 -}
355 -
356 -func (s *RelaySet) RecordBootstrapDiscoveryFailure(relayURL string, err error, now time.Time) {
357 - state := s.MarkRelayFailure(relayURL, now)
358 - if statusCode, code, unavailable := DiscoveryUnavailableStatus(err); unavailable {
359 - if state.ConsecutiveFailures > 1 {
360 - return
361 - }
362 - event := log.Info().Str("relay", relayURL)
363 - if statusCode > 0 {
364 - event = event.Int("status_code", statusCode)
365 - }
366 - if code != "" {
367 - event = event.Str("code", code)
368 - }
369 - event.Msg("bootstrap relay discovery unavailable; peer may have discovery disabled")
370 - return
371 - }
372 -
373 - log.Warn().
374 - Err(err).
375 - Str("relay", relayURL).
376 - Msg("bootstrap relay discovery failed")
377 -}
378 -
379 -func (s *RelaySet) AdvertisedDescriptors() []types.RelayDescriptor {
255 +func (s *RelaySet) ActiveRelayDescriptors() []types.RelayDescriptor {
256 if s == nil {
257 return nil
258 }
259 s.mu.RLock()
260 defer s.mu.RUnlock()
385 - if len(s.relays) == 0 {
386 - return nil
387 - }
388 -
389 - now := time.Now().UTC()
390 - out := make([]types.RelayDescriptor, 0, len(s.relays))
391 - for _, view := range s.relays {
392 - state := s.localByURL[view.Descriptor.APIHTTPSAddr]
393 - if !state.Advertised || relayExpiredAt(view, state, now) || view.Descriptor.APIHTTPSAddr == "" {
394 - continue
395 - }
396 - out = append(out, view.Descriptor)
397 - }
398 - if len(out) == 0 {
261 + if len(s.activeRelays) == 0 {
262 return nil
263 }
401 - sort.Slice(out, func(i, j int) bool {
402 - return out[i].APIHTTPSAddr < out[j].APIHTTPSAddr
403 - })
404 - return out
264 + return append([]types.RelayDescriptor(nil), s.activeRelays...)
265 }
266
407 -func (s *RelaySet) SyncableDescriptors() []types.RelayDescriptor {
267 +func (s *RelaySet) confirmableDescriptors() []types.RelayDescriptor {
268 if s == nil {
269 return nil
270 }
@@ -415,13 +275,24 @@ func (s *RelaySet) SyncableDescriptors() []types.RelayDescriptor {
275 }
276
277 now := time.Now().UTC()
278 + bootstrapRelayURLs := s.bootstrapRelayURLSetLocked()
279 out := make([]types.RelayDescriptor, 0, len(s.relays))
419 - for _, view := range s.relays {
420 - state := s.localByURL[view.Descriptor.APIHTTPSAddr]
421 - if state.Bootstrap || relayExpiredAt(view, state, now) || !view.Descriptor.SupportsOverlayPeer {
280 + for _, desc := range s.relays {
281 + if s.isSelfRelayDescriptorLocked(desc) {
282 continue
283 }
424 - out = append(out, view.Descriptor)
284 + relayURL := strings.TrimSpace(desc.APIHTTPSAddr)
285 + if relayURL == "" {
286 + continue
287 + }
288 + if _, ok := bootstrapRelayURLs[relayURL]; ok {
289 + continue
290 + }
291 + state := s.localByURL[relayURL]
292 + if state.Banned || relayExpiredAt(desc, state, now) {
293 + continue
294 + }
295 + out = append(out, desc)
296 }
297 if len(out) == 0 {
298 return nil
@@ -432,261 +303,163 @@ func (s *RelaySet) SyncableDescriptors() []types.RelayDescriptor {
303 return out
304 }
305
435 -func (s *RelaySet) Snapshot() map[string]types.RelayState {
306 +func (s *RelaySet) BanRelayURL(relayURL string) {
307 if s == nil {
437 - return nil
438 - }
439 - s.mu.RLock()
440 - defer s.mu.RUnlock()
441 - if len(s.relays) == 0 {
442 - return nil
308 + return
309 }
310 + s.mu.Lock()
311 + defer s.mu.Unlock()
312
445 - now := time.Now().UTC()
446 - snapshot := make(map[string]types.RelayState, len(s.relays))
447 - for relayKey, view := range s.relays {
448 - localState := s.localByURL[view.Descriptor.APIHTTPSAddr]
449 - snapshot[relayKey] = types.RelayState{
450 - Descriptor: view.Descriptor,
451 - Bootstrap: localState.Bootstrap,
452 - Advertised: localState.Advertised,
453 - Expired: relayExpiredAt(view, localState, now),
454 - FirstSeenAt: view.FirstSeenAt,
455 - LastSeenAt: view.LastSeenAt,
456 - ConsecutiveFailures: localState.ConsecutiveFailures,
457 - }
458 - }
459 - return snapshot
313 + relayURL = strings.TrimSpace(relayURL)
314 + if relayURL == "" {
315 + return
316 + }
317 + state := s.localByURL[relayURL]
318 + if state.Banned {
319 + return
320 + }
321 + state.Banned = true
322 + s.localByURL[relayURL] = state
323 + s.syncActiveLocked(time.Now().UTC())
324 }
325
462 -func (s *RelaySet) ReplaceKnownRelayURLs(relayURLs []string) {
326 +func (s *RelaySet) SetBootstrapRelayURLs(relayURLs []string) {
327 if s == nil {
328 return
329 }
330 s.mu.Lock()
331 defer s.mu.Unlock()
332 +
333 filtered := make([]string, 0, len(relayURLs))
334 + seen := make(map[string]struct{}, len(relayURLs))
335 for _, relayURL := range relayURLs {
336 relayURL = strings.TrimSpace(relayURL)
337 if relayURL == "" {
338 continue
339 }
474 - duplicate := false
475 - for _, existing := range filtered {
476 - if existing == relayURL {
477 - duplicate = true
478 - break
479 - }
340 + if s.isSelfRelayURLLocked(relayURL) {
341 + continue
342 }
481 - if duplicate {
343 + if _, ok := seen[relayURL]; ok {
344 continue
345 }
346 + seen[relayURL] = struct{}{}
347 filtered = append(filtered, relayURL)
348 + if _, ok := s.localByURL[relayURL]; !ok {
349 + s.localByURL[relayURL] = RelayLocalState{}
350 + }
351 }
486 - s.knownRelayURLs = append([]string(nil), filtered...)
487 -}
352
489 -func (s *RelaySet) registerDescriptor(desc types.RelayDescriptor, now time.Time) (string, bool, bool, error) {
490 - if s == nil {
491 - return "", false, false, nil
353 + for _, relayURL := range s.bootstrapRelayURLs {
354 + if _, ok := seen[relayURL]; ok {
355 + continue
356 + }
357 + if _, ok := s.relayKeysByURL[relayURL]; ok {
358 + continue
359 + }
360 + delete(s.localByURL, relayURL)
361 }
362 + s.bootstrapRelayURLs = filtered
363 + s.syncActiveLocked(time.Now().UTC())
364 +}
365 +
366 +func (s *RelaySet) registerDescriptorLocked(desc types.RelayDescriptor) error {
367 normalized, err := NormalizeDescriptor(desc)
368 if err != nil {
495 - return "", false, false, err
369 + return err
370 }
371 relayKey := normalized.Key()
372 if relayKey == "" {
499 - return "", false, false, errors.New("descriptor identity is required")
373 + return errors.New("descriptor identity is required")
374 }
375 if knownRelayKey, ok := s.relayKeysByURL[normalized.APIHTTPSAddr]; ok && knownRelayKey != relayKey {
502 - return "", false, false, errors.New("descriptor identity does not match known relay url")
376 + return errors.New("descriptor identity does not match known relay url")
377 }
378
505 - if now.IsZero() {
506 - now = time.Now().UTC()
507 - }
508 -
509 - view, ok := s.relays[relayKey]
510 - added := !ok
511 - if !ok {
512 - view.FirstSeenAt = now
513 - }
514 - previousURL := view.Descriptor.APIHTTPSAddr
515 - previousDescriptor := view.Descriptor
516 - view.Descriptor = normalized
517 - view.LastSeenAt = now
518 - s.relays[relayKey] = view
519 - s.relayKeysByURL[normalized.APIHTTPSAddr] = relayKey
379 + previous := s.relays[relayKey]
380 + previousURL := strings.TrimSpace(previous.APIHTTPSAddr)
381 if previousURL != "" && previousURL != normalized.APIHTTPSAddr {
382 delete(s.relayKeysByURL, previousURL)
522 - }
523 -
524 - changed := added || !reflect.DeepEqual(previousDescriptor, normalized)
525 - return relayKey, added, changed, nil
526 -}
527 -
528 -func relayDiscoveryURLs(selfDescriptor types.RelayDescriptor, relayDescriptors []types.RelayDescriptor) []string {
529 - relayURLs := make([]string, 0, 1+len(relayDescriptors))
530 - if apiURL := selfDescriptor.APIHTTPSAddr; apiURL != "" {
531 - relayURLs = append(relayURLs, apiURL)
532 - }
533 - for _, relayDescriptor := range relayDescriptors {
534 - if apiURL := relayDescriptor.APIHTTPSAddr; apiURL != "" {
535 - relayURLs = append(relayURLs, apiURL)
383 + if _, ok := s.relayKeysByURL[previousURL]; !ok {
384 + keepBootstrapState := false
385 + for _, bootstrapURL := range s.bootstrapRelayURLs {
386 + if bootstrapURL == previousURL {
387 + keepBootstrapState = true
388 + break
389 + }
390 + }
391 + if !keepBootstrapState {
392 + delete(s.localByURL, previousURL)
393 + }
394 }
395 }
538 - if len(relayURLs) == 0 {
539 - return nil
396 +
397 + s.relays[relayKey] = normalized
398 + s.relayKeysByURL[normalized.APIHTTPSAddr] = relayKey
399 + if _, ok := s.localByURL[normalized.APIHTTPSAddr]; !ok {
400 + s.localByURL[normalized.APIHTTPSAddr] = RelayLocalState{}
401 }
541 - return relayURLs
402 + return nil
403 }
404
544 -func (s *RelaySet) applyDiscoveryDescriptors(targetIdentity types.Identity, targetURL string, selfDescriptor types.RelayDescriptor, relayDescriptors []types.RelayDescriptor, now time.Time) (relaySetChanged bool, addedRelayCount int, err error) {
405 +func (s *RelaySet) ApplyRelayDiscoveryResponse(targetIdentity types.Identity, targetURL string, resp types.DiscoveryResponse, now time.Time) error {
406 if s == nil {
546 - return false, 0, nil
407 + return nil
408 }
409 if strings.TrimSpace(targetIdentity.Name) == "" && strings.TrimSpace(targetIdentity.Address) == "" {
549 - return false, 0, errors.New("target relay identity is required")
550 - }
551 - if now.IsZero() {
552 - now = time.Now().UTC()
553 - }
554 - if err := ValidateDescriptorTarget(selfDescriptor, targetIdentity, targetURL); err != nil {
555 - return false, 0, err
556 - }
557 -
558 - apply := func(desc types.RelayDescriptor, advertise, countAdded bool) error {
559 - _, added, descriptorChanged, err := s.registerDescriptor(desc, now)
560 - if err != nil {
561 - return err
562 - }
563 - localState := s.localByURL[desc.APIHTTPSAddr]
564 - wasAdvertised := localState.Advertised
565 - wasExpired := localState.Expired
566 - if advertise {
567 - localState.Advertised = true
568 - }
569 - localState.Expired = false
570 - s.localByURL[desc.APIHTTPSAddr] = localState
571 -
572 - changed := added || descriptorChanged || advertise && !wasAdvertised || wasExpired
573 - if added && countAdded {
574 - addedRelayCount++
575 - }
576 - if changed {
577 - relaySetChanged = true
578 - }
579 - return nil
410 + return errors.New("target relay identity is required")
411 }
412
582 - if err := apply(selfDescriptor, true, false); err != nil {
583 - return false, 0, err
413 + selfDescriptor, relayDescriptors, err := ValidateRelayDiscoveryResponse(resp, now)
414 + if err != nil {
415 + return err
416 }
585 - for _, relayDescriptor := range relayDescriptors {
586 - if err := apply(relayDescriptor, false, true); err != nil {
587 - return false, 0, err
588 - }
417 + if err := ValidateDescriptorTarget(selfDescriptor, targetIdentity, targetURL); err != nil {
418 + return err
419 }
590 - state := s.localByURL[selfDescriptor.APIHTTPSAddr]
591 - state.Reachable = true
592 - state.ConsecutiveFailures = 0
593 - state.LastSuccessAt = now
594 - s.localByURL[selfDescriptor.APIHTTPSAddr] = state
595 - s.logStatusChange()
596 - return relaySetChanged, addedRelayCount, nil
597 -}
420
599 -func (s *RelaySet) ApplyRelayDiscoveryResponse(targetIdentity types.Identity, targetURL string, resp types.DiscoveryResponse, now time.Time) (relayURLs []string, relaySetChanged bool, addedRelayCount int, warnErr error, err error) {
600 - selfDescriptor, relayDescriptors, validateErr := ValidateRelayDiscoveryResponse(resp, now)
601 - warnErr = validateErr
602 - if selfDescriptor.Key() == "" {
603 - return nil, false, 0, warnErr, validateErr
604 - }
421 s.mu.Lock()
606 - relaySetChanged, addedRelayCount, err = s.applyDiscoveryDescriptors(targetIdentity, targetURL, selfDescriptor, relayDescriptors, now)
607 - s.mu.Unlock()
608 - if err != nil {
609 - return nil, false, 0, warnErr, err
610 - }
611 - return relayDiscoveryURLs(selfDescriptor, relayDescriptors), relaySetChanged, addedRelayCount, warnErr, nil
612 -}
422 + defer s.mu.Unlock()
423
614 -func (s *RelaySet) ApplyOverlayRelayDiscoveryResponse(targetIdentity types.Identity, targetURL string, resp types.DiscoveryResponse, now time.Time) (relayURLs []string, relaySetChanged bool, addedRelayCount int, warnErr error, err error) {
615 - selfDescriptor, relayDescriptors, validateErr := ValidateRelayDiscoveryResponse(resp, now)
616 - warnErr = validateErr
617 - if selfDescriptor.Key() == "" {
618 - return nil, false, 0, warnErr, validateErr
619 - }
620 - if err := RequireOverlayRelayDescriptor(selfDescriptor); err != nil {
621 - return nil, false, 0, warnErr, err
424 + if err := s.registerDescriptorLocked(selfDescriptor); err != nil {
425 + return err
426 }
427 + selfState := s.localByURL[selfDescriptor.APIHTTPSAddr]
428 + selfState.Advertised = true
429 + selfState.Expired = false
430 + selfState.ConsecutiveFailures = 0
431 + s.localByURL[selfDescriptor.APIHTTPSAddr] = selfState
432
624 - filteredRelayDescriptors := make([]types.RelayDescriptor, 0, len(relayDescriptors))
433 for _, relayDescriptor := range relayDescriptors {
626 - if err := RequireOverlayRelayDescriptor(relayDescriptor); err != nil {
627 - if warnErr == nil {
628 - warnErr = err
629 - }
434 + if s.isSelfRelayDescriptorLocked(relayDescriptor) {
435 continue
436 }
632 - filteredRelayDescriptors = append(filteredRelayDescriptors, relayDescriptor)
633 - }
634 -
635 - s.mu.Lock()
636 - relaySetChanged, addedRelayCount, err = s.applyDiscoveryDescriptors(targetIdentity, targetURL, selfDescriptor, filteredRelayDescriptors, now)
637 - s.mu.Unlock()
638 - if err != nil {
639 - return nil, false, 0, warnErr, err
640 - }
641 - return relayDiscoveryURLs(selfDescriptor, filteredRelayDescriptors), relaySetChanged, addedRelayCount, warnErr, nil
642 -}
643 -
644 -func (s *RelaySet) RegisterBootstrapRelayURLs(inputs []string) ([]string, error) {
645 - if s == nil || len(inputs) == 0 {
646 - return nil, nil
647 - }
648 -
649 - normalized, err := utils.NormalizeRelayURLs(inputs...)
650 - if err != nil {
651 - return nil, err
652 - }
653 - normalized, err = utils.ExcludeLocalRelayURLs(normalized...)
654 - if err != nil {
655 - return nil, err
656 - }
657 - if len(normalized) == 0 {
658 - return nil, nil
659 - }
660 - s.mu.Lock()
661 - defer s.mu.Unlock()
662 -
663 - existing := make(map[string]struct{}, len(s.knownRelayURLs))
664 - for _, relayURL := range s.knownRelayURLs {
665 - existing[relayURL] = struct{}{}
666 - }
667 - added := make([]string, 0, len(normalized))
668 - for _, relayURL := range normalized {
669 - if _, ok := existing[relayURL]; ok {
437 + if err := s.registerDescriptorLocked(relayDescriptor); err != nil {
438 + log.Warn().
439 + Err(err).
440 + Str("relay", relayDescriptor.APIHTTPSAddr).
441 + Msg("skipping conflicting discovery relay hint")
442 continue
443 }
672 - existing[relayURL] = struct{}{}
673 - s.knownRelayURLs = append(s.knownRelayURLs, relayURL)
674 - added = append(added, relayURL)
675 - }
676 - for _, relayURL := range normalized {
677 - state := s.localByURL[relayURL]
678 - state.Bootstrap = true
679 - state.Reachable = false
680 - s.localByURL[relayURL] = state
681 - }
682 - s.logStatusChange()
683 - if len(added) == 0 {
684 - return nil, nil
685 - }
686 - return added, nil
444 + state := s.localByURL[relayDescriptor.APIHTTPSAddr]
445 + switch {
446 + case state.Expired:
447 + // Fresh hint re-enables direct confirmation but must not restore
448 + // advertisement until the relay confirms itself again.
449 + state.Advertised = false
450 + state.Expired = false
451 + state.ConsecutiveFailures = 0
452 + case !state.Advertised:
453 + state.Expired = false
454 + state.ConsecutiveFailures = 0
455 + }
456 + s.localByURL[relayDescriptor.APIHTTPSAddr] = state
457 + }
458 + s.syncActiveLocked(now)
459 + return nil
460 }
461
689 -func (s *RelaySet) RecordDiscoveryFailure(identity types.Identity, relayURL string, err error, recoveryFailures int, now time.Time) (expired bool, expireReason string, consecutiveFailures int) {
462 +func (s *RelaySet) RecordDiscoveryFailure(identity types.Identity, relayURL string, err error) (expired bool, expireReason string, consecutiveFailures int) {
463 if s == nil {
464 return false, "", 0
465 }
@@ -694,34 +467,31 @@ func (s *RelaySet) RecordDiscoveryFailure(identity types.Identity, relayURL stri
467 if relayKey == "" {
468 return false, "", 0
469 }
697 - relayURL = strings.TrimSpace(relayURL)
698 - if relayURL == "" {
699 - return false, "", 0
700 - }
701 - if now.IsZero() {
702 - now = time.Now().UTC()
703 - }
470
471 s.mu.Lock()
472 defer s.mu.Unlock()
473
708 - view, ok := s.relays[relayKey]
474 + desc, ok := s.relays[relayKey]
475 if !ok {
476 return false, "", 0
477 }
478 + relayURL = strings.TrimSpace(relayURL)
479 + if relayURL == "" || s.relayKeysByURL[relayURL] != relayKey {
480 + relayURL = desc.APIHTTPSAddr
481 + }
482 + if relayURL == "" {
483 + return false, "", 0
484 + }
485
713 - localState := s.localByURL[relayURL]
714 - localState.Reachable = false
715 - localState.ConsecutiveFailures++
716 - localState.LastFailureAt = now
717 - s.localByURL[relayURL] = localState
718 - s.logStatusChange()
719 - if !localState.Expired && localState.ConsecutiveFailures >= recoveryFailures {
720 - state := s.localByURL[view.Descriptor.APIHTTPSAddr]
486 + state := s.localByURL[relayURL]
487 + state.ConsecutiveFailures++
488 + s.localByURL[relayURL] = state
489 + if !state.Expired && state.ConsecutiveFailures >= discoveryRecoveryFailures {
490 state.Expired = true
722 - s.localByURL[view.Descriptor.APIHTTPSAddr] = state
723 - s.logStatusChange()
724 - return true, "recovery", localState.ConsecutiveFailures
491 + state.Advertised = false
492 + s.localByURL[relayURL] = state
493 + s.syncActiveLocked(time.Now().UTC())
494 + return true, "recovery", state.ConsecutiveFailures
495 }
496
497 var apiErr *types.APIRequestError
@@ -729,11 +499,114 @@ func (s *RelaySet) RecordDiscoveryFailure(identity types.Identity, relayURL stri
499 (apiErr.StatusCode == http.StatusForbidden ||
500 apiErr.StatusCode == http.StatusNotFound ||
501 apiErr.StatusCode == http.StatusGone) {
732 - state := s.localByURL[view.Descriptor.APIHTTPSAddr]
502 state.Expired = true
734 - s.localByURL[view.Descriptor.APIHTTPSAddr] = state
735 - s.logStatusChange()
736 - return true, "status", localState.ConsecutiveFailures
503 + state.Advertised = false
504 + s.localByURL[relayURL] = state
505 + s.syncActiveLocked(time.Now().UTC())
506 + return true, "status", state.ConsecutiveFailures
507 + }
508 + return false, "", state.ConsecutiveFailures
509 +}
510 +
511 +func logBootstrapDiscoveryFailure(relayURL string, err error) {
512 + if statusCode, code, unavailable := DiscoveryUnavailableStatus(err); unavailable {
513 + event := log.Info().Str("relay", relayURL)
514 + if statusCode > 0 {
515 + event = event.Int("status_code", statusCode)
516 + }
517 + if code != "" {
518 + event = event.Str("code", code)
519 + }
520 + event.Msg("bootstrap relay discovery unavailable")
521 + return
522 + }
523 +
524 + log.Warn().
525 + Err(err).
526 + Str("relay", relayURL).
527 + Msg("bootstrap relay discovery failed")
528 +}
529 +
530 +func logDirectDiscoveryFailure(relayURL string, err error, expired bool, expireReason string, consecutiveFailures int) {
531 + event := log.Warn().
532 + Err(err).
533 + Str("relay", relayURL)
534 + if expired {
535 + event = event.
536 + Bool("expired", true).
537 + Str("reason", expireReason)
538 + if consecutiveFailures > 0 {
539 + event = event.Int("consecutive_failures", consecutiveFailures)
540 + }
541 + }
542 + event.Msg("direct relay discovery failed")
543 +}
544 +
545 +func (s *RelaySet) refresh(ctx context.Context, rootCAPEM []byte) {
546 + if s == nil {
547 + return
548 + }
549 +
550 + for _, bootstrap := range s.bootstrapDescriptors() {
551 + resp, err := DiscoverRelayDiscovery(ctx, bootstrap.APIHTTPSAddr, rootCAPEM, nil)
552 + if err != nil {
553 + if ctx.Err() != nil {
554 + return
555 + }
556 + logBootstrapDiscoveryFailure(bootstrap.APIHTTPSAddr, err)
557 + continue
558 + }
559 + if err := s.ApplyRelayDiscoveryResponse(bootstrap.Identity, bootstrap.APIHTTPSAddr, resp, time.Now().UTC()); err != nil {
560 + log.Warn().
561 + Err(err).
562 + Str("relay", bootstrap.APIHTTPSAddr).
563 + Msg("bootstrap relay discovery failed")
564 + }
565 + }
566 + if ctx.Err() != nil {
567 + return
568 + }
569 +
570 + for _, relay := range s.confirmableDescriptors() {
571 + resp, err := DiscoverRelayDiscovery(ctx, relay.APIHTTPSAddr, rootCAPEM, nil)
572 + if err != nil {
573 + if ctx.Err() != nil {
574 + return
575 + }
576 + expired, expireReason, consecutiveFailures := s.RecordDiscoveryFailure(relay.Identity, relay.APIHTTPSAddr, err)
577 + logDirectDiscoveryFailure(relay.APIHTTPSAddr, err, expired, expireReason, consecutiveFailures)
578 + continue
579 + }
580 + if err := s.ApplyRelayDiscoveryResponse(relay.Identity, relay.APIHTTPSAddr, resp, time.Now().UTC()); err != nil {
581 + expired, expireReason, consecutiveFailures := s.RecordDiscoveryFailure(relay.Identity, relay.APIHTTPSAddr, err)
582 + logDirectDiscoveryFailure(relay.APIHTTPSAddr, err, expired, expireReason, consecutiveFailures)
583 + }
584 + }
585 +}
586 +
587 +func (s *RelaySet) RunLoop(ctx context.Context, rootCAPEM []byte, syncRuntime func() error) error {
588 + if s == nil {
589 + return nil
590 + }
591 +
592 + ticker := time.NewTicker(types.DiscoveryPollInterval)
593 + defer ticker.Stop()
594 +
595 + for {
596 + s.refresh(ctx, rootCAPEM)
597 + if ctx.Err() != nil {
598 + return nil
599 + }
600 + if syncRuntime != nil {
601 + if err := syncRuntime(); err != nil {
602 + return err
603 + }
604 + }
605 +
606 + select {
607 + case <-ctx.Done():
608 + return nil
609 + case <-ticker.C:
610 + }
611 }
738 - return false, "", localState.ConsecutiveFailures
612 }
portal/server.go
+30 -234
@@ -22,43 +22,35 @@ import (
22 "github.com/gosuda/portal/v2/portal/keyless"
23 "github.com/gosuda/portal/v2/portal/policy"
24 "github.com/gosuda/portal/v2/portal/transport"
25 - "github.com/gosuda/portal/v2/portal/wireguard"
25 "github.com/gosuda/portal/v2/types"
26 "github.com/gosuda/portal/v2/utils"
27 )
28
29 const (
31 - defaultLeaseTTL = 30 * time.Second
32 - defaultClaimTimeout = 10 * time.Second
33 - defaultIdleKeepalive = 15 * time.Second
34 - defaultReadyQueueLimit = 8
35 - defaultClientHelloWait = 2 * time.Second
36 - defaultControlBodyLimit = 4 << 20
37 - defaultWGRecoveryFailures = 3
30 + defaultLeaseTTL = 30 * time.Second
31 + defaultClaimTimeout = 10 * time.Second
32 + defaultIdleKeepalive = 15 * time.Second
33 + defaultReadyQueueLimit = 8
34 + defaultClientHelloWait = 2 * time.Second
35 + defaultControlBodyLimit = 4 << 20
36 )
37
38 type ServerConfig struct {
41 - PortalURL string
42 - IdentityPath string
43 - Bootstraps []string
44 - WireGuardPrivateKey string
45 - DiscoveryPort int
46 - WireGuardPublicKey string
47 - WireGuardEndpoint string
48 - OverlayIPv4 string
49 - OverlayCIDRs []string
50 - ACME acme.Config
51 - APIPort int
52 - SNIPort int
53 - APIListenAddr string
54 - SNIListenAddr string
55 - TrustedProxyCIDRs string
56 - TrustProxyHeaders bool
57 - DiscoveryEnabled bool
58 - MinPort int
59 - MaxPort int
60 - UDPEnabled bool
61 - TCPEnabled bool
39 + PortalURL string
40 + IdentityPath string
41 + Bootstraps []string
42 + ACME acme.Config
43 + APIPort int
44 + SNIPort int
45 + APIListenAddr string
46 + SNIListenAddr string
47 + TrustedProxyCIDRs string
48 + TrustProxyHeaders bool
49 + DiscoveryEnabled bool
50 + MinPort int
51 + MaxPort int
52 + UDPEnabled bool
53 + TCPEnabled bool
54 }
55
56 type Server struct {
@@ -68,14 +60,12 @@ type Server struct {
60 apiTLSClose io.Closer
61 acmeManager *acme.Manager
62 quicTunnel *quic.Listener
71 - overlay *wireguard.Overlay
63 cancel context.CancelFunc
64 group *errgroup.Group
65 registry *leaseRegistry
66 ports *transport.PortAllocator
67 tcpPorts *transport.PortAllocator
68 identity types.Identity
78 - wgConfig wireguard.Config
69 cfg ServerConfig
70 trustedProxyCIDRs []*net.IPNet
71 relaySet *discovery.RelaySet
@@ -101,31 +91,6 @@ func NewServer(cfg ServerConfig) (*Server, error) {
91 return nil, fmt.Errorf("normalize bootstraps: %w", err)
92 }
93 cfg.Bootstraps = bootstraps
104 - generatedWireGuardPrivateKey := ""
105 - if cfg.DiscoveryEnabled && strings.TrimSpace(cfg.WireGuardPrivateKey) == "" {
106 - generatedWireGuardPrivateKey, err = utils.GenerateWireGuardPrivateKey()
107 - if err != nil {
108 - return nil, err
109 - }
110 - cfg.WireGuardPrivateKey = generatedWireGuardPrivateKey
111 - }
112 - wgConfig, err := wireguard.NormalizeConfig(rootHost, wireguard.Config{
113 - PrivateKey: cfg.WireGuardPrivateKey,
114 - PublicKey: cfg.WireGuardPublicKey,
115 - Endpoint: cfg.WireGuardEndpoint,
116 - OverlayIPv4: cfg.OverlayIPv4,
117 - OverlayCIDRs: cfg.OverlayCIDRs,
118 - ListenPort: cfg.DiscoveryPort,
119 - })
120 - if err != nil {
121 - return nil, err
122 - }
123 - if generatedWireGuardPrivateKey != "" {
124 - log.Warn().
125 - Str("wireguard_public_key", wgConfig.PublicKey).
126 - Str("wireguard_private_key", generatedWireGuardPrivateKey).
127 - Msg("generated wireguard private key; set WIREGUARD_PRIVATE_KEY to preserve relay identity")
128 - }
94
95 transportEnabled := cfg.UDPEnabled || cfg.TCPEnabled
96 hasPortRange := cfg.MinPort > 0 && cfg.MaxPort > 0
@@ -159,6 +124,11 @@ func NewServer(cfg ServerConfig) (*Server, error) {
124 Str("address", identity.Address).
125 Msg("generated relay identity and saved it to disk")
126 }
127 + selfRelayURL, err := utils.NormalizeRelayURL(cfg.PortalURL)
128 + if err != nil {
129 + return nil, fmt.Errorf("normalize portal url: %w", err)
130 + }
131 + cfg.Bootstraps = utils.RemoveRelayURL(cfg.Bootstraps, selfRelayURL)
132
133 tcpPortMin, tcpPortMax := 0, 0
134 if cfg.TCPEnabled {
@@ -179,16 +149,15 @@ func NewServer(cfg ServerConfig) (*Server, error) {
149 ports: ports,
150 tcpPorts: tcpPorts,
151 identity: identity,
182 - wgConfig: wgConfig,
152 trustedProxyCIDRs: trustedProxyCIDRs,
153 }
154
155 if cfg.DiscoveryEnabled {
156 s.relaySet = discovery.NewRelaySet()
188 - _, err = s.relaySet.RegisterBootstrapRelayURLs(cfg.Bootstraps)
189 - if err != nil {
190 - return nil, err
157 + if err := s.relaySet.SetSelfRelay(identity, selfRelayURL); err != nil {
158 + return nil, fmt.Errorf("set self relay: %w", err)
159 }
160 + s.relaySet.SetBootstrapRelayURLs(cfg.Bootstraps)
161 }
162
163 return s, nil
@@ -237,25 +206,11 @@ func (s *Server) Start(ctx context.Context, apiMux *http.ServeMux) error {
206 s.cancel = cancel
207 s.group = group
208
240 - if s.wgConfig.PrivateKey != "" {
241 - if err := s.startOverlay(); err != nil {
242 - acmeManager.Stop()
243 - _ = apiServer.Close()
244 - _ = apiCloser.Close()
245 - _ = sniListener.Close()
246 - cancel()
247 - return err
248 - }
249 - }
250 -
209 group.Go(s.runAPIServer)
252 - if s.overlay != nil {
253 - group.Go(s.overlay.Serve)
254 - }
210 group.Go(func() error { return s.runSNIListener(groupCtx) })
211 group.Go(func() error { return s.runLeaseJanitor(groupCtx, 5*time.Second) })
212 if s.cfg.DiscoveryEnabled {
258 - group.Go(func() error { return s.runRelayDiscoveryLoop(groupCtx) })
213 + group.Go(func() error { return s.relaySet.RunLoop(groupCtx, nil, nil) })
214 }
215 s.acmeManager.Start(serverCtx)
216
@@ -279,7 +234,6 @@ func (s *Server) Start(ctx context.Context, apiMux *http.ServeMux) error {
234 Int("min_port", s.cfg.MinPort).
235 Int("max_port", s.cfg.MaxPort).
236 Bool("discovery_enabled", s.cfg.DiscoveryEnabled).
282 - Bool("wireguard_enabled", s.wgConfig.PrivateKey != "").
237 Bool("udp_enabled", s.cfg.UDPEnabled).
238 Bool("tcp_enabled", s.cfg.TCPEnabled)
239 if s.quicTunnel != nil {
@@ -341,11 +295,6 @@ func (s *Server) Shutdown(ctx context.Context) error {
295 shutdownErr = err
296 }
297 }
344 - if s.overlay != nil {
345 - if err := s.overlay.Shutdown(ctx); err != nil && shutdownErr == nil {
346 - shutdownErr = err
347 - }
348 - }
298 if s.apiTLSClose != nil {
299 _ = s.apiTLSClose.Close()
300 }
@@ -608,159 +557,6 @@ func (s *Server) runQUICTunnelListener(listener *quic.Listener) error {
557 }
558 }
559
611 -func (s *Server) startOverlay() error {
612 - peerMux := http.NewServeMux()
613 - peerMux.HandleFunc(types.PathRoot, s.handleRoot)
614 - peerMux.HandleFunc(types.PathHealthz, s.handleHealthz)
615 - peerMux.HandleFunc(types.PathDiscovery, func(w http.ResponseWriter, r *http.Request) {
616 - if !s.cfg.DiscoveryEnabled {
617 - http.NotFound(w, r)
618 - return
619 - }
620 - s.handleRelayDiscovery(w, r)
621 - })
622 -
623 - overlay, err := wireguard.NewOverlay(s.wgConfig, peerMux)
624 - if err != nil {
625 - return fmt.Errorf("start wireguard overlay: %w", err)
626 - }
627 -
628 - if err := overlay.Sync(s.identity.Key(), s.relaySet.Snapshot()); err != nil {
629 - _ = overlay.Shutdown(context.Background())
630 - return fmt.Errorf("sync wireguard peers: %w", err)
631 - }
632 -
633 - s.overlay = overlay
634 - return nil
635 -}
636 -
637 -func (s *Server) runRelayDiscoveryLoop(ctx context.Context) error {
638 - ticker := time.NewTicker(types.DiscoveryPollInterval)
639 - defer ticker.Stop()
640 -
641 - for {
642 - bootstraps := s.relaySet.BootstrapDescriptors()
643 -
644 - for _, bootstrap := range bootstraps {
645 - resp, err := discovery.DiscoverRelayDiscovery(ctx, bootstrap.APIHTTPSAddr, nil, nil)
646 - if err != nil {
647 - if ctx.Err() != nil {
648 - return nil
649 - }
650 - s.relaySet.RecordBootstrapDiscoveryFailure(bootstrap.APIHTTPSAddr, err, time.Now().UTC())
651 - continue
652 - }
653 -
654 - now := time.Now().UTC()
655 - var relaySetChanged bool
656 - var warnErr error
657 - _, relaySetChanged, _, warnErr, err = s.relaySet.ApplyRelayDiscoveryResponse(bootstrap.Identity, bootstrap.APIHTTPSAddr, resp, now)
658 - if relaySetChanged && s.overlay != nil {
659 - if syncErr := s.overlay.Sync(s.identity.Key(), s.relaySet.Snapshot()); syncErr != nil {
660 - if warnErr == nil {
661 - warnErr = syncErr
662 - }
663 - }
664 - }
665 - if err != nil {
666 - s.relaySet.MarkRelayFailure(bootstrap.APIHTTPSAddr, time.Now().UTC())
667 - log.Warn().
668 - Err(err).
669 - Str("relay", bootstrap.APIHTTPSAddr).
670 - Msg("bootstrap relay discovery failed")
671 - continue
672 - }
673 -
674 - if warnErr != nil {
675 - log.Warn().
676 - Err(warnErr).
677 - Str("relay", bootstrap.APIHTTPSAddr).
678 - Msg("bootstrap relay discovery completed with warnings")
679 - }
680 - }
681 - if ctx.Err() != nil {
682 - return nil
683 - }
684 -
685 - if s.overlay != nil {
686 - overlayClient := s.overlay.Client()
687 - syncableRelays := s.relaySet.SyncableDescriptors()
688 -
689 - for _, relay := range syncableRelays {
690 - var failureErr error
691 -
692 - if err := discovery.RequireOverlayRelayDescriptor(relay); err != nil {
693 - failureErr = err
694 - } else {
695 - discoverURL := "http://" + net.JoinHostPort(relay.OverlayIPv4, fmt.Sprintf("%d", wireguard.DefaultPeerAPIHTTPPort))
696 - resp, err := discovery.DiscoverRelayDiscovery(ctx, discoverURL, nil, overlayClient)
697 - if err != nil {
698 - if ctx.Err() != nil {
699 - return nil
700 - }
701 - failureErr = err
702 - } else {
703 - now := time.Now().UTC()
704 - var relaySetChanged bool
705 - var warnErr error
706 - var snapshot map[string]types.RelayState
707 - _, relaySetChanged, _, warnErr, err = s.relaySet.ApplyOverlayRelayDiscoveryResponse(relay.Identity, relay.APIHTTPSAddr, resp, now)
708 - if relaySetChanged {
709 - snapshot = s.relaySet.Snapshot()
710 - if syncErr := s.overlay.Sync(s.identity.Key(), snapshot); syncErr != nil {
711 - if warnErr == nil {
712 - warnErr = syncErr
713 - }
714 - }
715 - }
716 - if err != nil {
717 - failureErr = err
718 - } else {
719 - if warnErr != nil {
720 - log.Warn().
721 - Err(warnErr).
722 - Str("relay", relay.APIHTTPSAddr).
723 - Msg("overlay relay discovery completed with warnings")
724 - continue
725 - }
726 -
727 - continue
728 - }
729 - }
730 - }
731 - expired, expireReason, consecutiveFailures := s.relaySet.RecordDiscoveryFailure(relay.Identity, relay.APIHTTPSAddr, failureErr, defaultWGRecoveryFailures, time.Now().UTC())
732 - if expired {
733 - if syncErr := s.overlay.Sync(s.identity.Key(), s.relaySet.Snapshot()); syncErr != nil && failureErr == nil {
734 - failureErr = syncErr
735 - }
736 - }
737 -
738 - event := log.Warn().
739 - Err(failureErr).
740 - Str("relay", relay.APIHTTPSAddr)
741 - if expired {
742 - event = event.
743 - Bool("expired", true).
744 - Str("reason", expireReason)
745 - if consecutiveFailures > 0 {
746 - event = event.Int("consecutive_failures", consecutiveFailures)
747 - }
748 - }
749 - event.Msg("overlay relay discovery failed")
750 - }
751 - }
752 - if ctx.Err() != nil {
753 - return nil
754 - }
755 -
756 - select {
757 - case <-ctx.Done():
758 - return nil
759 - case <-ticker.C:
760 - }
761 - }
762 -}
763 -
560 func BridgeConns(left, right net.Conn) {
561 defer left.Close()
562 defer right.Close()
portal/server_test.go
+319 -204
@@ -10,9 +10,9 @@ import (
10 "crypto/x509/pkix"
11 "encoding/json"
12 "encoding/pem"
13 + "errors"
14 "io"
15 "math/big"
15 - "net"
16 "net/http"
17 "os"
18 "path/filepath"
@@ -32,31 +32,15 @@ func mustRelayDescriptor(t *testing.T, relayURL string) types.RelayDescriptor {
32 t.Helper()
33
34 now := time.Now().UTC()
35 - wireGuardPrivateKey, err := utils.NormalizeWireGuardPrivateKey(strings.Repeat("44", 32))
36 - if err != nil {
37 - t.Fatalf("NormalizeWireGuardPrivateKey() error = %v", err)
38 - }
39 - wireGuardPublicKey, err := utils.WireGuardPublicKeyFromPrivate(wireGuardPrivateKey)
40 - if err != nil {
41 - t.Fatalf("WireGuardPublicKeyFromPrivate() error = %v", err)
42 - }
43 - overlayIPv4, err := utils.DeriveWireGuardOverlayIPv4(wireGuardPublicKey)
44 - if err != nil {
45 - t.Fatalf("DeriveWireGuardOverlayIPv4() error = %v", err)
46 - }
35 desc, err := discovery.NormalizeDescriptor(types.RelayDescriptor{
36 Identity: types.Identity{
37 Name: utils.PortalRootHost(relayURL),
38 },
51 - Sequence: uint64(now.UnixMilli()),
52 - Version: 1,
53 - IssuedAt: now,
54 - ExpiresAt: now.Add(time.Hour),
55 - APIHTTPSAddr: relayURL,
56 - WireGuardPublicKey: wireGuardPublicKey,
57 - WireGuardEndpoint: net.JoinHostPort(utils.PortalRootHost(relayURL), "51820"),
58 - OverlayIPv4: overlayIPv4,
59 - SupportsOverlayPeer: true,
39 + Sequence: uint64(now.UnixMilli()),
40 + Version: 1,
41 + IssuedAt: now,
42 + ExpiresAt: now.Add(time.Hour),
43 + APIHTTPSAddr: relayURL,
44 })
45 if err != nil {
46 t.Fatalf("NormalizeDescriptor() error = %v", err)
@@ -111,7 +95,7 @@ func writeManualRelayCertificate(t *testing.T, keyDir, baseDomain string) {
95 }
96 }
97
114 -func TestNewServerGeneratesWireGuardWhenDiscoveryEnabled(t *testing.T) {
98 +func TestNewServerInitializesRelaySetWhenDiscoveryEnabled(t *testing.T) {
99 t.Parallel()
100
101 server, err := NewServer(ServerConfig{
@@ -122,17 +106,8 @@ func TestNewServerGeneratesWireGuardWhenDiscoveryEnabled(t *testing.T) {
106 if err != nil {
107 t.Fatalf("NewServer() error = %v", err)
108 }
125 - if server.wgConfig.PrivateKey == "" {
126 - t.Fatal("WireGuardPrivateKey = empty, want generated key")
127 - }
128 - if server.wgConfig.PublicKey == "" {
129 - t.Fatal("WireGuardPublicKey = empty, want derived key")
130 - }
131 - if server.wgConfig.Endpoint == "" {
132 - t.Fatal("WireGuardEndpoint = empty, want derived endpoint")
133 - }
134 - if server.wgConfig.OverlayIPv4 == "" {
135 - t.Fatal("OverlayIPv4 = empty, want derived overlay address")
109 + if server.relaySet == nil {
110 + t.Fatal("relaySet = nil, want discovery relay set")
111 }
112 }
113
@@ -439,63 +414,6 @@ func TestServerStartRejectsMismatchedACMEBaseDomain(t *testing.T) {
414 }
415 }
416
442 -func TestNewServerDerivesWireGuardConfigFromPrivateKey(t *testing.T) {
443 - t.Parallel()
444 -
445 - server, err := NewServer(ServerConfig{
446 - PortalURL: "https://portal.example.com",
447 - IdentityPath: tempIdentityPath(t),
448 - WireGuardPrivateKey: strings.Repeat("33", 32),
449 - DiscoveryPort: 41011,
450 - })
451 - if err != nil {
452 - t.Fatalf("NewServer() error = %v", err)
453 - }
454 -
455 - if server.wgConfig.PrivateKey == "" {
456 - t.Fatal("WireGuardPrivateKey = empty, want normalized key")
457 - }
458 - if server.wgConfig.PublicKey == "" {
459 - t.Fatal("WireGuardPublicKey = empty, want derived key")
460 - }
461 - if server.wgConfig.Endpoint != net.JoinHostPort("portal.example.com", "41011") {
462 - t.Fatalf("WireGuardEndpoint = %q, want %q", server.wgConfig.Endpoint, net.JoinHostPort("portal.example.com", "41011"))
463 - }
464 - if server.wgConfig.OverlayIPv4 == "" {
465 - t.Fatal("OverlayIPv4 = empty, want derived overlay address")
466 - }
467 - if err := utils.ValidateWireGuardEndpoint(server.wgConfig.Endpoint); err != nil {
468 - t.Fatalf("ValidateWireGuardEndpoint() error = %v", err)
469 - }
470 - if err := utils.ValidateOverlayIPv4(server.wgConfig.OverlayIPv4); err != nil {
471 - t.Fatalf("ValidateOverlayIPv4() error = %v", err)
472 - }
473 -
474 - wantOverlay, err := utils.DeriveWireGuardOverlayIPv4(server.wgConfig.PublicKey)
475 - if err != nil {
476 - t.Fatalf("DeriveWireGuardOverlayIPv4() error = %v", err)
477 - }
478 - if server.wgConfig.OverlayIPv4 != wantOverlay {
479 - t.Fatalf("OverlayIPv4 = %q, want %q", server.wgConfig.OverlayIPv4, wantOverlay)
480 - }
481 -}
482 -
483 -func TestNewServerIgnoresDiscoveryPortWithoutWireGuardKey(t *testing.T) {
484 - t.Parallel()
485 -
486 - server, err := NewServer(ServerConfig{
487 - PortalURL: "https://portal.example.com",
488 - IdentityPath: tempIdentityPath(t),
489 - DiscoveryPort: 51820,
490 - })
491 - if err != nil {
492 - t.Fatalf("NewServer() error = %v", err)
493 - }
494 - if server.wgConfig.Endpoint != "" {
495 - t.Fatalf("WireGuardEndpoint = %q, want empty without wireguard key", server.wgConfig.Endpoint)
496 - }
497 -}
498 -
417 func TestRegisterLeaseDerivesFixedHostnameFromName(t *testing.T) {
418 t.Parallel()
419
@@ -590,60 +508,80 @@ func TestRegisterLeaseBuildsUDPEnabledRuntime(t *testing.T) {
508 }
509 }
510
593 -func TestServerUpsertDiscoverySeedURLsSkipsLocalRelayHosts(t *testing.T) {
511 +func TestServerSetBootstrapRelayURLsAllowsLoopbackButSkipsSelfRelay(t *testing.T) {
512 t.Parallel()
513
514 server, err := NewServer(ServerConfig{
597 - PortalURL: "https://portal.example.com",
598 - IdentityPath: tempIdentityPath(t),
599 - Bootstraps: []string{"https://bootstrap.example.com"},
600 - WireGuardPrivateKey: strings.Repeat("23", 32),
601 - DiscoveryPort: 41022,
602 - DiscoveryEnabled: true,
515 + PortalURL: "https://relay-a.example.com",
516 + IdentityPath: tempIdentityPath(t),
517 + Bootstraps: []string{"https://bootstrap.example.com"},
518 + DiscoveryEnabled: true,
519 })
520 if err != nil {
521 t.Fatalf("NewServer() error = %v", err)
522 }
523
608 - added, err := server.relaySet.RegisterBootstrapRelayURLs([]string{
524 + server.relaySet.SetBootstrapRelayURLs([]string{
525 + "https://bootstrap.example.com",
526 "https://localhost:4017",
527 "https://relay-a.example.com",
611 - "https://127.0.0.1:4017",
528 + "https://relay-b.example.com",
529 })
613 - if err != nil {
614 - t.Fatalf("UpsertSeedURLs() error = %v", err)
530 + advertisedDescriptors := server.relaySet.ActiveRelayDescriptors()
531 + knownURLs := append([]string(nil), server.relaySet.ActiveRelayURLs()...)
532 + sort.Strings(knownURLs)
533 + if !reflect.DeepEqual(knownURLs, []string{
534 + "https://bootstrap.example.com",
535 + "https://localhost:4017",
536 + "https://relay-b.example.com",
537 + }) {
538 + t.Fatalf("ActiveRelayURLs() = %v, want loopback kept and self filtered", knownURLs)
539 }
616 -
617 - if !reflect.DeepEqual(added, []string{"https://relay-a.example.com"}) {
618 - t.Fatalf("UpsertSeedURLs() added = %v, want [%q]", added, "https://relay-a.example.com")
540 + if len(advertisedDescriptors) != 0 {
541 + t.Fatalf("advertised count = %d, want 0 before direct confirmation", len(advertisedDescriptors))
542 }
620 - knownRelayURLs, err := utils.ExcludeLocalRelayURLs("https://bootstrap.example.com", "https://relay-a.example.com")
543 +}
544 +
545 +func TestServerDiscoverySkipsSelfRelayHint(t *testing.T) {
546 + t.Parallel()
547 +
548 + server, err := NewServer(ServerConfig{
549 + PortalURL: "https://portal.example.com",
550 + IdentityPath: tempIdentityPath(t),
551 + Bootstraps: []string{"https://bootstrap.example.com"},
552 + DiscoveryEnabled: true,
553 + })
554 if err != nil {
622 - t.Fatalf("ExcludeLocalRelayURLs() error = %v", err)
623 - }
624 - bootstrapDescriptors := server.relaySet.BootstrapDescriptors()
625 - syncableDescriptors := server.relaySet.SyncableDescriptors()
626 - advertisedDescriptors := server.relaySet.AdvertisedDescriptors()
627 - knownURLs := make([]string, 0, len(bootstrapDescriptors))
628 - for _, descriptor := range bootstrapDescriptors {
629 - if strings.TrimSpace(descriptor.APIHTTPSAddr) == "" {
630 - continue
631 - }
632 - knownURLs = append(knownURLs, descriptor.APIHTTPSAddr)
555 + t.Fatalf("NewServer() error = %v", err)
556 }
634 - sort.Strings(knownURLs)
635 - knownURLs, err = utils.ExcludeLocalRelayURLs(knownURLs...)
557 +
558 + now := time.Now().UTC()
559 + bootstrapDesc := mustRelayDescriptor(t, "https://bootstrap.example.com")
560 + selfHint, err := discovery.NormalizeDescriptor(types.RelayDescriptor{
561 + Identity: server.identity.Copy(),
562 + Sequence: uint64(now.UnixMilli()),
563 + Version: 1,
564 + IssuedAt: now,
565 + ExpiresAt: now.Add(time.Hour),
566 + APIHTTPSAddr: "https://self-mirror.example.com",
567 + })
568 if err != nil {
637 - t.Fatalf("ExcludeLocalRelayURLs() known error = %v", err)
569 + t.Fatalf("NormalizeDescriptor() self hint error = %v", err)
570 }
639 - if !reflect.DeepEqual(knownURLs, knownRelayURLs) {
640 - t.Fatalf("BootstrapDescriptors() = %v, want [%q %q]", knownURLs, "https://bootstrap.example.com", "https://relay-a.example.com")
641 - }
642 - if len(syncableDescriptors) != 0 {
643 - t.Fatalf("syncable count = %d, want 0 before direct confirmation", len(syncableDescriptors))
571 +
572 + if err := server.relaySet.ApplyRelayDiscoveryResponse(
573 + bootstrapDesc.Identity,
574 + bootstrapDesc.APIHTTPSAddr,
575 + types.DiscoveryResponse{ProtocolVersion: types.ProtocolVersion, Self: bootstrapDesc, Relays: []types.RelayDescriptor{selfHint}},
576 + now,
577 + ); err != nil {
578 + t.Fatalf("ApplyRelayDiscoveryResponse() error = %v", err)
579 }
645 - if len(advertisedDescriptors) != 0 {
646 - t.Fatalf("advertised count = %d, want 0 before direct confirmation", len(advertisedDescriptors))
580 +
581 + knownURLs := append([]string(nil), server.relaySet.ActiveRelayURLs()...)
582 + sort.Strings(knownURLs)
583 + if !reflect.DeepEqual(knownURLs, []string{"https://bootstrap.example.com"}) {
584 + t.Fatalf("ActiveRelayURLs() = %v, want self hint excluded", knownURLs)
585 }
586 }
587
@@ -651,12 +589,10 @@ func TestServerRecordVerifiedDiscoveryPeerRequiresDirectConfirmation(t *testing.
589 t.Parallel()
590
591 server, err := NewServer(ServerConfig{
654 - PortalURL: "https://portal.example.com",
655 - IdentityPath: tempIdentityPath(t),
656 - Bootstraps: []string{"https://bootstrap.example.com"},
657 - WireGuardPrivateKey: strings.Repeat("24", 32),
658 - DiscoveryPort: 41023,
659 - DiscoveryEnabled: true,
592 + PortalURL: "https://portal.example.com",
593 + IdentityPath: tempIdentityPath(t),
594 + Bootstraps: []string{"https://bootstrap.example.com"},
595 + DiscoveryEnabled: true,
596 })
597 if err != nil {
598 t.Fatalf("NewServer() error = %v", err)
@@ -665,90 +601,120 @@ func TestServerRecordVerifiedDiscoveryPeerRequiresDirectConfirmation(t *testing.
601 bootstrapDesc := mustRelayDescriptor(t, "https://bootstrap.example.com")
602 relayADesc := mustRelayDescriptor(t, "https://relay-a.example.com")
603
668 - applyDiscovery := func(targetIdentity types.Identity, targetURL string, resp types.DiscoveryResponse, requireSelfOverlay bool) (bool, int, error, error) {
604 + applyDiscovery := func(targetIdentity types.Identity, targetURL string, resp types.DiscoveryResponse) error {
605 now := time.Now().UTC()
670 - if requireSelfOverlay {
671 - _, updated, added, warnErr, err := server.relaySet.ApplyOverlayRelayDiscoveryResponse(targetIdentity, targetURL, resp, now)
672 - return updated, added, warnErr, err
673 - }
674 - _, updated, added, warnErr, err := server.relaySet.ApplyRelayDiscoveryResponse(targetIdentity, targetURL, resp, now)
675 - return updated, added, warnErr, err
606 + return server.relaySet.ApplyRelayDiscoveryResponse(targetIdentity, targetURL, resp, now)
607 }
608
678 - resultUpdated, resultAdded, warnErr, err := applyDiscovery(
609 + err = applyDiscovery(
610 bootstrapDesc.Identity,
611 bootstrapDesc.APIHTTPSAddr,
612 types.DiscoveryResponse{ProtocolVersion: types.ProtocolVersion, Self: bootstrapDesc},
682 - false,
613 )
614 if err != nil {
615 t.Fatalf("applyRelayDiscoveryResponse() bootstrap error = %v", err)
616 }
687 - if warnErr != nil {
688 - t.Fatalf("applyRelayDiscoveryResponse() bootstrap warn = %v, want nil", warnErr)
689 - }
690 - if resultAdded != 0 {
691 - t.Fatalf("applyRelayDiscoveryResponse() bootstrap added = %d, want 0 for seeded bootstrap", resultAdded)
692 - }
693 - if !resultUpdated {
694 - t.Fatal("applyRelayDiscoveryResponse() bootstrap updated = false, want true")
695 - }
617
697 - resultUpdated, resultAdded, warnErr, err = applyDiscovery(
618 + err = applyDiscovery(
619 bootstrapDesc.Identity,
620 bootstrapDesc.APIHTTPSAddr,
621 types.DiscoveryResponse{ProtocolVersion: types.ProtocolVersion, Self: bootstrapDesc, Relays: []types.RelayDescriptor{relayADesc}},
701 - false,
622 )
623 if err != nil {
624 t.Fatalf("applyRelayDiscoveryResponse() hinted error = %v", err)
625 }
706 - if warnErr != nil {
707 - t.Fatalf("applyRelayDiscoveryResponse() hinted warn = %v, want nil", warnErr)
708 - }
709 - if resultAdded != 1 {
710 - t.Fatalf("applyRelayDiscoveryResponse() hinted added = %d, want 1", resultAdded)
711 - }
712 - if !resultUpdated {
713 - t.Fatal("applyRelayDiscoveryResponse() hinted updated = false, want true")
714 - }
715 - snapshot := server.relaySet.Snapshot()
716 - if len(snapshot) != 2 {
717 - t.Fatalf("Snapshot() size = %d, want 2 after hinted relay registration", len(snapshot))
626 + knownURLs := append([]string(nil), server.relaySet.ActiveRelayURLs()...)
627 + sort.Strings(knownURLs)
628 + if !reflect.DeepEqual(knownURLs, []string{"https://bootstrap.example.com"}) {
629 + t.Fatalf("ActiveRelayURLs() = %v, want [%q]", knownURLs, "https://bootstrap.example.com")
630 }
719 - bootstrapDescriptors := server.relaySet.BootstrapDescriptors()
720 - knownURLs := make([]string, 0, len(bootstrapDescriptors))
721 - for _, descriptor := range bootstrapDescriptors {
631 + advertisedDescriptors := server.relaySet.ActiveRelayDescriptors()
632 + advertisedURLs := make([]string, 0, len(advertisedDescriptors))
633 + for _, descriptor := range advertisedDescriptors {
634 if strings.TrimSpace(descriptor.APIHTTPSAddr) == "" {
635 continue
636 }
725 - knownURLs = append(knownURLs, descriptor.APIHTTPSAddr)
637 + advertisedURLs = append(advertisedURLs, descriptor.APIHTTPSAddr)
638 }
727 - sort.Strings(knownURLs)
728 - knownURLs, err = utils.ExcludeLocalRelayURLs(knownURLs...)
729 - if err != nil {
730 - t.Fatalf("ExcludeLocalRelayURLs() known error = %v", err)
639 + sort.Strings(advertisedURLs)
640 + if !reflect.DeepEqual(advertisedURLs, []string{"https://bootstrap.example.com"}) {
641 + t.Fatalf("ActiveRelayDescriptors() = %v, want [%q]", advertisedURLs, "https://bootstrap.example.com")
642 }
732 - if !reflect.DeepEqual(knownURLs, []string{"https://bootstrap.example.com"}) {
733 - t.Fatalf("BootstrapDescriptors() = %v, want [%q]", knownURLs, "https://bootstrap.example.com")
643 +
644 + err = applyDiscovery(
645 + relayADesc.Identity,
646 + relayADesc.APIHTTPSAddr,
647 + types.DiscoveryResponse{ProtocolVersion: types.ProtocolVersion, Self: relayADesc},
648 + )
649 + if err != nil {
650 + t.Fatalf("applyRelayDiscoveryResponse() confirm error = %v", err)
651 }
735 - syncableDescriptors := server.relaySet.SyncableDescriptors()
736 - syncableURLs := make([]string, 0, len(syncableDescriptors))
737 - for _, descriptor := range syncableDescriptors {
652 + advertisedDescriptors = server.relaySet.ActiveRelayDescriptors()
653 + advertisedURLs = advertisedURLs[:0]
654 + for _, descriptor := range advertisedDescriptors {
655 if strings.TrimSpace(descriptor.APIHTTPSAddr) == "" {
656 continue
657 }
741 - syncableURLs = append(syncableURLs, descriptor.APIHTTPSAddr)
658 + advertisedURLs = append(advertisedURLs, descriptor.APIHTTPSAddr)
659 + }
660 + sort.Strings(advertisedURLs)
661 + if !reflect.DeepEqual(advertisedURLs, []string{"https://bootstrap.example.com", "https://relay-a.example.com"}) {
662 + t.Fatalf("ActiveRelayDescriptors() = %v, want [%q %q]", advertisedURLs, "https://bootstrap.example.com", "https://relay-a.example.com")
663 }
743 - sort.Strings(syncableURLs)
744 - syncableURLs, err = utils.ExcludeLocalRelayURLs(syncableURLs...)
664 +}
665 +
666 +func TestServerRecordVerifiedDiscoveryPeerExpiresAfterRepeatedDirectFailures(t *testing.T) {
667 + t.Parallel()
668 +
669 + server, err := NewServer(ServerConfig{
670 + PortalURL: "https://portal.example.com",
671 + IdentityPath: tempIdentityPath(t),
672 + Bootstraps: []string{"https://bootstrap.example.com"},
673 + DiscoveryEnabled: true,
674 + })
675 if err != nil {
746 - t.Fatalf("ExcludeLocalRelayURLs() syncable error = %v", err)
676 + t.Fatalf("NewServer() error = %v", err)
677 }
748 - if !reflect.DeepEqual(syncableURLs, []string{"https://relay-a.example.com"}) {
749 - t.Fatalf("SyncablePeerDescriptors() = %v, want [%q]", syncableURLs, "https://relay-a.example.com")
678 +
679 + bootstrapDesc := mustRelayDescriptor(t, "https://bootstrap.example.com")
680 + relayADesc := mustRelayDescriptor(t, "https://relay-a.example.com")
681 + now := time.Now().UTC()
682 +
683 + if err := server.relaySet.ApplyRelayDiscoveryResponse(
684 + bootstrapDesc.Identity,
685 + bootstrapDesc.APIHTTPSAddr,
686 + types.DiscoveryResponse{ProtocolVersion: types.ProtocolVersion, Self: bootstrapDesc, Relays: []types.RelayDescriptor{relayADesc}},
687 + now,
688 + ); err != nil {
689 + t.Fatalf("ApplyRelayDiscoveryResponse() bootstrap error = %v", err)
690 }
751 - advertisedDescriptors := server.relaySet.AdvertisedDescriptors()
691 + if err := server.relaySet.ApplyRelayDiscoveryResponse(
692 + relayADesc.Identity,
693 + relayADesc.APIHTTPSAddr,
694 + types.DiscoveryResponse{ProtocolVersion: types.ProtocolVersion, Self: relayADesc},
695 + now.Add(time.Second),
696 + ); err != nil {
697 + t.Fatalf("ApplyRelayDiscoveryResponse() direct confirm error = %v", err)
698 + }
699 +
700 + for attempt := 1; attempt <= 3; attempt++ {
701 + expired, _, consecutiveFailures := server.relaySet.RecordDiscoveryFailure(
702 + relayADesc.Identity,
703 + relayADesc.APIHTTPSAddr,
704 + errors.New("direct discovery failed"),
705 + )
706 + if consecutiveFailures != attempt {
707 + t.Fatalf("RecordDiscoveryFailure() consecutive = %d, want %d", consecutiveFailures, attempt)
708 + }
709 + if attempt < 3 && expired {
710 + t.Fatalf("RecordDiscoveryFailure() expired early on attempt %d", attempt)
711 + }
712 + if attempt == 3 && !expired {
713 + t.Fatalf("RecordDiscoveryFailure() expired = false on attempt %d, want true", attempt)
714 + }
715 + }
716 +
717 + advertisedDescriptors := server.relaySet.ActiveRelayDescriptors()
718 advertisedURLs := make([]string, 0, len(advertisedDescriptors))
719 for _, descriptor := range advertisedDescriptors {
720 if strings.TrimSpace(descriptor.APIHTTPSAddr) == "" {
@@ -757,33 +723,168 @@ func TestServerRecordVerifiedDiscoveryPeerRequiresDirectConfirmation(t *testing.
723 advertisedURLs = append(advertisedURLs, descriptor.APIHTTPSAddr)
724 }
725 sort.Strings(advertisedURLs)
760 - advertisedURLs, err = utils.ExcludeLocalRelayURLs(advertisedURLs...)
761 - if err != nil {
762 - t.Fatalf("ExcludeLocalRelayURLs() advertised error = %v", err)
763 - }
726 if !reflect.DeepEqual(advertisedURLs, []string{"https://bootstrap.example.com"}) {
765 - t.Fatalf("AdvertisedDescriptors() = %v, want [%q]", advertisedURLs, "https://bootstrap.example.com")
727 + t.Fatalf("ActiveRelayDescriptors() = %v, want [%q] after relay expiry", advertisedURLs, "https://bootstrap.example.com")
728 + }
729 +}
730 +
731 +func TestServerBootstrapHintDoesNotResetDirectFailureBudget(t *testing.T) {
732 + t.Parallel()
733 +
734 + server, err := NewServer(ServerConfig{
735 + PortalURL: "https://portal.example.com",
736 + IdentityPath: tempIdentityPath(t),
737 + Bootstraps: []string{"https://bootstrap.example.com"},
738 + DiscoveryEnabled: true,
739 + })
740 + if err != nil {
741 + t.Fatalf("NewServer() error = %v", err)
742 }
743
768 - resultUpdated, resultAdded, warnErr, err = applyDiscovery(
744 + bootstrapDesc := mustRelayDescriptor(t, "https://bootstrap.example.com")
745 + relayADesc := mustRelayDescriptor(t, "https://relay-a.example.com")
746 + now := time.Now().UTC()
747 +
748 + if err := server.relaySet.ApplyRelayDiscoveryResponse(
749 + bootstrapDesc.Identity,
750 + bootstrapDesc.APIHTTPSAddr,
751 + types.DiscoveryResponse{ProtocolVersion: types.ProtocolVersion, Self: bootstrapDesc, Relays: []types.RelayDescriptor{relayADesc}},
752 + now,
753 + ); err != nil {
754 + t.Fatalf("ApplyRelayDiscoveryResponse() bootstrap error = %v", err)
755 + }
756 + if err := server.relaySet.ApplyRelayDiscoveryResponse(
757 relayADesc.Identity,
758 relayADesc.APIHTTPSAddr,
759 types.DiscoveryResponse{ProtocolVersion: types.ProtocolVersion, Self: relayADesc},
772 - true,
760 + now.Add(time.Second),
761 + ); err != nil {
762 + t.Fatalf("ApplyRelayDiscoveryResponse() direct confirm error = %v", err)
763 + }
764 +
765 + for attempt := 1; attempt <= 2; attempt++ {
766 + expired, _, consecutiveFailures := server.relaySet.RecordDiscoveryFailure(
767 + relayADesc.Identity,
768 + relayADesc.APIHTTPSAddr,
769 + errors.New("direct discovery failed"),
770 + )
771 + if expired {
772 + t.Fatalf("RecordDiscoveryFailure() expired early on attempt %d", attempt)
773 + }
774 + if consecutiveFailures != attempt {
775 + t.Fatalf("RecordDiscoveryFailure() consecutive = %d, want %d", consecutiveFailures, attempt)
776 + }
777 +
778 + if err := server.relaySet.ApplyRelayDiscoveryResponse(
779 + bootstrapDesc.Identity,
780 + bootstrapDesc.APIHTTPSAddr,
781 + types.DiscoveryResponse{ProtocolVersion: types.ProtocolVersion, Self: bootstrapDesc, Relays: []types.RelayDescriptor{relayADesc}},
782 + now.Add(time.Duration(attempt+1)*time.Second),
783 + ); err != nil {
784 + t.Fatalf("ApplyRelayDiscoveryResponse() hinted refresh error = %v", err)
785 + }
786 +
787 + advertisedDescriptors := server.relaySet.ActiveRelayDescriptors()
788 + advertisedURLs := make([]string, 0, len(advertisedDescriptors))
789 + for _, descriptor := range advertisedDescriptors {
790 + if strings.TrimSpace(descriptor.APIHTTPSAddr) == "" {
791 + continue
792 + }
793 + advertisedURLs = append(advertisedURLs, descriptor.APIHTTPSAddr)
794 + }
795 + sort.Strings(advertisedURLs)
796 + if !reflect.DeepEqual(advertisedURLs, []string{"https://bootstrap.example.com", "https://relay-a.example.com"}) {
797 + t.Fatalf("ActiveRelayDescriptors() = %v, want relay to remain advertised before expiry", advertisedURLs)
798 + }
799 + }
800 +
801 + expired, _, consecutiveFailures := server.relaySet.RecordDiscoveryFailure(
802 + relayADesc.Identity,
803 + relayADesc.APIHTTPSAddr,
804 + errors.New("direct discovery failed"),
805 )
806 + if !expired {
807 + t.Fatal("RecordDiscoveryFailure() expired = false on final attempt, want true")
808 + }
809 + if consecutiveFailures != 3 {
810 + t.Fatalf("RecordDiscoveryFailure() consecutive = %d, want 3", consecutiveFailures)
811 + }
812 +}
813 +
814 +func TestServerExpiredDiscoveryPeerNeedsFreshDirectConfirmation(t *testing.T) {
815 + t.Parallel()
816 +
817 + server, err := NewServer(ServerConfig{
818 + PortalURL: "https://portal.example.com",
819 + IdentityPath: tempIdentityPath(t),
820 + Bootstraps: []string{"https://bootstrap.example.com"},
821 + DiscoveryEnabled: true,
822 + })
823 if err != nil {
775 - t.Fatalf("applyRelayDiscoveryResponse() confirm error = %v", err)
824 + t.Fatalf("NewServer() error = %v", err)
825 }
777 - if warnErr != nil {
778 - t.Fatalf("applyRelayDiscoveryResponse() confirm warn = %v, want nil", warnErr)
826 +
827 + bootstrapDesc := mustRelayDescriptor(t, "https://bootstrap.example.com")
828 + relayADesc := mustRelayDescriptor(t, "https://relay-a.example.com")
829 + now := time.Now().UTC()
830 +
831 + if err := server.relaySet.ApplyRelayDiscoveryResponse(
832 + bootstrapDesc.Identity,
833 + bootstrapDesc.APIHTTPSAddr,
834 + types.DiscoveryResponse{ProtocolVersion: types.ProtocolVersion, Self: bootstrapDesc, Relays: []types.RelayDescriptor{relayADesc}},
835 + now,
836 + ); err != nil {
837 + t.Fatalf("ApplyRelayDiscoveryResponse() bootstrap error = %v", err)
838 }
780 - if resultAdded != 0 {
781 - t.Fatalf("applyRelayDiscoveryResponse() confirm added = %d, want 0", resultAdded)
839 + if err := server.relaySet.ApplyRelayDiscoveryResponse(
840 + relayADesc.Identity,
841 + relayADesc.APIHTTPSAddr,
842 + types.DiscoveryResponse{ProtocolVersion: types.ProtocolVersion, Self: relayADesc},
843 + now.Add(time.Second),
844 + ); err != nil {
845 + t.Fatalf("ApplyRelayDiscoveryResponse() direct confirm error = %v", err)
846 + }
847 +
848 + for attempt := 1; attempt <= 3; attempt++ {
849 + server.relaySet.RecordDiscoveryFailure(
850 + relayADesc.Identity,
851 + relayADesc.APIHTTPSAddr,
852 + errors.New("direct discovery failed"),
853 + )
854 }
783 - if !resultUpdated {
784 - t.Fatal("applyRelayDiscoveryResponse() confirm updated = false, want true")
855 +
856 + if err := server.relaySet.ApplyRelayDiscoveryResponse(
857 + bootstrapDesc.Identity,
858 + bootstrapDesc.APIHTTPSAddr,
859 + types.DiscoveryResponse{ProtocolVersion: types.ProtocolVersion, Self: bootstrapDesc, Relays: []types.RelayDescriptor{relayADesc}},
860 + now.Add(5*time.Second),
861 + ); err != nil {
862 + t.Fatalf("ApplyRelayDiscoveryResponse() fresh bootstrap error = %v", err)
863 }
786 - advertisedDescriptors = server.relaySet.AdvertisedDescriptors()
864 +
865 + advertisedDescriptors := server.relaySet.ActiveRelayDescriptors()
866 + advertisedURLs := make([]string, 0, len(advertisedDescriptors))
867 + for _, descriptor := range advertisedDescriptors {
868 + if strings.TrimSpace(descriptor.APIHTTPSAddr) == "" {
869 + continue
870 + }
871 + advertisedURLs = append(advertisedURLs, descriptor.APIHTTPSAddr)
872 + }
873 + sort.Strings(advertisedURLs)
874 + if !reflect.DeepEqual(advertisedURLs, []string{"https://bootstrap.example.com"}) {
875 + t.Fatalf("ActiveRelayDescriptors() = %v, want relay to stay hidden until reconfirmed", advertisedURLs)
876 + }
877 +
878 + if err := server.relaySet.ApplyRelayDiscoveryResponse(
879 + relayADesc.Identity,
880 + relayADesc.APIHTTPSAddr,
881 + types.DiscoveryResponse{ProtocolVersion: types.ProtocolVersion, Self: relayADesc},
882 + now.Add(6*time.Second),
883 + ); err != nil {
884 + t.Fatalf("ApplyRelayDiscoveryResponse() reconfirm error = %v", err)
885 + }
886 +
887 + advertisedDescriptors = server.relaySet.ActiveRelayDescriptors()
888 advertisedURLs = advertisedURLs[:0]
889 for _, descriptor := range advertisedDescriptors {
890 if strings.TrimSpace(descriptor.APIHTTPSAddr) == "" {
@@ -792,12 +893,26 @@ func TestServerRecordVerifiedDiscoveryPeerRequiresDirectConfirmation(t *testing.
893 advertisedURLs = append(advertisedURLs, descriptor.APIHTTPSAddr)
894 }
895 sort.Strings(advertisedURLs)
795 - advertisedURLs, err = utils.ExcludeLocalRelayURLs(advertisedURLs...)
896 + if !reflect.DeepEqual(advertisedURLs, []string{"https://bootstrap.example.com", "https://relay-a.example.com"}) {
897 + t.Fatalf("ActiveRelayDescriptors() = %v, want relay restored after direct confirmation", advertisedURLs)
898 + }
899 +}
900 +
901 +func TestNewServerFiltersSelfBootstrapURLFromConfig(t *testing.T) {
902 + t.Parallel()
903 +
904 + server, err := NewServer(ServerConfig{
905 + PortalURL: "https://portal.example.com",
906 + IdentityPath: tempIdentityPath(t),
907 + Bootstraps: []string{"https://bootstrap.example.com", "https://portal.example.com", "https://localhost:4017"},
908 + DiscoveryEnabled: true,
909 + })
910 if err != nil {
797 - t.Fatalf("ExcludeLocalRelayURLs() advertised second error = %v", err)
911 + t.Fatalf("NewServer() error = %v", err)
912 }
799 - if !reflect.DeepEqual(advertisedURLs, []string{"https://bootstrap.example.com", "https://relay-a.example.com"}) {
800 - t.Fatalf("AdvertisedDescriptors() = %v, want [%q %q]", advertisedURLs, "https://bootstrap.example.com", "https://relay-a.example.com")
913 +
914 + if !reflect.DeepEqual(server.cfg.Bootstraps, []string{"https://bootstrap.example.com", "https://localhost:4017"}) {
915 + t.Fatalf("cfg.Bootstraps = %v, want self bootstrap filtered and loopback kept", server.cfg.Bootstraps)
916 }
917 }
918
portal/transport/stream_client.go
+1 -13
@@ -46,14 +46,12 @@ func (s *ClientStream) RunLoop(
46 ctx context.Context,
47 open func(context.Context) (net.Conn, error),
48 currentTLSConfig func() *tls.Config,
49 - onReady func(),
50 - onInactive func(),
49 retry func(context.Context, string, error, int) bool,
50 ) {
51 var retries int
52
53 for {
56 - claimed, err := s.runSession(ctx, open, currentTLSConfig, onReady)
54 + claimed, err := s.runSession(ctx, open, currentTLSConfig)
55 switch {
56 case err == nil:
57 retries = 0
@@ -63,9 +61,6 @@ func (s *ClientStream) RunLoop(
61 retries = 0
62 default:
63 retries++
66 - if s.ActiveSessions() == 0 && onInactive != nil {
67 - onInactive()
68 - }
64 if retry == nil || !retry(ctx, "reverse session connect", err, retries) {
65 return
66 }
@@ -103,7 +98,6 @@ func (s *ClientStream) runSession(
98 ctx context.Context,
99 open func(context.Context) (net.Conn, error),
100 currentTLSConfig func() *tls.Config,
106 - onReady func(),
101 ) (bool, error) {
102 conn, err := open(ctx)
103 if err != nil {
@@ -129,18 +123,12 @@ func (s *ClientStream) runSession(
123 _ = conn.Close()
124 return true, err
125 }
132 - if onReady != nil {
133 - onReady()
134 - }
126 return true, nil
127 case types.MarkerRawStart:
128 if err := s.activateRaw(ctx, conn); err != nil {
129 _ = conn.Close()
130 return true, err
131 }
141 - if onReady != nil {
142 - onReady()
143 - }
132 return true, nil
133 default:
134 _ = conn.Close()
portal/wireguard/overlay.go deleted
-190
@@ -1,190 +0,0 @@
1 -package wireguard
2 -
3 -import (
4 - "context"
5 - "errors"
6 - "fmt"
7 - "net"
8 - "net/http"
9 - "sort"
10 - "strings"
11 - "time"
12 -
13 - "github.com/gosuda/portal/v2/types"
14 - "github.com/gosuda/portal/v2/utils"
15 -)
16 -
17 -type Config struct {
18 - PrivateKey string
19 - PublicKey string
20 - Endpoint string
21 - OverlayIPv4 string
22 - OverlayCIDRs []string
23 - ListenPort int
24 -}
25 -
26 -func NormalizeConfig(rootHost string, cfg Config) (Config, error) {
27 - configured := strings.TrimSpace(cfg.PrivateKey) != "" ||
28 - strings.TrimSpace(cfg.PublicKey) != "" ||
29 - strings.TrimSpace(cfg.Endpoint) != "" ||
30 - strings.TrimSpace(cfg.OverlayIPv4) != "" ||
31 - len(cfg.OverlayCIDRs) > 0
32 - if !configured {
33 - return cfg, nil
34 - }
35 -
36 - if strings.TrimSpace(cfg.PrivateKey) == "" {
37 - return Config{}, errors.New("wireguard private key is required when relay overlay is enabled")
38 - }
39 -
40 - privateKey, err := utils.NormalizeWireGuardPrivateKey(cfg.PrivateKey)
41 - if err != nil {
42 - return Config{}, fmt.Errorf("normalize wireguard private key: %w", err)
43 - }
44 - publicKey, err := utils.WireGuardPublicKeyFromPrivate(privateKey)
45 - if err != nil {
46 - return Config{}, fmt.Errorf("derive wireguard public key: %w", err)
47 - }
48 - if configuredPublicKey := strings.TrimSpace(cfg.PublicKey); configuredPublicKey != "" && configuredPublicKey != publicKey {
49 - return Config{}, errors.New("wireguard public key does not match private key")
50 - }
51 -
52 - cfg.PrivateKey = privateKey
53 - cfg.PublicKey = publicKey
54 - cfg.ListenPort = utils.IntOrDefault(cfg.ListenPort, DefaultListenPort)
55 - if len(cfg.OverlayCIDRs) > 0 {
56 - cfg.OverlayCIDRs, err = utils.NormalizeOverlayCIDRs(cfg.OverlayCIDRs)
57 - if err != nil {
58 - return Config{}, fmt.Errorf("normalize overlay cidrs: %w", err)
59 - }
60 - }
61 - if strings.TrimSpace(cfg.Endpoint) == "" {
62 - cfg.Endpoint = net.JoinHostPort(rootHost, fmt.Sprintf("%d", cfg.ListenPort))
63 - }
64 - if strings.TrimSpace(cfg.OverlayIPv4) == "" {
65 - cfg.OverlayIPv4, err = utils.DeriveWireGuardOverlayIPv4(cfg.PublicKey)
66 - if err != nil {
67 - return Config{}, fmt.Errorf("derive overlay ipv4: %w", err)
68 - }
69 - }
70 - if err := utils.ValidateWireGuardEndpoint(cfg.Endpoint); err != nil {
71 - return Config{}, err
72 - }
73 - if err := utils.ValidateOverlayIPv4(cfg.OverlayIPv4); err != nil {
74 - return Config{}, err
75 - }
76 - return cfg, nil
77 -}
78 -
79 -type Overlay struct {
80 - stack *stack
81 - listener net.Listener
82 - server *http.Server
83 -}
84 -
85 -func NewOverlay(cfg Config, handler http.Handler) (*Overlay, error) {
86 - stack, err := newStack(cfg)
87 - if err != nil {
88 - return nil, err
89 - }
90 -
91 - listener, err := stack.ListenTCP(DefaultPeerAPIHTTPPort)
92 - if err != nil {
93 - _ = stack.Close()
94 - return nil, err
95 - }
96 -
97 - server := &http.Server{
98 - Handler: handler,
99 - ReadHeaderTimeout: 10 * time.Second,
100 - }
101 -
102 - return &Overlay{
103 - stack: stack,
104 - listener: listener,
105 - server: server,
106 - }, nil
107 -}
108 -
109 -func (o *Overlay) Serve() error {
110 - if o == nil || o.server == nil || o.listener == nil {
111 - return nil
112 - }
113 -
114 - err := o.server.Serve(o.listener)
115 - if errors.Is(err, http.ErrServerClosed) || errors.Is(err, net.ErrClosed) {
116 - return nil
117 - }
118 - return err
119 -}
120 -
121 -func (o *Overlay) Shutdown(ctx context.Context) error {
122 - if o == nil {
123 - return nil
124 - }
125 -
126 - var shutdownErr error
127 - if o.server != nil {
128 - err := o.server.Shutdown(ctx)
129 - if err != nil && !errors.Is(err, http.ErrServerClosed) {
130 - shutdownErr = errors.Join(shutdownErr, err)
131 - }
132 - }
133 - if o.listener != nil {
134 - err := o.listener.Close()
135 - if err != nil && !errors.Is(err, net.ErrClosed) {
136 - shutdownErr = errors.Join(shutdownErr, err)
137 - }
138 - }
139 - if o.stack != nil {
140 - shutdownErr = errors.Join(shutdownErr, o.stack.Close())
141 - }
142 - return shutdownErr
143 -}
144 -
145 -func (o *Overlay) Client() *http.Client {
146 - if o == nil || o.stack == nil {
147 - return nil
148 - }
149 - return &http.Client{
150 - Transport: &http.Transport{
151 - DialContext: o.stack.DialContext,
152 - ForceAttemptHTTP2: false,
153 - },
154 - }
155 -}
156 -
157 -func (o *Overlay) Sync(selfIdentityKey string, snapshot map[string]types.RelayState) error {
158 - if o == nil || o.stack == nil {
159 - return nil
160 - }
161 - return o.stack.ApplyPeers(peersForSnapshot(selfIdentityKey, snapshot))
162 -}
163 -
164 -func peersForSnapshot(selfIdentityKey string, snapshot map[string]types.RelayState) []types.DesiredPeer {
165 - peers := make([]types.DesiredPeer, 0, len(snapshot))
166 - for _, state := range snapshot {
167 - if state.Expired {
168 - continue
169 - }
170 - desc := state.Descriptor
171 - if desc.Key() == selfIdentityKey || !desc.SupportsOverlayPeer {
172 - continue
173 - }
174 - if desc.WireGuardPublicKey == "" || desc.WireGuardEndpoint == "" || desc.OverlayIPv4 == "" {
175 - continue
176 - }
177 -
178 - allowedIPs := []string{desc.OverlayIPv4 + "/32"}
179 - allowedIPs = append(allowedIPs, desc.OverlayCIDRs...)
180 - peers = append(peers, types.DesiredPeer{
181 - WireGuardPublicKey: desc.WireGuardPublicKey,
182 - WireGuardEndpoint: desc.WireGuardEndpoint,
183 - AllowedIPs: allowedIPs,
184 - })
185 - }
186 - sort.Slice(peers, func(i, j int) bool {
187 - return peers[i].WireGuardPublicKey < peers[j].WireGuardPublicKey
188 - })
189 - return peers
190 -}
portal/wireguard/stack.go deleted
-248
@@ -1,248 +0,0 @@
1 -package wireguard
2 -
3 -import (
4 - "context"
5 - "errors"
6 - "fmt"
7 - "net"
8 - "net/netip"
9 - "strconv"
10 - "strings"
11 - "sync"
12 - "time"
13 -
14 - "golang.zx2c4.com/wireguard/conn"
15 - "golang.zx2c4.com/wireguard/device"
16 - "golang.zx2c4.com/wireguard/tun/netstack"
17 -
18 - "github.com/gosuda/portal/v2/types"
19 - "github.com/gosuda/portal/v2/utils"
20 -)
21 -
22 -const (
23 - DefaultMTU = 1420
24 - DefaultListenPort = 51820
25 - DefaultPeerAPIHTTPPort = 7777
26 - DefaultPersistentKeepalive = 25
27 - defaultEndpointResolveTTL = 3 * time.Second
28 -)
29 -
30 -type stack struct {
31 - device *device.Device
32 - net *netstack.Net
33 - overlayIP netip.Addr
34 -
35 - mu sync.Mutex
36 - closed bool
37 - peerEndpoints map[string]string
38 -}
39 -
40 -func newStack(cfg Config) (*stack, error) {
41 - canonicalPrivateKey, err := utils.NormalizeWireGuardPrivateKey(cfg.PrivateKey)
42 - if err != nil {
43 - return nil, fmt.Errorf("normalize wireguard private key: %w", err)
44 - }
45 -
46 - listenPort, err := utils.WireGuardListenPort(cfg.Endpoint)
47 - if err != nil {
48 - return nil, err
49 - }
50 -
51 - overlayIP, err := netip.ParseAddr(cfg.OverlayIPv4)
52 - if err != nil || !overlayIP.Is4() {
53 - return nil, errors.New("overlay ipv4 must be a valid IPv4 address")
54 - }
55 -
56 - tunDevice, network, err := netstack.CreateNetTUN([]netip.Addr{overlayIP}, nil, DefaultMTU)
57 - if err != nil {
58 - return nil, fmt.Errorf("create netstack tun: %w", err)
59 - }
60 -
61 - wgDevice := device.NewDevice(tunDevice, conn.NewDefaultBind(), device.NewLogger(device.LogLevelError, "portal-wg"))
62 - privateKeyHex, err := utils.WireGuardKeyHex(canonicalPrivateKey)
63 - if err != nil {
64 - wgDevice.Close()
65 - <-wgDevice.Wait()
66 - return nil, err
67 - }
68 -
69 - config := fmt.Sprintf("private_key=%s\nlisten_port=%d\n", privateKeyHex, listenPort)
70 - if err := wgDevice.IpcSet(config); err != nil {
71 - wgDevice.Close()
72 - <-wgDevice.Wait()
73 - return nil, fmt.Errorf("configure wireguard device: %w", err)
74 - }
75 - if err := wgDevice.Up(); err != nil {
76 - wgDevice.Close()
77 - <-wgDevice.Wait()
78 - return nil, fmt.Errorf("bring wireguard device up: %w", err)
79 - }
80 -
81 - return &stack{
82 - device: wgDevice,
83 - net: network,
84 - overlayIP: overlayIP,
85 - peerEndpoints: map[string]string{},
86 - }, nil
87 -}
88 -
89 -func (s *stack) ListenTCP(port int) (net.Listener, error) {
90 - if s == nil || s.net == nil {
91 - return nil, errors.New("wireguard is not initialized")
92 - }
93 - return s.net.ListenTCP(&net.TCPAddr{
94 - IP: net.ParseIP(s.overlayIP.String()),
95 - Port: port,
96 - })
97 -}
98 -
99 -func (s *stack) DialContext(ctx context.Context, network, address string) (net.Conn, error) {
100 - if s == nil || s.net == nil {
101 - return nil, errors.New("wireguard is not initialized")
102 - }
103 - switch network {
104 - case "tcp", "tcp4", "tcp6":
105 - default:
106 - return nil, fmt.Errorf("unsupported network %q", network)
107 - }
108 -
109 - host, portText, err := net.SplitHostPort(address)
110 - if err != nil {
111 - return nil, err
112 - }
113 - ip, err := netip.ParseAddr(strings.Trim(host, "[]"))
114 - if err != nil {
115 - return nil, err
116 - }
117 - port, err := strconv.Atoi(portText)
118 - if err != nil || port <= 0 || port > 65535 {
119 - return nil, errors.New("invalid tcp port")
120 - }
121 - return s.net.DialContextTCPAddrPort(ctx, netip.AddrPortFrom(ip, uint16(port)))
122 -}
123 -
124 -func (s *stack) ApplyPeers(peers []types.DesiredPeer) error {
125 - if s == nil || s.device == nil {
126 - return errors.New("wireguard is not initialized")
127 - }
128 -
129 - var builder strings.Builder
130 - builder.WriteString("replace_peers=true\n")
131 - var warnErr error
132 - nextPeerEndpoints := map[string]string{}
133 -
134 - for _, peer := range peers {
135 - peerKey := strings.TrimSpace(peer.WireGuardPublicKey)
136 - publicKeyHex, err := utils.WireGuardKeyHex(peer.WireGuardPublicKey)
137 - if err != nil {
138 - return fmt.Errorf("normalize peer %q public key: %w", peerKey, err)
139 - }
140 -
141 - resolvedEndpoint := ""
142 - if endpoint := peer.WireGuardEndpoint; endpoint != "" {
143 - resolvedEndpoint, err = resolvePeerEndpoint(endpoint)
144 - if err != nil {
145 - s.mu.Lock()
146 - currentEndpoint := s.peerEndpoints[publicKeyHex]
147 - s.mu.Unlock()
148 - if currentEndpoint != "" {
149 - warnErr = errors.Join(warnErr, fmt.Errorf("resolve peer %q endpoint: %w; using current endpoint %q", peerKey, err, currentEndpoint))
150 - resolvedEndpoint = currentEndpoint
151 - } else {
152 - warnErr = errors.Join(warnErr, fmt.Errorf("resolve peer %q endpoint: %w", peerKey, err))
153 - continue
154 - }
155 - }
156 - }
157 -
158 - builder.WriteString("public_key=")
159 - builder.WriteString(publicKeyHex)
160 - builder.WriteByte('\n')
161 - if resolvedEndpoint != "" {
162 - builder.WriteString("endpoint=")
163 - builder.WriteString(resolvedEndpoint)
164 - builder.WriteByte('\n')
165 - nextPeerEndpoints[publicKeyHex] = resolvedEndpoint
166 - }
167 -
168 - allowedIPs := utils.NormalizeIPPrefixes(peer.AllowedIPs)
169 - for _, allowedIP := range allowedIPs {
170 - builder.WriteString("allowed_ip=")
171 - builder.WriteString(allowedIP)
172 - builder.WriteByte('\n')
173 - }
174 - if DefaultPersistentKeepalive > 0 {
175 - builder.WriteString("persistent_keepalive_interval=")
176 - builder.WriteString(strconv.Itoa(DefaultPersistentKeepalive))
177 - builder.WriteByte('\n')
178 - }
179 - }
180 -
181 - if err := s.device.IpcSet(builder.String()); err != nil {
182 - return err
183 - }
184 - s.mu.Lock()
185 - s.peerEndpoints = nextPeerEndpoints
186 - s.mu.Unlock()
187 - return warnErr
188 -}
189 -
190 -func resolvePeerEndpoint(raw string) (string, error) {
191 - endpoint := strings.TrimSpace(raw)
192 - if endpoint == "" {
193 - return "", errors.New("wireguard endpoint is required")
194 - }
195 -
196 - host, port, err := net.SplitHostPort(endpoint)
197 - if err != nil {
198 - return "", err
199 - }
200 -
201 - host = strings.Trim(host, "[]")
202 - if host == "" {
203 - return "", errors.New("wireguard endpoint host is required")
204 - }
205 -
206 - if ip, err := netip.ParseAddr(host); err == nil {
207 - return net.JoinHostPort(ip.String(), port), nil
208 - }
209 -
210 - ctx, cancel := context.WithTimeout(context.Background(), defaultEndpointResolveTTL)
211 - defer cancel()
212 -
213 - addrs, err := net.DefaultResolver.LookupNetIP(ctx, "ip", host)
214 - if err != nil {
215 - return "", fmt.Errorf("lookup %q: %w", host, err)
216 - }
217 - if len(addrs) == 0 {
218 - return "", fmt.Errorf("lookup %q: no IP addresses found", host)
219 - }
220 -
221 - selected := addrs[0]
222 - for _, addr := range addrs {
223 - if addr.Is4() {
224 - selected = addr
225 - break
226 - }
227 - }
228 - return net.JoinHostPort(selected.String(), port), nil
229 -}
230 -
231 -func (s *stack) Close() error {
232 - if s == nil || s.device == nil {
233 - return nil
234 - }
235 -
236 - s.mu.Lock()
237 - if s.closed {
238 - s.mu.Unlock()
239 - return nil
240 - }
241 - s.closed = true
242 - device := s.device
243 - s.mu.Unlock()
244 -
245 - device.Close()
246 - <-device.Wait()
247 - return nil
248 -}
portal/wireguard/stack_test.go deleted
-181
@@ -1,181 +0,0 @@
1 -package wireguard
2 -
3 -import (
4 - "encoding/base64"
5 - "net"
6 - "strings"
7 - "testing"
8 -
9 - "github.com/gosuda/portal/v2/types"
10 - "github.com/gosuda/portal/v2/utils"
11 -)
12 -
13 -func TestNormalizePrivateKeyAndPublicKeyFromPrivate(t *testing.T) {
14 - t.Parallel()
15 -
16 - privateKey, err := utils.NormalizeWireGuardPrivateKey("1111111111111111111111111111111111111111111111111111111111111111")
17 - if err != nil {
18 - t.Fatalf("NormalizeWireGuardPrivateKey() error = %v", err)
19 - }
20 - if _, err := base64.StdEncoding.DecodeString(privateKey); err != nil {
21 - t.Fatalf("NormalizeWireGuardPrivateKey() returned non-base64 key: %v", err)
22 - }
23 -
24 - publicKey, err := utils.WireGuardPublicKeyFromPrivate(privateKey)
25 - if err != nil {
26 - t.Fatalf("WireGuardPublicKeyFromPrivate() error = %v", err)
27 - }
28 - decoded, err := base64.StdEncoding.DecodeString(publicKey)
29 - if err != nil {
30 - t.Fatalf("WireGuardPublicKeyFromPrivate() returned non-base64 key: %v", err)
31 - }
32 - if len(decoded) != 32 {
33 - t.Fatalf("public key length = %d, want 32", len(decoded))
34 - }
35 -}
36 -
37 -func TestStackStartAndClose(t *testing.T) {
38 - t.Parallel()
39 -
40 - privateKey, err := utils.NormalizeWireGuardPrivateKey("2222222222222222222222222222222222222222222222222222222222222222")
41 - if err != nil {
42 - t.Fatalf("NormalizeWireGuardPrivateKey() error = %v", err)
43 - }
44 -
45 - port := reserveUDPPort(t)
46 - stack, err := newStack(Config{
47 - PrivateKey: privateKey,
48 - Endpoint: net.JoinHostPort("127.0.0.1", port),
49 - OverlayIPv4: "10.77.0.1",
50 - })
51 - if err != nil {
52 - t.Fatalf("newStack() error = %v", err)
53 - }
54 - t.Cleanup(func() {
55 - if err := stack.Close(); err != nil {
56 - t.Fatalf("Close() error = %v", err)
57 - }
58 - })
59 -}
60 -
61 -func TestResolvePeerEndpointPreservesIPLiteral(t *testing.T) {
62 - t.Parallel()
63 -
64 - got, err := resolvePeerEndpoint("127.0.0.1:51820")
65 - if err != nil {
66 - t.Fatalf("resolvePeerEndpoint() error = %v", err)
67 - }
68 - if got != "127.0.0.1:51820" {
69 - t.Fatalf("resolvePeerEndpoint() = %q, want %q", got, "127.0.0.1:51820")
70 - }
71 -}
72 -
73 -func TestResolvePeerEndpointResolvesHostname(t *testing.T) {
74 - t.Parallel()
75 -
76 - got, err := resolvePeerEndpoint("localhost:51820")
77 - if err != nil {
78 - t.Fatalf("resolvePeerEndpoint() error = %v", err)
79 - }
80 -
81 - host, port, err := net.SplitHostPort(got)
82 - if err != nil {
83 - t.Fatalf("SplitHostPort() error = %v", err)
84 - }
85 - if port != "51820" {
86 - t.Fatalf("port = %q, want %q", port, "51820")
87 - }
88 - if strings.EqualFold(host, "localhost") {
89 - t.Fatalf("host = %q, want resolved IP literal", host)
90 - }
91 -
92 - ip := net.ParseIP(host)
93 - if ip == nil {
94 - t.Fatalf("host = %q, want valid IP literal", host)
95 - }
96 - if !ip.IsLoopback() {
97 - t.Fatalf("host = %q, want loopback IP", host)
98 - }
99 -}
100 -
101 -func TestApplyPeersKeepsCurrentEndpointOnResolveFailure(t *testing.T) {
102 - t.Parallel()
103 -
104 - privateKey, err := utils.NormalizeWireGuardPrivateKey("3333333333333333333333333333333333333333333333333333333333333333")
105 - if err != nil {
106 - t.Fatalf("NormalizeWireGuardPrivateKey() error = %v", err)
107 - }
108 -
109 - port := reserveUDPPort(t)
110 - stack, err := newStack(Config{
111 - PrivateKey: privateKey,
112 - Endpoint: net.JoinHostPort("127.0.0.1", port),
113 - OverlayIPv4: "10.77.0.1",
114 - })
115 - if err != nil {
116 - t.Fatalf("newStack() error = %v", err)
117 - }
118 - t.Cleanup(func() {
119 - if err := stack.Close(); err != nil {
120 - t.Fatalf("Close() error = %v", err)
121 - }
122 - })
123 -
124 - peerPrivateKey, err := utils.NormalizeWireGuardPrivateKey("4444444444444444444444444444444444444444444444444444444444444444")
125 - if err != nil {
126 - t.Fatalf("NormalizeWireGuardPrivateKey() error = %v", err)
127 - }
128 - peerPublicKey, err := utils.WireGuardPublicKeyFromPrivate(peerPrivateKey)
129 - if err != nil {
130 - t.Fatalf("WireGuardPublicKeyFromPrivate() error = %v", err)
131 - }
132 - peerPublicKeyHex, err := utils.WireGuardKeyHex(peerPublicKey)
133 - if err != nil {
134 - t.Fatalf("WireGuardKeyHex() error = %v", err)
135 - }
136 -
137 - peer := types.DesiredPeer{
138 - WireGuardPublicKey: peerPublicKey,
139 - WireGuardEndpoint: "127.0.0.1:51820",
140 - AllowedIPs: []string{"10.77.0.2/32"},
141 - }
142 - if err := stack.ApplyPeers([]types.DesiredPeer{peer}); err != nil {
143 - t.Fatalf("ApplyPeers() initial error = %v", err)
144 - }
145 -
146 - peer.WireGuardEndpoint = "peer.invalid:51820"
147 - err = stack.ApplyPeers([]types.DesiredPeer{peer})
148 - if err == nil {
149 - t.Fatal("ApplyPeers() warning error = nil, want resolve warning")
150 - }
151 - if !strings.Contains(err.Error(), "using current endpoint") {
152 - t.Fatalf("ApplyPeers() warning = %q, want current endpoint fallback", err)
153 - }
154 -
155 - config, err := stack.device.IpcGet()
156 - if err != nil {
157 - t.Fatalf("IpcGet() error = %v", err)
158 - }
159 - if !strings.Contains(config, "public_key="+peerPublicKeyHex+"\n") {
160 - t.Fatalf("IpcGet() = %q, want peer public key %q", config, peerPublicKeyHex)
161 - }
162 - if !strings.Contains(config, "endpoint=127.0.0.1:51820\n") {
163 - t.Fatalf("IpcGet() = %q, want endpoint %q", config, "127.0.0.1:51820")
164 - }
165 -}
166 -
167 -func reserveUDPPort(t *testing.T) string {
168 - t.Helper()
169 -
170 - conn, err := net.ListenPacket("udp4", "127.0.0.1:0")
171 - if err != nil {
172 - t.Fatalf("ListenPacket() error = %v", err)
173 - }
174 - defer conn.Close()
175 -
176 - _, port, err := net.SplitHostPort(conn.LocalAddr().String())
177 - if err != nil {
178 - t.Fatalf("SplitHostPort() error = %v", err)
179 - }
180 - return port
181 -}
sdk/expose.go
+29 -78
@@ -24,15 +24,14 @@ type Exposure struct {
24 cancel context.CancelFunc
25 done <-chan struct{}
26
27 - identity types.Identity
28 - TargetAddr string
29 - UDPAddr string
30 - udpEnabled bool
31 - tcpEnabled bool
32 - banMITM bool
33 - metadata types.LeaseMetadata
34 - rootCAPEM []byte
35 - discoveryEnabled bool
27 + identity types.Identity
28 + TargetAddr string
29 + UDPAddr string
30 + udpEnabled bool
31 + tcpEnabled bool
32 + banMITM bool
33 + metadata types.LeaseMetadata
34 + rootCAPEM []byte
35
36 accepted chan net.Conn
37 datagrams chan types.DatagramFrame
@@ -96,33 +95,36 @@ func Expose(ctx context.Context, cfg ExposeConfig) (*Exposure, error) {
95
96 exposureCtx, cancel := context.WithCancel(ctx)
97 exposure := &Exposure{
99 - cancel: cancel,
100 - done: exposureCtx.Done(),
101 - identity: identity,
102 - TargetAddr: targetAddr,
103 - UDPAddr: udpAddr,
104 - udpEnabled: cfg.UDPEnabled,
105 - tcpEnabled: cfg.TCPEnabled,
106 - banMITM: cfg.BanMITM,
107 - metadata: cfg.Metadata.Copy(),
108 - rootCAPEM: append([]byte(nil), cfg.RootCAPEM...),
109 - discoveryEnabled: cfg.Discovery,
110 - accepted: make(chan net.Conn, max(len(relayURLs)*defaultReadyTarget*2, 1)),
111 - datagrams: make(chan types.DatagramFrame, max(len(relayURLs)*32, 1)),
112 - relaySet: discovery.NewRelaySet(),
113 - relayListeners: make(map[string]*Listener, len(relayURLs)),
98 + cancel: cancel,
99 + done: exposureCtx.Done(),
100 + identity: identity,
101 + TargetAddr: targetAddr,
102 + UDPAddr: udpAddr,
103 + udpEnabled: cfg.UDPEnabled,
104 + tcpEnabled: cfg.TCPEnabled,
105 + banMITM: cfg.BanMITM,
106 + metadata: cfg.Metadata.Copy(),
107 + rootCAPEM: append([]byte(nil), cfg.RootCAPEM...),
108 + accepted: make(chan net.Conn, max(len(relayURLs)*defaultReadyTarget*2, 1)),
109 + datagrams: make(chan types.DatagramFrame, max(len(relayURLs)*32, 1)),
110 + relaySet: discovery.NewRelaySet(),
111 + relayListeners: make(map[string]*Listener, len(relayURLs)),
112 }
113
114 if len(relayURLs) > 0 {
117 - exposure.relaySet.ReplaceKnownRelayURLs(relayURLs)
115 + exposure.relaySet.SetBootstrapRelayURLs(relayURLs)
116 if err := exposure.reconcileRelayListeners(true); err != nil {
117 _ = exposure.Close()
118 return nil, err
119 }
120 }
121
124 - if exposure.discoveryEnabled {
125 - go exposure.runRelayDiscoveryLoop(exposureCtx)
122 + if cfg.Discovery {
123 + go func() {
124 + _ = exposure.relaySet.RunLoop(exposureCtx, exposure.rootCAPEM, func() error {
125 + return exposure.reconcileRelayListeners(false)
126 + })
127 + }()
128 }
129 go func() {
130 <-exposure.done
@@ -246,55 +248,6 @@ func (e *Exposure) Close() error {
248 return closeErr
249 }
250
249 -func (e *Exposure) runRelayDiscoveryLoop(ctx context.Context) {
250 - for {
251 - relayURLs := append([]string(nil), e.relaySet.ActiveRelayURLs()...)
252 - if len(relayURLs) > 0 {
253 - var discoveredRelayURLs []string
254 -
255 - for _, relayURL := range relayURLs {
256 - resp, err := discovery.DiscoverRelayDiscovery(ctx, relayURL, e.rootCAPEM, nil)
257 - if err != nil {
258 - if ctx.Err() != nil {
259 - return
260 - }
261 - continue
262 - }
263 -
264 - now := time.Now().UTC()
265 - targetDescriptor, err := discovery.SeedDescriptor(relayURL)
266 - if err != nil {
267 - continue
268 - }
269 - var descriptorRelayURLs []string
270 - descriptorRelayURLs, _, _, _, err = e.relaySet.ApplyRelayDiscoveryResponse(targetDescriptor.Identity, relayURL, resp, now)
271 - if err != nil {
272 - continue
273 - }
274 -
275 - if len(discoveredRelayURLs) == 0 {
276 - discoveredRelayURLs = append([]string(nil), relayURLs...)
277 - }
278 - discoveredRelayURLs, err = utils.MergeRelayURLs(discoveredRelayURLs, nil, descriptorRelayURLs)
279 - if err != nil {
280 - continue
281 - }
282 - }
283 -
284 - if len(discoveredRelayURLs) > 0 {
285 - if e.relaySet == nil {
286 - e.relaySet = discovery.NewRelaySet()
287 - }
288 - e.relaySet.ReplaceKnownRelayURLs(discoveredRelayURLs)
289 - _ = e.reconcileRelayListeners(false)
290 - }
291 - }
292 - if !utils.SleepOrDone(ctx, types.DiscoveryPollInterval) {
293 - return
294 - }
295 - }
296 -}
297 -
251 func (e *Exposure) reconcileRelayListeners(failOnError bool) error {
252 if e.relaySet == nil {
253 e.relaySet = discovery.NewRelaySet()
@@ -354,7 +307,6 @@ func (e *Exposure) reconcileRelayListeners(failOnError bool) error {
307 if e.relayListeners == nil {
308 e.relayListeners = make(map[string]*Listener, 1)
309 }
357 - e.relaySet.MarkRelayUnreachable(relayURL)
310 if _, exists := e.relayListeners[relayURL]; exists {
311 e.listenerMu.Unlock()
312 _ = listener.Close()
@@ -417,7 +369,6 @@ func (e *Exposure) runListenerAcceptLoop(listener *Listener) {
369 if listener.closed() || errors.Is(err, net.ErrClosed) {
370 return
371 }
420 - e.relaySet.MarkRelayFailure(relayURL, time.Now().UTC())
372 log.Warn().Err(err).Str("relay_url", relayURL).Msg("exposure listener accept failed")
373 return
374 }
sdk/expose_test.go
+8 -8
@@ -48,13 +48,13 @@ func TestExposureBanRelayURLMovesRelay(t *testing.T) {
48 relaySet: discovery.NewRelaySet(),
49 relayListeners: make(map[string]*Listener, 2),
50 }
51 - exposure.relaySet.ReplaceKnownRelayURLs([]string{relayA, relayB})
51 + exposure.relaySet.SetBootstrapRelayURLs([]string{relayA, relayB})
52 exposure.relayListeners = map[string]*Listener{
53 relayA: listener,
54 relayB: {},
55 }
56
57 - exposure.relaySet.BanRelayURL(relayA, "mitm")
57 + exposure.relaySet.BanRelayURL(relayA)
58 exposure.listenerMu.Lock()
59 delete(exposure.relayListeners, relayA)
60 exposure.listenerMu.Unlock()
@@ -85,12 +85,12 @@ func TestExposureSetRelayURLsSkipsBannedRelay(t *testing.T) {
85 relaySet: discovery.NewRelaySet(),
86 relayListeners: make(map[string]*Listener, 1),
87 }
88 - exposure.relaySet.BanRelayURL(relayB, "test")
88 + exposure.relaySet.BanRelayURL(relayB)
89 exposure.relayListeners = map[string]*Listener{
90 relayA: {},
91 }
92
93 - exposure.relaySet.ReplaceKnownRelayURLs([]string{relayA, relayB})
93 + exposure.relaySet.SetBootstrapRelayURLs([]string{relayA, relayB})
94 if err := exposure.reconcileRelayListeners(false); err != nil {
95 t.Fatalf("reconcileRelayListeners() error = %v", err)
96 }
@@ -123,7 +123,7 @@ func TestExposureSetRelayURLsRemovesStaleListener(t *testing.T) {
123 relaySet: discovery.NewRelaySet(),
124 relayListeners: make(map[string]*Listener, 2),
125 }
126 - exposure.relaySet.ReplaceKnownRelayURLs([]string{relayA, relayB})
126 + exposure.relaySet.SetBootstrapRelayURLs([]string{relayA, relayB})
127 exposure.relayListeners = map[string]*Listener{
128 relayA: {
129 api: &apiClient{baseURL: relayAURL},
@@ -135,7 +135,7 @@ func TestExposureSetRelayURLsRemovesStaleListener(t *testing.T) {
135 },
136 }
137
138 - exposure.relaySet.ReplaceKnownRelayURLs([]string{relayB})
138 + exposure.relaySet.SetBootstrapRelayURLs([]string{relayB})
139 if err := exposure.reconcileRelayListeners(false); err != nil {
140 t.Fatalf("reconcileRelayListeners() error = %v", err)
141 }
@@ -166,12 +166,12 @@ func TestExposurePinDiscoveredDescriptorAllowsURLChangeForSameIdentity(t *testin
166 exposure := &Exposure{relaySet: discovery.NewRelaySet()}
167 desc := mustRelayDescriptor(t, "relay-a", "https://relay-a.example")
168
169 - if _, _, _, _, err := exposure.relaySet.ApplyRelayDiscoveryResponse(desc.Identity, desc.APIHTTPSAddr, types.DiscoveryResponse{ProtocolVersion: types.ProtocolVersion, Self: desc}, time.Now().UTC()); err != nil {
169 + if err := exposure.relaySet.ApplyRelayDiscoveryResponse(desc.Identity, desc.APIHTTPSAddr, types.DiscoveryResponse{ProtocolVersion: types.ProtocolVersion, Self: desc}, time.Now().UTC()); err != nil {
170 t.Fatalf("ApplyRelayDiscoveryResponse() error = %v", err)
171 }
172
173 changedURL := mustRelayDescriptor(t, desc.Name, "https://relay-b.example")
174 - _, _, _, _, err := exposure.relaySet.ApplyRelayDiscoveryResponse(desc.Identity, "", types.DiscoveryResponse{ProtocolVersion: types.ProtocolVersion, Self: changedURL}, time.Now().UTC())
174 + err := exposure.relaySet.ApplyRelayDiscoveryResponse(desc.Identity, "", types.DiscoveryResponse{ProtocolVersion: types.ProtocolVersion, Self: changedURL}, time.Now().UTC())
175 if err != nil {
176 t.Fatalf("ApplyRelayDiscoveryResponse() error = %v, want nil for same relay identity", err)
177 }
sdk/listener.go
+1 -21
@@ -137,8 +137,6 @@ func (l *Listener) runStartup(ctx context.Context, readyTarget int) {
137 defer l.mu.Unlock()
138 return l.tlsConfig
139 },
140 - func() { l.markReachable() },
141 - func() { l.markUnreachable() },
140 l.retryOrClose,
141 )
142 }
@@ -509,10 +507,6 @@ func (l *Listener) retryOrClose(ctx context.Context, operation string, err error
507 Str("address", l.Address()).
508 Logger()
509
512 - if operation == "lease registration" {
513 - l.markUnreachable()
514 - }
515 -
510 if l.retryCount > 0 && retries > l.retryCount {
511 if operation != "lease renewal" {
512 logger.Error().
@@ -558,25 +552,11 @@ func (l *Listener) ban() {
552 return
553 }
554 if l.relaySet != nil && l.api != nil && l.api.baseURL != nil {
561 - l.relaySet.BanRelayURL(l.api.baseURL.String(), "mitm")
555 + l.relaySet.BanRelayURL(l.api.baseURL.String())
556 }
557 _ = l.Close()
558 }
559
566 -func (l *Listener) markReachable() {
567 - if l == nil || l.relaySet == nil || l.api == nil || l.api.baseURL == nil {
568 - return
569 - }
570 - l.relaySet.MarkRelayReachable(l.api.baseURL.String(), time.Now().UTC())
571 -}
572 -
573 -func (l *Listener) markUnreachable() {
574 - if l == nil || l.relaySet == nil || l.api == nil || l.api.baseURL == nil {
575 - return
576 - }
577 - l.relaySet.MarkRelayUnreachable(l.api.baseURL.String())
578 -}
579 -
560 func (l *Listener) BanMITM() bool {
561 if l == nil {
562 return false
sdk/mitm_test.go
+2 -4
@@ -221,9 +221,8 @@ func TestMITMProbeDetectionBansListener(t *testing.T) {
221 registered: make(chan struct{}),
222 banMITM: true,
223 }
224 - listener.relaySet.ReplaceKnownRelayURLs([]string{relayURL.String()})
224 + listener.relaySet.SetBootstrapRelayURLs([]string{relayURL.String()})
225 listener.mitmManager = newMITMManager(context.Background(), listener)
226 - listener.markReachable()
226
227 listener.mitmManager.logResult(MITMProbeReport{
228 RelayURL: relayURL.String(),
@@ -255,9 +254,8 @@ func TestMITMProbeDetectionWarnsWithoutBanningListener(t *testing.T) {
254 registered: make(chan struct{}),
255 banMITM: false,
256 }
258 - listener.relaySet.ReplaceKnownRelayURLs([]string{relayURL.String()})
257 + listener.relaySet.SetBootstrapRelayURLs([]string{relayURL.String()})
258 listener.mitmManager = newMITMManager(context.Background(), listener)
260 - listener.markReachable()
259
260 listener.mitmManager.logResult(MITMProbeReport{
261 RelayURL: relayURL.String(),
types/identity.go
+2 -24
@@ -87,34 +87,12 @@ type RelayDescriptor struct {
87 APIHTTPSAddr string `json:"api_https_addr"`
88 IngressTLSAddr string `json:"ingress_tls_addr,omitempty"`
89
90 - WireGuardPublicKey string `json:"wireguard_public_key,omitempty"`
91 - WireGuardEndpoint string `json:"wireguard_endpoint,omitempty"`
92 - OverlayIPv4 string `json:"overlay_ipv4,omitempty"`
93 - OverlayCIDRs []string `json:"overlay_cidrs,omitempty"`
94 -
95 - SupportsUDP bool `json:"supports_udp,omitempty"`
96 - SupportsTCP bool `json:"supports_tcp,omitempty"`
97 - SupportsOverlayPeer bool `json:"supports_overlay_peer,omitempty"`
90 + SupportsUDP bool `json:"supports_udp,omitempty"`
91 + SupportsTCP bool `json:"supports_tcp,omitempty"`
92 }
93
94 const DiscoveryPollInterval = 1 * time.Minute
95
102 -type RelayState struct {
103 - Descriptor RelayDescriptor `json:"descriptor"`
104 - Bootstrap bool `json:"bootstrap,omitempty"`
105 - Advertised bool `json:"advertised,omitempty"`
106 - Expired bool `json:"expired,omitempty"`
107 - FirstSeenAt time.Time `json:"first_seen_at"`
108 - LastSeenAt time.Time `json:"last_seen_at"`
109 - ConsecutiveFailures int `json:"consecutive_failures,omitempty"`
110 -}
111 -
112 -type DesiredPeer struct {
113 - WireGuardPublicKey string `json:"wireguard_public_key"`
114 - WireGuardEndpoint string `json:"wireguard_endpoint"`
115 - AllowedIPs []string `json:"allowed_ips,omitempty"`
116 -}
117 -
96 type DNSSECStatus struct {
97 State string `json:"state,omitempty"`
98 DSRecord string `json:"ds_record,omitempty"`
utils/crypto.go
+4 -183
@@ -1,22 +1,15 @@
1 package utils
2
3 import (
4 - "crypto/rand"
4 "crypto/sha256"
6 - "encoding/base64"
5 "encoding/hex"
6 "errors"
7 "fmt"
10 - "net"
11 - "net/netip"
12 - "sort"
13 - "strconv"
8 "strings"
9
10 "github.com/decred/dcrd/dcrec/secp256k1/v4"
17 - secp256k1ecdsa "github.com/decred/dcrd/dcrec/secp256k1/v4/ecdsa"
11 + "github.com/decred/dcrd/dcrec/secp256k1/v4/ecdsa"
12 "github.com/gosuda/portal/v2/types"
19 - "golang.org/x/crypto/curve25519"
13 "golang.org/x/crypto/sha3"
14 )
15
@@ -100,7 +93,7 @@ func SignEthereumPersonalMessage(message, privateKeyHex string) (string, error)
93 _, _ = hasher.Write(data)
94 hash := hasher.Sum(nil)
95
103 - compactSignature := secp256k1ecdsa.SignCompact(privateKey, hash, false)
96 + compactSignature := ecdsa.SignCompact(privateKey, hash, false)
97 if len(compactSignature) != 65 {
98 return "", errors.New("invalid compact signature length")
99 }
@@ -147,7 +140,7 @@ func SignSHA256Secp256k1DER(payload []byte, privateKeyHex string) (string, error
140 }
141
142 hash := sha256.Sum256(payload)
150 - signature := secp256k1ecdsa.Sign(privateKey, hash[:])
143 + signature := ecdsa.Sign(privateKey, hash[:])
144 return hex.EncodeToString(signature.Serialize()), nil
145 }
146
@@ -167,7 +160,7 @@ func VerifySHA256Secp256k1DER(payload []byte, publicKeyHex, signatureHex string)
160 if err != nil {
161 return errors.New("signature must be hex encoded")
162 }
170 - signature, err := secp256k1ecdsa.ParseDERSignature(sigBytes)
163 + signature, err := ecdsa.ParseDERSignature(sigBytes)
164 if err != nil {
165 return fmt.Errorf("parse signature: %w", err)
166 }
@@ -236,175 +229,3 @@ func ParseSecp256k1PrivateKeyHex(raw string, requireNonZero bool) (*secp256k1.Pr
229 }
230 return key, privateKeyHex, nil
231 }
239 -
240 -func NormalizeWireGuardPrivateKey(raw string) (string, error) {
241 - key, err := decodeWireGuardKey(raw)
242 - if err != nil {
243 - return "", err
244 - }
245 - clampWireGuardPrivateKey(&key)
246 - return base64.StdEncoding.EncodeToString(key[:]), nil
247 -}
248 -
249 -func GenerateWireGuardPrivateKey() (string, error) {
250 - var key [32]byte
251 - if _, err := rand.Read(key[:]); err != nil {
252 - return "", fmt.Errorf("generate wireguard private key: %w", err)
253 - }
254 - clampWireGuardPrivateKey(&key)
255 - return base64.StdEncoding.EncodeToString(key[:]), nil
256 -}
257 -
258 -func WireGuardPublicKeyFromPrivate(raw string) (string, error) {
259 - privateKey, err := decodeWireGuardKey(raw)
260 - if err != nil {
261 - return "", err
262 - }
263 - clampWireGuardPrivateKey(&privateKey)
264 - var publicKey [32]byte
265 - curve25519.ScalarBaseMult(&publicKey, &privateKey)
266 - return base64.StdEncoding.EncodeToString(publicKey[:]), nil
267 -}
268 -
269 -func WireGuardKeyHex(raw string) (string, error) {
270 - key, err := decodeWireGuardKey(raw)
271 - if err != nil {
272 - return "", err
273 - }
274 - return hex.EncodeToString(key[:]), nil
275 -}
276 -
277 -func ValidateWireGuardPublicKey(raw string) error {
278 - key := strings.TrimSpace(raw)
279 - if key == "" {
280 - return errors.New("wireguard_public_key is required")
281 - }
282 - decoded, err := base64.StdEncoding.DecodeString(key)
283 - if err != nil {
284 - return errors.New("wireguard_public_key must be base64 encoded")
285 - }
286 - if len(decoded) != 32 {
287 - return errors.New("wireguard_public_key must be 32 bytes")
288 - }
289 - return nil
290 -}
291 -
292 -func ValidateWireGuardEndpoint(raw string) error {
293 - endpoint := strings.TrimSpace(raw)
294 - if endpoint == "" {
295 - return errors.New("wireguard_endpoint is required")
296 - }
297 - host, port, err := net.SplitHostPort(endpoint)
298 - if err != nil {
299 - return errors.New("wireguard_endpoint must be host:port")
300 - }
301 - if strings.TrimSpace(host) == "" {
302 - return errors.New("wireguard_endpoint host is required")
303 - }
304 - portNum, err := strconv.Atoi(port)
305 - if err != nil || portNum <= 0 || portNum > 65535 {
306 - return errors.New("wireguard_endpoint port is invalid")
307 - }
308 - return nil
309 -}
310 -
311 -func WireGuardListenPort(rawEndpoint string) (int, error) {
312 - endpoint := strings.TrimSpace(rawEndpoint)
313 - if endpoint == "" {
314 - return 0, errors.New("wireguard endpoint is required")
315 - }
316 - _, portText, err := net.SplitHostPort(endpoint)
317 - if err != nil {
318 - return 0, errors.New("wireguard endpoint must be host:port")
319 - }
320 - port, err := strconv.Atoi(portText)
321 - if err != nil || port <= 0 || port > 65535 {
322 - return 0, errors.New("wireguard endpoint port is invalid")
323 - }
324 - return port, nil
325 -}
326 -
327 -func ValidateOverlayIPv4(raw string) error {
328 - ipText := strings.TrimSpace(raw)
329 - if ipText == "" {
330 - return errors.New("overlay_ipv4 is required")
331 - }
332 - ip := net.ParseIP(ipText)
333 - if ip == nil || ip.To4() == nil {
334 - return errors.New("overlay_ipv4 must be a valid IPv4 address")
335 - }
336 - return nil
337 -}
338 -
339 -func NormalizeOverlayCIDRs(inputs []string) ([]string, error) {
340 - if len(inputs) == 0 {
341 - return nil, nil
342 - }
343 - seen := make(map[string]struct{}, len(inputs))
344 - out := make([]string, 0, len(inputs))
345 - for _, input := range inputs {
346 - input = strings.TrimSpace(input)
347 - if input == "" {
348 - continue
349 - }
350 - _, network, err := net.ParseCIDR(input)
351 - if err != nil {
352 - return nil, fmt.Errorf("invalid overlay cidr %q", input)
353 - }
354 - normalized := network.String()
355 - if _, ok := seen[normalized]; ok {
356 - continue
357 - }
358 - seen[normalized] = struct{}{}
359 - out = append(out, normalized)
360 - }
361 - sort.Strings(out)
362 - return out, nil
363 -}
364 -
365 -func DeriveWireGuardOverlayIPv4(publicKey string) (string, error) {
366 - decoded, err := base64.StdEncoding.DecodeString(strings.TrimSpace(publicKey))
367 - if err != nil {
368 - return "", errors.New("wireguard public key must be base64 encoded")
369 - }
370 - if len(decoded) != 32 {
371 - return "", errors.New("wireguard public key must be 32 bytes")
372 - }
373 -
374 - sum := sha256.Sum256(decoded)
375 - return netip.AddrFrom4([4]byte{
376 - 100,
377 - 64 + (sum[0] & 0x3f),
378 - sum[1],
379 - 1 + (sum[2] % 254),
380 - }).String(), nil
381 -}
382 -
383 -func decodeWireGuardKey(raw string) ([32]byte, error) {
384 - var key [32]byte
385 - value := strings.TrimSpace(raw)
386 - if value == "" {
387 - return key, errors.New("wireguard key is required")
388 - }
389 -
390 - var decoded []byte
391 - var err error
392 - if len(value) == 64 && !strings.Contains(value, "=") {
393 - decoded, err = hex.DecodeString(value)
394 - } else {
395 - decoded, err = base64.StdEncoding.DecodeString(value)
396 - }
397 - if err != nil {
398 - return key, errors.New("wireguard key must be base64 or hex encoded")
399 - }
400 - if len(decoded) != len(key) {
401 - return key, errors.New("wireguard key must be 32 bytes")
402 - }
403 - copy(key[:], decoded)
404 - return key, nil
405 -}
406 -
407 -func clampWireGuardPrivateKey(key *[32]byte) {
408 - key[0] &= 248
409 - key[31] = (key[31] & 127) | 64
410 -}