Skip to content

Commit 56db3f3

Browse files
committed
perf(spanner): support PartialResultSet.last with background stream draining
When Cloud Spanner finishes transmitting query results, it marks `last = True` on the final PartialResultSet chunk. Previously, the client blocked synchronously waiting for gRPC trailers and EOF frames over the wire before returning the final rows. Additionally, abandoning streams early caused gRPC's C-core finalizer to mark them as CANCELLED upon garbage collection. This change enables immediate return upon observing `last = True` and offloads trailing metadata consumption to the background so streams complete cleanly with status OK: - In sync mode, completed streams are handed off to `_BoundedStreamDrainer`, which uses a bounded queue and daemon worker threads to drain to EOF. If the queue is full or the interpreter is shutting down, it falls back to inline draining. Process fork safety is ensured via `os.register_at_fork`. - In async mode, trailing frames are drained via a background `asyncio.create_task` with strong reference retention to prevent premature task garbage collection and clean cancellation handling. - Transaction precommit tokens, query stats, and metadata present on the final chunk are captured before handing the stream off to background draining.
1 parent 59943e8 commit 56db3f3

8 files changed

Lines changed: 1211 additions & 149 deletions

File tree

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

Lines changed: 42 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,6 @@
11
import asyncio
22
import inspect
3+
import os
34
import time
45

56
from google.api_core.exceptions import Aborted
@@ -153,3 +154,44 @@ def _create_experimental_host_transport(
153154
client_key,
154155
interceptors=interceptors,
155156
)
157+
158+
159+
_PENDING_DRAIN_TASKS = set()
160+
if hasattr(os, "register_at_fork"):
161+
os.register_at_fork(after_in_child=_PENDING_DRAIN_TASKS.clear)
162+
163+
164+
def _drain_stream(iterator):
165+
"""Drain an async stream iterator to EOF in the background.
166+
167+
Called when PartialResultSet.last is True to allow the caller to return immediately
168+
while consuming trailing gRPC metadata so the stream terminates cleanly with status OK.
169+
"""
170+
if iterator is None:
171+
return
172+
173+
async def _drain():
174+
try:
175+
async for _ in iterator:
176+
pass
177+
except asyncio.CancelledError:
178+
if hasattr(iterator, "cancel"):
179+
try:
180+
iterator.cancel()
181+
except Exception:
182+
pass
183+
raise
184+
except Exception:
185+
pass
186+
187+
try:
188+
task = asyncio.create_task(_drain())
189+
_PENDING_DRAIN_TASKS.add(task)
190+
task.add_done_callback(_PENDING_DRAIN_TASKS.discard)
191+
except RuntimeError:
192+
# Event loop may be closed or not running.
193+
if hasattr(iterator, "cancel"):
194+
try:
195+
iterator.cancel()
196+
except Exception:
197+
pass

packages/google-cloud-spanner/google/cloud/spanner_v1/_async/snapshot.py

Lines changed: 106 additions & 83 deletions
Original file line numberDiff line numberDiff line change
@@ -28,7 +28,7 @@
2828
from google.protobuf.struct_pb2 import Struct
2929

3030
from google.cloud.aio._cross_sync import CrossSync
31-
from google.cloud.spanner_v1._async._helpers import _retry
31+
from google.cloud.spanner_v1._async._helpers import _drain_stream, _retry
3232
from google.cloud.spanner_v1._async.streamed import StreamedResultSet
3333
from google.cloud.spanner_v1._helpers import (
3434
AtomicCounter,
@@ -114,94 +114,117 @@ async def _restart_on_unavailable(
114114
attempt = 1
115115
nth_request = getattr(request_id_manager, "_next_nth_request", 0)
116116
current_request_id = None
117-
118-
while True:
119-
try:
120-
# Get results iterator.
121-
if iterator is None:
122-
with (
123-
trace_call(
124-
trace_name,
125-
session,
126-
attributes,
127-
observability_options=observability_options,
128-
metadata=metadata,
129-
) as span,
130-
MetricsCapture(resource_info),
131-
):
132-
(
133-
call_metadata,
134-
current_request_id,
135-
) = request_id_manager.metadata_and_request_id(
136-
nth_request,
137-
attempt,
138-
metadata,
139-
span,
140-
)
141-
iterator = await CrossSync.run_if_async(
142-
method,
143-
request=request,
144-
metadata=call_metadata,
145-
)
146-
147-
# Add items from iterator to buffer.
148-
item: PartialResultSet
149-
async for item in iterator:
150-
item_buffer.append(item)
151-
152-
# Update the transaction from the response.
117+
stream_finished = False
118+
119+
try:
120+
while True:
121+
try:
122+
# Get results iterator.
123+
if iterator is None:
124+
with (
125+
trace_call(
126+
trace_name,
127+
session,
128+
attributes,
129+
observability_options=observability_options,
130+
metadata=metadata,
131+
) as span,
132+
MetricsCapture(resource_info),
133+
):
134+
(
135+
call_metadata,
136+
current_request_id,
137+
) = request_id_manager.metadata_and_request_id(
138+
nth_request,
139+
attempt,
140+
metadata,
141+
span,
142+
)
143+
iterator = await CrossSync.run_if_async(
144+
method,
145+
request=request,
146+
metadata=call_metadata,
147+
)
148+
149+
# Add items from iterator to buffer.
150+
item: PartialResultSet
151+
async for item in iterator:
152+
item_buffer.append(item)
153+
154+
# Update the transaction from the response.
155+
if transaction is not None:
156+
transaction._update_for_result_set_pb(item)
157+
if (
158+
item._pb is not None
159+
and item._pb.HasField("precommit_token")
160+
and transaction is not None
161+
):
162+
await transaction._update_for_precommit_token_pb(
163+
item.precommit_token
164+
)
165+
166+
try:
167+
item_is_last = item.last
168+
except AttributeError:
169+
item_is_last = False
170+
171+
if item_is_last:
172+
stream_finished = True
173+
_drain_stream(iterator)
174+
iterator = None
175+
break
176+
177+
if item.resume_token:
178+
resume_token = item.resume_token
179+
break
180+
181+
except ServiceUnavailable:
182+
del item_buffer[:]
183+
request.resume_token = resume_token
153184
if transaction is not None:
154-
transaction._update_for_result_set_pb(item)
155-
if (
156-
item._pb is not None
157-
and item._pb.HasField("precommit_token")
158-
and transaction is not None
159-
):
160-
await transaction._update_for_precommit_token_pb(
161-
item.precommit_token
162-
)
163-
164-
if item.resume_token:
165-
resume_token = item.resume_token
166-
break
167-
168-
except ServiceUnavailable:
169-
del item_buffer[:]
170-
request.resume_token = resume_token
171-
if transaction is not None:
172-
transaction_selector = transaction._build_transaction_selector_pb()
173-
request.transaction = transaction_selector
174-
attempt += 1
175-
iterator = None
176-
continue
177-
178-
except InternalServerError as exc:
179-
resumable_error = any(
180-
resumable_message in exc.message
181-
for resumable_message in _STREAM_RESUMPTION_INTERNAL_ERROR_MESSAGES
182-
)
183-
if not resumable_error:
185+
transaction_selector = transaction._build_transaction_selector_pb()
186+
request.transaction = transaction_selector
187+
attempt += 1
188+
iterator = None
189+
continue
190+
191+
except InternalServerError as exc:
192+
resumable_error = any(
193+
resumable_message in exc.message
194+
for resumable_message in _STREAM_RESUMPTION_INTERNAL_ERROR_MESSAGES
195+
)
196+
if not resumable_error:
197+
raise _augment_error_with_request_id(exc, current_request_id)
198+
del item_buffer[:]
199+
request.resume_token = resume_token
200+
if transaction is not None:
201+
transaction_selector = transaction._build_transaction_selector_pb()
202+
attempt += 1
203+
request.transaction = transaction_selector
204+
iterator = None
205+
continue
206+
207+
except Exception as exc:
208+
# Augment any other exception with the request ID
184209
raise _augment_error_with_request_id(exc, current_request_id)
185-
del item_buffer[:]
186-
request.resume_token = resume_token
187-
if transaction is not None:
188-
transaction_selector = transaction._build_transaction_selector_pb()
189-
attempt += 1
190-
request.transaction = transaction_selector
191-
iterator = None
192-
continue
193210

194-
except Exception as exc:
195-
# Augment any other exception with the request ID
196-
raise _augment_error_with_request_id(exc, current_request_id)
211+
if len(item_buffer) == 0:
212+
iterator = None
213+
break
197214

198-
if len(item_buffer) == 0:
199-
break
215+
for item in item_buffer:
216+
yield item
200217

201-
for item in item_buffer:
202-
yield item
218+
del item_buffer[:]
203219

204-
del item_buffer[:]
220+
if stream_finished:
221+
break
222+
finally:
223+
if iterator is not None and hasattr(iterator, "cancel"):
224+
try:
225+
iterator.cancel()
226+
except Exception:
227+
pass
205228

206229

207230
class _SnapshotBase(_SessionWrapper):

0 commit comments

Comments
 (0)