diff --git a/fluxer_messages/src/router_impl.rs b/fluxer_messages/src/router_impl.rs index 0d54a1d63..4bdcebdca 100644 --- a/fluxer_messages/src/router_impl.rs +++ b/fluxer_messages/src/router_impl.rs @@ -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" } diff --git a/fluxer_svc/src/router.rs b/fluxer_svc/src/router.rs index 5ad0b1a07..45507249f 100644 --- a/fluxer_svc/src/router.rs +++ b/fluxer_svc/src/router.rs @@ -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 { @@ -34,6 +37,11 @@ pub trait RouterService: Send + Sync + 'static { fn l1_invalidate(&self, _key: &str) {} } +fn shard_subject(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( transport: &impl Transport, service: &S, @@ -41,8 +49,7 @@ async fn forward_to_shard( request: &S::Request, route_key: &str, ) -> anyhow::Result> { - 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( .await } +async fn forward_to_shard_verbatim( + transport: &impl Transport, + service: &S, + ring: &HashRing, + request: &S::Request, + route_key: &str, + payload: &[u8], +) -> anyhow::Result> { + 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::(transport, service, ring, request, route_key).await?; + match rmp_serde::from_slice::(&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::(&legacy_bytes).is_ok() { + Ok(legacy_bytes) + } else { + Err(anyhow::anyhow!( + "failed to decode legacy shard response: {err}" + )) + } + } + } +} + +async fn dispatch_to_shard( + transport: &impl Transport, + service: &S, + ring: &HashRing, + request: &S::Request, + route_key: &str, + payload: &[u8], +) -> anyhow::Result> { + if S::CACHES_RESPONSES { + forward_to_shard::(transport, service, ring, request, route_key).await + } else { + forward_to_shard_verbatim::(transport, service, ring, request, route_key, payload).await + } +} + async fn handle_router_request( msg: T::Message, transport: T, @@ -117,27 +176,34 @@ async fn handle_router_request( 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::( + dispatch_to_shard::( &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::( + dispatch_to_shard::( &transport, service.as_ref(), ring.as_ref(), &request, &route_key, + msg.payload(), ) .await }; @@ -146,28 +212,36 @@ async fn handle_router_request( metrics.record_request_duration(elapsed); match coalesce_result { - Ok(response_bytes) => match rmp_serde::from_slice::(&response_bytes) { - Ok(response) => { - service.l1_insert(&request, &response); + Ok(response_bytes) => { + if !S::CACHES_RESPONSES { if msg.has_reply() { - let json = serde_json::to_vec(&response).unwrap_or_default(); - let _ = reply_message(&msg, &transport, &json).await; + let _ = reply_message(&msg, &transport, &response_bytes).await; } + return; } - Err(err) => { - debug!(error = %err, "failed to decode shard response"); - if msg.has_reply() { - if serde_json::from_slice::(&response_bytes).is_ok() { - let _ = reply_message(&msg, &transport, &response_bytes).await; - return; + match rmp_serde::from_slice::(&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; + } + } + Err(err) => { + debug!(error = %err, "failed to decode shard response"); + if msg.has_reply() { + if serde_json::from_slice::(&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) => { 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 { + 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::(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::(&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(); diff --git a/fluxer_svc/src/shard.rs b/fluxer_svc/src/shard.rs index 6b24e974b..0ceb8e670 100644 --- a/fluxer_svc/src/shard.rs +++ b/fluxer_svc/src/shard.rs @@ -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(self, payload: &[u8]) -> anyhow::Result { + match self { + Self::Json => Ok(serde_json::from_slice(payload)?), + Self::MessagePack => Ok(rmp_serde::from_slice(payload)?), + } + } + + fn encode(self, value: &T) -> anyhow::Result> { + 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 { + 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::(&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();