mirror of
https://github.com/fluxerapp/fluxer
synced 2026-10-07 19:22:14 +09:00
fix(fluxer_svc): drain router requests concurrently (#1221)
This commit is contained in:
+290
-115
@@ -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(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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};
|
||||
|
||||
Reference in New Issue
Block a user