420 lines
9.3 KiB
Go
420 lines
9.3 KiB
Go
package main
|
|
|
|
import (
|
|
"context"
|
|
"crypto/rand"
|
|
"crypto/rsa"
|
|
"database/sql"
|
|
"encoding/base64"
|
|
"encoding/json"
|
|
"errors"
|
|
"fmt"
|
|
"io"
|
|
"math/big"
|
|
"net/http"
|
|
"net/url"
|
|
"strings"
|
|
"sync"
|
|
"time"
|
|
|
|
"github.com/golang-jwt/jwt/v5"
|
|
)
|
|
|
|
type contextKey string
|
|
|
|
const userContextKey contextKey = "auth_user"
|
|
|
|
const (
|
|
sessionCookieName = "session_token"
|
|
stateCookieName = "oidc_state"
|
|
)
|
|
|
|
type User struct {
|
|
Sub string
|
|
Email string
|
|
Name string
|
|
}
|
|
|
|
type Auth struct {
|
|
issuer string
|
|
audience string
|
|
jwksURL string
|
|
authURL string
|
|
tokenURL string
|
|
clientID string
|
|
clientSecret string
|
|
redirectURL string
|
|
logoutURL string
|
|
devMode bool
|
|
db *sql.DB
|
|
|
|
mu sync.RWMutex
|
|
keyByKid map[string]*rsa.PublicKey
|
|
lastSyncAt time.Time
|
|
}
|
|
|
|
func NewAuth(cfg Config) (*Auth, error) {
|
|
a := &Auth{
|
|
issuer: cfg.OIDCIssuer,
|
|
audience: cfg.OIDCAudience,
|
|
jwksURL: cfg.OIDCJWKSURL,
|
|
authURL: cfg.OIDCAuthURL,
|
|
tokenURL: cfg.OIDCTokenURL,
|
|
clientID: cfg.OIDCClientID,
|
|
clientSecret: cfg.OIDCClientSecret,
|
|
redirectURL: cfg.OIDCRedirect,
|
|
logoutURL: cfg.OIDCLogoutURL,
|
|
devMode: cfg.DevMode,
|
|
keyByKid: map[string]*rsa.PublicKey{},
|
|
lastSyncAt: time.Time{},
|
|
}
|
|
if a.devMode {
|
|
return a, nil
|
|
}
|
|
if err := a.refreshKeys(context.Background()); err != nil {
|
|
return nil, err
|
|
}
|
|
return a, nil
|
|
}
|
|
|
|
func (a *Auth) WithDB(db *sql.DB) {
|
|
a.db = db
|
|
}
|
|
|
|
func (a *Auth) HandleLogin(w http.ResponseWriter, r *http.Request) {
|
|
if a.devMode {
|
|
http.Redirect(w, r, "/dashboard", http.StatusFound)
|
|
return
|
|
}
|
|
state, err := randomToken(24)
|
|
if err != nil {
|
|
http.Error(w, "failed to initialize login", http.StatusInternalServerError)
|
|
return
|
|
}
|
|
http.SetCookie(w, &http.Cookie{
|
|
Name: stateCookieName,
|
|
Value: state,
|
|
Path: "/",
|
|
HttpOnly: true,
|
|
Secure: false,
|
|
SameSite: http.SameSiteLaxMode,
|
|
MaxAge: 600,
|
|
})
|
|
|
|
q := url.Values{}
|
|
q.Set("client_id", a.clientID)
|
|
q.Set("response_type", "code")
|
|
q.Set("scope", "openid profile email")
|
|
q.Set("redirect_uri", a.redirectURL)
|
|
q.Set("state", state)
|
|
http.Redirect(w, r, a.authURL+"?"+q.Encode(), http.StatusFound)
|
|
}
|
|
|
|
func (a *Auth) HandleCallback(w http.ResponseWriter, r *http.Request) {
|
|
if a.devMode {
|
|
http.Redirect(w, r, "/dashboard", http.StatusFound)
|
|
return
|
|
}
|
|
|
|
stateCookie, err := r.Cookie(stateCookieName)
|
|
if err != nil {
|
|
http.Error(w, "missing state cookie", http.StatusBadRequest)
|
|
return
|
|
}
|
|
stateParam := r.URL.Query().Get("state")
|
|
if stateParam == "" || stateParam != stateCookie.Value {
|
|
http.Error(w, "invalid state", http.StatusBadRequest)
|
|
return
|
|
}
|
|
|
|
code := r.URL.Query().Get("code")
|
|
if code == "" {
|
|
http.Error(w, "missing authorization code", http.StatusBadRequest)
|
|
return
|
|
}
|
|
|
|
token, err := a.exchangeCode(r.Context(), code)
|
|
if err != nil {
|
|
http.Error(w, "token exchange failed", http.StatusUnauthorized)
|
|
return
|
|
}
|
|
|
|
if _, err := a.userFromToken(token); err != nil {
|
|
http.Error(w, "invalid access token", http.StatusUnauthorized)
|
|
return
|
|
}
|
|
|
|
http.SetCookie(w, &http.Cookie{
|
|
Name: sessionCookieName,
|
|
Value: token,
|
|
Path: "/",
|
|
HttpOnly: true,
|
|
Secure: false,
|
|
SameSite: http.SameSiteLaxMode,
|
|
MaxAge: 3600,
|
|
})
|
|
http.SetCookie(w, &http.Cookie{
|
|
Name: stateCookieName,
|
|
Value: "",
|
|
Path: "/",
|
|
HttpOnly: true,
|
|
Secure: false,
|
|
SameSite: http.SameSiteLaxMode,
|
|
MaxAge: -1,
|
|
})
|
|
http.Redirect(w, r, "/dashboard", http.StatusFound)
|
|
}
|
|
|
|
func (a *Auth) HandleLogout(w http.ResponseWriter, r *http.Request) {
|
|
http.SetCookie(w, &http.Cookie{
|
|
Name: sessionCookieName,
|
|
Value: "",
|
|
Path: "/",
|
|
HttpOnly: true,
|
|
Secure: false,
|
|
SameSite: http.SameSiteLaxMode,
|
|
MaxAge: -1,
|
|
})
|
|
if !a.devMode && a.logoutURL != "" {
|
|
http.Redirect(w, r, a.logoutURL, http.StatusFound)
|
|
return
|
|
}
|
|
http.Redirect(w, r, "/", http.StatusFound)
|
|
}
|
|
|
|
func (a *Auth) RequireAuth(next http.Handler) http.Handler {
|
|
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
var (
|
|
user User
|
|
err error
|
|
)
|
|
|
|
if a.devMode {
|
|
user = User{
|
|
Sub: "demo-user-001",
|
|
Email: "demo@example.com",
|
|
Name: "Demo User",
|
|
}
|
|
} else {
|
|
cookie, cErr := r.Cookie(sessionCookieName)
|
|
if cErr != nil || strings.TrimSpace(cookie.Value) == "" {
|
|
http.Redirect(w, r, "/", http.StatusFound)
|
|
return
|
|
}
|
|
user, err = a.userFromToken(cookie.Value)
|
|
if err != nil {
|
|
http.Redirect(w, r, "/", http.StatusFound)
|
|
return
|
|
}
|
|
}
|
|
|
|
if a.db != nil {
|
|
_, _ = a.db.ExecContext(r.Context(), `
|
|
INSERT INTO users (sub, email, name)
|
|
VALUES ($1, $2, $3)
|
|
ON CONFLICT (sub) DO UPDATE
|
|
SET email = EXCLUDED.email, name = EXCLUDED.name
|
|
`, user.Sub, user.Email, user.Name)
|
|
}
|
|
|
|
ctx := context.WithValue(r.Context(), userContextKey, user)
|
|
next.ServeHTTP(w, r.WithContext(ctx))
|
|
})
|
|
}
|
|
|
|
func MustUserFromContext(ctx context.Context) User {
|
|
v := ctx.Value(userContextKey)
|
|
if v == nil {
|
|
return User{}
|
|
}
|
|
if u, ok := v.(User); ok {
|
|
return u
|
|
}
|
|
return User{}
|
|
}
|
|
|
|
func (a *Auth) exchangeCode(ctx context.Context, code string) (string, error) {
|
|
form := url.Values{}
|
|
form.Set("grant_type", "authorization_code")
|
|
form.Set("code", code)
|
|
form.Set("client_id", a.clientID)
|
|
form.Set("client_secret", a.clientSecret)
|
|
form.Set("redirect_uri", a.redirectURL)
|
|
|
|
req, err := http.NewRequestWithContext(ctx, http.MethodPost, a.tokenURL, strings.NewReader(form.Encode()))
|
|
if err != nil {
|
|
return "", err
|
|
}
|
|
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
|
|
|
|
resp, err := http.DefaultClient.Do(req)
|
|
if err != nil {
|
|
return "", err
|
|
}
|
|
defer resp.Body.Close()
|
|
if resp.StatusCode != http.StatusOK {
|
|
body, _ := io.ReadAll(io.LimitReader(resp.Body, 512))
|
|
return "", fmt.Errorf("token endpoint failed: %s %s", resp.Status, string(body))
|
|
}
|
|
|
|
var tr struct {
|
|
AccessToken string `json:"access_token"`
|
|
}
|
|
if err := json.NewDecoder(resp.Body).Decode(&tr); err != nil {
|
|
return "", err
|
|
}
|
|
if tr.AccessToken == "" {
|
|
return "", errors.New("missing access_token")
|
|
}
|
|
return tr.AccessToken, nil
|
|
}
|
|
|
|
func (a *Auth) userFromToken(token string) (User, error) {
|
|
parser := jwt.NewParser(jwt.WithValidMethods([]string{"RS256"}))
|
|
parsed, err := parser.Parse(token, func(t *jwt.Token) (any, error) {
|
|
kid, _ := t.Header["kid"].(string)
|
|
if kid == "" {
|
|
return nil, errors.New("missing kid")
|
|
}
|
|
key, err := a.lookupKey(context.Background(), kid)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
return key, nil
|
|
})
|
|
if err != nil || !parsed.Valid {
|
|
return User{}, errors.New("invalid token")
|
|
}
|
|
|
|
claims, ok := parsed.Claims.(jwt.MapClaims)
|
|
if !ok {
|
|
return User{}, errors.New("invalid claims")
|
|
}
|
|
if iss, _ := claims["iss"].(string); iss != a.issuer {
|
|
return User{}, errors.New("invalid issuer")
|
|
}
|
|
if !audienceMatches(claims["aud"], a.audience) {
|
|
return User{}, errors.New("invalid audience")
|
|
}
|
|
|
|
sub, _ := claims["sub"].(string)
|
|
email, _ := claims["email"].(string)
|
|
name, _ := claims["name"].(string)
|
|
if sub == "" {
|
|
return User{}, errors.New("missing sub")
|
|
}
|
|
return User{Sub: sub, Email: email, Name: name}, nil
|
|
}
|
|
|
|
func (a *Auth) lookupKey(ctx context.Context, kid string) (*rsa.PublicKey, error) {
|
|
a.mu.RLock()
|
|
key := a.keyByKid[kid]
|
|
last := a.lastSyncAt
|
|
a.mu.RUnlock()
|
|
if key != nil {
|
|
return key, nil
|
|
}
|
|
if time.Since(last) > time.Minute {
|
|
if err := a.refreshKeys(ctx); err != nil {
|
|
return nil, err
|
|
}
|
|
a.mu.RLock()
|
|
defer a.mu.RUnlock()
|
|
if refreshed := a.keyByKid[kid]; refreshed != nil {
|
|
return refreshed, nil
|
|
}
|
|
}
|
|
return nil, fmt.Errorf("no key for kid %s", kid)
|
|
}
|
|
|
|
type jwksDoc struct {
|
|
Keys []struct {
|
|
Kid string `json:"kid"`
|
|
Kty string `json:"kty"`
|
|
N string `json:"n"`
|
|
E string `json:"e"`
|
|
} `json:"keys"`
|
|
}
|
|
|
|
func (a *Auth) refreshKeys(ctx context.Context) error {
|
|
req, err := http.NewRequestWithContext(ctx, http.MethodGet, a.jwksURL, nil)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
resp, err := http.DefaultClient.Do(req)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
defer resp.Body.Close()
|
|
if resp.StatusCode != http.StatusOK {
|
|
return fmt.Errorf("jwks request failed: %s", resp.Status)
|
|
}
|
|
|
|
var doc jwksDoc
|
|
if err := json.NewDecoder(resp.Body).Decode(&doc); err != nil {
|
|
return err
|
|
}
|
|
newKeys := map[string]*rsa.PublicKey{}
|
|
for _, k := range doc.Keys {
|
|
if k.Kty != "RSA" || k.Kid == "" || k.N == "" || k.E == "" {
|
|
continue
|
|
}
|
|
pub, err := rsaKeyFromJWK(k.N, k.E)
|
|
if err != nil {
|
|
continue
|
|
}
|
|
newKeys[k.Kid] = pub
|
|
}
|
|
if len(newKeys) == 0 {
|
|
return errors.New("no usable jwks keys")
|
|
}
|
|
a.mu.Lock()
|
|
a.keyByKid = newKeys
|
|
a.lastSyncAt = time.Now()
|
|
a.mu.Unlock()
|
|
return nil
|
|
}
|
|
|
|
func rsaKeyFromJWK(nB64, eB64 string) (*rsa.PublicKey, error) {
|
|
nBytes, err := base64.RawURLEncoding.DecodeString(nB64)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
eBytes, err := base64.RawURLEncoding.DecodeString(eB64)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
n := new(big.Int).SetBytes(nBytes)
|
|
e := 0
|
|
for _, b := range eBytes {
|
|
e = e<<8 + int(b)
|
|
}
|
|
if e == 0 {
|
|
return nil, errors.New("invalid exponent")
|
|
}
|
|
return &rsa.PublicKey{N: n, E: e}, nil
|
|
}
|
|
|
|
func audienceMatches(audValue any, expected string) bool {
|
|
switch v := audValue.(type) {
|
|
case string:
|
|
return v == expected
|
|
case []any:
|
|
for _, item := range v {
|
|
s, ok := item.(string)
|
|
if ok && s == expected {
|
|
return true
|
|
}
|
|
}
|
|
}
|
|
return false
|
|
}
|
|
|
|
func randomToken(n int) (string, error) {
|
|
b := make([]byte, n)
|
|
if _, err := rand.Read(b); err != nil {
|
|
return "", err
|
|
}
|
|
return base64.RawURLEncoding.EncodeToString(b), nil
|
|
}
|
|
|