[Cast] Lower cpu usage by optimizing socket select() call This patch optimizes a check for Writable handles, ensuring that we don't pass a handle to the write argument for select() if we don't have anything to write. This lowers the amount of cycles spent polling the select loop. This is a not-insignificant change, involving updating the SocketHandleWaiter interface and its subscribers to allow the waiter to poll for whether the socket has pending write events or not. Change-Id: Icbd38a67fbc1fa4856f9605d6306971568dcca5e Reviewed-on: https://chromium-review.googlesource.com/c/openscreen/+/6556928 Reviewed-by: Muyao Xu <muyaoxu@google.com> Commit-Queue: Jordan Bayles <jophba@chromium.org>
diff --git a/platform/BUILD.gn b/platform/BUILD.gn index e22250e..401f344 100644 --- a/platform/BUILD.gn +++ b/platform/BUILD.gn
@@ -268,6 +268,7 @@ if (is_posix) { sources += [ "impl/logging_unittest.cc", + "impl/platform_client_posix_unittest.cc", "impl/scoped_pipe_unittest.cc", "impl/socket_address_posix_unittest.cc", "impl/socket_handle_waiter_posix_unittest.cc",
diff --git a/platform/impl/platform_client_posix.cc b/platform/impl/platform_client_posix.cc index 0b49213..c041ebe 100644 --- a/platform/impl/platform_client_posix.cc +++ b/platform/impl/platform_client_posix.cc
@@ -4,6 +4,7 @@ #include "platform/impl/platform_client_posix.h" +#include <chrono> #include <functional> #include <utility> #include <vector> @@ -104,16 +105,41 @@ } void PlatformClientPosix::RunNetworkLoopUntilStopped() { - auto last = Clock::now(); +#if OSP_DCHECK_IS_ON() + Clock::time_point last_time = Clock::now(); + int iterations = 0; +#endif while (networking_loop_running_.load()) { +#if OSP_DCHECK_IS_ON() + ++iterations; + const Clock::time_point current_time = Clock::now(); + const Clock::duration delta = current_time - last_time; + if (delta > std::chrono::seconds(1)) { + OSP_DCHECK_GT(iterations, 0); + OSP_VLOG << "network loop execution time averaged " + << (delta / iterations) << " over the last second."; + last_time = current_time; + iterations = 0; + } +#endif if (!waiter_created_.load()) { std::this_thread::sleep_for(networking_loop_timeout_); continue; } - socket_handle_waiter()->ProcessHandles(networking_loop_timeout_); - auto now = Clock::now(); - OSP_LOG_ERROR << __func__ << ": loop took " << (now - last); - last = now; + const Error process_error = + socket_handle_waiter()->ProcessHandles(networking_loop_timeout_); + + // We may receive an "again" error code if there were no sockets to process. + if (process_error.code() == Error::Code::kAgain) { + std::this_thread::sleep_for(networking_loop_timeout_); + continue; + + // If there is a socket error it should be handled elsewhere. Just log + // the error here. + } else if (!process_error.ok()) { + OSP_LOG_ERROR << "error occurred while processing handles. error=" + << process_error; + } } }
diff --git a/platform/impl/platform_client_posix_unittest.cc b/platform/impl/platform_client_posix_unittest.cc new file mode 100644 index 0000000..20a1572 --- /dev/null +++ b/platform/impl/platform_client_posix_unittest.cc
@@ -0,0 +1,142 @@ +// Copyright 2025 The Chromium Authors +// Use of this source code is governed by a BSD-style license that can be +// found in the LICENSE file. + +#include "platform/impl/platform_client_posix.h" + +#include <chrono> +#include <memory> +#include <thread> +#include <utility> + +#include "gmock/gmock.h" +#include "gtest/gtest.h" +#include "platform/api/time.h" +#include "platform/impl/task_runner.h" +#include "platform/impl/tls_data_router_posix.h" +#include "platform/impl/udp_socket_reader_posix.h" +#include "platform/test/fake_clock.h" + +namespace openscreen { + +using ::testing::_; +using ::testing::Return; + +// Default timeout for operations in tests. +const Clock::duration kDefaultTestTimeout = std::chrono::milliseconds(10); + +// Mock for TaskRunner to inject and verify interactions. +class MockTaskRunnerImpl final : public TaskRunnerImpl { + public: + explicit MockTaskRunnerImpl(ClockNowFunctionPtr now_function) + : TaskRunnerImpl(now_function) {} + ~MockTaskRunnerImpl() override = default; + + MOCK_METHOD(void, PostPackagedTask, (TaskRunner::Task task), (override)); + MOCK_METHOD(void, + PostPackagedTaskWithDelay, + (TaskRunner::Task task, Clock::duration delay), + (override)); + MOCK_METHOD(bool, IsRunningOnTaskRunner, (), (override)); + MOCK_METHOD(void, RunUntilStopped, (), (override)); + MOCK_METHOD(void, RequestStopSoon, (), (override)); +}; + +class PlatformClientPosixTest : public ::testing::Test { + protected: + PlatformClientPosixTest() { + // Ensure no instance exists before each test. + if (PlatformClientPosix::GetInstance()) { + PlatformClientPosix::ShutDown(); + } + OSP_CHECK_EQ(PlatformClientPosix::GetInstance(), nullptr); + } + + ~PlatformClientPosixTest() { + // Ensure cleanup if a test fails to call ShutDown. + if (PlatformClientPosix::GetInstance()) { + PlatformClientPosix::ShutDown(); + } + OSP_CHECK_EQ(PlatformClientPosix::GetInstance(), nullptr); + } +}; + +TEST_F(PlatformClientPosixTest, CreateAndShutdown_DefaultTaskRunner) { + EXPECT_EQ(PlatformClientPosix::GetInstance(), nullptr); + + PlatformClientPosix::Create(kDefaultTestTimeout); + PlatformClientPosix* instance = PlatformClientPosix::GetInstance(); + EXPECT_NE(instance, nullptr); + EXPECT_NE(&instance->GetTaskRunner(), nullptr); + + PlatformClientPosix::ShutDown(); + EXPECT_EQ(PlatformClientPosix::GetInstance(), nullptr); +} + +TEST_F(PlatformClientPosixTest, CreateAndShutdown_ProvidedTaskRunner) { + EXPECT_EQ(PlatformClientPosix::GetInstance(), nullptr); + auto mock_task_runner = std::make_unique<MockTaskRunnerImpl>(&FakeClock::now); + + // When the PlatformClientPosix is shut down, it will request its TaskRunner + // to stop. + EXPECT_CALL(*mock_task_runner, RequestStopSoon()).Times(1); + EXPECT_CALL(*mock_task_runner, RunUntilStopped()).Times(0); + + const TaskRunner* mock_task_runner_ptr = mock_task_runner.get(); + PlatformClientPosix::Create(kDefaultTestTimeout, std::move(mock_task_runner)); + PlatformClientPosix* instance = PlatformClientPosix::GetInstance(); + EXPECT_NE(instance, nullptr); + + // Check that GetTaskRunner returns a reference to our mock. + // Since GetTaskRunner() returns TaskRunner&, and we passed TaskRunnerImpl*, + // we can check if the address matches. + EXPECT_EQ(&instance->GetTaskRunner(), mock_task_runner_ptr); + + PlatformClientPosix::ShutDown(); + EXPECT_EQ(PlatformClientPosix::GetInstance(), nullptr); +} + +TEST_F(PlatformClientPosixTest, + GetInstance_ReturnsNullBeforeCreateAndAfterShutdown) { + EXPECT_EQ(PlatformClientPosix::GetInstance(), nullptr); + + PlatformClientPosix::Create(kDefaultTestTimeout); + EXPECT_NE(PlatformClientPosix::GetInstance(), nullptr); + + PlatformClientPosix::ShutDown(); + EXPECT_EQ(PlatformClientPosix::GetInstance(), nullptr); +} + +TEST_F(PlatformClientPosixTest, ComponentInitialization_UdpSocketReader) { + PlatformClientPosix::Create(kDefaultTestTimeout); + PlatformClientPosix* instance = PlatformClientPosix::GetInstance(); + EXPECT_NE(instance, nullptr); + + // First call should initialize the UdpSocketReader. + UdpSocketReaderPosix* reader1 = instance->udp_socket_reader(); + EXPECT_NE(reader1, nullptr); + + // Second call should return the same instance. + UdpSocketReaderPosix* reader2 = instance->udp_socket_reader(); + EXPECT_EQ(reader1, reader2); + + PlatformClientPosix::ShutDown(); +} + +TEST_F(PlatformClientPosixTest, ComponentInitialization_TlsDataRouter) { + PlatformClientPosix::Create(kDefaultTestTimeout); + PlatformClientPosix* instance = PlatformClientPosix::GetInstance(); + EXPECT_NE(instance, nullptr); + + // First call should initialize the TlsDataRouter. + TlsDataRouterPosix* router1 = instance->tls_data_router(); + EXPECT_NE(router1, nullptr); + + // Second call should return the same instance. + TlsDataRouterPosix* router2 = instance->tls_data_router(); + EXPECT_EQ(router1, router2); + + PlatformClientPosix::ShutDown(); +} + +} // namespace openscreen
diff --git a/platform/impl/socket_handle_waiter.cc b/platform/impl/socket_handle_waiter.cc index c1426f1..096019b 100644 --- a/platform/impl/socket_handle_waiter.cc +++ b/platform/impl/socket_handle_waiter.cc
@@ -7,6 +7,7 @@ #include <algorithm> #include <atomic> +#include "platform/impl/socket_handle_posix.h" #include "util/osp_logging.h" #include "util/std_util.h" @@ -106,20 +107,32 @@ Error SocketHandleWaiter::ProcessHandles(Clock::duration timeout) { Clock::time_point start_time = now_function_(); - std::vector<ReadyHandle> handles; + std::vector<HandleWithFlags> handles; { std::lock_guard<std::mutex> lock(mutex_); handles_being_deleted_.clear(); handle_deletion_block_.notify_all(); handles.reserve(handle_mappings_.size()); - for (const auto& pair : handle_mappings_) { - handles.push_back({.handle = pair.first, .flags = pair.second.flags}); + for (auto& pair : handle_mappings_) { + uint32_t flags = pair.second.flags; + // Remove the write flag if there is no pending write. + if (flags & kWritable) { + const bool has_pending_write = + pair.second.subscriber->HasPendingWrite(pair.first); + if (!has_pending_write) { + flags &= ~kWritable; + } + } + handles.push_back(HandleWithFlags{.handle = pair.first, .flags = flags}); } } + if (handles.empty()) { + return Error::Code::kAgain; + } Clock::time_point current_time = now_function_(); Clock::duration remaining_timeout = timeout - (current_time - start_time); - ErrorOr<std::vector<ReadyHandle>> changed_handles = + ErrorOr<std::vector<HandleWithFlags>> changed_handles = AwaitSocketsReady(handles, remaining_timeout); std::vector<HandleWithSubscription> ready_handles;
diff --git a/platform/impl/socket_handle_waiter.h b/platform/impl/socket_handle_waiter.h index 1cb3470..a8016ac 100644 --- a/platform/impl/socket_handle_waiter.h +++ b/platform/impl/socket_handle_waiter.h
@@ -26,11 +26,16 @@ public: using SocketHandleRef = std::reference_wrapper<const SocketHandle>; + // Used to manage what types of events subscribers are subscribed to. enum Flags { kReadable = 1 << 0, - kWriteable = 1 << 1, + kWritable = 1 << 1, }; + // Common flag configurations. + static inline constexpr uint32_t kReadWriteFlags = + Flags::kReadable | Flags::kWritable; + class Subscriber { public: virtual ~Subscriber() = default; @@ -38,6 +43,15 @@ // Provides a socket handle to the subscriber which has data waiting to be // processed. virtual void ProcessReadyHandle(SocketHandleRef handle, uint32_t flags) = 0; + + // Method used to optimize event notifications. Generally speaking, + // sockets are ready for writing very often, causing the network event + // loop to be really busy -- a select() call may complete as frequently as + // every few nanoseconds -- so we really only want to be notified that a + // socket is ready for writing when we actually have something to write. + // + // NOTE: this is only used if the subscriber is subscribed to write events. + virtual bool HasPendingWrite(SocketHandleRef handle) = 0; }; explicit SocketHandleWaiter(ClockNowFunctionPtr now_function); @@ -67,11 +81,11 @@ OSP_DISALLOW_COPY_AND_ASSIGN(SocketHandleWaiter); // Gets all socket handles to process, checks them for readable data, and - // handles any changes that have occured. + // handles any changes that have occurred. Error ProcessHandles(Clock::duration timeout); protected: - struct ReadyHandle { + struct HandleWithFlags { SocketHandleRef handle; uint32_t flags; }; @@ -79,8 +93,15 @@ // Waits until data is available in one of the provided sockets or the // provided timeout has passed - whichever is first. If any sockets have data // available, they are returned. - virtual ErrorOr<std::vector<ReadyHandle>> AwaitSocketsReady( - const std::vector<ReadyHandle>& sockets, + // + // NOTE: The handle `flags` are checked against the subscriber's + // HasPendingWrite() method to ensure that the kWritable flag is only passed + // if there is a pending write before this method is called. The subscriber + // may be deleted while this method is being invoked, however the handle + // itself is guaranteed to not be deleted until the invocation of this method + // has been completed. + virtual ErrorOr<std::vector<HandleWithFlags>> AwaitSocketsReady( + const std::vector<HandleWithFlags>& sockets, const Clock::duration& timeout) = 0; private: @@ -92,7 +113,7 @@ }; struct HandleWithSubscription { - ReadyHandle ready_handle; + HandleWithFlags ready_handle; // Reference to the original subscription in the unordered map, so // we can keep track of when we updated this socket handle. SocketSubscription* subscription;
diff --git a/platform/impl/socket_handle_waiter_posix.cc b/platform/impl/socket_handle_waiter_posix.cc index 2fca5d1..1463e6f 100644 --- a/platform/impl/socket_handle_waiter_posix.cc +++ b/platform/impl/socket_handle_waiter_posix.cc
@@ -23,9 +23,9 @@ SocketHandleWaiterPosix::~SocketHandleWaiterPosix() = default; -ErrorOr<std::vector<SocketHandleWaiterPosix::ReadyHandle>> +ErrorOr<std::vector<SocketHandleWaiterPosix::HandleWithFlags>> SocketHandleWaiterPosix::AwaitSocketsReady( - const std::vector<SocketHandleWaiterPosix::ReadyHandle>& sockets, + const std::vector<SocketHandleWaiterPosix::HandleWithFlags>& sockets, const Clock::duration& timeout) { int max_fd = -1; fd_set read_handles{}; @@ -33,14 +33,18 @@ FD_ZERO(&read_handles); FD_ZERO(&write_handles); - for (const ReadyHandle& ready : sockets) { - if (ready.flags & Flags::kReadable) { - FD_SET(ready.handle.get().fd, &read_handles); + for (const HandleWithFlags& hwf : sockets) { + if (hwf.flags & Flags::kReadable) { + FD_SET(hwf.handle.get().fd, &read_handles); } - if (ready.flags & Flags::kWriteable) { - FD_SET(ready.handle.get().fd, &write_handles); + + // Only add the socket to the write_handles list if it is configured for + // write events and also has a pending write. This keeps us from polling + // select every few nanoseconds. + if (hwf.flags & Flags::kWritable) { + FD_SET(hwf.handle.get().fd, &write_handles); } - max_fd = std::max(max_fd, ready.handle.get().fd); + max_fd = std::max(max_fd, hwf.handle.get().fd); } if (max_fd < 0) { return Error::Code::kIOFailure; @@ -53,7 +57,7 @@ // level-triggered so incomplete reads/writes by the caller are fine and will // be picked up again on the next select() call. For more information, see: // http://man7.org/linux/man-pages/man2/select.2.html - int max_fd_to_watch = max_fd + 1; + const int max_fd_to_watch = max_fd + 1; const int rv = select(max_fd_to_watch, &read_handles, &write_handles, nullptr, &tv); if (rv == -1) { @@ -65,20 +69,19 @@ return Error::Code::kAgain; } - std::vector<ReadyHandle> changed_handles; - for (const ReadyHandle& ready : sockets) { + std::vector<HandleWithFlags> changed_handles; + for (const HandleWithFlags& hwf : sockets) { uint32_t flags = 0; - if (FD_ISSET(ready.handle.get().fd, &read_handles)) { + if (FD_ISSET(hwf.handle.get().fd, &read_handles)) { flags |= Flags::kReadable; } - if (FD_ISSET(ready.handle.get().fd, &write_handles)) { - flags |= Flags::kWriteable; + if (FD_ISSET(hwf.handle.get().fd, &write_handles)) { + flags |= Flags::kWritable; } if (flags) { - changed_handles.push_back({ready.handle, flags}); + changed_handles.push_back({hwf.handle, flags}); } } - return changed_handles; }
diff --git a/platform/impl/socket_handle_waiter_posix.h b/platform/impl/socket_handle_waiter_posix.h index 97d52d9..43e3cf1 100644 --- a/platform/impl/socket_handle_waiter_posix.h +++ b/platform/impl/socket_handle_waiter_posix.h
@@ -18,6 +18,7 @@ class SocketHandleWaiterPosix : public SocketHandleWaiter { public: using SocketHandleRef = SocketHandleWaiter::SocketHandleRef; + using HandleWithFlags = SocketHandleWaiter::HandleWithFlags; explicit SocketHandleWaiterPosix(ClockNowFunctionPtr now_function); ~SocketHandleWaiterPosix() override; @@ -30,10 +31,8 @@ void RequestStopSoon(); protected: - using SocketHandleWaiter::ReadyHandle; - - ErrorOr<std::vector<ReadyHandle>> AwaitSocketsReady( - const std::vector<ReadyHandle>& sockets, + ErrorOr<std::vector<HandleWithFlags>> AwaitSocketsReady( + const std::vector<HandleWithFlags>& sockets, const Clock::duration& timeout) override; private:
diff --git a/platform/impl/socket_handle_waiter_posix_unittest.cc b/platform/impl/socket_handle_waiter_posix_unittest.cc index 3e4e3ad..d7279ab 100644 --- a/platform/impl/socket_handle_waiter_posix_unittest.cc +++ b/platform/impl/socket_handle_waiter_posix_unittest.cc
@@ -5,7 +5,9 @@ #include "platform/impl/socket_handle_waiter_posix.h" #include <sys/socket.h> +#include <unistd.h> // For pipe() and close() +#include <cerrno> // For errno #include <chrono> #include <iostream> #include <thread> @@ -13,11 +15,14 @@ #include "gmock/gmock.h" #include "gtest/gtest.h" #include "platform/impl/socket_handle_posix.h" +#include "platform/impl/socket_handle_waiter.h" #include "platform/impl/timeval_posix.h" #include "platform/test/fake_clock.h" using ::testing::_; using ::testing::ByMove; +using ::testing::Gt; +using ::testing::IsEmpty; using ::testing::Return; namespace openscreen { @@ -26,27 +31,73 @@ class MockSubscriber : public SocketHandleWaiter::Subscriber { public: using SocketHandleRef = SocketHandleWaiter::SocketHandleRef; - MOCK_METHOD2(ProcessReadyHandle, void(SocketHandleRef, uint32_t)); + MOCK_METHOD(void, ProcessReadyHandle, (SocketHandleRef, uint32_t)); + MOCK_METHOD(bool, HasPendingWrite, (SocketHandleRef)); }; class TestingSocketHandleWaiter : public SocketHandleWaiter { public: + // These protected fields need to be public for testing below. using SocketHandleRef = SocketHandleWaiter::SocketHandleRef; - using ReadyHandle = SocketHandleWaiter::ReadyHandle; + using HandleWithFlags = SocketHandleWaiter::HandleWithFlags; TestingSocketHandleWaiter() : SocketHandleWaiter(&FakeClock::now) {} - MOCK_METHOD2( - AwaitSocketsReady, - ErrorOr<std::vector<ReadyHandle>>(const std::vector<ReadyHandle>&, - const Clock::duration&)); + MOCK_METHOD(ErrorOr<std::vector<HandleWithFlags>>, + AwaitSocketsReady, + (const std::vector<HandleWithFlags>&, const Clock::duration&), + (override)); FakeClock fake_clock{Clock::time_point{Clock::duration{1234567}}}; }; } // namespace -TEST(SocketHandleWaiterTest, BubblesUpAwaitSocketsReadyErrors) { +// Test fixture for tests that need an instance of SocketHandleWaiterPosix with +// an actual pipe. +class SocketHandleWaiterPosixInstanceTest : public ::testing::Test { + protected: + SocketHandleWaiterPosixInstanceTest() + : clock_(Clock::time_point{Clock::duration{1234567}}), + waiter_(&clock_.now) {} + + void TearDown() override { + // Clean up any FDs created in tests if they weren't closed properly. + // This is more of a safeguard. + for (int fd : fds_to_close_) { + close(fd); + } + fds_to_close_.clear(); + } + + // Helper to create a pipe and register fds for cleanup. + void CreatePipe(int pipe_fds[2]) { + ASSERT_NE(-1, pipe(pipe_fds)) + << "Failed to create pipe: " << strerror(errno); + fds_to_close_.push_back(pipe_fds[0]); + fds_to_close_.push_back(pipe_fds[1]); + } + + void ClosePipe(int pipe_fds[2]) { + auto remove_fd = [this](int fd_to_remove) { + fds_to_close_.erase( + std::remove(fds_to_close_.begin(), fds_to_close_.end(), fd_to_remove), + fds_to_close_.end()); + }; + + close(pipe_fds[0]); + remove_fd(pipe_fds[0]); + close(pipe_fds[1]); + remove_fd(pipe_fds[1]); + } + + FakeClock clock_; + SocketHandleWaiterPosix waiter_; // The actual class under test + MockSubscriber subscriber_; + std::vector<int> fds_to_close_; +}; + +TEST(SocketHandleWaiterBaseTest, BubblesUpAwaitSocketsReadyErrors) { MockSubscriber subscriber; TestingSocketHandleWaiter waiter; SocketHandle handle0(0); @@ -55,12 +106,13 @@ const SocketHandle& handle0_ref = handle0; const SocketHandle& handle1_ref = handle1; const SocketHandle& handle2_ref = handle2; - constexpr uint32_t rw_flags = SocketHandleWaiter::Flags::kReadable | - SocketHandleWaiter::Flags::kWriteable; - waiter.Subscribe(&subscriber, std::cref(handle0_ref), rw_flags); - waiter.Subscribe(&subscriber, std::cref(handle1_ref), rw_flags); - waiter.Subscribe(&subscriber, std::cref(handle2_ref), rw_flags); + waiter.Subscribe(&subscriber, std::cref(handle0_ref), + SocketHandleWaiter::kReadWriteFlags); + waiter.Subscribe(&subscriber, std::cref(handle1_ref), + SocketHandleWaiter::kReadWriteFlags); + waiter.Subscribe(&subscriber, std::cref(handle2_ref), + SocketHandleWaiter::kReadWriteFlags); Error::Code response = Error::Code::kAgain; EXPECT_CALL(subscriber, ProcessReadyHandle(_, _)).Times(0); EXPECT_CALL(waiter, AwaitSocketsReady(_, _)) @@ -68,7 +120,7 @@ waiter.ProcessHandles(Clock::duration{0}); } -TEST(SocketHandleWaiterTest, WatchedSocketsReturnedToCorrectSubscribers) { +TEST(SocketHandleWaiterBaseTest, WatchedSocketsReturnedToCorrectSubscribers) { MockSubscriber subscriber; MockSubscriber subscriber2; TestingSocketHandleWaiter waiter; @@ -81,32 +133,265 @@ const SocketHandle& handle2_ref = handle2; const SocketHandle& handle3_ref = handle3; - constexpr uint32_t r_flags = SocketHandleWaiter::Flags::kReadable; - constexpr uint32_t w_flags = SocketHandleWaiter::Flags::kWriteable; - constexpr uint32_t rw_flags = SocketHandleWaiter::Flags::kReadable | - SocketHandleWaiter::Flags::kWriteable; + waiter.Subscribe(&subscriber, std::cref(handle0_ref), + SocketHandleWaiter::kReadWriteFlags); + waiter.Subscribe(&subscriber, std::cref(handle2_ref), + SocketHandleWaiter::kReadWriteFlags); + waiter.Subscribe(&subscriber2, std::cref(handle1_ref), + SocketHandleWaiter::kReadWriteFlags); + waiter.Subscribe(&subscriber2, std::cref(handle3_ref), + SocketHandleWaiter::kReadWriteFlags); - waiter.Subscribe(&subscriber, std::cref(handle0_ref), rw_flags); - waiter.Subscribe(&subscriber, std::cref(handle2_ref), rw_flags); - waiter.Subscribe(&subscriber2, std::cref(handle1_ref), rw_flags); - waiter.Subscribe(&subscriber2, std::cref(handle3_ref), rw_flags); - - EXPECT_CALL(subscriber, ProcessReadyHandle(std::cref(handle0_ref), r_flags)) + EXPECT_CALL(subscriber, ProcessReadyHandle(std::cref(handle0_ref), + SocketHandleWaiter::kReadable)) .Times(1); - EXPECT_CALL(subscriber, ProcessReadyHandle(std::cref(handle2_ref), w_flags)) + EXPECT_CALL(subscriber, ProcessReadyHandle(std::cref(handle2_ref), + SocketHandleWaiter::kWritable)) .Times(1); - EXPECT_CALL(subscriber2, ProcessReadyHandle(std::cref(handle1_ref), r_flags)) + EXPECT_CALL(subscriber2, ProcessReadyHandle(std::cref(handle1_ref), + SocketHandleWaiter::kReadable)) .Times(1); - EXPECT_CALL(subscriber2, ProcessReadyHandle(std::cref(handle3_ref), rw_flags)) + EXPECT_CALL(subscriber2, + ProcessReadyHandle(std::cref(handle3_ref), + SocketHandleWaiter::kReadWriteFlags)) .Times(1); EXPECT_CALL(waiter, AwaitSocketsReady(_, _)) .WillOnce( - Return(ByMove(std::vector<TestingSocketHandleWaiter::ReadyHandle>{ - {std::cref(handle0_ref), r_flags}, - {std::cref(handle1_ref), r_flags}, - {std::cref(handle2_ref), w_flags}, - {std::cref(handle3_ref), rw_flags}}))); + Return(ByMove(std::vector<TestingSocketHandleWaiter::HandleWithFlags>{ + {std::cref(handle0_ref), SocketHandleWaiter::kReadable}, + {std::cref(handle1_ref), SocketHandleWaiter::kReadable}, + {std::cref(handle2_ref), SocketHandleWaiter::kWritable}, + {std::cref(handle3_ref), SocketHandleWaiter::kReadWriteFlags}}))); waiter.ProcessHandles(Clock::duration{0}); } +TEST(SocketHandleWaiterBaseTest, HandlesNoSubscriptions) { + TestingSocketHandleWaiter waiter; + + EXPECT_CALL(waiter, AwaitSocketsReady(_, _)).Times(0); + Error result = waiter.ProcessHandles(Clock::duration{0}); + EXPECT_EQ(result.code(), Error::Code::kAgain); +} + +TEST(SocketHandleWaiterBaseTest, UnsubscribeRemovesHandle) { + MockSubscriber subscriber; + TestingSocketHandleWaiter waiter; + SocketHandle handle(123); + + waiter.Subscribe(&subscriber, handle, SocketHandleWaiter::Flags::kReadable); + waiter.Unsubscribe(&subscriber, handle); + + EXPECT_CALL(waiter, AwaitSocketsReady(_, _)).Times(0); + Error result = waiter.ProcessHandles(Clock::duration{0}); + EXPECT_EQ(result.code(), Error::Code::kAgain); +} + +TEST(SocketHandleWaiterBaseTest, + UnsubscribeAllRemovesOnlySubscriptionsForProvidedSubscriber) { + MockSubscriber subscriber1; + MockSubscriber subscriber2; + TestingSocketHandleWaiter waiter; + SocketHandle handle0(0); + SocketHandle handle1(1); + SocketHandle handle2(2); + + waiter.Subscribe(&subscriber1, handle0, SocketHandleWaiter::kReadable); + waiter.Subscribe(&subscriber1, handle1, SocketHandleWaiter::kReadable); + waiter.Subscribe(&subscriber2, handle2, SocketHandleWaiter::kReadable); + + waiter.UnsubscribeAll(&subscriber1); + + EXPECT_CALL(subscriber1, ProcessReadyHandle(_, _)).Times(0); + EXPECT_CALL(subscriber2, ProcessReadyHandle(std::cref(handle2), + SocketHandleWaiter::kReadable)) + .Times(1); + + EXPECT_CALL( + waiter, + AwaitSocketsReady( + testing::Truly( + [&](const std::vector<TestingSocketHandleWaiter::HandleWithFlags>& + handle_list) { + return handle_list.size() == 1 && + handle_list[0].handle.get() == handle2 && + handle_list[0].flags == SocketHandleWaiter::kReadable; + }), + _)) + .WillOnce( + Return(ByMove(std::vector<TestingSocketHandleWaiter::HandleWithFlags>{ + {handle2, SocketHandleWaiter::kReadable}}))); + + waiter.ProcessHandles(Clock::duration{0}); +} + +TEST(SocketHandleWaiterBaseTest, + AwaitSocketsReadyCalledWithCorrectSubscribedFlags) { + MockSubscriber subscriber; + TestingSocketHandleWaiter waiter; + SocketHandle handle0(0); + constexpr uint32_t subscribed_flags = SocketHandleWaiter::Flags::kReadable; + // AwaitSocketsReady might report more flags if the underlying mechanism + // does. + constexpr uint32_t reported_flags = SocketHandleWaiter::Flags::kReadable | + SocketHandleWaiter::Flags::kWritable; + + waiter.Subscribe(&subscriber, handle0, subscribed_flags); + + EXPECT_CALL( + waiter, + AwaitSocketsReady( + testing::Truly( + [&](const std::vector<TestingSocketHandleWaiter::HandleWithFlags>& + handle_list) { + return handle_list.size() == 1 && + handle_list[0].handle.get() == handle0 && + handle_list[0].flags == subscribed_flags; + }), + _)) + .WillOnce( + Return(ByMove(std::vector<TestingSocketHandleWaiter::HandleWithFlags>{ + {handle0, reported_flags}}))); + + // The subscriber receives the flags reported by AwaitSocketsReady. + EXPECT_CALL(subscriber, + ProcessReadyHandle(std::cref(handle0), reported_flags)) + .Times(1); + + waiter.ProcessHandles(Clock::duration{0}); +} + +TEST(SocketHandleWaiterBaseTest, SubsequentSubscribeForSameHandleIsIgnored) { + MockSubscriber subscriber1; + MockSubscriber subscriber2; + TestingSocketHandleWaiter waiter; + SocketHandle handle0(0); + + waiter.Subscribe(&subscriber1, handle0, SocketHandleWaiter::Flags::kReadable); + // This second subscribe for the same handle should be ignored. + waiter.Subscribe(&subscriber2, handle0, SocketHandleWaiter::Flags::kWritable); + + EXPECT_CALL( + waiter, + AwaitSocketsReady( + testing::Truly( + [&](const std::vector<TestingSocketHandleWaiter::HandleWithFlags>& + handle_list) { + return handle_list.size() == 1 && + handle_list[0].handle.get() == handle0 && + handle_list[0].flags == + SocketHandleWaiter::Flags::kReadable; + }), + _)) + .WillOnce(Return(Error::Code::kAgain)); + + EXPECT_CALL(subscriber1, ProcessReadyHandle(_, _)).Times(0); + EXPECT_CALL(subscriber2, ProcessReadyHandle(_, _)).Times(0); + + waiter.ProcessHandles(Clock::duration{0}); +} + +TEST(SocketHandleWaiterBaseTest, OnHandleDeletionRemovesHandle) { + MockSubscriber subscriber; + TestingSocketHandleWaiter waiter; + SocketHandle handle0(0); + + waiter.Subscribe(&subscriber, handle0, SocketHandleWaiter::Flags::kReadable); + waiter.OnHandleDeletion(&subscriber, handle0, true); + + EXPECT_CALL(waiter, AwaitSocketsReady(IsEmpty(), _)).Times(0); + const Error result = waiter.ProcessHandles(Clock::duration{0}); + EXPECT_EQ(result.code(), Error::Code::kAgain); +} + +TEST(SocketHandleWaiterBaseTest, + ProcessHandlesIgnoresUnmappedHandlesFromAwaitSocketsReady) { + MockSubscriber subscriber; + TestingSocketHandleWaiter waiter; + SocketHandle handle0(0); // Subscribed + SocketHandle handle1(1); // Not subscribed, but returned by AwaitSocketsReady + + waiter.Subscribe(&subscriber, handle0, SocketHandleWaiter::kReadable); + + EXPECT_CALL(waiter, AwaitSocketsReady(_, _)) + .WillOnce( + Return(ByMove(std::vector<TestingSocketHandleWaiter::HandleWithFlags>{ + {handle0, SocketHandleWaiter::kReadable}, // Expected + {handle1, + SocketHandleWaiter::kReadable} // Unexpected by subscription + }))); + + EXPECT_CALL(subscriber, ProcessReadyHandle(std::cref(handle0), + SocketHandleWaiter::kReadable)) + .Times(1); + EXPECT_CALL(subscriber, ProcessReadyHandle(std::cref(handle1), _)) + .Times(0); // Should be ignored + + waiter.ProcessHandles(Clock::duration{0}); +} + +TEST_F(SocketHandleWaiterPosixInstanceTest, WriteNotProcessedIfNoPendingWrite) { + int pipe_fds[2]; + CreatePipe(pipe_fds); + SocketHandle write_end_handle(pipe_fds[1]); + + // Subscribe for write events on the write-end of the pipe. + waiter_.Subscribe(&subscriber_, std::cref(write_end_handle), + SocketHandleWaiter::Flags::kWritable); + + // Simulate that the subscriber has no data to write. + EXPECT_CALL(subscriber_, HasPendingWrite(std::cref(write_end_handle))) + .WillRepeatedly(Return(false)); + + // Expect ProcessReadyHandle NOT to be called with the kWritable flag. + // Since we only subscribed for kWritable, and HasPendingWrite is false, + // it should not be called at all for this handle's write event. + EXPECT_CALL(subscriber_, ProcessReadyHandle(std::cref(write_end_handle), _)) + .Times(0); + + const Error result = waiter_.ProcessHandles(std::chrono::milliseconds(10)); + + // If no FDs were ready (or ready but filtered out), ProcessHandles returns + // kAgain. + EXPECT_TRUE(result.ok() || result.code() == Error::Code::kAgain) + << "ProcessHandles returned: " << result; + + waiter_.Unsubscribe(&subscriber_, std::cref(write_end_handle)); + ClosePipe(pipe_fds); +} + +TEST_F(SocketHandleWaiterPosixInstanceTest, + ReadProcessedWriteIgnoredIfNoPendingWrite) { + int pipe_fds[2]; + CreatePipe(pipe_fds); + SocketHandle read_end_handle(pipe_fds[0]); + + // Subscribe for read and write events on the read-end of the pipe. + waiter_.Subscribe(&subscriber_, std::cref(read_end_handle), + SocketHandleWaiter::Flags::kReadable | + SocketHandleWaiter::Flags::kWritable); + + // Simulate that the subscriber has no data to write. + EXPECT_CALL(subscriber_, HasPendingWrite(std::cref(read_end_handle))) + .WillRepeatedly(Return(false)); + + // Make the read-end readable. + constexpr const char kTestBuf[] = "test"; + ASSERT_THAT(write(pipe_fds[1], kTestBuf, sizeof(kTestBuf) - 1), + Gt(ssize_t{0})); + + // Expect ProcessReadyHandle to be called only with kReadable. + EXPECT_CALL(subscriber_, + ProcessReadyHandle(std::cref(read_end_handle), + SocketHandleWaiter::Flags::kReadable)) + .Times(1); + + const Error result = waiter_.ProcessHandles(std::chrono::milliseconds(10)); + EXPECT_TRUE(result.ok()) << "ProcessHandles returned: " << result; + + std::array<uint8_t, 16> drain; + read(pipe_fds[0], drain.data(), drain.size()); + waiter_.Unsubscribe(&subscriber_, std::cref(read_end_handle)); + ClosePipe(pipe_fds); +} + } // namespace openscreen
diff --git a/platform/impl/task_runner.h b/platform/impl/task_runner.h index 470fd2e..4b061b3 100644 --- a/platform/impl/task_runner.h +++ b/platform/impl/task_runner.h
@@ -20,7 +20,7 @@ namespace openscreen { -class TaskRunnerImpl final : public TaskRunner { +class TaskRunnerImpl : public TaskRunner { public: using Task = TaskRunner::Task; @@ -50,20 +50,20 @@ Clock::duration waiter_timeout = std::chrono::milliseconds(100)); // TaskRunner overrides - ~TaskRunnerImpl() final; - void PostPackagedTask(Task task) final; - void PostPackagedTaskWithDelay(Task task, Clock::duration delay) final; - bool IsRunningOnTaskRunner() final; + ~TaskRunnerImpl(); + void PostPackagedTask(Task task) override; + void PostPackagedTaskWithDelay(Task task, Clock::duration delay) override; + bool IsRunningOnTaskRunner() override; // Blocks the current thread, executing tasks from the queue with the desired // timing; and does not return until some time after RequestStopSoon() is // called. - void RunUntilStopped(); + virtual void RunUntilStopped(); // Blocks the current thread, executing tasks from the queue with the desired // timing; and does not return until some time after the current process is // signaled with SIGINT or SIGTERM, or after RequestStopSoon() is called. - void RunUntilSignaled(); + virtual void RunUntilSignaled(); // Thread-safe method for requesting the TaskRunner to stop running after all // non-delayed tasks in the queue have run. This behavior allows final @@ -71,7 +71,7 @@ // // If any non-delayed tasks post additional non-delayed tasks, those will be // run as well before returning. - void RequestStopSoon(); + virtual void RequestStopSoon(); private: #if defined(ENABLE_TRACE_LOGGING)
diff --git a/platform/impl/tls_connection_posix.h b/platform/impl/tls_connection_posix.h index 2540ce4..07cbf67 100644 --- a/platform/impl/tls_connection_posix.h +++ b/platform/impl/tls_connection_posix.h
@@ -40,6 +40,8 @@ // automatically by TlsConnectionFactoryPosix after the handshake completes. void RegisterConnectionWithDataRouter(PlatformClientPosix* platform_client); + bool HasPendingWrite() const { return !buffer_.GetReadableRegion().empty(); } + const SocketHandle& socket_handle() const { return socket_->socket_handle(); } protected:
diff --git a/platform/impl/tls_data_router_posix.cc b/platform/impl/tls_data_router_posix.cc index a94d3cf..70ea83e 100644 --- a/platform/impl/tls_data_router_posix.cc +++ b/platform/impl/tls_data_router_posix.cc
@@ -7,6 +7,7 @@ #include <memory> #include <utility> +#include "platform/impl/socket_handle_waiter.h" #include "platform/impl/stream_socket_posix.h" #include "platform/impl/tls_connection_posix.h" #include "util/osp_logging.h" @@ -32,8 +33,7 @@ // We care about both read and write events waiter_->Subscribe(this, connection->socket_handle(), - SocketHandleWaiter::Flags::kReadable | - SocketHandleWaiter::Flags::kWriteable); + SocketHandleWaiter::kReadWriteFlags); } void TlsDataRouterPosix::DeregisterConnection(TlsConnectionPosix* connection) { @@ -65,8 +65,7 @@ // We care about both read and write events waiter_->Subscribe(this, socket_ptr->socket_handle(), - SocketHandleWaiter::Flags::kReadable | - SocketHandleWaiter::Flags::kWriteable); + SocketHandleWaiter::kReadWriteFlags); } void TlsDataRouterPosix::DeregisterAcceptObserver(SocketObserver* observer) { @@ -112,7 +111,7 @@ if (flags & SocketHandleWaiter::Flags::kReadable) { connection->TryReceiveMessage(); } - if (flags & SocketHandleWaiter::Flags::kWriteable) { + if (flags & SocketHandleWaiter::Flags::kWritable) { connection->SendAvailableBytes(); } return; @@ -121,6 +120,23 @@ } } +bool TlsDataRouterPosix::HasPendingWrite( + SocketHandleWaiter::SocketHandleRef handle) { + { + std::lock_guard<std::mutex> lock(connections_mutex_); + for (TlsConnectionPosix* connection : connections_) { + if (connection->socket_handle() == handle) { + return connection->HasPendingWrite(); + } + } + } + + // If we don't have the socket in the connections list, it's either + // an accept socket or a socket in the process of being destroyed. In either + // case, this is not an error and we can safely report no pending writes. + return false; +} + bool TlsDataRouterPosix::HasTimedOut(Clock::time_point start_time, Clock::duration timeout) { return now_function_() - start_time > timeout;
diff --git a/platform/impl/tls_data_router_posix.h b/platform/impl/tls_data_router_posix.h index e84b8b3..2ebf9af 100644 --- a/platform/impl/tls_data_router_posix.h +++ b/platform/impl/tls_data_router_posix.h
@@ -69,6 +69,7 @@ // SocketHandleWaiter::Subscriber overrides. void ProcessReadyHandle(SocketHandleWaiter::SocketHandleRef handle, uint32_t flags) override; + bool HasPendingWrite(SocketHandleWaiter::SocketHandleRef handle) override; OSP_DISALLOW_COPY_AND_ASSIGN(TlsDataRouterPosix);
diff --git a/platform/impl/tls_data_router_posix_unittest.cc b/platform/impl/tls_data_router_posix_unittest.cc index 41605a8..d32577a 100644 --- a/platform/impl/tls_data_router_posix_unittest.cc +++ b/platform/impl/tls_data_router_posix_unittest.cc
@@ -20,14 +20,14 @@ class MockNetworkWaiter final : public SocketHandleWaiter { public: - using ReadyHandle = SocketHandleWaiter::ReadyHandle; + using HandleWithFlags = SocketHandleWaiter::HandleWithFlags; MockNetworkWaiter() : SocketHandleWaiter(&FakeClock::now) {} - MOCK_METHOD2( - AwaitSocketsReady, - ErrorOr<std::vector<ReadyHandle>>(const std::vector<ReadyHandle>&, - const Clock::duration&)); + MOCK_METHOD(ErrorOr<std::vector<HandleWithFlags>>, + AwaitSocketsReady, + (const std::vector<HandleWithFlags>&, const Clock::duration&), + (override)); }; class MockSocket : public StreamSocketPosix { @@ -130,7 +130,7 @@ EXPECT_CALL(connection3, SendAvailableBytes()).Times(0); EXPECT_CALL(connection3, TryReceiveMessage()).Times(0); network_manager()->ProcessReadyHandle(connection2.socket_handle(), - SocketHandleWaiter::Flags::kWriteable); + SocketHandleWaiter::Flags::kWritable); } TEST_F(TlsNetworkingManagerPosixTest, DeregisterTlsConnection) {
diff --git a/platform/impl/tls_write_buffer.cc b/platform/impl/tls_write_buffer.cc index 7b0defb..ecd811d 100644 --- a/platform/impl/tls_write_buffer.cc +++ b/platform/impl/tls_write_buffer.cc
@@ -52,7 +52,7 @@ return true; } -ByteView TlsWriteBuffer::GetReadableRegion() { +ByteView TlsWriteBuffer::GetReadableRegion() const { const size_t current_read_bytes = bytes_read_so_far_.load(std::memory_order_relaxed); const size_t currently_written_bytes =
diff --git a/platform/impl/tls_write_buffer.h b/platform/impl/tls_write_buffer.h index 6a581c2..7765c87 100644 --- a/platform/impl/tls_write_buffer.h +++ b/platform/impl/tls_write_buffer.h
@@ -29,7 +29,7 @@ // Returns a subset of the readable region of data. At time of reading, more // data may be available for reading than what is represented in this Span. - ByteView GetReadableRegion(); + ByteView GetReadableRegion() const; // Marks the provided number of bytes as consumed by the consumer thread. void Consume(size_t byte_count);
diff --git a/platform/impl/udp_socket_reader_posix.cc b/platform/impl/udp_socket_reader_posix.cc index 1048068..0808a68 100644 --- a/platform/impl/udp_socket_reader_posix.cc +++ b/platform/impl/udp_socket_reader_posix.cc
@@ -23,19 +23,22 @@ void UdpSocketReaderPosix::ProcessReadyHandle(SocketHandleRef handle, uint32_t flags) { - if (flags & SocketHandleWaiter::Flags::kReadable) { - std::lock_guard<std::mutex> lock(mutex_); - // NOTE: Because sockets_ is expected to remain small, the performance here - // is better than using an unordered_set. - for (UdpSocketPosix* socket : sockets_) { - if (socket->GetHandle() == handle) { - socket->ReceiveMessage(); - break; - } + OSP_CHECK(flags & SocketHandleWaiter::Flags::kReadable); + std::lock_guard<std::mutex> lock(mutex_); + // NOTE: Because sockets_ is expected to remain small, the performance here + // is better than using an unordered_set. + for (UdpSocketPosix* socket : sockets_) { + if (socket->GetHandle() == handle) { + socket->ReceiveMessage(); + break; } } } +bool UdpSocketReaderPosix::HasPendingWrite(SocketHandleRef handle) { + OSP_NOTREACHED(); +} + void UdpSocketReaderPosix::OnCreate(UdpSocket* socket) { UdpSocketPosix* read_socket = static_cast<UdpSocketPosix*>(socket); { @@ -44,7 +47,7 @@ } // We only care about read events. waiter_.Subscribe(this, std::cref(read_socket->GetHandle()), - SocketHandleWaiter::Flags::kReadable); + SocketHandleWaiter::kReadable); } void UdpSocketReaderPosix::OnDestroy(UdpSocket* socket) {
diff --git a/platform/impl/udp_socket_reader_posix.h b/platform/impl/udp_socket_reader_posix.h index 223bf4b..c2ad317 100644 --- a/platform/impl/udp_socket_reader_posix.h +++ b/platform/impl/udp_socket_reader_posix.h
@@ -48,6 +48,9 @@ // SocketHandleWaiter::Subscriber overrides. void ProcessReadyHandle(SocketHandleRef handle, uint32_t flags) override; + // NOTE: we don't subscribe to write events from the socket handle waiter. + bool HasPendingWrite(SocketHandleRef handle) override; + OSP_DISALLOW_COPY_AND_ASSIGN(UdpSocketReaderPosix); protected:
diff --git a/platform/impl/udp_socket_reader_posix_unittest.cc b/platform/impl/udp_socket_reader_posix_unittest.cc index 6e3a60f..69a59f7 100644 --- a/platform/impl/udp_socket_reader_posix_unittest.cc +++ b/platform/impl/udp_socket_reader_posix_unittest.cc
@@ -46,15 +46,15 @@ // Mock event waiter class MockNetworkWaiter final : public SocketHandleWaiter { public: - using ReadyHandle = SocketHandleWaiter::ReadyHandle; + using HandleWithFlags = SocketHandleWaiter::HandleWithFlags; MockNetworkWaiter() : SocketHandleWaiter(&FakeClock::now) {} ~MockNetworkWaiter() override = default; - MOCK_METHOD2( - AwaitSocketsReady, - ErrorOr<std::vector<ReadyHandle>>(const std::vector<ReadyHandle>&, - const Clock::duration&)); + MOCK_METHOD(ErrorOr<std::vector<HandleWithFlags>>, + AwaitSocketsReady, + (const std::vector<HandleWithFlags>&, const Clock::duration&), + (override)); FakeClock fake_clock{Clock::time_point{Clock::duration{1234567}}}; };