358 lines
9.4 KiB
Go
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
|
|
},
|
|
}
|