Skip to content

Commit fb1a6a2

Browse files
committed
Fix DISTINCT ON for PostgreSQL: resolve source_field/relations, include Meta.ordering in validation, defer backend check.
1 parent 0b7bb86 commit fb1a6a2

2 files changed

Lines changed: 140 additions & 57 deletions

File tree

tests/test_distinct.py

Lines changed: 61 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,9 +1,10 @@
11
import pytest
22

3-
from tests.testmodels import Tournament
3+
from tests.testmodels import Author, Book, DefaultOrdered, SourceFieldPk, Tournament
44
from tortoise.contrib import test
55
from tortoise.contrib.test.condition import NotIn
66
from tortoise.exceptions import OperationalError
7+
from tortoise.functions import Count
78

89
# ---------------------------------------------------------------------------
910
# Basic DISTINCT (all databases)
@@ -185,6 +186,55 @@ async def test_distinct_on_only_with_order_by(db):
185186
assert [(t.name, t.desc) for t in tournaments] == [("1", "b"), ("2", "c")]
186187

187188

189+
@test.requireCapability(dialect="postgres")
190+
@pytest.mark.asyncio
191+
async def test_distinct_on_filter_by_model(db):
192+
tournament_1 = await Tournament.create(name="1", desc="a")
193+
await Tournament.create(name="1", desc="b")
194+
tournament_2 = await Tournament.create(name="2", desc="c")
195+
tournaments = await Tournament.filter(name__in=["1", "2"]).distinct("name")
196+
assert [tournament_1, tournament_2] == tournaments
197+
198+
199+
@test.requireCapability(dialect="postgres")
200+
@pytest.mark.asyncio
201+
async def test_distinct_on_source_field(db):
202+
await SourceFieldPk.create(name="1")
203+
await SourceFieldPk.create(name="2")
204+
await SourceFieldPk.all().distinct("id")
205+
206+
207+
@test.requireCapability(dialect="postgres")
208+
@pytest.mark.asyncio
209+
async def test_distinct_on_annotate_by_model(db):
210+
await Tournament.create(name="1", desc="a")
211+
await Tournament.create(name="1", desc="b")
212+
await Tournament.create(name="2", desc="c")
213+
tournaments = (
214+
await Tournament.annotate(count_name=Count("name")).distinct("name").order_by("name")
215+
)
216+
assert [1, 1] == [tournaments[0].count_name, tournaments[0].count_name]
217+
218+
219+
@test.requireCapability(dialect="postgres")
220+
@pytest.mark.asyncio
221+
async def test_distinct_on_default_ordered(db):
222+
await DefaultOrdered.create(one="1", second=1)
223+
await DefaultOrdered.create(one="2", second=2)
224+
await DefaultOrdered.all().distinct("one")
225+
226+
227+
@test.requireCapability(dialect="postgres")
228+
@pytest.mark.asyncio
229+
async def test_distinct_on_by_relation(db):
230+
author_1 = await Author.create(name="1")
231+
author_2 = await Author.create(name="1")
232+
await Book.create(name="1", rating=1, subject="1", author=author_1)
233+
await Book.create(name="2", rating=2, subject="2", author=author_2)
234+
books = await Book.all().distinct("author__name")
235+
assert len(books) == 1
236+
237+
188238
# ---------------------------------------------------------------------------
189239
# DISTINCT ON validation errors
190240
# ---------------------------------------------------------------------------
@@ -202,4 +252,13 @@ async def test_distinct_on_invalid_order_by(db):
202252
@pytest.mark.asyncio
203253
async def test_distinct_on_not_supported_outside_postgres(db):
204254
with pytest.raises(OperationalError):
205-
Tournament.all().distinct("name")
255+
await Tournament.all().distinct("name")
256+
257+
258+
@test.requireCapability(dialect="postgres")
259+
@pytest.mark.asyncio
260+
async def test_distinct_on_invalid_default_ordered(db):
261+
await DefaultOrdered.create(one="1", second=1)
262+
await DefaultOrdered.create(one="2", second=2)
263+
with pytest.raises(OperationalError):
264+
await DefaultOrdered.all().distinct("second")

tortoise/queryset.py

Lines changed: 79 additions & 55 deletions
Original file line numberDiff line numberDiff line change
@@ -102,7 +102,7 @@ def _choose_db(self, for_write: bool = False) -> BaseDBAsyncClient:
102102
db = router.db_for_read(self.model)
103103
return db or self.model._meta.db
104104

105-
def _apply_db(self, db: BaseDBAsyncClient) -> None:
105+
def _apply_db(self, db: BaseDBAsyncClient | None) -> None:
106106
"""
107107
Set the database connection for this query and update the query builder dialect.
108108
@@ -291,6 +291,63 @@ def resolve_ordering(
291291

292292
self.query = self.query.orderby(field, order=ordering[1])
293293

294+
def resolve_distinct(
295+
self,
296+
distinct: bool,
297+
distinct_on: list[str],
298+
orderings: Iterable[tuple[str, str | Order]],
299+
annotations: dict[str, Term | Expression],
300+
) -> None:
301+
if not distinct:
302+
return
303+
if not orderings and self.model._meta.ordering and not annotations:
304+
orderings = self.model._meta.ordering
305+
self.query._distinct = True
306+
if distinct_on:
307+
if not isinstance(self.query, PostgreSQLQueryBuilder):
308+
raise OperationalError("DISTINCT ON is only supported by PostgreSQL")
309+
ordering_fields = [ordering[0] for ordering in orderings]
310+
len_ordering_fields = len(ordering_fields)
311+
for i, field in enumerate(distinct_on):
312+
if ordering_fields and (i >= len_ordering_fields or ordering_fields[i] != field):
313+
raise OperationalError(
314+
f"DISTINCT ON fields must match the leading ORDER BY fields. "
315+
f"Expected ORDER BY to start with {distinct_on!r}."
316+
)
317+
self.query._distinct_on = []
318+
distinct_on_by_source_field = []
319+
for field_name in distinct_on:
320+
field_object = self.model._meta.fields_map.get(field_name)
321+
part_after = field_name
322+
related_table = self.model._meta.basetable
323+
related_model: type[Model] = self.model
324+
while part_after:
325+
related_field_name, __, part_after = part_after.partition("__")
326+
if related_field_name in related_model._meta.fetch_fields:
327+
related_field = cast(
328+
RelationalField, self.model._meta.fields_map[related_field_name]
329+
)
330+
related_table = self._join_table_by_field(
331+
related_table, related_field_name, related_field
332+
)
333+
related_model = related_field.model
334+
else:
335+
field_object = related_model._meta.fields_map.get(related_field_name)
336+
337+
if not field_object:
338+
raise FieldError(
339+
f"Unknown field {related_field_name} for model {related_model.__name__}"
340+
)
341+
related_table_field = related_table[
342+
field_object.source_field or related_field_name
343+
]
344+
if func := field_object.get_for_dialect(
345+
related_model._meta.db.capabilities.dialect, "function_cast"
346+
):
347+
related_table_field = func(field_object, related_table_field)
348+
distinct_on_by_source_field.append(related_table_field)
349+
self.query.distinct_on(*distinct_on_by_source_field)
350+
294351
def _resolve_annotate(self, fields_for_select: Collection[str] | None = None) -> bool:
295352
if not self._annotations:
296353
return False
@@ -405,7 +462,7 @@ def _clone(self) -> QuerySet[MODEL]:
405462
queryset._prefetch_queries = copy(self._prefetch_queries)
406463
queryset._single = self._single
407464
queryset._raise_does_not_exist = self._raise_does_not_exist
408-
queryset._db = self._db
465+
queryset._apply_db(self._db)
409466
queryset._limit = self._limit
410467
queryset._offset = self._offset
411468
queryset._fields_for_select = self._fields_for_select
@@ -627,16 +684,10 @@ def distinct(self, *args: str) -> QuerySet[MODEL]:
627684
628685
:param args: Field names for ``DISTINCT ON`` (PostgreSQL only). Omit for plain
629686
``DISTINCT``.
630-
:raises OperationalError: If field arguments are given on a non-PostgreSQL database,
631-
or if ``ORDER BY`` is specified but does not start with the ``DISTINCT ON`` fields.
632687
"""
633688
queryset = self._clone()
634689
queryset._distinct = True
635-
if args:
636-
if isinstance(self.query, PostgreSQLQueryBuilder):
637-
queryset._distinct_on = list(args)
638-
else:
639-
raise OperationalError("DISTINCT ON is only supported by PostgreSQL")
690+
queryset._distinct_on = list(args)
640691
return queryset
641692

642693
def union(self, *other_qs: QuerySet[Model], all: bool = False) -> UnionQuery[MODEL]:
@@ -1314,25 +1365,16 @@ def _make_query(self) -> None:
13141365
self._fields_for_select,
13151366
)
13161367
self.resolve_filters()
1368+
self.resolve_distinct(
1369+
self._distinct,
1370+
self._distinct_on,
1371+
self._orderings,
1372+
self._annotations,
1373+
)
13171374
if self._limit is not None:
13181375
self.query._limit = self.query._wrapper_cls(self._limit)
13191376
if self._offset is not None:
13201377
self.query._offset = self.query._wrapper_cls(self._offset)
1321-
if self._distinct:
1322-
self.query._distinct = True
1323-
if isinstance(self.query, PostgreSQLQueryBuilder) and self._distinct_on:
1324-
ordering_fields = [ordering[0] for ordering in self._orderings]
1325-
len_ordering_fields = len(ordering_fields)
1326-
for i, field in enumerate(self._distinct_on):
1327-
if ordering_fields and (
1328-
i >= len_ordering_fields or ordering_fields[i] != field
1329-
):
1330-
raise OperationalError(
1331-
f"DISTINCT ON fields must match the leading ORDER BY fields. "
1332-
f"Expected ORDER BY to start with {self._distinct_on!r}."
1333-
)
1334-
self.query._distinct_on = []
1335-
self.query.distinct_on(*self._distinct_on)
13361378
if self._select_for_update:
13371379
self.query = self.query.for_update(
13381380
self._select_for_update_nowait,
@@ -1828,25 +1870,16 @@ def _make_query(self) -> None:
18281870
fields_for_select=self._fields_for_select_list,
18291871
)
18301872
self.resolve_filters(self._fields_to_select_sql)
1873+
self.resolve_distinct(
1874+
self._distinct,
1875+
self._distinct_on,
1876+
self._orderings,
1877+
self._annotations,
1878+
)
18311879
if self._limit:
18321880
self.query._limit = self.query._wrapper_cls(self._limit)
18331881
if self._offset:
18341882
self.query._offset = self.query._wrapper_cls(self._offset)
1835-
if self._distinct:
1836-
self.query._distinct = True
1837-
if isinstance(self.query, PostgreSQLQueryBuilder) and self._distinct_on:
1838-
ordering_fields = [ordering[0] for ordering in self._orderings]
1839-
len_ordering_fields = len(ordering_fields)
1840-
for i, field in enumerate(self._distinct_on):
1841-
if ordering_fields and (
1842-
i >= len_ordering_fields or ordering_fields[i] != field
1843-
):
1844-
raise OperationalError(
1845-
f"DISTINCT ON fields must match the leading ORDER BY fields. "
1846-
f"Expected ORDER BY to start with {self._distinct_on!r}."
1847-
)
1848-
self.query._distinct_on = []
1849-
self.query.distinct_on(*self._distinct_on)
18501883
if self._group_bys:
18511884
self.query._groupbys = self._resolve_group_bys(*self._group_bys)
18521885

@@ -1966,6 +1999,12 @@ def _make_query(self) -> None:
19661999
fields_for_select=self._fields_for_select.keys(),
19672000
)
19682001
self.resolve_filters()
2002+
self.resolve_distinct(
2003+
self._distinct,
2004+
self._distinct_on,
2005+
self._orderings,
2006+
self._annotations,
2007+
)
19692008

19702009
# remove annotations that are not in fields_for_select
19712010
self.query._selects = [
@@ -1976,21 +2015,6 @@ def _make_query(self) -> None:
19762015
self.query._limit = self.query._wrapper_cls(self._limit)
19772016
if self._offset:
19782017
self.query._offset = self.query._wrapper_cls(self._offset)
1979-
if self._distinct:
1980-
self.query._distinct = True
1981-
if isinstance(self.query, PostgreSQLQueryBuilder) and self._distinct_on:
1982-
ordering_fields = [ordering[0] for ordering in self._orderings]
1983-
len_ordering_fields = len(ordering_fields)
1984-
for i, field in enumerate(self._distinct_on):
1985-
if ordering_fields and (
1986-
i >= len_ordering_fields or ordering_fields[i] != field
1987-
):
1988-
raise OperationalError(
1989-
f"DISTINCT ON fields must match the leading ORDER BY fields. "
1990-
f"Expected ORDER BY to start with {self._distinct_on!r}."
1991-
)
1992-
self.query._distinct_on = []
1993-
self.query.distinct_on(*self._distinct_on)
19942018
if self._group_bys:
19952019
self.query._groupbys = self._resolve_group_bys(*self._group_bys)
19962020

@@ -2521,7 +2545,7 @@ def _clone(self) -> UnionQuery[MODEL]:
25212545
union._models = self._models
25222546
union._union_query = None
25232547
union._selects = self._selects
2524-
union._db = self._db
2548+
union._apply_db(self._db)
25252549
union._qs = self._qs
25262550
union._all = self._all
25272551
union._orderings = self._orderings

0 commit comments

Comments
 (0)