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"
|
||||
|
||||
"go-backend/internal/database"
|
||||
"go-backend/internal/database/sqlc"
|
||||
"go-backend/internal/handlers"
|
||||
"go-backend/internal/services"
|
||||
|
||||
|
|
@ -25,8 +26,11 @@ func main() {
|
|||
}
|
||||
defer db.Close()
|
||||
|
||||
userService := services.NewUserService(db)
|
||||
authService := services.NewAuthService(db)
|
||||
queries := sqlc.New(db)
|
||||
|
||||
userService := services.NewUserService(db, queries)
|
||||
authService := services.NewAuthService(db, queries)
|
||||
|
||||
handler := handlers.NewHandler(userService, authService)
|
||||
|
||||
// Create router
|
||||
|
|
|
|||
|
|
@ -3,27 +3,13 @@ package auth
|
|||
import (
|
||||
"crypto/rand"
|
||||
"crypto/sha256"
|
||||
"crypto/subtle"
|
||||
"database/sql"
|
||||
"strings"
|
||||
"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
|
||||
|
||||
func generateSecureRandomString() (string, error) {
|
||||
func GenerateSecureRandomString() (string, error) {
|
||||
// Human readable alphabet (a-z, 0-9 without l, o, 0, 1 to avoid confusion)
|
||||
alphabet := "abcdefghijklmnpqrstuvwxyz23456789"
|
||||
|
||||
|
|
@ -43,99 +29,13 @@ func generateSecureRandomString() (string, error) {
|
|||
return id.String(), nil
|
||||
}
|
||||
|
||||
func hashSecret(secret string) []byte {
|
||||
func HashSecret(secret string) []byte {
|
||||
hash := sha256.Sum256([]byte(secret))
|
||||
return hash[:]
|
||||
}
|
||||
|
||||
func CreateSession(db *sql.DB, userID string) (*SessionWithToken, error) {
|
||||
id, err := generateSecureRandomString()
|
||||
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 CheckExpiration(expirarion time.Time) bool {
|
||||
return time.Since(expirarion).Seconds() >= sessionExpiresInSeconds
|
||||
}
|
||||
|
||||
func ParseCookies(cookieHeader string) map[string]string {
|
||||
|
|
|
|||
|
|
@ -4,7 +4,6 @@ import (
|
|||
"database/sql"
|
||||
"fmt"
|
||||
"os"
|
||||
"time"
|
||||
|
||||
_ "github.com/lib/pq"
|
||||
)
|
||||
|
|
@ -33,69 +32,9 @@ func Connect() (*sql.DB, error) {
|
|||
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 {
|
||||
if value := os.Getenv(key); value != "" {
|
||||
return value
|
||||
}
|
||||
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
|
||||
|
||||
import (
|
||||
"go-backend/internal/auth"
|
||||
"go-backend/internal/database"
|
||||
"context"
|
||||
"go-backend/internal/database/sqlc"
|
||||
"go-backend/internal/services"
|
||||
)
|
||||
|
||||
// UserService and AuthService are defined here, on the consumer side, so
|
||||
|
|
@ -10,13 +11,14 @@ import (
|
|||
// in fakes without touching the services package.
|
||||
|
||||
type UserService interface {
|
||||
FindByID(id string) (*database.User, error)
|
||||
FindByUsername(username string) (*database.User, error)
|
||||
FindByID(ctx context.Context, id string) (*sqlc.UserAuth, error)
|
||||
FindByUsername(ctx context.Context, username string) (*sqlc.UserAuth, error)
|
||||
CreateUser(ctx context.Context, userInput sqlc.InsertUserParams) (*sqlc.UserAuth, error)
|
||||
}
|
||||
|
||||
type AuthService interface {
|
||||
CreateSession(userID string) (*auth.SessionWithToken, error)
|
||||
ValidateSessionToken(token string) (*auth.Session, error)
|
||||
CreateSession(ctx context.Context, userID string) (*services.SessionWithToken, error)
|
||||
ValidateSessionToken(ctx context.Context, token string) (*sqlc.UserSession, error)
|
||||
}
|
||||
|
||||
type Handler struct {
|
||||
|
|
|
|||
|
|
@ -33,7 +33,7 @@ func (h *Handler) LoginHandler(w http.ResponseWriter, r *http.Request) {
|
|||
return
|
||||
}
|
||||
|
||||
user, err := h.users.FindByUsername(req.Username)
|
||||
user, err := h.users.FindByUsername(r.Context(), req.Username)
|
||||
if err != nil {
|
||||
log.Printf("Error finding user: %v", err)
|
||||
http.Error(w, "Internal server error", http.StatusInternalServerError)
|
||||
|
|
@ -50,7 +50,7 @@ func (h *Handler) LoginHandler(w http.ResponseWriter, r *http.Request) {
|
|||
return
|
||||
}
|
||||
|
||||
session, err := h.auth.CreateSession(user.ID)
|
||||
session, err := h.auth.CreateSession(r.Context(), user.ID)
|
||||
if err != nil {
|
||||
log.Printf("Error creating session: %v", err)
|
||||
http.Error(w, "Internal server error", http.StatusInternalServerError)
|
||||
|
|
|
|||
|
|
@ -10,6 +10,7 @@ import (
|
|||
|
||||
"go-backend/internal/auth"
|
||||
"go-backend/internal/database"
|
||||
"go-backend/internal/services"
|
||||
|
||||
"golang.org/x/crypto/argon2"
|
||||
)
|
||||
|
|
@ -32,7 +33,7 @@ func (f *fakeUsers) FindByID(id string) (*database.User, error) {
|
|||
}
|
||||
|
||||
type fakeAuth struct {
|
||||
session *auth.SessionWithToken
|
||||
session *services.SessionWithToken
|
||||
err error
|
||||
}
|
||||
|
||||
|
|
@ -69,7 +70,7 @@ func TestLoginHandler(t *testing.T) {
|
|||
name: "successful login sets session cookie",
|
||||
body: `{"username":"alice","password":"correct-horse"}`,
|
||||
users: &fakeUsers{user: validUser},
|
||||
auth: &fakeAuth{session: &auth.SessionWithToken{Token: "abc.def"}},
|
||||
auth: &fakeAuth{session: &services.SessionWithToken{Token: "abc.def"}},
|
||||
wantStatus: http.StatusOK,
|
||||
},
|
||||
{
|
||||
|
|
|
|||
|
|
@ -2,12 +2,14 @@ package handlers
|
|||
|
||||
import (
|
||||
"encoding/json"
|
||||
"go-backend/internal/database/sqlc"
|
||||
"net/http"
|
||||
)
|
||||
|
||||
type SignUpRequest struct {
|
||||
Username string `json:"username"`
|
||||
Password string `json:"password"`
|
||||
Email string `json:"email"`
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
|
||||
if req.Username == "" || req.Password == "" {
|
||||
if req.Username == "" || req.Password == "" || req.Email == "" {
|
||||
http.Error(w, "Invalid request body", http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
|
||||
user, err := h.users.FindByUsername(req.Username)
|
||||
user, err := h.users.FindByUsername(r.Context(), req.Username)
|
||||
|
||||
if err != nil {
|
||||
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)
|
||||
}
|
||||
|
||||
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
|
||||
session, err := h.auth.ValidateSessionToken(sessionToken)
|
||||
session, err := h.auth.ValidateSessionToken(r.Context(), sessionToken)
|
||||
if err != nil {
|
||||
http.Error(w, "Internal server error", http.StatusInternalServerError)
|
||||
return
|
||||
|
|
@ -33,7 +33,7 @@ func (h *Handler) UserHandler(w http.ResponseWriter, r *http.Request) {
|
|||
}
|
||||
|
||||
// Get user info
|
||||
user, err := h.users.FindByID(session.UserID)
|
||||
user, err := h.users.FindByID(r.Context(), session.UserID)
|
||||
if err != nil {
|
||||
http.Error(w, "Internal server error", http.StatusInternalServerError)
|
||||
return
|
||||
|
|
|
|||
|
|
@ -1,25 +1,96 @@
|
|||
package services
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/subtle"
|
||||
"database/sql"
|
||||
"errors"
|
||||
"strings"
|
||||
|
||||
"go-backend/internal/auth"
|
||||
"go-backend/internal/database/sqlc"
|
||||
)
|
||||
|
||||
type AuthService struct {
|
||||
db *sql.DB
|
||||
queries *sqlc.Queries
|
||||
}
|
||||
|
||||
func NewAuthService(db *sql.DB) *AuthService {
|
||||
func NewAuthService(db *sql.DB, queries *sqlc.Queries) *AuthService {
|
||||
return &AuthService{
|
||||
db: db,
|
||||
queries: queries,
|
||||
}
|
||||
}
|
||||
|
||||
func (s *AuthService) CreateSession(userID string) (*auth.SessionWithToken, error) {
|
||||
return auth.CreateSession(s.db, userID)
|
||||
type SessionWithToken struct {
|
||||
Session sqlc.UserSession
|
||||
Token string
|
||||
}
|
||||
|
||||
func (s *AuthService) ValidateSessionToken(token string) (*auth.Session, error) {
|
||||
return auth.ValidateSessionToken(s.db, token)
|
||||
func (s *AuthService) CreateSession(ctx context.Context, userID string) (*SessionWithToken, error) {
|
||||
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
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"errors"
|
||||
|
||||
"go-backend/internal/database"
|
||||
"go-backend/internal/database/sqlc"
|
||||
)
|
||||
|
||||
type UserService struct {
|
||||
db *sql.DB
|
||||
queries *sqlc.Queries
|
||||
}
|
||||
|
||||
func NewUserService(db *sql.DB) *UserService {
|
||||
func NewUserService(db *sql.DB, queries *sqlc.Queries) *UserService {
|
||||
return &UserService{
|
||||
db: db,
|
||||
queries: queries,
|
||||
}
|
||||
}
|
||||
|
||||
func (s *UserService) FindByID(id string) (*database.User, error) {
|
||||
return database.FindUserByID(s.db, id)
|
||||
func (s *UserService) FindByID(ctx context.Context, id string) (*sqlc.UserAuth, error) {
|
||||
user, err := s.queries.SelectUserById(ctx, id)
|
||||
|
||||
if err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return nil, nil
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
|
||||
func (s *UserService) FindByUsername(username string) (*database.User, error) {
|
||||
return database.FindUserByUsername(s.db, username)
|
||||
return &user, nil
|
||||
}
|
||||
|
||||
func (s *UserService) FindByUsername(ctx context.Context, username string) (*sqlc.UserAuth, error) {
|
||||
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,
|
||||
secret_hash BYTEA NOT NULL,
|
||||
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);
|
||||
|
|
|
|||
Loading…
Reference in a new issue