165 lines
4.5 KiB
Go
165 lines
4.5 KiB
Go
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
|
||
}
|