use std::{collections::HashMap, sync::Arc};
use async_trait::async_trait;
use crate::{observability::ops_stats::OpsStatsEventObserver, statsig_metadata::StatsigMetadata};
use super::ops_stats::OpsStatsEvent;
static HIGH_CARDINALITY_TAGS: &[&str] = &["lcut", "prev_lcut"];
static SDK_EXCEPTION_COUNT_EXTRA_TAGS: &[&str] = &["request_path", "status_code"];
#[derive(Clone)]
pub enum MetricType {
Increment,
Gauge,
Dist,
}
#[derive(Clone)]
pub struct ObservabilityEvent {
pub metric_type: MetricType,
pub metric_name: String,
pub value: f64,
pub tags: Option<HashMap<String, String>>,
}
impl ObservabilityEvent {
pub fn new_event(
metric_type: MetricType,
metric_name: String,
value: f64,
tags: Option<HashMap<String, String>>,
) -> OpsStatsEvent {
OpsStatsEvent::Observability(ObservabilityEvent {
metric_type,
metric_name,
value,
tags,
})
}
}
fn add_default_metric_tags(mut tags: HashMap<String, String>) -> HashMap<String, String> {
let metadata = StatsigMetadata::get_metadata();
tags.insert("sdk_type".to_string(), metadata.sdk_type);
tags.insert("sdk_version".to_string(), metadata.sdk_version);
tags
}
pub trait ObservabilityClient: Send + Sync + 'static + OpsStatsEventObserver {
fn init(&self);
fn increment(&self, metric_name: String, value: f64, tags: Option<HashMap<String, String>>);
fn gauge(&self, metric_name: String, value: f64, tags: Option<HashMap<String, String>>);
fn dist(&self, metric_name: String, value: f64, tags: Option<HashMap<String, String>>);
fn error(&self, tag: String, error: String);
fn should_enable_high_cardinality_for_this_tag(&self, tag: String) -> Option<bool>;
fn to_ops_stats_event_observer(self: Arc<Self>) -> Arc<dyn OpsStatsEventObserver>;
}
#[async_trait]
impl<T: ObservabilityClient> OpsStatsEventObserver for T {
async fn handle_event(&self, event: OpsStatsEvent) {
match event {
OpsStatsEvent::Observability(data) => {
let tags = data
.tags
.unwrap_or_default()
.into_iter()
.filter(|(key, _)| {
if HIGH_CARDINALITY_TAGS.contains(&key.as_str()) {
self.should_enable_high_cardinality_for_this_tag(key.to_string())
.unwrap_or_default()
} else {
true
}
})
.collect();
let tags = Some(add_default_metric_tags(tags));
let metric_name = format!("statsig.sdk.{}", data.metric_name.clone());
match data.metric_type {
MetricType::Increment => self.increment(metric_name, data.value, tags),
MetricType::Gauge => self.gauge(metric_name, data.value, tags),
MetricType::Dist => self.dist(metric_name, data.value, tags),
};
}
OpsStatsEvent::SDKError(error) => {
self.error(error.tag.clone(), error.info.to_string());
let mut tags = HashMap::from([
("tag".to_string(), error.tag),
("exception".to_string(), error.exception),
]);
if let Some(extra) = error.extra {
for (key, value) in extra {
if SDK_EXCEPTION_COUNT_EXTRA_TAGS.contains(&key.as_str()) {
tags.insert(key, value);
}
}
}
self.increment(
"statsig.sdk.sdk_exceptions_count".to_string(),
1.0,
Some(add_default_metric_tags(tags)),
);
}
_ => {}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn add_default_metric_tags_attaches_sdk_metadata() {
let metadata = StatsigMetadata::get_metadata();
let tags = add_default_metric_tags(HashMap::new());
assert_eq!(tags.get("sdk_type"), Some(&metadata.sdk_type));
assert_eq!(tags.get("sdk_version"), Some(&metadata.sdk_version));
}
#[test]
fn add_default_metric_tags_overwrites_call_site_sdk_metadata() {
let metadata = StatsigMetadata::get_metadata();
let tags = add_default_metric_tags(HashMap::from([
("sdk_type".to_string(), "wrong".to_string()),
("sdk_version".to_string(), "wrong".to_string()),
]));
assert_eq!(tags.get("sdk_type"), Some(&metadata.sdk_type));
assert_eq!(tags.get("sdk_version"), Some(&metadata.sdk_version));
}
}