diff --git a/internal/auth/auth.go b/internal/auth/auth.go index dd22e51..069ca66 100644 --- a/internal/auth/auth.go +++ b/internal/auth/auth.go @@ -1,6 +1,8 @@ package auth import ( + "crypto/rand" + "encoding/hex" "errors" "fmt" "github.com/golang-jwt/jwt/v5" @@ -84,3 +86,13 @@ func GetBearerToken(headers http.Header) (string, error) { return auth, nil } + +func MakeRefreshToken() (string, error) { + key := make([]byte, 32) + _, err := rand.Read(key) + if err != nil { + return "", err + } + token := hex.EncodeToString(key) + return token, nil +} diff --git a/internal/dataMaps/maps.go b/internal/dataMaps/maps.go index b3db4ae..9e7f6b3 100644 --- a/internal/dataMaps/maps.go +++ b/internal/dataMaps/maps.go @@ -21,9 +21,10 @@ type User struct { } type UserWithToken struct { - ID uuid.UUID `json:"id"` - CreatedAt time.Time `json:"created_at"` - UpdatedAt time.Time `json:"updated_at"` - Email string `json:"email"` - Token string `json:"token"` + ID uuid.UUID `json:"id"` + CreatedAt time.Time `json:"created_at"` + UpdatedAt time.Time `json:"updated_at"` + Email string `json:"email"` + Token string `json:"token"` + RefreshToken string `json:"refresh_token"` } diff --git a/internal/database/CreateRefreshToken.sql.go b/internal/database/CreateRefreshToken.sql.go new file mode 100644 index 0000000..5cba4b7 --- /dev/null +++ b/internal/database/CreateRefreshToken.sql.go @@ -0,0 +1,39 @@ +// Code generated by sqlc. DO NOT EDIT. +// versions: +// sqlc v1.28.0 +// source: CreateRefreshToken.sql + +package database + +import ( + "context" + "time" + + "github.com/google/uuid" +) + +const createRefreshToken = `-- name: CreateRefreshToken :one +INSERT INTO refresh_tokens (token, created_at, updated_at, user_id, expires_at, revoked_at) +VALUES ($1, NOW(), NOW(), $2, $3, null) +RETURNING token, created_at, updated_at, user_id, expires_at, revoked_at +` + +type CreateRefreshTokenParams struct { + Token string + UserID uuid.UUID + ExpiresAt time.Time +} + +func (q *Queries) CreateRefreshToken(ctx context.Context, arg CreateRefreshTokenParams) (RefreshToken, error) { + row := q.db.QueryRowContext(ctx, createRefreshToken, arg.Token, arg.UserID, arg.ExpiresAt) + var i RefreshToken + err := row.Scan( + &i.Token, + &i.CreatedAt, + &i.UpdatedAt, + &i.UserID, + &i.ExpiresAt, + &i.RevokedAt, + ) + return i, err +} diff --git a/internal/database/GetRefreshTokenByToken.sql.go b/internal/database/GetRefreshTokenByToken.sql.go new file mode 100644 index 0000000..75770f9 --- /dev/null +++ b/internal/database/GetRefreshTokenByToken.sql.go @@ -0,0 +1,28 @@ +// Code generated by sqlc. DO NOT EDIT. +// versions: +// sqlc v1.28.0 +// source: GetRefreshTokenByToken.sql + +package database + +import ( + "context" +) + +const getRefreshTokenByToken = `-- name: GetRefreshTokenByToken :one +SELECT token, created_at, updated_at, user_id, expires_at, revoked_at FROM refresh_tokens WHERE token = $1 AND revoked_at is null +` + +func (q *Queries) GetRefreshTokenByToken(ctx context.Context, token string) (RefreshToken, error) { + row := q.db.QueryRowContext(ctx, getRefreshTokenByToken, token) + var i RefreshToken + err := row.Scan( + &i.Token, + &i.CreatedAt, + &i.UpdatedAt, + &i.UserID, + &i.ExpiresAt, + &i.RevokedAt, + ) + return i, err +} diff --git a/internal/database/GetUserByID.sql.go b/internal/database/GetUserByID.sql.go new file mode 100644 index 0000000..fc1239f --- /dev/null +++ b/internal/database/GetUserByID.sql.go @@ -0,0 +1,29 @@ +// Code generated by sqlc. DO NOT EDIT. +// versions: +// sqlc v1.28.0 +// source: GetUserByID.sql + +package database + +import ( + "context" + + "github.com/google/uuid" +) + +const getUserByID = `-- name: GetUserByID :one +SELECT id, created_at, updated_at, email, hashed_password FROM users WHERE id = $1 +` + +func (q *Queries) GetUserByID(ctx context.Context, id uuid.UUID) (User, error) { + row := q.db.QueryRowContext(ctx, getUserByID, id) + var i User + err := row.Scan( + &i.ID, + &i.CreatedAt, + &i.UpdatedAt, + &i.Email, + &i.HashedPassword, + ) + return i, err +} diff --git a/internal/database/RevokeRefreshTokenByToken.sql.go b/internal/database/RevokeRefreshTokenByToken.sql.go new file mode 100644 index 0000000..3f66910 --- /dev/null +++ b/internal/database/RevokeRefreshTokenByToken.sql.go @@ -0,0 +1,19 @@ +// Code generated by sqlc. DO NOT EDIT. +// versions: +// sqlc v1.28.0 +// source: RevokeRefreshTokenByToken.sql + +package database + +import ( + "context" +) + +const revokeRefreshTokenByToken = `-- name: RevokeRefreshTokenByToken :exec +UPDATE refresh_tokens SET revoked_at = CURRENT_TIMESTAMP WHERE token = $1 +` + +func (q *Queries) RevokeRefreshTokenByToken(ctx context.Context, token string) error { + _, err := q.db.ExecContext(ctx, revokeRefreshTokenByToken, token) + return err +} diff --git a/internal/database/models.go b/internal/database/models.go index e10c8df..787dda5 100644 --- a/internal/database/models.go +++ b/internal/database/models.go @@ -5,6 +5,7 @@ package database import ( + "database/sql" "time" "github.com/google/uuid" @@ -18,6 +19,15 @@ type Chirp struct { UserID uuid.UUID } +type RefreshToken struct { + Token string + CreatedAt time.Time + UpdatedAt time.Time + UserID uuid.UUID + ExpiresAt time.Time + RevokedAt sql.NullTime +} + type User struct { ID uuid.UUID CreatedAt time.Time diff --git a/main.go b/main.go index b016bc6..e677df7 100644 --- a/main.go +++ b/main.go @@ -146,7 +146,7 @@ func postChirpsHandler(cfg *apiConfig) http.HandlerFunc { httpResponse.JSONHandler(w, http.StatusUnauthorized, `{"error": "Unauthorized"}`) } - uuid, err := auth.ValidateJWT(token, cfg.secret) + userId, err := auth.ValidateJWT(token, cfg.secret) params := parameters{} err = json.NewDecoder(r.Body).Decode(¶ms) @@ -155,6 +155,12 @@ func postChirpsHandler(cfg *apiConfig) http.HandlerFunc { return } + _, err = cfg.db.GetUserByID(context.Background(), userId) + if err != nil { + httpResponse.JSONHandler(w, http.StatusUnauthorized, `{"error": "Unauthorized"}`) + return + } + chirp, err := validateChirp(params.Body) if err != nil { httpResponse.JSONHandler(w, http.StatusBadRequest, fmt.Sprintf(`{"error": "%s"}`, err.Error())) @@ -163,7 +169,7 @@ func postChirpsHandler(cfg *apiConfig) http.HandlerFunc { newChirp, err := cfg.db.CreateChirp(context.Background(), database.CreateChirpParams{ Body: chirp, - UserID: uuid, + UserID: userId, }) if err != nil { httpResponse.JSONHandler(w, http.StatusInternalServerError, fmt.Sprintf(`{"error": "%s"}`, err.Error())) @@ -186,12 +192,71 @@ func postChirpsHandler(cfg *apiConfig) http.HandlerFunc { httpResponse.JSONHandler(w, http.StatusCreated, string(jsonData)) } } + +func refreshTokenHandler(cfg *apiConfig) http.HandlerFunc { + return func(w http.ResponseWriter, r *http.Request) { + token, err := auth.GetBearerToken(r.Header) + if err != nil { + httpResponse.JSONHandler(w, http.StatusUnauthorized, `{"error": "Unauthorized"}`) + return + } + + find, err := cfg.db.GetRefreshTokenByToken(context.Background(), token) + if err != nil { + httpResponse.JSONHandler(w, http.StatusUnauthorized, `{"error": "Unauthorized"}`) + return + } + + _, err = cfg.db.GetUserByID(context.Background(), find.UserID) + if err != nil { + httpResponse.JSONHandler(w, http.StatusUnauthorized, `{"error": "Unauthorized"}`) + return + } + + if find.RevokedAt.Valid { + httpResponse.JSONHandler(w, http.StatusUnauthorized, `{"error": "Unauthorized"}`) + return + } + + newToken, err := auth.MakeJWT(find.UserID, cfg.secret, time.Duration(3600)*time.Second) + if err != nil { + httpResponse.JSONHandler(w, http.StatusInternalServerError, fmt.Sprintf(`{"error": "%s"}`, err.Error())) + return + } + + httpResponse.JSONHandler(w, http.StatusOK, fmt.Sprintf(`{"token": "%s"}`, newToken)) + } +} + +func revokeRefreshTokenHandler(cfg *apiConfig) http.HandlerFunc { + return func(w http.ResponseWriter, r *http.Request) { + token, err := auth.GetBearerToken(r.Header) + if err != nil { + httpResponse.JSONHandler(w, http.StatusUnauthorized, `{"error": "Unauthorized"}`) + return + } + + find, err := cfg.db.GetRefreshTokenByToken(context.Background(), token) + if err != nil { + httpResponse.JSONHandler(w, http.StatusUnauthorized, `{"error": "Unauthorized"}`) + return + } + + err = cfg.db.RevokeRefreshTokenByToken(context.Background(), find.Token) + if err != nil { + httpResponse.JSONHandler(w, http.StatusInternalServerError, fmt.Sprintf(`{"error": "%s"}`, err.Error())) + return + } + + httpResponse.JSONHandler(w, http.StatusNoContent, "") + } +} + func loginHandler(cfg *apiConfig) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { type parameters struct { - Email string `json:"email"` - Password string `json:"password"` - ExpiresInSeconds int `json:"expires_in_seconds"` + Email string `json:"email"` + Password string `json:"password"` } params := parameters{} @@ -213,10 +278,6 @@ func loginHandler(cfg *apiConfig) http.HandlerFunc { return } - if params.ExpiresInSeconds == 0 || params.ExpiresInSeconds > 3600 { - params.ExpiresInSeconds = 3600 - } - userJSON := dataMaps.UserWithToken{ ID: user.ID, CreatedAt: user.CreatedAt, @@ -224,7 +285,7 @@ func loginHandler(cfg *apiConfig) http.HandlerFunc { Email: user.Email, } - expiresIn := time.Duration(params.ExpiresInSeconds) * time.Second + expiresIn := time.Duration(3600) * time.Second userJSON.Token, err = auth.MakeJWT(user.ID, cfg.secret, expiresIn) if err != nil { @@ -232,6 +293,22 @@ func loginHandler(cfg *apiConfig) http.HandlerFunc { return } + userJSON.RefreshToken, err = auth.MakeRefreshToken() + if err != nil { + httpResponse.JSONHandler(w, http.StatusInternalServerError, fmt.Sprintf(`{"error": "%s"}`, err.Error())) + return + } + + _, err = cfg.db.CreateRefreshToken(context.Background(), database.CreateRefreshTokenParams{ + Token: userJSON.RefreshToken, + UserID: userJSON.ID, + ExpiresAt: time.Now().UTC().Add(time.Duration(60*24*3600) * time.Second), + }) + if err != nil { + httpResponse.JSONHandler(w, http.StatusInternalServerError, fmt.Sprintf(`{"error": "%s"}`, err.Error())) + return + } + jsonData, err := json.Marshal(userJSON) if err != nil { httpResponse.JSONHandler(w, http.StatusInternalServerError, fmt.Sprintf(`{"error": "%s"}`, err.Error())) @@ -327,6 +404,8 @@ func main() { mux.HandleFunc("POST /api/reset", resetMetricsHandler(apiCfg)) mux.HandleFunc("POST /api/users", postUsersHandler(apiCfg)) mux.HandleFunc("POST /api/login", loginHandler(apiCfg)) + mux.HandleFunc("POST /api/refresh", refreshTokenHandler(apiCfg)) + mux.HandleFunc("POST /api/revoke", revokeRefreshTokenHandler(apiCfg)) mux.HandleFunc("POST /api/chirps", postChirpsHandler(apiCfg)) mux.HandleFunc("GET /api/chirps", getAllChirpsHandler(apiCfg)) mux.HandleFunc("GET /api/chirps/{chirpID}", getChirpByIDHandler(apiCfg)) diff --git a/sql/queries/CreateRefreshToken.sql b/sql/queries/CreateRefreshToken.sql new file mode 100644 index 0000000..c6bd251 --- /dev/null +++ b/sql/queries/CreateRefreshToken.sql @@ -0,0 +1,4 @@ +-- name: CreateRefreshToken :one +INSERT INTO refresh_tokens (token, created_at, updated_at, user_id, expires_at, revoked_at) +VALUES ($1, NOW(), NOW(), $2, $3, null) +RETURNING *; \ No newline at end of file diff --git a/sql/queries/GetRefreshTokenByToken.sql b/sql/queries/GetRefreshTokenByToken.sql new file mode 100644 index 0000000..9a67fa1 --- /dev/null +++ b/sql/queries/GetRefreshTokenByToken.sql @@ -0,0 +1,2 @@ +-- name: GetRefreshTokenByToken :one +SELECT * FROM refresh_tokens WHERE token = $1 AND revoked_at is null; \ No newline at end of file diff --git a/sql/queries/GetUserByID.sql b/sql/queries/GetUserByID.sql new file mode 100644 index 0000000..37dcc53 --- /dev/null +++ b/sql/queries/GetUserByID.sql @@ -0,0 +1,2 @@ +-- name: GetUserByID :one +SELECT * FROM users WHERE id = $1; \ No newline at end of file diff --git a/sql/queries/RevokeRefreshTokenByToken.sql b/sql/queries/RevokeRefreshTokenByToken.sql new file mode 100644 index 0000000..9ff2235 --- /dev/null +++ b/sql/queries/RevokeRefreshTokenByToken.sql @@ -0,0 +1,2 @@ +-- name: RevokeRefreshTokenByToken :exec +UPDATE refresh_tokens SET revoked_at = CURRENT_TIMESTAMP WHERE token = $1; \ No newline at end of file diff --git a/sql/schema/004_refresh_tokens.sql b/sql/schema/004_refresh_tokens.sql new file mode 100644 index 0000000..3a2d6e6 --- /dev/null +++ b/sql/schema/004_refresh_tokens.sql @@ -0,0 +1,12 @@ +-- +goose Up +CREATE TABLE refresh_tokens( + token TEXT PRIMARY KEY, + created_at TIMESTAMP not null DEFAULT CURRENT_TIMESTAMP, + updated_at TIMESTAMP not null DEFAULT CURRENT_TIMESTAMP, + user_id UUID not null REFERENCES users(id) ON DELETE CASCADE, + expires_at timestamp not null, + revoked_at timestamp default null +); + +-- +goose Down +DROP TABLE refresh_tokens; \ No newline at end of file