naco-api/internal/server/metricsMiddleware.go
2026-08-03 11:08:53 +00:00

58 lines
1.6 KiB
Go

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