【问题标题】:Thread-safe Cached Enumerator - lock with yield线程安全的缓存枚举器 - 带收益的锁
【发布时间】:2020-02-20 19:00:21
【问题描述】:

我有一个自定义的“CachedEnumerable”类(灵感来自 Caching IEnumerable),我需要为我的 asp.net 核心 Web 应用程序设置线程安全。

以下 Enumerator 线程的实现是否安全? (对 IList _cache 的所有其他读/写都被适当锁定)(可能与 Does the C# Yield free a lock? 相关)

更具体地说,如果有 2 个线程访问枚举器,我如何防止一个线程递增“索引”导致第二个枚举线程从 _cache 获取错误的元素(即索引 + 1 处的元素而不是在索引处)?这种竞争条件真的值得关注吗?

public IEnumerator<T> GetEnumerator()
{
    var index = 0;

    while (true)
    {
        T current;
        lock (_enumeratorLock)
        {
            if (index >= _cache.Count && !MoveNext()) break;
            current = _cache[index];
            index++;
        }
        yield return current;
    }
}

我的 CachedEnumerable 版本的完整代码:

 public class CachedEnumerable<T> : IDisposable, IEnumerable<T>
    {
        IEnumerator<T> _enumerator;
        private IList<T> _cache = new List<T>();
        public bool CachingComplete { get; private set; } = false;

        public CachedEnumerable(IEnumerable<T> enumerable)
        {
            switch (enumerable)
            {
                case CachedEnumerable<T> cachedEnumerable: //This case is actually dealt with by the extension method.
                    _cache = cachedEnumerable._cache;
                    CachingComplete = cachedEnumerable.CachingComplete;
                    _enumerator = cachedEnumerable.GetEnumerator();

                    break;
                case IList<T> list:
                    //_cache = list; //without clone...
                    //Clone:
                    _cache = new T[list.Count];
                    list.CopyTo((T[]) _cache, 0);
                    CachingComplete = true;
                    break;
                default:
                    _enumerator = enumerable.GetEnumerator();
                    break;
            }
        }

        public CachedEnumerable(IEnumerator<T> enumerator)
        {
            _enumerator = enumerator;
        }

        private int CurCacheCount
        {
            get
            {
                lock (_enumeratorLock)
                {
                    return _cache.Count;
                }
            }
        }

        public IEnumerator<T> GetEnumerator()
        {
            var index = 0;

            while (true)
            {
                T current;
                lock (_enumeratorLock)
                {
                    if (index >= _cache.Count && !MoveNext()) break;
                    current = _cache[index];
                    index++;
                }
                yield return current;
            }
        }

        //private readonly AsyncLock _enumeratorLock = new AsyncLock();
        private readonly object _enumeratorLock = new object();

        private bool MoveNext()
        {
            if (CachingComplete) return false;

            if (_enumerator != null && _enumerator.MoveNext()) //The null check should have been unnecessary b/c of the lock...
            {
                _cache.Add(_enumerator.Current);
                return true;
            }
            else
            {
                CachingComplete = true;
                DisposeWrappedEnumerator(); //Release the enumerator, as it is no longer needed.
            }

            return false;
        }

        public T ElementAt(int index)
        {
            lock (_enumeratorLock)
            {
                if (index < _cache.Count)
                {
                    return _cache[index];
                }
            }

            EnumerateUntil(index);

            lock (_enumeratorLock)
            {
                if (_cache.Count <= index) throw new ArgumentOutOfRangeException(nameof(index));
                return _cache[index];
            }
        }


        public bool TryGetElementAt(int index, out T value)
        {
            lock (_enumeratorLock)
            {
                value = default;
                if (index < CurCacheCount)
                {
                    value = _cache[index];
                    return true;
                }
            }

            EnumerateUntil(index);

            lock (_enumeratorLock)
            {
                if (_cache.Count <= index) return false;
                value = _cache[index];
            }

            return true;
        }

        private void EnumerateUntil(int index)
        {
            while (true)
            {
                lock (_enumeratorLock)
                {
                    if (_cache.Count > index || !MoveNext()) break;
                }
            }
        }


        public void Dispose()
        {
            DisposeWrappedEnumerator();
        }

        private void DisposeWrappedEnumerator()
        {
            if (_enumerator != null)
            {
                _enumerator.Dispose();
                _enumerator = null;
                if (_cache is List<T> list)
                {
                    list.Trim();
                }
            }
        }

        IEnumerator IEnumerable.GetEnumerator()
        {
            return GetEnumerator();
        }

        public int CachedCount
        {
            get
            {
                lock (_enumeratorLock)
                {
                    return _cache.Count;
                }
            }
        }

        public int Count()
        {
            if (CachingComplete)
            {
                return _cache.Count;
            }

            EnsureCachingComplete();

            return _cache.Count;
        }

        private void EnsureCachingComplete()
        {
            if (CachingComplete)
            {
                return;
            }

            //Enumerate the rest of the collection
            while (!CachingComplete)
            {
                lock (_enumeratorLock)
                {
                    if (!MoveNext()) break;
                }
            }
        }

        public T[] ToArray()
        {
            EnsureCachingComplete();
            //Once Caching is complete, we don't need to lock
            if (!(_cache is T[] array))
            {
                array = _cache.ToArray();
                _cache = array;
            }

            return array;
        }

        public T this[int index] => ElementAt(index);
    }

    public static CachedEnumerable<T> Cached<T>(this IEnumerable<T> source)
    {
        //no gain in caching a cache.
        if (source is CachedEnumerable<T> cached)
        {
            return cached;
        }

        return new CachedEnumerable<T>(source);
    }
}

基本用法:(虽然不是一个有意义的用例)

var cached = expensiveEnumerable.Cached();
foreach (var element in cached) {
   Console.WriteLine(element);
}

更新

我根据@Theodors 答案https://stackoverflow.com/a/58547863/5683904 测试了当前实现,并确认 (AFAICT) 在使用 foreach 枚举时它是线程安全的而不创建重复值 (Thread-safe Cached Enumerator - lock with yield):

class Program
{
    static async Task Main(string[] args)
    {
        var enumerable = Enumerable.Range(0, 1_000_000);
        var cachedEnumerable = new CachedEnumerable<int>(enumerable);
        var c = new ConcurrentDictionary<int, List<int>>();
        var tasks = Enumerable.Range(1, 100).Select(id => Test(id, cachedEnumerable, c));
        Task.WaitAll(tasks.ToArray());
        foreach (var keyValuePair in c)
        {
            var hasDuplicates = keyValuePair.Value.Distinct().Count() != keyValuePair.Value.Count;
            Console.WriteLine($"Task #{keyValuePair.Key} count: {keyValuePair.Value.Count}. Has duplicates? {hasDuplicates}");
        }
    }

    static async Task Test(int id, IEnumerable<int> cache, ConcurrentDictionary<int, List<int>> c)
    {
        foreach (var i in cache)
        {
            //await Task.Delay(10);
            c.AddOrUpdate(id, v => new List<int>() {i}, (k, v) =>
            {
                v.Add(i);
                return v;
            });
        }
    }
}

【问题讨论】:

  • 线程安全与否,如果在枚举时修改_cache,枚举将根本不正确。如果_cache 以线程安全的方式初始化一次,则不需要用锁保护进一步的访问。如果不是,那么无论如何,所有的赌注都会被取消。
  • 我提供了 CachedEnumerable 的完整代码。 _cache 是一个私有字段,只能通过增长来修改。它提供对缓存中元素的只读随机访问。
  • 所有这些东西似乎都被荒谬地过度设计了,而只是将一个可枚举的.ToList() 松掉并使用它,顺便说一句,这可以显然线程-安全的。与此相比,这些恶作剧可能会节省第一个请求的一些内存/响应能力,但您不仅会失去对行为进行推理的能力。不得不在互联网上随机询问是否线程安全从来都不是一个好兆头。
  • 我只用它来记忆昂贵的计算,这些计算是进程和内存密集型的,是延迟生成和多次使用的。
  • 你能告诉我们这是如何使用的(生产者和消费者双方)吗?

标签: c# multithreading linq


【解决方案1】:

您的类不是线程安全的,因为共享状态在您的类内未受保护的区域中发生了变异。未受保护的区域是:

  1. 构造函数
  2. Dispose 方法

共享状态为:

  1. _enumerator 私有字段
  2. _cache 私有字段
  3. CachingComplete 公共财产

关于您的课程的其他一些问题:

  1. 实现IDisposable 为调用者创建了处置您的类的责任。 IEnumerables 不需要是一次性的。相反IEnumerators 是一次性的,但它们的自动处理有语言支持(foreach 语句的功能)。
  2. 您的课程提供了IEnumerableElementAtCount 等)所不期望的扩展功能。也许您打算改为实现CachedList?如果不实现IList&lt;T&gt; 接口,像Count()ToArray() 这样的LINQ 方法将无法利用您的扩展功能,并且会像使用普通的IEnumerables 一样使用慢速路径。

更新:我刚刚注意到另一个线程安全问题。这一项与public IEnumerator&lt;T&gt; GetEnumerator() 方法有关。枚举器是编译器生成的,因为该方法是一个迭代器(使用yield return)。编译器生成的枚举器不是线程安全的。以这段代码为例:

var enumerable = Enumerable.Range(0, 1_000_000);
var cachedEnumerable = new CachedEnumerable<int>(enumerable);
var enumerator = cachedEnumerable.GetEnumerator();
var tasks = Enumerable.Range(1, 4).Select(id => Task.Run(() =>
{
    int count = 0;
    while (enumerator.MoveNext())
    {
        count++;
    }
    Console.WriteLine($"Task #{id} count: {count}");
})).ToArray();
Task.WaitAll(tasks);

四个线程同时使用相同的IEnumerator。可枚举项有 1,000,000 项。您可能期望每个线程会枚举约 250,000 个项目,但事实并非如此。

输出:

任务 #1 计数:0
任务 #4 计数:0
任务 #3 计数:0
任务 #2 计数:1000000

while (enumerator.MoveNext()) 行中的MoveNext 不是您的保险箱MoveNext。它是编译器生成的不安全MoveNext。虽然不安全,但它includes a mechanism intended probably for dealing with exceptions 在调用外部提供的代码之前暂时将枚举器标记为已完成。因此,当多个线程同时调用MoveNext 时,除第一个线程之外的所有线程都将获得false 的返回值,并在完成零循环后立即终止枚举。要解决这个问题,您可能必须编写自己的 IEnumerator 类。


更新:其实我最后一点关于线程安全枚举有点不公平,因为使用IEnumerator接口进行枚举本质上是不安全的操作,如果没有线程安全的合作是不可能解决的调用代码。这是因为获取下一个元素不是原子操作,因为它涉及两个步骤(调用MoveNext() + 读取Current)。因此,您的线程安全问题仅限于保护类的内部状态(字段_enumerator_cacheCachingComplete)。这些仅在构造函数和Dispose 方法中不受保护,但我认为您的类的正常使用可能不会遵循创建会导致内部状态损坏的竞争条件的代码路径。

就我个人而言,我也更愿意处理这些代码路径,而且我不会让它随心所欲。


更新:我为IAsyncEnumerables 编写了一个缓存,以演示另一种技术。源IAsyncEnumerable 的枚举不是由调用者驱动的,使用锁或信号量来获得独占访问,而是由一个单独的工作任务驱动。第一个调用者启动工作任务。每个调用者首先产生所有已缓存的项目,然后等待更多项目,或等待没有更多项目的通知。作为通知机制,我使用了TaskCompletionSource&lt;bool&gt;lock 仍然用于确保对共享资源的所有访问都是同步的。

public class CachedAsyncEnumerable<T> : IAsyncEnumerable<T>
{
    private readonly object _locker = new object();
    private IAsyncEnumerable<T> _source;
    private Task _sourceEnumerationTask;
    private List<T> _buffer;
    private TaskCompletionSource<bool> _moveNextTCS;
    private Exception _sourceEnumerationException;
    private int _sourceEnumerationVersion; // Incremented on exception

    public CachedAsyncEnumerable(IAsyncEnumerable<T> source)
    {
        _source = source ?? throw new ArgumentNullException(nameof(source));
    }

    public async IAsyncEnumerator<T> GetAsyncEnumerator(
        CancellationToken cancellationToken = default)
    {
        lock (_locker)
        {
            if (_sourceEnumerationTask == null)
            {
                _buffer = new List<T>();
                _moveNextTCS = new TaskCompletionSource<bool>();
                _sourceEnumerationTask = Task.Run(
                    () => EnumerateSourceAsync(cancellationToken));
            }
        }
        int index = 0;
        int localVersion = -1;
        while (true)
        {
            T current = default;
            Task<bool> moveNextTask = null;
            lock (_locker)
            {
                if (localVersion == -1)
                {
                    localVersion = _sourceEnumerationVersion;
                }
                else if (_sourceEnumerationVersion != localVersion)
                {
                    ExceptionDispatchInfo
                        .Capture(_sourceEnumerationException).Throw();
                }
                if (index < _buffer.Count)
                {
                    current = _buffer[index];
                    index++;
                }
                else
                {
                    moveNextTask = _moveNextTCS.Task;
                }
            }
            if (moveNextTask == null)
            {
                yield return current;
                continue;
            }
            var moved = await moveNextTask;
            if (!moved) yield break;
            lock (_locker)
            {
                current = _buffer[index];
                index++;
            }
            yield return current;
        }
    }

    private async Task EnumerateSourceAsync(CancellationToken cancellationToken)
    {
        TaskCompletionSource<bool> localMoveNextTCS;
        try
        {
            await foreach (var item in _source.WithCancellation(cancellationToken))
            {
                lock (_locker)
                {
                    _buffer.Add(item);
                    localMoveNextTCS = _moveNextTCS;
                    _moveNextTCS = new TaskCompletionSource<bool>();
                }
                localMoveNextTCS.SetResult(true);
            }
            lock (_locker)
            {
                localMoveNextTCS = _moveNextTCS;
                _buffer.TrimExcess();
                _source = null;
            }
            localMoveNextTCS.SetResult(false);
        }
        catch (Exception ex)
        {
            lock (_locker)
            {
                localMoveNextTCS = _moveNextTCS;
                _sourceEnumerationException = ex;
                _sourceEnumerationVersion++;
                _sourceEnumerationTask = null;
            }
            localMoveNextTCS.SetException(ex);
        }
    }
}

此实现遵循处理异常的特定策略。如果在枚举源IAsyncEnumerable时发生异常,该异常将传播到所有当前调用者,当前使用的IAsyncEnumerator将被丢弃,不完整的缓存数据也将被丢弃。当接收到下一个枚举请求时,新的工作任务可能会在稍后再次启动。

【讨论】:

  • 增加了一项观察。
  • 哇!谢谢 - 这正是我计划编写的测试 (stackoverflow.com/questions/58541336/…)。我担心产量扩展到什么,但我以前从未遇到过这一点。关于您的其他观点:我没有实现 IList,因为它在 Linq 或 Json.Net 完全实现列表的其他地方引起了问题,即使我特别希望它是懒惰的。
  • 不过,我用 foreach 对其进行了测试,它工作正常。您指出的问题是 IEnumerator 不是线程安全的。 var enumerable = Enumerable.Range(0, 10); var cachedEnumerable = new HelpersAndExtentions.CachedEnumerable(enumerable); var tasks = Enumerable.Range(1, 4).Select(id => Task.Run(() => { foreach (var i in cachedEnumerable) { Console.WriteLine($"Task #{id} count: {i} "); } })).ToArray(); Task.WaitAll(tasks);
  • 关于为什么CachedEnumerable 没有实现IList&lt;T&gt; 的解释很好!我再次更正了我的答案,因为我认为我之前的观点是不公平的。
  • 我添加了一个CachedAsyncEnumerable 实现,以演示解决相同问题的另一种方法。
【解决方案2】:

对缓存的访问,是的,它是线程安全的,每次只有一个线程可以从_cache对象中读取。

但是这样你不能保证所有线程都按照它们访问 GetEnumerator 的顺序获取元素。

检查这两个例子,如果行为是你所期望的,你可以这样使用锁。

示例 1:

THREAD1 调用 GetEnumerator

THREAD1 初始化T电流;

THREAD2 调用 GetEnumerator

THREAD2 初始化T电流;

线程2锁线程

线程 1 等待

THREAD2 从缓存中安全读取_cache[0]

THREAD2 索引++

线程 2 解锁

线程 1 锁定

THREAD1 从缓存中安全读取_cache[1]

THREAD1 i++

线程 1 解锁

THREAD2 产生返回电流;

THREAD1 产生返回电流;


示例 2:

THREAD2 初始化T电流;

线程2锁线程

THREAD2 从缓存中安全读取

线程 2 解锁

THREAD1 初始化T电流;

线程 1 锁定线程

THREAD1 从缓存中安全读取

线程 1 解锁

THREAD1 产生返回电流;

THREAD2 产生返回电流;

【讨论】:

  • 我只检查了你要求的方法,我没有检查是否所有的类都是线程安全的,
  • 在示例 1 中,THREAD 1 是否应该返回 _cache[0],但由于 THREAD 2 增加了索引,它是否会返回 current = _cache[1]?这是我最担心的错误。
  • 我认为这不是问题所在,该类本身是线程安全的,但您无法通过 Getenumerator 内的同步来确保正确的顺序。如果您想保留订单,调用者的责任是做适当的同步。如果两个线程同时访问,你怎么能确定之前是哪个线程?
  • 无论如何看看referencesource.microsoft.com/#mscorlib/system/Collections/…,它可以为您提供更好的方法来实现线程安全的GetEnumerator 而不会出现死锁。如您所见,它使用 SpinWait 而不是 lock 语句。
  • 我不太确定您所做的测试,我认为代码线程安全,因为始终枚举所有元素而不会引发异常和冲突。但是由于不同的线程同时执行,它们不能返回每个线程的所有值。使用 lock 你只保证每个元素只被不同的线程处理一次。我认为@Theodor Zoulias 的回答更加完整和可靠,因为他还分析了编译器生成的代码并展示了一个可重现的示例。
猜你喜欢
  • 2014-09-09
  • 1970-01-01
  • 1970-01-01
  • 1970-01-01
  • 2014-08-13
  • 1970-01-01
  • 1970-01-01
  • 2018-10-02
  • 1970-01-01
相关资源
最近更新 更多