diff --git a/internal/api/handler/auth.go b/internal/api/handler/auth.go index 0573cbd..efdab88 100644 --- a/internal/api/handler/auth.go +++ b/internal/api/handler/auth.go @@ -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) } diff --git a/internal/api/middleware/jwt.go b/internal/api/middleware/jwt.go index 7615a87..4b83ede 100644 --- a/internal/api/middleware/jwt.go +++ b/internal/api/middleware/jwt.go @@ -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") diff --git a/internal/api/request/auth.go b/internal/api/request/auth.go index a504d5a..e69fe74 100644 --- a/internal/api/request/auth.go +++ b/internal/api/request/auth.go @@ -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 +} diff --git a/internal/database/migrations/000001_create_users.down.sql b/internal/database/migrations/000001_create_users.down.sql index 0e03010..7150d35 100644 --- a/internal/database/migrations/000001_create_users.down.sql +++ b/internal/database/migrations/000001_create_users.down.sql @@ -1,2 +1,2 @@ -DROP TABLE IF EXISTS sessions; +DROP TABLE IF EXISTS refresh_tokens; DROP TABLE IF EXISTS users; diff --git a/internal/database/migrations/000001_create_users.up.sql b/internal/database/migrations/000001_create_users.up.sql index 76b4892..b8f889e 100644 --- a/internal/database/migrations/000001_create_users.up.sql +++ b/internal/database/migrations/000001_create_users.up.sql @@ -6,11 +6,13 @@ CREATE TABLE IF NOT EXISTS users ( 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(), 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, expires_at TIMESTAMPTZ NOT NULL, created_at TIMESTAMPTZ NOT NULL DEFAULT NOW() ); + +CREATE INDEX idx_refresh_tokens_hash ON refresh_tokens(token_hash); diff --git a/internal/domain/auth.go b/internal/domain/auth.go index 9e7b04d..75e89e8 100644 --- a/internal/domain/auth.go +++ b/internal/domain/auth.go @@ -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[:]) +} diff --git a/internal/repository/auth.go b/internal/repository/auth.go index 88dfc01..1c05973 100644 --- a/internal/repository/auth.go +++ b/internal/repository/auth.go @@ -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 +}