Something went wrong. Try again.
Tea journaling on ATProto (alpha)
Something went wrong. Try again.
7.5 kB · 261 lines
Go
123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262package middleware
import ( "context" "crypto/rand" "encoding/base64" "net/http" "strings" "sync" "time")
type cspNonceKeyType struct{}
var cspNonceKey = cspNonceKeyType{}
func generateNonce() (string, error) { b := make([]byte, 16) if _, err := rand.Read(b); err != nil { return "", err } return base64.StdEncoding.EncodeToString(b), nil}
func CSPNonceFromContext(ctx context.Context) string { if v, ok := ctx.Value(cspNonceKey).(string); ok { return v } return ""}
// SecurityHeadersMiddleware adds security headers to all responsesfunc SecurityHeadersMiddleware(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { nonce, err := generateNonce() if err != nil { http.Error(w, "Internal Server Error", http.StatusInternalServerError) return }
r = r.WithContext(context.WithValue(r.Context(), cspNonceKey, nonce))
// Prevent clickjacking w.Header().Set("X-Frame-Options", "DENY")
// Prevent MIME type sniffing w.Header().Set("X-Content-Type-Options", "nosniff")
// XSS protection (legacy but still useful for older browsers) w.Header().Set("X-XSS-Protection", "1; mode=block")
// Control referrer information w.Header().Set("Referrer-Policy", "strict-origin-when-cross-origin")
// Permissions policy - disable unnecessary features w.Header().Set("Permissions-Policy", "geolocation=(), microphone=(), camera=()")
// Content Security Policy // Allows: self for scripts/styles, inline styles (for Tailwind), inline HTMX/Alpine // Note: unsafe-eval required for Alpine.js standard build (CSP build has CDN MIME type issues) // Note: form-action allows https: for OAuth redirects to external authorization servers // TODO: set nonce/hash on unsafe tags -- needs to be set in elements as well csp := strings.Join([]string{ "default-src 'self'", "script-src 'self' 'unsafe-eval' 'nonce-" + nonce + "'", "style-src 'self' 'unsafe-inline'", // unsafe-inline needed for Tailwind "img-src 'self' https: data:", // Allow external images (avatars) and data URIs "font-src 'self'", "connect-src 'self' https:", // Allow connections to external APIs (OAuth, PDS) "frame-ancestors 'none'", "base-uri 'self'", "form-action 'self' https:", // Allow form submissions to external OAuth servers }, "; ") w.Header().Set("Content-Security-Policy", csp)
next.ServeHTTP(w, r) })}
// RateLimiter implements a simple per-IP rate limiter using token bucket algorithmtype RateLimiter struct { mu sync.Mutex visitors map[string]*visitor rate int // requests per window window time.Duration // time window cleanup time.Duration // cleanup interval for old entries}
type visitor struct { tokens int lastReset time.Time}
// NewRateLimiter creates a new rate limiter// rate: number of requests allowed per window// window: time window for rate limitingfunc NewRateLimiter(rate int, window time.Duration) *RateLimiter { rl := &RateLimiter{ visitors: make(map[string]*visitor), rate: rate, window: window, cleanup: window * 2, }
// Start cleanup goroutine go rl.cleanupLoop()
return rl}
func (rl *RateLimiter) cleanupLoop() { ticker := time.NewTicker(rl.cleanup) defer ticker.Stop()
for range ticker.C { rl.mu.Lock() now := time.Now() for ip, v := range rl.visitors { if now.Sub(v.lastReset) > rl.cleanup { delete(rl.visitors, ip) } } rl.mu.Unlock() }}
// Allow checks if a request from the given IP is allowedfunc (rl *RateLimiter) Allow(ip string) bool { rl.mu.Lock() defer rl.mu.Unlock()
now := time.Now() v, exists := rl.visitors[ip]
if !exists { rl.visitors[ip] = &visitor{ tokens: rl.rate - 1, // Use one token lastReset: now, } return true }
// Reset tokens if window has passed if now.Sub(v.lastReset) >= rl.window { v.tokens = rl.rate - 1 v.lastReset = now return true }
// Check if tokens available if v.tokens > 0 { v.tokens-- return true }
return false}
// RateLimitConfig holds configuration for rate limiting different endpoint typestype RateLimitConfig struct { // AuthLimiter for login/auth endpoints (stricter) AuthLimiter *RateLimiter // APILimiter for general API endpoints APILimiter *RateLimiter // GlobalLimiter for all other requests GlobalLimiter *RateLimiter}
// NewDefaultRateLimitConfig creates rate limiters with sensible defaultsfunc NewDefaultRateLimitConfig() *RateLimitConfig { return &RateLimitConfig{ AuthLimiter: NewRateLimiter(5, time.Minute), // 5 auth attempts per minute APILimiter: NewRateLimiter(60, time.Minute), // 60 API calls per minute GlobalLimiter: NewRateLimiter(120, time.Minute), // 120 requests per minute }}
// RateLimitMiddleware creates a rate limiting middlewarefunc RateLimitMiddleware(config *RateLimitConfig) func(http.Handler) http.Handler { return func(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { path := r.URL.Path
// Static assets are served from in-memory caches behind ETag and // don't represent the threat model rate-limiting protects against // (auth brute force, API scraping). A page load fans out into // 15+ revalidation requests in dev mode, which would otherwise // burn through the global bucket in a handful of refreshes. if strings.HasPrefix(path, "/static/") || path == "/favicon.ico" { next.ServeHTTP(w, r) return }
ip := GetClientIP(r)
var limiter *RateLimiter
// Select appropriate limiter based on path switch { case strings.HasPrefix(path, "/auth/") || path == "/login" || path == "/oauth/callback": limiter = config.AuthLimiter case strings.HasPrefix(path, "/api/"): limiter = config.APILimiter default: limiter = config.GlobalLimiter }
if !limiter.Allow(ip) { w.Header().Set("Retry-After", "60") http.Error(w, "Too many requests", http.StatusTooManyRequests) return }
next.ServeHTTP(w, r) }) }}
// RequireHTMXMiddleware ensures that certain API routes are only accessible via HTMX requests.// This prevents direct browser access to internal API endpoints that return fragments or JSON.// Routes that need to be publicly accessible (like /api/resolve-handle) should not use this middleware.func RequireHTMXMiddleware(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { // Check for HTMX request header if r.Header.Get("HX-Request") != "true" { http.NotFound(w, r) return } next.ServeHTTP(w, r) })}
// MaxBodySize limits the size of request bodiesconst ( MaxJSONBodySize = 1 << 20 // 1 MB for JSON requests MaxFormBodySize = 1 << 20 // 1 MB for form submissions)
// LimitBodyMiddleware limits request body size to prevent DoSfunc LimitBodyMiddleware(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { if r.Body != nil { contentType := r.Header.Get("Content-Type") var maxSize int64
switch { case strings.HasPrefix(contentType, "application/json"): maxSize = MaxJSONBodySize case strings.HasPrefix(contentType, "application/x-www-form-urlencoded"), strings.HasPrefix(contentType, "multipart/form-data"): maxSize = MaxFormBodySize default: maxSize = MaxJSONBodySize // Default limit }
r.Body = http.MaxBytesReader(w, r.Body, maxSize) }
next.ServeHTTP(w, r) })}