Skip to content

Commit a495bbb

Browse files
committed
fix(metrics): correct GFE metrics extraction and enable by default
1 parent f492d3d commit a495bbb

8 files changed

Lines changed: 705 additions & 24 deletions

File tree

packages/google-cloud-spanner/google/cloud/spanner_v1/metrics/metrics_interceptor.py

Lines changed: 354 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -14,14 +14,19 @@
1414

1515
"""Interceptor for collecting Cloud Spanner metrics."""
1616

17+
import inspect
18+
import logging
1719
import re
18-
from typing import Dict
20+
from typing import Any, Dict
1921

22+
import grpc
2023
from grpc_interceptor import ClientInterceptor
2124

2225
from .constants import GOOGLE_CLOUD_RESOURCE_KEY, SPANNER_METHOD_PREFIX
2326
from .spanner_metrics_tracer_factory import SpannerMetricsTracerFactory
2427

28+
logger = logging.getLogger(__name__)
29+
2530

2631
class MetricsInterceptor(ClientInterceptor):
2732
"""Interceptor that collects metrics for Cloud Spanner operations."""
@@ -115,17 +120,358 @@ def intercept(self, invoked_method, request_or_iterator, call_details):
115120
self._set_metrics_tracer_attributes(resources)
116121

117122
## Format method to be be spanner.<method name>
118-
method_name = call_details.method.removeprefix(SPANNER_METHOD_PREFIX).replace(
119-
"/", "."
120-
)
123+
method_str = call_details.method
124+
if isinstance(method_str, bytes):
125+
method_str = method_str.decode("utf-8")
126+
method_name = method_str.removeprefix(SPANNER_METHOD_PREFIX).replace("/", ".")
121127

122128
tracer.set_method(method_name)
123129
tracer.record_attempt_start()
124130
response = invoked_method(request_or_iterator, call_details)
125-
tracer.record_attempt_completion()
126131

127-
# Process and send GFE metrics if enabled
128-
if tracer.gfe_enabled:
129-
metadata = response.initial_metadata()
132+
return _wrap_response(response, tracer)
133+
134+
135+
def _wrap_response(response: Any, tracer: Any) -> Any:
136+
"""Wraps the response if it is streaming, or records metrics immediately if unary."""
137+
if hasattr(response, "__next__"):
138+
return _StreamingResponseWrapper(response, tracer)
139+
else:
140+
# Unary call: execute completion and record metrics immediately
141+
try:
142+
tracer.record_attempt_completion()
143+
metadata = []
144+
if hasattr(response, "initial_metadata"):
145+
try:
146+
metadata.extend(response.initial_metadata() or [])
147+
except Exception as e:
148+
logger.warning(f"Failed to retrieve initial metadata: {e}")
149+
if hasattr(response, "trailing_metadata"):
150+
try:
151+
metadata.extend(response.trailing_metadata() or [])
152+
except Exception as e:
153+
logger.warning(f"Failed to retrieve trailing metadata: {e}")
130154
tracer.record_gfe_metrics(metadata)
155+
except Exception as e:
156+
logger.warning(f"Failed to record metrics: {e}")
131157
return response
158+
159+
160+
class AsyncMetricsInterceptor(
161+
grpc.aio.UnaryUnaryClientInterceptor,
162+
grpc.aio.UnaryStreamClientInterceptor,
163+
grpc.aio.StreamUnaryClientInterceptor,
164+
grpc.aio.StreamStreamClientInterceptor,
165+
):
166+
"""Async Interceptor that collects metrics for Cloud Spanner operations."""
167+
168+
async def intercept_unary_unary(self, continuation, client_call_details, request):
169+
return await self._async_intercept(continuation, client_call_details, request)
170+
171+
async def intercept_unary_stream(self, continuation, client_call_details, request):
172+
return await self._async_intercept(continuation, client_call_details, request)
173+
174+
async def intercept_stream_unary(
175+
self, continuation, client_call_details, request_iterator
176+
):
177+
return await self._async_intercept(
178+
continuation, client_call_details, request_iterator
179+
)
180+
181+
async def intercept_stream_stream(
182+
self, continuation, client_call_details, request_iterator
183+
):
184+
return await self._async_intercept(
185+
continuation, client_call_details, request_iterator
186+
)
187+
188+
async def _async_intercept(
189+
self,
190+
continuation: Any,
191+
call_details: grpc.ClientCallDetails,
192+
request_or_iterator: Any,
193+
) -> Any:
194+
# Implementation for async interceptor
195+
factory = SpannerMetricsTracerFactory()
196+
tracer = SpannerMetricsTracerFactory.get_current_tracer()
197+
if tracer is None or not factory.enabled:
198+
return await continuation(call_details, request_or_iterator)
199+
200+
if not (
201+
tracer.client_attributes.get("project_id")
202+
and tracer.client_attributes.get("instance_id")
203+
and tracer.client_attributes.get("database")
204+
):
205+
resources = MetricsInterceptor._extract_resource_from_path(
206+
call_details.metadata
207+
)
208+
MetricsInterceptor._set_metrics_tracer_attributes(resources)
209+
210+
method_str = call_details.method
211+
if isinstance(method_str, bytes):
212+
method_str = method_str.decode("utf-8")
213+
method_name = method_str.removeprefix(SPANNER_METHOD_PREFIX).replace("/", ".")
214+
215+
tracer.set_method(method_name)
216+
tracer.record_attempt_start()
217+
response = await continuation(call_details, request_or_iterator)
218+
219+
if hasattr(response, "__anext__"):
220+
return _AsyncStreamingResponseWrapper(response, tracer)
221+
else:
222+
return _AsyncUnaryResponseWrapper(response, tracer)
223+
224+
225+
class _StreamingResponseWrapper:
226+
"""Wrapper for streaming RPC response iterators to defer metrics recording."""
227+
228+
def __init__(self, response, tracer):
229+
self._response = response
230+
self._tracer = tracer
231+
self._metrics_recorded = False
232+
self._iterator = None
233+
234+
def __iter__(self):
235+
self._iterator = iter(self._response)
236+
return self
237+
238+
def __next__(self):
239+
if self._iterator is None:
240+
self._iterator = iter(self._response)
241+
try:
242+
return next(self._iterator)
243+
except StopIteration:
244+
self._record_metrics()
245+
raise
246+
except Exception:
247+
self._record_metrics()
248+
raise
249+
250+
def _record_metrics(self):
251+
if self._metrics_recorded:
252+
return
253+
self._metrics_recorded = True
254+
try:
255+
self._tracer.record_attempt_completion()
256+
metadata = []
257+
if hasattr(self._response, "initial_metadata"):
258+
try:
259+
metadata.extend(self._response.initial_metadata() or [])
260+
except Exception as e:
261+
logger.warning(f"Failed to retrieve initial metadata: {e}")
262+
if hasattr(self._response, "trailing_metadata"):
263+
try:
264+
metadata.extend(self._response.trailing_metadata() or [])
265+
except Exception as e:
266+
logger.warning(f"Failed to retrieve trailing metadata: {e}")
267+
self._tracer.record_gfe_metrics(metadata)
268+
except Exception as e:
269+
logger.warning(f"Failed to record metrics: {e}")
270+
271+
def __del__(self):
272+
try:
273+
self._record_metrics()
274+
except Exception:
275+
pass
276+
277+
def __getattr__(self, name):
278+
return getattr(self._response, name)
279+
280+
281+
class _AsyncUnaryResponseWrapper(grpc.aio.UnaryUnaryCall):
282+
"""Wrapper for async unary RPC response to defer metrics recording until awaited."""
283+
284+
def __init__(self, response, tracer):
285+
self._response = response
286+
self._tracer = tracer
287+
self._metrics_recorded = False
288+
289+
def add_done_callback(self, *args, **kwargs):
290+
return getattr(self._response, "add_done_callback")(*args, **kwargs)
291+
292+
def cancel(self, *args, **kwargs):
293+
return getattr(self._response, "cancel")(*args, **kwargs)
294+
295+
def cancelled(self, *args, **kwargs):
296+
return getattr(self._response, "cancelled")(*args, **kwargs)
297+
298+
def code(self, *args, **kwargs):
299+
return getattr(self._response, "code")(*args, **kwargs)
300+
301+
def details(self, *args, **kwargs):
302+
return getattr(self._response, "details")(*args, **kwargs)
303+
304+
def done(self, *args, **kwargs):
305+
return getattr(self._response, "done")(*args, **kwargs)
306+
307+
def initial_metadata(self, *args, **kwargs):
308+
return getattr(self._response, "initial_metadata")(*args, **kwargs)
309+
310+
def time_remaining(self, *args, **kwargs):
311+
return getattr(self._response, "time_remaining")(*args, **kwargs)
312+
313+
def trailing_metadata(self, *args, **kwargs):
314+
return getattr(self._response, "trailing_metadata")(*args, **kwargs)
315+
316+
def wait_for_connection(self, *args, **kwargs):
317+
return getattr(self._response, "wait_for_connection")(*args, **kwargs)
318+
319+
def __await__(self):
320+
async def _wait():
321+
try:
322+
return await self._response
323+
finally:
324+
await self._record_metrics()
325+
326+
return _wait().__await__()
327+
328+
async def _record_metrics(self):
329+
if self._metrics_recorded:
330+
return
331+
self._metrics_recorded = True
332+
try:
333+
self._tracer.record_attempt_completion()
334+
metadata = []
335+
if hasattr(self._response, "initial_metadata"):
336+
try:
337+
res = self._response.initial_metadata()
338+
if inspect.isawaitable(res):
339+
res = await res
340+
metadata.extend(res or [])
341+
except Exception as e:
342+
logger.warning(f"Failed to retrieve initial metadata: {e}")
343+
if hasattr(self._response, "trailing_metadata"):
344+
try:
345+
res = self._response.trailing_metadata()
346+
if inspect.isawaitable(res):
347+
res = await res
348+
metadata.extend(res or [])
349+
except Exception as e:
350+
logger.warning(f"Failed to retrieve trailing metadata: {e}")
351+
self._tracer.record_gfe_metrics(metadata)
352+
except Exception as e:
353+
logger.warning(f"Failed to record metrics: {e}")
354+
355+
def __del__(self):
356+
if not self._metrics_recorded:
357+
self._metrics_recorded = True
358+
try:
359+
self._tracer.record_attempt_completion()
360+
except Exception:
361+
pass
362+
363+
def __getattr__(self, name):
364+
return getattr(self._response, name)
365+
366+
367+
class _AsyncStreamingResponseWrapper(
368+
grpc.aio.UnaryStreamCall,
369+
grpc.aio.StreamUnaryCall,
370+
grpc.aio.StreamStreamCall,
371+
):
372+
"""Wrapper for async streaming RPC response iterators to defer metrics recording."""
373+
374+
def __init__(self, response, tracer):
375+
self._response = response
376+
self._tracer = tracer
377+
self._metrics_recorded = False
378+
self._iterator = None
379+
380+
def add_done_callback(self, *args, **kwargs):
381+
return getattr(self._response, "add_done_callback")(*args, **kwargs)
382+
383+
def cancel(self, *args, **kwargs):
384+
return getattr(self._response, "cancel")(*args, **kwargs)
385+
386+
def cancelled(self, *args, **kwargs):
387+
return getattr(self._response, "cancelled")(*args, **kwargs)
388+
389+
def code(self, *args, **kwargs):
390+
return getattr(self._response, "code")(*args, **kwargs)
391+
392+
def details(self, *args, **kwargs):
393+
return getattr(self._response, "details")(*args, **kwargs)
394+
395+
def done(self, *args, **kwargs):
396+
return getattr(self._response, "done")(*args, **kwargs)
397+
398+
def initial_metadata(self, *args, **kwargs):
399+
return getattr(self._response, "initial_metadata")(*args, **kwargs)
400+
401+
def time_remaining(self, *args, **kwargs):
402+
return getattr(self._response, "time_remaining")(*args, **kwargs)
403+
404+
def trailing_metadata(self, *args, **kwargs):
405+
return getattr(self._response, "trailing_metadata")(*args, **kwargs)
406+
407+
def wait_for_connection(self, *args, **kwargs):
408+
return getattr(self._response, "wait_for_connection")(*args, **kwargs)
409+
410+
def read(self, *args, **kwargs):
411+
return getattr(self._response, "read")(*args, **kwargs)
412+
413+
def write(self, *args, **kwargs):
414+
return getattr(self._response, "write")(*args, **kwargs)
415+
416+
def done_writing(self, *args, **kwargs):
417+
return getattr(self._response, "done_writing")(*args, **kwargs)
418+
419+
def __aiter__(self):
420+
if hasattr(self._response, "__aiter__"):
421+
self._iterator = self._response.__aiter__()
422+
else:
423+
self._iterator = self._response
424+
return self
425+
426+
async def __anext__(self):
427+
if self._iterator is None:
428+
if hasattr(self._response, "__aiter__"):
429+
self._iterator = self._response.__aiter__()
430+
else:
431+
self._iterator = self._response
432+
try:
433+
return await self._iterator.__anext__()
434+
except StopAsyncIteration:
435+
await self._record_metrics()
436+
raise
437+
except Exception:
438+
await self._record_metrics()
439+
raise
440+
441+
async def _record_metrics(self):
442+
if self._metrics_recorded:
443+
return
444+
self._metrics_recorded = True
445+
try:
446+
self._tracer.record_attempt_completion()
447+
metadata = []
448+
if hasattr(self._response, "initial_metadata"):
449+
try:
450+
res = self._response.initial_metadata()
451+
if inspect.isawaitable(res):
452+
res = await res
453+
metadata.extend(res or [])
454+
except Exception as e:
455+
logger.warning(f"Failed to retrieve initial metadata: {e}")
456+
if hasattr(self._response, "trailing_metadata"):
457+
try:
458+
res = self._response.trailing_metadata()
459+
if inspect.isawaitable(res):
460+
res = await res
461+
metadata.extend(res or [])
462+
except Exception as e:
463+
logger.warning(f"Failed to retrieve trailing metadata: {e}")
464+
self._tracer.record_gfe_metrics(metadata)
465+
except Exception as e:
466+
logger.warning(f"Failed to record metrics: {e}")
467+
468+
def __del__(self):
469+
if not self._metrics_recorded:
470+
self._metrics_recorded = True
471+
try:
472+
self._tracer.record_attempt_completion()
473+
except Exception:
474+
pass
475+
476+
def __getattr__(self, name):
477+
return getattr(self._response, name)

0 commit comments

Comments
 (0)