105 lines
2.6 KiB
Go
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
|
|
}
|