Skip to content

自定义 Mediator 实现 ​

🔧 扩展 MediatR 核心功能,实现定制化需求


📖 概述 ​

虽然 MediatR 的默认实现已经非常强大,但在某些场景下,你可能需要自定义行为,例如:

  • 自定义通知发布策略(容错、重试、顺序控制)
  • 扩展请求路由机制
  • 添加诊断和监控
  • 实现特殊的异常处理策略

🎯 继承 Mediator 类 ​

基础自定义 ​

csharp
public class CustomMediator : Mediator
{
    private readonly ILogger<CustomMediator> _logger;
    private readonly DiagnosticSource _diagnosticSource;

    public CustomMediator(
        IServiceProvider serviceFactory, 
        ILogger<CustomMediator> logger,
        DiagnosticSource diagnosticSource) 
        : base(serviceFactory)
    {
        _logger = logger;
        _diagnosticSource = diagnosticSource;
    }

    // 自定义请求发送
    protected override async Task<TResponse> SendImpl<TRequest, TResponse>(
        IRequest<TResponse> request, 
        CancellationToken cancellationToken)
    {
        _logger.LogInformation("发送请求: {RequestType}", typeof(TRequest).Name);

        // 可以添加额外的逻辑
        if (_diagnosticSource.IsEnabled("MediatR.Request.Start"))
        {
            _diagnosticSource.Write("MediatR.Request.Start", new { Request = request });
        }

        return await base.SendImpl<TRequest, TResponse>(request, cancellationToken);
    }

    // 自定义通知发布
    protected override async Task PublishCore(
        IEnumerable<NotificationHandlerExecutor> handlerExecutors, 
        INotification notification, 
        CancellationToken cancellationToken)
    {
        _logger.LogInformation("发布通知: {NotificationType}", notification.GetType().Name);

        // 使用自定义策略发布
        await PublishWithRetry(handlerExecutors, notification, cancellationToken);
    }

    private async Task PublishWithRetry(
        IEnumerable<NotificationHandlerExecutor> handlers, 
        INotification notification, 
        CancellationToken ct)
    {
        foreach (var handler in handlers)
        {
            var retryCount = 3;
            for (int i = 0; i < retryCount; i++)
            {
                try
                {
                    await handler.HandlerCallback(notification, ct);
                    break; // 成功则跳出
                }
                catch (Exception ex) when (i < retryCount - 1)
                {
                    _logger.LogWarning(ex, 
                        "处理器执行失败,第 {RetryCount} 次重试", i + 1);
                    await Task.Delay(TimeSpan.FromSeconds(Math.Pow(2, i)), ct); // 指数退避
                }
            }
        }
    }
}

注册自定义 Mediator ​

csharp
builder.Services.AddScoped<IMediator, CustomMediator>();
builder.Services.AddScoped<ISender>(sp => sp.GetRequiredService<IMediator>());
builder.Services.AddScoped<IPublisher>(sp => sp.GetRequiredService<IMediator>());

🔄 自定义发布策略 ​

策略 1:容错发布(推荐) ​

csharp
public class ResilientPublisher : Mediator
{
    private readonly ILogger<ResilientPublisher> _logger;

    public ResilientPublisher(IServiceProvider serviceFactory, ILogger<ResilientPublisher> logger) 
        : base(serviceFactory)
    {
        _logger = logger;
    }

    protected override async Task PublishCore(
        IEnumerable<NotificationHandlerExecutor> handlerExecutors, 
        INotification notification, 
        CancellationToken cancellationToken)
    {
        var tasks = handlerExecutors.Select(async executor =>
        {
            try
            {
                await executor.HandlerCallback(notification, cancellationToken);
            }
            catch (Exception ex)
            {
                _logger.LogError(ex, 
                    "通知处理器失败: {HandlerType}. 通知将继续发送给其他处理器.",
                    executor.HandlerInstance.GetType().Name);
                
                // 不抛出异常,继续执行其他处理器
            }
        });

        await Task.WhenAll(tasks);
    }
}

策略 2:超时控制 ​

csharp
public class TimeoutMediator : Mediator
{
    private readonly TimeSpan _defaultTimeout = TimeSpan.FromSeconds(30);

    public TimeoutMediator(IServiceProvider serviceFactory) : base(serviceFactory) { }

    protected override async Task<TResponse> SendImpl<TRequest, TResponse>(
        IRequest<TResponse> request, 
        CancellationToken cancellationToken)
    {
        using var cts = CancellationTokenSource.CreateLinkedTokenSource(cancellationToken);
        cts.CancelAfter(_defaultTimeout);

        try
        {
            return await base.SendImpl<TRequest, TResponse>(request, cts.Token);
        }
        catch (OperationCanceledException) when (!cancellationToken.IsCancellationRequested)
        {
            throw new TimeoutException($"请求处理超时: {typeof(TRequest).Name}");
        }
    }
}

策略 3:优先级队列 ​

csharp
public interface IPrioritizedNotification
{
    int Priority { get; } // 值越小优先级越高
}

public class PriorityMediator : Mediator
{
    public PriorityMediator(IServiceProvider serviceFactory) : base(serviceFactory) { }

    protected override async Task PublishCore(
        IEnumerable<NotificationHandlerExecutor> handlerExecutors, 
        INotification notification, 
        CancellationToken cancellationToken)
    {
        // 如果通知实现了优先级接口,按优先级排序
        if (notification is IPrioritizedNotification)
        {
            var sortedHandlers = handlerExecutors.OrderBy(h => 
                ((IPrioritizedNotification)notification).Priority);

            foreach (var handler in sortedHandlers)
            {
                await handler.HandlerCallback(notification, cancellationToken);
            }
        }
        else
        {
            // 默认并行执行
            await base.PublishCore(handlerExecutors, notification, cancellationToken);
        }
    }
}

🛠️ 扩展请求路由机制 ​

基于特性的路由 ​

csharp
[AttributeUsage(AttributeTargets.Class)]
public class RouteToAttribute : Attribute
{
    public Type HandlerType { get; }

    public RouteToAttribute(Type handlerType)
    {
        HandlerType = handlerType;
    }
}

// 使用特性指定处理器
[RouteTo(typeof(SpecialOrderHandler))]
public class SpecialOrderCommand : IRequest<OrderResult>
{
    public string Data { get; set; }
}

// 自定义 Mediator 解析特性
public class RoutedMediator : Mediator
{
    public RoutedMediator(IServiceProvider serviceFactory) : base(serviceFactory) { }

    protected override async Task<TResponse> SendImpl<TRequest, TResponse>(
        IRequest<TResponse> request, 
        CancellationToken cancellationToken)
    {
        // 检查是否有路由特性
        var routeAttribute = typeof(TRequest).GetCustomAttribute<RouteToAttribute>();
        
        if (routeAttribute != null)
        {
            // 使用指定的处理器
            var handler = ActivatorUtilities.CreateInstance(
                ServiceProvider, 
                routeAttribute.HandlerType);
            
            var handleMethod = handler.GetType()
                .GetMethod("Handle");
            
            return await (Task<TResponse>)handleMethod.Invoke(handler, new object[] { request, cancellationToken });
        }

        // 默认行为
        return await base.SendImpl<TRequest, TResponse>(request, cancellationToken);
    }
}

📊 诊断与监控 ​

集成 DiagnosticSource ​

csharp
public class DiagnosticMediator : Mediator
{
    private readonly DiagnosticListener _diagnosticListener;

    public DiagnosticMediator(
        IServiceProvider serviceFactory, 
        DiagnosticListener diagnosticListener) 
        : base(serviceFactory)
    {
        _diagnosticListener = diagnosticListener;
    }

    protected override async Task<TResponse> SendImpl<TRequest, TResponse>(
        IRequest<TResponse> request, 
        CancellationToken cancellationToken)
    {
        var requestId = Guid.NewGuid();
        var requestName = typeof(TRequest).Name;

        // 开始事件
        if (_diagnosticListener.IsEnabled("MediatR.Request.Start"))
        {
            _diagnosticListener.Write("MediatR.Request.Start", new 
            { 
                RequestId = requestId,
                RequestName = requestName,
                Timestamp = DateTime.UtcNow
            });
        }

        var stopwatch = Stopwatch.StartNew();

        try
        {
            var response = await base.SendImpl<TRequest, TResponse>(request, cancellationToken);

            stopwatch.Stop();

            // 成功事件
            if (_diagnosticListener.IsEnabled("MediatR.Request.Success"))
            {
                _diagnosticListener.Write("MediatR.Request.Success", new 
                { 
                    RequestId = requestId,
                    Duration = stopwatch.ElapsedMilliseconds,
                    Response = response
                });
            }

            return response;
        }
        catch (Exception ex)
        {
            stopwatch.Stop();

            // 失败事件
            if (_diagnosticListener.IsEnabled("MediatR.Request.Error"))
            {
                _diagnosticListener.Write("MediatR.Request.Error", new 
                { 
                    RequestId = requestId,
                    Duration = stopwatch.ElapsedMilliseconds,
                    Exception = ex
                });
            }

            throw;
        }
    }
}

订阅诊断事件 ​

csharp
public class MediatRDiagnosticsObserver : IObserver<KeyValuePair<string, object>>
{
    private readonly ILogger<MediatRDiagnosticsObserver> _logger;

    public void OnNext(KeyValuePair<string, object> kvp)
    {
        switch (kvp.Key)
        {
            case "MediatR.Request.Start":
                _logger.LogInformation("请求开始: {Data}", kvp.Value);
                break;
            case "MediatR.Request.Success":
                _logger.LogInformation("请求成功: {Data}", kvp.Value);
                break;
            case "MediatR.Request.Error":
                _logger.LogError("请求失败: {Data}", kvp.Value);
                break;
        }
    }

    public void OnCompleted() { }
    public void OnError(Exception error) { }
}

// 注册观察者
var diagnosticListener = new DiagnosticListener("MediatR");
diagnosticListener.Subscribe(new MediatRDiagnosticsObserver(logger));

🎯 实战案例:完整的自定义 Mediator ​

csharp
public class EnterpriseMediator : Mediator
{
    private readonly ILogger<EnterpriseMediator> _logger;
    private readonly DiagnosticListener _diagnostics;
    private readonly ActivitySource _activitySource;

    public EnterpriseMediator(
        IServiceProvider serviceFactory,
        ILogger<EnterpriseMediator> logger,
        DiagnosticListener diagnostics,
        ActivitySource activitySource) 
        : base(serviceFactory)
    {
        _logger = logger;
        _diagnostics = diagnostics;
        _activitySource = activitySource;
    }

    protected override async Task<TResponse> SendImpl<TRequest, TResponse>(
        IRequest<TResponse> request, 
        CancellationToken cancellationToken)
    {
        using var activity = _activitySource.StartActivity(
            $"MediatR.{typeof(TRequest).Name}",
            ActivityKind.Internal);

        activity?.SetTag("request.type", typeof(TRequest).FullName);

        try
        {
            var response = await base.SendImpl<TRequest, TResponse>(request, cancellationToken);
            
            activity?.SetStatus(ActivityStatusCode.Ok);
            return response;
        }
        catch (Exception ex)
        {
            activity?.SetStatus(ActivityStatusCode.Error, ex.Message);
            activity?.RecordException(ex);
            throw;
        }
    }

    protected override async Task PublishCore(
        IEnumerable<NotificationHandlerExecutor> handlerExecutors, 
        INotification notification, 
        CancellationToken cancellationToken)
    {
        using var activity = _activitySource.StartActivity(
            $"MediatR.{notification.GetType().Name}",
            ActivityKind.Internal);

        var tasks = handlerExecutors.Select(async executor =>
        {
            using var handlerActivity = _activitySource.StartActivity(
                $"Handler.{executor.HandlerInstance.GetType().Name}");

            try
            {
                await executor.HandlerCallback(notification, cancellationToken);
                handlerActivity?.SetStatus(ActivityStatusCode.Ok);
            }
            catch (Exception ex)
            {
                handlerActivity?.SetStatus(ActivityStatusCode.Error, ex.Message);
                handlerActivity?.RecordException(ex);
                
                _logger.LogError(ex, "处理器失败: {HandlerType}", 
                    executor.HandlerInstance.GetType().Name);
            }
        });

        await Task.WhenAll(tasks);
    }
}

🎓 总结 ​

何时自定义 Mediator? ​

场景推荐方案
通知容错处理重写 PublishCore
超时控制重写 SendImpl
诊断监控集成 DiagnosticSource
分布式追踪集成 ActivitySource
自定义路由重写 SendImpl + 特性

最佳实践 ​

✅ 推荐:

  • 保持自定义逻辑简洁
  • 记录详细的日志
  • 正确处理异常
  • 编写单元测试

❌ 避免:

  • 不要在 Mediator 中执行业务逻辑
  • 不要过度复杂化
  • 不要忘记调用 base 方法

💡 提示:自定义 Mediator 是高级特性,大多数场景下默认实现已足够!

Released under the CC BY-SA 4.0 License.