task/migrate-x402-v2 #1
2 changed files with 38 additions and 27 deletions
migrate middlewares
commit
3716462cb8
|
|
@ -5,8 +5,7 @@ import (
|
|||
"net/http"
|
||||
"time"
|
||||
|
||||
"github.com/mark3labs/mcp-go/mcp"
|
||||
"github.com/mark3labs/mcp-go/server"
|
||||
"github.com/modelcontextprotocol/go-sdk/mcp"
|
||||
log "github.com/nacorid/logger"
|
||||
"github.com/nacorid/naco-api/internal/utils"
|
||||
)
|
||||
|
|
@ -23,15 +22,18 @@ func NewMetricsHook() *MetricsMiddleware {
|
|||
}
|
||||
}
|
||||
|
||||
func (mh *MetricsMiddleware) OnCall(next server.ToolHandlerFunc) server.ToolHandlerFunc {
|
||||
return func(ctx context.Context, request mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
start := time.Now()
|
||||
mh.requestsBucket.Increment()
|
||||
mh.totalRequests.Increment()
|
||||
result, err := next(ctx, request)
|
||||
duration := time.Since(start)
|
||||
log.DebugWithContext(ctx, "MetricsMiddleware: MCP Handler", "duration", duration)
|
||||
return result, err
|
||||
func (mh *MetricsMiddleware) OnCall(next mcp.MethodHandler) mcp.MethodHandler {
|
||||
return func(ctx context.Context, method string, request mcp.Request) (mcp.Result, error) {
|
||||
if method == "tools/call" {
|
||||
start := time.Now()
|
||||
mh.requestsBucket.Increment()
|
||||
mh.totalRequests.Increment()
|
||||
result, err := next(ctx, method, request)
|
||||
duration := time.Since(start)
|
||||
log.DebugWithContext(ctx, "MetricsMiddleware: MCP Handler", "duration", duration)
|
||||
return result, err
|
||||
}
|
||||
return next(ctx, method, request)
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -8,8 +8,8 @@ import (
|
|||
"strings"
|
||||
"sync"
|
||||
|
||||
"github.com/mark3labs/mcp-go/mcp"
|
||||
"github.com/mark3labs/mcp-go/server"
|
||||
"github.com/modelcontextprotocol/go-sdk/mcp"
|
||||
log "github.com/nacorid/logger"
|
||||
"golang.org/x/time/rate"
|
||||
)
|
||||
|
||||
|
|
@ -45,23 +45,32 @@ func (m *RateLimitMiddleware) getLimiter(sessionID string) *rate.Limiter {
|
|||
return limiter
|
||||
}
|
||||
|
||||
func (m *RateLimitMiddleware) OnCall(next server.ToolHandlerFunc) server.ToolHandlerFunc {
|
||||
return func(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
if slices.Contains(m.request, req.Params.Name) {
|
||||
var sessionID string
|
||||
session := server.ClientSessionFromContext(ctx)
|
||||
if session == nil {
|
||||
sessionID = getClientIP(req.Header)
|
||||
} else {
|
||||
sessionID = session.SessionID()
|
||||
}
|
||||
limiter := m.getLimiter(sessionID)
|
||||
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())
|
||||
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, req)
|
||||
return next(ctx, method, req)
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
Loading…
Reference in a new issue