diff --git a/.gitignore b/.gitignore index 573a4c08f..8749c5989 100644 --- a/.gitignore +++ b/.gitignore @@ -11,6 +11,7 @@ __pycache__/ # Coverage .coverage +coverage.xml # Tox .tox diff --git a/kasa/device_factory.py b/kasa/device_factory.py index ecb0d0a13..e1b27df05 100644 --- a/kasa/device_factory.py +++ b/kasa/device_factory.py @@ -34,6 +34,7 @@ KlapTransportV2, LinkieTransportV2, SslTransport, + TpapTransport, XorTransport, ) from .transports.sslaestransport import SslAesTransport @@ -196,6 +197,16 @@ def get_protocol(config: DeviceConfig, *, strict: bool = False) -> BaseProtocol protocol_name = ctype.device_family.value.split(".")[0] _LOGGER.debug("Finding protocol for %s", ctype.device_family) + if ctype.encryption_type is DeviceEncryptionType.Tpap and ( + ctype.device_family + in { + DeviceFamily.SmartIpCamera, + DeviceFamily.SmartTapoDoorbell, + } + or (ctype.device_family is DeviceFamily.SmartTapoHub and ctype.https) + ): + return SmartCamProtocol(transport=TpapTransport(config=config)) + if ctype.device_family in { DeviceFamily.SmartIpCamera, DeviceFamily.SmartTapoDoorbell, @@ -233,8 +244,11 @@ def get_protocol(config: DeviceConfig, *, strict: bool = False) -> BaseProtocol "SMART.KLAP": (SmartProtocol, KlapTransportV2), "SMART.KLAP.HTTPS": (SmartProtocol, KlapTransportV2), # H200 is device family SMART.TAPOHUB and uses SmartCamProtocol so use - # https to distuingish from SmartProtocol devices + # https to distinguish from SmartProtocol devices "SMART.AES.HTTPS": (SmartCamProtocol, SslAesTransport), + # TPAP devices (SMART.* with encrypt_type TPAP and TPAP/HTTPS). + "SMART.TPAP": (SmartProtocol, TpapTransport), + "SMART.TPAP.HTTPS": (SmartProtocol, TpapTransport), } if not (prot_tran_cls := supported_device_protocols.get(protocol_transport_key)): return None diff --git a/kasa/deviceconfig.py b/kasa/deviceconfig.py index ff0fbf8fe..182848160 100644 --- a/kasa/deviceconfig.py +++ b/kasa/deviceconfig.py @@ -62,6 +62,7 @@ class DeviceEncryptionType(Enum): Klap = "KLAP" Aes = "AES" Xor = "XOR" + Tpap = "TPAP" class DeviceFamily(Enum): diff --git a/kasa/exceptions.py b/kasa/exceptions.py index 1c764ad7a..df2b8cc6b 100644 --- a/kasa/exceptions.py +++ b/kasa/exceptions.py @@ -124,6 +124,7 @@ def from_int(value: int) -> SmartErrorCode: ACCOUNT_ERROR = -2101 STAT_ERROR = -2201 STAT_SAVE_ERROR = -2202 + STAT_ACCESS_ERROR = -2203 DST_ERROR = -2301 DST_SAVE_ERROR = -2302 @@ -190,6 +191,7 @@ def from_int(value: int) -> SmartErrorCode: SmartErrorCode.SESSION_TIMEOUT_ERROR, SmartErrorCode.SESSION_EXPIRED, SmartErrorCode.INVALID_NONCE, + SmartErrorCode.STAT_ACCESS_ERROR, ] SMART_AUTHENTICATION_ERRORS = [ diff --git a/kasa/transports/__init__.py b/kasa/transports/__init__.py index 192b4156a..0836b1081 100644 --- a/kasa/transports/__init__.py +++ b/kasa/transports/__init__.py @@ -6,17 +6,19 @@ from .linkietransport import LinkieTransportV2 from .sslaestransport import SslAesTransport from .ssltransport import SslTransport +from .tpaptransport import TpapTransport from .xortransport import XorEncryption, XorTransport __all__ = [ - "AesTransport", "AesEncyptionSession", - "SslTransport", - "SslAesTransport", + "AesTransport", "BaseTransport", "KlapTransport", "KlapTransportV2", "LinkieTransportV2", - "XorTransport", + "SslAesTransport", + "SslTransport", + "TpapTransport", "XorEncryption", + "XorTransport", ] diff --git a/kasa/transports/tpaptransport.py b/kasa/transports/tpaptransport.py new file mode 100644 index 000000000..3a6471417 --- /dev/null +++ b/kasa/transports/tpaptransport.py @@ -0,0 +1,1434 @@ +"""Implementation of the TP-Link TPAP transport.""" + +from __future__ import annotations + +import asyncio +import base64 +import hashlib +import hmac +import logging +import secrets +import ssl +import struct +from datetime import UTC, datetime +from typing import TYPE_CHECKING, Any + +from cryptography import x509 +from cryptography.exceptions import InvalidSignature +from cryptography.hazmat.primitives import hashes, serialization +from cryptography.hazmat.primitives.asymmetric import ec, padding, rsa +from cryptography.hazmat.primitives.ciphers import algorithms +from cryptography.hazmat.primitives.ciphers.aead import AESCCM, ChaCha20Poly1305 +from cryptography.hazmat.primitives.cmac import CMAC +from cryptography.hazmat.primitives.kdf.hkdf import HKDF +from ecdsa import NIST256p, NIST384p, NIST521p, ellipticcurve +from ecdsa.curves import Curve +from ecdsa.ellipticcurve import CurveFp, PointJacobi +from passlib.hash import md5_crypt, sha256_crypt +from yarl import URL + +from kasa.credentials import Credentials +from kasa.deviceconfig import DeviceConfig, DeviceFamily +from kasa.exceptions import ( + SMART_AUTHENTICATION_ERRORS, + SMART_RETRYABLE_ERRORS, + AuthenticationError, + DeviceError, + KasaException, + SmartErrorCode, + _ConnectionError, + _RetryableError, +) +from kasa.httpclient import HttpClient +from kasa.json import dumps as json_dumps +from kasa.json import loads as json_loads + +from .basetransport import BaseTransport + +_LOGGER = logging.getLogger(__name__) + + +class TpapEncryptionSession: + """Class for a TPAP encryption session.""" + + PAKE_CONTEXT_TAG = b"PAKE V1" + TAG_LEN = 16 + NONCE_LEN = 12 + CIPHER_PARAMETERS = { + "aes_128_ccm": ( + b"tp-kdf-salt-aes128-key", + b"tp-kdf-info-aes128-key", + b"tp-kdf-salt-aes128-iv", + b"tp-kdf-info-aes128-iv", + 16, + ), + "aes_256_ccm": ( + b"tp-kdf-salt-aes256-key", + b"tp-kdf-info-aes256-key", + b"tp-kdf-salt-aes256-iv", + b"tp-kdf-info-aes256-iv", + 32, + ), + "chacha20_poly1305": ( + b"tp-kdf-salt-chacha20-key", + b"tp-kdf-info-chacha20-key", + b"tp-kdf-salt-chacha20-iv", + b"tp-kdf-info-chacha20-iv", + 32, + ), + } + + def __init__(self, transport: TpapTransport) -> None: + self._transport = transport + self._handshake_lock = asyncio.Lock() + self._device_mac: str = "" + self._tpap_tls: int | None = None + self._tpap_port: int | None = None + self._tpap_dac: bool = False + self._tpap_pake: list[int] = [] + self._tpap_user_hash_type: int | None = None + self._session_id: str | None = None + self._sequence: int | None = None + self._ds_url: URL | None = None + self._cipher_id: str = "aes_128_ccm" + self._hkdf_hash: str = "SHA256" + self._key: bytes | None = None + self._base_nonce: bytes | None = None + self._shared_key: bytes | None = None + self._expected_dev_confirm: str | None = None + self._dac_nonce_base64: str | None = None + self._user_random: str | None = None + self.reset() + + @property + def _uses_camera_auth(self) -> bool: + device_family = self._transport._config.connection_type.device_family + return device_family in self._transport.CAMERA_AUTH_DEVICE_FAMILIES + + @property + def _uses_robot_tpap_auth(self) -> bool: + device_family = self._transport._config.connection_type.device_family + return device_family in self._transport.ROBOT_TPAP_DEVICE_FAMILIES + + @property + def tls_mode(self) -> int | None: + """The discovered TLS mode.""" + return self._tpap_tls + + @property + def ds_url(self) -> URL | None: + """The secure DS endpoint for the current session.""" + return self._ds_url + + @property + def device_mac(self) -> str: + """The discovered device MAC.""" + return self._device_mac + + @property + def is_established(self) -> bool: + """Return true if the session is established.""" + return ( + self._session_id is not None + and self._sequence is not None + and self._ds_url is not None + and self._key is not None + and self._base_nonce is not None + ) + + def _invalidate_session(self) -> None: + """Reset live session state while preserving discovered metadata.""" + self._session_id = None + self._sequence = None + self._ds_url = None + self._cipher_id = "aes_128_ccm" + self._hkdf_hash = "SHA256" + self._key = None + self._base_nonce = None + self._shared_key = None + self._expected_dev_confirm = None + self._dac_nonce_base64 = None + self._user_random = None + + def reset(self) -> None: + """Reset discovered metadata and session state.""" + self._transport._ssl_context = None + self._transport._app_url = self._transport._get_initial_app_url() + self._device_mac = self._transport._known_device_mac + self._tpap_tls = self._transport._known_tpap_tls + self._tpap_port = self._transport._known_tpap_port + self._tpap_dac = self._transport._known_tpap_dac + self._tpap_pake = list(self._transport._known_tpap_pake) + self._tpap_user_hash_type = self._transport._known_tpap_user_hash_type + self._invalidate_session() + + @staticmethod + def _parse_optional_int(value: Any) -> int | None: + if value is None: + return None + try: + return int(value) + except (TypeError, ValueError): + return None + + @staticmethod + def _require_result_dict(response: dict[str, Any]) -> dict[str, Any]: + result = response.get("result") + if not isinstance(result, dict): + raise KasaException("TPAP response missing result object") + return result + + async def perform_handshake(self) -> None: + """Perform the handshake.""" + async with self._handshake_lock: + if self.is_established: + return + + self.reset() + _LOGGER.debug( + "TPAP: starting handshake with %s", + self._transport._host, + ) + + await self._discover() + await self._perform_auth_handshake() + + _LOGGER.debug("TPAP: handshake complete with %s", self._transport._host) + + async def _discover(self) -> None: + body = {"method": "login", "params": {"sub_method": "discover"}} + status, data = await self._transport._http_client.post( + self._transport._app_url.with_path("/"), + json=body, + headers=self._transport.COMMON_HEADERS, + ssl=await self._transport.get_ssl_context(), + ) + if status != 200 or not isinstance(data, dict): + raise KasaException( + f"TPAP discover failed for {self._transport._host}: " + f"{status} {type(data)}" + ) + + self._handle_response_error_code(data, "discover") + result = self._require_result_dict(data) + tpap = result.get("tpap") + if not isinstance(tpap, dict): + raise KasaException("TPAP discover response missing tpap object") + + self._device_mac = str(result.get("mac") or "") + self._tpap_tls = self._parse_optional_int(tpap.get("tls")) + self._tpap_port = self._parse_optional_int(tpap.get("port")) + self._tpap_dac = bool(tpap.get("dac")) + self._tpap_pake = list(tpap.get("pake") or []) + self._tpap_user_hash_type = self._parse_optional_int(tpap.get("user_hash_type")) + + self._transport._known_device_mac = self._device_mac + self._transport._known_tpap_tls = self._tpap_tls + self._transport._known_tpap_port = self._tpap_port + self._transport._known_tpap_dac = self._tpap_dac + self._transport._known_tpap_pake = list(self._tpap_pake) + self._transport._known_tpap_user_hash_type = self._tpap_user_hash_type + self._update_transport_url() + + # Discover runs before we know the real TLS mode, so rebuild for auth. + self._transport._ssl_context = None + + async def _login(self, params: dict[str, Any], *, step_name: str) -> dict[str, Any]: + body = {"method": "login", "params": params} + ssl_context = await self._transport.get_ssl_context() + status, data = await self._transport._http_client.post( + self._transport._app_url.with_path("/"), + json=body, + headers=self._transport.COMMON_HEADERS, + ssl=ssl_context, + ) + if status != 200 or not isinstance(data, dict): + raise KasaException( + f"TPAP {step_name} failed for {self._transport._host}: " + f"{status} {type(data)}" + ) + + self._handle_response_error_code(data, step_name) + return self._require_result_dict(data) + + def _update_transport_url(self) -> None: + self._transport._app_url = self._transport._build_app_url( + tls_mode=self._tpap_tls, + port=self._tpap_port, + ) + + def _handle_response_error_code( + self, response: dict[str, Any], action: str + ) -> None: + """Handle response errors to request reauth etc.""" + error_code_raw = response.get("error_code") + try: + error_code = SmartErrorCode.from_int(error_code_raw) + except (TypeError, ValueError): + _LOGGER.warning( + "Device %s received unknown error code: %s", + self._transport._host, + error_code_raw, + ) + error_code = SmartErrorCode.INTERNAL_UNKNOWN_ERROR + + if error_code is SmartErrorCode.SUCCESS: + return + + full = ( + f"TPAP {action} failed for {self._transport._host}: " + f"{error_code.name}({error_code.value})" + ) + if error_code in SMART_RETRYABLE_ERRORS: + raise _RetryableError(full, error_code=error_code) + if error_code in SMART_AUTHENTICATION_ERRORS: + self._invalidate_session() + raise AuthenticationError(full, error_code=error_code) + raise DeviceError(full, error_code=error_code) + + async def _perform_auth_handshake(self) -> None: + passcode_type = self._get_passcode_type() + if passcode_type is None: + raise AuthenticationError( + f"TPAP: no supported passcode type for {self._transport._host}" + ) + + candidate_secrets = self._get_candidate_secrets() + if not candidate_secrets: + raise AuthenticationError( + f"TPAP: no credential candidates available for {self._transport._host}" + ) + + register_username = self._get_register_username() + candidate_count = len(candidate_secrets) + last_error: KasaException | None = None + + for attempt, candidate_secret in enumerate(candidate_secrets, start=1): + self._shared_key = None + self._expected_dev_confirm = None + self._dac_nonce_base64 = None + self._user_random = None + self._user_random = self._base64(secrets.token_bytes(32)) + register_params = { + "sub_method": "pake_register", + "username": register_username, + "user_random": self._user_random, + "cipher_suites": [1], + "encryption": ["aes_128_ccm"], + "passcode_type": passcode_type, + "stok": None, + } + + try: + register_result = await self._login( + register_params, step_name="pake_register" + ) + credentials_string = self._resolve_credentials( + register_result, + candidate_secret, + passcode_type=passcode_type, + ) + share_params = self._build_share_params_from_register( + register_result, credentials_string + ) + if self._use_dac_certification(): + self._dac_nonce_base64 = self._base64(secrets.token_bytes(32)) + share_params["dac_nonce"] = self._dac_nonce_base64 + + share_result = await self._login(share_params, step_name="pake_share") + self._establish_session_from_share_result(share_result) + return + except (_RetryableError, _ConnectionError): + raise + except KasaException as exc: + last_error = exc + if attempt < candidate_count: + _LOGGER.debug( + "TPAP: credential candidate %d/%d failed for %s: %s", + attempt, + candidate_count, + self._transport._host, + exc, + ) + + if last_error is not None: + if self._uses_camera_auth and 2 in self._tpap_pake: + _LOGGER.debug( + "TPAP: all password-based camera candidates failed for %s", + self._transport._host, + ) + raise last_error + + raise KasaException( # pragma: no cover + "TPAP: handshake did not produce a session" + ) + + @staticmethod + def _md5_hex(value: str) -> str: + return hashlib.md5(value.encode()).hexdigest() # noqa: S324 + + @staticmethod + def _sha256_hex_upper(value: str) -> str: + return hashlib.sha256(value.encode()).hexdigest().upper() # noqa: S324 + + def _get_register_username(self) -> str: + return ( + self._sha256_hex_upper("admin") + if self._tpap_user_hash_type == 1 + else self._md5_hex("admin") + ) + + @staticmethod + def _base64(value: bytes) -> str: + return base64.b64encode(value).decode() + + @staticmethod + def _unbase64(value: str) -> bytes: + return base64.b64decode(value) + + @staticmethod + def _sec1_to_xy(sec1: bytes, curve: ec.EllipticCurve) -> tuple[int, int]: + public_key = ec.EllipticCurvePublicKey.from_encoded_point(curve, sec1) + numbers = public_key.public_numbers() + return numbers.x, numbers.y + + @staticmethod + def _xy_to_uncompressed(x: int, y: int, curve: ec.EllipticCurve) -> bytes: + numbers = ec.EllipticCurvePublicNumbers(x, y, curve) + public_key = numbers.public_key() + return public_key.public_bytes( + encoding=serialization.Encoding.X962, + format=serialization.PublicFormat.UncompressedPoint, + ) + + @staticmethod + def _len8le(value: bytes) -> bytes: + return len(value).to_bytes(8, "little") + value + + @staticmethod + def _encode_w(value: int) -> bytes: + minimal_length = 1 if value == 0 else (value.bit_length() + 7) // 8 + unsigned = value.to_bytes(minimal_length, "big", signed=False) + if minimal_length % 2 == 0: + return unsigned + if unsigned[0] & 0x80: + return b"\x00" + unsigned + return unsigned + + @staticmethod + def _hash(algorithm: str, data: bytes) -> bytes: + if algorithm.upper() == "SHA512": + return hashlib.sha512(data).digest() + return hashlib.sha256(data).digest() + + @staticmethod + def _hkdf_expand(label: str, prk: bytes, digest_len: int, algorithm: str) -> bytes: + hkdf_algorithm = ( + hashes.SHA512() if algorithm.upper() == "SHA512" else hashes.SHA256() + ) + zero_salt = b"\x00" * digest_len + return HKDF( + algorithm=hkdf_algorithm, + length=digest_len, + salt=zero_salt, + info=label.encode(), + ).derive(prk) + + @staticmethod + def _hmac(algorithm: str, key: bytes, data: bytes) -> bytes: + digest = hashlib.sha512 if algorithm.upper() == "SHA512" else hashlib.sha256 + return hmac.new(key, data, digest).digest() + + @staticmethod + def _cmac_aes(key: bytes, data: bytes) -> bytes: + cmac = CMAC(algorithms.AES(key)) + cmac.update(data) + return cmac.finalize() + + @staticmethod + def _pbkdf2_sha256( + password: bytes, salt: bytes, iterations: int, length: int + ) -> bytes: + return hashlib.pbkdf2_hmac("sha256", password, salt, iterations, length) + + @classmethod + def _derive_ab( + cls, credentials: bytes, salt: bytes, iterations: int, hash_len: int = 32 + ) -> tuple[int, int]: + i_d = hash_len + 8 + derived = cls._pbkdf2_sha256(credentials, salt, iterations, 2 * i_d) + return ( + int.from_bytes(derived[:i_d], "big"), + int.from_bytes(derived[i_d:], "big"), + ) + + @staticmethod + def _sha1_hex(value: str) -> str: + return hashlib.sha1(value.encode()).hexdigest() # noqa: S324 + + @classmethod + def _authkey_mask(cls, passcode: str, tmpkey: str, dictionary: str) -> str: + masked = [] + max_length = max(len(tmpkey), len(passcode)) + for index in range(max_length): + lhs = ord(passcode[index]) if index < len(passcode) else 0xBB + rhs = ord(tmpkey[index]) if index < len(tmpkey) else 0xBB + masked.append(dictionary[(lhs ^ rhs) % len(dictionary)]) + return "".join(masked) + + @classmethod + def _sha1_username_mac_shadow( + cls, username: str, mac12hex: str, password: str + ) -> str: + if ( + not username + or len(mac12hex) != 12 + or not all(char in "0123456789abcdefABCDEF" for char in mac12hex) + ): + return password + + mac = ":".join(mac12hex[index : index + 2] for index in range(0, 12, 2)).upper() + return cls._sha1_hex(cls._md5_hex(username) + "_" + mac) + + @classmethod + def _md5_crypt(cls, password: str, prefix: str) -> str | None: + if not prefix or not prefix.startswith("$1$") or len(password) > 30000: + return None + + spec = prefix[3:] + if "$" in spec: + spec = spec.split("$", 1)[0] + return md5_crypt.using(salt=spec[:8]).hash(password) + + @classmethod + def _sha256_crypt( + cls, password: str, prefix: str, rounds_from_params: int | None = None + ) -> str | None: + if not prefix: + return None + + default_rounds = 5000 + min_rounds = 1000 + max_rounds = 999_999_999 + + spec = prefix[3:] if prefix.startswith("$5$") else prefix + rounds: int | None = None + + if spec.startswith("rounds="): + rounds_part, _, salt_part = spec.partition("$") + try: + rounds = int(rounds_part.split("=", 1)[1]) + except ValueError: + rounds = default_rounds + rounds = max(min_rounds, min(max_rounds, rounds)) + salt = salt_part + else: + salt = spec.split("$", 1)[0] if "$" in spec else spec + + if rounds_from_params is not None: + try: + parsed_rounds = int(rounds_from_params) + except (TypeError, ValueError): + parsed_rounds = default_rounds + rounds = max(min_rounds, min(max_rounds, parsed_rounds)) + + salt = salt[:16] + if rounds is not None: + return sha256_crypt.using(rounds=rounds, salt=salt).hash(password) + return sha256_crypt.using(salt=salt).hash(password) + + @classmethod + def _build_credentials( + cls, extra_crypt: dict | None, username: str, passcode: str, mac_no_colon: str + ) -> str: + if not extra_crypt: + return f"{username}/{passcode}" if username else passcode + + crypt_type = (extra_crypt.get("type") or "").lower() + params = extra_crypt.get("params") + if not isinstance(params, dict): + params = {} + + if crypt_type == "password_shadow": + try: + passwd_id = int(params.get("passwd_id", 0)) + except (TypeError, ValueError): + _LOGGER.debug("TPAP: invalid passwd_id, using passcode") + return passcode + prefix = str(params.get("passwd_prefix", "") or "") + if passwd_id == 1: + return cls._md5_crypt(passcode, prefix) or passcode + if passwd_id == 2: + return cls._sha1_hex(passcode) + if passwd_id == 3: + return cls._sha1_username_mac_shadow(username, mac_no_colon, passcode) + if passwd_id == 5: + return ( + cls._sha256_crypt( + passcode, + prefix, + rounds_from_params=params.get("passwd_rounds"), + ) + or passcode + ) + return passcode + + if crypt_type == "password_authkey": + tmpkey = str(params.get("authkey_tmpkey", "") or "") + dictionary = str(params.get("authkey_dictionary", "") or "") + if tmpkey and dictionary: + return cls._authkey_mask(passcode, tmpkey, dictionary) + return passcode + + if crypt_type == "password_sha_with_salt": + try: + sha_name = int(params.get("sha_name", -1)) + except (TypeError, ValueError): + _LOGGER.debug("TPAP: invalid sha_name, using passcode") + return passcode + sha_salt_b64 = str(params.get("sha_salt", "") or "") + username_hint = "admin" if sha_name == 0 else "user" + try: + decoded_salt = base64.b64decode(sha_salt_b64).decode() + except Exception: + _LOGGER.debug("TPAP: invalid base64 salt, using passcode") + return passcode + return hashlib.sha256( + (username_hint + decoded_salt + passcode).encode() + ).hexdigest() + + return f"{username}/{passcode}" if username else passcode + + def _suite_hash_name(self, suite_type: int) -> str: + return "SHA512" if suite_type in (2, 4, 5, 7, 9) else "SHA256" + + def _suite_mac_is_cmac(self, suite_type: int) -> bool: + return suite_type in (8, 9) + + def _use_dac_certification(self) -> bool: + return self._tpap_tls == 0 and self._tpap_dac + + @staticmethod + def _mac_pass_from_device_mac(mac_colon: str) -> str: + mac_hex = mac_colon.replace(":", "").replace("-", "") + try: + mac_bytes = bytes.fromhex(mac_hex) + except ValueError as exc: + raise KasaException( + "Invalid device MAC for TPAP default passcode derivation" + ) from exc + if len(mac_bytes) < 6: + raise KasaException( + "Device MAC is too short for TPAP default passcode derivation" + ) + seed = b"GqY5o136oa4i6VprTlMW2DpVXxmfW8" + ikm = seed + mac_bytes[3:6] + mac_bytes[0:3] + return ( + HKDF( + algorithm=hashes.SHA256(), + length=32, + salt=b"tp-kdf-salt-default-passcode", + info=b"tp-kdf-info-default-passcode", + ) + .derive(ikm) + .hex() + .upper() + ) + + def _get_passcode_type(self) -> str | None: + passcode_type_order: tuple[tuple[tuple[int, ...], str], ...] + default_passcode_type: str | None + + if self._uses_camera_auth: + passcode_type_order = ( + ((2, 1), "userpw"), + ((0,), "default_userpw"), + ((3,), "shared_token"), + ) + default_passcode_type = None + elif self._uses_robot_tpap_auth: + passcode_type_order = ( + ((0,), "default_userpw"), + ((2,), "userpw"), + ((3,), "shared_token"), + ) + default_passcode_type = "default_userpw" + else: + passcode_type_order = ( + ((0,), "default_userpw"), + ((2, 5), "userpw"), + ((3,), "shared_token"), + ) + default_passcode_type = "default_userpw" + + pake = set(self._tpap_pake) + for pake_values, passcode_type in passcode_type_order: + if pake.intersection(pake_values): + return passcode_type + return default_passcode_type + + def _get_candidate_secrets(self, passcode_type: str | None = None) -> list[str]: + passcode_type = passcode_type or self._get_passcode_type() + if passcode_type is None: + return [] + if passcode_type == "default_userpw": + return ( + [self._mac_pass_from_device_mac(self._device_mac)] + if self._device_mac + else [] + ) + creds = self._transport._config.credentials + password = (creds.password if creds else "") or "" + if not self._uses_camera_auth: + return [password] + if passcode_type == "shared_token": + return [self._md5_hex(password)] + if 2 not in self._tpap_pake: + return [password] + return list( + dict.fromkeys( + [ + self._md5_hex(password), + self._sha256_hex_upper(password), + ] + ) + ) + + def _resolve_credentials( + self, + register_result: dict[str, Any], + candidate_secret: str, + *, + passcode_type: str | None = None, + ) -> str: + if (passcode_type or self._get_passcode_type()) == "default_userpw": + return candidate_secret + extra_crypt_value = register_result.get("extra_crypt") + extra_crypt = extra_crypt_value if isinstance(extra_crypt_value, dict) else {} + if self._uses_camera_auth and not extra_crypt: + return candidate_secret + creds = self._transport._config.credentials + username = ( + "" if self._uses_camera_auth else (creds.username if creds else "") or "" + ) + mac_no_colon = self._device_mac.replace(":", "").replace("-", "") + return self._build_credentials( + extra_crypt, + username, + candidate_secret, + mac_no_colon, + ) + + @staticmethod + def _suite_parameters( + suite_type: int, + ) -> tuple[bytes, bytes, Curve, ec.EllipticCurve]: + if suite_type in (1, 2, 8, 9): + return ( + bytes.fromhex( + "02886e2f97ace46e55ba9dd7242579f2993b64e16ef3dcab95afd497333d8fa12f" + ), + bytes.fromhex( + "03d8bbd6c639c62937b04d997f38c3770719c629d7014d49a24b4f98baa1292b49" + ), + NIST256p, + ec.SECP256R1(), + ) + if suite_type in (3, 4): + return ( + bytes.fromhex( + "030ff0895ae5ebf6187080a82d82b42e2765e3b2f8749c7e05eba366434b363d3dc36f15314739074d2eb8613fceec2853" + ), + bytes.fromhex( + "02c72cf2e390853a1c1c4ad816a62fd15824f56078918f43f922ca21518f9c543bb252c5490214cf9aa3f0baab4b665c10" + ), + NIST384p, + ec.SECP384R1(), + ) + if suite_type == 5: + return ( + bytes.fromhex( + "02003f06f38131b2ba2600791e82488e8d20ab889af753a41806c5db18d37d85608cfae06b82e4a72cd744c719193562a653ea1f119eef9356907edc9b56979962d7aa" + ), + bytes.fromhex( + "0200c7924b9ec017f3094562894336a53c50167ba8c5963876880542bc669e494b2532d76c5b53dfb349fdf69154b9e0048c58a42e8ed04cef052a3bc349d95575cd25" + ), + NIST521p, + ec.SECP521R1(), + ) + raise KasaException(f"Unsupported TPAP suite type: {suite_type}") + + def _build_share_params_from_register( + self, register_result: dict[str, Any], credentials_string: str + ) -> dict[str, Any]: + if self._user_random is None: + raise KasaException("TPAP user random not initialized") + + dev_random = str(register_result.get("dev_random") or "") + dev_salt = str(register_result.get("dev_salt") or "") + dev_share = str(register_result.get("dev_share") or "") + for field, value in ( + ("dev_random", dev_random), + ("dev_salt", dev_salt), + ("dev_share", dev_share), + ): + if not value: + raise KasaException(f"TPAP register response missing {field}") + + suite_type_value = register_result.get("cipher_suites") + if suite_type_value is None: + raise KasaException("TPAP register response has invalid cipher_suites") + try: + suite_type = int(suite_type_value) + except (TypeError, ValueError) as exc: + raise KasaException( + "TPAP register response has invalid cipher_suites" + ) from exc + + iterations_value = register_result.get("iterations") + if iterations_value is None: + raise KasaException("TPAP register response has invalid iterations") + try: + iterations = int(iterations_value) + except (TypeError, ValueError) as exc: + raise KasaException( + "TPAP register response has invalid iterations" + ) from exc + + if iterations <= 0: + raise KasaException("TPAP register response has invalid iterations") + + encryption = str(register_result.get("encryption") or "") + if not encryption: + raise KasaException("TPAP register response missing encryption") + chosen_cipher = self._normalize_cipher_id(encryption) + if chosen_cipher not in self.CIPHER_PARAMETERS: + raise KasaException(f"Unsupported TPAP session cipher: {encryption}") + + self._cipher_id = chosen_cipher + self._hkdf_hash = self._suite_hash_name(suite_type) + + m_comp, n_comp, nist, crypto_curve = self._suite_parameters(suite_type) + curve: CurveFp = nist.curve + generator: PointJacobi = nist.generator + order = generator.order() + g_point = generator + + m_x, m_y = self._sec1_to_xy(m_comp, crypto_curve) + n_x, n_y = self._sec1_to_xy(n_comp, crypto_curve) + m_point = ellipticcurve.Point(curve, m_x, m_y, order) + n_point = ellipticcurve.Point(curve, n_x, n_y, order) + + credential_bytes = credentials_string.encode() + a_value, b_value = self._derive_ab( + credential_bytes, self._unbase64(dev_salt), iterations, 32 + ) + w_value = a_value % order + h_value = b_value % order + x_value = secrets.randbelow(order - 1) + 1 + + l_point = x_value * g_point + w_value * m_point + l_encoded = self._xy_to_uncompressed(l_point.x(), l_point.y(), crypto_curve) + + device_share_bytes = self._unbase64(dev_share) + r_x, r_y = self._sec1_to_xy(device_share_bytes, crypto_curve) + r_point = ellipticcurve.Point(curve, r_x, r_y, order) + r_encoded = self._xy_to_uncompressed(r_point.x(), r_point.y(), crypto_curve) + + r_prime = r_point + (-(w_value * n_point)) + z_point = x_value * r_prime + v_point = (h_value % order) * r_prime + + z_encoded = self._xy_to_uncompressed(z_point.x(), z_point.y(), crypto_curve) + v_encoded = self._xy_to_uncompressed(v_point.x(), v_point.y(), crypto_curve) + m_encoded = self._xy_to_uncompressed(m_point.x(), m_point.y(), crypto_curve) + n_encoded = self._xy_to_uncompressed(n_point.x(), n_point.y(), crypto_curve) + + context_hash = self._hash( + self._hkdf_hash, + self.PAKE_CONTEXT_TAG + + self._unbase64(self._user_random) + + self._unbase64(dev_random), + ) + w_encoded = self._encode_w(w_value) + + transcript = ( + self._len8le(context_hash) + + self._len8le(b"") + + self._len8le(b"") + + self._len8le(m_encoded) + + self._len8le(n_encoded) + + self._len8le(l_encoded) + + self._len8le(r_encoded) + + self._len8le(z_encoded) + + self._len8le(v_encoded) + + self._len8le(w_encoded) + ) + + transcript_hash = self._hash(self._hkdf_hash, transcript) + digest_len = 64 if self._hkdf_hash.upper() == "SHA512" else 32 + mac_len = 16 if self._suite_mac_is_cmac(suite_type) else 32 + confirmation_keys = self._hkdf_expand( + "ConfirmationKeys", transcript_hash, mac_len * 2, self._hkdf_hash + ) + key_confirm_a = confirmation_keys[:mac_len] + key_confirm_b = confirmation_keys[mac_len : mac_len * 2] + self._shared_key = self._hkdf_expand( + "SharedKey", transcript_hash, digest_len, self._hkdf_hash + ) + + if self._suite_mac_is_cmac(suite_type): + user_confirm = self._cmac_aes(key_confirm_a, r_encoded) + expected_dev_confirm = self._cmac_aes(key_confirm_b, l_encoded) + else: + user_confirm = self._hmac(self._hkdf_hash, key_confirm_a, r_encoded) + expected_dev_confirm = self._hmac(self._hkdf_hash, key_confirm_b, l_encoded) + + self._expected_dev_confirm = self._base64(expected_dev_confirm) + return { + "sub_method": "pake_share", + "user_share": self._base64(l_encoded), + "user_confirm": self._base64(user_confirm), + } + + def _verify_dac_proof(self, share_result: dict[str, Any]) -> None: + """Verify DAC certificate chain and proof signature.""" + try: + dac_ca = str(share_result.get("dac_ca") or "") + dac_ica = str(share_result.get("dac_ica") or "") + dac_proof = share_result.get("dac_proof") + if not ( + dac_ca and dac_proof and self._shared_key and self._dac_nonce_base64 + ): + return + if not isinstance(dac_proof, str): + raise KasaException("Invalid DAC proof type") + + ca_cert = self._transport._load_certificate_value(dac_ca) + ica_cert = ( + self._transport._load_certificate_value(dac_ica) if dac_ica else None + ) + self._transport._verify_dac_certificate_chain(ca_cert, ica_cert) + message = self._shared_key + self._unbase64(self._dac_nonce_base64) + signature = self._unbase64(dac_proof) + public_key = ca_cert.public_key() + if not isinstance(public_key, ec.EllipticCurvePublicKey): + raise KasaException( + "Unsupported DAC proof public key type: " + f"{type(public_key).__name__}" + ) + public_key.verify(signature, message, ec.ECDSA(hashes.SHA256())) + except InvalidSignature as exc: + _LOGGER.error("TPAP: invalid DAC proof signature") + raise KasaException("Invalid DAC proof signature") from exc + except Exception as exc: + _LOGGER.error("TPAP: DAC verification failed: %s", exc) + raise KasaException(f"DAC verification failed: {exc}") from exc + + def _establish_session_from_share_result( + self, share_result: dict[str, Any] + ) -> None: + dev_confirm = str(share_result.get("dev_confirm") or "").lower() + if not dev_confirm: + raise KasaException("TPAP share response missing dev_confirm") + if dev_confirm != (self._expected_dev_confirm or "").lower(): + raise KasaException("TPAP confirmation mismatch") + + if self._use_dac_certification(): + self._verify_dac_proof(share_result) + + session_id = str( + share_result.get("sessionId") or share_result.get("stok") or "" + ) + if not session_id: + _LOGGER.error("TPAP: missing session ID from device") + raise KasaException("Missing session fields from device") + if self._shared_key is None: + raise KasaException("TPAP shared key was not derived") + start_seq = share_result.get("start_seq") + if start_seq is None: + raise KasaException("Missing session fields from device") + try: + sequence = int(start_seq) + except (TypeError, ValueError) as exc: + raise KasaException("Invalid session fields from device") from exc + + self._key, self._base_nonce = self.key_nonce_from_shared( + self._shared_key, self._cipher_id, hkdf_hash=self._hkdf_hash + ) + self._session_id = session_id + self._sequence = sequence + self._ds_url = URL(f"{self._transport._app_url}/stok={self._session_id}/ds") + + @classmethod + def _normalize_cipher_id(cls, cipher_id: str) -> str: + return cipher_id.lower().replace("-", "_") + + @classmethod + def _cipher_parameters( + cls, cipher_id: str + ) -> tuple[bytes, bytes, bytes, bytes, int]: + normalized = cls._normalize_cipher_id(cipher_id) + try: + return cls.CIPHER_PARAMETERS[normalized] + except KeyError as exc: + raise KasaException( + f"Unsupported TPAP session cipher: {cipher_id}" + ) from exc + + @staticmethod + def _hkdf( + master: bytes, *, salt: bytes, info: bytes, length: int, algo: str = "SHA256" + ) -> bytes: + algorithm = hashes.SHA256() if algo.upper() == "SHA256" else hashes.SHA512() + return HKDF(algorithm=algorithm, length=length, salt=salt, info=info).derive( + master + ) + + @staticmethod + def _nonce_from_base(base_nonce: bytes, seq: int) -> bytes: + if len(base_nonce) < 4: + raise ValueError("base nonce too short") + return base_nonce[:-4] + struct.pack(">I", seq) + + @classmethod + def key_nonce_from_shared( + cls, shared_key: bytes, cipher_id: str, hkdf_hash: str = "SHA256" + ) -> tuple[bytes, bytes]: + """Derive the session key and base nonce.""" + key_salt, key_info, nonce_salt, nonce_info, key_len = cls._cipher_parameters( + cipher_id + ) + return ( + cls._hkdf( + shared_key, + salt=key_salt, + info=key_info, + length=key_len, + algo=hkdf_hash, + ), + cls._hkdf( + shared_key, + salt=nonce_salt, + info=nonce_info, + length=cls.NONCE_LEN, + algo=hkdf_hash, + ), + ) + + @classmethod + def _encrypt_payload( + cls, cipher_id: str, key: bytes, base_nonce: bytes, plaintext: bytes, seq: int + ) -> bytes: + nonce = cls._nonce_from_base(base_nonce, seq) + normalized = cls._normalize_cipher_id(cipher_id) + if normalized.startswith("aes_"): + return AESCCM(key, tag_length=cls.TAG_LEN).encrypt(nonce, plaintext, None) + return ChaCha20Poly1305(key).encrypt(nonce, plaintext, None) + + @classmethod + def _decrypt_payload( + cls, + cipher_id: str, + key: bytes, + base_nonce: bytes, + ciphertext_and_tag: bytes, + seq: int, + ) -> bytes: + nonce = cls._nonce_from_base(base_nonce, seq) + normalized = cls._normalize_cipher_id(cipher_id) + if normalized.startswith("aes_"): + return AESCCM(key, tag_length=cls.TAG_LEN).decrypt( + nonce, ciphertext_and_tag, None + ) + return ChaCha20Poly1305(key).decrypt(nonce, ciphertext_and_tag, None) + + @classmethod + def sec_encrypt( + cls, + cipher_id: str, + key: bytes, + base_nonce: bytes, + plaintext: bytes, + seq: int = 1, + ) -> tuple[bytes, bytes]: + """Encrypt the message.""" + combined = cls._encrypt_payload(cipher_id, key, base_nonce, plaintext, seq) + return combined[: -cls.TAG_LEN], combined[-cls.TAG_LEN :] + + @classmethod + def sec_decrypt( + cls, + cipher_id: str, + key: bytes, + base_nonce: bytes, + ciphertext: bytes, + tag: bytes, + seq: int = 1, + ) -> bytes: + """Decrypt the message.""" + return cls._decrypt_payload(cipher_id, key, base_nonce, ciphertext + tag, seq) + + def _require_established_session(self) -> tuple[str, int, URL, bytes, bytes]: + if not self.is_established: + raise KasaException("TPAP transport is not established") + if TYPE_CHECKING: + assert self._sequence is not None + assert self._ds_url is not None + assert self._key is not None + assert self._base_nonce is not None + + return ( + self._cipher_id, + self._sequence, + self._ds_url, + self._key, + self._base_nonce, + ) + + def encrypt(self, payload: bytes | str) -> tuple[bytes, int]: + """Encrypt the message.""" + cipher_id, seq, _, key, base_nonce = self._require_established_session() + plaintext = payload.encode() if isinstance(payload, str) else payload + encrypted = self._encrypt_payload(cipher_id, key, base_nonce, plaintext, seq) + self._sequence = seq + 1 + return struct.pack(">I", seq) + encrypted, seq + + def advance(self, seq: int) -> None: + """Advance the request sequence.""" + if self._sequence == seq: + self._sequence = seq + 1 + + def decrypt(self, payload: bytes, request_seq: int) -> bytes: + """Decrypt the message.""" + cipher_id, _, _, key, base_nonce = self._require_established_session() + if len(payload) < 4 + self.TAG_LEN: + raise KasaException("TPAP response too short") + + response_seq = struct.unpack(">I", payload[:4])[0] + if response_seq != request_seq: + _LOGGER.debug( + "Device returned unexpected rseq %d (expected %d)", + response_seq, + request_seq, + ) + return self._decrypt_payload( + cipher_id, key, base_nonce, payload[4:], response_seq + ) + + +class TpapTransport(BaseTransport): + """Implementation of the TPAP encryption protocol.""" + + DEFAULT_PORT: int = 80 + DEFAULT_HTTPS_PORT: int = 4433 + CAMERA_AUTH_DEVICE_FAMILIES = ( + DeviceFamily.SmartIpCamera, + DeviceFamily.SmartTapoDoorbell, + ) + ROBOT_TPAP_DEVICE_FAMILIES = {DeviceFamily.SmartTapoRobovac} + CIPHERS = ":".join( + [ + "ECDHE-ECDSA-AES256-GCM-SHA384", + "ECDHE-ECDSA-AES256-SHA384", + "ECDHE-ECDSA-AES256-SHA", + "ECDHE-ECDSA-AES128-GCM-SHA256", + "ECDHE-ECDSA-AES128-SHA256", + "ECDHE-ECDSA-AES128-SHA", + "ECDHE-RSA-AES256-GCM-SHA384", + "ECDHE-RSA-AES256-SHA384", + "ECDHE-RSA-AES256-SHA", + "ECDHE-RSA-AES128-GCM-SHA256", + "ECDHE-RSA-AES128-SHA256", + "ECDHE-RSA-AES128-SHA", + ] + ) + COMMON_HEADERS = {"Content-Type": "application/json"} + TPAP_ROOT_CA_PEM = """ +-----BEGIN CERTIFICATE----- +MIICNzCCAdygAwIBAgIUNLD7w5j5WU/efCe8bqkfGSRGgLYwCgYIKoZIzj0EAwIw +ezEnMCUGA1UEAwweVFAtTElOSyBTWVNURU1TIERFVklDRSBST09UIENBMR0wGwYD +VQQKDBRUUC1MSU5LIFNZU1RFTVMgSU5DLjEPMA0GA1UEBwwGSXJ2aW5lMRMwEQYD +VQQIDApDYWxpZm9ybmlhMQswCQYDVQQGEwJVUzAgFw0yNDExMjIwMjU3NDhaGA8y +MDU0MTExNTAyNTc0OFowezEnMCUGA1UEAwweVFAtTElOSyBTWVNURU1TIERFVklD +RSBST09UIENBMR0wGwYDVQQKDBRUUC1MSU5LIFNZU1RFTVMgSU5DLjEPMA0GA1UE +BwwGSXJ2aW5lMRMwEQYDVQQIDApDYWxpZm9ybmlhMQswCQYDVQQGEwJVUzBZMBMG +ByqGSM49AgEGCCqGSM49AwEHA0IABLwo8H9H6BoJDvcoewi4wPrPryVXir4z4yXV +n29R5XCAcFfKk06pYPupG6pjaKOLKWXnaOdPZThDFxwGLo3urV2jPDA6MAsGA1Ud +DwQEAwIBhjAMBgNVHRMEBTADAQH/MB0GA1UdDgQWBBRivfUtiHYsZBOKo80uZEwk +XhBkdDAKBggqhkjOPQQDAgNJADBGAiEA+7j5jemtXcGYN0unH+9rjVhVAL7WrsOi +5rbc0IIvD6MCIQCZuGGssu4Ygt2V8Vr0QF2fO9wxfNB3aRRMYQ+6lMrLGA== +-----END CERTIFICATE----- +""".strip() + + def __init__(self, *, config: DeviceConfig) -> None: + """Create the transport.""" + super().__init__(config=config) + if self._credentials is None and self._credentials_hash: + try: + decoded_hash = json_loads( + base64.b64decode(self._credentials_hash.encode()).decode() + ) + username = decoded_hash["un"] + password = decoded_hash["pwd"] + self._credentials = Credentials(username, password) + self._config.credentials = self._credentials + except Exception as ex: + _LOGGER.debug("Unable to decode stored TPAP credentials_hash: %s", ex) + self._http_client: HttpClient = HttpClient(self._config) + self._ssl_context: ssl.SSLContext | bool | None = None + protocol = "https" if config.connection_type.https else "http" + self._bootstrap_url = URL(f"{protocol}://{self._host}:{self._port}") + self._app_url = self._bootstrap_url + self._known_device_mac = "" + self._known_tpap_tls: int | None = None + self._known_tpap_port: int | None = None + self._known_tpap_dac = False + self._known_tpap_pake: list[int] = [] + self._known_tpap_user_hash_type: int | None = None + self._send_lock: asyncio.Lock = asyncio.Lock() + self._encryption_session = TpapEncryptionSession(self) + + @property + def default_port(self) -> int: + """Default port for the transport.""" + config = self._config + if port := config.connection_type.http_port: + return port + if config.connection_type.https: + return self.DEFAULT_HTTPS_PORT + return self.DEFAULT_PORT + + @property + def credentials_hash(self) -> str | None: + """The hashed credentials used by the transport.""" + if self._credentials and self._credentials.username: + credentials_hash = { + "un": self._credentials.username, + "pwd": self._credentials.password, + } + return base64.b64encode(json_dumps(credentials_hash).encode()).decode() + return self._config.credentials_hash + + def _build_app_url(self, *, tls_mode: int | None, port: int | None) -> URL: + scheme = "https" if tls_mode in (1, 2) else "http" + if port and port > 0: + resolved_port = port + elif scheme == "https": + resolved_port = self.DEFAULT_HTTPS_PORT + else: + resolved_port = self._port + return URL.build( + scheme=scheme, + host=self._host, + port=resolved_port, + ) + + def _get_initial_app_url(self) -> URL: + if not (self._known_tpap_port and self._known_tpap_port > 0) and ( + self._known_tpap_tls not in (1, 2) + ): + return self._bootstrap_url + + return self._build_app_url( + tls_mode=self._known_tpap_tls, + port=self._known_tpap_port, + ) + + @classmethod + def _load_root_ca_certificate(cls) -> x509.Certificate: + return x509.load_pem_x509_certificate(cls.TPAP_ROOT_CA_PEM.encode()) + + @classmethod + def _load_certificate_value(cls, certificate_value: str) -> x509.Certificate: + raw_value = certificate_value.strip() + if not raw_value: + raise KasaException("Empty certificate value") + + candidates: list[bytes] = [raw_value.encode()] + decoded_candidate: bytes | None = None + try: + decoded_candidate = base64.b64decode(raw_value, validate=True) + except Exception: + decoded_candidate = None + if decoded_candidate is not None: + candidates.insert(0, decoded_candidate) + + last_error: Exception | None = None + for candidate in candidates: + try: + if b"-----BEGIN CERTIFICATE-----" in candidate: + return x509.load_pem_x509_certificate(candidate) + return x509.load_der_x509_certificate(candidate) + except Exception as exc: + last_error = exc + + raise KasaException("Invalid certificate value") from last_error + + @staticmethod + def _verify_certificate_validity(certificate: x509.Certificate) -> None: + now = datetime.now(UTC) + if hasattr(certificate, "not_valid_before_utc"): + not_before = certificate.not_valid_before_utc + not_after = certificate.not_valid_after_utc + else: + not_before = certificate.not_valid_before.replace(tzinfo=UTC) + not_after = certificate.not_valid_after.replace(tzinfo=UTC) + if now < not_before or now > not_after: + raise KasaException("Certificate is outside its validity period") + + @staticmethod + def _verify_certificate_signature( + certificate: x509.Certificate, issuer: x509.Certificate + ) -> None: + public_key = issuer.public_key() + signature_hash = certificate.signature_hash_algorithm + if signature_hash is None: + raise KasaException("Certificate signature hash algorithm is unavailable") + if isinstance(public_key, ec.EllipticCurvePublicKey): + public_key.verify( + certificate.signature, + certificate.tbs_certificate_bytes, + ec.ECDSA(signature_hash), + ) + return + if isinstance(public_key, rsa.RSAPublicKey): + public_key.verify( + certificate.signature, + certificate.tbs_certificate_bytes, + padding.PKCS1v15(), + signature_hash, + ) + return + raise KasaException( + f"Unsupported DAC issuer public key type: {type(public_key).__name__}" + ) + + @classmethod + def _verify_dac_certificate_chain( + cls, + dac_ca_certificate: x509.Certificate, + dac_ica_certificate: x509.Certificate | None, + ) -> None: + try: + root_certificate = cls._load_root_ca_certificate() + cls._verify_certificate_validity(dac_ca_certificate) + if dac_ica_certificate is not None: + cls._verify_certificate_validity(dac_ica_certificate) + cls._verify_certificate_signature( + dac_ca_certificate, dac_ica_certificate + ) + cls._verify_certificate_signature(dac_ica_certificate, root_certificate) + else: + cls._verify_certificate_signature(dac_ca_certificate, root_certificate) + except Exception as exc: + raise KasaException( + f"DAC certificate chain verification failed: {exc}" + ) from exc + + @staticmethod + def _load_json_dict(payload: bytes) -> dict[str, Any]: + response_data = json_loads(payload.decode()) + if not isinstance(response_data, dict): + raise KasaException("Unexpected TPAP JSON response body type") + return response_data + + @staticmethod + def _should_retry_live_session(exc: Exception) -> bool: + if isinstance(exc, _ConnectionError): + return "Connection reset" in str(exc) + + if not isinstance(exc, _RetryableError): + return False + + return exc.error_code in { + SmartErrorCode.SESSION_TIMEOUT_ERROR, + SmartErrorCode.SESSION_EXPIRED, + SmartErrorCode.INVALID_NONCE, + SmartErrorCode.TRANSPORT_NOT_AVAILABLE_ERROR, + SmartErrorCode.STAT_ACCESS_ERROR, + } + + async def get_ssl_context(self) -> ssl.SSLContext | bool: + """Get or create the SSL context.""" + if self._ssl_context is None: + loop = asyncio.get_running_loop() + self._ssl_context = await loop.run_in_executor( + None, self._create_ssl_context + ) + return self._ssl_context + + def _create_ssl_context(self) -> ssl.SSLContext | bool: + tls_mode = self._encryption_session.tls_mode + if tls_mode == 0: + return False + + context = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT) + context.set_ciphers(self.CIPHERS) + context.check_hostname = False + + if tls_mode in (None, 1): + context.verify_mode = ssl.CERT_NONE + return context + + context.verify_mode = ssl.CERT_REQUIRED + context.load_verify_locations(cadata=self.TPAP_ROOT_CA_PEM) + return context + + async def send(self, request: str) -> dict[str, Any]: + """Send the request.""" + try: + return await self._send_once(request) + except Exception as exc: + if not self._should_retry_live_session(exc): + raise + + _LOGGER.debug( + "TPAP: resetting live session and retrying after error: %s", + exc, + ) + await self.reset() + return await self._send_once(request) + + async def _send_once(self, request: str) -> dict[str, Any]: + """Send a single request.""" + if not self._encryption_session.is_established: + await self._encryption_session.perform_handshake() + + ds_url = self._encryption_session.ds_url + if ds_url is None: + raise KasaException("TPAP transport is not established") + + async with self._send_lock: + payload, seq = self._encryption_session.encrypt(request) + headers = {"Content-Type": "application/octet-stream"} + ssl_context = await self.get_ssl_context() + status, data = await self._http_client.post( + ds_url, + data=payload, + headers=headers, + ssl=ssl_context, + ) + if status != 200: + raise KasaException( + f"TPAP secure request failed for {self._host}: status {status}" + ) + + if isinstance(data, bytes | bytearray): + plaintext = self._encryption_session.decrypt(bytes(data), seq) + return self._load_json_dict(plaintext) + + if isinstance(data, dict): + self._encryption_session._handle_response_error_code(data, "request") + return data + + raise KasaException("Unexpected TPAP response body type") + + async def close(self) -> None: + """Close the http client and reset internal state.""" + await self.reset() + await self._http_client.close() + + async def reset(self) -> None: + """Reset internal transport state.""" + self._encryption_session.reset() diff --git a/pyproject.toml b/pyproject.toml index 2866b9b2e..7af410971 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -11,7 +11,9 @@ dependencies = [ "cryptography>=1.9", "aiohttp>=3", "tzdata>=2024.2 ; platform_system == 'Windows'", - "mashumaro>=3.20" + "mashumaro>=3.20", + "ecdsa>=0.19.1", + "passlib>=1.7.4", ] classifiers = [ @@ -195,3 +197,11 @@ module = [ "devtools.create_module_fixtures" ] disable_error_code = "import-not-found,import-untyped" + +[[tool.mypy.overrides]] +module = ["ecdsa", "ecdsa.*"] +ignore_missing_imports = true + +[[tool.mypy.overrides]] +module = ["passlib", "passlib.*"] +ignore_missing_imports = true diff --git a/tests/test_device_factory.py b/tests/test_device_factory.py index 19ccfb73d..52c070c91 100644 --- a/tests/test_device_factory.py +++ b/tests/test_device_factory.py @@ -45,6 +45,7 @@ LinkieTransportV2, SslAesTransport, SslTransport, + TpapTransport, XorTransport, ) @@ -241,6 +242,18 @@ async def test_device_class_from_unknown_family(caplog): SslAesTransport, id="smartcam-doorbell", ), + pytest.param( + CP(DF.SmartIpCamera, ET.Tpap, https=True), + SmartCamProtocol, + TpapTransport, + id="smartcam-tpap", + ), + pytest.param( + CP(DF.SmartTapoHub, ET.Tpap, https=True), + SmartCamProtocol, + TpapTransport, + id="smartcam-hub-tpap", + ), pytest.param( CP(DF.IotIpCamera, ET.Aes, https=True), IotProtocol, @@ -283,6 +296,12 @@ async def test_device_class_from_unknown_family(caplog): KlapTransportV2, id="smart-chime", ), + pytest.param( + CP(DF.SmartTapoPlug, ET.Tpap, https=False), + SmartProtocol, + TpapTransport, + id="smart-tpap", + ), ], ) async def test_get_protocol( diff --git a/tests/transports/test_tpaptransport.py b/tests/transports/test_tpaptransport.py new file mode 100644 index 000000000..dc3ec5529 --- /dev/null +++ b/tests/transports/test_tpaptransport.py @@ -0,0 +1,2100 @@ +from __future__ import annotations + +import base64 +import hashlib +import json +import logging +import ssl +import struct +from datetime import UTC, datetime, timedelta +from types import SimpleNamespace +from typing import Any, cast + +import pytest +from cryptography import x509 +from cryptography.hazmat.primitives import hashes, serialization +from cryptography.hazmat.primitives.asymmetric import ec, rsa +from cryptography.x509.oid import NameOID +from yarl import URL + +import kasa.transports.tpaptransport as tp +from kasa.credentials import Credentials +from kasa.deviceconfig import DeviceConfig, DeviceFamily +from kasa.exceptions import ( + AuthenticationError, + DeviceError, + KasaException, + SmartErrorCode, + _ConnectionError, + _RetryableError, +) + + +def _discover_response( + *, + mac: str = "AA:BB:CC:DD:EE:FF", + tls: int = 2, + port: int = 4567, + pake: list[int] | None = None, + user_hash_type: int | None = None, + dac: bool = False, +) -> dict[str, Any]: + tpap_info: dict[str, Any] = { + "dac": dac, + "tls": tls, + "port": port, + "pake": pake or [], + } + if user_hash_type is not None: + tpap_info["user_hash_type"] = user_hash_type + + return {"error_code": 0, "result": {"mac": mac, "tpap": tpap_info}} + + +def _register_result( + *, + extra_crypt: dict[str, Any] | None = None, + cipher_suites: int = 2, + iterations: int = 100, + encryption: str = "aes_128_ccm", +) -> dict[str, Any]: + return { + "dev_random": base64.b64encode(b"\x00" * 16).decode(), + "dev_salt": base64.b64encode(b"\x11" * 16).decode(), + "dev_share": base64.b64encode(_p256_pub_uncompressed()).decode(), + "cipher_suites": cipher_suites, + "iterations": iterations, + "encryption": encryption, + "extra_crypt": extra_crypt or {}, + } + + +def _share_result( + session: tp.TpapEncryptionSession, + *, + session_id: str = "STOK", + start_seq: int = 7, +) -> dict[str, Any]: + assert session._expected_dev_confirm is not None + return { + "dev_confirm": session._expected_dev_confirm.lower(), + "sessionId": session_id, + "start_seq": start_seq, + } + + +def _make_discover_post( + response: dict[str, Any], +) -> Any: + async def post( + url: URL, + *, + json: dict[str, Any] | None = None, + data: bytes | None = None, + headers: dict[str, str] | None = None, + ssl: ssl.SSLContext | bool | None = None, + ) -> tuple[int, dict[str, Any]]: + del url, data, headers, ssl + assert json is not None + assert json["params"]["sub_method"] == "discover" + return 200, response + + return post + + +def _make_handshake_login( + session: tp.TpapEncryptionSession, + *, + capture_register: dict[str, Any] | None = None, + register_result: dict[str, Any] | None = None, + session_id: str = "STOK", + start_seq: int = 7, +) -> Any: + async def fake_login(params: dict[str, Any], *, step_name: str) -> dict[str, Any]: + if step_name == "pake_register": + if capture_register is not None: + capture_register.update(params) + return register_result or _register_result() + + assert step_name == "pake_share" + return _share_result(session, session_id=session_id, start_seq=start_seq) + + return fake_login + + +def _p256_pub_uncompressed() -> bytes: + private_key = ec.derive_private_key( + int.from_bytes(b"\x01" * 32, "big"), + ec.SECP256R1(), + ) + return private_key.public_key().public_bytes( + serialization.Encoding.X962, + serialization.PublicFormat.UncompressedPoint, + ) + + +def _establish_session( + transport: tp.TpapTransport, + session: tp.TpapEncryptionSession, + *, + session_id: str = "SID", + start_seq: int = 10, +) -> None: + key, base_nonce = tp.TpapEncryptionSession.key_nonce_from_shared( + b"shared-secret", "aes_128_ccm" + ) + session._cipher_id = "aes_128_ccm" + session._key = key + session._base_nonce = base_nonce + session._session_id = session_id + session._sequence = start_seq + session._ds_url = URL(f"{transport._app_url}/stok={session_id}/ds") + + +def _make_established_transport() -> tuple[tp.TpapTransport, tp.TpapEncryptionSession]: + transport = _make_tpap_transport("host") + session = transport._encryption_session + _establish_session(transport, session) + return transport, session + + +def _make_tpap_transport( + host: str = "tpap-host", + *, + family: DeviceFamily | None = None, + credentials: Credentials | None = None, +) -> tp.TpapTransport: + config = DeviceConfig(host) + if family is not None: + config.connection_type.device_family = family + if credentials is not None: + config.credentials = credentials + return tp.TpapTransport(config=config) + + +def _make_camera_tpap_transport(host: str = "cam-host") -> tp.TpapTransport: + return _make_tpap_transport( + host, + family=DeviceFamily.SmartIpCamera, + credentials=Credentials("user", "pass"), + ) + + +def _build_certificate( + private_key: ec.EllipticCurvePrivateKey | rsa.RSAPrivateKey, + subject_common_name: str, + issuer_common_name: str, + issuer_private_key: ec.EllipticCurvePrivateKey | rsa.RSAPrivateKey, + *, + is_ca: bool = False, + subject_alt_names: list[x509.GeneralName] | None = None, +) -> x509.Certificate: + now = datetime.now(UTC).replace(tzinfo=None) + builder = ( + x509.CertificateBuilder() + .subject_name( + x509.Name([x509.NameAttribute(NameOID.COMMON_NAME, subject_common_name)]) + ) + .issuer_name( + x509.Name([x509.NameAttribute(NameOID.COMMON_NAME, issuer_common_name)]) + ) + .public_key(private_key.public_key()) + .serial_number(x509.random_serial_number()) + .not_valid_before(now - timedelta(days=1)) + .not_valid_after(now + timedelta(days=30)) + .add_extension( + x509.BasicConstraints(ca=is_ca, path_length=None), + critical=True, + ) + ) + if subject_alt_names: + builder = builder.add_extension( + x509.SubjectAlternativeName(subject_alt_names), + critical=False, + ) + return builder.sign(issuer_private_key, hashes.SHA256()) + + +def test_session_cipher_helpers_roundtrip() -> None: + aes_key, aes_nonce = tp.TpapEncryptionSession.key_nonce_from_shared( + b"k" * 16, "aes_256_ccm" + ) + aes_ct, aes_tag = tp.TpapEncryptionSession.sec_encrypt( + "aes_256_ccm", aes_key, aes_nonce, b"hello", seq=3 + ) + assert ( + tp.TpapEncryptionSession.sec_decrypt( + "aes_256_ccm", + aes_key, + aes_nonce, + aes_ct, + aes_tag, + seq=3, + ) + == b"hello" + ) + + chacha_key, chacha_nonce = tp.TpapEncryptionSession.key_nonce_from_shared( + b"c" * 32, "chacha20_poly1305" + ) + chacha_ct, chacha_tag = tp.TpapEncryptionSession.sec_encrypt( + "chacha20_poly1305", + chacha_key, + chacha_nonce, + b"world", + seq=5, + ) + assert ( + tp.TpapEncryptionSession.sec_decrypt( + "chacha20_poly1305", + chacha_key, + chacha_nonce, + chacha_ct, + chacha_tag, + seq=5, + ) + == b"world" + ) + + +async def test_session_encrypt_and_decrypt_roundtrip() -> None: + transport, session = _make_established_transport() + + payload, seq = session.encrypt(b'{"ok": true}') + assert seq == 10 + assert session._sequence == 11 + assert session.decrypt(payload, seq) == b'{"ok": true}' + assert str(session.ds_url) == f"{transport._app_url}/stok=SID/ds" + + +@pytest.mark.parametrize( + ( + "family", + "discover_response", + "expected_username", + "expected_passcode_type", + "expected_tls_mode", + "expected_scheme", + "expected_port", + ), + [ + pytest.param( + None, + _discover_response(pake=[2], user_hash_type=0), + tp.TpapEncryptionSession._md5_hex("admin"), + "userpw", + 2, + "https", + 4567, + id="iot-md5", + ), + pytest.param( + None, + _discover_response(pake=[2], user_hash_type=1), + tp.TpapEncryptionSession._sha256_hex_upper("admin"), + "userpw", + 2, + "https", + 4567, + id="iot-sha256", + ), + pytest.param( + None, + _discover_response(tls=0, port=80), + tp.TpapEncryptionSession._md5_hex("admin"), + "default_userpw", + 0, + "http", + 80, + id="generic-http", + ), + pytest.param( + DeviceFamily.SmartTapoHub, + _discover_response(tls=0, port=80, pake=[2], user_hash_type=0, dac=True), + tp.TpapEncryptionSession._md5_hex("admin"), + "userpw", + 0, + "http", + 80, + id="tapo-hub-http-userpw", + ), + pytest.param( + DeviceFamily.SmartIpCamera, + _discover_response(pake=[2], user_hash_type=0), + tp.TpapEncryptionSession._md5_hex("admin"), + "userpw", + 2, + "https", + 4567, + id="camera-admin", + ), + ], +) +async def test_session_perform_handshake_registers_expected_auth_values( + monkeypatch: pytest.MonkeyPatch, + family: DeviceFamily | None, + discover_response: dict[str, Any], + expected_username: str, + expected_passcode_type: str, + expected_tls_mode: int, + expected_scheme: str, + expected_port: int, +) -> None: + transport = _make_tpap_transport( + "handshake-host", + family=family, + credentials=Credentials("user", "pass"), + ) + session = transport._encryption_session + captured_register: dict[str, Any] = {} + + transport._http_client.post = _make_discover_post( # type: ignore[assignment] + discover_response + ) + monkeypatch.setattr( + session, + "_login", + _make_handshake_login(session, capture_register=captured_register), + raising=True, + ) + + await session.perform_handshake() + + assert captured_register["username"] == expected_username + assert captured_register["encryption"] == ["aes_128_ccm"] + assert captured_register["passcode_type"] == expected_passcode_type + assert session.is_established is True + assert session.tls_mode == expected_tls_mode + assert transport._app_url.scheme == expected_scheme + assert transport._app_url.host == "handshake-host" + assert transport._app_url.port == expected_port + assert session.ds_url is not None + assert session.ds_url.scheme == expected_scheme + assert session.ds_url.host == "handshake-host" + assert session.ds_url.port == expected_port + assert session.ds_url.path == "/stok=STOK/ds" + + +async def test_session_camera_auth_uses_device_family() -> None: + camera_transport = _make_camera_tpap_transport() + + hub_transport = _make_tpap_transport("hub-host", family=DeviceFamily.SmartTapoHub) + robot_transport = _make_tpap_transport( + "robot-host", family=DeviceFamily.SmartTapoRobovac + ) + iot_transport = _make_tpap_transport() + + assert camera_transport._encryption_session._uses_camera_auth is True + assert hub_transport._encryption_session._uses_camera_auth is False + assert robot_transport._encryption_session._uses_camera_auth is False + assert robot_transport._encryption_session._uses_robot_tpap_auth is True + assert iot_transport._encryption_session._uses_camera_auth is False + + +async def test_smartcam_session_builds_password_candidates_without_lat() -> None: + transport = _make_camera_tpap_transport() + session = transport._encryption_session + session._tpap_pake = [2] + + assert session._get_candidate_secrets() == [ + tp.TpapEncryptionSession._md5_hex("pass"), + tp.TpapEncryptionSession._sha256_hex_upper("pass"), + ] + + +@pytest.mark.asyncio +async def test_smartcam_session_retries_next_candidate_on_handshake_failure( + monkeypatch: pytest.MonkeyPatch, +) -> None: + transport = _make_camera_tpap_transport() + session = transport._encryption_session + session._tpap_pake = [2] + attempted_candidates: list[str] = [] + + async def fake_login(params: dict[str, Any], *, step_name: str) -> dict[str, Any]: + del params + if step_name == "pake_register": + return {"extra_crypt": {}} + return {} + + def fake_iter_candidates() -> list[str]: + return ["first", "second"] + + def fake_build_share_params( + register_result: dict[str, Any], credentials_string: str + ) -> dict[str, Any]: + del register_result + attempted_candidates.append(credentials_string) + return {} + + def fake_establish_session(share_result: dict[str, Any]) -> None: + del share_result + if attempted_candidates[-1] == "first": + raise KasaException("bad candidate") + _establish_session(transport, session, session_id="CAM-SID", start_seq=4) + + monkeypatch.setattr(session, "_login", fake_login, raising=True) + monkeypatch.setattr( + session, "_get_candidate_secrets", fake_iter_candidates, raising=True + ) + monkeypatch.setattr( + session, + "_build_share_params_from_register", + fake_build_share_params, + raising=True, + ) + monkeypatch.setattr( + session, + "_establish_session_from_share_result", + fake_establish_session, + raising=True, + ) + + await session._perform_auth_handshake() + + assert attempted_candidates == ["first", "second"] + assert session._session_id == "CAM-SID" + + +@pytest.mark.asyncio +async def test_smartcam_session_raises_when_no_candidates( + monkeypatch: pytest.MonkeyPatch, +) -> None: + transport = _make_camera_tpap_transport() + session = transport._encryption_session + session._tpap_pake = [2] + + monkeypatch.setattr(session, "_get_candidate_secrets", lambda: [], raising=True) + + with pytest.raises(AuthenticationError, match="no credential candidates"): + await session._perform_auth_handshake() + + +@pytest.mark.asyncio +async def test_smartcam_session_raises_when_no_supported_passcode_type() -> None: + transport = _make_camera_tpap_transport() + session = transport._encryption_session + session._tpap_pake = [9] + + with pytest.raises(AuthenticationError, match="no supported passcode type"): + await session._perform_auth_handshake() + + +@pytest.mark.asyncio +async def test_smartcam_session_reraises_last_candidate_error( + caplog: pytest.LogCaptureFixture, + monkeypatch: pytest.MonkeyPatch, +) -> None: + transport = _make_camera_tpap_transport() + session = transport._encryption_session + session._tpap_pake = [2] + + async def fake_login(params: dict[str, Any], *, step_name: str) -> dict[str, Any]: + del params + if step_name == "pake_register": + return {"extra_crypt": {}} + return {} + + monkeypatch.setattr(session, "_login", fake_login, raising=True) + monkeypatch.setattr( + session, "_get_candidate_secrets", lambda: ["only"], raising=True + ) + monkeypatch.setattr( + session, + "_build_share_params_from_register", + lambda register_result, credentials_string: {}, + raising=True, + ) + monkeypatch.setattr( + session, + "_establish_session_from_share_result", + lambda share_result: (_ for _ in ()).throw(KasaException("last failure")), + raising=True, + ) + + with ( + caplog.at_level(logging.DEBUG), + pytest.raises(KasaException, match="last failure"), + ): + await session._perform_auth_handshake() + + assert "all password-based camera candidates failed" in caplog.text + + +@pytest.mark.asyncio +async def test_generic_tpap_session_reraises_last_candidate_error_without_hint( + caplog: pytest.LogCaptureFixture, + monkeypatch: pytest.MonkeyPatch, +) -> None: + config = DeviceConfig("tpap-host") + config.credentials = Credentials("user", "pass") + transport = tp.TpapTransport(config=config) + session = transport._encryption_session + + async def fake_login(params: dict[str, Any], *, step_name: str) -> dict[str, Any]: + del params + if step_name == "pake_register": + return {"extra_crypt": {}} + return {} + + monkeypatch.setattr(session, "_login", fake_login, raising=True) + monkeypatch.setattr( + session, "_get_candidate_secrets", lambda: ["only"], raising=True + ) + monkeypatch.setattr( + session, + "_build_share_params_from_register", + lambda register_result, credentials_string: {}, + raising=True, + ) + monkeypatch.setattr( + session, + "_establish_session_from_share_result", + lambda share_result: (_ for _ in ()).throw(KasaException("last failure")), + raising=True, + ) + + with ( + caplog.at_level(logging.DEBUG), + pytest.raises(KasaException, match="last failure"), + ): + await session._perform_auth_handshake() + + assert "all password-based camera candidates failed" not in caplog.text + + +@pytest.mark.asyncio +async def test_smartcam_session_propagates_retryable_error( + monkeypatch: pytest.MonkeyPatch, +) -> None: + transport = _make_camera_tpap_transport() + session = transport._encryption_session + session._tpap_pake = [2] + attempts: list[str] = [] + + async def fake_login(params: dict[str, Any], *, step_name: str) -> dict[str, Any]: + del params + if step_name == "pake_register": + return {"extra_crypt": {}} + raise _RetryableError("retry me", error_code=SmartErrorCode.SESSION_EXPIRED) + + def fake_build_share_params( + register_result: dict[str, Any], credentials_string: str + ) -> dict[str, Any]: + del register_result + attempts.append(credentials_string) + return {} + + monkeypatch.setattr(session, "_login", fake_login, raising=True) + monkeypatch.setattr( + session, + "_get_candidate_secrets", + lambda: ["first", "second"], + raising=True, + ) + monkeypatch.setattr( + session, + "_build_share_params_from_register", + fake_build_share_params, + raising=True, + ) + + with pytest.raises(_RetryableError, match="retry me"): + await session._perform_auth_handshake() + + assert attempts == ["first"] + + +@pytest.mark.asyncio +async def test_iot_session_adds_dac_nonce_when_required( + monkeypatch: pytest.MonkeyPatch, +) -> None: + transport = _make_tpap_transport( + family=DeviceFamily.SmartTapoPlug, + credentials=Credentials("user", "pass"), + ) + session = transport._encryption_session + session._tpap_pake = [5] + session._tpap_tls = 0 + session._tpap_dac = True + captured_share_params: dict[str, Any] | None = None + + async def fake_login(params: dict[str, Any], *, step_name: str) -> dict[str, Any]: + nonlocal captured_share_params + if step_name == "pake_register": + return {"extra_crypt": {}} + captured_share_params = params.copy() + return {} + + monkeypatch.setattr(session, "_login", fake_login, raising=True) + monkeypatch.setattr( + session, "_get_candidate_secrets", lambda: ["candidate"], raising=True + ) + monkeypatch.setattr( + session, + "_build_share_params_from_register", + lambda register_result, credentials_string: {}, + raising=True, + ) + monkeypatch.setattr( + session, + "_establish_session_from_share_result", + lambda share_result: _establish_session( + transport, session, session_id="DAC-SID", start_seq=2 + ), + raising=True, + ) + + await session._perform_auth_handshake() + + assert captured_share_params is not None + assert captured_share_params["dac_nonce"] + assert len(base64.b64decode(captured_share_params["dac_nonce"])) == 32 + assert session._session_id == "DAC-SID" + + +@pytest.mark.parametrize( + ("tls", "dac", "pake", "expected"), + [ + pytest.param(0, True, [2], True, id="plain-tpap-dac"), + pytest.param(0, False, [2], False, id="plain-tpap-no-dac"), + pytest.param(1, True, [2], False, id="tls-tpap-dac"), + ], +) +async def test_use_dac_certification_uses_discovered_dac_support( + tls: int, dac: bool, pake: list[int], expected: bool +) -> None: + transport = _make_tpap_transport("tpap-host") + session = transport._encryption_session + session._tpap_pake = pake + session._tpap_tls = tls + session._tpap_dac = dac + + assert session._use_dac_certification() is expected + + +@pytest.mark.parametrize( + ("family", "pake", "expected"), + [ + pytest.param( + DeviceFamily.SmartIpCamera, + [1], + "userpw", + id="camera-setup-code", + ), + pytest.param( + DeviceFamily.SmartIpCamera, + [0, 2], + "userpw", + id="camera-pake-two-before-zero", + ), + pytest.param(None, [0], "default_userpw", id="default-pake-zero"), + pytest.param(None, [1], "default_userpw", id="iot-pake-one"), + pytest.param(None, [5], "userpw", id="iot-pake-five"), + pytest.param(None, [2, 3], "userpw", id="iot-userpw-before-shared-token"), + pytest.param(DeviceFamily.SmartTapoHub, [2], "userpw", id="hub-userpw"), + pytest.param( + DeviceFamily.SmartTapoRobovac, + [2, 3], + "userpw", + id="robot-userpw-before-shared-token", + ), + pytest.param( + DeviceFamily.SmartTapoRobovac, + [5], + "default_userpw", + id="robot-pake-five-default", + ), + pytest.param(DeviceFamily.SmartTapoBulb, [5], "userpw", id="bulb-pake-five"), + pytest.param( + DeviceFamily.SmartTapoRobovac, + [9], + "default_userpw", + id="robot-unknown-pake", + ), + pytest.param( + DeviceFamily.SmartIpCamera, + [3], + "shared_token", + id="camera-shared-token", + ), + pytest.param(None, [], "default_userpw", id="default-missing-pake"), + pytest.param(DeviceFamily.SmartIpCamera, [], None, id="camera-missing-pake"), + ], +) +async def test_get_passcode_type( + family: DeviceFamily | None, pake: list[int], expected: str | None +) -> None: + transport = _make_tpap_transport("tpap-host", family=family) + session = transport._encryption_session + session._tpap_pake = pake + + assert session._get_passcode_type() == expected + + +@pytest.mark.parametrize( + ("pake", "expected"), + [ + pytest.param([1], ["pass"], id="setup-code-uses-raw-password"), + pytest.param( + [3], [tp.TpapEncryptionSession._md5_hex("pass")], id="shared-token" + ), + pytest.param([9], [], id="unknown-pake"), + pytest.param([], [], id="missing-pake"), + ], +) +async def test_camera_candidate_secrets(pake: list[int], expected: list[str]) -> None: + transport = _make_camera_tpap_transport() + session = transport._encryption_session + session._tpap_pake = pake + + assert session._get_candidate_secrets() == expected + + +async def test_smartcam_candidate_builder_dedupes_duplicates( + monkeypatch: pytest.MonkeyPatch, +) -> None: + transport = _make_camera_tpap_transport() + session = transport._encryption_session + session._tpap_pake = [2] + monkeypatch.setattr(session, "_md5_hex", lambda value: "same", raising=False) + monkeypatch.setattr( + session, "_sha256_hex_upper", lambda value: "same", raising=False + ) + + assert session._get_candidate_secrets() == ["same"] + + +async def test_smartcam_resolve_credentials_applies_extra_crypt() -> None: + transport = _make_camera_tpap_transport() + session = transport._encryption_session + session._tpap_pake = [2] + register_result = { + "extra_crypt": {"type": "password_shadow", "params": {"passwd_id": 2}} + } + + assert session._resolve_credentials(register_result, "candidate") == ( + tp.TpapEncryptionSession._sha1_hex("candidate") + ) + + +async def test_default_passcode_resolve_credentials_returns_candidate() -> None: + transport = _make_tpap_transport() + session = transport._encryption_session + session._tpap_pake = [0] + session._device_mac = "AA:BB:CC:DD:EE:FF" + + assert session._resolve_credentials({}, "candidate") == "candidate" + + +async def test_tls2_ssl_context_loads_root_ca( + monkeypatch: pytest.MonkeyPatch, +) -> None: + root_key = ec.generate_private_key(ec.SECP256R1()) + root_cert = _build_certificate(root_key, "root", "root", root_key, is_ca=True) + + monkeypatch.setattr( + tp.TpapTransport, + "TPAP_ROOT_CA_PEM", + root_cert.public_bytes(serialization.Encoding.PEM).decode(), + raising=True, + ) + + transport = tp.TpapTransport(config=DeviceConfig("tls-host")) + transport._encryption_session._tpap_tls = 2 + context = transport._create_ssl_context() + + assert isinstance(context, ssl.SSLContext) + assert context.verify_mode == ssl.CERT_REQUIRED + assert context.get_ca_certs() + + +async def test_dac_verification_checks_chain_and_signature( + monkeypatch: pytest.MonkeyPatch, +) -> None: + root_key = ec.generate_private_key(ec.SECP256R1()) + root_cert = _build_certificate(root_key, "root", "root", root_key, is_ca=True) + root_pem = root_cert.public_bytes(serialization.Encoding.PEM).decode() + + ica_key = ec.generate_private_key(ec.SECP256R1()) + ica_cert = _build_certificate(ica_key, "ica", "root", root_key, is_ca=True) + dac_key = ec.generate_private_key(ec.SECP256R1()) + dac_cert = _build_certificate(dac_key, "dac", "ica", ica_key) + + monkeypatch.setattr( + tp.TpapTransport, + "TPAP_ROOT_CA_PEM", + root_pem, + raising=True, + ) + + transport = tp.TpapTransport(config=DeviceConfig("dac-host")) + session = transport._encryption_session + session._shared_key = b"shared-key" + nonce = b"dac-nonce" + session._dac_nonce_base64 = base64.b64encode(nonce).decode() + + proof = dac_key.sign(session._shared_key + nonce, ec.ECDSA(hashes.SHA256())) + share_result = { + "dac_ca": base64.b64encode( + dac_cert.public_bytes(serialization.Encoding.PEM) + ).decode(), + "dac_ica": base64.b64encode( + ica_cert.public_bytes(serialization.Encoding.PEM) + ).decode(), + "dac_proof": base64.b64encode(proof).decode(), + } + + session._verify_dac_proof(share_result) + + share_result["dac_proof"] = base64.b64encode(b"bad-proof").decode() + with pytest.raises(KasaException, match="Invalid DAC proof signature"): + session._verify_dac_proof(share_result) + + +@pytest.mark.asyncio +async def test_transport_send_happy_path() -> None: + transport, session = _make_established_transport() + + async def post( + url: URL, + *, + json: dict[str, Any] | None = None, + data: bytes | None = None, + headers: dict[str, str] | None = None, + ssl: ssl.SSLContext | bool | None = None, + ) -> tuple[int, bytes]: + del url, json, headers, ssl + assert data is not None + return 200, data + + transport._http_client.post = post # type: ignore[assignment] + + out = await transport.send('{"result": {"ok": true}}') + + assert out["result"]["ok"] is True + assert session._sequence == 11 + + +@pytest.mark.parametrize( + ("first_failure", "session_id", "start_seq"), + [ + pytest.param("session-expired", "SID-RETRY", 20, id="session-expired"), + pytest.param("connection-reset", "SID-CONN", 30, id="connection-reset"), + ], +) +async def test_transport_send_retries_live_session_failures( + monkeypatch: pytest.MonkeyPatch, + first_failure: str, + session_id: str, + start_seq: int, +) -> None: + transport, session = _make_established_transport() + request_calls = 0 + + async def fake_handshake() -> None: + _establish_session( + transport, session, session_id=session_id, start_seq=start_seq + ) + + async def post( + url: URL, + *, + json: dict[str, Any] | None = None, + data: bytes | None = None, + headers: dict[str, str] | None = None, + ssl: ssl.SSLContext | bool | None = None, + ) -> tuple[int, dict[str, Any] | bytes]: + nonlocal request_calls + del url, json, headers, ssl + request_calls += 1 + if request_calls == 1: + if first_failure == "session-expired": + return 200, {"error_code": SmartErrorCode.SESSION_EXPIRED.value} + raise _ConnectionError("Connection reset by peer") + assert data is not None + return 200, data + + transport._http_client.post = post # type: ignore[assignment] + monkeypatch.setattr(session, "perform_handshake", fake_handshake, raising=True) + + out = await transport.send('{"result": {"retry": true}}') + + assert out["result"]["retry"] is True + assert request_calls == 2 + assert session._session_id == session_id + assert session._sequence == start_seq + 1 + + +@pytest.mark.asyncio +async def test_transport_reset_clears_session_state() -> None: + transport, session = _make_established_transport() + + await transport.reset() + + assert session.is_established is False + assert transport._app_url == transport._bootstrap_url + with pytest.raises(KasaException, match="TPAP transport is not established"): + session.encrypt(b"{}") + + +@pytest.mark.asyncio +async def test_transport_reset_preserves_discovered_transport_identity() -> None: + transport = tp.TpapTransport(config=DeviceConfig("tpap-host")) + session = transport._encryption_session + session._device_mac = "AA:BB:CC:DD:EE:FF" + session._tpap_tls = 2 + session._tpap_port = 4567 + session._tpap_dac = True + session._tpap_pake = [0, 2] + session._tpap_user_hash_type = 1 + transport._known_device_mac = session._device_mac + transport._known_tpap_tls = session._tpap_tls + transport._known_tpap_port = session._tpap_port + transport._known_tpap_dac = session._tpap_dac + transport._known_tpap_pake = list(session._tpap_pake) + transport._known_tpap_user_hash_type = session._tpap_user_hash_type + + await transport.reset() + + assert transport._app_url == URL("https://tpap-host:4567") + assert session.device_mac == "AA:BB:CC:DD:EE:FF" + assert session.tls_mode == 2 + assert session._tpap_port == 4567 + assert session._tpap_dac is True + assert session._tpap_pake == [0, 2] + assert session._tpap_user_hash_type == 1 + + +# -------------------------- +# Discovery and Login +# -------------------------- + + +@pytest.mark.asyncio +async def test_perform_handshake_is_noop_when_session_already_established() -> None: + transport, session = _make_established_transport() + + await session.perform_handshake() + + assert session.is_established is True + + +@pytest.mark.asyncio +async def test_perform_handshake_restarts_when_session_was_invalidated( + monkeypatch: pytest.MonkeyPatch, +) -> None: + transport, session = _make_established_transport() + session._invalidate_session() + discover_called = False + + async def fake_discover() -> None: + nonlocal discover_called + discover_called = True + + async def fake_auth_handshake() -> None: + _establish_session(transport, session, session_id="SID-NEW", start_seq=14) + + monkeypatch.setattr(session, "_discover", fake_discover, raising=True) + monkeypatch.setattr( + session, "_perform_auth_handshake", fake_auth_handshake, raising=True + ) + + await session.perform_handshake() + + assert discover_called is True + assert session.is_established is True + assert session._session_id == "SID-NEW" + + +@pytest.mark.asyncio +async def test_discover_raises_on_bad_status_or_body() -> None: + transport = tp.TpapTransport(config=DeviceConfig("discover-host")) + session = transport._encryption_session + + async def post( + url: URL, + *, + json: dict[str, Any] | None = None, + data: bytes | None = None, + headers: dict[str, str] | None = None, + ssl: ssl.SSLContext | bool | None = None, + ) -> tuple[int, bytes]: + del url, json, data, headers, ssl + return 500, b"bad" + + transport._http_client.post = post # type: ignore[assignment] + + with pytest.raises(KasaException, match="TPAP discover failed"): + await session._discover() + + +@pytest.mark.asyncio +async def test_discover_parses_invalid_numeric_fields_as_none() -> None: + transport = tp.TpapTransport(config=DeviceConfig("discover-host")) + session = transport._encryption_session + + async def post( + url: URL, + *, + json: dict[str, Any] | None = None, + data: bytes | None = None, + headers: dict[str, str] | None = None, + ssl: ssl.SSLContext | bool | None = None, + ) -> tuple[int, dict[str, Any]]: + del url, json, data, headers, ssl + return ( + 200, + { + "error_code": 0, + "result": { + "mac": "AA:BB:CC:DD:EE:FF", + "tpap": { + "tls": "bad", + "port": "bad", + "dac": True, + "pake": [2], + "user_hash_type": "bad", + }, + }, + }, + ) + + transport._http_client.post = post # type: ignore[assignment] + + await session._discover() + + assert session.device_mac == "AA:BB:CC:DD:EE:FF" + assert session.tls_mode is None + assert session._tpap_port is None + assert session._tpap_user_hash_type is None + assert str(transport._app_url) == str(transport._bootstrap_url) + + +@pytest.mark.asyncio +async def test_discover_propagates_device_error_codes() -> None: + transport = tp.TpapTransport(config=DeviceConfig("discover-host")) + session = transport._encryption_session + + async def post( + url: URL, + *, + json: dict[str, Any] | None = None, + data: bytes | None = None, + headers: dict[str, str] | None = None, + ssl: ssl.SSLContext | bool | None = None, + ) -> tuple[int, dict[str, Any]]: + del url, json, data, headers, ssl + return ( + 200, + {"error_code": SmartErrorCode.SESSION_EXPIRED.value}, + ) + + transport._http_client.post = post # type: ignore[assignment] + + with pytest.raises(_RetryableError): + await session._discover() + + +@pytest.mark.asyncio +async def test_discover_requires_result_and_tpap_objects() -> None: + transport = tp.TpapTransport(config=DeviceConfig("discover-host")) + session = transport._encryption_session + + async def post_without_result( + url: URL, + *, + json: dict[str, Any] | None = None, + data: bytes | None = None, + headers: dict[str, str] | None = None, + ssl: ssl.SSLContext | bool | None = None, + ) -> tuple[int, dict[str, Any]]: + del url, json, data, headers, ssl + return 200, {"error_code": 0} + + transport._http_client.post = post_without_result # type: ignore[assignment] + + with pytest.raises(KasaException, match="missing result object"): + await session._discover() + + async def post_without_tpap( + url: URL, + *, + json: dict[str, Any] | None = None, + data: bytes | None = None, + headers: dict[str, str] | None = None, + ssl: ssl.SSLContext | bool | None = None, + ) -> tuple[int, dict[str, Any]]: + del url, json, data, headers, ssl + return 200, {"error_code": 0, "result": {}} + + transport._http_client.post = post_without_tpap # type: ignore[assignment] + + with pytest.raises(KasaException, match="missing tpap object"): + await session._discover() + + +@pytest.mark.asyncio +async def test_login_raises_on_bad_status_or_body() -> None: + transport = tp.TpapTransport(config=DeviceConfig("login-host")) + session = transport._encryption_session + + async def post( + url: URL, + *, + json: dict[str, Any] | None = None, + data: bytes | None = None, + headers: dict[str, str] | None = None, + ssl: ssl.SSLContext | bool | None = None, + ) -> tuple[int, bytes]: + del url, json, data, headers, ssl + return 500, b"bad" + + transport._http_client.post = post # type: ignore[assignment] + + with pytest.raises(KasaException, match="TPAP pake_register failed"): + await session._login({}, step_name="pake_register") + + +@pytest.mark.asyncio +async def test_login_propagates_error_code_handling() -> None: + transport = tp.TpapTransport(config=DeviceConfig("login-host")) + session = transport._encryption_session + + async def post( + url: URL, + *, + json: dict[str, Any] | None = None, + data: bytes | None = None, + headers: dict[str, str] | None = None, + ssl: ssl.SSLContext | bool | None = None, + ) -> tuple[int, dict[str, Any]]: + del url, json, data, headers, ssl + return ( + 200, + {"error_code": SmartErrorCode.LOGIN_ERROR.value}, + ) + + transport._http_client.post = post # type: ignore[assignment] + + with pytest.raises(AuthenticationError, match="TPAP pake_register failed"): + await session._login({}, step_name="pake_register") + + async def post_without_result( + url: URL, + *, + json: dict[str, Any] | None = None, + data: bytes | None = None, + headers: dict[str, str] | None = None, + ssl: ssl.SSLContext | bool | None = None, + ) -> tuple[int, dict[str, Any]]: + del url, json, data, headers, ssl + return 200, {"error_code": 0} + + transport._http_client.post = post_without_result # type: ignore[assignment] + with pytest.raises(KasaException, match="missing result object"): + await session._login({}, step_name="pake_register") + + +@pytest.mark.asyncio +async def test_login_tls2_uses_post() -> None: + transport = tp.TpapTransport(config=DeviceConfig("login-host")) + session = transport._encryption_session + session._tpap_tls = 2 + session._update_transport_url() + + async def post( + url: URL, + *, + params: dict[str, Any] | None = None, + data: bytes | None = None, + json: dict[str, Any] | None = None, + headers: dict[str, str] | None = None, + cookies_dict: dict[str, str] | None = None, + ssl: ssl.SSLContext | bool | None = None, + ) -> tuple[int, dict[str, Any]]: + del params, data, cookies_dict, ssl + assert url == URL("https://login-host:4433/") + assert json == {"method": "login", "params": {}} + assert headers == transport.COMMON_HEADERS + return 200, {"error_code": 0, "result": {}} + + transport._http_client.post = post # type: ignore[assignment] + + assert await session._login({}, step_name="pake_register") == {} + + +@pytest.mark.asyncio +async def test_update_transport_url_uses_https_default_or_bootstrap_port() -> None: + transport = tp.TpapTransport(config=DeviceConfig("tpap-host")) + session = transport._encryption_session + + session._tpap_tls = 1 + session._tpap_port = None + session._update_transport_url() + assert transport._app_url == URL("https://tpap-host:4433") + + session._tpap_tls = 0 + session._tpap_port = None + session._update_transport_url() + assert str(transport._app_url) == "http://tpap-host" + await transport.close() + + +@pytest.mark.asyncio +async def test_private_handle_response_error_code_covers_invalid_retry_auth_and_device() -> ( + None +): + transport = tp.TpapTransport(config=DeviceConfig("tpap-host")) + session = transport._encryption_session + session._tpap_tls = 2 + session._tpap_port = 4567 + transport._app_url = URL("https://tpap-host:4567") + _establish_session(transport, session, session_id="SID-AUTH", start_seq=4) + + with pytest.raises(DeviceError) as invalid_code: + session._handle_response_error_code({"error_code": "not-an-int"}, "ignored") + assert invalid_code.value.error_code is SmartErrorCode.INTERNAL_UNKNOWN_ERROR + + with pytest.raises(_RetryableError): + session._handle_response_error_code( + {"error_code": SmartErrorCode.SESSION_EXPIRED.value}, "retry" + ) + with pytest.raises(_RetryableError): + session._handle_response_error_code( + {"error_code": SmartErrorCode.STAT_ACCESS_ERROR.value}, "retry" + ) + + with pytest.raises(AuthenticationError): + session._handle_response_error_code( + {"error_code": SmartErrorCode.LOGIN_ERROR.value}, "auth" + ) + + assert session.is_established is False + assert session.tls_mode == 2 + assert session._tpap_port == 4567 + assert transport._app_url == URL("https://tpap-host:4567") + + with pytest.raises(DeviceError): + session._handle_response_error_code( + {"error_code": SmartErrorCode.DEVICE_ERROR.value}, "device" + ) + await transport.close() + + +# -------------------------- +# Helper and Credential Functions +# -------------------------- + + +def test_encode_w_and_hash_helpers_cover_edge_cases() -> None: + assert tp.TpapEncryptionSession._encode_w(0) == b"\x00" + assert tp.TpapEncryptionSession._encode_w(0x80) == b"\x00\x80" + assert ( + tp.TpapEncryptionSession._hash("SHA512", b"data") + == hashlib.sha512(b"data").digest() + ) + assert len(tp.TpapEncryptionSession._cmac_aes(b"\x00" * 16, b"data")) == 16 + + +def test_authkey_shadow_and_crypt_helpers() -> None: + masked = tp.TpapEncryptionSession._authkey_mask("abc", "xy", "0123456789") + assert len(masked) == 3 + assert ( + tp.TpapEncryptionSession._sha1_username_mac_shadow("", "AABBCCDDEEFF", "p") + == "p" + ) + assert tp.TpapEncryptionSession._sha1_username_mac_shadow( + "user", "AABBCCDDEEFF", "p" + ) == tp.TpapEncryptionSession._sha1_hex( + tp.TpapEncryptionSession._md5_hex("user") + "_AA:BB:CC:DD:EE:FF" + ) + assert tp.TpapEncryptionSession._md5_crypt("pass", "") is None + assert tp.TpapEncryptionSession._md5_crypt("pass", "$1$salt$") is not None + assert tp.TpapEncryptionSession._md5_crypt("pass", "$1$salt") is not None + assert tp.TpapEncryptionSession._sha256_crypt("pass", "") is None + assert ( + tp.TpapEncryptionSession._sha256_crypt("pass", "$5$rounds=oops$salt") + is not None + ) + assert ( + tp.TpapEncryptionSession._sha256_crypt( + "pass", "$5$salt$", rounds_from_params=2000 + ) + is not None + ) + assert ( + tp.TpapEncryptionSession._sha256_crypt( + "pass", + "$5$salt$", + rounds_from_params=cast(Any, "bad-rounds"), + ) + is not None + ) + + +@pytest.mark.parametrize( + ("extra_crypt", "username", "passcode", "mac_no_colon", "expected"), + [ + (None, "user", "pass", "AABBCCDDEEFF", "user/pass"), + ({}, "", "pass", "AABBCCDDEEFF", "pass"), + ( + {"type": "password_shadow", "params": {"passwd_id": 2}}, + "user", + "pass", + "AABBCCDDEEFF", + tp.TpapEncryptionSession._sha1_hex("pass"), + ), + ( + { + "type": "password_shadow", + "params": {"passwd_id": 3}, + }, + "user", + "pass", + "AABBCCDDEEFF", + tp.TpapEncryptionSession._sha1_username_mac_shadow( + "user", "AABBCCDDEEFF", "pass" + ), + ), + ( + { + "type": "password_authkey", + "params": { + "authkey_tmpkey": "xy", + "authkey_dictionary": "0123456789", + }, + }, + "user", + "pass", + "AABBCCDDEEFF", + tp.TpapEncryptionSession._authkey_mask("pass", "xy", "0123456789"), + ), + ( + { + "type": "password_sha_with_salt", + "params": { + "sha_name": 0, + "sha_salt": base64.b64encode(b"salt").decode(), + }, + }, + "user", + "pass", + "AABBCCDDEEFF", + hashlib.sha256(b"adminsaltpass").hexdigest(), + ), + ( + {"type": "unknown", "params": {}}, + "user", + "pass", + "AABBCCDDEEFF", + "user/pass", + ), + ], +) +def test_build_credentials_variants( + extra_crypt: dict[str, Any] | None, + username: str, + passcode: str, + mac_no_colon: str, + expected: str, +) -> None: + assert ( + tp.TpapEncryptionSession._build_credentials( + extra_crypt, username, passcode, mac_no_colon + ) + == expected + ) + + +@pytest.mark.parametrize( + ("extra_crypt", "expected"), + [ + pytest.param( + { + "type": "password_shadow", + "params": {"passwd_id": 1, "passwd_prefix": "$1$salt$"}, + }, + "not-pass", + id="shadow-md5-prefix", + ), + pytest.param( + { + "type": "password_shadow", + "params": {"passwd_id": 5, "passwd_prefix": "$5$salt$"}, + }, + "not-none", + id="shadow-sha256-prefix", + ), + pytest.param( + {"type": "password_authkey", "params": {}}, + "pass", + id="authkey-missing-params", + ), + pytest.param( + {"type": "password_shadow", "params": {"passwd_id": 99}}, + "pass", + id="shadow-unknown-id", + ), + pytest.param( + {"type": "password_shadow", "params": "not-a-dict"}, + "pass", + id="shadow-invalid-params", + ), + pytest.param( + {"type": "password_shadow", "params": {"passwd_id": "bad"}}, + "pass", + id="shadow-invalid-passwd-id", + ), + pytest.param( + { + "type": "password_sha_with_salt", + "params": {"sha_name": "bad", "sha_salt": "c2FsdA=="}, + }, + "pass", + id="sha-with-salt-invalid-sha-name", + ), + ], +) +def test_build_credentials_fallback_paths( + extra_crypt: dict[str, Any], expected: str +) -> None: + result = tp.TpapEncryptionSession._build_credentials( + extra_crypt, + "user", + "pass", + "AABBCCDDEEFF", + ) + + if expected == "not-pass": + assert result != "pass" + elif expected == "not-none": + assert result is not None + else: + assert result == expected + + +def test_mac_pass_from_device_mac_validates_input() -> None: + with pytest.raises(KasaException, match="Invalid device MAC"): + tp.TpapEncryptionSession._mac_pass_from_device_mac("not-a-mac") + + with pytest.raises(KasaException, match="too short"): + tp.TpapEncryptionSession._mac_pass_from_device_mac("AA:BB:CC:DD:EE") + + +# -------------------------- +# Register, Share, and Suite Handling +# -------------------------- + + +@pytest.mark.parametrize( + ("suite_type", "curve_name"), + [ + (3, "NIST384p"), + (5, "NIST521p"), + ], +) +def test_suite_parameters_support_additional_curves( + suite_type: int, curve_name: str +) -> None: + _, _, curve, _ = tp.TpapEncryptionSession._suite_parameters(suite_type) + assert curve.name == curve_name + + +def test_suite_parameters_reject_unsupported_suite() -> None: + with pytest.raises(KasaException, match="Unsupported TPAP suite type"): + tp.TpapEncryptionSession._suite_parameters(999) + + +async def test_build_share_params_from_register_requires_user_random() -> None: + transport = _make_tpap_transport() + session = transport._encryption_session + + with pytest.raises(KasaException, match="user random not initialized"): + session._build_share_params_from_register({}, "secret") + + +@pytest.mark.parametrize( + ("overrides", "match"), + [ + pytest.param({"dev_random": ""}, "missing dev_random", id="missing-dev-random"), + pytest.param({"dev_salt": ""}, "missing dev_salt", id="missing-dev-salt"), + pytest.param({"dev_share": ""}, "missing dev_share", id="missing-dev-share"), + pytest.param( + {"cipher_suites": "bad"}, + "has invalid cipher_suites", + id="invalid-cipher-suites-str", + ), + pytest.param( + {"cipher_suites": None}, + "has invalid cipher_suites", + id="invalid-cipher-suites-none", + ), + pytest.param( + {"iterations": 0}, "has invalid iterations", id="invalid-iterations-0" + ), + pytest.param( + {"iterations": None}, + "has invalid iterations", + id="invalid-iterations-none", + ), + pytest.param( + {"iterations": "bad"}, + "has invalid iterations", + id="invalid-iterations-str", + ), + pytest.param( + {"encryption": "bogus-cipher"}, + "Unsupported TPAP session cipher", + id="unsupported-cipher", + ), + pytest.param({"encryption": ""}, "missing encryption", id="missing-encryption"), + ], +) +async def test_build_share_params_from_register_validates_required_fields( + overrides: dict[str, Any], match: str +) -> None: + transport = _make_tpap_transport() + session = transport._encryption_session + session._user_random = base64.b64encode(b"\x01" * 16).decode() + + with pytest.raises(KasaException, match=match): + session._build_share_params_from_register( + {**_register_result(), **overrides}, + "secret", + ) + + +async def test_build_share_params_from_register_uses_cmac_suites() -> None: + transport = _make_tpap_transport() + session = transport._encryption_session + session._user_random = base64.b64encode(b"\x01" * 16).decode() + + share_params = session._build_share_params_from_register( + { + "dev_random": base64.b64encode(b"\x00" * 16).decode(), + "dev_salt": base64.b64encode(b"\x11" * 16).decode(), + "dev_share": base64.b64encode(_p256_pub_uncompressed()).decode(), + "cipher_suites": 8, + "iterations": 100, + "encryption": "aes_128_ccm", + }, + "secret", + ) + + assert session._expected_dev_confirm is not None + assert share_params["user_confirm"] + + +@pytest.mark.asyncio +async def test_verify_dac_proof_returns_early_without_required_fields() -> None: + transport = tp.TpapTransport(config=DeviceConfig("tpap-host")) + session = transport._encryption_session + session._verify_dac_proof({}) + await transport.close() + + +@pytest.mark.asyncio +async def test_verify_dac_proof_wraps_non_signature_errors() -> None: + transport = tp.TpapTransport(config=DeviceConfig("tpap-host")) + session = transport._encryption_session + session._shared_key = b"shared" + session._dac_nonce_base64 = base64.b64encode(b"nonce").decode() + + with pytest.raises(KasaException, match="DAC verification failed"): + session._verify_dac_proof({"dac_ca": "not-a-cert", "dac_proof": "not-b64"}) + await transport.close() + + +@pytest.mark.asyncio +async def test_verify_dac_proof_rejects_invalid_proof_type_and_public_key( + monkeypatch: pytest.MonkeyPatch, +) -> None: + transport = tp.TpapTransport(config=DeviceConfig("tpap-host")) + session = transport._encryption_session + session._shared_key = b"shared" + session._dac_nonce_base64 = base64.b64encode(b"nonce").decode() + + cert_key = ec.generate_private_key(ec.SECP256R1()) + cert = _build_certificate(cert_key, "root", "root", cert_key, is_ca=True) + + monkeypatch.setattr( + transport, + "_load_certificate_value", + lambda value: cert, + raising=True, + ) + monkeypatch.setattr( + transport, + "_verify_dac_certificate_chain", + lambda dac_ca_certificate, dac_ica_certificate: None, + raising=True, + ) + + with pytest.raises(KasaException, match="Invalid DAC proof type"): + session._verify_dac_proof({"dac_ca": "cert", "dac_proof": 1}) + + bad_cert = cast(Any, SimpleNamespace(public_key=lambda: object())) + monkeypatch.setattr( + transport, + "_load_certificate_value", + lambda value: bad_cert, + raising=True, + ) + with pytest.raises(KasaException, match="Unsupported DAC proof public key type"): + session._verify_dac_proof( + {"dac_ca": "cert", "dac_proof": base64.b64encode(b"proof").decode()} + ) + await transport.close() + + +@pytest.mark.asyncio +async def test_establish_session_from_share_result_error_paths( + monkeypatch: pytest.MonkeyPatch, +) -> None: + transport = tp.TpapTransport(config=DeviceConfig("tpap-host")) + session = transport._encryption_session + session._expected_dev_confirm = "expected" + + with pytest.raises(KasaException, match="missing dev_confirm"): + session._establish_session_from_share_result({}) + + with pytest.raises(KasaException, match="confirmation mismatch"): + session._establish_session_from_share_result({"dev_confirm": "wrong"}) + + monkeypatch.setattr(session, "_use_dac_certification", lambda: True, raising=True) + verified: list[dict[str, Any]] = [] + monkeypatch.setattr( + session, "_verify_dac_proof", lambda share_result: verified.append(share_result) + ) + session._shared_key = b"shared" + session._establish_session_from_share_result( + {"dev_confirm": "expected", "stok": "STOK", "start_seq": 3} + ) + assert verified + assert session._session_id == "STOK" + assert session._sequence == 3 + + session._expected_dev_confirm = "expected" + session._shared_key = b"shared" + with pytest.raises(KasaException, match="Missing session fields"): + session._establish_session_from_share_result({"dev_confirm": "expected"}) + + session._expected_dev_confirm = "expected" + session._shared_key = b"shared" + with pytest.raises(KasaException, match="Missing session fields"): + session._establish_session_from_share_result( + {"dev_confirm": "expected", "sessionId": "SID"} + ) + + session._expected_dev_confirm = "expected" + session._shared_key = b"shared" + with pytest.raises(KasaException, match="Invalid session fields"): + session._establish_session_from_share_result( + { + "dev_confirm": "expected", + "sessionId": "SID", + "start_seq": "bad", + } + ) + + session._expected_dev_confirm = "expected" + session._shared_key = None + with pytest.raises(KasaException, match="shared key was not derived"): + session._establish_session_from_share_result( + {"dev_confirm": "expected", "sessionId": "SID", "start_seq": 1} + ) + await transport.close() + + +# -------------------------- +# Certificate and TLS Helpers +# -------------------------- + + +def test_cipher_and_nonce_helpers() -> None: + with pytest.raises(KasaException, match="Unsupported TPAP session cipher"): + tp.TpapEncryptionSession._cipher_parameters("bogus") + with pytest.raises(ValueError, match="base nonce too short"): + tp.TpapEncryptionSession._nonce_from_base(b"\x00\x01\x02", 1) + + +def test_load_certificate_value_variants() -> None: + cert_key = ec.generate_private_key(ec.SECP256R1()) + cert = _build_certificate(cert_key, "leaf", "leaf", cert_key) + pem = cert.public_bytes(serialization.Encoding.PEM).decode() + der_b64 = base64.b64encode(cert.public_bytes(serialization.Encoding.DER)).decode() + + assert tp.TpapTransport._load_certificate_value(pem).subject == cert.subject + assert tp.TpapTransport._load_certificate_value(der_b64).subject == cert.subject + + with pytest.raises(KasaException, match="Empty certificate value"): + tp.TpapTransport._load_certificate_value(" ") + with pytest.raises(KasaException, match="Invalid certificate value"): + tp.TpapTransport._load_certificate_value("totally-invalid") + + +def test_verify_certificate_validity_handles_naive_datetimes() -> None: + now = datetime.now(UTC).replace(tzinfo=None) + valid = SimpleNamespace( + not_valid_before=now - timedelta(days=1), + not_valid_after=now + timedelta(days=1), + ) + tp.TpapTransport._verify_certificate_validity(cast(Any, valid)) + + expired = SimpleNamespace( + not_valid_before=now - timedelta(days=3), + not_valid_after=now - timedelta(days=2), + ) + with pytest.raises(KasaException, match="outside its validity period"): + tp.TpapTransport._verify_certificate_validity(cast(Any, expired)) + + +def test_verify_certificate_signature_variants() -> None: + ec_root = ec.generate_private_key(ec.SECP256R1()) + ec_leaf = ec.generate_private_key(ec.SECP256R1()) + ec_cert = _build_certificate(ec_leaf, "leaf", "root", ec_root) + ec_issuer = _build_certificate(ec_root, "root", "root", ec_root, is_ca=True) + tp.TpapTransport._verify_certificate_signature(ec_cert, ec_issuer) + + rsa_root = rsa.generate_private_key(public_exponent=65537, key_size=2048) + rsa_leaf = rsa.generate_private_key(public_exponent=65537, key_size=2048) + rsa_cert = _build_certificate(rsa_leaf, "leaf", "root", rsa_root) + rsa_issuer = _build_certificate(rsa_root, "root", "root", rsa_root, is_ca=True) + tp.TpapTransport._verify_certificate_signature(rsa_cert, rsa_issuer) + + no_hash_cert = SimpleNamespace(signature_hash_algorithm=None) + with pytest.raises(KasaException, match="hash algorithm is unavailable"): + tp.TpapTransport._verify_certificate_signature( + cast(Any, no_hash_cert), + cast(Any, SimpleNamespace(public_key=lambda: ec_root.public_key())), + ) + + bad_issuer = SimpleNamespace(public_key=lambda: object()) + with pytest.raises(KasaException, match="Unsupported DAC issuer public key type"): + tp.TpapTransport._verify_certificate_signature(ec_cert, cast(Any, bad_issuer)) + + +def test_verify_dac_certificate_chain_variants( + monkeypatch: pytest.MonkeyPatch, +) -> None: + root_key = ec.generate_private_key(ec.SECP256R1()) + root_cert = _build_certificate(root_key, "root", "root", root_key, is_ca=True) + dac_key = ec.generate_private_key(ec.SECP256R1()) + dac_cert = _build_certificate(dac_key, "dac", "root", root_key) + + monkeypatch.setattr( + tp.TpapTransport, + "_load_root_ca_certificate", + classmethod(lambda cls: root_cert), + raising=True, + ) + tp.TpapTransport._verify_dac_certificate_chain(dac_cert, None) + + monkeypatch.setattr( + tp.TpapTransport, + "_load_root_ca_certificate", + classmethod(lambda cls: (_ for _ in ()).throw(ValueError("boom"))), + raising=True, + ) + with pytest.raises( + KasaException, match="DAC certificate chain verification failed" + ): + tp.TpapTransport._verify_dac_certificate_chain(dac_cert, None) + + +# -------------------------- +# Transport and Payload Handling +# -------------------------- + + +@pytest.mark.asyncio +async def test_transport_properties_and_initial_url_helpers() -> None: + config = DeviceConfig("tpap-host") + config.credentials_hash = "hash" + config.connection_type.http_port = 8080 + transport = tp.TpapTransport(config=config) + transport._known_tpap_tls = 2 + + assert transport.default_port == 8080 + assert transport.credentials_hash == "hash" + assert transport._get_initial_app_url() == URL("https://tpap-host:4433") + await transport.close() + + https_transport = tp.TpapTransport(config=DeviceConfig("secure-host")) + https_transport._config.connection_type.https = True + assert https_transport.default_port == tp.TpapTransport.DEFAULT_HTTPS_PORT + await https_transport.close() + + +@pytest.mark.asyncio +async def test_transport_credentials_hash_encodes_live_credentials() -> None: + transport = tp.TpapTransport( + config=DeviceConfig("tpap-host", credentials=Credentials("user", "pass")) + ) + + credentials_hash = transport.credentials_hash + + assert credentials_hash is not None + decoded_hash = json.loads(base64.b64decode(credentials_hash.encode()).decode()) + assert decoded_hash == {"un": "user", "pwd": "pass"} + await transport.close() + + +@pytest.mark.asyncio +async def test_transport_decodes_credentials_hash_for_tpap_handshake() -> None: + seed_transport = tp.TpapTransport( + config=DeviceConfig("tpap-host", credentials=Credentials("user", "pass")) + ) + credentials_hash = seed_transport.credentials_hash + assert credentials_hash is not None + + config = DeviceConfig("tpap-host", credentials_hash=credentials_hash) + transport = tp.TpapTransport(config=config) + session = transport._encryption_session + session._tpap_pake = [2] + + assert transport._credentials == Credentials("user", "pass") + assert config.credentials == transport._credentials + assert session._get_candidate_secrets() == ["pass"] + await seed_transport.close() + await transport.close() + + +@pytest.mark.asyncio +async def test_transport_ignores_malformed_credentials_hash() -> None: + transport = tp.TpapTransport( + config=DeviceConfig("tpap-host", credentials_hash="not-valid-base64-json") + ) + + assert transport._credentials is None + assert transport._config.credentials is None + assert transport.credentials_hash == "not-valid-base64-json" + await transport.close() + + +@pytest.mark.parametrize( + ("exc", "expected"), + [ + pytest.param( + _ConnectionError("Connection reset by peer"), + True, + id="connection-reset", + ), + pytest.param( + _ConnectionError("different connection issue"), + False, + id="other-connection-error", + ), + pytest.param(KasaException("x"), False, id="generic-kasa"), + pytest.param( + _RetryableError("x", error_code=SmartErrorCode.LOGIN_ERROR), + False, + id="retryable-non-session-error", + ), + pytest.param( + _RetryableError("x", error_code=SmartErrorCode.STAT_ACCESS_ERROR), + True, + id="stat-access-error", + ), + ], +) +def test_should_retry_live_session_variants(exc: Exception, expected: bool) -> None: + assert tp.TpapTransport._should_retry_live_session(exc) is expected + + +@pytest.mark.asyncio +async def test_get_ssl_context_caches_result(monkeypatch: pytest.MonkeyPatch) -> None: + transport = tp.TpapTransport(config=DeviceConfig("tpap-host")) + created = 0 + + def fake_create_ssl_context() -> bool: + nonlocal created + created += 1 + return False + + monkeypatch.setattr(transport, "_create_ssl_context", fake_create_ssl_context) + + assert await transport.get_ssl_context() is False + assert await transport.get_ssl_context() is False + assert created == 1 + + +async def test_create_ssl_context_for_tls0_and_tls1() -> None: + transport = tp.TpapTransport(config=DeviceConfig("tpap-host")) + session = transport._encryption_session + + session._tpap_tls = 0 + assert transport._create_ssl_context() is False + + session._tpap_tls = 1 + context = transport._create_ssl_context() + assert isinstance(context, ssl.SSLContext) + assert context.verify_mode == ssl.CERT_NONE + + +def test_payload_and_sequence_helpers() -> None: + key, base_nonce = tp.TpapEncryptionSession.key_nonce_from_shared( + b"shared-secret", "aes_128_ccm" + ) + session = object.__new__(tp.TpapEncryptionSession) + session._cipher_id = "aes_128_ccm" + session._sequence = 5 + session._ds_url = URL("https://tpap-host/stok=SID/ds") + session._key = key + session._base_nonce = base_nonce + session._session_id = "SID" + + payload, seq = session.encrypt("hello") + assert seq == 5 + + session._sequence = 5 + session.advance(5) + assert session._sequence == 6 + session.advance(3) + assert session._sequence == 6 + + with pytest.raises(KasaException, match="response too short"): + session.decrypt(b"\x00\x00\x00\x01", 1) + + encrypted = tp.TpapEncryptionSession._encrypt_payload( + "aes_128_ccm", key, base_nonce, b"world", 8 + ) + wrapped = struct.pack(">I", 8) + encrypted + assert session.decrypt(wrapped, 7) == b"world" + + +@pytest.mark.asyncio +async def test_send_reraises_non_retryable_error() -> None: + transport = tp.TpapTransport(config=DeviceConfig("tpap-host")) + + async def fake_send_once(request: str) -> dict[str, Any]: + del request + raise KasaException("boom") + + transport._send_once = fake_send_once # type: ignore[assignment] + + with pytest.raises(KasaException, match="boom"): + await transport.send("{}") + + +@pytest.mark.asyncio +async def test_send_once_error_paths(monkeypatch: pytest.MonkeyPatch) -> None: + transport, session = _make_established_transport() + + async def fake_handshake() -> None: + session._ds_url = None + + session._invalidate_session() + monkeypatch.setattr(session, "perform_handshake", fake_handshake, raising=True) + with pytest.raises(KasaException, match="not established"): + await transport._send_once("{}") + + _establish_session(transport, session) + + async def post_bad_status( + url: URL, + *, + json: dict[str, Any] | None = None, + data: bytes | None = None, + headers: dict[str, str] | None = None, + ssl: ssl.SSLContext | bool | None = None, + ) -> tuple[int, bytes]: + del url, json, data, headers, ssl + return 500, b"bad" + + transport._http_client.post = post_bad_status # type: ignore[assignment] + with pytest.raises(KasaException, match="secure request failed.*status 500"): + await transport._send_once("{}") + + async def post_dict_success( + url: URL, + *, + json: dict[str, Any] | None = None, + data: bytes | None = None, + headers: dict[str, str] | None = None, + ssl: ssl.SSLContext | bool | None = None, + ) -> tuple[int, dict[str, Any]]: + del url, json, data, headers, ssl + return 200, {"error_code": 0, "result": {"ok": True}} + + transport._http_client.post = post_dict_success # type: ignore[assignment] + assert (await transport._send_once("{}"))["result"]["ok"] is True + + async def post_weird_type( + url: URL, + *, + json: dict[str, Any] | None = None, + data: bytes | None = None, + headers: dict[str, str] | None = None, + ssl: ssl.SSLContext | bool | None = None, + ) -> tuple[int, int]: + del url, json, data, headers, ssl + return 200, 123 + + transport._http_client.post = post_weird_type # type: ignore[assignment] + with pytest.raises(KasaException, match="Unexpected TPAP response body type"): + await transport._send_once("{}") + + +@pytest.mark.asyncio +async def test_send_once_tls2_uses_post() -> None: + transport, session = _make_established_transport() + session._tpap_tls = 2 + + async def post( + url: URL, + *, + params: dict[str, Any] | None = None, + data: bytes | None = None, + json: dict[str, Any] | None = None, + headers: dict[str, str] | None = None, + cookies_dict: dict[str, str] | None = None, + ssl: ssl.SSLContext | bool | None = None, + ) -> tuple[int, bytes]: + del params, json, cookies_dict, ssl + assert url == session.ds_url + assert headers == {"Content-Type": "application/octet-stream"} + assert data is not None + return 200, data + + transport._http_client.post = post # type: ignore[assignment] + + out = await transport._send_once('{"result": {"ok": true}}') + assert out["result"]["ok"] is True + + +@pytest.mark.asyncio +async def test_transport_close_resets_and_closes_http_client( + monkeypatch: pytest.MonkeyPatch, +) -> None: + transport, session = _make_established_transport() + closed = False + + async def fake_close() -> None: + nonlocal closed + closed = True + + monkeypatch.setattr(transport._http_client, "close", fake_close, raising=True) + + await transport.close() + + assert closed is True + assert session.is_established is False + assert ( + tp.TpapEncryptionSession._build_credentials( + { + "type": "password_sha_with_salt", + "params": {"sha_name": 1, "sha_salt": "not-b64"}, + }, + "user", + "pass", + "AABBCCDDEEFF", + ) + == "pass" + ) + + +def test_transport_response_helpers_validate_json_payloads() -> None: + with pytest.raises(KasaException, match="Unexpected TPAP JSON response body type"): + tp.TpapTransport._load_json_dict(b"[]") diff --git a/uv.lock b/uv.lock index e32696bd6..7332b06d4 100644 --- a/uv.lock +++ b/uv.lock @@ -556,6 +556,18 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/26/87/f238c0670b94533ac0353a4e2a1a771a0cc73277b88bff23d3ae35a256c1/docutils-0.20.1-py3-none-any.whl", hash = "sha256:96f387a2c5562db4476f09f13bbab2192e764cac08ebbf3a34a95d9b1e4a59d6", size = 572666, upload-time = "2023-05-16T23:39:15.976Z" }, ] +[[package]] +name = "ecdsa" +version = "0.19.2" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "six" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/25/ca/8de7744cb3bc966c85430ca2d0fcaeea872507c6a4cf6e007f7fe269ed9d/ecdsa-0.19.2.tar.gz", hash = "sha256:62635b0ac1ca2e027f82122b5b81cb706edc38cd91c63dda28e4f3455a2bf930", size = 202432, upload-time = "2026-03-26T09:58:17.675Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/51/79/119091c98e2bf49e24ed9f3ae69f816d715d2904aefa6a2baa039a2ba0b0/ecdsa-0.19.2-py2.py3-none-any.whl", hash = "sha256:840f5dc5e375c68f36c1a7a5b9caad28f95daa65185c9253c0c08dd952bb7399", size = 150818, upload-time = "2026-03-26T09:58:15.808Z" }, +] + [[package]] name = "execnet" version = "2.1.2" @@ -1292,6 +1304,15 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/99/5d/8268b644392ee874ee82a635cd0df1773de230bde356c38de28e298392cc/parso-0.8.7-py2.py3-none-any.whl", hash = "sha256:a8926eb2a1b915486941fdbd31e86a4baf88fe8c210f25f2f35ecec5b574ca1c", size = 107025, upload-time = "2026-05-01T23:12:58.867Z" }, ] +[[package]] +name = "passlib" +version = "1.7.4" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/b6/06/9da9ee59a67fae7761aab3ccc84fa4f3f33f125b370f1ccdb915bf967c11/passlib-1.7.4.tar.gz", hash = "sha256:defd50f72b65c5402ab2c573830a6978e5f202ad0d984793c8dde2c4152ebe04", size = 689844, upload-time = "2020-10-08T19:00:52.121Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/3b/a4/ab6b7589382ca3df236e03faa71deac88cae040af60c071a78d254a62172/passlib-1.7.4-py2.py3-none-any.whl", hash = "sha256:aa6bca462b8d8bda89c70b382f0c298a20b5560af6cbfa2dce410c0a2fb669f1", size = 525554, upload-time = "2020-10-08T19:00:49.856Z" }, +] + [[package]] name = "pathspec" version = "1.1.1" @@ -1642,7 +1663,9 @@ dependencies = [ { name = "aiohttp" }, { name = "asyncclick" }, { name = "cryptography" }, + { name = "ecdsa" }, { name = "mashumaro" }, + { name = "passlib" }, { name = "tzdata", marker = "sys_platform == 'win32'" }, ] @@ -1690,9 +1713,11 @@ requires-dist = [ { name = "aiohttp", specifier = ">=3" }, { name = "asyncclick", specifier = ">=8.1.7" }, { name = "cryptography", specifier = ">=1.9" }, + { name = "ecdsa", specifier = ">=0.19.1" }, { name = "kasa-crypt", marker = "extra == 'speedups'", specifier = ">=0.2.0" }, { name = "mashumaro", specifier = ">=3.20" }, { name = "orjson", marker = "extra == 'speedups'", specifier = ">=3.11.1" }, + { name = "passlib", specifier = ">=1.7.4" }, { name = "ptpython", marker = "extra == 'shell'" }, { name = "rich", marker = "extra == 'shell'" }, { name = "tzdata", marker = "sys_platform == 'win32'", specifier = ">=2024.2" },