diff --git a/CHANGELOG.md b/CHANGELOG.md index 11958a0..352bd84 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -2,6 +2,29 @@ ## [Unreleased] +### Fixed +- `SpaceControls.display_setpoint_str()` — removed dead unreachable code in the + fallback branch; `temperature_setpoint_c` is typed `float` so the `None`-guard + lines were never executed (confirmed by coverage) +- `SystemSnapshot.apply_outdoor_unit()` — refactored to collect field patches into + an `updates` dict and call `dataclasses.replace()` once, consistent with all other + `apply_*` methods; previously called `replace()` twice in sequence, creating an + unnecessary intermediate object +- `QuiltClient.close()` now sets `self._token = None`; previously left a stale token + accessible via `get_current_token()` after the channel was closed + +### Changed +- `invoke_refresh_callback` (formerly `_invoke_refresh_callback`) extracted from + `transport.py` and `services/streaming.py` into a single shared implementation in + `tokens.py`; the streaming copy lacked the `WeakKeyDictionary` signature cache, + causing `inspect.signature()` to be called on every token-refresh event +- `FanSpeed.to_wire()` and `LouverAngle.to_wire()` now reference module-level + constant dicts (`_FAN_SPEED_WIRE_MAP`, `_LOUVER_ANGLE_WIRE_MAP`) instead of + re-allocating the mapping on every call +- `_id_variants()` moved from `models/system.py` into `models/_helpers.py` and + reused by `lookup_hardware()`; eliminates duplicated ID-normalisation logic +- `QuiltClient.invalidate_snapshot()` log level changed from `WARNING` to `DEBUG` + ## [0.5.0] - 2026-06-04 ### Added protocol support diff --git a/src/quilt_hp/client.py b/src/quilt_hp/client.py index f8693dd..824925b 100644 --- a/src/quilt_hp/client.py +++ b/src/quilt_hp/client.py @@ -273,7 +273,7 @@ async def get_snapshot(self, system_id: str | None = None) -> SystemSnapshot: def invalidate_snapshot(self) -> None: """Discard the cached snapshot so the next call fetches fresh data.""" - logger.warning("Invalidating snapshot cache") + logger.debug("Invalidating snapshot cache") self._snapshot_cache = None self._snapshot_cached_at = 0.0 @@ -640,6 +640,7 @@ async def close(self) -> None: if self._channel is not None: await self._channel.close() self._channel = None + self._token = None self._hds = None self._sysinfo = None self._user_svc = None diff --git a/src/quilt_hp/models/_helpers.py b/src/quilt_hp/models/_helpers.py index ca7e313..f9d754d 100644 --- a/src/quilt_hp/models/_helpers.py +++ b/src/quilt_hp/models/_helpers.py @@ -1,22 +1,43 @@ from __future__ import annotations +def _id_variant_keys(raw: str) -> tuple[str, ...]: + """Return ID variant keys in deterministic priority order (exact → tail → casefold).""" + tail_slash = raw.rsplit("/", 1)[-1] + tail_colon = raw.rsplit(":", 1)[-1] + return ( + raw, + tail_slash, + tail_colon, + raw.casefold(), + tail_slash.casefold(), + tail_colon.casefold(), + ) + + +def _id_variants(value: str | None) -> set[str]: + """Return raw and normalized ID variants for matching resource IDs.""" + if not value: + return set() + raw = value.strip() + if not raw: + return set() + return {v for v in _id_variant_keys(raw) if v} + + def lookup_hardware(hw_map: dict[str, object], hardware_id: str | None) -> object | None: - """Resolve hardware objects across common ID formats.""" + """Resolve hardware objects across common ID formats. + + Keys are tried in deterministic priority order: exact → tail (after last + ``/`` or ``:`` separator) → casefold variants, matching the behaviour of + the original implementation. + """ if not hardware_id: return None raw = hardware_id.strip() if not raw: return None - keys = ( - raw, - raw.rsplit("/", 1)[-1], - raw.rsplit(":", 1)[-1], - raw.casefold(), - raw.rsplit("/", 1)[-1].casefold(), - raw.rsplit(":", 1)[-1].casefold(), - ) - for key in keys: + for key in _id_variant_keys(raw): hw = hw_map.get(key) if hw is not None: return hw diff --git a/src/quilt_hp/models/enums.py b/src/quilt_hp/models/enums.py index 093ab2b..fdc60ea 100644 --- a/src/quilt_hp/models/enums.py +++ b/src/quilt_hp/models/enums.py @@ -59,15 +59,7 @@ def __str__(self) -> str: def to_wire(self) -> tuple[int, float]: """Return (fan_speed_mode, fan_speed_percent) for the wire protocol.""" - _MAP: dict[FanSpeed, tuple[int, float]] = { - FanSpeed.AUTO: (1, 0.0), # FAN_SPEED_MODE_AUTO - FanSpeed.QUIET: (2, 0.20), # FAN_SPEED_MODE_SETPOINT - FanSpeed.LOW: (2, 0.40), - FanSpeed.MEDIUM: (2, 0.60), - FanSpeed.HIGH: (2, 0.80), - FanSpeed.BLAST: (2, 1.00), - } - return _MAP[self] + return _FAN_SPEED_WIRE_MAP[self.value] @classmethod def from_wire(cls, mode: int, percent: float) -> FanSpeed: @@ -85,6 +77,16 @@ def from_wire(cls, mode: int, percent: float) -> FanSpeed: return cls.BLAST +_FAN_SPEED_WIRE_MAP: dict[int, tuple[int, float]] = { + 0: (1, 0.0), # AUTO → FAN_SPEED_MODE_AUTO + 1: (2, 0.20), # QUIET → FAN_SPEED_MODE_SETPOINT + 2: (2, 0.40), # LOW + 3: (2, 0.60), # MEDIUM + 4: (2, 0.80), # HIGH + 5: (2, 1.00), # BLAST +} + + class LouverMode(IntEnum): """Indoor unit louver mode.""" @@ -129,7 +131,7 @@ def __str__(self) -> str: def to_wire(self) -> float: """Return the louver_fixed_position float for the wire.""" - return {1: 0.20, 2: 0.40, 3: 0.60, 4: 0.80, 5: 1.00}[self.value] + return _LOUVER_ANGLE_WIRE_MAP[self.value] @classmethod def from_wire(cls, position: float) -> LouverAngle: @@ -145,6 +147,9 @@ def from_wire(cls, position: float) -> LouverAngle: return cls.ANGLE5 +_LOUVER_ANGLE_WIRE_MAP: dict[int, float] = {1: 0.20, 2: 0.40, 3: 0.60, 4: 0.80, 5: 1.00} + + class LightPreset(IntEnum): """Built-in LED color presets (RGBW packed int32).""" diff --git a/src/quilt_hp/models/space.py b/src/quilt_hp/models/space.py index ae77041..042087b 100644 --- a/src/quilt_hp/models/space.py +++ b/src/quilt_hp/models/space.py @@ -85,12 +85,7 @@ def fmt(val_c: float) -> str: return fmt(self.heating_setpoint_c) if mode == HVACMode.AUTO: return f"{fmt(self.heating_setpoint_c)}–{fmt(self.cooling_setpoint_c)}" - best = self.temperature_setpoint_c - if best is None: - best = self.cooling_setpoint_c - if best is None: - best = self.heating_setpoint_c - return fmt(best) if best is not None else "--" + return fmt(self.temperature_setpoint_c) @property def has_standby_sentinel_setpoints(self) -> bool: diff --git a/src/quilt_hp/models/system.py b/src/quilt_hp/models/system.py index 82f0573..5e7214e 100644 --- a/src/quilt_hp/models/system.py +++ b/src/quilt_hp/models/system.py @@ -5,6 +5,7 @@ from dataclasses import dataclass from typing import Any, cast +from quilt_hp.models._helpers import _id_variants from quilt_hp.models.comfort import ComfortSetting from quilt_hp.models.controller import Controller from quilt_hp.models.enums import ( @@ -24,21 +25,6 @@ from quilt_hp.models.space import Space -def _id_variants(value: str | None) -> set[str]: - """Return raw and normalized ID variants for matching resource IDs.""" - if not value: - return set() - raw = value.strip() - if not raw: - return set() - tail_slash = raw.rsplit("/", 1)[-1] - tail_colon = raw.rsplit(":", 1)[-1] - variants = {raw, tail_slash, tail_colon, raw.casefold()} - variants.add(tail_slash.casefold()) - variants.add(tail_colon.casefold()) - return {v for v in variants if v} - - @dataclass(slots=True) class Location: """A Quilt location with global settings like schedule execution state.""" @@ -283,17 +269,17 @@ def apply_outdoor_unit(self, odu: OutdoorUnit) -> OutdoorUnit: for i, u in enumerate(self.outdoor_units): if u.id == odu.id: + updates: dict[str, Any] = {} # Preserve hvac_state when stream diff has a default-zero state if not odu.hvac_state and u.hvac_state: - odu = replace(odu, hvac_state=u.hvac_state) + updates["hvac_state"] = u.hvac_state # Preserve hardware info — stream diffs are parsed without hw_map if odu.model_sku is None and u.model_sku is not None: - odu = replace( - odu, - model_sku=u.model_sku, - serial_number=u.serial_number, - firmware_version=u.firmware_version, - ) + updates["model_sku"] = u.model_sku + updates["serial_number"] = u.serial_number + updates["firmware_version"] = u.firmware_version + if updates: + odu = replace(odu, **updates) self.outdoor_units[i] = odu return odu self.outdoor_units.append(odu) diff --git a/src/quilt_hp/services/streaming.py b/src/quilt_hp/services/streaming.py index 7506309..984f97e 100644 --- a/src/quilt_hp/services/streaming.py +++ b/src/quilt_hp/services/streaming.py @@ -9,7 +9,6 @@ import asyncio import contextlib -import inspect import logging import time from collections.abc import AsyncIterator, Awaitable, Callable, Sequence @@ -30,7 +29,7 @@ from quilt_hp.models.sensor import ControllerRemoteSensor, RemoteSensor from quilt_hp.models.software_update import SoftwareUpdateInfo from quilt_hp.models.space import Space -from quilt_hp.tokens import TokenRefreshContext, TokenRefreshReason +from quilt_hp.tokens import TokenRefreshContext, TokenRefreshReason, invoke_refresh_callback logger = logging.getLogger(__name__) @@ -60,21 +59,6 @@ def Subscribe( type _AnyCallback = Callable[[Any], Awaitable[None] | None] -async def _invoke_refresh_callback( - refresh_callback: RefreshCallback, context: TokenRefreshContext -) -> None: - try: - has_params = bool(inspect.signature(refresh_callback).parameters) - except TypeError: - has_params = False - except ValueError: - has_params = False - if has_params: - await cast("Callable[[TokenRefreshContext], Awaitable[None]]", refresh_callback)(context) - return - await cast("Callable[[], Awaitable[None]]", refresh_callback)() - - def _parse_varint(data: bytes, pos: int) -> tuple[int, int]: """Parse a protobuf varint from raw bytes.""" result, shift = 0, 0 @@ -624,7 +608,7 @@ async def _run_stream_with_reconnect(self) -> None: source="streaming", attempt=attempt + 1, ) - await _invoke_refresh_callback(self._authenticate, context) + await invoke_refresh_callback(self._authenticate, context) except Exception: logger.exception("Token refresh failed; giving up stream") self._error = exc diff --git a/src/quilt_hp/tokens.py b/src/quilt_hp/tokens.py index 210e88d..b9fbd2a 100644 --- a/src/quilt_hp/tokens.py +++ b/src/quilt_hp/tokens.py @@ -7,13 +7,20 @@ from __future__ import annotations +import inspect import time +import weakref +from collections.abc import Awaitable, Callable from dataclasses import dataclass from enum import StrEnum -from typing import Protocol +from typing import Protocol, cast _TOKEN_BUFFER_S = 300 # treat tokens as expired 5 min before actual expiry +# Cache whether a refresh callback accepts a TokenRefreshContext argument, +# so inspect.signature is only called once per unique callable. +_REFRESH_CALLBACK_HAS_PARAMS: weakref.WeakKeyDictionary[object, bool] = weakref.WeakKeyDictionary() + @dataclass(slots=True) class CachedTokens: @@ -117,3 +124,36 @@ def on_refresh_failure( ) -> RefreshFailureAction: """Return fallback strategy when refresh fails.""" ... + + +type _RefreshCallback = ( + Callable[[], Awaitable[None]] | Callable[[TokenRefreshContext], Awaitable[None]] +) + + +async def invoke_refresh_callback( + refresh_callback: _RefreshCallback, context: TokenRefreshContext +) -> None: + """Invoke a refresh callback, passing context only if it accepts a parameter. + + Whether each callback accepts a ``TokenRefreshContext`` argument is cached + per-callable in a WeakKeyDictionary so that ``inspect.signature`` is only + called once per unique callback object. + """ + try: + has_params = _REFRESH_CALLBACK_HAS_PARAMS.get(refresh_callback) + except TypeError: + has_params = None # non-weakrefable callable — skip cache + if has_params is None: + try: + has_params = bool(inspect.signature(refresh_callback).parameters) + except TypeError, ValueError: + has_params = False + try: + _REFRESH_CALLBACK_HAS_PARAMS[refresh_callback] = has_params + except TypeError: + pass # non-weakrefable callable — skip caching + if has_params: + await cast("Callable[[TokenRefreshContext], Awaitable[None]]", refresh_callback)(context) + return + await cast("Callable[[], Awaitable[None]]", refresh_callback)() diff --git a/src/quilt_hp/transport.py b/src/quilt_hp/transport.py index 7adb1e8..f84655e 100644 --- a/src/quilt_hp/transport.py +++ b/src/quilt_hp/transport.py @@ -2,11 +2,8 @@ from __future__ import annotations -import inspect import logging -import weakref from collections.abc import Awaitable, Callable -from typing import cast import grpc import grpc.aio @@ -18,7 +15,12 @@ grpc_host, ) from quilt_hp.exceptions import QuiltAuthError -from quilt_hp.tokens import CurrentTokenProvider, TokenRefreshContext, TokenRefreshReason +from quilt_hp.tokens import ( + CurrentTokenProvider, + TokenRefreshContext, + TokenRefreshReason, + invoke_refresh_callback, +) type RefreshCallback = ( Callable[[], Awaitable[None]] | Callable[[TokenRefreshContext], Awaitable[None]] @@ -26,7 +28,6 @@ type TokenProviderLike = Callable[[], str] | CurrentTokenProvider logger = logging.getLogger(__name__) -_REFRESH_CALLBACK_HAS_PARAMS: weakref.WeakKeyDictionary[object, bool] = weakref.WeakKeyDictionary() def _resolve_token_provider(token_provider: TokenProviderLike) -> Callable[[], str]: @@ -35,28 +36,6 @@ def _resolve_token_provider(token_provider: TokenProviderLike) -> Callable[[], s return token_provider.get_current_token -async def _invoke_refresh_callback( - refresh_callback: RefreshCallback, context: TokenRefreshContext -) -> None: - try: - has_params = _REFRESH_CALLBACK_HAS_PARAMS.get(refresh_callback) - except TypeError: - has_params = None # non-weakrefable callable — skip cache - if has_params is None: - try: - has_params = bool(inspect.signature(refresh_callback).parameters) - except TypeError, ValueError: - has_params = False - try: - _REFRESH_CALLBACK_HAS_PARAMS[refresh_callback] = has_params - except TypeError: - pass # non-weakrefable callable — skip caching - if has_params: - await cast("Callable[[TokenRefreshContext], Awaitable[None]]", refresh_callback)(context) - return - await cast("Callable[[], Awaitable[None]]", refresh_callback)() - - class _AuthInterceptor( grpc.aio.UnaryUnaryClientInterceptor, # type: ignore[misc] grpc.aio.UnaryStreamClientInterceptor, # type: ignore[misc] @@ -109,7 +88,7 @@ async def _refresh_and_retry( or the credentials are otherwise invalid. """ if self._refresh_callback is not None: - await _invoke_refresh_callback( + await invoke_refresh_callback( self._refresh_callback, TokenRefreshContext( reason=TokenRefreshReason.TRANSPORT_UNAUTHENTICATED, diff --git a/tests/test_streaming.py b/tests/test_streaming.py index c4d4871..90071ab 100644 --- a/tests/test_streaming.py +++ b/tests/test_streaming.py @@ -12,10 +12,10 @@ NotifierStream, _dispatch, _get_len_field, - _invoke_refresh_callback, _parse_varint, ) from quilt_hp.tokens import TokenRefreshContext, TokenRefreshReason +from quilt_hp.tokens import invoke_refresh_callback as _invoke_refresh_callback class _FakeRpcError(grpc.aio.AioRpcError): diff --git a/tests/test_transport.py b/tests/test_transport.py index 6f07318..a460a34 100644 --- a/tests/test_transport.py +++ b/tests/test_transport.py @@ -6,7 +6,7 @@ from quilt_hp import transport from quilt_hp.const import APP_VERSION, Environment, grpc_host -from quilt_hp.tokens import TokenRefreshContext, TokenRefreshReason +from quilt_hp.tokens import TokenRefreshContext, TokenRefreshReason, invoke_refresh_callback def test_grpc_host_prod() -> None: @@ -45,7 +45,7 @@ async def _with_context(context: TokenRefreshContext) -> None: reason=TokenRefreshReason.TRANSPORT_UNAUTHENTICATED, source="test", ) - await transport._invoke_refresh_callback(_with_context, context) + await invoke_refresh_callback(_with_context, context) assert captured == [context] @@ -60,5 +60,5 @@ async def _legacy() -> None: reason=TokenRefreshReason.TRANSPORT_UNAUTHENTICATED, source="test", ) - await transport._invoke_refresh_callback(_legacy, context) + await invoke_refresh_callback(_legacy, context) assert calls == ["called"] diff --git a/tests/test_transport_interceptor_extra.py b/tests/test_transport_interceptor_extra.py index 1c54604..06eb950 100644 --- a/tests/test_transport_interceptor_extra.py +++ b/tests/test_transport_interceptor_extra.py @@ -6,7 +6,7 @@ import grpc import pytest -from quilt_hp import transport +from quilt_hp import tokens, transport from quilt_hp.const import Environment from quilt_hp.exceptions import QuiltAuthError @@ -32,12 +32,12 @@ async def test_invoke_refresh_callback_handles_signature_fallback( async def _legacy() -> None: called.append("legacy") - monkeypatch.setattr(transport.inspect, "signature", MagicMock(side_effect=TypeError("bad"))) + monkeypatch.setattr(tokens.inspect, "signature", MagicMock(side_effect=TypeError("bad"))) - await transport._invoke_refresh_callback( + await tokens.invoke_refresh_callback( _legacy, - transport.TokenRefreshContext( - reason=transport.TokenRefreshReason.TRANSPORT_UNAUTHENTICATED, + tokens.TokenRefreshContext( + reason=tokens.TokenRefreshReason.TRANSPORT_UNAUTHENTICATED, source="test", ), ) @@ -48,10 +48,10 @@ async def _legacy() -> None: async def test_invoke_refresh_callback_caches_signature( monkeypatch: pytest.MonkeyPatch, ) -> None: - transport._REFRESH_CALLBACK_HAS_PARAMS.clear() - called: list[transport.TokenRefreshContext] = [] + tokens._REFRESH_CALLBACK_HAS_PARAMS.clear() + called: list[tokens.TokenRefreshContext] = [] - async def _with_context(context: transport.TokenRefreshContext) -> None: + async def _with_context(context: tokens.TokenRefreshContext) -> None: called.append(context) signature = inspect.Signature( @@ -63,14 +63,14 @@ async def _with_context(context: transport.TokenRefreshContext) -> None: ] ) inspect_signature = MagicMock(return_value=signature) - monkeypatch.setattr(transport.inspect, "signature", inspect_signature) + monkeypatch.setattr(tokens.inspect, "signature", inspect_signature) - context = transport.TokenRefreshContext( - reason=transport.TokenRefreshReason.TRANSPORT_UNAUTHENTICATED, + context = tokens.TokenRefreshContext( + reason=tokens.TokenRefreshReason.TRANSPORT_UNAUTHENTICATED, source="test", ) - await transport._invoke_refresh_callback(_with_context, context) - await transport._invoke_refresh_callback(_with_context, context) + await tokens.invoke_refresh_callback(_with_context, context) + await tokens.invoke_refresh_callback(_with_context, context) assert called == [context, context] assert inspect_signature.call_count == 1