156 lines
3.2 KiB
Go
156 lines
3.2 KiB
Go
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) {
|
|
// Human readable alphabet (a-z, 0-9 without l, o, 0, 1 to avoid confusion)
|
|
alphabet := "abcdefghijklmnpqrstuvwxyz23456789"
|
|
|
|
// Generate 24 bytes = 192 bits of entropy.
|
|
// We're only going to use 5 bits per byte so the total entropy will be 192 * 5 / 8 = 120 bits
|
|
bytes := make([]byte, 24)
|
|
_, err := rand.Read(bytes)
|
|
if err != nil {
|
|
return "", err
|
|
}
|
|
|
|
var id strings.Builder
|
|
for _, b := range bytes {
|
|
// >> 3 "removes" the right-most 3 bits of the byte
|
|
id.WriteByte(alphabet[b>>3])
|
|
}
|
|
return id.String(), nil
|
|
}
|
|
|
|
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 ParseCookies(cookieHeader string) map[string]string {
|
|
cookies := make(map[string]string)
|
|
if cookieHeader == "" {
|
|
return cookies
|
|
}
|
|
|
|
pairs := strings.Split(cookieHeader, ";")
|
|
for _, pair := range pairs {
|
|
pair = strings.TrimSpace(pair)
|
|
parts := strings.SplitN(pair, "=", 2)
|
|
if len(parts) == 2 {
|
|
cookies[parts[0]] = parts[1]
|
|
}
|
|
}
|
|
return cookies
|
|
}
|