Files
chorus/portal/session/manager.go
T

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
}