Coverage for polar/meter/aggregation.py: 87%
74 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 12:42 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 12:42 +0000
1from enum import StrEnum
2from typing import Annotated, Any, Literal
4from pydantic import AfterValidator, BaseModel, Discriminator, TypeAdapter
5from sqlalchemy import (
6 ColumnExpressionArgument,
7 Dialect,
8 Float,
9 TypeDecorator,
10 false,
11 func,
12 true,
13)
14from sqlalchemy.dialects.postgresql import JSONB
17class AggregationFunction(StrEnum):
18 cnt = "count" # `count` is a reserved keyword, so we use `cnt` as key
19 sum = "sum"
20 max = "max"
21 min = "min"
22 avg = "avg"
23 unique = "unique"
25 def get_sql_function(self, attr: Any) -> Any:
26 match self:
27 case AggregationFunction.cnt:
28 return func.count(attr)
29 case AggregationFunction.sum:
30 return func.sum(attr)
31 case AggregationFunction.max:
32 return func.max(attr)
33 case AggregationFunction.min:
34 return func.min(attr)
35 case AggregationFunction.avg:
36 return func.avg(attr)
37 case AggregationFunction.unique: 37 ↛ exitline 37 didn't return from function 'get_sql_function' because the pattern on line 37 always matched
38 return func.count(func.distinct(attr))
41class CountAggregation(BaseModel):
42 func: Literal[AggregationFunction.cnt] = AggregationFunction.cnt
44 def get_sql_column(self, model: type[Any]) -> Any:
45 return self.func.get_sql_function(model.id)
47 def get_sql_clause(self, model: type[Any]) -> ColumnExpressionArgument[bool]:
48 return true()
50 def is_summable(self) -> bool:
51 """
52 Whether this aggregation can be computed separately across different price groups
53 and then summed together. Count aggregations are summable.
54 """
55 return True
58def _strip_metadata_prefix(value: str) -> str:
59 prefix = "metadata."
60 return value[len(prefix) :] if value.startswith(prefix) else value
63class PropertyAggregation(BaseModel):
64 func: Literal[
65 AggregationFunction.sum,
66 AggregationFunction.max,
67 AggregationFunction.min,
68 AggregationFunction.avg,
69 ]
70 property: Annotated[str, AfterValidator(_strip_metadata_prefix)]
72 def get_sql_column(self, model: type[Any]) -> Any:
73 if self.property in model._filterable_fields: 73 ↛ 74line 73 didn't jump to line 74 because the condition on line 73 was never true
74 _, attr = model._filterable_fields[self.property]
75 attr = func.cast(attr, Float)
76 else:
77 attr = model.user_metadata[self.property].as_float()
79 return self.func.get_sql_function(attr)
81 def get_sql_clause(self, model: type[Any]) -> ColumnExpressionArgument[bool]:
82 if self.property in model._filterable_fields: 82 ↛ 83line 82 didn't jump to line 83 because the condition on line 82 was never true
83 allowed_type, _ = model._filterable_fields[self.property]
84 return true() if allowed_type is int else false()
86 return func.jsonb_typeof(model.user_metadata[self.property]) == "number"
88 def is_summable(self) -> bool:
89 """
90 Whether this aggregation can be computed separately across different groups
91 and then summed together. Only SUM is summable; MAX, MIN, AVG are not.
92 """
93 return self.func == AggregationFunction.sum
96class UniqueAggregation(BaseModel):
97 func: Literal[AggregationFunction.unique] = AggregationFunction.unique
98 property: Annotated[str, AfterValidator(_strip_metadata_prefix)]
100 def get_sql_column(self, model: type[Any]) -> Any:
101 attr = model.user_metadata[self.property]
102 return self.func.get_sql_function(attr)
104 def get_sql_clause(self, model: type[Any]) -> ColumnExpressionArgument[bool]:
105 return true()
107 def is_summable(self) -> bool:
108 """
109 Whether this aggregation can be computed separately across different groups
110 and then summed together. Unique count is not summable (same unique value
111 could appear in multiple groups).
112 """
113 return False
116_Aggregation = CountAggregation | PropertyAggregation | UniqueAggregation
117Aggregation = Annotated[_Aggregation, Discriminator("func")]
118AggregationTypeAdapter: TypeAdapter[Aggregation] = TypeAdapter(Aggregation)
121class AggregationType(TypeDecorator[Any]):
122 impl = JSONB
123 cache_ok = True
125 def process_bind_param(self, value: Any, dialect: Dialect) -> Any:
126 if isinstance(value, _Aggregation): 126 ↛ 128line 126 didn't jump to line 128 because the condition on line 126 was always true
127 return value.model_dump()
128 return value
130 def process_result_value(self, value: str | None, dialect: Dialect) -> Any:
131 if value is not None:
132 return AggregationTypeAdapter.validate_python(value)
133 return value