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 }