package httpapi import ( "crypto/sha256" "crypto/subtle" "log/slog" "net/http" "strings" ) // accessLevel says which credential a route requires. type accessLevel int const ( // accessOpen: no credential required (registration, heartbeat, fetching // the assignment, /healthz). accessOpen accessLevel = iota + 1 // accessAgent: requires the static agent token (validator-agents and // probers writing results/events). accessAgent // accessAdmin: requires the administrator token (/api/v1/admin/*). accessAdmin ) func (a accessLevel) String() string { switch a { case accessOpen: return "open" case accessAgent: return "agent" case accessAdmin: return "admin" } return "invalid" } // Authenticator holds the static bearer tokens. An empty token disables the // check for that level (backward compatibility: the API starts open with a // warning, see cmd/control-api). type Authenticator struct { AdminToken string AgentToken string Log *slog.Logger } // WithAuth sets the admin and agent tokens. Call it before Handler(). func (s *Server) WithAuth(admin, agent string) *Server { s.Auth = Authenticator{AdminToken: admin, AgentToken: agent} return s } func (a Authenticator) tokenFor(level accessLevel) string { switch level { case accessAdmin: return a.AdminToken case accessAgent: return a.AgentToken } return "" } // require wraps h so that it is only reached with the credential demanded by // level. Open routes, and levels with an empty configured token, pass through. func (a Authenticator) require(level accessLevel, h http.HandlerFunc) http.HandlerFunc { if level == accessOpen { return h } if level != accessAdmin && level != accessAgent { // An unclassified route must never be served. panic("httpapi: route registered without a valid access level") } return func(w http.ResponseWriter, r *http.Request) { want := a.tokenFor(level) if want == "" { h(w, r) return } got, ok := bearerToken(r) if !ok || !tokensEqual(got, want) { if a.Log != nil { a.Log.Warn("unauthorized request", "level", level.String(), "method", r.Method, "path", r.URL.Path, "remote", r.RemoteAddr) } w.Header().Set("WWW-Authenticate", "Bearer") writeError(w, http.StatusUnauthorized, "unauthorized") return } h(w, r) } } // bearerToken extracts the token from "Authorization: Bearer ". func bearerToken(r *http.Request) (string, bool) { h := r.Header.Get("Authorization") const prefix = "bearer " if len(h) <= len(prefix) || !strings.EqualFold(h[:len(prefix)], prefix) { return "", false } tok := strings.TrimSpace(h[len(prefix):]) return tok, tok != "" } // tokensEqual compares in constant time; hashing first hides the length. func tokensEqual(got, want string) bool { g := sha256.Sum256([]byte(got)) w := sha256.Sum256([]byte(want)) return subtle.ConstantTimeCompare(g[:], w[:]) == 1 }