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 }