fix(fluxer_svc): drain router requests concurrently (#1221)

This commit is contained in:
Hampus
2026-06-30 01:47:23 +02:00
committed by GitHub
parent 194483e98a
commit 31665225eb
2 changed files with 353 additions and 116 deletions
+290 -115
View File
@@ -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<S: RouterService>(
.await
}
async fn handle_router_request<S, T>(
msg: T::Message,
transport: T,
service: Arc<S>,
ring: Arc<HashRing>,
inflight: Cache<InflightKey, Vec<u8>>,
metrics: Arc<ServiceMetrics>,
) 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::<S>(
&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::<S>(
&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::<S::Response>(&response_bytes) {
Ok(response) => {
if let Ok(req) = serde_json::from_slice::<S::Request>(&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::<serde_json::Value>(&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<S>(
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::<S>(&transport, &service, &ring, &request, &route_key)
.await
})
.await
.map_err(|err| anyhow::anyhow!("{err}"))
} else {
forward_to_shard::<S>(&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::<S::Response>(&response_bytes) {
Ok(response) => {
if let Ok(req) = serde_json::from_slice::<S::Request>(&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::<serde_json::Value>(&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::<S, _>(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::<MockResponse>(&response_a).unwrap(),
MockResponse {
key: "ok".to_owned()
}
);
assert_eq!(
serde_json::from_slice::<MockResponse>(&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(),
}
}
}
+62
View File
@@ -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,8 +97,13 @@ impl NatsTransport {
warn!("NATS server entering lame duck mode");
}
async_nats::Event::SlowConsumer(sid) => {
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};