Add authentication: admin/agent bearer tokens for the API, login for the dashboard
control-api: every route now carries a mandatory access level (admin / agent / open) in a route table. All /api/v1/admin/* require the admin token; the write calls of validator-agent and prober (self-check, events, results, complete) require a separate static agent token; register, heartbeat and fetching the assignment stay open. Tokens come from env vars, are compared in constant time and never logged. An empty token leaves that level open with a startup warning (backward compatible). validator-agent / prober: apiclient sends the agent token only to control-api. admin-dashboard: login/password (from env) with a stateless HMAC session cookie, Origin-based CSRF check, per-IP brute-force throttle, HX-Redirect for htmx polls, logout in the sidebar; the dashboard calls control-api with the admin token. Login page layout fixed after review. Also: env plumbing in docker-compose/rxprod-compose/systemd/config examples, e2e script with token assertions, tests, docs (API, SETUP, USAGE, DASHBOARD, README), plan and review under docs/changes/, bin/ rebuilt with new SHA256SUMS. Co-Authored-By: Claude Sonnet 5.5 <noreply@anthropic.com>
This commit is contained in:
1 parent
972ad47d0c
commit
debf2afed2
67 files changed
+2050
-105
No files matched your search
@@ -53,6 +53,13 @@ func New(cfg *config.ValidatorAgent, log *slog.Logger) *Agent {
|
||||
}
|
||||
}
|
||||
|
||||
// WithToken sets the bearer token sent to the Control API (and only to it:
|
||||
// the IP-echo lookup and all check traffic use separate clients).
|
||||
func (a *Agent) WithToken(token string) *Agent {
|
||||
a.client.Token = token
|
||||
return a
|
||||
}
|
||||
|
||||
// Run registers with the Control API and polls forever until ctx is
|
||||
// cancelled.
|
||||
func (a *Agent) Run(ctx context.Context) error {
|
||||
|
||||
@@ -98,3 +98,34 @@ func TestRegisterWithRetryStopsOnCancel(t *testing.T) {
|
||||
t.Fatalf("expected context.Canceled, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestTokenGoesOnlyToControlAPI: the bearer token authenticates calls to
|
||||
// control-api, and must never leak to the external IP-echo service.
|
||||
func TestTokenGoesOnlyToControlAPI(t *testing.T) {
|
||||
var apiAuth, echoAuth atomic.Value
|
||||
api := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
apiAuth.Store(r.Header.Get("Authorization"))
|
||||
w.Write([]byte(`{"ok":true,"poll_interval_seconds":5}`))
|
||||
}))
|
||||
defer api.Close()
|
||||
echo := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
echoAuth.Store(r.Header.Get("Authorization"))
|
||||
w.Write([]byte("203.0.113.7"))
|
||||
}))
|
||||
defer echo.Close()
|
||||
|
||||
a := New(&config.ValidatorAgent{ValidatorID: "val-1", ControlAPIURL: api.URL}, testLogger()).WithToken("agent-secret")
|
||||
ctx := context.Background()
|
||||
if err := a.registerWithRetry(ctx); err != nil {
|
||||
t.Fatalf("register: %v", err)
|
||||
}
|
||||
if _, err := fetchIPEcho(ctx, echo.URL); err != nil {
|
||||
t.Fatalf("fetchIPEcho: %v", err)
|
||||
}
|
||||
if got, _ := apiAuth.Load().(string); got != "Bearer agent-secret" {
|
||||
t.Fatalf("control-api Authorization = %q, want bearer token", got)
|
||||
}
|
||||
if got, _ := echoAuth.Load().(string); got != "" {
|
||||
t.Fatalf("IP-echo request carried Authorization %q, want none", got)
|
||||
}
|
||||
}
|
||||
@@ -17,6 +17,9 @@ import (
|
||||
type Client struct {
|
||||
BaseURL string
|
||||
HTTPClient *http.Client
|
||||
// Token, when non-empty, is sent as "Authorization: Bearer <Token>" on
|
||||
// every request to the Control API.
|
||||
Token string
|
||||
}
|
||||
|
||||
func New(baseURL string, timeout time.Duration) *Client {
|
||||
@@ -42,6 +45,9 @@ func (c *Client) Do(ctx context.Context, method, path string, body, out interfac
|
||||
if body != nil {
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
}
|
||||
if c.Token != "" {
|
||||
req.Header.Set("Authorization", "Bearer "+c.Token)
|
||||
}
|
||||
|
||||
resp, err := c.HTTPClient.Do(req)
|
||||
if err != nil {
|
||||
|
||||
@@ -0,0 +1,30 @@
|
||||
package apiclient
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestDoSetsBearerOnlyWhenTokenSet(t *testing.T) {
|
||||
var got []string
|
||||
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
got = append(got, r.Header.Get("Authorization"))
|
||||
w.WriteHeader(http.StatusNoContent)
|
||||
}))
|
||||
defer ts.Close()
|
||||
|
||||
c := New(ts.URL, 5*time.Second)
|
||||
if _, err := c.Do(context.Background(), http.MethodGet, "/x", nil, nil); err != nil {
|
||||
t.Fatalf("do without token: %v", err)
|
||||
}
|
||||
c.Token = "secret-agent-token"
|
||||
if _, err := c.Do(context.Background(), http.MethodPost, "/x", map[string]string{"a": "b"}, nil); err != nil {
|
||||
t.Fatalf("do with token: %v", err)
|
||||
}
|
||||
if len(got) != 2 || got[0] != "" || got[1] != "Bearer secret-agent-token" {
|
||||
t.Fatalf("Authorization headers = %q, want [\"\" \"Bearer secret-agent-token\"]", got)
|
||||
}
|
||||
}
|
||||
@@ -25,6 +25,15 @@ type ControlAPI struct {
|
||||
Targets map[string][]string `yaml:"targets"`
|
||||
Inbound InboundConfig `yaml:"inbound_checks"`
|
||||
IPAddresses []string `yaml:"ip_addresses"`
|
||||
Auth ControlAPIAuth `yaml:"auth"`
|
||||
}
|
||||
|
||||
// ControlAPIAuth names the environment variables control-api reads its two
|
||||
// static bearer tokens from. Values are never stored in the config. An unset
|
||||
// (empty) token leaves that access level open, with a startup warning.
|
||||
type ControlAPIAuth struct {
|
||||
AdminTokenEnv string `yaml:"admin_token_env"` // default CONTROL_API_ADMIN_TOKEN — protects /api/v1/admin/*
|
||||
AgentTokenEnv string `yaml:"agent_token_env"` // default CONTROL_API_AGENT_TOKEN — protects agent/prober write calls
|
||||
}
|
||||
|
||||
type ServerConfig struct {
|
||||
@@ -155,6 +164,12 @@ func LoadControlAPI(path string) (*ControlAPI, error) {
|
||||
if c.OpenStack.PasswordEnv == "" {
|
||||
c.OpenStack.PasswordEnv = "OS_PASSWORD"
|
||||
}
|
||||
if c.Auth.AdminTokenEnv == "" {
|
||||
c.Auth.AdminTokenEnv = "CONTROL_API_ADMIN_TOKEN"
|
||||
}
|
||||
if c.Auth.AgentTokenEnv == "" {
|
||||
c.Auth.AgentTokenEnv = "CONTROL_API_AGENT_TOKEN"
|
||||
}
|
||||
if c.Orchestrator.PollIntervalSeconds == 0 {
|
||||
c.Orchestrator.PollIntervalSeconds = 5
|
||||
}
|
||||
@@ -182,8 +197,11 @@ func LoadControlAPI(path string) (*ControlAPI, error) {
|
||||
// ---- validator-agent ----
|
||||
|
||||
type ValidatorAgent struct {
|
||||
ValidatorID string `yaml:"validator_id"`
|
||||
ControlAPIURL string `yaml:"control_api_url"`
|
||||
ValidatorID string `yaml:"validator_id"`
|
||||
ControlAPIURL string `yaml:"control_api_url"`
|
||||
// ControlAPITokenEnv is the name of the env var holding the agent bearer
|
||||
// token (default CONTROL_API_AGENT_TOKEN). Empty value = no token sent.
|
||||
ControlAPITokenEnv string `yaml:"control_api_token_env"`
|
||||
PollIntervalSeconds int `yaml:"poll_interval_seconds"`
|
||||
SelfCheck SelfCheckCfg `yaml:"self_check"`
|
||||
Checks AgentChecks `yaml:"checks"`
|
||||
@@ -223,6 +241,9 @@ func LoadValidatorAgent(path string) (*ValidatorAgent, error) {
|
||||
if c.PollIntervalSeconds == 0 {
|
||||
c.PollIntervalSeconds = 5
|
||||
}
|
||||
if c.ControlAPITokenEnv == "" {
|
||||
c.ControlAPITokenEnv = "CONTROL_API_AGENT_TOKEN"
|
||||
}
|
||||
if c.SelfCheck.TimeoutSeconds == 0 {
|
||||
c.SelfCheck.TimeoutSeconds = 10
|
||||
}
|
||||
@@ -253,8 +274,11 @@ func LoadValidatorAgent(path string) (*ValidatorAgent, error) {
|
||||
// ---- prober ----
|
||||
|
||||
type Prober struct {
|
||||
SiteID string `yaml:"site_id"`
|
||||
ControlAPIURL string `yaml:"control_api_url"`
|
||||
SiteID string `yaml:"site_id"`
|
||||
ControlAPIURL string `yaml:"control_api_url"`
|
||||
// ControlAPITokenEnv is the name of the env var holding the agent bearer
|
||||
// token (default CONTROL_API_AGENT_TOKEN). Empty value = no token sent.
|
||||
ControlAPITokenEnv string `yaml:"control_api_token_env"`
|
||||
PollIntervalSeconds int `yaml:"poll_interval_seconds"`
|
||||
Checks ProberChecks `yaml:"checks"`
|
||||
}
|
||||
@@ -273,6 +297,9 @@ func LoadProber(path string) (*Prober, error) {
|
||||
if c.PollIntervalSeconds == 0 {
|
||||
c.PollIntervalSeconds = 5
|
||||
}
|
||||
if c.ControlAPITokenEnv == "" {
|
||||
c.ControlAPITokenEnv = "CONTROL_API_AGENT_TOKEN"
|
||||
}
|
||||
if c.Checks.TCPTimeoutSeconds == 0 {
|
||||
c.Checks.TCPTimeoutSeconds = 5
|
||||
}
|
||||
@@ -300,11 +327,25 @@ type AdminDashboard struct {
|
||||
Server ServerConfig `yaml:"server"`
|
||||
ControlAPI DashboardControlAPIConfig `yaml:"control_api"`
|
||||
Overview DashboardOverviewConfig `yaml:"overview"`
|
||||
Auth DashboardAuthConfig `yaml:"auth"`
|
||||
}
|
||||
|
||||
// DashboardAuthConfig names the env vars holding the single administrator's
|
||||
// login, password and the session-cookie HMAC key. If username or password
|
||||
// is empty at runtime, login is not required (with a startup warning).
|
||||
type DashboardAuthConfig struct {
|
||||
UsernameEnv string `yaml:"username_env"` // default ADMIN_DASHBOARD_USERNAME
|
||||
PasswordEnv string `yaml:"password_env"` // default ADMIN_DASHBOARD_PASSWORD
|
||||
SessionSecretEnv string `yaml:"session_secret_env"` // default ADMIN_DASHBOARD_SESSION_SECRET
|
||||
SessionTTLMinutes int `yaml:"session_ttl_minutes"` // default 480
|
||||
}
|
||||
|
||||
type DashboardControlAPIConfig struct {
|
||||
BaseURL string `yaml:"base_url"`
|
||||
TimeoutSeconds int `yaml:"timeout_seconds"`
|
||||
// TokenEnv is the name of the env var holding control-api's admin bearer
|
||||
// token (default ADMIN_DASHBOARD_CONTROL_API_TOKEN).
|
||||
TokenEnv string `yaml:"token_env"`
|
||||
}
|
||||
|
||||
// DashboardOverviewConfig configures the overview page's "текущая
|
||||
@@ -331,6 +372,21 @@ func LoadAdminDashboard(path string) (*AdminDashboard, error) {
|
||||
if c.ControlAPI.TimeoutSeconds == 0 {
|
||||
c.ControlAPI.TimeoutSeconds = 10
|
||||
}
|
||||
if c.ControlAPI.TokenEnv == "" {
|
||||
c.ControlAPI.TokenEnv = "ADMIN_DASHBOARD_CONTROL_API_TOKEN"
|
||||
}
|
||||
if c.Auth.UsernameEnv == "" {
|
||||
c.Auth.UsernameEnv = "ADMIN_DASHBOARD_USERNAME"
|
||||
}
|
||||
if c.Auth.PasswordEnv == "" {
|
||||
c.Auth.PasswordEnv = "ADMIN_DASHBOARD_PASSWORD"
|
||||
}
|
||||
if c.Auth.SessionSecretEnv == "" {
|
||||
c.Auth.SessionSecretEnv = "ADMIN_DASHBOARD_SESSION_SECRET"
|
||||
}
|
||||
if c.Auth.SessionTTLMinutes == 0 {
|
||||
c.Auth.SessionTTLMinutes = 480
|
||||
}
|
||||
if c.Overview.LastCompletedCount == 0 {
|
||||
c.Overview.LastCompletedCount = 20
|
||||
}
|
||||
|
||||
@@ -0,0 +1,392 @@
|
||||
package dashboard
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/hmac"
|
||||
"crypto/rand"
|
||||
"crypto/sha256"
|
||||
"crypto/subtle"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"log/slog"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"reflect"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
const (
|
||||
sessionCookieName = "session"
|
||||
defaultSessionTTL = 8 * time.Hour
|
||||
|
||||
// Brute-force throttle: maxLoginFailures failures from one client IP
|
||||
// within loginFailureWindow lock that IP out until the oldest failure
|
||||
// leaves the window.
|
||||
maxLoginFailures = 5
|
||||
loginFailureWindow = 10 * time.Minute
|
||||
)
|
||||
|
||||
// authState holds everything the login/session machinery needs. A nil-safe
|
||||
// zero value is never used: Server.auth is always set by New, with
|
||||
// enabled=false when no credentials were configured.
|
||||
type authState struct {
|
||||
enabled bool
|
||||
user [sha256.Size]byte // sha256(username)
|
||||
pass [sha256.Size]byte // sha256(password)
|
||||
key []byte // HMAC key for session cookies
|
||||
ttl time.Duration
|
||||
now func() time.Time
|
||||
throttle *loginThrottle
|
||||
}
|
||||
|
||||
func newAuthState(cfg Config, log *slog.Logger) *authState {
|
||||
a := &authState{now: time.Now, throttle: newLoginThrottle(), ttl: cfg.SessionTTL}
|
||||
if a.ttl <= 0 {
|
||||
a.ttl = defaultSessionTTL
|
||||
}
|
||||
if cfg.Username == "" || cfg.Password == "" {
|
||||
log.Warn("dashboard login is disabled: username/password are not set, anyone who can reach this port has full access")
|
||||
return a
|
||||
}
|
||||
a.enabled = true
|
||||
a.user = sha256.Sum256([]byte(cfg.Username))
|
||||
a.pass = sha256.Sum256([]byte(cfg.Password))
|
||||
if cfg.SessionSecret != "" {
|
||||
a.key = []byte(cfg.SessionSecret)
|
||||
} else {
|
||||
a.key = make([]byte, 32)
|
||||
if _, err := rand.Read(a.key); err != nil {
|
||||
panic("dashboard: crypto/rand failed: " + err.Error())
|
||||
}
|
||||
log.Warn("dashboard session secret is not set: using a random one, sessions are reset on every restart")
|
||||
}
|
||||
return a
|
||||
}
|
||||
|
||||
// checkCredentials compares both fields in constant time and always
|
||||
// evaluates both, so neither the length nor which field was wrong leaks.
|
||||
func (a *authState) checkCredentials(user, pass string) bool {
|
||||
u := sha256.Sum256([]byte(user))
|
||||
p := sha256.Sum256([]byte(pass))
|
||||
uOK := subtle.ConstantTimeCompare(u[:], a.user[:])
|
||||
pOK := subtle.ConstantTimeCompare(p[:], a.pass[:])
|
||||
return uOK&pOK == 1
|
||||
}
|
||||
|
||||
// ---- session cookie ----
|
||||
|
||||
type sessionPayload struct {
|
||||
User string `json:"u"`
|
||||
Exp int64 `json:"exp"`
|
||||
}
|
||||
|
||||
func (a *authState) sign(msg string) []byte {
|
||||
m := hmac.New(sha256.New, a.key)
|
||||
m.Write([]byte(msg))
|
||||
return m.Sum(nil)
|
||||
}
|
||||
|
||||
func (a *authState) issue(user string) (value string, exp time.Time) {
|
||||
exp = a.now().Add(a.ttl)
|
||||
raw, _ := json.Marshal(sessionPayload{User: user, Exp: exp.Unix()})
|
||||
p := base64.RawURLEncoding.EncodeToString(raw)
|
||||
return p + "." + base64.RawURLEncoding.EncodeToString(a.sign(p)), exp
|
||||
}
|
||||
|
||||
// verify returns the session's user if value is an untampered, unexpired
|
||||
// cookie issued with this server's key.
|
||||
func (a *authState) verify(value string) (string, bool) {
|
||||
p, sig, ok := strings.Cut(value, ".")
|
||||
if !ok {
|
||||
return "", false
|
||||
}
|
||||
got, err := base64.RawURLEncoding.DecodeString(sig)
|
||||
if err != nil || !hmac.Equal(got, a.sign(p)) {
|
||||
return "", false
|
||||
}
|
||||
raw, err := base64.RawURLEncoding.DecodeString(p)
|
||||
if err != nil {
|
||||
return "", false
|
||||
}
|
||||
var sp sessionPayload
|
||||
if json.Unmarshal(raw, &sp) != nil || sp.User == "" || a.now().Unix() >= sp.Exp {
|
||||
return "", false
|
||||
}
|
||||
return sp.User, true
|
||||
}
|
||||
|
||||
func isHTTPS(r *http.Request) bool {
|
||||
return r.TLS != nil || strings.EqualFold(r.Header.Get("X-Forwarded-Proto"), "https")
|
||||
}
|
||||
|
||||
func (a *authState) setCookie(w http.ResponseWriter, r *http.Request, user string) {
|
||||
value, exp := a.issue(user)
|
||||
http.SetCookie(w, &http.Cookie{
|
||||
Name: sessionCookieName, Value: value, Path: "/", Expires: exp,
|
||||
MaxAge: int(a.ttl.Seconds()), HttpOnly: true, Secure: isHTTPS(r), SameSite: http.SameSiteStrictMode,
|
||||
})
|
||||
}
|
||||
|
||||
func clearCookie(w http.ResponseWriter, r *http.Request) {
|
||||
http.SetCookie(w, &http.Cookie{
|
||||
Name: sessionCookieName, Value: "", Path: "/", MaxAge: -1,
|
||||
HttpOnly: true, Secure: isHTTPS(r), SameSite: http.SameSiteStrictMode,
|
||||
})
|
||||
}
|
||||
|
||||
// ---- request context ----
|
||||
|
||||
type ctxKey struct{}
|
||||
|
||||
func userFromRequest(r *http.Request) string {
|
||||
u, _ := r.Context().Value(ctxKey{}).(string)
|
||||
return u
|
||||
}
|
||||
|
||||
// ---- middleware ----
|
||||
|
||||
// authMiddleware gates every route except the login page, logout and static
|
||||
// assets. With login disabled it is a pass-through (no CSRF check either).
|
||||
func (s *Server) authMiddleware(next http.Handler) http.Handler {
|
||||
a := s.auth
|
||||
if !a.enabled {
|
||||
return next
|
||||
}
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if !safeMethod(r.Method) && !sameOrigin(r) {
|
||||
s.Log.Warn("rejected cross-origin request", "method", r.Method, "path", r.URL.Path)
|
||||
http.Error(w, "forbidden: cross-origin request", http.StatusForbidden)
|
||||
return
|
||||
}
|
||||
if isOpenPath(r) {
|
||||
next.ServeHTTP(w, r)
|
||||
return
|
||||
}
|
||||
if c, err := r.Cookie(sessionCookieName); err == nil {
|
||||
if user, ok := a.verify(c.Value); ok {
|
||||
next.ServeHTTP(w, r.WithContext(context.WithValue(r.Context(), ctxKey{}, user)))
|
||||
return
|
||||
}
|
||||
}
|
||||
// htmx polls fragments in the background: answer with an
|
||||
// HX-Redirect instead of a redirect whose login page would be
|
||||
// swapped into the fragment.
|
||||
if r.Header.Get("HX-Request") == "true" {
|
||||
w.Header().Set("HX-Redirect", "/login")
|
||||
http.Error(w, "unauthorized", http.StatusUnauthorized)
|
||||
return
|
||||
}
|
||||
target := "/login"
|
||||
if r.Method == http.MethodGet {
|
||||
if n := r.URL.RequestURI(); n != "/" {
|
||||
target += "?next=" + url.QueryEscape(n)
|
||||
}
|
||||
}
|
||||
http.Redirect(w, r, target, http.StatusSeeOther)
|
||||
})
|
||||
}
|
||||
|
||||
func isOpenPath(r *http.Request) bool {
|
||||
switch r.URL.Path {
|
||||
case "/login":
|
||||
return r.Method == http.MethodGet || r.Method == http.MethodPost
|
||||
case "/logout":
|
||||
return r.Method == http.MethodPost
|
||||
}
|
||||
return strings.HasPrefix(r.URL.Path, "/static/") && (r.Method == http.MethodGet || r.Method == http.MethodHead)
|
||||
}
|
||||
|
||||
func safeMethod(m string) bool {
|
||||
return m == http.MethodGet || m == http.MethodHead || m == http.MethodOptions
|
||||
}
|
||||
|
||||
// sameOrigin implements the CSRF check: Origin (or, when absent, Referer)
|
||||
// must name this very host. Browsers always send Origin on cross-site
|
||||
// POSTs; a request with neither header is refused.
|
||||
func sameOrigin(r *http.Request) bool {
|
||||
h := r.Header.Get("Origin")
|
||||
if h == "" {
|
||||
h = r.Header.Get("Referer")
|
||||
}
|
||||
if h == "" || h == "null" {
|
||||
return false
|
||||
}
|
||||
u, err := url.Parse(h)
|
||||
if err != nil || u.Host == "" {
|
||||
return false
|
||||
}
|
||||
return strings.EqualFold(u.Host, r.Host)
|
||||
}
|
||||
|
||||
// safeNext returns next only if it is a same-origin relative path, so the
|
||||
// post-login redirect can never leave the site.
|
||||
func safeNext(next string) string {
|
||||
if next == "" || next[0] != '/' || strings.HasPrefix(next, "//") || strings.HasPrefix(next, "/\\") {
|
||||
return "/"
|
||||
}
|
||||
for _, c := range next {
|
||||
if c < 0x20 || c == 0x7f || c == '\\' {
|
||||
return "/"
|
||||
}
|
||||
}
|
||||
u, err := url.Parse(next)
|
||||
if err != nil || u.Scheme != "" || u.Host != "" {
|
||||
return "/"
|
||||
}
|
||||
return next
|
||||
}
|
||||
|
||||
// ---- brute-force throttle ----
|
||||
|
||||
type loginThrottle struct {
|
||||
mu sync.Mutex
|
||||
failures map[string][]time.Time
|
||||
}
|
||||
|
||||
func newLoginThrottle() *loginThrottle {
|
||||
return &loginThrottle{failures: map[string][]time.Time{}}
|
||||
}
|
||||
|
||||
// prune drops expired failures (all clients) — caller holds mu. The map is
|
||||
// bounded by the number of distinct IPs that failed in the last window.
|
||||
func (t *loginThrottle) prune(now time.Time) {
|
||||
for ip, fs := range t.failures {
|
||||
i := 0
|
||||
for i < len(fs) && now.Sub(fs[i]) >= loginFailureWindow {
|
||||
i++
|
||||
}
|
||||
if i == len(fs) {
|
||||
delete(t.failures, ip)
|
||||
} else if i > 0 {
|
||||
t.failures[ip] = fs[i:]
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// blocked reports whether ip is locked out, and for how long.
|
||||
func (t *loginThrottle) blocked(ip string, now time.Time) (bool, time.Duration) {
|
||||
t.mu.Lock()
|
||||
defer t.mu.Unlock()
|
||||
t.prune(now)
|
||||
fs := t.failures[ip]
|
||||
if len(fs) < maxLoginFailures {
|
||||
return false, 0
|
||||
}
|
||||
return true, fs[0].Add(loginFailureWindow).Sub(now)
|
||||
}
|
||||
|
||||
func (t *loginThrottle) fail(ip string, now time.Time) {
|
||||
t.mu.Lock()
|
||||
defer t.mu.Unlock()
|
||||
t.prune(now)
|
||||
t.failures[ip] = append(t.failures[ip], now)
|
||||
}
|
||||
|
||||
func (t *loginThrottle) clear(ip string) {
|
||||
t.mu.Lock()
|
||||
defer t.mu.Unlock()
|
||||
delete(t.failures, ip)
|
||||
}
|
||||
|
||||
func clientIP(r *http.Request) string {
|
||||
host, _, err := net.SplitHostPort(r.RemoteAddr)
|
||||
if err != nil {
|
||||
return r.RemoteAddr
|
||||
}
|
||||
return host
|
||||
}
|
||||
|
||||
// ---- handlers ----
|
||||
|
||||
type loginPageData struct {
|
||||
Error string
|
||||
Next string
|
||||
}
|
||||
|
||||
func (s *Server) renderLogin(w http.ResponseWriter, status int, data loginPageData) {
|
||||
w.Header().Set("Content-Type", "text/html; charset=utf-8")
|
||||
w.Header().Set("Cache-Control", "no-store")
|
||||
w.WriteHeader(status)
|
||||
if err := s.tmpl.ExecuteTemplate(w, "login_page", data); err != nil {
|
||||
s.Log.Error("render login", "err", err)
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Server) handleLoginPage(w http.ResponseWriter, r *http.Request) {
|
||||
if !s.auth.enabled {
|
||||
http.Redirect(w, r, "/", http.StatusSeeOther)
|
||||
return
|
||||
}
|
||||
if c, err := r.Cookie(sessionCookieName); err == nil {
|
||||
if _, ok := s.auth.verify(c.Value); ok {
|
||||
http.Redirect(w, r, safeNext(r.URL.Query().Get("next")), http.StatusSeeOther)
|
||||
return
|
||||
}
|
||||
}
|
||||
s.renderLogin(w, http.StatusOK, loginPageData{Next: safeNext(r.URL.Query().Get("next"))})
|
||||
}
|
||||
|
||||
func (s *Server) handleLoginSubmit(w http.ResponseWriter, r *http.Request) {
|
||||
a := s.auth
|
||||
if !a.enabled {
|
||||
http.Redirect(w, r, "/", http.StatusSeeOther)
|
||||
return
|
||||
}
|
||||
r.Body = http.MaxBytesReader(w, r.Body, 4096)
|
||||
if err := r.ParseForm(); err != nil {
|
||||
http.Error(w, "bad request", http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
next := safeNext(r.PostForm.Get("next"))
|
||||
ip := clientIP(r)
|
||||
now := a.now()
|
||||
|
||||
if blocked, wait := a.throttle.blocked(ip, now); blocked {
|
||||
secs := int(wait.Seconds()) + 1
|
||||
w.Header().Set("Retry-After", strconv.Itoa(secs))
|
||||
s.Log.Warn("login throttled", "remote", ip)
|
||||
s.renderLogin(w, http.StatusTooManyRequests, loginPageData{Next: next, Error: "Слишком много попыток входа. Повторите позже."})
|
||||
return
|
||||
}
|
||||
user := r.PostForm.Get("username")
|
||||
if !a.checkCredentials(user, r.PostForm.Get("password")) {
|
||||
a.throttle.fail(ip, now)
|
||||
s.Log.Warn("login failed", "remote", ip)
|
||||
s.renderLogin(w, http.StatusOK, loginPageData{Next: next, Error: "Неверный логин или пароль"})
|
||||
return
|
||||
}
|
||||
a.throttle.clear(ip)
|
||||
a.setCookie(w, r, user)
|
||||
http.Redirect(w, r, next, http.StatusSeeOther)
|
||||
}
|
||||
|
||||
func (s *Server) handleLogout(w http.ResponseWriter, r *http.Request) {
|
||||
clearCookie(w, r)
|
||||
http.Redirect(w, r, "/login", http.StatusSeeOther)
|
||||
}
|
||||
|
||||
// withAuthInfo returns a copy of the page-data struct data with its embedded
|
||||
// PageData's AuthEnabled/User filled in, so the sidebar can show the logout
|
||||
// control. Data without an embedded PageData is returned unchanged.
|
||||
func (s *Server) withAuthInfo(r *http.Request, data interface{}) interface{} {
|
||||
if !s.auth.enabled || data == nil {
|
||||
return data
|
||||
}
|
||||
v := reflect.ValueOf(data)
|
||||
if v.Kind() != reflect.Struct {
|
||||
return data
|
||||
}
|
||||
cp := reflect.New(v.Type()).Elem()
|
||||
cp.Set(v)
|
||||
pd := cp.FieldByName("PageData")
|
||||
if !pd.IsValid() || pd.Type() != reflect.TypeOf(PageData{}) {
|
||||
return data
|
||||
}
|
||||
pd.FieldByName("AuthEnabled").SetBool(true)
|
||||
pd.FieldByName("User").SetString(userFromRequest(r))
|
||||
return cp.Interface()
|
||||
}
|
||||
@@ -0,0 +1,422 @@
|
||||
package dashboard
|
||||
|
||||
import (
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"net/url"
|
||||
"os"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
const (
|
||||
testUser = "admin"
|
||||
testPass = "correct horse battery staple"
|
||||
testSecret = "test-session-secret"
|
||||
)
|
||||
|
||||
func newAuthTestServer(t *testing.T, caURL string, mutate func(*Config)) (*Server, *httptest.Server) {
|
||||
t.Helper()
|
||||
cfg := Config{
|
||||
ControlAPIBaseURL: caURL,
|
||||
ControlAPITimeout: 5 * time.Second,
|
||||
LastCompletedCount: 20,
|
||||
OverviewPollIntervalS: 5,
|
||||
ControlAPIToken: "ca-admin-token",
|
||||
Username: testUser,
|
||||
Password: testPass,
|
||||
SessionSecret: testSecret,
|
||||
SessionTTL: time.Hour,
|
||||
}
|
||||
if mutate != nil {
|
||||
mutate(&cfg)
|
||||
}
|
||||
log := slog.New(slog.NewTextHandler(os.Stderr, &slog.HandlerOptions{Level: slog.LevelError}))
|
||||
srv, err := New(cfg, log)
|
||||
if err != nil {
|
||||
t.Fatalf("new dashboard server: %v", err)
|
||||
}
|
||||
ts := httptest.NewServer(srv.Handler())
|
||||
t.Cleanup(ts.Close)
|
||||
return srv, ts
|
||||
}
|
||||
|
||||
// noFollow is a client that returns redirects as-is.
|
||||
func noFollow(ts *httptest.Server) *http.Client {
|
||||
c := *ts.Client()
|
||||
c.CheckRedirect = func(*http.Request, []*http.Request) error { return http.ErrUseLastResponse }
|
||||
return &c
|
||||
}
|
||||
|
||||
type reqOpts struct {
|
||||
method string
|
||||
path string
|
||||
form url.Values
|
||||
cookie string
|
||||
headers map[string]string
|
||||
}
|
||||
|
||||
func doReq(t *testing.T, ts *httptest.Server, o reqOpts) (*http.Response, string) {
|
||||
t.Helper()
|
||||
if o.method == "" {
|
||||
o.method = http.MethodGet
|
||||
}
|
||||
var body io.Reader
|
||||
if o.form != nil {
|
||||
body = strings.NewReader(o.form.Encode())
|
||||
}
|
||||
req, err := http.NewRequest(o.method, ts.URL+o.path, body)
|
||||
if err != nil {
|
||||
t.Fatalf("new request: %v", err)
|
||||
}
|
||||
if o.form != nil {
|
||||
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
|
||||
}
|
||||
if o.cookie != "" {
|
||||
req.AddCookie(&http.Cookie{Name: sessionCookieName, Value: o.cookie})
|
||||
}
|
||||
for k, v := range o.headers {
|
||||
req.Header.Set(k, v)
|
||||
}
|
||||
resp, err := noFollow(ts).Do(req)
|
||||
if err != nil {
|
||||
t.Fatalf("%s %s: %v", o.method, o.path, err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
b, _ := io.ReadAll(resp.Body)
|
||||
return resp, string(b)
|
||||
}
|
||||
|
||||
func sessionCookie(resp *http.Response) *http.Cookie {
|
||||
for _, c := range resp.Cookies() {
|
||||
if c.Name == sessionCookieName {
|
||||
return c
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func loginForm(user, pass string) url.Values {
|
||||
return url.Values{"username": {user}, "password": {pass}}
|
||||
}
|
||||
|
||||
// login performs a successful login and returns the session cookie value.
|
||||
func login(t *testing.T, ts *httptest.Server) string {
|
||||
t.Helper()
|
||||
resp, _ := doReq(t, ts, reqOpts{method: http.MethodPost, path: "/login", form: loginForm(testUser, testPass),
|
||||
headers: map[string]string{"Origin": ts.URL}})
|
||||
c := sessionCookie(resp)
|
||||
if resp.StatusCode != http.StatusSeeOther || c == nil {
|
||||
t.Fatalf("login: status=%d cookie=%v, want 303 + session cookie", resp.StatusCode, c)
|
||||
}
|
||||
return c.Value
|
||||
}
|
||||
|
||||
func TestUnauthenticatedRedirectsToLogin(t *testing.T) {
|
||||
_, caURL := newFakeControlAPI(t)
|
||||
_, ts := newAuthTestServer(t, caURL, nil)
|
||||
|
||||
resp, _ := doReq(t, ts, reqOpts{path: "/ips"})
|
||||
if resp.StatusCode != http.StatusSeeOther {
|
||||
t.Fatalf("status=%d, want 303", resp.StatusCode)
|
||||
}
|
||||
if loc := resp.Header.Get("Location"); loc != "/login?next=%2Fips" {
|
||||
t.Fatalf("Location=%q, want /login?next=%%2Fips", loc)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUnauthenticatedHTMXGets401WithHXRedirect(t *testing.T) {
|
||||
_, caURL := newFakeControlAPI(t)
|
||||
_, ts := newAuthTestServer(t, caURL, nil)
|
||||
|
||||
resp, _ := doReq(t, ts, reqOpts{path: "/overview/fragment", headers: map[string]string{"HX-Request": "true"}})
|
||||
if resp.StatusCode != http.StatusUnauthorized || resp.Header.Get("HX-Redirect") != "/login" {
|
||||
t.Fatalf("status=%d HX-Redirect=%q, want 401 + /login", resp.StatusCode, resp.Header.Get("HX-Redirect"))
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoginSuccessSetsCookieAndGrantsAccess(t *testing.T) {
|
||||
_, caURL := newFakeControlAPI(t)
|
||||
_, ts := newAuthTestServer(t, caURL, nil)
|
||||
|
||||
resp, _ := doReq(t, ts, reqOpts{method: http.MethodPost, path: "/login",
|
||||
form: url.Values{"username": {testUser}, "password": {testPass}, "next": {"/ips"}},
|
||||
headers: map[string]string{"Origin": ts.URL}})
|
||||
if resp.StatusCode != http.StatusSeeOther || resp.Header.Get("Location") != "/ips" {
|
||||
t.Fatalf("status=%d Location=%q, want 303 /ips", resp.StatusCode, resp.Header.Get("Location"))
|
||||
}
|
||||
c := sessionCookie(resp)
|
||||
if c == nil {
|
||||
t.Fatal("no session cookie")
|
||||
}
|
||||
if !c.HttpOnly || c.SameSite != http.SameSiteStrictMode || c.Path != "/" || c.MaxAge <= 0 {
|
||||
t.Fatalf("cookie flags: httponly=%v samesite=%v path=%q maxage=%d", c.HttpOnly, c.SameSite, c.Path, c.MaxAge)
|
||||
}
|
||||
if c.Secure {
|
||||
t.Fatal("cookie must not be Secure over plain HTTP")
|
||||
}
|
||||
|
||||
resp, body := doReq(t, ts, reqOpts{path: "/overview", cookie: c.Value})
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
t.Fatalf("authenticated GET /overview: status=%d", resp.StatusCode)
|
||||
}
|
||||
if !strings.Contains(body, "Выйти") || !strings.Contains(body, testUser) {
|
||||
t.Fatal("sidebar must show the user and the logout button when auth is enabled")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCookieSecureBehindHTTPSProxy(t *testing.T) {
|
||||
_, caURL := newFakeControlAPI(t)
|
||||
_, ts := newAuthTestServer(t, caURL, nil)
|
||||
resp, _ := doReq(t, ts, reqOpts{method: http.MethodPost, path: "/login", form: loginForm(testUser, testPass),
|
||||
headers: map[string]string{"Origin": ts.URL, "X-Forwarded-Proto": "https"}})
|
||||
if c := sessionCookie(resp); c == nil || !c.Secure {
|
||||
t.Fatalf("cookie=%v, want Secure with X-Forwarded-Proto=https", c)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoginWrongPasswordShowsErrorWithoutCookie(t *testing.T) {
|
||||
_, caURL := newFakeControlAPI(t)
|
||||
_, ts := newAuthTestServer(t, caURL, nil)
|
||||
|
||||
for _, f := range []url.Values{loginForm(testUser, "nope"), loginForm("nobody", testPass)} {
|
||||
resp, body := doReq(t, ts, reqOpts{method: http.MethodPost, path: "/login", form: f,
|
||||
headers: map[string]string{"Origin": ts.URL}})
|
||||
if resp.StatusCode != http.StatusOK || sessionCookie(resp) != nil {
|
||||
t.Fatalf("status=%d cookie=%v, want 200 and no cookie", resp.StatusCode, sessionCookie(resp))
|
||||
}
|
||||
if !strings.Contains(body, "Неверный логин или пароль") {
|
||||
t.Fatal("login page must show the error")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestTamperedAndExpiredCookiesRejected(t *testing.T) {
|
||||
_, caURL := newFakeControlAPI(t)
|
||||
srv, ts := newAuthTestServer(t, caURL, nil)
|
||||
good := login(t, ts)
|
||||
|
||||
payload, sig, _ := strings.Cut(good, ".")
|
||||
for name, v := range map[string]string{
|
||||
"tampered payload": "AAAA" + payload[4:] + "." + sig,
|
||||
"tampered signature": payload + "." + sig[:len(sig)-2] + "AA",
|
||||
"no signature": payload,
|
||||
"garbage": "x.y",
|
||||
"signed with new key": otherKeyCookie(t),
|
||||
} {
|
||||
resp, _ := doReq(t, ts, reqOpts{path: "/overview", cookie: v})
|
||||
if resp.StatusCode != http.StatusSeeOther {
|
||||
t.Fatalf("%s: status=%d, want 303 to /login", name, resp.StatusCode)
|
||||
}
|
||||
}
|
||||
|
||||
// Expiry: advance the server's clock beyond the TTL.
|
||||
srv.auth.now = func() time.Time { return time.Now().Add(2 * time.Hour) }
|
||||
resp, _ := doReq(t, ts, reqOpts{path: "/overview", cookie: good})
|
||||
if resp.StatusCode != http.StatusSeeOther {
|
||||
t.Fatalf("expired cookie: status=%d, want 303", resp.StatusCode)
|
||||
}
|
||||
}
|
||||
|
||||
func otherKeyCookie(t *testing.T) string {
|
||||
t.Helper()
|
||||
a := newAuthState(Config{Username: "u", Password: "p", SessionSecret: "another-secret"},
|
||||
slog.New(slog.NewTextHandler(io.Discard, nil)))
|
||||
v, _ := a.issue(testUser)
|
||||
return v
|
||||
}
|
||||
|
||||
func TestLogoutClearsCookie(t *testing.T) {
|
||||
_, caURL := newFakeControlAPI(t)
|
||||
_, ts := newAuthTestServer(t, caURL, nil)
|
||||
good := login(t, ts)
|
||||
|
||||
resp, _ := doReq(t, ts, reqOpts{method: http.MethodPost, path: "/logout", cookie: good,
|
||||
headers: map[string]string{"Origin": ts.URL}})
|
||||
c := sessionCookie(resp)
|
||||
if resp.StatusCode != http.StatusSeeOther || resp.Header.Get("Location") != "/login" || c == nil || c.MaxAge >= 0 || c.Value != "" {
|
||||
t.Fatalf("logout: status=%d loc=%q cookie=%+v", resp.StatusCode, resp.Header.Get("Location"), c)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCSRFOriginCheck(t *testing.T) {
|
||||
_, caURL := newFakeControlAPI(t)
|
||||
_, ts := newAuthTestServer(t, caURL, nil)
|
||||
good := login(t, ts)
|
||||
|
||||
cases := map[string]map[string]string{
|
||||
"foreign origin": {"Origin": "http://evil.example"},
|
||||
"null origin": {"Origin": "null"},
|
||||
"no origin/referer": {},
|
||||
"foreign referer": {"Referer": "http://evil.example/page"},
|
||||
}
|
||||
for name, h := range cases {
|
||||
resp, _ := doReq(t, ts, reqOpts{method: http.MethodPost, path: "/ips/clear", cookie: good, headers: h})
|
||||
if resp.StatusCode != http.StatusForbidden {
|
||||
t.Fatalf("%s: status=%d, want 403", name, resp.StatusCode)
|
||||
}
|
||||
}
|
||||
// Login itself is covered too.
|
||||
resp, _ := doReq(t, ts, reqOpts{method: http.MethodPost, path: "/login", form: loginForm(testUser, testPass),
|
||||
headers: map[string]string{"Origin": "http://evil.example"}})
|
||||
if resp.StatusCode != http.StatusForbidden || sessionCookie(resp) != nil {
|
||||
t.Fatalf("cross-origin login: status=%d, want 403 and no cookie", resp.StatusCode)
|
||||
}
|
||||
// Same-origin Origin, and Referer fallback, pass.
|
||||
for name, h := range map[string]map[string]string{
|
||||
"origin": {"Origin": ts.URL},
|
||||
"referer": {"Referer": ts.URL + "/ips"},
|
||||
} {
|
||||
resp, _ := doReq(t, ts, reqOpts{method: http.MethodPost, path: "/ips/clear", cookie: good, headers: h})
|
||||
if resp.StatusCode == http.StatusForbidden || resp.StatusCode == http.StatusSeeOther {
|
||||
t.Fatalf("same-origin %s: status=%d, want request to be served", name, resp.StatusCode)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoginThrottle(t *testing.T) {
|
||||
_, caURL := newFakeControlAPI(t)
|
||||
_, ts := newAuthTestServer(t, caURL, nil)
|
||||
|
||||
for i := 0; i < maxLoginFailures; i++ {
|
||||
resp, _ := doReq(t, ts, reqOpts{method: http.MethodPost, path: "/login", form: loginForm(testUser, "bad"),
|
||||
headers: map[string]string{"Origin": ts.URL}})
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
t.Fatalf("attempt %d: status=%d, want 200", i+1, resp.StatusCode)
|
||||
}
|
||||
}
|
||||
// Now locked out, even with the right password.
|
||||
resp, _ := doReq(t, ts, reqOpts{method: http.MethodPost, path: "/login", form: loginForm(testUser, testPass),
|
||||
headers: map[string]string{"Origin": ts.URL}})
|
||||
if resp.StatusCode != http.StatusTooManyRequests || resp.Header.Get("Retry-After") == "" || sessionCookie(resp) != nil {
|
||||
t.Fatalf("status=%d Retry-After=%q, want 429 with Retry-After and no cookie", resp.StatusCode, resp.Header.Get("Retry-After"))
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoginThrottleClearedBySuccessAndExpires(t *testing.T) {
|
||||
th := newLoginThrottle()
|
||||
now := time.Now()
|
||||
for i := 0; i < maxLoginFailures-1; i++ {
|
||||
th.fail("1.2.3.4", now)
|
||||
}
|
||||
if b, _ := th.blocked("1.2.3.4", now); b {
|
||||
t.Fatal("blocked before reaching the limit")
|
||||
}
|
||||
th.fail("1.2.3.4", now)
|
||||
if b, _ := th.blocked("1.2.3.4", now); !b {
|
||||
t.Fatal("not blocked at the limit")
|
||||
}
|
||||
if b, _ := th.blocked("1.2.3.4", now.Add(loginFailureWindow+time.Second)); b {
|
||||
t.Fatal("still blocked after the window")
|
||||
}
|
||||
if len(th.failures) != 0 {
|
||||
t.Fatal("expired entries must be pruned")
|
||||
}
|
||||
th.fail("5.6.7.8", now)
|
||||
th.clear("5.6.7.8")
|
||||
if len(th.failures) != 0 {
|
||||
t.Fatal("clear must drop the entry")
|
||||
}
|
||||
}
|
||||
|
||||
func TestStaticAndLoginPageOpen(t *testing.T) {
|
||||
_, caURL := newFakeControlAPI(t)
|
||||
_, ts := newAuthTestServer(t, caURL, nil)
|
||||
|
||||
if resp, _ := doReq(t, ts, reqOpts{path: "/static/dashboard.css"}); resp.StatusCode != http.StatusOK {
|
||||
t.Fatalf("/static/dashboard.css: status=%d, want 200", resp.StatusCode)
|
||||
}
|
||||
resp, body := doReq(t, ts, reqOpts{path: "/login"})
|
||||
if resp.StatusCode != http.StatusOK || !strings.Contains(body, `name="password"`) {
|
||||
t.Fatalf("/login: status=%d", resp.StatusCode)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoginNextOpenRedirectRejected(t *testing.T) {
|
||||
for next, want := range map[string]string{
|
||||
"/ips": "/ips",
|
||||
"/ips?q=1": "/ips?q=1",
|
||||
"//evil.example": "/",
|
||||
"/\\evil.example": "/",
|
||||
"http://evil.example": "/",
|
||||
"https://evil.example/x": "/",
|
||||
"javascript:alert(1)": "/",
|
||||
"": "/",
|
||||
"/a\r\nSet-Cookie: x": "/",
|
||||
} {
|
||||
if got := safeNext(next); got != want {
|
||||
t.Fatalf("safeNext(%q)=%q, want %q", next, got, want)
|
||||
}
|
||||
}
|
||||
|
||||
_, caURL := newFakeControlAPI(t)
|
||||
_, ts := newAuthTestServer(t, caURL, nil)
|
||||
resp, _ := doReq(t, ts, reqOpts{method: http.MethodPost, path: "/login",
|
||||
form: url.Values{"username": {testUser}, "password": {testPass}, "next": {"//evil.example"}},
|
||||
headers: map[string]string{"Origin": ts.URL}})
|
||||
if resp.Header.Get("Location") != "/" {
|
||||
t.Fatalf("Location=%q, want /", resp.Header.Get("Location"))
|
||||
}
|
||||
}
|
||||
|
||||
func TestControlAPITokenForwarded(t *testing.T) {
|
||||
fake, caURL := newFakeControlAPI(t)
|
||||
_, ts := newAuthTestServer(t, caURL, nil)
|
||||
good := login(t, ts)
|
||||
|
||||
if resp, _ := doReq(t, ts, reqOpts{path: "/overview", cookie: good}); resp.StatusCode != http.StatusOK {
|
||||
t.Fatalf("GET /overview: status=%d", resp.StatusCode)
|
||||
}
|
||||
fake.authMu.Lock()
|
||||
defer fake.authMu.Unlock()
|
||||
if len(fake.authHeaders) == 0 {
|
||||
t.Fatal("control-api was never called")
|
||||
}
|
||||
for _, h := range fake.authHeaders {
|
||||
if h != "Bearer ca-admin-token" {
|
||||
t.Fatalf("control-api Authorization=%q, want Bearer ca-admin-token", h)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthDisabledWithZeroConfig(t *testing.T) {
|
||||
fake, caURL := newFakeControlAPI(t)
|
||||
ts := newTestServer(t, caURL) // zero auth fields
|
||||
resp, body := doReq(t, ts, reqOpts{path: "/overview"})
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
t.Fatalf("status=%d, want 200 without login", resp.StatusCode)
|
||||
}
|
||||
if strings.Contains(body, "/logout") {
|
||||
t.Fatal("logout control must be hidden when auth is disabled")
|
||||
}
|
||||
if resp, _ := doReq(t, ts, reqOpts{path: "/login"}); resp.StatusCode != http.StatusSeeOther || resp.Header.Get("Location") != "/" {
|
||||
t.Fatalf("/login with auth disabled: status=%d loc=%q, want redirect to /", resp.StatusCode, resp.Header.Get("Location"))
|
||||
}
|
||||
fake.authMu.Lock()
|
||||
defer fake.authMu.Unlock()
|
||||
for _, h := range fake.authHeaders {
|
||||
if h != "" {
|
||||
t.Fatalf("no token configured but Authorization=%q was sent", h)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthDisabledWhenOnlyUsernameSet(t *testing.T) {
|
||||
_, caURL := newFakeControlAPI(t)
|
||||
_, ts := newAuthTestServer(t, caURL, func(c *Config) { c.Password = "" })
|
||||
if resp, _ := doReq(t, ts, reqOpts{path: "/overview"}); resp.StatusCode != http.StatusOK {
|
||||
t.Fatalf("status=%d, want 200 (auth disabled without password)", resp.StatusCode)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEmptySessionSecretUsesRandomKey(t *testing.T) {
|
||||
_, caURL := newFakeControlAPI(t)
|
||||
_, ts := newAuthTestServer(t, caURL, func(c *Config) { c.SessionSecret = "" })
|
||||
good := login(t, ts)
|
||||
if resp, _ := doReq(t, ts, reqOpts{path: "/overview", cookie: good}); resp.StatusCode != http.StatusOK {
|
||||
t.Fatalf("status=%d, want 200 with random-key session", resp.StatusCode)
|
||||
}
|
||||
}
|
||||
@@ -39,6 +39,8 @@ func (e *apiErr) Error() string {
|
||||
type client struct {
|
||||
baseURL string
|
||||
http *http.Client
|
||||
// token, when non-empty, is sent to control-api as a Bearer credential.
|
||||
token string
|
||||
}
|
||||
|
||||
func newClient(baseURL string, timeout time.Duration) *client {
|
||||
@@ -61,6 +63,9 @@ func (c *client) do(ctx context.Context, method, path string, body, out interfac
|
||||
if body != nil {
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
}
|
||||
if c.token != "" {
|
||||
req.Header.Set("Authorization", "Bearer "+c.token)
|
||||
}
|
||||
|
||||
resp, err := c.http.Do(req)
|
||||
if err != nil {
|
||||
|
||||
@@ -46,6 +46,10 @@ type fakeControlAPI struct {
|
||||
// autoCycleDown makes all four endpoints answer 500 (unavailable API).
|
||||
autoCycle autoCycleDTO
|
||||
autoCycleDown bool
|
||||
|
||||
// authHeaders records the Authorization header of every request served.
|
||||
authMu sync.Mutex
|
||||
authHeaders []string
|
||||
}
|
||||
|
||||
func newFakeControlAPI(t *testing.T) (*fakeControlAPI, string) {
|
||||
@@ -532,7 +536,12 @@ func (f *fakeControlAPI) handler() http.Handler {
|
||||
writeJSON(w, http.StatusOK, map[string]bool{"ok": true})
|
||||
})
|
||||
|
||||
return mux
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
f.authMu.Lock()
|
||||
f.authHeaders = append(f.authHeaders, r.Header.Get("Authorization"))
|
||||
f.authMu.Unlock()
|
||||
mux.ServeHTTP(w, r)
|
||||
})
|
||||
}
|
||||
|
||||
func (f *fakeControlAPI) findIP(addr string) int {
|
||||
|
||||
@@ -31,7 +31,7 @@ func (s *Server) handleCheckTypesPage(w http.ResponseWriter, r *http.Request) {
|
||||
data, err := s.loadCheckTypesPage(r)
|
||||
data.ActiveNav = "check-types"
|
||||
data.Banner = bannerFor(err)
|
||||
s.renderPage(w, "checktypes_page", data)
|
||||
s.renderPage(w, r, "checktypes_page", data)
|
||||
}
|
||||
|
||||
func (s *Server) renderCheckTypesTable(w http.ResponseWriter, r *http.Request, actionErr error) {
|
||||
|
||||
@@ -25,7 +25,7 @@ func (s *Server) handleIPsPage(w http.ResponseWriter, r *http.Request) {
|
||||
data := ipsPageData{Items: items, FIPSettleSeconds: settings.FIPSettleSeconds}
|
||||
data.ActiveNav = "ips"
|
||||
data.Banner = bannerFor(err)
|
||||
s.renderPage(w, "ips_page", data)
|
||||
s.renderPage(w, r, "ips_page", data)
|
||||
}
|
||||
|
||||
func (s *Server) handleIPDetail(w http.ResponseWriter, r *http.Request) {
|
||||
@@ -34,7 +34,7 @@ func (s *Server) handleIPDetail(w http.ResponseWriter, r *http.Request) {
|
||||
data := ipDetailData{Detail: detail}
|
||||
data.ActiveNav = "ips"
|
||||
data.Banner = bannerFor(err)
|
||||
s.renderPage(w, "ip_detail_page", data)
|
||||
s.renderPage(w, r, "ip_detail_page", data)
|
||||
}
|
||||
|
||||
// renderIPsTable re-fetches the current queue and renders the ips_table
|
||||
|
||||
@@ -57,7 +57,7 @@ func (s *Server) handleOverview(w http.ResponseWriter, r *http.Request) {
|
||||
data, err := s.loadOverview(r)
|
||||
data.ActiveNav = "overview"
|
||||
data.Banner = bannerFor(err)
|
||||
s.renderPage(w, "overview_page", data)
|
||||
s.renderPage(w, r, "overview_page", data)
|
||||
}
|
||||
|
||||
// handleOverviewFragment serves both the recurring poll and every
|
||||
|
||||
@@ -34,7 +34,7 @@ func (s *Server) handleRegistryPage(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
data.ActiveNav = "registry"
|
||||
data.Banner = bannerFor(err)
|
||||
s.renderPage(w, "registry_page", data)
|
||||
s.renderPage(w, r, "registry_page", data)
|
||||
}
|
||||
|
||||
// filterRegistryItems narrows items to those whose address contains q
|
||||
@@ -68,5 +68,5 @@ func (s *Server) handleRegistryDetail(w http.ResponseWriter, r *http.Request) {
|
||||
data := registryDetailData{History: history}
|
||||
data.ActiveNav = "registry"
|
||||
data.Banner = bannerFor(err)
|
||||
s.renderPage(w, "registry_detail_page", data)
|
||||
s.renderPage(w, r, "registry_detail_page", data)
|
||||
}
|
||||
@@ -29,7 +29,7 @@ func (s *Server) handleSettingsPage(w http.ResponseWriter, r *http.Request) {
|
||||
data := settingsPageData{Settings: settings, Inbound: inbound, AutoCycle: autoCycle}
|
||||
data.ActiveNav = "settings"
|
||||
data.Banner = bannerFor(err)
|
||||
s.renderPage(w, "settings_page", data)
|
||||
s.renderPage(w, r, "settings_page", data)
|
||||
}
|
||||
|
||||
// renderSettingsForm re-fetches the current settings and renders the
|
||||
|
||||
@@ -16,7 +16,7 @@ func (s *Server) handleSitesPage(w http.ResponseWriter, r *http.Request) {
|
||||
data := sitesPageData{Items: items}
|
||||
data.ActiveNav = "sites"
|
||||
data.Banner = bannerFor(err)
|
||||
s.renderPage(w, "sites_page", data)
|
||||
s.renderPage(w, r, "sites_page", data)
|
||||
}
|
||||
|
||||
func (s *Server) renderSitesTable(w http.ResponseWriter, r *http.Request, actionErr error) {
|
||||
|
||||
@@ -83,7 +83,7 @@ func (s *Server) handleTargetsPage(w http.ResponseWriter, r *http.Request) {
|
||||
data, err := s.loadTargetsPage(r)
|
||||
data.ActiveNav = "targets"
|
||||
data.Banner = bannerFor(err)
|
||||
s.renderPage(w, "targets_page", data)
|
||||
s.renderPage(w, r, "targets_page", data)
|
||||
}
|
||||
|
||||
// renderTargetsTable renders the #targets-table-wrap swap target plus,
|
||||
|
||||
@@ -15,7 +15,7 @@ func (s *Server) handleValidatorsPage(w http.ResponseWriter, r *http.Request) {
|
||||
data := validatorsPageData{Items: items}
|
||||
data.ActiveNav = "validators"
|
||||
data.Banner = bannerFor(err)
|
||||
s.renderPage(w, "validators_page", data)
|
||||
s.renderPage(w, r, "validators_page", data)
|
||||
}
|
||||
|
||||
func (s *Server) renderValidatorsTable(w http.ResponseWriter, r *http.Request, actionErr error) {
|
||||
|
||||
@@ -155,6 +155,10 @@ type bannerData struct {
|
||||
type PageData struct {
|
||||
Banner bannerData
|
||||
ActiveNav string
|
||||
// AuthEnabled/User are filled in by renderPage (see withAuthInfo) so the
|
||||
// sidebar can show the signed-in user and the logout button.
|
||||
AuthEnabled bool
|
||||
User string
|
||||
}
|
||||
|
||||
func bannerFor(err error) bannerData {
|
||||
@@ -172,8 +176,9 @@ func bannerFor(err error) bannerData {
|
||||
// navigation. actionErr (if any — e.g. the primary control-api call for
|
||||
// this page failed) is surfaced via the embedded PageData.Banner, which
|
||||
// the caller must have already set via bannerFor.
|
||||
func (s *Server) renderPage(w http.ResponseWriter, name string, data interface{}) {
|
||||
func (s *Server) renderPage(w http.ResponseWriter, r *http.Request, name string, data interface{}) {
|
||||
w.Header().Set("Content-Type", "text/html; charset=utf-8")
|
||||
data = s.withAuthInfo(r, data)
|
||||
if err := s.tmpl.ExecuteTemplate(w, name, data); err != nil {
|
||||
s.Log.Error("render page", "template", name, "err", err)
|
||||
}
|
||||
|
||||
@@ -3,6 +3,12 @@ package dashboard
|
||||
import "net/http"
|
||||
|
||||
func (s *Server) routes(mux *http.ServeMux) {
|
||||
// /login, /logout and /static/ are the only routes reachable without a
|
||||
// session — see authMiddleware in auth.go.
|
||||
mux.HandleFunc("GET /login", s.handleLoginPage)
|
||||
mux.HandleFunc("POST /login", s.handleLoginSubmit)
|
||||
mux.HandleFunc("POST /logout", s.handleLogout)
|
||||
|
||||
mux.HandleFunc("GET /{$}", s.handleIndex)
|
||||
|
||||
mux.HandleFunc("GET /overview", s.handleOverview)
|
||||
|
||||
@@ -18,6 +18,19 @@ type Config struct {
|
||||
ControlAPITimeout time.Duration
|
||||
LastCompletedCount int
|
||||
OverviewPollIntervalS int
|
||||
|
||||
// ControlAPIToken is the admin bearer token sent to control-api. Empty =
|
||||
// no Authorization header.
|
||||
ControlAPIToken string
|
||||
// Username/Password are the single dashboard administrator's login. If
|
||||
// either is empty, login is DISABLED (every page is open).
|
||||
Username string
|
||||
Password string
|
||||
// SessionSecret is the HMAC key for session cookies; empty = random per
|
||||
// process start (sessions are lost on restart).
|
||||
SessionSecret string
|
||||
// SessionTTL is the session lifetime (default 8h).
|
||||
SessionTTL time.Duration
|
||||
}
|
||||
|
||||
type Server struct {
|
||||
@@ -25,6 +38,7 @@ type Server struct {
|
||||
Cfg Config
|
||||
tmpl *template.Template
|
||||
Log *slog.Logger
|
||||
auth *authState
|
||||
}
|
||||
|
||||
func New(cfg Config, log *slog.Logger) (*Server, error) {
|
||||
@@ -32,18 +46,22 @@ func New(cfg Config, log *slog.Logger) (*Server, error) {
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
ca := newClient(cfg.ControlAPIBaseURL, cfg.ControlAPITimeout)
|
||||
ca.token = cfg.ControlAPIToken
|
||||
return &Server{
|
||||
CA: newClient(cfg.ControlAPIBaseURL, cfg.ControlAPITimeout),
|
||||
CA: ca,
|
||||
Cfg: cfg,
|
||||
tmpl: tmpl,
|
||||
Log: log,
|
||||
auth: newAuthState(cfg, log),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (s *Server) Handler() http.Handler {
|
||||
mux := http.NewServeMux()
|
||||
s.routes(mux)
|
||||
return loggingMiddleware(s.Log, mux)
|
||||
// Never log cookies, Authorization or form bodies here.
|
||||
return loggingMiddleware(s.Log, s.authMiddleware(mux))
|
||||
}
|
||||
|
||||
func loggingMiddleware(log *slog.Logger, next http.Handler) http.Handler {
|
||||
|
||||
@@ -337,7 +337,7 @@ nav.nav-groups { display: flex; flex-direction: column; gap: 1px; }
|
||||
}
|
||||
.form-check { display: flex; align-items: center; gap: 6px; font-size: 13px; }
|
||||
|
||||
input[type="text"], input[type="number"], input[type="search"], textarea, select {
|
||||
input[type="text"], input[type="password"], input[type="number"], input[type="search"], textarea, select {
|
||||
font-family: var(--font-ui);
|
||||
font-size: 13.5px;
|
||||
padding: 7px 10px;
|
||||
@@ -544,3 +544,18 @@ tr.edit-row textarea { flex: 1 1 auto; }
|
||||
Height is set by JS (see templates/targets.html), computed from the
|
||||
field's own line-height, so it stays correct if font-size ever changes. */
|
||||
textarea.autosize { overflow-y: hidden; resize: none; }
|
||||
|
||||
/* ---------- login page & sidebar user ---------- */
|
||||
.login-shell { min-height: 100vh; display: flex; align-items: center; justify-content: center; padding: 24px 16px; position: relative; }
|
||||
.login-card { width: 100%; max-width: 380px; }
|
||||
.login-card h1 { font-size: 18px; margin: 0 0 4px; }
|
||||
.login-card .panel-body { display: flex; flex-direction: column; gap: 14px; }
|
||||
/* .field is flex: 1 1 220px (for rows); in this column the basis would become a 220px height. */
|
||||
.login-card .field { flex: 0 0 auto; }
|
||||
.sidebar-user {
|
||||
margin-top: auto; margin-bottom: 8px;
|
||||
display: flex; align-items: center; justify-content: space-between; gap: 8px;
|
||||
font-family: var(--font-display); font-size: 11px; color: var(--text-muted);
|
||||
}
|
||||
.sidebar-user-name { min-width: 0; overflow: hidden; text-overflow: ellipsis; white-space: nowrap; }
|
||||
.sidebar-user + .sidebar-foot { margin-top: 0; }
|
||||
@@ -78,6 +78,13 @@
|
||||
</a>
|
||||
</nav>
|
||||
|
||||
{{if .AuthEnabled}}
|
||||
<form class="sidebar-user" method="post" action="/logout">
|
||||
<span class="sidebar-user-name" title="{{.User}}">{{.User}}</span>
|
||||
<button type="submit" class="btn btn-ghost btn-sm">Выйти</button>
|
||||
</form>
|
||||
{{end}}
|
||||
|
||||
<div class="sidebar-foot">
|
||||
<span class="sidebar-foot-status"><span class="pulse-dot"></span>control-api</span>
|
||||
<button type="button" class="theme-toggle" id="themeToggle" aria-label="Переключить тему" title="Переключить тему">
|
||||
|
||||
@@ -0,0 +1,29 @@
|
||||
{{define "login_page"}}
|
||||
<!doctype html>
|
||||
<html lang="ru">
|
||||
<head>{{template "html_head" .}}</head>
|
||||
<body>
|
||||
<div class="bg-grid"></div>
|
||||
<main class="login-shell">
|
||||
<div class="login-card">
|
||||
<form class="panel" method="post" action="/login" autocomplete="on">
|
||||
<div class="panel-head"><h2>Cloud IP Validator — вход</h2></div>
|
||||
<div class="panel-body">
|
||||
{{if .Error}}<div class="alert alert-warning" role="alert">{{.Error}}</div>{{end}}
|
||||
<input type="hidden" name="next" value="{{.Next}}">
|
||||
<div class="field">
|
||||
<label for="login-username">Логин</label>
|
||||
<input type="text" id="login-username" name="username" autocomplete="username" autofocus required>
|
||||
</div>
|
||||
<div class="field">
|
||||
<label for="login-password">Пароль</label>
|
||||
<input type="password" id="login-password" name="password" autocomplete="current-password" required>
|
||||
</div>
|
||||
<button type="submit" class="btn btn-primary btn-block">Войти</button>
|
||||
</div>
|
||||
</form>
|
||||
</div>
|
||||
</main>
|
||||
</body>
|
||||
</html>
|
||||
{{end}}
|
||||
@@ -0,0 +1,108 @@
|
||||
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
|
||||
}
|
||||
@@ -0,0 +1,181 @@
|
||||
package httpapi
|
||||
|
||||
import (
|
||||
"context"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"cloudipvalidator/internal/config"
|
||||
"cloudipvalidator/internal/db"
|
||||
"cloudipvalidator/internal/openstack"
|
||||
"cloudipvalidator/internal/orchestrator"
|
||||
)
|
||||
|
||||
const (
|
||||
testAdminToken = "admin-token-for-tests"
|
||||
testAgentToken = "agent-token-for-tests"
|
||||
)
|
||||
|
||||
func newAuthTestServer(t *testing.T, admin, agent string) (*Server, *httptest.Server) {
|
||||
t.Helper()
|
||||
ctx := context.Background()
|
||||
d, err := db.Open(ctx, filepath.Join(t.TempDir(), "test.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open db: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { d.Close() })
|
||||
cfg := &config.ControlAPI{
|
||||
Orchestrator: config.OrchestratorConfig{
|
||||
PollIntervalSeconds: 1, SelfCheckTimeoutSeconds: 10, MaxSelfCheckRetries: 3,
|
||||
CheckingWindowSeconds: 120, MaxRetries: 3, LeaseTTLSeconds: 180, HeartbeatTimeoutSeconds: 30,
|
||||
},
|
||||
Inbound: config.InboundConfig{Ports: []int{22}, ICMP: true},
|
||||
}
|
||||
if err := d.BootstrapFromConfig(ctx, cfg); err != nil {
|
||||
t.Fatalf("bootstrap: %v", err)
|
||||
}
|
||||
log := slog.New(slog.NewTextHandler(os.Stderr, &slog.HandlerOptions{Level: slog.LevelError}))
|
||||
srv := New(d, orchestrator.New(d, openstack.NewMockClient(), cfg, log), log).WithAuth(admin, agent)
|
||||
ts := httptest.NewServer(srv.Handler())
|
||||
t.Cleanup(ts.Close)
|
||||
return srv, ts
|
||||
}
|
||||
|
||||
// concretePath fills the {wildcards} of a ServeMux pattern with dummy values.
|
||||
func concretePath(pattern string) (method, path string) {
|
||||
method, path, _ = strings.Cut(pattern, " ")
|
||||
var b strings.Builder
|
||||
for {
|
||||
i := strings.IndexByte(path, '{')
|
||||
if i < 0 {
|
||||
b.WriteString(path)
|
||||
break
|
||||
}
|
||||
j := strings.IndexByte(path, '}')
|
||||
b.WriteString(path[:i])
|
||||
b.WriteString("1")
|
||||
path = path[j+1:]
|
||||
}
|
||||
return method, b.String()
|
||||
}
|
||||
|
||||
func callWithToken(t *testing.T, ts *httptest.Server, pattern, token string) *http.Response {
|
||||
t.Helper()
|
||||
method, path := concretePath(pattern)
|
||||
req, err := http.NewRequest(method, ts.URL+path, strings.NewReader("{}"))
|
||||
if err != nil {
|
||||
t.Fatalf("new request: %v", err)
|
||||
}
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
if token != "" {
|
||||
req.Header.Set("Authorization", "Bearer "+token)
|
||||
}
|
||||
resp, err := ts.Client().Do(req)
|
||||
if err != nil {
|
||||
t.Fatalf("%s: %v", pattern, err)
|
||||
}
|
||||
resp.Body.Close()
|
||||
return resp
|
||||
}
|
||||
|
||||
func TestRouteTableIsClassified(t *testing.T) {
|
||||
srv, _ := newAuthTestServer(t, "", "")
|
||||
counts := map[string]int{}
|
||||
for _, rt := range srv.Routes() {
|
||||
counts[rt.Access]++
|
||||
if rt.Access == "invalid" {
|
||||
t.Fatalf("route %q has no valid access level", rt.Pattern)
|
||||
}
|
||||
if strings.Contains(rt.Pattern, "/api/v1/admin/") && rt.Access != "admin" {
|
||||
t.Fatalf("admin route %q is %s, want admin", rt.Pattern, rt.Access)
|
||||
}
|
||||
}
|
||||
if counts["admin"] != 33 || counts["agent"] != 5 || counts["open"] != 7 {
|
||||
t.Fatalf("access counts = %v, want admin=33 agent=5 open=7", counts)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthMatrixOverRouteTable(t *testing.T) {
|
||||
srv, ts := newAuthTestServer(t, testAdminToken, testAgentToken)
|
||||
for _, rt := range srv.Routes() {
|
||||
rt := rt
|
||||
t.Run(rt.Pattern, func(t *testing.T) {
|
||||
tokens := []struct {
|
||||
name, token string
|
||||
allowed bool
|
||||
}{
|
||||
{"none", "", rt.Access == "open"},
|
||||
{"wrong", "definitely-wrong", rt.Access == "open"},
|
||||
{"admin token", testAdminToken, rt.Access == "open" || rt.Access == "admin"},
|
||||
{"agent token", testAgentToken, rt.Access == "open" || rt.Access == "agent"},
|
||||
}
|
||||
for _, tc := range tokens {
|
||||
resp := callWithToken(t, ts, rt.Pattern, tc.token)
|
||||
if tc.allowed && resp.StatusCode == http.StatusUnauthorized {
|
||||
t.Fatalf("%s: got 401, want request to pass auth", tc.name)
|
||||
}
|
||||
if !tc.allowed {
|
||||
if resp.StatusCode != http.StatusUnauthorized {
|
||||
t.Fatalf("%s: status=%d, want 401", tc.name, resp.StatusCode)
|
||||
}
|
||||
if resp.Header.Get("WWW-Authenticate") != "Bearer" {
|
||||
t.Fatalf("%s: missing WWW-Authenticate: Bearer", tc.name)
|
||||
}
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestEmptyTokensLeaveEverythingOpen(t *testing.T) {
|
||||
srv, ts := newAuthTestServer(t, "", "")
|
||||
for _, rt := range srv.Routes() {
|
||||
if resp := callWithToken(t, ts, rt.Pattern, ""); resp.StatusCode == http.StatusUnauthorized {
|
||||
t.Fatalf("%s: got 401 with no tokens configured", rt.Pattern)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestOnlyConfiguredLevelIsEnforced(t *testing.T) {
|
||||
_, ts := newAuthTestServer(t, testAdminToken, "")
|
||||
if resp := callWithToken(t, ts, "GET /api/v1/admin/status", ""); resp.StatusCode != http.StatusUnauthorized {
|
||||
t.Fatalf("admin without token: status=%d, want 401", resp.StatusCode)
|
||||
}
|
||||
if resp := callWithToken(t, ts, "POST /api/v1/agents/{id}/results", ""); resp.StatusCode == http.StatusUnauthorized {
|
||||
t.Fatalf("agent route must stay open while agent token is unset")
|
||||
}
|
||||
}
|
||||
|
||||
func TestHealthzAlwaysOpenAndBearerParsing(t *testing.T) {
|
||||
_, ts := newAuthTestServer(t, testAdminToken, testAgentToken)
|
||||
if resp := callWithToken(t, ts, "GET /healthz", ""); resp.StatusCode != http.StatusOK {
|
||||
t.Fatalf("healthz: status=%d, want 200", resp.StatusCode)
|
||||
}
|
||||
// Raw token without the Bearer scheme must not authenticate.
|
||||
req, _ := http.NewRequest(http.MethodGet, ts.URL+"/api/v1/admin/status", nil)
|
||||
req.Header.Set("Authorization", testAdminToken)
|
||||
resp, err := ts.Client().Do(req)
|
||||
if err != nil {
|
||||
t.Fatalf("request: %v", err)
|
||||
}
|
||||
resp.Body.Close()
|
||||
if resp.StatusCode != http.StatusUnauthorized {
|
||||
t.Fatalf("scheme-less token: status=%d, want 401", resp.StatusCode)
|
||||
}
|
||||
// Lowercase scheme is accepted (RFC 7235: case-insensitive).
|
||||
req, _ = http.NewRequest(http.MethodGet, ts.URL+"/api/v1/admin/status", nil)
|
||||
req.Header.Set("Authorization", "bearer "+testAdminToken)
|
||||
resp, err = ts.Client().Do(req)
|
||||
if err != nil {
|
||||
t.Fatalf("request: %v", err)
|
||||
}
|
||||
resp.Body.Close()
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
t.Fatalf("lowercase bearer: status=%d, want 200", resp.StatusCode)
|
||||
}
|
||||
}
|
||||
+75
-53
@@ -2,61 +2,83 @@ package httpapi
|
||||
|
||||
import "net/http"
|
||||
|
||||
func (s *Server) routes(mux *http.ServeMux) {
|
||||
mux.HandleFunc("GET /healthz", s.handleHealthz)
|
||||
// route is one entry of the Control API route table. The access level is
|
||||
// mandatory: every route must state who may call it, so a new endpoint cannot
|
||||
// be exposed by accident. The same table is used by the tests.
|
||||
type route struct {
|
||||
Pattern string
|
||||
Handler http.HandlerFunc
|
||||
Access accessLevel
|
||||
}
|
||||
|
||||
mux.HandleFunc("POST /api/v1/agents/register", s.handleAgentRegister)
|
||||
mux.HandleFunc("POST /api/v1/agents/{id}/heartbeat", s.handleAgentHeartbeat)
|
||||
mux.HandleFunc("GET /api/v1/agents/{id}/assignment", s.handleAgentAssignment)
|
||||
mux.HandleFunc("POST /api/v1/agents/{id}/self-check", s.handleAgentSelfCheck)
|
||||
mux.HandleFunc("POST /api/v1/agents/{id}/events", s.handleAgentEvent)
|
||||
mux.HandleFunc("POST /api/v1/agents/{id}/results", s.handleAgentResults)
|
||||
mux.HandleFunc("POST /api/v1/agents/{id}/complete", s.handleAgentComplete)
|
||||
// RouteInfo is the exported, handler-free view of a route table entry.
|
||||
type RouteInfo struct {
|
||||
Pattern string
|
||||
Access string // "open", "agent" or "admin"
|
||||
}
|
||||
|
||||
mux.HandleFunc("POST /api/v1/probers/register", s.handleProberRegister)
|
||||
mux.HandleFunc("POST /api/v1/probers/{site_id}/heartbeat", s.handleProberHeartbeat)
|
||||
mux.HandleFunc("GET /api/v1/probers/{site_id}/assignments", s.handleProberAssignments)
|
||||
mux.HandleFunc("POST /api/v1/probers/{site_id}/results", s.handleProberResults)
|
||||
// Routes returns the access classification of every registered route.
|
||||
func (s *Server) Routes() []RouteInfo {
|
||||
table := s.routeTable()
|
||||
out := make([]RouteInfo, 0, len(table))
|
||||
for _, rt := range table {
|
||||
out = append(out, RouteInfo{Pattern: rt.Pattern, Access: rt.Access.String()})
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
mux.HandleFunc("GET /api/v1/admin/status", s.handleAdminStatus)
|
||||
mux.HandleFunc("GET /api/v1/admin/ips", s.handleAdminIPs)
|
||||
mux.HandleFunc("POST /api/v1/admin/ips", s.handleAdminSubmitIPs)
|
||||
mux.HandleFunc("POST /api/v1/admin/ips/scan", s.handleAdminScanFloatingIPs)
|
||||
mux.HandleFunc("GET /api/v1/admin/ips/{ip}", s.handleAdminIPDetail)
|
||||
mux.HandleFunc("POST /api/v1/admin/ips/{ip}/cancel", s.handleAdminCancelIP)
|
||||
mux.HandleFunc("DELETE /api/v1/admin/ips/{ip}", s.handleAdminDeleteIP)
|
||||
mux.HandleFunc("POST /api/v1/admin/ips/delete", s.handleAdminDeleteIPs)
|
||||
mux.HandleFunc("POST /api/v1/admin/ips/clear", s.handleAdminClearQueue)
|
||||
mux.HandleFunc("GET /api/v1/admin/validators", s.handleAdminValidators)
|
||||
|
||||
mux.HandleFunc("GET /api/v1/admin/auto-cycle", s.handleAdminGetAutoCycle)
|
||||
mux.HandleFunc("PUT /api/v1/admin/auto-cycle", s.handleAdminPutAutoCycle)
|
||||
mux.HandleFunc("POST /api/v1/admin/auto-cycle/start", s.handleAdminStartAutoCycle)
|
||||
mux.HandleFunc("POST /api/v1/admin/auto-cycle/stop", s.handleAdminStopAutoCycle)
|
||||
|
||||
mux.HandleFunc("GET /api/v1/admin/registry", s.handleAdminRegistry)
|
||||
mux.HandleFunc("GET /api/v1/admin/registry/{ip}", s.handleAdminRegistryHistory)
|
||||
|
||||
mux.HandleFunc("GET /api/v1/admin/config/validators", s.handleConfigListValidators)
|
||||
mux.HandleFunc("POST /api/v1/admin/config/validators", s.handleConfigCreateValidator)
|
||||
mux.HandleFunc("PUT /api/v1/admin/config/validators/{id}", s.handleConfigUpdateValidator)
|
||||
mux.HandleFunc("DELETE /api/v1/admin/config/validators/{id}", s.handleConfigDeleteValidator)
|
||||
|
||||
mux.HandleFunc("GET /api/v1/admin/config/sites", s.handleConfigListSites)
|
||||
mux.HandleFunc("PUT /api/v1/admin/config/sites/{index}", s.handleConfigPutSite)
|
||||
mux.HandleFunc("DELETE /api/v1/admin/config/sites/{index}", s.handleConfigDeleteSite)
|
||||
|
||||
mux.HandleFunc("GET /api/v1/admin/config/targets", s.handleConfigListTargets)
|
||||
mux.HandleFunc("PUT /api/v1/admin/config/targets/{group}", s.handleConfigPutTargetGroup)
|
||||
mux.HandleFunc("DELETE /api/v1/admin/config/targets/{group}", s.handleConfigDeleteTargetGroup)
|
||||
|
||||
mux.HandleFunc("GET /api/v1/admin/config/check-types", s.handleConfigListCheckTypes)
|
||||
mux.HandleFunc("PUT /api/v1/admin/config/check-types/{name}", s.handleConfigPutCheckType)
|
||||
mux.HandleFunc("DELETE /api/v1/admin/config/check-types/{name}", s.handleConfigDeleteCheckType)
|
||||
|
||||
mux.HandleFunc("GET /api/v1/admin/config/orchestrator", s.handleConfigGetOrchestratorSettings)
|
||||
mux.HandleFunc("PUT /api/v1/admin/config/orchestrator", s.handleConfigPutOrchestratorSettings)
|
||||
func (s *Server) routes(mux *http.ServeMux) {
|
||||
for _, rt := range s.routeTable() {
|
||||
mux.HandleFunc(rt.Pattern, s.Auth.require(rt.Access, rt.Handler))
|
||||
}
|
||||
}
|
||||
|
||||
mux.HandleFunc("GET /api/v1/admin/config/inbound-checks", s.handleConfigGetInboundChecks)
|
||||
mux.HandleFunc("PUT /api/v1/admin/config/inbound-checks", s.handleConfigPutInboundChecks)
|
||||
func (s *Server) routeTable() []route {
|
||||
return []route{
|
||||
{"GET /healthz", s.handleHealthz, accessOpen},
|
||||
{"POST /api/v1/agents/register", s.handleAgentRegister, accessOpen},
|
||||
{"POST /api/v1/agents/{id}/heartbeat", s.handleAgentHeartbeat, accessOpen},
|
||||
{"GET /api/v1/agents/{id}/assignment", s.handleAgentAssignment, accessOpen},
|
||||
{"POST /api/v1/agents/{id}/self-check", s.handleAgentSelfCheck, accessAgent},
|
||||
{"POST /api/v1/agents/{id}/events", s.handleAgentEvent, accessAgent},
|
||||
{"POST /api/v1/agents/{id}/results", s.handleAgentResults, accessAgent},
|
||||
{"POST /api/v1/agents/{id}/complete", s.handleAgentComplete, accessAgent},
|
||||
{"POST /api/v1/probers/register", s.handleProberRegister, accessOpen},
|
||||
{"POST /api/v1/probers/{site_id}/heartbeat", s.handleProberHeartbeat, accessOpen},
|
||||
{"GET /api/v1/probers/{site_id}/assignments", s.handleProberAssignments, accessOpen},
|
||||
{"POST /api/v1/probers/{site_id}/results", s.handleProberResults, accessAgent},
|
||||
{"GET /api/v1/admin/status", s.handleAdminStatus, accessAdmin},
|
||||
{"GET /api/v1/admin/ips", s.handleAdminIPs, accessAdmin},
|
||||
{"POST /api/v1/admin/ips", s.handleAdminSubmitIPs, accessAdmin},
|
||||
{"POST /api/v1/admin/ips/scan", s.handleAdminScanFloatingIPs, accessAdmin},
|
||||
{"GET /api/v1/admin/ips/{ip}", s.handleAdminIPDetail, accessAdmin},
|
||||
{"POST /api/v1/admin/ips/{ip}/cancel", s.handleAdminCancelIP, accessAdmin},
|
||||
{"DELETE /api/v1/admin/ips/{ip}", s.handleAdminDeleteIP, accessAdmin},
|
||||
{"POST /api/v1/admin/ips/delete", s.handleAdminDeleteIPs, accessAdmin},
|
||||
{"POST /api/v1/admin/ips/clear", s.handleAdminClearQueue, accessAdmin},
|
||||
{"GET /api/v1/admin/validators", s.handleAdminValidators, accessAdmin},
|
||||
{"GET /api/v1/admin/auto-cycle", s.handleAdminGetAutoCycle, accessAdmin},
|
||||
{"PUT /api/v1/admin/auto-cycle", s.handleAdminPutAutoCycle, accessAdmin},
|
||||
{"POST /api/v1/admin/auto-cycle/start", s.handleAdminStartAutoCycle, accessAdmin},
|
||||
{"POST /api/v1/admin/auto-cycle/stop", s.handleAdminStopAutoCycle, accessAdmin},
|
||||
{"GET /api/v1/admin/registry", s.handleAdminRegistry, accessAdmin},
|
||||
{"GET /api/v1/admin/registry/{ip}", s.handleAdminRegistryHistory, accessAdmin},
|
||||
{"GET /api/v1/admin/config/validators", s.handleConfigListValidators, accessAdmin},
|
||||
{"POST /api/v1/admin/config/validators", s.handleConfigCreateValidator, accessAdmin},
|
||||
{"PUT /api/v1/admin/config/validators/{id}", s.handleConfigUpdateValidator, accessAdmin},
|
||||
{"DELETE /api/v1/admin/config/validators/{id}", s.handleConfigDeleteValidator, accessAdmin},
|
||||
{"GET /api/v1/admin/config/sites", s.handleConfigListSites, accessAdmin},
|
||||
{"PUT /api/v1/admin/config/sites/{index}", s.handleConfigPutSite, accessAdmin},
|
||||
{"DELETE /api/v1/admin/config/sites/{index}", s.handleConfigDeleteSite, accessAdmin},
|
||||
{"GET /api/v1/admin/config/targets", s.handleConfigListTargets, accessAdmin},
|
||||
{"PUT /api/v1/admin/config/targets/{group}", s.handleConfigPutTargetGroup, accessAdmin},
|
||||
{"DELETE /api/v1/admin/config/targets/{group}", s.handleConfigDeleteTargetGroup, accessAdmin},
|
||||
{"GET /api/v1/admin/config/check-types", s.handleConfigListCheckTypes, accessAdmin},
|
||||
{"PUT /api/v1/admin/config/check-types/{name}", s.handleConfigPutCheckType, accessAdmin},
|
||||
{"DELETE /api/v1/admin/config/check-types/{name}", s.handleConfigDeleteCheckType, accessAdmin},
|
||||
{"GET /api/v1/admin/config/orchestrator", s.handleConfigGetOrchestratorSettings, accessAdmin},
|
||||
{"PUT /api/v1/admin/config/orchestrator", s.handleConfigPutOrchestratorSettings, accessAdmin},
|
||||
{"GET /api/v1/admin/config/inbound-checks", s.handleConfigGetInboundChecks, accessAdmin},
|
||||
{"PUT /api/v1/admin/config/inbound-checks", s.handleConfigPutInboundChecks, accessAdmin},
|
||||
}
|
||||
}
|
||||
@@ -19,6 +19,7 @@ type Server struct {
|
||||
DB *db.DB
|
||||
Orch *orchestrator.Orchestrator
|
||||
Log *slog.Logger
|
||||
Auth Authenticator
|
||||
}
|
||||
|
||||
func New(d *db.DB, o *orchestrator.Orchestrator, log *slog.Logger) *Server {
|
||||
@@ -26,6 +27,7 @@ func New(d *db.DB, o *orchestrator.Orchestrator, log *slog.Logger) *Server {
|
||||
}
|
||||
|
||||
func (s *Server) Handler() http.Handler {
|
||||
s.Auth.Log = s.Log
|
||||
mux := http.NewServeMux()
|
||||
s.routes(mux)
|
||||
return loggingMiddleware(s.Log, mux)
|
||||
|
||||
@@ -44,6 +44,13 @@ func New(cfg *config.Prober, log *slog.Logger) *Prober {
|
||||
}
|
||||
}
|
||||
|
||||
// WithToken sets the bearer token sent to the Control API (and only to it:
|
||||
// probes use their own dialers).
|
||||
func (p *Prober) WithToken(token string) *Prober {
|
||||
p.client.Token = token
|
||||
return p
|
||||
}
|
||||
|
||||
func (p *Prober) Run(ctx context.Context) error {
|
||||
if err := p.registerWithRetry(ctx); err != nil {
|
||||
return err
|
||||
|
||||
@@ -208,3 +208,31 @@ func TestExtraChecksForPort(t *testing.T) {
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestWithTokenSendsBearerToControlAPI checks the prober's control-api
|
||||
// client carries the agent token on every request.
|
||||
func TestWithTokenSendsBearerToControlAPI(t *testing.T) {
|
||||
var mu sync.Mutex
|
||||
var auths []string
|
||||
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
mu.Lock()
|
||||
auths = append(auths, r.Header.Get("Authorization"))
|
||||
mu.Unlock()
|
||||
w.Write([]byte(`[]`))
|
||||
}))
|
||||
defer ts.Close()
|
||||
|
||||
p := New(&config.Prober{SiteID: "site-1", ControlAPIURL: ts.URL}, testLogger()).WithToken("agent-secret")
|
||||
p.pollOnce(context.Background())
|
||||
|
||||
mu.Lock()
|
||||
defer mu.Unlock()
|
||||
if len(auths) == 0 {
|
||||
t.Fatal("no requests reached control-api")
|
||||
}
|
||||
for _, a := range auths {
|
||||
if a != "Bearer agent-secret" {
|
||||
t.Fatalf("Authorization = %q, want bearer token", a)
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in new issue
Block a user