@@ -240,7 +240,11 @@ func (s *Server) handleRenew(w http.ResponseWriter, r *http.Request) {
240240 return
241241 }
242242
243- resp , err := s .renewLease (req , clientIP )
243+ ttl := s .cfg .LeaseTTL
244+ if req .TTL > 0 {
245+ ttl = time .Duration (req .TTL ) * time .Second
246+ }
247+ record , err := s .registry .Renew (strings .TrimSpace (req .LeaseID ), req .ReverseToken , ttl , clientIP , utils .SanitizeReportedIP (req .ReportedIP ))
244248 if err != nil {
245249 status , code := http .StatusBadRequest , types .APIErrorCodeInvalidRequest
246250 if errors .Is (err , errLeaseNotFound ) {
@@ -256,7 +260,7 @@ func (s *Server) handleRenew(w http.ResponseWriter, r *http.Request) {
256260 return
257261 }
258262
259- utils .WriteAPIData (w , http .StatusOK , resp )
263+ utils .WriteAPIData (w , http .StatusOK , types. RenewResponse { LeaseID : record . ID , ExpiresAt : record . ExpiresAt } )
260264}
261265
262266func (s * Server ) handleUnregister (w http.ResponseWriter , r * http.Request ) {
@@ -271,7 +275,8 @@ func (s *Server) handleUnregister(w http.ResponseWriter, r *http.Request) {
271275 return
272276 }
273277
274- if err := s .unregisterLease (req ); err != nil {
278+ record , err := s .registry .Unregister (strings .TrimSpace (req .LeaseID ), req .ReverseToken )
279+ if err != nil {
275280 status , code := http .StatusBadRequest , types .APIErrorCodeInvalidRequest
276281 if errors .Is (err , errLeaseNotFound ) {
277282 status , code = http .StatusNotFound , types .APIErrorCodeLeaseNotFound
@@ -282,6 +287,9 @@ func (s *Server) handleUnregister(w http.ResponseWriter, r *http.Request) {
282287 utils .WriteAPIError (w , status , code , err .Error ())
283288 return
284289 }
290+ if record != nil {
291+ record .Close ()
292+ }
285293
286294 utils .WriteAPIOK (w , http .StatusOK )
287295}
@@ -304,30 +312,22 @@ func (s *Server) handleConnect(w http.ResponseWriter, r *http.Request) {
304312 return
305313 }
306314
307- lease , err := s .registry .FindByID (leaseID )
308- if err == nil && ! s .registry .policy .IsLeaseRoutable (lease .ID ) {
309- err = errLeaseRejected
310- }
311- if err == nil && ! utils .TokenMatches (lease .ReverseToken , token ) {
312- err = errUnauthorized
313- }
314- if err == nil && lease .stream == nil {
315- err = errTransportMismatch
316- }
317- switch {
318- case errors .Is (err , errLeaseNotFound ):
315+ lease , err := s .admitLeaseByID (leaseID , token , false )
316+ switch err {
317+ case nil :
318+ case errLeaseNotFound :
319319 utils .WriteAPIError (w , http .StatusNotFound , types .APIErrorCodeLeaseNotFound , err .Error ())
320320 return
321- case errors . Is ( err , errLeaseRejected ) :
321+ case errLeaseRejected :
322322 utils .WriteAPIError (w , http .StatusForbidden , types .APIErrorCodeLeaseRejected , "lease is not approved for routing" )
323323 return
324- case errors . Is ( err , errUnauthorized ) :
324+ case errUnauthorized :
325325 utils .WriteAPIError (w , http .StatusForbidden , types .APIErrorCodeUnauthorized , err .Error ())
326326 return
327- case errors . Is ( err , errTransportMismatch ) :
327+ case errTransportMismatch :
328328 utils .WriteAPIError (w , http .StatusConflict , types .APIErrorCodeTransportMismatch , "lease does not support stream transport" )
329329 return
330- case err != nil :
330+ default :
331331 utils .WriteAPIError (w , http .StatusBadRequest , types .APIErrorCodeInvalidRequest , err .Error ())
332332 return
333333 }
@@ -396,34 +396,26 @@ func (s *Server) handleQUICTunnelConn(conn *quic.Conn) {
396396 return
397397 }
398398
399- lease , err := s .registry .FindByID (msg .LeaseID )
400- if err == nil && ! s .registry .policy .IsLeaseRoutable (lease .ID ) {
401- err = errLeaseRejected
402- }
403- if err == nil && ! utils .TokenMatches (lease .ReverseToken , msg .ReverseToken ) {
404- err = errUnauthorized
405- }
406- if err == nil && (lease .stream == nil || lease .datagram == nil ) {
407- err = errTransportMismatch
408- }
409- switch {
410- case errors .Is (err , errLeaseNotFound ):
399+ lease , err := s .admitLeaseByID (msg .LeaseID , msg .ReverseToken , true )
400+ switch err {
401+ case nil :
402+ case errLeaseNotFound :
411403 _ = json .NewEncoder (stream ).Encode (types.QUICControlResponse {OK : false , Error : types .APIErrorCodeLeaseNotFound })
412404 _ = conn .CloseWithError (1 , "lease not found" )
413405 return
414- case errors .Is (err , errUnauthorized ):
415- _ = json .NewEncoder (stream ).Encode (types.QUICControlResponse {OK : false , Error : types .APIErrorCodeUnauthorized })
416- _ = conn .CloseWithError (1 , "unauthorized" )
417- return
418- case errors .Is (err , errLeaseRejected ):
406+ case errLeaseRejected :
419407 _ = json .NewEncoder (stream ).Encode (types.QUICControlResponse {OK : false , Error : types .APIErrorCodeLeaseRejected })
420408 _ = conn .CloseWithError (1 , "lease rejected" )
421409 return
422- case errors .Is (err , errTransportMismatch ):
410+ case errUnauthorized :
411+ _ = json .NewEncoder (stream ).Encode (types.QUICControlResponse {OK : false , Error : types .APIErrorCodeUnauthorized })
412+ _ = conn .CloseWithError (1 , "unauthorized" )
413+ return
414+ case errTransportMismatch :
423415 _ = json .NewEncoder (stream ).Encode (types.QUICControlResponse {OK : false , Error : types .APIErrorCodeTransportMismatch })
424416 _ = conn .CloseWithError (1 , "transport mismatch" )
425417 return
426- case err != nil :
418+ default :
427419 _ = json .NewEncoder (stream ).Encode (types.QUICControlResponse {OK : false , Error : types .APIErrorCodeInvalidRequest })
428420 _ = conn .CloseWithError (1 , "invalid control message" )
429421 return
@@ -445,6 +437,23 @@ func (s *Server) handleQUICTunnelConn(conn *quic.Conn) {
445437 Msg ("quic tunnel connected" )
446438}
447439
440+ func (s * Server ) admitLeaseByID (leaseID , token string , requireDatagram bool ) (* leaseRecord , error ) {
441+ lease , err := s .registry .FindByID (leaseID )
442+ if err != nil {
443+ return nil , err
444+ }
445+ if ! s .registry .policy .IsLeaseRoutable (lease .ID ) {
446+ return nil , errLeaseRejected
447+ }
448+ if ! utils .TokenMatches (lease .ReverseToken , token ) {
449+ return nil , errUnauthorized
450+ }
451+ if lease .stream == nil || (requireDatagram && lease .datagram == nil ) {
452+ return nil , errTransportMismatch
453+ }
454+ return lease , nil
455+ }
456+
448457func (s * Server ) registerLease (req types.RegisterRequest , clientIP string ) (types.RegisterResponse , error ) {
449458 name , err := utils .NormalizeDNSLabel (req .Name )
450459 if err != nil {
@@ -541,34 +550,6 @@ func (s *Server) registerLease(req types.RegisterRequest, clientIP string) (type
541550 return resp , nil
542551}
543552
544- func (s * Server ) renewLease (req types.RenewRequest , clientIP string ) (types.RenewResponse , error ) {
545- if s .registry .policy .IPFilter ().IsIPBanned (clientIP ) {
546- return types.RenewResponse {}, errIPBanned
547- }
548-
549- ttl := s .cfg .LeaseTTL
550- if req .TTL > 0 {
551- ttl = time .Duration (req .TTL ) * time .Second
552- }
553- record , err := s .registry .Renew (strings .TrimSpace (req .LeaseID ), req .ReverseToken , ttl , clientIP , utils .SanitizeReportedIP (req .ReportedIP ))
554- if err != nil {
555- return types.RenewResponse {}, err
556- }
557-
558- return types.RenewResponse {LeaseID : record .ID , ExpiresAt : record .ExpiresAt }, nil
559- }
560-
561- func (s * Server ) unregisterLease (req types.UnregisterRequest ) error {
562- record , err := s .registry .Unregister (strings .TrimSpace (req .LeaseID ), req .ReverseToken )
563- if err != nil {
564- return err
565- }
566- if record != nil {
567- record .Close ()
568- }
569- return nil
570- }
571-
572553func (s * Server ) runAPIServer () error {
573554 err := s .apiServer .Serve (s .apiListener )
574555 if err == nil || errors .Is (err , http .ErrServerClosed ) || errors .Is (err , net .ErrClosed ) {
0 commit comments