Skip to main content

dynamo_runtime/pipeline/network/egress/
addressed_router.rs

1// SPDX-FileCopyrightText: Copyright (c) 2024-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2// SPDX-License-Identifier: Apache-2.0
3
4use 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
47/// Stream transformation helper that:
48/// - decodes a response byte stream from network into the fully-shaped `ManyOut<U>`
49/// - emits TTFT and transport-roundtrip metrics on first response
50/// - hands off the `InflightGuard` to a stream-lifetime `InflightDecStream` so
51///   the inflight gauge stays accurate for the whole response lifetime.
52fn 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
140/// Build the request control message, and serialize for transfer.
141///
142/// `request` provides the optional unary request payload. Should set for
143/// SingleIn generation.
144/// `send_conn_info` provides the connection info for the request stream.
145/// Should set for ManyIn generation.
146fn build_request_envelope<T>(
147    context: &context::Context<()>,
148    recv_conn_info: ConnectionInfo,
149    send_conn_info: Option<ConnectionInfo>,
150    request: Option<&T>,
151    payload_codec: RequestPlanePayloadCodec,
152) -> Result<bytes::Bytes, Error>
153where
154    T: serde::Serialize,
155{
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
204fn payload_codec_for_worker(instance: Option<&Instance>) -> RequestPlanePayloadCodec {
205    instance
206        .and_then(|instance| instance.request_plane_codec)
207        .unwrap_or(RequestPlanePayloadCodec::Json)
208}
209
210/// Await the network request-stream dial-in (if `request_stream_provider` is `Some`)
211/// and spawn a detached task that forwards every item from `input_stream` onto
212/// the request stream. Returns once the forwarder is spawned; `Err` if request-stream
213/// dial-in fails.
214async fn spawn_request_stream_forwarder<T>(
215    request_stream_provider: Option<StreamProvider<StreamSender>>,
216    mut input_stream: crate::engine::DataStream<T>,
217    engine_ctx: Arc<dyn crate::engine::AsyncEngineContext>,
218    payload_codec: RequestPlanePayloadCodec,
219) -> Result<(), Error>
220where
221    T: serde::Serialize + Send + 'static,
222{
223    let Some(provider) = request_stream_provider else {
224        return Ok(());
225    };
226
227    let request_sender = match provider.await {
228        Ok(Ok(sender)) => sender,
229        Ok(Err(e)) => {
230            return Err(anyhow::anyhow!(
231                DynamoError::builder()
232                    .error_type(ErrorType::CannotConnect)
233                    .message(format!("Worker dial-in failed for request stream: {e}"))
234                    .build()
235            ));
236        }
237        Err(_) => {
238            return Err(anyhow::anyhow!(
239                DynamoError::builder()
240                    .error_type(ErrorType::Disconnected)
241                    .message("Worker disconnected before request stream was established")
242                    .build()
243            ));
244        }
245    };
246
247    // The task exits on stream end, context kill/stop, send error (worker
248    // dropped its receiver), or local serialize failure. On any exit
249    // `request_sender` drops and triggers transport shutdown (see server.rs for details)
250    // which closes the upstream mpsc, triggering the server-side handler to emit
251    // `Sentinel`, which signals the worker's reader to end cleanly.
252    tokio::spawn(async move {
253        loop {
254            let item = tokio::select! {
255                biased;
256                _ = engine_ctx.killed() => break,
257                _ = engine_ctx.stopped() => break,
258                item = input_stream.next() => match item {
259                    Some(item) => item,
260                    None => break,
261                },
262            };
263            let bytes = match payload_codec.encode(&item) {
264                Ok(b) => b,
265                Err(e) => {
266                    // Stream-side framing failure: the engine sees a
267                    // partial input, so kill the context to abort both
268                    // directions consistently rather than silently
269                    // dropping frames.
270                    tracing::error!(
271                        error = %e,
272                        codec = payload_codec.name(),
273                        "failed to serialize bidirectional request frame; killing context"
274                    );
275                    engine_ctx.kill();
276                    break;
277                }
278            };
279            if request_sender.send(bytes.into()).await.is_err() {
280                tracing::debug!("worker request-stream receiver dropped; forwarder exiting");
281                break;
282            }
283        }
284    });
285
286    Ok(())
287}
288
289/// RAII guard that decrements REQUEST_PLANE_INFLIGHT on drop unless disarmed.
290/// Protects against gauge leaks when `?` operators cause early returns between
291/// the increment and `InflightDecStream` construction.
292struct InflightGuard {
293    armed: bool,
294}
295
296impl InflightGuard {
297    fn new() -> Self {
298        Self { armed: true }
299    }
300
301    /// Consume the guard without decrementing. Call this when `InflightDecStream`
302    /// takes over responsibility for the decrement.
303    fn disarm(mut self) {
304        self.armed = false;
305    }
306}
307
308impl Drop for InflightGuard {
309    fn drop(&mut self) {
310        if self.armed {
311            REQUEST_PLANE_INFLIGHT.dec();
312        }
313    }
314}
315
316/// Wrapper that decrements request-plane inflight gauge when the stream is dropped.
317struct InflightDecStream<S> {
318    inner: S,
319}
320
321impl<S, T> Stream for InflightDecStream<S>
322where
323    S: Stream<Item = T> + Unpin,
324{
325    type Item = T;
326
327    fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
328        Pin::new(&mut self.inner).poll_next(cx)
329    }
330}
331
332impl<S> Drop for InflightDecStream<S> {
333    fn drop(&mut self) {
334        REQUEST_PLANE_INFLIGHT.dec();
335    }
336}
337
338/// Extract the TCP stream subject from a [`ConnectionInfo`], if it carries a
339/// well-formed [`tcp::TcpStreamConnectionInfo`]. Used for the pre-dispatch
340/// tombstone check.
341fn subject_of(conn_info: &ConnectionInfo) -> Option<String> {
342    serde_json::from_str::<tcp::TcpStreamConnectionInfo>(&conn_info.info)
343        .ok()
344        .map(|ci| ci.subject)
345}
346
347pub struct AddressedRequest<T> {
348    request: T,
349    address: String,
350    /// Carries endpoint name + instance_id so cancellation is scoped to the
351    /// exact (endpoint, instance) pair, not all endpoints on the same runtime.
352    instance: Option<Instance>,
353}
354
355impl<T> AddressedRequest<T> {
356    pub fn new(request: T, address: String) -> Self {
357        Self {
358            request,
359            address,
360            instance: None,
361        }
362    }
363
364    pub fn with_instance(request: T, address: String, instance: Instance) -> Self {
365        Self {
366            request,
367            address,
368            instance: Some(instance),
369        }
370    }
371
372    pub fn for_instance(request: T, instance: Instance) -> Self {
373        let address = instance.transport.address().to_string();
374        Self::with_instance(request, address, instance)
375    }
376
377    pub(crate) fn into_parts(self) -> (T, String, Option<Instance>) {
378        (self.request, self.address, self.instance)
379    }
380}
381
382pub struct AddressedPushRouter {
383    // Request transport (unified trait object - works with all transports)
384    req_client: Arc<dyn RequestPlaneClient>,
385
386    // Response transport (TCP streaming - unchanged)
387    resp_transport: Arc<tcp::server::TcpStreamServer>,
388}
389
390impl AddressedPushRouter {
391    /// Create a new router with a request plane client
392    ///
393    /// This is the unified constructor that works with any transport type.
394    /// The client is provided as a trait object, hiding the specific implementation.
395    pub fn new(
396        req_client: Arc<dyn RequestPlaneClient>,
397        resp_transport: Arc<tcp::server::TcpStreamServer>,
398    ) -> Result<Arc<Self>> {
399        Ok(Arc::new(Self {
400            req_client,
401            resp_transport,
402        }))
403    }
404
405    pub async fn from_runtime_provider(
406        provider: &impl DistributedRuntimeProvider,
407    ) -> Result<Arc<Self>> {
408        let manager = provider.drt().network_manager();
409        let req_client = manager.create_client()?;
410        let resp_transport = provider.drt().tcp_server().await?;
411
412        tracing::debug!(
413            transport = req_client.transport_name(),
414            "Creating AddressedPushRouter with request plane client"
415        );
416
417        Self::new(req_client, resp_transport)
418    }
419
420    /// Cancel all pending response-stream registrations for an instance.
421    pub async fn cancel_instance_streams(&self, instance_id: &EndpointInstanceId) -> usize {
422        self.resp_transport
423            .cancel_instance_streams(instance_id)
424            .await
425    }
426
427    /// Clear the tombstone after an instance reappears in discovery.
428    pub async fn clear_instance_tombstone(&self, instance_id: &EndpointInstanceId) {
429        self.resp_transport
430            .clear_instance_tombstone(instance_id)
431            .await
432    }
433
434    /// Bidirectional generation. Note that it doesn't implement the AsyncEngine trait directly
435    /// because there is no trivial way to wrap (instance and address) into ManyIn style.
436    /// May wrap as SingleIn<AddressedStreamRequest<T>> and unwrap here but really just syntax
437    /// sugar, so we just do it inline here. Will consider only if we do want to call this from
438    /// typed erased AsyncEngine impls.
439    pub async fn generate_bidirectional<T, U>(
440        &self,
441        instance: Instance,
442        address: String,
443        input: ManyIn<T>,
444    ) -> Result<ManyOut<U>, Error>
445    where
446        T: Data + Serialize,
447        U: Data + for<'de> Deserialize<'de> + MaybeError,
448    {
449        let (request_stream, context) = input.into_parts();
450        let input_stream = request_stream.take().ok_or_else(|| {
451            anyhow::anyhow!("RequestStream::take called twice on bidirectional dispatch input")
452        })?;
453
454        self.dispatch_and_finalize::<T, U>(
455            &context,
456            address,
457            Some(&instance),
458            None,
459            Some(input_stream),
460        )
461        .await
462    }
463
464    /// Shared dispatch core for both unary and bidirectional requests. Wire
465    /// shape is inferred from the inputs:
466    ///   - `input_stream = Some(_)` + `request = None` → bidirectional,
467    ///     header-only envelope. The worker dials back for both halves and
468    ///     pulls request frames off the spawned forwarder.
469    ///   - `input_stream = None` + `request = Some(_)` → unary, two-part
470    ///     `[ctrl, data]` envelope. The payload travels in the data part.
471    async fn dispatch_and_finalize<T, U>(
472        &self,
473        context: &context::Context<()>,
474        address: String,
475        instance: Option<&Instance>,
476        request: Option<&T>,
477        input_stream: Option<crate::engine::DataStream<T>>,
478    ) -> Result<ManyOut<U>, Error>
479    where
480        T: Data + Serialize,
481        U: Data + for<'de> Deserialize<'de> + MaybeError,
482    {
483        let engine_ctx = context.context();
484
485        let queue_start = Instant::now();
486        REQUEST_PLANE_INFLIGHT.inc();
487        let inflight_guard = InflightGuard::new();
488
489        let enable_request_stream = input_stream.is_some();
490        let payload_codec = payload_codec_for_worker(instance);
491
492        // Hold the `RegisteredStream` as their RAII cleanup stays armed while held,
493        // which simplifies the cancellation of registration on error. Each side is
494        // disarmed by `into_parts()` on awaiting stream provider: past that point the
495        // subject is reaped by the worker's dial-in (instance healthy) or the discovery
496        // watcher (instance dropped), so no cleanup is owed.
497        let (send_registered, recv_registered) = self
498            .register_streams(engine_ctx.clone(), enable_request_stream, true)
499            .await?;
500        let recv_registered = recv_registered.ok_or_else(|| {
501            anyhow::anyhow!("response stream registration missing despite enable_response_stream")
502        })?;
503
504        // Tombstone check: if discovery already removed the worker, fail fast
505        // with a migratable error rather than writing to the request plane.
506        // Dropping the held registrations on this return runs their cleanup.
507        let recv_subject = subject_of(&recv_registered.connection_info);
508        let send_subject = send_registered
509            .as_ref()
510            .and_then(|r| subject_of(&r.connection_info));
511        if let (Some(subject), Some(inst)) = (&recv_subject, instance)
512            && !self
513                .resp_transport
514                .associate_instance(
515                    subject,
516                    send_subject.as_deref(),
517                    &inst.endpoint_instance_id(),
518                )
519                .await
520        {
521            return Err(anyhow::anyhow!(
522                DynamoError::builder()
523                    .error_type(ErrorType::Disconnected)
524                    .message("Worker removed before request could be sent (tombstoned instance)")
525                    .build()
526            ));
527        }
528
529        let buffer = build_request_envelope(
530            context,
531            recv_registered.connection_info.clone(),
532            send_registered.as_ref().map(|r| r.connection_info.clone()),
533            request,
534            payload_codec,
535        )?;
536        REQUEST_PLANE_QUEUE_SECONDS.observe(queue_start.elapsed().as_secs_f64());
537
538        let tx_start = Instant::now();
539        let request_plane_response = self.dispatch_buffer(address, buffer, context.id()).await?;
540        REQUEST_PLANE_SEND_SECONDS.observe(tx_start.elapsed().as_secs_f64());
541
542        // A worker rejection surfaces on the request-plane ACK, not the response
543        // stream. Short-circuit before waiting on a response-plane connection the
544        // worker will never open; returning early drops `recv_registered` and
545        // `inflight_guard` (their Drop cleans up).
546        if let Some(err) = detect_worker_rejection_response(&request_plane_response) {
547            tracing::warn!(
548                request_id = context.id(),
549                worker_response = %err.to_string(),
550                "Request rejected by worker"
551            );
552            return Err(err.into());
553        }
554
555        // Spawn the forwarder before awaiting the response prologue so request
556        // frames pre-load into the worker's input buffer while the engine
557        // initialises in parallel. The response provider only resolves after
558        // `engine.generate()` returns; awaiting it second avoids stalling the
559        // request-side handshake on engine setup latency.
560        if let Some(stream) = input_stream {
561            let request_stream_provider = send_registered.map(|r| {
562                let (_conn_info, provider) = r.into_parts();
563                provider
564            });
565            spawn_request_stream_forwarder(
566                request_stream_provider,
567                stream,
568                engine_ctx.clone(),
569                payload_codec,
570            )
571            .await?;
572        }
573
574        let _nvtx_wait = dynamo_nvtx_range!("transport.tcp.wait_backend");
575        tracing::trace!(request_id = context.id(), "awaiting transport handshake");
576
577        // Disarms the recv-side cleanup; see the holding rationale above.
578        let (_recv_conn_info, response_stream_provider) = recv_registered.into_parts();
579
580        // RecvError → migratable Disconnected (watcher cancelled the subject
581        // or the worker died before establishing the response stream).
582        let response_stream = match response_stream_provider.await {
583            Ok(Ok(stream)) => stream,
584            Ok(Err(e)) => {
585                // generate() failed before any response bytes; migrate via
586                // CannotConnect since the dominant cause is a worker-local
587                // setup/version issue. The wire prologue carries only an
588                // opaque string today, so app-level rejections also retry
589                // -- safe because no side effects are visible yet. Follow-up:
590                // structured prologue error type for finer routing.
591                return Err(anyhow::anyhow!(
592                    DynamoError::builder()
593                        .error_type(ErrorType::CannotConnect)
594                        .message(format!(
595                            "Worker generate() failed before response stream: {e}"
596                        ))
597                        .build()
598                ));
599            }
600            Err(_recv_err) => {
601                // oneshot dropped: either the discovery watcher cancelled
602                // this subject or the worker died mid-handshake.
603                return Err(anyhow::anyhow!(
604                    DynamoError::builder()
605                        .error_type(ErrorType::Disconnected)
606                        .message("Worker disconnected before response stream was established")
607                        .build()
608                ));
609            }
610        };
611        drop(_nvtx_wait);
612
613        Ok(decode_response_stream(
614            response_stream.rx,
615            engine_ctx,
616            queue_start,
617            tx_start,
618            inflight_guard,
619            payload_codec,
620        ))
621    }
622
623    /// Register the requested halves of a data-plane stream with the response
624    /// transport. Returns `(send_stream, recv_stream)` mirroring the
625    /// `PendingConnections::into_parts` shape — either side is `None` when not
626    /// requested. Asserts post-registration that the transport produced
627    /// exactly the requested shape; a mismatch is a transport-layer bug, not
628    /// a runtime error path.
629    async fn register_streams(
630        &self,
631        engine_ctx: Arc<dyn crate::engine::AsyncEngineContext>,
632        enable_request_stream: bool,
633        enable_response_stream: bool,
634    ) -> Result<
635        (
636            Option<RegisteredStream<StreamSender>>,
637            Option<RegisteredStream<StreamReceiver>>,
638        ),
639        Error,
640    > {
641        let options = StreamOptions::builder()
642            .context(engine_ctx)
643            .enable_request_stream(enable_request_stream)
644            .enable_response_stream(enable_response_stream)
645            .build()?;
646
647        let pending: PendingConnections = self.resp_transport.register(options).await;
648        let (send_stream, recv_stream) = pending.into_parts();
649
650        // Transport-layer invariant: the data plane produces exactly the halves
651        // we requested. A mismatch is a bug in the transport, not a runtime
652        // error path, so assert only in debug builds rather than panicking prod.
653        debug_assert_eq!(
654            send_stream.is_some(),
655            enable_request_stream,
656            "data-plane registration: request-stream presence does not match request"
657        );
658        debug_assert_eq!(
659            recv_stream.is_some(),
660            enable_response_stream,
661            "data-plane registration: response-stream presence does not match request"
662        );
663
664        Ok((send_stream, recv_stream))
665    }
666
667    /// Build standard request-plane headers (trace propagation, request-id,
668    /// frontend send-timestamp) and write the encoded buffer through the
669    /// request-plane client.
670    ///
671    /// Returns the request-plane ACK bytes (empty `TcpResponseMessage` on the
672    /// success path; a rejection-marker payload when the worker rejects the
673    /// request — see [`detect_worker_rejection_response`]).
674    async fn dispatch_buffer(
675        &self,
676        address: String,
677        buffer: bytes::Bytes,
678        request_id: &str,
679    ) -> Result<bytes::Bytes, Error> {
680        let mut headers = std::collections::HashMap::new();
681        inject_trace_headers_into_map(&mut headers);
682        headers.insert("request-id".to_string(), request_id.to_string());
683        let send_ts_ns = std::time::SystemTime::now()
684            .duration_since(std::time::UNIX_EPOCH)
685            .unwrap_or_default()
686            .as_nanos() as u64;
687        headers.insert("x-frontend-send-ts-ns".to_string(), send_ts_ns.to_string());
688
689        let _nvtx_send = dynamo_nvtx_range!("transport.tcp.send");
690        let ack = self
691            .req_client
692            .send_request(address, buffer, headers)
693            .await?;
694        drop(_nvtx_send);
695        Ok(ack)
696    }
697}
698
699/// Map a worker rejection ACK to the corresponding typed error. `None` for
700/// normal responses, including the empty "queued" ACK.
701fn detect_worker_rejection_response(res_bytes: &[u8]) -> Option<DynamoError> {
702    const OVERLOAD_PREFIX: &[u8] = b"Server overloaded:";
703    const UNAVAILABLE_PREFIX: &[u8] = b"Server unavailable:";
704
705    let error_type = if res_bytes.starts_with(OVERLOAD_PREFIX) {
706        ErrorType::ResourceExhausted
707    } else if res_bytes.starts_with(UNAVAILABLE_PREFIX) {
708        ErrorType::Unavailable
709    } else {
710        return None;
711    };
712
713    let msg = String::from_utf8_lossy(res_bytes).into_owned();
714    Some(
715        DynamoError::builder()
716            .error_type(error_type)
717            .message(msg)
718            .build(),
719    )
720}
721
722#[cfg(test)]
723mod rejection_detection_tests {
724    use super::*;
725
726    #[test]
727    fn overload_payload_maps_to_resource_exhausted() {
728        let err = detect_worker_rejection_response(b"Server overloaded: worker at capacity")
729            .expect("should detect overload");
730        assert_eq!(err.error_type(), ErrorType::ResourceExhausted);
731    }
732
733    #[test]
734    fn empty_ack_is_not_overload() {
735        // The success-path ACK is empty; misreading it as overload breaks every request.
736        assert!(detect_worker_rejection_response(b"").is_none());
737        assert!(detect_worker_rejection_response(br#"{"data":"chunk"}"#).is_none());
738    }
739
740    #[test]
741    fn detected_overload_satisfies_http_529_gate() {
742        // request_was_rejected (http/service/metrics.rs) → 529 keys on ResourceExhausted.
743        let err =
744            detect_worker_rejection_response(b"Server overloaded: test").expect("should detect");
745        let any_err: anyhow::Error = err.into();
746        assert!(crate::error::match_error_chain(
747            any_err.as_ref(),
748            &[ErrorType::ResourceExhausted],
749            &[]
750        ));
751    }
752}
753
754#[async_trait::async_trait]
755impl<T, U> AsyncEngine<SingleIn<AddressedRequest<T>>, ManyOut<U>, Error> for AddressedPushRouter
756where
757    T: Data + Serialize,
758    U: Data + for<'de> Deserialize<'de> + MaybeError,
759{
760    async fn generate(&self, request: SingleIn<AddressedRequest<T>>) -> Result<ManyOut<U>, Error> {
761        let (addressed_request, context) = request.transfer(());
762        let (request, address, instance_info) = addressed_request.into_parts();
763
764        self.dispatch_and_finalize::<T, U>(
765            &context,
766            address,
767            instance_info.as_ref(),
768            Some(&request),
769            None,
770        )
771        .await
772    }
773}
774
775#[cfg(test)]
776mod tests {
777    use super::{
778        CONTROL_MESSAGE_MAX_BYTES, ConnectionInfo, RequestControlMessage, RequestPlanePayloadCodec,
779        RequestType, ResponseType, TwoPartCodec, build_request_envelope, payload_codec_for_worker,
780        serialize_control_message,
781    };
782    use crate::{
783        component::{Instance, TransportType},
784        pipeline::Context,
785    };
786    use serde::{Deserialize, Serialize};
787    use std::collections::BTreeMap;
788
789    fn base_control_message(metadata: BTreeMap<String, String>) -> RequestControlMessage {
790        RequestControlMessage {
791            id: "request-123".to_string(),
792            request_type: RequestType::SingleIn,
793            response_type: ResponseType::ManyOut,
794            payload_codec: RequestPlanePayloadCodec::Json,
795            connection_info: ConnectionInfo {
796                transport: "tcp".to_string(),
797                info: "{}".to_string(),
798            },
799            metadata,
800            frontend_send_ts_ns: None,
801            request_stream_connection_info: None,
802        }
803    }
804
805    #[derive(Debug, Deserialize, Serialize, PartialEq, Eq)]
806    struct TestRequest {
807        value: u64,
808    }
809
810    #[test]
811    fn legacy_worker_without_codec_metadata_receives_json() {
812        let worker = Instance {
813            component: "worker".to_string(),
814            endpoint: "generate".to_string(),
815            namespace: "default".to_string(),
816            instance_id: 42,
817            transport: TransportType::Nats("worker.generate".to_string()),
818            device_type: None,
819            request_plane_codec: None,
820        };
821        let payload_codec = payload_codec_for_worker(Some(&worker));
822        assert_eq!(payload_codec, RequestPlanePayloadCodec::Json);
823
824        let request = TestRequest { value: 123 };
825        let buffer = build_request_envelope(
826            &Context::new(()),
827            ConnectionInfo {
828                transport: "tcp".to_string(),
829                info: "{}".to_string(),
830            },
831            None,
832            Some(&request),
833            payload_codec,
834        )
835        .expect("legacy-worker request envelope should encode");
836        let message = TwoPartCodec::default()
837            .decode_message(buffer)
838            .expect("request envelope should decode");
839
840        let control: RequestControlMessage = serde_json::from_slice(&message.header).unwrap();
841        assert_eq!(control.payload_codec, RequestPlanePayloadCodec::Json);
842        assert_eq!(
843            serde_json::from_slice::<TestRequest>(&message.data).unwrap(),
844            request
845        );
846    }
847
848    #[test]
849    fn serialize_control_message_succeeds_under_limit() {
850        let mut metadata = BTreeMap::new();
851        metadata.insert("x-tiny-blob".to_string(), "alpha".to_string());
852
853        let ctrl = serialize_control_message(&base_control_message(metadata))
854            .expect("control message should serialize under the limit");
855        assert!(ctrl.len() <= CONTROL_MESSAGE_MAX_BYTES);
856    }
857
858    #[test]
859    fn serialize_control_message_errors_over_limit() {
860        let mut metadata = BTreeMap::new();
861        metadata.insert(
862            "x-large-blob".to_string(),
863            "x".repeat(CONTROL_MESSAGE_MAX_BYTES),
864        );
865
866        let err = serialize_control_message(&base_control_message(metadata))
867            .expect_err("oversized control message should fail")
868            .to_string();
869        assert!(err.contains("request control message too large"));
870        assert!(err.contains(&CONTROL_MESSAGE_MAX_BYTES.to_string()));
871    }
872}