Extract calendar validators and dict factories into helper module (#183433)

Co-authored-by: Claude Opus 5 <noreply@anthropic.com>
This commit is contained in:
epenet
2026-09-29 14:10:09 +02:00
committed by GitHub
co-authored by Claude Opus 5
parent 9ce2523ba8
commit 745fe40dc6
2 changed files with 244 additions and 215 deletions
+38 -215
View File
@@ -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
]
}
+206
View File
@@ -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()}