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 Response = MessageResponse;
const CACHES_RESPONSES: bool = false;
fn service_name(&self) -> &str {
"messages"
}
+202 -6
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_MAX_ENTRIES: u64 = 10_000;
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);
pub trait RouterService: Send + Sync + 'static {
type Request: serde::Serialize + serde::de::DeserializeOwned + 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 route_key(request: &Self::Request) -> 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 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>(
transport: &impl Transport,
service: &S,
@@ -41,8 +49,7 @@ async fn forward_to_shard<S: RouterService>(
request: &S::Request,
route_key: &str,
) -> anyhow::Result<Vec<u8>> {
let shard_id = ring.owner(route_key);
let shard_subject = format!("svc.{}.shard.{shard_id}", service.service_name());
let shard_subject = shard_subject(service, ring, route_key);
let msgpack_payload = rmp_serde::to_vec_named(request)
.map_err(|e| anyhow::anyhow!("failed to encode request as msgpack: {e}"))?;
@@ -52,6 +59,58 @@ async fn forward_to_shard<S: RouterService>(
.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>(
msg: T::Message,
transport: T,
@@ -117,27 +176,34 @@ async fn handle_router_request<S, T>(
let forward_ring = ring.clone();
let forward_route_key = route_key.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);
inflight
.try_get_with(inflight_key, async move {
forward_to_shard::<S>(
dispatch_to_shard::<S>(
&forward_transport,
forward_service.as_ref(),
forward_ring.as_ref(),
&forward_request,
&forward_route_key,
&forward_payload,
)
.await
})
.await
.map_err(|err| anyhow::anyhow!("{err}"))
} else {
forward_to_shard::<S>(
dispatch_to_shard::<S>(
&transport,
service.as_ref(),
ring.as_ref(),
&request,
&route_key,
msg.payload(),
)
.await
};
@@ -146,7 +212,14 @@ async fn handle_router_request<S, T>(
metrics.record_request_duration(elapsed);
match coalesce_result {
Ok(response_bytes) => match rmp_serde::from_slice::<S::Response>(&response_bytes) {
Ok(response_bytes) => {
if !S::CACHES_RESPONSES {
if msg.has_reply() {
let _ = reply_message(&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() {
@@ -167,7 +240,8 @@ async fn handle_router_request<S, T>(
let _ = reply_message(&msg, &transport, &error_response).await;
}
}
},
}
}
Err(err) => {
debug!(error = %err, "shard request failed (coalesced)");
metrics.record_request_error();
@@ -601,6 +675,128 @@ mod tests {
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]
async fn router_sheds_requests_when_permits_are_exhausted() {
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;
#[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 {
type Request: 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();
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,
Err(err) => {
warn!(error = %err, "failed to decode shard request");
warn!(error = %err, ?encoding, "failed to decode shard request");
metrics.record_request_error();
reply_shard_error(&msg, &transport, "shard_request_decode_error")
.await;
@@ -134,7 +164,7 @@ where
let elapsed = (now_ms() - request_start).max(0) as u64;
metrics.record_request_duration(elapsed);
if msg.has_reply() {
match rmp_serde::to_vec_named(&response) {
match encoding.encode(&response) {
Ok(response_bytes) => {
if let Err(err) =
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]
async fn shard_sheds_requests_when_permits_are_exhausted() {
let transport = InMemoryTransport::new();