Files
cloud-ip-validator/internal/httpapi/auth.go
T

108 lines
2.8 KiB
Go
Raw Normal View History

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 <token>".
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
}