diff --git a/.github/workflows/lint.yaml b/.github/workflows/lint.yaml index faba17e..690ea7a 100644 --- a/.github/workflows/lint.yaml +++ b/.github/workflows/lint.yaml @@ -76,6 +76,11 @@ jobs: pip install ruff pip install -r ./homeway/requirements.txt pip install "zstandard>=0.21.0,<0.23.0" + - name: ๐Ÿงช Running unit tests + env: + PYTHONPATH: ${{ github.workspace }}/homeway + PYTHONDONTWRITEBYTECODE: "1" + run: python -m unittest discover -s tests -p "test_*.py" -v # This crazy script is needed to ensure that we always run pylint # on all of the files in the repo and the path is set correctly for it to work. - name: ๐Ÿ‘– Running Pylint diff --git a/homeway/homeway/interfaces.py b/homeway/homeway/interfaces.py index b6bbaf6..0b38523 100644 --- a/homeway/homeway/interfaces.py +++ b/homeway/homeway/interfaces.py @@ -5,9 +5,8 @@ from .buffer import Buffer from .httpresult import HttpResult -from .Proto import WebStreamMsg - if TYPE_CHECKING: + from .Proto import WebStreamMsg from .compression import CompressionResult # @@ -248,7 +247,7 @@ class IWebStreamHelper(ABC): # This function should throw on critical errors, that will reset the connection. # Returning true will case the websocket to close on return. @abstractmethod - def IncomingServerMessage(self, webStreamMsg:WebStreamMsg.WebStreamMsg) -> bool: + def IncomingServerMessage(self, webStreamMsg:"WebStreamMsg.WebStreamMsg") -> bool: pass diff --git a/homeway/homeway_linuxhost/ha/configmanager.py b/homeway/homeway_linuxhost/ha/configmanager.py index 58bb939..d8bcc06 100644 --- a/homeway/homeway_linuxhost/ha/configmanager.py +++ b/homeway/homeway_linuxhost/ha/configmanager.py @@ -51,8 +51,13 @@ def __init__(self, logger: logging.Logger) -> None: self.HaConnection: Optional[Connection] = None self.RestartRequired: bool = False self.HttpConfigUpdateStateLock = threading.Lock() + self.HttpConfigUpdateWakeEvent = threading.Event() self.HttpConfigUpdateThreadRunning = False self.HttpConfigUpdateRequested = False + # Normalized fingerprint of an HTTP config staged by Homeway and awaiting + # confirmation. Never use a bare pending state here, since it might belong to a + # user changing unrelated settings in the Home Assistant UI. + self.PendingHttpConfigToPromote: Optional[Dict[str, Any]] = None CommandHandler.Get().RegisterConfigManager(self) @@ -226,6 +231,7 @@ def _UpdateAssistantConfigIfNeeded(self, configFilePath: str) -> bool: def _OnHaConnected(self) -> None: with self.HttpConfigUpdateStateLock: self.HttpConfigUpdateRequested = True + self.HttpConfigUpdateWakeEvent.set() if self.HttpConfigUpdateThreadRunning: return self.HttpConfigUpdateThreadRunning = True @@ -240,11 +246,16 @@ def _UpdateHttpConfig_Thread(self) -> None: while True: with self.HttpConfigUpdateStateLock: self.HttpConfigUpdateRequested = False + self.HttpConfigUpdateWakeEvent.clear() shouldRetry = self._UpdateHttpConfigViaApiIfNeeded() if (shouldRetry and retryCount < ConfigManager.c_HttpConfigUpdateRetryCount): retryCount += 1 - time.sleep(ConfigManager.c_HttpConfigUpdateRetryDelaySec) + # A reconnect means the API is available again, so retry immediately + # instead of waiting out the normal startup delay. + self.HttpConfigUpdateWakeEvent.wait( + ConfigManager.c_HttpConfigUpdateRetryDelaySec + ) continue with self.HttpConfigUpdateStateLock: @@ -284,6 +295,48 @@ def _UpdateHttpConfigViaApiIfNeeded(self) -> bool: activeConfigType = result.get("active_config_type", "stable") activePendingConfig = activeConfigType == "pending" + pendingConfig = result.get("pending") + pendingConfigOwnedByHomeway = ( + isinstance(pendingConfig, dict) + and self._IsPendingHttpConfigOwnedByHomeway(pendingConfig) + ) + + if self._HasPendingHttpConfigOwnedByHomeway(): + if pendingConfigOwnedByHomeway: + assert isinstance(pendingConfig, dict) + pendingConfig = cast(Dict[str, Any], pendingConfig) + if pendingConfig.get("error") is not None: + self.Logger.warning( + "Homeway's pending HTTP config was rejected by Home Assistant; it will not be auto-confirmed." + ) + self._ClearPendingHttpConfigOwnedByHomeway() + haConnection.ClearServerRestartExpected() + return False + + if activePendingConfig: + # Seeing our exact config active proves the API-triggered restart + # completed, even if its configure response was lost with the socket. + self.RestartRequired = False + else: + # Configure stores the pending slot before its queued restart runs. If a + # retry reaches the old server during that window, wait for the restart + # rather than staging the same config again. + self.Logger.info( + "Homeway's HTTP config is pending; waiting for Home Assistant to restart." + ) + return True + else: + # The request definitively did not leave our config pending (or another + # operation replaced it), so it is no longer safe to auto-promote. + self._ClearPendingHttpConfigOwnedByHomeway() + haConnection.ClearServerRestartExpected() + + if isinstance(pendingConfig, dict) and not pendingConfigOwnedByHomeway: + self.Logger.info( + "Home Assistant has a pending HTTP config not staged by Homeway; leaving it for the user to confirm." + ) + return False + sourceConfig = result.get("pending" if activePendingConfig else "stable") if activeConfigType == "default" or activeConfigType == "default_legacy_port": sourceConfig = result.get("default") @@ -321,18 +374,34 @@ def _UpdateHttpConfigViaApiIfNeeded(self) -> bool: return False self.Logger.info("Updating Home Assistant HTTP trusted proxy config through the WebSocket API.") + # Arm both the ownership check and quick reconnect path before configure. HA can + # close the websocket for its queued restart before its response reaches us. + self._SetPendingHttpConfigOwnedByHomeway(config) + haConnection.SetServerRestartExpected() response = haConnection.SendAndReceiveMsg( {"type": "http/config/configure", "config": config} ) shouldRetry, configureResult = self._GetHttpConfigApiResult("update", response) if configureResult is None: + if response is not None: + # A definite API error means configure did not queue a restart. A missing + # response is ambiguous, so retain the state for reconnect recovery. + self._ClearPendingHttpConfigOwnedByHomeway() + haConnection.ClearServerRestartExpected() return shouldRetry if configureResult.get("restart", False): # The API restart also applies any assistant YAML changes waiting for a restart. self.RestartRequired = False self.Logger.info("Home Assistant is restarting to apply the HTTP trusted proxy config.") + # Best-effort confirmation on the current websocket prevents HA's frontend from + # observing the pending slot and leaving its dialog open. This intentionally + # bypasses the post-restart trial for Homeway's narrow proxy-only change; the + # exact-config ownership check and reconnect path handle an interrupted request. + return self._PromotePendingHttpConfig(haConnection) else: + self._ClearPendingHttpConfigOwnedByHomeway() + haConnection.ClearServerRestartExpected() self.Logger.info("Home Assistant accepted the HTTP trusted proxy config without requiring a restart.") return False @@ -342,10 +411,55 @@ def _PromotePendingHttpConfig(self, haConnection: Connection) -> bool: response = haConnection.SendAndReceiveMsg({"type": "http/config/promote"}) shouldRetry, result = self._GetHttpConfigApiResult("confirm", response, allowEmptyResult=True) if result is not None: + self._ClearPendingHttpConfigOwnedByHomeway() + haConnection.ClearServerRestartExpected() self.Logger.info("Home Assistant HTTP trusted proxy config confirmed.") return shouldRetry + def _SetPendingHttpConfigOwnedByHomeway(self, config: Dict[str, Any]) -> None: + normalizedConfig = self._NormalizeHttpConfigForComparison(config) + with self.HttpConfigUpdateStateLock: + self.PendingHttpConfigToPromote = normalizedConfig + + + def _ClearPendingHttpConfigOwnedByHomeway(self) -> None: + with self.HttpConfigUpdateStateLock: + self.PendingHttpConfigToPromote = None + + + def _HasPendingHttpConfigOwnedByHomeway(self) -> bool: + with self.HttpConfigUpdateStateLock: + return self.PendingHttpConfigToPromote is not None + + + def _IsPendingHttpConfigOwnedByHomeway(self, config: Dict[str, Any]) -> bool: + normalizedConfig = self._NormalizeHttpConfigForComparison(config) + with self.HttpConfigUpdateStateLock: + return normalizedConfig == self.PendingHttpConfigToPromote + + + @staticmethod + def _NormalizeHttpConfigForComparison(config: Dict[str, Any]) -> Dict[str, Any]: + normalizedConfig = dict(config) + for key in ConfigManager.c_HttpConfigMetaKeys: + normalizedConfig.pop(key, None) + + # HA's storage schema canonicalizes host addresses to /32 or /128 networks. + # Normalize both our outbound fingerprint and the returned pending config so the + # ownership check survives that representation change. + trustedProxies = normalizedConfig.get("trusted_proxies") + if isinstance(trustedProxies, list): + normalizedTrustedProxies: List[Any] = [] + for trustedProxy in trustedProxies: + try: + normalizedTrustedProxies.append(str(ip_network(trustedProxy))) + except ValueError: + normalizedTrustedProxies.append(trustedProxy) + normalizedConfig["trusted_proxies"] = normalizedTrustedProxies + return normalizedConfig + + def _GetHttpConfigApiResult( self, operation: str, diff --git a/homeway/homeway_linuxhost/ha/connection.py b/homeway/homeway_linuxhost/ha/connection.py index 90b0002..8ee9e71 100644 --- a/homeway/homeway_linuxhost/ha/connection.py +++ b/homeway/homeway_linuxhost/ha/connection.py @@ -20,6 +20,11 @@ class Connection(IHomeAssistantWebSocket): # For debugging, it's too chatty to enable always. c_LogWsMessages = False + # During an intentional HA restart, retry quickly so config confirmation and the other + # local integrations are restored before remote clients complete their reconnect flow. + c_ExpectedServerRestartReconnectDelaySec = 1.0 + c_ExpectedServerRestartReconnectMaxAttempts = 60 + def __init__(self, logger:logging.Logger, eventHandler:EventHandler) -> None: self.Logger = logger @@ -41,6 +46,12 @@ def __init__(self, logger:logging.Logger, eventHandler:EventHandler) -> None: # If set, when the websocket is connected, we should send the HA restart command. self.IssueRestartOnConnect = False + # Set when HA tells us it is about to restart. This avoids applying the normal + # connection-error backoff to an expected, short-lived disconnect. + self.ServerRestartExpectedLock = threading.Lock() + self.ServerRestartExpected = threading.Event() + self.ServerRestartReconnectAttempt = 0 + # Allows for blocking message send responses. self.PendingContextsLock = threading.Lock() self.PendingContexts: Dict[int, PendingContexts] = {} @@ -68,6 +79,35 @@ def GetHomeAssistantVersionString(self) -> Optional[str]: return self.HaVersionString + # Indicates that Home Assistant is about to restart intentionally. + def SetServerRestartExpected(self) -> None: + with self.ServerRestartExpectedLock: + self.ServerRestartReconnectAttempt = 0 + self.ServerRestartExpected.set() + + + # Clears the intentional restart reconnect state when HA confirms no restart is coming. + def ClearServerRestartExpected(self) -> None: + with self.ServerRestartExpectedLock: + self.ServerRestartReconnectAttempt = 0 + self.ServerRestartExpected.clear() + + + def _ShouldUseExpectedServerRestartReconnect(self) -> bool: + with self.ServerRestartExpectedLock: + if self.ServerRestartExpected.is_set() is False: + return False + if ( + self.ServerRestartReconnectAttempt + >= Connection.c_ExpectedServerRestartReconnectMaxAttempts + ): + self.ServerRestartReconnectAttempt = 0 + self.ServerRestartExpected.clear() + return False + self.ServerRestartReconnectAttempt += 1 + return True + + # Issues the restart command to Home Assistant. def RestartHa(self) -> None: if self.IsConnected: @@ -130,16 +170,32 @@ def _OnConnected(self) -> None: def ConnectionThread(self): while True: # Reset the state vars - self.IsConnected = False - self.Ws = None - self.MsgId = 1 + with self.PendingContextsLock: + self.IsConnected = False + self.Ws = None + pendingContexts = list(self.PendingContexts.values()) + self.PendingContexts.clear() + for pendingContext in pendingContexts: + pendingContext.Event.set() + with self.MsgIdLock: + self.MsgId = 1 # If this isn't the first connection, sleep a bit before trying again. if self.ConId != 0: - self.BackoffCounter += 1 - self.BackoffCounter = min(self.BackoffCounter, 12) - self.Logger.error(f"{self._getLogTag()} sleeping before trying the HA connection again.") - time.sleep(5 * self.BackoffCounter) + if self._ShouldUseExpectedServerRestartReconnect(): + if self.ServerRestartReconnectAttempt == 1: + self.Logger.info( + f"{self._getLogTag()} HA restart expected; reconnecting without normal backoff." + ) + else: + time.sleep(Connection.c_ExpectedServerRestartReconnectDelaySec) + else: + self.BackoffCounter += 1 + self.BackoffCounter = min(self.BackoffCounter, 12) + self.Logger.error(f"{self._getLogTag()} sleeping before trying the HA connection again.") + # Wake immediately if another thread learns that this disconnect is + # an expected server restart while normal backoff is in progress. + self.ServerRestartExpected.wait(5 * self.BackoffCounter) self.ConId += 1 try: @@ -159,14 +215,16 @@ def ConnectionThread(self): # If we got auth from the env var, we running in the add on and use this address. uri = f"{(ServerInfo.GetApiServerBaseUrl('ws'))}/api/websocket" self.Logger.info(f"{self._getLogTag()} Starting connection to [{uri}]") - self.Ws = Client(uri, onWsOpen=self.Opened, onWsData=self._OnData, onWsClose=self.Closed) + ws = Client(uri, onWsOpen=self.Opened, onWsData=self._OnData, onWsClose=self.Closed) + with self.PendingContextsLock: + self.Ws = ws # It's important that we disable cert checks since the server might have a self signed cert or cert for a hostname that we aren't using. # This is safe to do, since the connection will be localhost or on the local LAN - self.Ws.SetDisableCertCheck(True) + ws.SetDisableCertCheck(True) # Run until success or failure. - self.Ws.RunUntilClosed() + ws.RunUntilClosed() self.Logger.info(f"{self._getLogTag()} Loop restarting.") @@ -185,10 +243,29 @@ def Opened(self, ws: IWebSocketClient): # Called when the websocket is closed. def Closed(self, ws: IWebSocketClient): self.Logger.info(f"{self._getLogTag()} Websocket closed") + # A restart can close the socket while a request is waiting for its result. Wake + # those callers immediately rather than making them wait for the full timeout. + # Detaching the map also prevents response IDs from the next websocket session + # colliding with requests left over from this one. + with self.PendingContextsLock: + # Ignore a delayed callback from an older websocket. It must not tear down the + # state or wake requests belonging to the current connection. + if ws is not self.Ws: + return + self.IsConnected = False + self.Ws = None + pendingContexts = list(self.PendingContexts.values()) + self.PendingContexts.clear() + for pendingContext in pendingContexts: + pendingContext.Event.set() def _OnData(self, ws: IWebSocketClient, buffer: Buffer, msgType: WebSocketOpCode) -> None: try: + # Ignore messages delivered late by a websocket from an older connection loop. + if ws is not self.Ws: + return + jsonStr = buffer.GetBytesLike().decode() jsonObj: Dict[str, Any] = json.loads(jsonStr) if self.Logger.isEnabledFor(logging.DEBUG) and Connection.c_LogWsMessages: @@ -200,7 +277,10 @@ def _OnData(self, ws: IWebSocketClient, buffer: Buffer, msgType: WebSocketOpCode # Check if this is the auth response. if "type" in jsonObj and jsonObj["type"] == "auth_ok": # Auth success! - self.IsConnected = True + with self.PendingContextsLock: + if ws is not self.Ws: + return + self.IsConnected = True self.BackoffCounter = 0 self._OnConnected() return @@ -247,6 +327,8 @@ def _OnData(self, ws: IWebSocketClient, buffer: Buffer, msgType: WebSocketOpCode # Check if there's a pending context for this message. # It's ok if there's no pending context, since we might have sent a message that we don't care about the response. with self.PendingContextsLock: + if ws is not self.Ws: + return if msgId in self.PendingContexts: # If we find a mach, set the response and signal the event. pendingContext = self.PendingContexts[msgId] @@ -278,31 +360,32 @@ def SendMsg( ignoreConnectionState: bool = False, timeoutSec: float = 10.0, ) -> Optional[Dict[str, Any]]: - # Check the connection state. - if ignoreConnectionState is False: - if self.IsConnected is False: - self.Logger.error(f"{self._getLogTag()} message tried to be sent while we weren't authed.") - return None - # Capture and check the websocket. - ws = self.Ws - if ws is None: - self.Logger.error(f"{self._getLogTag()} message tried to be sent while we weren't connected.") - return None - msgId = 0 pendingContext: Optional[PendingContexts] = None try: - # Add the id field to all messages that are post auth. - if self.IsConnected: - with self.MsgIdLock: - msgId = self.MsgId - msg["id"] = msgId - self.MsgId += 1 - - # Create a pending context - if waitForResponse: - pendingContext = PendingContexts() - with self.PendingContextsLock: + # Check the connection, capture its websocket, and register the pending + # context under the same lock Closed() uses to tear them down. Otherwise a + # close between the state check and registration can strand this call until + # its full timeout and let its message ID leak into the next websocket. + with self.PendingContextsLock: + if ignoreConnectionState is False and self.IsConnected is False: + self.Logger.error(f"{self._getLogTag()} message tried to be sent while we weren't authed.") + return None + ws = self.Ws + if ws is None: + self.Logger.error(f"{self._getLogTag()} message tried to be sent while we weren't connected.") + return None + + # Add the id field to all messages that are post auth. + if self.IsConnected: + with self.MsgIdLock: + msgId = self.MsgId + msg["id"] = msgId + self.MsgId += 1 + + # Create a pending context. + if waitForResponse: + pendingContext = PendingContexts() self.PendingContexts[msgId] = pendingContext # Dump the message @@ -334,7 +417,10 @@ def SendMsg( # If we have a pending context, make sure to remove it. if pendingContext is not None: with self.PendingContextsLock: - del self.PendingContexts[msgId] + # Closed() can already have removed this context. Only remove the + # current object so a reused message ID can never delete a newer call. + if self.PendingContexts.get(msgId) is pendingContext: + del self.PendingContexts[msgId] return None diff --git a/pyrightconfig.json b/pyrightconfig.json index 0be544c..ee433ac 100644 --- a/pyrightconfig.json +++ b/pyrightconfig.json @@ -7,6 +7,13 @@ "**/debugprofiler.py" ], + "executionEnvironments": [ + { + "root": "tests", + "extraPaths": ["homeway"] + } + ], + // Set the version to the oldest we support, so we can catch issues. "pythonVersion": "3.7", "pythonPlatform": "Linux", diff --git a/tests/test_configmanager.py b/tests/test_configmanager.py new file mode 100644 index 0000000..00e46e1 --- /dev/null +++ b/tests/test_configmanager.py @@ -0,0 +1,361 @@ +"""Tests for Home Assistant HTTP configuration management.""" + +# pylint: disable=import-error,no-name-in-module,protected-access + +import logging +import threading +import unittest +from typing import Any, Dict, List, Optional, cast + +from homeway_linuxhost.ha.configmanager import ConfigManager +from homeway_linuxhost.ha.connection import Connection + + +class FakeConnection: + """Minimal Home Assistant WebSocket connection fake.""" + + def __init__( + self, + responses: List[Optional[Dict[str, Any]]], + version: str = "2026.8.1", + ) -> None: + self.Responses = list(responses) + self.Messages: List[Dict[str, Any]] = [] + self.Version = version + self.ServerRestartExpected = False + + def GetHomeAssistantVersionString(self) -> Optional[str]: + return self.Version + + def SetServerRestartExpected(self) -> None: + self.ServerRestartExpected = True + + def ClearServerRestartExpected(self) -> None: + self.ServerRestartExpected = False + + def SendAndReceiveMsg( + self, msg: Dict[str, Any], timeoutSec: float = 10.0 + ) -> Optional[Dict[str, Any]]: + del timeoutSec + self.Messages.append(msg) + return self.Responses.pop(0) + +def _HttpConfigResponse( + stable: Dict[str, Any], + activeConfigType: str = "stable", + pending: Optional[Dict[str, Any]] = None, +) -> Dict[str, Any]: + return { + "success": True, + "result": { + "stable": stable, + "pending": pending, + "active_config_type": activeConfigType, + "default": stable, + }, + } + + +def _CreateManager(connection: FakeConnection) -> ConfigManager: + manager = object.__new__(ConfigManager) + manager.Logger = logging.getLogger("test_configmanager") + manager.HaConnection = cast(Connection, connection) + manager.RestartRequired = True + manager.HttpConfigUpdateStateLock = threading.Lock() + manager.HttpConfigUpdateWakeEvent = threading.Event() + manager.HttpConfigUpdateThreadRunning = False + manager.HttpConfigUpdateRequested = False + manager.PendingHttpConfigToPromote = None + return manager + + +class ConfigManagerHttpConfigTests(unittest.TestCase): + """Exercise the HA 2026.8 HTTP config API lifecycle.""" + + def test_confirms_staged_config_before_queued_restart(self) -> None: + stable: Dict[str, Any] = { + "server_port": 8123, + "trusted_proxies": ["10.0.0.0/24"], + } + connection = FakeConnection( + [ + _HttpConfigResponse(stable), + {"success": True, "result": {"restart": True}}, + {"success": True, "result": None}, + ] + ) + manager = _CreateManager(connection) + + shouldRetry = manager._UpdateHttpConfigViaApiIfNeeded() + + self.assertFalse(shouldRetry) + self.assertEqual( + [msg["type"] for msg in connection.Messages], + ["http/config", "http/config/configure", "http/config/promote"], + ) + self.assertFalse(connection.ServerRestartExpected) + self.assertFalse(manager.RestartRequired) + + def test_retries_confirmation_if_restart_closes_socket_first(self) -> None: + stable: Dict[str, Any] = { + "server_port": 8123, + "trusted_proxies": [], + } + connection = FakeConnection( + [ + _HttpConfigResponse(stable), + {"success": True, "result": {"restart": True}}, + None, + ] + ) + manager = _CreateManager(connection) + + shouldRetry = manager._UpdateHttpConfigViaApiIfNeeded() + + self.assertTrue(shouldRetry) + self.assertEqual(connection.Messages[-1]["type"], "http/config/promote") + self.assertTrue(connection.ServerRestartExpected) + + def test_lost_configure_response_keeps_restart_recovery_armed(self) -> None: + stable: Dict[str, Any] = { + "server_port": 8123, + "trusted_proxies": [], + } + connection = FakeConnection([_HttpConfigResponse(stable), None]) + manager = _CreateManager(connection) + + shouldRetry = manager._UpdateHttpConfigViaApiIfNeeded() + + self.assertTrue(shouldRetry) + self.assertTrue(connection.ServerRestartExpected) + self.assertTrue(manager._HasPendingHttpConfigOwnedByHomeway()) + self.assertEqual( + [msg["type"] for msg in connection.Messages], + ["http/config", "http/config/configure"], + ) + + def test_configure_api_error_clears_restart_recovery_state(self) -> None: + stable: Dict[str, Any] = { + "server_port": 8123, + "trusted_proxies": [], + } + connection = FakeConnection( + [ + _HttpConfigResponse(stable), + { + "success": False, + "error": {"code": "not_running", "message": "starting"}, + }, + ] + ) + manager = _CreateManager(connection) + + shouldRetry = manager._UpdateHttpConfigViaApiIfNeeded() + + self.assertTrue(shouldRetry) + self.assertFalse(connection.ServerRestartExpected) + self.assertFalse(manager._HasPendingHttpConfigOwnedByHomeway()) + + def test_does_not_confirm_when_update_needs_no_restart(self) -> None: + stable: Dict[str, Any] = { + "server_port": 8123, + "trusted_proxies": [], + } + connection = FakeConnection( + [ + _HttpConfigResponse(stable), + {"success": True, "result": {"restart": False}}, + ] + ) + manager = _CreateManager(connection) + + shouldRetry = manager._UpdateHttpConfigViaApiIfNeeded() + + self.assertFalse(shouldRetry) + self.assertEqual( + [msg["type"] for msg in connection.Messages], + ["http/config", "http/config/configure"], + ) + self.assertFalse(connection.ServerRestartExpected) + + def test_promotes_only_matching_pending_config_after_reconnect(self) -> None: + pending: Dict[str, Any] = { + "server_port": 8123, + "use_x_forwarded_for": True, + "trusted_proxies": ["172.30.32.0/23", "127.0.0.1/32", "::1/128"], + "created_at": "2026-08-14T00:00:00+00:00", + "error": None, + "error_message": None, + } + connection = FakeConnection( + [ + _HttpConfigResponse( + pending, activeConfigType="pending", pending=pending + ), + {"success": True, "result": None}, + ] + ) + manager = _CreateManager(connection) + manager._SetPendingHttpConfigOwnedByHomeway( + { + "server_port": 8123, + "use_x_forwarded_for": True, + "trusted_proxies": ["172.30.32.0/23", "127.0.0.1", "::1"], + } + ) + + shouldRetry = manager._UpdateHttpConfigViaApiIfNeeded() + + self.assertFalse(shouldRetry) + self.assertEqual( + [msg["type"] for msg in connection.Messages], + ["http/config", "http/config/promote"], + ) + self.assertFalse(manager._HasPendingHttpConfigOwnedByHomeway()) + self.assertFalse(manager.RestartRequired) + + def test_leaves_unowned_pending_config_for_user_confirmation(self) -> None: + pending: Dict[str, Any] = { + "server_port": 8443, + "use_x_forwarded_for": True, + "trusted_proxies": ["172.30.32.0/23", "127.0.0.1/32", "::1/128"], + "created_at": "2026-08-14T00:00:00+00:00", + "error": None, + "error_message": None, + } + connection = FakeConnection( + [ + _HttpConfigResponse( + pending, activeConfigType="pending", pending=pending + ) + ] + ) + manager = _CreateManager(connection) + + shouldRetry = manager._UpdateHttpConfigViaApiIfNeeded() + + self.assertFalse(shouldRetry) + self.assertEqual( + [msg["type"] for msg in connection.Messages], ["http/config"] + ) + + def test_does_not_overwrite_unowned_config_waiting_for_restart(self) -> None: + stable: Dict[str, Any] = { + "server_port": 8123, + "trusted_proxies": [], + } + pending: Dict[str, Any] = { + "server_port": 8443, + "trusted_proxies": [], + "error": None, + } + connection = FakeConnection( + [_HttpConfigResponse(stable, activeConfigType="stable", pending=pending)] + ) + manager = _CreateManager(connection) + manager._SetPendingHttpConfigOwnedByHomeway( + { + "server_port": 8123, + "use_x_forwarded_for": True, + "trusted_proxies": ["172.30.32.0/23"], + } + ) + connection.SetServerRestartExpected() + + shouldRetry = manager._UpdateHttpConfigViaApiIfNeeded() + + self.assertFalse(shouldRetry) + self.assertEqual( + [msg["type"] for msg in connection.Messages], ["http/config"] + ) + self.assertFalse(manager._HasPendingHttpConfigOwnedByHomeway()) + self.assertFalse(connection.ServerRestartExpected) + + def test_waits_until_owned_pending_config_is_active(self) -> None: + stable: Dict[str, Any] = { + "server_port": 8123, + "trusted_proxies": [], + } + pending: Dict[str, Any] = { + "server_port": 8123, + "use_x_forwarded_for": True, + "trusted_proxies": ["172.30.32.0/23", "127.0.0.1/32", "::1/128"], + "error": None, + } + connection = FakeConnection( + [_HttpConfigResponse(stable, activeConfigType="stable", pending=pending)] + ) + manager = _CreateManager(connection) + manager._SetPendingHttpConfigOwnedByHomeway( + { + "server_port": 8123, + "use_x_forwarded_for": True, + "trusted_proxies": ["172.30.32.0/23", "127.0.0.1", "::1"], + } + ) + + shouldRetry = manager._UpdateHttpConfigViaApiIfNeeded() + + self.assertTrue(shouldRetry) + self.assertEqual( + [msg["type"] for msg in connection.Messages], ["http/config"] + ) + self.assertTrue(manager._HasPendingHttpConfigOwnedByHomeway()) + self.assertTrue(manager.RestartRequired) + + def test_rejected_owned_pending_config_is_not_promoted(self) -> None: + stable: Dict[str, Any] = { + "server_port": 8123, + "trusted_proxies": [], + } + pending: Dict[str, Any] = { + "server_port": 8123, + "use_x_forwarded_for": True, + "trusted_proxies": ["172.30.32.0/23", "127.0.0.1/32", "::1/128"], + "error": "bind_failed", + "error_message": "Address is already in use", + } + connection = FakeConnection( + [_HttpConfigResponse(stable, activeConfigType="stable", pending=pending)] + ) + manager = _CreateManager(connection) + manager._SetPendingHttpConfigOwnedByHomeway( + { + "server_port": 8123, + "use_x_forwarded_for": True, + "trusted_proxies": ["172.30.32.0/23", "127.0.0.1", "::1"], + } + ) + connection.SetServerRestartExpected() + + shouldRetry = manager._UpdateHttpConfigViaApiIfNeeded() + + self.assertFalse(shouldRetry) + self.assertEqual( + [msg["type"] for msg in connection.Messages], ["http/config"] + ) + self.assertFalse(manager._HasPendingHttpConfigOwnedByHomeway()) + self.assertFalse(connection.ServerRestartExpected) + self.assertTrue(manager.RestartRequired) + + def test_reconnect_wakes_worker_without_blindly_promoting(self) -> None: + connection = FakeConnection([]) + manager = _CreateManager(connection) + manager.HttpConfigUpdateThreadRunning = True + manager._SetPendingHttpConfigOwnedByHomeway( + { + "server_port": 8123, + "use_x_forwarded_for": True, + "trusted_proxies": ["172.30.32.0/23"], + } + ) + + manager._OnHaConnected() + + self.assertEqual(connection.Messages, []) + self.assertTrue(manager.HttpConfigUpdateRequested) + self.assertTrue(manager.HttpConfigUpdateWakeEvent.is_set()) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_ha_connection.py b/tests/test_ha_connection.py new file mode 100644 index 0000000..59215e6 --- /dev/null +++ b/tests/test_ha_connection.py @@ -0,0 +1,157 @@ +"""Tests for the persistent Home Assistant WebSocket connection.""" + +# pylint: disable=import-error,no-name-in-module,protected-access + +import logging +import threading +import unittest +from typing import Any, Dict, List, Optional, cast +from unittest.mock import patch + +from homeway_linuxhost.ha.connection import Connection, PendingContexts +from homeway_linuxhost.ha.eventhandler import EventHandler +from homeway.buffer import Buffer +from homeway.interfaces import IWebSocketClient, WebSocketOpCode + + +class FakeWebSocket: + """WebSocket fake that records when a message is sent.""" + + def __init__(self) -> None: + self.MessageSent = threading.Event() + + def Send( + self, + buffer: Buffer, + msgStartOffsetBytes: Optional[int] = None, + msgSize: Optional[int] = None, + isData: bool = True, + ) -> None: + del buffer, msgStartOffsetBytes, msgSize, isData + self.MessageSent.set() + + +class HomeAssistantConnectionTests(unittest.TestCase): + """Exercise connection teardown behavior.""" + + def test_close_wakes_pending_send_and_receive(self) -> None: + connection = Connection( + logging.getLogger("test_ha_connection"), cast(EventHandler, object()) + ) + websocket = FakeWebSocket() + connection.Ws = cast(Any, websocket) + connection.IsConnected = True + responses: List[Optional[Dict[str, Any]]] = [] + + worker = threading.Thread( + target=lambda: responses.append( + connection.SendAndReceiveMsg( + {"type": "test/pending"}, timeoutSec=5.0 + ) + ), + daemon=True, + ) + worker.start() + self.assertTrue(websocket.MessageSent.wait(1.0)) + + connection.Closed(cast(IWebSocketClient, websocket)) + worker.join(1.0) + + self.assertFalse(worker.is_alive()) + self.assertEqual(responses, [None]) + self.assertFalse(connection.IsConnected) + self.assertIsNone(connection.Ws) + self.assertEqual(connection.PendingContexts, {}) + + def test_stale_close_does_not_reset_current_connection(self) -> None: + connection = Connection( + logging.getLogger("test_ha_connection"), cast(EventHandler, object()) + ) + currentWebsocket = FakeWebSocket() + oldWebsocket = FakeWebSocket() + connection.Ws = cast(Any, currentWebsocket) + connection.IsConnected = True + + connection.Closed(cast(IWebSocketClient, oldWebsocket)) + + self.assertTrue(connection.IsConnected) + self.assertIs(connection.Ws, currentWebsocket) + + def test_stale_data_does_not_authenticate_current_connection(self) -> None: + connection = Connection( + logging.getLogger("test_ha_connection"), cast(EventHandler, object()) + ) + currentWebsocket = FakeWebSocket() + oldWebsocket = FakeWebSocket() + connection.Ws = cast(Any, currentWebsocket) + + connection._OnData( + cast(IWebSocketClient, oldWebsocket), + Buffer(b'{"type":"auth_ok"}'), + WebSocketOpCode.TEXT, + ) + + self.assertFalse(connection.IsConnected) + + def test_result_that_becomes_stale_cannot_complete_a_new_request(self) -> None: + connection = Connection( + logging.getLogger("test_ha_connection"), cast(EventHandler, object()) + ) + oldWebsocket = FakeWebSocket() + currentWebsocket = FakeWebSocket() + connection.Ws = cast(Any, oldWebsocket) + connection.IsConnected = True + pendingContext = PendingContexts() + connection.PendingContexts[1] = pendingContext + + class SwitchingBuffer: + """Switch connections after the callback's initial identity check.""" + + def GetBytesLike(self) -> bytes: + with connection.PendingContextsLock: + connection.Ws = cast(Any, currentWebsocket) + return b'{"id":1,"type":"result","success":true,"result":{}}' + + connection._OnData( + cast(IWebSocketClient, oldWebsocket), + cast(Buffer, SwitchingBuffer()), + WebSocketOpCode.TEXT, + ) + + self.assertFalse(pendingContext.Event.is_set()) + self.assertIsNone(pendingContext.Response) + + def test_auth_keeps_expected_restart_armed_until_config_is_resolved(self) -> None: + connection = Connection( + logging.getLogger("test_ha_connection"), cast(EventHandler, object()) + ) + websocket = FakeWebSocket() + connection.Ws = cast(Any, websocket) + connection.SetServerRestartExpected() + + with patch.object(connection, "_OnConnected"): + connection._OnData( + cast(IWebSocketClient, websocket), + Buffer(b'{"type":"auth_ok"}'), + WebSocketOpCode.TEXT, + ) + + self.assertTrue(connection.IsConnected) + self.assertTrue(connection.ServerRestartExpected.is_set()) + + def test_expected_restart_reconnect_state_has_a_bounded_budget(self) -> None: + connection = Connection( + logging.getLogger("test_ha_connection"), cast(EventHandler, object()) + ) + connection.SetServerRestartExpected() + + for _ in range(Connection.c_ExpectedServerRestartReconnectMaxAttempts): + self.assertTrue(connection._ShouldUseExpectedServerRestartReconnect()) + + self.assertFalse(connection._ShouldUseExpectedServerRestartReconnect()) + self.assertFalse(connection.ServerRestartExpected.is_set()) + self.assertEqual(connection.ServerRestartReconnectAttempt, 0) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_installer_imports.py b/tests/test_installer_imports.py new file mode 100644 index 0000000..9be20d0 --- /dev/null +++ b/tests/test_installer_imports.py @@ -0,0 +1,48 @@ +"""Smoke tests for the installer's and service's different Python package roots.""" + +import importlib.util +import os +from pathlib import Path +import subprocess +import sys +import unittest + + +class InstallerImportTests(unittest.TestCase): + """Import entry point dependencies without starting installation or services.""" + + def _CheckImports(self, workingDirectory: Path, code: str) -> None: + env = os.environ.copy() + # Match each launcher's package root, regardless of the test runner's path. + env["PYTHONPATH"] = str(workingDirectory) + result = subprocess.run( + [sys.executable, "-B", "-c", code], + cwd=workingDirectory, + env=env, + capture_output=True, + text=True, + timeout=30, + check=False, + ) + self.assertEqual(result.returncode, 0, result.stdout + result.stderr) + + @unittest.skipUnless(importlib.util.find_spec("pwd"), "Installer requires Unix") + def test_installer_imports_from_repository_root(self) -> None: + """Install and update must load before any installer actions can run.""" + self._CheckImports( + Path(__file__).resolve().parents[1], + "from homeway_installer.Installer import Installer", + ) + + def test_service_imports_from_addon_root(self) -> None: + """The service uses the inner homeway directory as its package root.""" + self._CheckImports( + Path(__file__).resolve().parents[1] / "homeway", + "from homeway.interfaces import IWebStreamHelper; " + "from homeway.websocketimpl import Client; " + "from homeway.Proto.WebStreamMsg import WebStreamMsg", + ) + + +if __name__ == "__main__": + unittest.main()