Skip to content

Commit 8cb3c6a

Browse files
authored
perf(spanner): optimize query option merging and prevent in-place mutation (#18358)
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 58031be commit 8cb3c6a

2 files changed

Lines changed: 207 additions & 26 deletions

File tree

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

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

162162

163+
_VALID_QUERY_OPTIONS_KEYS = frozenset(
164+
ExecuteSqlRequest.QueryOptions._meta.fields.keys()
165+
)
166+
167+
168+
def _to_query_options(options):
169+
"""Normalize dict or QueryOptions to a non-empty QueryOptions, or None.
170+
171+
:type options:
172+
:class:`~google.cloud.spanner_v1.types.ExecuteSqlRequest.QueryOptions`
173+
or :class:`dict` or None
174+
:param options: Query options to normalize.
175+
176+
:rtype:
177+
:class:`~google.cloud.spanner_v1.types.ExecuteSqlRequest.QueryOptions`
178+
or None
179+
:returns:
180+
A non-empty QueryOptions instance, or None if options is empty or None.
181+
182+
:raises TypeError:
183+
If options is not a QueryOptions, dict, or None.
184+
:raises ValueError:
185+
If options is a dict containing unknown fields.
186+
"""
187+
if options is None:
188+
return None
189+
if isinstance(options, dict):
190+
if options.keys() <= _VALID_QUERY_OPTIONS_KEYS and not any(options.values()):
191+
return None
192+
options = ExecuteSqlRequest.QueryOptions(options)
193+
elif not isinstance(options, ExecuteSqlRequest.QueryOptions):
194+
raise TypeError(
195+
f"query_options must be a QueryOptions or dict, got {type(options).__name__}"
196+
)
197+
return options if type(options).pb(options).ByteSize() > 0 else None
198+
199+
163200
def _merge_query_options(base, merge):
164201
"""Merge higher precedence QueryOptions with current QueryOptions.
165202
@@ -182,23 +219,20 @@ def _merge_query_options(base, merge):
182219
QueryOptions object formed by merging the two given QueryOptions.
183220
If the resultant object only has empty fields, returns None.
184221
"""
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:
222+
if base is None and merge is None:
201223
return None
224+
225+
base = _to_query_options(base)
226+
merge = _to_query_options(merge)
227+
if base is None:
228+
return merge
229+
if merge is None:
230+
return base
231+
232+
combined = ExecuteSqlRequest.QueryOptions()
233+
combined_pb = type(combined).pb(combined)
234+
combined_pb.CopyFrom(type(base).pb(base))
235+
combined_pb.MergeFrom(type(merge).pb(merge))
202236
return combined
203237

204238

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

Lines changed: 157 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -22,7 +22,50 @@
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+
for invalid_value in ("invalid", 123, "", [], False):
58+
with self.subTest(invalid_value=invalid_value):
59+
with self.assertRaises(TypeError):
60+
self._callFUT(invalid_value)
61+
62+
def test_unknown_key_with_empty_value_raises_error(self):
63+
with self.assertRaises(ValueError):
64+
self._callFUT({"optmizer_version": ""})
65+
66+
def test_unknown_key_with_non_empty_value_raises_error(self):
67+
with self.assertRaises(ValueError):
68+
self._callFUT({"optmizer_version": "1"})
2669

2770

2871
class Test_merge_query_options(unittest.TestCase):
@@ -37,8 +80,6 @@ def test_base_none_and_merge_none(self):
3780
self.assertIsNone(result)
3881

3982
def test_base_dict_and_merge_none(self):
40-
from google.cloud.spanner_v1 import ExecuteSqlRequest
41-
4283
base = {
4384
"optimizer_version": "2",
4485
"optimizer_statistics_package": "auto_20191128_14_47_22UTC",
@@ -52,16 +93,12 @@ def test_base_dict_and_merge_none(self):
5293
self.assertEqual(result, expected)
5394

5495
def test_base_empty_and_merge_empty(self):
55-
from google.cloud.spanner_v1 import ExecuteSqlRequest
56-
5796
base = ExecuteSqlRequest.QueryOptions()
5897
merge = ExecuteSqlRequest.QueryOptions()
5998
result = self._callFUT(base, merge)
6099
self.assertIsNone(result)
61100

62101
def test_base_none_merge_object(self):
63-
from google.cloud.spanner_v1 import ExecuteSqlRequest
64-
65102
base = None
66103
merge = ExecuteSqlRequest.QueryOptions(
67104
optimizer_version="3",
@@ -71,28 +108,138 @@ def test_base_none_merge_object(self):
71108
self.assertEqual(result, merge)
72109

73110
def test_base_none_merge_dict(self):
74-
from google.cloud.spanner_v1 import ExecuteSqlRequest
75-
76111
base = None
77112
merge = {"optimizer_version": "3"}
78113
expected = ExecuteSqlRequest.QueryOptions(optimizer_version="3")
79114
result = self._callFUT(base, merge)
80115
self.assertEqual(result, expected)
81116

82117
def test_base_object_merge_dict(self):
83-
from google.cloud.spanner_v1 import ExecuteSqlRequest
118+
base = ExecuteSqlRequest.QueryOptions(
119+
optimizer_version="1",
120+
optimizer_statistics_package="auto_20191128_14_47_22UTC",
121+
)
122+
merge = {"optimizer_version": "3"}
123+
expected = ExecuteSqlRequest.QueryOptions(
124+
optimizer_version="3",
125+
optimizer_statistics_package="auto_20191128_14_47_22UTC",
126+
)
127+
result = self._callFUT(base, merge)
128+
self.assertEqual(result, expected)
129+
130+
def test_base_object_and_merge_none(self):
131+
base = ExecuteSqlRequest.QueryOptions(
132+
optimizer_version="2",
133+
optimizer_statistics_package="auto_20191128_14_47_22UTC",
134+
)
135+
result = self._callFUT(base, None)
136+
self.assertEqual(result, base)
137+
138+
def test_base_empty_object_and_merge_none(self):
139+
base = ExecuteSqlRequest.QueryOptions()
140+
result = self._callFUT(base, None)
141+
self.assertIsNone(result)
142+
143+
def test_base_none_merge_empty_object(self):
144+
merge = ExecuteSqlRequest.QueryOptions()
145+
result = self._callFUT(None, merge)
146+
self.assertIsNone(result)
84147

148+
def test_base_object_not_mutated_on_merge(self):
85149
base = ExecuteSqlRequest.QueryOptions(
86150
optimizer_version="1",
87151
optimizer_statistics_package="auto_20191128_14_47_22UTC",
88152
)
89153
merge = {"optimizer_version": "3"}
154+
result = self._callFUT(base, merge)
90155
expected = ExecuteSqlRequest.QueryOptions(
91156
optimizer_version="3",
92157
optimizer_statistics_package="auto_20191128_14_47_22UTC",
93158
)
159+
self.assertEqual(result, expected)
160+
self.assertEqual(base.optimizer_version, "1")
161+
162+
def test_base_dict_merge_dict(self):
163+
base = {"optimizer_version": "1"}
164+
merge = {"optimizer_statistics_package": "auto_20191128_14_47_22UTC"}
165+
expected = ExecuteSqlRequest.QueryOptions(
166+
optimizer_version="1",
167+
optimizer_statistics_package="auto_20191128_14_47_22UTC",
168+
)
169+
result = self._callFUT(base, merge)
170+
self.assertEqual(result, expected)
171+
172+
def test_base_dict_override_dict(self):
173+
base = {
174+
"optimizer_version": "1",
175+
"optimizer_statistics_package": "pkg1",
176+
}
177+
merge = {"optimizer_version": "2"}
178+
expected = ExecuteSqlRequest.QueryOptions(
179+
optimizer_version="2",
180+
optimizer_statistics_package="pkg1",
181+
)
182+
result = self._callFUT(base, merge)
183+
self.assertEqual(result, expected)
184+
185+
def test_base_dict_empty_merge_none(self):
186+
result = self._callFUT({}, None)
187+
self.assertIsNone(result)
188+
189+
def test_base_none_merge_dict_empty(self):
190+
result = self._callFUT(None, {})
191+
self.assertIsNone(result)
192+
193+
def test_base_empty_dict_merge_empty_dict(self):
194+
result = self._callFUT({}, {})
195+
self.assertIsNone(result)
196+
197+
def test_base_empty_dict_merge_object(self):
198+
merge = ExecuteSqlRequest.QueryOptions(optimizer_version="1")
199+
result = self._callFUT({}, merge)
200+
self.assertEqual(result, merge)
201+
202+
def test_base_object_merge_empty_dict(self):
203+
base = ExecuteSqlRequest.QueryOptions(optimizer_version="1")
204+
result = self._callFUT(base, {})
205+
self.assertEqual(result, base)
206+
207+
def test_base_object_merge_object(self):
208+
base = ExecuteSqlRequest.QueryOptions(
209+
optimizer_version="1",
210+
optimizer_statistics_package="pkg1",
211+
)
212+
merge = ExecuteSqlRequest.QueryOptions(optimizer_version="2")
94213
result = self._callFUT(base, merge)
214+
expected = ExecuteSqlRequest.QueryOptions(
215+
optimizer_version="2",
216+
optimizer_statistics_package="pkg1",
217+
)
95218
self.assertEqual(result, expected)
219+
self.assertEqual(base.optimizer_version, "1")
220+
self.assertEqual(base.optimizer_statistics_package, "pkg1")
221+
self.assertEqual(merge.optimizer_version, "2")
222+
self.assertEqual(merge.optimizer_statistics_package, "")
223+
224+
def test_invalid_type_raises_error(self):
225+
for invalid_value in ("invalid", 123, "", [], False):
226+
with self.subTest(invalid_value=invalid_value):
227+
with self.assertRaises(TypeError):
228+
self._callFUT(invalid_value, None)
229+
with self.assertRaises(TypeError):
230+
self._callFUT(None, invalid_value)
231+
232+
def test_unknown_key_in_base_raises_error(self):
233+
with self.assertRaises(ValueError):
234+
self._callFUT({"optmizer_version": ""}, None)
235+
with self.assertRaises(ValueError):
236+
self._callFUT({"optmizer_version": "1"}, None)
237+
238+
def test_unknown_key_in_merge_raises_error(self):
239+
with self.assertRaises(ValueError):
240+
self._callFUT(None, {"optmizer_version": ""})
241+
with self.assertRaises(ValueError):
242+
self._callFUT(None, {"optmizer_version": "1"})
96243

97244

98245
class Test_get_cloud_region(unittest.TestCase):

0 commit comments

Comments
 (0)