1from abc import ABC, abstractmethod
2from collections import defaultdict
3from typing import Any, ClassVar, Literal
4
5from plain.admin.dates import DatetimeRangeAliases
6from plain.postgres import Model
7from plain.postgres.aggregates import Count
8from plain.postgres.functions import (
9 TruncDate,
10 TruncMonth,
11)
12
13from ..field_refs import (
14 FieldRef,
15 converge_declared_field,
16 field_lookup_path,
17 instance_field_ref,
18)
19from .base import Card
20
21
22class ChartCard(Card, ABC):
23 template_name = "admin/cards/chart.html"
24
25 def get_template_context(self) -> dict[str, Any]:
26 context = super().get_template_context()
27 context["chart_data"] = self.get_chart_data()
28 return context
29
30 @abstractmethod
31 def get_chart_data(self) -> dict: ...
32
33
34class TrendCard(ChartCard):
35 """
36 A card that renders a trend chart.
37 Primarily intended for use with models, but it can also be customized.
38 """
39
40 model: type[Model] | None = None
41 # Field references (`JobResult.created_at`, `JobResult.status`) or lookup
42 # paths. A reference has to belong to `model`.
43 datetime_field: FieldRef | None = None
44 group_field: FieldRef | None = None
45 group_labels: ClassVar[dict[str, str] | None] = None
46 # CSS color values resolved by charts.js. `var(--chart-N)` reads the
47 # admin's chart palette so charts retheme automatically (incl. dark mode).
48 default_group_colors: tuple[str, ...] = (
49 "var(--chart-1)",
50 "var(--chart-2)",
51 "var(--chart-3)",
52 "var(--chart-4)",
53 "var(--chart-5)",
54 )
55 group_colors: ClassVar[dict[str, str] | None] = None
56 aggregates: tuple[Literal["sum", "avg", "max"], ...] = ("sum",)
57 default_filter = DatetimeRangeAliases.SINCE_30_DAYS_AGO
58
59 filters = DatetimeRangeAliases
60
61 def __init_subclass__(cls, **kwargs: Any) -> None:
62 super().__init_subclass__(**kwargs)
63 # Normalize the declared field references now, so a reference to
64 # another model's field is a TypeError at class definition rather than
65 # when the card is first rendered, and so no `Field` is left sitting
66 # on the class as a live descriptor.
67 converge_declared_field(cls, "datetime_field", model=cls.model)
68 converge_declared_field(cls, "group_field", model=cls.model)
69
70 def get_datetime_field(self) -> str | None:
71 """`datetime_field` as a lookup path, checked against the card's model.
72
73 Read off the instance, so a card that sets its own field takes effect.
74 """
75 ref = instance_field_ref(self, "datetime_field")
76 if ref is None:
77 return None
78 return field_lookup_path(
79 ref,
80 model=self.model,
81 declared_as=f"{type(self).__qualname__}.datetime_field",
82 )
83
84 def get_group_field(self) -> str | None:
85 """`group_field` as a lookup path, checked against the card's model."""
86 ref = instance_field_ref(self, "group_field")
87 if ref is None:
88 return None
89 return field_lookup_path(
90 ref,
91 model=self.model,
92 declared_as=f"{type(self).__qualname__}.group_field",
93 )
94
95 def get_current_filter(self) -> str:
96 if s := super().get_current_filter():
97 return s
98 return self.default_filter.value
99
100 def get_trend_data(self) -> dict[str, int] | dict[str, dict[str, int]]:
101 """Return trend data, optionally grouped by group_field.
102
103 Without group_field: {date_str: count}
104 With group_field: {group_label: {date_str: count}}
105 """
106 datetime_field = self.get_datetime_field()
107 if self.model is None or datetime_field is None:
108 raise NotImplementedError(
109 "model and datetime_field must be set, or get_trend_data must be overridden"
110 )
111 group_field = self.get_group_field()
112
113 datetime_range = DatetimeRangeAliases.to_range(self.get_current_filter())
114 filter_kwargs = {f"{datetime_field}__range": datetime_range.as_tuple()}
115
116 if datetime_range.total_days() < 300:
117 truncator = TruncDate
118 iterator = datetime_range.iter_days
119 else:
120 truncator = TruncMonth
121 iterator = datetime_range.iter_months
122
123 value_fields = ["chart_date"]
124 if group_field:
125 value_fields.append(group_field)
126
127 rows = (
128 self.model.query.filter(**filter_kwargs)
129 .annotate(chart_date=truncator(datetime_field))
130 .values(*value_fields)
131 .annotate(chart_date_count=Count("id"))
132 )
133
134 dates = list(iterator())
135
136 if not group_field:
137 date_values: defaultdict[Any, int] = defaultdict(int)
138 for row in rows:
139 date_values[row["chart_date"]] = row["chart_date_count"]
140 return {date.strftime("%Y-%m-%d"): date_values[date] for date in dates}
141
142 groups: dict[str, defaultdict[Any, int]] = defaultdict(lambda: defaultdict(int))
143 for row in rows:
144 raw = row[group_field]
145 raw_value = "Unknown" if raw is None else str(raw)
146 groups[raw_value][row["chart_date"]] = row["chart_date_count"]
147
148 return {
149 group: {date.strftime("%Y-%m-%d"): counts[date] for date in dates}
150 for group, counts in sorted(groups.items())
151 }
152
153 def get_chart_data(self) -> dict:
154 data = self.get_trend_data()
155
156 if self.get_group_field():
157 return self._build_grouped_chart(data)
158
159 return self._build_single_chart(data)
160
161 def _build_single_chart(self, data: dict) -> dict:
162 return {
163 "type": "bar",
164 "data": {
165 "labels": list(data.keys()),
166 "datasets": [
167 {
168 "label": self.title,
169 "data": list(data.values()),
170 "backgroundColor": "var(--chart-1)",
171 "borderRadius": {"topLeft": 2, "topRight": 2},
172 "borderSkipped": False,
173 "categoryPercentage": 0.9,
174 "barPercentage": 1.0,
175 },
176 ],
177 },
178 **self._chart_options(stacked=False),
179 "plain": self._plain_meta(),
180 }
181
182 def _build_grouped_chart(self, data: dict) -> dict:
183 if not data:
184 return self._build_single_chart({})
185
186 labels = list(next(iter(data.values())).keys())
187
188 group_labels = self.group_labels or {}
189
190 datasets = []
191 for i, (raw_name, date_counts) in enumerate(data.items()):
192 display_name = group_labels.get(raw_name, raw_name)
193 if self.group_colors and raw_name in self.group_colors:
194 color = self.group_colors[raw_name]
195 else:
196 color = self.default_group_colors[i % len(self.default_group_colors)]
197 datasets.append(
198 {
199 "label": str(display_name),
200 "data": list(date_counts.values()),
201 "backgroundColor": color,
202 "categoryPercentage": 0.9,
203 "barPercentage": 1.0,
204 }
205 )
206
207 return {
208 "type": "bar",
209 "data": {
210 "labels": labels,
211 "datasets": datasets,
212 },
213 **self._chart_options(stacked=True),
214 "plain": self._plain_meta(),
215 }
216
217 def _plain_meta(self) -> dict:
218 return {
219 "aggregates": list(self.aggregates),
220 }
221
222 def _chart_options(self, *, stacked: bool) -> dict:
223 return {
224 "options": {
225 "responsive": True,
226 "maintainAspectRatio": False,
227 "animation": {
228 "duration": 600,
229 "easing": "easeOutQuart",
230 },
231 "interaction": {
232 "mode": "index",
233 "intersect": False,
234 "axis": "x",
235 },
236 "plugins": {
237 "legend": {"display": False},
238 "tooltip": {"enabled": False},
239 },
240 "scales": {
241 "x": {
242 "display": False,
243 "grid": {"display": False},
244 "stacked": stacked,
245 },
246 "y": {
247 "beginAtZero": True,
248 "display": False,
249 "stacked": stacked,
250 },
251 },
252 "layout": {
253 "padding": {"top": 4, "bottom": 0, "left": 0, "right": 0},
254 },
255 },
256 }