#include #include #include #include #include // NOTE: if atomic operations on std::shared_ptr<> are lock-free, // then compare/exchange would be enough to implement the lock-free // mpmc queue template class queue { private: struct node; struct counted_node_ptr { int external_count; node* ptr; }; std::atomic head; std::atomic tail; struct node_counter { unsigned internal_count:30; unsigned external_counters:2; // you need only 2 bits because there are at most two such // counters (next and tail) }; struct node { std::atomic data; std::atomic count; std::atomic next; node() { node_counter new_count; new_count.internal_count = 0; new_count.external_counters = 2; // because every node starts out referenced from tail and // from the next pointer of the previous node once you've // actually aadded it to the queue count.store(new_count); counted_node_ptr new_next; new_next.ptr = nullptr; new_next.external_count = 0; next.store(new_next); } void release_ref() { node_counter old_counter= count.load(std::memory_order_relaxed); node_counter new_counter; do { new_counter=old_counter; --new_counter.internal_count; } while(!count.compare_exchange_strong( old_counter,new_counter, std::memory_order_acquire,std::memory_order_relaxed)); if(!new_counter.internal_count && !new_counter.external_counters) { delete this; } } }; void set_new_tail(counted_node_ptr &old_tail, counted_node_ptr const &new_tail) { node* const current_tail_ptr = old_tail.ptr; while(!tail.compare_exchange_weak(old_tail,new_tail) && old_tail.ptr == current_tail_ptr); if(old_tail.ptr == current_tail_ptr) free_external_counter(old_tail); else current_tail_ptr->release_ref(); } static void increase_external_count( std::atomic& counter, counted_node_ptr& old_counter) { counted_node_ptr new_counter; do { new_counter = old_counter; ++new_counter.external_count; } while(!counter.compare_exchange_strong( old_counter,new_counter, std::memory_order_acquire,std::memory_order_relaxed)); old_counter.external_count = new_counter.external_count; } static void free_external_counter(counted_node_ptr &old_node_ptr) { node* const ptr = old_node_ptr.ptr; int const count_increase = old_node_ptr.external_count-2; node_counter old_counter = ptr->count.load(std::memory_order_relaxed); node_counter new_counter; do { new_counter=old_counter; --new_counter.external_counters; new_counter.internal_count+=count_increase; } while(!ptr->count.compare_exchange_strong( old_counter,new_counter, std::memory_order_acquire,std::memory_order_relaxed)); if(!new_counter.internal_count && !new_counter.external_counters) { delete ptr; } } public: queue() : head() , tail() {} queue(const queue& other)=delete; queue& operator=(const queue& other)=delete; ~queue() { /* while (node* const old_head = head.load()) { head.store(old_head->next); delete old_head; } */ } void push(T new_value) { std::unique_ptr new_data(new T(new_value)); counted_node_ptr new_next; new_next.ptr = new node; new_next.external_count=1; counted_node_ptr old_tail=tail.load(); for(;;) { increase_external_count(tail,old_tail); T* old_data=nullptr; if(old_tail.ptr->data.compare_exchange_strong( old_data,new_data.get())) { counted_node_ptr old_next={0}; if(!old_tail.ptr->next.compare_exchange_strong( old_next,new_next)) { delete new_next.ptr; new_next=old_next; } set_new_tail(old_tail, new_next); new_data.release(); break; } else { counted_node_ptr old_next={0}; if(old_tail.ptr->next.compare_exchange_strong( old_next,new_next)) { old_next=new_next; new_next.ptr=new node; } set_new_tail(old_tail, old_next); } } } std::unique_ptr pop() { counted_node_ptr old_head = head.load(std::memory_order_relaxed); for(;;) { increase_external_count(head, old_head); node* const ptr = old_head.ptr; if (ptr==tail.load().ptr) { ptr->release_ref(); return std::unique_ptr(); } counted_node_ptr next = ptr->next.load(); if (head.compare_exchange_strong(old_head, next)) { T* const res=ptr->data.exchange(nullptr); free_external_counter(old_head); return std::unique_ptr(res); } ptr->release_ref(); } } }; void push(queue* q) { for (int i = 0; i < 10; ++i) { printf("pushing %d\n", i); q->push(i); } } void pop(queue* q) { int i = 0; while (i < 10) { std::shared_ptr p = q->pop(); if (p) { printf("poping %d\n", *p); ++i; } } } int main() { queue q; std::thread t1(push, &q); std::thread t2(pop, &q); t1.join(); t2.join(); return 0; }