Skip to main content

openvm_stark_sdk/bench/
mod.rs

1use std::{collections::BTreeMap, ffi::OsStr};
2
3#[cfg(feature = "prometheus")]
4use metrics_exporter_prometheus::PrometheusBuilder;
5use metrics_tracing_context::MetricsLayer;
6use metrics_util::{
7    debugging::{DebugValue, DebuggingRecorder, Snapshot},
8    CompositeKey, MetricKind,
9};
10use serde_json::json;
11use tracing_forest::ForestLayer;
12use tracing_subscriber::{layer::SubscriberExt, EnvFilter, Registry};
13#[cfg(feature = "metrics")]
14use {
15    crate::metrics_tracing::TimingMetricsLayer, metrics_tracing_context::TracingContextLayer,
16    metrics_util::layers::Layer,
17};
18
19#[cfg(feature = "nvtx")]
20use crate::nvtx_tracing::NvtxLayer;
21
22/// Run a function with metric collection enabled. The metrics will be written to a file specified
23/// by an environment variable which name is `output_path_envar`.
24pub fn run_with_metric_collection<R>(
25    output_path_envar: impl AsRef<OsStr>,
26    f: impl FnOnce() -> R,
27) -> R {
28    let file = std::env::var(output_path_envar).map(|path| std::fs::File::create(path).unwrap());
29    // Set up tracing:
30    let env_filter =
31        EnvFilter::try_from_default_env().unwrap_or_else(|_| EnvFilter::new("info,p3_=warn"));
32    // Plonky3 logging is more verbose, so we set default to debug.
33    let subscriber = Registry::default()
34        .with(env_filter)
35        .with(ForestLayer::default())
36        .with(MetricsLayer::new());
37    #[cfg(feature = "metrics")]
38    let subscriber = subscriber.with(TimingMetricsLayer::new());
39    #[cfg(feature = "nvtx")]
40    let subscriber = subscriber.with(NvtxLayer::new(Default::default()));
41    // Prepare tracing.
42    tracing::subscriber::set_global_default(subscriber).unwrap();
43
44    // Prepare metrics.
45    let recorder = DebuggingRecorder::new();
46    let snapshotter = recorder.snapshotter();
47    #[cfg(feature = "metrics")]
48    {
49        let recorder = TracingContextLayer::all().layer(recorder);
50        // Install the registry as the global recorder
51        metrics::set_global_recorder(recorder).unwrap();
52    }
53    let res = f();
54
55    if let Ok(file) = file {
56        serde_json::to_writer_pretty(&file, &serialize_metric_snapshot(snapshotter.snapshot()))
57            .unwrap();
58    }
59    res
60}
61
62/// Run a function with metric exporter enabled. The metrics will be served on the port specified
63/// by an environment variable which name is `metrics_port_envar`.
64#[cfg(feature = "prometheus")]
65pub fn run_with_metric_exporter<R>(
66    metrics_port_envar: impl AsRef<OsStr>,
67    f: impl FnOnce() -> R,
68) -> R {
69    // Get the port from environment variable or use a default
70    let metrics_port = std::env::var(metrics_port_envar)
71        .map(|port| port.parse::<u16>().unwrap_or(9091))
72        .unwrap();
73    let endpoint = format!("http://127.0.0.1:{}/metrics/job/stark-sdk", metrics_port);
74
75    // Clear metrics before pushing to the push gateway
76    let status = std::process::Command::new("curl")
77        .args(["-X", "DELETE", &endpoint])
78        .status()
79        .expect("Failed to clear metrics");
80    if status.success() {
81        println!("Metrics cleared successfully");
82    }
83
84    // Install the default crypto provider
85    rustls::crypto::aws_lc_rs::default_provider()
86        .install_default()
87        .expect("Failed to install default crypto provider");
88
89    // Set up Prometheus recorder and exporter
90    let builder = PrometheusBuilder::new()
91        .with_push_gateway(endpoint, std::time::Duration::from_secs(60), None, None)
92        .expect("Push gateway endpoint should be valid");
93
94    let recorder = if let Ok(handle) = tokio::runtime::Handle::try_current() {
95        let (recorder, exporter) = {
96            let _g = handle.enter();
97            builder.build().unwrap()
98        };
99        handle.spawn(exporter);
100        recorder
101    } else {
102        let thread_name = "metrics-exporter-prometheus-push-gateway";
103        let runtime = tokio::runtime::Builder::new_current_thread()
104            .enable_all()
105            .build()
106            .unwrap();
107        let (recorder, exporter) = {
108            let _g = runtime.enter();
109            builder.build().unwrap()
110        };
111        std::thread::Builder::new()
112            .name(thread_name.to_string())
113            .spawn(move || runtime.block_on(exporter))
114            .unwrap();
115        recorder
116    };
117
118    // Set up tracing:
119    let env_filter =
120        EnvFilter::try_from_default_env().unwrap_or_else(|_| EnvFilter::new("info,p3_=warn"));
121    // Plonky3 logging is more verbose, so we set default to debug.
122    let subscriber = Registry::default()
123        .with(env_filter)
124        .with(ForestLayer::default())
125        .with(MetricsLayer::new());
126    #[cfg(feature = "metrics")]
127    let subscriber = subscriber.with(TimingMetricsLayer::new());
128    #[cfg(feature = "nvtx")]
129    let subscriber = subscriber.with(NvtxLayer::new(Default::default()));
130    // Prepare tracing.
131    tracing::subscriber::set_global_default(subscriber).unwrap();
132
133    // Prepare metrics
134    let recorder = TracingContextLayer::all().layer(recorder);
135    // Install the registry as the global recorder
136    metrics::set_global_recorder(recorder).unwrap();
137
138    // Run the actual function
139    let res = f();
140    std::thread::sleep(std::time::Duration::from_secs(80));
141    println!(
142        "Metrics available at http://127.0.0.1:{}/metrics/job/stark-sdk",
143        metrics_port
144    );
145    res
146}
147
148/// Serialize a gauge/counter metric into a JSON object. The object has the following structure:
149/// {
150///    "metric": <Metric Name>,
151///    "labels": [
152///       (<key1>, <value1>),
153///       (<key2>, <value2>),
154///     ],
155///    "value": <float value if gauge | integer value if counter>
156/// }
157fn serialize_metric(ckey: CompositeKey, value: DebugValue) -> serde_json::Value {
158    let (_kind, key) = ckey.into_parts();
159    let (key_name, labels) = key.into_parts();
160    let value = match value {
161        DebugValue::Gauge(v) => v.into_inner().to_string(),
162        DebugValue::Counter(v) => v.to_string(),
163        DebugValue::Histogram(_) => todo!("Histograms not supported yet."),
164    };
165    let labels = labels
166        .into_iter()
167        .map(|label| {
168            let (k, v) = label.into_parts();
169            (k.as_ref().to_owned(), v.as_ref().to_owned())
170        })
171        .collect::<Vec<_>>();
172
173    json!({
174        "metric": key_name.as_str(),
175        "labels": labels,
176        "value": value,
177    })
178}
179
180/// Serialize a metric snapshot into a JSON object. The object has the following structure:
181/// {
182///   "gauge": [
183///     {
184///         "metric": <Metric Name>,
185///         "labels": [
186///             (<key1>, <value1>),
187///             (<key2>, <value2>),
188///         ],
189///         "value": <float value>
190///     },
191///     ...
192///   ],
193///   ...
194/// }
195pub fn serialize_metric_snapshot(snapshot: Snapshot) -> serde_json::Value {
196    let mut ret = BTreeMap::<_, Vec<serde_json::Value>>::new();
197    for (ckey, _, _, value) in snapshot.into_vec() {
198        match ckey.kind() {
199            MetricKind::Gauge => {
200                ret.entry("gauge")
201                    .or_default()
202                    .push(serialize_metric(ckey, value));
203            }
204            MetricKind::Counter => {
205                ret.entry("counter")
206                    .or_default()
207                    .push(serialize_metric(ckey, value));
208            }
209            MetricKind::Histogram => todo!(),
210        }
211    }
212    json!(ret)
213}