Files
pure-note/internal/middleware/ratelimit.go
T
wangairnan 3105b0415c fix(middleware): Origin 同源校验增加 scheme 比对
http/https 不再视为同源;可信反代后采信 X-Forwarded-Proto(round2 P2-11),补测试。
2026-09-08 17:32:57 +08:00

165 lines
4.5 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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
}