feat: added refresh tokens
This commit is contained in:
@@ -25,7 +25,7 @@ func NewAuthHandler(authService *domain.AuthService) *AuthHandler {
|
|||||||
// @Accept json
|
// @Accept json
|
||||||
// @Produce json
|
// @Produce json
|
||||||
// @Param payload body request.LoginPayload true "Login payload"
|
// @Param payload body request.LoginPayload true "Login payload"
|
||||||
// @Success 200 {object} domain.AuthToken
|
// @Success 200 {object} domain.TokenPair
|
||||||
// @Error 400 {object} response.ErrorResponse "failed to decode request body"
|
// @Error 400 {object} response.ErrorResponse "failed to decode request body"
|
||||||
// @Error 500 {object} response.ErrorResponse "failed to login"
|
// @Error 500 {object} response.ErrorResponse "failed to login"
|
||||||
// @Router /api/v1/login [post]
|
// @Router /api/v1/login [post]
|
||||||
@@ -39,12 +39,45 @@ func (h *AuthHandler) Login(w http.ResponseWriter, r *http.Request) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
token, err := h.AuthService.Login(ctx, payload.Email, payload.Password)
|
tokens, err := h.AuthService.Login(ctx, payload.Email, payload.Password)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
slog.ErrorContext(ctx, "failed to login", slog.Any("error", err))
|
slog.ErrorContext(ctx, "failed to login", slog.Any("error", err))
|
||||||
response.Error(ctx, w, http.StatusInternalServerError, "failed to login")
|
response.Error(ctx, w, http.StatusInternalServerError, "failed to login")
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
response.JSON(ctx, w, http.StatusOK, token)
|
response.JSON(ctx, w, http.StatusOK, tokens)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Refresh handles refreshing an access token
|
||||||
|
//
|
||||||
|
// @Summary Refresh route
|
||||||
|
// @Description Token issuing with refresh token
|
||||||
|
// @Tags Authorization
|
||||||
|
// @Accept json
|
||||||
|
// @Produce json
|
||||||
|
// @Param payload body request.RefreshPayload true "Refresh payload"
|
||||||
|
// @Success 200 {object} domain.TokenPair
|
||||||
|
// @Error 400 {object} response.ErrorResponse "failed to decode request body"
|
||||||
|
// @Error 500 {object} response.ErrorResponse "failed to refresh token"
|
||||||
|
// @Router /api/v1/refresh [post]
|
||||||
|
func (h *AuthHandler) Refresh(w http.ResponseWriter, r *http.Request) {
|
||||||
|
ctx := r.Context()
|
||||||
|
|
||||||
|
payload, err := request.DecodeRefresh(r)
|
||||||
|
if err != nil {
|
||||||
|
slog.ErrorContext(ctx, "failed to decode refresh payload", slog.Any("error", err))
|
||||||
|
response.Error(ctx, w, http.StatusBadRequest, "failed to decode request body")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
tokens, err := h.AuthService.Refresh(ctx, payload.RefreshToken)
|
||||||
|
if err != nil {
|
||||||
|
slog.ErrorContext(ctx, "failed to refresh token", slog.Any("error", err))
|
||||||
|
// TODO: Add error definition to repository, respond with Unauthorized when token is invalid
|
||||||
|
response.Error(ctx, w, http.StatusInternalServerError, "failed to refresh token")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
response.JSON(ctx, w, http.StatusOK, tokens)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -28,7 +28,7 @@ func WithJWT(authService *domain.AuthService) func(http.HandlerFunc) http.Handle
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
authContext, err := authService.ParseToken(ctx, parts[1])
|
authContext, err := authService.ParseAccessToken(ctx, parts[1])
|
||||||
if err != nil {
|
if err != nil {
|
||||||
slog.ErrorContext(ctx, "invalid token", slog.Any("error", err))
|
slog.ErrorContext(ctx, "invalid token", slog.Any("error", err))
|
||||||
response.Error(ctx, w, http.StatusUnauthorized, "invalid or expired token")
|
response.Error(ctx, w, http.StatusUnauthorized, "invalid or expired token")
|
||||||
|
|||||||
@@ -11,6 +11,10 @@ type LoginPayload struct {
|
|||||||
Password string `json:"password"`
|
Password string `json:"password"`
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type RefreshPayload struct {
|
||||||
|
RefreshToken string `json:"refresh_token"`
|
||||||
|
}
|
||||||
|
|
||||||
func DecodeLogin(r *http.Request) (LoginPayload, error) {
|
func DecodeLogin(r *http.Request) (LoginPayload, error) {
|
||||||
var payload LoginPayload
|
var payload LoginPayload
|
||||||
|
|
||||||
@@ -23,3 +27,16 @@ func DecodeLogin(r *http.Request) (LoginPayload, error) {
|
|||||||
|
|
||||||
return payload, nil
|
return payload, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func DecodeRefresh(r *http.Request) (RefreshPayload, error) {
|
||||||
|
var payload RefreshPayload
|
||||||
|
|
||||||
|
decoder := json.NewDecoder(r.Body)
|
||||||
|
decoder.DisallowUnknownFields()
|
||||||
|
|
||||||
|
if err := decoder.Decode(&payload); err != nil {
|
||||||
|
return payload, fmt.Errorf("error decoding refresh payload: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return payload, nil
|
||||||
|
}
|
||||||
|
|||||||
+74
-14
@@ -2,6 +2,9 @@ package domain
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"crypto/rand"
|
||||||
|
"crypto/sha256"
|
||||||
|
"encoding/hex"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"log/slog"
|
"log/slog"
|
||||||
@@ -12,7 +15,10 @@ import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
type AuthRepository interface {
|
type AuthRepository interface {
|
||||||
GetByEmail(ctx context.Context, email string) (*AuthUser, error)
|
GetUserById(ctx context.Context, id string) (*AuthUser, error)
|
||||||
|
GetUserByEmail(ctx context.Context, email string) (*AuthUser, error)
|
||||||
|
StoreRefreshToken(ctx context.Context, userID string, tokenHash string, expiresAt time.Time) error
|
||||||
|
ConsumeRefreshToken(ctx context.Context, tokenHash string) (string, error)
|
||||||
}
|
}
|
||||||
|
|
||||||
type AuthService struct {
|
type AuthService struct {
|
||||||
@@ -20,8 +26,9 @@ type AuthService struct {
|
|||||||
jwtSecret []byte
|
jwtSecret []byte
|
||||||
}
|
}
|
||||||
|
|
||||||
type AuthToken struct {
|
type TokenPair struct {
|
||||||
Token string `json:"token"`
|
AccessToken string `json:"access_token"`
|
||||||
|
RefreshToken string `json:"refresh_token"`
|
||||||
}
|
}
|
||||||
|
|
||||||
type AuthUser struct {
|
type AuthUser struct {
|
||||||
@@ -54,7 +61,7 @@ func NewAuthService(r AuthRepository, jwtSecret []byte) *AuthService {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (s *AuthService) authenticate(ctx context.Context, email string, password string) (*AuthUser, error) {
|
func (s *AuthService) authenticate(ctx context.Context, email string, password string) (*AuthUser, error) {
|
||||||
user, err := s.repo.GetByEmail(ctx, email)
|
user, err := s.repo.GetUserByEmail(ctx, email)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return user, err
|
return user, err
|
||||||
}
|
}
|
||||||
@@ -70,20 +77,29 @@ func (s *AuthService) authenticate(ctx context.Context, email string, password s
|
|||||||
return user, nil
|
return user, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *AuthService) Login(ctx context.Context, email string, password string) (AuthToken, error) {
|
func (s *AuthService) Login(ctx context.Context, email string, password string) (*TokenPair, error) {
|
||||||
var authToken AuthToken
|
|
||||||
|
|
||||||
user, err := s.authenticate(ctx, email, password)
|
user, err := s.authenticate(ctx, email, password)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return authToken, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
authToken.Token, err = s.GenerateToken(user)
|
return s.issueTokens(ctx, user)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *AuthService) Refresh(ctx context.Context, refreshToken string) (*TokenPair, error) {
|
||||||
|
tokenHash := hashToken(refreshToken)
|
||||||
|
|
||||||
|
userID, err := s.repo.ConsumeRefreshToken(ctx, tokenHash)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return authToken, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
return authToken, nil
|
authUser, err := s.repo.GetUserById(ctx, userID)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
return s.issueTokens(ctx, authUser)
|
||||||
}
|
}
|
||||||
|
|
||||||
type JWTClaims struct {
|
type JWTClaims struct {
|
||||||
@@ -97,13 +113,35 @@ type contextKey string
|
|||||||
|
|
||||||
const AuthContextKey = contextKey("auth")
|
const AuthContextKey = contextKey("auth")
|
||||||
|
|
||||||
func (s *AuthService) GenerateToken(authUser *AuthUser) (string, error) {
|
func (s *AuthService) issueTokens(ctx context.Context, authUser *AuthUser) (*TokenPair, error) {
|
||||||
|
accessToken, err := s.GenerateAccessToken(authUser)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to issue access token: %s", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
refreshToken, err := s.GenerateRefreshToken()
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to issue refresh token: %s", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
err = s.repo.StoreRefreshToken(ctx, authUser.ID, hashToken(refreshToken), time.Now().Add(time.Hour*24*7))
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to store refresh token: %s", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return &TokenPair{
|
||||||
|
AccessToken: accessToken,
|
||||||
|
RefreshToken: refreshToken,
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *AuthService) GenerateAccessToken(authUser *AuthUser) (string, error) {
|
||||||
claims := JWTClaims{
|
claims := JWTClaims{
|
||||||
UserID: authUser.ID,
|
UserID: authUser.ID,
|
||||||
Email: authUser.Email,
|
Email: authUser.Email,
|
||||||
Name: authUser.Name,
|
Name: authUser.Name,
|
||||||
}
|
}
|
||||||
claims.ExpiresAt = jwt.NewNumericDate(time.Now().Add(time.Hour * 24))
|
claims.ExpiresAt = jwt.NewNumericDate(time.Now().Add(time.Minute * 15))
|
||||||
|
|
||||||
token := jwt.NewWithClaims(jwt.SigningMethodHS256, claims)
|
token := jwt.NewWithClaims(jwt.SigningMethodHS256, claims)
|
||||||
|
|
||||||
@@ -115,7 +153,16 @@ func (s *AuthService) GenerateToken(authUser *AuthUser) (string, error) {
|
|||||||
return tokenString, nil
|
return tokenString, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *AuthService) ParseToken(ctx context.Context, tokenString string) (*AuthContext, error) {
|
func (s *AuthService) GenerateRefreshToken() (string, error) {
|
||||||
|
rawToken, err := generateSecureToken(32)
|
||||||
|
if err != nil {
|
||||||
|
return "", fmt.Errorf("failed to generate refresh token: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return rawToken, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *AuthService) ParseAccessToken(ctx context.Context, tokenString string) (*AuthContext, error) {
|
||||||
claims := &JWTClaims{}
|
claims := &JWTClaims{}
|
||||||
token, err := jwt.ParseWithClaims(tokenString, claims, func(token *jwt.Token) (any, error) {
|
token, err := jwt.ParseWithClaims(tokenString, claims, func(token *jwt.Token) (any, error) {
|
||||||
if _, ok := token.Method.(*jwt.SigningMethodHMAC); !ok {
|
if _, ok := token.Method.(*jwt.SigningMethodHMAC); !ok {
|
||||||
@@ -134,3 +181,16 @@ func (s *AuthService) ParseToken(ctx context.Context, tokenString string) (*Auth
|
|||||||
Name: claims.Name,
|
Name: claims.Name,
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func generateSecureToken(n int) (string, error) {
|
||||||
|
bytes := make([]byte, n)
|
||||||
|
if _, err := rand.Read(bytes); err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
return hex.EncodeToString(bytes), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func hashToken(token string) string {
|
||||||
|
hash := sha256.Sum256([]byte(token))
|
||||||
|
return hex.EncodeToString(hash[:])
|
||||||
|
}
|
||||||
|
|||||||
@@ -3,7 +3,9 @@ package repository
|
|||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"database/sql"
|
"database/sql"
|
||||||
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"time"
|
||||||
|
|
||||||
"github.com/robindittmar/dttmr-api/internal/domain"
|
"github.com/robindittmar/dttmr-api/internal/domain"
|
||||||
)
|
)
|
||||||
@@ -16,7 +18,21 @@ func NewAuthRepo(db *sql.DB) *AuthRepo {
|
|||||||
return &AuthRepo{db: db}
|
return &AuthRepo{db: db}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *AuthRepo) GetByEmail(ctx context.Context, email string) (*domain.AuthUser, error) {
|
func (r *AuthRepo) GetUserById(ctx context.Context, id string) (*domain.AuthUser, error) {
|
||||||
|
user := &domain.AuthUser{}
|
||||||
|
|
||||||
|
err := r.db.QueryRowContext(ctx,
|
||||||
|
"SELECT id, email, name, password_hash FROM users WHERE id = $1",
|
||||||
|
id,
|
||||||
|
).Scan(&user.ID, &user.Email, &user.Name, &user.PasswordHash)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to get user: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return user, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *AuthRepo) GetUserByEmail(ctx context.Context, email string) (*domain.AuthUser, error) {
|
||||||
user := &domain.AuthUser{}
|
user := &domain.AuthUser{}
|
||||||
|
|
||||||
err := r.db.QueryRowContext(ctx,
|
err := r.db.QueryRowContext(ctx,
|
||||||
@@ -29,3 +45,30 @@ func (r *AuthRepo) GetByEmail(ctx context.Context, email string) (*domain.AuthUs
|
|||||||
|
|
||||||
return user, nil
|
return user, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (r *AuthRepo) StoreRefreshToken(ctx context.Context, userID string, tokenHash string, expiresAt time.Time) error {
|
||||||
|
_, err := r.db.ExecContext(ctx,
|
||||||
|
"INSERT INTO refresh_tokens (user_id, token_hash, expires_at) VALUES ($1, $2, $3)",
|
||||||
|
userID, tokenHash, expiresAt,
|
||||||
|
)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("failed to store refresh token: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
func (r *AuthRepo) ConsumeRefreshToken(ctx context.Context, tokenHash string) (string, error) {
|
||||||
|
var userID string
|
||||||
|
err := r.db.QueryRowContext(ctx,
|
||||||
|
"DELETE FROM refresh_tokens WHERE token_hash = $1 AND expires_at > NOW() RETURNING user_id",
|
||||||
|
tokenHash,
|
||||||
|
).Scan(&userID)
|
||||||
|
if err != nil {
|
||||||
|
if errors.Is(err, sql.ErrNoRows) {
|
||||||
|
return "", fmt.Errorf("refresh token not found or expired: %w", err)
|
||||||
|
}
|
||||||
|
return "", fmt.Errorf("failed to consume refresh token: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return userID, nil
|
||||||
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user