task/migrate-x402-v2 #1

Merged
nacorid merged 6 commits from task/migrate-x402-v2 into main 2026-07-30 08:13:05 +00:00
2 changed files with 38 additions and 27 deletions
Showing only changes of commit 3716462cb8 - Show all commits

migrate middlewares

Nacorid 2026-07-18 23:28:02 +00:00
Signed by: nacorid
SSH key fingerprint: SHA256:zAJkAgjXXOAJqP6R2fp8eKCNlnKgpf33G/Baa1xtNGA

View file

@ -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)
}
}

View file

@ -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)
}
}