implement auth using Auth0

This commit is contained in:
2025-12-29 23:49:23 -07:00
parent c9100383e4
commit cc74842e63
18 changed files with 797 additions and 212 deletions
+313
View File
@@ -0,0 +1,313 @@
package authentication
import (
"context"
"crypto/rand"
"encoding/json"
"errors"
"fmt"
"net/url"
"time"
"github.com/coreos/go-oidc/v3/oidc"
"github.com/jackc/pgx/v5"
"github.com/jackc/pgx/v5/pgtype"
"github.com/jackc/pgx/v5/pgxpool"
"golang.org/x/oauth2"
"ruben/inventory2/internal/consts"
)
// TODO: move these to a config?
const (
// The URL of our Auth0 Tenant Domain.
// If you're using a Custom Domain, be sure to set this to that value instead.
AUTH0_DOMAIN = "dev-uq3gqy5bdnwxmr6d.us.auth0.com"
// Our Auth0 application"s Client ID.
AUTH0_CLIENT_ID = "JEjrXTQ9fxlTLgp9RgTIACpUk8a2lqNT"
// Our Auth0 application"s Client Secret.
AUTH0_CLIENT_SECRET = "83U-iWdVaNnwk9XDzteo_2VMyOq_l1siKYqg1_2E7jCzgL8MnkaxlysPMcPMGlxA"
// The Callback URL of our application.
AUTH0_CALLBACK_URL = "https://inventory-plus-plus.com/login/callback"
)
type (
// Authenticator is used to authenticate our users.
Authenticator struct {
*oidc.Provider
oauth2.Config
db *pgxpool.Pool
}
// AccessTokenClaims is the claims Auth0 provides in access tokens
AccessTokenClaims struct {
Audience string `json:"aud"`
Expires int64 `json:"exp"`
FamilyName string `json:"family_name"`
GivenName string `json:"given_name"`
IssuedAt int64 `json:"iat"`
Issuer string `json:"iss"`
Name string `json:"name"`
Nickname string `json:"nickname"`
Picture string `json:"picture"`
SessionID string `json:"sid"`
Subject string `json:"sub"`
UpdatedAt time.Time `json:"updated_at"`
}
)
// New instantiates the *Authenticator.
func New(ctx context.Context, db *pgxpool.Pool) (*Authenticator, error) {
provider, err := oidc.NewProvider(
ctx,
"https://"+AUTH0_DOMAIN+"/",
)
if err != nil {
return nil, err
}
return &Authenticator{
Provider: provider,
Config: oauth2.Config{
ClientID: AUTH0_CLIENT_ID,
ClientSecret: AUTH0_CLIENT_SECRET,
RedirectURL: AUTH0_CALLBACK_URL,
Endpoint: provider.Endpoint(),
Scopes: []string{oidc.ScopeOpenID, "profile"},
},
db: db,
}, nil
}
// VerifyIDToken verifies that an *oauth2.Token is a valid *oidc.IDToken.
func (a *Authenticator) VerifyIDToken(ctx context.Context, token *oauth2.Token) (*oidc.IDToken, error) {
rawIDToken, ok := token.Extra("id_token").(string)
if !ok {
return nil, errors.New("no id_token field in oauth2 token")
}
oidcConfig := &oidc.Config{
ClientID: a.ClientID,
}
return a.Verifier(oidcConfig).Verify(ctx, rawIDToken)
}
// NewState creates a new state for logging in, saving it in the database.
func (a *Authenticator) NewState(ctx context.Context) ([32]byte, error) {
state, err := generateRandomState()
if err != nil {
return state, fmt.Errorf("failed to generate random state: %w", err)
}
if _, err = a.db.Exec(
ctx,
"INSERT INTO oauth_login_states (state) VALUES (@state)",
pgx.NamedArgs{
"state": state[:],
},
); err != nil {
return state, fmt.Errorf("failed to execute query: %w", err)
}
return state, nil
}
func generateRandomState() ([32]byte, error) {
var b [32]byte
_, err := rand.Read(b[:])
return b, err
}
// GetStateExpiration get's the oauth state's expiration
// func (a *Authenticator) GetStateExpiration(ctx context.Context, state [32]byte) (time.Time, error) {
func (a *Authenticator) GetStateExpiration(ctx context.Context, state string) (time.Time, error) {
rows, err := a.db.Query(
ctx,
`SELECT expiration from oauth_login_states WHERE state = ('\x' || @state)::BYTEA`,
pgx.NamedArgs{
//"state": state[:],
"state": state,
},
)
if err != nil {
return time.Time{}, fmt.Errorf("failed to perform query: %w", err)
}
exp, err := pgx.CollectExactlyOneRow(rows, pgx.RowTo[time.Time])
if err != nil {
if errors.Is(err, pgx.ErrNoRows) {
return time.Time{}, consts.ErrNotFound
}
return time.Time{}, fmt.Errorf("failed to scan row: %w", err)
}
return exp, nil
}
func (a *Authenticator) GetLogoutURL(requestHost string) *url.URL {
return &url.URL{
Scheme: "https",
Host: AUTH0_DOMAIN,
Path: "/v2/logout",
RawQuery: url.Values{
"returnTo": {
(&url.URL{
Scheme: "https",
Host: requestHost,
}).String(),
},
"client_id": {AUTH0_CLIENT_ID},
}.Encode(),
}
}
func (a *Authenticator) Exchange(ctx context.Context, state, code string) (accessToken string, expiration time.Time, err error) {
// validate state
if exp, err := a.GetStateExpiration(ctx, state); errors.Is(err, consts.ErrNotFound) {
return "", time.Time{}, fmt.Errorf("invalid state: %w", consts.ErrNotFound)
} else if err != nil {
return "", time.Time{}, fmt.Errorf("failed to load state expiration: %w", err)
} else if exp.Before(time.Now()) {
return "", time.Time{}, fmt.Errorf("invalid state: state expired")
}
// obtain token and profile
token, err := a.Config.Exchange(ctx, code)
if err != nil {
return "", time.Time{}, fmt.Errorf("failed to exchange an authorization code for a token: %w", err)
}
fmt.Println("TOKEN:", token)
idToken, err := a.VerifyIDToken(ctx, token)
if err != nil {
return "", time.Time{}, fmt.Errorf("failed to verify ID Token: %w", err)
}
fmt.Println("ID TOKEN:", idToken)
// claims []byte
var claims map[string]any
if err := idToken.Claims(&claims); err != nil {
return "", time.Time{}, fmt.Errorf("Failed to obtain id token claims: %w", err)
}
fmt.Println("CUSTOM CLAIMS / PROFILE:", claims)
claimsJSON, _ := json.Marshal(claims)
if _, err := a.db.Exec(
ctx,
`
INSERT INTO oauth_tokens (
access_token,
token_type,
refresh_token,
expiry,
id_token_issuer,
id_token_audience,
id_token_subject,
id_token_expiry,
id_token_issued_at,
id_token_nonce,
id_token_access_token_hash,
claims
)
VALUES (
@access_token,
@token_type,
@refresh_token,
@expiry,
@id_token_issuer,
@id_token_audience,
@id_token_subject,
@id_token_expiry,
@id_token_issued_at,
@id_token_nonce,
@id_token_access_token_hash,
@claims
)
`,
pgx.NamedArgs{
"access_token": token.AccessToken,
"token_type": token.TokenType,
"refresh_token": token.RefreshToken,
"expiry": token.Expiry,
"id_token_issuer": idToken.Issuer,
"id_token_audience": pgtype.FlatArray[string](idToken.Audience),
"id_token_subject": idToken.Subject,
"id_token_expiry": idToken.Expiry,
"id_token_issued_at": idToken.IssuedAt,
"id_token_nonce": idToken.Nonce,
"id_token_access_token_hash": idToken.AccessTokenHash,
"claims": json.RawMessage(claimsJSON),
},
); err != nil {
return "", time.Time{}, fmt.Errorf("failed to perform query to save tokens: %w", err)
}
return token.AccessToken, token.Expiry.UTC(), nil
}
// TODO: need to automatically clean up expired tokens
func (a *Authenticator) DeleteOAuthTokens(ctx context.Context, accessToken string) error {
_, err := a.db.Exec(ctx, "DELETE FROM oauth_tokens WHERE access_token = @access_token", pgx.NamedArgs{"access_token": accessToken})
if err != nil {
return fmt.Errorf("failed to perform query: %w", err)
}
return nil
}
func (a *Authenticator) GetAccessTokenClaimsAndExpiration(ctx context.Context, accessToken string) (claims AccessTokenClaims, expiration time.Time, err error) {
rows, err := a.db.Query(
ctx,
`
SELECT
expiry,
claims
FROM
oauth_tokens
WHERE
access_token = @access_token
`,
pgx.NamedArgs{
"access_token": accessToken,
},
)
if err != nil {
return AccessTokenClaims{}, time.Time{}, fmt.Errorf("failed to perform query: %w", err)
}
type Row struct {
Expiry time.Time
Claims json.RawMessage
}
r, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByNameLax[Row])
if err != nil {
if errors.Is(err, pgx.ErrNoRows) {
return AccessTokenClaims{}, time.Time{}, consts.ErrNotFound
}
return AccessTokenClaims{}, time.Time{}, fmt.Errorf("failed to scan row: %w", err)
}
if err := json.Unmarshal(r.Claims, &claims); err != nil {
return AccessTokenClaims{}, time.Time{}, fmt.Errorf("failed to scan claims json: %w", err)
}
return claims, r.Expiry, nil
}