From 31665225ebead7d39adf7da263194e9ad09ca486 Mon Sep 17 00:00:00 2001 From: Hampus Date: Tue, 30 Jun 2026 01:47:23 +0200 Subject: [PATCH] fix(fluxer_svc): drain router requests concurrently (#1221) --- fluxer_svc/src/router.rs | 405 ++++++++++++++++++++++++++---------- fluxer_svc/src/transport.rs | 64 +++++- 2 files changed, 353 insertions(+), 116 deletions(-) diff --git a/fluxer_svc/src/router.rs b/fluxer_svc/src/router.rs index 53cda0ee6..93ee9671a 100644 --- a/fluxer_svc/src/router.rs +++ b/fluxer_svc/src/router.rs @@ -7,6 +7,7 @@ use crate::transport::{Transport, TransportMessage, TransportSubscriber, reply_m use moka::future::Cache; use std::sync::Arc; use std::time::Duration; +use tokio::sync::Semaphore; use tokio::task::JoinSet; use tracing::{debug, info, warn}; @@ -51,6 +52,137 @@ async fn forward_to_shard( .await } +async fn handle_router_request( + msg: T::Message, + transport: T, + service: Arc, + ring: Arc, + inflight: Cache>, + metrics: Arc, +) where + S: RouterService, + T: Transport, +{ + let request_start = now_ms(); + metrics.record_request(); + if msg.payload().len() > MAX_ROUTER_REQUEST_BYTES { + warn!( + payload_bytes = msg.payload().len(), + max_payload_bytes = MAX_ROUTER_REQUEST_BYTES, + "rejecting oversized router request" + ); + metrics.record_request_error(); + if msg.has_reply() { + let error_response = + serde_json::to_vec(&serde_json::json!({"error": "request_too_large"})) + .unwrap_or_default(); + let _ = reply_message(&msg, &transport, &error_response).await; + } + return; + } + let payload = msg.payload().to_vec(); + + let request: S::Request = match serde_json::from_slice(&payload) { + Ok(r) => r, + Err(err) => { + warn!(error = %err, "failed to decode incoming request"); + metrics.record_request_error(); + if msg.has_reply() { + let error_response = + serde_json::to_vec(&serde_json::json!({"error": "decode_error"})) + .unwrap_or_default(); + let _ = reply_message(&msg, &transport, &error_response).await; + } + return; + } + }; + + if let Some(cached) = service.l1_lookup(&request) { + metrics.record_cache_hit(); + let elapsed = (now_ms() - request_start).max(0) as u64; + metrics.record_request_duration(elapsed); + if msg.has_reply() { + let response_bytes = serde_json::to_vec(&cached).unwrap_or_default(); + let _ = reply_message(&msg, &transport, &response_bytes).await; + } + return; + } + + let route_key = S::route_key(&request); + metrics.record_cache_miss(); + metrics.record_shard_forward(); + + let coalesce_result = if let Some(coalesce_key) = S::coalesce_key(&request) { + let forward_transport = transport.clone(); + let forward_service = service.clone(); + let forward_ring = ring.clone(); + let forward_route_key = route_key.clone(); + let inflight_key = (route_key, coalesce_key); + inflight + .try_get_with(inflight_key, async move { + forward_to_shard::( + &forward_transport, + forward_service.as_ref(), + forward_ring.as_ref(), + &request, + &forward_route_key, + ) + .await + }) + .await + .map_err(|err| anyhow::anyhow!("{err}")) + } else { + forward_to_shard::( + &transport, + service.as_ref(), + ring.as_ref(), + &request, + &route_key, + ) + .await + }; + + let elapsed = (now_ms() - request_start).max(0) as u64; + metrics.record_request_duration(elapsed); + + match coalesce_result { + Ok(response_bytes) => match rmp_serde::from_slice::(&response_bytes) { + Ok(response) => { + if let Ok(req) = serde_json::from_slice::(&payload) { + service.l1_insert(&req, &response); + } + if msg.has_reply() { + let json = serde_json::to_vec(&response).unwrap_or_default(); + let _ = reply_message(&msg, &transport, &json).await; + } + } + Err(err) => { + debug!(error = %err, "failed to decode shard response"); + if msg.has_reply() { + if serde_json::from_slice::(&response_bytes).is_ok() { + let _ = reply_message(&msg, &transport, &response_bytes).await; + return; + } + let error_response = + serde_json::to_vec(&serde_json::json!({"error": "shard_decode_error"})) + .unwrap_or_default(); + let _ = reply_message(&msg, &transport, &error_response).await; + } + } + }, + Err(err) => { + debug!(error = %err, "shard request failed (coalesced)"); + metrics.record_request_error(); + if msg.has_reply() { + let error_response = + serde_json::to_vec(&serde_json::json!({"error": "shard_unavailable"})) + .unwrap_or_default(); + let _ = reply_message(&msg, &transport, &error_response).await; + } + } + } +} + pub async fn run_router( config: &ServiceConfig, service: S, @@ -89,12 +221,17 @@ where let req_service = service.clone(); let req_queue = queue_group.clone(); let req_metrics = metrics.clone(); + let req_permits = Arc::new(Semaphore::new(config.max_concurrent_requests)); tasks.spawn(async move { loop { let mut sub = req_transport .subscribe_queue(&request_subject, &req_queue) .await?; - info!(subject = request_subject, "router listening for requests"); + info!( + subject = request_subject, + max_concurrent_requests = req_permits.available_permits(), + "router listening for requests" + ); loop { let msg = tokio::select! { @@ -111,123 +248,20 @@ where } }; + let permit = match req_permits.clone().acquire_owned().await { + Ok(permit) => permit, + Err(_) => return anyhow::Ok(()), + }; let transport = req_transport.clone(); let service = req_service.clone(); let ring = ring.clone(); let inflight = inflight.clone(); let metrics = req_metrics.clone(); - let request_start = now_ms(); - metrics.record_request(); - if msg.payload().len() > MAX_ROUTER_REQUEST_BYTES { - warn!( - payload_bytes = msg.payload().len(), - max_payload_bytes = MAX_ROUTER_REQUEST_BYTES, - "rejecting oversized router request" - ); - metrics.record_request_error(); - if msg.has_reply() { - let error_response = - serde_json::to_vec(&serde_json::json!({"error": "request_too_large"})) - .unwrap_or_default(); - let _ = reply_message(&msg, &transport, &error_response).await; - } - continue; - } - let payload = msg.payload().to_vec(); - - let request: S::Request = match serde_json::from_slice(&payload) { - Ok(r) => r, - Err(err) => { - warn!(error = %err, "failed to decode incoming request"); - metrics.record_request_error(); - if msg.has_reply() { - let error_response = - serde_json::to_vec(&serde_json::json!({"error": "decode_error"})) - .unwrap_or_default(); - let _ = reply_message(&msg, &transport, &error_response).await; - } - continue; - } - }; - - if let Some(cached) = service.l1_lookup(&request) { - metrics.record_cache_hit(); - let elapsed = (now_ms() - request_start).max(0) as u64; - metrics.record_request_duration(elapsed); - if msg.has_reply() { - let response_bytes = serde_json::to_vec(&cached).unwrap_or_default(); - let _ = reply_message(&msg, &transport, &response_bytes).await; - } - continue; - } - - let route_key = S::route_key(&request); - metrics.record_cache_miss(); - metrics.record_shard_forward(); - - let coalesce_result = if let Some(coalesce_key) = S::coalesce_key(&request) { - let transport = transport.clone(); - let service = service.clone(); - let ring = ring.clone(); - let route_key = route_key.clone(); - let inflight_key = (route_key.clone(), coalesce_key); - inflight - .try_get_with(inflight_key, async move { - forward_to_shard::(&transport, &service, &ring, &request, &route_key) - .await - }) - .await - .map_err(|err| anyhow::anyhow!("{err}")) - } else { - forward_to_shard::(&transport, &service, &ring, &request, &route_key).await - }; - - let elapsed = (now_ms() - request_start).max(0) as u64; - metrics.record_request_duration(elapsed); - - match coalesce_result { - Ok(response_bytes) => { - match rmp_serde::from_slice::(&response_bytes) { - Ok(response) => { - if let Ok(req) = serde_json::from_slice::(&payload) { - service.l1_insert(&req, &response); - } - if msg.has_reply() { - let json = serde_json::to_vec(&response).unwrap_or_default(); - let _ = reply_message(&msg, &transport, &json).await; - } - } - Err(err) => { - debug!(error = %err, "failed to decode shard response"); - if msg.has_reply() { - if serde_json::from_slice::(&response_bytes) - .is_ok() - { - let _ = - reply_message(&msg, &transport, &response_bytes).await; - continue; - } - let error_response = serde_json::to_vec( - &serde_json::json!({"error": "shard_decode_error"}), - ) - .unwrap_or_default(); - let _ = reply_message(&msg, &transport, &error_response).await; - } - } - } - } - Err(err) => { - debug!(error = %err, "shard request failed (coalesced)"); - metrics.record_request_error(); - if msg.has_reply() { - let error_response = serde_json::to_vec( - &serde_json::json!({"error": "shard_unavailable"}), - ) - .unwrap_or_default(); - let _ = reply_message(&msg, &transport, &error_response).await; - } - } - } + tokio::spawn(async move { + let _permit = permit; + handle_router_request::(msg, transport, service, ring, inflight, metrics) + .await; + }); } } }); @@ -286,15 +320,21 @@ where #[cfg(test)] mod tests { use super::*; + use crate::config::{DatabaseBackend, Mode, ServiceConfig}; + use crate::transport::{InMemoryTransport, Transport, TransportSubscriber, reply_message}; use serde::{Deserialize, Serialize}; + use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering}; + use tokio::sync::Notify; #[derive(Serialize, Deserialize)] struct MockRequest { key: String, } - #[derive(Clone, Serialize, Deserialize)] - struct MockResponse; + #[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] + struct MockResponse { + key: String, + } struct MockRouter; @@ -318,4 +358,139 @@ mod tests { }; assert_eq!(MockRouter::coalesce_key(&request), None); } + + #[tokio::test] + async fn router_forwards_uncached_requests_concurrently() { + let transport = InMemoryTransport::new(); + let mut shard_sub = transport.subscribe("svc.mock.shard.0").await.unwrap(); + + let started = Arc::new(AtomicUsize::new(0)); + let both_started = Arc::new(Notify::new()); + let released = Arc::new(AtomicBool::new(false)); + let release = Arc::new(Notify::new()); + + let shard_transport = transport.clone(); + let shard_started = started.clone(); + let shard_both_started = both_started.clone(); + let shard_released = released.clone(); + let shard_release = release.clone(); + let shard_task = tokio::spawn(async move { + for _ in 0..2 { + let msg = shard_sub.next().await.unwrap(); + let reply_transport = shard_transport.clone(); + let reply_released = shard_released.clone(); + let reply_release = shard_release.clone(); + tokio::spawn(async move { + while !reply_released.load(Ordering::SeqCst) { + reply_release.notified().await; + } + let response = MockResponse { + key: "ok".to_owned(), + }; + let response_bytes = rmp_serde::to_vec_named(&response).unwrap(); + reply_message(&msg, &reply_transport, &response_bytes) + .await + .unwrap(); + }); + + if shard_started.fetch_add(1, Ordering::SeqCst) + 1 == 2 { + shard_both_started.notify_waiters(); + } + } + }); + + let router_config = test_config(2); + let router_transport = transport.clone(); + let router_task = + tokio::spawn( + async move { run_router(&router_config, MockRouter, router_transport).await }, + ); + + tokio::time::sleep(Duration::from_millis(25)).await; + + let request_a = serde_json::to_vec(&MockRequest { + key: "a".to_owned(), + }) + .unwrap(); + let request_b = serde_json::to_vec(&MockRequest { + key: "b".to_owned(), + }) + .unwrap(); + + let client_a = { + let transport = transport.clone(); + tokio::spawn(async move { + transport + .request("svc.mock", &request_a, Duration::from_secs(1)) + .await + }) + }; + let client_b = { + let transport = transport.clone(); + tokio::spawn(async move { + transport + .request("svc.mock", &request_b, Duration::from_secs(1)) + .await + }) + }; + + tokio::time::timeout(Duration::from_millis(250), async { + while started.load(Ordering::SeqCst) < 2 { + both_started.notified().await; + } + }) + .await + .expect("router should forward both requests before the first shard reply is released"); + + released.store(true, Ordering::SeqCst); + release.notify_waiters(); + + let response_a = client_a.await.unwrap().unwrap(); + let response_b = client_b.await.unwrap().unwrap(); + assert_eq!( + serde_json::from_slice::(&response_a).unwrap(), + MockResponse { + key: "ok".to_owned() + } + ); + assert_eq!( + serde_json::from_slice::(&response_b).unwrap(), + MockResponse { + key: "ok".to_owned() + } + ); + + shard_task.await.unwrap(); + router_task.abort(); + } + + fn test_config(max_concurrent_requests: usize) -> ServiceConfig { + ServiceConfig { + service_name: "mock".to_owned(), + mode: Mode::Router, + database_backend: DatabaseBackend::Postgres, + shard_id: 0, + shard_count: 1, + listen_addr: "127.0.0.1:0".parse().unwrap(), + nats_url: "memory".to_owned(), + cache_max_entries: 100, + cache_ttl: Duration::from_secs(30), + cache_hard_ttl: Duration::from_secs(600), + max_concurrent_requests, + scylla_hosts: Vec::new(), + scylla_keyspace: "fluxer".to_owned(), + scylla_username: None, + scylla_password: None, + postgres_url: None, + postgres_host: "127.0.0.1".to_owned(), + postgres_port: 5432, + postgres_database: "fluxer".to_owned(), + postgres_username: "fluxer".to_owned(), + postgres_password: Some("fluxer".to_owned()), + postgres_ssl: false, + postgres_ssl_ca: None, + postgres_max_connections: 1, + postgres_kv_table: "fluxer_kv".to_owned(), + } + } } diff --git a/fluxer_svc/src/transport.rs b/fluxer_svc/src/transport.rs index bfb7d1d50..ce7b4f41d 100644 --- a/fluxer_svc/src/transport.rs +++ b/fluxer_svc/src/transport.rs @@ -3,10 +3,14 @@ use async_trait::async_trait; use bytes::Bytes; use std::sync::Arc; +use std::sync::atomic::{AtomicU64, Ordering}; use std::time::Duration; use tokio::sync::Notify; use tracing::{info, warn}; +const SLOW_CONSUMER_LOG_INTERVAL_MS: u64 = 1_000; +const SLOW_CONSUMER_LOG_NEVER_MS: u64 = u64::MAX; + #[async_trait] pub trait Transport: Clone + Send + Sync + 'static { type Message: TransportMessage + Send + Sync + 'static; @@ -76,8 +80,10 @@ impl NatsTransport { let reconnect_notify = Arc::new(Notify::new()); let event_notify = reconnect_notify.clone(); + let slow_consumer_last_log_ms = Arc::new(AtomicU64::new(SLOW_CONSUMER_LOG_NEVER_MS)); let options = async_nats::ConnectOptions::new().event_callback(move |event| { let notify = event_notify.clone(); + let slow_consumer_last_log_ms = slow_consumer_last_log_ms.clone(); async move { match event { async_nats::Event::Connected => { @@ -91,7 +97,12 @@ impl NatsTransport { warn!("NATS server entering lame duck mode"); } async_nats::Event::SlowConsumer(sid) => { - warn!(subscription_id = sid, "NATS slow consumer detected"); + if should_log_slow_consumer( + crate::metrics::now_ms().max(0) as u64, + &slow_consumer_last_log_ms, + ) { + warn!(subscription_id = sid, "NATS slow consumer detected"); + } } other => { info!(event = %other, "NATS event"); @@ -143,6 +154,23 @@ impl NatsTransport { } } +fn should_log_slow_consumer(now_ms: u64, last_log_ms: &AtomicU64) -> bool { + let mut last = last_log_ms.load(Ordering::Relaxed); + loop { + if last != SLOW_CONSUMER_LOG_NEVER_MS + && now_ms.saturating_sub(last) < SLOW_CONSUMER_LOG_INTERVAL_MS + { + return false; + } + + match last_log_ms.compare_exchange_weak(last, now_ms, Ordering::Relaxed, Ordering::Relaxed) + { + Ok(_) => return true, + Err(actual) => last = actual, + } + } +} + #[async_trait] impl Transport for NatsTransport { type Message = NatsMessage; @@ -231,6 +259,40 @@ impl NatsMessage { } } +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn slow_consumer_logs_initial_event_then_throttles() { + let last_log_ms = AtomicU64::new(SLOW_CONSUMER_LOG_NEVER_MS); + + assert!(should_log_slow_consumer(10, &last_log_ms)); + assert!(!should_log_slow_consumer(1_009, &last_log_ms)); + assert!(should_log_slow_consumer(1_010, &last_log_ms)); + } + + #[test] + fn slow_consumer_throttle_allows_one_racing_logger() { + let last_log_ms = Arc::new(AtomicU64::new(SLOW_CONSUMER_LOG_NEVER_MS)); + let logged = Arc::new(AtomicU64::new(0)); + + std::thread::scope(|scope| { + for _ in 0..8 { + let last_log_ms = last_log_ms.clone(); + let logged = logged.clone(); + scope.spawn(move || { + if should_log_slow_consumer(42, &last_log_ms) { + logged.fetch_add(1, Ordering::Relaxed); + } + }); + } + }); + + assert_eq!(logged.load(Ordering::Relaxed), 1); + } +} + #[cfg(debug_assertions)] mod in_memory { use super::{Transport, TransportMessage, TransportSubscriber};