turnstile-caddy/session.go
Lukas Schaefer d38dff43a1
All checks were successful
/ resolve-go (push) Successful in 3s
/ test (push) Successful in 6s
/ cleanup (push) Successful in 1s
Implement remember me support to survive browser restarts
Signed-off-by: Lukas Schaefer <lukas@lschaefer.xyz>
2026-09-12 23:24:56 -04:00

358 lines
9.4 KiB
Go

package turnstile
import (
"crypto/hmac"
"crypto/rand"
"crypto/sha256"
"encoding/base64"
"encoding/hex"
"encoding/json"
"fmt"
"io"
"net/http"
"net/url"
"strconv"
"strings"
"time"
"github.com/caddyserver/caddy/v2"
"github.com/caddyserver/caddy/v2/modules/caddyhttp"
"go.uber.org/zap"
)
const (
sessionPassCookieVersion = "4"
forgejoSettingsPath = "/user/settings"
forgejoUserPath = "/api/v1/user"
)
func normalizeUpstream(raw string) (string, error) {
if raw == "" {
return "", fmt.Errorf("turnstile: upstream is required")
}
if !strings.Contains(raw, "://") {
raw = "http://" + raw
}
u, err := url.Parse(raw)
if err != nil {
return "", fmt.Errorf("turnstile: invalid upstream: %w", err)
}
if u.Host == "" {
return "", fmt.Errorf("turnstile: invalid upstream: missing host")
}
if u.Path != "" && u.Path != "/" {
return "", fmt.Errorf("turnstile: upstream must not include a path")
}
return u.Scheme + "://" + u.Host, nil
}
type sessionPassSigner struct {
secret []byte
name string
ttl time.Duration
}
func newSessionPassSigner(secret, name string, ttl time.Duration) (*sessionPassSigner, error) {
if secret == "" {
return nil, fmt.Errorf("cookie_secret is required")
}
if name == "" {
name = "__turnstile_session_pass"
}
if ttl < time.Minute {
ttl = time.Hour
}
return &sessionPassSigner{
secret: []byte(secret),
name: name,
ttl: ttl,
}, nil
}
func credentialBinding(secret []byte, token string) string {
if token == "" {
return ""
}
mac := hmac.New(sha256.New, secret)
mac.Write([]byte(token))
return hex.EncodeToString(mac.Sum(nil))
}
func bindingMatches(secret []byte, token, binding string) bool {
if token == "" || binding == "" {
return false
}
expected := credentialBinding(secret, token)
return hmac.Equal([]byte(expected), []byte(binding))
}
func (s *sessionPassSigner) set(w http.ResponseWriter, r *http.Request, sessionToken, rememberToken, username string) error {
if sessionToken == "" && rememberToken == "" {
return fmt.Errorf("session pass requires a session or remember-me token")
}
if username == "" {
return fmt.Errorf("session pass requires a username")
}
nonce := make([]byte, 16)
if _, err := rand.Read(nonce); err != nil {
return err
}
exp := time.Now().Add(s.ttl).Unix()
sessionBinding := credentialBinding(s.secret, sessionToken)
rememberBinding := credentialBinding(s.secret, rememberToken)
payload := fmt.Sprintf("%s|%d|%s|%s|%s|%s", sessionPassCookieVersion, exp, hex.EncodeToString(nonce), sessionBinding, rememberBinding, username)
mac := hmac.New(sha256.New, s.secret)
mac.Write([]byte(payload))
sig := hex.EncodeToString(mac.Sum(nil))
value := base64.RawURLEncoding.EncodeToString([]byte(payload + "|" + sig))
http.SetCookie(w, &http.Cookie{
Name: s.name,
Value: value,
Path: "/",
HttpOnly: true,
Secure: r.TLS != nil || strings.EqualFold(r.Header.Get("X-Forwarded-Proto"), "https"),
SameSite: http.SameSiteLaxMode,
Expires: time.Unix(exp, 0),
MaxAge: int(s.ttl.Seconds()),
})
return nil
}
func (s *sessionPassSigner) valid(r *http.Request, sessionToken, rememberToken string) (string, bool) {
if sessionToken == "" && rememberToken == "" {
return "", false
}
cookie, err := r.Cookie(s.name)
if err != nil || cookie.Value == "" {
return "", false
}
raw, err := base64.RawURLEncoding.DecodeString(cookie.Value)
if err != nil {
return "", false
}
parts := strings.Split(string(raw), "|")
if len(parts) != 7 {
return "", false
}
ver, expStr, nonce, sessionBinding, rememberBinding, username, sig := parts[0], parts[1], parts[2], parts[3], parts[4], parts[5], parts[6]
if ver != sessionPassCookieVersion || nonce == "" || sig == "" || username == "" {
return "", false
}
if sessionBinding == "" && rememberBinding == "" {
return "", false
}
exp, err := strconv.ParseInt(expStr, 10, 64)
if err != nil || time.Now().Unix() > exp {
return "", false
}
matched := false
if sessionToken != "" {
// Session cookie present: only it may validate. A changed SID must
// not fall through to remember-me; caller re-verifies the new session.
matched = bindingMatches(s.secret, sessionToken, sessionBinding)
} else {
matched = bindingMatches(s.secret, rememberToken, rememberBinding)
}
if !matched {
return "", false
}
payload := fmt.Sprintf("%s|%s|%s|%s|%s|%s", ver, expStr, nonce, sessionBinding, rememberBinding, username)
mac := hmac.New(sha256.New, s.secret)
mac.Write([]byte(payload))
expected := hex.EncodeToString(mac.Sum(nil))
if !hmac.Equal([]byte(expected), []byte(sig)) {
return "", false
}
return username, true
}
func (m *Middleware) cookieValue(r *http.Request, name string) string {
if name == "" {
return ""
}
cookie, err := r.Cookie(name)
if err != nil || cookie.Value == "" {
return ""
}
return cookie.Value
}
func (m *Middleware) sessionToken(r *http.Request) string {
return m.cookieValue(r, m.ForgejoCookieName)
}
func (m *Middleware) rememberToken(r *http.Request) string {
return m.cookieValue(r, m.ForgejoRememberCookieName)
}
func setSessionUserID(r *http.Request, username string) {
if username == "" {
return
}
caddyhttp.SetVar(r.Context(), "user_id", username)
if repl, ok := r.Context().Value(caddy.ReplacerCtxKey).(*caddy.Replacer); ok {
repl.Set("http.auth.user.id", username)
}
}
func forgejoAuthorizationHeaderToken(r *http.Request) (string, bool) {
auth := strings.TrimSpace(r.Header.Get("Authorization"))
if auth == "" {
return "", false
}
lower := strings.ToLower(auth)
if !strings.HasPrefix(lower, "token ") && !strings.HasPrefix(lower, "bearer ") {
return "", false
}
parts := strings.SplitN(auth, " ", 2)
if len(parts) != 2 || parts[1] == "" {
return "", false
}
return auth, true
}
func (m *Middleware) ensureSessionUserID(w http.ResponseWriter, r *http.Request) bool {
if m.ForgejoCookieName == "" || m.sessionPasses == nil {
return false
}
if authHeader, ok := forgejoAuthorizationHeaderToken(r); ok {
if username, ok := m.sessionPasses.valid(r, authHeader, ""); ok {
setSessionUserID(r, username)
return true
}
username, ok, err := m.verifyForgejoAuthorizationToken(r, authHeader)
if err != nil {
m.logger.Warn("forgejo authorization verify error", zap.Error(err))
return false
}
if !ok {
return false
}
if err := m.sessionPasses.set(w, r, authHeader, "", username); err != nil {
m.logger.Warn("session pass cookie error", zap.Error(err))
}
setSessionUserID(r, username)
return true
}
sessionToken := m.sessionToken(r)
rememberToken := m.rememberToken(r)
if username, ok := m.sessionPasses.valid(r, sessionToken, rememberToken); ok {
setSessionUserID(r, username)
return true
}
if sessionToken == "" {
return false
}
username, ok, err := m.verifyForgejoSession(r, sessionToken)
if err != nil {
m.logger.Warn("session verify error", zap.Error(err))
return false
}
if !ok {
return false
}
if err := m.sessionPasses.set(w, r, sessionToken, rememberToken, username); err != nil {
m.logger.Warn("session pass cookie error", zap.Error(err))
}
setSessionUserID(r, username)
return true
}
func parseForgejoAPIUser(body []byte) (string, bool) {
var user struct {
Username string `json:"username"`
}
if err := json.Unmarshal(body, &user); err != nil || user.Username == "" {
return "", false
}
return user.Username, true
}
func (m *Middleware) verifyForgejoAuthorizationToken(r *http.Request, authHeader string) (string, bool, error) {
req, err := http.NewRequestWithContext(r.Context(), http.MethodGet, m.Upstream+forgejoUserPath, nil)
if err != nil {
return "", false, err
}
if host := r.Host; host != "" {
req.Host = host
}
req.Header.Set("Authorization", authHeader)
resp, err := forgejoVerifyClient.Do(req)
if err != nil {
return "", false, err
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
return "", false, nil
}
body, err := io.ReadAll(io.LimitReader(resp.Body, 1<<20))
if err != nil {
return "", false, err
}
username, ok := parseForgejoAPIUser(body)
if !ok {
return "", false, nil
}
return username, true, nil
}
func parseForgejoSettingsUser(body []byte) (string, bool) {
html := string(body)
if !strings.Contains(html, `class="page-content user settings profile"`) {
return "", false
}
if !strings.Contains(html, `action="/user/settings"`) {
return "", false
}
const prefix = `name="name" value="`
i := strings.Index(html, prefix)
if i < 0 {
return "", false
}
rest := html[i+len(prefix):]
j := strings.IndexByte(rest, '"')
if j <= 0 {
return "", false
}
return rest[:j], true
}
func (m *Middleware) verifyForgejoSession(r *http.Request, sessionToken string) (string, bool, error) {
req, err := http.NewRequestWithContext(r.Context(), http.MethodGet, m.Upstream+forgejoSettingsPath, nil)
if err != nil {
return "", false, err
}
if host := r.Host; host != "" {
req.Host = host
}
req.AddCookie(&http.Cookie{Name: m.ForgejoCookieName, Value: sessionToken})
resp, err := forgejoVerifyClient.Do(req)
if err != nil {
return "", false, err
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
return "", false, nil
}
body, err := io.ReadAll(io.LimitReader(resp.Body, 1<<20))
if err != nil {
return "", false, err
}
username, ok := parseForgejoSettingsUser(body)
if !ok {
return "", false, nil
}
return username, true, nil
}
var forgejoVerifyClient = &http.Client{
Timeout: 10 * time.Second,
CheckRedirect: func(*http.Request, []*http.Request) error {
return http.ErrUseLastResponse
},
}