main
go 259 lines 6.56 KB
Raw
1 package discovery
2
3 import (
4 "context"
5 "crypto/tls"
6 "net/http"
7 "net/url"
8 "time"
9
10 "github.com/rs/zerolog/log"
11
12 "github.com/gosuda/portal-tunnel/v2/types"
13 "github.com/gosuda/portal-tunnel/v2/utils"
14 )
15
16 const (
17 defaultRequestTimeout = 15 * time.Second
18 DiscoveryPollInterval = 30 * time.Second
19 defaultRecoveryFailures = 5
20 )
21
22 type OverlayRuntime interface {
23 DiscoverRelay(context.Context, types.RelayDescriptor) (types.DiscoveryResponse, error)
24 Sync([]types.RelayDescriptor) error
25 }
26
27 type Refresher struct {
28 relaySet *RelaySet
29 httpClient *http.Client
30 overlay OverlayRuntime
31 directRecoveryFailures int
32 lastAnnounceSuccess map[string]bool
33 }
34
35 func NewRefresher(relaySet *RelaySet, overlay OverlayRuntime) *Refresher {
36 return &Refresher{
37 relaySet: relaySet,
38 httpClient: utils.NewHTTPClient(
39 utils.WithHTTPTLSConfig(&tls.Config{
40 MinVersion: tls.VersionTLS12,
41 NextProtos: []string{"http/1.1"},
42 }),
43 utils.WithoutHTTP2(),
44 utils.WithHTTPTimeout(defaultRequestTimeout),
45 ),
46 overlay: overlay,
47 directRecoveryFailures: defaultRecoveryFailures,
48 lastAnnounceSuccess: make(map[string]bool),
49 }
50 }
51
52 func (r *Refresher) Refresh(ctx context.Context, self *types.RelayDescriptor) error {
53 if r.overlay != nil {
54 if err := r.refreshOverlay(ctx); err != nil && ctx.Err() == nil {
55 log.Warn().
56 Err(err).
57 Msg("overlay discovery failed")
58 }
59 if ctx.Err() != nil {
60 return ctx.Err()
61 }
62 }
63 if err := r.refreshHTTPS(ctx); err != nil {
64 return err
65 }
66 if self != nil {
67 if err := r.announceSelf(ctx, *self); err != nil {
68 return err
69 }
70 }
71 return nil
72 }
73
74 func (r *Refresher) announceSelf(ctx context.Context, descriptor types.RelayDescriptor) error {
75 req := types.DiscoveryAnnounceRequest{
76 ProtocolVersion: types.DiscoveryVersion,
77 Descriptor: descriptor,
78 }
79
80 for _, relayURL := range r.relaySet.BootstrapRelayURLs() {
81 if relayURL == descriptor.APIHTTPSAddr {
82 continue
83 }
84 baseURL, err := url.Parse(relayURL)
85 if err != nil {
86 if r.shouldLogAnnounce(relayURL, false) {
87 log.Warn().
88 Err(err).
89 Str("relay", relayURL).
90 Msg("relay discovery announce target skipped")
91 }
92 continue
93 }
94
95 if err := utils.HTTPDoAPIPath(ctx, r.httpClient, baseURL, http.MethodPost, types.PathDiscoveryAnnounce, req, nil, nil); err != nil {
96 if ctx.Err() != nil {
97 return ctx.Err()
98 }
99 if r.shouldLogAnnounce(relayURL, false) {
100 log.Warn().
101 Err(err).
102 Str("relay", relayURL).
103 Msg("relay discovery announce failed")
104 }
105 continue
106 }
107 if r.shouldLogAnnounce(relayURL, true) {
108 log.Info().
109 Str("relay", relayURL).
110 Msg("relay discovery announce succeeded")
111 }
112 }
113 return nil
114 }
115
116 func (r *Refresher) shouldLogAnnounce(relayURL string, success bool) bool {
117 previous, ok := r.lastAnnounceSuccess[relayURL]
118 if ok && previous == success {
119 return false
120 }
121 r.lastAnnounceSuccess[relayURL] = success
122 return true
123 }
124
125 func (r *Refresher) refreshHTTPS(ctx context.Context) error {
126 now := time.Now().UTC()
127 for _, state := range r.relaySet.refreshCandidates(now) {
128 relayURL := state.Descriptor.APIHTTPSAddr
129
130 recoveryFailures := r.directRecoveryFailures
131 if state.Bootstrap {
132 recoveryFailures = 0
133 }
134
135 baseURL, err := url.Parse(relayURL)
136 if err != nil {
137 if recoveryFailures > 0 {
138 r.logDiscoveryFailure(relayURL, relayURL, recoveryFailures, err)
139 }
140 continue
141 }
142 client := r.httpClient
143 var closeClient func()
144 if utils.IsLocalRelayHost(baseURL.Hostname()) {
145 _, localClient, transport, err := utils.NewHTTPTLSClient(ctx, baseURL, defaultRequestTimeout)
146 if err != nil {
147 if recoveryFailures > 0 {
148 r.logDiscoveryFailure(relayURL, relayURL, recoveryFailures, err)
149 }
150 continue
151 }
152 client = localClient
153 closeClient = transport.CloseIdleConnections
154 }
155
156 startedAt := time.Now()
157 var resp types.DiscoveryResponse
158 if err := utils.HTTPDoAPIPath(ctx, client, baseURL, http.MethodGet, types.PathDiscovery, nil, nil, &resp); err != nil {
159 if closeClient != nil {
160 closeClient()
161 }
162 if ctx.Err() != nil {
163 return ctx.Err()
164 }
165 if recoveryFailures > 0 {
166 r.logDiscoveryFailure(relayURL, relayURL, recoveryFailures, err)
167 }
168 continue
169 }
170 if closeClient != nil {
171 closeClient()
172 }
173 measuredAt := time.Now().UTC()
174
175 if _, err := r.relaySet.ApplyRelayDiscoveryResponse(relayURL, resp, measuredAt); err != nil {
176 if recoveryFailures > 0 {
177 r.logDiscoveryFailure(relayURL, relayURL, recoveryFailures, err)
178 }
179 continue
180 }
181 r.relaySet.RecordDiscoveryRTT(relayURL, time.Since(startedAt), measuredAt)
182 }
183 return nil
184 }
185
186 func (r *Refresher) refreshOverlay(ctx context.Context) error {
187 now := time.Now().UTC()
188 states := r.relaySet.overlayPeerRelayStates(now)
189 if len(states) == 0 {
190 return nil
191 }
192 descriptors := make([]types.RelayDescriptor, 0, len(states))
193 for _, state := range states {
194 descriptors = append(descriptors, state.Descriptor)
195 }
196 if err := r.overlay.Sync(descriptors); err != nil {
197 return err
198 }
199 relaySetChanged := false
200 for _, state := range r.relaySet.overlayRefreshCandidates(now) {
201 relay := state.Descriptor
202 recoveryFailures := r.directRecoveryFailures
203 if state.Bootstrap {
204 recoveryFailures = 0
205 }
206 startedAt := time.Now()
207 resp, err := r.overlay.DiscoverRelay(ctx, relay)
208 if err != nil {
209 if ctx.Err() != nil {
210 return ctx.Err()
211 }
212 if recoveryFailures > 0 {
213 r.logDiscoveryFailure(relay.APIHTTPSAddr, relay.APIHTTPSAddr, recoveryFailures, err)
214 }
215 continue
216 }
217
218 measuredAt := time.Now().UTC()
219 changed, err := r.relaySet.ApplyRelayDiscoveryResponse(relay.APIHTTPSAddr, resp, measuredAt)
220 if err != nil {
221 if recoveryFailures > 0 {
222 r.logDiscoveryFailure(relay.APIHTTPSAddr, relay.APIHTTPSAddr, recoveryFailures, err)
223 }
224 continue
225 }
226 r.relaySet.RecordDiscoveryRTT(relay.APIHTTPSAddr, time.Since(startedAt), measuredAt)
227 if changed {
228 relaySetChanged = true
229 }
230 }
231 if !relaySetChanged {
232 return nil
233 }
234 if err := r.overlay.Sync(r.relaySet.OverlayPeerDescriptor()); err != nil {
235 return err
236 }
237 return nil
238 }
239
240 func (r *Refresher) logDiscoveryFailure(targetRelayURL, sourceURL string, recoveryFailures int, err error) {
241 backedOff, backoffReason, failureCount := r.relaySet.RecordDiscoveryFailure(targetRelayURL, recoveryFailures)
242 if !backedOff {
243 return
244 }
245
246 event := log.Warn().
247 Err(err).
248 Str("relay", sourceURL).
249 Bool("backed_off", true).
250 Str("reason", backoffReason)
251 if failureCount > 0 {
252 event = event.Int("discovery_failures", failureCount)
253 }
254 if backoffReason == "unhealthy" {
255 event.Msg("discovery source removed from relay pool")
256 return
257 }
258 event.Msg("discovery source retry delayed")
259 }