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 }, }