diff --git a/homeassistant/components/usage_prediction/__init__.py b/homeassistant/components/usage_prediction/__init__.py index 72198e42022f..f74e500896f5 100644 --- a/homeassistant/components/usage_prediction/__init__.py +++ b/homeassistant/components/usage_prediction/__init__.py @@ -4,6 +4,8 @@ import asyncio from datetime import timedelta from typing import Any +import probatio + from homeassistant.components import websocket_api from homeassistant.core import HomeAssistant from homeassistant.helpers import config_validation as cv @@ -11,7 +13,7 @@ from homeassistant.helpers.typing import ConfigType from homeassistant.util import dt as dt_util from . import common_control -from .const import DATA_CACHE, DOMAIN +from .const import DATA_CACHE, DEFAULT_LIMIT, DOMAIN from .models import EntityUsageDataCache, EntityUsagePredictions CONFIG_SCHEMA = cv.empty_config_schema(DOMAIN) @@ -26,7 +28,14 @@ async def async_setup(hass: HomeAssistant, config: ConfigType) -> bool: return True -@websocket_api.websocket_command({"type": f"{DOMAIN}/common_control"}) +@websocket_api.websocket_command( + { + probatio.Required("type"): f"{DOMAIN}/common_control", + probatio.Optional("limit", default=DEFAULT_LIMIT): probatio.All( + int, probatio.Range(min=1) + ), + } +) @websocket_api.async_response async def ws_common_control( hass: HomeAssistant, @@ -39,7 +48,7 @@ async def ws_common_control( connection.send_result( msg["id"], { - "entities": getattr(result, time_category), + "entities": getattr(result, time_category)[: msg["limit"]], }, ) diff --git a/homeassistant/components/usage_prediction/common_control.py b/homeassistant/components/usage_prediction/common_control.py index d3f1da724035..543d9de2af57 100644 --- a/homeassistant/components/usage_prediction/common_control.py +++ b/homeassistant/components/usage_prediction/common_control.py @@ -25,8 +25,6 @@ _LOGGER = logging.getLogger(__name__) # Time categories for usage patterns TIME_CATEGORIES = ["morning", "afternoon", "evening", "night"] -RESULTS_TO_INCLUDE = 8 - # Rows fetched per round trip while streaming the events query QUERY_YIELD_PER = 4096 @@ -106,19 +104,10 @@ async def async_predict_common_control( ) return EntityUsagePredictions( - morning=[ - ent_id for (ent_id, _) in results["morning"].most_common(RESULTS_TO_INCLUDE) - ], - afternoon=[ - ent_id - for (ent_id, _) in results["afternoon"].most_common(RESULTS_TO_INCLUDE) - ], - evening=[ - ent_id for (ent_id, _) in results["evening"].most_common(RESULTS_TO_INCLUDE) - ], - night=[ - ent_id for (ent_id, _) in results["night"].most_common(RESULTS_TO_INCLUDE) - ], + morning=[ent_id for (ent_id, _) in results["morning"].most_common()], + afternoon=[ent_id for (ent_id, _) in results["afternoon"].most_common()], + evening=[ent_id for (ent_id, _) in results["evening"].most_common()], + night=[ent_id for (ent_id, _) in results["night"].most_common()], ) diff --git a/homeassistant/components/usage_prediction/const.py b/homeassistant/components/usage_prediction/const.py index 65aeb1773fea..665e465090d6 100644 --- a/homeassistant/components/usage_prediction/const.py +++ b/homeassistant/components/usage_prediction/const.py @@ -8,6 +8,8 @@ from .models import EntityUsageDataCache, EntityUsagePredictions DOMAIN = "usage_prediction" +DEFAULT_LIMIT = 8 + DATA_CACHE: HassKey[ dict[str, asyncio.Task[EntityUsagePredictions] | EntityUsageDataCache] ] = HassKey("usage_prediction") diff --git a/tests/components/usage_prediction/test_common_control.py b/tests/components/usage_prediction/test_common_control.py index bb6b1a41bf39..f408de5c72ba 100644 --- a/tests/components/usage_prediction/test_common_control.py +++ b/tests/components/usage_prediction/test_common_control.py @@ -11,6 +11,7 @@ from homeassistant.components.usage_prediction.common_control import ( async_predict_common_control, time_category, ) +from homeassistant.components.usage_prediction.const import DEFAULT_LIMIT from homeassistant.components.usage_prediction.models import EntityUsagePredictions from homeassistant.const import EVENT_CALL_SERVICE from homeassistant.core import Context, HomeAssistant @@ -292,34 +293,18 @@ async def test_old_events_excluded(hass: HomeAssistant) -> None: @pytest.mark.usefixtures("recorder_mock") -async def test_entities_limit(hass: HomeAssistant) -> None: - """Test that only top entities are returned per time category.""" +async def test_more_than_default_limit_predicted(hass: HomeAssistant) -> None: + """Test more entities are predicted than a client gets by default.""" user_id = str(uuid.uuid4()) + entity_ids = [f"light.light_{index}" for index in range(DEFAULT_LIMIT + 2)] - hass.states.async_set("light.most_used", "off") - hass.states.async_set("light.second", "off") - hass.states.async_set("light.third", "off") - hass.states.async_set("light.fourth", "off") - hass.states.async_set("light.fifth", "off") - hass.states.async_set("light.sixth", "off") - hass.states.async_set("light.seventh", "off") + for entity_id in entity_ids: + hass.states.async_set(entity_id, "off") - # Create more than 5 different entities in morning with freeze_time("2023-07-01 08:00:00"): - # Create entities with different frequencies - entities_with_counts = [ - ("light.most_used", 10), - ("light.second", 8), - ("light.third", 6), - ("light.fourth", 4), - ("light.fifth", 2), - ("light.sixth", 1), - ("light.seventh", 1), - ] - - for entity_id, count in entities_with_counts: + # Distinct counts so the expected order is deterministic + for count, entity_id in enumerate(reversed(entity_ids), start=1): for _ in range(count): - # Use different context for each call hass.bus.async_fire( EVENT_CALL_SERVICE, { @@ -333,26 +318,11 @@ async def test_entities_limit(hass: HomeAssistant) -> None: await async_wait_recording_done(hass) - with ( - freeze_time("2023-07-02 10:00:00"), - patch( - "homeassistant.components.usage_prediction.common_control.RESULTS_TO_INCLUDE", - 5, - ), - ): # Next day, so events are recent + with freeze_time("2023-07-02 10:00:00"): # Next day, so events are recent results = await async_predict_common_control(hass, user_id) - # Should be the top 5 most used (08:00 UTC = 00:00 local = night) - assert results.night == [ - "light.most_used", - "light.second", - "light.third", - "light.fourth", - "light.fifth", - ] - assert results.morning == [] - assert results.afternoon == [] - assert results.evening == [] + # 08:00 UTC = 00:00 local = night + assert results == EntityUsagePredictions(night=entity_ids) @pytest.mark.usefixtures("recorder_mock") diff --git a/tests/components/usage_prediction/test_websocket.py b/tests/components/usage_prediction/test_websocket.py index 198c50553a9f..e7d32794f6c6 100644 --- a/tests/components/usage_prediction/test_websocket.py +++ b/tests/components/usage_prediction/test_websocket.py @@ -3,12 +3,14 @@ from collections.abc import Generator from copy import deepcopy from datetime import datetime, timedelta +from typing import Any from unittest.mock import Mock, patch from freezegun import freeze_time import pytest from homeassistant.components.usage_prediction import DOMAIN +from homeassistant.components.usage_prediction.const import DEFAULT_LIMIT from homeassistant.components.usage_prediction.models import EntityUsagePredictions from homeassistant.core import HomeAssistant from homeassistant.setup import async_setup_component @@ -17,8 +19,11 @@ from homeassistant.util import dt as dt_util from tests.common import MockUser from tests.typing import WebSocketGenerator +# Morning in the test time zone NOW = datetime(2026, 8, 26, 15, 0, 0, tzinfo=dt_util.UTC) +MORNING_ENTITIES = [f"light.morning_{index}" for index in range(60)] + @pytest.fixture def mock_predict_common_control() -> Generator[Mock]: @@ -114,3 +119,97 @@ async def test_caching_behavior( assert msg["result"] == {"entities": ["light.kitchen", "light.bla"]} # Should now be 2 (new database call) assert mock_predict_common_control.call_count == 2 + + +@pytest.mark.usefixtures("recorder_mock") +@pytest.mark.parametrize( + ("extra_msg", "expected_entities"), + [ + pytest.param({}, MORNING_ENTITIES[:DEFAULT_LIMIT], id="default"), + pytest.param({"limit": 3}, MORNING_ENTITIES[:3], id="fewer"), + pytest.param({"limit": 100}, MORNING_ENTITIES, id="more_than_predicted"), + ], +) +async def test_common_control_limit( + hass: HomeAssistant, + hass_ws_client: WebSocketGenerator, + mock_predict_common_control: Mock, + extra_msg: dict[str, Any], + expected_entities: list[str], +) -> None: + """Test the client can ask for how many entities it wants.""" + mock_predict_common_control.return_value = EntityUsagePredictions( + morning=MORNING_ENTITIES + ) + assert await async_setup_component(hass, DOMAIN, {}) + + client = await hass_ws_client(hass) + + with freeze_time(NOW): + await client.send_json( + {"id": 1, "type": "usage_prediction/common_control"} | extra_msg + ) + msg = await client.receive_json() + + assert msg["success"] is True + assert msg["result"] == {"entities": expected_entities} + + +@pytest.mark.usefixtures("recorder_mock") +@pytest.mark.parametrize( + "limit", + [ + pytest.param(0, id="below_minimum"), + pytest.param(-1, id="negative"), + pytest.param("3", id="string"), + pytest.param(3.5, id="float"), + ], +) +async def test_common_control_invalid_limit( + hass: HomeAssistant, + hass_ws_client: WebSocketGenerator, + mock_predict_common_control: Mock, + limit: float | str, +) -> None: + """Test an invalid limit is rejected without predicting.""" + assert await async_setup_component(hass, DOMAIN, {}) + + client = await hass_ws_client(hass) + + await client.send_json( + {"id": 1, "type": "usage_prediction/common_control", "limit": limit} + ) + msg = await client.receive_json() + + assert msg["success"] is False + assert msg["error"]["code"] == "invalid_format" + assert mock_predict_common_control.call_count == 0 + + +@pytest.mark.usefixtures("recorder_mock") +async def test_common_control_limit_served_from_cache( + hass: HomeAssistant, + hass_ws_client: WebSocketGenerator, + mock_predict_common_control: Mock, +) -> None: + """Test a later call with a different limit is served from the cache.""" + mock_predict_common_control.return_value = EntityUsagePredictions( + morning=MORNING_ENTITIES + ) + assert await async_setup_component(hass, DOMAIN, {}) + + client = await hass_ws_client(hass) + + with freeze_time(NOW): + await client.send_json( + {"id": 1, "type": "usage_prediction/common_control", "limit": 3} + ) + first = await client.receive_json() + await client.send_json( + {"id": 2, "type": "usage_prediction/common_control", "limit": 20} + ) + second = await client.receive_json() + + assert first["result"] == {"entities": MORNING_ENTITIES[:3]} + assert second["result"] == {"entities": MORNING_ENTITIES[:20]} + assert mock_predict_common_control.call_count == 1