naco-api/internal/server/rateLimitMiddleware.go
2026-07-29 22:09:01 +00:00

105 lines
2.6 KiB
Go

package server
import (
"context"
"fmt"
"net/http"
"slices"
"strings"
"sync"
"github.com/modelcontextprotocol/go-sdk/mcp"
log "github.com/nacorid/logger"
"golang.org/x/time/rate"
)
type RateLimitMiddleware struct {
limiters map[string]*rate.Limiter
mutex sync.RWMutex
rate rate.Limit
burst int
request []string
}
func NewRateLimitMiddleware(requestsPerSecond float64, burst int, request []string) *RateLimitMiddleware {
return &RateLimitMiddleware{
limiters: make(map[string]*rate.Limiter),
rate: rate.Limit(requestsPerSecond),
burst: burst,
request: request,
}
}
func (m *RateLimitMiddleware) getLimiter(sessionID string) *rate.Limiter {
m.mutex.RLock()
limiter, exists := m.limiters[sessionID]
m.mutex.RUnlock()
if !exists {
m.mutex.Lock()
limiter = rate.NewLimiter(m.rate, m.burst)
m.limiters[sessionID] = limiter
m.mutex.Unlock()
}
return limiter
}
func (m *RateLimitMiddleware) OnCall(next mcp.MethodHandler) mcp.MethodHandler {
return func(ctx context.Context, method string, req mcp.Request) (mcp.Result, error) {
if method == "tools/call" {
callReq, ok := req.(*mcp.CallToolRequest)
if ok && slices.Contains(m.request, callReq.Params.Name) {
var sessionID string
session := req.GetSession()
if session == nil || session.ID() == "" {
if extra := req.GetExtra(); extra != nil {
sessionID = getClientIP(req.GetExtra().Header)
} else {
log.WarnWithContext(ctx, "No session or extra information found in request; using 'unknown' as session ID")
sessionID = "unknown"
}
} else {
sessionID = session.ID()
}
limiter := m.getLimiter(sessionID)
if !limiter.Allow() {
return nil, fmt.Errorf("rate limit exceeded for session %s\nRatelimit: %v requests per second", sessionID, limiter.Limit())
}
}
return next(ctx, method, req)
}
return next(ctx, method, req)
}
}
func (m *RateLimitMiddleware) OnCallHTTP(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
toolName := strings.TrimPrefix(r.URL.Path, "/")
if slices.Contains(m.request, toolName) {
sessionID := getClientIP(r.Header)
limiter := m.getLimiter(sessionID)
if !limiter.Allow() {
http.Error(w, "rate limit exceeded", http.StatusTooManyRequests)
return
}
}
next.ServeHTTP(w, r)
})
}
func getClientIP(h http.Header) string {
addr := "@"
forwarded := h.Get("X-Forwarded-For")
if forwarded != "" {
ip := strings.Split(forwarded, ",")[0]
addr = strings.TrimSpace(ip)
} else {
realIP := h.Get("X-Real-Ip")
if realIP != "" {
addr = realIP
}
}
return addr
}