Merge pull request 'Added refresh tokens' (#9) from dev into main

This commit was merged in pull request #9.
This commit is contained in:
2026-08-03 11:44:38 +02:00
7 changed files with 177 additions and 22 deletions
+36 -3
View File
@@ -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)
} }
+1 -1
View File
@@ -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")
+17
View File
@@ -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
}
@@ -1,2 +1,2 @@
DROP TABLE IF EXISTS sessions; DROP TABLE IF EXISTS refresh_tokens;
DROP TABLE IF EXISTS users; DROP TABLE IF EXISTS users;
@@ -6,11 +6,13 @@ CREATE TABLE IF NOT EXISTS users (
created_at TIMESTAMPTZ NOT NULL DEFAULT NOW() created_at TIMESTAMPTZ NOT NULL DEFAULT NOW()
); );
CREATE TABLE IF NOT EXISTS sessions ( CREATE TABLE IF NOT EXISTS refresh_tokens (
id UUID PRIMARY KEY DEFAULT gen_random_uuid(), id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
user_id UUID NOT NULL REFERENCES users(id) ON DELETE CASCADE, user_id UUID NOT NULL REFERENCES users(id) ON DELETE CASCADE,
token_id UUID UNIQUE NOT NULL, token_hash VARCHAR(64) UNIQUE NOT NULL,
is_revoked BOOLEAN NOT NULL DEFAULT FALSE, is_revoked BOOLEAN NOT NULL DEFAULT FALSE,
expires_at TIMESTAMPTZ NOT NULL, expires_at TIMESTAMPTZ NOT NULL,
created_at TIMESTAMPTZ NOT NULL DEFAULT NOW() created_at TIMESTAMPTZ NOT NULL DEFAULT NOW()
); );
CREATE INDEX idx_refresh_tokens_hash ON refresh_tokens(token_hash);
+74 -14
View File
@@ -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[:])
}
+44 -1
View File
@@ -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
}