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 }