init
This commit is contained in:
commit
b15b95781c
108 changed files with 14802 additions and 0 deletions
284
internal/db/sessions.go
Normal file
284
internal/db/sessions.go
Normal file
|
|
@ -0,0 +1,284 @@
|
|||
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
|
||||
}
|
||||
Loading…
Add table
Add a link
Reference in a new issue