1use serde::Deserialize;
4use serde::Deserializer;
5use serde::Serialize;
6use serde::de;
7use serde::de::Unexpected;
8use serde::de::Visitor;
9use std::collections::HashMap;
10use std::fmt;
11use std::sync::{Arc, RwLock};
12use tracing_core::Event;
13use tracing_core::Metadata;
14use tracing_subscriber::filter;
15use tracing_subscriber::layer;
16
17use super::ConfigError;
18
19#[derive(Default, Debug, Serialize, Copy, Clone, PartialEq, PartialOrd)]
21#[repr(u8)]
22pub enum TelemetryLevel {
23 OFF = 0,
25 ERROR = 1,
27 WARN = 2,
29 INFO = 3,
31 DEBUG = 4,
33 #[default]
35 TRACE = 5,
36}
37
38impl From<TelemetryLevel> for filter::LevelFilter {
39 fn from(val: TelemetryLevel) -> Self {
40 match val {
41 TelemetryLevel::OFF => filter::LevelFilter::OFF,
42 TelemetryLevel::ERROR => filter::LevelFilter::ERROR,
43 TelemetryLevel::WARN => filter::LevelFilter::WARN,
44 TelemetryLevel::INFO => filter::LevelFilter::INFO,
45 TelemetryLevel::DEBUG => filter::LevelFilter::DEBUG,
46 TelemetryLevel::TRACE => filter::LevelFilter::TRACE,
47 }
48 }
49}
50
51impl From<TelemetryLevel> for &str {
52 fn from(val: TelemetryLevel) -> Self {
53 match val {
54 TelemetryLevel::OFF => "off",
55 TelemetryLevel::ERROR => "error",
56 TelemetryLevel::WARN => "warn",
57 TelemetryLevel::INFO => "info",
58 TelemetryLevel::DEBUG => "debug",
59 TelemetryLevel::TRACE => "trace",
60 }
61 }
62}
63
64impl TryFrom<&str> for TelemetryLevel {
65 type Error = ConfigError;
66
67 fn try_from(value: &str) -> Result<Self, Self::Error> {
68 match value.to_lowercase().as_str() {
69 "off" => Ok(TelemetryLevel::OFF),
70 "error" => Ok(TelemetryLevel::ERROR),
71 "warn" => Ok(TelemetryLevel::WARN),
72 "info" => Ok(TelemetryLevel::INFO),
73 "debug" => Ok(TelemetryLevel::DEBUG),
74 "trace" => Ok(TelemetryLevel::TRACE),
75 _ => Err(ConfigError::WrongValue(
76 "TelemetryLevel".into(),
77 value.to_string(),
78 )),
79 }
80 }
81}
82
83impl Visitor<'_> for TelemetryLevel {
84 type Value = TelemetryLevel;
85
86 fn expecting(&self, formatter: &mut fmt::Formatter) -> fmt::Result {
87 formatter.write_str(
88 "Telemetry Level from values: off[0], error[1], warn[2], info[3], debug[4], trace[5]",
89 )
90 }
91
92 fn visit_str<E>(self, s: &str) -> Result<Self::Value, E>
93 where
94 E: de::Error,
95 {
96 match s.to_lowercase().as_str() {
97 "off" => Ok(TelemetryLevel::OFF),
98 "error" => Ok(TelemetryLevel::ERROR),
99 "warn" => Ok(TelemetryLevel::WARN),
100 "info" => Ok(TelemetryLevel::INFO),
101 "debug" => Ok(TelemetryLevel::DEBUG),
102 "trace" => Ok(TelemetryLevel::TRACE),
103 _ => Err(de::Error::invalid_value(Unexpected::Str(s), &self)),
104 }
105 }
106
107 fn visit_i64<E>(self, value: i64) -> Result<Self::Value, E>
108 where
109 E: de::Error,
110 {
111 match value {
112 0 => Ok(TelemetryLevel::OFF),
113 1 => Ok(TelemetryLevel::ERROR),
114 2 => Ok(TelemetryLevel::WARN),
115 3 => Ok(TelemetryLevel::INFO),
116 4 => Ok(TelemetryLevel::DEBUG),
117 5 => Ok(TelemetryLevel::TRACE),
118 _ => Err(de::Error::invalid_value(Unexpected::Signed(value), &self)),
119 }
120 }
121
122 fn visit_u64<E>(self, value: u64) -> Result<Self::Value, E>
123 where
124 E: de::Error,
125 {
126 self.visit_i64(value as i64)
127 }
128}
129
130impl<'de> Deserialize<'de> for TelemetryLevel {
131 fn deserialize<D>(deserializer: D) -> Result<TelemetryLevel, D::Error>
132 where
133 D: Deserializer<'de>,
134 {
135 deserializer.deserialize_any(TelemetryLevel::default())
136 }
137}
138
139#[derive(Debug, Clone)]
157pub struct TelemetryFilter {
158 inner: Arc<RwLock<TelemetryFilterInner>>,
159 max_level: Option<filter::LevelFilter>,
160}
161
162#[derive(Debug, Clone)]
163struct TelemetryFilterInner {
164 proc_levels: HashMap<String, filter::LevelFilter>,
165 level: filter::LevelFilter,
166}
167
168impl TelemetryFilter {
169 pub fn new(level: filter::LevelFilter) -> TelemetryFilter {
171 TelemetryFilter {
172 inner: Arc::new(RwLock::new(TelemetryFilterInner {
173 proc_levels: HashMap::new(),
174 level,
175 })),
176 max_level: None,
177 }
178 }
179
180 pub fn clone_with_level(&self, level: TelemetryLevel) -> TelemetryFilter {
182 TelemetryFilter {
183 inner: self.inner.clone(),
184 max_level: Some(level.into()),
185 }
186 }
187
188 pub fn set_level(&self, level: filter::LevelFilter) {
190 let mut level_changed = false;
191 if let Ok(mut inner) = self.inner.write() {
192 level_changed = inner.level != level;
193 inner.level = level;
194 }
195
196 if level_changed {
197 tracing_core::callsite::rebuild_interest_cache();
198 }
199 }
200
201 pub fn level(&self) -> filter::LevelFilter {
203 let mut level = self
204 .inner
205 .read()
206 .map(|inner| inner.level)
207 .unwrap_or(filter::LevelFilter::OFF);
208
209 if let Some(max_level) = self.max_level
210 && max_level < level
211 {
212 level = max_level;
213 }
214
215 level
216 }
217
218 pub fn add_proc_filter(&self, proc_name: String, level: filter::LevelFilter) {
220 if let Ok(mut inner) = self.inner.write() {
221 inner.proc_levels.insert(proc_name, level);
222 }
223 }
224
225 fn is_enabled(&self, metadata: &Metadata<'_>) -> bool {
226 let Ok(inner) = self.inner.read() else {
227 return false;
228 };
229
230 let mut level = if let Some(value) = inner.proc_levels.get(metadata.name()) {
231 *value
232 } else if let Some(value) = inner.proc_levels.get(metadata.target()) {
233 *value
234 } else {
235 inner.level
236 };
237
238 if let Some(max_level) = self.max_level
239 && max_level < level
240 {
241 level = max_level;
242 }
243
244 metadata.level() <= &level
245 }
246}
247
248impl Default for TelemetryFilter {
249 fn default() -> TelemetryFilter {
250 TelemetryFilter::new(filter::LevelFilter::TRACE)
251 }
252}
253
254impl<S> layer::Filter<S> for TelemetryFilter {
255 fn enabled(&self, metadata: &Metadata<'_>, _: &layer::Context<'_, S>) -> bool {
256 self.is_enabled(metadata)
257 }
258
259 fn event_enabled(&self, event: &Event<'_>, _: &layer::Context<'_, S>) -> bool {
260 self.is_enabled(event.metadata())
261 }
262
263 fn max_level_hint(&self) -> Option<filter::LevelFilter> {
264 Some(self.level())
265 }
266}
267
268#[cfg(test)]
269mod tests {
270 use super::*;
271
272 #[test]
273 fn telemetry_level() {
274 assert!(
275 TelemetryLevel::try_from("warn").expect("Warn Telemetry level should exist")
276 < TelemetryLevel::INFO,
277 "{:?} < Info",
278 TelemetryLevel::try_from("warn")
279 );
280 assert_eq!(
281 "The config parameter TelemetryLevel have an incorrect value `wrong`".to_owned(),
282 TelemetryLevel::try_from("wrong")
283 .expect_err("Wrong Telemetry level shouldn't exist")
284 .to_string()
285 );
286
287 assert_eq!(
288 filter::LevelFilter::DEBUG,
289 filter::LevelFilter::from(TelemetryLevel::DEBUG)
290 );
291 }
292}