refact: consoliate api codes to utils
Kim committed
Apr 1, 2026 at 14:26 UTC
4839d4fd46770fec4f9d1a2ea93e3dc0dbfe3f92
11 files changed
+295
-274
cmd/relay-server/admin.go
+36
-48
@@ -145,11 +145,10 @@ func (f *Frontend) serveAdmin(w http.ResponseWriter, r *http.Request) {
145
f.handleLogin(w, r)
146
return
147
}
148
- utils.WriteAPIError(w, http.StatusMethodNotAllowed, types.APIErrorCodeMethodNotAllowed, "method not allowed")
148
+ utils.MethodNotAllowedError().Write(w)
149
return
150
case types.PathAdminLogout:
151
- if r.Method != http.MethodPost {
152
- utils.WriteAPIError(w, http.StatusMethodNotAllowed, types.APIErrorCodeMethodNotAllowed, "method not allowed")
151
+ if !utils.RequireMethod(w, r, http.MethodPost) {
152
return
153
}
154
if cookie, err := r.Cookie(cookieName); err == nil && cookie.Value != "" {
@@ -164,11 +163,10 @@ func (f *Frontend) serveAdmin(w http.ResponseWriter, r *http.Request) {
163
SameSite: http.SameSiteStrictMode,
164
MaxAge: -1,
165
})
167
- utils.WriteAPIData(w, http.StatusOK, map[string]any{})
166
+ utils.WriteAPIEmpty(w, http.StatusOK)
167
return
168
case types.PathAdminAuthStatus:
170
- if r.Method != http.MethodGet {
171
- utils.WriteAPIError(w, http.StatusMethodNotAllowed, types.APIErrorCodeMethodNotAllowed, "method not allowed")
169
+ if !utils.RequireMethod(w, r, http.MethodGet) {
170
return
171
}
172
utils.WriteAPIData(w, http.StatusOK, types.AdminAuthStatusResponse{
@@ -184,18 +182,12 @@ func (f *Frontend) serveAdmin(w http.ResponseWriter, r *http.Request) {
182
}
183
184
runtime := f.server.PolicyRuntime()
187
- methodNotAllowed := func() {
188
- utils.WriteAPIError(w, http.StatusMethodNotAllowed, types.APIErrorCodeMethodNotAllowed, "method not allowed")
189
- }
190
- writeOK := func() {
191
- f.saveAdminState(runtime)
192
- utils.WriteAPIData(w, http.StatusOK, map[string]any{})
193
- }
185
+ methodNotAllowed := utils.MethodNotAllowedError()
186
+ invalidRequestBody := utils.InvalidRequestMessage("invalid request body")
187
188
switch path {
189
case types.PathAdminSnapshot:
197
- if r.Method != http.MethodGet {
198
- methodNotAllowed()
190
+ if !utils.RequireMethod(w, r, http.MethodGet) {
191
return
192
}
193
utils.WriteAPIData(w, http.StatusOK, types.AdminSnapshotResponse{
@@ -208,13 +200,11 @@ func (f *Frontend) serveAdmin(w http.ResponseWriter, r *http.Request) {
200
},
201
})
202
case types.PathAdminLandingPage:
211
- if r.Method != http.MethodPost {
212
- methodNotAllowed()
203
+ if !utils.RequireMethod(w, r, http.MethodPost) {
204
return
205
}
215
- var req types.AdminLandingPageSettingsRequest
216
- if err := utils.DecodeJSONBody(w, r, &req, 1<<16); err != nil {
217
- utils.WriteAPIError(w, http.StatusBadRequest, types.APIErrorCodeInvalidRequest, "invalid request body")
206
+ req, ok := utils.DecodeJSONRequestAs[types.AdminLandingPageSettingsRequest](w, r, 1<<16, invalidRequestBody)
207
+ if !ok {
208
return
209
}
210
f.setLandingPageEnabled(req.Enabled)
@@ -223,13 +213,11 @@ func (f *Frontend) serveAdmin(w http.ResponseWriter, r *http.Request) {
213
Enabled: f.isLandingPageEnabled(),
214
})
215
case types.PathAdminUDP:
226
- if r.Method != http.MethodPost {
227
- methodNotAllowed()
216
+ if !utils.RequireMethod(w, r, http.MethodPost) {
217
return
218
}
230
- var req types.AdminUDPSettingsRequest
231
- if err := utils.DecodeJSONBody(w, r, &req, 1<<16); err != nil {
232
- utils.WriteAPIError(w, http.StatusBadRequest, types.APIErrorCodeInvalidRequest, "invalid request body")
219
+ req, ok := utils.DecodeJSONRequestAs[types.AdminUDPSettingsRequest](w, r, 1<<16, invalidRequestBody)
220
+ if !ok {
221
return
222
}
223
if req.MaxLeases < 0 {
@@ -243,13 +231,11 @@ func (f *Frontend) serveAdmin(w http.ResponseWriter, r *http.Request) {
231
MaxLeases: runtime.UDPMaxLeases(),
232
})
233
case types.PathAdminApproval:
246
- if r.Method != http.MethodPost {
247
- methodNotAllowed()
234
+ if !utils.RequireMethod(w, r, http.MethodPost) {
235
return
236
}
250
- var req types.AdminApprovalModeRequest
251
- if err := utils.DecodeJSONBody(w, r, &req, 1<<16); err != nil {
252
- utils.WriteAPIError(w, http.StatusBadRequest, types.APIErrorCodeInvalidRequest, "invalid request body")
237
+ req, ok := utils.DecodeJSONRequestAs[types.AdminApprovalModeRequest](w, r, 1<<16, invalidRequestBody)
238
+ if !ok {
239
return
240
}
241
if err := runtime.Approver().SetMode(policy.Mode(strings.TrimSpace(req.Mode))); err != nil {
@@ -284,16 +270,16 @@ func (f *Frontend) serveAdmin(w http.ResponseWriter, r *http.Request) {
270
case http.MethodDelete:
271
runtime.UnbanLease(leaseID)
272
default:
287
- methodNotAllowed()
273
+ methodNotAllowed.Write(w)
274
return
275
}
290
- writeOK()
276
+ f.saveAdminState(runtime)
277
+ utils.WriteAPIEmpty(w, http.StatusOK)
278
case "bps":
279
switch r.Method {
280
case http.MethodPost:
294
- var req types.AdminBPSRequest
295
- if err := utils.DecodeJSONBody(w, r, &req, 1<<16); err != nil {
296
- utils.WriteAPIError(w, http.StatusBadRequest, types.APIErrorCodeInvalidRequest, "invalid request body")
281
+ req, ok := utils.DecodeJSONRequestAs[types.AdminBPSRequest](w, r, 1<<16, invalidRequestBody)
282
+ if !ok {
283
return
284
}
285
if req.BPS <= 0 {
@@ -304,10 +290,11 @@ func (f *Frontend) serveAdmin(w http.ResponseWriter, r *http.Request) {
290
case http.MethodDelete:
291
runtime.BPSManager().DeleteLeaseBPS(leaseID)
292
default:
307
- methodNotAllowed()
293
+ methodNotAllowed.Write(w)
294
return
295
}
310
- writeOK()
296
+ f.saveAdminState(runtime)
297
+ utils.WriteAPIEmpty(w, http.StatusOK)
298
case "approve":
299
approver := runtime.Approver()
300
switch r.Method {
@@ -317,10 +304,11 @@ func (f *Frontend) serveAdmin(w http.ResponseWriter, r *http.Request) {
304
case http.MethodDelete:
305
approver.Revoke(leaseID)
306
default:
320
- methodNotAllowed()
307
+ methodNotAllowed.Write(w)
308
return
309
}
323
- writeOK()
310
+ f.saveAdminState(runtime)
311
+ utils.WriteAPIEmpty(w, http.StatusOK)
312
case "deny":
313
approver := runtime.Approver()
314
switch r.Method {
@@ -329,10 +317,11 @@ func (f *Frontend) serveAdmin(w http.ResponseWriter, r *http.Request) {
317
case http.MethodDelete:
318
approver.Undeny(leaseID)
319
default:
332
- methodNotAllowed()
320
+ methodNotAllowed.Write(w)
321
return
322
}
335
- writeOK()
323
+ f.saveAdminState(runtime)
324
+ utils.WriteAPIEmpty(w, http.StatusOK)
325
default:
326
http.NotFound(w, r)
327
}
@@ -356,10 +345,11 @@ func (f *Frontend) serveAdmin(w http.ResponseWriter, r *http.Request) {
345
case http.MethodDelete:
346
filter.UnbanIP(rawIP)
347
default:
359
- methodNotAllowed()
348
+ methodNotAllowed.Write(w)
349
return
350
}
362
- writeOK()
351
+ f.saveAdminState(runtime)
352
+ utils.WriteAPIEmpty(w, http.StatusOK)
353
default:
354
http.NotFound(w, r)
355
}
@@ -367,8 +357,7 @@ func (f *Frontend) serveAdmin(w http.ResponseWriter, r *http.Request) {
357
}
358
359
func (f *Frontend) handleLogin(w http.ResponseWriter, r *http.Request) {
370
- if r.Method != http.MethodPost {
371
- utils.WriteAPIError(w, http.StatusMethodNotAllowed, types.APIErrorCodeMethodNotAllowed, "method not allowed")
360
+ if !utils.RequireMethod(w, r, http.MethodPost) {
361
return
362
}
363
if !f.auth.AuthEnabled() {
@@ -376,9 +365,8 @@ func (f *Frontend) handleLogin(w http.ResponseWriter, r *http.Request) {
365
return
366
}
367
379
- var req types.AdminLoginRequest
380
- if err := utils.DecodeJSONBody(w, r, &req, 1<<16); err != nil {
381
- utils.WriteAPIError(w, http.StatusBadRequest, types.APIErrorCodeInvalidRequest, "invalid request body")
368
+ req, ok := utils.DecodeJSONRequestAs[types.AdminLoginRequest](w, r, 1<<16, utils.InvalidRequestMessage("invalid request body"))
369
+ if !ok {
370
return
371
}
372
if !f.auth.ValidateKey(req.Key) {
cmd/relay-server/frontend.go
+2
-3
@@ -205,14 +205,13 @@ func (f *Frontend) injectServerData(htmlContent string) string {
205
}
206
207
func (f *Frontend) serveTunnelStatus(w http.ResponseWriter, r *http.Request) {
208
- if r.Method != http.MethodGet {
209
- utils.WriteAPIError(w, http.StatusMethodNotAllowed, types.APIErrorCodeMethodNotAllowed, "method not allowed")
208
+ if !utils.RequireMethod(w, r, http.MethodGet) {
209
return
210
}
211
212
hostname := strings.ToLower(strings.TrimSpace(r.URL.Query().Get("hostname")))
213
if hostname == "" {
215
- utils.WriteAPIError(w, http.StatusBadRequest, types.APIErrorCodeInvalidRequest, "hostname is required")
214
+ utils.InvalidRequestMessage("hostname is required").Write(w)
215
return
216
}
217
portal/acme/route53/provider.go
+4
-8
@@ -171,7 +171,10 @@ func upsertARecord(ctx context.Context, client *awsroute53.Client, hostedZoneID,
171
return errors.New("hosted zone id is required")
172
}
173
174
- fqdn := ensureTrailingDot(utils.NormalizeHostname(name))
174
+ fqdn := utils.NormalizeHostname(name)
175
+ if !strings.HasSuffix(fqdn, ".") {
176
+ fqdn += "."
177
+ }
178
recordSet := &types.ResourceRecordSet{
179
Name: aws.String(fqdn),
180
Type: types.RRTypeA,
@@ -251,10 +254,3 @@ func normalizeZoneID(raw string) string {
254
trimmed := strings.TrimSpace(raw)
255
return strings.TrimPrefix(trimmed, "/hostedzone/")
256
}
254
-
255
-func ensureTrailingDot(name string) string {
256
- if strings.HasSuffix(name, ".") {
257
- return name
258
- }
259
- return name + "."
260
-}
portal/api_server.go
+61
-70
@@ -112,9 +112,28 @@ func (s *Server) handleHealthz(w http.ResponseWriter, _ *http.Request) {
112
utils.WriteAPIData(w, http.StatusOK, map[string]any{"status": "ok"})
113
}
114
115
+func (s *Server) extractAllowedClientIP(w http.ResponseWriter, r *http.Request) (string, bool) {
116
+ clientIP := policy.ExtractClientIP(r, s.cfg.TrustProxyHeaders, s.trustedProxyCIDRs)
117
+ if !s.registry.policy.IPFilter().IsIPBanned(clientIP) {
118
+ return clientIP, true
119
+ }
120
+ utils.WriteAPIError(w, http.StatusForbidden, types.APIErrorCodeIPBanned, "request denied because source IP is banned")
121
+ return "", false
122
+}
123
+
124
+func leaseLookupError(err error) utils.APIErrorResponse {
125
+ if errors.Is(err, errLeaseNotFound) {
126
+ return utils.APIErrorResponse{
127
+ Status: http.StatusNotFound,
128
+ Code: types.APIErrorCodeLeaseNotFound,
129
+ Message: err.Error(),
130
+ }
131
+ }
132
+ return utils.InvalidRequestError(err)
133
+}
134
+
135
func (s *Server) handleRelayDiscovery(w http.ResponseWriter, r *http.Request) {
116
- if r.Method != http.MethodGet {
117
- utils.WriteAPIError(w, http.StatusMethodNotAllowed, types.APIErrorCodeMethodNotAllowed, "method not allowed")
136
+ if !utils.RequireMethod(w, r, http.MethodGet) {
137
return
138
}
139
@@ -164,8 +183,7 @@ func (s *Server) handleRelayDiscovery(w http.ResponseWriter, r *http.Request) {
183
func (s *Server) handleDomain(w http.ResponseWriter, r *http.Request) {
184
w.Header().Set("Access-Control-Allow-Origin", "*")
185
167
- if r.Method != http.MethodGet {
168
- utils.WriteAPIError(w, http.StatusMethodNotAllowed, types.APIErrorCodeMethodNotAllowed, "method not allowed")
186
+ if !utils.RequireMethod(w, r, http.MethodGet) {
187
return
188
}
189
@@ -176,20 +194,17 @@ func (s *Server) handleDomain(w http.ResponseWriter, r *http.Request) {
194
}
195
196
func (s *Server) handleRegister(w http.ResponseWriter, r *http.Request) {
179
- if r.Method != http.MethodPost {
180
- utils.WriteAPIError(w, http.StatusMethodNotAllowed, types.APIErrorCodeMethodNotAllowed, "method not allowed")
197
+ if !utils.RequireMethod(w, r, http.MethodPost) {
198
return
199
}
200
184
- clientIP := policy.ExtractClientIP(r, s.cfg.TrustProxyHeaders, s.trustedProxyCIDRs)
185
- if s.registry.policy.IPFilter().IsIPBanned(clientIP) {
186
- utils.WriteAPIError(w, http.StatusForbidden, types.APIErrorCodeIPBanned, "request denied because source IP is banned")
201
+ clientIP, ok := s.extractAllowedClientIP(w, r)
202
+ if !ok {
203
return
204
}
205
190
- var req types.RegisterRequest
191
- if err := utils.DecodeJSONBody(w, r, &req, defaultControlBodyLimit); err != nil {
192
- utils.WriteAPIError(w, http.StatusBadRequest, types.APIErrorCodeInvalidJSON, err.Error())
206
+ req, ok := utils.DecodeJSONRequest[types.RegisterRequest](w, r, defaultControlBodyLimit)
207
+ if !ok {
208
return
209
}
210
@@ -199,7 +214,7 @@ func (s *Server) handleRegister(w http.ResponseWriter, r *http.Request) {
214
case errors.Is(err, auth.ErrInvalidSignature):
215
utils.WriteAPIError(w, http.StatusForbidden, types.APIErrorCodeUnauthorized, err.Error())
216
default:
202
- utils.WriteAPIError(w, http.StatusBadRequest, types.APIErrorCodeInvalidRequest, err.Error())
217
+ utils.InvalidRequestError(err).Write(w)
218
}
219
return
220
}
@@ -220,7 +235,7 @@ func (s *Server) handleRegister(w http.ResponseWriter, r *http.Request) {
235
case errors.Is(err, errUDPCapacityExceeded):
236
utils.WriteAPIError(w, http.StatusServiceUnavailable, types.APIErrorCodeUDPCapacityExceeded, err.Error())
237
default:
223
- utils.WriteAPIError(w, http.StatusBadRequest, types.APIErrorCodeInvalidRequest, err.Error())
238
+ utils.InvalidRequestError(err).Write(w)
239
}
240
return
241
}
@@ -229,20 +244,16 @@ func (s *Server) handleRegister(w http.ResponseWriter, r *http.Request) {
244
}
245
246
func (s *Server) handleRegisterChallenge(w http.ResponseWriter, r *http.Request) {
232
- if r.Method != http.MethodPost {
233
- utils.WriteAPIError(w, http.StatusMethodNotAllowed, types.APIErrorCodeMethodNotAllowed, "method not allowed")
247
+ if !utils.RequireMethod(w, r, http.MethodPost) {
248
return
249
}
250
237
- clientIP := policy.ExtractClientIP(r, s.cfg.TrustProxyHeaders, s.trustedProxyCIDRs)
238
- if s.registry.policy.IPFilter().IsIPBanned(clientIP) {
239
- utils.WriteAPIError(w, http.StatusForbidden, types.APIErrorCodeIPBanned, "request denied because source IP is banned")
251
+ if _, ok := s.extractAllowedClientIP(w, r); !ok {
252
return
253
}
254
243
- var req types.RegisterChallengeRequest
244
- if err := utils.DecodeJSONBody(w, r, &req, defaultControlBodyLimit); err != nil {
245
- utils.WriteAPIError(w, http.StatusBadRequest, types.APIErrorCodeInvalidJSON, err.Error())
255
+ req, ok := utils.DecodeJSONRequest[types.RegisterChallengeRequest](w, r, defaultControlBodyLimit)
256
+ if !ok {
257
return
258
}
259
@@ -275,7 +286,7 @@ func (s *Server) handleRegisterChallenge(w http.ResponseWriter, r *http.Request)
286
case errors.Is(err, errUDPCapacityExceeded):
287
utils.WriteAPIError(w, http.StatusServiceUnavailable, types.APIErrorCodeUDPCapacityExceeded, err.Error())
288
default:
278
- utils.WriteAPIError(w, http.StatusBadRequest, types.APIErrorCodeInvalidRequest, err.Error())
289
+ utils.InvalidRequestError(err).Write(w)
290
}
291
return
292
}
@@ -284,20 +295,17 @@ func (s *Server) handleRegisterChallenge(w http.ResponseWriter, r *http.Request)
295
}
296
297
func (s *Server) handleRenew(w http.ResponseWriter, r *http.Request) {
287
- if r.Method != http.MethodPost {
288
- utils.WriteAPIError(w, http.StatusMethodNotAllowed, types.APIErrorCodeMethodNotAllowed, "method not allowed")
298
+ if !utils.RequireMethod(w, r, http.MethodPost) {
299
return
300
}
301
292
- clientIP := policy.ExtractClientIP(r, s.cfg.TrustProxyHeaders, s.trustedProxyCIDRs)
293
- if s.registry.policy.IPFilter().IsIPBanned(clientIP) {
294
- utils.WriteAPIError(w, http.StatusForbidden, types.APIErrorCodeIPBanned, "request denied because source IP is banned")
302
+ clientIP, ok := s.extractAllowedClientIP(w, r)
303
+ if !ok {
304
return
305
}
306
298
- var req types.RenewRequest
299
- if err := utils.DecodeJSONBody(w, r, &req, defaultControlBodyLimit); err != nil {
300
- utils.WriteAPIError(w, http.StatusBadRequest, types.APIErrorCodeInvalidJSON, err.Error())
307
+ req, ok := utils.DecodeJSONRequest[types.RenewRequest](w, r, defaultControlBodyLimit)
308
+ if !ok {
309
return
310
}
311
@@ -313,12 +321,7 @@ func (s *Server) handleRenew(w http.ResponseWriter, r *http.Request) {
321
}
322
record, err := s.registry.Renew(strings.TrimSpace(req.LeaseID), ttl, clientIP, utils.SanitizeReportedIP(req.ReportedIP))
323
if err != nil {
316
- switch {
317
- case errors.Is(err, errLeaseNotFound):
318
- utils.WriteAPIError(w, http.StatusNotFound, types.APIErrorCodeLeaseNotFound, err.Error())
319
- default:
320
- utils.WriteAPIError(w, http.StatusBadRequest, types.APIErrorCodeInvalidRequest, err.Error())
321
- }
324
+ leaseLookupError(err).Write(w)
325
return
326
}
327
nextAccessToken, _, err := auth.IssueLeaseAccessToken(s.ownerIdentity.PrivateKey, s.ownerIdentity.Address, s.cfg.PortalURL, claims.Subject, record.ID, ttl)
@@ -335,14 +338,12 @@ func (s *Server) handleRenew(w http.ResponseWriter, r *http.Request) {
338
}
339
340
func (s *Server) handleUnregister(w http.ResponseWriter, r *http.Request) {
338
- if r.Method != http.MethodPost {
339
- utils.WriteAPIError(w, http.StatusMethodNotAllowed, types.APIErrorCodeMethodNotAllowed, "method not allowed")
341
+ if !utils.RequireMethod(w, r, http.MethodPost) {
342
return
343
}
344
343
- var req types.UnregisterRequest
344
- if err := utils.DecodeJSONBody(w, r, &req, defaultControlBodyLimit); err != nil {
345
- utils.WriteAPIError(w, http.StatusBadRequest, types.APIErrorCodeInvalidJSON, err.Error())
345
+ req, ok := utils.DecodeJSONRequest[types.UnregisterRequest](w, r, defaultControlBodyLimit)
346
+ if !ok {
347
return
348
}
349
if _, err := auth.VerifyLeaseAccessToken(req.AccessToken, s.ownerIdentity.PublicKey, s.cfg.PortalURL, strings.TrimSpace(req.LeaseID), time.Now().UTC()); err != nil {
@@ -352,24 +353,18 @@ func (s *Server) handleUnregister(w http.ResponseWriter, r *http.Request) {
353
354
record, err := s.registry.Unregister(strings.TrimSpace(req.LeaseID))
355
if err != nil {
355
- switch {
356
- case errors.Is(err, errLeaseNotFound):
357
- utils.WriteAPIError(w, http.StatusNotFound, types.APIErrorCodeLeaseNotFound, err.Error())
358
- default:
359
- utils.WriteAPIError(w, http.StatusBadRequest, types.APIErrorCodeInvalidRequest, err.Error())
360
- }
356
+ leaseLookupError(err).Write(w)
357
return
358
}
359
if record != nil {
360
record.Close()
361
}
362
367
- utils.WriteAPIData(w, http.StatusOK, map[string]any{})
363
+ utils.WriteAPIEmpty(w, http.StatusOK)
364
}
365
366
func (s *Server) handleConnect(w http.ResponseWriter, r *http.Request) {
371
- if r.Method != http.MethodGet {
372
- utils.WriteAPIError(w, http.StatusMethodNotAllowed, types.APIErrorCodeMethodNotAllowed, "method not allowed")
367
+ if !utils.RequireMethod(w, r, http.MethodGet) {
368
return
369
}
370
if r.ProtoMajor != 1 {
@@ -379,29 +374,25 @@ func (s *Server) handleConnect(w http.ResponseWriter, r *http.Request) {
374
375
leaseID := strings.TrimSpace(r.URL.Query().Get("lease_id"))
376
token := strings.TrimSpace(r.Header.Get(types.HeaderAccessToken))
382
- clientIP := policy.ExtractClientIP(r, s.cfg.TrustProxyHeaders, s.trustedProxyCIDRs)
383
- if s.registry.policy.IPFilter().IsIPBanned(clientIP) {
384
- utils.WriteAPIError(w, http.StatusForbidden, types.APIErrorCodeIPBanned, "request denied because source IP is banned")
377
+ clientIP, ok := s.extractAllowedClientIP(w, r)
378
+ if !ok {
379
return
380
}
381
382
lease, err := s.admitLeaseByID(leaseID, token, false)
389
- switch {
390
- case err == nil:
391
- case errors.Is(err, errLeaseNotFound):
392
- utils.WriteAPIError(w, http.StatusNotFound, types.APIErrorCodeLeaseNotFound, err.Error())
393
- return
394
- case errors.Is(err, errLeaseRejected):
395
- utils.WriteAPIError(w, http.StatusForbidden, types.APIErrorCodeLeaseRejected, "lease is not approved for routing")
396
- return
397
- case errors.Is(err, errUnauthorized):
398
- utils.WriteAPIError(w, http.StatusForbidden, types.APIErrorCodeUnauthorized, err.Error())
399
- return
400
- case errors.Is(err, errTransportMismatch):
401
- utils.WriteAPIError(w, http.StatusConflict, types.APIErrorCodeTransportMismatch, "lease does not support stream transport")
402
- return
403
- default:
404
- utils.WriteAPIError(w, http.StatusBadRequest, types.APIErrorCodeInvalidRequest, err.Error())
383
+ if err != nil {
384
+ switch {
385
+ case errors.Is(err, errLeaseNotFound):
386
+ utils.WriteAPIError(w, http.StatusNotFound, types.APIErrorCodeLeaseNotFound, err.Error())
387
+ case errors.Is(err, errLeaseRejected):
388
+ utils.WriteAPIError(w, http.StatusForbidden, types.APIErrorCodeLeaseRejected, "lease is not approved for routing")
389
+ case errors.Is(err, errUnauthorized):
390
+ utils.WriteAPIError(w, http.StatusForbidden, types.APIErrorCodeUnauthorized, err.Error())
391
+ case errors.Is(err, errTransportMismatch):
392
+ utils.WriteAPIError(w, http.StatusConflict, types.APIErrorCodeTransportMismatch, "lease does not support stream transport")
393
+ default:
394
+ utils.InvalidRequestError(err).Write(w)
395
+ }
396
return
397
}
398
portal/discovery/discovery.go
+1
-3
@@ -157,8 +157,6 @@ func DiscoverRelayDiscovery(ctx context.Context, baseURL string, rootCAPEM []byt
157
return types.DiscoveryResponse{}, fmt.Errorf("parse discovery base url: %w", err)
158
}
159
160
- requestURL := parsedBaseURL.ResolveReference(&url.URL{Path: types.PathDiscovery})
161
-
160
client := httpClient
161
if client == nil {
162
_, client, err = keyless.NewRelayHTTPClient(ctx, parsedBaseURL, rootCAPEM, defaultRequestTimeout)
@@ -173,7 +171,7 @@ func DiscoverRelayDiscovery(ctx context.Context, baseURL string, rootCAPEM []byt
171
}
172
173
var resp types.DiscoveryResponse
176
- if err := utils.HTTPDoAPI(ctx, client, http.MethodGet, requestURL.String(), nil, nil, &resp); err != nil {
174
+ if err := utils.HTTPDoAPIPath(ctx, client, parsedBaseURL, http.MethodGet, types.PathDiscovery, nil, nil, &resp); err != nil {
175
return types.DiscoveryResponse{}, err
176
}
177
return resp, nil
portal/server.go
+1
@@ -527,6 +527,7 @@ func (s *Server) startOverlay() error {
527
s.overlay = overlay
528
return nil
529
}
530
+
531
func (s *Server) runRelayDiscoveryLoop(ctx context.Context) error {
532
ticker := time.NewTicker(types.DiscoveryPollInterval)
533
defer ticker.Stop()
sdk/api_client.go
+6
-7
@@ -103,7 +103,7 @@ func (a *apiClient) registerLease(ctx context.Context, ttl time.Duration, udpEna
103
}
104
105
var challenge types.RegisterChallengeResponse
106
- if err := utils.HTTPDoAPI(ctx, a.httpClient, http.MethodPost, a.baseURL.ResolveReference(&url.URL{Path: types.PathSDKRegisterChallenge}).String(), types.RegisterChallengeRequest{
106
+ if err := utils.HTTPDoAPIPath(ctx, a.httpClient, a.baseURL, http.MethodPost, types.PathSDKRegisterChallenge, types.RegisterChallengeRequest{
107
Name: a.name,
108
Metadata: a.metadata.Copy(),
109
OwnerAddress: a.ownerAddress,
@@ -119,7 +119,7 @@ func (a *apiClient) registerLease(ctx context.Context, ttl time.Duration, udpEna
119
}
120
121
var resp types.RegisterResponse
122
- if err := utils.HTTPDoAPI(ctx, a.httpClient, http.MethodPost, a.baseURL.ResolveReference(&url.URL{Path: types.PathSDKRegister}).String(), types.RegisterRequest{
122
+ if err := utils.HTTPDoAPIPath(ctx, a.httpClient, a.baseURL, http.MethodPost, types.PathSDKRegister, types.RegisterRequest{
123
ChallengeID: challenge.ChallengeID,
124
SIWEMessage: challenge.SIWEMessage,
125
SIWESignature: signature,
@@ -173,7 +173,7 @@ func (a *apiClient) reportedIP(ctx context.Context) string {
173
174
func (a *apiClient) ensureCompatible(ctx context.Context, httpClient *http.Client) error {
175
var resp types.DomainResponse
176
- if err := utils.HTTPDoAPI(ctx, httpClient, http.MethodGet, a.baseURL.ResolveReference(&url.URL{Path: types.PathSDKDomain}).String(), nil, nil, &resp); err != nil {
176
+ if err := utils.HTTPDoAPIPath(ctx, httpClient, a.baseURL, http.MethodGet, types.PathSDKDomain, nil, nil, &resp); err != nil {
177
err = fmt.Errorf("check relay compatibility: %w", err)
178
var netErr net.Error
179
var apiErr *types.APIRequestError
@@ -204,7 +204,7 @@ func (a *apiClient) renewLease(ctx context.Context, leaseID string, ttl time.Dur
204
}
205
206
var resp types.RenewResponse
207
- if err := utils.HTTPDoAPI(ctx, a.httpClient, http.MethodPost, a.baseURL.ResolveReference(&url.URL{Path: types.PathSDKRenew}).String(), types.RenewRequest{
207
+ if err := utils.HTTPDoAPIPath(ctx, a.httpClient, a.baseURL, http.MethodPost, types.PathSDKRenew, types.RenewRequest{
208
LeaseID: leaseID,
209
AccessToken: accessToken,
210
TTL: int(ttl / time.Second),
@@ -229,7 +229,7 @@ func (a *apiClient) unregisterLease(ctx context.Context, leaseID string) error {
229
a.mu.RLock()
230
accessToken := a.accessToken
231
a.mu.RUnlock()
232
- return utils.HTTPDoAPI(ctx, a.httpClient, http.MethodPost, a.baseURL.ResolveReference(&url.URL{Path: types.PathSDKUnregister}).String(), types.UnregisterRequest{
232
+ return utils.HTTPDoAPIPath(ctx, a.httpClient, a.baseURL, http.MethodPost, types.PathSDKUnregister, types.UnregisterRequest{
233
LeaseID: leaseID,
234
AccessToken: accessToken,
235
}, nil, nil)
@@ -250,8 +250,7 @@ func (a *apiClient) openReverseSession(ctx context.Context, leaseID string) (net
250
return nil, err
251
}
252
253
- connectRef, _ := url.Parse(types.PathSDKConnect)
254
- connectURL := a.baseURL.ResolveReference(connectRef)
253
+ connectURL := utils.ResolveAPIURL(a.baseURL, types.PathSDKConnect)
254
query := connectURL.Query()
255
query.Set("lease_id", leaseID)
256
connectURL.RawQuery = query.Encode()
utils/api.go
+174
-21
@@ -1,52 +1,159 @@
1
package utils
2
3
import (
4
+ "bytes"
5
+ "context"
6
"encoding/json"
7
"fmt"
8
"io"
9
"net/http"
10
+ "net/url"
11
"strings"
12
13
"github.com/gosuda/portal/v2/types"
14
)
15
13
-func WriteAPIEnvelope[T any](w http.ResponseWriter, status int, envelope types.APIEnvelope[T]) {
16
+type APIErrorResponse struct {
17
+ Status int
18
+ Code string
19
+ Message string
20
+}
21
+
22
+func (resp APIErrorResponse) Write(w http.ResponseWriter) {
23
+ WriteAPIError(w, resp.Status, resp.Code, resp.Message)
24
+}
25
+
26
+func WriteAPIData(w http.ResponseWriter, status int, data any) {
27
w.Header().Set("Content-Type", "application/json")
28
w.WriteHeader(status)
16
- _ = json.NewEncoder(w).Encode(envelope)
29
+ _ = json.NewEncoder(w).Encode(types.APIEnvelope[any]{OK: true, Data: data})
30
}
31
19
-func WriteAPIData(w http.ResponseWriter, status int, data any) {
20
- WriteAPIEnvelope(w, status, types.APIEnvelope[any]{OK: true, Data: data})
32
+func WriteAPIEmpty(w http.ResponseWriter, status int) {
33
+ WriteAPIData(w, status, map[string]any{})
34
}
35
36
func WriteAPIError(w http.ResponseWriter, status int, code, message string) {
24
- WriteAPIEnvelope(w, status, types.APIEnvelope[any]{
37
+ w.Header().Set("Content-Type", "application/json")
38
+ w.WriteHeader(status)
39
+ _ = json.NewEncoder(w).Encode(types.APIEnvelope[any]{
40
OK: false,
41
Error: &types.APIError{Code: code, Message: message},
42
})
43
}
44
30
-func DecodeAPIEnvelope[T any](r io.Reader) (types.APIEnvelope[T], error) {
31
- var envelope types.APIEnvelope[T]
32
- if err := json.NewDecoder(r).Decode(&envelope); err != nil {
33
- return types.APIEnvelope[T]{}, err
45
+func MethodNotAllowedError() APIErrorResponse {
46
+ return APIErrorResponse{
47
+ Status: http.StatusMethodNotAllowed,
48
+ Code: types.APIErrorCodeMethodNotAllowed,
49
+ Message: "method not allowed",
50
}
35
- return envelope, nil
51
}
52
38
-func NewAPIRequestError(statusCode int, apiErr *types.APIError) *types.APIRequestError {
39
- if apiErr == nil {
53
+func InvalidRequestError(err error) APIErrorResponse {
54
+ return APIErrorResponse{
55
+ Status: http.StatusBadRequest,
56
+ Code: types.APIErrorCodeInvalidRequest,
57
+ Message: err.Error(),
58
+ }
59
+}
60
+
61
+func InvalidRequestMessage(message string) APIErrorResponse {
62
+ return APIErrorResponse{
63
+ Status: http.StatusBadRequest,
64
+ Code: types.APIErrorCodeInvalidRequest,
65
+ Message: message,
66
+ }
67
+}
68
+
69
+func RequireMethod(w http.ResponseWriter, r *http.Request, method string) bool {
70
+ if r.Method == method {
71
+ return true
72
+ }
73
+ MethodNotAllowedError().Write(w)
74
+ return false
75
+}
76
+
77
+func ResolveAPIURL(baseURL *url.URL, path string) *url.URL {
78
+ ref := &url.URL{Path: path}
79
+ if baseURL == nil {
80
+ return ref
81
+ }
82
+ return baseURL.ResolveReference(ref)
83
+}
84
+
85
+func httpDo(ctx context.Context, client *http.Client, method, rawURL string, body io.Reader, headers http.Header) (*http.Response, error) {
86
+ if client == nil {
87
+ client = http.DefaultClient
88
+ }
89
+
90
+ req, err := http.NewRequestWithContext(ctx, method, rawURL, body)
91
+ if err != nil {
92
+ return nil, err
93
+ }
94
+ for key, values := range headers {
95
+ for _, value := range values {
96
+ req.Header.Add(key, value)
97
+ }
98
+ }
99
+ return client.Do(req)
100
+}
101
+
102
+func HTTPDoJSON(ctx context.Context, client *http.Client, method, rawURL string, payload any, headers http.Header, out any) error {
103
+ body, reqHeaders, err := httpJSONRequest(payload, headers)
104
+ if err != nil {
105
+ return err
106
+ }
107
+
108
+ resp, err := httpDo(ctx, client, method, rawURL, body, reqHeaders)
109
+ if err != nil {
110
+ return err
111
+ }
112
+ defer resp.Body.Close()
113
+
114
+ if out == nil {
115
+ return nil
116
+ }
117
+ return json.NewDecoder(resp.Body).Decode(out)
118
+}
119
+
120
+func HTTPDoAPIPath(ctx context.Context, client *http.Client, baseURL *url.URL, method, path string, payload any, headers http.Header, out any) error {
121
+ body, reqHeaders, err := httpJSONRequest(payload, headers)
122
+ if err != nil {
123
+ return err
124
+ }
125
+
126
+ resp, err := httpDo(ctx, client, method, ResolveAPIURL(baseURL, path).String(), body, reqHeaders)
127
+ if err != nil {
128
+ return err
129
+ }
130
+ defer resp.Body.Close()
131
+
132
+ if resp.StatusCode < http.StatusOK || resp.StatusCode >= http.StatusMultipleChoices {
133
+ return DecodeAPIRequestError(resp)
134
+ }
135
+
136
+ var envelope types.APIEnvelope[json.RawMessage]
137
+ if err := json.NewDecoder(resp.Body).Decode(&envelope); err != nil {
138
+ return fmt.Errorf("decode response: %w", err)
139
+ }
140
+ if !envelope.OK {
141
+ if envelope.Error == nil {
142
+ return &types.APIRequestError{
143
+ StatusCode: resp.StatusCode,
144
+ Message: fmt.Sprintf("api request failed with status %d", resp.StatusCode),
145
+ }
146
+ }
147
return &types.APIRequestError{
41
- StatusCode: statusCode,
42
- Message: fmt.Sprintf("api request failed with status %d", statusCode),
148
+ StatusCode: resp.StatusCode,
149
+ Code: envelope.Error.Code,
150
+ Message: envelope.Error.Message,
151
}
152
}
45
- return &types.APIRequestError{
46
- StatusCode: statusCode,
47
- Code: apiErr.Code,
48
- Message: apiErr.Message,
153
+ if out == nil {
154
+ return nil
155
}
156
+ return json.Unmarshal(envelope.Data, out)
157
}
158
159
func DecodeAPIRequestError(resp *http.Response) error {
@@ -57,7 +164,17 @@ func DecodeAPIRequestError(resp *http.Response) error {
164
body, _ := io.ReadAll(io.LimitReader(resp.Body, 8<<10))
165
var envelope types.APIEnvelope[json.RawMessage]
166
if err := json.Unmarshal(body, &envelope); err == nil && !envelope.OK {
60
- return NewAPIRequestError(resp.StatusCode, envelope.Error)
167
+ if envelope.Error == nil {
168
+ return &types.APIRequestError{
169
+ StatusCode: resp.StatusCode,
170
+ Message: fmt.Sprintf("api request failed with status %d", resp.StatusCode),
171
+ }
172
+ }
173
+ return &types.APIRequestError{
174
+ StatusCode: resp.StatusCode,
175
+ Code: envelope.Error.Code,
176
+ Message: envelope.Error.Message,
177
+ }
178
}
179
180
return &types.APIRequestError{
@@ -66,8 +183,44 @@ func DecodeAPIRequestError(resp *http.Response) error {
183
}
184
}
185
69
-func DecodeJSONBody(w http.ResponseWriter, r *http.Request, dst any, maxBytes int64) error {
186
+func DecodeJSONRequest[T any](w http.ResponseWriter, r *http.Request, maxBytes int64) (T, bool) {
187
+ var dst T
188
+ r.Body = http.MaxBytesReader(w, r.Body, maxBytes)
189
+ defer r.Body.Close()
190
+ if err := json.NewDecoder(r.Body).Decode(&dst); err != nil {
191
+ WriteAPIError(w, http.StatusBadRequest, types.APIErrorCodeInvalidJSON, err.Error())
192
+ return dst, false
193
+ }
194
+ return dst, true
195
+}
196
+
197
+func DecodeJSONRequestAs[T any](w http.ResponseWriter, r *http.Request, maxBytes int64, invalid APIErrorResponse) (T, bool) {
198
+ var dst T
199
r.Body = http.MaxBytesReader(w, r.Body, maxBytes)
200
defer r.Body.Close()
72
- return json.NewDecoder(r.Body).Decode(dst)
201
+ if err := json.NewDecoder(r.Body).Decode(&dst); err != nil {
202
+ invalid.Write(w)
203
+ return dst, false
204
+ }
205
+ return dst, true
206
+}
207
+
208
+func httpJSONRequest(payload any, headers http.Header) (io.Reader, http.Header, error) {
209
+ reqHeaders := make(http.Header, len(headers))
210
+ for key, values := range headers {
211
+ reqHeaders[key] = append([]string(nil), values...)
212
+ }
213
+
214
+ if payload == nil {
215
+ return nil, reqHeaders, nil
216
+ }
217
+
218
+ buf, err := json.Marshal(payload)
219
+ if err != nil {
220
+ return nil, nil, fmt.Errorf("marshal payload: %w", err)
221
+ }
222
+ if reqHeaders.Get("Content-Type") == "" {
223
+ reqHeaders.Set("Content-Type", "application/json")
224
+ }
225
+ return bytes.NewReader(buf), reqHeaders, nil
226
}
utils/api_test.go
+5
-5
@@ -1,7 +1,7 @@
1
package utils
2
3
import (
4
- "bytes"
4
+ "encoding/json"
5
"errors"
6
"io"
7
"net/http"
@@ -22,12 +22,12 @@ func TestWriteAPIDataAndDecodeEnvelope(t *testing.T) {
22
t.Fatalf("WriteAPIData() status = %d, want %d", rec.Code, http.StatusCreated)
23
}
24
25
- envelope, err := DecodeAPIEnvelope[map[string]string](bytes.NewReader(rec.Body.Bytes()))
26
- if err != nil {
27
- t.Fatalf("DecodeAPIEnvelope() error = %v", err)
25
+ var envelope types.APIEnvelope[map[string]string]
26
+ if err := json.NewDecoder(rec.Body).Decode(&envelope); err != nil {
27
+ t.Fatalf("json.Decode() error = %v", err)
28
}
29
if !envelope.OK || envelope.Data["status"] != "ok" {
30
- t.Fatalf("DecodeAPIEnvelope() = %+v, want ok envelope", envelope)
30
+ t.Fatalf("decoded envelope = %+v, want ok envelope", envelope)
31
}
32
}
33
utils/http.go
deleted
-105
@@ -1,105 +0,0 @@
1
-package utils
2
-
3
-import (
4
- "bytes"
5
- "context"
6
- "encoding/json"
7
- "fmt"
8
- "io"
9
- "net/http"
10
-)
11
-
12
-func HTTPDo(ctx context.Context, client *http.Client, method, rawURL string, body io.Reader, headers http.Header) (*http.Response, error) {
13
- if client == nil {
14
- client = http.DefaultClient
15
- }
16
-
17
- req, err := http.NewRequestWithContext(ctx, method, rawURL, body)
18
- if err != nil {
19
- return nil, err
20
- }
21
- for key, values := range headers {
22
- for _, value := range values {
23
- req.Header.Add(key, value)
24
- }
25
- }
26
- return client.Do(req)
27
-}
28
-
29
-func HTTPDoJSON(ctx context.Context, client *http.Client, method, rawURL string, payload any, headers http.Header, out any) error {
30
- body, reqHeaders, err := httpJSONRequest(payload, headers)
31
- if err != nil {
32
- return err
33
- }
34
-
35
- resp, err := HTTPDo(ctx, client, method, rawURL, body, reqHeaders)
36
- if err != nil {
37
- return err
38
- }
39
- defer resp.Body.Close()
40
-
41
- if out == nil {
42
- return nil
43
- }
44
- return json.NewDecoder(resp.Body).Decode(out)
45
-}
46
-
47
-func HTTPDoAPI(ctx context.Context, client *http.Client, method, rawURL string, payload any, headers http.Header, out any) error {
48
- body, reqHeaders, err := httpJSONRequest(payload, headers)
49
- if err != nil {
50
- return err
51
- }
52
-
53
- resp, err := HTTPDo(ctx, client, method, rawURL, body, reqHeaders)
54
- if err != nil {
55
- return err
56
- }
57
- defer resp.Body.Close()
58
-
59
- if resp.StatusCode < http.StatusOK || resp.StatusCode >= http.StatusMultipleChoices {
60
- return DecodeAPIRequestError(resp)
61
- }
62
-
63
- envelope, err := DecodeAPIEnvelope[json.RawMessage](resp.Body)
64
- if err != nil {
65
- return fmt.Errorf("decode response: %w", err)
66
- }
67
- if !envelope.OK {
68
- return NewAPIRequestError(resp.StatusCode, envelope.Error)
69
- }
70
- if out == nil {
71
- return nil
72
- }
73
- return json.Unmarshal(envelope.Data, out)
74
-}
75
-
76
-func HTTPReadString(resp *http.Response, limit int64) (string, error) {
77
- if resp == nil {
78
- return "", nil
79
- }
80
- body, err := io.ReadAll(io.LimitReader(resp.Body, limit))
81
- if err != nil {
82
- return "", err
83
- }
84
- return string(body), nil
85
-}
86
-
87
-func httpJSONRequest(payload any, headers http.Header) (io.Reader, http.Header, error) {
88
- reqHeaders := make(http.Header, len(headers))
89
- for key, values := range headers {
90
- reqHeaders[key] = append([]string(nil), values...)
91
- }
92
-
93
- if payload == nil {
94
- return nil, reqHeaders, nil
95
- }
96
-
97
- buf, err := json.Marshal(payload)
98
- if err != nil {
99
- return nil, nil, fmt.Errorf("marshal payload: %w", err)
100
- }
101
- if reqHeaders.Get("Content-Type") == "" {
102
- reqHeaders.Set("Content-Type", "application/json")
103
- }
104
- return bytes.NewReader(buf), reqHeaders, nil
105
-}
utils/network.go
+5
-4
@@ -4,6 +4,7 @@ import (
4
"context"
5
"encoding/json"
6
"errors"
7
+ "io"
8
"net"
9
"net/http"
10
"strings"
@@ -71,14 +72,14 @@ func resolvePublicIP(ctx context.Context, totalTimeout, attemptTimeout time.Dura
72
}
73
74
requestCtx, cancelRequest := context.WithTimeout(ctx, requestTimeout)
74
- resp, err := HTTPDo(requestCtx, client, http.MethodGet, endpoint, nil, headers)
75
+ resp, err := httpDo(requestCtx, client, http.MethodGet, endpoint, nil, headers)
76
cancelRequest()
77
if err != nil {
78
lastErr = err
79
continue
80
}
81
81
- body, readErr := HTTPReadString(resp, 256)
82
+ limitedBody, readErr := io.ReadAll(io.LimitReader(resp.Body, 256))
83
_ = resp.Body.Close()
84
if resp.StatusCode != http.StatusOK {
85
lastErr = errors.New(resp.Status)
@@ -89,7 +90,7 @@ func resolvePublicIP(ctx context.Context, totalTimeout, attemptTimeout time.Dura
90
continue
91
}
92
92
- candidate := SanitizeReportedIP(body)
93
+ candidate := SanitizeReportedIP(string(limitedBody))
94
if candidate == "" {
95
lastErr = errors.New("invalid public ip response")
96
continue
@@ -134,7 +135,7 @@ func ResolvePortalRelayURLs(ctx context.Context, explicit []string, includeDefau
135
var registry struct {
136
Relays []string `json:"relays"`
137
}
137
- resp, err := HTTPDo(ctx, client, http.MethodGet, types.PortalRelayRegistryURL, nil, nil)
138
+ resp, err := httpDo(ctx, client, http.MethodGet, types.PortalRelayRegistryURL, nil, nil)
139
if err != nil {
140
return explicit, nil
141
}