Spaces:
Running
Running
| """ | |
| modules/audio_device_control.py — Phase 4 (System Control): Default audio I/O device switching (Windows). | |
| Implements CoreAudio endpoint enumeration + default device switching using comtypes. | |
| No fake outputs: returns (ok, message). | |
| References: | |
| - IMMDeviceEnumerator + MMDevice API | |
| - IPolicyConfig (undocumented but widely used for SetDefaultEndpoint) | |
| """ | |
| from __future__ import annotations | |
| from dataclasses import dataclass | |
| from typing import Literal | |
| import ctypes | |
| from ctypes import POINTER, wintypes | |
| import comtypes # type: ignore | |
| from comtypes import GUID, HRESULT, IUnknown # type: ignore | |
| from comtypes.client import CreateObject # type: ignore | |
| DataFlow = Literal["render", "capture"] | |
| class AudioDevResult: | |
| ok: bool | |
| message: str | |
| # --- CoreAudio constants --- | |
| eRender = 0 | |
| eCapture = 1 | |
| eAll = 2 | |
| eConsole = 0 | |
| eMultimedia = 1 | |
| eCommunications = 2 | |
| DEVICE_STATE_ACTIVE = 0x00000001 | |
| STGM_READ = 0x00000000 | |
| class PROPERTYKEY(ctypes.Structure): | |
| _fields_ = [("fmtid", GUID), ("pid", wintypes.DWORD)] | |
| class PROPVARIANT(ctypes.Structure): | |
| _fields_ = [ | |
| ("vt", wintypes.USHORT), | |
| ("wReserved1", wintypes.USHORT), | |
| ("wReserved2", wintypes.USHORT), | |
| ("wReserved3", wintypes.USHORT), | |
| ("p", ctypes.c_void_p), | |
| ("p2", ctypes.c_void_p), | |
| ] | |
| PKEY_Device_FriendlyName = PROPERTYKEY( | |
| GUID("{A45C254E-DF1C-4EFD-8020-67D146A850E0}"), | |
| 14, | |
| ) | |
| class IPropertyStore(IUnknown): | |
| _iid_ = GUID("{886D8EEB-8CF2-4446-8D02-CDBA1DBDCF99}") | |
| _methods_ = [ | |
| comtypes.COMMETHOD([], HRESULT, "GetCount", (["out"], POINTER(wintypes.DWORD), "cProps")), | |
| comtypes.COMMETHOD( | |
| [], HRESULT, "GetAt", (["in"], wintypes.DWORD, "iProp"), (["out"], POINTER(PROPERTYKEY), "pkey") | |
| ), | |
| comtypes.COMMETHOD( | |
| [], HRESULT, "GetValue", (["in"], POINTER(PROPERTYKEY), "key"), (["out"], POINTER(PROPVARIANT), "pv") | |
| ), | |
| ] | |
| class IMMDevice(IUnknown): | |
| _iid_ = GUID("{D666063F-1587-4E43-81F1-B948E807363F}") | |
| _methods_ = [ | |
| comtypes.COMMETHOD([], HRESULT, "Activate"), | |
| comtypes.COMMETHOD([], HRESULT, "OpenPropertyStore", (["in"], wintypes.DWORD, "stgmAccess"), (["out"], POINTER(POINTER(IPropertyStore)), "ppProperties")), | |
| comtypes.COMMETHOD([], HRESULT, "GetId", (["out"], POINTER(wintypes.LPWSTR), "ppstrId")), | |
| comtypes.COMMETHOD([], HRESULT, "GetState", (["out"], POINTER(wintypes.DWORD), "pdwState")), | |
| ] | |
| class IMMDeviceCollection(IUnknown): | |
| _iid_ = GUID("{0BD7A1BE-7A1A-44DB-8397-C0F6C1F041C5}") | |
| _methods_ = [ | |
| comtypes.COMMETHOD([], HRESULT, "GetCount", (["out"], POINTER(wintypes.UINT), "pcDevices")), | |
| comtypes.COMMETHOD([], HRESULT, "Item", (["in"], wintypes.UINT, "nDevice"), (["out"], POINTER(POINTER(IMMDevice)), "ppDevice")), | |
| ] | |
| class IMMDeviceEnumerator(IUnknown): | |
| _iid_ = GUID("{A95664D2-9614-4F35-A746-DE8DB63617E6}") | |
| _methods_ = [ | |
| comtypes.COMMETHOD( | |
| [], | |
| HRESULT, | |
| "EnumAudioEndpoints", | |
| (["in"], wintypes.DWORD, "dataFlow"), | |
| (["in"], wintypes.DWORD, "dwStateMask"), | |
| (["out"], POINTER(POINTER(IMMDeviceCollection)), "ppDevices"), | |
| ), | |
| comtypes.COMMETHOD([], HRESULT, "GetDefaultAudioEndpoint"), | |
| comtypes.COMMETHOD([], HRESULT, "GetDevice", (["in"], wintypes.LPCWSTR, "pwstrId"), (["out"], POINTER(POINTER(IMMDevice)), "ppDevice")), | |
| comtypes.COMMETHOD([], HRESULT, "RegisterEndpointNotificationCallback"), | |
| comtypes.COMMETHOD([], HRESULT, "UnregisterEndpointNotificationCallback"), | |
| ] | |
| # --- PolicyConfig (SetDefaultEndpoint) --- | |
| class IPolicyConfig(IUnknown): | |
| _iid_ = GUID("{F8679F50-850A-41CF-9C72-430F290290C8}") | |
| _methods_ = [ | |
| comtypes.COMMETHOD([], HRESULT, "GetMixFormat"), | |
| comtypes.COMMETHOD([], HRESULT, "GetDeviceFormat"), | |
| comtypes.COMMETHOD([], HRESULT, "SetDeviceFormat"), | |
| comtypes.COMMETHOD([], HRESULT, "GetProcessingPeriod"), | |
| comtypes.COMMETHOD([], HRESULT, "SetProcessingPeriod"), | |
| comtypes.COMMETHOD([], HRESULT, "GetShareMode"), | |
| comtypes.COMMETHOD([], HRESULT, "SetShareMode"), | |
| comtypes.COMMETHOD([], HRESULT, "GetPropertyValue"), | |
| comtypes.COMMETHOD([], HRESULT, "SetPropertyValue"), | |
| comtypes.COMMETHOD([], HRESULT, "SetDefaultEndpoint", (["in"], wintypes.LPCWSTR, "pwstrDeviceId"), (["in"], wintypes.DWORD, "role")), | |
| comtypes.COMMETHOD([], HRESULT, "SetEndpointVisibility"), | |
| ] | |
| CLSID_MMDeviceEnumerator = GUID("{BCDE0395-E52F-467C-8E3D-C4579291692E}") | |
| CLSID_PolicyConfigClient = GUID("{870AF99C-171D-4F9E-AF0D-E63DF40C2BC9}") | |
| def _get_friendly_name(dev: IMMDevice) -> str: | |
| store_ptr = POINTER(IPropertyStore)() | |
| hr = dev.OpenPropertyStore(STGM_READ, ctypes.byref(store_ptr)) | |
| if hr != 0 or not store_ptr: | |
| return "" | |
| pv = PROPVARIANT() | |
| hr2 = store_ptr.GetValue(ctypes.byref(PKEY_Device_FriendlyName), ctypes.byref(pv)) | |
| if hr2 != 0: | |
| return "" | |
| # VT_LPWSTR = 31 | |
| if pv.vt != 31 or not pv.p: | |
| return "" | |
| return ctypes.wstring_at(pv.p) | |
| def _enum(flow: DataFlow) -> list[tuple[str, str]]: | |
| """ | |
| Returns list of (device_id, friendly_name) for active endpoints. | |
| """ | |
| enumerator = CreateObject(CLSID_MMDeviceEnumerator, interface=IMMDeviceEnumerator) | |
| flow_id = eRender if flow == "render" else eCapture | |
| coll_ptr = POINTER(IMMDeviceCollection)() | |
| hr = enumerator.EnumAudioEndpoints(flow_id, DEVICE_STATE_ACTIVE, ctypes.byref(coll_ptr)) | |
| if hr != 0 or not coll_ptr: | |
| return [] | |
| count = wintypes.UINT() | |
| coll_ptr.GetCount(ctypes.byref(count)) | |
| out: list[tuple[str, str]] = [] | |
| for i in range(int(count.value)): | |
| dev_ptr = POINTER(IMMDevice)() | |
| coll_ptr.Item(i, ctypes.byref(dev_ptr)) | |
| if not dev_ptr: | |
| continue | |
| did = wintypes.LPWSTR() | |
| dev_ptr.GetId(ctypes.byref(did)) | |
| name = _get_friendly_name(dev_ptr) | |
| out.append((str(did), name)) | |
| return out | |
| def list_devices(flow: DataFlow) -> AudioDevResult: | |
| devs = _enum(flow) | |
| if not devs: | |
| return AudioDevResult(False, "No audio devices found.") | |
| lines = [f"{name} | {did}" for did, name in devs] | |
| return AudioDevResult(True, "\n".join(lines[:50])) | |
| def set_default_device(flow: DataFlow, name_substring: str) -> AudioDevResult: | |
| s = (name_substring or "").strip().lower() | |
| if not s: | |
| return AudioDevResult(False, "Missing device name.") | |
| devs = _enum(flow) | |
| if not devs: | |
| return AudioDevResult(False, "No audio devices found.") | |
| target_id = None | |
| target_name = None | |
| for did, name in devs: | |
| if s in (name or "").lower(): | |
| target_id = did | |
| target_name = name | |
| break | |
| if not target_id: | |
| return AudioDevResult(False, f"No {flow} device matching '{name_substring}'.") | |
| try: | |
| client = CreateObject(CLSID_PolicyConfigClient, interface=IPolicyConfig) | |
| for role in (eConsole, eMultimedia, eCommunications): | |
| client.SetDefaultEndpoint(target_id, role) | |
| return AudioDevResult(True, f"Default {flow} device set to {target_name}.") | |
| except Exception as e: | |
| return AudioDevResult(False, f"Audio device switch failed: {e}") | |
| def set_default_output(name_substring: str) -> AudioDevResult: | |
| return set_default_device("render", name_substring) | |
| def set_default_input(name_substring: str) -> AudioDevResult: | |
| return set_default_device("capture", name_substring) | |