From 745fe40dc61c04bdffbfde2bcb16d79b265e21cc Mon Sep 17 00:00:00 2001 From: epenet <6771947+epenet@users.noreply.github.com> Date: Tue, 29 Sep 2026 14:10:09 +0200 Subject: [PATCH] Extract calendar validators and dict factories into helper module (#183433) Co-authored-by: Claude Opus 5 --- homeassistant/components/calendar/__init__.py | 253 +++--------------- homeassistant/components/calendar/helper.py | 206 ++++++++++++++ 2 files changed, 244 insertions(+), 215 deletions(-) create mode 100644 homeassistant/components/calendar/helper.py diff --git a/homeassistant/components/calendar/__init__.py b/homeassistant/components/calendar/__init__.py index 0fabb010dc50..0d1f4f0d2465 100644 --- a/homeassistant/components/calendar/__init__.py +++ b/homeassistant/components/calendar/__init__.py @@ -4,13 +4,11 @@ from collections.abc import Callable, Iterable import dataclasses import datetime from http import HTTPStatus -from itertools import groupby import logging import re from typing import Any, Final, cast, final, override from aiohttp import web -from dateutil.rrule import rrulestr import probatio from homeassistant.auth.models import User @@ -66,11 +64,26 @@ from .const import ( EVENT_TIME_FIELDS, EVENT_TYPES, EVENT_UID, - LIST_EVENT_FIELDS, CalendarEntityFeature, CalendarEntityStateAttribute, CalendarEventStatus, ) +from .helper import ( + MIN_EVENT_DURATION, + MIN_NEW_EVENT_DURATION, + api_event_dict_factory, + as_local_timezone, + empty_as_none, + event_dict_factory, + get_datetime_local, + has_consistent_timezone, + has_min_duration, + has_positive_interval, + has_same_type, + has_timezone, + list_events_dict_factory, + validate_rrule, +) # mypy: disallow-any-generics @@ -82,145 +95,6 @@ PLATFORM_SCHEMA_BASE = cv.PLATFORM_SCHEMA_BASE SCAN_INTERVAL = datetime.timedelta(seconds=60) EVENT_LISTENER_DEBOUNCE_COOLDOWN = 1.0 # seconds -# Don't support rrules more often than daily -VALID_FREQS = {"DAILY", "WEEKLY", "MONTHLY", "YEARLY"} - -# Ensure events created in Home Assistant have a positive duration -MIN_NEW_EVENT_DURATION = datetime.timedelta(seconds=1) - -# Events must have a non-negative duration e.g. Google Calendar can create zero -# duration events in the UI. -MIN_EVENT_DURATION = datetime.timedelta(seconds=0) - - -def _has_timezone(*keys: Any) -> Callable[[dict[str, Any]], dict[str, Any]]: - """Assert that all datetime values have a timezone.""" - - def validate(obj: dict[str, Any]) -> dict[str, Any]: - """Validate that all datetime values have a timezone.""" - for k in keys: - if ( - (value := obj.get(k)) - and isinstance(value, datetime.datetime) - and value.tzinfo is None - ): - raise probatio.Invalid("Expected all values to have a timezone") - return obj - - return validate - - -def _has_consistent_timezone(*keys: Any) -> Callable[[dict[str, Any]], dict[str, Any]]: - """Verify that all datetime values have a consistent timezone.""" - - def validate(obj: dict[str, Any]) -> dict[str, Any]: - """Test that all keys that are datetime values have the same timezone.""" - tzinfos = [] - for key in keys: - if not (value := obj.get(key)) or not isinstance(value, datetime.datetime): - return obj - tzinfos.append(value.tzinfo) - uniq_values = groupby(tzinfos) - if len(list(uniq_values)) > 1: - raise probatio.Invalid("Expected all values to have the same timezone") - return obj - - return validate - - -def _as_local_timezone(*keys: Any) -> Callable[[dict[str, Any]], dict[str, Any]]: - """Convert all datetime values to the local timezone.""" - - def validate(obj: dict[str, Any]) -> dict[str, Any]: - """Convert all keys that are datetime values to local timezone.""" - for k in keys: - if (value := obj.get(k)) and isinstance(value, datetime.datetime): - obj[k] = dt_util.as_local(value) - return obj - - return validate - - -def _has_min_duration( - start_key: str, end_key: str, min_duration: datetime.timedelta -) -> Callable[[dict[str, Any]], dict[str, Any]]: - """Verify that the time span between start and end has a minimum duration.""" - - def validate(obj: dict[str, Any]) -> dict[str, Any]: - if (start := obj.get(start_key)) and (end := obj.get(end_key)): - duration = end - start - if duration < min_duration: - raise probatio.Invalid( - "Expected minimum event duration" - f" of {min_duration} ({start}, {end})" - ) - return obj - - return validate - - -def _has_positive_interval( - start_key: str, end_key: str, duration_key: str -) -> Callable[[dict[str, Any]], dict[str, Any]]: - """Verify that the time span between start and end is greater than zero.""" - - def validate(obj: dict[str, Any]) -> dict[str, Any]: - if (duration := obj.get(duration_key)) is not None: - if duration <= datetime.timedelta(seconds=0): - raise probatio.Invalid(f"Expected positive duration ({duration})") - return obj - - if (start := obj.get(start_key)) and (end := obj.get(end_key)): - if start >= end: - raise probatio.Invalid( - f"Expected end time to be after start time ({start}, {end})" - ) - return obj - - return validate - - -def _has_same_type(*keys: Any) -> Callable[[dict[str, Any]], dict[str, Any]]: - """Verify that all values are of the same type.""" - - def validate(obj: dict[str, Any]) -> dict[str, Any]: - """Test that all keys in the dict have values of the same type.""" - uniq_values = groupby(type(obj[k]) for k in keys) - if len(list(uniq_values)) > 1: - raise probatio.Invalid(f"Expected all values to be the same type: {keys}") - return obj - - return validate - - -def _validate_rrule(value: Any) -> str: - """Validate a recurrence rule string.""" - if value is None: - raise probatio.Invalid("rrule value is None") - - if not isinstance(value, str): - raise probatio.Invalid("rrule value expected a string") - - try: - rrulestr(value) - except ValueError as err: - raise probatio.Invalid(f"Invalid rrule '{value}': {err}") from err - - # Example format: FREQ=DAILY;UNTIL=... - rule_parts = dict(s.split("=", 1) for s in value.split(";")) - if not (freq := rule_parts.get("FREQ")): - raise probatio.Invalid("rrule did not contain FREQ") - - if freq not in VALID_FREQS: - raise probatio.Invalid(f"Invalid frequency for rule: {value}") - - return str(value) - - -def _empty_as_none(value: str | None) -> str | None: - """Convert any empty string values to None.""" - return value or None - CREATE_EVENT_SERVICE = "create_event" CREATE_EVENT_SCHEMA = probatio.All( @@ -255,10 +129,10 @@ CREATE_EVENT_SCHEMA = probatio.All( ), }, ), - _has_consistent_timezone(EVENT_START_DATETIME, EVENT_END_DATETIME), - _as_local_timezone(EVENT_START_DATETIME, EVENT_END_DATETIME), - _has_min_duration(EVENT_START_DATE, EVENT_END_DATE, MIN_NEW_EVENT_DURATION), - _has_min_duration(EVENT_START_DATETIME, EVENT_END_DATETIME, MIN_NEW_EVENT_DURATION), + has_consistent_timezone(EVENT_START_DATETIME, EVENT_END_DATETIME), + as_local_timezone(EVENT_START_DATETIME, EVENT_END_DATETIME), + has_min_duration(EVENT_START_DATE, EVENT_END_DATE, MIN_NEW_EVENT_DURATION), + has_min_duration(EVENT_START_DATETIME, EVENT_END_DATETIME, MIN_NEW_EVENT_DURATION), ) WEBSOCKET_EVENT_SCHEMA = probatio.Schema( @@ -269,12 +143,12 @@ WEBSOCKET_EVENT_SCHEMA = probatio.Schema( probatio.Required(EVENT_SUMMARY): cv.string, probatio.Optional(EVENT_DESCRIPTION): cv.string, probatio.Optional(EVENT_LOCATION): cv.string, - probatio.Optional(EVENT_RRULE): _validate_rrule, + probatio.Optional(EVENT_RRULE): validate_rrule, }, - _has_same_type(EVENT_START, EVENT_END), - _has_consistent_timezone(EVENT_START, EVENT_END), - _as_local_timezone(EVENT_START, EVENT_END), - _has_min_duration(EVENT_START, EVENT_END, MIN_NEW_EVENT_DURATION), + has_same_type(EVENT_START, EVENT_END), + has_consistent_timezone(EVENT_START, EVENT_END), + as_local_timezone(EVENT_START, EVENT_END), + has_min_duration(EVENT_START, EVENT_END, MIN_NEW_EVENT_DURATION), ) ) @@ -285,12 +159,12 @@ CALENDAR_EVENT_SCHEMA = probatio.Schema( probatio.Required("start"): probatio.Any(cv.date, cv.datetime), probatio.Required("end"): probatio.Any(cv.date, cv.datetime), probatio.Required(EVENT_SUMMARY): cv.string, - probatio.Optional(EVENT_RRULE): _validate_rrule, + probatio.Optional(EVENT_RRULE): validate_rrule, }, - _has_same_type("start", "end"), - _has_timezone("start", "end"), - _as_local_timezone("start", "end"), - _has_min_duration("start", "end", MIN_EVENT_DURATION), + has_same_type("start", "end"), + has_timezone("start", "end"), + as_local_timezone("start", "end"), + has_min_duration("start", "end", MIN_EVENT_DURATION), ), extra=probatio.ALLOW_EXTRA, ) @@ -308,7 +182,7 @@ SERVICE_GET_EVENTS_SCHEMA: Final = probatio.All( ), } ), - _has_positive_interval(EVENT_START_DATETIME, EVENT_END_DATETIME, EVENT_DURATION), + has_positive_interval(EVENT_START_DATETIME, EVENT_END_DATETIME, EVENT_DURATION), ) @@ -385,12 +259,12 @@ class CalendarEvent: @property def start_datetime_local(self) -> datetime.datetime: """Return event start time as a local datetime.""" - return _get_datetime_local(self.start) + return get_datetime_local(self.start) @property def end_datetime_local(self) -> datetime.datetime: """Return event end time as a local datetime.""" - return _get_datetime_local(self.end) + return get_datetime_local(self.end) @property def all_day(self) -> bool: @@ -400,7 +274,7 @@ class CalendarEvent: def as_dict(self) -> dict[str, Any]: """Return a dict representation of the event.""" return { - **dataclasses.asdict(self, dict_factory=_event_dict_factory), + **dataclasses.asdict(self, dict_factory=event_dict_factory), "all_day": self.all_day, } @@ -428,57 +302,6 @@ class CalendarEvent: self.end = self.start + datetime.timedelta(days=1) -def _event_dict_factory(obj: Iterable[tuple[str, Any]]) -> dict[str, str]: - """Convert CalendarEvent dataclass items to dictionary of attributes.""" - result: dict[str, str] = {} - for name, value in obj: - if isinstance(value, (datetime.datetime, datetime.date)): - result[name] = value.isoformat() - elif value is not None: - result[name] = str(value) - return result - - -def _api_event_dict_factory(obj: Iterable[tuple[str, Any]]) -> dict[str, Any]: - """Convert CalendarEvent dataclass items to the API format.""" - result: dict[str, Any] = {} - for name, value in obj: - if isinstance(value, datetime.datetime): - result[name] = {"dateTime": dt_util.as_local(value).isoformat()} - elif isinstance(value, datetime.date): - result[name] = {"date": value.isoformat()} - else: - result[name] = value - return result - - -def _list_events_dict_factory( - obj: Iterable[tuple[str, Any]], -) -> dict[str, JsonValueType]: - """Convert CalendarEvent dataclass items to dictionary of attributes.""" - return { - name: value - for name, value in _event_dict_factory(obj).items() - if name in LIST_EVENT_FIELDS and value is not None - } - - -def _get_datetime_local( - dt_or_d: datetime.datetime | datetime.date, -) -> datetime.datetime: - """Convert a calendar event date/datetime to a datetime if needed.""" - if isinstance(dt_or_d, datetime.datetime): - return dt_util.as_local(dt_or_d) - return dt_util.start_of_local_day(dt_or_d) - - -def _get_api_date(dt_or_d: datetime.datetime | datetime.date) -> dict[str, str]: - """Convert a calendar event date/datetime to a datetime if needed.""" - if isinstance(dt_or_d, datetime.datetime): - return {"dateTime": dt_util.as_local(dt_or_d).isoformat()} - return {"date": dt_or_d.isoformat()} - - def extract_offset(summary: str, offset_prefix: str) -> tuple[str, datetime.timedelta]: """Extract the offset from the event summary. @@ -838,7 +661,7 @@ class CalendarEventView(http.HomeAssistantView): return self.json( [ - dataclasses.asdict(event, dict_factory=_api_event_dict_factory) + dataclasses.asdict(event, dict_factory=api_event_dict_factory) for event in calendar_event_list ] ) @@ -915,7 +738,7 @@ async def handle_calendar_event_create( probatio.Required("entity_id"): cv.entity_id, probatio.Required(EVENT_UID): cv.string, probatio.Optional(EVENT_RECURRENCE_ID): probatio.Any( - probatio.All(cv.string, _empty_as_none), None + probatio.All(cv.string, empty_as_none), None ), probatio.Optional(EVENT_RECURRENCE_RANGE): cv.string, } @@ -962,7 +785,7 @@ async def handle_calendar_event_delete( probatio.Required("entity_id"): cv.entity_id, probatio.Required(EVENT_UID): cv.string, probatio.Optional(EVENT_RECURRENCE_ID): probatio.Any( - probatio.All(cv.string, _empty_as_none), None + probatio.All(cv.string, empty_as_none), None ), probatio.Optional(EVENT_RECURRENCE_RANGE): cv.string, probatio.Required(CONF_EVENT): WEBSOCKET_EVENT_SCHEMA, @@ -1125,7 +948,7 @@ async def async_get_events_service( ) return { "events": [ - dataclasses.asdict(event, dict_factory=_list_events_dict_factory) + dataclasses.asdict(event, dict_factory=list_events_dict_factory) for event in calendar_event_list ] } diff --git a/homeassistant/components/calendar/helper.py b/homeassistant/components/calendar/helper.py new file mode 100644 index 000000000000..4a343294c16a --- /dev/null +++ b/homeassistant/components/calendar/helper.py @@ -0,0 +1,206 @@ +"""Helper functions for the calendar integration.""" + +from collections.abc import Callable, Iterable +import datetime +from itertools import groupby +from typing import Any + +from dateutil.rrule import rrulestr +import probatio + +from homeassistant.util import dt as dt_util +from homeassistant.util.json import JsonValueType + +from .const import LIST_EVENT_FIELDS + +# Don't support rrules more often than daily +VALID_FREQS = {"DAILY", "WEEKLY", "MONTHLY", "YEARLY"} + + +# Ensure events created in Home Assistant have a positive duration +MIN_NEW_EVENT_DURATION = datetime.timedelta(seconds=1) + + +# Events must have a non-negative duration e.g. Google Calendar can create zero +# duration events in the UI. +MIN_EVENT_DURATION = datetime.timedelta(seconds=0) + + +def has_timezone(*keys: Any) -> Callable[[dict[str, Any]], dict[str, Any]]: + """Assert that all datetime values have a timezone.""" + + def validate(obj: dict[str, Any]) -> dict[str, Any]: + """Validate that all datetime values have a timezone.""" + for k in keys: + if ( + (value := obj.get(k)) + and isinstance(value, datetime.datetime) + and value.tzinfo is None + ): + raise probatio.Invalid("Expected all values to have a timezone") + return obj + + return validate + + +def has_consistent_timezone(*keys: Any) -> Callable[[dict[str, Any]], dict[str, Any]]: + """Verify that all datetime values have a consistent timezone.""" + + def validate(obj: dict[str, Any]) -> dict[str, Any]: + """Test that all keys that are datetime values have the same timezone.""" + tzinfos = [] + for key in keys: + if not (value := obj.get(key)) or not isinstance(value, datetime.datetime): + return obj + tzinfos.append(value.tzinfo) + uniq_values = groupby(tzinfos) + if len(list(uniq_values)) > 1: + raise probatio.Invalid("Expected all values to have the same timezone") + return obj + + return validate + + +def as_local_timezone(*keys: Any) -> Callable[[dict[str, Any]], dict[str, Any]]: + """Convert all datetime values to the local timezone.""" + + def validate(obj: dict[str, Any]) -> dict[str, Any]: + """Convert all keys that are datetime values to local timezone.""" + for k in keys: + if (value := obj.get(k)) and isinstance(value, datetime.datetime): + obj[k] = dt_util.as_local(value) + return obj + + return validate + + +def has_min_duration( + start_key: str, end_key: str, min_duration: datetime.timedelta +) -> Callable[[dict[str, Any]], dict[str, Any]]: + """Verify that the time span between start and end has a minimum duration.""" + + def validate(obj: dict[str, Any]) -> dict[str, Any]: + if (start := obj.get(start_key)) and (end := obj.get(end_key)): + duration = end - start + if duration < min_duration: + raise probatio.Invalid( + "Expected minimum event duration" + f" of {min_duration} ({start}, {end})" + ) + return obj + + return validate + + +def has_positive_interval( + start_key: str, end_key: str, duration_key: str +) -> Callable[[dict[str, Any]], dict[str, Any]]: + """Verify that the time span between start and end is greater than zero.""" + + def validate(obj: dict[str, Any]) -> dict[str, Any]: + if (duration := obj.get(duration_key)) is not None: + if duration <= datetime.timedelta(seconds=0): + raise probatio.Invalid(f"Expected positive duration ({duration})") + return obj + + if (start := obj.get(start_key)) and (end := obj.get(end_key)): + if start >= end: + raise probatio.Invalid( + f"Expected end time to be after start time ({start}, {end})" + ) + return obj + + return validate + + +def has_same_type(*keys: Any) -> Callable[[dict[str, Any]], dict[str, Any]]: + """Verify that all values are of the same type.""" + + def validate(obj: dict[str, Any]) -> dict[str, Any]: + """Test that all keys in the dict have values of the same type.""" + uniq_values = groupby(type(obj[k]) for k in keys) + if len(list(uniq_values)) > 1: + raise probatio.Invalid(f"Expected all values to be the same type: {keys}") + return obj + + return validate + + +def validate_rrule(value: Any) -> str: + """Validate a recurrence rule string.""" + if value is None: + raise probatio.Invalid("rrule value is None") + + if not isinstance(value, str): + raise probatio.Invalid("rrule value expected a string") + + try: + rrulestr(value) + except ValueError as err: + raise probatio.Invalid(f"Invalid rrule '{value}': {err}") from err + + # Example format: FREQ=DAILY;UNTIL=... + rule_parts = dict(s.split("=", 1) for s in value.split(";")) + if not (freq := rule_parts.get("FREQ")): + raise probatio.Invalid("rrule did not contain FREQ") + + if freq not in VALID_FREQS: + raise probatio.Invalid(f"Invalid frequency for rule: {value}") + + return str(value) + + +def empty_as_none(value: str | None) -> str | None: + """Convert any empty string values to None.""" + return value or None + + +def event_dict_factory(obj: Iterable[tuple[str, Any]]) -> dict[str, str]: + """Convert CalendarEvent dataclass items to dictionary of attributes.""" + result: dict[str, str] = {} + for name, value in obj: + if isinstance(value, (datetime.datetime, datetime.date)): + result[name] = value.isoformat() + elif value is not None: + result[name] = str(value) + return result + + +def api_event_dict_factory(obj: Iterable[tuple[str, Any]]) -> dict[str, Any]: + """Convert CalendarEvent dataclass items to the API format.""" + result: dict[str, Any] = {} + for name, value in obj: + if isinstance(value, datetime.datetime): + result[name] = {"dateTime": dt_util.as_local(value).isoformat()} + elif isinstance(value, datetime.date): + result[name] = {"date": value.isoformat()} + else: + result[name] = value + return result + + +def list_events_dict_factory( + obj: Iterable[tuple[str, Any]], +) -> dict[str, JsonValueType]: + """Convert CalendarEvent dataclass items to dictionary of attributes.""" + return { + name: value + for name, value in event_dict_factory(obj).items() + if name in LIST_EVENT_FIELDS and value is not None + } + + +def get_datetime_local( + dt_or_d: datetime.datetime | datetime.date, +) -> datetime.datetime: + """Convert a calendar event date/datetime to a datetime if needed.""" + if isinstance(dt_or_d, datetime.datetime): + return dt_util.as_local(dt_or_d) + return dt_util.start_of_local_day(dt_or_d) + + +def _get_api_date(dt_or_d: datetime.datetime | datetime.date) -> dict[str, str]: + """Convert a calendar event date/datetime to a datetime if needed.""" + if isinstance(dt_or_d, datetime.datetime): + return {"dateTime": dt_util.as_local(dt_or_d).isoformat()} + return {"date": dt_or_d.isoformat()}