From 745828301c936ddd1dac49b2611bdf1a4477f9ab Mon Sep 17 00:00:00 2001 From: Etienne Perot Date: Wed, 27 Nov 2024 17:10:00 -0800 Subject: [PATCH] Kubernetes tests: Let the benchmark metric recorder be overridden via context PiperOrigin-RevId: 700842559 --- test/kubernetes/benchmetric/BUILD | 2 +- test/kubernetes/benchmetric/benchmetric.go | 12 ++++++++++++ 2 files changed, 13 insertions(+), 1 deletion(-) diff --git a/test/kubernetes/benchmetric/BUILD b/test/kubernetes/benchmetric/BUILD index f247e5b9f..a31504a38 100644 --- a/test/kubernetes/benchmetric/BUILD +++ b/test/kubernetes/benchmetric/BUILD @@ -2,7 +2,7 @@ load("//tools:defs.bzl", "go_library") package( default_applicable_licenses = ["//:license"], - default_visibility = ["//test/kubernetes:__subpackages__"], + default_visibility = ["//:sandbox"], licenses = ["notice"], ) diff --git a/test/kubernetes/benchmetric/benchmetric.go b/test/kubernetes/benchmetric/benchmetric.go index e6e2c02a1..3f6641442 100644 --- a/test/kubernetes/benchmetric/benchmetric.go +++ b/test/kubernetes/benchmetric/benchmetric.go @@ -178,8 +178,20 @@ var ( recorderOnce sync.Once ) +type recorderContextKeyType int + +const recorderContextKey recorderContextKeyType = iota + +// WithRecorder returns a context with the given `Recorder`. +func WithRecorder(ctx context.Context, recorder Recorder) context.Context { + return context.WithValue(ctx, recorderContextKey, recorder) +} + // GetRecorder returns the benchmark's `Recorder` singleton. func GetRecorder(ctx context.Context) (Recorder, error) { + if ctx.Value(recorderContextKey) != nil { + return ctx.Value(recorderContextKey).(Recorder), nil + } recorderOnce.Do(func() { recorder, recorderErr = recorderFn(ctx) })