Sanitize MCP errors to avoid leaking internal details
Signed-off-by: Sacha Al Himdani <sacha@getprobo.com>
This commit is contained in:
@@ -20,49 +20,73 @@ import (
|
||||
"fmt"
|
||||
"runtime/debug"
|
||||
|
||||
"github.com/modelcontextprotocol/go-sdk/mcp"
|
||||
"go.gearno.de/kit/log"
|
||||
mcpgenmcp "go.probo.inc/mcpgen/mcp"
|
||||
"go.probo.inc/probo/pkg/coredata"
|
||||
"go.probo.inc/probo/pkg/iam"
|
||||
"go.probo.inc/probo/pkg/validator"
|
||||
)
|
||||
|
||||
func RecoveryMiddleware(logger *log.Logger) func(mcp.MethodHandler) mcp.MethodHandler {
|
||||
return func(next mcp.MethodHandler) mcp.MethodHandler {
|
||||
return func(ctx context.Context, method string, req mcp.Request) (result mcp.Result, err error) {
|
||||
defer func() {
|
||||
if r := recover(); r != nil {
|
||||
err = convertPanicToError(ctx, logger, r)
|
||||
result = nil
|
||||
}
|
||||
}()
|
||||
|
||||
result, err = next(ctx, method, req)
|
||||
return
|
||||
// NewRecoverFunc returns a RecoverFunc for the generated MCP server that
|
||||
// classifies panics into safe client-facing errors and logs unknown errors.
|
||||
func NewRecoverFunc(logger *log.Logger) mcpgenmcp.RecoverFunc {
|
||||
return func(ctx context.Context, r any) error {
|
||||
if r == nil {
|
||||
logger.ErrorCtx(ctx, "nil panic in MCP tool handler")
|
||||
return fmt.Errorf("internal server error")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func convertPanicToError(ctx context.Context, logger *log.Logger, panicValue any) error {
|
||||
if panicValue == nil {
|
||||
logger.ErrorCtx(ctx, "nil panic in MCP method handler")
|
||||
if err, ok := r.(error); ok {
|
||||
return sanitizeError(ctx, logger, err)
|
||||
}
|
||||
|
||||
logger.ErrorCtx(
|
||||
ctx,
|
||||
"unexpected panic in MCP tool handler",
|
||||
log.Any("panic", r),
|
||||
log.String("stack", string(debug.Stack())),
|
||||
)
|
||||
|
||||
return fmt.Errorf("internal server error")
|
||||
}
|
||||
}
|
||||
|
||||
// sanitizeError classifies known error types and returns a clear message for
|
||||
// those. Unknown errors are logged and replaced with a generic internal error
|
||||
// to avoid leaking implementation details to the client.
|
||||
func sanitizeError(ctx context.Context, logger *log.Logger, err error) error {
|
||||
var permissionDeniedErr *iam.ErrInsufficientPermissions
|
||||
if errTyped, ok := panicValue.(error); ok && errors.As(errTyped, &permissionDeniedErr) {
|
||||
return fmt.Errorf("permission denied: %s", permissionDeniedErr.Error())
|
||||
if errors.As(err, &permissionDeniedErr) {
|
||||
return fmt.Errorf("permission denied")
|
||||
}
|
||||
|
||||
if err, ok := panicValue.(error); ok {
|
||||
return err
|
||||
var assumptionRequiredErr *iam.ErrAssumptionRequired
|
||||
if errors.As(err, &assumptionRequiredErr) {
|
||||
return fmt.Errorf("assumption required")
|
||||
}
|
||||
|
||||
// Log unexpected panics with stack trace
|
||||
logger.ErrorCtx(
|
||||
ctx,
|
||||
"unexpected panic in MCP method handler",
|
||||
log.Any("panic", panicValue),
|
||||
log.String("stack", string(debug.Stack())),
|
||||
)
|
||||
if errors.Is(err, coredata.ErrResourceNotFound) {
|
||||
return fmt.Errorf("resource not found")
|
||||
}
|
||||
|
||||
if errors.Is(err, coredata.ErrResourceAlreadyExists) {
|
||||
return fmt.Errorf("resource already exists")
|
||||
}
|
||||
|
||||
if errors.Is(err, coredata.ErrResourceInUse) {
|
||||
return fmt.Errorf("resource is in use")
|
||||
}
|
||||
|
||||
var validationErrors validator.ValidationErrors
|
||||
if errors.As(err, &validationErrors) {
|
||||
return validationErrors
|
||||
}
|
||||
|
||||
var validationError *validator.ValidationError
|
||||
if errors.As(err, &validationError) {
|
||||
return validationError
|
||||
}
|
||||
|
||||
logger.ErrorCtx(ctx, "internal error in MCP tool handler", log.Error(err))
|
||||
return fmt.Errorf("internal server error")
|
||||
}
|
||||
|
||||
@@ -50,11 +50,9 @@ func NewMux(logger *log.Logger, proboSvc *probo.Service, iamSvc *iam.Service, ac
|
||||
logger: logger,
|
||||
}
|
||||
|
||||
mcpServer := server.New(resolver)
|
||||
mcpServer := server.New(resolver, mcpgenmcp.WithRecoverFunc(mcputils.NewRecoverFunc(logger)))
|
||||
|
||||
// Add panic recovery middleware to handle panics in goroutines spawned by MCP SDK
|
||||
mcpServer.AddReceivingMiddleware(mcputils.LoggingMiddleware(logger))
|
||||
mcpServer.AddReceivingMiddleware(mcputils.RecoveryMiddleware(logger))
|
||||
|
||||
getServer := func(r *http.Request) *mcp.Server { return mcpServer }
|
||||
eventStore := mcp.NewMemoryEventStore(nil)
|
||||
|
||||
Reference in New Issue
Block a user