thread pool update + refactor

This commit is contained in:
MihailRis
2024-04-11 15:48:27 +03:00
parent 1549a02731
commit 90d0b54a69
17 changed files with 271 additions and 103 deletions
+74 -20
View File
@@ -4,6 +4,7 @@
#include <queue>
#include <atomic>
#include <thread>
#include <chrono>
#include <iostream>
#include <functional>
#include <condition_variable>
@@ -14,8 +15,9 @@
namespace util {
template<class T>
template<class J, class T>
struct ThreadPoolResult {
J job;
std::condition_variable& variable;
int workerIndex;
bool& locked;
@@ -34,7 +36,7 @@ template<class T, class R>
class ThreadPool : public Task {
debug::Logger logger;
std::queue<T> jobs;
std::queue<ThreadPoolResult<R>> results;
std::queue<ThreadPoolResult<T, R>> results;
std::mutex resultsMutex;
std::vector<std::thread> threads;
std::condition_variable jobsMutexCondition;
@@ -46,7 +48,9 @@ class ThreadPool : public Task {
std::atomic<int> busyWorkers = 0;
std::atomic<uint> jobsDone = 0;
bool working = true;
bool failed = false;
bool standaloneResults = true;
bool stopOnFail = true;
void threadLoop(int index, std::shared_ptr<Worker<T, R>> worker) {
std::condition_variable variable;
@@ -59,7 +63,7 @@ class ThreadPool : public Task {
jobsMutexCondition.wait(lock, [this] {
return !jobs.empty() || !working;
});
if (!working) {
if (!working || failed) {
break;
}
job = jobs.front();
@@ -71,7 +75,7 @@ class ThreadPool : public Task {
R result = (*worker)(job);
{
std::lock_guard<std::mutex> lock(resultsMutex);
results.push(ThreadPoolResult<R> {variable, index, locked, result});
results.push(ThreadPoolResult<T, R> {job, variable, index, locked, result});
if (!standaloneResults) {
locked = true;
}
@@ -88,6 +92,10 @@ class ThreadPool : public Task {
if (onJobFailed) {
onJobFailed(job);
}
if (stopOnFail) {
std::lock_guard<std::mutex> lock(jobsMutex);
failed = true;
}
logger.error() << "uncaught exception: " << err.what();
}
jobsDone++;
@@ -109,6 +117,10 @@ public:
terminate();
}
bool isActive() const override {
return working;
}
void terminate() override {
if (!working) {
return;
@@ -120,7 +132,7 @@ public:
{
std::lock_guard<std::mutex> lock(resultsMutex);
while (!results.empty()) {
ThreadPoolResult<R> entry = results.front();
ThreadPoolResult<T,R> entry = results.front();
results.pop();
if (!standaloneResults) {
entry.locked = false;
@@ -136,24 +148,54 @@ public:
}
void update() override {
std::lock_guard<std::mutex> lock(resultsMutex);
while (!results.empty()) {
ThreadPoolResult<R> entry = results.front();
results.pop();
resultConsumer(entry.entry);
if (!standaloneResults) {
entry.locked = false;
entry.variable.notify_all();
}
if (!working) {
return;
}
if (failed) {
throw std::runtime_error("some job failed");
}
if (onComplete && busyWorkers == 0) {
std::lock_guard<std::mutex> lock(jobsMutex);
if (jobs.empty()) {
onComplete();
bool complete = false;
{
std::lock_guard<std::mutex> lock(resultsMutex);
while (!results.empty()) {
ThreadPoolResult<T,R> entry = results.front();
results.pop();
try {
resultConsumer(entry.entry);
} catch (std::exception& err) {
logger.error() << err.what();
if (onJobFailed) {
onJobFailed(entry.job);
}
if (stopOnFail) {
std::lock_guard<std::mutex> lock(jobsMutex);
failed = true;
complete = false;
}
break;
}
if (!standaloneResults) {
entry.locked = false;
entry.variable.notify_all();
}
}
if (onComplete && busyWorkers == 0) {
std::lock_guard<std::mutex> lock(jobsMutex);
if (jobs.empty()) {
onComplete();
complete = true;
}
}
}
if (failed) {
throw std::runtime_error("some job failed");
}
if (complete) {
terminate();
}
}
@@ -170,6 +212,10 @@ public:
standaloneResults = flag;
}
void setStopOnFail(bool flag) {
stopOnFail = flag;
}
/// @brief onJobFailed called on exception thrown in worker thread.
/// Use engine.postRunnable when calling terminate()
void setOnJobFailed(consumer<T&> callback) {
@@ -189,6 +235,14 @@ public:
uint getWorkDone() const override {
return jobsDone;
}
virtual void waitForEnd() override {
using namespace std::chrono_literals;
while (working) {
std::this_thread::sleep_for(2ms);
update();
}
}
};
} // namespace util