feat: added refresh tokens
This commit is contained in:
@@ -25,7 +25,7 @@ func NewAuthHandler(authService *domain.AuthService) *AuthHandler {
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @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 500 {object} response.ErrorResponse "failed to login"
|
||||
// @Router /api/v1/login [post]
|
||||
@@ -39,12 +39,45 @@ func (h *AuthHandler) Login(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
token, err := h.AuthService.Login(ctx, payload.Email, payload.Password)
|
||||
tokens, err := h.AuthService.Login(ctx, payload.Email, payload.Password)
|
||||
if err != nil {
|
||||
slog.ErrorContext(ctx, "failed to login", slog.Any("error", err))
|
||||
response.Error(ctx, w, http.StatusInternalServerError, "failed to login")
|
||||
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
|
||||
}
|
||||
|
||||
authContext, err := authService.ParseToken(ctx, parts[1])
|
||||
authContext, err := authService.ParseAccessToken(ctx, parts[1])
|
||||
if err != nil {
|
||||
slog.ErrorContext(ctx, "invalid token", slog.Any("error", err))
|
||||
response.Error(ctx, w, http.StatusUnauthorized, "invalid or expired token")
|
||||
|
||||
@@ -11,6 +11,10 @@ type LoginPayload struct {
|
||||
Password string `json:"password"`
|
||||
}
|
||||
|
||||
type RefreshPayload struct {
|
||||
RefreshToken string `json:"refresh_token"`
|
||||
}
|
||||
|
||||
func DecodeLogin(r *http.Request) (LoginPayload, error) {
|
||||
var payload LoginPayload
|
||||
|
||||
@@ -23,3 +27,16 @@ func DecodeLogin(r *http.Request) (LoginPayload, error) {
|
||||
|
||||
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 (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
@@ -12,7 +15,10 @@ import (
|
||||
)
|
||||
|
||||
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 {
|
||||
@@ -20,8 +26,9 @@ type AuthService struct {
|
||||
jwtSecret []byte
|
||||
}
|
||||
|
||||
type AuthToken struct {
|
||||
Token string `json:"token"`
|
||||
type TokenPair struct {
|
||||
AccessToken string `json:"access_token"`
|
||||
RefreshToken string `json:"refresh_token"`
|
||||
}
|
||||
|
||||
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) {
|
||||
user, err := s.repo.GetByEmail(ctx, email)
|
||||
user, err := s.repo.GetUserByEmail(ctx, email)
|
||||
if err != nil {
|
||||
return user, err
|
||||
}
|
||||
@@ -70,20 +77,29 @@ func (s *AuthService) authenticate(ctx context.Context, email string, password s
|
||||
return user, nil
|
||||
}
|
||||
|
||||
func (s *AuthService) Login(ctx context.Context, email string, password string) (AuthToken, error) {
|
||||
var authToken AuthToken
|
||||
|
||||
func (s *AuthService) Login(ctx context.Context, email string, password string) (*TokenPair, error) {
|
||||
user, err := s.authenticate(ctx, email, password)
|
||||
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 {
|
||||
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 {
|
||||
@@ -97,13 +113,35 @@ type contextKey string
|
||||
|
||||
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{
|
||||
UserID: authUser.ID,
|
||||
Email: authUser.Email,
|
||||
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)
|
||||
|
||||
@@ -115,7 +153,16 @@ func (s *AuthService) GenerateToken(authUser *AuthUser) (string, error) {
|
||||
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{}
|
||||
token, err := jwt.ParseWithClaims(tokenString, claims, func(token *jwt.Token) (any, error) {
|
||||
if _, ok := token.Method.(*jwt.SigningMethodHMAC); !ok {
|
||||
@@ -134,3 +181,16 @@ func (s *AuthService) ParseToken(ctx context.Context, tokenString string) (*Auth
|
||||
Name: claims.Name,
|
||||
}, 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 (
|
||||
"context"
|
||||
"database/sql"
|
||||
"errors"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/robindittmar/dttmr-api/internal/domain"
|
||||
)
|
||||
@@ -16,7 +18,21 @@ func NewAuthRepo(db *sql.DB) *AuthRepo {
|
||||
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{}
|
||||
|
||||
err := r.db.QueryRowContext(ctx,
|
||||
@@ -29,3 +45,30 @@ func (r *AuthRepo) GetByEmail(ctx context.Context, email string) (*domain.AuthUs
|
||||
|
||||
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