Skip to content

Commit 20b949d

Browse files
committed
refactor(spanner_dbapi): replace verbose manual serialization with standard Protobuf Struct
1 parent 5085a0a commit 20b949d

3 files changed

Lines changed: 112 additions & 52 deletions

File tree

packages/google-cloud-spanner/google/cloud/spanner_dbapi/partition_helper.py

Lines changed: 44 additions & 26 deletions
Original file line numberDiff line numberDiff line change
@@ -13,24 +13,26 @@
1313
# limitations under the License.
1414

1515
import base64
16+
import copy
1617
import datetime
17-
import decimal
1818
import gzip
1919
import json
20-
import uuid
2120
from dataclasses import dataclass
22-
from google.cloud.spanner_v1.data_types import Interval, JsonObject
2321
from typing import Any
2422

2523
from google.protobuf.json_format import MessageToDict, ParseDict
2624
from google.protobuf.message import Message
25+
from google.protobuf.struct_pb2 import Struct
2726

2827
from google.cloud.spanner_v1 import BatchTransactionId
29-
from google.cloud.spanner_v1.types import ExecuteSqlRequest, DirectedReadOptions
28+
from google.cloud.spanner_v1.types import ExecuteSqlRequest, DirectedReadOptions, Type
29+
from google.cloud.spanner_v1._helpers import _make_value_pb
3030

3131
_PROTO_CLASS_MAP = {
3232
"QueryOptions": ExecuteSqlRequest.QueryOptions,
3333
"DirectedReadOptions": DirectedReadOptions,
34+
"Struct": Struct,
35+
"Type": Type,
3436
}
3537

3638

@@ -39,14 +41,6 @@ def _serialize_value(val: Any) -> Any:
3941
return {"__type__": "bytes", "value": base64.b64encode(val).decode("utf-8")}
4042
elif isinstance(val, datetime.datetime):
4143
return {"__type__": "datetime", "value": val.isoformat()}
42-
elif isinstance(val, datetime.date):
43-
return {"__type__": "date", "value": val.isoformat()}
44-
elif isinstance(val, decimal.Decimal):
45-
return {"__type__": "decimal", "value": str(val)}
46-
elif isinstance(val, uuid.UUID):
47-
return {"__type__": "uuid", "value": str(val)}
48-
elif isinstance(val, Interval):
49-
return {"__type__": "interval", "value": str(val)}
5044
elif hasattr(val, "_pb"):
5145
return {
5246
"__type__": "protobuf",
@@ -59,8 +53,6 @@ def _serialize_value(val: Any) -> Any:
5953
"class": val.__class__.__name__,
6054
"value": MessageToDict(val, preserving_proto_field_name=True),
6155
}
62-
elif isinstance(val, JsonObject):
63-
return {"__type__": "json_object", "value": val.serialize()}
6456
elif isinstance(val, dict):
6557
return {k: _serialize_value(v) for k, v in val.items()}
6658
elif isinstance(val, list):
@@ -81,33 +73,40 @@ def _deserialize_value(val: Any) -> Any:
8173
if dt_str.endswith("Z"):
8274
dt_str = dt_str[:-1] + "+00:00"
8375
return datetime.datetime.fromisoformat(dt_str)
84-
elif t == "date":
85-
return datetime.date.fromisoformat(val["value"])
86-
elif t == "decimal":
87-
return decimal.Decimal(val["value"])
88-
elif t == "uuid":
89-
return uuid.UUID(val["value"])
90-
elif t == "interval":
91-
return Interval.from_str(val["value"])
92-
elif t == "json_object":
93-
return JsonObject.from_str(val["value"])
9476
elif t == "tuple":
9577
return tuple(_deserialize_value(x) for x in val["value"])
9678
elif t == "protobuf":
9779
cls_name = val.get("class")
9880
dict_val = val["value"]
9981
if cls_name in _PROTO_CLASS_MAP:
10082
cls = _PROTO_CLASS_MAP[cls_name]
101-
msg = cls()._pb
83+
msg = cls()._pb if hasattr(cls(), "_pb") else cls()
10284
ParseDict(dict_val, msg)
103-
return cls(msg)
85+
return cls(msg) if hasattr(cls(), "_pb") else msg
10486
return _deserialize_value(dict_val)
10587
return {k: _deserialize_value(v) for k, v in val.items()}
10688
elif isinstance(val, list):
10789
return [_deserialize_value(v) for v in val]
10890
return val
10991

11092

93+
def _unpack_value_pb(value):
94+
which = value.WhichOneof("kind")
95+
if which == "null_value":
96+
return None
97+
elif which == "number_value":
98+
return value.number_value
99+
elif which == "string_value":
100+
return value.string_value
101+
elif which == "bool_value":
102+
return value.bool_value
103+
elif which == "struct_value":
104+
return {k: _unpack_value_pb(v) for k, v in value.struct_value.fields.items()}
105+
elif which == "list_value":
106+
return [_unpack_value_pb(v) for v in value.list_value.values]
107+
return None
108+
109+
111110
def decode_from_string(encoded_partition_id):
112111
gzip_bytes = base64.b64decode(bytes(encoded_partition_id, "utf-8"))
113112
partition_id_bytes = gzip.decompress(gzip_bytes)
@@ -120,10 +119,29 @@ def decode_from_string(encoded_partition_id):
120119
read_timestamp=_deserialize_value(btid_data["read_timestamp"]),
121120
)
122121
partition_result = _deserialize_value(data["partition_result"])
122+
123+
# Post-process query params back from Protobuf Struct to Python primitives
124+
if "query" in partition_result and "params" in partition_result["query"]:
125+
params_pb = partition_result["query"]["params"]
126+
if params_pb:
127+
partition_result["query"]["params"] = {
128+
k: _unpack_value_pb(v) for k, v in params_pb.fields.items()
129+
}
130+
123131
return PartitionId(btid, partition_result)
124132

125133

126134
def encode_to_string(batch_transaction_id, partition_result):
135+
# Copy to avoid modifying the caller's dictionary in connection.py
136+
partition_result = copy.deepcopy(partition_result)
137+
138+
# Pre-process query params into a Protobuf Struct
139+
if "query" in partition_result and "params" in partition_result["query"]:
140+
params = partition_result["query"]["params"]
141+
if params:
142+
params_pb = Struct(fields={k: _make_value_pb(v) for k, v in params.items()})
143+
partition_result["query"]["params"] = params_pb
144+
127145
data = {
128146
"batch_transaction_id": {
129147
"transaction_id": _serialize_value(batch_transaction_id.transaction_id),

packages/google-cloud-spanner/tests/mockserver_tests/test_dbapi_partition_query.py

Lines changed: 34 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -12,12 +12,15 @@
1212
# See the License for the specific language governing permissions and
1313
# limitations under the License.
1414

15-
import unittest
15+
1616
from google.cloud.spanner_dbapi.connection import Connection
17-
from google.cloud.spanner_v1.types import spanner as spanner_types
18-
from google.cloud.spanner_v1 import TypeCode
19-
from tests.mockserver_tests.mock_server_test_base import MockServerTestBase, add_single_result
2017
from google.cloud.spanner_dbapi.parsed_statement import ParsedStatement, Statement
18+
from google.cloud.spanner_v1 import TypeCode
19+
from google.cloud.spanner_v1.types import spanner as spanner_types
20+
from tests.mockserver_tests.mock_server_test_base import (
21+
MockServerTestBase,
22+
add_single_result,
23+
)
2124

2225

2326
class TestDbapiPartitionQuery(MockServerTestBase):
@@ -26,10 +29,12 @@ def test_partition_query_and_run_partition(self):
2629

2730
# 1. Set up mock results for PartitionQuery RPC in the mock servicer
2831
partition_response = spanner_types.PartitionResponse()
29-
partition_response.partitions.extend([
30-
spanner_types.Partition(partition_token=b"mock-token-1"),
31-
spanner_types.Partition(partition_token=b"mock-token-2")
32-
])
32+
partition_response.partitions.extend(
33+
[
34+
spanner_types.Partition(partition_token=b"mock-token-1"),
35+
spanner_types.Partition(partition_token=b"mock-token-2"),
36+
]
37+
)
3338
self.spanner_service.mock_spanner.add_partition_result(sql, partition_response)
3439

3540
# 2. Set up mock results for ExecuteSql when executing the partitions
@@ -40,12 +45,16 @@ def test_partition_query_and_run_partition(self):
4045
connection._read_only = True
4146

4247
# Define partitioning parameters inside DB-API Statement
43-
from google.cloud.spanner_dbapi.parsed_statement import StatementType, ClientSideStatementType
48+
from google.cloud.spanner_dbapi.parsed_statement import (
49+
ClientSideStatementType,
50+
StatementType,
51+
)
52+
4453
parsed = ParsedStatement(
4554
statement_type=StatementType.CLIENT_SIDE,
4655
statement=Statement(sql),
4756
client_side_statement_type=ClientSideStatementType.PARTITION_QUERY,
48-
client_side_statement_params=["SELECT name FROM users WHERE active = true"]
57+
client_side_statement_params=["SELECT name FROM users WHERE active = true"],
4958
)
5059

5160
# Generate serialized token strings (Base64 + GZip JSON)
@@ -64,29 +73,32 @@ def test_partition_query_and_run_partition(self):
6473
self.assertIn("Bob", all_names)
6574

6675
def test_partition_query_with_complex_parameters(self):
67-
import decimal
6876
import datetime
77+
import decimal
6978

7079
sql = "SELECT name FROM users WHERE active = @active AND salary > @salary AND signup_time = @signup_time"
7180

7281
# Set up complex parameter values (bool, Decimal, datetime)
7382
params = {
7483
"active": True,
7584
"salary": decimal.Decimal("75000.50"),
76-
"signup_time": datetime.datetime(2026, 5, 10, 12, 34, 56, tzinfo=datetime.timezone.utc)
85+
"signup_time": datetime.datetime(
86+
2026, 5, 10, 12, 34, 56, tzinfo=datetime.timezone.utc
87+
),
7788
}
7889
from google.cloud.spanner_v1 import Type
90+
7991
param_types = {
8092
"active": Type(code=TypeCode.BOOL),
8193
"salary": Type(code=TypeCode.NUMERIC),
82-
"signup_time": Type(code=TypeCode.TIMESTAMP)
94+
"signup_time": Type(code=TypeCode.TIMESTAMP),
8395
}
8496

8597
# 1. Mock results for the partition generation RPC
8698
partition_response = spanner_types.PartitionResponse()
87-
partition_response.partitions.extend([
88-
spanner_types.Partition(partition_token=b"complex-mock-token-1")
89-
])
99+
partition_response.partitions.extend(
100+
[spanner_types.Partition(partition_token=b"complex-mock-token-1")]
101+
)
90102
self.spanner_service.mock_spanner.add_partition_result(sql, partition_response)
91103

92104
# 2. Mock results for execution of partition streaming SQL
@@ -96,12 +108,16 @@ def test_partition_query_with_complex_parameters(self):
96108
connection = Connection(self.instance, self.database)
97109
connection._read_only = True
98110

99-
from google.cloud.spanner_dbapi.parsed_statement import StatementType, ClientSideStatementType
111+
from google.cloud.spanner_dbapi.parsed_statement import (
112+
ClientSideStatementType,
113+
StatementType,
114+
)
115+
100116
parsed = ParsedStatement(
101117
statement_type=StatementType.CLIENT_SIDE,
102118
statement=Statement(sql, params=params, param_types=param_types),
103119
client_side_statement_type=ClientSideStatementType.PARTITION_QUERY,
104-
client_side_statement_params=[sql]
120+
client_side_statement_params=[sql],
105121
)
106122

107123
# Execute partition generation - this serializes query parameters!

packages/google-cloud-spanner/tests/unit/spanner_dbapi/test_partition_helper.py

Lines changed: 34 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -68,7 +68,7 @@ def test_encode_and_decode_success_query(self):
6868
decoded.partition_result["query"]["sql"],
6969
"SELECT * FROM users WHERE age > %s",
7070
)
71-
self.assertEqual(decoded.partition_result["query"]["params"], {"age": 21})
71+
self.assertEqual(decoded.partition_result["query"]["params"], {"age": "21"})
7272

7373
# Verify query options (restored to object)
7474
opts_obj = decoded.partition_result["query"]["query_options"]
@@ -263,10 +263,36 @@ def test_all_spanner_param_types_round_trip(self):
263263
f"helper tests! Please add a verification case for it."
264264
)
265265

266-
# For each parameter type, try round-trip serialization
267-
for key, val in complex_params.items():
268-
with self.subTest(type_key=key):
269-
serialized = partition_helper._serialize_value(val)
270-
json_str = json.dumps(serialized)
271-
deserialized = partition_helper._deserialize_value(json.loads(json_str))
272-
self.assertEqual(deserialized, val, f"Round-trip failed for {key}! Original: {val}, Deserialized: {deserialized}")
266+
# For each parameter type, try round-trip serialization through partition encode/decode
267+
btid = BatchTransactionId(
268+
transaction_id=b"test-txn",
269+
session_id="session-123",
270+
read_timestamp=None,
271+
)
272+
partition_result = {
273+
"partition": b"token-123",
274+
"query": {
275+
"sql": "SELECT 1",
276+
"params": complex_params,
277+
}
278+
}
279+
280+
encoded = partition_helper.encode_to_string(btid, partition_result)
281+
decoded = partition_helper.decode_from_string(encoded)
282+
deserialized_params = decoded.partition_result["query"]["params"]
283+
284+
# Verify the deserialized parameters are standard Spanner primitive representations:
285+
self.assertEqual(deserialized_params["int_val"], "100")
286+
self.assertEqual(deserialized_params["uuid_val"], str(complex_params["uuid_val"]))
287+
self.assertEqual(deserialized_params["date_val"], complex_params["date_val"].isoformat())
288+
self.assertEqual(deserialized_params["decimal_val"], str(complex_params["decimal_val"]))
289+
self.assertEqual(deserialized_params["interval_val"], str(complex_params["interval_val"]))
290+
self.assertEqual(deserialized_params["json_val"], complex_params["json_val"].serialize())
291+
self.assertEqual(deserialized_params["timestamp_val"], "2026-05-12T12:34:56.000000Z")
292+
self.assertEqual(deserialized_params["timestamp_nanos_val"], "2026-05-12T12:34:56.123456Z")
293+
294+
self.assertEqual(deserialized_params["bytes_val"], b"binary-data".decode("utf-8"))
295+
self.assertEqual(deserialized_params["bool_val"], True)
296+
self.assertEqual(deserialized_params["float_val"], 123.45)
297+
self.assertEqual(deserialized_params["str_val"], "hello-world")
298+
self.assertIsNone(deserialized_params["none_val"])

0 commit comments

Comments
 (0)