diff --git a/docs/advanced_features/router.md b/docs/advanced_features/router.md index 388b86cda..0be734063 100644 --- a/docs/advanced_features/router.md +++ b/docs/advanced_features/router.md @@ -438,6 +438,14 @@ python -m sglang_router.launch_router \ --request-id-headers x-request-id x-trace-id ``` +Enable opentelmetry tracing: +```bash +python -m sglang_router.launch_router \ + --worker-urls http://worker1:8000 \ + --enable-trace \ + --otlp-traces-endpoint 0.0.0.0:4317 +``` + --- ## Troubleshooting diff --git a/sgl-model-gateway/Cargo.toml b/sgl-model-gateway/Cargo.toml index c112e1f82..3f2718645 100644 --- a/sgl-model-gateway/Cargo.toml +++ b/sgl-model-gateway/Cargo.toml @@ -55,6 +55,10 @@ tracing = "0.1" tracing-subscriber = { version = "0.3", features = ["env-filter", "json", "chrono"] } tracing-log = "0.2" tracing-appender = "0.2.3" +opentelemetry = "0.27" +opentelemetry_sdk = { version = "0.27", features = ["trace", "rt-tokio"] } +opentelemetry-otlp = { version = "0.27", features = ["trace", "grpc-tonic"] } +tracing-opentelemetry = "0.28" chrono = "0.4" kube = { version = "1.1.0", features = ["runtime", "derive"] } k8s-openapi = { version = "0.25.0", features = ["v1_33"] } @@ -131,6 +135,9 @@ tempfile = "3.8" lazy_static = "1.4" wasm-encoder = "0.242" npyz = { version = "0.8", features = ["npz"] } # For reading numpy .npz files in golden tests +opentelemetry-proto = { version = "0.27", features = ["gen-tonic"] } +tonic-v12 = { version = "0.12.3", package = "tonic" } +serial_test = "3.0" [[bench]] name = "request_processing" diff --git a/sgl-model-gateway/README.md b/sgl-model-gateway/README.md index dee4d0d3a..f695b2c87 100644 --- a/sgl-model-gateway/README.md +++ b/sgl-model-gateway/README.md @@ -9,7 +9,7 @@ High-performance model routing control and data plane for large-scale LLM deploy - Multi-model inference gateway mode (`--enable-igw`) that runs several routers at once and applies per-model policies. - Conversation, response, and chat-history connectors that centralize state at the router, enabling compliant sharing across models/MCP loops with in-memory, no-op, or Oracle ATP storage options. - Built-in reliability primitives: retries with exponential backoff, circuit breakers, token-bucket rate limiting, and queuing. -- First-class observability with structured logging and Prometheus metrics. +- First-class observability with structured logging, OpenTelemetry trace and Prometheus metrics. ### Architecture at a Glance **Control Plane** @@ -562,6 +562,7 @@ Only one of `--oracle-dsn` or `--oracle-tns-alias` should be supplied. - **Prometheus Metrics**: Enable with `--prometheus-host`/`--prometheus-port` (defaults to `0.0.0.0:29000`). Metrics cover request latency, retry behavior, circuit breaker states, worker health/load, queue depth, PD pipeline stats, tokenizer timings, and MCP activity. - **Request IDs**: Configurable headers via `--request-id-headers`; responses include `x-request-id`. - **CORS**: Set `--cors-allowed-origins` for browser access. +- **Request Tracing via OpenTelemetry**: Enable with `--enable-trace` and set opentelemetry collector endpoint with `--otlp-traces-endpoint :`. ## Security diff --git a/sgl-model-gateway/bindings/python/sglang_router/launch_router.py b/sgl-model-gateway/bindings/python/sglang_router/launch_router.py index 8da5191c5..506842f84 100644 --- a/sgl-model-gateway/bindings/python/sglang_router/launch_router.py +++ b/sgl-model-gateway/bindings/python/sglang_router/launch_router.py @@ -41,10 +41,6 @@ def launch_router(args: argparse.Namespace) -> Optional[Router]: mini_lb = MiniLoadBalancer(router_args) mini_lb.start() else: - # TODO: support tracing for router(Rust). - del router_args.enable_trace - del router_args.otlp_traces_endpoint - if Router is None: raise RuntimeError("Rust Router is not installed") router_args._validate_router_args() diff --git a/sgl-model-gateway/bindings/python/src/lib.rs b/sgl-model-gateway/bindings/python/src/lib.rs index 8a321cc95..9181abb25 100644 --- a/sgl-model-gateway/bindings/python/src/lib.rs +++ b/sgl-model-gateway/bindings/python/src/lib.rs @@ -224,6 +224,8 @@ struct Router { client_cert_path: Option, client_key_path: Option, ca_cert_paths: Vec, + enable_trace: bool, + otlp_traces_endpoint: String, } impl Router { @@ -309,6 +311,11 @@ impl Router { _ => None, }; + let trace_config = Some(config::TraceConfig { + enable_trace: self.enable_trace, + otlp_traces_endpoint: self.otlp_traces_endpoint.clone(), + }); + let history_backend = match self.history_backend { HistoryBackendType::Memory => config::HistoryBackend::Memory, HistoryBackendType::None => config::HistoryBackend::None, @@ -376,6 +383,7 @@ impl Router { .maybe_api_key(self.api_key.as_ref()) .maybe_discovery(discovery) .maybe_metrics(metrics) + .maybe_trace(trace_config) .maybe_log_dir(self.log_dir.as_ref()) .maybe_log_level(self.log_level.as_ref()) .maybe_request_id_headers(self.request_id_headers.clone()) @@ -477,6 +485,8 @@ impl Router { client_cert_path = None, client_key_path = None, ca_cert_paths = vec![], + enable_trace = false, + otlp_traces_endpoint = String::from("localhost:4317"), ))] #[allow(clippy::too_many_arguments)] fn new( @@ -552,6 +562,8 @@ impl Router { client_cert_path: Option, client_key_path: Option, ca_cert_paths: Vec, + enable_trace: bool, + otlp_traces_endpoint: String, ) -> PyResult { let mut all_urls = worker_urls.clone(); @@ -641,6 +653,8 @@ impl Router { client_cert_path, client_key_path, ca_cert_paths, + enable_trace, + otlp_traces_endpoint, }) } diff --git a/sgl-model-gateway/src/config/builder.rs b/sgl-model-gateway/src/config/builder.rs index b13d6d8a6..ac3591222 100644 --- a/sgl-model-gateway/src/config/builder.rs +++ b/sgl-model-gateway/src/config/builder.rs @@ -1,7 +1,7 @@ use super::{ CircuitBreakerConfig, ConfigError, ConfigResult, DiscoveryConfig, HealthCheckConfig, HistoryBackend, MetricsConfig, OracleConfig, PolicyConfig, PostgresConfig, RetryConfig, - RouterConfig, RoutingMode, TokenizerCacheConfig, + RouterConfig, RoutingMode, TokenizerCacheConfig, TraceConfig, }; use crate::{core::ConnectionMode, mcp::McpConfig}; @@ -297,6 +297,24 @@ impl RouterConfigBuilder { self } + // ===================== Otel Trace ==================== + + pub fn enable_trace>(mut self, endpoint: S) -> Self { + self.config.trace_config = Some(TraceConfig { + enable_trace: true, + otlp_traces_endpoint: endpoint.into(), + }); + self + } + + pub fn disable_trace(mut self) -> Self { + self.config.trace_config = Some(TraceConfig { + enable_trace: false, + otlp_traces_endpoint: "".to_string(), + }); + self + } + // ==================== Logging ==================== pub fn log_dir>(mut self, dir: S) -> Self { @@ -461,6 +479,11 @@ impl RouterConfigBuilder { self } + pub fn maybe_trace(mut self, trace_config: Option) -> Self { + self.config.trace_config = trace_config; + self + } + pub fn maybe_log_dir(mut self, dir: Option>) -> Self { self.config.log_dir = dir.map(|d| d.into()); self @@ -703,11 +726,13 @@ mod tests { .to_builder() .port(4000) .enable_metrics("0.0.0.0", 29000) + .enable_trace("localhost:4317") .build() .unwrap(); assert_eq!(modified.port, 4000); assert!(modified.metrics.is_some()); + assert!(modified.trace_config.is_some()); } /// Test complex routing mode helper method diff --git a/sgl-model-gateway/src/config/types.rs b/sgl-model-gateway/src/config/types.rs index 898013c68..b12802e29 100644 --- a/sgl-model-gateway/src/config/types.rs +++ b/sgl-model-gateway/src/config/types.rs @@ -23,6 +23,7 @@ pub struct RouterConfig { pub api_key: Option, pub discovery: Option, pub metrics: Option, + pub trace_config: Option, pub log_dir: Option, pub log_level: Option, pub request_id_headers: Option>, @@ -461,6 +462,21 @@ impl Default for MetricsConfig { } } +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct TraceConfig { + pub enable_trace: bool, + pub otlp_traces_endpoint: String, +} + +impl Default for TraceConfig { + fn default() -> Self { + Self { + enable_trace: false, + otlp_traces_endpoint: "localhost:4317".to_string(), + } + } +} + impl Default for RouterConfig { fn default() -> Self { Self { @@ -478,6 +494,7 @@ impl Default for RouterConfig { api_key: None, discovery: None, metrics: None, + trace_config: None, log_dir: None, log_level: None, request_id_headers: None, @@ -544,6 +561,14 @@ impl RouterConfig { self.metrics.is_some() } + /// Check if tracing is enabled + pub fn has_tracing(&self) -> bool { + match &self.trace_config { + Some(trace_config) => trace_config.enable_trace, + None => false, + } + } + /// Compute the effective retry config considering disable flag pub fn effective_retry_config(&self) -> RetryConfig { let mut cfg = self.retry.clone(); @@ -588,6 +613,7 @@ mod tests { assert_eq!(config.worker_startup_check_interval_secs, 30); assert!(config.discovery.is_none()); assert!(config.metrics.is_none()); + assert!(config.trace_config.is_none()); assert!(config.log_dir.is_none()); assert!(config.log_level.is_none()); } @@ -636,6 +662,7 @@ mod tests { assert_eq!(config.log_level, deserialized.log_level); assert!(deserialized.discovery.is_none()); assert!(deserialized.metrics.is_none()); + assert!(deserialized.trace_config.is_none()); } #[test] @@ -891,6 +918,25 @@ mod tests { assert_eq!(config.host, "0.0.0.0"); } + #[test] + fn test_trace_config_default() { + let config = TraceConfig::default(); + + assert!(!config.enable_trace); + assert_eq!(config.otlp_traces_endpoint, "localhost:4317"); + } + + #[test] + fn test_trace_config_custom() { + let config = TraceConfig { + enable_trace: true, + otlp_traces_endpoint: "otel-collector:4317".to_string(), + }; + + assert!(config.enable_trace); + assert_eq!(config.otlp_traces_endpoint, "otel-collector:4317"); + } + #[test] fn test_mode_type() { let config = RouterConfig::builder() @@ -932,6 +978,17 @@ mod tests { assert!(config.has_metrics()); } + #[test] + fn test_has_tracing() { + let config = RouterConfig::default(); + assert!(!config.has_tracing()); + + let config = RouterConfig::builder() + .enable_trace("localhost:4317") + .build_unchecked(); + assert!(config.has_tracing()); + } + #[test] fn test_large_worker_lists() { let large_urls: Vec = (0..1000).map(|i| format!("http://worker{}", i)).collect(); @@ -1016,6 +1073,7 @@ mod tests { ..Default::default() }) .enable_metrics("0.0.0.0", 9090) + .enable_trace("localhost:4317") .log_dir("/var/log/sglang") .log_level("info") .max_concurrent_requests(64) @@ -1026,6 +1084,7 @@ mod tests { assert_eq!(config.policy.name(), "power_of_two"); assert!(config.has_service_discovery()); assert!(config.has_metrics()); + assert!(config.has_tracing()); } #[test] @@ -1055,6 +1114,7 @@ mod tests { ..Default::default() }) .metrics_config(MetricsConfig::default()) + .enable_trace("localhost:4317") .log_level("debug") .max_concurrent_requests(64) .build_unchecked(); @@ -1064,6 +1124,7 @@ mod tests { assert_eq!(config.policy.name(), "cache_aware"); assert!(config.has_service_discovery()); assert!(config.has_metrics()); + assert!(config.has_tracing()); } #[test] @@ -1092,6 +1153,7 @@ mod tests { bootstrap_port_annotation: "mycompany.io/bootstrap".to_string(), }) .enable_metrics("::", 9999) // IPv6 any + .enable_trace("localhost:4317") .log_dir("/opt/logs/sglang") .log_level("trace") .max_concurrent_requests(64) @@ -1099,6 +1161,7 @@ mod tests { assert!(config.has_service_discovery()); assert!(config.has_metrics()); + assert!(config.has_tracing()); assert_eq!(config.mode_type(), "regular"); let json = serde_json::to_string_pretty(&config).unwrap(); diff --git a/sgl-model-gateway/src/config/validation.rs b/sgl-model-gateway/src/config/validation.rs index ed906ed79..eda1000c8 100644 --- a/sgl-model-gateway/src/config/validation.rs +++ b/sgl-model-gateway/src/config/validation.rs @@ -18,6 +18,10 @@ impl ConfigValidator { Self::validate_metrics(metrics)?; } + if let Some(trace_config) = &config.trace_config { + Self::validate_trace(trace_config)?; + } + Self::validate_compatibility(config)?; let retry_cfg = config.effective_retry_config(); @@ -357,6 +361,46 @@ impl ConfigValidator { Ok(()) } + fn validate_trace(trace_config: &TraceConfig) -> ConfigResult<()> { + if !trace_config.enable_trace { + return Ok(()); + } + + let endpoint = &trace_config.otlp_traces_endpoint; + + let Some((host, port_str)) = endpoint.rsplit_once(':') else { + return Err(ConfigError::InvalidValue { + field: "trace_config.otlp_traces_endpoint".to_string(), + value: endpoint.clone(), + reason: + "expected format :, e.g., otel-collector:4317 or 127.0.0.1:4317" + .to_string(), + }); + }; + + if host.is_empty() { + return Err(ConfigError::InvalidValue { + field: "trace_config.otlp_traces_endpoint".to_string(), + value: endpoint.clone(), + reason: "host part cannot be empty".to_string(), + }); + } + + // check port: must be 1~65535 + match port_str.parse::() { + Ok(p) if p > 0 => (), // valid port + _ => { + return Err(ConfigError::InvalidValue { + field: "trace_config.otlp_traces_endpoint".to_string(), + value: endpoint.clone(), + reason: "port must be a number between 1 and 65535".to_string(), + }); + } + }; + + Ok(()) + } + fn validate_retry(retry: &RetryConfig) -> ConfigResult<()> { if retry.max_retries < 1 { return Err(ConfigError::InvalidValue { diff --git a/sgl-model-gateway/src/lib.rs b/sgl-model-gateway/src/lib.rs index ae57f7c19..8d1144dce 100644 --- a/sgl-model-gateway/src/lib.rs +++ b/sgl-model-gateway/src/lib.rs @@ -9,6 +9,7 @@ pub mod mcp; pub mod metrics; pub mod middleware; pub mod multimodal; +pub mod otel_trace; pub mod policies; pub mod protocols; pub mod reasoning_parser; diff --git a/sgl-model-gateway/src/logging.rs b/sgl-model-gateway/src/logging.rs index baf1ddd3b..fa8e57157 100644 --- a/sgl-model-gateway/src/logging.rs +++ b/sgl-model-gateway/src/logging.rs @@ -10,6 +10,8 @@ use tracing_subscriber::{ fmt::time::ChronoUtc, layer::SubscriberExt, util::SubscriberInitExt, EnvFilter, Layer, }; +use crate::{config::TraceConfig, otel_trace::get_otel_layer}; + #[derive(Debug, Clone)] pub struct LoggingConfig { pub level: Level, @@ -38,7 +40,7 @@ pub struct LogGuard { _file_guard: Option, } -pub fn init_logging(config: LoggingConfig) -> LogGuard { +pub fn init_logging(config: LoggingConfig, otel_layer_config: Option) -> LogGuard { let _ = LogTracer::init(); let level_filter = match config.level { @@ -121,6 +123,19 @@ pub fn init_logging(config: LoggingConfig) -> LogGuard { layers.push(file_layer); } + if let Some(otel_layer_config) = &otel_layer_config { + if otel_layer_config.enable_trace { + match get_otel_layer() { + Ok(otel_layer) => { + layers.push(otel_layer); + } + Err(e) => { + eprintln!("Failed to initialize OpenTelemetry: {}", e); + } + } + } + } + let _ = tracing_subscriber::registry() .with(env_filter) .with(layers) diff --git a/sgl-model-gateway/src/main.rs b/sgl-model-gateway/src/main.rs index 10f88fc7d..97f24f711 100644 --- a/sgl-model-gateway/src/main.rs +++ b/sgl-model-gateway/src/main.rs @@ -5,10 +5,11 @@ use sgl_model_gateway::{ config::{ CircuitBreakerConfig, ConfigError, ConfigResult, DiscoveryConfig, HealthCheckConfig, HistoryBackend, MetricsConfig, OracleConfig, PolicyConfig, PostgresConfig, RetryConfig, - RouterConfig, RoutingMode, TokenizerCacheConfig, + RouterConfig, RoutingMode, TokenizerCacheConfig, TraceConfig, }, core::ConnectionMode, metrics::PrometheusConfig, + otel_trace::{is_otel_enabled, shutdown_otel}, server::{self, ServerConfig}, service_discovery::ServiceDiscoveryConfig, version, @@ -348,6 +349,12 @@ struct CliArgs { #[arg(long, default_value_t = false)] enable_wasm: bool, + + #[arg(long, default_value_t = false)] + enable_trace: bool, + + #[arg(long, default_value = "localhost:4317")] + otlp_traces_endpoint: String, } enum OracleConnectSource { @@ -539,6 +546,11 @@ impl CliArgs { host: self.prometheus_host.clone(), }); + let trace_config = Some(TraceConfig { + enable_trace: self.enable_trace, + otlp_traces_endpoint: self.otlp_traces_endpoint.clone(), + }); + let mut all_urls = Vec::new(); match &mode { RoutingMode::Regular { worker_urls } => { @@ -624,6 +636,7 @@ impl CliArgs { .maybe_api_key(self.api_key.as_ref()) .maybe_discovery(discovery) .maybe_metrics(metrics) + .maybe_trace(trace_config) .maybe_log_dir(self.log_dir.as_ref()) .maybe_request_id_headers( (!self.request_id_headers.is_empty()).then(|| self.request_id_headers.clone()), @@ -768,6 +781,8 @@ Provide --worker-urls or PD flags as usual.", let server_config = cli_args.to_server_config(router_config); let runtime = tokio::runtime::Runtime::new()?; runtime.block_on(async move { server::startup(server_config).await })?; - + if is_otel_enabled() { + shutdown_otel(); + } Ok(()) } diff --git a/sgl-model-gateway/src/middleware.rs b/sgl-model-gateway/src/middleware.rs index bf8884822..7ea623c03 100644 --- a/sgl-model-gateway/src/middleware.rs +++ b/sgl-model-gateway/src/middleware.rs @@ -207,6 +207,7 @@ impl MakeSpan for RequestSpan { // Don't try to extract request ID here - it won't be available yet // The RequestIdLayer runs after TraceLayer creates the span info_span!( + target: "sgl_model_gateway::otel-trace", "http_request", method = %request.method(), uri = %request.uri(), @@ -215,6 +216,7 @@ impl MakeSpan for RequestSpan { status_code = Empty, latency = Empty, error = Empty, + module = "sglang::router_rs" ) } } diff --git a/sgl-model-gateway/src/otel_trace.rs b/sgl-model-gateway/src/otel_trace.rs new file mode 100644 index 000000000..6043d80c7 --- /dev/null +++ b/sgl-model-gateway/src/otel_trace.rs @@ -0,0 +1,229 @@ +use std::{ + collections::HashSet, + sync::{ + atomic::{AtomicBool, Ordering}, + OnceLock, + }, + time::Duration, +}; + +use anyhow::Result; +use axum::http::{HeaderMap, HeaderName, HeaderValue}; +use opentelemetry::{global, trace::TracerProvider as _, KeyValue}; +use opentelemetry_otlp::WithExportConfig; +use opentelemetry_sdk::{ + propagation::TraceContextPropagator, + runtime, + trace::{BatchConfigBuilder, BatchSpanProcessor, Tracer as SdkTracer, TracerProvider}, + Resource, +}; +use tokio::task::spawn_blocking; +use tracing::{Metadata, Subscriber}; +use tracing_opentelemetry::{self, OpenTelemetrySpanExt}; +use tracing_subscriber::{ + layer::{Context, Filter}, + Layer, +}; + +use crate::routers::http::events::get_module_path as http_router_get_module_path; + +static ENABLED: AtomicBool = AtomicBool::new(false); + +// global tracer +static TRACER: OnceLock = OnceLock::new(); +static PROVIDER: OnceLock = OnceLock::new(); + +pub struct CustomOtelFilter { + allowed_targets: HashSet, +} + +impl CustomOtelFilter { + pub fn new() -> Self { + let mut allowed_targets = HashSet::new(); + allowed_targets.insert("sgl_model_gateway::otel-trace".to_string()); + allowed_targets.insert(http_router_get_module_path().to_string()); + + Self { allowed_targets } + } +} + +impl Filter for CustomOtelFilter +where + S: Subscriber, +{ + fn enabled(&self, meta: &Metadata<'_>, _cx: &Context<'_, S>) -> bool { + self.allowed_targets.contains(meta.target()) + } + + fn callsite_enabled(&self, meta: &'static Metadata<'static>) -> tracing::subscriber::Interest { + if self.allowed_targets.contains(meta.target()) { + tracing::subscriber::Interest::always() + } else { + tracing::subscriber::Interest::never() + } + } +} + +impl Default for CustomOtelFilter { + fn default() -> Self { + Self::new() + } +} + +/// init OpenTelemetry connection +pub fn otel_tracing_init(enable: bool, otlp_endpoint: Option<&str>) -> Result<()> { + if !enable { + ENABLED.store(false, Ordering::Relaxed); + return Ok(()); + } + + let endpoint = otlp_endpoint.unwrap_or("localhost:4317"); + let endpoint = if !endpoint.starts_with("http://") && !endpoint.starts_with("https://") { + format!("http://{}", endpoint) + } else { + endpoint.to_string() + }; + + let result = std::panic::catch_unwind(|| -> Result<()> { + global::set_text_map_propagator(TraceContextPropagator::new()); + + let exporter = opentelemetry_otlp::SpanExporter::builder() + .with_tonic() + .with_endpoint(endpoint) + .with_protocol(opentelemetry_otlp::Protocol::Grpc) + .build()?; + + let batch_config = BatchConfigBuilder::default() + .with_scheduled_delay(Duration::from_millis(500)) + .with_max_export_batch_size(64) + .build(); + + let span_processor = BatchSpanProcessor::builder(exporter, runtime::Tokio) + .with_batch_config(batch_config) + .build(); + + let resource = Resource::default().merge(&Resource::new(vec![KeyValue::new( + "service.name", + "sgl-router", + )])); + + let provider = TracerProvider::builder() + .with_span_processor(span_processor) + .with_resource(resource) + .build(); + PROVIDER + .set(provider.clone()) + .map_err(|_| anyhow::anyhow!("Provider already initialized"))?; + + let tracer = provider.tracer("sgl-router"); + + TRACER + .set(tracer) + .map_err(|_| anyhow::anyhow!("Tracer already initialized"))?; + + let _ = global::set_tracer_provider(provider); + + ENABLED.store(true, Ordering::Relaxed); + + Ok(()) + }); + + match result { + Ok(Ok(())) => { + eprintln!( + "[tracing] OpenTelemetry initialized successfully, enabled: {}", + ENABLED.load(Ordering::Relaxed) + ); + Ok(()) + } + Ok(Err(e)) => { + eprintln!("[tracing] Failed to initialize OTLP tracer: {}", e); + ENABLED.store(false, Ordering::Relaxed); + Err(e) + } + Err(_) => { + eprintln!("[tracing] Panic during OpenTelemetry initialization"); + ENABLED.store(false, Ordering::Relaxed); + Err(anyhow::anyhow!("Panic during initialization")) + } + } +} + +pub fn get_otel_layer() -> Result + Send + Sync + 'static>, &'static str> +where + S: Subscriber + for<'a> tracing_subscriber::registry::LookupSpan<'a> + Send + Sync, +{ + if !is_otel_enabled() { + return Err("OpenTelemetry is not enabled"); + } + + let tracer = TRACER + .get() + .ok_or("Tracer not initialized. Call otel_tracing_init first.")? + .clone(); + + let custom_filter = CustomOtelFilter::new(); + + let layer = tracing_opentelemetry::layer() + .with_tracer(tracer) + .with_filter(custom_filter); + + Ok(Box::new(layer)) +} + +pub fn is_otel_enabled() -> bool { + ENABLED.load(Ordering::Relaxed) +} + +pub async fn flush_spans_async() -> Result<()> { + if !is_otel_enabled() { + return Ok(()); + } + + if let Some(provider) = PROVIDER.get() { + let provider = provider.clone(); + + spawn_blocking(move || provider.force_flush()) + .await + .map_err(|e| { + anyhow::anyhow!("Failed to join blocking task for flushing spans: {}", e) + })?; + + Ok(()) + } else { + Err(anyhow::anyhow!("Provider not initialized")) + } +} + +pub fn shutdown_otel() { + if ENABLED.load(Ordering::Relaxed) { + global::shutdown_tracer_provider(); + ENABLED.store(false, Ordering::Relaxed); + eprintln!("[tracing] OpenTelemetry shut down"); + } +} + +pub fn inject_trace_context_http(headers: &mut HeaderMap) -> Result<()> { + if !is_otel_enabled() { + return Err(anyhow::anyhow!("OTEL not enabled")); + } + + let context = tracing::Span::current().context(); + + struct HeaderInjector<'a>(&'a mut HeaderMap); + + impl<'a> opentelemetry::propagation::Injector for HeaderInjector<'a> { + fn set(&mut self, key: &str, value: String) { + if let Ok(header_name) = HeaderName::from_bytes(key.as_bytes()) { + if let Ok(header_value) = HeaderValue::from_str(&value) { + self.0.insert(header_name, header_value); + } + } + } + } + + global::get_text_map_propagator(|propagator| { + propagator.inject_context(&context, &mut HeaderInjector(headers)); + }); + Ok(()) +} diff --git a/sgl-model-gateway/src/routers/http/events.rs b/sgl-model-gateway/src/routers/http/events.rs new file mode 100644 index 000000000..096c49a75 --- /dev/null +++ b/sgl-model-gateway/src/routers/http/events.rs @@ -0,0 +1,69 @@ +//! request events for observability and monitoring + +use tracing::{debug, event, Level}; + +use crate::otel_trace::is_otel_enabled; + +pub fn get_module_path() -> &'static str { + module_path!() +} + +pub trait Event { + fn emit(&self); +} + +#[derive(Debug)] +pub struct RequestPDSentEvent { + pub prefill_url: String, + pub decode_url: String, +} + +impl Event for RequestPDSentEvent { + fn emit(&self) { + if !is_otel_enabled() { + debug!( + "Sending concurrent requests to prefill={} decode={}", + self.prefill_url, self.decode_url + ); + } else { + event!( + Level::INFO, + prefill_url = %self.prefill_url, + decode_url = %self.decode_url, + "Sending concurrent requests" + ); + } + } +} + +#[derive(Debug)] +pub struct RequestSentEvent { + pub url: String, +} + +impl Event for RequestSentEvent { + fn emit(&self) { + if !is_otel_enabled() { + debug!("Sending request to {}", self.url); + } else { + event!( + Level::INFO, + url = %self.url, + "Sending requests" + ); + } + } +} + +#[derive(Debug)] +pub struct RequestReceivedEvent {} + +impl Event for RequestReceivedEvent { + fn emit(&self) { + if !is_otel_enabled() { + debug!("Received concurrent requests"); + } else { + event!(Level::INFO, "Received concurrent requests"); + } + } +} diff --git a/sgl-model-gateway/src/routers/http/mod.rs b/sgl-model-gateway/src/routers/http/mod.rs index 3f31b6f86..5beafc8d1 100644 --- a/sgl-model-gateway/src/routers/http/mod.rs +++ b/sgl-model-gateway/src/routers/http/mod.rs @@ -1,5 +1,6 @@ //! HTTP router implementations +pub mod events; pub mod pd_router; pub mod pd_types; pub mod router; diff --git a/sgl-model-gateway/src/routers/http/pd_router.rs b/sgl-model-gateway/src/routers/http/pd_router.rs index eb4dc8172..98d464c85 100644 --- a/sgl-model-gateway/src/routers/http/pd_router.rs +++ b/sgl-model-gateway/src/routers/http/pd_router.rs @@ -14,13 +14,17 @@ use serde_json::{json, Value}; use tokio_stream::wrappers::UnboundedReceiverStream; use tracing::{debug, error, warn}; -use super::pd_types::api_path; +use super::{ + events::{self, Event}, + pd_types::api_path, +}; use crate::{ config::types::RetryConfig, core::{ is_retryable_status, RetryExecutor, Worker, WorkerLoadGuard, WorkerRegistry, WorkerType, }, metrics::RouterMetrics, + otel_trace::inject_trace_context_http, policies::{LoadBalancingPolicy, PolicyRegistry}, protocols::{ chat::{ChatCompletionRequest, ChatMessage, MessageContent}, @@ -391,6 +395,12 @@ impl PDRouter { None }; + let mut headers_with_trace = headers.cloned().unwrap_or_default(); + let headers = match inject_trace_context_http(&mut headers_with_trace) { + Ok(()) => Some(&headers_with_trace), + Err(_) => headers, + }; + // Build both requests let prefill_request = self.build_post_with_headers( &self.client, @@ -410,15 +420,16 @@ impl PDRouter { ); // Send both requests concurrently and wait for both - debug!( - "Sending concurrent requests to prefill={} decode={}", - prefill.url(), - decode.url() - ); + events::RequestPDSentEvent { + prefill_url: prefill.url().to_string(), + decode_url: decode.url().to_string(), + } + .emit(); let (prefill_result, decode_result) = tokio::join!(prefill_request.send(), decode_request.send()); - debug!("Received responses from both servers"); + + events::RequestReceivedEvent {}.emit(); let duration = start_time.elapsed(); RouterMetrics::record_pd_request_duration(context.route, duration); @@ -873,7 +884,11 @@ impl PDRouter { // Whitelist important end-to-end headers, skip hop-by-hop let forward = matches!( name_lc.as_str(), - "authorization" | "x-request-id" | "x-correlation-id" + "authorization" + | "x-request-id" + | "x-correlation-id" + | "traceparent" // W3C Trace Context + | "tracestate" // W3C Trace Context ) || name_lc.starts_with("x-request-id-"); if forward { if let Ok(val) = value.to_str() { diff --git a/sgl-model-gateway/src/routers/http/router.rs b/sgl-model-gateway/src/routers/http/router.rs index ca9de0b92..d7ca65b10 100644 --- a/sgl-model-gateway/src/routers/http/router.rs +++ b/sgl-model-gateway/src/routers/http/router.rs @@ -15,12 +15,14 @@ use reqwest::Client; use tokio_stream::wrappers::UnboundedReceiverStream; use tracing::{debug, error}; +use super::events::{self, Event}; use crate::{ config::types::RetryConfig, core::{ is_retryable_status, ConnectionMode, RetryExecutor, Worker, WorkerRegistry, WorkerType, }, metrics::RouterMetrics, + otel_trace::inject_trace_context_http, policies::PolicyRegistry, protocols::{ chat::ChatCompletionRequest, @@ -212,6 +214,16 @@ impl Router { None }; + events::RequestSentEvent { + url: worker.url().to_string(), + } + .emit(); + let mut headers_with_trace = headers.cloned().unwrap_or_default(); + let headers = match inject_trace_context_http(&mut headers_with_trace) { + Ok(()) => Some(&headers_with_trace), + Err(_) => headers, + }; + let response = self .send_typed_request( headers, @@ -223,6 +235,8 @@ impl Router { ) .await; + events::RequestReceivedEvent {}.emit(); + worker.record_outcome(response.status().is_success()); // For retryable failures, we need to decrement load since send_typed_request diff --git a/sgl-model-gateway/src/server.rs b/sgl-model-gateway/src/server.rs index fdf90bdd6..dde36f5a3 100644 --- a/sgl-model-gateway/src/server.rs +++ b/sgl-model-gateway/src/server.rs @@ -35,6 +35,7 @@ use crate::{ logging::{self, LoggingConfig}, metrics::{self, PrometheusConfig}, middleware::{self, AuthConfig, QueuedRequest}, + otel_trace, protocols::{ chat::ChatCompletionRequest, classify::ClassifyRequest, @@ -703,25 +704,35 @@ pub fn build_app( pub async fn startup(config: ServerConfig) -> Result<(), Box> { static LOGGING_INITIALIZED: AtomicBool = AtomicBool::new(false); + if let Some(trace_config) = &config.router_config.trace_config { + otel_trace::otel_tracing_init( + trace_config.enable_trace, + Some(&trace_config.otlp_traces_endpoint), + )?; + } + let _log_guard = if !LOGGING_INITIALIZED.swap(true, Ordering::SeqCst) { - Some(logging::init_logging(LoggingConfig { - level: config - .log_level - .as_deref() - .and_then(|s| match s.to_uppercase().parse::() { - Ok(l) => Some(l), - Err(_) => { - warn!("Invalid log level string: '{s}'. Defaulting to INFO."); - None - } - }) - .unwrap_or(Level::INFO), - json_format: false, - log_dir: config.log_dir.clone(), - colorize: true, - log_file_name: "sgl-model-gateway".to_string(), - log_targets: None, - })) + Some(logging::init_logging( + LoggingConfig { + level: config + .log_level + .as_deref() + .and_then(|s| match s.to_uppercase().parse::() { + Ok(l) => Some(l), + Err(_) => { + warn!("Invalid log level string: '{s}'. Defaulting to INFO."); + None + } + }) + .unwrap_or(Level::INFO), + json_format: false, + log_dir: config.log_dir.clone(), + colorize: true, + log_file_name: "sgl-model-gateway".to_string(), + log_targets: None, + }, + config.router_config.trace_config.clone(), + )) } else { None }; diff --git a/sgl-model-gateway/tests/otel_tracing_test.rs b/sgl-model-gateway/tests/otel_tracing_test.rs new file mode 100644 index 000000000..8453aebea --- /dev/null +++ b/sgl-model-gateway/tests/otel_tracing_test.rs @@ -0,0 +1,249 @@ +mod common; + +use std::{ + sync::{ + atomic::{AtomicUsize, Ordering}, + Arc, + }, + time::Duration, +}; + +use axum::{body::Body, extract::Request, http::StatusCode}; +use common::mock_worker::{HealthStatus, MockWorker, MockWorkerConfig, WorkerType}; +use opentelemetry_proto::tonic::collector::trace::v1::{ + trace_service_server::{TraceService, TraceServiceServer}, + ExportTraceServiceRequest, ExportTraceServiceResponse, +}; +use portpicker::pick_unused_port; +use serde_json::json; +use serial_test::serial; +use sgl_model_gateway::{ + config::{RouterConfig, TraceConfig}, + core::Job, + logging, otel_trace, + routers::RouterFactory, +}; +use tokio::sync::oneshot; +use tonic_v12::{transport::Server, Request as TonicRequest, Response, Status}; +use tower::ServiceExt; +use tracing::info_span; + +#[derive(Clone)] +struct TestOtelCollector { + span_count: Arc, +} + +impl TestOtelCollector { + fn new() -> Self { + Self { + span_count: Arc::new(AtomicUsize::new(0)), + } + } + + fn get_span_count(&self) -> usize { + self.span_count.load(Ordering::SeqCst) + } +} + +#[tonic_v12::async_trait] +impl TraceService for TestOtelCollector { + async fn export( + &self, + request: TonicRequest, + ) -> Result, Status> { + let req = request.into_inner(); + + let mut total_spans = 0; + + for resource_span in &req.resource_spans { + for scope_span in &resource_span.scope_spans { + total_spans += scope_span.spans.len(); + } + } + + self.span_count.fetch_add(total_spans, Ordering::SeqCst); + + Ok(Response::new(ExportTraceServiceResponse { + partial_success: None, + })) + } +} + +async fn start_collector( + port: u16, + shutdown_rx: oneshot::Receiver<()>, +) -> Result> { + let addr = format!("0.0.0.0:{}", port).parse()?; + let collector = TestOtelCollector::new(); + let collector_clone = collector.clone(); + + tokio::spawn(async move { + let _ = Server::builder() + .add_service(TraceServiceServer::new(collector_clone)) + .serve_with_shutdown(addr, async { + shutdown_rx.await.ok(); + }) + .await; + }); + + tokio::time::sleep(Duration::from_millis(200)).await; + + Ok(collector) +} + +#[tokio::test] +#[serial] +async fn test_router_with_tracing() { + // 1. Start the OTLP collector + let port = pick_unused_port().expect("Failed to pick unused port"); + let (shutdown_tx, shutdown_rx) = oneshot::channel(); + let collector = start_collector(port, shutdown_rx) + .await + .expect("Failed to start collector"); + let collector_endpoint = format!("0.0.0.0:{}", port); + println!("OTLP Collector started on: {}", collector_endpoint); + + // 2. create the mock worker + let mut mock_worker = MockWorker::new(MockWorkerConfig { + port: 0, + worker_type: WorkerType::Regular, + health_status: HealthStatus::Healthy, + response_delay_ms: 0, + fail_rate: 0.0, + }); + + let worker_url = mock_worker.start().await.unwrap(); + tokio::time::sleep(Duration::from_millis(200)).await; + println!("Mock worker started on: {}", worker_url); + + // 3. create router config and enable tracing + let router_config = RouterConfig::builder() + .regular_mode(vec![worker_url.clone()]) + .random_policy() + .host("0.0.0.0") + .port(0) + .max_payload_size(256 * 1024 * 1024) + .request_timeout_secs(60) + .worker_startup_timeout_secs(1) + .worker_startup_check_interval_secs(1) + .max_concurrent_requests(64) + .queue_timeout_secs(60) + .enable_trace(&collector_endpoint) + .build_unchecked(); + + // 4. Initialize the OTLP client + let init_result = otel_trace::otel_tracing_init(true, Some(&collector_endpoint)); + assert!( + init_result.is_ok(), + "Failed to initialize OTEL: {:?}", + init_result.err() + ); + println!("OpenTelemetry initialized successfully"); + + let trace_config = TraceConfig { + enable_trace: true, + otlp_traces_endpoint: collector_endpoint.clone(), + }; + let _log_guard = logging::init_logging( + logging::LoggingConfig { + level: tracing::Level::INFO, + json_format: false, + log_dir: None, + colorize: false, + log_file_name: "test-otel".to_string(), + log_targets: Some(vec!["sgl_model_gateway".to_string()]), + }, + Some(trace_config), + ); + println!("Logging initialized with OTEL layer"); + + // 5. Create a span and sleep for a while + let _span = info_span!(target: "sgl_model_gateway::otel-trace", "test_router_with_tracing"); + tokio::time::sleep(Duration::from_secs(1)).await; + drop(_span); + + // 6. create app context and router + let app_context = common::create_test_context(router_config.clone()).await; + + // 7. initialize worker + let job_queue = app_context + .worker_job_queue + .get() + .expect("JobQueue should be initialized"); + + let job = Job::InitializeWorkersFromConfig { + router_config: Box::new(router_config.clone()), + }; + + job_queue + .submit(job) + .await + .expect("Failed to submit worker init job"); + + // 8. wait for worker initialization + tokio::time::sleep(Duration::from_millis(1000)).await; + println!("Workers initialized"); + + // 9. create router + let router = RouterFactory::create_router(&app_context) + .await + .expect("Failed to create router"); + + println!("Router created"); + + // 10. create app (middleware::create_logging_layer() will use the already initialized OTEL layer) + let app = + common::test_app::create_test_app_with_context(Arc::from(router), app_context.clone()); + + println!("App created with logging middleware"); + + // 10. send request + let request_body = json!({ + "model": "test-model", + "messages": [ + {"role": "user", "content": "Hello, test OpenTelemetry tracing!"} + ], + "temperature": 0.7, + "max_tokens": 50 + }); + + let request = Request::builder() + .method("POST") + .uri("/v1/chat/completions") + .header("content-type", "application/json") + .body(Body::from(request_body.to_string())) + .unwrap(); + + println!("Sending request to router..."); + let response = app.oneshot(request).await.unwrap(); + + assert_eq!(response.status(), StatusCode::OK, "Request should succeed"); + + println!("Request completed successfully"); + drop(response); + + // 11. Wait for spans to be exported + match otel_trace::flush_spans_async().await { + Ok(_) => println!("Spans flushed successfully"), + Err(e) => println!("Failed to flush spans: {:?}", e), + } + + // 12. Verify that the spans were exported to the collector + let span_count = collector.get_span_count(); + println!("Total spans received by collector: {}", span_count); + + assert!( + span_count == 2, + "Expected to receive at least 2 span, but got {}. \ + This indicates that tracing data is not being exported to the OTLP collector.", + span_count + ); + + println!("Test passed! Collector received {} spans", span_count); + + // 13. cleanup + let _ = shutdown_tx.send(()); + mock_worker.stop().await; + + println!("Cleanup completed"); +}