Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
35 changes: 35 additions & 0 deletions tests/test_queryset.py
Original file line number Diff line number Diff line change
Expand Up @@ -1156,3 +1156,38 @@ async def test_union_with_annotate_raises(db):

with pytest.raises(ParamsError, match="Union queries do not support annotations"):
await qs1.union(qs2)


@pytest.mark.asyncio
async def test_delete_filter_by_related_field(db):
"""Deleting through a filter on a related field must stay valid (#283)."""
author1 = await Author.create(name="Conan Doyle")
author2 = await Author.create(name="Ursula Le Guin")
await Book.create(name="The Hound", author=author1, rating=5)
await Book.create(name="The Sign of Four", author=author1, rating=4)
await Book.create(name="A Wizard of Earthsea", author=author2, rating=3)

deleted = await Book.filter(author__name="Conan Doyle").delete()

assert deleted == 2
assert await Book.filter(author__name="Conan Doyle").count() == 0
# negative control: rows belonging to the other author are untouched
assert await Book.filter(author__name="Ursula Le Guin").count() == 1
assert await Author.all().count() == 2


@pytest.mark.asyncio
async def test_update_filter_by_related_field(db):
"""Updating through a filter on a related field must stay valid (#283)."""
author1 = await Author.create(name="Conan Doyle")
author2 = await Author.create(name="Ursula Le Guin")
book1 = await Book.create(name="The Hound", author=author1, rating=5)
await Book.create(name="A Wizard of Earthsea", author=author2, rating=3)

updated = await Book.filter(author__name="Conan Doyle").update(rating=1)

assert updated == 1
assert (await Book.get(id=book1.id)).rating == 1
# negative control: the other author's book keeps its rating
other = await Book.get(name="A Wizard of Earthsea")
assert other.rating == 3
46 changes: 43 additions & 3 deletions tortoise/queryset.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,7 @@
from pypika_tortoise.analytics import Count
from pypika_tortoise.functions import Cast
from pypika_tortoise.queries import QueryBuilder, _SetOperation
from pypika_tortoise.terms import Case, Field, Star, Term, ValueWrapper
from pypika_tortoise.terms import Case, Field, Star, Term, Tuple, ValueWrapper

from tortoise.backends.base.client import BaseDBAsyncClient, Capabilities
from tortoise.exceptions import (
Expand Down Expand Up @@ -164,6 +164,29 @@ def resolve_filters(self, fields_for_select: Collection[str] | None = None) -> N
*[self.model._meta.basetable[field] for field in self.model._meta.db_fields]
)

def _pk_in_subquery(self, table: Table) -> Term:
"""Build ``pk IN (SELECT "_t".pk FROM (SELECT pk ...) AS "_t")`` from the current
filter-resolved select query.

Filters on related fields add JOINs to the query, but most backends do not allow
JOINs in DELETE/UPDATE statements, so the target rows are selected by primary key
in a subquery instead.
"""
pk_cols = [
table[self.model._meta.fields_db_projection[name]]
for name, field in self.model._meta.fields_map.items()
if field.pk
]
inner = copy(self.query)
inner._selects = list(pk_cols)
alias = "_t"
wrapped = self.model._meta.db.query_class.from_(inner.as_(alias)).select(
*[Table(alias)[col.name] for col in pk_cols]
)
if len(pk_cols) == 1:
return pk_cols[0].isin(wrapped)
return Tuple(*pk_cols).isin(wrapped)

def _join_table_by_field(
self, table: Table, related_field_name: str, related_field: RelationalField
) -> Table:
Expand Down Expand Up @@ -1349,12 +1372,24 @@ def __init__(

def _make_query(self) -> None:
table = self.model._meta.basetable
self.query = self._db.query_class.update(table)
# Resolve filters on a SELECT-shaped query first: filters on related fields add
# JOINs, which are not valid in UPDATE statements on most backends, so the target
# rows are selected by primary key through a subquery instead (see _pk_in_subquery).
self.query = copy(self.model._meta.basequery)
if self.capabilities.support_update_limit_order_by and self._limit:
self.query._limit = self.query._wrapper_cls(self._limit)
self.resolve_ordering(self.model, table, self._orderings, self._annotations)

self.resolve_filters()
if self.query._joins:
self.query = self._db.query_class.update(table).where(self._pk_in_subquery(table))
else:
update_query = self._db.query_class.update(table)
update_query._wheres = self.query._wheres
update_query._havings = self.query._havings
update_query._orderbys = self.query._orderbys
update_query._limit = self.query._limit
self.query = update_query
for key, value in self.update_kwargs.items():
field_object = self.model._meta.fields_map.get(key)
if not field_object:
Expand Down Expand Up @@ -1427,16 +1462,21 @@ def __init__(
self._orderings = orderings

def _make_query(self) -> None:
table = self.model._meta.basetable
self.query = copy(self.model._meta.basequery)
if self.capabilities.support_update_limit_order_by and self._limit:
self.query._limit = self.query._wrapper_cls(self._limit)
self.resolve_ordering(
model=self.model,
table=self.model._meta.basetable,
table=table,
orderings=self._orderings,
annotations=self._annotations,
)
self.resolve_filters()
if self.query._joins:
self.query = self.model._meta.db.query_class.from_(table).where(
self._pk_in_subquery(table)
)
self.query._delete_from = True
return

Expand Down
Loading