Threading api refactor (#1955)

Refactor the multi-threading api to support
using custom user-provided thread factory
instead of always spawning POSIX Threads.
diff --git a/docs/user_guide.md b/docs/user_guide.md
index b3c1cce..0bcfe15 100644
--- a/docs/user_guide.md
+++ b/docs/user_guide.md
@@ -863,6 +863,46 @@
 
 Without `UseRealTime`, CPU time is used by default.
 
+### Manual Multithreaded Benchmarks
+
+Google/benchmark uses `std::thread` as multithreading environment per default.
+If you want to use another multithreading environment (e.g. OpenMP), you can provide
+a factory function to your benchmark using the `ThreadRunner` function.
+The factory function takes the number of threads as argument and creates a custom class
+derived from `benchmark::ThreadRunnerBase`.
+This custom class must override the function
+`void RunThreads(const std::function<void(int)>& fn)`.
+`RunThreads` is called by the main thread and spawns the requested number of threads.
+Each spawned thread must call `fn(thread_index)`, where `thread_index` is its own
+thread index. Before `RunThreads` returns, all spawned threads must be joined.
+```c++
+class OpenMPThreadRunner : public benchmark::ThreadRunnerBase
+{
+  OpenMPThreadRunner(int num_threads)
+  : num_threads_(num_threads)
+  {}
+
+  void RunThreads(const std::function<void(int)>& fn) final
+  {
+#pragma omp parallel num_threads(num_threads_)
+    fn(omp_get_thread_num());
+  }
+
+private:
+  int num_threads_;
+};
+
+BENCHMARK(BM_MultiThreaded)
+  ->ThreadRunner([](int num_threads) {
+    return std::make_unique<OpenMPThreadRunner>(num_threads);
+  })
+  ->Threads(1)->Threads(2)->Threads(4);
+```
+The above example creates a parallel OpenMP region before it enters `BM_MultiThreaded`.
+The actual benchmark code can remain the same and is therefore not tied to a specific
+thread runner. The measurement does not include the time for creating and joining the
+threads.
+
 <a name="cpu-timers" />
 
 ## CPU Timers
diff --git a/include/benchmark/benchmark.h b/include/benchmark/benchmark.h
index 624ab29..5e40975 100644
--- a/include/benchmark/benchmark.h
+++ b/include/benchmark/benchmark.h
@@ -1093,8 +1093,18 @@
   return StateIterator();
 }
 
+// Base class for user-defined multi-threading
+struct ThreadRunnerBase {
+  virtual ~ThreadRunnerBase() {}
+  virtual void RunThreads(const std::function<void(int)>& fn) = 0;
+};
+
 namespace internal {
 
+// Define alias of ThreadRunner factory function type
+using threadrunner_factory =
+    std::function<std::unique_ptr<ThreadRunnerBase>(int)>;
+
 typedef void(Function)(State&);
 
 // ------------------------------------------------------
@@ -1299,6 +1309,9 @@
   // Equivalent to ThreadRange(NumCPUs(), NumCPUs())
   Benchmark* ThreadPerCpu();
 
+  // Sets a user-defined threadrunner (see ThreadRunnerBase)
+  Benchmark* ThreadRunner(threadrunner_factory&& factory);
+
   virtual void Run(State& state) = 0;
 
   TimeUnit GetTimeUnit() const;
@@ -1340,6 +1353,8 @@
   callback_function setup_;
   callback_function teardown_;
 
+  threadrunner_factory threadrunner_;
+
   BENCHMARK_DISALLOW_COPY_AND_ASSIGN(Benchmark);
 };
 
diff --git a/src/benchmark_api_internal.h b/src/benchmark_api_internal.h
index 82ab71f..efa0602 100644
--- a/src/benchmark_api_internal.h
+++ b/src/benchmark_api_internal.h
@@ -41,6 +41,9 @@
   int threads() const { return threads_; }
   void Setup() const;
   void Teardown() const;
+  const auto& GetUserThreadRunnerFactory() const {
+    return benchmark_.threadrunner_;
+  }
 
   State Run(IterationCount iters, int thread_id, internal::ThreadTimer* timer,
             internal::ThreadManager* manager,
diff --git a/src/benchmark_register.cc b/src/benchmark_register.cc
index 8b94540..d8cefe4 100644
--- a/src/benchmark_register.cc
+++ b/src/benchmark_register.cc
@@ -484,6 +484,11 @@
   return this;
 }
 
+Benchmark* Benchmark::ThreadRunner(threadrunner_factory&& factory) {
+  threadrunner_ = std::move(factory);
+  return this;
+}
+
 void Benchmark::SetName(const std::string& name) { name_ = name; }
 
 const char* Benchmark::GetName() const { return name_.c_str(); }
diff --git a/src/benchmark_runner.cc b/src/benchmark_runner.cc
index 55ad69c..427bd85 100644
--- a/src/benchmark_runner.cc
+++ b/src/benchmark_runner.cc
@@ -34,6 +34,7 @@
 #include <cstdio>
 #include <cstdlib>
 #include <fstream>
+#include <functional>
 #include <iostream>
 #include <limits>
 #include <memory>
@@ -182,6 +183,38 @@
   return iters_or_time.iters;
 }
 
+class ThreadRunnerDefault : public ThreadRunnerBase {
+ public:
+  explicit ThreadRunnerDefault(int num_threads)
+      : pool(static_cast<size_t>(num_threads - 1)) {}
+
+  void RunThreads(const std::function<void(int)>& fn) final {
+    // Run all but one thread in separate threads
+    for (std::size_t ti = 0; ti < pool.size(); ++ti) {
+      pool[ti] = std::thread(fn, static_cast<int>(ti + 1));
+    }
+    // And run one thread here directly.
+    // (If we were asked to run just one thread, we don't create new threads.)
+    // Yes, we need to do this here *after* we start the separate threads.
+    fn(0);
+
+    // The main thread has finished. Now let's wait for the other threads.
+    for (std::thread& thread : pool) {
+      thread.join();
+    }
+  }
+
+ private:
+  std::vector<std::thread> pool;
+};
+
+std::unique_ptr<ThreadRunnerBase> GetThreadRunner(
+    const threadrunner_factory& userThreadRunnerFactory, int num_threads) {
+  return userThreadRunnerFactory
+             ? userThreadRunnerFactory(num_threads)
+             : std::make_unique<ThreadRunnerDefault>(num_threads);
+}
+
 }  // end namespace
 
 BenchTimeType ParseBenchMinTime(const std::string& value) {
@@ -258,7 +291,8 @@
       has_explicit_iteration_count(b.iterations() != 0 ||
                                    parsed_benchtime_flag.tag ==
                                        BenchTimeType::ITERS),
-      pool(static_cast<size_t>(b.threads() - 1)),
+      thread_runner(
+          GetThreadRunner(b.GetUserThreadRunnerFactory(), b.threads())),
       iters(FLAGS_benchmark_dry_run
                 ? 1
                 : (has_explicit_iteration_count
@@ -289,22 +323,10 @@
   std::unique_ptr<internal::ThreadManager> manager;
   manager.reset(new internal::ThreadManager(b.threads()));
 
-  // Run all but one thread in separate threads
-  for (std::size_t ti = 0; ti < pool.size(); ++ti) {
-    pool[ti] = std::thread(&RunInThread, &b, iters, static_cast<int>(ti + 1),
-                           manager.get(), perf_counters_measurement_ptr,
-                           /*profiler_manager=*/nullptr);
-  }
-  // And run one thread here directly.
-  // (If we were asked to run just one thread, we don't create new threads.)
-  // Yes, we need to do this here *after* we start the separate threads.
-  RunInThread(&b, iters, 0, manager.get(), perf_counters_measurement_ptr,
-              /*profiler_manager=*/nullptr);
-
-  // The main thread has finished. Now let's wait for the other threads.
-  for (std::thread& thread : pool) {
-    thread.join();
-  }
+  thread_runner->RunThreads([&](int thread_idx) {
+    RunInThread(&b, iters, thread_idx, manager.get(),
+                perf_counters_measurement_ptr, /*profiler_manager=*/nullptr);
+  });
 
   IterationResults i;
   // Acquire the measurements/counters from the manager, UNDER THE LOCK!
diff --git a/src/benchmark_runner.h b/src/benchmark_runner.h
index bc76c81..9a2231a 100644
--- a/src/benchmark_runner.h
+++ b/src/benchmark_runner.h
@@ -15,6 +15,7 @@
 #ifndef BENCHMARK_RUNNER_H_
 #define BENCHMARK_RUNNER_H_
 
+#include <memory>
 #include <thread>
 #include <vector>
 
@@ -89,7 +90,7 @@
 
   int num_repetitions_done = 0;
 
-  std::vector<std::thread> pool;
+  std::unique_ptr<ThreadRunnerBase> thread_runner;
 
   IterationCount iters;  // preserved between repetitions!
   // So only the first repetition has to find/calculate it,
diff --git a/test/CMakeLists.txt b/test/CMakeLists.txt
index df575d9..bde248f 100644
--- a/test/CMakeLists.txt
+++ b/test/CMakeLists.txt
@@ -189,6 +189,9 @@
 compile_output_test(internal_threading_test)
 benchmark_add_test(NAME internal_threading_test COMMAND internal_threading_test --benchmark_min_time=0.01s)
 
+compile_output_test(manual_threading_test)
+benchmark_add_test(NAME manual_threading_test COMMAND manual_threading_test --benchmark_min_time=0.01s)
+
 compile_output_test(report_aggregates_only_test)
 benchmark_add_test(NAME report_aggregates_only_test COMMAND report_aggregates_only_test --benchmark_min_time=0.01s)
 
diff --git a/test/manual_threading_test.cc b/test/manual_threading_test.cc
new file mode 100644
index 0000000..e85d495
--- /dev/null
+++ b/test/manual_threading_test.cc
@@ -0,0 +1,174 @@
+
+#include <memory>
+#undef NDEBUG
+
+#include <chrono>
+#include <thread>
+
+#include "../src/timers.h"
+#include "benchmark/benchmark.h"
+
+namespace {
+
+const std::chrono::duration<double, std::milli> time_frame(50);
+const double time_frame_in_sec(
+    std::chrono::duration_cast<std::chrono::duration<double, std::ratio<1, 1>>>(
+        time_frame)
+        .count());
+
+void MyBusySpinwait() {
+  const auto start = benchmark::ChronoClockNow();
+
+  while (true) {
+    const auto now = benchmark::ChronoClockNow();
+    const auto elapsed = now - start;
+
+    if (std::chrono::duration<double, std::chrono::seconds::period>(elapsed) >=
+        time_frame) {
+      return;
+    }
+  }
+}
+
+int numRunThreadsCalled_ = 0;
+
+class ManualThreadRunner : public benchmark::ThreadRunnerBase {
+ public:
+  explicit ManualThreadRunner(int num_threads)
+      : pool(static_cast<size_t>(num_threads - 1)) {}
+
+  void RunThreads(const std::function<void(int)>& fn) final {
+    for (std::size_t ti = 0; ti < pool.size(); ++ti) {
+      pool[ti] = std::thread(fn, static_cast<int>(ti + 1));
+    }
+
+    fn(0);
+
+    for (std::thread& thread : pool) {
+      thread.join();
+    }
+
+    ++numRunThreadsCalled_;
+  }
+
+ private:
+  std::vector<std::thread> pool;
+};
+
+// ========================================================================= //
+// --------------------------- TEST CASES BEGIN ---------------------------- //
+// ========================================================================= //
+
+// ========================================================================= //
+// BM_ManualThreading
+// Creation of threads is done before the start of the measurement,
+// joining after the finish of the measurement.
+void BM_ManualThreading(benchmark::State& state) {
+  for (auto _ : state) {
+    MyBusySpinwait();
+    state.SetIterationTime(time_frame_in_sec);
+  }
+  state.counters["invtime"] =
+      benchmark::Counter{1, benchmark::Counter::kIsRate};
+}
+
+}  // end namespace
+
+BENCHMARK(BM_ManualThreading)
+    ->Iterations(1)
+    ->ThreadRunner([](int num_threads) {
+      return std::make_unique<ManualThreadRunner>(num_threads);
+    })
+    ->Threads(1);
+BENCHMARK(BM_ManualThreading)
+    ->Iterations(1)
+    ->ThreadRunner([](int num_threads) {
+      return std::make_unique<ManualThreadRunner>(num_threads);
+    })
+    ->Threads(1)
+    ->UseRealTime();
+BENCHMARK(BM_ManualThreading)
+    ->Iterations(1)
+    ->ThreadRunner([](int num_threads) {
+      return std::make_unique<ManualThreadRunner>(num_threads);
+    })
+    ->Threads(1)
+    ->UseManualTime();
+BENCHMARK(BM_ManualThreading)
+    ->Iterations(1)
+    ->ThreadRunner([](int num_threads) {
+      return std::make_unique<ManualThreadRunner>(num_threads);
+    })
+    ->Threads(1)
+    ->MeasureProcessCPUTime();
+BENCHMARK(BM_ManualThreading)
+    ->Iterations(1)
+    ->ThreadRunner([](int num_threads) {
+      return std::make_unique<ManualThreadRunner>(num_threads);
+    })
+    ->Threads(1)
+    ->MeasureProcessCPUTime()
+    ->UseRealTime();
+BENCHMARK(BM_ManualThreading)
+    ->Iterations(1)
+    ->ThreadRunner([](int num_threads) {
+      return std::make_unique<ManualThreadRunner>(num_threads);
+    })
+    ->Threads(1)
+    ->MeasureProcessCPUTime()
+    ->UseManualTime();
+
+BENCHMARK(BM_ManualThreading)
+    ->Iterations(1)
+    ->ThreadRunner([](int num_threads) {
+      return std::make_unique<ManualThreadRunner>(num_threads);
+    })
+    ->Threads(2);
+BENCHMARK(BM_ManualThreading)
+    ->Iterations(1)
+    ->ThreadRunner([](int num_threads) {
+      return std::make_unique<ManualThreadRunner>(num_threads);
+    })
+    ->Threads(2)
+    ->UseRealTime();
+BENCHMARK(BM_ManualThreading)
+    ->Iterations(1)
+    ->ThreadRunner([](int num_threads) {
+      return std::make_unique<ManualThreadRunner>(num_threads);
+    })
+    ->Threads(2)
+    ->UseManualTime();
+BENCHMARK(BM_ManualThreading)
+    ->Iterations(1)
+    ->ThreadRunner([](int num_threads) {
+      return std::make_unique<ManualThreadRunner>(num_threads);
+    })
+    ->Threads(2)
+    ->MeasureProcessCPUTime();
+BENCHMARK(BM_ManualThreading)
+    ->Iterations(1)
+    ->ThreadRunner([](int num_threads) {
+      return std::make_unique<ManualThreadRunner>(num_threads);
+    })
+    ->Threads(2)
+    ->MeasureProcessCPUTime()
+    ->UseRealTime();
+BENCHMARK(BM_ManualThreading)
+    ->Iterations(1)
+    ->ThreadRunner([](int num_threads) {
+      return std::make_unique<ManualThreadRunner>(num_threads);
+    })
+    ->Threads(2)
+    ->MeasureProcessCPUTime()
+    ->UseManualTime();
+
+// ========================================================================= //
+// ---------------------------- TEST CASES END ----------------------------- //
+// ========================================================================= //
+
+int main(int argc, char* argv[]) {
+  benchmark::Initialize(&argc, argv);
+  benchmark::RunSpecifiedBenchmarks();
+  benchmark::Shutdown();
+  assert(numRunThreadsCalled_ > 0);
+}