Skip to content

Commit b98075d

Browse files
JaeHyuck Sajacobtylerwalls
authored andcommitted
Refs #36822 -- Hoisted bulk_batch_size() implementations to base backend.
1 parent 07a1640 commit b98075d

7 files changed

Lines changed: 72 additions & 95 deletions

File tree

django/db/backends/base/operations.py

Lines changed: 13 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -2,12 +2,14 @@
22
import decimal
33
import json
44
from importlib import import_module
5+
from itertools import chain
56

67
import sqlparse
78

89
from django.conf import settings
910
from django.db import NotSupportedError, transaction
1011
from django.db.models.expressions import Col
12+
from django.db.models.fields.composite import CompositePrimaryKey
1113
from django.utils import timezone
1214
from django.utils.duration import duration_microseconds
1315
from django.utils.encoding import force_str
@@ -78,7 +80,17 @@ def bulk_batch_size(self, fields, objs):
7880
are the fields going to be inserted in the batch, the objs contains
7981
all the objects to be inserted.
8082
"""
81-
return len(objs)
83+
if self.connection.features.max_query_params is None or not fields:
84+
return len(objs)
85+
86+
return self.connection.features.max_query_params // len(
87+
list(
88+
chain.from_iterable(
89+
field.fields if isinstance(field, CompositePrimaryKey) else [field]
90+
for field in fields
91+
)
92+
)
93+
)
8294

8395
def format_for_duration_arithmetic(self, sql):
8496
raise NotImplementedError(

django/db/backends/oracle/operations.py

Lines changed: 0 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -1,15 +1,13 @@
11
import datetime
22
import uuid
33
from functools import lru_cache
4-
from itertools import chain
54

65
from django.conf import settings
76
from django.db import NotSupportedError
87
from django.db.backends.base.operations import BaseDatabaseOperations
98
from django.db.backends.utils import split_tzname_delta, strip_quotes, truncate_name
109
from django.db.models import (
1110
AutoField,
12-
CompositePrimaryKey,
1311
Exists,
1412
ExpressionWrapper,
1513
Lookup,
@@ -707,18 +705,6 @@ def subtract_temporals(self, internal_type, lhs, rhs):
707705
)
708706
return super().subtract_temporals(internal_type, lhs, rhs)
709707

710-
def bulk_batch_size(self, fields, objs):
711-
"""Oracle restricts the number of parameters in a query."""
712-
fields = list(
713-
chain.from_iterable(
714-
field.fields if isinstance(field, CompositePrimaryKey) else [field]
715-
for field in fields
716-
)
717-
)
718-
if fields:
719-
return self.connection.features.max_query_params // len(fields)
720-
return len(objs)
721-
722708
def conditional_expression_supported_in_where_clause(self, expression):
723709
"""
724710
Oracle supports only EXISTS(...) or filters in the WHERE clause, others

django/db/backends/sqlite3/operations.py

Lines changed: 0 additions & 20 deletions
Original file line numberDiff line numberDiff line change
@@ -29,26 +29,6 @@ class DatabaseOperations(BaseDatabaseOperations):
2929
# SQLite. Use JSON_TYPE() instead.
3030
jsonfield_datatype_values = frozenset(["null", "false", "true"])
3131

32-
def bulk_batch_size(self, fields, objs):
33-
"""
34-
SQLite has a variable limit defined by SQLITE_LIMIT_VARIABLE_NUMBER
35-
(reflected in max_query_params).
36-
"""
37-
fields = list(
38-
chain.from_iterable(
39-
(
40-
field.fields
41-
if isinstance(field, models.CompositePrimaryKey)
42-
else [field]
43-
)
44-
for field in fields
45-
)
46-
)
47-
if fields:
48-
return self.connection.features.max_query_params // len(fields)
49-
else:
50-
return len(objs)
51-
5232
def check_expression_support(self, expression):
5333
bad_fields = (models.DateField, models.DateTimeField, models.TimeField)
5434
bad_aggregates = (models.Sum, models.Avg, models.Variance, models.StdDev)

tests/backends/base/test_operations.py

Lines changed: 52 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,7 @@
11
import decimal
22

33
from django.core.management.color import no_style
4-
from django.db import NotSupportedError, connection, transaction
4+
from django.db import NotSupportedError, connection, models, transaction
55
from django.db.backends.base.operations import BaseDatabaseOperations
66
from django.db.models import DurationField
77
from django.db.models.expressions import Col
@@ -11,10 +11,11 @@
1111
TransactionTestCase,
1212
override_settings,
1313
skipIfDBFeature,
14+
skipUnlessDBFeature,
1415
)
1516
from django.utils import timezone
1617

17-
from ..models import Author, Book
18+
from ..models import Author, Book, Person
1819

1920

2021
class SimpleDatabaseOperationTests(SimpleTestCase):
@@ -201,6 +202,55 @@ def test_subtract_temporals(self):
201202
with self.assertRaisesMessage(NotSupportedError, msg):
202203
self.ops.subtract_temporals(duration_field_internal_type, None, None)
203204

205+
@skipUnlessDBFeature("max_query_params")
206+
def test_bulk_batch_size_limited(self):
207+
max_query_params = connection.features.max_query_params
208+
objects = range(max_query_params + 1)
209+
first_name_field = Person._meta.get_field("first_name")
210+
last_name_field = Person._meta.get_field("last_name")
211+
composite_pk = models.CompositePrimaryKey("first_name", "last_name")
212+
composite_pk.fields = [first_name_field, last_name_field]
213+
214+
self.assertEqual(connection.ops.bulk_batch_size([], objects), len(objects))
215+
self.assertEqual(
216+
connection.ops.bulk_batch_size([first_name_field], objects),
217+
max_query_params,
218+
)
219+
self.assertEqual(
220+
connection.ops.bulk_batch_size(
221+
[first_name_field, last_name_field], objects
222+
),
223+
max_query_params // 2,
224+
)
225+
self.assertEqual(
226+
connection.ops.bulk_batch_size([composite_pk, first_name_field], objects),
227+
max_query_params // 3,
228+
)
229+
230+
@skipIfDBFeature("max_query_params")
231+
def test_bulk_batch_size_unlimited(self):
232+
objects = range(2**16 + 1)
233+
first_name_field = Person._meta.get_field("first_name")
234+
last_name_field = Person._meta.get_field("last_name")
235+
composite_pk = models.CompositePrimaryKey("first_name", "last_name")
236+
composite_pk.fields = [first_name_field, last_name_field]
237+
238+
self.assertEqual(connection.ops.bulk_batch_size([], objects), len(objects))
239+
self.assertEqual(
240+
connection.ops.bulk_batch_size([first_name_field], objects),
241+
len(objects),
242+
)
243+
self.assertEqual(
244+
connection.ops.bulk_batch_size(
245+
[first_name_field, last_name_field], objects
246+
),
247+
len(objects),
248+
)
249+
self.assertEqual(
250+
connection.ops.bulk_batch_size([composite_pk, first_name_field], objects),
251+
len(objects),
252+
)
253+
204254

205255
class SqlFlushTests(TransactionTestCase):
206256
available_apps = ["backends"]

tests/backends/oracle/test_operations.py

Lines changed: 1 addition & 26 deletions
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,7 @@
11
import unittest
22

33
from django.core.management.color import no_style
4-
from django.db import connection, models
4+
from django.db import connection
55
from django.test import TransactionTestCase
66

77
from ..models import Person, Tag
@@ -17,31 +17,6 @@ def test_sequence_name_truncation(self):
1717
)
1818
self.assertEqual(seq_name, "SCHEMA_AUTHORWITHEVENLOB0B8_SQ")
1919

20-
def test_bulk_batch_size(self):
21-
# Oracle restricts the number of parameters in a query.
22-
objects = range(2**16)
23-
self.assertEqual(connection.ops.bulk_batch_size([], objects), len(objects))
24-
# Each field is a parameter for each object.
25-
first_name_field = Person._meta.get_field("first_name")
26-
last_name_field = Person._meta.get_field("last_name")
27-
self.assertEqual(
28-
connection.ops.bulk_batch_size([first_name_field], objects),
29-
connection.features.max_query_params,
30-
)
31-
self.assertEqual(
32-
connection.ops.bulk_batch_size(
33-
[first_name_field, last_name_field],
34-
objects,
35-
),
36-
connection.features.max_query_params // 2,
37-
)
38-
composite_pk = models.CompositePrimaryKey("first_name", "last_name")
39-
composite_pk.fields = [first_name_field, last_name_field]
40-
self.assertEqual(
41-
connection.ops.bulk_batch_size([composite_pk, first_name_field], objects),
42-
connection.features.max_query_params // 3,
43-
)
44-
4520
def test_sql_flush(self):
4621
statements = connection.ops.sql_flush(
4722
no_style(),

tests/backends/sqlite/test_operations.py

Lines changed: 1 addition & 24 deletions
Original file line numberDiff line numberDiff line change
@@ -2,7 +2,7 @@
22
import unittest
33

44
from django.core.management.color import no_style
5-
from django.db import connection, models
5+
from django.db import connection
66
from django.test import TestCase
77

88
from ..models import Person, Tag
@@ -88,29 +88,6 @@ def test_sql_flush_sequences_allow_cascade(self):
8888
statements[-1],
8989
)
9090

91-
def test_bulk_batch_size(self):
92-
self.assertEqual(connection.ops.bulk_batch_size([], [Person()]), 1)
93-
first_name_field = Person._meta.get_field("first_name")
94-
last_name_field = Person._meta.get_field("last_name")
95-
self.assertEqual(
96-
connection.ops.bulk_batch_size([first_name_field], [Person()]),
97-
connection.features.max_query_params,
98-
)
99-
self.assertEqual(
100-
connection.ops.bulk_batch_size(
101-
[first_name_field, last_name_field], [Person()]
102-
),
103-
connection.features.max_query_params // 2,
104-
)
105-
composite_pk = models.CompositePrimaryKey("first_name", "last_name")
106-
composite_pk.fields = [first_name_field, last_name_field]
107-
self.assertEqual(
108-
connection.ops.bulk_batch_size(
109-
[composite_pk, first_name_field], [Person()]
110-
),
111-
connection.features.max_query_params // 3,
112-
)
113-
11491
def test_bulk_batch_size_respects_variable_limit(self):
11592
first_name_field = Person._meta.get_field("first_name")
11693
last_name_field = Person._meta.get_field("last_name")

tests/composite_pk/tests.py

Lines changed: 5 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -149,21 +149,18 @@ def test_in_bulk(self):
149149

150150
def test_in_bulk_batching(self):
151151
Comment.objects.all().delete()
152-
batching_required = connection.features.max_query_params is not None
153-
expected_queries = 2 if batching_required else 1
152+
num_objects = 10
153+
connection.features.__dict__.pop("max_query_params", None)
154154
with unittest.mock.patch.object(
155-
type(connection.features), "max_query_params", 10
155+
type(connection.features), "max_query_params", num_objects
156156
):
157-
num_requiring_batching = (
158-
connection.ops.bulk_batch_size([Comment._meta.pk], []) + 1
159-
)
160157
comments = [
161158
Comment(id=i, tenant=self.tenant, user=self.user)
162-
for i in range(1, num_requiring_batching + 1)
159+
for i in range(1, num_objects + 1)
163160
]
164161
Comment.objects.bulk_create(comments)
165162
id_list = list(Comment.objects.values_list("pk", flat=True))
166-
with self.assertNumQueries(expected_queries):
163+
with self.assertNumQueries(2):
167164
comment_dict = Comment.objects.in_bulk(id_list=id_list)
168165
self.assertQuerySetEqual(comment_dict, id_list)
169166

0 commit comments

Comments
 (0)