package server import ( "context" "fmt" "net/http" "time" "github.com/modelcontextprotocol/go-sdk/mcp" log "github.com/nacorid/logger" "github.com/nacorid/naco-api/internal/utils" ) type MetricsMiddleware struct { requestsBucket utils.RollingStats totalRequests utils.Counter } func NewMetricsHook() *MetricsMiddleware { return &MetricsMiddleware{ requestsBucket: *utils.NewRollingStats(60, 1*time.Second), totalRequests: *utils.NewCounter(), } } 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 := fmt.Sprintf("%.2fms", float64(time.Since(start))/float64(time.Millisecond)) log.DebugWithContext(ctx, "MetricsMiddleware: MCP Handler", "duration", duration) return result, err } return next(ctx, method, request) } } func (mh *MetricsMiddleware) OnCallHTTP(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { start := time.Now() mh.requestsBucket.Increment() mh.totalRequests.Increment() next.ServeHTTP(w, r) duration := fmt.Sprintf("%.2fms", float64(time.Since(start))/float64(time.Millisecond)) log.DebugWithContext(r.Context(), "MetricsMiddleware: HTTP Handler", "duration", duration) }) } func (mh *MetricsMiddleware) GetRPS() float64 { return mh.requestsBucket.GetRPS() } func (mh *MetricsMiddleware) GetTotalRequests() int64 { return mh.totalRequests.Get() }