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 }