tests: add type annotations to top-level test files (#1688)

Add type annotations (parameter types and `-> None` return types) to all
top-level test functions across 22 test files. This enables mypy to
check test function bodies, catching type errors that were previously
hidden.
This commit is contained in:
ZeliardM
2026-10-04 10:36:01 -04:00
committed by GitHub
parent eed64fe013
commit b36c795cbc
22 changed files with 557 additions and 392 deletions

View File

@@ -15,6 +15,7 @@ import pytest
from asyncclick.testing import CliRunner from asyncclick.testing import CliRunner
from kasa import ( from kasa import (
Device,
DeviceConfig, DeviceConfig,
SmartProtocol, SmartProtocol,
) )
@@ -68,14 +69,14 @@ async def _close_transport_and_http_sessions(monkeypatch):
await session.close() await session.close()
def load_fixture(foldername, filename): def load_fixture(foldername: str, filename: str) -> str:
"""Load a fixture.""" """Load a fixture."""
path = Path(Path(__file__).parent / "fixtures" / foldername / filename) path = Path(Path(__file__).parent / "fixtures" / foldername / filename)
with path.open() as fdp: with path.open() as fdp:
return fdp.read() return fdp.read()
async def handle_turn_on(dev, turn_on): async def handle_turn_on(dev: Device, turn_on: bool) -> None:
if turn_on: if turn_on:
await dev.turn_on() await dev.turn_on()
else: else:
@@ -110,15 +111,16 @@ def dummy_protocol():
yield protocol yield protocol
def pytest_configure(): def pytest_configure() -> None:
pytest.fixtures_missing_methods = {} pytest.fixtures_missing_methods = {} # type: ignore[attr-defined]
def pytest_sessionfinish(session, exitstatus): def pytest_sessionfinish(session: pytest.Session, exitstatus: int) -> None:
if not pytest.fixtures_missing_methods: fixtures_missing: dict = getattr(pytest, "fixtures_missing_methods", {})
if not fixtures_missing:
return return
msg = "\n" msg = "\n"
for fixture, methods in sorted(pytest.fixtures_missing_methods.items()): for fixture, methods in sorted(fixtures_missing.items()):
method_list = ", ".join(methods) method_list = ", ".join(methods)
msg += f"Fixture {fixture} missing: {method_list}\n" msg += f"Fixture {fixture} missing: {method_list}\n"
@@ -128,7 +130,7 @@ def pytest_sessionfinish(session, exitstatus):
) )
def pytest_addoption(parser): def pytest_addoption(parser: pytest.Parser) -> None:
parser.addoption( parser.addoption(
"--ip", action="store", default=None, help="run against device on given ip" "--ip", action="store", default=None, help="run against device on given ip"
) )
@@ -140,7 +142,9 @@ def pytest_addoption(parser):
) )
def pytest_collection_modifyitems(config, items): def pytest_collection_modifyitems(
config: pytest.Config, items: list[pytest.Item]
) -> None:
if not config.getoption("--ip"): if not config.getoption("--ip"):
print("Testing against fixtures.") print("Testing against fixtures.")
# pytest_socket doesn't work properly in windows with asyncio # pytest_socket doesn't work properly in windows with asyncio
@@ -161,11 +165,11 @@ def pytest_collection_modifyitems(config, items):
@pytest.fixture(autouse=True, scope="session") @pytest.fixture(autouse=True, scope="session")
def asyncio_sleep_fixture(request): # noqa: PT004 def asyncio_sleep_fixture(request: pytest.FixtureRequest): # noqa: PT004
"""Patch sleep to prevent tests actually waiting.""" """Patch sleep to prevent tests actually waiting."""
orig_asyncio_sleep = asyncio.sleep orig_asyncio_sleep = asyncio.sleep
async def _asyncio_sleep(*_, **__): async def _asyncio_sleep(*_, **__) -> None:
await orig_asyncio_sleep(0) await orig_asyncio_sleep(0)
if request.config.getoption("--ip"): if request.config.getoption("--ip"):
@@ -176,7 +180,7 @@ def asyncio_sleep_fixture(request): # noqa: PT004
@pytest.fixture(autouse=True, scope="session") @pytest.fixture(autouse=True, scope="session")
def mock_datagram_endpoint(request): # noqa: PT004 def mock_datagram_endpoint(request: pytest.FixtureRequest):
"""Mock create_datagram_endpoint so it doesn't perform io.""" """Mock create_datagram_endpoint so it doesn't perform io."""
async def _create_datagram_endpoint(protocol_factory, *_, **__): async def _create_datagram_endpoint(protocol_factory, *_, **__):

View File

@@ -215,7 +215,7 @@ def parametrize_subtract(params: pytest.MarkDecorator, subtract: pytest.MarkDeco
def parametrize( def parametrize(
desc, desc: str,
*, *,
model_filter=None, model_filter=None,
protocol_filter=None, protocol_filter=None,
@@ -395,7 +395,7 @@ chime_smart = parametrize(
vacuum = parametrize("vacuums", device_type_filter=[DeviceType.Vacuum]) vacuum = parametrize("vacuums", device_type_filter=[DeviceType.Vacuum])
def check_categories(): def check_categories() -> None:
"""Check that every fixture file is categorized.""" """Check that every fixture file is categorized."""
categorized_fixtures = set( categorized_fixtures = set(
dimmer_iot.args[1] dimmer_iot.args[1]
@@ -428,7 +428,7 @@ def check_categories():
check_categories() check_categories()
def device_for_fixture_name(model, protocol): def device_for_fixture_name(model: str, protocol: str):
if protocol in {"SMART", "SMART.CHILD"}: if protocol in {"SMART", "SMART.CHILD"}:
return SmartDevice return SmartDevice
elif protocol in {"SMARTCAM", "SMARTCAM.CHILD"}: elif protocol in {"SMARTCAM", "SMARTCAM.CHILD"}:
@@ -467,7 +467,9 @@ async def _update_and_close(d) -> Device:
return d return d
async def _discover_update_and_close(ip, username, password) -> Device: async def _discover_update_and_close(
ip: str, username: str | None, password: str | None
) -> Device:
if username and password: if username and password:
credentials = Credentials(username=username, password=password) credentials = Credentials(username=username, password=password)
else: else:
@@ -477,7 +479,7 @@ async def _discover_update_and_close(ip, username, password) -> Device:
async def get_device_for_fixture( async def get_device_for_fixture(
fixture_data: FixtureInfo, *, verbatim=False, update_after_init=True fixture_data: FixtureInfo, *, verbatim: bool = False, update_after_init: bool = True
) -> Device: ) -> Device:
# if the wanted file is not an absolute path, prepend the fixtures directory # if the wanted file is not an absolute path, prepend the fixtures directory
@@ -502,7 +504,9 @@ async def get_device_for_fixture(
fixture_data.data, fixture_data.name, verbatim=verbatim fixture_data.data, fixture_data.name, verbatim=verbatim
) )
else: else:
d.protocol = FakeIotProtocol(fixture_data.data, verbatim=verbatim) d.protocol = FakeIotProtocol(
fixture_data.data, fixture_data.name, verbatim=verbatim
)
discovery_data = None discovery_data = None
if "discovery_result" in fixture_data.data: if "discovery_result" in fixture_data.data:
@@ -520,21 +524,21 @@ async def get_device_for_fixture(
return d return d
async def get_device_for_fixture_protocol(fixture, protocol): async def get_device_for_fixture_protocol(fixture: str, protocol: str):
finfo = FixtureInfo(name=fixture, protocol=protocol, data={}) finfo = FixtureInfo(name=fixture, protocol=protocol, data={})
for fixture_info in FIXTURE_DATA: for fixture_info in FIXTURE_DATA:
if finfo == fixture_info: if finfo == fixture_info:
return await get_device_for_fixture(fixture_info) return await get_device_for_fixture(fixture_info)
def get_fixture_info(fixture, protocol): def get_fixture_info(fixture: str, protocol: str):
finfo = FixtureInfo(name=fixture, protocol=protocol, data={}) finfo = FixtureInfo(name=fixture, protocol=protocol, data={})
for fixture_info in FIXTURE_DATA: for fixture_info in FIXTURE_DATA:
if finfo == fixture_info: if finfo == fixture_info:
return fixture_info return fixture_info
def get_nearest_fixture_to_ip(dev): def get_nearest_fixture_to_ip(dev: Device):
if isinstance(dev, SmartDevice): if isinstance(dev, SmartDevice):
protocol_fixtures = filter_fixtures("", protocol_filter={"SMART"}) protocol_fixtures = filter_fixtures("", protocol_filter={"SMART"})
elif isinstance(dev, SmartCamDevice): elif isinstance(dev, SmartCamDevice):
@@ -572,7 +576,7 @@ def get_nearest_fixture_to_ip(dev):
@pytest.fixture(params=filter_fixtures("main devices"), ids=idgenerator) @pytest.fixture(params=filter_fixtures("main devices"), ids=idgenerator)
async def dev(request) -> AsyncGenerator[Device, None]: async def dev(request: pytest.FixtureRequest) -> AsyncGenerator[Device, None]:
"""Device fixture. """Device fixture.
Provides a device (given --ip) or parametrized fixture for the supported devices. Provides a device (given --ip) or parametrized fixture for the supported devices.

View File

@@ -8,6 +8,7 @@ from json import dumps as json_dumps
from typing import Any, TypedDict from typing import Any, TypedDict
import pytest import pytest
from pytest_mock import MockerFixture
from kasa.transports.xortransport import XorEncryption from kasa.transports.xortransport import XorEncryption
@@ -48,8 +49,8 @@ UNSUPPORTED_HOMEWIFISYSTEM = {
def _make_unsupported( def _make_unsupported(
device_family, device_family: str,
encrypt_type, encrypt_type: str,
*, *,
https: bool = False, https: bool = False,
omit_keys: dict[str, Any] | None = None, omit_keys: dict[str, Any] | None = None,
@@ -112,7 +113,7 @@ UNSUPPORTED_DEVICES = {
def parametrize_discovery( def parametrize_discovery(
desc, *, data_root_filter=None, protocol_filter=None, model_filter=None desc: str, *, data_root_filter=None, protocol_filter=None, model_filter=None
): ):
filtered_fixtures = filter_fixtures( filtered_fixtures = filter_fixtures(
desc, desc,
@@ -141,7 +142,7 @@ smart_discovery = parametrize_discovery("smart discovery", protocol_filter={"SMA
), ),
ids=idgenerator, ids=idgenerator,
) )
async def discovery_mock(request, mocker): async def discovery_mock(request: pytest.FixtureRequest, mocker: MockerFixture):
"""Mock discovery and patch protocol queries to use Fake protocols.""" """Mock discovery and patch protocol queries to use Fake protocols."""
fi: FixtureInfo = request.param fi: FixtureInfo = request.param
fixture_info = FixtureInfo(fi.name, fi.protocol, copy.deepcopy(fi.data)) fixture_info = FixtureInfo(fi.name, fi.protocol, copy.deepcopy(fi.data))
@@ -238,7 +239,7 @@ def create_discovery_mock(ip: str, fixture_data: dict):
return dm return dm
def patch_discovery(fixture_infos: dict[str, FixtureInfo], mocker): def patch_discovery(fixture_infos: dict[str, FixtureInfo], mocker: MockerFixture):
"""Mock discovery and patch protocol queries to use Fake protocols.""" """Mock discovery and patch protocol queries to use Fake protocols."""
discovery_mocks = { discovery_mocks = {
ip: create_discovery_mock(ip, fixture_info.data) ip: create_discovery_mock(ip, fixture_info.data)
@@ -271,7 +272,7 @@ def patch_discovery(fixture_infos: dict[str, FixtureInfo], mocker):
await exception_queue.put(None) await exception_queue.put(None)
callback_queue.task_done() callback_queue.task_done()
async def wait_for_coro(): async def wait_for_coro() -> None:
await callback_queue.join() await callback_queue.join()
if ex := exception_queue.get_nowait(): if ex := exception_queue.get_nowait():
raise ex raise ex
@@ -286,7 +287,7 @@ def patch_discovery(fixture_infos: dict[str, FixtureInfo], mocker):
) )
# do_discover_mock # do_discover_mock
async def mock_discover(self): async def mock_discover(self) -> None:
"""Call datagram_received for all mock fixtures. """Call datagram_received for all mock fixtures.
Handles test cases modifying the ip and hostname of the first fixture Handles test cases modifying the ip and hostname of the first fixture
@@ -333,7 +334,7 @@ def patch_discovery(fixture_infos: dict[str, FixtureInfo], mocker):
mocker.patch("kasa.IotProtocol.query", _query) mocker.patch("kasa.IotProtocol.query", _query)
mocker.patch("kasa.SmartProtocol.query", _query) mocker.patch("kasa.SmartProtocol.query", _query)
def _getaddrinfo(host, *_, **__): def _getaddrinfo(host: str, *_, **__):
nonlocal first_host, first_ip nonlocal first_host, first_ip
first_host = host # Store the hostname used by discover single first_host = host # Store the hostname used by discover single
first_ip = list(discovery_mocks.values())[ first_ip = list(discovery_mocks.values())[
@@ -358,7 +359,7 @@ def patch_discovery(fixture_infos: dict[str, FixtureInfo], mocker):
), ),
ids=idgenerator, ids=idgenerator,
) )
def discovery_data(request, mocker): def discovery_data(request: pytest.FixtureRequest, mocker: MockerFixture):
"""Return raw discovery file contents as JSON. Used for discovery tests.""" """Return raw discovery file contents as JSON. Used for discovery tests."""
fixture_info = request.param fixture_info = request.param
fixture_data = copy.deepcopy(fixture_info.data) fixture_data = copy.deepcopy(fixture_info.data)
@@ -383,12 +384,12 @@ def discovery_data(request, mocker):
@pytest.fixture( @pytest.fixture(
params=UNSUPPORTED_DEVICES.values(), ids=list(UNSUPPORTED_DEVICES.keys()) params=UNSUPPORTED_DEVICES.values(), ids=list(UNSUPPORTED_DEVICES.keys())
) )
def unsupported_device_info(request, mocker): def unsupported_device_info(request: pytest.FixtureRequest, mocker: MockerFixture):
"""Return unsupported devices for cli and discovery tests.""" """Return unsupported devices for cli and discovery tests."""
discovery_data = request.param discovery_data = request.param
host = "127.0.0.1" host = "127.0.0.1"
async def mock_discover(self): async def mock_discover(self) -> None:
if discovery_data: if discovery_data:
data = ( data = (
b"\x02\x00\x00\x01\x01[\x00\x00\x00\x00\x00\x00W\xcev\xf8" b"\x02\x00\x00\x01\x01[\x00\x00\x00\x00\x00\x00W\xcev\xf8"

View File

@@ -1,5 +1,6 @@
import copy import copy
import logging import logging
from typing import Any
from kasa.deviceconfig import DeviceConfig from kasa.deviceconfig import DeviceConfig
from kasa.protocols import IotProtocol from kasa.protocols import IotProtocol
@@ -8,7 +9,7 @@ from kasa.transports.basetransport import BaseTransport
_LOGGER = logging.getLogger(__name__) _LOGGER = logging.getLogger(__name__)
def get_realtime(obj, x, *args): def get_realtime(obj, x: dict, *args):
return { return {
"current": 0.268587, "current": 0.268587,
"voltage": 125.836131, "voltage": 125.836131,
@@ -17,7 +18,7 @@ def get_realtime(obj, x, *args):
} }
def get_monthstat(obj, x, *args): def get_monthstat(obj, x: dict, *args):
if x["year"] < 2016: if x["year"] < 2016:
return {"month_list": []} return {"month_list": []}
@@ -29,7 +30,7 @@ def get_monthstat(obj, x, *args):
} }
def get_daystat(obj, x, *args): def get_daystat(obj, x: dict, *args):
if x["year"] < 2016: if x["year"] < 2016:
return {"day_list": []} return {"day_list": []}
@@ -48,11 +49,11 @@ emeter_support = {
} }
def get_realtime_units(obj, x, *args): def get_realtime_units(obj, x: dict, *args):
return {"power_mw": 10800} return {"power_mw": 10800}
def get_monthstat_units(obj, x, *args): def get_monthstat_units(obj, x: dict, *args):
if x["year"] < 2016: if x["year"] < 2016:
return {"month_list": []} return {"month_list": []}
@@ -64,7 +65,7 @@ def get_monthstat_units(obj, x, *args):
} }
def get_daystat_units(obj, x, *args): def get_daystat_units(obj, x: dict, *args):
if x["year"] < 2016: if x["year"] < 2016:
return {"day_list": []} return {"day_list": []}
@@ -89,11 +90,11 @@ emeter_commands = {
} }
def error(msg="default msg"): def error(msg: str = "default msg"):
return {"err_code": -1323, "msg": msg} return {"err_code": -1323, "msg": msg}
def success(res): def success(res: dict):
if res: if res:
res.update({"err_code": 0}) res.update({"err_code": 0})
else: else:
@@ -223,7 +224,9 @@ DEFAULT_BEHAVIOR = {
class FakeIotProtocol(IotProtocol): class FakeIotProtocol(IotProtocol):
def __init__(self, info, fixture_name=None, *, verbatim=False): def __init__(
self, info: dict[str, Any], fixture_name: str, *, verbatim: bool = False
) -> None:
super().__init__( super().__init__(
transport=FakeIotTransport(info, fixture_name, verbatim=verbatim), transport=FakeIotTransport(info, fixture_name, verbatim=verbatim),
) )
@@ -235,7 +238,9 @@ class FakeIotProtocol(IotProtocol):
class FakeIotTransport(BaseTransport): class FakeIotTransport(BaseTransport):
def __init__(self, info, fixture_name=None, *, verbatim=False): def __init__(
self, info: dict[str, Any], fixture_name: str, *, verbatim: bool = False
) -> None:
super().__init__(config=DeviceConfig("127.0.0.123")) super().__init__(config=DeviceConfig("127.0.0.123"))
info = copy.deepcopy(info) info = copy.deepcopy(info)
self.discovery_data = info self.discovery_data = info
@@ -292,7 +297,7 @@ class FakeIotTransport(BaseTransport):
def credentials_hash(self) -> None: def credentials_hash(self) -> None:
return None return None
def set_alias(self, x, child_ids=None): def set_alias(self, x: dict, child_ids: list | None = None) -> None:
if child_ids is None: if child_ids is None:
child_ids = [] child_ids = []
_LOGGER.debug("Setting alias to %s, child_ids: %s", x["alias"], child_ids) _LOGGER.debug("Setting alias to %s, child_ids: %s", x["alias"], child_ids)
@@ -303,7 +308,7 @@ class FakeIotTransport(BaseTransport):
else: else:
self.proto["system"]["get_sysinfo"]["alias"] = x["alias"] self.proto["system"]["get_sysinfo"]["alias"] = x["alias"]
def set_relay_state(self, x, child_ids=None): def set_relay_state(self, x: dict, child_ids: list | None = None) -> None:
if child_ids is None: if child_ids is None:
child_ids = [] child_ids = []
_LOGGER.debug("Setting relay state to %s", x["state"]) _LOGGER.debug("Setting relay state to %s", x["state"])
@@ -317,19 +322,19 @@ class FakeIotTransport(BaseTransport):
else: else:
self.proto["system"]["get_sysinfo"]["relay_state"] = x["state"] self.proto["system"]["get_sysinfo"]["relay_state"] = x["state"]
def set_led_off(self, x, *args): def set_led_off(self, x: dict, *args) -> None:
_LOGGER.debug("Setting led off to %s", x) _LOGGER.debug("Setting led off to %s", x)
self.proto["system"]["get_sysinfo"]["led_off"] = x["off"] self.proto["system"]["get_sysinfo"]["led_off"] = x["off"]
def set_mac(self, x, *args): def set_mac(self, x: dict, *args) -> None:
_LOGGER.debug("Setting mac to %s", x) _LOGGER.debug("Setting mac to %s", x)
self.proto["system"]["get_sysinfo"]["mac"] = x["mac"] self.proto["system"]["get_sysinfo"]["mac"] = x["mac"]
def set_hs220_brightness(self, x, *args): def set_hs220_brightness(self, x: dict, *args) -> None:
_LOGGER.debug("Setting brightness to %s", x) _LOGGER.debug("Setting brightness to %s", x)
self.proto["system"]["get_sysinfo"]["brightness"] = x["brightness"] self.proto["system"]["get_sysinfo"]["brightness"] = x["brightness"]
def set_hs220_dimmer_transition(self, x, *args): def set_hs220_dimmer_transition(self, x: dict, *args) -> None:
_LOGGER.debug("Setting dimmer transition to %s", x) _LOGGER.debug("Setting dimmer transition to %s", x)
brightness = x["brightness"] brightness = x["brightness"]
if brightness == 0: if brightness == 0:
@@ -338,11 +343,11 @@ class FakeIotTransport(BaseTransport):
self.proto["system"]["get_sysinfo"]["relay_state"] = 1 self.proto["system"]["get_sysinfo"]["relay_state"] = 1
self.proto["system"]["get_sysinfo"]["brightness"] = x["brightness"] self.proto["system"]["get_sysinfo"]["brightness"] = x["brightness"]
def set_lighting_effect(self, effect, *args): def set_lighting_effect(self, effect: dict, *args) -> None:
_LOGGER.debug("Setting light effect to %s", effect) _LOGGER.debug("Setting light effect to %s", effect)
self.proto["system"]["get_sysinfo"]["lighting_effect_state"] = dict(effect) self.proto["system"]["get_sysinfo"]["lighting_effect_state"] = dict(effect)
def transition_light_state(self, state_changes, *args): def transition_light_state(self, state_changes: dict, *args) -> None:
# Setting the light state on a device will turn off any active lighting effects. # Setting the light state on a device will turn off any active lighting effects.
# Unless it's just the brightness in which case it will update the brightness for # Unless it's just the brightness in which case it will update the brightness for
# the lighting effect # the lighting effect
@@ -388,13 +393,13 @@ class FakeIotTransport(BaseTransport):
_LOGGER.debug("New light state: %s", new_state) _LOGGER.debug("New light state: %s", new_state)
self.proto["system"]["get_sysinfo"]["light_state"] = new_state self.proto["system"]["get_sysinfo"]["light_state"] = new_state
def set_preferred_state(self, new_state, *args): def set_preferred_state(self, new_state: dict, *args) -> None:
"""Implement set_preferred_state.""" """Implement set_preferred_state."""
self.proto["system"]["get_sysinfo"]["preferred_state"][new_state["index"]] = ( self.proto["system"]["get_sysinfo"]["preferred_state"][new_state["index"]] = (
new_state new_state
) )
def light_state(self, x, *args): def light_state(self, x: dict, *args):
light_state = self.proto["system"]["get_sysinfo"]["light_state"] light_state = self.proto["system"]["get_sysinfo"]["light_state"]
# Our tests have light state off, so we simply return the dft_on_state when device is on. # Our tests have light state off, so we simply return the dft_on_state when device is on.
_LOGGER.debug("reporting light state: %s", light_state) _LOGGER.debug("reporting light state: %s", light_state)
@@ -404,7 +409,7 @@ class FakeIotTransport(BaseTransport):
else: else:
return light_state return light_state
def set_time(self, new_state: dict, *args): def set_time(self, new_state: dict, *args) -> None:
"""Implement set_time.""" """Implement set_time."""
mods = [ mods = [
v v
@@ -478,7 +483,7 @@ class FakeIotTransport(BaseTransport):
"smartlife.iot.common.schedule": SCHEDULE_MODULE, "smartlife.iot.common.schedule": SCHEDULE_MODULE,
} }
async def send(self, request, port=9999): async def send(self, request, port: int | None = 9999):
if not self.verbatim: if not self.verbatim:
return await self._send(request, port) return await self._send(request, port)
@@ -491,7 +496,7 @@ class FakeIotTransport(BaseTransport):
response.update({"err_msg": "module not support"}) response.update({"err_msg": "module not support"})
return copy.deepcopy(response) return copy.deepcopy(response)
async def _send(self, request, port=9999): async def _send(self, request, port: int | None = 9999):
proto = self.proto proto = self.proto
# collect child ids from context # collect child ids from context
try: try:
@@ -500,13 +505,13 @@ class FakeIotTransport(BaseTransport):
except KeyError: except KeyError:
child_ids = [] child_ids = []
def get_response_for_module(target): def get_response_for_module(target: str):
if target not in proto: if target not in proto:
return error(msg="target not found") return error(msg="target not found")
if "err_code" in proto[target] and proto[target]["err_code"] != 0: if "err_code" in proto[target] and proto[target]["err_code"] != 0:
return {target: proto[target]} return {target: proto[target]}
def get_response_for_command(cmd): def get_response_for_command(cmd: str):
if cmd not in proto[target]: if cmd not in proto[target]:
return error(msg=f"command {cmd} not found") return error(msg=f"command {cmd} not found")
@@ -528,7 +533,7 @@ class FakeIotTransport(BaseTransport):
from collections import defaultdict from collections import defaultdict
cmd_responses = defaultdict(dict) cmd_responses: dict = defaultdict(dict)
for cmd in request[target]: for cmd in request[target]:
cmd_responses[target][cmd] = get_response_for_command(cmd) cmd_responses[target][cmd] = get_response_for_command(cmd)

View File

@@ -13,7 +13,14 @@ from kasa.transports.basetransport import BaseTransport
class FakeSmartProtocol(SmartProtocol): class FakeSmartProtocol(SmartProtocol):
def __init__(self, info, fixture_name, *, is_child=False, verbatim=False): def __init__(
self,
info: dict,
fixture_name: str,
*,
is_child: bool = False,
verbatim: bool = False,
) -> None:
super().__init__( super().__init__(
transport=FakeSmartTransport( transport=FakeSmartTransport(
info, fixture_name, is_child=is_child, verbatim=verbatim info, fixture_name, is_child=is_child, verbatim=verbatim
@@ -198,7 +205,7 @@ class FakeSmartTransport(BaseTransport):
), ),
} }
def _missing_result(self, method): def _missing_result(self, method: str):
"""Check the FIXTURE_MISSING_MAP for responses. """Check the FIXTURE_MISSING_MAP for responses.
Fixtures generated prior to a query being supported by dump_devinfo Fixtures generated prior to a query being supported by dump_devinfo
@@ -242,7 +249,10 @@ class FakeSmartTransport(BaseTransport):
@staticmethod @staticmethod
def _get_child_protocols( def _get_child_protocols(
parent_fixture_info, parent_fixture_name, child_devices_key, verbatim parent_fixture_info: dict,
parent_fixture_name: str,
child_devices_key: str,
verbatim: bool,
): ):
child_infos = parent_fixture_info.get(child_devices_key, {}).get( child_infos = parent_fixture_info.get(child_devices_key, {}).get(
"child_device_list", [] "child_device_list", []
@@ -250,13 +260,14 @@ class FakeSmartTransport(BaseTransport):
if not child_infos: if not child_infos:
return return
found_child_fixture_infos = [] found_child_fixture_infos = []
child_protocols = {} child_protocols: dict[str, SmartProtocol] = {}
# imported here to avoid circular import # imported here to avoid circular import
from .conftest import filter_fixtures from .conftest import filter_fixtures
def try_get_child_fixture_info(child_dev_info, protocol): def try_get_child_fixture_info(child_dev_info: dict, protocol: str):
hw_version = child_dev_info["hw_ver"] hw_version = child_dev_info["hw_ver"]
sw_version = child_dev_info.get("sw_ver", child_dev_info.get("fw_ver")) sw_version = child_dev_info.get("sw_ver", child_dev_info.get("fw_ver"))
assert isinstance(sw_version, str)
sw_version = sw_version.split(" ")[0] sw_version = sw_version.split(" ")[0]
model = child_dev_info.get("device_model", child_dev_info.get("model")) model = child_dev_info.get("device_model", child_dev_info.get("model"))
assert sw_version assert sw_version
@@ -436,7 +447,7 @@ class FakeSmartTransport(BaseTransport):
raise NotImplementedError(f"Method {child_method} not implemented for children") raise NotImplementedError(f"Method {child_method} not implemented for children")
def _get_on_off_gradually_info(self, info, params): def _get_on_off_gradually_info(self, info: dict, params: dict | None):
if self.components["on_off_gradually"] == 1: if self.components["on_off_gradually"] == 1:
info["get_on_off_gradually_info"] = {"enable": True} info["get_on_off_gradually_info"] = {"enable": True}
else: else:
@@ -446,7 +457,7 @@ class FakeSmartTransport(BaseTransport):
} }
return copy.deepcopy(info["get_on_off_gradually_info"]) return copy.deepcopy(info["get_on_off_gradually_info"])
def _set_on_off_gradually_info(self, info, params): def _set_on_off_gradually_info(self, info: dict, params: dict):
# Child devices can have the required properties directly in info # Child devices can have the required properties directly in info
# the _handle_control_child_missing directly passes in get_device_info # the _handle_control_child_missing directly passes in get_device_info
@@ -488,7 +499,7 @@ class FakeSmartTransport(BaseTransport):
] ]
return {"error_code": 0} return {"error_code": 0}
def _set_dynamic_light_effect(self, info, params): def _set_dynamic_light_effect(self, info: dict, params: dict) -> None:
"""Set or remove values as per the device behaviour.""" """Set or remove values as per the device behaviour."""
info["get_device_info"]["dynamic_light_effect_enable"] = params["enable"] info["get_device_info"]["dynamic_light_effect_enable"] = params["enable"]
info["get_dynamic_light_effect_rules"]["enable"] = params["enable"] info["get_dynamic_light_effect_rules"]["enable"] = params["enable"]
@@ -501,7 +512,7 @@ class FakeSmartTransport(BaseTransport):
if "current_rule_id" in info["get_dynamic_light_effect_rules"]: if "current_rule_id" in info["get_dynamic_light_effect_rules"]:
del info["get_dynamic_light_effect_rules"]["current_rule_id"] del info["get_dynamic_light_effect_rules"]["current_rule_id"]
def _set_edit_dynamic_light_effect_rule(self, info, params): def _set_edit_dynamic_light_effect_rule(self, info: dict, params: dict) -> None:
"""Edit dynamic light effect rule.""" """Edit dynamic light effect rule."""
rules = info["get_dynamic_light_effect_rules"]["rule_list"] rules = info["get_dynamic_light_effect_rules"]["rule_list"]
for rule in rules: for rule in rules:
@@ -511,7 +522,7 @@ class FakeSmartTransport(BaseTransport):
raise Exception("Unable to find rule with id") raise Exception("Unable to find rule with id")
def _set_light_strip_effect(self, info, params): def _set_light_strip_effect(self, info: dict, params: dict) -> None:
"""Set or remove values as per the device behaviour.""" """Set or remove values as per the device behaviour."""
# Brightness is not always available # Brightness is not always available
if (brightness := params.get("brightness")) is not None: if (brightness := params.get("brightness")) is not None:
@@ -522,12 +533,12 @@ class FakeSmartTransport(BaseTransport):
info["get_device_info"]["lighting_effect"]["id"] = params["id"] info["get_device_info"]["lighting_effect"]["id"] = params["id"]
info["get_lighting_effect"] = copy.deepcopy(params) info["get_lighting_effect"] = copy.deepcopy(params)
def _set_led_info(self, info, params): def _set_led_info(self, info: dict, params: dict) -> None:
"""Set or remove values as per the device behaviour.""" """Set or remove values as per the device behaviour."""
info["get_led_info"]["led_status"] = params["led_rule"] != "never" info["get_led_info"]["led_status"] = params["led_rule"] != "never"
info["get_led_info"]["led_rule"] = params["led_rule"] info["get_led_info"]["led_rule"] = params["led_rule"]
def _set_preset_rules(self, info, params): def _set_preset_rules(self, info: dict, params: dict):
"""Set or remove values as per the device behaviour.""" """Set or remove values as per the device behaviour."""
if "brightness" not in info["get_preset_rules"]: if "brightness" not in info["get_preset_rules"]:
return {"error_code": SmartErrorCode.PARAMS_ERROR} return {"error_code": SmartErrorCode.PARAMS_ERROR}
@@ -541,7 +552,7 @@ class FakeSmartTransport(BaseTransport):
] ]
return {"error_code": 0} return {"error_code": 0}
def _set_child_preset_rules(self, info, params): def _set_child_preset_rules(self, info: dict, params: dict):
"""Set or remove values as per the device behaviour.""" """Set or remove values as per the device behaviour."""
# So far the only child device with light preset (KS240) has the # So far the only child device with light preset (KS240) has the
# data available to read in the device_info. If a child device # data available to read in the device_info. If a child device
@@ -551,14 +562,14 @@ class FakeSmartTransport(BaseTransport):
info["preset_state"] = [{"brightness": b} for b in params["brightness"]] info["preset_state"] = [{"brightness": b} for b in params["brightness"]]
return {"error_code": 0} return {"error_code": 0}
def _edit_preset_rules(self, info, params): def _edit_preset_rules(self, info: dict, params: dict):
"""Set or remove values as per the device behaviour.""" """Set or remove values as per the device behaviour."""
if "states" not in info["get_preset_rules"] is None: if info["get_preset_rules"].get("states") is None:
return {"error_code": SmartErrorCode.PARAMS_ERROR} return {"error_code": SmartErrorCode.PARAMS_ERROR}
info["get_preset_rules"]["states"][params["index"]] = params["state"] info["get_preset_rules"]["states"][params["index"]] = params["state"]
return {"error_code": 0} return {"error_code": 0}
def _set_temperature_unit(self, info, params): def _set_temperature_unit(self, info: dict, params: dict):
"""Set or remove values as per the device behaviour.""" """Set or remove values as per the device behaviour."""
unit = params["temp_unit"] unit = params["temp_unit"]
if unit not in {"celsius", "fahrenheit"}: if unit not in {"celsius", "fahrenheit"}:
@@ -579,7 +590,7 @@ class FakeSmartTransport(BaseTransport):
return {"error_code": 0} return {"error_code": 0}
def _hub_remove_device(self, info, params): def _hub_remove_device(self, info: dict, params: dict):
"""Remove hub device.""" """Remove hub device."""
items_to_remove = [dev["device_id"] for dev in params["child_device_list"]] items_to_remove = [dev["device_id"] for dev in params["child_device_list"]]
children = info["get_child_device_list"]["child_device_list"] children = info["get_child_device_list"]["child_device_list"]
@@ -590,10 +601,10 @@ class FakeSmartTransport(BaseTransport):
return {"error_code": 0} return {"error_code": 0}
def get_child_device_queries(self, method, params): def get_child_device_queries(self, method: str, params: dict):
return self._get_method_from_info(method, params) return self._get_method_from_info(method, params)
def _get_method_from_info(self, method, params): def _get_method_from_info(self, method: str, params: dict | None):
result = copy.deepcopy(self.info[method]) result = copy.deepcopy(self.info[method])
if result and "start_index" in result and "sum" in result: if result and "start_index" in result and "sum" in result:
list_key = next( list_key = next(

View File

@@ -13,7 +13,14 @@ from .fakeprotocol_smart import FakeSmartTransport
class FakeSmartCamProtocol(SmartCamProtocol): class FakeSmartCamProtocol(SmartCamProtocol):
def __init__(self, info, fixture_name, *, is_child=False, verbatim=False): def __init__(
self,
info: dict,
fixture_name: str,
*,
is_child: bool = False,
verbatim: bool = False,
) -> None:
super().__init__( super().__init__(
transport=FakeSmartCamTransport( transport=FakeSmartCamTransport(
info, fixture_name, is_child=is_child, verbatim=verbatim info, fixture_name, is_child=is_child, verbatim=verbatim
@@ -29,15 +36,15 @@ class FakeSmartCamProtocol(SmartCamProtocol):
class FakeSmartCamTransport(BaseTransport): class FakeSmartCamTransport(BaseTransport):
def __init__( def __init__(
self, self,
info, info: dict,
fixture_name, fixture_name: str,
*, *,
list_return_size=10, list_return_size=10,
is_child=False, is_child: bool = False,
get_child_fixtures=True, get_child_fixtures=True,
verbatim=False, verbatim: bool = False,
components_not_included=False, components_not_included=False,
): ) -> None:
super().__init__( super().__init__(
config=DeviceConfig( config=DeviceConfig(
"127.0.0.123", "127.0.0.123",
@@ -125,7 +132,7 @@ class FakeSmartCamTransport(BaseTransport):
} }
@staticmethod @staticmethod
def _get_param_set_value(info: dict, set_keys: list[str], value): def _get_param_set_value(info: dict, set_keys: list[str], value: dict) -> None:
cifp = info.get(CHILD_INFO_FROM_PARENT) cifp = info.get(CHILD_INFO_FROM_PARENT)
for key in set_keys[:-1]: for key in set_keys[:-1]:
@@ -205,7 +212,7 @@ class FakeSmartCamTransport(BaseTransport):
], ],
} }
def _hub_remove_device(self, info, params): def _hub_remove_device(self, info: dict, params: dict):
"""Remove hub device.""" """Remove hub device."""
items_to_remove = [dev["device_id"] for dev in params["child_device_list"]] items_to_remove = [dev["device_id"] for dev in params["child_device_list"]]
children = info["getChildDeviceList"]["child_device_list"] children = info["getChildDeviceList"]["child_device_list"]
@@ -225,10 +232,10 @@ class FakeSmartCamTransport(BaseTransport):
next(it, None) next(it, None)
return next(it) return next(it)
def get_child_device_queries(self, method, params): def get_child_device_queries(self, method: str, params: dict | None):
return self._get_method_from_info(method, params) return self._get_method_from_info(method, params)
def _get_method_from_info(self, method, params): def _get_method_from_info(self, method: str, params: dict | None):
result = copy.deepcopy(self.info[method]) result = copy.deepcopy(self.info[method])
if "start_index" in result and "sum" in result: if "start_index" in result and "sum" in result:
list_key = next( list_key = next(

View File

@@ -9,6 +9,7 @@ from pathlib import Path
from typing import NamedTuple from typing import NamedTuple
import pytest import pytest
from pytest_mock import MockerFixture
from kasa.device_type import DeviceType from kasa.device_type import DeviceType
from kasa.iot import IotDevice from kasa.iot import IotDevice
@@ -105,7 +106,7 @@ FIXTURE_DATA: list[FixtureInfo] = get_fixture_infos()
def filter_fixtures( def filter_fixtures(
desc, desc: str,
*, *,
data_root_filter: str | None = None, data_root_filter: str | None = None,
protocol_filter: set[str] | None = None, protocol_filter: set[str] | None = None,
@@ -221,7 +222,7 @@ def filter_fixtures(
params=filter_fixtures("all fixture infos"), params=filter_fixtures("all fixture infos"),
ids=idgenerator, ids=idgenerator,
) )
def fixture_info(request, mocker): def fixture_info(request: pytest.FixtureRequest, mocker: MockerFixture):
"""Return raw discovery file contents as JSON. Used for discovery tests.""" """Return raw discovery file contents as JSON. Used for discovery tests."""
fixture_info = request.param fixture_info = request.param
fixture_data = copy.deepcopy(fixture_info.data) fixture_data = copy.deepcopy(fixture_info.data)

View File

@@ -18,14 +18,14 @@ from tests.device_fixtures import (
@bulb @bulb
async def test_state_attributes(dev: Device): async def test_state_attributes(dev: Device) -> None:
assert "Cloud connection" in dev.state_information assert "Cloud connection" in dev.state_information
assert isinstance(dev.state_information["Cloud connection"], bool) assert isinstance(dev.state_information["Cloud connection"], bool)
@color_bulb @color_bulb
@turn_on @turn_on
async def test_hsv(dev: Device, turn_on): async def test_hsv(dev: Device, turn_on: bool) -> None:
light = dev.modules.get(Module.Light) light = dev.modules.get(Module.Light)
assert light assert light
await handle_turn_on(dev, turn_on) await handle_turn_on(dev, turn_on)
@@ -105,8 +105,14 @@ async def test_hsv(dev: Device, turn_on):
], ],
) )
async def test_invalid_hsv( async def test_invalid_hsv(
dev: Device, turn_on, hue, sat, brightness, exception_cls, error dev: Device,
): turn_on: bool,
hue,
sat,
brightness,
exception_cls: type[Exception],
error: str,
) -> None:
light = dev.modules.get(Module.Light) light = dev.modules.get(Module.Light)
assert light assert light
await handle_turn_on(dev, turn_on) await handle_turn_on(dev, turn_on)
@@ -117,7 +123,7 @@ async def test_invalid_hsv(
@color_bulb @color_bulb
@pytest.mark.skip("requires color feature") @pytest.mark.skip("requires color feature")
async def test_color_state_information(dev: Device): async def test_color_state_information(dev: Device) -> None:
light = dev.modules.get(Module.Light) light = dev.modules.get(Module.Light)
assert light assert light
assert "HSV" in dev.state_information assert "HSV" in dev.state_information
@@ -125,7 +131,7 @@ async def test_color_state_information(dev: Device):
@non_color_bulb @non_color_bulb
async def test_hsv_on_non_color(dev: Device): async def test_hsv_on_non_color(dev: Device) -> None:
light = dev.modules.get(Module.Light) light = dev.modules.get(Module.Light)
assert light assert light
assert not light.has_feature("hsv") assert not light.has_feature("hsv")
@@ -138,7 +144,7 @@ async def test_hsv_on_non_color(dev: Device):
@variable_temp @variable_temp
@pytest.mark.skip("requires colortemp module") @pytest.mark.skip("requires colortemp module")
async def test_variable_temp_state_information(dev: Device): async def test_variable_temp_state_information(dev: Device) -> None:
light = dev.modules.get(Module.Light) light = dev.modules.get(Module.Light)
assert light assert light
assert "Color temperature" in dev.state_information assert "Color temperature" in dev.state_information
@@ -147,7 +153,7 @@ async def test_variable_temp_state_information(dev: Device):
@variable_temp @variable_temp
@turn_on @turn_on
async def test_try_set_colortemp(dev: Device, turn_on): async def test_try_set_colortemp(dev: Device, turn_on: bool) -> None:
light = dev.modules.get(Module.Light) light = dev.modules.get(Module.Light)
assert light assert light
await handle_turn_on(dev, turn_on) await handle_turn_on(dev, turn_on)
@@ -157,7 +163,7 @@ async def test_try_set_colortemp(dev: Device, turn_on):
@variable_temp @variable_temp
async def test_out_of_range_temperature(dev: Device): async def test_out_of_range_temperature(dev: Device) -> None:
light = dev.modules.get(Module.Light) light = dev.modules.get(Module.Light)
assert light assert light
with pytest.raises( with pytest.raises(
@@ -171,7 +177,7 @@ async def test_out_of_range_temperature(dev: Device):
@non_variable_temp @non_variable_temp
async def test_non_variable_temp(dev: Device): async def test_non_variable_temp(dev: Device) -> None:
light = dev.modules.get(Module.Light) light = dev.modules.get(Module.Light)
assert light assert light
with pytest.raises(KasaException): with pytest.raises(KasaException):
@@ -182,7 +188,7 @@ async def test_non_variable_temp(dev: Device):
@bulb @bulb
def test_device_type_bulb(dev: Device): def test_device_type_bulb(dev: Device) -> None:
assert dev.device_type in {DeviceType.Bulb, DeviceType.LightStrip} assert dev.device_type in {DeviceType.Bulb, DeviceType.LightStrip}
@@ -218,7 +224,7 @@ def test_device_type_bulb(dev: Device):
@bulb @bulb
async def test_deprecated_light_is_has_attributes( async def test_deprecated_light_is_has_attributes(
dev: Device, attribute: str, use_msg: str, use_fn: Callable[[Device, Module], bool] dev: Device, attribute: str, use_msg: str, use_fn: Callable[[Device, Module], bool]
): ) -> None:
light = dev.modules.get(Module.Light) light = dev.modules.get(Module.Light)
assert light assert light
@@ -230,7 +236,7 @@ async def test_deprecated_light_is_has_attributes(
@bulb @bulb
async def test_deprecated_light_valid_temperature_range(dev: Device): async def test_deprecated_light_valid_temperature_range(dev: Device) -> None:
light = dev.modules.get(Module.Light) light = dev.modules.get(Module.Light)
assert light assert light

View File

@@ -3,10 +3,11 @@ from datetime import UTC, datetime
import pytest import pytest
from freezegun.api import FrozenDateTimeFactory from freezegun.api import FrozenDateTimeFactory
from pytest_mock import MockerFixture
from kasa import Device from kasa import Device
from kasa.device_type import DeviceType from kasa.device_type import DeviceType
from kasa.protocols.smartprotocol import _ChildProtocolWrapper from kasa.protocols.smartprotocol import SmartProtocol, _ChildProtocolWrapper
from kasa.smart.smartchilddevice import SmartChildDevice from kasa.smart.smartchilddevice import SmartChildDevice
from kasa.smart.smartdevice import NON_HUB_PARENT_ONLY_MODULES, SmartDevice from kasa.smart.smartdevice import NON_HUB_PARENT_ONLY_MODULES, SmartDevice
@@ -30,7 +31,9 @@ has_children = parametrize_combine([has_children_smart, strip_iot])
@strip_smart @strip_smart
def test_childdevice_init(dev, dummy_protocol, mocker): def test_childdevice_init(
dev, dummy_protocol: SmartProtocol, mocker: MockerFixture
) -> None:
"""Test that child devices get initialized and use protocol wrapper.""" """Test that child devices get initialized and use protocol wrapper."""
assert len(dev.children) > 0 assert len(dev.children) > 0
@@ -42,7 +45,9 @@ def test_childdevice_init(dev, dummy_protocol, mocker):
@strip_smart @strip_smart
async def test_childdevice_update(dev, dummy_protocol, mocker): async def test_childdevice_update(
dev, dummy_protocol: SmartProtocol, mocker: MockerFixture
) -> None:
"""Test that parent update updates children.""" """Test that parent update updates children."""
child_info = dev.internal_state["get_child_device_list"] child_info = dev.internal_state["get_child_device_list"]
child_list = child_info["child_device_list"] child_list = child_info["child_device_list"]
@@ -101,7 +106,9 @@ async def test_childdevice_properties(dev: SmartChildDevice):
@non_hub_parent_smart @non_hub_parent_smart
async def test_parent_only_modules(dev, dummy_protocol, mocker): async def test_parent_only_modules(
dev, dummy_protocol: SmartProtocol, mocker: MockerFixture
) -> None:
"""Test that parent only modules are not available on children.""" """Test that parent only modules are not available on children."""
for child in dev.children: for child in dev.children:
for module in NON_HUB_PARENT_ONLY_MODULES: for module in NON_HUB_PARENT_ONLY_MODULES:
@@ -109,7 +116,7 @@ async def test_parent_only_modules(dev, dummy_protocol, mocker):
@has_children @has_children
async def test_parent_property(dev: Device): async def test_parent_property(dev: Device) -> None:
"""Test a child device exposes it's parent.""" """Test a child device exposes it's parent."""
if not dev.children: if not dev.children:
pytest.skip(f"Device {dev} fixture does not have any children") pytest.skip(f"Device {dev} fixture does not have any children")
@@ -121,7 +128,7 @@ async def test_parent_property(dev: Device):
@has_children_smart @has_children_smart
@pytest.mark.requires_dummy @pytest.mark.requires_dummy
async def test_child_time(dev: Device, freezer: FrozenDateTimeFactory): async def test_child_time(dev: Device, freezer: FrozenDateTimeFactory) -> None:
"""Test a child device gets the time from it's parent module. """Test a child device gets the time from it's parent module.
This is excluded from real device testing as the test often fail if the This is excluded from real device testing as the test often fail if the
@@ -137,11 +144,11 @@ async def test_child_time(dev: Device, freezer: FrozenDateTimeFactory):
@pytest.mark.xdist_group(name="caplog") @pytest.mark.xdist_group(name="caplog")
async def test_child_device_type_unknown(caplog): async def test_child_device_type_unknown(caplog: pytest.LogCaptureFixture) -> None:
"""Test for device type when category is unknown.""" """Test for device type when category is unknown."""
class DummyDevice(SmartChildDevice): class DummyDevice(SmartChildDevice):
def __init__(self): def __init__(self) -> None:
super().__init__( super().__init__(
SmartDevice("127.0.0.1"), SmartDevice("127.0.0.1"),
{"device_id": "1", "category": "foobar"}, {"device_id": "1", "category": "foobar"},

View File

@@ -64,7 +64,7 @@ from .conftest import (
pytestmark = [pytest.mark.requires_dummy] pytestmark = [pytest.mark.requires_dummy]
async def test_help(runner): async def test_help(runner: CliRunner) -> None:
"""Test that all the lazy modules are correctly names.""" """Test that all the lazy modules are correctly names."""
res = await runner.invoke(cli, ["--help"]) res = await runner.invoke(cli, ["--help"])
assert res.exit_code == 0, "--help failed, check lazy module names" assert res.exit_code == 0, "--help failed, check lazy module names"
@@ -77,7 +77,9 @@ async def test_help(runner):
pytest.param("SMART.TAPOPLUG", None, id="Only device_family"), pytest.param("SMART.TAPOPLUG", None, id="Only device_family"),
], ],
) )
async def test_update_called_by_cli(dev, mocker, runner, device_family, encrypt_type): async def test_update_called_by_cli(
dev: Device, mocker: MockerFixture, runner: CliRunner, device_family, encrypt_type
) -> None:
"""Test that device update is called on main.""" """Test that device update is called on main."""
update = mocker.patch.object(dev, "update") update = mocker.patch.object(dev, "update")
@@ -107,7 +109,7 @@ async def test_update_called_by_cli(dev, mocker, runner, device_family, encrypt_
update.assert_called() update.assert_called()
async def test_list_devices(discovery_mock, runner): async def test_list_devices(discovery_mock, runner: CliRunner) -> None:
"""Test that device update is called on main.""" """Test that device update is called on main."""
res = await runner.invoke( res = await runner.invoke(
cli, cli,
@@ -128,7 +130,9 @@ async def test_list_devices(discovery_mock, runner):
assert row in res.output assert row in res.output
async def test_discover_raw(discovery_mock, runner, mocker): async def test_discover_raw(
discovery_mock, runner: CliRunner, mocker: MockerFixture
) -> None:
"""Test the discover raw command.""" """Test the discover raw command."""
redact_spy = mocker.patch("kasa.cli.discover.redact_data", side_effect=redact_data) redact_spy = mocker.patch("kasa.cli.discover.redact_data", side_effect=redact_data)
res = await runner.invoke( res = await runner.invoke(
@@ -169,7 +173,13 @@ async def test_discover_raw(discovery_mock, runner, mocker):
], ],
) )
@new_discovery @new_discovery
async def test_list_update_failed(discovery_mock, mocker, runner, exception, expected): async def test_list_update_failed(
discovery_mock,
mocker: MockerFixture,
runner: CliRunner,
exception: type[Exception],
expected: str,
) -> None:
"""Test that device update is called on main.""" """Test that device update is called on main."""
device_class = Discover._get_device_class(discovery_mock.discovery_data) device_class = Discover._get_device_class(discovery_mock.discovery_data)
mocker.patch.object( mocker.patch.object(
@@ -196,7 +206,9 @@ async def test_list_update_failed(discovery_mock, mocker, runner, exception, exp
assert row in res.output.replace("\n", "") assert row in res.output.replace("\n", "")
async def test_list_unsupported(unsupported_device_info, runner): async def test_list_unsupported(
unsupported_device_info: dict, runner: CliRunner
) -> None:
"""Test that device update is called on main.""" """Test that device update is called on main."""
res = await runner.invoke( res = await runner.invoke(
cli, cli,
@@ -213,14 +225,14 @@ async def test_list_unsupported(unsupported_device_info, runner):
assert row in res.output assert row in res.output
async def test_sysinfo(dev: Device, runner): async def test_sysinfo(dev: Device, runner: CliRunner) -> None:
res = await runner.invoke(sysinfo, obj=dev) res = await runner.invoke(sysinfo, obj=dev)
assert "System info" in res.output assert "System info" in res.output
assert dev.model in res.output assert dev.model in res.output
@turn_on @turn_on
async def test_state(dev, turn_on, runner): async def test_state(dev: Device, turn_on: bool, runner: CliRunner) -> None:
await handle_turn_on(dev, turn_on) await handle_turn_on(dev, turn_on)
await dev.update() await dev.update()
res = await runner.invoke(state, obj=dev) res = await runner.invoke(state, obj=dev)
@@ -232,7 +244,7 @@ async def test_state(dev, turn_on, runner):
@turn_on @turn_on
async def test_toggle(dev, turn_on, runner): async def test_toggle(dev: Device, turn_on: bool, runner: CliRunner) -> None:
if isinstance(dev, SmartCamDevice) and dev.device_type == DeviceType.Hub: if isinstance(dev, SmartCamDevice) and dev.device_type == DeviceType.Hub:
pytest.skip(reason="Hub cannot toggle state") pytest.skip(reason="Hub cannot toggle state")
@@ -245,7 +257,7 @@ async def test_toggle(dev, turn_on, runner):
assert dev.is_on != turn_on assert dev.is_on != turn_on
async def test_alias(dev, runner): async def test_alias(dev: Device, runner: CliRunner) -> None:
res = await runner.invoke(alias, obj=dev) res = await runner.invoke(alias, obj=dev)
assert f"Alias: {dev.alias}" in res.output assert f"Alias: {dev.alias}" in res.output
@@ -263,7 +275,9 @@ async def test_alias(dev, runner):
await dev.set_alias(old_alias or "") await dev.set_alias(old_alias or "")
async def test_raw_command(dev, mocker, runner): async def test_raw_command(
dev: Device, mocker: MockerFixture, runner: CliRunner
) -> None:
update = mocker.patch.object(dev, "update") update = mocker.patch.object(dev, "update")
from kasa.smart import SmartDevice from kasa.smart import SmartDevice
@@ -291,17 +305,17 @@ async def test_raw_command(dev, mocker, runner):
assert "Usage" in res.output assert "Usage" in res.output
async def test_command_with_child(dev, mocker, runner): async def test_command_with_child(dev, mocker: MockerFixture, runner: CliRunner):
"""Test 'command' command with --child.""" """Test 'command' command with --child."""
update_mock = mocker.patch.object(dev, "update") update_mock = mocker.patch.object(dev, "update")
# create_autospec for device slows tests way too much, so we use a dummy here # create_autospec for device slows tests way too much, so we use a dummy here
class DummyDevice(dev.__class__): class DummyDevice(dev.__class__):
def __init__(self): def __init__(self) -> None:
super().__init__("127.0.0.1") super().__init__("127.0.0.1")
# device_type and _info initialised for repr # device_type and _info initialised for repr
self._device_type = Device.Type.StripSocket self._device_type = Device.Type.StripSocket
self._info = {} self._info: dict = {}
async def _query_helper(*_, **__): async def _query_helper(*_, **__):
return {"dummy": "response"} return {"dummy": "response"}
@@ -324,7 +338,7 @@ async def test_command_with_child(dev, mocker, runner):
@device_smart @device_smart
async def test_reboot(dev, mocker, runner): async def test_reboot(dev: Device, mocker: MockerFixture, runner: CliRunner) -> None:
"""Test that reboot works on SMART devices.""" """Test that reboot works on SMART devices."""
query_mock = mocker.patch.object(dev.protocol, "query") query_mock = mocker.patch.object(dev.protocol, "query")
@@ -338,7 +352,9 @@ async def test_reboot(dev, mocker, runner):
@device_smart @device_smart
async def test_factory_reset(dev, mocker, runner): async def test_factory_reset(
dev: Device, mocker: MockerFixture, runner: CliRunner
) -> None:
"""Test that factory reset works on SMART devices.""" """Test that factory reset works on SMART devices."""
query_mock = mocker.patch.object(dev.protocol, "query") query_mock = mocker.patch.object(dev.protocol, "query")
@@ -353,7 +369,7 @@ async def test_factory_reset(dev, mocker, runner):
@device_smart @device_smart
async def test_wifi_scan(dev, runner): async def test_wifi_scan(dev: Device, runner: CliRunner) -> None:
res = await runner.invoke(wifi, ["scan"], obj=dev) res = await runner.invoke(wifi, ["scan"], obj=dev)
assert res.exit_code == 0 assert res.exit_code == 0
@@ -361,7 +377,7 @@ async def test_wifi_scan(dev, runner):
@parametrize_combine([device_smart, device_iot]) @parametrize_combine([device_smart, device_iot])
async def test_wifi_join(dev, mocker, runner): async def test_wifi_join(dev: Device, mocker: MockerFixture, runner: CliRunner) -> None:
update = mocker.patch.object(dev, "update") update = mocker.patch.object(dev, "update")
res = await runner.invoke( res = await runner.invoke(
wifi, wifi,
@@ -378,7 +394,9 @@ async def test_wifi_join(dev, mocker, runner):
@parametrize_combine([device_smart, device_iot]) @parametrize_combine([device_smart, device_iot])
async def test_wifi_join_missing_keytype(dev, mocker, runner): async def test_wifi_join_missing_keytype(
dev: Device, mocker: MockerFixture, runner: CliRunner
) -> None:
"""Test that missing keytype raises KasaException and CLI echoes the message.""" """Test that missing keytype raises KasaException and CLI echoes the message."""
update = mocker.patch.object(dev, "update") update = mocker.patch.object(dev, "update")
res = await runner.invoke( res = await runner.invoke(
@@ -396,7 +414,9 @@ async def test_wifi_join_missing_keytype(dev, mocker, runner):
@device_smartcam @device_smartcam
async def test_wifi_join_smartcam(dev, mocker, runner): async def test_wifi_join_smartcam(
dev: Device, mocker: MockerFixture, runner: CliRunner
) -> None:
update = mocker.patch.object(dev, "update") update = mocker.patch.object(dev, "update")
res = await runner.invoke( res = await runner.invoke(
wifi, wifi,
@@ -413,7 +433,7 @@ async def test_wifi_join_smartcam(dev, mocker, runner):
@device_smart @device_smart
async def test_wifi_join_no_creds(dev, runner): async def test_wifi_join_no_creds(dev: Device, runner: CliRunner) -> None:
dev.protocol._transport._credentials = None dev.protocol._transport._credentials = None
res = await runner.invoke( res = await runner.invoke(
wifi, wifi,
@@ -426,7 +446,9 @@ async def test_wifi_join_no_creds(dev, runner):
@device_smart @device_smart
async def test_wifi_join_exception(dev, mocker, runner): async def test_wifi_join_exception(
dev: Device, mocker: MockerFixture, runner: CliRunner
) -> None:
mocker.patch.object(dev.protocol, "query", side_effect=DeviceError(error_code=9999)) mocker.patch.object(dev.protocol, "query", side_effect=DeviceError(error_code=9999))
res = await runner.invoke( res = await runner.invoke(
wifi, wifi,
@@ -439,7 +461,7 @@ async def test_wifi_join_exception(dev, mocker, runner):
@device_smart @device_smart
async def test_update_credentials(dev, runner): async def test_update_credentials(dev: Device, runner: CliRunner) -> None:
res = await runner.invoke( res = await runner.invoke(
update_credentials, update_credentials,
["--username", "foo", "--password", "bar"], ["--username", "foo", "--password", "bar"],
@@ -454,7 +476,7 @@ async def test_update_credentials(dev, runner):
) )
async def test_time_get(dev, runner): async def test_time_get(dev: Device, runner: CliRunner) -> None:
"""Test time get command.""" """Test time get command."""
res = await runner.invoke( res = await runner.invoke(
time, time,
@@ -464,7 +486,7 @@ async def test_time_get(dev, runner):
assert "Current time: " in res.output assert "Current time: " in res.output
async def test_time_sync(dev, mocker, runner): async def test_time_sync(dev: Device, mocker: MockerFixture, runner: CliRunner) -> None:
"""Test time sync command.""" """Test time sync command."""
update = mocker.patch.object(dev, "update") update = mocker.patch.object(dev, "update")
set_time_mock = mocker.spy(dev.modules[Module.Time], "set_time") set_time_mock = mocker.spy(dev.modules[Module.Time], "set_time")
@@ -482,7 +504,7 @@ async def test_time_sync(dev, mocker, runner):
@parametrize_combine([device_smart, device_iot]) @parametrize_combine([device_smart, device_iot])
async def test_time_set(dev: Device, mocker, runner): async def test_time_set(dev: Device, mocker: MockerFixture, runner: CliRunner) -> None:
"""Test time set command.""" """Test time set command."""
time_mod = dev.modules[Module.Time] time_mod = dev.modules[Module.Time]
set_time_mock = mocker.spy(time_mod, "set_time") set_time_mock = mocker.spy(time_mod, "set_time")
@@ -524,7 +546,7 @@ async def test_time_set(dev: Device, mocker, runner):
assert "New time: " in res.output assert "New time: " in res.output
async def test_emeter(dev: Device, mocker, runner): async def test_emeter(dev: Device, mocker: MockerFixture, runner: CliRunner) -> None:
mocker.patch("kasa.Discover.discover_single", return_value=dev) mocker.patch("kasa.Discover.discover_single", return_value=dev)
base_cmd = ["--host", "dummy", "energy"] base_cmd = ["--host", "dummy", "energy"]
res = await runner.invoke(cli, base_cmd, obj=dev) res = await runner.invoke(cli, base_cmd, obj=dev)
@@ -554,9 +576,9 @@ async def test_emeter(dev: Device, mocker, runner):
child_status.assert_called() child_status.assert_called()
assert child_status.call_count == 1 assert child_status.call_count == 1
res = await runner.invoke( child_alias = dev.children[0].alias
cli, [*base_cmd, "--name", dev.children[0].alias], obj=dev assert child_alias is not None
) res = await runner.invoke(cli, [*base_cmd, "--name", child_alias], obj=dev)
assert "Voltage: 122.066 V" in res.output assert "Voltage: 122.066 V" in res.output
assert child_status.call_count == 2 assert child_status.call_count == 2
@@ -583,7 +605,7 @@ async def test_emeter(dev: Device, mocker, runner):
daily.assert_called_with(year=1900, month=12) daily.assert_called_with(year=1900, month=12)
async def test_brightness(dev: Device, runner): async def test_brightness(dev: Device, runner: CliRunner) -> None:
res = await runner.invoke(brightness, obj=dev) res = await runner.invoke(brightness, obj=dev)
if not (light := dev.modules.get(Module.Light)) or not light.has_feature( if not (light := dev.modules.get(Module.Light)) or not light.has_feature(
"brightness" "brightness"
@@ -602,7 +624,7 @@ async def test_brightness(dev: Device, runner):
assert "Brightness: 12" in res.output assert "Brightness: 12" in res.output
async def test_color_temperature(dev: Device, runner): async def test_color_temperature(dev: Device, runner: CliRunner) -> None:
res = await runner.invoke(temperature, obj=dev) res = await runner.invoke(temperature, obj=dev)
if not (light := dev.modules.get(Module.Light)) or not ( if not (light := dev.modules.get(Module.Light)) or not (
color_temp_feat := light.get_feature("color_temp") color_temp_feat := light.get_feature("color_temp")
@@ -637,7 +659,7 @@ async def test_color_temperature(dev: Device, runner):
assert res.exit_code == 2 assert res.exit_code == 2
async def test_color_hsv(dev: Device, runner: CliRunner): async def test_color_hsv(dev: Device, runner: CliRunner) -> None:
res = await runner.invoke(hsv, obj=dev) res = await runner.invoke(hsv, obj=dev)
if not (light := dev.modules.get(Module.Light)) or not light.has_feature("hsv"): if not (light := dev.modules.get(Module.Light)) or not light.has_feature("hsv"):
assert "Device does not support colors" in res.output assert "Device does not support colors" in res.output
@@ -656,7 +678,7 @@ async def test_color_hsv(dev: Device, runner: CliRunner):
assert res.exit_code == 2 assert res.exit_code == 2
async def test_light_effect(dev: Device, runner: CliRunner): async def test_light_effect(dev: Device, runner: CliRunner) -> None:
res = await runner.invoke(effect, obj=dev) res = await runner.invoke(effect, obj=dev)
if not (light_effect := dev.modules.get(Module.LightEffect)): if not (light_effect := dev.modules.get(Module.LightEffect)):
assert "Device does not support effects" in res.output assert "Device does not support effects" in res.output
@@ -682,7 +704,7 @@ async def test_light_effect(dev: Device, runner: CliRunner):
assert res.exit_code == 2 assert res.exit_code == 2
async def test_light_preset(dev: Device, runner: CliRunner): async def test_light_preset(dev: Device, runner: CliRunner) -> None:
res = await runner.invoke(presets, obj=dev) res = await runner.invoke(presets, obj=dev)
if not (light_preset := dev.modules.get(Module.LightPreset)): if not (light_preset := dev.modules.get(Module.LightPreset)):
assert "Device does not support light presets" in res.output assert "Device does not support light presets" in res.output
@@ -725,7 +747,7 @@ async def test_light_preset(dev: Device, runner: CliRunner):
assert "Need to supply at least one option to modify." in res.output assert "Need to supply at least one option to modify." in res.output
async def test_led(dev: Device, runner: CliRunner): async def test_led(dev: Device, runner: CliRunner) -> None:
res = await runner.invoke(led, obj=dev) res = await runner.invoke(led, obj=dev)
if not (led_module := dev.modules.get(Module.Led)): if not (led_module := dev.modules.get(Module.Led)):
assert "Device does not support led" in res.output assert "Device does not support led" in res.output
@@ -748,7 +770,9 @@ async def test_led(dev: Device, runner: CliRunner):
assert led_module.led is False assert led_module.led is False
async def test_json_output(dev: Device, mocker, runner): async def test_json_output(
dev: Device, mocker: MockerFixture, runner: CliRunner
) -> None:
"""Test that the json output produces correct output.""" """Test that the json output produces correct output."""
mocker.patch("kasa.Discover.discover_single", return_value=dev) mocker.patch("kasa.Discover.discover_single", return_value=dev)
# These will mock the features to avoid accessing non-existing ones # These will mock the features to avoid accessing non-existing ones
@@ -761,13 +785,15 @@ async def test_json_output(dev: Device, mocker, runner):
@new_discovery @new_discovery
async def test_credentials(discovery_mock, mocker, runner): async def test_credentials(
discovery_mock, mocker: MockerFixture, runner: CliRunner
) -> None:
"""Test credentials are passed correctly from cli to device.""" """Test credentials are passed correctly from cli to device."""
# Patch state to echo username and password # Patch state to echo username and password
pass_dev = click.make_pass_decorator(Device) pass_dev = click.make_pass_decorator(Device) # type: ignore[type-abstract]
@pass_dev @pass_dev
async def _state(dev: Device): async def _state(dev: Device) -> None:
if dev.credentials: if dev.credentials:
click.echo( click.echo(
f"Username:{dev.credentials.username} Password:{dev.credentials.password}" f"Username:{dev.credentials.username} Password:{dev.credentials.password}"
@@ -776,29 +802,32 @@ async def test_credentials(discovery_mock, mocker, runner):
mocker.patch("kasa.cli.device.state", new=_state) mocker.patch("kasa.cli.device.state", new=_state)
dr = DiscoveryResult.from_dict(discovery_mock.discovery_data["result"]) dr = DiscoveryResult.from_dict(discovery_mock.discovery_data["result"])
assert dr.mgt_encrypt_schm is not None
cli_args: list[str] = [
"--host",
"127.0.0.123",
"--username",
"foo",
"--password",
"bar",
"--device-family",
dr.device_type,
]
if dr.mgt_encrypt_schm.encrypt_type is not None:
cli_args += ["--encrypt-type", dr.mgt_encrypt_schm.encrypt_type]
cli_args += ["--login-version", str(dr.mgt_encrypt_schm.lv or 1)]
res = await runner.invoke( res = await runner.invoke(
cli, cli,
[ cli_args,
"--host",
"127.0.0.123",
"--username",
"foo",
"--password",
"bar",
"--device-family",
dr.device_type,
"--encrypt-type",
dr.mgt_encrypt_schm.encrypt_type,
"--login-version",
dr.mgt_encrypt_schm.lv or 1,
],
) )
assert res.exit_code == 0 assert res.exit_code == 0
assert "Username:foo Password:bar\n" in res.output assert "Username:foo Password:bar\n" in res.output
async def test_without_device_type(dev, mocker, runner): async def test_without_device_type(
dev: Device, mocker: MockerFixture, runner: CliRunner
) -> None:
"""Test connecting without the device type.""" """Test connecting without the device type."""
discovery_mock = mocker.patch( discovery_mock = mocker.patch(
"kasa.discover.Discover.discover_single", return_value=dev "kasa.discover.Discover.discover_single", return_value=dev
@@ -833,7 +862,7 @@ async def test_without_device_type(dev, mocker, runner):
@pytest.mark.parametrize("auth_param", ["--username", "--password"]) @pytest.mark.parametrize("auth_param", ["--username", "--password"])
async def test_invalid_credential_params(auth_param, runner): async def test_invalid_credential_params(auth_param: str, runner: CliRunner) -> None:
"""Test for handling only one of username or password supplied.""" """Test for handling only one of username or password supplied."""
res = await runner.invoke( res = await runner.invoke(
cli, cli,
@@ -853,7 +882,7 @@ async def test_invalid_credential_params(auth_param, runner):
) )
async def test_duplicate_target_device(runner): async def test_duplicate_target_device(runner: CliRunner) -> None:
"""Test that defining both --host or --alias gives an error.""" """Test that defining both --host or --alias gives an error."""
res = await runner.invoke( res = await runner.invoke(
cli, cli,
@@ -868,7 +897,9 @@ async def test_duplicate_target_device(runner):
assert "Error: Use either --alias or --host, not both." in res.output assert "Error: Use either --alias or --host, not both." in res.output
async def test_discover(discovery_mock, mocker, runner): async def test_discover(
discovery_mock, mocker: MockerFixture, runner: CliRunner
) -> None:
"""Test discovery output.""" """Test discovery output."""
# These will mock the features to avoid accessing non-existing # These will mock the features to avoid accessing non-existing
mocker.patch("kasa.device.Device.features", return_value={}) mocker.patch("kasa.device.Device.features", return_value={})
@@ -878,7 +909,7 @@ async def test_discover(discovery_mock, mocker, runner):
cli, cli,
[ [
"--discovery-timeout", "--discovery-timeout",
0, "0",
"--username", "--username",
"foo", "foo",
"--password", "--password",
@@ -890,7 +921,9 @@ async def test_discover(discovery_mock, mocker, runner):
assert res.exit_code == 0 assert res.exit_code == 0
async def test_discover_host(discovery_mock, mocker, runner): async def test_discover_host(
discovery_mock, mocker: MockerFixture, runner: CliRunner
) -> None:
"""Test discovery output.""" """Test discovery output."""
# These will mock the features to avoid accessing non-existing # These will mock the features to avoid accessing non-existing
mocker.patch("kasa.device.Device.features", return_value={}) mocker.patch("kasa.device.Device.features", return_value={})
@@ -900,7 +933,7 @@ async def test_discover_host(discovery_mock, mocker, runner):
cli, cli,
[ [
"--discovery-timeout", "--discovery-timeout",
0, "0",
"--host", "--host",
"127.0.0.123", "127.0.0.123",
"--username", "--username",
@@ -913,13 +946,15 @@ async def test_discover_host(discovery_mock, mocker, runner):
assert res.exit_code == 0 assert res.exit_code == 0
async def test_discover_unsupported(unsupported_device_info, runner): async def test_discover_unsupported(
unsupported_device_info: dict, runner: CliRunner
) -> None:
"""Test discovery output.""" """Test discovery output."""
res = await runner.invoke( res = await runner.invoke(
cli, cli,
[ [
"--discovery-timeout", "--discovery-timeout",
0, "0",
"--username", "--username",
"foo", "foo",
"--password", "--password",
@@ -932,7 +967,9 @@ async def test_discover_unsupported(unsupported_device_info, runner):
assert "== Unsupported device ==" in res.output assert "== Unsupported device ==" in res.output
async def test_host_unsupported(unsupported_device_info, runner): async def test_host_unsupported(
unsupported_device_info: dict, runner: CliRunner
) -> None:
"""Test discovery output.""" """Test discovery output."""
host = "127.0.0.1" host = "127.0.0.1"
@@ -954,7 +991,9 @@ async def test_host_unsupported(unsupported_device_info, runner):
@new_discovery @new_discovery
async def test_discover_auth_failed(discovery_mock, mocker, runner): async def test_discover_auth_failed(
discovery_mock, mocker: MockerFixture, runner: CliRunner
) -> None:
"""Test discovery output.""" """Test discovery output."""
host = "127.0.0.1" host = "127.0.0.1"
discovery_mock.ip = host discovery_mock.ip = host
@@ -968,7 +1007,7 @@ async def test_discover_auth_failed(discovery_mock, mocker, runner):
cli, cli,
[ [
"--discovery-timeout", "--discovery-timeout",
0, "0",
"--username", "--username",
"foo", "foo",
"--password", "--password",
@@ -984,7 +1023,9 @@ async def test_discover_auth_failed(discovery_mock, mocker, runner):
@new_discovery @new_discovery
async def test_host_auth_failed(discovery_mock, mocker, runner): async def test_host_auth_failed(
discovery_mock, mocker: MockerFixture, runner: CliRunner
) -> None:
"""Test discovery output.""" """Test discovery output."""
host = "127.0.0.1" host = "127.0.0.1"
discovery_mock.ip = host discovery_mock.ip = host
@@ -1012,17 +1053,20 @@ async def test_host_auth_failed(discovery_mock, mocker, runner):
@pytest.mark.parametrize("device_type", TYPES) @pytest.mark.parametrize("device_type", TYPES)
async def test_type_param(device_type, mocker, runner): async def test_type_param(
device_type: str, mocker: MockerFixture, runner: CliRunner
) -> None:
"""Test for handling only one of username or password supplied.""" """Test for handling only one of username or password supplied."""
result_device = FileNotFoundError result_device: type[FileNotFoundError] | Device = FileNotFoundError
pass_dev = click.make_pass_decorator(Device) pass_dev = click.make_pass_decorator(Device) # type: ignore[type-abstract]
@pass_dev @pass_dev
async def _state(dev: Device): async def _state(dev: Device) -> None:
nonlocal result_device nonlocal result_device
result_device = dev result_device = dev
mocker.patch("kasa.cli.device.state", new=_state) mocker.patch("kasa.cli.device.state", new=_state)
expected_type: type[Device]
if device_type == "camera": if device_type == "camera":
expected_type = SmartCamDevice expected_type = SmartCamDevice
elif device_type == "smart": elif device_type == "smart":
@@ -1047,7 +1091,10 @@ async def test_type_param(device_type, mocker, runner):
], ],
) )
async def test_type_camera_login_version( async def test_type_camera_login_version(
cli_login_version, expected_login_version, mocker, runner cli_login_version: int | None,
expected_login_version: int,
mocker: MockerFixture,
runner: CliRunner,
): ):
"""Test that --type camera respects an explicitly provided --login-version.""" """Test that --type camera respects an explicitly provided --login-version."""
from kasa.deviceconfig import DeviceConfig from kasa.deviceconfig import DeviceConfig
@@ -1078,7 +1125,7 @@ async def test_type_camera_login_version(
@pytest.mark.skip( @pytest.mark.skip(
"Skip until pytest-asyncio supports pytest 8.0, https://github.com/pytest-dev/pytest-asyncio/issues/737" "Skip until pytest-asyncio supports pytest 8.0, https://github.com/pytest-dev/pytest-asyncio/issues/737"
) )
async def test_shell(dev: Device, mocker, runner): async def test_shell(dev: Device, mocker: MockerFixture, runner: CliRunner) -> None:
"""Test that the shell commands tries to embed a shell.""" """Test that the shell commands tries to embed a shell."""
mocker.patch("kasa.Discover.discover", return_value=[dev]) mocker.patch("kasa.Discover.discover", return_value=[dev])
# repl = mocker.patch("ptpython.repl") # repl = mocker.patch("ptpython.repl")
@@ -1092,7 +1139,7 @@ async def test_shell(dev: Device, mocker, runner):
embed.assert_called() embed.assert_called()
async def test_errors(mocker, runner): async def test_errors(mocker: MockerFixture, runner: CliRunner) -> None:
err = KasaException("Foobar") err = KasaException("Foobar")
# Test masking # Test masking
@@ -1137,7 +1184,7 @@ async def test_errors(mocker, runner):
assert "Raised error:" not in res.output assert "Raised error:" not in res.output
async def test_feature(mocker, runner): async def test_feature(mocker: MockerFixture, runner: CliRunner) -> None:
"""Test feature command.""" """Test feature command."""
dummy_device = await get_device_for_fixture_protocol( dummy_device = await get_device_for_fixture_protocol(
"P300(EU)_1.0_1.0.13.json", "SMART" "P300(EU)_1.0_1.0.13.json", "SMART"
@@ -1154,7 +1201,9 @@ async def test_feature(mocker, runner):
assert res.exit_code == 0 assert res.exit_code == 0
async def test_features_all(discovery_mock, mocker, runner): async def test_features_all(
discovery_mock, mocker: MockerFixture, runner: CliRunner
) -> None:
"""Test feature command on all fixtures.""" """Test feature command on all fixtures."""
res = await runner.invoke( res = await runner.invoke(
cli, cli,
@@ -1168,7 +1217,7 @@ async def test_features_all(discovery_mock, mocker, runner):
assert res.exit_code == 0 assert res.exit_code == 0
async def test_feature_single(mocker, runner): async def test_feature_single(mocker: MockerFixture, runner: CliRunner) -> None:
"""Test feature command returning single value.""" """Test feature command returning single value."""
dummy_device = await get_device_for_fixture_protocol( dummy_device = await get_device_for_fixture_protocol(
"P300(EU)_1.0_1.0.13.json", "SMART" "P300(EU)_1.0_1.0.13.json", "SMART"
@@ -1184,7 +1233,7 @@ async def test_feature_single(mocker, runner):
assert res.exit_code == 0 assert res.exit_code == 0
async def test_feature_missing(mocker, runner): async def test_feature_missing(mocker: MockerFixture, runner: CliRunner) -> None:
"""Test feature command returning single value.""" """Test feature command returning single value."""
dummy_device = await get_device_for_fixture_protocol( dummy_device = await get_device_for_fixture_protocol(
"P300(EU)_1.0_1.0.13.json", "SMART" "P300(EU)_1.0_1.0.13.json", "SMART"
@@ -1200,7 +1249,7 @@ async def test_feature_missing(mocker, runner):
assert res.exit_code == 1 assert res.exit_code == 1
async def test_feature_set(mocker, runner): async def test_feature_set(mocker: MockerFixture, runner: CliRunner) -> None:
"""Test feature command's set value.""" """Test feature command's set value."""
dummy_device = await get_device_for_fixture_protocol( dummy_device = await get_device_for_fixture_protocol(
"P300(EU)_1.0_1.0.13.json", "SMART" "P300(EU)_1.0_1.0.13.json", "SMART"
@@ -1219,7 +1268,7 @@ async def test_feature_set(mocker, runner):
assert res.exit_code == 0 assert res.exit_code == 0
async def test_feature_set_child(mocker, runner): async def test_feature_set_child(mocker: MockerFixture, runner: CliRunner) -> None:
"""Test feature command's set value.""" """Test feature command's set value."""
dummy_device = await get_device_for_fixture_protocol( dummy_device = await get_device_for_fixture_protocol(
"P300(EU)_1.0_1.0.13.json", "SMART" "P300(EU)_1.0_1.0.13.json", "SMART"
@@ -1254,7 +1303,7 @@ async def test_feature_set_child(mocker, runner):
assert res.exit_code == 0 assert res.exit_code == 0
async def test_feature_set_unquoted(mocker, runner): async def test_feature_set_unquoted(mocker: MockerFixture, runner: CliRunner) -> None:
"""Test feature command's set value.""" """Test feature command's set value."""
dummy_device = await get_device_for_fixture_protocol( dummy_device = await get_device_for_fixture_protocol(
"ES20M(US)_1.0_1.0.11.json", "IOT" "ES20M(US)_1.0_1.0.11.json", "IOT"
@@ -1273,7 +1322,7 @@ async def test_feature_set_unquoted(mocker, runner):
assert res.exit_code != 0 assert res.exit_code != 0
async def test_feature_set_badquoted(mocker, runner): async def test_feature_set_badquoted(mocker: MockerFixture, runner: CliRunner) -> None:
"""Test feature command's set value.""" """Test feature command's set value."""
dummy_device = await get_device_for_fixture_protocol( dummy_device = await get_device_for_fixture_protocol(
"ES20M(US)_1.0_1.0.11.json", "IOT" "ES20M(US)_1.0_1.0.11.json", "IOT"
@@ -1292,7 +1341,7 @@ async def test_feature_set_badquoted(mocker, runner):
assert res.exit_code != 0 assert res.exit_code != 0
async def test_feature_set_goodquoted(mocker, runner): async def test_feature_set_goodquoted(mocker: MockerFixture, runner: CliRunner) -> None:
"""Test feature command's set value.""" """Test feature command's set value."""
dummy_device = await get_device_for_fixture_protocol( dummy_device = await get_device_for_fixture_protocol(
"ES20M(US)_1.0_1.0.11.json", "IOT" "ES20M(US)_1.0_1.0.11.json", "IOT"
@@ -1313,7 +1362,7 @@ async def test_feature_set_goodquoted(mocker, runner):
async def test_cli_child_commands( async def test_cli_child_commands(
dev: Device, runner: CliRunner, mocker: MockerFixture dev: Device, runner: CliRunner, mocker: MockerFixture
): ) -> None:
if not dev.children: if not dev.children:
res = await runner.invoke(alias, ["--child-index", "0"], obj=dev) res = await runner.invoke(alias, ["--child-index", "0"], obj=dev)
assert f"Device: {dev.host} does not have children" in res.output assert f"Device: {dev.host} does not have children" in res.output
@@ -1418,7 +1467,9 @@ async def test_cli_child_commands(
assert dev.children[0].update == child_update_method assert dev.children[0].update == child_update_method
async def test_discover_config(dev: Device, mocker, runner): async def test_discover_config(
dev: Device, mocker: MockerFixture, runner: CliRunner
) -> None:
"""Test that device config is returned.""" """Test that device config is returned."""
host = "127.0.0.1" host = "127.0.0.1"
mocker.patch("kasa.device_factory._connect", side_effect=[Exception, dev]) mocker.patch("kasa.device_factory._connect", side_effect=[Exception, dev])
@@ -1451,7 +1502,9 @@ async def test_discover_config(dev: Device, mocker, runner):
) )
async def test_discover_config_invalid(mocker, runner): async def test_discover_config_invalid(
mocker: MockerFixture, runner: CliRunner
) -> None:
"""Test the device config command with invalids.""" """Test the device config command with invalids."""
host = "127.0.0.1" host = "127.0.0.1"
mocker.patch("kasa.discover.Discover.try_connect_all", return_value=None) mocker.patch("kasa.discover.Discover.try_connect_all", return_value=None)

View File

@@ -3,6 +3,7 @@ import inspect
import pkgutil import pkgutil
import sys import sys
from datetime import UTC, datetime, timedelta, timezone from datetime import UTC, datetime, timedelta, timezone
from types import ModuleType
from unittest.mock import AsyncMock from unittest.mock import AsyncMock
from zoneinfo import ZoneInfo from zoneinfo import ZoneInfo
@@ -77,9 +78,9 @@ temp_control_smart = parametrize(
interfaces = pytest.mark.parametrize("interface", kasa.interfaces.__all__) interfaces = pytest.mark.parametrize("interface", kasa.interfaces.__all__)
def _get_subclasses(of_class, package): def _get_subclasses(of_class: type, package: ModuleType) -> set[type]:
"""Get all the subclasses of a given class.""" """Get all the subclasses of a given class."""
subclasses = set() subclasses: set[type] = set()
# iter_modules returns ModuleInfo: (module_finder, name, ispkg) # iter_modules returns ModuleInfo: (module_finder, name, ispkg)
for _, modname, ispkg in pkgutil.iter_modules(package.__path__): for _, modname, ispkg in pkgutil.iter_modules(package.__path__):
importlib.import_module("." + modname, package=package.__name__) importlib.import_module("." + modname, package=package.__name__)
@@ -100,7 +101,7 @@ def _get_subclasses(of_class, package):
@interfaces @interfaces
def test_feature_attributes(interface): def test_feature_attributes(interface: str) -> None:
"""Test that all common derived classes define the FeatureAttributes.""" """Test that all common derived classes define the FeatureAttributes."""
klass = getattr(kasa.interfaces, interface) klass = getattr(kasa.interfaces, interface)
@@ -126,7 +127,7 @@ def test_feature_attributes(interface):
@led @led
async def test_led_module(dev: Device, mocker: MockerFixture): async def test_led_module(dev: Device, mocker: MockerFixture) -> None:
"""Test fan speed feature.""" """Test fan speed feature."""
led_module = dev.modules.get(Module.Led) led_module = dev.modules.get(Module.Led)
assert led_module assert led_module
@@ -153,7 +154,7 @@ async def test_led_module(dev: Device, mocker: MockerFixture):
@light_effect @light_effect
async def test_light_effect_module(dev: Device, mocker: MockerFixture): async def test_light_effect_module(dev: Device, mocker: MockerFixture) -> None:
"""Test fan speed feature.""" """Test fan speed feature."""
light_effect_module = dev.modules[Module.LightEffect] light_effect_module = dev.modules[Module.LightEffect]
assert light_effect_module assert light_effect_module
@@ -205,7 +206,7 @@ async def test_light_effect_module(dev: Device, mocker: MockerFixture):
@light_effect @light_effect
async def test_light_effect_brightness(dev: Device, mocker: MockerFixture): async def test_light_effect_brightness(dev: Device, mocker: MockerFixture) -> None:
"""Test that light module uses light_effect for brightness when active.""" """Test that light module uses light_effect for brightness when active."""
light_module = dev.modules[Module.Light] light_module = dev.modules[Module.Light]
@@ -230,7 +231,7 @@ async def test_light_effect_brightness(dev: Device, mocker: MockerFixture):
@dimmable @dimmable
async def test_light_brightness(dev: Device): async def test_light_brightness(dev: Device) -> None:
"""Test brightness setter and getter.""" """Test brightness setter and getter."""
assert isinstance(dev, Device) assert isinstance(dev, Device)
light = next(get_parent_and_child_modules(dev, Module.Light)) light = next(get_parent_and_child_modules(dev, Module.Light))
@@ -253,7 +254,7 @@ async def test_light_brightness(dev: Device):
@variable_temp @variable_temp
async def test_light_color_temp(dev: Device): async def test_light_color_temp(dev: Device) -> None:
"""Test color temp setter and getter.""" """Test color temp setter and getter."""
assert isinstance(dev, Device) assert isinstance(dev, Device)
@@ -292,7 +293,7 @@ async def test_light_color_temp(dev: Device):
@light @light
async def test_light_set_state(dev: Device): async def test_light_set_state(dev: Device) -> None:
"""Test brightness setter and getter.""" """Test brightness setter and getter."""
assert isinstance(dev, Device) assert isinstance(dev, Device)
light = next(get_parent_and_child_modules(dev, Module.Light)) light = next(get_parent_and_child_modules(dev, Module.Light))
@@ -319,7 +320,7 @@ async def test_light_set_state(dev: Device):
@light_preset @light_preset
async def test_light_preset_module(dev: Device, mocker: MockerFixture): async def test_light_preset_module(dev: Device, mocker: MockerFixture) -> None:
"""Test light preset module.""" """Test light preset module."""
preset_mod = next(get_parent_and_child_modules(dev, Module.LightPreset)) preset_mod = next(get_parent_and_child_modules(dev, Module.LightPreset))
assert preset_mod assert preset_mod
@@ -370,7 +371,7 @@ async def test_light_preset_module(dev: Device, mocker: MockerFixture):
@light_preset @light_preset
async def test_light_preset_save(dev: Device, mocker: MockerFixture): async def test_light_preset_save(dev: Device, mocker: MockerFixture) -> None:
"""Test saving a new preset value.""" """Test saving a new preset value."""
preset_mod = next(get_parent_and_child_modules(dev, Module.LightPreset)) preset_mod = next(get_parent_and_child_modules(dev, Module.LightPreset))
assert preset_mod assert preset_mod
@@ -393,7 +394,7 @@ async def test_light_preset_save(dev: Device, mocker: MockerFixture):
@temp_control_smart @temp_control_smart
async def test_thermostat(dev: Device, mocker: MockerFixture): async def test_thermostat(dev: Device, mocker: MockerFixture) -> None:
"""Test saving a new preset value.""" """Test saving a new preset value."""
therm_mod = next(get_parent_and_child_modules(dev, Module.Thermostat)) therm_mod = next(get_parent_and_child_modules(dev, Module.Thermostat))
assert therm_mod assert therm_mod
@@ -426,7 +427,7 @@ async def test_thermostat(dev: Device, mocker: MockerFixture):
@time @time
async def test_set_time(dev: Device): async def test_set_time(dev: Device) -> None:
"""Test setting the device time.""" """Test setting the device time."""
time_mod = dev.modules[Module.Time] time_mod = dev.modules[Module.Time]
@@ -463,7 +464,9 @@ async def test_set_time(dev: Device):
assert time_mod.time == original_time assert time_mod.time == original_time
async def test_time_post_update_no_time_uses_utc_unit(monkeypatch: pytest.MonkeyPatch): async def test_time_post_update_no_time_uses_utc_unit(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""If neither get_timezone nor get_time are present, timezone falls back to UTC.""" """If neither get_timezone nor get_time are present, timezone falls back to UTC."""
from kasa.iot.modules.time import Time as TimeModule from kasa.iot.modules.time import Time as TimeModule
@@ -476,7 +479,7 @@ async def test_time_post_update_no_time_uses_utc_unit(monkeypatch: pytest.Monkey
async def test_time_post_update_uses_offset_when_index_missing_unit( async def test_time_post_update_uses_offset_when_index_missing_unit(
monkeypatch: pytest.MonkeyPatch, mocker: MockerFixture monkeypatch: pytest.MonkeyPatch, mocker: MockerFixture
): ) -> None:
"""When index present but zone not on host, fall back to offset-based guess.""" """When index present but zone not on host, fall back to offset-based guess."""
from zoneinfo import ZoneInfoNotFoundError from zoneinfo import ZoneInfoNotFoundError
@@ -547,7 +550,7 @@ async def test_time_post_update_unsynced_clock_uses_utc(
assert time_mod.time.year == 2000 assert time_mod.time.year == 2000
async def test_time_get_time_exception_returns_none_unit(mocker: MockerFixture): async def test_time_get_time_exception_returns_none_unit(mocker: MockerFixture) -> None:
"""Cover Time.get_time exception path (unit test of iot Time).""" """Cover Time.get_time exception path (unit test of iot Time)."""
from kasa.iot.modules.time import Time as TimeModule from kasa.iot.modules.time import Time as TimeModule
@@ -557,7 +560,7 @@ async def test_time_get_time_exception_returns_none_unit(mocker: MockerFixture):
assert await TimeModule.get_time(inst) is None assert await TimeModule.get_time(inst) is None
async def test_time_get_time_success_unit(mocker: MockerFixture): async def test_time_get_time_success_unit(mocker: MockerFixture) -> None:
"""Cover the success path of Time.get_time.""" """Cover the success path of Time.get_time."""
from kasa.iot.modules.time import Time as TimeModule from kasa.iot.modules.time import Time as TimeModule
@@ -589,7 +592,7 @@ async def test_time_get_time_success_unit(mocker: MockerFixture):
async def test_time_post_update_with_time_no_tz_uses_guess_unit( async def test_time_post_update_with_time_no_tz_uses_guess_unit(
monkeypatch: pytest.MonkeyPatch, mocker: MockerFixture monkeypatch: pytest.MonkeyPatch, mocker: MockerFixture
): ) -> None:
"""When get_time is present but get_timezone is missing, use offset-based guess (dst_expected None).""" """When get_time is present but get_timezone is missing, use offset-based guess (dst_expected None)."""
from kasa.iot.modules.time import Time as TimeModule from kasa.iot.modules.time import Time as TimeModule
@@ -620,7 +623,7 @@ async def test_time_post_update_with_time_no_tz_uses_guess_unit(
async def test_time_set_time_wraps_exception_unit( async def test_time_set_time_wraps_exception_unit(
monkeypatch: pytest.MonkeyPatch, mocker: MockerFixture monkeypatch: pytest.MonkeyPatch, mocker: MockerFixture
): ) -> None:
"""Cover exception wrapping in Time.set_time (unit test of iot Time).""" """Cover exception wrapping in Time.set_time (unit test of iot Time)."""
from kasa.iot.modules.time import Time as TimeModule from kasa.iot.modules.time import Time as TimeModule
@@ -673,7 +676,7 @@ async def test_smart_time_set_time_no_region_added_when_tzname_none_unit(
async def test_smartcam_time_post_update_fallback_parses_timezone_str_unit( async def test_smartcam_time_post_update_fallback_parses_timezone_str_unit(
monkeypatch: pytest.MonkeyPatch, mocker: MockerFixture monkeypatch: pytest.MonkeyPatch, mocker: MockerFixture
): ) -> None:
"""Exercise smartcam Time._post_update_hook fallback when ZoneInfo not found, parsing 'timezone' string.""" """Exercise smartcam Time._post_update_hook fallback when ZoneInfo not found, parsing 'timezone' string."""
from zoneinfo import ZoneInfoNotFoundError from zoneinfo import ZoneInfoNotFoundError

View File

@@ -34,7 +34,7 @@ from kasa.smart import SmartChildDevice, SmartDevice
from kasa.smartcam import SmartCamChild, SmartCamDevice from kasa.smartcam import SmartCamChild, SmartCamDevice
def _get_subclasses(of_class): def _get_subclasses(of_class: type):
package = sys.modules["kasa"] package = sys.modules["kasa"]
subclasses = set() subclasses = set()
for _, modname, _ in pkgutil.iter_modules(package.__path__): for _, modname, _ in pkgutil.iter_modules(package.__path__):
@@ -44,6 +44,7 @@ def _get_subclasses(of_class):
if ( if (
inspect.isclass(obj) inspect.isclass(obj)
and issubclass(obj, of_class) and issubclass(obj, of_class)
and module.__package__ is not None
and module.__package__ != "kasa" and module.__package__ != "kasa"
and module.__package__ != "kasa.interfaces" and module.__package__ != "kasa.interfaces"
): ):
@@ -56,12 +57,12 @@ device_classes = pytest.mark.parametrize(
) )
async def test_device_id(dev: Device): async def test_device_id(dev: Device) -> None:
"""Test all devices have a device id.""" """Test all devices have a device id."""
assert dev.device_id assert dev.device_id
async def test_alias(dev): async def test_alias(dev: Device) -> None:
test_alias = "TEST1234" test_alias = "TEST1234"
original = dev.alias original = dev.alias
@@ -77,7 +78,7 @@ async def test_alias(dev):
@device_classes @device_classes
async def test_device_class_ctors(device_class_name_obj): async def test_device_class_ctors(device_class_name_obj) -> None:
"""Make sure constructor api not broken for new and existing SmartDevices.""" """Make sure constructor api not broken for new and existing SmartDevices."""
host = "127.0.0.2" host = "127.0.0.2"
port = 1234 port = 1234
@@ -112,7 +113,7 @@ async def test_device_class_ctors(device_class_name_obj):
@device_classes @device_classes
async def test_device_class_repr(device_class_name_obj): async def test_device_class_repr(device_class_name_obj) -> None:
"""Test device repr when update() not called and no discovery info.""" """Test device repr when update() not called and no discovery info."""
host = "127.0.0.2" host = "127.0.0.2"
port = 1234 port = 1234
@@ -155,16 +156,16 @@ async def test_device_class_repr(device_class_name_obj):
assert repr(dev) == expected_repr assert repr(dev) == expected_repr
async def test_create_device_with_timeout(): async def test_create_device_with_timeout() -> None:
"""Make sure timeout is passed to the protocol.""" """Make sure timeout is passed to the protocol."""
host = "127.0.0.1" host = "127.0.0.1"
dev = IotDevice(host, config=DeviceConfig(host, timeout=100)) dev: Device = IotDevice(host, config=DeviceConfig(host, timeout=100))
assert dev.protocol._transport._timeout == 100 assert dev.protocol._transport._timeout == 100
dev = SmartDevice(host, config=DeviceConfig(host, timeout=100)) dev = SmartDevice(host, config=DeviceConfig(host, timeout=100))
assert dev.protocol._transport._timeout == 100 assert dev.protocol._transport._timeout == 100
async def test_create_thin_wrapper(): async def test_create_thin_wrapper() -> None:
"""Make sure thin wrapper is created with the correct device type.""" """Make sure thin wrapper is created with the correct device type."""
mock = AsyncMock() mock = AsyncMock()
config = DeviceConfig( config = DeviceConfig(
@@ -186,7 +187,7 @@ async def test_create_thin_wrapper():
@pytest.mark.parametrize( @pytest.mark.parametrize(
("device_class", "use_class"), kasa.deprecated_smart_devices.items() ("device_class", "use_class"), kasa.deprecated_smart_devices.items()
) )
def test_deprecated_devices(device_class, use_class): def test_deprecated_devices(device_class, use_class) -> None:
package_name = ".".join(use_class.__module__.split(".")[:-1]) package_name = ".".join(use_class.__module__.split(".")[:-1])
msg = f"{device_class} is deprecated, use {use_class.__name__} from package {package_name} instead" msg = f"{device_class} is deprecated, use {use_class.__name__} from package {package_name} instead"
with pytest.deprecated_call(match=msg): with pytest.deprecated_call(match=msg):
@@ -201,7 +202,7 @@ def test_deprecated_devices(device_class, use_class):
@pytest.mark.parametrize( @pytest.mark.parametrize(
("deprecated_class", "use_class"), kasa.deprecated_classes.items() ("deprecated_class", "use_class"), kasa.deprecated_classes.items()
) )
def test_deprecated_classes(deprecated_class, use_class): def test_deprecated_classes(deprecated_class, use_class) -> None:
msg = f"{deprecated_class} is deprecated, use {use_class.__name__} instead" msg = f"{deprecated_class} is deprecated, use {use_class.__name__} instead"
with pytest.deprecated_call(match=msg): with pytest.deprecated_call(match=msg):
getattr(kasa, deprecated_class) getattr(kasa, deprecated_class)
@@ -234,7 +235,7 @@ deprecated_is_light_function_smart_module = {
def test_deprecated_device_type_attributes(dev: SmartDevice): def test_deprecated_device_type_attributes(dev: SmartDevice):
"""Test deprecated attributes on all devices.""" """Test deprecated attributes on all devices."""
def _test_attr(attribute): def _test_attr(attribute: str):
msg = f"{attribute} is deprecated" msg = f"{attribute} is deprecated"
if module := Device._deprecated_device_type_attributes[attribute][0]: if module := Device._deprecated_device_type_attributes[attribute][0]:
msg += f", use: {module} in device.modules instead" msg += f", use: {module} in device.modules instead"
@@ -249,8 +250,13 @@ def test_deprecated_device_type_attributes(dev: SmartDevice):
async def _test_attribute( async def _test_attribute(
dev: Device, attribute_name, is_expected, module_name, *args, will_raise=False dev: Device,
): attribute_name: str,
is_expected: bool,
module_name: str | None,
*args: object,
will_raise: type[Exception] | None = None,
) -> None:
will_warn = is_expected or attribute_name in deprecated_warns_before_attribute_error will_warn = is_expected or attribute_name in deprecated_warns_before_attribute_error
if will_raise: if will_raise:
@@ -278,7 +284,7 @@ async def _test_attribute(
assert attribute_val is not None assert attribute_val is not None
async def test_deprecated_light_effect_attributes(dev: Device): async def test_deprecated_light_effect_attributes(dev: Device) -> None:
light_effect = dev.modules.get(Module.LightEffect) light_effect = dev.modules.get(Module.LightEffect)
await _test_attribute(dev, "effect", bool(light_effect), "LightEffect") await _test_attribute(dev, "effect", bool(light_effect), "LightEffect")
@@ -299,7 +305,7 @@ async def test_deprecated_light_effect_attributes(dev: Device):
) )
async def test_deprecated_light_attributes(dev: Device): async def test_deprecated_light_attributes(dev: Device) -> None:
light = dev.modules.get(Module.Light) light = dev.modules.get(Module.Light)
await _test_attribute(dev, "is_dimmable", bool(light), "Light") await _test_attribute(dev, "is_dimmable", bool(light), "Light")
@@ -330,7 +336,7 @@ async def test_deprecated_light_attributes(dev: Device):
await _test_attribute(dev, "has_effects", bool(light), "Light") await _test_attribute(dev, "has_effects", bool(light), "Light")
async def test_deprecated_other_attributes(dev: Device): async def test_deprecated_other_attributes(dev: Device) -> None:
led_module = dev.modules.get(Module.Led) led_module = dev.modules.get(Module.Led)
await _test_attribute(dev, "led", bool(led_module), "Led") await _test_attribute(dev, "led", bool(led_module), "Led")
@@ -338,7 +344,7 @@ async def test_deprecated_other_attributes(dev: Device):
await _test_attribute(dev, "supported_modules", True, None) await _test_attribute(dev, "supported_modules", True, None)
async def test_deprecated_emeter_attributes(dev: Device): async def test_deprecated_emeter_attributes(dev: Device) -> None:
energy_module = dev.modules.get(Module.Energy) energy_module = dev.modules.get(Module.Energy)
await _test_attribute(dev, "get_emeter_realtime", bool(energy_module), "Energy") await _test_attribute(dev, "get_emeter_realtime", bool(energy_module), "Energy")
@@ -350,7 +356,7 @@ async def test_deprecated_emeter_attributes(dev: Device):
await _test_attribute(dev, "get_emeter_monthly", bool(energy_module), "Energy") await _test_attribute(dev, "get_emeter_monthly", bool(energy_module), "Energy")
async def test_deprecated_light_preset_attributes(dev: Device): async def test_deprecated_light_preset_attributes(dev: Device) -> None:
preset = dev.modules.get(Module.LightPreset) preset = dev.modules.get(Module.LightPreset)
exc: type[AttributeError] | type[KasaException] | None = ( exc: type[AttributeError] | type[KasaException] | None = (
@@ -381,7 +387,7 @@ async def test_deprecated_light_preset_attributes(dev: Device):
async def test_device_type_aliases(): async def test_device_type_aliases():
"""Test that the device type aliases in Device work.""" """Test that the device type aliases in Device work."""
def _mock_connect(config, *args, **kwargs): def _mock_connect(config: DeviceConfig, *args, **kwargs):
mock = AsyncMock() mock = AsyncMock()
mock.config = config mock.config = config
return mock return mock
@@ -402,7 +408,7 @@ async def test_device_type_aliases():
assert DeviceType.Dimmer == Device.Type.Dimmer assert DeviceType.Dimmer == Device.Type.Dimmer
async def test_device_timezones(): async def test_device_timezones() -> None:
"""Test the timezone data is good.""" """Test the timezone data is good."""
# Check all indexes return a zoneinfo # Check all indexes return a zoneinfo
for i in range(110): for i in range(110):

View File

@@ -11,6 +11,7 @@ from typing import cast
import aiohttp import aiohttp
import pytest # type: ignore # https://github.com/pytest-dev/pytest/issues/3342 import pytest # type: ignore # https://github.com/pytest-dev/pytest/issues/3342
from pytest_mock import MockerFixture
from kasa import ( from kasa import (
BaseProtocol, BaseProtocol,
@@ -57,7 +58,7 @@ from .conftest import DISCOVERY_MOCK_IP
pytestmark = [pytest.mark.requires_dummy] pytestmark = [pytest.mark.requires_dummy]
def _get_connection_type_device_class(discovery_info): def _get_connection_type_device_class(discovery_info: dict):
if "result" in discovery_info: if "result" in discovery_info:
device_class = Discover._get_device_class(discovery_info) device_class = Discover._get_device_class(discovery_info)
dr = DiscoveryResult.from_dict(discovery_info["result"]) dr = DiscoveryResult.from_dict(discovery_info["result"])
@@ -74,8 +75,8 @@ def _get_connection_type_device_class(discovery_info):
async def test_connect( async def test_connect(
discovery_mock, discovery_mock,
mocker, mocker: MockerFixture,
): ) -> None:
"""Test that if the protocol is passed in it gets set correctly.""" """Test that if the protocol is passed in it gets set correctly."""
host = DISCOVERY_MOCK_IP host = DISCOVERY_MOCK_IP
ctype, device_class = _get_connection_type_device_class( ctype, device_class = _get_connection_type_device_class(
@@ -102,7 +103,9 @@ async def test_connect(
@pytest.mark.parametrize("custom_port", [123, None]) @pytest.mark.parametrize("custom_port", [123, None])
async def test_connect_custom_port(discovery_mock, mocker, custom_port): async def test_connect_custom_port(
discovery_mock, mocker: MockerFixture, custom_port: int | None
) -> None:
"""Make sure that connect returns an initialized SmartDevice instance.""" """Make sure that connect returns an initialized SmartDevice instance."""
host = DISCOVERY_MOCK_IP host = DISCOVERY_MOCK_IP
@@ -127,7 +130,7 @@ async def test_connect_custom_port(discovery_mock, mocker, custom_port):
async def test_connect_logs_connect_time( async def test_connect_logs_connect_time(
discovery_mock, discovery_mock,
caplog: pytest.LogCaptureFixture, caplog: pytest.LogCaptureFixture,
): ) -> None:
"""Test that the connect time is logged when debug logging is enabled.""" """Test that the connect time is logged when debug logging is enabled."""
discovery_data = discovery_mock.discovery_data discovery_data = discovery_mock.discovery_data
ctype, _ = _get_connection_type_device_class(discovery_data) ctype, _ = _get_connection_type_device_class(discovery_data)
@@ -143,7 +146,7 @@ async def test_connect_logs_connect_time(
assert "seconds to update" in caplog.text assert "seconds to update" in caplog.text
async def test_connect_query_fails(discovery_mock, mocker): async def test_connect_query_fails(discovery_mock, mocker: MockerFixture) -> None:
"""Make sure that connect fails when query fails.""" """Make sure that connect fails when query fails."""
host = DISCOVERY_MOCK_IP host = DISCOVERY_MOCK_IP
discovery_data = discovery_mock.discovery_data discovery_data = discovery_mock.discovery_data
@@ -162,7 +165,7 @@ async def test_connect_query_fails(discovery_mock, mocker):
assert close_mock.call_count == 1 assert close_mock.call_count == 1
async def test_connect_http_client(discovery_mock, mocker): async def test_connect_http_client(discovery_mock, mocker: MockerFixture) -> None:
"""Make sure that discover_single returns an initialized SmartDevice instance.""" """Make sure that discover_single returns an initialized SmartDevice instance."""
host = DISCOVERY_MOCK_IP host = DISCOVERY_MOCK_IP
discovery_data = discovery_mock.discovery_data discovery_data = discovery_mock.discovery_data
@@ -175,7 +178,8 @@ async def test_connect_http_client(discovery_mock, mocker):
) )
dev = await connect(config=config) dev = await connect(config=config)
if ctype.encryption_type != DeviceEncryptionType.Xor: if ctype.encryption_type != DeviceEncryptionType.Xor:
assert dev.protocol._transport._http_client.client != http_client http_client_wrap = dev.protocol._transport._http_client # type: ignore[attr-defined]
assert http_client_wrap.client != http_client
await dev.disconnect() await dev.disconnect()
config = DeviceConfig( config = DeviceConfig(
@@ -186,12 +190,13 @@ async def test_connect_http_client(discovery_mock, mocker):
) )
dev = await connect(config=config) dev = await connect(config=config)
if ctype.encryption_type != DeviceEncryptionType.Xor: if ctype.encryption_type != DeviceEncryptionType.Xor:
assert dev.protocol._transport._http_client.client == http_client http_client_wrap = dev.protocol._transport._http_client # type: ignore[attr-defined]
assert http_client_wrap.client == http_client
await dev.disconnect() await dev.disconnect()
await http_client.close() await http_client.close()
async def test_device_types(dev: Device): async def test_device_types(dev: Device) -> None:
await dev.update() await dev.update()
if isinstance(dev, SmartCamDevice): if isinstance(dev, SmartCamDevice):
res = SmartCamDevice._get_device_type_from_sysinfo(dev.sys_info) res = SmartCamDevice._get_device_type_from_sysinfo(dev.sys_info)
@@ -208,7 +213,9 @@ async def test_device_types(dev: Device):
@pytest.mark.xdist_group(name="caplog") @pytest.mark.xdist_group(name="caplog")
async def test_device_class_from_unknown_family(caplog): async def test_device_class_from_unknown_family(
caplog: pytest.LogCaptureFixture,
) -> None:
"""Verify that unknown SMART devices yield a warning and fallback to SmartDevice.""" """Verify that unknown SMART devices yield a warning and fallback to SmartDevice."""
dummy_name = "SMART.foo" dummy_name = "SMART.foo"
with caplog.at_level(logging.DEBUG): with caplog.at_level(logging.DEBUG):
@@ -303,7 +310,7 @@ async def test_get_protocol(
conn_params: DeviceConnectionParameters, conn_params: DeviceConnectionParameters,
expected_protocol: type[BaseProtocol], expected_protocol: type[BaseProtocol],
expected_transport: type[BaseTransport], expected_transport: type[BaseTransport],
): ) -> None:
"""Test get_protocol returns the right protocol.""" """Test get_protocol returns the right protocol."""
config = DeviceConfig("127.0.0.1", connection_type=conn_params) config = DeviceConfig("127.0.0.1", connection_type=conn_params)
protocol = get_protocol(config) protocol = get_protocol(config)

View File

@@ -1,7 +1,7 @@
from kasa.device_type import DeviceType from kasa.device_type import DeviceType
async def test_device_type_from_value(): async def test_device_type_from_value() -> None:
"""Make sure that every device type can be created from its value.""" """Make sure that every device type can be created from its value."""
for name in DeviceType: for name in DeviceType:
assert DeviceType.from_value(name.value) is not None assert DeviceType.from_value(name.value) is not None

View File

@@ -2,6 +2,7 @@ import json
from dataclasses import replace from dataclasses import replace
from json import dumps as json_dumps from json import dumps as json_dumps
from json import loads as json_loads from json import loads as json_loads
from unittest.mock import MagicMock
import aiohttp import aiohttp
import pytest import pytest
@@ -32,7 +33,7 @@ CAMERA_AES_CONFIG = DeviceConfig(
) )
async def test_serialization(): async def test_serialization() -> None:
"""Test device config serialization.""" """Test device config serialization."""
config = DeviceConfig(host="Foo", http_client=aiohttp.ClientSession()) config = DeviceConfig(host="Foo", http_client=aiohttp.ClientSession())
config_dict = config.to_dict() config_dict = config.to_dict()
@@ -52,7 +53,7 @@ async def test_serialization():
], ],
ids=lambda arg: arg.split("_")[-1] if isinstance(arg, str) else "", ids=lambda arg: arg.split("_")[-1] if isinstance(arg, str) else "",
) )
async def test_deserialization(fixture_name: str, expected_value: DeviceConfig): async def test_deserialization(fixture_name: str, expected_value: DeviceConfig) -> None:
"""Test device config deserialization.""" """Test device config deserialization."""
dict_val = json.loads(load_fixture("serialization", fixture_name)) dict_val = json.loads(load_fixture("serialization", fixture_name))
config = DeviceConfig.from_dict(dict_val) config = DeviceConfig.from_dict(dict_val)
@@ -60,17 +61,19 @@ async def test_deserialization(fixture_name: str, expected_value: DeviceConfig):
assert expected_value.to_dict() == dict_val assert expected_value.to_dict() == dict_val
async def test_serialization_http_client(): async def test_serialization_http_client() -> None:
"""Test that the http client does not try to serialize.""" """Test that an HTTP client is excluded from serialization."""
dict_val = json.loads(load_fixture("serialization", "deviceconfig_plug-klap.json")) dict_val = json.loads(load_fixture("serialization", "deviceconfig_plug-klap.json"))
config = replace(PLUG_KLAP_CONFIG, http_client=object()) config = replace(
PLUG_KLAP_CONFIG, http_client=MagicMock(spec=aiohttp.ClientSession)
)
assert config.http_client assert config.http_client
assert config.to_dict() == dict_val assert config.to_dict() == dict_val
async def test_conn_param_no_https(): async def test_conn_param_no_https() -> None:
"""Test no https in connection param defaults to False.""" """Test no https in connection param defaults to False."""
dict_val = { dict_val = {
"device_family": "SMART.TAPOPLUG", "device_family": "SMART.TAPOPLUG",
@@ -90,12 +93,12 @@ async def test_conn_param_no_https():
], ],
ids=["invalid-dict", "not-dict"], ids=["invalid-dict", "not-dict"],
) )
def test_deserialization_errors(input_value, expected_error): def test_deserialization_errors(input_value, expected_error: type[Exception]) -> None:
with pytest.raises(expected_error): with pytest.raises(expected_error):
DeviceConfig.from_dict(input_value) DeviceConfig.from_dict(input_value)
async def test_credentials_hash(): async def test_credentials_hash() -> None:
config = DeviceConfig( config = DeviceConfig(
host="Foo", host="Foo",
http_client=aiohttp.ClientSession(), http_client=aiohttp.ClientSession(),
@@ -109,7 +112,7 @@ async def test_credentials_hash():
assert config2.credentials is None assert config2.credentials is None
async def test_blank_credentials_hash(): async def test_blank_credentials_hash() -> None:
config = DeviceConfig( config = DeviceConfig(
host="Foo", host="Foo",
http_client=aiohttp.ClientSession(), http_client=aiohttp.ClientSession(),
@@ -123,7 +126,7 @@ async def test_blank_credentials_hash():
assert config2.credentials is None assert config2.credentials is None
async def test_exclude_credentials(): async def test_exclude_credentials() -> None:
config = DeviceConfig( config = DeviceConfig(
host="Foo", host="Foo",
http_client=aiohttp.ClientSession(), http_client=aiohttp.ClientSession(),

View File

@@ -35,7 +35,7 @@ iot_fixtures = parametrize(
) )
async def test_fixture_names(fixture_info: FixtureInfo): async def test_fixture_names(fixture_info: FixtureInfo) -> None:
"""Test that device info gets the right fixture names.""" """Test that device info gets the right fixture names."""
if fixture_info.protocol in {"SMARTCAM"}: if fixture_info.protocol in {"SMARTCAM"}:
device_info = SmartCamDevice._get_device_info( device_info = SmartCamDevice._get_device_info(
@@ -58,7 +58,7 @@ async def test_fixture_names(fixture_info: FixtureInfo):
@smart_fixtures @smart_fixtures
async def test_smart_fixtures(fixture_info: FixtureInfo): async def test_smart_fixtures(fixture_info: FixtureInfo) -> None:
"""Test that smart fixtures are created the same.""" """Test that smart fixtures are created the same."""
dev = await get_device_for_fixture(fixture_info, verbatim=True) dev = await get_device_for_fixture(fixture_info, verbatim=True)
assert isinstance(dev, SmartDevice) assert isinstance(dev, SmartDevice)
@@ -74,7 +74,7 @@ async def test_smart_fixtures(fixture_info: FixtureInfo):
assert fixture_info.data == fixture_result.data assert fixture_info.data == fixture_result.data
def _normalize_child_device_ids(info: dict): def _normalize_child_device_ids(info: dict) -> None:
"""Scrubbed child device ids in hubs may not match ids in child fixtures. """Scrubbed child device ids in hubs may not match ids in child fixtures.
Different hub fixtures could create the same child fixture so we scrub Different hub fixtures could create the same child fixture so we scrub
@@ -91,7 +91,7 @@ def _normalize_child_device_ids(info: dict):
@smartcam_fixtures @smartcam_fixtures
async def test_smartcam_fixtures(fixture_info: FixtureInfo): async def test_smartcam_fixtures(fixture_info: FixtureInfo) -> None:
"""Test that smartcam fixtures are created the same.""" """Test that smartcam fixtures are created the same."""
dev = await get_device_for_fixture(fixture_info, verbatim=True) dev = await get_device_for_fixture(fixture_info, verbatim=True)
assert isinstance(dev, SmartCamDevice) assert isinstance(dev, SmartCamDevice)
@@ -136,7 +136,7 @@ async def test_smartcam_fixtures(fixture_info: FixtureInfo):
@iot_fixtures @iot_fixtures
async def test_iot_fixtures(fixture_info: FixtureInfo): async def test_iot_fixtures(fixture_info: FixtureInfo) -> None:
"""Test that iot fixtures are created the same.""" """Test that iot fixtures are created the same."""
# Iot fixtures often do not have enough data to perform a device update() # Iot fixtures often do not have enough data to perform a device update()
# without missing info being added to suppress the update # without missing info being added to suppress the update

View File

@@ -14,6 +14,7 @@ import aiohttp
import pytest # type: ignore # https://github.com/pytest-dev/pytest/issues/3342 import pytest # type: ignore # https://github.com/pytest-dev/pytest/issues/3342
from cryptography.hazmat.primitives import hashes from cryptography.hazmat.primitives import hashes
from cryptography.hazmat.primitives.asymmetric import padding as asymmetric_padding from cryptography.hazmat.primitives.asymmetric import padding as asymmetric_padding
from pytest_mock import MockerFixture
from kasa import ( from kasa import (
Credentials, Credentials,
@@ -81,7 +82,7 @@ UNSUPPORTED = {
@wallswitch_iot @wallswitch_iot
async def test_type_detection_switch(dev: Device): async def test_type_detection_switch(dev: Device) -> None:
d = Discover._get_device_class(dev._last_update)("localhost") d = Discover._get_device_class(dev._last_update)("localhost")
with pytest.deprecated_call(match="use device_type property instead"): with pytest.deprecated_call(match="use device_type property instead"):
assert d.is_wallswitch assert d.is_wallswitch
@@ -89,13 +90,13 @@ async def test_type_detection_switch(dev: Device):
@plug_iot @plug_iot
async def test_type_detection_plug(dev: Device): async def test_type_detection_plug(dev: Device) -> None:
d = Discover._get_device_class(dev._last_update)("localhost") d = Discover._get_device_class(dev._last_update)("localhost")
assert d.device_type == DeviceType.Plug assert d.device_type == DeviceType.Plug
@bulb_iot @bulb_iot
async def test_type_detection_bulb(dev: Device): async def test_type_detection_bulb(dev: Device) -> None:
d = Discover._get_device_class(dev._last_update)("localhost") d = Discover._get_device_class(dev._last_update)("localhost")
# TODO: light_strip is a special case for now to force bulb tests on it # TODO: light_strip is a special case for now to force bulb tests on it
@@ -104,25 +105,25 @@ async def test_type_detection_bulb(dev: Device):
@strip_iot @strip_iot
async def test_type_detection_strip(dev: Device): async def test_type_detection_strip(dev: Device) -> None:
d = Discover._get_device_class(dev._last_update)("localhost") d = Discover._get_device_class(dev._last_update)("localhost")
assert d.device_type == DeviceType.Strip assert d.device_type == DeviceType.Strip
@dimmer_iot @dimmer_iot
async def test_type_detection_dimmer(dev: Device): async def test_type_detection_dimmer(dev: Device) -> None:
d = Discover._get_device_class(dev._last_update)("localhost") d = Discover._get_device_class(dev._last_update)("localhost")
assert d.device_type == DeviceType.Dimmer assert d.device_type == DeviceType.Dimmer
@lightstrip_iot @lightstrip_iot
async def test_type_detection_lightstrip(dev: Device): async def test_type_detection_lightstrip(dev: Device) -> None:
d = Discover._get_device_class(dev._last_update)("localhost") d = Discover._get_device_class(dev._last_update)("localhost")
assert d.device_type == DeviceType.LightStrip assert d.device_type == DeviceType.LightStrip
@pytest.mark.xdist_group(name="caplog") @pytest.mark.xdist_group(name="caplog")
async def test_type_unknown(caplog): async def test_type_unknown(caplog: pytest.LogCaptureFixture) -> None:
invalid_info = {"system": {"get_sysinfo": {"type": "nosuchtype"}}} invalid_info = {"system": {"get_sysinfo": {"type": "nosuchtype"}}}
assert Discover._get_device_class(invalid_info) is IotPlug assert Discover._get_device_class(invalid_info) is IotPlug
msg = "Unknown device type nosuchtype, falling back to plug" msg = "Unknown device type nosuchtype, falling back to plug"
@@ -130,7 +131,9 @@ async def test_type_unknown(caplog):
@pytest.mark.parametrize("custom_port", [123, None]) @pytest.mark.parametrize("custom_port", [123, None])
async def test_discover_single(discovery_mock, custom_port, mocker): async def test_discover_single(
discovery_mock, custom_port: int | None, mocker: MockerFixture
) -> None:
"""Make sure that discover_single returns an initialized SmartDevice instance.""" """Make sure that discover_single returns an initialized SmartDevice instance."""
host = "127.0.0.1" host = "127.0.0.1"
discovery_mock.ip = host discovery_mock.ip = host
@@ -178,7 +181,7 @@ async def test_discover_single(discovery_mock, custom_port, mocker):
assert x.config == config assert x.config == config
async def test_discover_single_hostname(discovery_mock, mocker): async def test_discover_single_hostname(discovery_mock, mocker: MockerFixture) -> None:
"""Make sure that discover_single returns an initialized SmartDevice instance.""" """Make sure that discover_single returns an initialized SmartDevice instance."""
host = "foobar" host = "foobar"
ip = "127.0.0.1" ip = "127.0.0.1"
@@ -198,7 +201,7 @@ async def test_discover_single_hostname(discovery_mock, mocker):
x = await Discover.discover_single(host, credentials=Credentials()) x = await Discover.discover_single(host, credentials=Credentials())
async def test_discover_credentials(mocker): async def test_discover_credentials(mocker: MockerFixture) -> None:
"""Make sure that discover gives credentials precedence over un and pw.""" """Make sure that discover gives credentials precedence over un and pw."""
host = "127.0.0.1" host = "127.0.0.1"
@@ -226,7 +229,7 @@ async def test_discover_credentials(mocker):
assert dp.mock_calls[3].kwargs["credentials"] is None assert dp.mock_calls[3].kwargs["credentials"] is None
async def test_discover_single_credentials(mocker): async def test_discover_single_credentials(mocker: MockerFixture) -> None:
"""Make sure that discover_single gives credentials precedence over un and pw.""" """Make sure that discover_single gives credentials precedence over un and pw."""
host = "127.0.0.1" host = "127.0.0.1"
@@ -254,7 +257,9 @@ async def test_discover_single_credentials(mocker):
assert dp.mock_calls[3].kwargs["credentials"] is None assert dp.mock_calls[3].kwargs["credentials"] is None
async def test_discover_single_unsupported(unsupported_device_info, mocker): async def test_discover_single_unsupported(
unsupported_device_info: dict, mocker: MockerFixture
) -> None:
"""Make sure that discover_single handles unsupported devices correctly.""" """Make sure that discover_single handles unsupported devices correctly."""
host = "127.0.0.1" host = "127.0.0.1"
@@ -265,7 +270,7 @@ async def test_discover_single_unsupported(unsupported_device_info, mocker):
await Discover.discover_single(host) await Discover.discover_single(host)
async def test_discover_single_no_response(mocker): async def test_discover_single_no_response(mocker: MockerFixture) -> None:
"""Make sure that discover_single handles no response correctly.""" """Make sure that discover_single handles no response correctly."""
host = "127.0.0.1" host = "127.0.0.1"
mocker.patch.object(_DiscoverProtocol, "do_discover") mocker.patch.object(_DiscoverProtocol, "do_discover")
@@ -285,7 +290,9 @@ INVALIDS = [
@pytest.mark.parametrize(("msg", "data"), INVALIDS) @pytest.mark.parametrize(("msg", "data"), INVALIDS)
async def test_discover_invalid_info(msg, data, mocker): async def test_discover_invalid_info(
msg: str, data: dict, mocker: MockerFixture
) -> None:
"""Make sure that invalid discovery information raises an exception.""" """Make sure that invalid discovery information raises an exception."""
host = "127.0.0.1" host = "127.0.0.1"
@@ -300,7 +307,7 @@ async def test_discover_invalid_info(msg, data, mocker):
await Discover.discover_single(host) await Discover.discover_single(host)
async def test_discover_send(mocker): async def test_discover_send(mocker: MockerFixture) -> None:
"""Test discovery parameters.""" """Test discovery parameters."""
discovery_timeout = 0 discovery_timeout = 0
discovery_ports = 3 discovery_ports = 3
@@ -312,7 +319,9 @@ async def test_discover_send(mocker):
assert transport.sendto.call_count == proto.discovery_packets * discovery_ports assert transport.sendto.call_count == proto.discovery_packets * discovery_ports
async def test_discover_datagram_received(mocker, discovery_data): async def test_discover_datagram_received(
mocker: MockerFixture, discovery_data: dict
) -> None:
"""Verify that datagram received fills discovered_devices.""" """Verify that datagram received fills discovered_devices."""
proto = _DiscoverProtocol() proto = _DiscoverProtocol()
@@ -338,7 +347,9 @@ async def test_discover_datagram_received(mocker, discovery_data):
@pytest.mark.parametrize(("msg", "data"), INVALIDS) @pytest.mark.parametrize(("msg", "data"), INVALIDS)
async def test_discover_invalid_responses(msg, data, mocker): async def test_discover_invalid_responses(
msg: str, data: dict, mocker: MockerFixture
) -> None:
"""Verify that we don't crash whole discovery if some devices in the network are sending unexpected data.""" """Verify that we don't crash whole discovery if some devices in the network are sending unexpected data."""
proto = _DiscoverProtocol() proto = _DiscoverProtocol()
mocker.patch("kasa.discover.json_loads", return_value=data) mocker.patch("kasa.discover.json_loads", return_value=data)
@@ -371,7 +382,9 @@ AUTHENTICATION_DATA_KLAP = {
@new_discovery @new_discovery
async def test_discover_single_authentication(discovery_mock, mocker): async def test_discover_single_authentication(
discovery_mock, mocker: MockerFixture
) -> None:
"""Make sure that discover_single handles authenticating devices correctly.""" """Make sure that discover_single handles authenticating devices correctly."""
host = "127.0.0.1" host = "127.0.0.1"
discovery_mock.ip = host discovery_mock.ip = host
@@ -398,7 +411,7 @@ async def test_discover_single_authentication(discovery_mock, mocker):
@new_discovery @new_discovery
async def test_device_update_from_new_discovery_info(discovery_mock): async def test_device_update_from_new_discovery_info(discovery_mock) -> None:
"""Make sure that new discovery devices update from discovery info correctly.""" """Make sure that new discovery devices update from discovery info correctly."""
discovery_data = discovery_mock.discovery_data discovery_data = discovery_mock.discovery_data
device_class = Discover._get_device_class(discovery_data) device_class = Discover._get_device_class(discovery_data)
@@ -420,7 +433,9 @@ async def test_device_update_from_new_discovery_info(discovery_mock):
assert device.modules assert device.modules
async def test_discover_single_http_client(discovery_mock, mocker): async def test_discover_single_http_client(
discovery_mock, mocker: MockerFixture
) -> None:
"""Make sure that discover_single returns an initialized SmartDevice instance.""" """Make sure that discover_single returns an initialized SmartDevice instance."""
host = "127.0.0.1" host = "127.0.0.1"
discovery_mock.ip = host discovery_mock.ip = host
@@ -437,7 +452,7 @@ async def test_discover_single_http_client(discovery_mock, mocker):
assert x.protocol._transport._http_client.client == http_client assert x.protocol._transport._http_client.client == http_client
async def test_discover_http_client(discovery_mock, mocker): async def test_discover_http_client(discovery_mock, mocker: MockerFixture) -> None:
"""Make sure that discover returns an initialized SmartDevice instance.""" """Make sure that discover returns an initialized SmartDevice instance."""
host = "127.0.0.1" host = "127.0.0.1"
discovery_mock.ip = host discovery_mock.ip = host
@@ -476,7 +491,7 @@ LEGACY_DISCOVER_DATA = {
class FakeDatagramTransport(asyncio.DatagramTransport): class FakeDatagramTransport(asyncio.DatagramTransport):
GHOST_PORT = 8888 GHOST_PORT = 8888
def __init__(self, dp, port, do_not_reply_count, unsupported=False): def __init__(self, dp, port, do_not_reply_count, unsupported=False) -> None:
self.dp = dp self.dp = dp
self.port = port self.port = port
self.do_not_reply_count = do_not_reply_count self.do_not_reply_count = do_not_reply_count
@@ -505,7 +520,9 @@ class FakeDatagramTransport(asyncio.DatagramTransport):
@pytest.mark.parametrize("port", [9999, 20002]) @pytest.mark.parametrize("port", [9999, 20002])
@pytest.mark.parametrize("do_not_reply_count", [0, 1, 2, 3, 4]) @pytest.mark.parametrize("do_not_reply_count", [0, 1, 2, 3, 4])
async def test_do_discover_drop_packets(mocker, port, do_not_reply_count): async def test_do_discover_drop_packets(
mocker: MockerFixture, port: int, do_not_reply_count: int
) -> None:
"""Make sure that _DiscoverProtocol handles authenticating devices correctly.""" """Make sure that _DiscoverProtocol handles authenticating devices correctly."""
host = "127.0.0.1" host = "127.0.0.1"
discovery_timeout = 0 discovery_timeout = 0
@@ -531,7 +548,9 @@ async def test_do_discover_drop_packets(mocker, port, do_not_reply_count):
[(FakeDatagramTransport.GHOST_PORT, True), (20002, False)], [(FakeDatagramTransport.GHOST_PORT, True), (20002, False)],
ids=["unknownport", "unsupporteddevice"], ids=["unknownport", "unsupporteddevice"],
) )
async def test_do_discover_invalid(mocker, port, will_timeout): async def test_do_discover_invalid(
mocker: MockerFixture, port: int, will_timeout: bool
) -> None:
"""Make sure that _DiscoverProtocol handles invalid devices correctly.""" """Make sure that _DiscoverProtocol handles invalid devices correctly."""
host = "127.0.0.1" host = "127.0.0.1"
discovery_timeout = 0 discovery_timeout = 0
@@ -550,7 +569,7 @@ async def test_do_discover_invalid(mocker, port, will_timeout):
assert dp.discover_task.cancelled() != will_timeout assert dp.discover_task.cancelled() != will_timeout
async def test_discover_propogates_task_exceptions(discovery_mock): async def test_discover_propogates_task_exceptions(discovery_mock) -> None:
"""Make sure that discover propogates callback exceptions.""" """Make sure that discover propogates callback exceptions."""
discovery_timeout = 0 discovery_timeout = 0
@@ -563,7 +582,7 @@ async def test_discover_propogates_task_exceptions(discovery_mock):
) )
async def test_do_discover_no_connection(mocker): async def test_do_discover_no_connection(mocker: MockerFixture) -> None:
"""Make sure that if the datagram connection doesnt start a TimeoutError is raised.""" """Make sure that if the datagram connection doesnt start a TimeoutError is raised."""
host = "127.0.0.1" host = "127.0.0.1"
discovery_timeout = 0 discovery_timeout = 0
@@ -581,7 +600,7 @@ async def test_do_discover_no_connection(mocker):
await dp.wait_for_discovery_to_complete() await dp.wait_for_discovery_to_complete()
async def test_do_discover_external_cancel(mocker): async def test_do_discover_external_cancel(mocker: MockerFixture) -> None:
"""Make sure that a cancel other than when target is discovered propogates.""" """Make sure that a cancel other than when target is discovered propogates."""
host = "127.0.0.1" host = "127.0.0.1"
discovery_timeout = 1 discovery_timeout = 1
@@ -601,7 +620,9 @@ async def test_do_discover_external_cancel(mocker):
@pytest.mark.xdist_group(name="caplog") @pytest.mark.xdist_group(name="caplog")
async def test_discovery_redaction(discovery_mock, caplog: pytest.LogCaptureFixture): async def test_discovery_redaction(
discovery_mock, caplog: pytest.LogCaptureFixture
) -> None:
"""Test query sensitive info redaction.""" """Test query sensitive info redaction."""
mac = "12:34:56:78:9A:BC" mac = "12:34:56:78:9A:BC"
@@ -636,7 +657,7 @@ async def test_discovery_redaction(discovery_mock, caplog: pytest.LogCaptureFixt
assert "12:34:56:00:00:00" in caplog.text assert "12:34:56:00:00:00" in caplog.text
async def test_discovery_decryption(): async def test_discovery_decryption() -> None:
"""Test discovery decryption.""" """Test discovery decryption."""
key = b"8\x89\x02\xfa\xf5Xs\x1c\xa1 H\x9a\x82\xc7\xd9\t" key = b"8\x89\x02\xfa\xf5Xs\x1c\xa1 H\x9a\x82\xc7\xd9\t"
iv = b"9=\xf8\x1bS\xcd0\xb5\x89i\xba\xfd^9\x9f\xfa" iv = b"9=\xf8\x1bS\xcd0\xb5\x89i\xba\xfd^9\x9f\xfa"
@@ -669,7 +690,7 @@ async def test_discovery_decryption():
assert dr.decrypted_data == data_dict assert dr.decrypted_data == data_dict
async def test_discover_try_connect_all(discovery_mock, mocker): async def test_discover_try_connect_all(discovery_mock, mocker: MockerFixture) -> None:
"""Test that device update is called on main.""" """Test that device update is called on main."""
if "result" in discovery_mock.discovery_data: if "result" in discovery_mock.discovery_data:
dev_class = get_device_class_from_family( dev_class = get_device_class_from_family(
@@ -703,7 +724,7 @@ async def test_discover_try_connect_all(discovery_mock, mocker):
return discovery_mock.query_data return discovery_mock.query_data
raise KasaException("Unable to execute query") raise KasaException("Unable to execute query")
async def _update(self, *args, **kwargs): async def _update(self, *args, **kwargs) -> None:
if ( if (
self.protocol.__class__ is protocol_class self.protocol.__class__ is protocol_class
and self.protocol._transport.__class__ is transport_class and self.protocol._transport.__class__ is transport_class
@@ -729,7 +750,7 @@ async def test_discover_try_connect_all(discovery_mock, mocker):
assert dev.protocol._transport._http_client.client == session assert dev.protocol._transport._http_client.client == session
async def test_discovery_device_repr(discovery_mock, mocker): async def test_discovery_device_repr(discovery_mock, mocker: MockerFixture) -> None:
"""Test that repr works when only discovery data is available.""" """Test that repr works when only discovery data is available."""
host = "foobar" host = "foobar"
ip = "127.0.0.1" ip = "127.0.0.1"

View File

@@ -1,10 +1,10 @@
import logging import logging
from unittest.mock import AsyncMock, patch from unittest.mock import AsyncMock, MagicMock, patch
import pytest import pytest
from pytest_mock import MockerFixture from pytest_mock import MockerFixture
from kasa import Device, Feature, KasaException from kasa import Device, Feature, KasaException, Module
from kasa.iot import IotStrip from kasa.iot import IotStrip
_LOGGER = logging.getLogger(__name__) _LOGGER = logging.getLogger(__name__)
@@ -32,7 +32,7 @@ def dummy_feature() -> Feature:
return feat return feat
def test_feature_api(dummy_feature: Feature): def test_feature_api(dummy_feature: Feature) -> None:
"""Test all properties of a dummy feature.""" """Test all properties of a dummy feature."""
assert dummy_feature.device is not None assert dummy_feature.device is not None
assert dummy_feature.name == "dummy_feature" assert dummy_feature.name == "dummy_feature"
@@ -47,7 +47,7 @@ def test_feature_api(dummy_feature: Feature):
@pytest.mark.parametrize( @pytest.mark.parametrize(
"read_only_type", [Feature.Type.Sensor, Feature.Type.BinarySensor] "read_only_type", [Feature.Type.Sensor, Feature.Type.BinarySensor]
) )
def test_feature_setter_on_sensor(read_only_type): def test_feature_setter_on_sensor(read_only_type: Feature.Type) -> None:
"""Test that creating a sensor feature with a setter causes an error.""" """Test that creating a sensor feature with a setter causes an error."""
with pytest.raises(ValueError, match="Invalid type for configurable feature"): with pytest.raises(ValueError, match="Invalid type for configurable feature"):
Feature( Feature(
@@ -60,14 +60,14 @@ def test_feature_setter_on_sensor(read_only_type):
) )
def test_feature_value(dummy_feature: Feature): def test_feature_value(dummy_feature: Feature) -> None:
"""Verify that property gets accessed on *value* access.""" """Verify that property gets accessed on *value* access."""
dummy_feature.attribute_getter = "test_prop" dummy_feature.attribute_getter = "test_prop"
dummy_feature.device.test_prop = "dummy" # type: ignore[attr-defined] dummy_feature.device.test_prop = "dummy" # type: ignore[attr-defined]
assert dummy_feature.value == "dummy" assert dummy_feature.value == "dummy"
def test_feature_value_container(mocker, dummy_feature: Feature): def test_feature_value_container(mocker: MockerFixture, dummy_feature: Feature) -> None:
"""Test that container's attribute is accessed when expected.""" """Test that container's attribute is accessed when expected."""
class DummyContainer: class DummyContainer:
@@ -86,13 +86,15 @@ def test_feature_value_container(mocker, dummy_feature: Feature):
mock_dev_prop.assert_not_called() mock_dev_prop.assert_not_called()
def test_feature_value_callable(dev, dummy_feature: Feature): def test_feature_value_callable(dev: Device, dummy_feature: Feature) -> None:
"""Verify that callables work as *attribute_getter*.""" """Verify that callables work as *attribute_getter*."""
dummy_feature.attribute_getter = lambda x: "dummy value" dummy_feature.attribute_getter = lambda x: "dummy value"
assert dummy_feature.value == "dummy value" assert dummy_feature.value == "dummy value"
async def test_feature_setter(dev, mocker, dummy_feature: Feature): async def test_feature_setter(
dev: Device, mocker: MockerFixture, dummy_feature: Feature
) -> None:
"""Verify that *set_value* calls the defined method.""" """Verify that *set_value* calls the defined method."""
mock_set_dummy = mocker.patch.object( mock_set_dummy = mocker.patch.object(
dummy_feature.device, "set_dummy", create=True, new_callable=AsyncMock dummy_feature.device, "set_dummy", create=True, new_callable=AsyncMock
@@ -102,14 +104,14 @@ async def test_feature_setter(dev, mocker, dummy_feature: Feature):
mock_set_dummy.assert_called_with("dummy value") mock_set_dummy.assert_called_with("dummy value")
async def test_feature_setter_read_only(dummy_feature): async def test_feature_setter_read_only(dummy_feature: Feature) -> None:
"""Verify that read-only feature raises an exception when trying to change it.""" """Verify that read-only feature raises an exception when trying to change it."""
dummy_feature.attribute_setter = None dummy_feature.attribute_setter = None
with pytest.raises(ValueError, match="Tried to set read-only feature"): with pytest.raises(ValueError, match="Tried to set read-only feature"):
await dummy_feature.set_value("value for read only feature") await dummy_feature.set_value("value for read only feature")
async def test_feature_action(mocker): async def test_feature_action(mocker: MockerFixture) -> None:
"""Test that setting value on button calls the setter.""" """Test that setting value on button calls the setter."""
feat = Feature( feat = Feature(
device=DummyDevice(), # type: ignore[arg-type] device=DummyDevice(), # type: ignore[arg-type]
@@ -129,7 +131,9 @@ async def test_feature_action(mocker):
@pytest.mark.xdist_group(name="caplog") @pytest.mark.xdist_group(name="caplog")
async def test_feature_choice_list(dummy_feature, caplog, mocker: MockerFixture): async def test_feature_choice_list(
dummy_feature: Feature, caplog: pytest.LogCaptureFixture, mocker: MockerFixture
) -> None:
"""Test the choice feature type.""" """Test the choice feature type."""
dummy_feature.type = Feature.Type.Choice dummy_feature.type = Feature.Type.Choice
dummy_feature.choices_getter = lambda: ["first", "second"] dummy_feature.choices_getter = lambda: ["first", "second"]
@@ -152,7 +156,7 @@ async def test_feature_choice_list(dummy_feature, caplog, mocker: MockerFixture)
@pytest.mark.parametrize("precision_hint", [1, 2, 3]) @pytest.mark.parametrize("precision_hint", [1, 2, 3])
async def test_precision_hint(dummy_feature, precision_hint): async def test_precision_hint(dummy_feature: Feature, precision_hint: int) -> None:
"""Test that precision hint works as expected.""" """Test that precision hint works as expected."""
dummy_value = 3.141593 dummy_value = 3.141593
dummy_feature.type = Feature.Type.Sensor dummy_feature.type = Feature.Type.Sensor
@@ -163,61 +167,73 @@ async def test_precision_hint(dummy_feature, precision_hint):
assert f"{round(dummy_value, precision_hint)} dummyunit" in repr(dummy_feature) assert f"{round(dummy_value, precision_hint)} dummyunit" in repr(dummy_feature)
async def test_feature_setters(dev: Device, mocker: MockerFixture): async def _test_feature_setter(
"""Test that all feature setters query something.""" dev: Device, feat: Feature, query_mock: MagicMock
) -> None:
"""Exercise one configurable feature setter."""
# setters that do not call set on the device itself. # setters that do not call set on the device itself.
internal_setters = {"pan_step", "tilt_step"} internal_setters = {"pan_step", "tilt_step"}
async def _test_feature(feat, query_mock): if feat.attribute_setter is None:
if feat.attribute_setter is None: return
return
# IotStrip makes calls via it's children # IotStrip makes calls via its children.
expecting_call = feat.id not in internal_setters and not isinstance( expecting_call = feat.id not in internal_setters and not isinstance(dev, IotStrip)
dev, IotStrip
)
if feat.type == Feature.Type.Number: if feat.type == Feature.Type.Number:
await feat.set_value(feat.minimum_value) await feat.set_value(feat.minimum_value)
elif feat.type == Feature.Type.Switch: elif feat.type == Feature.Type.Switch:
await feat.set_value(True) await feat.set_value(True)
elif feat.type == Feature.Type.Action: elif feat.type == Feature.Type.Action:
await feat.set_value("dummyvalue") await feat.set_value("dummyvalue")
elif feat.type == Feature.Type.Choice: elif feat.type == Feature.Type.Choice:
await feat.set_value(feat.choices[0]) choices = feat.choices
elif feat.type == Feature.Type.Unknown: if choices is None:
_LOGGER.warning("Feature '%s' has no type, cannot test the setter", feat) raise AssertionError(f"Choice feature {feat.id} has no choices")
expecting_call = False await feat.set_value(choices[0])
else: elif feat.type == Feature.Type.Unknown:
raise NotImplementedError(f"set_value not implemented for {feat.type}") _LOGGER.warning("Feature '%s' has no type, cannot test the setter", feat)
expecting_call = False
else:
raise NotImplementedError(f"set_value not implemented for {feat.type}")
if expecting_call: if expecting_call:
query_mock.assert_called() query_mock.assert_called()
async def _test_features(dev):
exceptions = []
for feat in dev.features.values():
try:
patch_dev = feat.container._device if feat.container else feat.device
with (
patch.object(patch_dev.protocol, "query", name=feat.id) as query,
# patch update in case feature setter does an update
patch.object(patch_dev, "update"),
):
await _test_feature(feat, query)
# we allow our own exceptions to avoid mocking valid responses
except KasaException:
pass
except Exception as ex:
ex.add_note(f"Exception when trying to set {feat} on {dev}")
exceptions.append(ex)
return exceptions async def _test_device_feature_setters(dev: Device) -> list[Exception]:
"""Exercise all configurable feature setters and collect unexpected errors."""
exceptions: list[Exception] = []
for feat in dev.features.values():
try:
if isinstance(feat.container, Module):
patch_dev = feat.container._device
elif feat.container is not None:
patch_dev = feat.container
else:
patch_dev = feat.device
with (
patch.object(patch_dev.protocol, "query", name=feat.id) as query,
# Patch update in case the feature setter invokes it.
patch.object(patch_dev, "update"),
):
await _test_feature_setter(dev, feat, query)
# Allow library exceptions to avoid mocking valid device responses.
except KasaException:
pass
except Exception as ex:
ex.add_note(f"Exception when trying to set {feat} on {dev}")
exceptions.append(ex)
exceptions = await _test_features(dev) return exceptions
async def test_feature_setters(dev: Device) -> None:
"""Test that all feature setters query something."""
exceptions = await _test_device_feature_setters(dev)
for child in dev.children: for child in dev.children:
exceptions.extend(await _test_features(child)) exceptions.extend(await _test_device_feature_setters(child))
if exceptions: if exceptions:
raise ExceptionGroup( raise ExceptionGroup(

View File

@@ -3,6 +3,8 @@ import re
import aiohttp import aiohttp
import pytest import pytest
from pytest_mock import MockerFixture
from yarl import URL
from kasa.deviceconfig import DeviceConfig from kasa.deviceconfig import DeviceConfig
from kasa.exceptions import ( from kasa.exceptions import (
@@ -61,9 +63,15 @@ from kasa.httpclient import HttpClient
), ),
) )
@pytest.mark.parametrize("mock_read", [False, True], ids=("post", "read")) @pytest.mark.parametrize("mock_read", [False, True], ids=("post", "read"))
async def test_httpclient_errors(mocker, error, error_raises, error_message, mock_read): async def test_httpclient_errors(
mocker: MockerFixture,
error: Exception,
error_raises: type[Exception],
error_message: str,
mock_read: bool,
) -> None:
class _mock_response: class _mock_response:
def __init__(self, status, error): def __init__(self, status, error) -> None:
self.status = status self.status = status
self.error = error self.error = error
self.call_count = 0 self.call_count = 0
@@ -71,10 +79,10 @@ async def test_httpclient_errors(mocker, error, error_raises, error_message, moc
async def __aenter__(self): async def __aenter__(self):
return self return self
async def __aexit__(self, exc_t, exc_v, exc_tb): async def __aexit__(self, exc_t, exc_v, exc_tb) -> None:
pass pass
async def read(self): async def read(self) -> bytes:
self.call_count += 1 self.call_count += 1
raise self.error raise self.error
@@ -99,7 +107,7 @@ async def test_httpclient_errors(mocker, error, error_raises, error_message, moc
+ re.escape(f", {repr(error)})") + re.escape(f", {repr(error)})")
) )
with pytest.raises(error_raises, match=error_message) as exc_info: with pytest.raises(error_raises, match=error_message) as exc_info:
await client.post("http://foobar") await client.post(URL("http://foobar"))
assert re.match(full_msg, str(exc_info.value)) assert re.match(full_msg, str(exc_info.value))
if mock_read: if mock_read:

View File

@@ -11,7 +11,7 @@ from .conftest import plug, plug_iot, plug_smart, switch_smart, wallswitch_iot
@plug_iot @plug_iot
async def test_plug_sysinfo(dev): async def test_plug_sysinfo(dev) -> None:
assert dev.sys_info is not None assert dev.sys_info is not None
SYSINFO_SCHEMA(dev.sys_info) SYSINFO_SCHEMA(dev.sys_info)
@@ -21,7 +21,7 @@ async def test_plug_sysinfo(dev):
@wallswitch_iot @wallswitch_iot
async def test_switch_sysinfo(dev): async def test_switch_sysinfo(dev) -> None:
assert dev.sys_info is not None assert dev.sys_info is not None
SYSINFO_SCHEMA(dev.sys_info) SYSINFO_SCHEMA(dev.sys_info)
@@ -31,7 +31,7 @@ async def test_switch_sysinfo(dev):
@plug_iot @plug_iot
async def test_plug_led(dev): async def test_plug_led(dev) -> None:
with pytest.deprecated_call(match="use: Module.Led in device.modules instead"): with pytest.deprecated_call(match="use: Module.Led in device.modules instead"):
original = dev.led original = dev.led
@@ -47,7 +47,7 @@ async def test_plug_led(dev):
@wallswitch_iot @wallswitch_iot
async def test_switch_led(dev): async def test_switch_led(dev) -> None:
with pytest.deprecated_call(match="use: Module.Led in device.modules instead"): with pytest.deprecated_call(match="use: Module.Led in device.modules instead"):
original = dev.led original = dev.led
@@ -63,7 +63,7 @@ async def test_switch_led(dev):
@plug_smart @plug_smart
async def test_plug_device_info(dev): async def test_plug_device_info(dev) -> None:
assert dev._info is not None assert dev._info is not None
assert dev.model is not None assert dev.model is not None
@@ -71,7 +71,7 @@ async def test_plug_device_info(dev):
@switch_smart @switch_smart
async def test_switch_device_info(dev): async def test_switch_device_info(dev) -> None:
assert dev._info is not None assert dev._info is not None
assert dev.model is not None assert dev.model is not None
@@ -81,5 +81,5 @@ async def test_switch_device_info(dev):
@plug @plug
def test_device_type_plug(dev): def test_device_type_plug(dev) -> None:
assert dev.device_type == DeviceType.Plug assert dev.device_type == DeviceType.Plug

View File

@@ -2,6 +2,7 @@ import asyncio
import pytest import pytest
import xdoctest import xdoctest
from pytest_mock import MockerFixture
from .conftest import ( from .conftest import (
get_device_for_fixture_protocol, get_device_for_fixture_protocol,
@@ -10,7 +11,7 @@ from .conftest import (
) )
def test_bulb_examples(mocker): def test_bulb_examples(mocker: MockerFixture) -> None:
"""Use KL130 (bulb with all features) to test the doctests.""" """Use KL130 (bulb with all features) to test the doctests."""
p = asyncio.run(get_device_for_fixture_protocol("KL130(US)_1.0_1.8.11.json", "IOT")) p = asyncio.run(get_device_for_fixture_protocol("KL130(US)_1.0_1.8.11.json", "IOT"))
asyncio.run(p.set_alias("Bedroom Bulb")) asyncio.run(p.set_alias("Bedroom Bulb"))
@@ -22,7 +23,7 @@ def test_bulb_examples(mocker):
assert not res["failed"] assert not res["failed"]
def test_iotdevice_examples(mocker): def test_iotdevice_examples(mocker: MockerFixture) -> None:
"""Use HS110 for emeter examples.""" """Use HS110 for emeter examples."""
p = asyncio.run(get_device_for_fixture_protocol("HS110(EU)_1.0_1.2.5.json", "IOT")) p = asyncio.run(get_device_for_fixture_protocol("HS110(EU)_1.0_1.2.5.json", "IOT"))
asyncio.run(p.set_alias("Bedroom Lamp Plug")) asyncio.run(p.set_alias("Bedroom Lamp Plug"))
@@ -34,7 +35,7 @@ def test_iotdevice_examples(mocker):
assert not res["failed"] assert not res["failed"]
def test_plug_examples(mocker): def test_plug_examples(mocker: MockerFixture) -> None:
"""Test plug examples.""" """Test plug examples."""
p = asyncio.run(get_device_for_fixture_protocol("HS110(EU)_1.0_1.2.5.json", "IOT")) p = asyncio.run(get_device_for_fixture_protocol("HS110(EU)_1.0_1.2.5.json", "IOT"))
asyncio.run(p.set_alias("Bedroom Lamp Plug")) asyncio.run(p.set_alias("Bedroom Lamp Plug"))
@@ -45,13 +46,13 @@ def test_plug_examples(mocker):
assert not res["failed"] assert not res["failed"]
def test_strip_examples(readmes_mock): def test_strip_examples(readmes_mock) -> None:
"""Test strip examples.""" """Test strip examples."""
res = xdoctest.doctest_module("kasa.iot.iotstrip", "all") res = xdoctest.doctest_module("kasa.iot.iotstrip", "all")
assert not res["failed"] assert not res["failed"]
def test_dimmer_examples(mocker): def test_dimmer_examples(mocker: MockerFixture) -> None:
"""Test dimmer examples.""" """Test dimmer examples."""
p = asyncio.run(get_device_for_fixture_protocol("HS220(US)_1.0_1.5.7.json", "IOT")) p = asyncio.run(get_device_for_fixture_protocol("HS220(US)_1.0_1.5.7.json", "IOT"))
mocker.patch("kasa.iot.iotdimmer.IotDimmer", return_value=p) mocker.patch("kasa.iot.iotdimmer.IotDimmer", return_value=p)
@@ -60,7 +61,7 @@ def test_dimmer_examples(mocker):
assert not res["failed"] assert not res["failed"]
def test_lightstrip_examples(mocker): def test_lightstrip_examples(mocker: MockerFixture) -> None:
"""Test lightstrip examples.""" """Test lightstrip examples."""
p = asyncio.run(get_device_for_fixture_protocol("KL430(US)_1.0_1.0.10.json", "IOT")) p = asyncio.run(get_device_for_fixture_protocol("KL430(US)_1.0_1.0.10.json", "IOT"))
asyncio.run(p.set_alias("Bedroom Lightstrip")) asyncio.run(p.set_alias("Bedroom Lightstrip"))
@@ -71,7 +72,7 @@ def test_lightstrip_examples(mocker):
assert not res["failed"] assert not res["failed"]
def test_discovery_examples(readmes_mock): def test_discovery_examples(readmes_mock) -> None:
"""Test discovery examples.""" """Test discovery examples."""
res = xdoctest.doctest_module("kasa.discover", "all") res = xdoctest.doctest_module("kasa.discover", "all")
assert res["n_passed"] > 0 assert res["n_passed"] > 0
@@ -79,7 +80,7 @@ def test_discovery_examples(readmes_mock):
assert not res["failed"] assert not res["failed"]
def test_deviceconfig_examples(readmes_mock): def test_deviceconfig_examples(readmes_mock) -> None:
"""Test discovery examples.""" """Test discovery examples."""
res = xdoctest.doctest_module("kasa.deviceconfig", "all") res = xdoctest.doctest_module("kasa.deviceconfig", "all")
assert res["n_passed"] > 0 assert res["n_passed"] > 0
@@ -87,7 +88,7 @@ def test_deviceconfig_examples(readmes_mock):
assert not res["failed"] assert not res["failed"]
def test_device_examples(readmes_mock): def test_device_examples(readmes_mock) -> None:
"""Test device examples.""" """Test device examples."""
res = xdoctest.doctest_module("kasa.device", "all") res = xdoctest.doctest_module("kasa.device", "all")
assert res["n_passed"] > 0 assert res["n_passed"] > 0
@@ -95,7 +96,7 @@ def test_device_examples(readmes_mock):
assert not res["failed"] assert not res["failed"]
def test_light_examples(readmes_mock): def test_light_examples(readmes_mock) -> None:
"""Test device examples.""" """Test device examples."""
res = xdoctest.doctest_module("kasa.interfaces.light", "all") res = xdoctest.doctest_module("kasa.interfaces.light", "all")
assert res["n_passed"] > 0 assert res["n_passed"] > 0
@@ -103,7 +104,7 @@ def test_light_examples(readmes_mock):
assert not res["failed"] assert not res["failed"]
def test_light_preset_examples(readmes_mock): def test_light_preset_examples(readmes_mock) -> None:
"""Test device examples.""" """Test device examples."""
res = xdoctest.doctest_module("kasa.interfaces.lightpreset", "all") res = xdoctest.doctest_module("kasa.interfaces.lightpreset", "all")
assert res["n_passed"] > 0 assert res["n_passed"] > 0
@@ -111,7 +112,7 @@ def test_light_preset_examples(readmes_mock):
assert not res["failed"] assert not res["failed"]
def test_light_effect_examples(readmes_mock): def test_light_effect_examples(readmes_mock) -> None:
"""Test device examples.""" """Test device examples."""
res = xdoctest.doctest_module("kasa.interfaces.lighteffect", "all") res = xdoctest.doctest_module("kasa.interfaces.lighteffect", "all")
assert res["n_passed"] > 0 assert res["n_passed"] > 0
@@ -119,7 +120,7 @@ def test_light_effect_examples(readmes_mock):
assert not res["failed"] assert not res["failed"]
def test_child_examples(readmes_mock): def test_child_examples(readmes_mock) -> None:
"""Test device examples.""" """Test device examples."""
res = xdoctest.doctest_module("kasa.smart.modules.childdevice", "all") res = xdoctest.doctest_module("kasa.smart.modules.childdevice", "all")
assert res["n_passed"] > 0 assert res["n_passed"] > 0
@@ -127,7 +128,7 @@ def test_child_examples(readmes_mock):
assert not res["failed"] assert not res["failed"]
def test_module_examples(readmes_mock): def test_module_examples(readmes_mock) -> None:
"""Test device examples.""" """Test device examples."""
res = xdoctest.doctest_module("kasa.module", "all") res = xdoctest.doctest_module("kasa.module", "all")
assert res["n_passed"] > 0 assert res["n_passed"] > 0
@@ -135,7 +136,7 @@ def test_module_examples(readmes_mock):
assert not res["failed"] assert not res["failed"]
def test_feature_examples(readmes_mock): def test_feature_examples(readmes_mock) -> None:
"""Test device examples.""" """Test device examples."""
res = xdoctest.doctest_module("kasa.feature", "all") res = xdoctest.doctest_module("kasa.feature", "all")
assert res["n_passed"] > 0 assert res["n_passed"] > 0
@@ -143,7 +144,7 @@ def test_feature_examples(readmes_mock):
assert not res["failed"] assert not res["failed"]
def test_tutorial_examples(readmes_mock): def test_tutorial_examples(readmes_mock) -> None:
"""Test discovery examples.""" """Test discovery examples."""
res = xdoctest.doctest_module("docs/tutorial.py", "all") res = xdoctest.doctest_module("docs/tutorial.py", "all")
assert res["n_passed"] > 0 assert res["n_passed"] > 0
@@ -151,7 +152,7 @@ def test_tutorial_examples(readmes_mock):
assert not res["failed"] assert not res["failed"]
def test_childsetup_examples(readmes_mock, mocker): def test_childsetup_examples(readmes_mock, mocker: MockerFixture) -> None:
"""Test device examples.""" """Test device examples."""
pair_resp = [ pair_resp = [
{ {
@@ -171,7 +172,7 @@ def test_childsetup_examples(readmes_mock, mocker):
@pytest.fixture @pytest.fixture
async def readmes_mock(mocker): async def readmes_mock(mocker: MockerFixture):
fixture_infos = { fixture_infos = {
"127.0.0.1": get_fixture_info("KP303(UK)_1.0_1.0.3.json", "IOT"), # Strip "127.0.0.1": get_fixture_info("KP303(UK)_1.0_1.0.3.json", "IOT"), # Strip
"127.0.0.2": get_fixture_info("HS110(EU)_1.0_1.2.5.json", "IOT"), # Plug "127.0.0.2": get_fixture_info("HS110(EU)_1.0_1.2.5.json", "IOT"), # Plug

View File

@@ -10,7 +10,7 @@ from .conftest import handle_turn_on, strip, strip_iot, turn_on
@strip @strip
@turn_on @turn_on
async def test_children_change_state(dev, turn_on): async def test_children_change_state(dev: Device, turn_on: bool) -> None:
await handle_turn_on(dev, turn_on) await handle_turn_on(dev, turn_on)
for plug in dev.children: for plug in dev.children:
orig_state = plug.is_on orig_state = plug.is_on
@@ -37,7 +37,7 @@ async def test_children_change_state(dev, turn_on):
@strip @strip
async def test_children_alias(dev): async def test_children_alias(dev: Device) -> None:
test_alias = "TEST1234" test_alias = "TEST1234"
for plug in dev.children: for plug in dev.children:
original = plug.alias original = plug.alias
@@ -45,13 +45,14 @@ async def test_children_alias(dev):
await dev.update() # TODO: set_alias does not call parent's update().. await dev.update() # TODO: set_alias does not call parent's update()..
assert plug.alias == test_alias assert plug.alias == test_alias
assert original is not None
await plug.set_alias(alias=original) await plug.set_alias(alias=original)
await dev.update() # TODO: set_alias does not call parent's update().. await dev.update() # TODO: set_alias does not call parent's update()..
assert plug.alias == original assert plug.alias == original
@strip @strip
async def test_children_on_since(dev): async def test_children_on_since(dev: Device) -> None:
on_sinces = [] on_sinces = []
for plug in dev.children: for plug in dev.children:
if plug.is_on: if plug.is_on:
@@ -69,7 +70,7 @@ async def test_children_on_since(dev):
@strip @strip
async def test_get_plug_by_name(dev: IotStrip): async def test_get_plug_by_name(dev: IotStrip) -> None:
name = dev.children[0].alias name = dev.children[0].alias
assert dev.get_plug_by_name(name) == dev.children[0] # type: ignore[arg-type] assert dev.get_plug_by_name(name) == dev.children[0] # type: ignore[arg-type]
@@ -78,7 +79,7 @@ async def test_get_plug_by_name(dev: IotStrip):
@strip @strip
async def test_get_plug_by_index(dev: IotStrip): async def test_get_plug_by_index(dev: IotStrip) -> None:
assert dev.get_plug_by_index(0) == dev.children[0] assert dev.get_plug_by_index(0) == dev.children[0]
with pytest.raises(KasaException): with pytest.raises(KasaException):
@@ -89,7 +90,7 @@ async def test_get_plug_by_index(dev: IotStrip):
@strip @strip
async def test_plug_features(dev: IotStrip): async def test_plug_features(dev: IotStrip) -> None:
"""Test the child plugs have default features.""" """Test the child plugs have default features."""
for child in dev.children: for child in dev.children:
assert "state" in child.features assert "state" in child.features
@@ -97,7 +98,7 @@ async def test_plug_features(dev: IotStrip):
@pytest.mark.skip("this test will wear out your relays") @pytest.mark.skip("this test will wear out your relays")
async def test_all_binary_states(dev): async def test_all_binary_states(dev) -> None:
# test every binary state # test every binary state
# TODO: this needs to be fixed, dev.plugs is not available for each device.. # TODO: this needs to be fixed, dev.plugs is not available for each device..
for state in range(2 ** len(dev.children)): for state in range(2 ** len(dev.children)):
@@ -142,7 +143,7 @@ async def test_all_binary_states(dev):
@strip @strip
def test_children_api(dev): def test_children_api(dev: Device) -> None:
"""Test the child device API.""" """Test the child device API."""
first = dev.children[0] first = dev.children[0]
first_by_get_child_device = dev.get_child_device(first.device_id) first_by_get_child_device = dev.get_child_device(first.device_id)
@@ -150,7 +151,7 @@ def test_children_api(dev):
@strip_iot @strip_iot
async def test_children_energy(dev: Device): async def test_children_energy(dev: Device) -> None:
if Module.Energy not in dev.modules: if Module.Energy not in dev.modules:
pytest.skip(f"skipping device {dev.model} does not support energy") pytest.skip(f"skipping device {dev.model} does not support energy")