package middleware import ( "net/http" "net/url" "strconv" "strings" "sync" "time" "golang.org/x/time/rate" ) // OriginCheck 对所有非 GET/HEAD/OPTIONS 请求执行 Origin/Referer→Host 校验 // (§9.1-T2,含 /api/auth/*):Origin 优先,缺省回退 Referer;两者皆缺亦拒绝。 func OriginCheck(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { if r.Method != http.MethodGet && r.Method != http.MethodHead && r.Method != http.MethodOptions { if !sameOrigin(r) { writeErr(w, http.StatusForbidden, "origin_forbidden", "跨站请求被拒绝(Origin/Referer 校验失败)") return } } next.ServeHTTP(w, r) }) } // sameOrigin 校验 Origin(或 Referer)与请求同源:host 一致且 scheme 与请求 // 实际 scheme 一致(TLS 直连为 https;可信反代后取 X-Forwarded-Proto,评审 round2 P2-11)。 func sameOrigin(r *http.Request) bool { raw := r.Header.Get("Origin") if raw == "" { raw = r.Header.Get("Referer") } if raw == "" { return false } u, err := url.Parse(raw) if err != nil || u.Host == "" || u.Scheme == "" { return false } scheme := "http" if r.TLS != nil { scheme = "https" } else if trueBehindProxy(r) && strings.EqualFold(r.Header.Get("X-Forwarded-Proto"), "https") { scheme = "https" } return strings.EqualFold(u.Scheme, scheme) && strings.EqualFold(u.Host, r.Host) } func writeErr(w http.ResponseWriter, status int, code, msg string) { w.Header().Set("Content-Type", "application/json; charset=utf-8") w.WriteHeader(status) // 与 httpapi 错误包络同构;此处避免包循环引用。 _, _ = w.Write([]byte(`{"error":{"code":"` + code + `","message":` + strconv.Quote(msg) + `}}`)) } // Limiter 有界令牌桶集合:per-key 限流,桶上限 maxEntries, // 超限时逐出最久未活跃且已耗尽的桶(防伪造 IP 撑爆内存,§7.2/§11)。 type Limiter struct { mu sync.Mutex buckets map[string]*bucket rate rate.Limit burst int maxEntries int ttl time.Duration lastEvictAt time.Time } type bucket struct { lim *rate.Limiter lastSeen time.Time } // NewLimiter 创建限流器:每秒 rate 个、burst 容量、maxEntries 桶上限、ttl 空闲逐出。 func NewLimiter(perSec float64, burst, maxEntries int, ttl time.Duration) *Limiter { return &Limiter{ buckets: make(map[string]*bucket), rate: rate.Limit(perSec), burst: burst, maxEntries: maxEntries, ttl: ttl, } } // get 取(或建)key 的桶。 func (l *Limiter) get(key string) *rate.Limiter { l.mu.Lock() defer l.mu.Unlock() now := time.Now() // 惰性清理:桶数达上限或距上次清理超过 ttl 时执行 if len(l.buckets) >= l.maxEntries || now.Sub(l.lastEvictAt) > l.ttl { l.evictLocked(now) } b, ok := l.buckets[key] if !ok { b = &bucket{lim: rate.NewLimiter(l.rate, l.burst)} l.buckets[key] = b } b.lastSeen = now return b.lim } // evictLocked 逐出过期桶;仍超限时逐出最久未活跃的桶。 func (l *Limiter) evictLocked(now time.Time) { l.lastEvictAt = now var oldestKey string var oldest time.Time for k, b := range l.buckets { if now.Sub(b.lastSeen) > l.ttl { delete(l.buckets, k) continue } if oldestKey == "" || b.lastSeen.Before(oldest) { oldestKey, oldest = k, b.lastSeen } } for len(l.buckets) >= l.maxEntries && oldestKey != "" { delete(l.buckets, oldestKey) oldestKey = "" for k, b := range l.buckets { if oldestKey == "" || b.lastSeen.Before(oldest) { oldestKey, oldest = k, b.lastSeen } } } } // Allow 消费 1 个 token;无可用 token 返回 false(不阻塞)。 func (l *Limiter) Allow(key string) bool { return l.get(key).Allow() } // Available 桶内剩余 token 是否 ≥1(登录防爆破的快速预检,不消费)。 func (l *Limiter) Available(key string) bool { return l.get(key).Tokens() >= 1 } // Remaining 桶内剩余 token 数(日志观测用,向下取整,不消费)。 func (l *Limiter) Remaining(key string) int { return int(l.get(key).Tokens()) } // RetryAfter 距下一个可用 token 的秒数(向上取整,至少 1)。 // 用 Reserve+Cancel 探测延迟,不真正消费 token。 func (l *Limiter) RetryAfter(key string) int { l.mu.Lock() b := l.buckets[key] l.mu.Unlock() if b == nil { return 1 } res := b.lim.Reserve() d := res.Delay() res.CancelAt(time.Now()) if d <= 0 { return 1 } s := int(d.Seconds()) if s < 1 { s = 1 } return s }