Coverage for core / src / sensorkit / auto / mode.py: 96%

92 statements  

« prev     ^ index     » next       coverage.py v7.13.5, created at 2026-09-02 00:03 +0000

1# SPDX-License-Identifier: Apache-2.0 

2from __future__ import annotations 

3 

4from abc import ABC, abstractmethod 

5from collections.abc import Iterable, Mapping 

6from datetime import UTC, datetime, timedelta 

7from typing import Annotated, Any, Literal, override 

8 

9from intervaltree import IntervalTree 

10from loguru import logger 

11from pydantic import AfterValidator, BaseModel, Field, model_validator 

12 

13from sensorkit.common.interval import intersect_interval_trees 

14from sensorkit.common.time import TimeRangeError, parse_time_delta, parse_time_range 

15 

16DAILY_PERIOD = timedelta(days=1) 

17 

18 

19class Mode(BaseModel): 

20 """A named operating mode with scheduling criteria and an associated sensor state.""" 

21 name: str 

22 description: str = "" 

23 state: Literal["operate", "standby"] = "operate" 

24 criteria: list[AnyCriterion] = Field(default_factory=list) 

25 context: dict = Field(default_factory=dict) 

26 

27 def evaluate(self, context: Mapping[str, Any]): 

28 """Evaluate all criteria and return an IntervalTree of common intervals.""" 

29 trees: list[IntervalTree] = [] 

30 

31 for criterion in self.criteria: 

32 # Evaluate the criterion, returning a tree of datetime intervals. 

33 tree = IntervalTree.from_tuples(criterion.evaluate(context)) 

34 

35 if tree.is_empty(): 

36 # We can short-circuit if any evaluation returns no intervals. 

37 return tree 

38 

39 # Merge overlaps to make the resulting tree disjoint. 

40 tree.merge_overlaps(strict=False) 

41 trees.append(tree) 

42 

43 # Return the intersection of all trees, applying AND semantics. 

44 return intersect_interval_trees(*trees) 

45 

46 def __lt__(self, other): 

47 # Operate states take priority over standby states. 

48 return self.state != other.state and self.state == "operate" 

49 

50 

51def _ensure_unique_mode_names(modes: list[Mode]): 

52 if len(modes) != len(set(mode.name for mode in modes)): 

53 raise ValueError("Mode names must be unique") 

54 

55 return modes 

56 

57 

58ModeList = Annotated[ 

59 list[Mode], 

60 AfterValidator(_ensure_unique_mode_names) 

61] 

62 

63type IntervalList = Iterable[tuple[datetime, datetime]] 

64type AnyCriterion = TimeRangeCriterion | TaskingAvailableCriterion | AfterActivityCriterion 

65 

66 

67class Criterion(BaseModel, ABC): 

68 """Abstract scheduling criterion that produces a list of valid time intervals.""" 

69 when: str 

70 

71 @abstractmethod 

72 def evaluate(self, context: Mapping[str, Any]) -> IntervalList: 

73 """Evaluate this criterion and return valid `(start, end)` datetime interval tuples.""" 

74 

75 @model_validator(mode="before") 

76 @classmethod 

77 def _compat(cls, data: Any): 

78 if isinstance(data, dict) and "kind" in data: 

79 logger.warning(f"update {data["kind"]} criterion to specify 'when' instead of 'kind'") 

80 assert "when" not in data 

81 data["when"] = data.pop("kind") 

82 

83 return data 

84 

85 

86class TimeRangeCriterion(Criterion): 

87 """Criterion satisfied during a configured daily time range (e.g. sunset to sunrise).""" 

88 when: Literal["time_range"] = "time_range" 

89 start: str 

90 end: str 

91 

92 @override 

93 def evaluate(self, context: Mapping[str, Any]): 

94 time_ref = context.get("time_ref") or datetime.now(UTC) 

95 symbol_handlers = context.get("time_range_parsers") 

96 horizon = time_ref + DAILY_PERIOD 

97 start, end = parse_time_range( 

98 self.start, 

99 self.end, 

100 time_ref=time_ref, 

101 symbol_handlers=symbol_handlers, 

102 ) 

103 

104 yield start, end 

105 

106 if end >= horizon: 

107 return 

108 

109 # A day-periodic spec's next occurrence starts one period on from this one, which 

110 # is always enough to reach the horizon. 

111 try: 

112 next_start, next_end = parse_time_range( 

113 self.start, 

114 self.end, 

115 time_ref=start + DAILY_PERIOD, 

116 symbol_handlers=symbol_handlers, 

117 ) 

118 except TimeRangeError: 

119 return 

120 

121 if next_end == end: 

122 # Correct the edge case where the input endpoints evaluate to the same time. 

123 next_start += DAILY_PERIOD 

124 next_end += DAILY_PERIOD 

125 

126 yield next_start, next_end 

127 

128 

129class TaskingAvailableCriterion(Criterion): 

130 """Criterion satisfied when a program has active offer intervals in the current schedule.""" 

131 when: Literal["tasking_available"] = "tasking_available" 

132 from_program: str | None = None 

133 

134 @override 

135 def evaluate(self, context: Mapping[str, Any]): 

136 offers: IntervalTree 

137 

138 if offers := context.get("scheduler_previous_combined"): 

139 if self.from_program is not None: 

140 return (iv for iv in offers if self.from_program in iv.data) 

141 else: 

142 return offers 

143 else: 

144 return [] 

145 

146 

147class AfterActivityCriterion(Criterion): 

148 """Criterion satisfied for a configured duration after each completed program offer window.""" 

149 when: Literal["after_activity"] = "after_activity" 

150 duration: str 

151 

152 @override 

153 def evaluate(self, context: Mapping[str, Any]): 

154 tree = IntervalTree() 

155 duration = parse_time_delta(self.duration) 

156 

157 offers: IntervalTree 

158 

159 if offers := context.get("scheduler_offer_history"): 

160 for iv in offers: 

161 tree.addi(iv.end, iv.end + duration) 

162 elif offers := context.get("scheduler_previous_combined"): 

163 for iv in offers: 

164 tree.addi(iv.end, iv.end + duration) 

165 

166 return tree