feat: add max active relays configuration and update relay selection logic

rabbitprincess committed Apr 12, 2026 at 21:58 UTC 86ddb2ff3c4402be4cb807afff32e154d29e2354
10 files changed +329 -173
cmd/demo-app/main.go
+29 -25
@@ -34,18 +34,19 @@ func main() {
34 }
35
36 type demoConfig struct {
37 - relayURLs string
38 - discovery bool
39 - banMITM bool
40 - identityPath string
41 - identityJSON string
42 - addr string
43 - name string
44 - desc string
45 - tags string
46 - owner string
47 - hide bool
48 - thumbnail string
37 + relayURLs string
38 + discovery bool
39 + banMITM bool
40 + identityPath string
41 + identityJSON string
42 + addr string
43 + name string
44 + desc string
45 + tags string
46 + owner string
47 + hide bool
48 + thumbnail string
49 + maxActiveRelays int
50 }
51
52 // registerConnectivityFlags registers the relay, discovery, identity, and
@@ -56,6 +57,7 @@ func registerConnectivityFlags(fs *flag.FlagSet, cfg *demoConfig, defaultRelays
57 utils.BoolFlagEnv(fs, &cfg.banMITM, "ban-mitm", false, "ban relay when the MITM self-probe detects TLS termination", "BAN_MITM")
58 utils.StringFlagEnv(fs, &cfg.identityPath, "identity-path", "identity.json", "identity json file path", "IDENTITY_PATH")
59 utils.StringFlagEnv(fs, &cfg.identityJSON, "identity-json", "", "identity json payload; overrides --identity-path contents and is persisted there when both are set", "IDENTITY_JSON")
60 + utils.IntFlagEnv(fs, &cfg.maxActiveRelays, "max-active-relays", 3, nil, "maximum number of active relays to keep connected", "MAX_ACTIVE_RELAYS")
61 utils.StringFlag(fs, &cfg.owner, "owner", "PortalApp Developer", "lease owner")
62 }
63
@@ -128,12 +130,13 @@ func runUDPCommand(args []string) error {
130
131 func runTCPDemo(ctx context.Context, cfg demoConfig) error {
132 exposure, err := sdk.Expose(ctx, sdk.ExposeConfig{
131 - RelayURLs: utils.SplitCSV(cfg.relayURLs),
132 - IdentityPath: cfg.identityPath,
133 - IdentityJSON: cfg.identityJSON,
134 - Name: cfg.name,
135 - BanMITM: cfg.banMITM,
136 - Discovery: cfg.discovery,
133 + RelayURLs: utils.SplitCSV(cfg.relayURLs),
134 + IdentityPath: cfg.identityPath,
135 + IdentityJSON: cfg.identityJSON,
136 + Name: cfg.name,
137 + BanMITM: cfg.banMITM,
138 + MaxActiveRelays: cfg.maxActiveRelays,
139 + Discovery: cfg.discovery,
140 Metadata: types.LeaseMetadata{
141 Description: cfg.desc,
142 Tags: utils.SplitCSV(cfg.tags),
@@ -170,13 +173,14 @@ func runTCPDemo(ctx context.Context, cfg demoConfig) error {
173
174 func runUDPDemo(ctx context.Context, cfg demoConfig) error {
175 exposure, err := sdk.Expose(ctx, sdk.ExposeConfig{
173 - RelayURLs: utils.SplitCSV(cfg.relayURLs),
174 - IdentityPath: cfg.identityPath,
175 - IdentityJSON: cfg.identityJSON,
176 - Name: cfg.name,
177 - UDPEnabled: true,
178 - BanMITM: cfg.banMITM,
179 - Discovery: cfg.discovery,
176 + RelayURLs: utils.SplitCSV(cfg.relayURLs),
177 + IdentityPath: cfg.identityPath,
178 + IdentityJSON: cfg.identityJSON,
179 + Name: cfg.name,
180 + UDPEnabled: true,
181 + BanMITM: cfg.banMITM,
182 + MaxActiveRelays: cfg.maxActiveRelays,
183 + Discovery: cfg.discovery,
184 Metadata: types.LeaseMetadata{
185 Description: cfg.desc,
186 Tags: utils.SplitCSV(cfg.tags),
cmd/portal-tunnel/main.go
+29 -26
@@ -33,22 +33,23 @@ func main() {
33 }
34
35 type exposeFlags struct {
36 - relayCSV string
37 - discovery bool
38 - banMITM bool
39 - identityPath string
40 - identityJSON string
41 - name string
42 - desc string
43 - tags string
44 - owner string
45 - thumbnail string
46 - hide bool
47 - targetAddr string
48 - httpRoutes []string
49 - udp bool
50 - udpAddr string
51 - tcp bool
36 + relayCSV string
37 + discovery bool
38 + banMITM bool
39 + identityPath string
40 + identityJSON string
41 + name string
42 + desc string
43 + tags string
44 + owner string
45 + thumbnail string
46 + hide bool
47 + targetAddr string
48 + httpRoutes []string
49 + udp bool
50 + udpAddr string
51 + tcp bool
52 + maxActiveRelays int
53 }
54
55 func runExposeCommand(args []string) error {
@@ -70,6 +71,7 @@ func runExposeCommand(args []string) error {
71 utils.BoolFlagEnv(fs, &flags.udp, "udp", false, "Enable public UDP relay in addition to the default TCP relay", "UDP_ENABLED")
72 utils.StringFlagEnv(fs, &flags.udpAddr, "udp-addr", "", "Local UDP target address for relayed datagrams (host:port or port only); defaults to the target when --udp is enabled", "UDP_ADDR")
73 utils.BoolFlagEnv(fs, &flags.tcp, "tcp", false, "Request a dedicated TCP port on the relay for raw TCP services (no TLS; e.g., Minecraft, game servers)", "TCP_ENABLED")
74 + utils.IntFlagEnv(fs, &flags.maxActiveRelays, "max-active-relays", 3, nil, "Maximum number of active relays to keep connected", "MAX_ACTIVE_RELAYS")
75
76 if err := utils.ParseFlagSet(fs, args, printExposeUsage); err != nil {
77 if errors.Is(err, flag.ErrHelp) {
@@ -99,16 +101,17 @@ func runExposeCommand(args []string) error {
101 defer stop()
102
103 exposure, err := sdk.Expose(ctx, sdk.ExposeConfig{
102 - RelayURLs: utils.SplitCSV(flags.relayCSV),
103 - IdentityPath: flags.identityPath,
104 - IdentityJSON: flags.identityJSON,
105 - Name: flags.name,
106 - TargetAddr: flags.targetAddr,
107 - UDPAddr: flags.udpAddr,
108 - UDPEnabled: flags.udp,
109 - TCPEnabled: flags.tcp,
110 - BanMITM: flags.banMITM,
111 - Discovery: flags.discovery,
104 + RelayURLs: utils.SplitCSV(flags.relayCSV),
105 + IdentityPath: flags.identityPath,
106 + IdentityJSON: flags.identityJSON,
107 + Name: flags.name,
108 + TargetAddr: flags.targetAddr,
109 + UDPAddr: flags.udpAddr,
110 + UDPEnabled: flags.udp,
111 + TCPEnabled: flags.tcp,
112 + BanMITM: flags.banMITM,
113 + MaxActiveRelays: flags.maxActiveRelays,
114 + Discovery: flags.discovery,
115 Metadata: types.LeaseMetadata{
116 Description: flags.desc,
117 Tags: utils.SplitCSV(flags.tags),
portal/api_server.go
+1 -1
@@ -206,7 +206,7 @@ func (s *Server) handleRelayDiscovery(w http.ResponseWriter, r *http.Request) {
206 Relays: nil,
207 }
208 if s.relaySet != nil {
209 - resp.Relays = s.relaySet.AdvertisedDescriptors()
209 + resp.Relays = s.relaySet.ConfirmedDescriptors()
210 }
211 utils.WriteAPIData(w, http.StatusOK, resp)
212 }
portal/discovery/policy.go
+106 -2
@@ -3,6 +3,8 @@ package discovery
3 import (
4 "errors"
5 "net/http"
6 + "sort"
7 + "strings"
8 "time"
9
10 "github.com/gosuda/portal-tunnel/v2/types"
@@ -10,7 +12,8 @@ import (
12
13 type RelayPolicy interface {
14 SelectActive([]RelayState) []RelayState
13 - SelectAdvertised([]RelayState) []RelayState
15 + SelectConfirmed([]RelayState) []RelayState
16 + SelectPriority([]RelayState, ClientState) []RelayState
17 OnConfirmed(RelayState) RelayState
18 OnHinted(RelayState) RelayState
19 OnFailure(RelayState, error, int) (RelayState, bool, string)
@@ -43,12 +46,113 @@ func (p DefaultRelayPolicy) SelectActive(states []RelayState) []RelayState {
46 })
47 }
48
46 -func (p DefaultRelayPolicy) SelectAdvertised(states []RelayState) []RelayState {
49 +func (p DefaultRelayPolicy) SelectConfirmed(states []RelayState) []RelayState {
50 return p.selectStates(states, func(state RelayState) bool {
51 return state.Confirmed
52 })
53 }
54
55 +func (p DefaultRelayPolicy) SelectPriority(states []RelayState, clientState ClientState) []RelayState {
56 + selected := p.SelectActive(states)
57 + if len(selected) == 0 {
58 + return nil
59 + }
60 +
61 + currentRelayURLs := make([]string, 0, len(clientState.ActiveRelayURLs))
62 + activeRelayURLSet := make(map[string]struct{}, len(clientState.ActiveRelayURLs))
63 + for _, relayURL := range clientState.ActiveRelayURLs {
64 + relayURL = strings.TrimSpace(relayURL)
65 + if relayURL == "" {
66 + continue
67 + }
68 + currentRelayURLs = append(currentRelayURLs, relayURL)
69 + activeRelayURLSet[relayURL] = struct{}{}
70 + }
71 +
72 + out := selected[:0]
73 + for _, state := range selected {
74 + if clientState.RequireUDP && state.hasDescriptor() && !state.Descriptor.SupportsUDP {
75 + continue
76 + }
77 + if clientState.RequireTCP && state.hasDescriptor() && !state.Descriptor.SupportsTCP {
78 + continue
79 + }
80 + out = append(out, state)
81 + }
82 + if len(out) == 0 {
83 + return nil
84 + }
85 +
86 + currentRelays := make([]RelayState, 0, len(currentRelayURLs))
87 + eligibleByURL := make(map[string]RelayState, len(out))
88 + for _, state := range out {
89 + eligibleByURL[strings.TrimSpace(state.Descriptor.APIHTTPSAddr)] = state
90 + }
91 +
92 + currentHealthy := len(currentRelayURLs) > 0
93 + for _, relayURL := range currentRelayURLs {
94 + state, ok := eligibleByURL[relayURL]
95 + if !ok {
96 + currentHealthy = false
97 + break
98 + }
99 + currentRelays = append(currentRelays, state)
100 + }
101 + if currentHealthy && (clientState.MaxActiveRelays <= 0 || len(currentRelays) >= clientState.MaxActiveRelays) {
102 + return currentRelays
103 + }
104 +
105 + sort.SliceStable(out, func(i, j int) bool {
106 + left := out[i]
107 + right := out[j]
108 +
109 + // 1. Prefer confirmed relays over candidates that are only known through bootstrap discovery.
110 + if left.Confirmed != right.Confirmed {
111 + return left.Confirmed
112 + }
113 +
114 + // 2. Prefer relays the client is already using so a healthy pool stays stable instead of churning.
115 + _, leftActive := activeRelayURLSet[strings.TrimSpace(left.Descriptor.APIHTTPSAddr)]
116 + _, rightActive := activeRelayURLSet[strings.TrimSpace(right.Descriptor.APIHTTPSAddr)]
117 + if leftActive != rightActive {
118 + return leftActive
119 + }
120 +
121 + // 3. Prefer bootstrap relays when the stronger signals above are equal.
122 + if left.Bootstrap != right.Bootstrap {
123 + return left.Bootstrap
124 + }
125 +
126 + // 4. Prefer relays reporting lower concurrent load.
127 + if left.Descriptor.Load != right.Descriptor.Load {
128 + return left.Descriptor.Load < right.Descriptor.Load
129 + }
130 +
131 + // 5. Prefer relays reporting lower traffic score after the more stable load signal ties.
132 + if left.Descriptor.LoadScore != right.Descriptor.LoadScore {
133 + return left.Descriptor.LoadScore < right.Descriptor.LoadScore
134 + }
135 +
136 + // 6. Prefer relays with measured discovery RTT, then prefer the lower RTT when both have measurements.
137 + leftHasRTT := !left.DiscoveryRTTAt.IsZero()
138 + rightHasRTT := !right.DiscoveryRTTAt.IsZero()
139 + if leftHasRTT != rightHasRTT {
140 + return leftHasRTT
141 + }
142 + if left.DiscoveryRTT != right.DiscoveryRTT {
143 + return left.DiscoveryRTT < right.DiscoveryRTT
144 + }
145 +
146 + // 7. Fall back to URL ordering so selection remains deterministic when all policy signals tie.
147 + return left.Descriptor.APIHTTPSAddr < right.Descriptor.APIHTTPSAddr
148 + })
149 +
150 + if clientState.MaxActiveRelays > 0 && len(out) > clientState.MaxActiveRelays {
151 + out = out[:clientState.MaxActiveRelays]
152 + }
153 + return out
154 +}
155 +
156 func (DefaultRelayPolicy) OnConfirmed(state RelayState) RelayState {
157 if state.Banned {
158 return state
portal/discovery/refresher.go
+8 -2
@@ -97,6 +97,7 @@ func (r *Refresher) refreshHTTPS(ctx context.Context) error {
97 continue
98 }
99
100 + startedAt := time.Now()
101 var resp types.DiscoveryResponse
102 if err := utils.HTTPDoAPIPath(ctx, r.httpClient, baseURL, http.MethodGet, types.PathDiscovery, nil, nil, &resp); err != nil {
103 if ctx.Err() != nil {
@@ -105,9 +106,11 @@ func (r *Refresher) refreshHTTPS(ctx context.Context) error {
106 continue
107 }
108
108 - if _, err := r.relaySet.ApplyRelayDiscoveryResponse(relay.Identity, relay.APIHTTPSAddr, resp, time.Now().UTC()); err != nil {
109 + measuredAt := time.Now().UTC()
110 + if _, err := r.relaySet.ApplyRelayDiscoveryResponse(relay.Identity, relay.APIHTTPSAddr, resp, measuredAt); err != nil {
111 continue
112 }
113 + r.relaySet.RecordDiscoveryRTT(relay.APIHTTPSAddr, time.Since(startedAt), measuredAt)
114 }
115 if err := ctx.Err(); err != nil {
116 return err
@@ -135,6 +138,7 @@ func (r *Refresher) refreshHTTPS(ctx context.Context) error {
138 continue
139 }
140
141 + startedAt := time.Now()
142 var resp types.DiscoveryResponse
143 if err := utils.HTTPDoAPIPath(ctx, r.httpClient, baseURL, http.MethodGet, types.PathDiscovery, nil, nil, &resp); err != nil {
144 if ctx.Err() != nil {
@@ -144,10 +148,12 @@ func (r *Refresher) refreshHTTPS(ctx context.Context) error {
148 continue
149 }
150
147 - if _, err := r.relaySet.ApplyRelayDiscoveryResponse(relay.Identity, relay.APIHTTPSAddr, resp, time.Now().UTC()); err != nil {
151 + measuredAt := time.Now().UTC()
152 + if _, err := r.relaySet.ApplyRelayDiscoveryResponse(relay.Identity, relay.APIHTTPSAddr, resp, measuredAt); err != nil {
153 r.logDirectDiscoveryFailure(relay, err, r.directRecoveryFailures)
154 continue
155 }
156 + r.relaySet.RecordDiscoveryRTT(relay.APIHTTPSAddr, time.Since(startedAt), measuredAt)
157 }
158 return nil
159 }
portal/discovery/relayset.go
+26 -6
@@ -15,7 +15,7 @@ import (
15
16 // RelaySet owns the shared relay discovery view: configured bootstrap relay URLs,
17 // the latest validated descriptor seen for each relay, and local runtime state
18 -// such as ban/reachability/failure tracking.
18 +// such as ban/reachability/failure tracking and observed discovery RTT.
19 type RelaySet struct {
20 mu sync.RWMutex
21 relays map[string]RelayState
@@ -102,6 +102,13 @@ func (s *RelaySet) ActiveRelays() []RelayState {
102 return s.policy.SelectActive(s.relayStatesLocked())
103 }
104
105 +func (s *RelaySet) PriorityRelays(clientState ClientState) []RelayState {
106 + s.mu.RLock()
107 + defer s.mu.RUnlock()
108 +
109 + return s.policy.SelectPriority(s.relayStatesLocked(), clientState)
110 +}
111 +
112 func (s *RelaySet) OverlayPeerStates() []RelayState {
113 s.mu.RLock()
114 states := s.relayStatesLocked()
@@ -126,11 +133,11 @@ func (s *RelaySet) OverlayPeerStates() []RelayState {
133 return out
134 }
135
129 -func (s *RelaySet) AdvertisedDescriptors() []types.RelayDescriptor {
136 +func (s *RelaySet) ConfirmedDescriptors() []types.RelayDescriptor {
137 s.mu.RLock()
138 defer s.mu.RUnlock()
139
133 - states := s.policy.SelectAdvertised(s.relayStatesLocked())
140 + states := s.policy.SelectConfirmed(s.relayStatesLocked())
141 out := make([]types.RelayDescriptor, 0, len(states))
142 for _, state := range states {
143 out = append(out, state.Descriptor)
@@ -138,9 +145,6 @@ func (s *RelaySet) AdvertisedDescriptors() []types.RelayDescriptor {
145 if len(out) == 0 {
146 return nil
147 }
141 - sort.Slice(out, func(i, j int) bool {
142 - return out[i].APIHTTPSAddr < out[j].APIHTTPSAddr
143 - })
148 return out
149 }
150
@@ -197,6 +201,8 @@ func (s *RelaySet) applyDiscoveredStateLocked(state RelayState, confirmed bool)
201 record.Reachable = existing.Reachable
202 record.Confirmed = existing.Confirmed
203 record.Banned = existing.Banned
204 + record.DiscoveryRTT = existing.DiscoveryRTT
205 + record.DiscoveryRTTAt = existing.DiscoveryRTTAt
206 record.consecutiveFailures = existing.consecutiveFailures
207 }
208 break
@@ -296,6 +302,20 @@ func (s *RelaySet) ApplyRelayDiscoveryResponse(targetIdentity types.Identity, ta
302 return relaySetChanged, nil
303 }
304
305 +func (s *RelaySet) RecordDiscoveryRTT(relayURL string, rtt time.Duration, measuredAt time.Time) {
306 + s.mu.Lock()
307 + defer s.mu.Unlock()
308 +
309 + state, ok := s.relays[relayURL]
310 + if !ok {
311 + return
312 + }
313 +
314 + state.DiscoveryRTT = rtt
315 + state.DiscoveryRTTAt = measuredAt
316 + s.relays[relayURL] = state
317 +}
318 +
319 func (s *RelaySet) RecordDiscoveryFailure(identity types.Identity, relayURL string, err error, recoveryFailures int) (expired bool, expireReason string, consecutiveFailures int) {
320 relayKey := identity.Key()
321 if relayKey == "" {
portal/discovery/relaystate.go
+15 -6
@@ -10,16 +10,25 @@ import (
10 )
11
12 type RelayState struct {
13 - Descriptor types.RelayDescriptor
14 - Bootstrap bool
15 - Reachable bool
16 - Confirmed bool
17 - Banned bool
18 - LastSeenAt time.Time
13 + Descriptor types.RelayDescriptor
14 + Bootstrap bool
15 + Reachable bool
16 + Confirmed bool
17 + Banned bool
18 + LastSeenAt time.Time
19 + DiscoveryRTT time.Duration
20 + DiscoveryRTTAt time.Time
21
22 consecutiveFailures int
23 }
24
25 +type ClientState struct {
26 + ActiveRelayURLs []string
27 + MaxActiveRelays int
28 + RequireUDP bool
29 + RequireTCP bool
30 +}
31 +
32 func newRelayState(desc types.RelayDescriptor, seenAt time.Time) (RelayState, error) {
33 state := RelayState{
34 Descriptor: desc,
portal/server_test.go
+15 -15
@@ -506,7 +506,7 @@ func TestServerSetBootstrapRelayURLsAllowsLoopbackButSkipsSelfRelay(t *testing.T
506 }); err != nil {
507 t.Fatalf("SetBootstrapRelayURLs() error = %v", err)
508 }
509 - advertisedDescriptors := server.relaySet.AdvertisedDescriptors()
509 + advertisedDescriptors := server.relaySet.ConfirmedDescriptors()
510 knownURLs := make([]string, 0)
511 for _, state := range server.relaySet.ActiveRelays() {
512 knownURLs = append(knownURLs, state.Descriptor.APIHTTPSAddr)
@@ -618,7 +618,7 @@ func TestServerRecordVerifiedDiscoveryPeerRequiresDirectConfirmation(t *testing.
618 if !reflect.DeepEqual(knownURLs, []string{"https://bootstrap.example.com"}) {
619 t.Fatalf("ActiveRelayURLs() = %v, want [%q]", knownURLs, "https://bootstrap.example.com")
620 }
621 - advertisedDescriptors := server.relaySet.AdvertisedDescriptors()
621 + advertisedDescriptors := server.relaySet.ConfirmedDescriptors()
622 advertisedURLs := make([]string, 0, len(advertisedDescriptors))
623 for _, descriptor := range advertisedDescriptors {
624 if strings.TrimSpace(descriptor.APIHTTPSAddr) == "" {
@@ -628,7 +628,7 @@ func TestServerRecordVerifiedDiscoveryPeerRequiresDirectConfirmation(t *testing.
628 }
629 sort.Strings(advertisedURLs)
630 if !reflect.DeepEqual(advertisedURLs, []string{"https://bootstrap.example.com"}) {
631 - t.Fatalf("AdvertisedDescriptors() = %v, want [%q]", advertisedURLs, "https://bootstrap.example.com")
631 + t.Fatalf("ConfirmedDescriptors() = %v, want [%q]", advertisedURLs, "https://bootstrap.example.com")
632 }
633
634 err = applyDiscovery(
@@ -639,7 +639,7 @@ func TestServerRecordVerifiedDiscoveryPeerRequiresDirectConfirmation(t *testing.
639 if err != nil {
640 t.Fatalf("applyRelayDiscoveryResponse() confirm error = %v", err)
641 }
642 - advertisedDescriptors = server.relaySet.AdvertisedDescriptors()
642 + advertisedDescriptors = server.relaySet.ConfirmedDescriptors()
643 advertisedURLs = advertisedURLs[:0]
644 for _, descriptor := range advertisedDescriptors {
645 if strings.TrimSpace(descriptor.APIHTTPSAddr) == "" {
@@ -649,7 +649,7 @@ func TestServerRecordVerifiedDiscoveryPeerRequiresDirectConfirmation(t *testing.
649 }
650 sort.Strings(advertisedURLs)
651 if !reflect.DeepEqual(advertisedURLs, []string{"https://bootstrap.example.com", "https://relay-a.example.com"}) {
652 - t.Fatalf("AdvertisedDescriptors() = %v, want [%q %q]", advertisedURLs, "https://bootstrap.example.com", "https://relay-a.example.com")
652 + t.Fatalf("ConfirmedDescriptors() = %v, want [%q %q]", advertisedURLs, "https://bootstrap.example.com", "https://relay-a.example.com")
653 }
654 }
655
@@ -693,7 +693,7 @@ func TestServerBannedDiscoveryPeerIsNotAdvertised(t *testing.T) {
693
694 server.relaySet.BanRelayURL(relayADesc.APIHTTPSAddr)
695
696 - advertisedDescriptors := server.relaySet.AdvertisedDescriptors()
696 + advertisedDescriptors := server.relaySet.ConfirmedDescriptors()
697 advertisedURLs := make([]string, 0, len(advertisedDescriptors))
698 for _, descriptor := range advertisedDescriptors {
699 if strings.TrimSpace(descriptor.APIHTTPSAddr) == "" {
@@ -703,7 +703,7 @@ func TestServerBannedDiscoveryPeerIsNotAdvertised(t *testing.T) {
703 }
704 sort.Strings(advertisedURLs)
705 if !reflect.DeepEqual(advertisedURLs, []string{"https://bootstrap.example.com"}) {
706 - t.Fatalf("AdvertisedDescriptors() = %v, want banned relay excluded", advertisedURLs)
706 + t.Fatalf("ConfirmedDescriptors() = %v, want banned relay excluded", advertisedURLs)
707 }
708 }
709
@@ -763,7 +763,7 @@ func TestServerRecordVerifiedDiscoveryPeerExpiresAfterRepeatedDirectFailures(t *
763 }
764 }
765
766 - advertisedDescriptors := server.relaySet.AdvertisedDescriptors()
766 + advertisedDescriptors := server.relaySet.ConfirmedDescriptors()
767 advertisedURLs := make([]string, 0, len(advertisedDescriptors))
768 for _, descriptor := range advertisedDescriptors {
769 if strings.TrimSpace(descriptor.APIHTTPSAddr) == "" {
@@ -773,7 +773,7 @@ func TestServerRecordVerifiedDiscoveryPeerExpiresAfterRepeatedDirectFailures(t *
773 }
774 sort.Strings(advertisedURLs)
775 if !reflect.DeepEqual(advertisedURLs, []string{"https://bootstrap.example.com"}) {
776 - t.Fatalf("AdvertisedDescriptors() = %v, want [%q] after relay expiry", advertisedURLs, "https://bootstrap.example.com")
776 + t.Fatalf("ConfirmedDescriptors() = %v, want [%q] after relay expiry", advertisedURLs, "https://bootstrap.example.com")
777 }
778 }
779
@@ -840,7 +840,7 @@ func TestServerBootstrapHintDoesNotResetDirectFailureBudget(t *testing.T) {
840 t.Fatalf("ApplyRelayDiscoveryResponse() hinted refresh error = %v", err)
841 }
842
843 - advertisedDescriptors := server.relaySet.AdvertisedDescriptors()
843 + advertisedDescriptors := server.relaySet.ConfirmedDescriptors()
844 advertisedURLs := make([]string, 0, len(advertisedDescriptors))
845 for _, descriptor := range advertisedDescriptors {
846 if strings.TrimSpace(descriptor.APIHTTPSAddr) == "" {
@@ -850,7 +850,7 @@ func TestServerBootstrapHintDoesNotResetDirectFailureBudget(t *testing.T) {
850 }
851 sort.Strings(advertisedURLs)
852 if !reflect.DeepEqual(advertisedURLs, []string{"https://bootstrap.example.com", "https://relay-a.example.com"}) {
853 - t.Fatalf("AdvertisedDescriptors() = %v, want relay to remain advertised before expiry", advertisedURLs)
853 + t.Fatalf("ConfirmedDescriptors() = %v, want relay to remain advertised before expiry", advertisedURLs)
854 }
855 }
856
@@ -926,7 +926,7 @@ func TestServerExpiredDiscoveryPeerNeedsFreshDirectConfirmation(t *testing.T) {
926 t.Fatalf("ApplyRelayDiscoveryResponse() fresh bootstrap error = %v", err)
927 }
928
929 - advertisedDescriptors := server.relaySet.AdvertisedDescriptors()
929 + advertisedDescriptors := server.relaySet.ConfirmedDescriptors()
930 advertisedURLs := make([]string, 0, len(advertisedDescriptors))
931 for _, descriptor := range advertisedDescriptors {
932 if strings.TrimSpace(descriptor.APIHTTPSAddr) == "" {
@@ -936,7 +936,7 @@ func TestServerExpiredDiscoveryPeerNeedsFreshDirectConfirmation(t *testing.T) {
936 }
937 sort.Strings(advertisedURLs)
938 if !reflect.DeepEqual(advertisedURLs, []string{"https://bootstrap.example.com"}) {
939 - t.Fatalf("AdvertisedDescriptors() = %v, want relay to stay hidden until reconfirmed", advertisedURLs)
939 + t.Fatalf("ConfirmedDescriptors() = %v, want relay to stay hidden until reconfirmed", advertisedURLs)
940 }
941
942 if err := applyRelay(
@@ -950,7 +950,7 @@ func TestServerExpiredDiscoveryPeerNeedsFreshDirectConfirmation(t *testing.T) {
950 t.Fatalf("ApplyRelayDiscoveryResponse() reconfirm error = %v", err)
951 }
952
953 - advertisedDescriptors = server.relaySet.AdvertisedDescriptors()
953 + advertisedDescriptors = server.relaySet.ConfirmedDescriptors()
954 advertisedURLs = advertisedURLs[:0]
955 for _, descriptor := range advertisedDescriptors {
956 if strings.TrimSpace(descriptor.APIHTTPSAddr) == "" {
@@ -960,7 +960,7 @@ func TestServerExpiredDiscoveryPeerNeedsFreshDirectConfirmation(t *testing.T) {
960 }
961 sort.Strings(advertisedURLs)
962 if !reflect.DeepEqual(advertisedURLs, []string{"https://bootstrap.example.com", "https://relay-a.example.com"}) {
963 - t.Fatalf("AdvertisedDescriptors() = %v, want relay restored after direct confirmation", advertisedURLs)
963 + t.Fatalf("ConfirmedDescriptors() = %v, want relay restored after direct confirmation", advertisedURLs)
964 }
965 }
966
sdk/expose.go
+100 -89
@@ -6,6 +6,7 @@ import (
6 "fmt"
7 "net"
8 "net/http"
9 + "sort"
10 "strings"
11 "sync"
12 "sync/atomic"
@@ -24,14 +25,15 @@ type Exposure struct {
25 cancel context.CancelFunc
26 done <-chan struct{}
27
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
28 + identity types.Identity
29 + TargetAddr string
30 + UDPAddr string
31 + udpEnabled bool
32 + tcpEnabled bool
33 + banMITM bool
34 + maxActiveRelays int
35 + metadata types.LeaseMetadata
36 + rootCAPEM []byte
37
38 accepted chan net.Conn
39 datagrams chan types.DatagramFrame
@@ -45,21 +47,22 @@ type Exposure struct {
47 }
48
49 type ExposeConfig struct {
48 - RelayURLs []string
49 - IdentityPath string
50 - IdentityJSON string
51 - Name string
52 - TargetAddr string
53 - UDPAddr string
54 - UDPEnabled bool
55 - TCPEnabled bool
56 - BanMITM bool
57 - Discovery bool
58 - Metadata types.LeaseMetadata
59 - RootCAPEM []byte
50 + RelayURLs []string
51 + IdentityPath string
52 + IdentityJSON string
53 + Name string
54 + TargetAddr string
55 + UDPAddr string
56 + UDPEnabled bool
57 + TCPEnabled bool
58 + BanMITM bool
59 + MaxActiveRelays int
60 + Discovery bool
61 + Metadata types.LeaseMetadata
62 + RootCAPEM []byte
63 }
64
62 -// Expose creates relay listeners for each normalized relay URL and exposes a
65 +// Expose creates relay listeners for the selected relay pool and exposes a
66 // dynamic listener hub for accepting traffic from all of them.
67 func Expose(ctx context.Context, cfg ExposeConfig) (*Exposure, error) {
68 relayURLs, err := utils.ResolvePortalRelayURLs(ctx, cfg.RelayURLs, cfg.Discovery)
@@ -100,20 +103,21 @@ func Expose(ctx context.Context, cfg ExposeConfig) (*Exposure, error) {
103
104 exposureCtx, cancel := context.WithCancel(ctx)
105 exposure := &Exposure{
103 - cancel: cancel,
104 - done: exposureCtx.Done(),
105 - identity: identity,
106 - TargetAddr: targetAddr,
107 - UDPAddr: udpAddr,
108 - udpEnabled: cfg.UDPEnabled,
109 - tcpEnabled: cfg.TCPEnabled,
110 - banMITM: cfg.BanMITM,
111 - metadata: cfg.Metadata.Copy(),
112 - rootCAPEM: append([]byte(nil), cfg.RootCAPEM...),
113 - accepted: make(chan net.Conn, max(len(relayURLs)*defaultReadyTarget*2, 1)),
114 - datagrams: make(chan types.DatagramFrame, max(len(relayURLs)*32, 1)),
115 - relaySet: relaySet,
116 - relayListeners: make(map[string]*Listener, len(relayURLs)),
106 + cancel: cancel,
107 + done: exposureCtx.Done(),
108 + identity: identity,
109 + TargetAddr: targetAddr,
110 + UDPAddr: udpAddr,
111 + udpEnabled: cfg.UDPEnabled,
112 + tcpEnabled: cfg.TCPEnabled,
113 + banMITM: cfg.BanMITM,
114 + maxActiveRelays: cfg.MaxActiveRelays,
115 + metadata: cfg.Metadata.Copy(),
116 + rootCAPEM: append([]byte(nil), cfg.RootCAPEM...),
117 + accepted: make(chan net.Conn, max(len(relayURLs)*defaultReadyTarget*2, 1)),
118 + datagrams: make(chan types.DatagramFrame, max(len(relayURLs)*32, 1)),
119 + relaySet: relaySet,
120 + relayListeners: make(map[string]*Listener, len(relayURLs)),
121 }
122
123 if len(relayURLs) > 0 {
@@ -159,21 +163,34 @@ func (e *Exposure) runDiscoveryLoop(ctx context.Context) {
163 }
164
165 func (e *Exposure) ActiveRelayURLs() []string {
162 - var out []string
163 - for _, state := range e.relaySet.ActiveRelays() {
164 - out = append(out, state.Descriptor.APIHTTPSAddr)
166 + return append([]string(nil), e.clientState().ActiveRelayURLs...)
167 +}
168 +
169 +func (e *Exposure) clientState() discovery.ClientState {
170 + state := discovery.ClientState{
171 + MaxActiveRelays: e.maxActiveRelays,
172 + RequireUDP: e.udpEnabled,
173 + RequireTCP: e.tcpEnabled,
174 + }
175 +
176 + e.listenerMu.RLock()
177 + state.ActiveRelayURLs = make([]string, 0, len(e.relayListeners))
178 + for relayURL := range e.relayListeners {
179 + state.ActiveRelayURLs = append(state.ActiveRelayURLs, relayURL)
180 }
166 - return out
181 + e.listenerMu.RUnlock()
182 + sort.Strings(state.ActiveRelayURLs)
183 + return state
184 }
185
186 func (e *Exposure) Addr() net.Addr {
170 - return listenerAddr("portal:exposure")
187 + if e.identity.Address == "" {
188 + return listenerAddr("portal:exposure")
189 + }
190 + return listenerAddr("portal:" + e.identity.Address)
191 }
192
193 func (e *Exposure) Identity() types.Identity {
174 - if e == nil {
175 - return types.Identity{}
176 - }
194 return e.identity.Copy()
195 }
196
@@ -214,16 +231,10 @@ func (e *Exposure) WaitDatagramReady(ctx context.Context) ([]string, error) {
231
232 for {
233 e.listenerMu.RLock()
217 - listeners := make([]*Listener, 0, len(e.relayListeners))
218 - for _, listener := range e.relayListeners {
219 - listeners = append(listeners, listener)
220 - }
221 - e.listenerMu.RUnlock()
222 -
223 - addrs := make([]string, 0, len(listeners))
234 + addrs := make([]string, 0, len(e.relayListeners))
235 seen := make(map[string]struct{})
236 resolvedWithoutDatagram := true
226 - for _, listener := range listeners {
237 + for _, listener := range e.relayListeners {
238 if listener == nil {
239 continue
240 }
@@ -239,6 +250,7 @@ func (e *Exposure) WaitDatagramReady(ctx context.Context) ([]string, error) {
250 resolvedWithoutDatagram = false
251 }
252 }
253 + e.listenerMu.RUnlock()
254 if len(addrs) > 0 {
255 return addrs, nil
256 }
@@ -257,21 +269,14 @@ func (e *Exposure) WaitDatagramReady(ctx context.Context) ([]string, error) {
269 }
270
271 func (e *Exposure) RunHTTP(ctx context.Context, handler http.Handler, localAddr string) error {
260 - var relayListener net.Listener
272 e.listenerMu.RLock()
262 - activeListeners := make([]*Listener, 0, len(e.relayListeners))
263 - for _, state := range e.relaySet.ActiveRelays() {
264 - listener, ok := e.relayListeners[state.Descriptor.APIHTTPSAddr]
265 - if !ok {
266 - continue
267 - }
268 - activeListeners = append(activeListeners, listener)
269 - }
273 + hasRelayListeners := len(e.relayListeners) > 0
274 e.listenerMu.RUnlock()
271 - if len(activeListeners) > 0 {
272 - relayListener = e
275 +
276 + if hasRelayListeners {
277 + return RunHTTP(ctx, e, handler, localAddr)
278 }
274 - return RunHTTP(ctx, relayListener, handler, localAddr)
279 + return RunHTTP(ctx, nil, handler, localAddr)
280 }
281
282 type exposureConn struct {
@@ -338,28 +343,26 @@ func (e *Exposure) Close() error {
343 e.cancel()
344 }
345
341 - e.listenerMu.RLock()
342 - relayURLs := make([]string, 0, len(e.relayListeners))
343 - listeners := make([]*Listener, 0, len(e.relayListeners))
344 - for relayURL, listener := range e.relayListeners {
345 - relayURLs = append(relayURLs, relayURL)
346 - listeners = append(listeners, listener)
347 - }
348 - e.listenerMu.RUnlock()
346 + e.listenerMu.Lock()
347 + relayListeners := e.relayListeners
348 + e.relayListeners = make(map[string]*Listener)
349 + e.listenerMu.Unlock()
350
350 - for _, listener := range listeners {
351 + relayURLs := make([]string, 0, len(relayListeners))
352 + for relayURL, listener := range relayListeners {
353 + relayURLs = append(relayURLs, relayURL)
354 if listener != nil {
355 closeErr = errors.Join(closeErr, listener.Close())
356 }
357 }
358
359 event := log.Info().
357 - Int("relay_count", len(listeners)).
360 + Int("relay_count", len(relayListeners)).
361 Strs("relays", relayURLs)
362 if closeErr != nil {
363 event = log.Warn().
364 Err(closeErr).
362 - Int("relay_count", len(listeners)).
365 + Int("relay_count", len(relayListeners)).
366 Strs("relays", relayURLs)
367 }
368 event.Msg("exposure closed")
@@ -368,30 +371,38 @@ func (e *Exposure) Close() error {
371 }
372
373 func (e *Exposure) reconcileRelayListeners(failOnError bool) error {
371 - e.listenerMu.Lock()
372 - currentRelayURLs := make([]string, 0, len(e.relayListeners))
373 - for relayURL := range e.relayListeners {
374 - currentRelayURLs = append(currentRelayURLs, relayURL)
374 + selectedRelays := e.relaySet.PriorityRelays(e.clientState())
375 + desiredRelayURLs := make(map[string]struct{}, len(selectedRelays))
376 + for _, state := range selectedRelays {
377 + desiredRelayURLs[state.Descriptor.APIHTTPSAddr] = struct{}{}
378 }
376 - activeRelayURLs := make([]string, 0, len(e.relayListeners))
377 - for _, state := range e.relaySet.ActiveRelays() {
378 - activeRelayURLs = append(activeRelayURLs, state.Descriptor.APIHTTPSAddr)
379 - }
380 - missingRelayURLs := utils.FilterRelayURLs(activeRelayURLs, currentRelayURLs)
381 - staleRelayURLs := utils.FilterRelayURLs(currentRelayURLs, activeRelayURLs)
382 - staleListeners := make([]*Listener, 0, len(staleRelayURLs))
383 - for _, relayURL := range staleRelayURLs {
384 - staleListeners = append(staleListeners, e.relayListeners[relayURL])
379 +
380 + e.listenerMu.Lock()
381 + staleRelayListeners := make(map[string]*Listener)
382 + for relayURL, listener := range e.relayListeners {
383 + if _, ok := desiredRelayURLs[relayURL]; ok {
384 + continue
385 + }
386 + staleRelayListeners[relayURL] = listener
387 delete(e.relayListeners, relayURL)
388 }
389 +
390 + missingRelayURLs := make([]string, 0, len(selectedRelays))
391 + for _, state := range selectedRelays {
392 + relayURL := state.Descriptor.APIHTTPSAddr
393 + if _, ok := e.relayListeners[relayURL]; ok {
394 + continue
395 + }
396 + missingRelayURLs = append(missingRelayURLs, relayURL)
397 + }
398 e.listenerMu.Unlock()
399
389 - for i, listener := range staleListeners {
400 + for relayURL, listener := range staleRelayListeners {
401 if listener == nil {
402 continue
403 }
404 if err := listener.Close(); err != nil && !errors.Is(err, net.ErrClosed) {
394 - log.Warn().Err(err).Str("relay_url", staleRelayURLs[i]).Msg("close stale relay listener")
405 + log.Warn().Err(err).Str("relay_url", relayURL).Msg("close stale relay listener")
406 }
407 }
408 for _, relayURL := range missingRelayURLs {
sdk/expose_test.go
-1
@@ -29,7 +29,6 @@ func mustRelayDescriptor(t *testing.T, relayName, relayURL string) types.RelayDe
29 Name: relayName,
30 },
31 RelayID: relayURL,
32 - Sequence: uint64(now.UnixMilli()),
32 Version: 1,
33 IssuedAt: now,
34 ExpiresAt: now.Add(time.Hour),