Skip to content

Commit fa341d0

Browse files
committed
feat: add DISTINCT ON support for PostgreSQL.
1 parent 179ea6f commit fa341d0

3 files changed

Lines changed: 384 additions & 24 deletions

File tree

tests/test_distinct.py

Lines changed: 204 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,204 @@
1+
import pytest
2+
3+
from tests.testmodels import Tournament
4+
from tortoise.contrib import test
5+
from tortoise.exceptions import OperationalError
6+
7+
# ---------------------------------------------------------------------------
8+
# Basic DISTINCT (all databases)
9+
# ---------------------------------------------------------------------------
10+
11+
12+
@pytest.mark.asyncio
13+
async def test_distinct_no_args(db):
14+
await Tournament.create(name="1", desc="a")
15+
await Tournament.create(name="1", desc="b")
16+
tournaments = await Tournament.all().distinct()
17+
assert len(tournaments) == 2
18+
19+
20+
# ---------------------------------------------------------------------------
21+
# DISTINCT ON (PostgreSQL only)
22+
# ---------------------------------------------------------------------------
23+
24+
25+
@test.requireCapability(dialect="postgres")
26+
@pytest.mark.asyncio
27+
async def test_distinct_on_single_field(db):
28+
tournament_1 = await Tournament.create(name="1", desc="1")
29+
await Tournament.create(name="1", desc="2")
30+
await Tournament.create(name="1", desc="3")
31+
32+
tournaments = await Tournament.all().distinct("name")
33+
assert tournaments == [tournament_1]
34+
35+
36+
@test.requireCapability(dialect="postgres")
37+
@pytest.mark.asyncio
38+
async def test_distinct_on_single_field_with_order_by(db):
39+
await Tournament.create(name="1", desc="1")
40+
await Tournament.create(name="1", desc="2")
41+
tournament_3 = await Tournament.create(name="1", desc="3")
42+
43+
tournaments = await Tournament.all().distinct("name").order_by("name", "-desc")
44+
assert tournaments == [tournament_3]
45+
46+
47+
@test.requireCapability(dialect="postgres")
48+
@pytest.mark.asyncio
49+
async def test_distinct_on_multiple_fields(db):
50+
tournament_1 = await Tournament.create(name="1", desc="a")
51+
await Tournament.create(name="1", desc="a")
52+
tournament_3 = await Tournament.create(name="2", desc="b")
53+
54+
tournaments = await Tournament.all().distinct("name", "desc").order_by("name", "desc")
55+
assert tournaments == [tournament_1, tournament_3]
56+
57+
58+
@test.requireCapability(dialect="postgres")
59+
@pytest.mark.asyncio
60+
async def test_distinct_on_values_list_single_field(db):
61+
"""values_list selects one field, same as DISTINCT ON field."""
62+
await Tournament.create(name="1", desc="a")
63+
await Tournament.create(name="1", desc="b")
64+
await Tournament.create(name="2", desc="c")
65+
66+
tournaments = await Tournament.all().distinct("name").values_list("name", flat=True)
67+
assert tournaments == ["1", "2"]
68+
69+
70+
@test.requireCapability(dialect="postgres")
71+
@pytest.mark.asyncio
72+
async def test_distinct_on_values_list_multiple_fields(db):
73+
await Tournament.create(name="1", desc="a")
74+
await Tournament.create(name="1", desc="b")
75+
await Tournament.create(name="2", desc="c")
76+
77+
tournaments = await Tournament.all().distinct("name").values_list("name", "desc")
78+
assert tournaments == [("1", "a"), ("2", "c")]
79+
80+
81+
@test.requireCapability(dialect="postgres")
82+
@pytest.mark.asyncio
83+
async def test_distinct_on_values_list_extra_fields(db):
84+
await Tournament.create(name="1", desc="a")
85+
await Tournament.create(name="1", desc="b")
86+
await Tournament.create(name="2", desc="c")
87+
88+
tournaments = await Tournament.all().distinct("name").values_list("desc", flat=True)
89+
assert tournaments == ["a", "c"]
90+
91+
92+
@test.requireCapability(dialect="postgres")
93+
@pytest.mark.asyncio
94+
async def test_distinct_on_values_list_extra_field_respects_order_by(db):
95+
await Tournament.create(name="1", desc="a")
96+
await Tournament.create(name="1", desc="b")
97+
await Tournament.create(name="2", desc="c")
98+
99+
tournaments = (
100+
await Tournament.all()
101+
.distinct("name")
102+
.order_by("name", "-desc")
103+
.values_list("desc", flat=True)
104+
)
105+
assert tournaments == ["b", "c"]
106+
107+
108+
@test.requireCapability(dialect="postgres")
109+
@pytest.mark.asyncio
110+
async def test_distinct_on_values_single_field(db):
111+
await Tournament.create(name="1", desc="a")
112+
await Tournament.create(name="1", desc="b")
113+
await Tournament.create(name="2", desc="c")
114+
115+
tournaments = await Tournament.all().distinct("name").values("name")
116+
assert tournaments == [{"name": "1"}, {"name": "2"}]
117+
118+
119+
@test.requireCapability(dialect="postgres")
120+
@pytest.mark.asyncio
121+
async def test_distinct_on_values_multiple_fields(db):
122+
await Tournament.create(name="1", desc="a")
123+
await Tournament.create(name="1", desc="b")
124+
await Tournament.create(name="2", desc="c")
125+
126+
tournaments = await Tournament.all().distinct("name").values("name", "desc")
127+
assert tournaments == [{"name": "1", "desc": "a"}, {"name": "2", "desc": "c"}]
128+
129+
130+
@test.requireCapability(dialect="postgres")
131+
@pytest.mark.asyncio
132+
async def test_distinct_on_values_extra_fields(db):
133+
await Tournament.create(name="1", desc="a")
134+
await Tournament.create(name="1", desc="b")
135+
await Tournament.create(name="2", desc="c")
136+
137+
tournaments = await Tournament.all().distinct("name").values("desc")
138+
assert tournaments == [{"desc": "a"}, {"desc": "c"}]
139+
140+
141+
@test.requireCapability(dialect="postgres")
142+
@pytest.mark.asyncio
143+
async def test_distinct_on_values_extra_field_respects_order_by(db):
144+
await Tournament.create(name="1", desc="a")
145+
await Tournament.create(name="1", desc="b")
146+
await Tournament.create(name="2", desc="c")
147+
148+
tournaments = await Tournament.all().distinct("name").order_by("name", "-desc").values("desc")
149+
assert tournaments == [{"desc": "b"}, {"desc": "c"}]
150+
151+
152+
@test.requireCapability(dialect="postgres")
153+
@pytest.mark.asyncio
154+
async def test_distinct_on_only_same_field(db):
155+
await Tournament.create(name="1", desc="a")
156+
await Tournament.create(name="1", desc="b")
157+
await Tournament.create(name="2", desc="c")
158+
159+
tournaments = await Tournament.all().distinct("name").only("name")
160+
assert [t.name for t in tournaments] == ["1", "2"]
161+
162+
163+
@test.requireCapability(dialect="postgres")
164+
@pytest.mark.asyncio
165+
async def test_distinct_on_only_extra_field(db):
166+
await Tournament.create(name="1", desc="a")
167+
await Tournament.create(name="1", desc="b")
168+
await Tournament.create(name="2", desc="c")
169+
170+
tournaments = await Tournament.all().distinct("name").only("name", "desc")
171+
assert [(t.name, t.desc) for t in tournaments] == [("1", "a"), ("2", "c")]
172+
173+
174+
@test.requireCapability(dialect="postgres")
175+
@pytest.mark.asyncio
176+
async def test_distinct_on_only_with_order_by(db):
177+
await Tournament.create(name="1", desc="a")
178+
await Tournament.create(name="1", desc="b")
179+
await Tournament.create(name="2", desc="c")
180+
181+
tournaments = (
182+
await Tournament.all().distinct("name").order_by("name", "-desc").only("name", "desc")
183+
)
184+
assert [(t.name, t.desc) for t in tournaments] == [("1", "b"), ("2", "c")]
185+
186+
187+
# ---------------------------------------------------------------------------
188+
# DISTINCT ON validation errors
189+
# ---------------------------------------------------------------------------
190+
191+
192+
@test.requireCapability(dialect="postgres")
193+
@pytest.mark.asyncio
194+
async def test_distinct_on_invalid_order_by(db):
195+
await Tournament.create(name="1")
196+
with pytest.raises(OperationalError):
197+
await Tournament.all().distinct("name").order_by("desc")
198+
199+
200+
@test.skipCapability(dialect="postgres")
201+
@pytest.mark.asyncio
202+
async def test_distinct_on_not_supported_outside_postgres(db):
203+
with pytest.raises(OperationalError):
204+
Tournament.all().distinct("name")

tortoise/contrib/test/__init__.py

Lines changed: 64 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -50,6 +50,7 @@ async def test_sqlite_only(db):
5050
"TortoiseContext",
5151
"tortoise_test_context",
5252
"requireCapability",
53+
"skipCapability",
5354
"truncate_all_models",
5455
"init_memory_sqlite",
5556
"SkipTest",
@@ -235,6 +236,69 @@ def skip_wrapper(*args: typing.Any, **kwargs: typing.Any) -> typing.Any:
235236
return decorator
236237

237238

239+
def skipCapability(
240+
connection_name: str = "models", **conditions: typing.Any
241+
) -> Callable[[_FT], _FT]:
242+
"""
243+
Skip a test if the specified capabilities are matched.
244+
245+
This is the inverse of :func:`requireCapability`.
246+
247+
Usage:
248+
249+
.. code-block:: python3
250+
251+
@skipCapability(dialect='postgres')
252+
@pytest.mark.asyncio
253+
async def test_skip_on_postgres(db):
254+
...
255+
256+
:param connection_name: name of the connection to retrieve capabilities from.
257+
:param conditions: capability tests — if all match, the test is skipped.
258+
"""
259+
260+
def decorator(test_item: _FT) -> _FT:
261+
if not isinstance(test_item, type):
262+
263+
def check_capabilities() -> None:
264+
db = get_connection(connection_name)
265+
if all(getattr(db.capabilities, key) == val for key, val in conditions.items()):
266+
raise SkipTest(f"Skipped because capabilities match: {conditions}")
267+
268+
if inspect.iscoroutinefunction(test_item):
269+
270+
@wraps(test_item)
271+
async def skip_wrapper(*args: typing.Any, **kwargs: typing.Any) -> typing.Any:
272+
check_capabilities()
273+
return await test_item(*args, **kwargs)
274+
275+
else:
276+
277+
@wraps(test_item)
278+
def skip_wrapper(*args: typing.Any, **kwargs: typing.Any) -> typing.Any:
279+
check_capabilities()
280+
return test_item(*args, **kwargs)
281+
282+
return cast(_FT, skip_wrapper)
283+
284+
# Assume a class is decorated
285+
funcs = {
286+
var: f
287+
for var in dir(test_item)
288+
if var.startswith("test_") and callable(f := getattr(test_item, var))
289+
}
290+
for name, func in funcs.items():
291+
setattr(
292+
test_item,
293+
name,
294+
skipCapability(connection_name=connection_name, **conditions)(func),
295+
)
296+
297+
return test_item
298+
299+
return decorator
300+
301+
238302
@typing.overload
239303
def init_memory_sqlite(models: ModulesConfigType | None = None) -> AsyncFuncDeco: ...
240304

0 commit comments

Comments
 (0)