Skip to main content

common_telemetry/
logging.rs

1// Copyright 2023 Greptime Team
2//
3// Licensed under the Apache License, Version 2.0 (the "License");
4// you may not use this file except in compliance with the License.
5// You may obtain a copy of the License at
6//
7//     http://www.apache.org/licenses/LICENSE-2.0
8//
9// Unless required by applicable law or agreed to in writing, software
10// distributed under the License is distributed on an "AS IS" BASIS,
11// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12// See the License for the specific language governing permissions and
13// limitations under the License.
14
15//! logging stuffs, inspired by databend
16mod file_retention;
17
18use std::collections::HashMap;
19use std::env;
20use std::io::IsTerminal;
21use std::sync::atomic::{AtomicBool, Ordering};
22use std::sync::{Arc, Mutex, Once};
23use std::time::Duration;
24
25use common_base::readable_size::ReadableSize;
26use common_base::serde::empty_string_as_default;
27use file_retention::{DirectoryRetention, LogFileKind, build_file_appender};
28use once_cell::sync::{Lazy, OnceCell};
29use opentelemetry::trace::TracerProvider;
30use opentelemetry::{KeyValue, global};
31use opentelemetry_otlp::{Protocol, SpanExporter, WithExportConfig, WithHttpConfig};
32use opentelemetry_sdk::propagation::TraceContextPropagator;
33use opentelemetry_sdk::trace::{Sampler, Tracer};
34use opentelemetry_semantic_conventions::resource;
35use serde::{Deserialize, Serialize};
36use tracing::callsite;
37use tracing::metadata::LevelFilter;
38use tracing_appender::non_blocking::WorkerGuard;
39use tracing_log::LogTracer;
40use tracing_subscriber::filter::{FilterFn, Targets};
41use tracing_subscriber::fmt::Layer;
42use tracing_subscriber::layer::{Layered, SubscriberExt};
43use tracing_subscriber::prelude::*;
44use tracing_subscriber::{EnvFilter, Registry, filter};
45
46use crate::tracing_sampler::{TracingSampleOptions, create_sampler};
47
48/// The default endpoint when use gRPC exporter protocol.
49pub const DEFAULT_OTLP_GRPC_ENDPOINT: &str = "http://localhost:4317";
50
51/// The default endpoint when use HTTP exporter protocol.
52pub const DEFAULT_OTLP_HTTP_ENDPOINT: &str = "http://localhost:4318/v1/traces";
53
54/// The default logs directory.
55pub const DEFAULT_LOGGING_DIR: &str = "logs";
56
57/// Handle for reloading log level
58pub static LOG_RELOAD_HANDLE: OnceCell<tracing_subscriber::reload::Handle<Targets, Registry>> =
59    OnceCell::new();
60
61type DynSubscriber = Layered<tracing_subscriber::reload::Layer<Targets, Registry>, Registry>;
62type OtelTraceLayer = tracing_opentelemetry::OpenTelemetryLayer<DynSubscriber, Tracer>;
63
64struct TraceLayerState {
65    enabled: AtomicBool,
66    layer: OnceCell<OtelTraceLayer>,
67}
68
69#[derive(Clone)]
70pub struct TraceReloadHandle {
71    inner: Arc<TraceLayerState>,
72}
73
74impl TraceReloadHandle {
75    fn new(inner: Arc<TraceLayerState>) -> Self {
76        Self { inner }
77    }
78
79    /// Enables or disables OTLP data collection, initializing it on first enable.
80    /// Disabling stops new spans, events, fields and links. Existing spans still
81    /// finish their lifecycle and export on close.
82    pub fn set_enabled(&self, enabled: bool) -> Result<(), &'static str> {
83        self.set_enabled_with(enabled, || {
84            get_or_init_tracer().map(|tracer| tracing_opentelemetry::layer().with_tracer(tracer))
85        })
86    }
87
88    fn set_enabled_with(
89        &self,
90        enabled: bool,
91        init: impl FnOnce() -> Result<OtelTraceLayer, &'static str>,
92    ) -> Result<(), &'static str> {
93        if enabled {
94            self.inner.layer.get_or_try_init(init)?;
95        }
96        self.inner.enabled.store(enabled, Ordering::Release);
97        callsite::rebuild_interest_cache();
98        Ok(())
99    }
100}
101
102/// An OTLP layer with a runtime switch and a stable address for downcasts.
103struct TraceLayer {
104    inner: Arc<TraceLayerState>,
105}
106
107impl TraceLayer {
108    fn new(initial: Option<OtelTraceLayer>) -> (Self, TraceReloadHandle) {
109        let inner = Arc::new(TraceLayerState {
110            enabled: AtomicBool::new(initial.is_some()),
111            layer: initial.map(OnceCell::with_value).unwrap_or_default(),
112        });
113        (
114            Self {
115                inner: inner.clone(),
116            },
117            TraceReloadHandle::new(inner),
118        )
119    }
120
121    fn with_layer<R>(&self, f: impl FnOnce(&OtelTraceLayer) -> R) -> Option<R> {
122        self.inner.layer.get().map(f)
123    }
124
125    fn is_enabled(&self) -> bool {
126        self.inner.enabled.load(Ordering::Acquire)
127    }
128}
129
130impl tracing_subscriber::Layer<DynSubscriber> for TraceLayer {
131    fn on_register_dispatch(&self, subscriber: &tracing::Dispatch) {
132        let _ = self.with_layer(|layer| layer.on_register_dispatch(subscriber));
133    }
134
135    fn register_callsite(
136        &self,
137        metadata: &'static tracing::Metadata<'static>,
138    ) -> tracing::subscriber::Interest {
139        self.with_layer(|layer| layer.register_callsite(metadata))
140            .unwrap_or_else(tracing::subscriber::Interest::always)
141    }
142
143    fn enabled(
144        &self,
145        metadata: &tracing::Metadata<'_>,
146        ctx: tracing_subscriber::layer::Context<'_, DynSubscriber>,
147    ) -> bool {
148        self.with_layer(|layer| layer.enabled(metadata, ctx))
149            .unwrap_or(true)
150    }
151
152    fn on_new_span(
153        &self,
154        attrs: &tracing::span::Attributes<'_>,
155        id: &tracing::span::Id,
156        ctx: tracing_subscriber::layer::Context<'_, DynSubscriber>,
157    ) {
158        if self.is_enabled() {
159            let _ = self.with_layer(|layer| layer.on_new_span(attrs, id, ctx));
160        }
161    }
162
163    fn max_level_hint(&self) -> Option<LevelFilter> {
164        self.with_layer(|layer| layer.max_level_hint()).flatten()
165    }
166
167    fn on_record(
168        &self,
169        span: &tracing::span::Id,
170        values: &tracing::span::Record<'_>,
171        ctx: tracing_subscriber::layer::Context<'_, DynSubscriber>,
172    ) {
173        if self.is_enabled() {
174            let _ = self.with_layer(|layer| layer.on_record(span, values, ctx));
175        }
176    }
177
178    fn on_follows_from(
179        &self,
180        span: &tracing::span::Id,
181        follows: &tracing::span::Id,
182        ctx: tracing_subscriber::layer::Context<'_, DynSubscriber>,
183    ) {
184        if self.is_enabled() {
185            let _ = self.with_layer(|layer| layer.on_follows_from(span, follows, ctx));
186        }
187    }
188
189    fn event_enabled(
190        &self,
191        event: &tracing::Event<'_>,
192        ctx: tracing_subscriber::layer::Context<'_, DynSubscriber>,
193    ) -> bool {
194        self.with_layer(|layer| layer.event_enabled(event, ctx))
195            .unwrap_or(true)
196    }
197
198    fn on_event(
199        &self,
200        event: &tracing::Event<'_>,
201        ctx: tracing_subscriber::layer::Context<'_, DynSubscriber>,
202    ) {
203        if self.is_enabled() {
204            let _ = self.with_layer(|layer| layer.on_event(event, ctx));
205        }
206    }
207
208    fn on_enter(
209        &self,
210        id: &tracing::span::Id,
211        ctx: tracing_subscriber::layer::Context<'_, DynSubscriber>,
212    ) {
213        let _ = self.with_layer(|layer| layer.on_enter(id, ctx));
214    }
215
216    fn on_exit(
217        &self,
218        id: &tracing::span::Id,
219        ctx: tracing_subscriber::layer::Context<'_, DynSubscriber>,
220    ) {
221        let _ = self.with_layer(|layer| layer.on_exit(id, ctx));
222    }
223
224    fn on_close(
225        &self,
226        id: tracing::span::Id,
227        ctx: tracing_subscriber::layer::Context<'_, DynSubscriber>,
228    ) {
229        let _ = self.with_layer(|layer| layer.on_close(id, ctx));
230    }
231
232    fn on_id_change(
233        &self,
234        old: &tracing::span::Id,
235        new: &tracing::span::Id,
236        ctx: tracing_subscriber::layer::Context<'_, DynSubscriber>,
237    ) {
238        let _ = self.with_layer(|layer| layer.on_id_change(old, new, ctx));
239    }
240
241    unsafe fn downcast_raw(&self, id: std::any::TypeId) -> Option<*const ()> {
242        // Keep downcasts available while disabled: an in-flight WithContext
243        // callback may still need the layer. OnceCell keeps both addresses valid
244        // for the subscriber's lifetime, even across concurrent toggles.
245        self.inner
246            .layer
247            .get()
248            .and_then(|layer| unsafe { layer.downcast_raw(id) })
249    }
250}
251
252/// Handle for reloading trace level
253pub static TRACE_RELOAD_HANDLE: OnceCell<TraceReloadHandle> = OnceCell::new();
254
255static TRACER: OnceCell<Mutex<TraceState>> = OnceCell::new();
256
257#[derive(Debug)]
258enum TraceState {
259    Ready(Tracer),
260    Deferred(TraceContext),
261}
262
263/// The logging options that used to initialize the logger.
264#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
265#[serde(default)]
266pub struct LoggingOptions {
267    /// The directory to store log files. If not set, logs will be written to stdout.
268    pub dir: String,
269
270    /// The log level that can be one of "trace", "debug", "info", "warn", "error". Default is "info".
271    pub level: Option<String>,
272
273    /// The log format that can be one of "json" or "text". Default is "text".
274    #[serde(default, deserialize_with = "empty_string_as_default")]
275    pub log_format: LogFormat,
276
277    /// The maximum number of log files set by default.
278    pub max_log_files: usize,
279
280    /// The maximum total size of managed log files in `dir`. Zero disables size-based retention.
281    pub max_log_dir_size: ReadableSize,
282
283    /// Whether to append logs to stdout. Default is true.
284    pub append_stdout: bool,
285
286    /// Whether to write logs to files in `dir`. Default is true.
287    pub enable_file_logging: bool,
288
289    /// Whether to enable tracing with OTLP. Default is false.
290    pub enable_otlp_tracing: bool,
291
292    /// The endpoint of OTLP.
293    pub otlp_endpoint: Option<String>,
294
295    /// The tracing sample ratio.
296    pub tracing_sample_ratio: Option<TracingSampleOptions>,
297
298    /// The protocol of OTLP export.
299    pub otlp_export_protocol: Option<OtlpExportProtocol>,
300
301    /// Additional HTTP headers for OTLP exporter.
302    #[serde(skip_serializing_if = "HashMap::is_empty")]
303    pub otlp_headers: HashMap<String, String>,
304
305    /// Whether to enable per-region metrics.
306    pub enable_per_region_metrics: bool,
307}
308
309/// The protocol of OTLP export.
310#[derive(Clone, Debug, Serialize, Deserialize, PartialEq)]
311#[serde(rename_all = "snake_case")]
312pub enum OtlpExportProtocol {
313    /// GRPC protocol.
314    Grpc,
315
316    /// HTTP protocol with binary protobuf.
317    Http,
318}
319
320/// The options of slow query.
321#[derive(Clone, Debug, Serialize, Deserialize, PartialEq)]
322#[serde(default)]
323pub struct SlowQueryOptions {
324    /// Whether to enable slow query log.
325    pub enable: bool,
326
327    /// The record type of slow queries.
328    #[serde(deserialize_with = "empty_string_as_default")]
329    pub record_type: SlowQueriesRecordType,
330
331    /// The threshold of slow queries.
332    #[serde(with = "humantime_serde")]
333    pub threshold: Duration,
334
335    /// The sample ratio of slow queries.
336    pub sample_ratio: f64,
337
338    /// The table TTL of `slow_queries` system table. Default is "90d".
339    /// It's used when `record_type` is `SystemTable`.
340    #[serde(with = "humantime_serde")]
341    pub ttl: Duration,
342}
343
344impl Default for SlowQueryOptions {
345    fn default() -> Self {
346        Self {
347            enable: true,
348            record_type: SlowQueriesRecordType::SystemTable,
349            threshold: Duration::from_secs(30),
350            sample_ratio: 1.0,
351            ttl: Duration::from_secs(90 * 86400),
352        }
353    }
354}
355
356#[derive(Clone, Debug, Serialize, Deserialize, Copy, PartialEq, Default)]
357#[serde(rename_all = "snake_case")]
358pub enum SlowQueriesRecordType {
359    /// Record the slow query in the system table.
360    #[default]
361    SystemTable,
362    /// Record the slow query in a specific logs file.
363    Log,
364}
365
366#[derive(Clone, Debug, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
367#[serde(rename_all = "snake_case")]
368pub enum LogFormat {
369    Json,
370    #[default]
371    Text,
372}
373
374#[derive(Clone, Debug)]
375struct TraceContext {
376    app_name: String,
377    node_id: String,
378    logging_opts: LoggingOptions,
379}
380
381impl Default for LoggingOptions {
382    fn default() -> Self {
383        Self {
384            // The directory path will be configured at application startup, typically using the data home directory as a base.
385            dir: "".to_string(),
386            level: None,
387            log_format: LogFormat::Text,
388            enable_otlp_tracing: false,
389            otlp_endpoint: None,
390            tracing_sample_ratio: None,
391            append_stdout: true,
392            enable_file_logging: true,
393            // Rotation hourly, 24 files per day, keeps info log files of 30 days
394            max_log_files: 720,
395            max_log_dir_size: ReadableSize::default(),
396            otlp_export_protocol: None,
397            otlp_headers: HashMap::new(),
398            enable_per_region_metrics: false,
399        }
400    }
401}
402
403#[derive(Default, Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
404pub struct TracingOptions {
405    #[cfg(feature = "tokio-console")]
406    pub tokio_console_addr: Option<String>,
407}
408
409/// Init tracing for unittest.
410/// Write logs to file `unittest`.
411pub fn init_default_ut_logging() {
412    static START: Once = Once::new();
413
414    START.call_once(|| {
415        let mut g = GLOBAL_UT_LOG_GUARD.as_ref().lock().unwrap();
416
417        // When running in Github's actions, env "UNITTEST_LOG_DIR" is set to a directory other
418        // than "/tmp".
419        // This is to fix the problem that the "/tmp" disk space of action runner's is small,
420        // if we write testing logs in it, actions would fail due to disk out of space error.
421        let dir =
422            env::var("UNITTEST_LOG_DIR").unwrap_or_else(|_| "/tmp/__unittest_logs".to_string());
423
424        let level = env::var("UNITTEST_LOG_LEVEL").unwrap_or_else(|_|
425            "debug,hyper=warn,tower=warn,datafusion=warn,reqwest=warn,sqlparser=warn,h2=info,opendal=info,rskafka=info".to_string()
426        );
427        let opts = LoggingOptions {
428            dir: dir.clone(),
429            level: Some(level),
430            ..Default::default()
431        };
432        *g = Some(init_global_logging(
433            "unittest",
434            &opts,
435            &TracingOptions::default(),
436            None,
437            None,
438        ));
439
440        crate::info!("logs dir = {}", dir);
441    });
442}
443
444static GLOBAL_UT_LOG_GUARD: Lazy<Arc<Mutex<Option<Vec<WorkerGuard>>>>> =
445    Lazy::new(|| Arc::new(Mutex::new(None)));
446
447const DEFAULT_LOG_TARGETS: &str = "info";
448
449#[allow(clippy::print_stdout)]
450pub fn init_global_logging(
451    app_name: &str,
452    opts: &LoggingOptions,
453    tracing_opts: &TracingOptions,
454    node_id: Option<String>,
455    slow_query_opts: Option<&SlowQueryOptions>,
456) -> Vec<WorkerGuard> {
457    static START: Once = Once::new();
458    let mut guards = vec![];
459    let node_id = node_id.unwrap_or_else(|| "none".to_string());
460
461    START.call_once(|| {
462        // Enable log compatible layer to convert log record to tracing span.
463        LogTracer::init().expect("log tracer must be valid");
464
465        // Configure the stdout logging layer.
466        let stdout_logging_layer = if opts.append_stdout {
467            let (writer, guard) = tracing_appender::non_blocking(std::io::stdout());
468            guards.push(guard);
469
470            if opts.log_format == LogFormat::Json {
471                Some(
472                    Layer::new()
473                        .json()
474                        .with_writer(writer)
475                        .with_ansi(std::io::stdout().is_terminal())
476                        .boxed(),
477                )
478            } else {
479                Some(
480                    Layer::new()
481                        .with_writer(writer)
482                        .with_ansi(std::io::stdout().is_terminal())
483                        .boxed(),
484                )
485            }
486        } else {
487            None
488        };
489
490        let file_logging_enabled = opts.enable_file_logging && !opts.dir.is_empty();
491
492        let retention = file_logging_enabled
493            .then(|| {
494                DirectoryRetention::new(opts.dir.clone(), opts.max_log_dir_size, opts.max_log_files)
495            })
496            .flatten();
497
498        // Configure the file logging layer with rolling policy.
499        let file_logging_layer = if file_logging_enabled {
500            let rolling_appender =
501                build_file_appender(opts, LogFileKind::Default, retention.as_ref());
502            let (writer, guard) = tracing_appender::non_blocking(rolling_appender);
503            guards.push(guard);
504
505            if opts.log_format == LogFormat::Json {
506                Some(
507                    Layer::new()
508                        .json()
509                        .with_writer(writer)
510                        .with_ansi(false)
511                        .boxed(),
512                )
513            } else {
514                Some(Layer::new().with_writer(writer).with_ansi(false).boxed())
515            }
516        } else {
517            None
518        };
519
520        // Configure the error file logging layer with rolling policy.
521        let err_file_logging_layer = if file_logging_enabled {
522            let rolling_appender =
523                build_file_appender(opts, LogFileKind::Error, retention.as_ref());
524            let (writer, guard) = tracing_appender::non_blocking(rolling_appender);
525            guards.push(guard);
526
527            if opts.log_format == LogFormat::Json {
528                Some(
529                    Layer::new()
530                        .json()
531                        .with_writer(writer)
532                        .with_ansi(false)
533                        .with_filter(filter::LevelFilter::ERROR)
534                        .boxed(),
535                )
536            } else {
537                Some(
538                    Layer::new()
539                        .with_writer(writer)
540                        .with_ansi(false)
541                        .with_filter(filter::LevelFilter::ERROR)
542                        .boxed(),
543                )
544            }
545        } else {
546            None
547        };
548
549        let slow_query_logging_layer =
550            build_slow_query_logger(opts, slow_query_opts, retention.as_ref(), &mut guards);
551
552        if let Some(retention) = &retention {
553            retention.initialize();
554        }
555
556        // resolve log level settings from:
557        // - options from command line or config files
558        // - environment variable: RUST_LOG
559        // - default settings
560        let filter = opts
561            .level
562            .as_deref()
563            .or(env::var(EnvFilter::DEFAULT_ENV).ok().as_deref())
564            .unwrap_or(DEFAULT_LOG_TARGETS)
565            .parse::<filter::Targets>()
566            .expect("error parsing log level string");
567
568        let (dyn_filter, reload_handle) = tracing_subscriber::reload::Layer::new(filter.clone());
569
570        LOG_RELOAD_HANDLE
571            .set(reload_handle)
572            .expect("reload handle already set, maybe init_global_logging get called twice?");
573
574        let mut initial_tracer = None;
575        let trace_state = if opts.enable_otlp_tracing {
576            let tracer = create_tracer(app_name, &node_id, opts);
577            initial_tracer = Some(tracer.clone());
578            TraceState::Ready(tracer)
579        } else {
580            TraceState::Deferred(TraceContext {
581                app_name: app_name.to_string(),
582                node_id: node_id.clone(),
583                logging_opts: opts.clone(),
584            })
585        };
586
587        TRACER
588            .set(Mutex::new(trace_state))
589            .expect("trace state already initialized");
590
591        let initial_trace_layer = initial_tracer
592            .as_ref()
593            .map(|tracer| tracing_opentelemetry::layer().with_tracer(tracer.clone()));
594
595        let (dyn_trace_layer, trace_reload_handle) = TraceLayer::new(initial_trace_layer);
596
597        TRACE_RELOAD_HANDLE
598            .set(trace_reload_handle)
599            .unwrap_or_else(|_| panic!("failed to set trace reload handle"));
600
601        // Must enable 'tokio_unstable' cfg to use this feature.
602        // For example: `RUSTFLAGS="--cfg tokio_unstable" cargo run -F common-telemetry/console -- standalone start`
603        #[cfg(feature = "tokio-console")]
604        let subscriber = {
605            let tokio_console_layer =
606                if let Some(tokio_console_addr) = &tracing_opts.tokio_console_addr {
607                    let addr: std::net::SocketAddr = tokio_console_addr.parse().unwrap_or_else(|e| {
608                    panic!("Invalid binding address '{tokio_console_addr}' for tokio-console: {e}");
609                });
610                    println!("tokio-console listening on {addr}");
611
612                    Some(
613                        console_subscriber::ConsoleLayer::builder()
614                            .server_addr(addr)
615                            .spawn(),
616                    )
617                } else {
618                    None
619                };
620
621            Registry::default()
622                .with(dyn_filter)
623                .with(dyn_trace_layer)
624                .with(tokio_console_layer)
625                .with(stdout_logging_layer)
626                .with(file_logging_layer)
627                .with(err_file_logging_layer)
628                .with(slow_query_logging_layer)
629        };
630
631        // consume the `tracing_opts` to avoid "unused" warnings.
632        let _ = tracing_opts;
633
634        #[cfg(not(feature = "tokio-console"))]
635        let subscriber = Registry::default()
636            .with(dyn_filter)
637            .with(dyn_trace_layer)
638            .with(stdout_logging_layer)
639            .with(file_logging_layer)
640            .with(err_file_logging_layer)
641            .with(slow_query_logging_layer);
642
643        global::set_text_map_propagator(TraceContextPropagator::new());
644
645        tracing::subscriber::set_global_default(subscriber)
646            .expect("error setting global tracing subscriber");
647    });
648
649    guards
650}
651
652fn create_tracer(app_name: &str, node_id: &str, opts: &LoggingOptions) -> Tracer {
653    let sampler = opts
654        .tracing_sample_ratio
655        .as_ref()
656        .map(create_sampler)
657        .map(Sampler::ParentBased)
658        .unwrap_or(Sampler::ParentBased(Box::new(Sampler::AlwaysOn)));
659
660    let resource = opentelemetry_sdk::Resource::builder_empty()
661        .with_attributes([
662            KeyValue::new(resource::SERVICE_NAME, app_name.to_string()),
663            KeyValue::new(resource::SERVICE_INSTANCE_ID, node_id.to_string()),
664            KeyValue::new(resource::SERVICE_VERSION, common_version::version()),
665            KeyValue::new(resource::PROCESS_PID, std::process::id().to_string()),
666        ])
667        .build();
668
669    opentelemetry_sdk::trace::SdkTracerProvider::builder()
670        .with_batch_exporter(build_otlp_exporter(opts))
671        .with_sampler(sampler)
672        .with_resource(resource)
673        .build()
674        .tracer("greptimedb")
675}
676
677/// Ensure that the OTLP tracer has been constructed, building it lazily if needed.
678pub fn get_or_init_tracer() -> Result<Tracer, &'static str> {
679    let state = TRACER.get().ok_or("trace state is not initialized")?;
680    let mut guard = state.lock().expect("trace state lock poisoned");
681
682    match &mut *guard {
683        TraceState::Ready(tracer) => Ok(tracer.clone()),
684        TraceState::Deferred(context) => {
685            let tracer = create_tracer(&context.app_name, &context.node_id, &context.logging_opts);
686            *guard = TraceState::Ready(tracer.clone());
687            Ok(tracer)
688        }
689    }
690}
691
692fn build_otlp_exporter(opts: &LoggingOptions) -> SpanExporter {
693    let protocol = opts
694        .otlp_export_protocol
695        .clone()
696        .unwrap_or(OtlpExportProtocol::Http);
697
698    let endpoint = opts
699        .otlp_endpoint
700        .as_ref()
701        .map(|e| {
702            if e.starts_with("http") {
703                e.clone()
704            } else {
705                format!("http://{}", e)
706            }
707        })
708        .unwrap_or_else(|| match protocol {
709            OtlpExportProtocol::Grpc => DEFAULT_OTLP_GRPC_ENDPOINT.to_string(),
710            OtlpExportProtocol::Http => DEFAULT_OTLP_HTTP_ENDPOINT.to_string(),
711        });
712
713    match protocol {
714        OtlpExportProtocol::Grpc => SpanExporter::builder()
715            .with_tonic()
716            .with_endpoint(endpoint)
717            .build()
718            .expect("Failed to create OTLP gRPC exporter "),
719
720        OtlpExportProtocol::Http => SpanExporter::builder()
721            .with_http()
722            .with_endpoint(endpoint)
723            .with_protocol(Protocol::HttpBinary)
724            .with_headers(opts.otlp_headers.clone())
725            .build()
726            .expect("Failed to create OTLP HTTP exporter "),
727    }
728}
729
730fn build_slow_query_logger<S>(
731    opts: &LoggingOptions,
732    slow_query_opts: Option<&SlowQueryOptions>,
733    retention: Option<&DirectoryRetention>,
734    guards: &mut Vec<WorkerGuard>,
735) -> Option<Box<dyn tracing_subscriber::Layer<S> + Send + Sync + 'static>>
736where
737    S: tracing::Subscriber
738        + Send
739        + 'static
740        + for<'span> tracing_subscriber::registry::LookupSpan<'span>,
741{
742    if let Some(slow_query_opts) = slow_query_opts {
743        if opts.enable_file_logging
744            && !opts.dir.is_empty()
745            && slow_query_opts.enable
746            && slow_query_opts.record_type == SlowQueriesRecordType::Log
747        {
748            let rolling_appender = build_file_appender(opts, LogFileKind::SlowQuery, retention);
749            let (writer, guard) = tracing_appender::non_blocking(rolling_appender);
750            guards.push(guard);
751
752            // Only logs if the field contains "slow".
753            let slow_query_filter = FilterFn::new(|metadata| {
754                metadata
755                    .fields()
756                    .iter()
757                    .any(|field| field.name().contains("slow"))
758            });
759
760            if opts.log_format == LogFormat::Json {
761                Some(
762                    Layer::new()
763                        .json()
764                        .with_writer(writer)
765                        .with_ansi(false)
766                        .with_filter(slow_query_filter)
767                        .boxed(),
768                )
769            } else {
770                Some(
771                    Layer::new()
772                        .with_writer(writer)
773                        .with_ansi(false)
774                        .with_filter(slow_query_filter)
775                        .boxed(),
776                )
777            }
778        } else {
779            None
780        }
781    } else {
782        None
783    }
784}
785
786#[cfg(test)]
787mod tests {
788    use std::sync::Barrier;
789    use std::sync::atomic::AtomicUsize;
790
791    use opentelemetry::trace::TraceContextExt;
792    use opentelemetry_sdk::trace::SdkTracerProvider;
793    use tracing_opentelemetry::OpenTelemetrySpanExt;
794
795    use super::*;
796
797    #[test]
798    fn test_trace_switch_preserves_active_spans() {
799        let provider = SdkTracerProvider::builder()
800            .with_sampler(Sampler::AlwaysOn)
801            .build();
802        for initially_enabled in [false, true] {
803            let new_layer =
804                || tracing_opentelemetry::layer().with_tracer(provider.tracer("switch"));
805            let (filter, _) = tracing_subscriber::reload::Layer::new(
806                Targets::new().with_default(tracing::Level::INFO),
807            );
808            let (layer, handle) = TraceLayer::new(initially_enabled.then(new_layer));
809            let dispatch = tracing::Dispatch::new(Registry::default().with(filter).with(layer));
810            tracing::dispatcher::with_default(&dispatch, || {
811                let initial = tracing::info_span!("initial");
812                assert!(!initial.is_disabled());
813                assert_eq!(
814                    initial.context().span().span_context().is_valid(),
815                    initially_enabled
816                );
817                handle.set_enabled_with(true, || Ok(new_layer())).unwrap();
818                let parent = tracing::info_span!("parent");
819                let parent_context = parent.context();
820                assert!(parent_context.span().span_context().is_valid());
821                let previous_context = opentelemetry::Context::current();
822                let entered = parent.enter();
823                handle.set_enabled(false).unwrap();
824                let disabled = tracing::info_span!("disabled");
825                assert!(!disabled.is_disabled());
826                assert!(!disabled.context().span().span_context().is_valid());
827                assert_eq!(
828                    parent.context().span().span_context(),
829                    parent_context.span().span_context()
830                );
831                assert!(dispatch.downcast_ref::<OtelTraceLayer>().is_some());
832                drop(entered);
833                assert_eq!(
834                    opentelemetry::Context::current().span().span_context(),
835                    previous_context.span().span_context()
836                );
837                let disabled_entered = disabled.enter();
838                handle
839                    .set_enabled_with(true, || panic!("layer must not be replaced"))
840                    .unwrap();
841                drop(disabled_entered);
842                let child = tracing::info_span!(parent: &parent, "child");
843                assert_eq!(
844                    child.context().span().span_context().trace_id(),
845                    parent_context.span().span_context().trace_id()
846                );
847                assert!(!disabled.context().span().span_context().is_valid());
848            });
849        }
850    }
851
852    #[test]
853    fn test_trace_switch_retries_initialization() {
854        let (filter, _) = tracing_subscriber::reload::Layer::new(
855            Targets::new().with_default(tracing::Level::INFO),
856        );
857        let (layer, handle) = TraceLayer::new(None);
858        let subscriber = Registry::default().with(filter).with(layer);
859        tracing::subscriber::with_default(subscriber, || {
860            handle
861                .set_enabled_with(false, || panic!("disabling must not initialize OTLP"))
862                .unwrap();
863            assert_eq!(
864                handle.set_enabled_with(true, || Err("initialization failed")),
865                Err("initialization failed")
866            );
867            let failed = tracing::info_span!("after_failed_enable");
868            assert!(!failed.context().span().span_context().is_valid());
869            let provider = SdkTracerProvider::builder()
870                .with_sampler(Sampler::AlwaysOn)
871                .build();
872            handle
873                .set_enabled_with(true, || {
874                    Ok(tracing_opentelemetry::layer().with_tracer(provider.tracer("retry")))
875                })
876                .unwrap();
877            let enabled = tracing::info_span!("after_successful_enable");
878            assert!(enabled.context().span().span_context().is_valid());
879        });
880    }
881
882    #[test]
883    fn test_trace_switch_stops_collecting_fields_when_disabled() {
884        struct CountFormatting(AtomicUsize);
885
886        impl std::fmt::Debug for CountFormatting {
887            fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
888                self.0.fetch_add(1, Ordering::Relaxed);
889                f.write_str("value")
890            }
891        }
892
893        let value = CountFormatting(AtomicUsize::new(0));
894        let provider = SdkTracerProvider::builder()
895            .with_sampler(Sampler::AlwaysOn)
896            .build();
897        let (filter, _) = tracing_subscriber::reload::Layer::new(
898            Targets::new().with_default(tracing::Level::INFO),
899        );
900        let (layer, handle) = TraceLayer::new(Some(
901            tracing_opentelemetry::layer().with_tracer(provider.tracer("fields")),
902        ));
903        tracing::subscriber::with_default(Registry::default().with(filter).with(layer), || {
904            let admitted = tracing::info_span!("admitted", field = tracing::field::Empty);
905            admitted.record("field", tracing::field::debug(&value));
906            admitted.in_scope(|| tracing::info!(value = ?value));
907            assert_eq!(value.0.load(Ordering::Relaxed), 2);
908
909            handle.set_enabled(false).unwrap();
910            let unadmitted = tracing::info_span!("unadmitted", field = tracing::field::Empty);
911            for span in [&admitted, &unadmitted] {
912                span.record("field", tracing::field::debug(&value));
913                span.in_scope(|| tracing::info!(value = ?value));
914            }
915            assert_eq!(value.0.load(Ordering::Relaxed), 2);
916
917            handle
918                .set_enabled_with(true, || panic!("layer must not be replaced"))
919                .unwrap();
920            admitted.record("field", tracing::field::debug(&value));
921            admitted.in_scope(|| tracing::info!(value = ?value));
922            assert_eq!(value.0.load(Ordering::Relaxed), 4);
923        });
924    }
925
926    #[test]
927    fn test_trace_switch_finishes_admitted_spans() {
928        #[derive(Debug)]
929        struct Exporter(std::sync::mpsc::Sender<String>);
930
931        impl opentelemetry_sdk::trace::SpanExporter for Exporter {
932            async fn export(
933                &self,
934                batch: Vec<opentelemetry_sdk::trace::SpanData>,
935            ) -> opentelemetry_sdk::error::OTelSdkResult {
936                for span in batch {
937                    self.0.send(span.name.into_owned()).unwrap();
938                }
939                Ok(())
940            }
941        }
942
943        let (sender, receiver) = std::sync::mpsc::channel();
944        let provider = SdkTracerProvider::builder()
945            .with_sampler(Sampler::AlwaysOn)
946            .with_simple_exporter(Exporter(sender))
947            .build();
948        let (filter, _) = tracing_subscriber::reload::Layer::new(
949            Targets::new().with_default(tracing::Level::INFO),
950        );
951        let (layer, handle) = TraceLayer::new(Some(
952            tracing_opentelemetry::layer().with_tracer(provider.tracer("export")),
953        ));
954        tracing::subscriber::with_default(Registry::default().with(filter).with(layer), || {
955            let active = tracing::info_span!("admitted");
956            handle.set_enabled(false).unwrap();
957            let disabled = tracing::info_span!("not_admitted");
958            drop(active);
959            drop(disabled);
960        });
961        provider.force_flush().unwrap();
962        assert_eq!(receiver.try_iter().collect::<Vec<_>>(), ["admitted"]);
963    }
964
965    #[test]
966    fn test_trace_switch_concurrent_context_access() {
967        let provider = SdkTracerProvider::builder()
968            .with_sampler(Sampler::AlwaysOn)
969            .build();
970        let (filter, _) = tracing_subscriber::reload::Layer::new(
971            Targets::new().with_default(tracing::Level::INFO),
972        );
973        let (layer, handle) = TraceLayer::new(Some(
974            tracing_opentelemetry::layer().with_tracer(provider.tracer("concurrent")),
975        ));
976        let dispatch = tracing::Dispatch::new(Registry::default().with(filter).with(layer));
977        let parent = tracing::dispatcher::with_default(&dispatch, || {
978            tracing::info_span!("concurrent_parent")
979        });
980        let context = parent.context();
981        let barrier = Barrier::new(4);
982        std::thread::scope(|scope| {
983            for _ in 0..3 {
984                scope.spawn(|| {
985                    tracing::dispatcher::with_default(&dispatch, || {
986                        let _entered = parent.enter();
987                        barrier.wait();
988                        for _ in 0..500 {
989                            assert_eq!(
990                                parent.context().span().span_context(),
991                                context.span().span_context()
992                            );
993                            {
994                                let child = tracing::info_span!("concurrent_child");
995                                let _entered = child.enter();
996                                let _ = child.context();
997                                tracing::info!("concurrent event");
998                            }
999                            assert_eq!(
1000                                opentelemetry::Context::current().span().span_context(),
1001                                context.span().span_context()
1002                            );
1003                        }
1004                    });
1005                });
1006            }
1007            barrier.wait();
1008            for i in 0..1000 {
1009                handle
1010                    .set_enabled_with(i % 2 == 0, || panic!("layer must not be replaced"))
1011                    .unwrap();
1012            }
1013        });
1014    }
1015
1016    #[test]
1017    fn test_logging_options_deserialization_default() {
1018        let json = r#"{}"#;
1019        let opts: LoggingOptions = serde_json::from_str(json).unwrap();
1020
1021        assert_eq!(opts.log_format, LogFormat::Text);
1022        assert_eq!(opts.dir, "");
1023        assert_eq!(opts.level, None);
1024        assert!(opts.append_stdout);
1025        assert!(opts.enable_file_logging);
1026    }
1027
1028    #[test]
1029    fn test_logging_options_deserialization_enable_file_logging() {
1030        let json = r#"{"enable_file_logging": false}"#;
1031        let opts: LoggingOptions = serde_json::from_str(json).unwrap();
1032
1033        assert!(!opts.enable_file_logging);
1034    }
1035
1036    #[test]
1037    fn test_logging_options_deserialization_max_log_dir_size() {
1038        let json = r#"{"max_log_dir_size": "1MiB"}"#;
1039        let opts: LoggingOptions = serde_json::from_str(json).unwrap();
1040
1041        assert_eq!(opts.max_log_dir_size, ReadableSize::mb(1));
1042    }
1043
1044    #[test]
1045    fn test_logging_options_deserialization_empty_log_format() {
1046        let json = r#"{"log_format": ""}"#;
1047        let opts: LoggingOptions = serde_json::from_str(json).unwrap();
1048
1049        // Empty string should use default (Text)
1050        assert_eq!(opts.log_format, LogFormat::Text);
1051    }
1052
1053    #[test]
1054    fn test_logging_options_deserialization_valid_log_format() {
1055        let json_format = r#"{"log_format": "json"}"#;
1056        let opts: LoggingOptions = serde_json::from_str(json_format).unwrap();
1057        assert_eq!(opts.log_format, LogFormat::Json);
1058
1059        let text_format = r#"{"log_format": "text"}"#;
1060        let opts: LoggingOptions = serde_json::from_str(text_format).unwrap();
1061        assert_eq!(opts.log_format, LogFormat::Text);
1062    }
1063
1064    #[test]
1065    fn test_logging_options_deserialization_missing_log_format() {
1066        let json = r#"{"dir": "/tmp/logs"}"#;
1067        let opts: LoggingOptions = serde_json::from_str(json).unwrap();
1068
1069        // Missing log_format should use default (Text)
1070        assert_eq!(opts.log_format, LogFormat::Text);
1071        assert_eq!(opts.dir, "/tmp/logs");
1072    }
1073
1074    #[test]
1075    fn test_slow_query_options_deserialization_default() {
1076        let json = r#"{"enable": true, "threshold": "30s"}"#;
1077        let opts: SlowQueryOptions = serde_json::from_str(json).unwrap();
1078
1079        assert_eq!(opts.record_type, SlowQueriesRecordType::SystemTable);
1080        assert!(opts.enable);
1081    }
1082
1083    #[test]
1084    fn test_slow_query_options_deserialization_empty_record_type() {
1085        let json = r#"{"enable": true, "record_type": "", "threshold": "30s"}"#;
1086        let opts: SlowQueryOptions = serde_json::from_str(json).unwrap();
1087
1088        // Empty string should use default (SystemTable)
1089        assert_eq!(opts.record_type, SlowQueriesRecordType::SystemTable);
1090        assert!(opts.enable);
1091    }
1092
1093    #[test]
1094    fn test_slow_query_options_deserialization_valid_record_type() {
1095        let system_table_json =
1096            r#"{"enable": true, "record_type": "system_table", "threshold": "30s"}"#;
1097        let opts: SlowQueryOptions = serde_json::from_str(system_table_json).unwrap();
1098        assert_eq!(opts.record_type, SlowQueriesRecordType::SystemTable);
1099
1100        let log_json = r#"{"enable": true, "record_type": "log", "threshold": "30s"}"#;
1101        let opts: SlowQueryOptions = serde_json::from_str(log_json).unwrap();
1102        assert_eq!(opts.record_type, SlowQueriesRecordType::Log);
1103    }
1104
1105    #[test]
1106    fn test_slow_query_options_deserialization_missing_record_type() {
1107        let json = r#"{"enable": false, "threshold": "30s"}"#;
1108        let opts: SlowQueryOptions = serde_json::from_str(json).unwrap();
1109
1110        // Missing record_type should use default (SystemTable)
1111        assert_eq!(opts.record_type, SlowQueriesRecordType::SystemTable);
1112        assert!(!opts.enable);
1113    }
1114
1115    #[test]
1116    fn test_otlp_export_protocol_deserialization_valid_values() {
1117        let grpc_json = r#""grpc""#;
1118        let protocol: OtlpExportProtocol = serde_json::from_str(grpc_json).unwrap();
1119        assert_eq!(protocol, OtlpExportProtocol::Grpc);
1120
1121        let http_json = r#""http""#;
1122        let protocol: OtlpExportProtocol = serde_json::from_str(http_json).unwrap();
1123        assert_eq!(protocol, OtlpExportProtocol::Http);
1124    }
1125
1126    #[test]
1127    fn test_logging_options_partial_eq_all_fields() {
1128        let base = LoggingOptions::default();
1129
1130        let mut log_format = base.clone();
1131        log_format.log_format = LogFormat::Json;
1132        assert_ne!(base, log_format);
1133
1134        let mut max_log_files = base.clone();
1135        max_log_files.max_log_files += 1;
1136        assert_ne!(base, max_log_files);
1137
1138        let mut max_log_dir_size = base.clone();
1139        max_log_dir_size.max_log_dir_size = ReadableSize::mb(1);
1140        assert_ne!(base, max_log_dir_size);
1141
1142        let mut otlp_export_protocol = base.clone();
1143        otlp_export_protocol.otlp_export_protocol = Some(OtlpExportProtocol::Http);
1144        assert_ne!(base, otlp_export_protocol);
1145
1146        let mut otlp_headers = base.clone();
1147        otlp_headers
1148            .otlp_headers
1149            .insert("key".to_string(), "value".to_string());
1150        assert_ne!(base, otlp_headers);
1151    }
1152}