mirror of
https://github.com/fluxerapp/fluxer
synced 2026-10-07 19:22:14 +09:00
refactor(svc): tidy the rust services and build tooling (#2735)
This commit is contained in:
@@ -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")
|
||||
|
||||
@@ -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
@@ -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
@@ -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::*;
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user