Move all inline sql queries to sqlc queries
This commit is contained in:
parent
6301c4f3b5
commit
6b8f02fa0f
18 changed files with 370 additions and 209 deletions
|
|
@ -6,6 +6,7 @@ import (
|
||||||
"os"
|
"os"
|
||||||
|
|
||||||
"go-backend/internal/database"
|
"go-backend/internal/database"
|
||||||
|
"go-backend/internal/database/sqlc"
|
||||||
"go-backend/internal/handlers"
|
"go-backend/internal/handlers"
|
||||||
"go-backend/internal/services"
|
"go-backend/internal/services"
|
||||||
|
|
||||||
|
|
@ -25,8 +26,11 @@ func main() {
|
||||||
}
|
}
|
||||||
defer db.Close()
|
defer db.Close()
|
||||||
|
|
||||||
userService := services.NewUserService(db)
|
queries := sqlc.New(db)
|
||||||
authService := services.NewAuthService(db)
|
|
||||||
|
userService := services.NewUserService(db, queries)
|
||||||
|
authService := services.NewAuthService(db, queries)
|
||||||
|
|
||||||
handler := handlers.NewHandler(userService, authService)
|
handler := handlers.NewHandler(userService, authService)
|
||||||
|
|
||||||
// Create router
|
// Create router
|
||||||
|
|
|
||||||
|
|
@ -3,27 +3,13 @@ package auth
|
||||||
import (
|
import (
|
||||||
"crypto/rand"
|
"crypto/rand"
|
||||||
"crypto/sha256"
|
"crypto/sha256"
|
||||||
"crypto/subtle"
|
|
||||||
"database/sql"
|
|
||||||
"strings"
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
)
|
)
|
||||||
|
|
||||||
type Session struct {
|
|
||||||
ID string
|
|
||||||
SecretHash []byte
|
|
||||||
CreatedAt time.Time
|
|
||||||
UserID string
|
|
||||||
}
|
|
||||||
|
|
||||||
type SessionWithToken struct {
|
|
||||||
Session
|
|
||||||
Token string
|
|
||||||
}
|
|
||||||
|
|
||||||
const sessionExpiresInSeconds = 60 * 60 * 24 // 1 day
|
const sessionExpiresInSeconds = 60 * 60 * 24 // 1 day
|
||||||
|
|
||||||
func generateSecureRandomString() (string, error) {
|
func GenerateSecureRandomString() (string, error) {
|
||||||
// Human readable alphabet (a-z, 0-9 without l, o, 0, 1 to avoid confusion)
|
// Human readable alphabet (a-z, 0-9 without l, o, 0, 1 to avoid confusion)
|
||||||
alphabet := "abcdefghijklmnpqrstuvwxyz23456789"
|
alphabet := "abcdefghijklmnpqrstuvwxyz23456789"
|
||||||
|
|
||||||
|
|
@ -43,99 +29,13 @@ func generateSecureRandomString() (string, error) {
|
||||||
return id.String(), nil
|
return id.String(), nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func hashSecret(secret string) []byte {
|
func HashSecret(secret string) []byte {
|
||||||
hash := sha256.Sum256([]byte(secret))
|
hash := sha256.Sum256([]byte(secret))
|
||||||
return hash[:]
|
return hash[:]
|
||||||
}
|
}
|
||||||
|
|
||||||
func CreateSession(db *sql.DB, userID string) (*SessionWithToken, error) {
|
func CheckExpiration(expirarion time.Time) bool {
|
||||||
id, err := generateSecureRandomString()
|
return time.Since(expirarion).Seconds() >= sessionExpiresInSeconds
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
secret, err := generateSecureRandomString()
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
secretHash := hashSecret(secret)
|
|
||||||
token := id + "." + secret
|
|
||||||
now := time.Now()
|
|
||||||
|
|
||||||
// Insert session into database
|
|
||||||
_, err = db.Exec(
|
|
||||||
"INSERT INTO user_session (id, secret_hash, user_id, created_at) VALUES ($1, $2, $3, $4)",
|
|
||||||
id, secretHash, userID, now.Unix(),
|
|
||||||
)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
return &SessionWithToken{
|
|
||||||
Session: Session{
|
|
||||||
ID: id,
|
|
||||||
SecretHash: secretHash,
|
|
||||||
CreatedAt: now,
|
|
||||||
UserID: userID,
|
|
||||||
},
|
|
||||||
Token: token,
|
|
||||||
}, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func ValidateSessionToken(db *sql.DB, token string) (*Session, error) {
|
|
||||||
tokenParts := strings.Split(token, ".")
|
|
||||||
if len(tokenParts) != 2 {
|
|
||||||
return nil, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
sessionID := tokenParts[0]
|
|
||||||
sessionSecret := tokenParts[1]
|
|
||||||
|
|
||||||
session, err := getSession(db, sessionID)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
if session == nil {
|
|
||||||
return nil, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
tokenSecretHash := hashSecret(sessionSecret)
|
|
||||||
if subtle.ConstantTimeCompare(tokenSecretHash, session.SecretHash) != 1 {
|
|
||||||
return nil, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
return session, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func getSession(db *sql.DB, sessionID string) (*Session, error) {
|
|
||||||
var session Session
|
|
||||||
var createdAtUnix int64
|
|
||||||
|
|
||||||
err := db.QueryRow(
|
|
||||||
"SELECT id, secret_hash, user_id, created_at FROM user_session WHERE id = $1",
|
|
||||||
sessionID,
|
|
||||||
).Scan(&session.ID, &session.SecretHash, &session.UserID, &createdAtUnix)
|
|
||||||
|
|
||||||
if err == sql.ErrNoRows {
|
|
||||||
return nil, nil
|
|
||||||
}
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
session.CreatedAt = time.Unix(createdAtUnix, 0)
|
|
||||||
|
|
||||||
// Check expiration
|
|
||||||
if time.Since(session.CreatedAt).Seconds() >= sessionExpiresInSeconds {
|
|
||||||
_, err = db.Exec("DELETE FROM user_session WHERE id = $1", sessionID)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
return nil, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
return &session, nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func ParseCookies(cookieHeader string) map[string]string {
|
func ParseCookies(cookieHeader string) map[string]string {
|
||||||
|
|
|
||||||
|
|
@ -4,7 +4,6 @@ import (
|
||||||
"database/sql"
|
"database/sql"
|
||||||
"fmt"
|
"fmt"
|
||||||
"os"
|
"os"
|
||||||
"time"
|
|
||||||
|
|
||||||
_ "github.com/lib/pq"
|
_ "github.com/lib/pq"
|
||||||
)
|
)
|
||||||
|
|
@ -33,69 +32,9 @@ func Connect() (*sql.DB, error) {
|
||||||
return db, nil
|
return db, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func FindUserByID(db *sql.DB, id string) (*User, error) {
|
|
||||||
var user User
|
|
||||||
var createdAtUnix, updatedAtUnix int64
|
|
||||||
|
|
||||||
err := db.QueryRow(`
|
|
||||||
SELECT id, username, email, password_hash,
|
|
||||||
EXTRACT(EPOCH FROM created_at)::bigint,
|
|
||||||
EXTRACT(EPOCH FROM updated_at)::bigint
|
|
||||||
FROM user_auth
|
|
||||||
WHERE id = $1
|
|
||||||
`, id).Scan(
|
|
||||||
&user.ID, &user.Username, &user.Email, &user.PasswordHash,
|
|
||||||
&createdAtUnix, &updatedAtUnix,
|
|
||||||
)
|
|
||||||
|
|
||||||
if err == sql.ErrNoRows {
|
|
||||||
return nil, nil
|
|
||||||
}
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
user.CreatedAt = timeFromUnix(createdAtUnix)
|
|
||||||
user.UpdatedAt = timeFromUnix(updatedAtUnix)
|
|
||||||
|
|
||||||
return &user, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func FindUserByUsername(db *sql.DB, username string) (*User, error) {
|
|
||||||
var user User
|
|
||||||
var createdAtUnix, updatedAtUnix int64
|
|
||||||
|
|
||||||
err := db.QueryRow(`
|
|
||||||
SELECT id, username, email, password_hash,
|
|
||||||
EXTRACT(EPOCH FROM created_at)::bigint,
|
|
||||||
EXTRACT(EPOCH FROM updated_at)::bigint
|
|
||||||
FROM user_auth
|
|
||||||
WHERE username = $1
|
|
||||||
`, username).Scan(
|
|
||||||
&user.ID, &user.Username, &user.Email, &user.PasswordHash,
|
|
||||||
&createdAtUnix, &updatedAtUnix,
|
|
||||||
)
|
|
||||||
|
|
||||||
if err == sql.ErrNoRows {
|
|
||||||
return nil, nil
|
|
||||||
}
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
user.CreatedAt = timeFromUnix(createdAtUnix)
|
|
||||||
user.UpdatedAt = timeFromUnix(updatedAtUnix)
|
|
||||||
|
|
||||||
return &user, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func getEnv(key, defaultValue string) string {
|
func getEnv(key, defaultValue string) string {
|
||||||
if value := os.Getenv(key); value != "" {
|
if value := os.Getenv(key); value != "" {
|
||||||
return value
|
return value
|
||||||
}
|
}
|
||||||
return defaultValue
|
return defaultValue
|
||||||
}
|
}
|
||||||
|
|
||||||
func timeFromUnix(timestamp int64) time.Time {
|
|
||||||
return time.Unix(timestamp, 0)
|
|
||||||
}
|
|
||||||
|
|
|
||||||
|
|
@ -1,12 +0,0 @@
|
||||||
package database
|
|
||||||
|
|
||||||
import "time"
|
|
||||||
|
|
||||||
type User struct {
|
|
||||||
ID string `json:"id"`
|
|
||||||
Username string `json:"username"`
|
|
||||||
Email string `json:"email"`
|
|
||||||
PasswordHash string `json:"-"` // Don't serialize password hash
|
|
||||||
CreatedAt time.Time `json:"created_at"`
|
|
||||||
UpdatedAt time.Time `json:"updated_at"`
|
|
||||||
}
|
|
||||||
12
internal/database/queries/sessions.sql
Normal file
12
internal/database/queries/sessions.sql
Normal file
|
|
@ -0,0 +1,12 @@
|
||||||
|
-- name: SelectSessionById :one
|
||||||
|
SELECT id, secret_hash, user_id, created_at
|
||||||
|
FROM user_session
|
||||||
|
WHERE id = $1;
|
||||||
|
|
||||||
|
-- name: InsertSession :one
|
||||||
|
INSERT INTO user_session (id, secret_hash, user_id)
|
||||||
|
VALUES ($1, $2, $3)
|
||||||
|
RETURNING id, secret_hash, user_id, created_at;
|
||||||
|
|
||||||
|
-- name: DeleteSessionById :exec
|
||||||
|
DELETE FROM user_session WHERE id = $1;
|
||||||
14
internal/database/queries/users.sql
Normal file
14
internal/database/queries/users.sql
Normal file
|
|
@ -0,0 +1,14 @@
|
||||||
|
-- name: SelectUserById :one
|
||||||
|
SELECT id, username, email, password_hash, created_at, updated_at
|
||||||
|
FROM user_auth
|
||||||
|
WHERE id = $1;
|
||||||
|
|
||||||
|
-- name: SelectUserByUsername :one
|
||||||
|
SELECT id, username, email, password_hash, created_at, updated_at
|
||||||
|
FROM user_auth
|
||||||
|
WHERE username = $1;
|
||||||
|
|
||||||
|
-- name: InsertUser :one
|
||||||
|
INSERT INTO user_auth (username, email, password_hash)
|
||||||
|
VALUES ($1, $2, $3)
|
||||||
|
RETURNING id, username, email, password_hash, created_at, updated_at;
|
||||||
31
internal/database/sqlc/db.go
Normal file
31
internal/database/sqlc/db.go
Normal file
|
|
@ -0,0 +1,31 @@
|
||||||
|
// Code generated by sqlc. DO NOT EDIT.
|
||||||
|
// versions:
|
||||||
|
// sqlc v1.29.0
|
||||||
|
|
||||||
|
package sqlc
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"database/sql"
|
||||||
|
)
|
||||||
|
|
||||||
|
type DBTX interface {
|
||||||
|
ExecContext(context.Context, string, ...interface{}) (sql.Result, error)
|
||||||
|
PrepareContext(context.Context, string) (*sql.Stmt, error)
|
||||||
|
QueryContext(context.Context, string, ...interface{}) (*sql.Rows, error)
|
||||||
|
QueryRowContext(context.Context, string, ...interface{}) *sql.Row
|
||||||
|
}
|
||||||
|
|
||||||
|
func New(db DBTX) *Queries {
|
||||||
|
return &Queries{db: db}
|
||||||
|
}
|
||||||
|
|
||||||
|
type Queries struct {
|
||||||
|
db DBTX
|
||||||
|
}
|
||||||
|
|
||||||
|
func (q *Queries) WithTx(tx *sql.Tx) *Queries {
|
||||||
|
return &Queries{
|
||||||
|
db: tx,
|
||||||
|
}
|
||||||
|
}
|
||||||
25
internal/database/sqlc/models.go
Normal file
25
internal/database/sqlc/models.go
Normal file
|
|
@ -0,0 +1,25 @@
|
||||||
|
// Code generated by sqlc. DO NOT EDIT.
|
||||||
|
// versions:
|
||||||
|
// sqlc v1.29.0
|
||||||
|
|
||||||
|
package sqlc
|
||||||
|
|
||||||
|
import (
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
type UserAuth struct {
|
||||||
|
ID string
|
||||||
|
Username string
|
||||||
|
Email string
|
||||||
|
PasswordHash string
|
||||||
|
CreatedAt time.Time
|
||||||
|
UpdatedAt time.Time
|
||||||
|
}
|
||||||
|
|
||||||
|
type UserSession struct {
|
||||||
|
ID string
|
||||||
|
SecretHash []byte
|
||||||
|
UserID string
|
||||||
|
CreatedAt time.Time
|
||||||
|
}
|
||||||
61
internal/database/sqlc/sessions.sql.go
Normal file
61
internal/database/sqlc/sessions.sql.go
Normal file
|
|
@ -0,0 +1,61 @@
|
||||||
|
// Code generated by sqlc. DO NOT EDIT.
|
||||||
|
// versions:
|
||||||
|
// sqlc v1.29.0
|
||||||
|
// source: sessions.sql
|
||||||
|
|
||||||
|
package sqlc
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
)
|
||||||
|
|
||||||
|
const deleteSessionById = `-- name: DeleteSessionById :exec
|
||||||
|
DELETE FROM user_session WHERE id = $1
|
||||||
|
`
|
||||||
|
|
||||||
|
func (q *Queries) DeleteSessionById(ctx context.Context, id string) error {
|
||||||
|
_, err := q.db.ExecContext(ctx, deleteSessionById, id)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
const insertSession = `-- name: InsertSession :one
|
||||||
|
INSERT INTO user_session (id, secret_hash, user_id)
|
||||||
|
VALUES ($1, $2, $3)
|
||||||
|
RETURNING id, secret_hash, user_id, created_at
|
||||||
|
`
|
||||||
|
|
||||||
|
type InsertSessionParams struct {
|
||||||
|
ID string
|
||||||
|
SecretHash []byte
|
||||||
|
UserID string
|
||||||
|
}
|
||||||
|
|
||||||
|
func (q *Queries) InsertSession(ctx context.Context, arg InsertSessionParams) (UserSession, error) {
|
||||||
|
row := q.db.QueryRowContext(ctx, insertSession, arg.ID, arg.SecretHash, arg.UserID)
|
||||||
|
var i UserSession
|
||||||
|
err := row.Scan(
|
||||||
|
&i.ID,
|
||||||
|
&i.SecretHash,
|
||||||
|
&i.UserID,
|
||||||
|
&i.CreatedAt,
|
||||||
|
)
|
||||||
|
return i, err
|
||||||
|
}
|
||||||
|
|
||||||
|
const selectSessionById = `-- name: SelectSessionById :one
|
||||||
|
SELECT id, secret_hash, user_id, created_at
|
||||||
|
FROM user_session
|
||||||
|
WHERE id = $1
|
||||||
|
`
|
||||||
|
|
||||||
|
func (q *Queries) SelectSessionById(ctx context.Context, id string) (UserSession, error) {
|
||||||
|
row := q.db.QueryRowContext(ctx, selectSessionById, id)
|
||||||
|
var i UserSession
|
||||||
|
err := row.Scan(
|
||||||
|
&i.ID,
|
||||||
|
&i.SecretHash,
|
||||||
|
&i.UserID,
|
||||||
|
&i.CreatedAt,
|
||||||
|
)
|
||||||
|
return i, err
|
||||||
|
}
|
||||||
76
internal/database/sqlc/users.sql.go
Normal file
76
internal/database/sqlc/users.sql.go
Normal file
|
|
@ -0,0 +1,76 @@
|
||||||
|
// Code generated by sqlc. DO NOT EDIT.
|
||||||
|
// versions:
|
||||||
|
// sqlc v1.29.0
|
||||||
|
// source: users.sql
|
||||||
|
|
||||||
|
package sqlc
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
)
|
||||||
|
|
||||||
|
const insertUser = `-- name: InsertUser :one
|
||||||
|
INSERT INTO user_auth (username, email, password_hash)
|
||||||
|
VALUES ($1, $2, $3)
|
||||||
|
RETURNING id, username, email, password_hash, created_at, updated_at
|
||||||
|
`
|
||||||
|
|
||||||
|
type InsertUserParams struct {
|
||||||
|
Username string
|
||||||
|
Email string
|
||||||
|
PasswordHash string
|
||||||
|
}
|
||||||
|
|
||||||
|
func (q *Queries) InsertUser(ctx context.Context, arg InsertUserParams) (UserAuth, error) {
|
||||||
|
row := q.db.QueryRowContext(ctx, insertUser, arg.Username, arg.Email, arg.PasswordHash)
|
||||||
|
var i UserAuth
|
||||||
|
err := row.Scan(
|
||||||
|
&i.ID,
|
||||||
|
&i.Username,
|
||||||
|
&i.Email,
|
||||||
|
&i.PasswordHash,
|
||||||
|
&i.CreatedAt,
|
||||||
|
&i.UpdatedAt,
|
||||||
|
)
|
||||||
|
return i, err
|
||||||
|
}
|
||||||
|
|
||||||
|
const selectUserById = `-- name: SelectUserById :one
|
||||||
|
SELECT id, username, email, password_hash, created_at, updated_at
|
||||||
|
FROM user_auth
|
||||||
|
WHERE id = $1
|
||||||
|
`
|
||||||
|
|
||||||
|
func (q *Queries) SelectUserById(ctx context.Context, id string) (UserAuth, error) {
|
||||||
|
row := q.db.QueryRowContext(ctx, selectUserById, id)
|
||||||
|
var i UserAuth
|
||||||
|
err := row.Scan(
|
||||||
|
&i.ID,
|
||||||
|
&i.Username,
|
||||||
|
&i.Email,
|
||||||
|
&i.PasswordHash,
|
||||||
|
&i.CreatedAt,
|
||||||
|
&i.UpdatedAt,
|
||||||
|
)
|
||||||
|
return i, err
|
||||||
|
}
|
||||||
|
|
||||||
|
const selectUserByUsername = `-- name: SelectUserByUsername :one
|
||||||
|
SELECT id, username, email, password_hash, created_at, updated_at
|
||||||
|
FROM user_auth
|
||||||
|
WHERE username = $1
|
||||||
|
`
|
||||||
|
|
||||||
|
func (q *Queries) SelectUserByUsername(ctx context.Context, username string) (UserAuth, error) {
|
||||||
|
row := q.db.QueryRowContext(ctx, selectUserByUsername, username)
|
||||||
|
var i UserAuth
|
||||||
|
err := row.Scan(
|
||||||
|
&i.ID,
|
||||||
|
&i.Username,
|
||||||
|
&i.Email,
|
||||||
|
&i.PasswordHash,
|
||||||
|
&i.CreatedAt,
|
||||||
|
&i.UpdatedAt,
|
||||||
|
)
|
||||||
|
return i, err
|
||||||
|
}
|
||||||
|
|
@ -1,8 +1,9 @@
|
||||||
package handlers
|
package handlers
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"go-backend/internal/auth"
|
"context"
|
||||||
"go-backend/internal/database"
|
"go-backend/internal/database/sqlc"
|
||||||
|
"go-backend/internal/services"
|
||||||
)
|
)
|
||||||
|
|
||||||
// UserService and AuthService are defined here, on the consumer side, so
|
// UserService and AuthService are defined here, on the consumer side, so
|
||||||
|
|
@ -10,13 +11,14 @@ import (
|
||||||
// in fakes without touching the services package.
|
// in fakes without touching the services package.
|
||||||
|
|
||||||
type UserService interface {
|
type UserService interface {
|
||||||
FindByID(id string) (*database.User, error)
|
FindByID(ctx context.Context, id string) (*sqlc.UserAuth, error)
|
||||||
FindByUsername(username string) (*database.User, error)
|
FindByUsername(ctx context.Context, username string) (*sqlc.UserAuth, error)
|
||||||
|
CreateUser(ctx context.Context, userInput sqlc.InsertUserParams) (*sqlc.UserAuth, error)
|
||||||
}
|
}
|
||||||
|
|
||||||
type AuthService interface {
|
type AuthService interface {
|
||||||
CreateSession(userID string) (*auth.SessionWithToken, error)
|
CreateSession(ctx context.Context, userID string) (*services.SessionWithToken, error)
|
||||||
ValidateSessionToken(token string) (*auth.Session, error)
|
ValidateSessionToken(ctx context.Context, token string) (*sqlc.UserSession, error)
|
||||||
}
|
}
|
||||||
|
|
||||||
type Handler struct {
|
type Handler struct {
|
||||||
|
|
|
||||||
|
|
@ -33,7 +33,7 @@ func (h *Handler) LoginHandler(w http.ResponseWriter, r *http.Request) {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
user, err := h.users.FindByUsername(req.Username)
|
user, err := h.users.FindByUsername(r.Context(), req.Username)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
log.Printf("Error finding user: %v", err)
|
log.Printf("Error finding user: %v", err)
|
||||||
http.Error(w, "Internal server error", http.StatusInternalServerError)
|
http.Error(w, "Internal server error", http.StatusInternalServerError)
|
||||||
|
|
@ -50,7 +50,7 @@ func (h *Handler) LoginHandler(w http.ResponseWriter, r *http.Request) {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
session, err := h.auth.CreateSession(user.ID)
|
session, err := h.auth.CreateSession(r.Context(), user.ID)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
log.Printf("Error creating session: %v", err)
|
log.Printf("Error creating session: %v", err)
|
||||||
http.Error(w, "Internal server error", http.StatusInternalServerError)
|
http.Error(w, "Internal server error", http.StatusInternalServerError)
|
||||||
|
|
|
||||||
|
|
@ -10,6 +10,7 @@ import (
|
||||||
|
|
||||||
"go-backend/internal/auth"
|
"go-backend/internal/auth"
|
||||||
"go-backend/internal/database"
|
"go-backend/internal/database"
|
||||||
|
"go-backend/internal/services"
|
||||||
|
|
||||||
"golang.org/x/crypto/argon2"
|
"golang.org/x/crypto/argon2"
|
||||||
)
|
)
|
||||||
|
|
@ -32,7 +33,7 @@ func (f *fakeUsers) FindByID(id string) (*database.User, error) {
|
||||||
}
|
}
|
||||||
|
|
||||||
type fakeAuth struct {
|
type fakeAuth struct {
|
||||||
session *auth.SessionWithToken
|
session *services.SessionWithToken
|
||||||
err error
|
err error
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -69,7 +70,7 @@ func TestLoginHandler(t *testing.T) {
|
||||||
name: "successful login sets session cookie",
|
name: "successful login sets session cookie",
|
||||||
body: `{"username":"alice","password":"correct-horse"}`,
|
body: `{"username":"alice","password":"correct-horse"}`,
|
||||||
users: &fakeUsers{user: validUser},
|
users: &fakeUsers{user: validUser},
|
||||||
auth: &fakeAuth{session: &auth.SessionWithToken{Token: "abc.def"}},
|
auth: &fakeAuth{session: &services.SessionWithToken{Token: "abc.def"}},
|
||||||
wantStatus: http.StatusOK,
|
wantStatus: http.StatusOK,
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
|
|
|
||||||
|
|
@ -2,12 +2,14 @@ package handlers
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
|
"go-backend/internal/database/sqlc"
|
||||||
"net/http"
|
"net/http"
|
||||||
)
|
)
|
||||||
|
|
||||||
type SignUpRequest struct {
|
type SignUpRequest struct {
|
||||||
Username string `json:"username"`
|
Username string `json:"username"`
|
||||||
Password string `json:"password"`
|
Password string `json:"password"`
|
||||||
|
Email string `json:"email"`
|
||||||
}
|
}
|
||||||
|
|
||||||
func (h *Handler) SignUpHandler(w http.ResponseWriter, r *http.Request) {
|
func (h *Handler) SignUpHandler(w http.ResponseWriter, r *http.Request) {
|
||||||
|
|
@ -25,12 +27,13 @@ func (h *Handler) SignUpHandler(w http.ResponseWriter, r *http.Request) {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
if req.Username == "" || req.Password == "" {
|
if req.Username == "" || req.Password == "" || req.Email == "" {
|
||||||
http.Error(w, "Invalid request body", http.StatusBadRequest)
|
http.Error(w, "Invalid request body", http.StatusBadRequest)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
user, err := h.users.FindByUsername(req.Username)
|
user, err := h.users.FindByUsername(r.Context(), req.Username)
|
||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
http.Error(w, "Internal server error", http.StatusInternalServerError)
|
http.Error(w, "Internal server error", http.StatusInternalServerError)
|
||||||
}
|
}
|
||||||
|
|
@ -39,4 +42,6 @@ func (h *Handler) SignUpHandler(w http.ResponseWriter, r *http.Request) {
|
||||||
http.Error(w, "Username already taken", http.StatusBadRequest)
|
http.Error(w, "Username already taken", http.StatusBadRequest)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
user, err = h.users.CreateUser(r.Context(), sqlc.InsertUserParams{Username: req.Username, Email: req.Email, PasswordHash: req.Password})
|
||||||
|
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -22,7 +22,7 @@ func (h *Handler) UserHandler(w http.ResponseWriter, r *http.Request) {
|
||||||
}
|
}
|
||||||
|
|
||||||
// Validate session
|
// Validate session
|
||||||
session, err := h.auth.ValidateSessionToken(sessionToken)
|
session, err := h.auth.ValidateSessionToken(r.Context(), sessionToken)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
http.Error(w, "Internal server error", http.StatusInternalServerError)
|
http.Error(w, "Internal server error", http.StatusInternalServerError)
|
||||||
return
|
return
|
||||||
|
|
@ -33,7 +33,7 @@ func (h *Handler) UserHandler(w http.ResponseWriter, r *http.Request) {
|
||||||
}
|
}
|
||||||
|
|
||||||
// Get user info
|
// Get user info
|
||||||
user, err := h.users.FindByID(session.UserID)
|
user, err := h.users.FindByID(r.Context(), session.UserID)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
http.Error(w, "Internal server error", http.StatusInternalServerError)
|
http.Error(w, "Internal server error", http.StatusInternalServerError)
|
||||||
return
|
return
|
||||||
|
|
|
||||||
|
|
@ -1,25 +1,96 @@
|
||||||
package services
|
package services
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
|
"crypto/subtle"
|
||||||
"database/sql"
|
"database/sql"
|
||||||
|
"errors"
|
||||||
|
"strings"
|
||||||
|
|
||||||
"go-backend/internal/auth"
|
"go-backend/internal/auth"
|
||||||
|
"go-backend/internal/database/sqlc"
|
||||||
)
|
)
|
||||||
|
|
||||||
type AuthService struct {
|
type AuthService struct {
|
||||||
db *sql.DB
|
db *sql.DB
|
||||||
|
queries *sqlc.Queries
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewAuthService(db *sql.DB) *AuthService {
|
func NewAuthService(db *sql.DB, queries *sqlc.Queries) *AuthService {
|
||||||
return &AuthService{
|
return &AuthService{
|
||||||
db: db,
|
db: db,
|
||||||
|
queries: queries,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *AuthService) CreateSession(userID string) (*auth.SessionWithToken, error) {
|
type SessionWithToken struct {
|
||||||
return auth.CreateSession(s.db, userID)
|
Session sqlc.UserSession
|
||||||
|
Token string
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *AuthService) ValidateSessionToken(token string) (*auth.Session, error) {
|
func (s *AuthService) CreateSession(ctx context.Context, userID string) (*SessionWithToken, error) {
|
||||||
return auth.ValidateSessionToken(s.db, token)
|
id, err := auth.GenerateSecureRandomString()
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
secret, err := auth.GenerateSecureRandomString()
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
secretHash := auth.HashSecret(secret)
|
||||||
|
token := id + "." + secret
|
||||||
|
|
||||||
|
session, err := s.queries.InsertSession(ctx, sqlc.InsertSessionParams{
|
||||||
|
ID: id,
|
||||||
|
SecretHash: secretHash,
|
||||||
|
UserID: userID,
|
||||||
|
})
|
||||||
|
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
return &SessionWithToken{
|
||||||
|
Session: session,
|
||||||
|
Token: token,
|
||||||
|
}, nil
|
||||||
|
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *AuthService) ValidateSessionToken(ctx context.Context, token string) (*sqlc.UserSession, error) {
|
||||||
|
tokenParts := strings.Split(token, ".")
|
||||||
|
if len(tokenParts) != 2 {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
sessionID := tokenParts[0]
|
||||||
|
sessionSecret := tokenParts[1]
|
||||||
|
|
||||||
|
session, err := s.queries.SelectSessionById(ctx, sessionID)
|
||||||
|
|
||||||
|
if err != nil {
|
||||||
|
if errors.Is(err, sql.ErrNoRows) {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
if auth.CheckExpiration(session.CreatedAt) {
|
||||||
|
err = s.queries.DeleteSessionById(ctx, sessionID)
|
||||||
|
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
tokenSecretHash := auth.HashSecret(sessionSecret)
|
||||||
|
if subtle.ConstantTimeCompare(tokenSecretHash, session.SecretHash) != 1 {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
return &session, nil
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -1,25 +1,57 @@
|
||||||
package services
|
package services
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
"database/sql"
|
"database/sql"
|
||||||
|
"errors"
|
||||||
|
|
||||||
"go-backend/internal/database"
|
"go-backend/internal/database/sqlc"
|
||||||
)
|
)
|
||||||
|
|
||||||
type UserService struct {
|
type UserService struct {
|
||||||
db *sql.DB
|
db *sql.DB
|
||||||
|
queries *sqlc.Queries
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewUserService(db *sql.DB) *UserService {
|
func NewUserService(db *sql.DB, queries *sqlc.Queries) *UserService {
|
||||||
return &UserService{
|
return &UserService{
|
||||||
db: db,
|
db: db,
|
||||||
|
queries: queries,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *UserService) FindByID(id string) (*database.User, error) {
|
func (s *UserService) FindByID(ctx context.Context, id string) (*sqlc.UserAuth, error) {
|
||||||
return database.FindUserByID(s.db, id)
|
user, err := s.queries.SelectUserById(ctx, id)
|
||||||
|
|
||||||
|
if err != nil {
|
||||||
|
if errors.Is(err, sql.ErrNoRows) {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
return &user, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *UserService) FindByUsername(username string) (*database.User, error) {
|
func (s *UserService) FindByUsername(ctx context.Context, username string) (*sqlc.UserAuth, error) {
|
||||||
return database.FindUserByUsername(s.db, username)
|
user, err := s.queries.SelectUserByUsername(ctx, username)
|
||||||
|
|
||||||
|
if err != nil {
|
||||||
|
if errors.Is(err, sql.ErrNoRows) {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
return &user, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *UserService) CreateUser(ctx context.Context, userInput sqlc.InsertUserParams) (*sqlc.UserAuth, error) {
|
||||||
|
user, err := s.queries.InsertUser(ctx, userInput)
|
||||||
|
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
return &user, nil
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -15,7 +15,7 @@ CREATE TABLE IF NOT EXISTS user_session (
|
||||||
id TEXT PRIMARY KEY,
|
id TEXT PRIMARY KEY,
|
||||||
secret_hash BYTEA NOT NULL,
|
secret_hash BYTEA NOT NULL,
|
||||||
user_id TEXT NOT NULL REFERENCES user_auth (id) ON DELETE CASCADE,
|
user_id TEXT NOT NULL REFERENCES user_auth (id) ON DELETE CASCADE,
|
||||||
created_at BIGINT NOT NULL
|
created_at TIMESTAMPTZ NOT NULL DEFAULT now()
|
||||||
);
|
);
|
||||||
|
|
||||||
CREATE INDEX IF NOT EXISTS idx_user_session_user_id ON user_session (user_id);
|
CREATE INDEX IF NOT EXISTS idx_user_session_user_id ON user_session (user_id);
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue