diff --git a/kasa/httpclient.py b/kasa/httpclient.py index 31d8dfbb..0a100d70 100644 --- a/kasa/httpclient.py +++ b/kasa/httpclient.py @@ -3,6 +3,7 @@ from __future__ import annotations import asyncio +import builtins import logging import ssl import time @@ -144,7 +145,9 @@ class HttpClient: raise _ConnectionError( f"Device connection error: {self._config.host}: {ex}", ex ) from ex - except (aiohttp.ServerTimeoutError, TimeoutError) as ex: + # TimeoutError is imported from kasa.exceptions and shadows the builtin, + # which is what aiohttp raises when ClientTimeout(total=...) expires. + except (aiohttp.ServerTimeoutError, builtins.TimeoutError) as ex: raise TimeoutError( "Unable to query the device, " + f"timed out: {self._config.host}: {ex}", diff --git a/kasa/protocols/smartprotocol.py b/kasa/protocols/smartprotocol.py index ad3e7331..364e795c 100644 --- a/kasa/protocols/smartprotocol.py +++ b/kasa/protocols/smartprotocol.py @@ -245,7 +245,16 @@ class SmartProtocol(BaseProtocol): batch_name, pf(smart_request), ) - response_step = await self._transport.send(smart_request) + try: + response_step = await self._transport.send(smart_request) + except TimeoutError as ex: + # P300 does not respond to some batched requests (e.g. both child + # list requests together) so disable batching. Batch size is never + # 1 here as that case sends single requests above. + self._multi_request_batch_size = 1 + raise _RetryableError( + "Timeout during multi request, multi requests disabled" + ) from ex if debug_enabled: if self._redact_data: data = redact_data(response_step, REDACTORS) diff --git a/tests/protocols/test_smartprotocol.py b/tests/protocols/test_smartprotocol.py index 9ccbc3cb..608be537 100644 --- a/tests/protocols/test_smartprotocol.py +++ b/tests/protocols/test_smartprotocol.py @@ -1,3 +1,4 @@ +import json import logging import pytest @@ -9,6 +10,7 @@ from kasa.exceptions import ( DeviceError, KasaException, SmartErrorCode, + TimeoutError, ) from kasa.protocols.smartcamprotocol import SmartCamProtocol from kasa.protocols.smartprotocol import SmartProtocol, _ChildProtocolWrapper @@ -235,6 +237,102 @@ async def test_smart_device_multiple_request_non_json_decode_failure( assert send_mock.call_count == 1 +async def test_smart_device_multiple_request_timeout( + dummy_protocol: SmartProtocol, mocker: MockerFixture +) -> None: + """Test that a timed out multi request disables batching and retries.""" + requests = {} + mock_responses = [] + for i in range(10): + method = f"get_method_{i}" + requests[method] = {"foo": "bar", "bar": "foo"} + mock_responses.append( + {"method": method, "result": {"great": "success"}, "error_code": 0} + ) + + send_mock = mocker.patch.object( + dummy_protocol._transport, + "send", + side_effect=[TimeoutError("Simulated timeout"), *mock_responses], + ) + mocker.patch("asyncio.sleep") + dummy_protocol._multi_request_batch_size = 5 + resp = await dummy_protocol.query(requests, retry_count=1) + assert dummy_protocol._multi_request_batch_size == 1 + assert resp == {method: {"great": "success"} for method in requests} + # Call count should be the timed out batch + number of requests + assert send_mock.call_count == len(requests) + 1 + + +async def test_smart_device_multiple_request_timeout_single_requests( + dummy_protocol: SmartProtocol, mocker: MockerFixture +) -> None: + """Test that timeouts on single requests are still retried and raised.""" + send_mock = mocker.patch.object( + dummy_protocol._transport, + "send", + side_effect=TimeoutError("Simulated timeout"), + ) + mocker.patch("asyncio.sleep") + dummy_protocol._multi_request_batch_size = 1 + with pytest.raises(TimeoutError): + await dummy_protocol.query(DUMMY_MULTIPLE_QUERY, retry_count=2) + assert dummy_protocol._multi_request_batch_size == 1 + assert send_mock.call_count == 3 + + +async def test_smart_device_multiple_request_child_lists_timeout( + dummy_protocol: SmartProtocol, mocker: MockerFixture +) -> None: + """Test devices that do not respond to batched child list requests. + + The P300(EU) 1.0.7 firmware never answers a multipleRequest containing both + get_child_device_component_list and get_child_device_list, while each + request sent on its own succeeds. + """ + child_list_methods = {"get_child_device_component_list", "get_child_device_list"} + results = { + "get_child_device_component_list": { + "child_component_list": [], + "start_index": 0, + "sum": 0, + }, + "get_child_device_list": { + "child_device_list": [], + "start_index": 0, + "sum": 0, + }, + } + + async def _send(request: str) -> dict: + req = json.loads(request) + if req["method"] != "multipleRequest": + return {"result": results[req["method"]], "error_code": 0} + methods = {r["method"] for r in req["params"]["requests"]} + if child_list_methods <= methods: + raise TimeoutError("Simulated timeout") + return { + "result": { + "responses": [ + {"method": m, "result": results[m], "error_code": 0} + for m in methods + ] + }, + "error_code": 0, + } + + send_mock = mocker.patch.object( + dummy_protocol._transport, "send", side_effect=_send + ) + mocker.patch("asyncio.sleep") + assert dummy_protocol._multi_request_batch_size == 5 + resp = await dummy_protocol.query(dict.fromkeys(child_list_methods)) + assert resp == results + assert dummy_protocol._multi_request_batch_size == 1 + # The timed out batch + one single request per method + assert send_mock.call_count == 3 + + async def test_childdevicewrapper_unwrapping( dummy_protocol: SmartProtocol, mocker: MockerFixture ) -> None: diff --git a/tests/test_httpclient.py b/tests/test_httpclient.py index 906b39ed..82115ce9 100644 --- a/tests/test_httpclient.py +++ b/tests/test_httpclient.py @@ -1,3 +1,4 @@ +import builtins import re import aiohttp @@ -35,6 +36,13 @@ from kasa.httpclient import HttpClient TimeoutError, "Unable to query the device, timed out: ", ), + # aiohttp raises the builtin TimeoutError when ClientTimeout(total=...) + # expires after the connection is established. + ( + builtins.TimeoutError(), + TimeoutError, + "Unable to query the device, timed out: ", + ), (Exception(), KasaException, "Unable to query the device: "), ( aiohttp.ServerFingerprintMismatch(b"exp", b"got", "host", 1), @@ -47,6 +55,7 @@ from kasa.httpclient import HttpClient "ClientOSError", "ServerTimeoutError", "TimeoutError", + "BuiltinTimeoutError", "Exception", "ServerFingerprintMismatch", ),