137 lines
3.7 KiB
Go
137 lines
3.7 KiB
Go
package session
|
|
|
|
import (
|
|
"crypto/hmac"
|
|
"crypto/rand"
|
|
"crypto/sha256"
|
|
"encoding/base64"
|
|
"errors"
|
|
"net/http"
|
|
"strings"
|
|
"sync"
|
|
"time"
|
|
)
|
|
|
|
const cookieName = "chorus_session"
|
|
const maxSessions = 10000
|
|
|
|
var ErrInvalidConfig = errors.New("session configuration is invalid")
|
|
var ErrCapacity = errors.New("session capacity is exhausted")
|
|
|
|
type State struct {
|
|
ID string
|
|
UserID uint64
|
|
CSRFToken string
|
|
ExpiresAt time.Time
|
|
}
|
|
type Manager struct {
|
|
mutex sync.RWMutex
|
|
sessions map[string]State
|
|
key []byte
|
|
ttl time.Duration
|
|
secure bool
|
|
now func() time.Time
|
|
}
|
|
|
|
func New(key []byte, ttl time.Duration, secure bool) (*Manager, error) {
|
|
if len(key) < 32 || ttl <= 0 {
|
|
return nil, ErrInvalidConfig
|
|
}
|
|
return &Manager{sessions: map[string]State{}, key: append([]byte(nil), key...), ttl: ttl, secure: secure, now: time.Now}, nil
|
|
}
|
|
|
|
func (m *Manager) Ensure(response http.ResponseWriter, request *http.Request) (State, error) {
|
|
if state, ok := m.Get(request); ok {
|
|
return state, nil
|
|
}
|
|
return m.rotate(response, "", 0)
|
|
}
|
|
func (m *Manager) Authenticate(response http.ResponseWriter, request *http.Request, userID uint64) (State, error) {
|
|
old, _ := m.Get(request)
|
|
return m.rotate(response, old.ID, userID)
|
|
}
|
|
func (m *Manager) Logout(response http.ResponseWriter, request *http.Request) (State, error) {
|
|
old, _ := m.Get(request)
|
|
return m.rotate(response, old.ID, 0)
|
|
}
|
|
func (m *Manager) Get(request *http.Request) (State, bool) {
|
|
cookie, err := request.Cookie(cookieName)
|
|
if err != nil {
|
|
return State{}, false
|
|
}
|
|
id, ok := m.verifyCookie(cookie.Value)
|
|
if !ok {
|
|
return State{}, false
|
|
}
|
|
m.mutex.RLock()
|
|
state, ok := m.sessions[id]
|
|
m.mutex.RUnlock()
|
|
if !ok || !state.ExpiresAt.After(m.now()) {
|
|
if ok {
|
|
m.delete(id)
|
|
}
|
|
return State{}, false
|
|
}
|
|
return state, true
|
|
}
|
|
func (m *Manager) ValidateCSRF(request *http.Request, state State) bool {
|
|
provided := request.Header.Get("X-CSRF-Token")
|
|
return state.CSRFToken != "" && hmac.Equal([]byte(state.CSRFToken), []byte(provided))
|
|
}
|
|
|
|
func (m *Manager) rotate(response http.ResponseWriter, oldID string, userID uint64) (State, error) {
|
|
id, err := randomToken(32)
|
|
if err != nil {
|
|
return State{}, err
|
|
}
|
|
csrf, err := randomToken(32)
|
|
if err != nil {
|
|
return State{}, err
|
|
}
|
|
state := State{ID: id, UserID: userID, CSRFToken: csrf, ExpiresAt: m.now().Add(m.ttl)}
|
|
m.mutex.Lock()
|
|
for id, existing := range m.sessions {
|
|
if !existing.ExpiresAt.After(m.now()) {
|
|
delete(m.sessions, id)
|
|
}
|
|
}
|
|
if oldID != "" {
|
|
delete(m.sessions, oldID)
|
|
}
|
|
if len(m.sessions) >= maxSessions {
|
|
m.mutex.Unlock()
|
|
return State{}, ErrCapacity
|
|
}
|
|
m.sessions[id] = state
|
|
m.mutex.Unlock()
|
|
http.SetCookie(response, &http.Cookie{Name: cookieName, Value: m.signCookie(id), Path: "/", HttpOnly: true, Secure: m.secure, SameSite: http.SameSiteLaxMode, Expires: state.ExpiresAt, MaxAge: int(m.ttl.Seconds())})
|
|
return state, nil
|
|
}
|
|
func (m *Manager) delete(id string) { m.mutex.Lock(); delete(m.sessions, id); m.mutex.Unlock() }
|
|
func (m *Manager) signCookie(id string) string {
|
|
return id + "." + base64.RawURLEncoding.EncodeToString(m.signature(id))
|
|
}
|
|
func (m *Manager) verifyCookie(value string) (string, bool) {
|
|
parts := strings.Split(value, ".")
|
|
if len(parts) != 2 {
|
|
return "", false
|
|
}
|
|
signature, err := base64.RawURLEncoding.DecodeString(parts[1])
|
|
if err != nil || !hmac.Equal(signature, m.signature(parts[0])) {
|
|
return "", false
|
|
}
|
|
return parts[0], true
|
|
}
|
|
func (m *Manager) signature(value string) []byte {
|
|
mac := hmac.New(sha256.New, m.key)
|
|
mac.Write([]byte(value))
|
|
return mac.Sum(nil)
|
|
}
|
|
func randomToken(size int) (string, error) {
|
|
value := make([]byte, size)
|
|
if _, err := rand.Read(value); err != nil {
|
|
return "", err
|
|
}
|
|
return base64.RawURLEncoding.EncodeToString(value), nil
|
|
}
|