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 }