diff --git a/src/platform/windows/audio.cpp b/src/platform/windows/audio.cpp index 9ace2a6b338..bc9d8acef2f 100644 --- a/src/platform/windows/audio.cpp +++ b/src/platform/windows/audio.cpp @@ -285,15 +285,51 @@ namespace platf::audio { */ class co_init_t: public deinit_t { public: - co_init_t() { - CoInitializeEx(nullptr, COINIT_MULTITHREADED | COINIT_SPEED_OVER_MEMORY); - } + co_init_t() = default; ~co_init_t() override { - CoUninitialize(); + if (SUCCEEDED(status_)) { + CoUninitialize(); + } + } + + /** + * @brief Return the result from `CoInitializeEx()`. + * @return COM initialization status for the current thread. + */ + [[nodiscard]] HRESULT status() const { + return status_; } + + private: + HRESULT status_ {CoInitializeEx(nullptr, COINIT_MULTITHREADED | COINIT_SPEED_OVER_MEMORY)}; }; + /** + * @brief Determine whether a COM initialization result leaves COM available to the caller. + * @param status Result returned by `CoInitializeEx()`. + * @return `true` for a successful initialization or an already initialized apartment. + */ + [[nodiscard]] bool is_com_available(HRESULT status) { + return SUCCEEDED(status) || status == RPC_E_CHANGED_MODE; + } + + /** + * @brief Ensure COM remains initialized for the lifetime of the current audio thread. + * @return `true` when COM is available on the current thread. + */ + [[nodiscard]] bool initialize_com_for_audio_thread() { + thread_local co_init_t co_init; + const auto status = co_init.status(); + + if (!is_com_available(status)) { + BOOST_LOG(error) << "Couldn't initialize COM for audio capture: [0x"sv << util::hex(status).to_string_view() << ']'; + return false; + } + + return true; + } + /** * @brief RAII wrapper that initializes and clears a Windows PROPVARIANT. */ @@ -604,6 +640,10 @@ namespace platf::audio { * @return 0 on success; nonzero or negative platform status on failure. */ int init(std::uint32_t sample_rate, std::uint32_t frame_size, std::uint32_t channels_out, bool continuous, device_t capture_device) { + if (!initialize_com_for_audio_thread()) { + return -1; + } + audio_event.reset(CreateEventA(nullptr, FALSE, FALSE, nullptr)); if (!audio_event) { BOOST_LOG(error) << "Couldn't create Event handle"sv; @@ -1385,6 +1425,10 @@ namespace platf::audio { * @return 0 on success; nonzero or negative platform status on failure. */ int init() { + if (!initialize_com_for_audio_thread()) { + return -1; + } + auto status = CoCreateInstance( CLSID_CPolicyConfigClient, nullptr, @@ -1428,6 +1472,15 @@ namespace platf::audio { #ifdef SUNSHINE_TESTS namespace tests { + /** + * @brief Evaluate a COM initialization status through the production acceptance policy. + * @param status Result returned by `CoInitializeEx()`. + * @return `true` when COM is available to the caller. + */ + bool com_is_available(HRESULT status) { + return is_com_available(status); + } + /** * @brief Resolve a sink through the production Windows endpoint lookup. * diff --git a/tests/unit/platform/windows/test_audio.cpp b/tests/unit/platform/windows/test_audio.cpp index 27cba090a0f..1c28f9c1d8f 100644 --- a/tests/unit/platform/windows/test_audio.cpp +++ b/tests/unit/platform/windows/test_audio.cpp @@ -18,6 +18,7 @@ #include "src/platform/common.h" namespace platf::audio::tests { + bool com_is_available(HRESULT status); bool sink_device_available(const std::string &sink, IMMDeviceEnumerator *device_enum); bool microphone_available(const std::string &assigned_sink, const std::string &configured_sink, IMMDeviceEnumerator *device_enum); bool capture_follows_default_device(IMMDeviceEnumerator *device_enum, IMMDevice *capture_device); @@ -243,6 +244,16 @@ namespace { }; } // namespace +TEST(WindowsAudioTest, AcceptsUsableComInitializationResults) { + EXPECT_TRUE(platf::audio::tests::com_is_available(S_OK)); + EXPECT_TRUE(platf::audio::tests::com_is_available(S_FALSE)); + EXPECT_TRUE(platf::audio::tests::com_is_available(RPC_E_CHANGED_MODE)); +} + +TEST(WindowsAudioTest, RejectsFailedComInitializationResult) { + EXPECT_FALSE(platf::audio::tests::com_is_available(E_FAIL)); +} + TEST(WindowsAudioTest, AssignedSinkTakesPriorityOverConfiguredSink) { fake_device_enumerator_t enumerator {L"assigned-id"};