diff --git a/internal/domains/authentication/auth.go b/internal/domains/authentication/auth.go index 5284f38..3bbb764 100644 --- a/internal/domains/authentication/auth.go +++ b/internal/domains/authentication/auth.go @@ -44,18 +44,19 @@ type ( // AccessTokenClaims is the claims Auth0 provides in access tokens AccessTokenClaims struct { - Audience string `json:"aud"` + Audience string `json:"aud"` // TODO: fill in from db Expires int64 `json:"exp"` + Expiration time.Time `json:"-"` // parsed Expires FamilyName string `json:"family_name"` GivenName string `json:"given_name"` - IssuedAt int64 `json:"iat"` - Issuer string `json:"iss"` + IssuedAt int64 `json:"iat"` // TODO: fill in from db + Issuer string `json:"iss"` // TODO: fill in from db 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"` + SessionID string `json:"sid"` // TODO: fill in from db + Subject string `json:"sub"` // TODO: fill in from db + UpdatedAt time.Time `json:"updated_at"` // TODO: fill in from db } ) diff --git a/internal/domains/authentication/store.go b/internal/domains/authentication/store.go index 1bed492..09bda1d 100644 --- a/internal/domains/authentication/store.go +++ b/internal/domains/authentication/store.go @@ -95,7 +95,7 @@ func (a *Authenticator) DeleteOAuthTokens(ctx context.Context, accessToken strin return nil } -func (a *Authenticator) GetAccessTokenClaimsAndExpiration(ctx context.Context, accessToken string) (claims AccessTokenClaims, expiration time.Time, err error) { +func (a *Authenticator) GetAccessTokenClaimsAndExpiration(ctx context.Context, accessToken string) (claims AccessTokenClaims, err error) { rows, err := a.db.Query( ctx, ` @@ -119,7 +119,7 @@ func (a *Authenticator) GetAccessTokenClaimsAndExpiration(ctx context.Context, a }, ) if err != nil { - return AccessTokenClaims{}, time.Time{}, fmt.Errorf("failed to perform query: %w", err) + return AccessTokenClaims{}, fmt.Errorf("failed to perform query: %w", err) } type Row struct { @@ -135,11 +135,13 @@ func (a *Authenticator) GetAccessTokenClaimsAndExpiration(ctx context.Context, a 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{}, consts.ErrNotFound } - return AccessTokenClaims{}, time.Time{}, fmt.Errorf("failed to scan row: %w", err) + return AccessTokenClaims{}, fmt.Errorf("failed to scan row: %w", err) } + claims.Expires = r.Expiry.Unix() + claims.Expiration = r.Expiry claims.Name = r.Id_token_custom_claims_name claims.Picture = r.Id_token_custom_claims_picture claims.Nickname = r.Id_token_custom_claims_nickname @@ -147,7 +149,7 @@ func (a *Authenticator) GetAccessTokenClaimsAndExpiration(ctx context.Context, a claims.FamilyName = r.Id_token_custom_claims_family_name claims.UpdatedAt = r.Id_token_custom_claims_updated_at - return claims, r.Expiry, nil + return claims, nil } func (a *Authenticator) getRefreshTokenForAccessToken(ctx context.Context, accessToken string) (refreshToken, tokenType string, err error) { diff --git a/internal/server/api/sse/router.go b/internal/server/api/sse/router.go index 08202d8..e8fd586 100644 --- a/internal/server/api/sse/router.go +++ b/internal/server/api/sse/router.go @@ -29,7 +29,6 @@ func Routes( r gin.IRouter, logger *logging.Logger, sq *sse.Queue, - auth *middleware.Auth, ) { s := &sseRouter{ log: logger, @@ -37,7 +36,7 @@ func Routes( users: make(map[string][maxNumOpenConnectionsPerUser]context.CancelFunc), } - r.GET("/", auth.AuthenticateAndAddIdentityGin(), response.Handler(s.serveEvents)) + r.GET("/", response.Handler(s.serveEvents)) } func (r *sseRouter) serveEvents(c *gin.Context) (response.Response, error) { diff --git a/internal/server/api/templates/router.go b/internal/server/api/templates/router.go index 6d65068..9a24705 100644 --- a/internal/server/api/templates/router.go +++ b/internal/server/api/templates/router.go @@ -23,13 +23,12 @@ import ( type ( webpageRouter struct { - log *logging.Logger - contentDir string - templater *templater.Templater - rawEvents *raw_events.Store - accts *accounts.Store - etsy *etsy_platform.Platform - authMiddleware *middleware.Auth + log *logging.Logger + contentDir string + templater *templater.Templater + rawEvents *raw_events.Store + accts *accounts.Store + etsy *etsy_platform.Platform } // ErrTemplateNotFound is returned if the reason the template failed to compile @@ -41,12 +40,12 @@ type ( func SetupRoutes( logger *logging.Logger, - r gin.IRouter, + r gin.IRoutes, contentDir string, rawEvents *raw_events.Store, accts *accounts.Store, etsy *etsy_platform.Platform, - authMiddleware *middleware.Auth, + auth *middleware.Auth, ) { s := &webpageRouter{ @@ -90,34 +89,13 @@ func SetupRoutes( } }, }), - rawEvents: rawEvents, - accts: accts, - etsy: etsy, - authMiddleware: authMiddleware, + rawEvents: rawEvents, + accts: accts, + etsy: etsy, } - authenticate := s.authMiddleware.AuthenticateAndAddIdentityToRequest() - - r.GET("/*rest", response.Handler(func(c *gin.Context) (response.Response, error) { - c.Request.URL.Path = c.Request.URL.Path[3:] - defer func() { - c.Request.URL.Path = "/ui" + c.Request.URL.Path - }() - - p := c.Request.URL.Path - if p == "/" || p == "" { - c, err := s.authMiddleware.AddIdentityToRequest(c) - if err != nil { - return nil, err - } - return s.serveTemplate(c) - } - - if res, err := authenticate(c); res != nil || err != nil { - return res, err - } - return s.serveTemplate(c) - })) + r.GET("", response.Handler(s.serveTemplate)) + r.GET("/*rest", response.Handler(auth.Authenticate()), response.Handler(s.serveTemplate)) } // GET / diff --git a/internal/server/middleware/auth.go b/internal/server/middleware/auth.go index 4a76b8b..ed2ba32 100644 --- a/internal/server/middleware/auth.go +++ b/internal/server/middleware/auth.go @@ -52,91 +52,69 @@ func NewAuth( } } -func (a *Auth) AddIdentity(fn response.HandlerFunc) response.HandlerFunc { - return func(c *gin.Context) (response.Response, error) { - c, err := a.AddIdentityToRequest(c) - if err != nil { - return nil, err - } - - return fn(c) +// AddIdentityToRequest will add an Identity to the context that can then be retrieved via GetIdentity. +func (a *Auth) AddIdentityToRequest(c *gin.Context) { + if err := a.addIdentityToRequest(c); err != nil { + c.Error(err) + c.Abort() } } -func (a *Auth) AddIdentityToRequest(c *gin.Context) (*gin.Context, error) { +func (a *Auth) addIdentityToRequest(c *gin.Context) error { r := c.Request ck, err := r.Cookie("access_token") if err != nil { - return c, nil + return nil } ctx := r.Context() accessToken := ck.Value - claims, expiration, err := a.auth.GetAccessTokenClaimsAndExpiration(ctx, accessToken) + claims, err := a.auth.GetAccessTokenClaimsAndExpiration(ctx, accessToken) if err != nil { if errors.Is(err, consts.ErrNotFound) { - return c, nil + return nil } - return c, response.Errorf("failed to load authentication details: %w", err) + return response.Errorf("failed to load authentication details: %w", err) } + expiration := claims.Expiration if expiration.Before(time.Now()) { - return c, nil + return nil } user, acct, err := a.accts.GetUserAndAccountByAccessToken(ctx, accessToken) if err != nil { - return c, response.Errorf("failed to load user and account defails: %w", err) + return response.Errorf("failed to load user and account defails: %w", err) } - c.Request = r.WithContext(SetIdentity(ctx, Identity{ + c.Request = r.WithContext(SetIdentity(c, Identity{ AccessToken: accessToken, Claims: claims, User: user, Account: acct, })) - return c, nil + return nil } -// TODO: use the new middleware pattern -// auth middleware to verify access_token cookie and set custom claims in the request context -func (a *Auth) AuthenticateAndAddIdentityGin(assertions ...AuthorizationAssertions) gin.HandlerFunc { - fn := a.AuthenticateAndAddIdentityToRequest(assertions...) - return response.Handler(fn) -} - -func (a *Auth) AuthenticateAndAddIdentityToRequest(assertions ...AuthorizationAssertions) response.HandlerFunc { +// Authenticate should only be used along with and after AddIdentityToRequest +// Typically used with response.Handler to make a gin.HandlerFunc. +func (a *Auth) Authenticate(assertions ...AuthorizationAssertions) func(c *gin.Context) (response.Response, error) { return func(c *gin.Context) (response.Response, error) { - r := c.Request - ck, err := r.Cookie("access_token") - if err != nil { - return response.TemporaryRedirect("/"). - JSON("no access_token cookie provided"), nil - } - - ctx := r.Context() - - accessToken := ck.Value - - claims, expiration, err := a.auth.GetAccessTokenClaimsAndExpiration(ctx, accessToken) - if err != nil { - if errors.Is(err, consts.ErrNotFound) { - u, err := a.newLoginURL(ctx, a.auth, r.URL.String()) - if err != nil { - return nil, response.Errorf("failed to generate login url: %w", err) - } - - return response.TemporaryRedirect(u), nil - } - - return nil, response.Errorf("failed to authenticate: %w", err) + id, ok := getIdentity(c) + if !ok { + return nil, response.Unauthorized(). + HTML([]byte(` +