v0.165.0
  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        }