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