284 lines
6.6 KiB
Go
284 lines
6.6 KiB
Go
package db
|
|
|
|
import (
|
|
"context"
|
|
"database/sql"
|
|
"errors"
|
|
"fmt"
|
|
"time"
|
|
)
|
|
|
|
type CreateSessionParams struct {
|
|
UserID int64
|
|
TokenHash string
|
|
ExpiresAt time.Time
|
|
IPAddress string
|
|
UserAgent string
|
|
}
|
|
|
|
func (r *SessionRepository) Create(ctx context.Context, params CreateSessionParams) (*Session, error) {
|
|
result, err := r.q.ExecContext(
|
|
ctx,
|
|
`INSERT INTO sessions (user_id, token_hash, expires_at, ip_address, user_agent) VALUES (?, ?, ?, ?, ?)`,
|
|
params.UserID,
|
|
params.TokenHash,
|
|
formatTimestamp(params.ExpiresAt),
|
|
params.IPAddress,
|
|
params.UserAgent,
|
|
)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("insert session: %w", err)
|
|
}
|
|
|
|
sessionID, err := result.LastInsertId()
|
|
if err != nil {
|
|
return nil, fmt.Errorf("load inserted session id: %w", err)
|
|
}
|
|
|
|
return r.GetByID(ctx, sessionID)
|
|
}
|
|
|
|
func (r *SessionRepository) GetByID(ctx context.Context, id int64) (*Session, error) {
|
|
session, err := scanSession(r.q.QueryRowContext(
|
|
ctx,
|
|
`SELECT id, user_id, token_hash, expires_at, last_seen_at, invalidated_at, ip_address, user_agent, created_at
|
|
FROM sessions
|
|
WHERE id = ?`,
|
|
id,
|
|
))
|
|
if err != nil {
|
|
if errors.Is(err, sql.ErrNoRows) {
|
|
return nil, ErrNotFound
|
|
}
|
|
|
|
return nil, fmt.Errorf("scan session by id: %w", err)
|
|
}
|
|
|
|
return session, nil
|
|
}
|
|
|
|
func (r *SessionRepository) GetActiveWithUserByTokenHash(ctx context.Context, tokenHash string, now time.Time) (*SessionWithUser, error) {
|
|
record, err := scanSessionWithUser(r.q.QueryRowContext(
|
|
ctx,
|
|
`SELECT
|
|
s.id,
|
|
s.user_id,
|
|
s.token_hash,
|
|
s.expires_at,
|
|
s.last_seen_at,
|
|
s.invalidated_at,
|
|
s.ip_address,
|
|
s.user_agent,
|
|
s.created_at,
|
|
u.id,
|
|
u.email,
|
|
u.password_hash,
|
|
u.role,
|
|
u.is_active,
|
|
u.created_at,
|
|
u.updated_at,
|
|
u.last_login_at
|
|
FROM sessions AS s
|
|
INNER JOIN users AS u ON u.id = s.user_id
|
|
WHERE s.token_hash = ?
|
|
AND s.invalidated_at IS NULL
|
|
AND s.expires_at > ?
|
|
AND u.is_active = 1
|
|
LIMIT 1`,
|
|
tokenHash,
|
|
formatTimestamp(now),
|
|
))
|
|
if err != nil {
|
|
if errors.Is(err, sql.ErrNoRows) {
|
|
return nil, ErrNotFound
|
|
}
|
|
|
|
return nil, fmt.Errorf("scan active session by token hash: %w", err)
|
|
}
|
|
|
|
return record, nil
|
|
}
|
|
|
|
func (r *SessionRepository) Touch(ctx context.Context, sessionID int64, seenAt time.Time) error {
|
|
result, err := r.q.ExecContext(
|
|
ctx,
|
|
`UPDATE sessions SET last_seen_at = ? WHERE id = ?`,
|
|
formatTimestamp(seenAt),
|
|
sessionID,
|
|
)
|
|
if err != nil {
|
|
return fmt.Errorf("update session last_seen_at: %w", err)
|
|
}
|
|
|
|
rowsAffected, err := result.RowsAffected()
|
|
if err != nil {
|
|
return fmt.Errorf("read affected session rows: %w", err)
|
|
}
|
|
|
|
if rowsAffected == 0 {
|
|
return ErrNotFound
|
|
}
|
|
|
|
return nil
|
|
}
|
|
|
|
func (r *SessionRepository) InvalidateByTokenHash(ctx context.Context, tokenHash string, invalidatedAt time.Time) error {
|
|
result, err := r.q.ExecContext(
|
|
ctx,
|
|
`UPDATE sessions
|
|
SET invalidated_at = ?
|
|
WHERE token_hash = ?
|
|
AND invalidated_at IS NULL`,
|
|
formatTimestamp(invalidatedAt),
|
|
tokenHash,
|
|
)
|
|
if err != nil {
|
|
return fmt.Errorf("invalidate session by token hash: %w", err)
|
|
}
|
|
|
|
rowsAffected, err := result.RowsAffected()
|
|
if err != nil {
|
|
return fmt.Errorf("read invalidated session rows: %w", err)
|
|
}
|
|
|
|
if rowsAffected == 0 {
|
|
return ErrNotFound
|
|
}
|
|
|
|
return nil
|
|
}
|
|
|
|
func scanSession(scanner rowScanner) (*Session, error) {
|
|
var (
|
|
session Session
|
|
expiresAtRaw string
|
|
lastSeenAtRaw sql.NullString
|
|
invalidatedAtRaw sql.NullString
|
|
createdAtRaw string
|
|
)
|
|
|
|
if err := scanner.Scan(
|
|
&session.ID,
|
|
&session.UserID,
|
|
&session.TokenHash,
|
|
&expiresAtRaw,
|
|
&lastSeenAtRaw,
|
|
&invalidatedAtRaw,
|
|
&session.IPAddress,
|
|
&session.UserAgent,
|
|
&createdAtRaw,
|
|
); err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
expiresAt, err := parseTimestamp(expiresAtRaw)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("parse session expires_at: %w", err)
|
|
}
|
|
|
|
lastSeenAt, err := parseNullableTimestamp(lastSeenAtRaw)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("parse session last_seen_at: %w", err)
|
|
}
|
|
|
|
invalidatedAt, err := parseNullableTimestamp(invalidatedAtRaw)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("parse session invalidated_at: %w", err)
|
|
}
|
|
|
|
createdAt, err := parseTimestamp(createdAtRaw)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("parse session created_at: %w", err)
|
|
}
|
|
|
|
session.ExpiresAt = expiresAt
|
|
session.LastSeenAt = lastSeenAt
|
|
session.InvalidatedAt = invalidatedAt
|
|
session.CreatedAt = createdAt
|
|
|
|
return &session, nil
|
|
}
|
|
|
|
func scanSessionWithUser(scanner rowScanner) (*SessionWithUser, error) {
|
|
var (
|
|
record SessionWithUser
|
|
sessionExpiresRaw string
|
|
sessionLastSeenRaw sql.NullString
|
|
sessionInvalidRaw sql.NullString
|
|
sessionCreatedRaw string
|
|
userRole string
|
|
userIsActive int
|
|
userCreatedRaw string
|
|
userUpdatedRaw string
|
|
userLastLoginRaw sql.NullString
|
|
)
|
|
|
|
if err := scanner.Scan(
|
|
&record.Session.ID,
|
|
&record.Session.UserID,
|
|
&record.Session.TokenHash,
|
|
&sessionExpiresRaw,
|
|
&sessionLastSeenRaw,
|
|
&sessionInvalidRaw,
|
|
&record.Session.IPAddress,
|
|
&record.Session.UserAgent,
|
|
&sessionCreatedRaw,
|
|
&record.User.ID,
|
|
&record.User.Email,
|
|
&record.User.PasswordHash,
|
|
&userRole,
|
|
&userIsActive,
|
|
&userCreatedRaw,
|
|
&userUpdatedRaw,
|
|
&userLastLoginRaw,
|
|
); err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
sessionExpiresAt, err := parseTimestamp(sessionExpiresRaw)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("parse joined session expires_at: %w", err)
|
|
}
|
|
|
|
sessionLastSeenAt, err := parseNullableTimestamp(sessionLastSeenRaw)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("parse joined session last_seen_at: %w", err)
|
|
}
|
|
|
|
sessionInvalidatedAt, err := parseNullableTimestamp(sessionInvalidRaw)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("parse joined session invalidated_at: %w", err)
|
|
}
|
|
|
|
sessionCreatedAt, err := parseTimestamp(sessionCreatedRaw)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("parse joined session created_at: %w", err)
|
|
}
|
|
|
|
userCreatedAt, err := parseTimestamp(userCreatedRaw)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("parse joined user created_at: %w", err)
|
|
}
|
|
|
|
userUpdatedAt, err := parseTimestamp(userUpdatedRaw)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("parse joined user updated_at: %w", err)
|
|
}
|
|
|
|
userLastLoginAt, err := parseNullableTimestamp(userLastLoginRaw)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("parse joined user last_login_at: %w", err)
|
|
}
|
|
|
|
record.Session.ExpiresAt = sessionExpiresAt
|
|
record.Session.LastSeenAt = sessionLastSeenAt
|
|
record.Session.InvalidatedAt = sessionInvalidatedAt
|
|
record.Session.CreatedAt = sessionCreatedAt
|
|
record.User.Role = UserRole(userRole)
|
|
record.User.IsActive = userIsActive == 1
|
|
record.User.CreatedAt = userCreatedAt
|
|
record.User.UpdatedAt = userUpdatedAt
|
|
record.User.LastLoginAt = userLastLoginAt
|
|
|
|
return &record, nil
|
|
}
|