diff --git a/httplib.h b/httplib.h index 667520a2..c04b2ffa 100644 --- a/httplib.h +++ b/httplib.h @@ -1643,6 +1643,8 @@ public: using Expect100ContinueHandler = std::function; + using StartHandler = std::function; + using WebSocketHandler = std::function; using SubProtocolSelector = @@ -1694,6 +1696,9 @@ public: Server &set_pre_request_handler(HandlerWithResponse handler); Server &set_expect_100_continue_handler(Expect100ContinueHandler handler); + + Server &set_start_handler(StartHandler handler); + Server &set_logger(Logger logger); Server &set_pre_compression_logger(Logger logger); Server &set_error_logger(ErrorLogger error_logger); @@ -1883,6 +1888,7 @@ private: Handler post_routing_handler_; HandlerWithResponse pre_request_handler_; Expect100ContinueHandler expect_100_continue_handler_; + StartHandler start_handler_; mutable std::mutex logger_mutex_; Logger logger_; @@ -11100,6 +11106,11 @@ Server::set_expect_100_continue_handler(Expect100ContinueHandler handler) { return *this; } +inline Server &Server::set_start_handler(StartHandler handler) { + start_handler_ = std::move(handler); + return *this; +} + inline Server &Server::set_address_family(int family) { address_family_ = family; return *this; @@ -11795,6 +11806,8 @@ inline bool Server::listen_internal() { is_running_ = true; auto se = detail::scope_exit([&]() { is_running_ = false; }); + if (start_handler_) { start_handler_(); } + { std::unique_ptr task_queue(new_task_queue()); diff --git a/test/test.cc b/test/test.cc index 8e2eb01e..fa4b0e04 100644 --- a/test/test.cc +++ b/test/test.cc @@ -810,6 +810,41 @@ TEST(ParseAcceptHeaderTest, ContentTypesPopulatedAndInvalidHeaderHandling) { } } +TEST(ServerStartHandlerTest, CalledOnceWhenReady) { + Server svr; + svr.Get("/", [](const Request & /*req*/, Response &res) { + res.set_content("ok", "text/plain"); + }); + + std::atomic start_count{0}; + std::atomic running_when_called{false}; + svr.set_start_handler([&]() { + running_when_called = svr.is_running(); + start_count++; + }); + + auto port = svr.bind_to_any_port(HOST); + std::thread t([&]() { svr.listen_after_bind(); }); + auto se = detail::scope_exit([&] { + svr.stop(); + t.join(); + ASSERT_FALSE(svr.is_running()); + }); + + svr.wait_until_ready(); + + // A successful request proves the accept loop is running, which the start + // handler precedes; so by now the handler must have run exactly once. + Client cli(HOST, port); + cli.set_connection_timeout(std::chrono::seconds(5)); + auto res = cli.Get("/"); + ASSERT_TRUE(res); + EXPECT_EQ(StatusCode::OK_200, res->status); + + EXPECT_EQ(1, start_count.load()); + EXPECT_TRUE(running_when_called.load()); +} + TEST(DivideTest, DivideStringTests) { auto divide = [](const std::string &str, char d) { std::string lhs;