1use std::collections::BTreeMap;
5use std::sync::Arc;
6use std::time::Instant;
7
8use super::unified_client::RequestPlaneClient;
9use super::*;
10use crate::component::Instance;
11use crate::discovery::EndpointInstanceId;
12use crate::dynamo_nvtx_range;
13use crate::engine::{AsyncEngine, AsyncEngineContextProvider, Data};
14use crate::error::{DynamoError, ErrorType};
15use crate::logging::inject_trace_headers_into_map;
16use crate::metrics::frontend_perf::STAGE_DURATION_SECONDS;
17use crate::metrics::request_plane::{
18 REQUEST_PLANE_INFLIGHT, REQUEST_PLANE_QUEUE_SECONDS, REQUEST_PLANE_ROUNDTRIP_TTFT_SECONDS,
19 REQUEST_PLANE_SEND_SECONDS,
20};
21use crate::pipeline::network::ConnectionInfo;
22use crate::pipeline::network::NetworkStreamWrapper;
23use crate::pipeline::network::PendingConnections;
24use crate::pipeline::network::RegisteredStream;
25use crate::pipeline::network::RequestControlMessage;
26use crate::pipeline::network::RequestPlanePayloadCodec;
27use crate::pipeline::network::RequestType;
28use crate::pipeline::network::ResponseType;
29use crate::pipeline::network::StreamOptions;
30use crate::pipeline::network::StreamProvider;
31use crate::pipeline::network::StreamReceiver;
32use crate::pipeline::network::StreamSender;
33use crate::pipeline::network::TwoPartCodec;
34use crate::pipeline::network::codec::TwoPartMessage;
35use crate::pipeline::network::tcp;
36use crate::pipeline::{ManyIn, ManyOut, PipelineError, ResponseStream, SingleIn};
37use crate::protocols::maybe_error::MaybeError;
38use crate::traits::DistributedRuntimeProvider;
39
40use anyhow::{Error, Result};
41use futures::stream::Stream;
42use std::pin::Pin;
43use std::task::{Context, Poll};
44use tokio_stream::{StreamExt, StreamNotifyClose, wrappers::ReceiverStream};
45use tracing::Instrument;
46
47fn decode_response_stream<U>(
53 response_rx: tokio::sync::mpsc::Receiver<bytes::Bytes>,
54 engine_ctx: Arc<dyn crate::engine::AsyncEngineContext>,
55 queue_start: Instant,
56 tx_start: Instant,
57 inflight_guard: InflightGuard,
58 payload_codec: RequestPlanePayloadCodec,
59) -> ManyOut<U>
60where
61 U: Data + for<'de> Deserialize<'de> + MaybeError,
62{
63 let engine_ctx_for_stream = engine_ctx.clone();
64 let mut is_complete_final = false;
65 let mut first_response = true;
66 let stream = StreamNotifyClose::new(ReceiverStream::new(response_rx)).filter_map(move |res| {
67 if let Some(res_bytes) = res {
68 if first_response {
69 first_response = false;
70 REQUEST_PLANE_ROUNDTRIP_TTFT_SECONDS.observe(tx_start.elapsed().as_secs_f64());
71 STAGE_DURATION_SECONDS
72 .with_label_values(&["transport_roundtrip"])
73 .observe(queue_start.elapsed().as_secs_f64());
74 }
75 if is_complete_final {
76 let err = DynamoError::msg(
77 "Response received after generation ended - this should never happen",
78 );
79 return Some(U::from_err(err));
80 }
81 match payload_codec.decode::<NetworkStreamWrapper<U>>(&res_bytes) {
82 Ok(item) => {
83 is_complete_final = item.complete_final;
84 if let Some(data) = item.data {
85 Some(data)
86 } else if is_complete_final {
87 None
88 } else {
89 let err =
90 DynamoError::msg("Empty response received - this should never happen");
91 Some(U::from_err(err))
92 }
93 }
94 Err(err) => {
95 let response_bytes_len = res_bytes.len();
96 tracing::warn!(
97 %err,
98 codec = payload_codec.name(),
99 response_bytes_len,
100 "failed deserializing request-plane response"
101 );
102 Some(U::from_err(DynamoError::msg(err.to_string())))
103 }
104 }
105 } else if is_complete_final {
106 None
107 } else if engine_ctx_for_stream.is_stopped() {
108 tracing::debug!("Request cancelled and then trying to read a response");
109 None
110 } else {
111 let err = DynamoError::builder()
112 .error_type(ErrorType::Disconnected)
113 .message("Stream ended before generation completed")
114 .build();
115 tracing::debug!("{err}");
116 Some(U::from_err(err))
117 }
118 });
119
120 inflight_guard.disarm();
121 let stream = InflightDecStream { inner: stream };
122 ResponseStream::new(Box::pin(stream), engine_ctx)
123}
124
125const CONTROL_MESSAGE_MAX_BYTES: usize = 128 * 1024;
126
127fn serialize_control_message(control_message: &RequestControlMessage) -> Result<Vec<u8>, Error> {
128 let ctrl = serde_json::to_vec(control_message)?;
129 if ctrl.len() > CONTROL_MESSAGE_MAX_BYTES {
130 return Err(PipelineError::Generic(format!(
131 "request control message too large: {} bytes exceeds limit {}",
132 ctrl.len(),
133 CONTROL_MESSAGE_MAX_BYTES
134 ))
135 .into());
136 }
137 Ok(ctrl)
138}
139
140fn build_request_envelope<T>(
147 context: &context::Context<()>,
148 recv_conn_info: ConnectionInfo,
149 send_conn_info: Option<ConnectionInfo>,
150 request: Option<&T>,
151) -> Result<bytes::Bytes, Error>
152where
153 T: serde::Serialize,
154{
155 let payload_codec = RequestPlanePayloadCodec::configured();
156 let request_id = context.id();
157 let request_type = if send_conn_info.is_some() {
158 RequestType::ManyIn
159 } else {
160 RequestType::SingleIn
161 };
162 let control_message = RequestControlMessage {
163 id: request_id.to_string(),
164 request_type,
165 response_type: ResponseType::ManyOut,
166 payload_codec,
167 connection_info: recv_conn_info,
168 metadata: context.metadata().clone(),
169 frontend_send_ts_ns: None,
170 request_stream_connection_info: send_conn_info,
171 };
172
173 let ctrl = serialize_control_message(&control_message)?;
174 let data: Option<Vec<u8>> = match request {
175 Some(req) => Some(payload_codec.encode(req)?),
176 None => None,
177 };
178
179 let msg = match data {
180 Some(d) => {
181 tracing::trace!(
182 request_id,
183 "packaging two-part message; ctrl: {} bytes, data: {} bytes",
184 ctrl.len(),
185 d.len(),
186 );
187 TwoPartMessage::from_parts(ctrl.into(), d.into())
188 }
189 None => {
190 tracing::trace!(
191 request_id,
192 "packaging bidirectional header-only envelope; ctrl: {} bytes",
193 ctrl.len(),
194 );
195 TwoPartMessage::from_header(ctrl.into())
196 }
197 };
198
199 let codec = TwoPartCodec::default();
200 let buffer = codec.encode_message(msg)?;
201 Ok(buffer)
202}
203
204async fn spawn_request_stream_forwarder<T>(
209 request_stream_provider: Option<StreamProvider<StreamSender>>,
210 mut input_stream: crate::engine::DataStream<T>,
211 engine_ctx: Arc<dyn crate::engine::AsyncEngineContext>,
212 payload_codec: RequestPlanePayloadCodec,
213) -> Result<(), Error>
214where
215 T: serde::Serialize + Send + 'static,
216{
217 let Some(provider) = request_stream_provider else {
218 return Ok(());
219 };
220
221 let request_sender = match provider.await {
222 Ok(Ok(sender)) => sender,
223 Ok(Err(e)) => {
224 return Err(anyhow::anyhow!(
225 DynamoError::builder()
226 .error_type(ErrorType::CannotConnect)
227 .message(format!("Worker dial-in failed for request stream: {e}"))
228 .build()
229 ));
230 }
231 Err(_) => {
232 return Err(anyhow::anyhow!(
233 DynamoError::builder()
234 .error_type(ErrorType::Disconnected)
235 .message("Worker disconnected before request stream was established")
236 .build()
237 ));
238 }
239 };
240
241 tokio::spawn(async move {
247 loop {
248 let item = tokio::select! {
249 biased;
250 _ = engine_ctx.killed() => break,
251 _ = engine_ctx.stopped() => break,
252 item = input_stream.next() => match item {
253 Some(item) => item,
254 None => break,
255 },
256 };
257 let bytes = match payload_codec.encode(&item) {
258 Ok(b) => b,
259 Err(e) => {
260 tracing::error!(
265 error = %e,
266 codec = payload_codec.name(),
267 "failed to serialize bidirectional request frame; killing context"
268 );
269 engine_ctx.kill();
270 break;
271 }
272 };
273 if request_sender.send(bytes.into()).await.is_err() {
274 tracing::debug!("worker request-stream receiver dropped; forwarder exiting");
275 break;
276 }
277 }
278 });
279
280 Ok(())
281}
282
283struct InflightGuard {
287 armed: bool,
288}
289
290impl InflightGuard {
291 fn new() -> Self {
292 Self { armed: true }
293 }
294
295 fn disarm(mut self) {
298 self.armed = false;
299 }
300}
301
302impl Drop for InflightGuard {
303 fn drop(&mut self) {
304 if self.armed {
305 REQUEST_PLANE_INFLIGHT.dec();
306 }
307 }
308}
309
310struct InflightDecStream<S> {
312 inner: S,
313}
314
315impl<S, T> Stream for InflightDecStream<S>
316where
317 S: Stream<Item = T> + Unpin,
318{
319 type Item = T;
320
321 fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
322 Pin::new(&mut self.inner).poll_next(cx)
323 }
324}
325
326impl<S> Drop for InflightDecStream<S> {
327 fn drop(&mut self) {
328 REQUEST_PLANE_INFLIGHT.dec();
329 }
330}
331
332fn subject_of(conn_info: &ConnectionInfo) -> Option<String> {
336 serde_json::from_str::<tcp::TcpStreamConnectionInfo>(&conn_info.info)
337 .ok()
338 .map(|ci| ci.subject)
339}
340
341pub struct AddressedRequest<T> {
342 request: T,
343 address: String,
344 instance: Option<Instance>,
347}
348
349impl<T> AddressedRequest<T> {
350 pub fn new(request: T, address: String) -> Self {
351 Self {
352 request,
353 address,
354 instance: None,
355 }
356 }
357
358 pub fn with_instance(request: T, address: String, instance: Instance) -> Self {
359 Self {
360 request,
361 address,
362 instance: Some(instance),
363 }
364 }
365
366 pub fn for_instance(request: T, instance: Instance) -> Self {
367 let address = instance.transport.address().to_string();
368 Self::with_instance(request, address, instance)
369 }
370
371 pub(crate) fn into_parts(self) -> (T, String, Option<Instance>) {
372 (self.request, self.address, self.instance)
373 }
374}
375
376pub struct AddressedPushRouter {
377 req_client: Arc<dyn RequestPlaneClient>,
379
380 resp_transport: Arc<tcp::server::TcpStreamServer>,
382}
383
384impl AddressedPushRouter {
385 pub fn new(
390 req_client: Arc<dyn RequestPlaneClient>,
391 resp_transport: Arc<tcp::server::TcpStreamServer>,
392 ) -> Result<Arc<Self>> {
393 Ok(Arc::new(Self {
394 req_client,
395 resp_transport,
396 }))
397 }
398
399 pub async fn from_runtime_provider(
400 provider: &impl DistributedRuntimeProvider,
401 ) -> Result<Arc<Self>> {
402 let manager = provider.drt().network_manager();
403 let req_client = manager.create_client()?;
404 let resp_transport = provider.drt().tcp_server().await?;
405
406 tracing::debug!(
407 transport = req_client.transport_name(),
408 "Creating AddressedPushRouter with request plane client"
409 );
410
411 Self::new(req_client, resp_transport)
412 }
413
414 pub async fn cancel_instance_streams(&self, instance_id: &EndpointInstanceId) -> usize {
416 self.resp_transport
417 .cancel_instance_streams(instance_id)
418 .await
419 }
420
421 pub async fn clear_instance_tombstone(&self, instance_id: &EndpointInstanceId) {
423 self.resp_transport
424 .clear_instance_tombstone(instance_id)
425 .await
426 }
427
428 pub async fn generate_bidirectional<T, U>(
434 &self,
435 instance: Instance,
436 address: String,
437 input: ManyIn<T>,
438 ) -> Result<ManyOut<U>, Error>
439 where
440 T: Data + Serialize,
441 U: Data + for<'de> Deserialize<'de> + MaybeError,
442 {
443 let (request_stream, context) = input.into_parts();
444 let input_stream = request_stream.take().ok_or_else(|| {
445 anyhow::anyhow!("RequestStream::take called twice on bidirectional dispatch input")
446 })?;
447
448 self.dispatch_and_finalize::<T, U>(
449 &context,
450 address,
451 Some(&instance),
452 None,
453 Some(input_stream),
454 )
455 .await
456 }
457
458 async fn dispatch_and_finalize<T, U>(
466 &self,
467 context: &context::Context<()>,
468 address: String,
469 instance: Option<&Instance>,
470 request: Option<&T>,
471 input_stream: Option<crate::engine::DataStream<T>>,
472 ) -> Result<ManyOut<U>, Error>
473 where
474 T: Data + Serialize,
475 U: Data + for<'de> Deserialize<'de> + MaybeError,
476 {
477 let engine_ctx = context.context();
478
479 let queue_start = Instant::now();
480 REQUEST_PLANE_INFLIGHT.inc();
481 let inflight_guard = InflightGuard::new();
482
483 let enable_request_stream = input_stream.is_some();
484 let payload_codec = RequestPlanePayloadCodec::configured();
485
486 let (send_registered, recv_registered) = self
492 .register_streams(engine_ctx.clone(), enable_request_stream, true)
493 .await?;
494 let recv_registered = recv_registered.ok_or_else(|| {
495 anyhow::anyhow!("response stream registration missing despite enable_response_stream")
496 })?;
497
498 let recv_subject = subject_of(&recv_registered.connection_info);
502 let send_subject = send_registered
503 .as_ref()
504 .and_then(|r| subject_of(&r.connection_info));
505 if let (Some(subject), Some(inst)) = (&recv_subject, instance)
506 && !self
507 .resp_transport
508 .associate_instance(
509 subject,
510 send_subject.as_deref(),
511 &inst.endpoint_instance_id(),
512 )
513 .await
514 {
515 return Err(anyhow::anyhow!(
516 DynamoError::builder()
517 .error_type(ErrorType::Disconnected)
518 .message("Worker removed before request could be sent (tombstoned instance)")
519 .build()
520 ));
521 }
522
523 let buffer = build_request_envelope(
524 context,
525 recv_registered.connection_info.clone(),
526 send_registered.as_ref().map(|r| r.connection_info.clone()),
527 request,
528 )?;
529 REQUEST_PLANE_QUEUE_SECONDS.observe(queue_start.elapsed().as_secs_f64());
530
531 let tx_start = Instant::now();
532 let request_plane_response = self.dispatch_buffer(address, buffer, context.id()).await?;
533 REQUEST_PLANE_SEND_SECONDS.observe(tx_start.elapsed().as_secs_f64());
534
535 if let Some(err) = detect_worker_rejection_response(&request_plane_response) {
540 tracing::warn!(
541 request_id = context.id(),
542 worker_response = %err.to_string(),
543 "Request rejected by worker"
544 );
545 return Err(err.into());
546 }
547
548 if let Some(stream) = input_stream {
554 let request_stream_provider = send_registered.map(|r| {
555 let (_conn_info, provider) = r.into_parts();
556 provider
557 });
558 spawn_request_stream_forwarder(
559 request_stream_provider,
560 stream,
561 engine_ctx.clone(),
562 payload_codec,
563 )
564 .await?;
565 }
566
567 let _nvtx_wait = dynamo_nvtx_range!("transport.tcp.wait_backend");
568 tracing::trace!(request_id = context.id(), "awaiting transport handshake");
569
570 let (_recv_conn_info, response_stream_provider) = recv_registered.into_parts();
572
573 let response_stream = match response_stream_provider.await {
576 Ok(Ok(stream)) => stream,
577 Ok(Err(e)) => {
578 return Err(anyhow::anyhow!(
585 DynamoError::builder()
586 .error_type(ErrorType::CannotConnect)
587 .message(format!(
588 "Worker generate() failed before response stream: {e}"
589 ))
590 .build()
591 ));
592 }
593 Err(_recv_err) => {
594 return Err(anyhow::anyhow!(
597 DynamoError::builder()
598 .error_type(ErrorType::Disconnected)
599 .message("Worker disconnected before response stream was established")
600 .build()
601 ));
602 }
603 };
604 drop(_nvtx_wait);
605
606 Ok(decode_response_stream(
607 response_stream.rx,
608 engine_ctx,
609 queue_start,
610 tx_start,
611 inflight_guard,
612 payload_codec,
613 ))
614 }
615
616 async fn register_streams(
623 &self,
624 engine_ctx: Arc<dyn crate::engine::AsyncEngineContext>,
625 enable_request_stream: bool,
626 enable_response_stream: bool,
627 ) -> Result<
628 (
629 Option<RegisteredStream<StreamSender>>,
630 Option<RegisteredStream<StreamReceiver>>,
631 ),
632 Error,
633 > {
634 let options = StreamOptions::builder()
635 .context(engine_ctx)
636 .enable_request_stream(enable_request_stream)
637 .enable_response_stream(enable_response_stream)
638 .build()?;
639
640 let pending: PendingConnections = self.resp_transport.register(options).await;
641 let (send_stream, recv_stream) = pending.into_parts();
642
643 debug_assert_eq!(
647 send_stream.is_some(),
648 enable_request_stream,
649 "data-plane registration: request-stream presence does not match request"
650 );
651 debug_assert_eq!(
652 recv_stream.is_some(),
653 enable_response_stream,
654 "data-plane registration: response-stream presence does not match request"
655 );
656
657 Ok((send_stream, recv_stream))
658 }
659
660 async fn dispatch_buffer(
668 &self,
669 address: String,
670 buffer: bytes::Bytes,
671 request_id: &str,
672 ) -> Result<bytes::Bytes, Error> {
673 let mut headers = std::collections::HashMap::new();
674 inject_trace_headers_into_map(&mut headers);
675 headers.insert("request-id".to_string(), request_id.to_string());
676 let send_ts_ns = std::time::SystemTime::now()
677 .duration_since(std::time::UNIX_EPOCH)
678 .unwrap_or_default()
679 .as_nanos() as u64;
680 headers.insert("x-frontend-send-ts-ns".to_string(), send_ts_ns.to_string());
681
682 let _nvtx_send = dynamo_nvtx_range!("transport.tcp.send");
683 let ack = self
684 .req_client
685 .send_request(address, buffer, headers)
686 .await?;
687 drop(_nvtx_send);
688 Ok(ack)
689 }
690}
691
692fn detect_worker_rejection_response(res_bytes: &[u8]) -> Option<DynamoError> {
695 const OVERLOAD_PREFIX: &[u8] = b"Server overloaded:";
696 const UNAVAILABLE_PREFIX: &[u8] = b"Server unavailable:";
697
698 let error_type = if res_bytes.starts_with(OVERLOAD_PREFIX) {
699 ErrorType::ResourceExhausted
700 } else if res_bytes.starts_with(UNAVAILABLE_PREFIX) {
701 ErrorType::Unavailable
702 } else {
703 return None;
704 };
705
706 let msg = String::from_utf8_lossy(res_bytes).into_owned();
707 Some(
708 DynamoError::builder()
709 .error_type(error_type)
710 .message(msg)
711 .build(),
712 )
713}
714
715#[cfg(test)]
716mod rejection_detection_tests {
717 use super::*;
718
719 #[test]
720 fn overload_payload_maps_to_resource_exhausted() {
721 let err = detect_worker_rejection_response(b"Server overloaded: worker at capacity")
722 .expect("should detect overload");
723 assert_eq!(err.error_type(), ErrorType::ResourceExhausted);
724 }
725
726 #[test]
727 fn empty_ack_is_not_overload() {
728 assert!(detect_worker_rejection_response(b"").is_none());
730 assert!(detect_worker_rejection_response(br#"{"data":"chunk"}"#).is_none());
731 }
732
733 #[test]
734 fn detected_overload_satisfies_http_529_gate() {
735 let err =
737 detect_worker_rejection_response(b"Server overloaded: test").expect("should detect");
738 let any_err: anyhow::Error = err.into();
739 assert!(crate::error::match_error_chain(
740 any_err.as_ref(),
741 &[ErrorType::ResourceExhausted],
742 &[]
743 ));
744 }
745}
746
747#[async_trait::async_trait]
748impl<T, U> AsyncEngine<SingleIn<AddressedRequest<T>>, ManyOut<U>, Error> for AddressedPushRouter
749where
750 T: Data + Serialize,
751 U: Data + for<'de> Deserialize<'de> + MaybeError,
752{
753 async fn generate(&self, request: SingleIn<AddressedRequest<T>>) -> Result<ManyOut<U>, Error> {
754 let (addressed_request, context) = request.transfer(());
755 let (request, address, instance_info) = addressed_request.into_parts();
756
757 self.dispatch_and_finalize::<T, U>(
758 &context,
759 address,
760 instance_info.as_ref(),
761 Some(&request),
762 None,
763 )
764 .await
765 }
766}
767
768#[cfg(test)]
769mod tests {
770 use super::{
771 CONTROL_MESSAGE_MAX_BYTES, ConnectionInfo, RequestControlMessage, RequestPlanePayloadCodec,
772 RequestType, ResponseType, serialize_control_message,
773 };
774 use std::collections::BTreeMap;
775
776 fn base_control_message(metadata: BTreeMap<String, String>) -> RequestControlMessage {
777 RequestControlMessage {
778 id: "request-123".to_string(),
779 request_type: RequestType::SingleIn,
780 response_type: ResponseType::ManyOut,
781 payload_codec: RequestPlanePayloadCodec::Json,
782 connection_info: ConnectionInfo {
783 transport: "tcp".to_string(),
784 info: "{}".to_string(),
785 },
786 metadata,
787 frontend_send_ts_ns: None,
788 request_stream_connection_info: None,
789 }
790 }
791
792 #[test]
793 fn serialize_control_message_succeeds_under_limit() {
794 let mut metadata = BTreeMap::new();
795 metadata.insert("x-tiny-blob".to_string(), "alpha".to_string());
796
797 let ctrl = serialize_control_message(&base_control_message(metadata))
798 .expect("control message should serialize under the limit");
799 assert!(ctrl.len() <= CONTROL_MESSAGE_MAX_BYTES);
800 }
801
802 #[test]
803 fn serialize_control_message_errors_over_limit() {
804 let mut metadata = BTreeMap::new();
805 metadata.insert(
806 "x-large-blob".to_string(),
807 "x".repeat(CONTROL_MESSAGE_MAX_BYTES),
808 );
809
810 let err = serialize_control_message(&base_control_message(metadata))
811 .expect_err("oversized control message should fail")
812 .to_string();
813 assert!(err.contains("request control message too large"));
814 assert!(err.contains(&CONTROL_MESSAGE_MAX_BYTES.to_string()));
815 }
816}