yep
This commit is contained in:
parent
4f6d7cfe60
commit
8fa8d04427
10 changed files with 231 additions and 71 deletions
217
auth.go
217
auth.go
|
|
@ -2,14 +2,17 @@ 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"
|
||||
|
|
@ -21,6 +24,11 @@ type contextKey string
|
|||
|
||||
const userContextKey contextKey = "auth_user"
|
||||
|
||||
const (
|
||||
sessionCookieName = "session_token"
|
||||
stateCookieName = "oidc_state"
|
||||
)
|
||||
|
||||
type User struct {
|
||||
Sub string
|
||||
Email string
|
||||
|
|
@ -28,11 +36,17 @@ type User struct {
|
|||
}
|
||||
|
||||
type Auth struct {
|
||||
issuer string
|
||||
audience string
|
||||
jwksURL string
|
||||
devMode bool
|
||||
db *sql.DB
|
||||
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
|
||||
|
|
@ -41,12 +55,18 @@ type Auth struct {
|
|||
|
||||
func NewAuth(cfg Config) (*Auth, error) {
|
||||
a := &Auth{
|
||||
issuer: cfg.OIDCIssuer,
|
||||
audience: cfg.OIDCAudience,
|
||||
jwksURL: cfg.OIDCJWKSURL,
|
||||
devMode: cfg.DevMode,
|
||||
keyByKid: map[string]*rsa.PublicKey{},
|
||||
lastSyncAt: time.Time{},
|
||||
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
|
||||
|
|
@ -61,12 +81,114 @@ 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",
|
||||
|
|
@ -74,9 +196,14 @@ func (a *Auth) RequireAuth(next http.Handler) http.Handler {
|
|||
Name: "Demo User",
|
||||
}
|
||||
} else {
|
||||
user, err = a.authenticateRequest(r)
|
||||
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.Error(w, "unauthorized", http.StatusUnauthorized)
|
||||
http.Redirect(w, r, "/", http.StatusFound)
|
||||
return
|
||||
}
|
||||
}
|
||||
|
|
@ -106,19 +233,50 @@ func MustUserFromContext(ctx context.Context) User {
|
|||
return User{}
|
||||
}
|
||||
|
||||
func (a *Auth) authenticateRequest(r *http.Request) (User, error) {
|
||||
token := extractBearer(r.Header.Get("Authorization"))
|
||||
if token == "" {
|
||||
return User{}, errors.New("missing token")
|
||||
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(r.Context(), kid)
|
||||
key, err := a.lookupKey(context.Background(), kid)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
|
@ -132,7 +290,6 @@ func (a *Auth) authenticateRequest(r *http.Request) (User, error) {
|
|||
if !ok {
|
||||
return User{}, errors.New("invalid claims")
|
||||
}
|
||||
|
||||
if iss, _ := claims["iss"].(string); iss != a.issuer {
|
||||
return User{}, errors.New("invalid issuer")
|
||||
}
|
||||
|
|
@ -189,7 +346,6 @@ func (a *Auth) refreshKeys(ctx context.Context) error {
|
|||
return err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return fmt.Errorf("jwks request failed: %s", resp.Status)
|
||||
}
|
||||
|
|
@ -212,7 +368,6 @@ func (a *Auth) refreshKeys(ctx context.Context) error {
|
|||
if len(newKeys) == 0 {
|
||||
return errors.New("no usable jwks keys")
|
||||
}
|
||||
|
||||
a.mu.Lock()
|
||||
a.keyByKid = newKeys
|
||||
a.lastSyncAt = time.Now()
|
||||
|
|
@ -240,17 +395,6 @@ func rsaKeyFromJWK(nB64, eB64 string) (*rsa.PublicKey, error) {
|
|||
return &rsa.PublicKey{N: n, E: e}, nil
|
||||
}
|
||||
|
||||
func extractBearer(value string) string {
|
||||
if value == "" {
|
||||
return ""
|
||||
}
|
||||
parts := strings.SplitN(value, " ", 2)
|
||||
if len(parts) != 2 || !strings.EqualFold(parts[0], "Bearer") {
|
||||
return ""
|
||||
}
|
||||
return strings.TrimSpace(parts[1])
|
||||
}
|
||||
|
||||
func audienceMatches(audValue any, expected string) bool {
|
||||
switch v := audValue.(type) {
|
||||
case string:
|
||||
|
|
@ -265,3 +409,12 @@ func audienceMatches(audValue any, expected string) bool {
|
|||
}
|
||||
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
|
||||
}
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue