Use an RAII wrapper to manage COM object references

This commit is contained in:
Chris Robinson
2020-09-05 12:32:41 -07:00
parent 88f7617807
commit 73f5331305
+111 -87
View File
@@ -129,6 +129,69 @@ inline ALuint RefTime2Samples(const ReferenceTime &val, ALuint srate)
}
template<typename T>
class ComPtr {
T *mPtr{nullptr};
public:
ComPtr() noexcept = default;
ComPtr(const ComPtr &rhs) : mPtr{rhs.mPtr} { if(mPtr) mPtr->AddRef(); }
ComPtr(ComPtr&& rhs) noexcept : mPtr{rhs.mPtr} { rhs.mPtr = nullptr; }
ComPtr(std::nullptr_t) noexcept { }
explicit ComPtr(T *ptr) noexcept : mPtr{ptr} { }
~ComPtr() { if(mPtr) mPtr->Release(); }
ComPtr& operator=(const ComPtr &rhs)
{
if(!rhs.mPtr)
{
if(mPtr)
mPtr->Release();
mPtr = nullptr;
}
else
{
rhs.mPtr->AddRef();
try {
if(mPtr)
mPtr->Release();
mPtr = rhs.mPtr;
}
catch(...) {
rhs.mPtr->Release();
throw;
}
}
return *this;
}
ComPtr& operator=(ComPtr&& rhs)
{
if(mPtr)
mPtr->Release();
mPtr = rhs.mPtr;
rhs.mPtr = nullptr;
return *this;
}
operator bool() const noexcept { return mPtr != nullptr; }
T& operator*() const noexcept { return *mPtr; }
T* operator->() const noexcept { return mPtr; }
T* get() const noexcept { return mPtr; }
T** getPtr() noexcept { return &mPtr; }
T* release() noexcept
{
T *ret{mPtr};
mPtr = nullptr;
return ret;
}
void swap(ComPtr &rhs) noexcept { std::swap(mPtr, rhs.mPtr); }
void swap(ComPtr&& rhs) noexcept { std::swap(mPtr, rhs.mPtr); }
};
class GuidPrinter {
char mMsg[64];
@@ -191,8 +254,8 @@ NameGUIDPair get_device_name_and_guid(IMMDevice *device)
std::string name{DEVNAME_HEAD};
std::string guid;
IPropertyStore *ps;
HRESULT hr = device->OpenPropertyStore(STGM_READ, &ps);
ComPtr<IPropertyStore> ps;
HRESULT hr = device->OpenPropertyStore(STGM_READ, ps.getPtr());
if(FAILED(hr))
{
WARN("OpenPropertyStore failed: 0x%08lx\n", hr);
@@ -229,15 +292,13 @@ NameGUIDPair get_device_name_and_guid(IMMDevice *device)
guid = "Unknown Device GUID";
}
ps->Release();
return std::make_pair(std::move(name), std::move(guid));
}
void get_device_formfactor(IMMDevice *device, EndpointFormFactor *formfactor)
{
IPropertyStore *ps;
HRESULT hr = device->OpenPropertyStore(STGM_READ, &ps);
ComPtr<IPropertyStore> ps;
HRESULT hr = device->OpenPropertyStore(STGM_READ, ps.getPtr());
if(FAILED(hr))
{
WARN("OpenPropertyStore failed: 0x%08lx\n", hr);
@@ -254,8 +315,6 @@ void get_device_formfactor(IMMDevice *device, EndpointFormFactor *formfactor)
*formfactor = UnknownFormFactor;
else
WARN("Unexpected PROPVARIANT type: 0x%04x\n", pvform->vt);
ps->Release();
}
@@ -288,7 +347,7 @@ WCHAR *get_device_id(IMMDevice *device)
{
WCHAR *devid;
HRESULT hr = device->GetId(&devid);
const HRESULT hr{device->GetId(&devid)};
if(FAILED(hr))
{
ERR("Failed to get device id: %lx\n", hr);
@@ -302,8 +361,8 @@ void probe_devices(IMMDeviceEnumerator *devenum, EDataFlow flowdir, al::vector<D
{
al::vector<DevMap>{}.swap(list);
IMMDeviceCollection *coll;
HRESULT hr{devenum->EnumAudioEndpoints(flowdir, DEVICE_STATE_ACTIVE, &coll)};
ComPtr<IMMDeviceCollection> coll;
HRESULT hr{devenum->EnumAudioEndpoints(flowdir, DEVICE_STATE_ACTIVE, coll.getPtr())};
if(FAILED(hr))
{
ERR("Failed to enumerate audio endpoints: 0x%08lx\n", hr);
@@ -315,34 +374,30 @@ void probe_devices(IMMDeviceEnumerator *devenum, EDataFlow flowdir, al::vector<D
if(SUCCEEDED(hr) && count > 0)
list.reserve(count);
IMMDevice *device{};
hr = devenum->GetDefaultAudioEndpoint(flowdir, eMultimedia, &device);
ComPtr<IMMDevice> device;
hr = devenum->GetDefaultAudioEndpoint(flowdir, eMultimedia, device.getPtr());
if(SUCCEEDED(hr))
{
if(WCHAR *devid{get_device_id(device)})
if(WCHAR *devid{get_device_id(device.get())})
{
add_device(device, devid, list);
add_device(device.get(), devid, list);
CoTaskMemFree(devid);
}
device->Release();
device = nullptr;
}
for(UINT i{0};i < count;++i)
{
hr = coll->Item(i, &device);
hr = coll->Item(i, device.getPtr());
if(FAILED(hr)) continue;
if(WCHAR *devid{get_device_id(device)})
if(WCHAR *devid{get_device_id(device.get())})
{
add_device(device, devid, list);
add_device(device.get(), devid, list);
CoTaskMemFree(devid);
}
device->Release();
device = nullptr;
}
coll->Release();
}
@@ -513,7 +568,7 @@ int WasapiProxy::messageHandler(std::promise<HRESULT> *promise)
{
TRACE("Starting message thread\n");
HRESULT cohr = CoInitializeEx(nullptr, COINIT_MULTITHREADED);
HRESULT cohr{CoInitializeEx(nullptr, COINIT_MULTITHREADED)};
if(FAILED(cohr))
{
WARN("Failed to initialize COM: 0x%08lx\n", cohr);
@@ -531,9 +586,7 @@ int WasapiProxy::messageHandler(std::promise<HRESULT> *promise)
CoUninitialize();
return 0;
}
auto Enumerator = static_cast<IMMDeviceEnumerator*>(ptr);
Enumerator->Release();
Enumerator = nullptr;
static_cast<IMMDeviceEnumerator*>(ptr)->Release();
CoUninitialize();
TRACE("Message thread initialization complete\n");
@@ -595,21 +648,19 @@ int WasapiProxy::messageHandler(std::promise<HRESULT> *promise)
if(++deviceCount == 1)
hr = cohr = CoInitializeEx(nullptr, COINIT_MULTITHREADED);
if(SUCCEEDED(hr))
hr = CoCreateInstance(CLSID_MMDeviceEnumerator, nullptr, CLSCTX_INPROC_SERVER, IID_IMMDeviceEnumerator, &ptr);
hr = CoCreateInstance(CLSID_MMDeviceEnumerator, nullptr, CLSCTX_INPROC_SERVER,
IID_IMMDeviceEnumerator, &ptr);
if(FAILED(hr))
msg.mPromise.set_value(hr);
else
{
Enumerator = static_cast<IMMDeviceEnumerator*>(ptr);
ComPtr<IMMDeviceEnumerator> enumerator{static_cast<IMMDeviceEnumerator*>(ptr)};
if(msg.mType == MsgType::EnumeratePlayback)
probe_devices(Enumerator, eRender, PlaybackDevices);
probe_devices(enumerator.get(), eRender, PlaybackDevices);
else if(msg.mType == MsgType::EnumerateCapture)
probe_devices(Enumerator, eCapture, CaptureDevices);
probe_devices(enumerator.get(), eCapture, CaptureDevices);
msg.mPromise.set_value(S_OK);
Enumerator->Release();
Enumerator = nullptr;
}
if(--deviceCount == 0 && SUCCEEDED(cohr))
@@ -650,9 +701,9 @@ struct WasapiPlayback final : public BackendBase, WasapiProxy {
std::wstring mDevId;
HRESULT mOpenStatus{E_FAIL};
IMMDevice *mMMDev{nullptr};
IAudioClient *mClient{nullptr};
IAudioRenderClient *mRender{nullptr};
ComPtr<IMMDevice> mMMDev{nullptr};
ComPtr<IAudioClient> mClient{nullptr};
ComPtr<IAudioRenderClient> mRender{nullptr};
HANDLE mNotifyEvent{nullptr};
UINT32 mFrameStep{0u};
@@ -680,7 +731,7 @@ WasapiPlayback::~WasapiPlayback()
FORCE_ALIGN int WasapiPlayback::mixerProc()
{
HRESULT hr = CoInitializeEx(nullptr, COINIT_MULTITHREADED);
HRESULT hr{CoInitializeEx(nullptr, COINIT_MULTITHREADED)};
if(FAILED(hr))
{
ERR("CoInitializeEx(nullptr, COINIT_MULTITHREADED) failed: 0x%08lx\n", hr);
@@ -798,43 +849,34 @@ void WasapiPlayback::open(const ALCchar *name)
HRESULT WasapiPlayback::openProxy()
{
void *ptr;
HRESULT hr{CoCreateInstance(CLSID_MMDeviceEnumerator, nullptr, CLSCTX_INPROC_SERVER, IID_IMMDeviceEnumerator, &ptr)};
HRESULT hr{CoCreateInstance(CLSID_MMDeviceEnumerator, nullptr, CLSCTX_INPROC_SERVER,
IID_IMMDeviceEnumerator, &ptr)};
if(SUCCEEDED(hr))
{
auto Enumerator = static_cast<IMMDeviceEnumerator*>(ptr);
ComPtr<IMMDeviceEnumerator> enumerator{static_cast<IMMDeviceEnumerator*>(ptr)};
if(mDevId.empty())
hr = Enumerator->GetDefaultAudioEndpoint(eRender, eMultimedia, &mMMDev);
hr = enumerator->GetDefaultAudioEndpoint(eRender, eMultimedia, mMMDev.getPtr());
else
hr = Enumerator->GetDevice(mDevId.c_str(), &mMMDev);
Enumerator->Release();
hr = enumerator->GetDevice(mDevId.c_str(), mMMDev.getPtr());
}
if(SUCCEEDED(hr))
hr = mMMDev->Activate(IID_IAudioClient, CLSCTX_INPROC_SERVER, nullptr, &ptr);
if(SUCCEEDED(hr))
{
mClient = static_cast<IAudioClient*>(ptr);
mClient = ComPtr<IAudioClient>{static_cast<IAudioClient*>(ptr)};
if(mDevice->DeviceName.empty())
mDevice->DeviceName = get_device_name_and_guid(mMMDev).first;
mDevice->DeviceName = get_device_name_and_guid(mMMDev.get()).first;
}
if(FAILED(hr))
{
if(mMMDev)
mMMDev->Release();
mMMDev = nullptr;
}
return hr;
}
void WasapiPlayback::closeProxy()
{
if(mClient)
mClient->Release();
mClient = nullptr;
if(mMMDev)
mMMDev->Release();
mMMDev = nullptr;
}
@@ -849,18 +891,16 @@ bool WasapiPlayback::reset()
HRESULT WasapiPlayback::resetProxy()
{
if(mClient)
mClient->Release();
mClient = nullptr;
void *ptr;
HRESULT hr = mMMDev->Activate(IID_IAudioClient, CLSCTX_INPROC_SERVER, nullptr, &ptr);
HRESULT hr{mMMDev->Activate(IID_IAudioClient, CLSCTX_INPROC_SERVER, nullptr, &ptr)};
if(FAILED(hr))
{
ERR("Failed to reactivate audio client: 0x%08lx\n", hr);
return hr;
}
mClient = static_cast<IAudioClient*>(ptr);
mClient = ComPtr<IAudioClient>{static_cast<IAudioClient*>(ptr)};
WAVEFORMATEX *wfx;
hr = mClient->GetMixFormat(&wfx);
@@ -1063,7 +1103,7 @@ HRESULT WasapiPlayback::resetProxy()
mFrameStep = OutputType.Format.nChannels;
EndpointFormFactor formfactor{UnknownFormFactor};
get_device_formfactor(mMMDev, &formfactor);
get_device_formfactor(mMMDev.get(), &formfactor);
mDevice->IsHeadphones = (mDevice->FmtChans == DevFmtStereo
&& (formfactor == Headphones || formfactor == Headset));
@@ -1116,7 +1156,7 @@ HRESULT WasapiPlayback::startProxy()
{
ResetEvent(mNotifyEvent);
HRESULT hr = mClient->Start();
HRESULT hr{mClient->Start()};
if(FAILED(hr))
{
ERR("Failed to start audio client: 0x%08lx\n", hr);
@@ -1127,13 +1167,12 @@ HRESULT WasapiPlayback::startProxy()
hr = mClient->GetService(IID_IAudioRenderClient, &ptr);
if(SUCCEEDED(hr))
{
mRender = static_cast<IAudioRenderClient*>(ptr);
mRender = ComPtr<IAudioRenderClient>{static_cast<IAudioRenderClient*>(ptr)};
try {
mKillNow.store(false, std::memory_order_release);
mThread = std::thread{std::mem_fn(&WasapiPlayback::mixerProc), this};
}
catch(...) {
mRender->Release();
mRender = nullptr;
ERR("Failed to start thread\n");
hr = E_FAIL;
@@ -1158,7 +1197,6 @@ void WasapiPlayback::stopProxy()
mKillNow.store(true, std::memory_order_release);
mThread.join();
mRender->Release();
mRender = nullptr;
mClient->Stop();
}
@@ -1199,9 +1237,9 @@ struct WasapiCapture final : public BackendBase, WasapiProxy {
std::wstring mDevId;
HRESULT mOpenStatus{E_FAIL};
IMMDevice *mMMDev{nullptr};
IAudioClient *mClient{nullptr};
IAudioCaptureClient *mCapture{nullptr};
ComPtr<IMMDevice> mMMDev{nullptr};
ComPtr<IAudioClient> mClient{nullptr};
ComPtr<IAudioCaptureClient> mCapture{nullptr};
HANDLE mNotifyEvent{nullptr};
ChannelConverter mChannelConv{};
@@ -1228,7 +1266,7 @@ WasapiCapture::~WasapiCapture()
FORCE_ALIGN int WasapiCapture::recordProc()
{
HRESULT hr = CoInitializeEx(nullptr, COINIT_MULTITHREADED);
HRESULT hr{CoInitializeEx(nullptr, COINIT_MULTITHREADED)};
if(FAILED(hr))
{
ERR("CoInitializeEx(nullptr, COINIT_MULTITHREADED) failed: 0x%08lx\n", hr);
@@ -1389,47 +1427,35 @@ HRESULT WasapiCapture::openProxy()
IID_IMMDeviceEnumerator, &ptr)};
if(SUCCEEDED(hr))
{
auto Enumerator = static_cast<IMMDeviceEnumerator*>(ptr);
ComPtr<IMMDeviceEnumerator> enumerator{static_cast<IMMDeviceEnumerator*>(ptr)};
if(mDevId.empty())
hr = Enumerator->GetDefaultAudioEndpoint(eCapture, eMultimedia, &mMMDev);
hr = enumerator->GetDefaultAudioEndpoint(eCapture, eMultimedia, mMMDev.getPtr());
else
hr = Enumerator->GetDevice(mDevId.c_str(), &mMMDev);
Enumerator->Release();
hr = enumerator->GetDevice(mDevId.c_str(), mMMDev.getPtr());
}
if(SUCCEEDED(hr))
hr = mMMDev->Activate(IID_IAudioClient, CLSCTX_INPROC_SERVER, nullptr, &ptr);
if(SUCCEEDED(hr))
{
mClient = static_cast<IAudioClient*>(ptr);
mClient = ComPtr<IAudioClient>{static_cast<IAudioClient*>(ptr)};
if(mDevice->DeviceName.empty())
mDevice->DeviceName = get_device_name_and_guid(mMMDev).first;
mDevice->DeviceName = get_device_name_and_guid(mMMDev.get()).first;
}
if(FAILED(hr))
{
if(mMMDev)
mMMDev->Release();
mMMDev = nullptr;
}
return hr;
}
void WasapiCapture::closeProxy()
{
if(mClient)
mClient->Release();
mClient = nullptr;
if(mMMDev)
mMMDev->Release();
mMMDev = nullptr;
}
HRESULT WasapiCapture::resetProxy()
{
if(mClient)
mClient->Release();
mClient = nullptr;
void *ptr;
@@ -1439,7 +1465,7 @@ HRESULT WasapiCapture::resetProxy()
ERR("Failed to reactivate audio client: 0x%08lx\n", hr);
return hr;
}
mClient = static_cast<IAudioClient*>(ptr);
mClient = ComPtr<IAudioClient>{static_cast<IAudioClient*>(ptr)};
// Make sure buffer is at least 100ms in size
ReferenceTime buf_time{ReferenceTime{seconds{mDevice->BufferSize}} / mDevice->Frequency};
@@ -1668,13 +1694,12 @@ HRESULT WasapiCapture::startProxy()
hr = mClient->GetService(IID_IAudioCaptureClient, &ptr);
if(SUCCEEDED(hr))
{
mCapture = static_cast<IAudioCaptureClient*>(ptr);
mCapture = ComPtr<IAudioCaptureClient>{static_cast<IAudioCaptureClient*>(ptr)};
try {
mKillNow.store(false, std::memory_order_release);
mThread = std::thread{std::mem_fn(&WasapiCapture::recordProc), this};
}
catch(...) {
mCapture->Release();
mCapture = nullptr;
ERR("Failed to start thread\n");
hr = E_FAIL;
@@ -1702,7 +1727,6 @@ void WasapiCapture::stopProxy()
mKillNow.store(true, std::memory_order_release);
mThread.join();
mCapture->Release();
mCapture = nullptr;
mClient->Stop();
mClient->Reset();