diff --git a/kasa/transports/sslaestransport.py b/kasa/transports/sslaestransport.py index a517ca48..ba7e8322 100644 --- a/kasa/transports/sslaestransport.py +++ b/kasa/transports/sslaestransport.py @@ -274,6 +274,20 @@ class SslAesTransport(BaseTransport): _LOGGER.debug(msg) raise _RetryableError(msg) + # Some devices answer 401 when the session has expired and they + # require a new handshake: reauthenticate and retry the request. + if status_code == 401: + _LOGGER.debug( + "Device %s replied with status 401 to passthrough, " + "session expired, handshake required", + self._host, + ) + self._state = TransportState.HANDSHAKE_REQUIRED + raise _RetryableError( + f"{self._host} responded with status 401 to passthrough, " + "session expired" + ) + if status_code != 200: raise KasaException( f"{self._host} responded with an unexpected " diff --git a/tests/transports/test_sslaestransport.py b/tests/transports/test_sslaestransport.py index eeaebfea..3142a2f0 100644 --- a/tests/transports/test_sslaestransport.py +++ b/tests/transports/test_sslaestransport.py @@ -797,3 +797,24 @@ class MockSslAesDevice: def put_next_response(self, request: dict | bytes) -> None: self._next_responses.append(request) + + +async def test_passthrough_401_requires_new_handshake(mocker): + """A 401 on passthrough means the session expired: retryable, new handshake.""" + host = "127.0.0.1" + mock_ssl_aes_device = MockSslAesDevice(host) + mocker.patch.object( + aiohttp.ClientSession, "post", side_effect=mock_ssl_aes_device.post + ) + transport = SslAesTransport( + config=DeviceConfig(host, credentials=Credentials(MOCK_USER, MOCK_PWD)) + ) + request = {"method": "getDeviceInfo", "params": None} + + await transport.perform_handshake() + assert transport._state is TransportState.ESTABLISHED + + mock_ssl_aes_device.status_code = 401 + with pytest.raises(_RetryableError, match="session expired"): + await transport.send(json_dumps(request)) + assert transport._state is TransportState.HANDSHAKE_REQUIRED