// 织忆 MemoryWeave — 认证 + Per-Agent 令牌桶限流 package middleware import ( "fmt" "log" "net/http" "os" "strings" "sync" "time" ) // ─── Per-Agent 令牌桶 ──────────────────────────────────────── type tokenBucket struct { tokens float64 lastRefill time.Time rate float64 // tokens/sec burst float64 mu sync.Mutex } func newTokenBucket(rate, burst float64) *tokenBucket { return &tokenBucket{ tokens: burst, lastRefill: time.Now(), rate: rate, burst: burst, } } func (tb *tokenBucket) allow() bool { tb.mu.Lock() defer tb.mu.Unlock() now := time.Now() elapsed := now.Sub(tb.lastRefill).Seconds() tb.tokens += elapsed * tb.rate if tb.tokens > tb.burst { tb.tokens = tb.burst } tb.lastRefill = now if tb.tokens >= 1.0 { tb.tokens -= 1.0 return true } return false } // ─── Per-Agent 配置 ────────────────────────────────────────── type AgentRateConfig struct { RecallQPS float64 CommitQPS float64 Burst float64 } var agentRates = map[string]AgentRateConfig{ "hermes": {RecallQPS: 10, CommitQPS: 2, Burst: 20}, "hermes-a06":{RecallQPS: 10, CommitQPS: 2, Burst: 20}, "openclaw": {RecallQPS: 10, CommitQPS: 2, Burst: 20}, "cron-job": {RecallQPS: 5, CommitQPS: 1, Burst: 10}, } var defaultRate = AgentRateConfig{RecallQPS: 5, CommitQPS: 1, Burst: 10} // ─── RateLimiter ───────────────────────────────────────────── type RateLimiter struct { mu sync.Mutex buckets map[string]*tokenBucket // agentID → bucket maxPerMinute int // 全局兜底 } func NewRateLimiter(maxPerMinute int) *RateLimiter { return &RateLimiter{ buckets: make(map[string]*tokenBucket), maxPerMinute: maxPerMinute, } } // Allow 检查 agent 是否允许请求 func (rl *RateLimiter) Allow(agentID string, endpointType string) bool { rl.mu.Lock() defer rl.mu.Unlock() // 获取 agent 配置 cfg, ok := agentRates[agentID] if !ok { cfg = defaultRate } // 根据端点类型选择 QPS var rate float64 switch endpointType { case "recall": rate = cfg.RecallQPS case "commit": rate = cfg.CommitQPS default: rate = cfg.RecallQPS } bucketKey := agentID + ":" + endpointType bucket, exists := rl.buckets[bucketKey] if !exists { bucket = newTokenBucket(rate, cfg.Burst) rl.buckets[bucketKey] = bucket } return bucket.allow() } // ─── 全局实例 ──────────────────────────────────────────────── var GlobalLimiter *RateLimiter func init() { GlobalLimiter = NewRateLimiter(120) } // ─── HTTP 中间件 ───────────────────────────────────────────── // Auth 认证中间件 (X-API-Key) func Auth(next http.Handler) http.Handler { apiKey := os.Getenv("API_KEY") if apiKey == "" { apiKey = "zhiyi-dev-key-2026" } return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { // /health, /metrics, /static/* 不需要认证 if r.URL.Path == "/health" || r.URL.Path == "/metrics" || strings.HasPrefix(r.URL.Path, "/static/") { next.ServeHTTP(w, r) return } key := r.Header.Get("X-API-Key") if key != apiKey { w.Header().Set("Content-Type", "application/json") w.WriteHeader(401) w.Write([]byte(`{"error":"unauthorized: invalid or missing X-API-Key"}`)) return } // 提取 agent ID(从 X-Agent-ID 头或 URL 路径) agentID := r.Header.Get("X-Agent-ID") if agentID == "" { agentID = "unknown" } // 确定端点类型 endpointType := "recall" if r.URL.Path == "/api/v1/commit" || r.URL.Path == "/api/v1/batch-commit" { endpointType = "commit" } // Per-agent 令牌桶检查 if !GlobalLimiter.Allow(agentID, endpointType) { w.Header().Set("Content-Type", "application/json") w.Header().Set("Retry-After", "1") w.Header().Set("X-RateLimit-Reset", fmt.Sprintf("%d", time.Now().Unix()+1)) w.WriteHeader(429) w.Write([]byte(fmt.Sprintf( `{"error":"rate_limit_exceeded","agent":"%s","retry_after":1}`, agentID, ))) log.Printf("[ratelimit] %s exceeded for agent %s", endpointType, agentID) return } next.ServeHTTP(w, r) }) } // ─── 兼容旧接口: RateLimit (全局兜底) ───────────────────────── // RateLimit 返回带 Retry-After 头的限流中间件(全局兜底) func RateLimit(maxPerMinute int) func(http.Handler) http.Handler { var lastReset time.Time var count int var mu sync.Mutex return func(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { mu.Lock() defer mu.Unlock() now := time.Now() if now.Sub(lastReset) > time.Minute { count = 0 lastReset = now } count++ resetTime := lastReset.Add(time.Minute).Unix() if count > maxPerMinute { w.Header().Set("Retry-After", fmt.Sprintf("%d", int(time.Until(lastReset.Add(time.Minute)).Seconds()+1))) w.Header().Set("X-RateLimit-Reset", fmt.Sprintf("%d", resetTime)) w.Header().Set("Content-Type", "application/json") w.WriteHeader(429) w.Write([]byte(`{"error":"rate_limit_exceeded","retry_after":1}`)) return } next.ServeHTTP(w, r) }) } }