Skip to content

Commit 3e476a1

Browse files
committed
perf(spanner): optimize query option merging and prevent in-place mutation
Short-circuit query option merging when options are unset to eliminate allocations on the hot path, make field handling generic, and merge into a fresh protobuf to avoid in-place mutation of base options.
1 parent f40cf77 commit 3e476a1

2 files changed

Lines changed: 178 additions & 26 deletions

File tree

packages/google-cloud-spanner/google/cloud/spanner_v1/_helpers.py

Lines changed: 43 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -160,6 +160,36 @@ def _try_to_coerce_bytes(bytestring):
160160
)
161161

162162

163+
def _to_query_options(options):
164+
"""Normalize dict or QueryOptions to a non-empty QueryOptions, or None.
165+
166+
:type options:
167+
:class:`~google.cloud.spanner_v1.types.ExecuteSqlRequest.QueryOptions`
168+
or :class:`dict` or None
169+
:param options: Query options to normalize.
170+
171+
:rtype:
172+
:class:`~google.cloud.spanner_v1.types.ExecuteSqlRequest.QueryOptions`
173+
or None
174+
:returns:
175+
A non-empty QueryOptions instance, or None if options is empty or None.
176+
177+
:raises TypeError:
178+
If options is not a QueryOptions, dict, or None.
179+
"""
180+
if options is None:
181+
return None
182+
if isinstance(options, dict):
183+
if not any(options.values()):
184+
return None
185+
options = ExecuteSqlRequest.QueryOptions(options)
186+
elif not isinstance(options, ExecuteSqlRequest.QueryOptions):
187+
raise TypeError(
188+
f"query_options must be a QueryOptions or dict, got {type(options).__name__}"
189+
)
190+
return options if type(options).pb(options).ByteSize() > 0 else None
191+
192+
163193
def _merge_query_options(base, merge):
164194
"""Merge higher precedence QueryOptions with current QueryOptions.
165195
@@ -182,23 +212,20 @@ def _merge_query_options(base, merge):
182212
QueryOptions object formed by merging the two given QueryOptions.
183213
If the resultant object only has empty fields, returns None.
184214
"""
185-
combined = base or ExecuteSqlRequest.QueryOptions()
186-
if isinstance(combined, dict):
187-
combined = ExecuteSqlRequest.QueryOptions(
188-
optimizer_version=combined.get("optimizer_version", ""),
189-
optimizer_statistics_package=combined.get(
190-
"optimizer_statistics_package", ""
191-
),
192-
)
193-
merge = merge or ExecuteSqlRequest.QueryOptions()
194-
if isinstance(merge, dict):
195-
merge = ExecuteSqlRequest.QueryOptions(
196-
optimizer_version=merge.get("optimizer_version", ""),
197-
optimizer_statistics_package=merge.get("optimizer_statistics_package", ""),
198-
)
199-
type(combined).pb(combined).MergeFrom(type(merge).pb(merge))
200-
if not combined.optimizer_version and not combined.optimizer_statistics_package:
215+
if base is None and merge is None:
201216
return None
217+
218+
base = _to_query_options(base)
219+
merge = _to_query_options(merge)
220+
if base is None:
221+
return merge
222+
if merge is None:
223+
return base
224+
225+
combined = ExecuteSqlRequest.QueryOptions()
226+
combined_pb = type(combined).pb(combined)
227+
combined_pb.CopyFrom(type(base).pb(base))
228+
combined_pb.MergeFrom(type(merge).pb(merge))
202229
return combined
203230

204231

packages/google-cloud-spanner/tests/unit/test__helpers.py

Lines changed: 135 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -22,7 +22,42 @@
2222
from opentelemetry.sdk.resources import Resource
2323
from opentelemetry.semconv.resource import ResourceAttributes
2424

25-
from google.cloud.spanner_v1 import TransactionOptions, _helpers
25+
from google.cloud.spanner_v1 import ExecuteSqlRequest, TransactionOptions, _helpers
26+
27+
28+
class Test_to_query_options(unittest.TestCase):
29+
def _callFUT(self, *args, **kw):
30+
from google.cloud.spanner_v1._helpers import _to_query_options
31+
32+
return _to_query_options(*args, **kw)
33+
34+
def test_none(self):
35+
self.assertIsNone(self._callFUT(None))
36+
37+
def test_empty_dict(self):
38+
self.assertIsNone(self._callFUT({}))
39+
40+
def test_dict_with_empty_values(self):
41+
self.assertIsNone(self._callFUT({"optimizer_version": ""}))
42+
43+
def test_valid_dict(self):
44+
expected = ExecuteSqlRequest.QueryOptions(optimizer_version="1")
45+
result = self._callFUT({"optimizer_version": "1"})
46+
self.assertEqual(result, expected)
47+
48+
def test_empty_proto_object(self):
49+
self.assertIsNone(self._callFUT(ExecuteSqlRequest.QueryOptions()))
50+
51+
def test_populated_proto_object(self):
52+
options = ExecuteSqlRequest.QueryOptions(optimizer_version="1")
53+
result = self._callFUT(options)
54+
self.assertEqual(result, options)
55+
56+
def test_invalid_type(self):
57+
with self.assertRaises(TypeError):
58+
self._callFUT("invalid")
59+
with self.assertRaises(TypeError):
60+
self._callFUT(123)
2661

2762

2863
class Test_merge_query_options(unittest.TestCase):
@@ -37,8 +72,6 @@ def test_base_none_and_merge_none(self):
3772
self.assertIsNone(result)
3873

3974
def test_base_dict_and_merge_none(self):
40-
from google.cloud.spanner_v1 import ExecuteSqlRequest
41-
4275
base = {
4376
"optimizer_version": "2",
4477
"optimizer_statistics_package": "auto_20191128_14_47_22UTC",
@@ -52,16 +85,12 @@ def test_base_dict_and_merge_none(self):
5285
self.assertEqual(result, expected)
5386

5487
def test_base_empty_and_merge_empty(self):
55-
from google.cloud.spanner_v1 import ExecuteSqlRequest
56-
5788
base = ExecuteSqlRequest.QueryOptions()
5889
merge = ExecuteSqlRequest.QueryOptions()
5990
result = self._callFUT(base, merge)
6091
self.assertIsNone(result)
6192

6293
def test_base_none_merge_object(self):
63-
from google.cloud.spanner_v1 import ExecuteSqlRequest
64-
6594
base = None
6695
merge = ExecuteSqlRequest.QueryOptions(
6796
optimizer_version="3",
@@ -71,29 +100,125 @@ def test_base_none_merge_object(self):
71100
self.assertEqual(result, merge)
72101

73102
def test_base_none_merge_dict(self):
74-
from google.cloud.spanner_v1 import ExecuteSqlRequest
75-
76103
base = None
77104
merge = {"optimizer_version": "3"}
78105
expected = ExecuteSqlRequest.QueryOptions(optimizer_version="3")
79106
result = self._callFUT(base, merge)
80107
self.assertEqual(result, expected)
81108

82109
def test_base_object_merge_dict(self):
83-
from google.cloud.spanner_v1 import ExecuteSqlRequest
110+
base = ExecuteSqlRequest.QueryOptions(
111+
optimizer_version="1",
112+
optimizer_statistics_package="auto_20191128_14_47_22UTC",
113+
)
114+
merge = {"optimizer_version": "3"}
115+
expected = ExecuteSqlRequest.QueryOptions(
116+
optimizer_version="3",
117+
optimizer_statistics_package="auto_20191128_14_47_22UTC",
118+
)
119+
result = self._callFUT(base, merge)
120+
self.assertEqual(result, expected)
84121

122+
def test_base_object_and_merge_none(self):
123+
base = ExecuteSqlRequest.QueryOptions(
124+
optimizer_version="2",
125+
optimizer_statistics_package="auto_20191128_14_47_22UTC",
126+
)
127+
result = self._callFUT(base, None)
128+
self.assertEqual(result, base)
129+
130+
def test_base_empty_object_and_merge_none(self):
131+
base = ExecuteSqlRequest.QueryOptions()
132+
result = self._callFUT(base, None)
133+
self.assertIsNone(result)
134+
135+
def test_base_none_merge_empty_object(self):
136+
merge = ExecuteSqlRequest.QueryOptions()
137+
result = self._callFUT(None, merge)
138+
self.assertIsNone(result)
139+
140+
def test_base_object_not_mutated_on_merge(self):
85141
base = ExecuteSqlRequest.QueryOptions(
86142
optimizer_version="1",
87143
optimizer_statistics_package="auto_20191128_14_47_22UTC",
88144
)
89145
merge = {"optimizer_version": "3"}
146+
result = self._callFUT(base, merge)
90147
expected = ExecuteSqlRequest.QueryOptions(
91148
optimizer_version="3",
92149
optimizer_statistics_package="auto_20191128_14_47_22UTC",
93150
)
151+
self.assertEqual(result, expected)
152+
self.assertEqual(base.optimizer_version, "1")
153+
154+
def test_base_dict_merge_dict(self):
155+
base = {"optimizer_version": "1"}
156+
merge = {"optimizer_statistics_package": "auto_20191128_14_47_22UTC"}
157+
expected = ExecuteSqlRequest.QueryOptions(
158+
optimizer_version="1",
159+
optimizer_statistics_package="auto_20191128_14_47_22UTC",
160+
)
94161
result = self._callFUT(base, merge)
95162
self.assertEqual(result, expected)
96163

164+
def test_base_dict_override_dict(self):
165+
base = {
166+
"optimizer_version": "1",
167+
"optimizer_statistics_package": "pkg1",
168+
}
169+
merge = {"optimizer_version": "2"}
170+
expected = ExecuteSqlRequest.QueryOptions(
171+
optimizer_version="2",
172+
optimizer_statistics_package="pkg1",
173+
)
174+
result = self._callFUT(base, merge)
175+
self.assertEqual(result, expected)
176+
177+
def test_base_dict_empty_merge_none(self):
178+
result = self._callFUT({}, None)
179+
self.assertIsNone(result)
180+
181+
def test_base_none_merge_dict_empty(self):
182+
result = self._callFUT(None, {})
183+
self.assertIsNone(result)
184+
185+
def test_base_empty_dict_merge_empty_dict(self):
186+
result = self._callFUT({}, {})
187+
self.assertIsNone(result)
188+
189+
def test_base_empty_dict_merge_object(self):
190+
merge = ExecuteSqlRequest.QueryOptions(optimizer_version="1")
191+
result = self._callFUT({}, merge)
192+
self.assertEqual(result, merge)
193+
194+
def test_base_object_merge_empty_dict(self):
195+
base = ExecuteSqlRequest.QueryOptions(optimizer_version="1")
196+
result = self._callFUT(base, {})
197+
self.assertEqual(result, base)
198+
199+
def test_base_object_merge_object(self):
200+
base = ExecuteSqlRequest.QueryOptions(
201+
optimizer_version="1",
202+
optimizer_statistics_package="pkg1",
203+
)
204+
merge = ExecuteSqlRequest.QueryOptions(optimizer_version="2")
205+
result = self._callFUT(base, merge)
206+
expected = ExecuteSqlRequest.QueryOptions(
207+
optimizer_version="2",
208+
optimizer_statistics_package="pkg1",
209+
)
210+
self.assertEqual(result, expected)
211+
self.assertEqual(base.optimizer_version, "1")
212+
self.assertEqual(base.optimizer_statistics_package, "pkg1")
213+
self.assertEqual(merge.optimizer_version, "2")
214+
self.assertEqual(merge.optimizer_statistics_package, "")
215+
216+
def test_invalid_type_raises_error(self):
217+
with self.assertRaises(TypeError):
218+
self._callFUT("invalid", None)
219+
with self.assertRaises(TypeError):
220+
self._callFUT(None, 123)
221+
97222

98223
class Test_get_cloud_region(unittest.TestCase):
99224
def setUp(self):

0 commit comments

Comments
 (0)