Files
pure-note/internal/middleware/ratelimit.go
T
wangairnan 247e88c4fb 后端基础层:模块定义、配置解析、SQLite 存储层、Argon2id 认证与 Markdown 渲染
- internal/config:serve/init/backup/gc 子命令参数解析(--dev 强制 loopback 守卫)
- internal/store:user_version 版本化迁移(仅追加式 + 越界拒启 + --allow-newer)、
  笔记/图片/引用/会话/设置 DAO、并集可见性查询、VACUUM INTO 在线备份、
  回收站 30 天 + 孤儿图 7 天宽限 gc
- internal/auth:Argon2id PHC 串(m=19456,t=2,p=1)、256bit token 与 SHA-256 摘要
- internal/markdown:goldmark(默认转义)+ bluemonday 双保险,摘要纯文本提取
- internal/middleware:安全头(CSP/HSTS/nosniff 等)、错误日志、
  有界令牌桶限流(per-IP 桶上限 + TTL 逐出)、Origin/Referer 同源校验
2026-09-08 08:14:12 +08:00

158 lines
4.1 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 与请求 Host 一致。
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 == "" {
return false
}
return 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
}