Skip to content

Commit 384e6be

Browse files
handle aggregation method on the whole collection (#1203)
1 parent 53123ce commit 384e6be

3 files changed

Lines changed: 194 additions & 8 deletions

File tree

beanie/odm/interfaces/aggregate.py

Lines changed: 108 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -4,6 +4,7 @@
44
from pydantic import BaseModel
55
from pymongo.asynchronous.client_session import AsyncClientSession
66

7+
from beanie.odm.fields import ExpressionField
78
from beanie.odm.queries.aggregation import AggregationQuery
89
from beanie.odm.queries.find import FindMany
910

@@ -68,3 +69,110 @@ def aggregate(
6869
ignore_cache=ignore_cache,
6970
**pymongo_kwargs,
7071
)
72+
73+
@classmethod
74+
async def sum(
75+
cls,
76+
field: Union[ExpressionField, float, int, str],
77+
session: Optional[AsyncClientSession] = None,
78+
ignore_cache: bool = False,
79+
) -> Optional[float]:
80+
"""
81+
Sum of values of the given field over the entire collection.
82+
83+
Example:
84+
85+
```python
86+
87+
class Sample(Document):
88+
price: int
89+
90+
sum_count = await Document.sum(Sample.price)
91+
92+
```
93+
94+
:param field: Union[ExpressionField, float, int, str]
95+
:param session: Optional[AsyncClientSession] - pymongo session
96+
:param ignore_cache: bool
97+
:return: float - sum. None if there are no items.
98+
"""
99+
return await cls.find_all().sum(field, session, ignore_cache)
100+
101+
@classmethod
102+
async def avg(
103+
cls,
104+
field: Union[ExpressionField, float, int, str],
105+
session: Optional[AsyncClientSession] = None,
106+
ignore_cache: bool = False,
107+
) -> Optional[float]:
108+
"""
109+
Average of values of the given field over the entire collection.
110+
111+
Example:
112+
113+
```python
114+
115+
class Sample(Document):
116+
price: int
117+
118+
avg_count = await Document.avg(Sample.price)
119+
```
120+
121+
:param field: Union[ExpressionField, float, int, str]
122+
:param session: Optional[AsyncClientSession] - pymongo session
123+
:param ignore_cache: bool
124+
:return: Optional[float] - avg. None if there are no items.
125+
"""
126+
return await cls.find_all().avg(field, session, ignore_cache)
127+
128+
@classmethod
129+
async def max(
130+
cls,
131+
field: Union[ExpressionField, str, Any],
132+
session: Optional[AsyncClientSession] = None,
133+
ignore_cache: bool = False,
134+
) -> Optional[float]:
135+
"""
136+
Max of the values of the given field over the entire collection.
137+
138+
Example:
139+
140+
```python
141+
142+
class Sample(Document):
143+
price: int
144+
145+
max_count = await Document.max(Sample.price)
146+
```
147+
148+
:param field: Union[ExpressionField, str, Any]
149+
:param session: Optional[AsyncClientSession] - pymongo session
150+
:return: float - max. None if there are no items.
151+
"""
152+
return await cls.find_all().max(field, session, ignore_cache)
153+
154+
@classmethod
155+
async def min(
156+
cls,
157+
field: Union[ExpressionField, str, Any],
158+
session: Optional[AsyncClientSession] = None,
159+
ignore_cache: bool = False,
160+
) -> Optional[float]:
161+
"""
162+
Min of the values of the given field over the entire collection.
163+
164+
Example:
165+
166+
```python
167+
168+
class Sample(Document):
169+
price: int
170+
171+
min_count = await Document.min(Sample.price)
172+
```
173+
174+
:param field: Union[ExpressionField, str, Any]
175+
:param session: Optional[AsyncClientSession] - pymongo session
176+
:return: float - min. None if there are no items.
177+
"""
178+
return await cls.find_all().min(field, session, ignore_cache)

beanie/odm/interfaces/aggregation_methods.py

Lines changed: 8 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -22,7 +22,7 @@ def aggregate(
2222

2323
async def sum(
2424
self,
25-
field: Union[str, ExpressionField],
25+
field: Union[ExpressionField, float, int, str],
2626
session: Optional[AsyncClientSession] = None,
2727
ignore_cache: bool = False,
2828
) -> Optional[float]:
@@ -41,7 +41,7 @@ class Sample(Document):
4141
4242
```
4343
44-
:param field: Union[str, ExpressionField]
44+
:param field: Union[ExpressionField, float, int, str]
4545
:param session: Optional[AsyncClientSession] - pymongo session
4646
:param ignore_cache: bool
4747
:return: float - sum. None if there are no items.
@@ -66,7 +66,7 @@ class Sample(Document):
6666

6767
async def avg(
6868
self,
69-
field,
69+
field: Union[ExpressionField, float, int, str],
7070
session: Optional[AsyncClientSession] = None,
7171
ignore_cache: bool = False,
7272
) -> Optional[float]:
@@ -84,7 +84,7 @@ class Sample(Document):
8484
avg_count = await Document.find(Sample.price <= 100).avg(Sample.count)
8585
```
8686
87-
:param field: Union[str, ExpressionField]
87+
:param field: Union[ExpressionField, float, int, str]
8888
:param session: Optional[AsyncClientSession] - pymongo session
8989
:param ignore_cache: bool
9090
:return: Optional[float] - avg. None if there are no items.
@@ -108,7 +108,7 @@ class Sample(Document):
108108

109109
async def max(
110110
self,
111-
field: Union[str, ExpressionField],
111+
field: Union[ExpressionField, str, Any],
112112
session: Optional[AsyncClientSession] = None,
113113
ignore_cache: bool = False,
114114
) -> Optional[float]:
@@ -126,7 +126,7 @@ class Sample(Document):
126126
max_count = await Document.find(Sample.price <= 100).max(Sample.count)
127127
```
128128
129-
:param field: Union[str, ExpressionField]
129+
:param field: Union[ExpressionField, str, Any]
130130
:param session: Optional[AsyncClientSession] - pymongo session
131131
:return: float - max. None if there are no items.
132132
"""
@@ -149,7 +149,7 @@ class Sample(Document):
149149

150150
async def min(
151151
self,
152-
field: Union[str, ExpressionField],
152+
field: Union[ExpressionField, str, Any],
153153
session: Optional[AsyncClientSession] = None,
154154
ignore_cache: bool = False,
155155
) -> Optional[float]:
@@ -167,7 +167,7 @@ class Sample(Document):
167167
min_count = await Document.find(Sample.price <= 100).min(Sample.count)
168168
```
169169
170-
:param field: Union[str, ExpressionField]
170+
:param field: Union[ExpressionField, str, Any]
171171
:param session: Optional[AsyncClientSession] - pymongo session
172172
:return: float - min. None if there are no items.
173173
"""

tests/odm/query/test_aggregate_methods.py

Lines changed: 78 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -93,3 +93,81 @@ async def test_min_without_docs(session):
9393
)
9494

9595
assert n is None
96+
97+
98+
async def test_all_sum(preset_documents, session):
99+
n = await Sample.sum(Sample.increment)
100+
101+
assert n == 45
102+
103+
n = await Sample.sum(Sample.increment, session=session)
104+
105+
assert n == 45
106+
107+
108+
async def test_all_sum_without_docs(session):
109+
n = await Sample.sum(Sample.increment)
110+
111+
assert n is None
112+
113+
n = await Sample.sum(Sample.increment, session=session)
114+
115+
assert n is None
116+
117+
118+
async def test_all_avg(preset_documents, session):
119+
n = await Sample.avg(Sample.increment)
120+
121+
assert n == 4.5
122+
n = await Sample.avg(Sample.increment, session=session)
123+
124+
assert n == 4.5
125+
126+
127+
async def test_all_avg_without_docs(session):
128+
n = await Sample.avg(Sample.increment)
129+
130+
assert n is None
131+
n = await Sample.avg(Sample.increment, session=session)
132+
133+
assert n is None
134+
135+
136+
async def test_all_max(preset_documents, session):
137+
n = await Sample.max(Sample.increment)
138+
139+
assert n == 9
140+
141+
n = await Sample.max(Sample.increment, session=session)
142+
143+
assert n == 9
144+
145+
146+
async def test_all_max_without_docs(session):
147+
n = await Sample.max(Sample.increment)
148+
149+
assert n is None
150+
151+
n = await Sample.max(Sample.increment, session=session)
152+
153+
assert n is None
154+
155+
156+
async def test_all_min(preset_documents, session):
157+
n = await Sample.min(Sample.increment)
158+
159+
assert n == 0
160+
161+
n = await Sample.min(Sample.increment, session=session)
162+
163+
assert n == 0
164+
165+
166+
async def test_all_min_without_docs(session):
167+
n = await Sample.min(Sample.increment)
168+
169+
assert n is None
170+
171+
n = await Sample.min(Sample.increment, session=session)
172+
173+
assert n is None

0 commit comments

Comments
 (0)