This commit is contained in:
@@ -0,0 +1,293 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"crypto/sha256"
|
||||
"crypto/subtle"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/coreos/go-oidc/v3/oidc"
|
||||
"golang.org/x/oauth2"
|
||||
)
|
||||
|
||||
const (
|
||||
defaultFlowLifetime = 10 * time.Minute
|
||||
)
|
||||
|
||||
type OIDCConfig struct {
|
||||
Issuer string
|
||||
ClientID string
|
||||
ClientSecret string
|
||||
RedirectURL string
|
||||
Scopes []string
|
||||
}
|
||||
|
||||
type Authorization struct {
|
||||
URL string
|
||||
State string
|
||||
Nonce string
|
||||
CodeVerifier string
|
||||
ExpiresAt time.Time
|
||||
}
|
||||
|
||||
func BeginAuthorization(endpoint oauth2.Endpoint, config OIDCConfig, now time.Time) (Authorization, error) {
|
||||
if endpoint.AuthURL == "" || config.ClientID == "" || config.RedirectURL == "" {
|
||||
return Authorization{}, errors.New("OIDC authorization configuration is incomplete")
|
||||
}
|
||||
state, err := randomToken()
|
||||
if err != nil {
|
||||
return Authorization{}, errors.New("generate authorization state")
|
||||
}
|
||||
nonce, err := randomToken()
|
||||
if err != nil {
|
||||
return Authorization{}, errors.New("generate authorization nonce")
|
||||
}
|
||||
verifier, err := randomToken()
|
||||
if err != nil {
|
||||
return Authorization{}, errors.New("generate PKCE verifier")
|
||||
}
|
||||
scopes := config.Scopes
|
||||
if len(scopes) == 0 {
|
||||
scopes = []string{oidc.ScopeOpenID, "profile", "email"}
|
||||
}
|
||||
oauthConfig := oauth2.Config{
|
||||
ClientID: config.ClientID,
|
||||
ClientSecret: config.ClientSecret,
|
||||
Endpoint: endpoint,
|
||||
RedirectURL: config.RedirectURL,
|
||||
Scopes: scopes,
|
||||
}
|
||||
authURL := oauthConfig.AuthCodeURL(state,
|
||||
oauth2.SetAuthURLParam("nonce", nonce),
|
||||
oauth2.SetAuthURLParam("code_challenge", pkceChallenge(verifier)),
|
||||
oauth2.SetAuthURLParam("code_challenge_method", "S256"),
|
||||
)
|
||||
return Authorization{URL: authURL, State: state, Nonce: nonce, CodeVerifier: verifier, ExpiresAt: now.Add(defaultFlowLifetime)}, nil
|
||||
}
|
||||
|
||||
func ValidateCallback(flow Authorization, state, code string, now time.Time) error {
|
||||
if flow.State == "" || subtle.ConstantTimeCompare([]byte(flow.State), []byte(state)) != 1 {
|
||||
return errors.New("OIDC state validation failed")
|
||||
}
|
||||
if flow.CodeVerifier == "" || flow.Nonce == "" {
|
||||
return errors.New("OIDC flow is incomplete")
|
||||
}
|
||||
if now.After(flow.ExpiresAt) {
|
||||
return errors.New("OIDC authorization expired")
|
||||
}
|
||||
if strings.TrimSpace(code) == "" {
|
||||
return errors.New("OIDC authorization code is required")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func Exchange(ctx context.Context, flow Authorization, config OIDCConfig, endpoint oauth2.Endpoint, state, code string) (*oauth2.Token, error) {
|
||||
if err := ValidateCallback(flow, state, code, time.Now()); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
oauthConfig := oauth2.Config{ClientID: config.ClientID, ClientSecret: config.ClientSecret, Endpoint: endpoint, RedirectURL: config.RedirectURL}
|
||||
return oauthConfig.Exchange(ctx, code, oauth2.SetAuthURLParam("code_verifier", flow.CodeVerifier))
|
||||
}
|
||||
|
||||
// Discovery is the provider metadata required to run one authorization code flow:
|
||||
// the authorization/token endpoints for BeginAuthorization and Exchange, and the
|
||||
// ID token verifier for VerifyIDToken. Resolve it once and reuse it.
|
||||
type Discovery struct {
|
||||
Endpoint oauth2.Endpoint
|
||||
Verifier *oidc.IDTokenVerifier
|
||||
}
|
||||
|
||||
func Discover(ctx context.Context, config OIDCConfig) (Discovery, error) {
|
||||
if config.Issuer == "" || config.ClientID == "" {
|
||||
return Discovery{}, errors.New("OIDC issuer and client ID are required")
|
||||
}
|
||||
provider, err := oidc.NewProvider(ctx, config.Issuer)
|
||||
if err != nil {
|
||||
return Discovery{}, fmt.Errorf("OIDC discovery failed")
|
||||
}
|
||||
return Discovery{Endpoint: provider.Endpoint(), Verifier: provider.Verifier(&oidc.Config{ClientID: config.ClientID})}, nil
|
||||
}
|
||||
|
||||
func NewVerifier(ctx context.Context, config OIDCConfig) (*oidc.IDTokenVerifier, error) {
|
||||
discovery, err := Discover(ctx, config)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return discovery.Verifier, nil
|
||||
}
|
||||
|
||||
func VerifyIDToken(ctx context.Context, verifier *oidc.IDTokenVerifier, rawToken, expectedNonce string) (*oidc.IDToken, error) {
|
||||
if verifier == nil || strings.TrimSpace(rawToken) == "" || expectedNonce == "" {
|
||||
return nil, errors.New("OIDC token verification input is incomplete")
|
||||
}
|
||||
token, err := verifier.Verify(ctx, rawToken)
|
||||
if err != nil {
|
||||
return nil, errors.New("OIDC token verification failed")
|
||||
}
|
||||
var claims struct {
|
||||
Nonce string `json:"nonce"`
|
||||
}
|
||||
if err := token.Claims(&claims); err != nil || subtle.ConstantTimeCompare([]byte(claims.Nonce), []byte(expectedNonce)) != 1 {
|
||||
return nil, errors.New("OIDC nonce validation failed")
|
||||
}
|
||||
return token, nil
|
||||
}
|
||||
|
||||
const (
|
||||
defaultGroupsClaim = "groups"
|
||||
maxIdentityGroups = 128
|
||||
)
|
||||
|
||||
// Identity is the bounded subset of verified ID token claims Pulse consumes.
|
||||
type Identity struct {
|
||||
Subject string
|
||||
Groups []string
|
||||
}
|
||||
|
||||
// ExtractIdentity reads the subject and the configured role claim from an already
|
||||
// verified ID token. The claim may be a list of strings or a single string; values
|
||||
// are trimmed, empty values dropped and the list bounded.
|
||||
func ExtractIdentity(token *oidc.IDToken, groupsClaim string) (Identity, error) {
|
||||
if token == nil {
|
||||
return Identity{}, errors.New("OIDC identity token is required")
|
||||
}
|
||||
if groupsClaim == "" {
|
||||
groupsClaim = defaultGroupsClaim
|
||||
}
|
||||
subject := strings.TrimSpace(token.Subject)
|
||||
if subject == "" {
|
||||
return Identity{}, errors.New("OIDC subject claim is required")
|
||||
}
|
||||
var claims map[string]json.RawMessage
|
||||
if err := token.Claims(&claims); err != nil {
|
||||
return Identity{}, errors.New("OIDC claims could not be read")
|
||||
}
|
||||
raw, ok := claims[groupsClaim]
|
||||
if !ok {
|
||||
return Identity{Subject: subject}, nil
|
||||
}
|
||||
groups, err := normalizeGroupClaim(raw)
|
||||
if err != nil {
|
||||
return Identity{}, err
|
||||
}
|
||||
return Identity{Subject: subject, Groups: groups}, nil
|
||||
}
|
||||
|
||||
func normalizeGroupClaim(raw json.RawMessage) ([]string, error) {
|
||||
var values []string
|
||||
if err := json.Unmarshal(raw, &values); err != nil {
|
||||
var single string
|
||||
if err := json.Unmarshal(raw, &single); err != nil {
|
||||
return nil, errors.New("OIDC role claim is malformed")
|
||||
}
|
||||
values = []string{single}
|
||||
}
|
||||
groups := make([]string, 0, len(values))
|
||||
for _, value := range values {
|
||||
trimmed := strings.TrimSpace(value)
|
||||
if trimmed == "" || len(groups) >= maxIdentityGroups {
|
||||
continue
|
||||
}
|
||||
groups = append(groups, trimmed)
|
||||
}
|
||||
return groups, nil
|
||||
}
|
||||
|
||||
func randomToken() (string, error) {
|
||||
bytes := make([]byte, 32)
|
||||
if _, err := rand.Read(bytes); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return base64.RawURLEncoding.EncodeToString(bytes), nil
|
||||
}
|
||||
|
||||
func pkceChallenge(verifier string) string {
|
||||
digest := sha256.Sum256([]byte(verifier))
|
||||
return base64.RawURLEncoding.EncodeToString(digest[:])
|
||||
}
|
||||
|
||||
type Role string
|
||||
|
||||
const (
|
||||
RoleViewer Role = "viewer"
|
||||
RoleOperator Role = "operator"
|
||||
RoleEditor Role = "editor"
|
||||
RoleAdministrator Role = "administrator"
|
||||
)
|
||||
|
||||
type Permission string
|
||||
|
||||
const (
|
||||
PermissionView Permission = "view"
|
||||
PermissionOperate Permission = "operate"
|
||||
PermissionEdit Permission = "edit"
|
||||
PermissionAdmin Permission = "admin"
|
||||
)
|
||||
|
||||
type Principal struct {
|
||||
Subject string
|
||||
Role Role
|
||||
}
|
||||
|
||||
func MapRoles(claims []string, mapping map[string]Role) (Role, error) {
|
||||
priority := map[Role]int{RoleViewer: 1, RoleOperator: 2, RoleEditor: 3, RoleAdministrator: 4}
|
||||
var selected Role
|
||||
for _, claim := range claims {
|
||||
role, ok := mapping[claim]
|
||||
if !ok || priority[role] <= priority[selected] {
|
||||
continue
|
||||
}
|
||||
selected = role
|
||||
}
|
||||
if selected == "" {
|
||||
return "", errors.New("no authorized Pulse role")
|
||||
}
|
||||
return selected, nil
|
||||
}
|
||||
|
||||
func Allows(role Role, permission Permission) bool {
|
||||
level := map[Role]int{RoleViewer: 1, RoleOperator: 2, RoleEditor: 3, RoleAdministrator: 4}[role]
|
||||
required := map[Permission]int{PermissionView: 1, PermissionOperate: 2, PermissionEdit: 3, PermissionAdmin: 4}[permission]
|
||||
return level > 0 && required > 0 && level >= required
|
||||
}
|
||||
|
||||
func Require(permission Permission, next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) {
|
||||
principal, ok := PrincipalFromContext(request.Context())
|
||||
if !ok {
|
||||
response.Header().Set("Cache-Control", "private, no-store")
|
||||
http.Error(response, "unauthorized", http.StatusUnauthorized)
|
||||
return
|
||||
}
|
||||
if !Allows(principal.Role, permission) {
|
||||
response.Header().Set("Cache-Control", "private, no-store")
|
||||
http.Error(response, "forbidden", http.StatusForbidden)
|
||||
return
|
||||
}
|
||||
next.ServeHTTP(response, request)
|
||||
})
|
||||
}
|
||||
|
||||
type contextKey struct{}
|
||||
|
||||
func WithPrincipal(ctx context.Context, principal Principal) context.Context {
|
||||
return context.WithValue(ctx, contextKey{}, principal)
|
||||
}
|
||||
|
||||
func PrincipalFromContext(ctx context.Context) (Principal, bool) {
|
||||
principal, ok := ctx.Value(contextKey{}).(Principal)
|
||||
return principal, ok && principal.Subject != ""
|
||||
}
|
||||
|
||||
type BreakGlassPolicy struct {
|
||||
Enabled bool
|
||||
}
|
||||
|
||||
func (policy BreakGlassPolicy) Allows() bool { return policy.Enabled }
|
||||
@@ -0,0 +1,227 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"net/url"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"golang.org/x/oauth2"
|
||||
)
|
||||
|
||||
func TestBeginAuthorizationUsesStateNonceAndPKCE(t *testing.T) {
|
||||
now := time.Date(2026, 8, 1, 12, 0, 0, 0, time.UTC)
|
||||
flow, err := BeginAuthorization(oauth2.Endpoint{AuthURL: "https://auth.example/authorize"}, OIDCConfig{ClientID: "pulse", RedirectURL: "https://pulse.example/callback"}, now)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
parsed, err := url.Parse(flow.URL)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
query := parsed.Query()
|
||||
for _, key := range []string{"state", "nonce", "code_challenge", "code_challenge_method"} {
|
||||
if query.Get(key) == "" {
|
||||
t.Fatalf("authorization URL missing %s", key)
|
||||
}
|
||||
}
|
||||
if query.Get("state") != flow.State || query.Get("nonce") != flow.Nonce || query.Get("code_challenge_method") != "S256" {
|
||||
t.Fatalf("authorization URL does not match flow: %s", flow.URL)
|
||||
}
|
||||
if query.Get("code_challenge") != pkceChallenge(flow.CodeVerifier) {
|
||||
t.Fatal("authorization URL has incorrect PKCE challenge")
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateCallbackRejectsStateNonceFlowAbuse(t *testing.T) {
|
||||
now := time.Now()
|
||||
flow := Authorization{State: "expected", Nonce: "nonce", CodeVerifier: "verifier", ExpiresAt: now.Add(time.Minute)}
|
||||
if err := ValidateCallback(flow, "wrong", "code", now); err == nil {
|
||||
t.Fatal("wrong state was accepted")
|
||||
}
|
||||
if err := ValidateCallback(flow, flow.State, "", now); err == nil {
|
||||
t.Fatal("empty code was accepted")
|
||||
}
|
||||
flow.ExpiresAt = now.Add(-time.Second)
|
||||
if err := ValidateCallback(flow, flow.State, "code", now); err == nil {
|
||||
t.Fatal("expired flow was accepted")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRoleMappingAndAuthorizationMatrix(t *testing.T) {
|
||||
mapping := map[string]Role{"pulse-view": RoleViewer, "pulse-operator": RoleOperator, "pulse-admin": RoleAdministrator}
|
||||
role, err := MapRoles([]string{"unrelated", "pulse-operator", "pulse-view"}, mapping)
|
||||
if err != nil || role != RoleOperator {
|
||||
t.Fatalf("role mapping = %q, %v", role, err)
|
||||
}
|
||||
if _, err := MapRoles([]string{"unrelated"}, mapping); err == nil {
|
||||
t.Fatal("unmapped claims were authorized")
|
||||
}
|
||||
for _, test := range []struct {
|
||||
role Role
|
||||
permission Permission
|
||||
allowed bool
|
||||
}{
|
||||
{RoleViewer, PermissionView, true}, {RoleViewer, PermissionEdit, false},
|
||||
{RoleOperator, PermissionOperate, true}, {RoleOperator, PermissionAdmin, false},
|
||||
{RoleEditor, PermissionEdit, true}, {RoleEditor, PermissionAdmin, false},
|
||||
{RoleAdministrator, PermissionAdmin, true},
|
||||
} {
|
||||
if got := Allows(test.role, test.permission); got != test.allowed {
|
||||
t.Errorf("Allows(%s, %s) = %v, want %v", test.role, test.permission, got, test.allowed)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestUnauthorizedPathsAreDenied(t *testing.T) {
|
||||
handler := Require(PermissionEdit, http.HandlerFunc(func(response http.ResponseWriter, _ *http.Request) { response.WriteHeader(http.StatusNoContent) }))
|
||||
for _, test := range []struct {
|
||||
name string
|
||||
ctx context.Context
|
||||
status int
|
||||
}{
|
||||
{"anonymous", context.Background(), http.StatusUnauthorized},
|
||||
{"viewer", WithPrincipal(context.Background(), Principal{Subject: "user-1", Role: RoleViewer}), http.StatusForbidden},
|
||||
{"editor", WithPrincipal(context.Background(), Principal{Subject: "user-1", Role: RoleEditor}), http.StatusNoContent},
|
||||
} {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
request := httptest.NewRequest(http.MethodGet, "/protected", nil).WithContext(test.ctx)
|
||||
response := httptest.NewRecorder()
|
||||
handler.ServeHTTP(response, request)
|
||||
if response.Code != test.status {
|
||||
t.Fatalf("status = %d, want %d", response.Code, test.status)
|
||||
}
|
||||
if test.status == http.StatusUnauthorized || test.status == http.StatusForbidden {
|
||||
if response.Header().Get("Cache-Control") != "private, no-store" {
|
||||
t.Fatalf("cache control = %q", response.Header().Get("Cache-Control"))
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestBreakGlassIsDisabledByDefault(t *testing.T) {
|
||||
if (BreakGlassPolicy{}).Allows() {
|
||||
t.Fatal("break-glass unexpectedly enabled")
|
||||
}
|
||||
}
|
||||
|
||||
func TestDiscoverResolvesEndpointAndVerifier(t *testing.T) {
|
||||
var issuer string
|
||||
server := httptest.NewServer(http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) {
|
||||
if request.URL.Path != "/.well-known/openid-configuration" {
|
||||
response.WriteHeader(http.StatusNotFound)
|
||||
return
|
||||
}
|
||||
response.Header().Set("Content-Type", "application/json")
|
||||
_, _ = response.Write([]byte(`{"issuer":"` + issuer + `","authorization_endpoint":"` + issuer + `/authorize","token_endpoint":"` + issuer + `/token","jwks_uri":"` + issuer + `/jwks","id_token_signing_alg_values_supported":["RS256"]}`))
|
||||
}))
|
||||
defer server.Close()
|
||||
issuer = server.URL
|
||||
|
||||
discovery, err := Discover(context.Background(), OIDCConfig{Issuer: issuer, ClientID: "pulse"})
|
||||
if err != nil {
|
||||
t.Fatalf("Discover: %v", err)
|
||||
}
|
||||
if discovery.Endpoint.AuthURL != issuer+"/authorize" || discovery.Endpoint.TokenURL != issuer+"/token" || discovery.Verifier == nil {
|
||||
t.Fatalf("discovery = %#v", discovery.Endpoint)
|
||||
}
|
||||
if _, err := NewVerifier(context.Background(), OIDCConfig{Issuer: issuer, ClientID: "pulse"}); err != nil {
|
||||
t.Fatalf("NewVerifier: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDiscoverRejectsIncompleteOrUnreachableIssuer(t *testing.T) {
|
||||
unreachable := httptest.NewServer(http.NewServeMux())
|
||||
unreachable.Close()
|
||||
for _, test := range []struct {
|
||||
name string
|
||||
config OIDCConfig
|
||||
}{
|
||||
{"missing issuer", OIDCConfig{ClientID: "pulse"}},
|
||||
{"missing client id", OIDCConfig{Issuer: "https://idp.example"}},
|
||||
{"unreachable issuer", OIDCConfig{Issuer: unreachable.URL, ClientID: "pulse"}},
|
||||
} {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
discovery, err := Discover(context.Background(), test.config)
|
||||
if err == nil {
|
||||
t.Fatal("incomplete configuration was accepted")
|
||||
}
|
||||
if discovery.Verifier != nil {
|
||||
t.Fatal("a verifier was returned with an error")
|
||||
}
|
||||
if strings.Contains(err.Error(), test.config.Issuer) && test.config.Issuer != "" {
|
||||
t.Fatalf("error leaks the issuer: %v", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestExtractIdentityRequiresToken(t *testing.T) {
|
||||
if _, err := ExtractIdentity(nil, "groups"); err == nil {
|
||||
t.Fatal("nil token was accepted")
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizeGroupClaimBoundsAndShapes(t *testing.T) {
|
||||
many, err := json.Marshal(make([]string, maxIdentityGroups+50))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for _, test := range []struct {
|
||||
name string
|
||||
raw string
|
||||
want []string
|
||||
wantErr bool
|
||||
}{
|
||||
{name: "list", raw: `["pulse-admin"," pulse-view ",""]`, want: []string{"pulse-admin", "pulse-view"}},
|
||||
{name: "single string", raw: `"pulse-admin"`, want: []string{"pulse-admin"}},
|
||||
{name: "empty list", raw: `[]`, want: []string{}},
|
||||
{name: "object", raw: `{"groups":["pulse-admin"]}`, wantErr: true},
|
||||
{name: "number", raw: `7`, wantErr: true},
|
||||
} {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
groups, err := normalizeGroupClaim(json.RawMessage(test.raw))
|
||||
if (err != nil) != test.wantErr {
|
||||
t.Fatalf("err = %v, wantErr = %v", err, test.wantErr)
|
||||
}
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
if len(groups) != len(test.want) {
|
||||
t.Fatalf("groups = %#v, want %#v", groups, test.want)
|
||||
}
|
||||
for index, value := range test.want {
|
||||
if groups[index] != value {
|
||||
t.Fatalf("groups = %#v, want %#v", groups, test.want)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
bounded, err := normalizeGroupClaim(many)
|
||||
if err != nil {
|
||||
t.Fatalf("normalizeGroupClaim: %v", err)
|
||||
}
|
||||
if len(bounded) != 0 {
|
||||
t.Fatalf("blank group values were kept: %d", len(bounded))
|
||||
}
|
||||
filled := make([]string, maxIdentityGroups+50)
|
||||
for index := range filled {
|
||||
filled[index] = "group"
|
||||
}
|
||||
encoded, err := json.Marshal(filled)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
capped, err := normalizeGroupClaim(encoded)
|
||||
if err != nil {
|
||||
t.Fatalf("normalizeGroupClaim: %v", err)
|
||||
}
|
||||
if len(capped) != maxIdentityGroups {
|
||||
t.Fatalf("groups = %d, want %d", len(capped), maxIdentityGroups)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,208 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"net/http"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
type session struct {
|
||||
principal Principal
|
||||
issuedAt time.Time
|
||||
expiresAt time.Time
|
||||
absoluteExpiresAt time.Time
|
||||
context context.Context
|
||||
cancel context.CancelFunc
|
||||
}
|
||||
|
||||
// SessionAuthentication carries the principal and the revocable lifetime of
|
||||
// the authenticated browser session. Long-lived transports must derive their
|
||||
// lifecycle from Context so logout and the absolute deadline remain effective
|
||||
// after an HTTP upgrade.
|
||||
type SessionAuthentication struct {
|
||||
Principal Principal
|
||||
Context context.Context
|
||||
}
|
||||
|
||||
type SessionManager struct {
|
||||
mu sync.Mutex
|
||||
sessions map[string]session
|
||||
CookieName string
|
||||
TTL time.Duration
|
||||
AbsoluteTTL time.Duration
|
||||
RenewBefore time.Duration
|
||||
Secure bool
|
||||
MaxSessions int
|
||||
MaxSessionsPerSubject int
|
||||
}
|
||||
|
||||
func NewSessionManager(cookieName string, ttl time.Duration, secure bool) *SessionManager {
|
||||
if cookieName == "" {
|
||||
cookieName = "pulse_session"
|
||||
}
|
||||
if ttl <= 0 {
|
||||
ttl = 8 * time.Hour
|
||||
}
|
||||
return &SessionManager{sessions: make(map[string]session), CookieName: cookieName, TTL: ttl, AbsoluteTTL: ttl, Secure: secure, MaxSessions: 4096, MaxSessionsPerSubject: 8}
|
||||
}
|
||||
|
||||
// NewSlidingSessionManager creates an idle-expiring browser session with a
|
||||
// separate absolute lifetime. Successful authenticated requests renew the idle
|
||||
// deadline once half of the idle lifetime has elapsed, but never beyond the
|
||||
// absolute deadline. Both lifetimes remain finite and the opaque token stays in
|
||||
// an HttpOnly cookie.
|
||||
func NewSlidingSessionManager(cookieName string, idleTTL, absoluteTTL time.Duration, secure bool) *SessionManager {
|
||||
manager := NewSessionManager(cookieName, idleTTL, secure)
|
||||
if absoluteTTL < idleTTL {
|
||||
absoluteTTL = idleTTL
|
||||
}
|
||||
manager.AbsoluteTTL = absoluteTTL
|
||||
manager.RenewBefore = idleTTL / 2
|
||||
return manager
|
||||
}
|
||||
|
||||
func (manager *SessionManager) Issue(response http.ResponseWriter, principal Principal, now time.Time) error {
|
||||
if principal.Subject == "" || !Allows(principal.Role, PermissionView) {
|
||||
return errors.New("session principal is invalid")
|
||||
}
|
||||
token, err := randomToken()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
absoluteExpiresAt := now.Add(manager.AbsoluteTTL)
|
||||
expiresAt := earliest(now.Add(manager.TTL), absoluteExpiresAt)
|
||||
sessionContext, cancel := context.WithDeadline(context.Background(), absoluteExpiresAt)
|
||||
manager.mu.Lock()
|
||||
manager.purgeExpiredLocked(now)
|
||||
manager.enforceSubjectLimitLocked(principal.Subject)
|
||||
if manager.MaxSessions > 0 && len(manager.sessions) >= manager.MaxSessions {
|
||||
manager.mu.Unlock()
|
||||
cancel()
|
||||
return errors.New("session capacity reached")
|
||||
}
|
||||
manager.sessions[hashToken(token)] = session{principal: principal, issuedAt: now, expiresAt: expiresAt, absoluteExpiresAt: absoluteExpiresAt, context: sessionContext, cancel: cancel}
|
||||
manager.mu.Unlock()
|
||||
manager.setCookie(response, token, now, expiresAt)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (manager *SessionManager) Principal(request *http.Request, now time.Time) (Principal, bool) {
|
||||
authentication, ok := manager.authenticate(nil, request, now)
|
||||
return authentication.Principal, ok
|
||||
}
|
||||
|
||||
// Authenticate validates the session and renews an active sliding session when
|
||||
// it enters its renewal window. The token is deliberately stable: concurrent
|
||||
// API requests cannot invalidate each other, while Clear still revokes it
|
||||
// immediately server-side.
|
||||
func (manager *SessionManager) Authenticate(response http.ResponseWriter, request *http.Request, now time.Time) (Principal, bool) {
|
||||
authentication, ok := manager.authenticate(response, request, now)
|
||||
return authentication.Principal, ok
|
||||
}
|
||||
|
||||
// AuthenticateSession validates and renews the cookie while exposing the
|
||||
// revocable session context to middleware that serves long-lived transports.
|
||||
func (manager *SessionManager) AuthenticateSession(response http.ResponseWriter, request *http.Request, now time.Time) (SessionAuthentication, bool) {
|
||||
return manager.authenticate(response, request, now)
|
||||
}
|
||||
|
||||
func (manager *SessionManager) authenticate(response http.ResponseWriter, request *http.Request, now time.Time) (SessionAuthentication, bool) {
|
||||
cookie, err := request.Cookie(manager.CookieName)
|
||||
if err != nil || cookie.Value == "" {
|
||||
return SessionAuthentication{}, false
|
||||
}
|
||||
manager.mu.Lock()
|
||||
defer manager.mu.Unlock()
|
||||
manager.purgeExpiredLocked(now)
|
||||
stored, ok := manager.sessions[hashToken(cookie.Value)]
|
||||
if !ok {
|
||||
return SessionAuthentication{}, false
|
||||
}
|
||||
if !now.Before(stored.expiresAt) || !now.Before(stored.absoluteExpiresAt) {
|
||||
stored.cancel()
|
||||
delete(manager.sessions, hashToken(cookie.Value))
|
||||
return SessionAuthentication{}, false
|
||||
}
|
||||
if response != nil && manager.RenewBefore > 0 && stored.expiresAt.Sub(now) <= manager.RenewBefore {
|
||||
renewed := earliest(now.Add(manager.TTL), stored.absoluteExpiresAt)
|
||||
if renewed.After(stored.expiresAt) {
|
||||
stored.expiresAt = renewed
|
||||
manager.sessions[hashToken(cookie.Value)] = stored
|
||||
manager.setCookie(response, cookie.Value, now, renewed)
|
||||
}
|
||||
}
|
||||
return SessionAuthentication{Principal: stored.principal, Context: stored.context}, true
|
||||
}
|
||||
|
||||
func (manager *SessionManager) Clear(response http.ResponseWriter, request *http.Request) {
|
||||
if cookie, err := request.Cookie(manager.CookieName); err == nil {
|
||||
manager.mu.Lock()
|
||||
key := hashToken(cookie.Value)
|
||||
if stored, ok := manager.sessions[key]; ok {
|
||||
stored.cancel()
|
||||
delete(manager.sessions, key)
|
||||
}
|
||||
manager.mu.Unlock()
|
||||
}
|
||||
http.SetCookie(response, &http.Cookie{Name: manager.CookieName, Value: "", Path: "/", MaxAge: -1, HttpOnly: true, Secure: manager.Secure, SameSite: http.SameSiteLaxMode})
|
||||
}
|
||||
|
||||
func (manager *SessionManager) purgeExpiredLocked(now time.Time) {
|
||||
for key, stored := range manager.sessions {
|
||||
if !now.Before(stored.expiresAt) || !now.Before(stored.absoluteExpiresAt) {
|
||||
stored.cancel()
|
||||
delete(manager.sessions, key)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (manager *SessionManager) enforceSubjectLimitLocked(subject string) {
|
||||
if manager.MaxSessionsPerSubject <= 0 {
|
||||
return
|
||||
}
|
||||
for {
|
||||
count := 0
|
||||
oldestKey := ""
|
||||
var oldest time.Time
|
||||
for key, stored := range manager.sessions {
|
||||
if stored.principal.Subject != subject {
|
||||
continue
|
||||
}
|
||||
count++
|
||||
if oldestKey == "" || stored.issuedAt.Before(oldest) {
|
||||
oldestKey = key
|
||||
oldest = stored.issuedAt
|
||||
}
|
||||
}
|
||||
if count < manager.MaxSessionsPerSubject || oldestKey == "" {
|
||||
return
|
||||
}
|
||||
stored := manager.sessions[oldestKey]
|
||||
stored.cancel()
|
||||
delete(manager.sessions, oldestKey)
|
||||
}
|
||||
}
|
||||
|
||||
func hashToken(token string) string {
|
||||
digest := sha256.Sum256([]byte(token))
|
||||
return hex.EncodeToString(digest[:])
|
||||
}
|
||||
|
||||
func (manager *SessionManager) setCookie(response http.ResponseWriter, token string, now, expiresAt time.Time) {
|
||||
maxAge := int(expiresAt.Sub(now).Seconds())
|
||||
if maxAge < 1 {
|
||||
maxAge = 1
|
||||
}
|
||||
http.SetCookie(response, &http.Cookie{Name: manager.CookieName, Value: token, Path: "/", Expires: expiresAt, MaxAge: maxAge, HttpOnly: true, Secure: manager.Secure, SameSite: http.SameSiteLaxMode})
|
||||
}
|
||||
|
||||
func earliest(first, second time.Time) time.Time {
|
||||
if first.Before(second) {
|
||||
return first
|
||||
}
|
||||
return second
|
||||
}
|
||||
@@ -0,0 +1,197 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestSessionIssueReadAndClear(t *testing.T) {
|
||||
manager := NewSessionManager("pulse_test_session", time.Hour, true)
|
||||
now := time.Now()
|
||||
response := httptest.NewRecorder()
|
||||
principal := Principal{Subject: "subject-1", Role: RoleViewer}
|
||||
if err := manager.Issue(response, principal, now); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if response.Header().Get("Set-Cookie") == "" {
|
||||
t.Fatal("session cookie was not set")
|
||||
}
|
||||
if !response.Result().Cookies()[0].HttpOnly || !response.Result().Cookies()[0].Secure {
|
||||
t.Fatal("session cookie is not hardened")
|
||||
}
|
||||
request := httptest.NewRequest("GET", "/", nil)
|
||||
for _, cookie := range response.Result().Cookies() {
|
||||
request.AddCookie(cookie)
|
||||
}
|
||||
got, ok := manager.Principal(request, now.Add(time.Minute))
|
||||
if !ok || got != principal {
|
||||
t.Fatalf("session principal = %#v, %v", got, ok)
|
||||
}
|
||||
clearResponse := httptest.NewRecorder()
|
||||
manager.Clear(clearResponse, request)
|
||||
if _, ok := manager.Principal(request, now.Add(time.Minute)); ok {
|
||||
t.Fatal("cleared session remained valid")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSessionExpires(t *testing.T) {
|
||||
manager := NewSessionManager("pulse_test_session", time.Minute, false)
|
||||
now := time.Now()
|
||||
response := httptest.NewRecorder()
|
||||
if err := manager.Issue(response, Principal{Subject: "subject-1", Role: RoleViewer}, now); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
request := httptest.NewRequest("GET", "/", nil)
|
||||
request.AddCookie(response.Result().Cookies()[0])
|
||||
if _, ok := manager.Principal(request, now.Add(2*time.Minute)); ok {
|
||||
t.Fatal("expired session remained valid")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSlidingSessionRenewsIdleDeadlineButHonorsAbsoluteExpiry(t *testing.T) {
|
||||
manager := NewSlidingSessionManager("pulse_test_session", time.Minute, 3*time.Minute, true)
|
||||
now := time.Now().UTC().Truncate(time.Second)
|
||||
issued := httptest.NewRecorder()
|
||||
if err := manager.Issue(issued, Principal{Subject: "wallboard", Role: RoleViewer}, now); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
cookie := issued.Result().Cookies()[0]
|
||||
request := httptest.NewRequest(http.MethodGet, "/api/v1/system/status", nil)
|
||||
request.AddCookie(cookie)
|
||||
|
||||
beforeWindow := httptest.NewRecorder()
|
||||
if _, ok := manager.Authenticate(beforeWindow, request, now.Add(20*time.Second)); !ok {
|
||||
t.Fatal("active session was rejected before renewal window")
|
||||
}
|
||||
if beforeWindow.Header().Get("Set-Cookie") != "" {
|
||||
t.Fatal("session renewed before entering the bounded renewal window")
|
||||
}
|
||||
|
||||
for _, offset := range []time.Duration{40 * time.Second, 80 * time.Second, 130 * time.Second} {
|
||||
response := httptest.NewRecorder()
|
||||
if _, ok := manager.Authenticate(response, request, now.Add(offset)); !ok {
|
||||
t.Fatalf("active session was rejected at %s", offset)
|
||||
}
|
||||
renewed := response.Result().Cookies()
|
||||
if len(renewed) != 1 || renewed[0].Value != cookie.Value || renewed[0].Expires.After(now.Add(3*time.Minute)) {
|
||||
t.Fatalf("unsafe renewal at %s: %#v", offset, renewed)
|
||||
}
|
||||
}
|
||||
|
||||
if _, ok := manager.Authenticate(httptest.NewRecorder(), request, now.Add(3*time.Minute)); ok {
|
||||
t.Fatal("sliding session exceeded its absolute expiry")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSlidingSessionConcurrentRenewalKeepsTokenUsable(t *testing.T) {
|
||||
manager := NewSlidingSessionManager("pulse_test_session", time.Minute, time.Hour, false)
|
||||
now := time.Now().UTC()
|
||||
issued := httptest.NewRecorder()
|
||||
if err := manager.Issue(issued, Principal{Subject: "wallboard", Role: RoleViewer}, now); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
cookie := issued.Result().Cookies()[0]
|
||||
const workers = 24
|
||||
var wait sync.WaitGroup
|
||||
errors := make(chan string, workers)
|
||||
for index := 0; index < workers; index++ {
|
||||
wait.Add(1)
|
||||
go func() {
|
||||
defer wait.Done()
|
||||
request := httptest.NewRequest(http.MethodGet, "/api/v1/dashboards", nil)
|
||||
request.AddCookie(cookie)
|
||||
if _, ok := manager.Authenticate(httptest.NewRecorder(), request, now.Add(40*time.Second)); !ok {
|
||||
errors <- "concurrent renewal rejected a valid token"
|
||||
}
|
||||
}()
|
||||
}
|
||||
wait.Wait()
|
||||
close(errors)
|
||||
for message := range errors {
|
||||
t.Error(message)
|
||||
}
|
||||
request := httptest.NewRequest(http.MethodGet, "/api/v1/dashboards", nil)
|
||||
request.AddCookie(cookie)
|
||||
if _, ok := manager.Principal(request, now.Add(90*time.Second)); !ok {
|
||||
t.Fatal("stable token was invalidated by concurrent renewal")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSessionAuthenticationContextIsRevokedByClear(t *testing.T) {
|
||||
manager := NewSlidingSessionManager("pulse_test_session", time.Minute, time.Hour, true)
|
||||
now := time.Now().UTC()
|
||||
issued := httptest.NewRecorder()
|
||||
if err := manager.Issue(issued, Principal{Subject: "viewer", Role: RoleViewer}, now); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
request := httptest.NewRequest(http.MethodGet, "/api/v1/live", nil)
|
||||
request.AddCookie(issued.Result().Cookies()[0])
|
||||
authentication, ok := manager.AuthenticateSession(httptest.NewRecorder(), request, now.Add(time.Second))
|
||||
if !ok || authentication.Context == nil {
|
||||
t.Fatal("session authentication context was not returned")
|
||||
}
|
||||
manager.Clear(httptest.NewRecorder(), request)
|
||||
select {
|
||||
case <-authentication.Context.Done():
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("cleared session context remained active")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSessionAuthenticationContextEndsAtAbsoluteExpiry(t *testing.T) {
|
||||
manager := NewSlidingSessionManager("pulse_test_session", 25*time.Millisecond, 25*time.Millisecond, true)
|
||||
now := time.Now().UTC()
|
||||
issued := httptest.NewRecorder()
|
||||
if err := manager.Issue(issued, Principal{Subject: "viewer", Role: RoleViewer}, now); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
request := httptest.NewRequest(http.MethodGet, "/api/v1/live", nil)
|
||||
request.AddCookie(issued.Result().Cookies()[0])
|
||||
authentication, ok := manager.AuthenticateSession(httptest.NewRecorder(), request, now)
|
||||
if !ok {
|
||||
t.Fatal("new session was rejected")
|
||||
}
|
||||
select {
|
||||
case <-authentication.Context.Done():
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("session context exceeded its absolute deadline")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSessionStoreEvictsOldestSessionsPerSubject(t *testing.T) {
|
||||
manager := NewSlidingSessionManager("pulse_test_session", time.Hour, 24*time.Hour, true)
|
||||
manager.MaxSessionsPerSubject = 3
|
||||
now := time.Now().UTC()
|
||||
for index := 0; index < 12; index++ {
|
||||
if err := manager.Issue(httptest.NewRecorder(), Principal{Subject: "viewer", Role: RoleViewer}, now.Add(time.Duration(index)*time.Second)); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
if got := len(manager.sessions); got != 3 {
|
||||
t.Fatalf("session store size = %d, want 3", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSessionIssuePurgesExpiredEntriesAndHonorsGlobalCapacity(t *testing.T) {
|
||||
manager := NewSessionManager("pulse_test_session", time.Minute, true)
|
||||
manager.MaxSessions = 2
|
||||
manager.MaxSessionsPerSubject = 2
|
||||
now := time.Now().UTC()
|
||||
for _, subject := range []string{"viewer-1", "viewer-2"} {
|
||||
if err := manager.Issue(httptest.NewRecorder(), Principal{Subject: subject, Role: RoleViewer}, now); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
if err := manager.Issue(httptest.NewRecorder(), Principal{Subject: "viewer-3", Role: RoleViewer}, now); err == nil {
|
||||
t.Fatal("session capacity was not enforced")
|
||||
}
|
||||
if err := manager.Issue(httptest.NewRecorder(), Principal{Subject: "viewer-3", Role: RoleViewer}, now.Add(2*time.Minute)); err != nil {
|
||||
t.Fatalf("expired sessions were not purged: %v", err)
|
||||
}
|
||||
if got := len(manager.sessions); got != 1 {
|
||||
t.Fatalf("session store size after purge = %d, want 1", got)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user