FazBrowse GitHub Viewer | Trending |
URL:
| Home
Tools: [Download Repo ZIP]   [Original HTTPS Page]

Fix issue 732 semaphore exceptions by Hatem-707 · Pull Request #766 · taskflow/taskflow · GitHub

Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension .cpp  (1) .hpp  (3) All 2 file types selected
Viewed files
Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Unified
Split
Hide whitespace
Diff view
Unified
Split
Hide whitespace
33 changes: 29 additions & 4 deletions taskflow/core/executor.hpp
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters. Learn more about bidirectional Unicode characters
Original file line number Diff line number Diff line change
Expand Up @@ -1643,9 +1643,25 @@ inline void Executor::_invoke(Worker& worker, Node* node) {
}

// if acquiring semaphore(s) exists, acquire them first
if(node->_semaphores && !node->_semaphores->to_acquire.empty()) {
SmallVector<Node*> waiters;
if(!node->_acquire_all(waiters)) {
if (node->_semaphores && !node->_semaphores->to_acquire.empty()) {
SmallVector<Node *> waiters;
bool acquired = false;
bool exception_thrown = true;

// Acquiring can throw exceptions now
TF_EXECUTOR_EXCEPTION_HANDLER(worker, node, {
acquired = node->_acquire_all(waiters);
exception_thrown = false;
});

if (exception_thrown) {
_tear_down_invoke(worker, node, cache);
TF_INVOKE_CONTINUATION();
return;
}

// the node moved to semaphore queue
if (!acquired) {
_bulk_schedule(worker, waiters.begin(), waiters.size());
return;
}
Expand Down Expand Up @@ -1838,8 +1854,10 @@ inline void Executor::_observer_epilogue(Worker& worker, Node* node) {
}
}


// Procedure: _process_exception
inline void Executor::_process_exception(Worker&, Node* node) {
inline void Executor::_process_exception(Worker &, Node *node) {


// Finds the anchor and mark the entire path with exception,
// so recursive tasks can be cancelled properly.
Expand All @@ -1859,6 +1877,13 @@ inline void Executor::_process_exception(Worker&, Node* node) {
// flag used to ensure execution is caught in a thread-safe manner
constexpr static auto flag = ESTATE::EXCEPTION | ESTATE::CAUGHT;

// tarverse up chain to release waiting tasks in semaphores' after marking the chain Cancelled
// queues after on catching an exception
SmallVector<Node *> waiters;
std::unordered_set<Node *> visited;
node->_handle_broken_semaphores(waiters, visited);
_bulk_schedule(waiters.begin(), waiters.size());

// The exception occurs under a blocking call (e.g., corun, join).
if(ea) {
// multiple tasks may throw, and we only take the first thrown exception
Expand Down
22 changes: 20 additions & 2 deletions taskflow/core/graph.hpp
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters. Learn more about bidirectional Unicode characters
Original file line number Diff line number Diff line change
Expand Up @@ -487,8 +487,10 @@ class Node : public NodeBase {

void _remove_successors(Node*);
void _remove_predecessors(Node*);
};

void _handle_broken_semaphores(SmallVector<Node *> &,
std::unordered_set<Node *> &);
};

// ----------------------------------------------------------------------------
// Definition for Node::Static
Expand Down Expand Up @@ -771,7 +773,23 @@ inline void Node::_release_all(SmallVector<Node*>& nodes) {
}
}


inline void
Node::_handle_broken_semaphores(SmallVector<Node *> &waiters,
std::unordered_set<Node *> &visited) {
if (visited.contains(this)) {
return;
}
visited.insert(this);
if (_semaphores) {
auto &to_acquire = _semaphores->to_acquire;
for (auto sem : to_acquire) {
sem->_break(waiters);
}
}
for (size_t i = _num_successors; i < _edges.size(); i++) {
_edges[i]->_handle_broken_semaphores(waiters, visited);
}
}

// ----------------------------------------------------------------------------
// ExplicitAnchorGuard
Expand Down
24 changes: 24 additions & 0 deletions taskflow/core/semaphore.hpp
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters. Learn more about bidirectional Unicode characters
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,8 @@

#include "declarations.hpp"
#include "../utility/small_vector.hpp"
#include "declarations.hpp"
#include "taskflow/core/error.hpp"

/**
@file semaphore.hpp
Expand Down Expand Up @@ -118,12 +120,14 @@ class Semaphore {

size_t _max_value{0};
size_t _cur_value{0};
bool _is_broken{false};

SmallVector<Node*> _waiters;

bool _try_acquire_or_wait(Node*);

void _release(SmallVector<Node*>&);
void _break(SmallVector<Node *> &);
};

inline Semaphore::Semaphore(size_t max_value) :
Expand All @@ -133,6 +137,9 @@ inline Semaphore::Semaphore(size_t max_value) :

inline bool Semaphore::_try_acquire_or_wait(Node* me) {
std::lock_guard<std::mutex> lock(_mtx);
if (_is_broken) {
TF_THROW("Can't Acquire a Broken Semaphore.");
}
if(_cur_value > 0) {
--_cur_value;
return true;
Expand Down Expand Up @@ -175,14 +182,31 @@ inline size_t Semaphore::value() const {
inline void Semaphore::reset() {
std::lock_guard<std::mutex> lock(_mtx);
_cur_value = _max_value;
_is_broken = false;
_waiters.clear();
}

inline void Semaphore::reset(size_t new_max_value) {
std::lock_guard<std::mutex> lock(_mtx);
_cur_value = (_max_value = new_max_value);
_is_broken = false;
_waiters.clear();
}

inline void Semaphore::_break(SmallVector<Node *> &dst) {
std::lock_guard<std::mutex> lock(_mtx);
if (_is_broken) {
return;
}
_is_broken = true;
if (dst.empty()) {
dst.swap(_waiters);
} else {
dst.reserve(dst.size() + _waiters.size());
dst.insert(dst.end(), _waiters.begin(), _waiters.end());
_waiters.clear();
}
}

} // end of namespace tf. ---------------------------------------------------

98 changes: 98 additions & 0 deletions unittests/test_semaphores.cpp
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters. Learn more about bidirectional Unicode characters
Original file line number Diff line number Diff line change
Expand Up @@ -409,3 +409,101 @@ TEST_CASE("Semaphore.Deadlock.8threads" * doctest::timeout(300)) {
deadlock(8);
}
*/

void exception_deadlock(size_t semaphore_size) {

tf::Executor executor(semaphore_size * 4);
tf::Taskflow taskflow;

tf::Semaphore semaphore(semaphore_size);

tf::Task init = taskflow.emplace([]() {});
for (size_t i = 0; i < semaphore_size * 4; ++i) {
tf::Task A = taskflow.emplace([]() {});
tf::Task B = taskflow.emplace([]() {
std::this_thread::sleep_for(std::chrono::seconds(1));
throw std::runtime_error("exception");
});
tf::Task C = taskflow.emplace([]() {});
tf::Task D = taskflow.emplace([]() {});

init.precede(A);
A.precede(B);
B.precede(C);
C.precede(D);

A.acquire(semaphore);
D.release(semaphore);
}

REQUIRE(semaphore.value() == semaphore_size);

auto future = executor.run(taskflow);

// when B throws the exception, D will not run and thus semaphore is not
// released
auto status = future.wait_for(std::chrono::seconds(2));
REQUIRE(status == std::future_status::ready);
REQUIRE_THROWS_WITH_AS(future.get(), "exception", std::runtime_error);
// No Release tasks ran
REQUIRE(semaphore.value() == 0);

// acquiring a broken semaphore should result in an instant exception
auto future_2 = executor.run(taskflow);
auto status_2 = future_2.wait_for(std::chrono::seconds(2));
REQUIRE(status_2 == std::future_status::ready);
REQUIRE_THROWS_AS(future_2.get() ,std::runtime_error);

// reseting the semaphore makes it valid again
semaphore.reset();
REQUIRE(semaphore.value() == semaphore_size);
auto future_3 = executor.run(taskflow);
auto status_3 = future_3.wait_for(std::chrono::seconds(2));
REQUIRE(status_3 == std::future_status::ready);
REQUIRE_THROWS_WITH_AS(future_3.get(), "exception", std::runtime_error);
}

TEST_CASE("Exception.Exception.Deadlock.4threads" * doctest::timeout(10)) {
exception_deadlock(4);
}

void exception_cyclicgraph(size_t executors) {

tf::Executor executor(executors);
tf::Semaphore sem(1);
tf::Taskflow taskflow;
tf::Task init = taskflow.emplace([]() {});
int counter = 0;
for (size_t i = 0; i < executors ; i++) {
tf::Task A = taskflow.emplace([&counter]() { counter++; });
tf::Task B = taskflow.emplace([&counter]() {
if (counter > 0) {
return 1;
} else {
return 0;
}
});
tf::Task C = taskflow.emplace([]() {
std::this_thread::sleep_for(std::chrono::seconds(1));
throw std::runtime_error("exception");
});
init.precede(A);
A.precede(B);
B.precede(A, C);
A.acquire(sem);
C.release(sem);
}

auto future = executor.run(taskflow);
REQUIRE_THROWS_WITH_AS(future.get(), "exception", std::runtime_error);
sem.reset();
REQUIRE(sem.value() == 1);
auto future_2 = executor.run(taskflow);
auto status_2 = future_2.wait_for(std::chrono::seconds(2));
REQUIRE(status_2 == std::future_status::ready);
REQUIRE_THROWS_WITH_AS(future_2.get(), "exception", std::runtime_error);
}

TEST_CASE("Semaphore.Exception.CyclicGraph.4threads" * doctest::timeout(10)) {
exception_cyclicgraph(4);
}
Loading

Back | FazBrowse Home | New Git URL