refactor(svc): tidy the rust services and build tooling (#2735)

This commit is contained in:
Hampus
2026-09-13 17:38:32 +02:00
committed by GitHub
parent 33737e0f79
commit 6af33c7188
41 changed files with 863 additions and 1304 deletions
+2 -1
View File
@@ -75,8 +75,9 @@ impl ServiceConfig {
};
let mode = match optional_from(&get, "FLUXER_SVC_MODE").as_deref() {
None | Some("router") => Mode::Router,
Some("shard") => Mode::Shard,
_ => Mode::Router,
Some(other) => anyhow::bail!("unsupported FLUXER_SVC_MODE: {other}"),
};
let shard_count = optional_from(&get, "FLUXER_SVC_SHARD_COUNT")
+4
View File
@@ -9,6 +9,10 @@ impl HashRing {
Self { shard_count }
}
pub fn shard_count(&self) -> u32 {
self.shard_count
}
pub fn owner(&self, route_key: &str) -> u32 {
let mut best_shard = 0u32;
let mut best_score = 0u64;
+154 -83
View File
@@ -2,11 +2,15 @@
use crate::config::ServiceConfig;
use crate::hash_ring::HashRing;
use crate::metrics::{ServiceMetrics, now_ms};
use crate::transport::{Transport, TransportMessage, TransportSubscriber, reply_message};
use crate::metrics::ServiceMetrics;
use crate::transport::{
Transport, TransportMessage, TransportSubscriber, reply_bytes, reply_json_error,
};
use anyhow::Context;
use futures::stream::{FuturesUnordered, StreamExt};
use moka::future::Cache;
use std::sync::Arc;
use std::time::Duration;
use std::time::{Duration, Instant};
use tokio::sync::{Semaphore, TryAcquireError};
use tokio::task::JoinSet;
use tracing::{debug, info, warn};
@@ -14,6 +18,7 @@ use tracing::{debug, info, warn};
pub(crate) const SHARD_REQUEST_TIMEOUT: Duration = Duration::from_secs(5);
const INFLIGHT_TTL: Duration = Duration::from_millis(200);
const INFLIGHT_MAX_ENTRIES: u64 = 10_000;
const MAX_BROADCAST_CONCURRENCY: usize = 32;
const MAX_ROUTER_REQUEST_BYTES: usize = 2 * 1024 * 1024;
const LEGACY_SHARD_DECODE_ERROR: &[u8] = br#"{"error":"shard_request_decode_error"}"#;
type InflightKey = (String, String);
@@ -30,6 +35,13 @@ pub trait RouterService: Send + Sync + 'static {
None
}
fn is_broadcast_request(_request: &Self::Request) -> bool {
false
}
fn is_broadcast_acknowledgement(_response: &Self::Response) -> bool {
false
}
fn l1_lookup(&self, _req: &Self::Request) -> Option<Self::Response> {
None
}
@@ -111,6 +123,88 @@ async fn dispatch_to_shard<S: RouterService>(
}
}
async fn forward_to_all_shards<S: RouterService>(
transport: &impl Transport,
service: &S,
ring: &HashRing,
request: &S::Request,
) -> anyhow::Result<Vec<u8>> {
let payload = rmp_serde::to_vec_named(request)
.context("failed to encode broadcast request as msgpack")?;
let payload = payload.as_slice();
let service_name = service.service_name();
let deadline = tokio::time::Instant::now() + SHARD_REQUEST_TIMEOUT;
let request_shard = |shard_id| async move {
let timeout = deadline.saturating_duration_since(tokio::time::Instant::now());
if timeout.is_zero() {
anyhow::bail!("broadcast deadline expired before shard {shard_id}");
}
let subject = format!("svc.{service_name}.shard.{shard_id}");
let response_bytes = transport
.request(&subject, payload, timeout)
.await
.with_context(|| format!("broadcast request to shard {shard_id} failed"))?;
let response = rmp_serde::from_slice::<S::Response>(&response_bytes)
.with_context(|| format!("invalid broadcast response from shard {shard_id}"))?;
if !S::is_broadcast_acknowledgement(&response) {
anyhow::bail!("unexpected broadcast acknowledgement from shard {shard_id}");
}
anyhow::Ok(response)
};
let mut shard_ids = 0..ring.shard_count();
let mut pending = FuturesUnordered::new();
for shard_id in shard_ids.by_ref().take(MAX_BROADCAST_CONCURRENCY) {
pending.push(request_shard(shard_id));
}
let mut acknowledgement = None;
let mut first_error = None;
while let Some(result) = pending.next().await {
match result {
Ok(response) => acknowledgement = Some(response),
Err(error) => {
first_error.get_or_insert(error);
}
}
if first_error.is_none()
&& let Some(shard_id) = shard_ids.next()
{
pending.push(request_shard(shard_id));
}
}
if let Some(error) = first_error {
return Err(error);
}
let acknowledgement = acknowledgement.context("broadcast request has no configured shards")?;
if S::CACHES_RESPONSES {
rmp_serde::to_vec_named(&acknowledgement)
.context("failed to encode broadcast acknowledgement as msgpack")
} else {
serde_json::to_vec(&acknowledgement)
.context("failed to encode broadcast acknowledgement as json")
}
}
async fn reply_json_response(
message: &impl TransportMessage,
transport: &impl Transport,
response: &impl serde::Serialize,
metrics: &ServiceMetrics,
) {
if !message.has_reply() {
return;
}
match serde_json::to_vec(response) {
Ok(payload) => reply_bytes(message, transport, &payload).await,
Err(error) => {
warn!(error = %error, subject = message.subject(), "failed to encode router response");
metrics.record_request_error();
reply_json_error(message, transport, "encode_error").await;
}
}
}
async fn handle_router_request<S, T>(
msg: T::Message,
transport: T,
@@ -122,7 +216,7 @@ async fn handle_router_request<S, T>(
S: RouterService,
T: Transport,
{
let request_start = now_ms();
let request_start = Instant::now();
metrics.record_request();
if msg.payload().len() > MAX_ROUTER_REQUEST_BYTES {
warn!(
@@ -131,12 +225,7 @@ async fn handle_router_request<S, T>(
"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;
}
reply_json_error(&msg, &transport, "request_too_large").await;
return;
}
let request: S::Request = match serde_json::from_slice(msg.payload()) {
@@ -144,25 +233,20 @@ async fn handle_router_request<S, T>(
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;
}
reply_json_error(&msg, &transport, "decode_error").await;
return;
}
};
let request = Arc::new(request);
let broadcast = S::is_broadcast_request(&request);
if let Some(cached) = service.l1_lookup(&request) {
if S::CACHES_RESPONSES
&& !broadcast
&& 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;
}
metrics.record_request_duration(request_start.elapsed().as_millis() as u64);
reply_json_response(&msg, &transport, &cached, &metrics).await;
return;
}
@@ -170,7 +254,9 @@ async fn handle_router_request<S, T>(
metrics.record_cache_miss();
metrics.record_shard_forward();
let coalesce_result = if let Some(coalesce_key) = S::coalesce_key(&request) {
let coalesce_result = if broadcast {
forward_to_all_shards::<S>(&transport, service.as_ref(), ring.as_ref(), &request).await
} else if let Some(coalesce_key) = S::coalesce_key(&request) {
let forward_transport = transport.clone();
let forward_service = service.clone();
let forward_ring = ring.clone();
@@ -208,36 +294,31 @@ async fn handle_router_request<S, T>(
.await
};
let elapsed = (now_ms() - request_start).max(0) as u64;
metrics.record_request_duration(elapsed);
metrics.record_request_duration(request_start.elapsed().as_millis() as u64);
match coalesce_result {
Ok(response_bytes) => {
if !S::CACHES_RESPONSES {
if msg.has_reply() {
let _ = reply_message(&msg, &transport, &response_bytes).await;
reply_bytes(&msg, &transport, &response_bytes).await;
}
return;
}
match rmp_serde::from_slice::<S::Response>(&response_bytes) {
Ok(response) => {
service.l1_insert(&request, &response);
if msg.has_reply() {
let json = serde_json::to_vec(&response).unwrap_or_default();
let _ = reply_message(&msg, &transport, &json).await;
if !broadcast {
service.l1_insert(&request, &response);
}
reply_json_response(&msg, &transport, &response, &metrics).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;
reply_bytes(&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;
reply_json_error(&msg, &transport, "shard_decode_error").await;
}
}
}
@@ -245,12 +326,7 @@ async fn handle_router_request<S, T>(
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;
}
reply_json_error(&msg, &transport, "shard_unavailable").await;
}
}
}
@@ -267,7 +343,6 @@ where
let ring = Arc::new(HashRing::new(config.shard_count));
let name = service.service_name().to_owned();
let request_subject = format!("svc.{name}");
let invalidate_subject = format!("svc.{name}.invalidate.>");
let queue_group = format!("{name}-router");
let metrics = Arc::new(ServiceMetrics::default());
@@ -295,6 +370,7 @@ where
let req_metrics = metrics.clone();
let req_permits = Arc::new(Semaphore::new(config.max_concurrent_requests));
tasks.spawn(async move {
let mut requests = JoinSet::new();
loop {
let mut sub = req_transport
.subscribe_queue(&request_subject, &req_queue)
@@ -307,6 +383,12 @@ where
loop {
let msg = tokio::select! {
result = requests.join_next(), if !requests.is_empty() => {
if let Err(err) = result.expect("nonempty router request set") {
warn!(error = %err, "router request task failed");
}
continue;
}
msg_opt = sub.next() => {
let Some(msg) = msg_opt else {
warn!("router request subscription stream ended, will re-subscribe");
@@ -314,24 +396,20 @@ where
};
msg
}
_ = req_transport.wait_for_reconnect() => {
info!("NATS reconnected, re-subscribing router request listener");
break;
}
};
while let Some(result) = requests.try_join_next() {
if let Err(err) = result {
warn!(error = %err, "router request task failed");
}
}
let permit = match req_permits.clone().try_acquire_owned() {
Ok(permit) => permit,
Err(TryAcquireError::NoPermits) => {
debug!("shedding router request, no permits available");
req_metrics.record_request();
req_metrics.record_request_error();
if msg.has_reply() {
let error_response =
serde_json::to_vec(&serde_json::json!({"error": "overloaded"}))
.unwrap_or_default();
let _ = reply_message(&msg, &req_transport, &error_response).await;
}
reply_json_error(&msg, &req_transport, "overloaded").await;
continue;
}
Err(TryAcquireError::Closed) => return anyhow::Ok(()),
@@ -341,7 +419,7 @@ where
let ring = ring.clone();
let inflight = inflight.clone();
let metrics = req_metrics.clone();
tokio::spawn(async move {
requests.spawn(async move {
let _permit = permit;
handle_router_request::<S, _>(msg, transport, service, ring, inflight, metrics)
.await;
@@ -350,39 +428,32 @@ where
}
});
let inv_transport = transport.clone();
let inv_service = service.clone();
tasks.spawn(async move {
loop {
let mut sub = inv_transport.subscribe(&invalidate_subject).await?;
info!(
subject = invalidate_subject,
"router listening for cache invalidations"
);
if S::CACHES_RESPONSES {
let inv_transport = transport.clone();
let inv_service = service.clone();
let invalidate_prefix = format!("svc.{name}.invalidate.");
let invalidate_subject = format!("{invalidate_prefix}>");
tasks.spawn(async move {
loop {
tokio::select! {
msg_opt = sub.next() => {
let Some(msg) = msg_opt else {
warn!("router invalidation subscription stream ended, will re-subscribe");
break;
};
let subject = msg.subject().to_owned();
let key = subject
.strip_prefix(&format!("svc.{name}.invalidate."))
.unwrap_or("");
if !key.is_empty() {
inv_service.l1_invalidate(key);
}
}
_ = inv_transport.wait_for_reconnect() => {
info!("NATS reconnected, re-subscribing router invalidation listener");
break;
let mut sub = inv_transport.subscribe(&invalidate_subject).await?;
info!(
subject = invalidate_subject,
"router listening for cache invalidations"
);
while let Some(msg) = sub.next().await {
if let Some(key) = msg
.subject()
.strip_prefix(&invalidate_prefix)
.filter(|key| !key.is_empty())
{
inv_service.l1_invalidate(key);
}
}
warn!("router invalidation subscription stream ended, will re-subscribe");
}
}
});
});
}
tokio::select! {
result = tasks.join_next() => {
+70 -54
View File
@@ -1,11 +1,15 @@
// SPDX-License-Identifier: AGPL-3.0-or-later
use crate::config::ServiceConfig;
use crate::metrics::{ServiceMetrics, now_ms};
use crate::transport::{Transport, TransportMessage, TransportSubscriber, reply_message};
use crate::metrics::ServiceMetrics;
use crate::transport::{
Transport, TransportMessage, TransportSubscriber, reply_bytes, reply_json_error,
};
use anyhow::Context;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use tokio::sync::{Semaphore, TryAcquireError};
use std::time::Instant;
use tokio::sync::{Semaphore, TryAcquireError, oneshot};
use tokio::task::JoinSet;
use tracing::{debug, info, warn};
@@ -93,9 +97,14 @@ where
let shard_is_serving = is_serving.clone();
let shard_permits = request_permits.clone();
let shard_metrics = metrics.clone();
tasks.spawn(async move {
loop {
let mut sub = shard_transport.subscribe(&shard_subject).await?;
let (stop_requests, mut stop_requests_rx) = oneshot::channel();
let request_task = tasks.spawn(async move {
let mut requests = JoinSet::new();
'listening: loop {
let mut sub = tokio::select! {
_ = &mut stop_requests_rx => break 'listening,
result = shard_transport.subscribe(&shard_subject) => result?,
};
info!(
subject = shard_subject,
shard_id,
@@ -105,12 +114,23 @@ where
loop {
tokio::select! {
_ = &mut stop_requests_rx => break 'listening,
result = requests.join_next(), if !requests.is_empty() => {
if let Err(err) = result.expect("nonempty shard request set") {
warn!(error = %err, "shard request task failed");
}
}
msg_opt = sub.next() => {
let Some(msg) = msg_opt else {
warn!("shard subscription stream ended, will re-subscribe");
break;
};
while let Some(result) = requests.try_join_next() {
if let Err(err) = result {
warn!(error = %err, "shard request task failed");
}
}
if !shard_is_serving.load(Ordering::SeqCst) {
continue;
}
@@ -133,27 +153,25 @@ where
debug!("shedding shard request, no permits available");
shard_metrics.record_request();
shard_metrics.record_request_error();
reply_shard_error(&msg, &transport, "overloaded").await;
reply_json_error(&msg, &transport, "overloaded").await;
continue;
}
Err(TryAcquireError::Closed) => return anyhow::Ok(()),
};
let raw_payload = msg.payload().to_vec();
tokio::spawn(async move {
requests.spawn(async move {
let _permit = permit;
if !is_serving.load(Ordering::SeqCst) {
return;
}
let request_start = now_ms();
let request_start = Instant::now();
metrics.record_request();
let encoding = WireEncoding::detect(&raw_payload);
let request: S::Request = match encoding.decode(&raw_payload) {
let encoding = WireEncoding::detect(msg.payload());
let request: S::Request = match encoding.decode(msg.payload()) {
Ok(r) => r,
Err(err) => {
warn!(error = %err, ?encoding, "failed to decode shard request");
metrics.record_request_error();
reply_shard_error(&msg, &transport, "shard_request_decode_error")
reply_json_error(&msg, &transport, "shard_request_decode_error")
.await;
return;
}
@@ -161,25 +179,20 @@ where
match service.handle(request).await {
Ok(response) => {
let elapsed = (now_ms() - request_start).max(0) as u64;
metrics.record_request_duration(elapsed);
metrics.record_request_duration(request_start.elapsed().as_millis() as u64);
if msg.has_reply() {
match encoding.encode(&response) {
Ok(response_bytes) => {
if let Err(err) =
reply_message(&msg, &transport, &response_bytes).await
{
debug!(
error = %err,
"failed to send shard reply"
);
}
reply_bytes(&msg, &transport, &response_bytes).await;
}
Err(err) => {
warn!(
error = %err,
?encoding,
"failed to encode shard response"
);
metrics.record_request_error();
reply_json_error(&msg, &transport, "encode_error").await;
}
}
}
@@ -187,20 +200,19 @@ where
Err(err) => {
warn!(error = %err, "shard handler returned error");
metrics.record_request_error();
let elapsed = (now_ms() - request_start).max(0) as u64;
metrics.record_request_duration(elapsed);
reply_shard_error(&msg, &transport, "shard_handler_error").await;
metrics.record_request_duration(request_start.elapsed().as_millis() as u64);
reply_json_error(&msg, &transport, "shard_handler_error").await;
}
}
});
}
_ = shard_transport.wait_for_reconnect() => {
info!("NATS reconnected, re-subscribing shard listener");
break;
}
}
}
}
while let Some(result) = requests.join_next().await {
result.context("shard request task failed while draining")?;
}
anyhow::Ok(())
});
tokio::select! {
@@ -217,20 +229,34 @@ where
is_serving.store(false, Ordering::SeqCst);
let max_permits = config.max_concurrent_requests;
let drain_permits = request_permits.clone();
crate::shutdown::drain_with_timeout(
async move {
if let Ok(_permit) = drain_permits.acquire_many(max_permits as u32).await {
info!(
max_concurrent_requests = max_permits,
"all in-flight requests drained"
);
let _ = stop_requests.send(());
let drain = async {
loop {
let (task_id, result) = tasks.join_next_with_id().await
.expect("running shard listener while draining")
.context("shard service task failed while draining")?;
result?;
if task_id == request_task.id() {
return anyhow::Ok(());
}
},
crate::shutdown::DEFAULT_DRAIN_TIMEOUT,
)
.await;
}
};
match tokio::time::timeout(crate::shutdown::DEFAULT_DRAIN_TIMEOUT, drain).await {
Ok(result) => {
result?;
info!(
max_concurrent_requests = config.max_concurrent_requests,
"all in-flight requests drained"
);
info!("graceful drain completed");
}
Err(_) => {
warn!(
timeout_secs = crate::shutdown::DEFAULT_DRAIN_TIMEOUT.as_secs(),
"drain timeout exceeded, proceeding with shutdown"
);
}
}
info!("shard shutdown complete");
Ok(())
@@ -238,16 +264,6 @@ where
}
}
async fn reply_shard_error(msg: &impl TransportMessage, transport: &impl Transport, code: &str) {
if !msg.has_reply() {
return;
}
let response = serde_json::to_vec(&serde_json::json!({ "error": code })).unwrap_or_default();
if let Err(err) = reply_message(msg, transport, &response).await {
debug!(error = %err, "failed to send shard error reply");
}
}
#[cfg(test)]
mod tests {
use super::*;
+24 -1
View File
@@ -6,7 +6,7 @@ use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering};
use std::time::Duration;
use tokio::sync::Notify;
use tracing::{info, warn};
use tracing::{debug, info, warn};
const NATS_SUBSCRIPTION_CAPACITY: usize = 8_192;
const SLOW_CONSUMER_LOG_INTERVAL_MS: u64 = 1_000;
@@ -62,6 +62,29 @@ where
Ok(())
}
pub(crate) async fn reply_bytes(
message: &impl TransportMessage,
transport: &impl Transport,
payload: &[u8],
) {
if let Err(error) = reply_message(message, transport, payload).await {
debug!(error = %error, subject = message.subject(), "failed to send service reply");
}
}
pub(crate) async fn reply_json_error(
message: &impl TransportMessage,
transport: &impl Transport,
code: &str,
) {
if !message.has_reply() {
return;
}
let payload = serde_json::to_vec(&serde_json::json!({ "error": code }))
.expect("service error responses contain only a JSON-serializable string");
reply_bytes(message, transport, &payload).await;
}
#[derive(Clone)]
pub struct NatsTransport {
client: async_nats::Client,