perf(svc): forward messages shard replies without transcoding (#2224)

This commit is contained in:
Hampus
2026-08-31 01:04:47 +02:00
committed by GitHub
parent 21b1e4e719
commit 3594cbd5ca
3 changed files with 324 additions and 23 deletions
+2
View File
@@ -21,6 +21,8 @@ impl RouterService for MessagesRouter {
type Request = MessageRequest; type Request = MessageRequest;
type Response = MessageResponse; type Response = MessageResponse;
const CACHES_RESPONSES: bool = false;
fn service_name(&self) -> &str { fn service_name(&self) -> &str {
"messages" "messages"
} }
+216 -20
View File
@@ -15,12 +15,15 @@ pub(crate) const SHARD_REQUEST_TIMEOUT: Duration = Duration::from_secs(5);
const INFLIGHT_TTL: Duration = Duration::from_millis(200); const INFLIGHT_TTL: Duration = Duration::from_millis(200);
const INFLIGHT_MAX_ENTRIES: u64 = 10_000; const INFLIGHT_MAX_ENTRIES: u64 = 10_000;
const MAX_ROUTER_REQUEST_BYTES: usize = 2 * 1024 * 1024; 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); type InflightKey = (String, String);
pub trait RouterService: Send + Sync + 'static { pub trait RouterService: Send + Sync + 'static {
type Request: serde::Serialize + serde::de::DeserializeOwned + Send + Sync + 'static; type Request: serde::Serialize + serde::de::DeserializeOwned + Send + Sync + 'static;
type Response: serde::Serialize + serde::de::DeserializeOwned + Clone + Send + Sync + 'static; type Response: serde::Serialize + serde::de::DeserializeOwned + Clone + Send + Sync + 'static;
const CACHES_RESPONSES: bool = true;
fn service_name(&self) -> &str; fn service_name(&self) -> &str;
fn route_key(request: &Self::Request) -> String; fn route_key(request: &Self::Request) -> String;
fn coalesce_key(_request: &Self::Request) -> Option<String> { fn coalesce_key(_request: &Self::Request) -> Option<String> {
@@ -34,6 +37,11 @@ pub trait RouterService: Send + Sync + 'static {
fn l1_invalidate(&self, _key: &str) {} fn l1_invalidate(&self, _key: &str) {}
} }
fn shard_subject<S: RouterService>(service: &S, ring: &HashRing, route_key: &str) -> String {
let shard_id = ring.owner(route_key);
format!("svc.{}.shard.{shard_id}", service.service_name())
}
async fn forward_to_shard<S: RouterService>( async fn forward_to_shard<S: RouterService>(
transport: &impl Transport, transport: &impl Transport,
service: &S, service: &S,
@@ -41,8 +49,7 @@ async fn forward_to_shard<S: RouterService>(
request: &S::Request, request: &S::Request,
route_key: &str, route_key: &str,
) -> anyhow::Result<Vec<u8>> { ) -> anyhow::Result<Vec<u8>> {
let shard_id = ring.owner(route_key); let shard_subject = shard_subject(service, ring, route_key);
let shard_subject = format!("svc.{}.shard.{shard_id}", service.service_name());
let msgpack_payload = rmp_serde::to_vec_named(request) let msgpack_payload = rmp_serde::to_vec_named(request)
.map_err(|e| anyhow::anyhow!("failed to encode request as msgpack: {e}"))?; .map_err(|e| anyhow::anyhow!("failed to encode request as msgpack: {e}"))?;
@@ -52,6 +59,58 @@ async fn forward_to_shard<S: RouterService>(
.await .await
} }
async fn forward_to_shard_verbatim<S: RouterService>(
transport: &impl Transport,
service: &S,
ring: &HashRing,
request: &S::Request,
route_key: &str,
payload: &[u8],
) -> anyhow::Result<Vec<u8>> {
let shard_subject = shard_subject(service, ring, route_key);
let response_bytes = transport
.request(&shard_subject, payload, SHARD_REQUEST_TIMEOUT)
.await?;
if response_bytes != LEGACY_SHARD_DECODE_ERROR {
return Ok(response_bytes);
}
warn!(
subject = shard_subject,
"shard rejected a pass-through request, retrying with the legacy msgpack encoding"
);
let legacy_bytes = forward_to_shard::<S>(transport, service, ring, request, route_key).await?;
match rmp_serde::from_slice::<S::Response>(&legacy_bytes) {
Ok(response) => serde_json::to_vec(&response).map_err(|err| {
anyhow::anyhow!("failed to encode legacy shard response as json: {err}")
}),
Err(err) => {
if serde_json::from_slice::<serde_json::Value>(&legacy_bytes).is_ok() {
Ok(legacy_bytes)
} else {
Err(anyhow::anyhow!(
"failed to decode legacy shard response: {err}"
))
}
}
}
}
async fn dispatch_to_shard<S: RouterService>(
transport: &impl Transport,
service: &S,
ring: &HashRing,
request: &S::Request,
route_key: &str,
payload: &[u8],
) -> anyhow::Result<Vec<u8>> {
if S::CACHES_RESPONSES {
forward_to_shard::<S>(transport, service, ring, request, route_key).await
} else {
forward_to_shard_verbatim::<S>(transport, service, ring, request, route_key, payload).await
}
}
async fn handle_router_request<S, T>( async fn handle_router_request<S, T>(
msg: T::Message, msg: T::Message,
transport: T, transport: T,
@@ -117,27 +176,34 @@ async fn handle_router_request<S, T>(
let forward_ring = ring.clone(); let forward_ring = ring.clone();
let forward_route_key = route_key.clone(); let forward_route_key = route_key.clone();
let forward_request = request.clone(); let forward_request = request.clone();
let forward_payload = if S::CACHES_RESPONSES {
Vec::new()
} else {
msg.payload().to_vec()
};
let inflight_key = (route_key, coalesce_key); let inflight_key = (route_key, coalesce_key);
inflight inflight
.try_get_with(inflight_key, async move { .try_get_with(inflight_key, async move {
forward_to_shard::<S>( dispatch_to_shard::<S>(
&forward_transport, &forward_transport,
forward_service.as_ref(), forward_service.as_ref(),
forward_ring.as_ref(), forward_ring.as_ref(),
&forward_request, &forward_request,
&forward_route_key, &forward_route_key,
&forward_payload,
) )
.await .await
}) })
.await .await
.map_err(|err| anyhow::anyhow!("{err}")) .map_err(|err| anyhow::anyhow!("{err}"))
} else { } else {
forward_to_shard::<S>( dispatch_to_shard::<S>(
&transport, &transport,
service.as_ref(), service.as_ref(),
ring.as_ref(), ring.as_ref(),
&request, &request,
&route_key, &route_key,
msg.payload(),
) )
.await .await
}; };
@@ -146,28 +212,36 @@ async fn handle_router_request<S, T>(
metrics.record_request_duration(elapsed); metrics.record_request_duration(elapsed);
match coalesce_result { match coalesce_result {
Ok(response_bytes) => match rmp_serde::from_slice::<S::Response>(&response_bytes) { Ok(response_bytes) => {
Ok(response) => { if !S::CACHES_RESPONSES {
service.l1_insert(&request, &response);
if msg.has_reply() { if msg.has_reply() {
let json = serde_json::to_vec(&response).unwrap_or_default(); let _ = reply_message(&msg, &transport, &response_bytes).await;
let _ = reply_message(&msg, &transport, &json).await;
} }
return;
} }
Err(err) => { match rmp_serde::from_slice::<S::Response>(&response_bytes) {
debug!(error = %err, "failed to decode shard response"); Ok(response) => {
if msg.has_reply() { service.l1_insert(&request, &response);
if serde_json::from_slice::<serde_json::Value>(&response_bytes).is_ok() { if msg.has_reply() {
let _ = reply_message(&msg, &transport, &response_bytes).await; let json = serde_json::to_vec(&response).unwrap_or_default();
return; 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;
} }
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) => { Err(err) => {
debug!(error = %err, "shard request failed (coalesced)"); debug!(error = %err, "shard request failed (coalesced)");
metrics.record_request_error(); metrics.record_request_error();
@@ -601,6 +675,128 @@ mod tests {
router_task.abort(); router_task.abort();
} }
struct PassThroughRouter;
impl RouterService for PassThroughRouter {
type Request = MockRequest;
type Response = MockResponse;
const CACHES_RESPONSES: bool = false;
fn service_name(&self) -> &str {
"passthrough-mock"
}
fn route_key(request: &MockRequest) -> String {
request.key.clone()
}
fn coalesce_key(request: &MockRequest) -> Option<String> {
Some(request.key.clone())
}
}
#[tokio::test]
async fn router_forwards_pass_through_requests_and_replies_verbatim() {
let transport = InMemoryTransport::new();
let mut shard_sub = transport
.subscribe("svc.passthrough-mock.shard.0")
.await
.unwrap();
let observed = Arc::new(Mutex::new(Vec::new()));
let shard_transport = transport.clone();
let shard_observed = observed.clone();
let shard_task = tokio::spawn(async move {
while let Some(msg) = shard_sub.next().await {
shard_observed.lock().unwrap().push(msg.payload().to_vec());
reply_message(&msg, &shard_transport, br#"{ "key" : "verbatim" }"#)
.await
.unwrap();
}
});
let router_config = test_config(4);
let router_transport = transport.clone();
let router_task = tokio::spawn(async move {
run_router(&router_config, PassThroughRouter, router_transport).await
});
tokio::time::sleep(Duration::from_millis(25)).await;
let request = serde_json::to_vec(&MockRequest {
key: "a".to_owned(),
})
.unwrap();
let response = transport
.request("svc.passthrough-mock", &request, Duration::from_secs(1))
.await
.unwrap();
assert_eq!(response, br#"{ "key" : "verbatim" }"#);
assert_eq!(observed.lock().unwrap().as_slice(), [request]);
shard_task.abort();
router_task.abort();
}
#[tokio::test]
async fn router_retries_pass_through_requests_that_legacy_shards_reject() {
let transport = InMemoryTransport::new();
let mut shard_sub = transport
.subscribe("svc.passthrough-mock.shard.0")
.await
.unwrap();
let shard_transport = transport.clone();
let shard_task = tokio::spawn(async move {
while let Some(msg) = shard_sub.next().await {
let response_bytes = match rmp_serde::from_slice::<MockRequest>(msg.payload()) {
Ok(_) => rmp_serde::to_vec_named(&MockResponse {
key: "legacy".to_owned(),
})
.unwrap(),
Err(_) => serde_json::to_vec(
&serde_json::json!({"error": "shard_request_decode_error"}),
)
.unwrap(),
};
reply_message(&msg, &shard_transport, &response_bytes)
.await
.unwrap();
}
});
let router_config = test_config(4);
let router_transport = transport.clone();
let router_task = tokio::spawn(async move {
run_router(&router_config, PassThroughRouter, router_transport).await
});
tokio::time::sleep(Duration::from_millis(25)).await;
let request = serde_json::to_vec(&MockRequest {
key: "a".to_owned(),
})
.unwrap();
let response = transport
.request("svc.passthrough-mock", &request, Duration::from_secs(1))
.await
.unwrap();
assert_eq!(
serde_json::from_slice::<MockResponse>(&response).unwrap(),
MockResponse {
key: "legacy".to_owned()
}
);
shard_task.abort();
router_task.abort();
}
#[tokio::test] #[tokio::test]
async fn router_sheds_requests_when_permits_are_exhausted() { async fn router_sheds_requests_when_permits_are_exhausted() {
let transport = InMemoryTransport::new(); let transport = InMemoryTransport::new();
+106 -3
View File
@@ -11,6 +11,35 @@ use tracing::{debug, info, warn};
const MAX_SHARD_REQUEST_BYTES: usize = 2 * 1024 * 1024; const MAX_SHARD_REQUEST_BYTES: usize = 2 * 1024 * 1024;
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
enum WireEncoding {
Json,
MessagePack,
}
impl WireEncoding {
fn detect(payload: &[u8]) -> Self {
match payload.iter().find(|byte| !byte.is_ascii_whitespace()) {
Some(b'{') => Self::Json,
_ => Self::MessagePack,
}
}
fn decode<T: serde::de::DeserializeOwned>(self, payload: &[u8]) -> anyhow::Result<T> {
match self {
Self::Json => Ok(serde_json::from_slice(payload)?),
Self::MessagePack => Ok(rmp_serde::from_slice(payload)?),
}
}
fn encode<T: serde::Serialize>(self, value: &T) -> anyhow::Result<Vec<u8>> {
match self {
Self::Json => Ok(serde_json::to_vec(value)?),
Self::MessagePack => Ok(rmp_serde::to_vec_named(value)?),
}
}
}
pub trait ShardService: Send + Sync + 'static { pub trait ShardService: Send + Sync + 'static {
type Request: serde::Serialize + serde::de::DeserializeOwned + Send + 'static; type Request: serde::Serialize + serde::de::DeserializeOwned + Send + 'static;
type Response: serde::Serialize + serde::de::DeserializeOwned + Send + 'static; type Response: serde::Serialize + serde::de::DeserializeOwned + Send + 'static;
@@ -118,10 +147,11 @@ where
} }
let request_start = now_ms(); let request_start = now_ms();
metrics.record_request(); metrics.record_request();
let request: S::Request = match rmp_serde::from_slice(&raw_payload) { let encoding = WireEncoding::detect(&raw_payload);
let request: S::Request = match encoding.decode(&raw_payload) {
Ok(r) => r, Ok(r) => r,
Err(err) => { Err(err) => {
warn!(error = %err, "failed to decode shard request"); warn!(error = %err, ?encoding, "failed to decode shard request");
metrics.record_request_error(); metrics.record_request_error();
reply_shard_error(&msg, &transport, "shard_request_decode_error") reply_shard_error(&msg, &transport, "shard_request_decode_error")
.await; .await;
@@ -134,7 +164,7 @@ where
let elapsed = (now_ms() - request_start).max(0) as u64; let elapsed = (now_ms() - request_start).max(0) as u64;
metrics.record_request_duration(elapsed); metrics.record_request_duration(elapsed);
if msg.has_reply() { if msg.has_reply() {
match rmp_serde::to_vec_named(&response) { match encoding.encode(&response) {
Ok(response_bytes) => { Ok(response_bytes) => {
if let Err(err) = if let Err(err) =
reply_message(&msg, &transport, &response_bytes).await reply_message(&msg, &transport, &response_bytes).await
@@ -255,6 +285,79 @@ mod tests {
} }
} }
struct EchoShard;
impl ShardService for EchoShard {
type Request = MockRequest;
type Response = MockResponse;
fn service_name(&self) -> &str {
"mock"
}
async fn handle(&self, request: MockRequest) -> anyhow::Result<MockResponse> {
Ok(MockResponse { key: request.key })
}
}
#[test]
fn wire_encoding_detects_json_and_msgpack_payloads() {
let json = serde_json::to_vec(&MockRequest {
key: "a".to_owned(),
})
.unwrap();
let msgpack = rmp_serde::to_vec_named(&MockRequest {
key: "a".to_owned(),
})
.unwrap();
assert_eq!(WireEncoding::detect(&json), WireEncoding::Json);
assert_eq!(
WireEncoding::detect(b" \n{\"key\":\"a\"}"),
WireEncoding::Json
);
assert_eq!(WireEncoding::detect(&msgpack), WireEncoding::MessagePack);
assert_eq!(WireEncoding::detect(b""), WireEncoding::MessagePack);
}
#[tokio::test]
async fn shard_replies_in_the_request_encoding() {
let transport = InMemoryTransport::new();
let config = test_config(4);
let shard_transport = transport.clone();
let shard_task =
tokio::spawn(async move { run_shard(&config, EchoShard, shard_transport).await });
tokio::time::sleep(Duration::from_millis(25)).await;
let msgpack_request = rmp_serde::to_vec_named(&MockRequest {
key: "a".to_owned(),
})
.unwrap();
let msgpack_response = transport
.request("svc.mock.shard.0", &msgpack_request, Duration::from_secs(1))
.await
.unwrap();
assert_eq!(
rmp_serde::from_slice::<MockResponse>(&msgpack_response)
.unwrap()
.key,
"a"
);
let json_request = serde_json::to_vec(&MockRequest {
key: "b".to_owned(),
})
.unwrap();
let json_response = transport
.request("svc.mock.shard.0", &json_request, Duration::from_secs(1))
.await
.unwrap();
assert_eq!(json_response, br#"{"key":"b"}"#);
shard_task.abort();
}
#[tokio::test] #[tokio::test]
async fn shard_sheds_requests_when_permits_are_exhausted() { async fn shard_sheds_requests_when_permits_are_exhausted() {
let transport = InMemoryTransport::new(); let transport = InMemoryTransport::new();