Skip to content

Commit bc4ca60

Browse files
committed
fix(metrics): correct GFE metrics extraction and enable by default
1 parent 4f21b8b commit bc4ca60

8 files changed

Lines changed: 707 additions & 25 deletions

File tree

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

Lines changed: 356 additions & 9 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."""
@@ -67,7 +72,8 @@ def _extract_resource_from_path(metadata: Dict[str, str]) -> Dict[str, str]:
6772
resources = MetricsInterceptor._parse_resource_path(path)
6873
return resources
6974

70-
def _set_metrics_tracer_attributes(self, resources: Dict[str, str]) -> None:
75+
@staticmethod
76+
def _set_metrics_tracer_attributes(resources: Dict[str, str]) -> None:
7177
"""
7278
Sets the metric tracer attributes based on the provided resources.
7379
@@ -115,17 +121,358 @@ def intercept(self, invoked_method, request_or_iterator, call_details):
115121
self._set_metrics_tracer_attributes(resources)
116122

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

122129
tracer.set_method(method_name)
123130
tracer.record_attempt_start()
124131
response = invoked_method(request_or_iterator, call_details)
125-
tracer.record_attempt_completion()
126132

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

0 commit comments

Comments
 (0)