From 55ec0bb5d789ff74b7a584fd8b64aa1ecaf81929 Mon Sep 17 00:00:00 2001 From: jthomson04 Date: Mon, 20 Jul 2026 15:03:54 -0700 Subject: [PATCH 1/3] feat(runtime): multiplex TCP response streams Signed-off-by: jthomson04 --- benchmarks/frontend/scripts/run_perf.sh | 67 +- lib/runtime/src/component/endpoint.rs | 1 + lib/runtime/src/config/environment_names.rs | 29 + lib/runtime/src/distributed.rs | 18 + lib/runtime/src/metrics.rs | 1 + lib/runtime/src/metrics/response_mux.rs | 447 ++++ lib/runtime/src/pipeline/network.rs | 248 +- .../network/egress/addressed_router.rs | 8 +- .../src/pipeline/network/egress/tcp_client.rs | 17 +- .../pipeline/network/ingress/push_handler.rs | 68 +- lib/runtime/src/pipeline/network/tcp.rs | 98 +- .../src/pipeline/network/tcp/client.rs | 35 +- lib/runtime/src/pipeline/network/tcp/mux.rs | 698 ++++++ .../src/pipeline/network/tcp/mux/client.rs | 2150 +++++++++++++++++ .../src/pipeline/network/tcp/server.rs | 704 +++++- 15 files changed, 4468 insertions(+), 121 deletions(-) create mode 100644 lib/runtime/src/metrics/response_mux.rs create mode 100644 lib/runtime/src/pipeline/network/tcp/mux.rs create mode 100644 lib/runtime/src/pipeline/network/tcp/mux/client.rs diff --git a/benchmarks/frontend/scripts/run_perf.sh b/benchmarks/frontend/scripts/run_perf.sh index 8d4cdc5105a1..c5fe01971ff7 100755 --- a/benchmarks/frontend/scripts/run_perf.sh +++ b/benchmarks/frontend/scripts/run_perf.sh @@ -64,6 +64,12 @@ BENCHMARK_DURATION="${BENCHMARK_DURATION:-}" # aiperf --benchmark-duration (sec REQUEST_RATE="${REQUEST_RATE:-}" # aiperf --request-rate (requests/sec) WARMUP_DURATION="${WARMUP_DURATION:-}" # aiperf --warmup-duration (seconds) WARMUP_COUNT="${WARMUP_COUNT:-}" # aiperf --warmup-request-count +FRONTEND_CORES="${FRONTEND_CORES:-}" # Optional taskset for frontend only +OTHER_CORES="${OTHER_CORES:-}" # Optional shorthand for mockers + aiperf +WORKER_CORES="${WORKER_CORES:-}" # Optional taskset for mockers +CLIENT_CORES="${CLIENT_CORES:-}" # Optional taskset for aiperf +AIPERF_WORKERS_MAX="${AIPERF_WORKERS_MAX:-}" # Optional client worker cap; aiperf auto-sizes by default +AIPERF_RECORD_PROCESSORS="${AIPERF_RECORD_PROCESSORS:-}" # Optional metrics process count; aiperf auto-sizes by default # Opt-out flags SKIP_BPF=false @@ -104,6 +110,12 @@ while [[ $# -gt 0 ]]; do --request-rate) REQUEST_RATE="$2"; shift 2 ;; --warmup-duration) WARMUP_DURATION="$2"; shift 2 ;; --warmup-count) WARMUP_COUNT="$2"; shift 2 ;; + --frontend-cores) FRONTEND_CORES="$2"; shift 2 ;; + --other-cores) OTHER_CORES="$2"; shift 2 ;; + --worker-cores) WORKER_CORES="$2"; shift 2 ;; + --client-cores) CLIENT_CORES="$2"; shift 2 ;; + --aiperf-workers-max) AIPERF_WORKERS_MAX="$2"; shift 2 ;; + --aiperf-record-processors) AIPERF_RECORD_PROCESSORS="$2"; shift 2 ;; --skip-bpf) SKIP_BPF=true; shift ;; --skip-nsys) SKIP_NSYS=true; shift ;; --skip-flamegraph) SKIP_FLAMEGRAPH=true; shift ;; @@ -137,6 +149,13 @@ Service Options: --request-rate N Target requests per second (aiperf --request-rate) --warmup-duration N aiperf warmup phase duration in seconds --warmup-count N aiperf warmup request count (default: concurrency) + --frontend-cores LIST Pin frontend to this taskset CPU list (for example 0-3) + --other-cores LIST Pin mockers and aiperf to this taskset CPU list (for example 4-23) + --worker-cores LIST Pin mockers to this taskset CPU list (overrides --other-cores) + --client-cores LIST Pin aiperf to this taskset CPU list (overrides --other-cores) + --aiperf-workers-max N Override aiperf's auto-sized client worker pool + --aiperf-record-processors N + Override aiperf's auto-sized metrics process pool Load Options: --concurrency N aiperf concurrency (default: 64) @@ -161,6 +180,20 @@ USAGE esac done +[[ -z "$WORKER_CORES" ]] && WORKER_CORES="$OTHER_CORES" +[[ -z "$CLIENT_CORES" ]] && CLIENT_CORES="$OTHER_CORES" +FRONTEND_CPU_CMD=() +WORKER_CPU_CMD=() +CLIENT_CPU_CMD=() +if [[ -n "$FRONTEND_CORES" || -n "$WORKER_CORES" || -n "$CLIENT_CORES" ]]; then + command -v taskset >/dev/null 2>&1 || { + echo "ERROR: taskset is required when CPU lists are configured"; exit 1; + } +fi +[[ -n "$FRONTEND_CORES" ]] && FRONTEND_CPU_CMD=(taskset -c "$FRONTEND_CORES") +[[ -n "$WORKER_CORES" ]] && WORKER_CPU_CMD=(taskset -c "$WORKER_CORES") +[[ -n "$CLIENT_CORES" ]] && CLIENT_CPU_CMD=(taskset -c "$CLIENT_CORES") + # Default model-name to model if not set [[ -z "$MODEL_NAME" ]] && MODEL_NAME="$MODEL" @@ -213,6 +246,8 @@ echo "╚═══════════════════════ echo "" echo "Output: $OUTPUT_DIR" echo "Tokenizer: ${TOKENIZER_BACKEND:-hf (default)}" +echo "CPU sets: frontend=${FRONTEND_CORES:-unrestricted} workers=${WORKER_CORES:-unrestricted} client=${CLIENT_CORES:-unrestricted}" +echo "Response: mux=${DYN_TCP_RESPONSE_MUX:-0} batch=${DYN_TCP_RESPONSE_BATCH_INTERVAL_MS:-5}ms window=${DYN_TCP_RESPONSE_CONNECTION_WINDOW_BYTES:-262144}B packet_metrics=${DYN_TCP_RESPONSE_PACKET_METRICS:-0}" echo "" # ─── Pre-flight: detect available tools ────────────────────────────────────── @@ -409,7 +444,8 @@ for MN in "${MODEL_NAMES[@]}"; do fi MN_SAFE="${MN//\//_}" - HF_HUB_OFFLINE=1 DYN_SYSTEM_PORT=$WORKER_PORT DYN_EVENT_PLANE="$EVENT_PLANE" python -m dynamo.mocker "${MOCKER_ARGS[@]}" \ + HF_HUB_OFFLINE=1 DYN_SYSTEM_PORT=$WORKER_PORT DYN_EVENT_PLANE="$EVENT_PLANE" \ + "${WORKER_CPU_CMD[@]}" python -m dynamo.mocker "${MOCKER_ARGS[@]}" \ > "$OUTPUT_DIR/logs/mocker_${MN_SAFE}_${i}.log" 2>&1 & ALL_PIDS+=($!) echo " Worker $WORKER_IDX ($MN #$i): PID ${ALL_PIDS[-1]}, port $WORKER_PORT" @@ -455,7 +491,7 @@ fi if [[ "$HAS_NSYS" == true ]]; then echo " (under nsys profiling)" - env "${FRONTEND_ENV[@]}" \ + env "${FRONTEND_ENV[@]}" "${FRONTEND_CPU_CMD[@]}" \ "$NSYS_CMD" profile \ --trace=osrt,nvtx \ --sample=cpu \ @@ -485,7 +521,7 @@ if [[ "$HAS_NSYS" == true ]]; then echo " nsys wrapper PID: $NSYS_WRAPPER_PID" echo " Frontend PID: $FRONTEND_PID" else - env "${FRONTEND_ENV[@]}" python -m dynamo.frontend \ + env "${FRONTEND_ENV[@]}" "${FRONTEND_CPU_CMD[@]}" python -m dynamo.frontend \ > "$OUTPUT_DIR/logs/frontend.log" 2>&1 & FRONTEND_PID=$! ALL_PIDS+=($FRONTEND_PID) @@ -725,6 +761,14 @@ else _WARMUP_ARGS+=(--warmup-request-count "$CONCURRENCY") fi +_AIPERF_WORKER_ARGS=() +if [[ -n "$AIPERF_WORKERS_MAX" ]]; then + _AIPERF_WORKER_ARGS+=(--workers-max "$AIPERF_WORKERS_MAX") +fi +if [[ -n "$AIPERF_RECORD_PROCESSORS" ]]; then + _AIPERF_WORKER_ARGS+=(--record-processors "$AIPERF_RECORD_PROCESSORS") +fi + # Build the list of models to target _AIPERF_MODELS=() if [[ "$AIPERF_TARGETS" == "all" && ${#MODEL_NAMES[@]} -gt 1 ]]; then @@ -751,7 +795,7 @@ for _AIPERF_MODEL in "${_AIPERF_MODELS[@]}"; do _AIPERF_TOK_ARGS=(--tokenizer "$MODEL") fi - HF_HUB_OFFLINE=1 aiperf profile --artifact-dir "$AIPERF_ARTIFACT_DIR" \ + HF_HUB_OFFLINE=1 "${CLIENT_CPU_CMD[@]}" aiperf profile --artifact-dir "$AIPERF_ARTIFACT_DIR" \ --model "$_AIPERF_MODEL" \ "${_AIPERF_TOK_ARGS[@]}" \ --endpoint-type chat \ @@ -772,8 +816,7 @@ for _AIPERF_MODEL in "${_AIPERF_MODELS[@]}"; do "${_WARMUP_ARGS[@]}" \ --num-dataset-entries 12800 \ --random-seed 100 \ - --workers-max "$CONCURRENCY" \ - --record-processors 32 \ + "${_AIPERF_WORKER_ARGS[@]}" \ --ui simple || echo "WARNING: aiperf failed for model ${_AIPERF_MODEL}" done @@ -913,6 +956,18 @@ cat > "$OUTPUT_DIR/config.json" <> = metrics_labels .as_ref() .map(|v| v.iter().map(|(k, v)| (k.as_str(), v.as_str())).collect()); + handler.set_response_mux_client(endpoint.drt().response_mux_client())?; // Add metrics to the handler. The endpoint provides additional information to the handler. handler.add_metrics(&endpoint, metrics_labels.as_deref())?; diff --git a/lib/runtime/src/config/environment_names.rs b/lib/runtime/src/config/environment_names.rs index 9d8a97797633..b6ba00fd6e7d 100644 --- a/lib/runtime/src/config/environment_names.rs +++ b/lib/runtime/src/config/environment_names.rs @@ -625,6 +625,28 @@ pub mod tcp_response_stream { /// Host/interface for the TCP response stream server. /// If unset, the server auto-detects a routable local IP. pub const DYN_TCP_RESPONSE_STREAM_HOST: &str = "DYN_TCP_RESPONSE_STREAM_HOST"; + + /// Enables the coordinated multiplexed TCP response transport. + pub const DYN_TCP_RESPONSE_MUX: &str = "DYN_TCP_RESPONSE_MUX"; + + /// Maximum time response data may wait for cross-stream batching. + pub const DYN_TCP_RESPONSE_BATCH_INTERVAL_MS: &str = "DYN_TCP_RESPONSE_BATCH_INTERVAL_MS"; + + /// Maximum encoded bytes in one response batch. + pub const DYN_TCP_RESPONSE_BATCH_MAX_BYTES: &str = "DYN_TCP_RESPONSE_BATCH_MAX_BYTES"; + + /// Maximum logical frames in one response batch. + pub const DYN_TCP_RESPONSE_BATCH_MAX_FRAMES: &str = "DYN_TCP_RESPONSE_BATCH_MAX_FRAMES"; + + /// Per-stream response flow-control window in bytes. + pub const DYN_TCP_RESPONSE_STREAM_WINDOW_BYTES: &str = "DYN_TCP_RESPONSE_STREAM_WINDOW_BYTES"; + + /// Per-connection response flow-control window in bytes. + pub const DYN_TCP_RESPONSE_CONNECTION_WINDOW_BYTES: &str = + "DYN_TCP_RESPONSE_CONNECTION_WINDOW_BYTES"; + + /// Enables diagnostic TCP_INFO data-segment accounting for response sockets. + pub const DYN_TCP_RESPONSE_PACKET_METRICS: &str = "DYN_TCP_RESPONSE_PACKET_METRICS"; } /// Event Plane transport environment variables @@ -872,6 +894,13 @@ mod tests { // TCP Response Stream tcp_response_stream::DYN_TCP_RESPONSE_STREAM_PORT, tcp_response_stream::DYN_TCP_RESPONSE_STREAM_HOST, + tcp_response_stream::DYN_TCP_RESPONSE_MUX, + tcp_response_stream::DYN_TCP_RESPONSE_BATCH_INTERVAL_MS, + tcp_response_stream::DYN_TCP_RESPONSE_BATCH_MAX_BYTES, + tcp_response_stream::DYN_TCP_RESPONSE_BATCH_MAX_FRAMES, + tcp_response_stream::DYN_TCP_RESPONSE_STREAM_WINDOW_BYTES, + tcp_response_stream::DYN_TCP_RESPONSE_CONNECTION_WINDOW_BYTES, + tcp_response_stream::DYN_TCP_RESPONSE_PACKET_METRICS, // Event Plane event_plane::DYN_EVENT_PLANE, event_plane::DYN_EVENT_PLANE_CODEC, diff --git a/lib/runtime/src/distributed.rs b/lib/runtime/src/distributed.rs index dfa0c9f13074..04aa7f73741c 100644 --- a/lib/runtime/src/distributed.rs +++ b/lib/runtime/src/distributed.rs @@ -50,6 +50,7 @@ pub struct DistributedRuntime { nats_client: Option, network_manager: Arc, tcp_server: Arc>>, + response_mux_client: Arc, system_status_server: Arc>>, request_plane: RequestPlaneMode, @@ -193,11 +194,20 @@ impl DistributedRuntime { request_plane, ); + let response_mux_config = + crate::pipeline::network::tcp::mux::initialize_response_mux_config()?; + let response_mux_client = + crate::pipeline::network::tcp::mux::client::ResponseMuxClientPool::new( + runtime.child_token(), + response_mux_config, + ); + let distributed_runtime = Self { runtime, network_manager: Arc::new(network_manager), nats_client, tcp_server: Arc::new(OnceCell::new()), + response_mux_client, system_status_server: Arc::new(OnceLock::new()), discovery_client, discovery_metadata, @@ -213,6 +223,8 @@ impl DistributedRuntime { event_transport_kind, }; + crate::metrics::response_mux::ensure_registered(&distributed_runtime.metrics_registry); + // Initialize the uptime gauge in SystemHealth distributed_runtime .system_health @@ -400,6 +412,12 @@ impl DistributedRuntime { .clone()) } + pub fn response_mux_client( + &self, + ) -> Arc { + self.response_mux_client.clone() + } + /// Get the network manager /// /// The network manager consolidates all network configuration and provides diff --git a/lib/runtime/src/metrics.rs b/lib/runtime/src/metrics.rs index 46e7fbb2a759..bc968e957594 100644 --- a/lib/runtime/src/metrics.rs +++ b/lib/runtime/src/metrics.rs @@ -9,6 +9,7 @@ pub mod frontend_perf; pub mod prometheus_names; pub mod request_plane; +pub mod response_mux; pub mod tokio_perf; pub mod transport_metrics; pub mod work_handler_perf; diff --git a/lib/runtime/src/metrics/response_mux.rs b/lib/runtime/src/metrics/response_mux.rs new file mode 100644 index 000000000000..4e384f6ea97d --- /dev/null +++ b/lib/runtime/src/metrics/response_mux.rs @@ -0,0 +1,447 @@ +// SPDX-FileCopyrightText: Copyright (c) 2024-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +//! Bounded-cardinality metrics for the multiplexed TCP response transport. + +use once_cell::sync::{Lazy, OnceCell}; +use prometheus::{ + Histogram, HistogramOpts, HistogramVec, IntCounter, IntCounterVec, IntGaugeVec, Opts, +}; + +use crate::MetricsRegistry; + +pub static ACTIVE_CONNECTIONS: Lazy = Lazy::new(|| { + IntGaugeVec::new( + Opts::new( + "dynamo_tcp_response_mux_active_connections", + "Active physical TCP response-mux connections", + ), + &["role"], + ) + .expect("response mux active connection gauge") +}); + +pub static CONNECTIONS_TOTAL: Lazy = Lazy::new(|| { + IntCounterVec::new( + Opts::new( + "dynamo_tcp_response_mux_connections_total", + "Physical TCP response-mux connection lifecycle events", + ), + &["role", "result"], + ) + .expect("response mux connection counter") +}); + +pub static ACTIVE_STREAMS: Lazy = Lazy::new(|| { + IntGaugeVec::new( + Opts::new( + "dynamo_tcp_response_mux_active_streams", + "Active logical response streams", + ), + &["role"], + ) + .expect("response mux active stream gauge") +}); + +pub static SETUP_SECONDS: Lazy = Lazy::new(|| { + Histogram::with_opts( + HistogramOpts::new( + "dynamo_tcp_response_mux_setup_seconds", + "Time from request dispatch to logical response-stream prologue", + ) + .buckets(vec![ + 0.0001, 0.0005, 0.001, 0.005, 0.01, 0.05, 0.1, 0.5, 1.0, 5.0, + ]), + ) + .expect("response mux setup histogram") +}); + +pub static FRAMES_TOTAL: Lazy = Lazy::new(|| { + IntCounterVec::new( + Opts::new( + "dynamo_tcp_response_mux_frames_total", + "Multiplexed response frames by wire direction and type", + ), + &["direction", "frame_type"], + ) + .expect("response mux frame counter") +}); + +/// Pre-bound frame counters for one wire direction. Keeping these handles next +/// to a connection avoids a label-map lookup for every generated token. +pub struct FrameCounters { + prologue: IntCounter, + data: IntCounter, + end: IntCounter, + stop: IntCounter, + kill: IntCounter, + window_update: IntCounter, + reset: IntCounter, + connection_ack: IntCounter, +} + +impl FrameCounters { + pub fn for_direction(direction: &str) -> Self { + let counter = |frame_type| { + FRAMES_TOTAL + .with_label_values(&[direction, frame_type]) + .clone() + }; + Self { + prologue: counter("prologue"), + data: counter("data"), + end: counter("end"), + stop: counter("stop"), + kill: counter("kill"), + window_update: counter("window_update"), + reset: counter("reset"), + connection_ack: counter("connection_ack"), + } + } + + pub fn inc(&self, frame_type: &str) { + match frame_type { + "prologue" => self.prologue.inc(), + "data" => self.data.inc(), + "end" => self.end.inc(), + "stop" => self.stop.inc(), + "kill" => self.kill.inc(), + "window_update" => self.window_update.inc(), + "reset" => self.reset.inc(), + "connection_ack" => self.connection_ack.inc(), + _ => debug_assert!(false, "unknown response mux frame type {frame_type}"), + } + } +} + +pub static WRITER_QUEUE_DEPTH: Lazy = Lazy::new(|| { + IntGaugeVec::new( + Opts::new( + "dynamo_tcp_response_mux_writer_queue_depth", + "Queued response-mux frames waiting for the shared writer", + ), + &["role"], + ) + .expect("response mux queue gauge") +}); + +pub static QUEUED_BYTES: Lazy = Lazy::new(|| { + IntGaugeVec::new( + Opts::new( + "dynamo_tcp_response_mux_queued_bytes", + "Encoded response bytes queued in connection writers", + ), + &["role"], + ) + .expect("response mux queued byte gauge") +}); + +pub static FRAMES_PER_WRITE: Lazy = Lazy::new(|| { + HistogramVec::new( + HistogramOpts::new( + "dynamo_tcp_response_mux_frames_per_write", + "Logical response-mux frames encoded into each physical write", + ) + .buckets(vec![1.0, 2.0, 4.0, 8.0, 16.0, 32.0, 64.0]), + &["role"], + ) + .expect("response mux frames-per-write histogram") +}); + +pub static BATCH_BYTES: Lazy = Lazy::new(|| { + HistogramVec::new( + HistogramOpts::new( + "dynamo_tcp_response_mux_batch_bytes", + "Encoded response bytes in each physical TCP write", + ) + .buckets(vec![ + 64.0, 256.0, 1024.0, 4096.0, 16_384.0, 65_536.0, 262_144.0, + ]), + &["role"], + ) + .expect("response mux batch byte histogram") +}); + +pub static BATCH_WAIT_SECONDS: Lazy = Lazy::new(|| { + HistogramVec::new( + HistogramOpts::new( + "dynamo_tcp_response_mux_batch_wait_seconds", + "Observed userspace wait from first selected data frame to write", + ) + .buckets(vec![0.0, 0.0001, 0.0005, 0.001, 0.002, 0.005, 0.01, 0.1]), + &["role"], + ) + .expect("response mux batch wait histogram") +}); + +pub static CONFIGURED_BATCH_INTERVAL_MS: Lazy = Lazy::new(|| { + IntGaugeVec::new( + Opts::new( + "dynamo_tcp_response_mux_configured_batch_interval_ms", + "Configured response data batching interval in milliseconds", + ), + &["role"], + ) + .expect("response mux configured batch interval gauge") +}); + +pub static WRITE_CALLS_TOTAL: Lazy = Lazy::new(|| { + IntCounterVec::new( + Opts::new( + "dynamo_tcp_response_mux_write_calls_total", + "Physical response-mux TCP write calls", + ), + &["role"], + ) + .expect("response mux write call counter") +}); + +pub static DATA_SEGMENTS_TOTAL: Lazy = Lazy::new(|| { + IntCounterVec::new( + Opts::new( + "dynamo_tcp_response_data_segments_total", + "Kernel TCP data segments sent on response sockets when diagnostic packet metrics are enabled", + ), + &["transport", "role"], + ) + .expect("response TCP data segment counter") +}); + +pub static QUEUE_RESIDENCE_SECONDS: Lazy = Lazy::new(|| { + HistogramVec::new( + HistogramOpts::new( + "dynamo_tcp_response_mux_queue_residence_seconds", + "Time logical response frames wait before their physical write", + ) + .buckets(vec![ + 0.000001, 0.00001, 0.0001, 0.001, 0.005, 0.01, 0.1, 1.0, + ]), + &["role"], + ) + .expect("response mux queue residence histogram") +}); + +pub static RESETS_TOTAL: Lazy = Lazy::new(|| { + IntCounterVec::new( + Opts::new( + "dynamo_tcp_response_mux_resets_total", + "Logical response streams reset by role and low-cardinality reason", + ), + &["role", "reason"], + ) + .expect("response mux reset counter") +}); + +pub static RECONNECTS_TOTAL: Lazy = Lazy::new(|| { + IntCounterVec::new( + Opts::new( + "dynamo_tcp_response_mux_reconnects_total", + "Replacement physical response-mux connections", + ), + &["role"], + ) + .expect("response mux reconnect counter") +}); + +pub static FLOW_CONTROL_STALL_SECONDS: Lazy = Lazy::new(|| { + Histogram::with_opts( + HistogramOpts::new( + "dynamo_tcp_response_mux_flow_control_stall_seconds", + "Time response producers wait for stream-local credits", + ) + .buckets(vec![0.000001, 0.00001, 0.0001, 0.001, 0.01, 0.1, 1.0, 10.0]), + ) + .expect("response mux flow-control histogram") +}); + +pub static CONNECTION_FLOW_CONTROL_STALL_SECONDS: Lazy = Lazy::new(|| { + Histogram::with_opts( + HistogramOpts::new( + "dynamo_tcp_response_mux_connection_flow_control_stall_seconds", + "Time response producers wait for physical-connection Data credits", + ) + .buckets(vec![0.000001, 0.00001, 0.0001, 0.001, 0.01, 0.1, 1.0, 10.0]), + ) + .expect("response mux connection flow-control histogram") +}); + +pub static WRITER_ADMISSION_STALL_SECONDS: Lazy = Lazy::new(|| { + Histogram::with_opts( + HistogramOpts::new( + "dynamo_tcp_response_mux_writer_admission_stall_seconds", + "Time response producers wait for their stream-local writer queue", + ) + .buckets(vec![0.000001, 0.00001, 0.0001, 0.001, 0.01, 0.1, 1.0, 10.0]), + ) + .expect("response mux writer-admission histogram") +}); + +pub static QUEUED_BYTE_ADMISSION_STALL_SECONDS: Lazy = Lazy::new(|| { + Histogram::with_opts( + HistogramOpts::new( + "dynamo_tcp_response_mux_queued_byte_admission_stall_seconds", + "Time response producers wait for connection-wide queued-byte capacity", + ) + .buckets(vec![0.000001, 0.00001, 0.0001, 0.001, 0.01, 0.1, 1.0, 10.0]), + ) + .expect("response mux queued-byte admission histogram") +}); + +pub static STREAM_WRITER_QUEUE_OCCUPANCY: Lazy = Lazy::new(|| { + Histogram::with_opts( + HistogramOpts::new( + "dynamo_tcp_response_mux_stream_writer_queue_occupancy", + "Stream-local writer queue occupancy after admission", + ) + .buckets(vec![1.0, 2.0, 4.0, 8.0]), + ) + .expect("response mux stream writer queue occupancy histogram") +}); + +pub static READY_STREAMS: Lazy = Lazy::new(|| { + IntGaugeVec::new( + Opts::new( + "dynamo_tcp_response_mux_ready_streams", + "Logical streams currently scheduled on a fair connection writer", + ), + &["role"], + ) + .expect("response mux ready stream gauge") +}); + +pub static PRIORITY_QUEUE_RESIDENCE_SECONDS: Lazy = Lazy::new(|| { + Histogram::with_opts( + HistogramOpts::new( + "dynamo_tcp_response_mux_priority_queue_residence_seconds", + "Time prologue and reset frames wait in the priority writer lane", + ) + .buckets(vec![0.000001, 0.00001, 0.0001, 0.001, 0.01, 0.1, 1.0]), + ) + .expect("response mux priority queue residence histogram") +}); + +pub static ROUND_ROBIN_TURNS_TOTAL: Lazy = Lazy::new(|| { + IntCounter::new( + "dynamo_tcp_response_mux_round_robin_turns_total", + "Frames selected through the fair per-stream writer ring", + ) + .expect("response mux round-robin turn counter") +}); + +pub static WINDOW_UPDATES_TOTAL: Lazy = Lazy::new(|| { + IntCounterVec::new( + Opts::new( + "dynamo_tcp_response_mux_window_updates_total", + "Response-mux window update frames", + ), + &["direction"], + ) + .expect("response mux window-update counter") +}); + +pub static CONNECTION_LOST_STREAMS_TOTAL: Lazy = Lazy::new(|| { + IntCounter::new( + "dynamo_tcp_response_mux_connection_lost_streams_total", + "Logical streams failed by physical response-mux connection loss", + ) + .expect("response mux connection-lost stream counter") +}); + +static REGISTERED: OnceCell<()> = OnceCell::new(); + +pub fn ensure_registered(registry: &MetricsRegistry) { + let _ = REGISTERED.get_or_init(|| { + registry.add_metric_or_warn( + Box::new(ACTIVE_CONNECTIONS.clone()), + "response_mux_active_connections", + ); + registry.add_metric_or_warn( + Box::new(CONNECTIONS_TOTAL.clone()), + "response_mux_connections_total", + ); + registry.add_metric_or_warn( + Box::new(ACTIVE_STREAMS.clone()), + "response_mux_active_streams", + ); + registry.add_metric_or_warn( + Box::new(SETUP_SECONDS.clone()), + "response_mux_setup_seconds", + ); + registry.add_metric_or_warn(Box::new(FRAMES_TOTAL.clone()), "response_mux_frames_total"); + registry.add_metric_or_warn( + Box::new(WRITER_QUEUE_DEPTH.clone()), + "response_mux_writer_queue_depth", + ); + registry.add_metric_or_warn(Box::new(QUEUED_BYTES.clone()), "response_mux_queued_bytes"); + registry.add_metric_or_warn( + Box::new(FRAMES_PER_WRITE.clone()), + "response_mux_frames_per_write", + ); + registry.add_metric_or_warn(Box::new(BATCH_BYTES.clone()), "response_mux_batch_bytes"); + registry.add_metric_or_warn( + Box::new(BATCH_WAIT_SECONDS.clone()), + "response_mux_batch_wait_seconds", + ); + registry.add_metric_or_warn( + Box::new(CONFIGURED_BATCH_INTERVAL_MS.clone()), + "response_mux_configured_batch_interval_ms", + ); + registry.add_metric_or_warn( + Box::new(WRITE_CALLS_TOTAL.clone()), + "response_mux_write_calls_total", + ); + registry.add_metric_or_warn( + Box::new(DATA_SEGMENTS_TOTAL.clone()), + "response_data_segments_total", + ); + registry.add_metric_or_warn( + Box::new(QUEUE_RESIDENCE_SECONDS.clone()), + "response_mux_queue_residence_seconds", + ); + registry.add_metric_or_warn(Box::new(RESETS_TOTAL.clone()), "response_mux_resets_total"); + registry.add_metric_or_warn( + Box::new(RECONNECTS_TOTAL.clone()), + "response_mux_reconnects_total", + ); + registry.add_metric_or_warn( + Box::new(FLOW_CONTROL_STALL_SECONDS.clone()), + "response_mux_flow_control_stall_seconds", + ); + registry.add_metric_or_warn( + Box::new(CONNECTION_FLOW_CONTROL_STALL_SECONDS.clone()), + "response_mux_connection_flow_control_stall_seconds", + ); + registry.add_metric_or_warn( + Box::new(WRITER_ADMISSION_STALL_SECONDS.clone()), + "response_mux_writer_admission_stall_seconds", + ); + registry.add_metric_or_warn( + Box::new(QUEUED_BYTE_ADMISSION_STALL_SECONDS.clone()), + "response_mux_queued_byte_admission_stall_seconds", + ); + registry.add_metric_or_warn( + Box::new(STREAM_WRITER_QUEUE_OCCUPANCY.clone()), + "response_mux_stream_writer_queue_occupancy", + ); + registry.add_metric_or_warn( + Box::new(READY_STREAMS.clone()), + "response_mux_ready_streams", + ); + registry.add_metric_or_warn( + Box::new(PRIORITY_QUEUE_RESIDENCE_SECONDS.clone()), + "response_mux_priority_queue_residence_seconds", + ); + registry.add_metric_or_warn( + Box::new(ROUND_ROBIN_TURNS_TOTAL.clone()), + "response_mux_round_robin_turns_total", + ); + registry.add_metric_or_warn( + Box::new(WINDOW_UPDATES_TOTAL.clone()), + "response_mux_window_updates_total", + ); + registry.add_metric_or_warn( + Box::new(CONNECTION_LOST_STREAMS_TOTAL.clone()), + "response_mux_connection_lost_streams_total", + ); + }); +} diff --git a/lib/runtime/src/pipeline/network.rs b/lib/runtime/src/pipeline/network.rs index 9256c3f4d676..53ff9d783548 100644 --- a/lib/runtime/src/pipeline/network.rs +++ b/lib/runtime/src/pipeline/network.rs @@ -15,14 +15,19 @@ pub mod manager; pub mod tcp; use crate::SystemHealth; -use std::sync::{Arc, OnceLock}; +use std::{ + future::Future, + pin::Pin, + sync::{Arc, OnceLock}, + task::{Context as TaskContext, Poll}, +}; use anyhow::Result; use async_trait::async_trait; use bytes::Bytes; use codec::{TwoPartCodec, TwoPartMessage, TwoPartMessageType}; use derive_builder::Builder; -use futures::StreamExt; +use futures::{Stream, StreamExt}; // io::Cursor, TryStreamExt use super::{AsyncEngine, AsyncEngineContext, AsyncEngineContextProvider, ResponseStream}; use serde::{Deserialize, Serialize, de::DeserializeOwned}; @@ -357,22 +362,60 @@ mod registered_stream_tests { // receiver, so the sender would have to await the prologue which if // was not an error, would indicate the RequestStreamReceiver is read // to receive data. +#[async_trait] +pub(crate) trait MultiplexedStreamSender: Send + Sync { + async fn send_data(&self, data: Bytes) -> Result<()>; + async fn send_prologue(&self, error: Option) -> Result<(), String>; + async fn finish(&self) -> Result<()>; +} + +enum StreamSenderInner { + Dedicated(tokio::sync::mpsc::Sender), + Multiplexed(Arc), +} + pub struct StreamSender { - tx: tokio::sync::mpsc::Sender, + inner: StreamSenderInner, prologue: Option, } impl StreamSender { + pub(crate) fn dedicated( + tx: tokio::sync::mpsc::Sender, + prologue: Option, + ) -> Self { + Self { + inner: StreamSenderInner::Dedicated(tx), + prologue, + } + } + + pub(crate) fn multiplexed(sender: Arc) -> Self { + Self { + inner: StreamSenderInner::Multiplexed(sender), + prologue: Some(ResponseStreamPrologue { error: None }), + } + } + pub async fn send(&self, data: Bytes) -> Result<()> { - Ok(self.tx.send(TwoPartMessage::from_data(data)).await?) + match &self.inner { + StreamSenderInner::Dedicated(tx) => { + Ok(tx.send(TwoPartMessage::from_data(data)).await?) + } + StreamSenderInner::Multiplexed(sender) => sender.send_data(data).await, + } } pub async fn send_control(&self, control: ControlMessage) -> Result<()> { - let bytes = serde_json::to_vec(&control)?; - Ok(self - .tx - .send(TwoPartMessage::from_header(bytes.into())) - .await?) + match &self.inner { + StreamSenderInner::Dedicated(tx) => { + let bytes = serde_json::to_vec(&control)?; + Ok(tx.send(TwoPartMessage::from_header(bytes.into())).await?) + } + StreamSenderInner::Multiplexed(_) => { + anyhow::bail!("generic control messages are not valid on a mux response sender") + } + } } #[allow(clippy::needless_update)] @@ -390,19 +433,187 @@ impl StreamSender { return Err("Invalid prologue".to_string()); } }; - self.tx - .send(TwoPartMessage::from_header(header_bytes)) - .await - .map_err(|e| e.to_string())?; + match &self.inner { + StreamSenderInner::Dedicated(tx) => tx + .send(TwoPartMessage::from_header(header_bytes)) + .await + .map_err(|e| e.to_string())?, + StreamSenderInner::Multiplexed(sender) => { + sender.send_prologue(prologue.error).await? + } + } } else { panic!("Prologue already sent; or not set; logic error"); } Ok(()) } + + /// Finish one logical response stream without closing a shared physical + /// connection. Dedicated request-stream senders retain their existing + /// drop-driven lifecycle and therefore have no explicit finish action. + pub async fn finish(&self) -> Result<()> { + match &self.inner { + StreamSenderInner::Dedicated(_) => Ok(()), + StreamSenderInner::Multiplexed(sender) => sender.finish().await, + } + } +} + +pub(crate) struct StreamRxItem { + bytes: Bytes, + credit_bytes: usize, +} + +impl StreamRxItem { + pub(crate) fn dedicated(bytes: Bytes) -> Self { + Self { + bytes, + credit_bytes: 0, + } + } + + pub(crate) fn multiplexed(bytes: Bytes, credit_bytes: usize) -> Self { + Self { + bytes, + credit_bytes, + } + } +} + +impl AsRef<[u8]> for StreamRxItem { + fn as_ref(&self) -> &[u8] { + self.bytes.as_ref() + } +} + +pub(crate) struct StreamReceiverHooks { + pub context: Arc, + pub window_update_threshold: usize, + pub on_window_update: Arc, + pub on_close: Arc, } pub struct StreamReceiver { - rx: tokio::sync::mpsc::Receiver, + rx: tokio::sync::mpsc::Receiver, + hooks: Option, + stop_wait: Option + Send>>>, + kill_wait: Option + Send>>>, + consumed_since_update: usize, + stop_sent: bool, + closed: bool, +} + +impl StreamReceiver { + pub(crate) fn dedicated(rx: tokio::sync::mpsc::Receiver) -> Self { + Self { + rx, + hooks: None, + stop_wait: None, + kill_wait: None, + consumed_since_update: 0, + stop_sent: false, + closed: false, + } + } + + pub(crate) fn multiplexed( + rx: tokio::sync::mpsc::Receiver, + hooks: StreamReceiverHooks, + ) -> Self { + let stop_context = hooks.context.clone(); + let kill_context = hooks.context.clone(); + Self { + rx, + hooks: Some(hooks), + stop_wait: Some(Box::pin(async move { stop_context.stopped().await })), + kill_wait: Some(Box::pin(async move { kill_context.killed().await })), + consumed_since_update: 0, + stop_sent: false, + closed: false, + } + } + + pub async fn recv(&mut self) -> Option { + futures::future::poll_fn(|cx| Pin::new(&mut *self).poll_next(cx)).await + } +} + +impl Stream for StreamReceiver { + type Item = Bytes; + + fn poll_next(mut self: Pin<&mut Self>, cx: &mut TaskContext<'_>) -> Poll> { + if self.closed { + return Poll::Ready(None); + } + let killed = self + .kill_wait + .as_mut() + .is_some_and(|wait| wait.as_mut().poll(cx).is_ready()); + if killed { + if let Some(hooks) = self.hooks.as_ref() { + (hooks.on_close)(ControlMessage::Kill); + } + self.closed = true; + return Poll::Ready(None); + } + + let stopped = !self.stop_sent + && self + .stop_wait + .as_mut() + .is_some_and(|wait| wait.as_mut().poll(cx).is_ready()); + if stopped { + if let Some(hooks) = self.hooks.as_ref() { + (hooks.on_close)(ControlMessage::Stop); + } + self.stop_sent = true; + } + + match Pin::new(&mut self.rx).poll_recv(cx) { + Poll::Ready(Some(item)) => { + let threshold = self + .hooks + .as_ref() + .map(|hooks| hooks.window_update_threshold); + if let Some(threshold) = threshold { + self.consumed_since_update = + self.consumed_since_update.saturating_add(item.credit_bytes); + if self.consumed_since_update >= threshold { + if let Some(hooks) = self.hooks.as_ref() { + (hooks.on_window_update)(self.consumed_since_update); + } + self.consumed_since_update = 0; + } + } + Poll::Ready(Some(item.bytes)) + } + Poll::Ready(None) => { + self.closed = true; + Poll::Ready(None) + } + Poll::Pending => Poll::Pending, + } + } +} + +impl Drop for StreamReceiver { + fn drop(&mut self) { + if let Some(hooks) = self.hooks.as_ref() { + let mut credits = self.consumed_since_update; + while let Ok(item) = self.rx.try_recv() { + credits = credits.saturating_add(item.credit_bytes); + } + if credits > 0 { + (hooks.on_window_update)(credits); + } + } + if !self.closed + && let Some(hooks) = self.hooks.as_ref() + { + (hooks.on_close)(ControlMessage::Kill); + self.closed = true; + } + } } /// Connection Info is encoded as JSON and then again serialized has part of the Transport @@ -732,6 +943,7 @@ where pub struct Ingress { segment: OnceLock>>, metrics: OnceLock>, + response_mux_client: OnceLock>, /// Endpoint-specific notifier for health check timer resets endpoint_health_check_notifier: OnceLock>, payload_adapter: Arc, @@ -769,6 +981,7 @@ where Arc::new(Self { segment: OnceLock::new(), metrics: OnceLock::new(), + response_mux_client: OnceLock::new(), endpoint_health_check_notifier: OnceLock::new(), payload_adapter: Arc::new(payload_adapter), }) @@ -842,6 +1055,13 @@ pub trait PushWorkHandler: Send + Sync { metrics_labels: Option<&[(&str, &str)]>, ) -> Result<()>; + fn set_response_mux_client( + &self, + _client: Arc, + ) -> Result<()> { + Ok(()) + } + /// Set the endpoint-specific notifier for health check timer resets fn set_endpoint_health_check_notifier( &self, diff --git a/lib/runtime/src/pipeline/network/egress/addressed_router.rs b/lib/runtime/src/pipeline/network/egress/addressed_router.rs index fcdcf3c0bf7e..3acd4889c353 100644 --- a/lib/runtime/src/pipeline/network/egress/addressed_router.rs +++ b/lib/runtime/src/pipeline/network/egress/addressed_router.rs @@ -41,7 +41,7 @@ use anyhow::{Error, Result}; use futures::stream::Stream; use std::pin::Pin; use std::task::{Context, Poll}; -use tokio_stream::{StreamExt, StreamNotifyClose, wrappers::ReceiverStream}; +use tokio_stream::{StreamExt, StreamNotifyClose}; use tracing::Instrument; /// Stream transformation helper that: @@ -50,7 +50,7 @@ use tracing::Instrument; /// - hands off the `InflightGuard` to a stream-lifetime `InflightDecStream` so /// the inflight gauge stays accurate for the whole response lifetime. fn decode_response_stream( - response_rx: tokio::sync::mpsc::Receiver, + response_rx: crate::pipeline::network::StreamReceiver, engine_ctx: Arc, queue_start: Instant, tx_start: Instant, @@ -63,7 +63,7 @@ where let engine_ctx_for_stream = engine_ctx.clone(); let mut is_complete_final = false; let mut first_response = true; - let stream = StreamNotifyClose::new(ReceiverStream::new(response_rx)).filter_map(move |res| { + let stream = StreamNotifyClose::new(response_rx).filter_map(move |res| { if let Some(res_bytes) = res { if first_response { first_response = false; @@ -604,7 +604,7 @@ impl AddressedPushRouter { drop(_nvtx_wait); Ok(decode_response_stream( - response_stream.rx, + response_stream, engine_ctx, queue_start, tx_start, diff --git a/lib/runtime/src/pipeline/network/egress/tcp_client.rs b/lib/runtime/src/pipeline/network/egress/tcp_client.rs index 0971d5409201..9a183960185c 100644 --- a/lib/runtime/src/pipeline/network/egress/tcp_client.rs +++ b/lib/runtime/src/pipeline/network/egress/tcp_client.rs @@ -158,7 +158,7 @@ struct PendingRequest { /// chunks in `write_buf`. This preserves the large-payload zero-copy path /// without turning deep drains into a heap-allocated iovec list or an /// unbounded all-available-request batch. -struct TcpWriteBuffer { +pub(crate) struct TcpWriteBuffer { // Moving or relocating these Bytes handles does not copy their backing buffers. write_buf: VecDeque, flattened_writes: BytesMut, @@ -166,7 +166,7 @@ struct TcpWriteBuffer { } impl TcpWriteBuffer { - fn new() -> Self { + pub(crate) fn new() -> Self { Self { write_buf: VecDeque::with_capacity(WRITE_VECTORED_CHUNKS), flattened_writes: BytesMut::new(), @@ -197,7 +197,7 @@ impl TcpWriteBuffer { self.write(frame.payload); } - fn write(&mut self, buf: Bytes) { + pub(crate) fn write(&mut self, buf: Bytes) { if buf.is_empty() { return; } @@ -220,12 +220,20 @@ impl TcpWriteBuffer { } async fn write_all(&mut self, writer: &mut W) -> io::Result + where + W: AsyncWrite + Unpin, + { + self.write_all_counted(writer).await.map(|(bytes, _)| bytes) + } + + pub(crate) async fn write_all_counted(&mut self, writer: &mut W) -> io::Result<(usize, u64)> where W: AsyncWrite + Unpin, { self.flush_flattened(); let mut total_written = 0usize; + let mut write_calls = 0_u64; while !self.write_buf.is_empty() { let n = { let mut writes = [IoSlice::new(b""); WRITE_VECTORED_CHUNKS]; @@ -247,10 +255,11 @@ impl TcpWriteBuffer { } total_written += n; + write_calls += 1; self.advance(n); } - Ok(total_written) + Ok((total_written, write_calls)) } fn advance(&mut self, mut n: usize) { diff --git a/lib/runtime/src/pipeline/network/ingress/push_handler.rs b/lib/runtime/src/pipeline/network/ingress/push_handler.rs index 9e6c351d3252..c016477e8cba 100644 --- a/lib/runtime/src/pipeline/network/ingress/push_handler.rs +++ b/lib/runtime/src/pipeline/network/ingress/push_handler.rs @@ -473,7 +473,7 @@ where // response-stream open subsequently fails, the forwarder task // spawned below exits cleanly when `frame_tx.send` observes the // dropped `frame_rx`. - let request_stream_recv = tcp::client::TcpClient::create_request_stream( + let mut request_stream_recv = tcp::client::TcpClient::create_request_stream( context_arc.clone(), req_stream_conn_info, None, @@ -496,8 +496,7 @@ where let forwarder_ctx = context_arc.clone(); let payload_adapter = self.payload_adapter.clone(); tokio::spawn(async move { - let mut rx = request_stream_recv.rx; - while let Some(bytes) = rx.recv().await { + while let Some(bytes) = request_stream_recv.recv().await { // Stop forwarding on either kill or soft-stop, matching the // send-side `spawn_request_stream_forwarder`. Without the // `stopped()` check, a `stop_generating()` would leave this @@ -596,23 +595,35 @@ where WORK_HANDLER_NETWORK_TRANSIT_SECONDS.observe(transit_ns as f64 / 1_000_000_000.0); } - // todo - eventually have a handler class which will returned an abstracted object, but for now, - // we only support tcp here, so we can just unwrap the connection info tracing::trace!("creating tcp response stream"); - let mut publisher = tcp::client::TcpClient::create_response_stream( - request.context(), - response_connection_info, - self.metrics().map(|m| m.cancellation_total.clone()), - ) - .await - .map_err(|e| { - if let Some(m) = self.metrics() { - m.error_counter - .with_label_values(&[work_handler::error_types::RESPONSE_STREAM]) - .inc(); + let mut publisher = + if response_connection_info.transport == tcp::TCP_RESPONSE_MUX_TRANSPORT { + self.response_mux_client + .get() + .ok_or_else(|| { + PipelineError::Generic( + "response mux client was not initialized for endpoint ingress" + .to_string(), + ) + })? + .create_response_stream(request.context(), response_connection_info) + .await + } else { + tcp::client::TcpClient::create_response_stream( + request.context(), + response_connection_info, + self.metrics().map(|m| m.cancellation_total.clone()), + ) + .await } - PipelineError::Generic(format!("Failed to create response stream: {e}")) - })?; + .map_err(|e| { + if let Some(m) = self.metrics() { + m.error_counter + .with_label_values(&[work_handler::error_types::RESPONSE_STREAM]) + .inc(); + } + PipelineError::Generic(format!("Failed to create response stream: {e}")) + })?; tracing::trace!("calling generate"); let stream = self @@ -662,6 +673,9 @@ where self.pump_response_stream(stream, &publisher, payload_codec) .await; + publisher.finish().await.map_err(|err| { + PipelineError::Generic(format!("Failed to finish response stream: {err}")) + })?; // Ensure the metrics guard is not dropped until the end of the function. // Drop fires "request completed" log via RAII. @@ -694,6 +708,15 @@ where Ok(()) } + fn set_response_mux_client( + &self, + client: Arc, + ) -> Result<()> { + self.response_mux_client + .set(client) + .map_err(|_| anyhow::anyhow!("Response mux client already set")) + } + async fn handle_payload( &self, payload: Bytes, @@ -726,6 +749,15 @@ where Ok(()) } + fn set_response_mux_client( + &self, + client: Arc, + ) -> Result<()> { + self.response_mux_client + .set(client) + .map_err(|_| anyhow::anyhow!("Response mux client already set")) + } + async fn handle_payload( &self, payload: Bytes, diff --git a/lib/runtime/src/pipeline/network/tcp.rs b/lib/runtime/src/pipeline/network/tcp.rs index 74bbc003a400..dd75352ef530 100644 --- a/lib/runtime/src/pipeline/network/tcp.rs +++ b/lib/runtime/src/pipeline/network/tcp.rs @@ -115,6 +115,7 @@ //! - Downstream writes: nothing. Its TCP write half is closed right after the CallHome handshake. pub mod client; +pub mod mux; pub mod server; pub mod test_utils; @@ -129,6 +130,59 @@ use super::{ }; const TCP_TRANSPORT: &str = "tcp_server"; +pub const TCP_RESPONSE_MUX_TRANSPORT: &str = "tcp_response_mux_v1"; + +/// Read Linux's kernel-maintained count of data-bearing TCP segments for one +/// socket. This is used only by opt-in benchmark diagnostics; keeping the +/// query here avoids coupling the transport to packet-capture privileges. +#[cfg(target_os = "linux")] +pub(crate) fn tcp_data_segments_out(stream: &tokio::net::TcpStream) -> Option { + use std::os::fd::AsRawFd; + + tcp_data_segments_out_fd(stream.as_raw_fd()) +} + +#[cfg(target_os = "linux")] +pub(crate) fn tcp_data_segments_out_fd(fd: std::os::fd::RawFd) -> Option { + #[repr(C)] + struct LinuxTcpInfoThroughDataSegments { + _header: [u8; 8], + _metrics: [u32; 24], + _rates_and_bytes: [u64; 4], + _segments_out: u32, + _segments_in: u32, + _notsent_bytes: u32, + _min_rtt: u32, + _data_segments_in: u32, + data_segments_out: u32, + } + + let mut info = std::mem::MaybeUninit::::zeroed(); + let mut len = std::mem::size_of::() as libc::socklen_t; + let result = unsafe { + libc::getsockopt( + fd, + libc::IPPROTO_TCP, + libc::TCP_INFO, + info.as_mut_ptr().cast(), + &mut len, + ) + }; + if result != 0 || len < std::mem::size_of::() as _ { + return None; + } + Some(unsafe { info.assume_init() }.data_segments_out as u64) +} + +#[cfg(not(target_os = "linux"))] +pub(crate) fn tcp_data_segments_out(_stream: &tokio::net::TcpStream) -> Option { + None +} + +#[cfg(not(target_os = "linux"))] +pub(crate) fn tcp_data_segments_out_fd(_fd: std::os::fd::RawFd) -> Option { + None +} #[derive(Debug, Clone, Serialize, Deserialize)] pub struct TcpStreamConnectionInfo { @@ -169,6 +223,42 @@ impl TryFrom for TcpStreamConnectionInfo { } } +/// Connection information for one logical response stream carried by the +/// frontend's persistent response-mux listener. +#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)] +pub struct ResponseMuxConnectionInfo { + pub address: String, + pub frontend_server_id: uuid::Uuid, + pub stream_id: uuid::Uuid, + pub context: String, + pub version: u8, +} + +impl From for ConnectionInfo { + fn from(info: ResponseMuxConnectionInfo) -> Self { + Self { + transport: TCP_RESPONSE_MUX_TRANSPORT.to_string(), + info: serde_json::to_string(&info) + .expect("Failed to serialize ResponseMuxConnectionInfo"), + } + } +} + +impl TryFrom for ResponseMuxConnectionInfo { + type Error = anyhow::Error; + + fn try_from(info: ConnectionInfo) -> Result { + if info.transport != TCP_RESPONSE_MUX_TRANSPORT { + return Err(anyhow::anyhow!( + "Invalid transport; response mux requires `{TCP_RESPONSE_MUX_TRANSPORT}`; got {}", + info.transport + )); + } + serde_json::from_str(&info.info) + .map_err(|e| anyhow::anyhow!("Failed to parse response mux connection info: {e}")) + } +} + /// First message sent over a CallHome stream which will map the newly created socket to a specific /// response data stream which was registered with the same subject. /// @@ -247,14 +337,14 @@ mod tests { send_stream.send(payload.into()).await.unwrap(); // [client] The client can now receive the response from the server - let data = recv_stream.rx.recv().await.unwrap(); + let data = recv_stream.recv().await.unwrap(); let recv_msg = serde_json::from_slice::(&data).unwrap(); assert_eq!(msg.foo, recv_msg.foo); // Dropping the upstream `StreamSender` should cleanly close the request // stream — the downstream receiver should observe `None`. drop(send_stream); - assert!(recv_stream.rx.recv().await.is_none()); + assert!(recv_stream.recv().await.is_none()); } #[tokio::test] @@ -306,7 +396,7 @@ mod tests { // [server] After client sends the prologue, the server can pick up its `StreamReceiver` half. let (_conn_info, stream_provider) = pending_connection.recv_stream.unwrap().into_parts(); - let recv_stream = stream_provider.await.unwrap(); + let mut recv_stream = stream_provider.await.unwrap(); // [client] The client can now send the response message to the server let msg = TestMessage { @@ -319,7 +409,7 @@ mod tests { // [server] The server can now receive the response message from the client - let data = recv_stream.unwrap().rx.recv().await.unwrap(); + let data = recv_stream.as_mut().unwrap().recv().await.unwrap(); let recv_msg = serde_json::from_slice::(&data).unwrap(); diff --git a/lib/runtime/src/pipeline/network/tcp/client.rs b/lib/runtime/src/pipeline/network/tcp/client.rs index 28f6ac9a5145..cf61a1b89222 100644 --- a/lib/runtime/src/pipeline/network/tcp/client.rs +++ b/lib/runtime/src/pipeline/network/tcp/client.rs @@ -17,7 +17,7 @@ use prometheus::IntCounter; use super::{CallHomeHandshake, ControlMessage, TcpStreamConnectionInfo}; use crate::engine::AsyncEngineContext; use crate::pipeline::network::{ - ConnectionInfo, ResponseStreamPrologue, StreamReceiver, StreamSender, + ConnectionInfo, ResponseStreamPrologue, StreamReceiver, StreamRxItem, StreamSender, codec::{TwoPartCodec, TwoPartMessage, TwoPartMessageType}, tcp::StreamType, }; @@ -87,6 +87,9 @@ impl TcpClient { } let stream = TcpClient::connect(&info.address).await?; + let packet_baseline = super::mux::response_packet_metrics_enabled() + .then(|| super::tcp_data_segments_out(&stream)) + .flatten(); let peer_port = stream.peer_addr().ok().map(|addr| addr.port()); let (read_half, write_half) = tokio::io::split(stream); @@ -152,6 +155,7 @@ impl TcpClient { monitor_context, peer_port, subject, + packet_baseline, ) .await; }); @@ -161,10 +165,7 @@ impl TcpClient { let prologue = Some(ResponseStreamPrologue { error: None }); // create the stream sender - let stream_sender = StreamSender { - tx: bytes_tx, - prologue, - }; + let stream_sender = StreamSender::dedicated(bytes_tx, prologue); Ok(stream_sender) } @@ -230,7 +231,7 @@ impl TcpClient { // never writes again, so close the write half immediately. drop(framed_writer); - let (bytes_tx, bytes_rx) = tokio::sync::mpsc::channel::(64); + let (bytes_tx, bytes_rx) = tokio::sync::mpsc::channel::(64); tokio::spawn(handle_request_reader( framed_reader, @@ -239,13 +240,13 @@ impl TcpClient { cancellation_counter, )); - Ok(StreamReceiver { rx: bytes_rx }) + Ok(StreamReceiver::dedicated(bytes_rx)) } } async fn handle_request_reader( mut framed_reader: FramedRead, TwoPartCodec>, - bytes_tx: tokio::sync::mpsc::Sender, + bytes_tx: tokio::sync::mpsc::Sender, context: Arc, cancellation_counter: Option, ) { @@ -319,7 +320,7 @@ async fn handle_request_reader( } } TwoPartMessageType::DataOnly(data) => { - if bytes_tx.send(data).await.is_err() { + if bytes_tx.send(StreamRxItem::dedicated(data)).await.is_err() { tracing::debug!("downstream consumer dropped; exiting request-stream reader"); break; } @@ -372,6 +373,7 @@ async fn wait_for_connection_tasks( context: Arc, peer_port: Option, subject: String, + packet_baseline: Option, ) -> Result<()> { // Await the reader first and abort the writer on reader Err — the // writer parks on `bytes_rx.recv()` and won't wake on its own. @@ -418,6 +420,13 @@ async fn wait_for_connection_tasks( }; let stream = reader.unsplit(writer); + if let Some(baseline) = packet_baseline + && let Some(current) = super::tcp_data_segments_out(&stream) + { + crate::metrics::response_mux::DATA_SEGMENTS_TOTAL + .with_label_values(&["dedicated", "worker"]) + .inc_by(current.saturating_sub(baseline)); + } wait_for_server_shutdown(stream, context).await } @@ -1061,6 +1070,7 @@ mod tests { monitor_context, None, "test-subject".to_string(), + None, ), ) .await; @@ -1120,6 +1130,7 @@ mod tests { context, None, "test-reader-panic".to_string(), + None, ), ) .await; @@ -1500,8 +1511,8 @@ mod tests { struct RequestReaderHarness { framed_server: FramedWrite, TwoPartCodec>, framed_reader: FramedRead, TwoPartCodec>, - bytes_tx: mpsc::Sender, - bytes_rx: mpsc::Receiver, + bytes_tx: mpsc::Sender, + bytes_rx: mpsc::Receiver, controller: Arc, } @@ -1512,7 +1523,7 @@ mod tests { let framed_reader = FramedRead::new(read_half, TwoPartCodec::default()); let framed_server = FramedWrite::new(server_write, TwoPartCodec::default()); - let (bytes_tx, bytes_rx) = mpsc::channel::(64); + let (bytes_tx, bytes_rx) = mpsc::channel::(64); let controller = Arc::new(Controller::default()); RequestReaderHarness { diff --git a/lib/runtime/src/pipeline/network/tcp/mux.rs b/lib/runtime/src/pipeline/network/tcp/mux.rs new file mode 100644 index 000000000000..6ca5a12c91bf --- /dev/null +++ b/lib/runtime/src/pipeline/network/tcp/mux.rs @@ -0,0 +1,698 @@ +// SPDX-FileCopyrightText: Copyright (c) 2024-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +//! Multiplexed TCP response-stream protocol. +//! +//! A short [`TwoPartCodec`] handshake validates the version and frontend +//! identity. The persistent connection then switches to [`MuxCodec`], whose +//! compact fixed-width header carries the frame kind and logical stream UUID. +//! Connection writers can batch those frames without changing their wire +//! representation. + +use std::{io, sync::OnceLock, time::Duration}; + +use bytes::{BufMut, Bytes, BytesMut}; +use serde::{Deserialize, Serialize}; +use uuid::Uuid; + +use crate::pipeline::{ + error::TwoPartCodecError, + network::codec::{TwoPartCodec, TwoPartMessage, TwoPartMessageType}, +}; +use tokio_util::codec::{Decoder, Encoder}; + +pub mod client; + +pub const RESPONSE_MUX_VERSION: u8 = 1; +pub const RESPONSE_MUX_POOL_SIZE: usize = 4; +pub const RESPONSE_MUX_WRITER_QUEUE: usize = 4096; +pub const RESPONSE_MUX_STREAM_WRITER_QUEUE: usize = 8; +pub const RESPONSE_MUX_IDLE_TTL_SECS: u64 = 300; +pub const RESPONSE_MUX_CONNECT_TIMEOUT_SECS: u64 = 5; + +pub const RESPONSE_MUX_DEFAULT_BATCH_INTERVAL_MS: u64 = 5; +pub const RESPONSE_MUX_MAX_BATCH_INTERVAL_MS: u64 = 100; +pub const RESPONSE_MUX_DEFAULT_BATCH_MAX_BYTES: usize = 65_536; +pub const RESPONSE_MUX_DEFAULT_BATCH_MAX_FRAMES: usize = 64; +pub const RESPONSE_MUX_DEFAULT_STREAM_WINDOW_BYTES: usize = 262_144; +pub const RESPONSE_MUX_DEFAULT_CONNECTION_WINDOW_BYTES: usize = 262_144; +pub const RESPONSE_MUX_CREDIT_UPDATE_BYTES: usize = 65_536; +pub const RESPONSE_MUX_CREDIT_UPDATE_INTERVAL: Duration = Duration::from_millis(1); +pub const RESPONSE_MUX_SCHEDULER_QUANTUM: usize = 8; + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct ResponseMuxConfig { + pub enabled: bool, + pub packet_metrics: bool, + pub batch_interval: Duration, + pub batch_max_bytes: usize, + pub batch_max_frames: usize, + pub stream_window_bytes: usize, + pub connection_window_bytes: usize, +} + +impl ResponseMuxConfig { + fn parse_with(mut read: impl FnMut(&str) -> Option) -> anyhow::Result { + use crate::config::environment_names::tcp_response_stream as env; + + fn parse( + read: &mut impl FnMut(&str) -> Option, + name: &str, + default: T, + ) -> anyhow::Result + where + T: std::str::FromStr, + T::Err: std::fmt::Display, + { + match read(name) { + None => Ok(default), + Some(value) => value + .parse::() + .map_err(|err| anyhow::anyhow!("invalid {name}={value:?}: {err}")), + } + } + + let parse_bool = |name, value: Option| match value.as_deref() { + None | Some("") | Some("0") | Some("false") => Ok(false), + Some("1") | Some("true") => Ok(true), + Some(value) => anyhow::bail!("invalid {name}={value:?}; expected 0, 1, false, or true"), + }; + let enabled = parse_bool(env::DYN_TCP_RESPONSE_MUX, read(env::DYN_TCP_RESPONSE_MUX))?; + let packet_metrics = parse_bool( + env::DYN_TCP_RESPONSE_PACKET_METRICS, + read(env::DYN_TCP_RESPONSE_PACKET_METRICS), + )?; + let interval_ms = parse( + &mut read, + env::DYN_TCP_RESPONSE_BATCH_INTERVAL_MS, + RESPONSE_MUX_DEFAULT_BATCH_INTERVAL_MS, + )?; + if interval_ms > RESPONSE_MUX_MAX_BATCH_INTERVAL_MS { + anyhow::bail!( + "{} must be at most {} ms; got {}", + env::DYN_TCP_RESPONSE_BATCH_INTERVAL_MS, + RESPONSE_MUX_MAX_BATCH_INTERVAL_MS, + interval_ms + ); + } + let batch_max_bytes = parse( + &mut read, + env::DYN_TCP_RESPONSE_BATCH_MAX_BYTES, + RESPONSE_MUX_DEFAULT_BATCH_MAX_BYTES, + )?; + let batch_max_frames = parse( + &mut read, + env::DYN_TCP_RESPONSE_BATCH_MAX_FRAMES, + RESPONSE_MUX_DEFAULT_BATCH_MAX_FRAMES, + )?; + let stream_window_bytes = parse( + &mut read, + env::DYN_TCP_RESPONSE_STREAM_WINDOW_BYTES, + RESPONSE_MUX_DEFAULT_STREAM_WINDOW_BYTES, + )?; + let connection_window_bytes = parse( + &mut read, + env::DYN_TCP_RESPONSE_CONNECTION_WINDOW_BYTES, + RESPONSE_MUX_DEFAULT_CONNECTION_WINDOW_BYTES, + )?; + for (name, value) in [ + (env::DYN_TCP_RESPONSE_BATCH_MAX_BYTES, batch_max_bytes), + (env::DYN_TCP_RESPONSE_BATCH_MAX_FRAMES, batch_max_frames), + ( + env::DYN_TCP_RESPONSE_STREAM_WINDOW_BYTES, + stream_window_bytes, + ), + ( + env::DYN_TCP_RESPONSE_CONNECTION_WINDOW_BYTES, + connection_window_bytes, + ), + ] { + if value == 0 { + anyhow::bail!("{name} must be greater than zero"); + } + } + for (name, value) in [ + ( + env::DYN_TCP_RESPONSE_STREAM_WINDOW_BYTES, + stream_window_bytes, + ), + ( + env::DYN_TCP_RESPONSE_CONNECTION_WINDOW_BYTES, + connection_window_bytes, + ), + ] { + if value > u32::MAX as usize { + anyhow::bail!("{name} must fit in an unsigned 32-bit credit update"); + } + } + Ok(Self { + enabled, + packet_metrics, + batch_interval: Duration::from_millis(interval_ms), + batch_max_bytes, + batch_max_frames, + stream_window_bytes, + connection_window_bytes, + }) + } + + pub fn from_env() -> anyhow::Result { + Self::parse_with(|name| std::env::var(name).ok()) + } +} + +static RESPONSE_MUX_CONFIG: OnceLock = OnceLock::new(); + +pub fn initialize_response_mux_config() -> anyhow::Result { + if let Some(config) = RESPONSE_MUX_CONFIG.get() { + return Ok(*config); + } + let config = ResponseMuxConfig::from_env()?; + let _ = RESPONSE_MUX_CONFIG.set(config); + Ok(*RESPONSE_MUX_CONFIG.get().expect("response mux config set")) +} + +pub fn response_packet_metrics_enabled() -> bool { + RESPONSE_MUX_CONFIG + .get() + .is_some_and(|config| config.packet_metrics) +} + +pub const MUX_HEADER_LEN: usize = 24; + +/// First header-only frame on a newly accepted TCP stream. +#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)] +#[serde(tag = "kind", rename_all = "snake_case")] +pub enum ConnectionHandshake { + /// Dedicated per-request upstream -> downstream request stream. + RequestStream { subject: String }, + /// Persistent connection carrying many downstream -> upstream responses. + ResponseMux { + version: u8, + frontend_server_id: Uuid, + connection_id: Uuid, + }, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +#[repr(u8)] +pub enum MuxFrameKind { + Prologue = 1, + Data = 2, + End = 3, + Stop = 4, + Kill = 5, + WindowUpdate = 6, + Reset = 7, + ConnectionAck = 8, +} + +impl TryFrom for MuxFrameKind { + type Error = io::Error; + + fn try_from(value: u8) -> Result { + match value { + 1 => Ok(Self::Prologue), + 2 => Ok(Self::Data), + 3 => Ok(Self::End), + 4 => Ok(Self::Stop), + 5 => Ok(Self::Kill), + 6 => Ok(Self::WindowUpdate), + 7 => Ok(Self::Reset), + 8 => Ok(Self::ConnectionAck), + _ => Err(io::Error::new( + io::ErrorKind::InvalidData, + format!("unknown response mux frame kind {value}"), + )), + } + } +} + +impl MuxFrameKind { + pub const fn metric_label(self) -> &'static str { + match self { + Self::Prologue => "prologue", + Self::Data => "data", + Self::End => "end", + Self::Stop => "stop", + Self::Kill => "kill", + Self::WindowUpdate => "window_update", + Self::Reset => "reset", + Self::ConnectionAck => "connection_ack", + } + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct MuxFrame { + pub kind: MuxFrameKind, + pub stream_id: Uuid, + pub payload: Bytes, +} + +impl MuxFrame { + pub fn new(kind: MuxFrameKind, stream_id: Uuid, payload: Bytes) -> Self { + Self { + kind, + stream_id, + payload, + } + } + + pub fn empty(kind: MuxFrameKind, stream_id: Uuid) -> Self { + Self::new(kind, stream_id, Bytes::new()) + } + + pub fn window_update(stream_id: Uuid, credits: u32) -> Self { + let mut payload = BytesMut::with_capacity(4); + payload.put_u32(credits); + Self::new(MuxFrameKind::WindowUpdate, stream_id, payload.freeze()) + } + + pub fn connection_ack(decoded_bytes: u64) -> Self { + Self::new( + MuxFrameKind::ConnectionAck, + Uuid::nil(), + decoded_bytes.to_be_bytes().to_vec().into(), + ) + } + + pub fn connection_ack_offset(&self) -> io::Result { + if self.kind != MuxFrameKind::ConnectionAck || self.payload.len() != 8 { + return Err(io::Error::new( + io::ErrorKind::InvalidData, + "response mux connection ACK must contain eight bytes", + )); + } + Ok(u64::from_be_bytes( + self.payload + .as_ref() + .try_into() + .expect("validated connection ACK length"), + )) + } + + pub fn encoded_len(&self) -> usize { + MUX_HEADER_LEN + self.payload.len() + } + + pub fn window_credits(&self) -> io::Result { + if self.kind != MuxFrameKind::WindowUpdate || self.payload.len() != 4 { + return Err(io::Error::new( + io::ErrorKind::InvalidData, + "credit update must contain exactly four payload bytes", + )); + } + Ok(u32::from_be_bytes(self.payload[..4].try_into().unwrap())) + } + + pub fn into_two_part(self) -> TwoPartMessage { + let mut header = BytesMut::with_capacity(20); + header.put_u8(self.kind as u8); + header.put_u8(0); // flags, reserved for future protocol use + header.put_u16(0); + header.extend_from_slice(self.stream_id.as_bytes()); + TwoPartMessage::new(header.freeze(), self.payload) + } + + /// Split the wire representation into a small header and the original + /// payload allocation so connection writers can coalesce headers while + /// retaining large payloads as `Bytes` for vectored I/O. + pub fn encode_parts(&self) -> io::Result<(Bytes, Bytes)> { + let mut header = BytesMut::with_capacity(MUX_HEADER_LEN); + MuxCodec::default().encode_header(self, &mut header)?; + Ok((header.freeze(), self.payload.clone())) + } + + pub fn try_from_two_part(message: TwoPartMessage) -> io::Result { + let (header, payload) = match message.into_message_type() { + TwoPartMessageType::HeaderOnly(header) => (header, Bytes::new()), + TwoPartMessageType::HeaderAndData(header, data) => (header, data), + TwoPartMessageType::DataOnly(_) | TwoPartMessageType::Empty => { + return Err(io::Error::new( + io::ErrorKind::InvalidData, + "response mux frame is missing its fixed header", + )); + } + }; + + if header.len() != 20 { + return Err(io::Error::new( + io::ErrorKind::InvalidData, + format!( + "invalid response mux header length {}, expected 20", + header.len() + ), + )); + } + if header[1..4] != [0, 0, 0] { + return Err(io::Error::new( + io::ErrorKind::InvalidData, + "response mux frame has unsupported flags", + )); + } + + let kind = MuxFrameKind::try_from(header[0])?; + let stream_id = Uuid::from_slice(&header[4..20]).map_err(|err| { + io::Error::new( + io::ErrorKind::InvalidData, + format!("invalid response mux stream UUID: {err}"), + ) + })?; + let is_connection_frame = kind == MuxFrameKind::ConnectionAck; + if stream_id.is_nil() != is_connection_frame { + return Err(io::Error::new( + io::ErrorKind::InvalidData, + "only connection-level frames must use the nil stream UUID", + )); + } + + Self::validate(Self::new(kind, stream_id, payload)) + } + + fn validate(frame: Self) -> io::Result { + let is_connection_frame = frame.kind == MuxFrameKind::ConnectionAck; + if frame.stream_id.is_nil() != is_connection_frame { + return Err(io::Error::new( + io::ErrorKind::InvalidData, + "only connection-level frames must use the nil stream UUID", + )); + } + match frame.kind { + MuxFrameKind::Stop | MuxFrameKind::Kill if !frame.payload.is_empty() => { + Err(io::Error::new( + io::ErrorKind::InvalidData, + "control frame must not contain a payload", + )) + } + MuxFrameKind::WindowUpdate => { + frame.window_credits()?; + Ok(frame) + } + MuxFrameKind::ConnectionAck => { + frame.connection_ack_offset()?; + Ok(frame) + } + _ => Ok(frame), + } + } +} + +/// Compact response-mux framing used after the versioned connection +/// handshake. Each frame is `payload_len:u32`, kind, flags, reserved, UUID, +/// then payload. The fixed header is 24 bytes. +#[derive(Clone, Debug)] +pub struct MuxCodec { + max_message_size: usize, +} + +impl Default for MuxCodec { + fn default() -> Self { + Self::new(crate::pipeline::network::get_tcp_max_message_size()) + } +} + +impl MuxCodec { + pub fn new(max_message_size: usize) -> Self { + Self { max_message_size } + } + + fn encode_header(&self, frame: &MuxFrame, dst: &mut BytesMut) -> io::Result<()> { + let payload_len = u32::try_from(frame.payload.len()).map_err(|_| { + io::Error::new( + io::ErrorKind::InvalidInput, + "response mux payload exceeds u32", + ) + })?; + let encoded_len = MUX_HEADER_LEN + .checked_add(frame.payload.len()) + .ok_or_else(|| io::Error::new(io::ErrorKind::InvalidInput, "frame size overflow"))?; + if encoded_len > self.max_message_size { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + format!( + "response mux frame size {encoded_len} exceeds maximum {}", + self.max_message_size + ), + )); + } + dst.reserve(MUX_HEADER_LEN); + dst.put_u32(payload_len); + dst.put_u8(frame.kind as u8); + dst.put_u8(0); + dst.put_u16(0); + dst.extend_from_slice(frame.stream_id.as_bytes()); + Ok(()) + } +} + +impl Encoder for MuxCodec { + type Error = io::Error; + + fn encode(&mut self, frame: MuxFrame, dst: &mut BytesMut) -> io::Result<()> { + self.encode_header(&frame, dst)?; + dst.extend_from_slice(&frame.payload); + Ok(()) + } +} + +impl Decoder for MuxCodec { + type Item = MuxFrame; + type Error = io::Error; + + fn decode(&mut self, src: &mut BytesMut) -> io::Result> { + if src.len() < MUX_HEADER_LEN { + return Ok(None); + } + let payload_len = u32::from_be_bytes(src[..4].try_into().unwrap()) as usize; + let encoded_len = MUX_HEADER_LEN.checked_add(payload_len).ok_or_else(|| { + io::Error::new( + io::ErrorKind::InvalidData, + "response mux frame size overflow", + ) + })?; + if encoded_len > self.max_message_size { + return Err(io::Error::new( + io::ErrorKind::InvalidData, + format!( + "response mux frame size {encoded_len} exceeds maximum {}", + self.max_message_size + ), + )); + } + if src.len() < encoded_len { + src.reserve(encoded_len - src.len()); + return Ok(None); + } + if src[5] != 0 || src[6..8] != [0, 0] { + return Err(io::Error::new( + io::ErrorKind::InvalidData, + "response mux frame has unsupported flags or reserved bits", + )); + } + let kind = MuxFrameKind::try_from(src[4])?; + let stream_id = Uuid::from_slice(&src[8..24]).map_err(|err| { + io::Error::new( + io::ErrorKind::InvalidData, + format!("invalid response mux stream UUID: {err}"), + ) + })?; + let mut encoded = src.split_to(encoded_len); + let payload = encoded.split_off(MUX_HEADER_LEN).freeze(); + MuxFrame::validate(MuxFrame::new(kind, stream_id, payload)).map(Some) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::pipeline::network::codec::TwoPartCodec; + use tokio_util::codec::{Decoder, Encoder}; + + fn encode(frame: MuxFrame) -> BytesMut { + let mut bytes = BytesMut::new(); + TwoPartCodec::default() + .encode(frame.into_two_part(), &mut bytes) + .unwrap(); + bytes + } + + fn encode_compact(frame: MuxFrame) -> BytesMut { + let mut bytes = BytesMut::new(); + MuxCodec::default().encode(frame, &mut bytes).unwrap(); + bytes + } + + #[test] + fn mux_frame_round_trip_preserves_uuid_and_payload() { + let stream_id = Uuid::new_v4(); + let frame = MuxFrame::new( + MuxFrameKind::Data, + stream_id, + Bytes::from_static(b"payload"), + ); + let codec = TwoPartCodec::default(); + let encoded = codec.encode_message(frame.clone().into_two_part()).unwrap(); + let decoded = codec.decode_message(encoded).unwrap(); + assert_eq!(MuxFrame::try_from_two_part(decoded).unwrap(), frame); + } + + #[test] + fn window_update_requires_four_bytes() { + let frame = MuxFrame::new( + MuxFrameKind::WindowUpdate, + Uuid::new_v4(), + Bytes::from_static(&[1, 2]), + ); + assert!(MuxFrame::try_from_two_part(frame.into_two_part()).is_err()); + } + + #[test] + fn connection_ack_round_trips_with_cumulative_byte_offset() { + let frame = MuxFrame::connection_ack(987_654); + let decoded = MuxFrame::try_from_two_part(frame.clone().into_two_part()).unwrap(); + assert_eq!(decoded, frame); + assert_eq!(decoded.connection_ack_offset().unwrap(), 987_654); + assert_eq!(decoded.encoded_len(), 32); + } + + #[test] + fn connection_ack_rejects_stream_uuid() { + let frame = MuxFrame::new( + MuxFrameKind::ConnectionAck, + Uuid::new_v4(), + 128_u64.to_be_bytes().to_vec().into(), + ); + assert!(MuxFrame::try_from_two_part(frame.into_two_part()).is_err()); + } + + #[test] + fn unknown_kind_is_rejected() { + let mut encoded = encode_compact(MuxFrame::empty(MuxFrameKind::End, Uuid::new_v4())); + encoded[4] = 99; + assert!(MuxCodec::default().decode(&mut encoded).is_err()); + } + + #[test] + fn partial_compact_frame_waits_for_remaining_bytes() { + let expected = MuxFrame::new( + MuxFrameKind::Data, + Uuid::new_v4(), + Bytes::from_static(b"partial"), + ); + let encoded = encode_compact(expected.clone()); + let split = encoded.len() / 2; + let mut input = BytesMut::from(&encoded[..split]); + let mut codec = MuxCodec::default(); + assert!(codec.decode(&mut input).unwrap().is_none()); + input.extend_from_slice(&encoded[split..]); + let decoded = codec.decode(&mut input).unwrap().unwrap(); + assert_eq!(decoded, expected); + } + + #[test] + fn concatenated_frames_decode_independently_and_route_by_uuid() { + let first = MuxFrame::new( + MuxFrameKind::Data, + Uuid::new_v4(), + Bytes::from_static(b"one"), + ); + let second = MuxFrame::empty(MuxFrameKind::End, Uuid::new_v4()); + let mut input = encode_compact(first.clone()); + input.extend_from_slice(&encode_compact(second.clone())); + let mut codec = MuxCodec::default(); + let decoded_first = codec.decode(&mut input).unwrap().unwrap(); + let decoded_second = codec.decode(&mut input).unwrap().unwrap(); + assert_eq!(decoded_first, first); + assert_eq!(decoded_second, second); + assert_ne!(decoded_first.stream_id, decoded_second.stream_id); + } + + #[test] + fn maximum_compact_size_is_enforced() { + let frame = MuxFrame::new( + MuxFrameKind::Data, + Uuid::new_v4(), + Bytes::from_static(b"bounded"), + ); + let exact_len = MUX_HEADER_LEN + frame.payload.len(); + let mut bytes = BytesMut::new(); + assert!( + MuxCodec::new(exact_len) + .encode(frame.clone(), &mut bytes) + .is_ok() + ); + assert!( + MuxCodec::new(exact_len - 1) + .encode(frame, &mut BytesMut::new()) + .is_err() + ); + } + + #[test] + fn malformed_outer_lengths_are_rejected() { + let mut input = BytesMut::new(); + input.put_u64(u64::MAX); + input.put_u64(1); + input.put_u64(0); + assert!(TwoPartCodec::default().decode(&mut input).is_err()); + } + + #[test] + fn compact_codec_rejects_flags_and_connection_uuid_mismatch() { + let mut flags = encode_compact(MuxFrame::empty(MuxFrameKind::End, Uuid::new_v4())); + flags[5] = 1; + assert!(MuxCodec::default().decode(&mut flags).is_err()); + + let mut invalid = encode_compact(MuxFrame::connection_ack(64)); + invalid[8..24].copy_from_slice(Uuid::new_v4().as_bytes()); + assert!(MuxCodec::default().decode(&mut invalid).is_err()); + } + + fn config(values: &[(&str, &str)]) -> anyhow::Result { + let values = values + .iter() + .map(|(key, value)| ((*key).to_string(), (*value).to_string())) + .collect::>(); + ResponseMuxConfig::parse_with(|name| values.get(name).cloned()) + } + + #[test] + fn response_mux_config_defaults_to_disabled_and_five_ms() { + let config = config(&[]).unwrap(); + assert!(!config.enabled); + assert!(!config.packet_metrics); + assert_eq!(config.batch_interval, Duration::from_millis(5)); + assert_eq!(config.batch_max_bytes, 65_536); + assert_eq!(config.batch_max_frames, 64); + assert_eq!(config.stream_window_bytes, 262_144); + assert_eq!(config.connection_window_bytes, 262_144); + } + + #[test] + fn response_mux_config_accepts_zero_delay_and_valid_overrides() { + use crate::config::environment_names::tcp_response_stream as env; + let config = config(&[ + (env::DYN_TCP_RESPONSE_MUX, "1"), + (env::DYN_TCP_RESPONSE_PACKET_METRICS, "true"), + (env::DYN_TCP_RESPONSE_BATCH_INTERVAL_MS, "0"), + (env::DYN_TCP_RESPONSE_BATCH_MAX_BYTES, "8192"), + (env::DYN_TCP_RESPONSE_BATCH_MAX_FRAMES, "8"), + ]) + .unwrap(); + assert!(config.enabled); + assert!(config.packet_metrics); + assert_eq!(config.batch_interval, Duration::ZERO); + assert_eq!(config.batch_max_bytes, 8192); + assert_eq!(config.batch_max_frames, 8); + } + + #[test] + fn response_mux_config_rejects_malformed_and_over_100_ms() { + use crate::config::environment_names::tcp_response_stream as env; + assert!(config(&[(env::DYN_TCP_RESPONSE_BATCH_INTERVAL_MS, "nope")]).is_err()); + assert!(config(&[(env::DYN_TCP_RESPONSE_BATCH_INTERVAL_MS, "101")]).is_err()); + assert!(config(&[(env::DYN_TCP_RESPONSE_BATCH_INTERVAL_MS, "100")]).is_ok()); + assert!(config(&[(env::DYN_TCP_RESPONSE_PACKET_METRICS, "maybe")]).is_err()); + } +} diff --git a/lib/runtime/src/pipeline/network/tcp/mux/client.rs b/lib/runtime/src/pipeline/network/tcp/mux/client.rs new file mode 100644 index 000000000000..810cd48e87a1 --- /dev/null +++ b/lib/runtime/src/pipeline/network/tcp/mux/client.rs @@ -0,0 +1,2150 @@ +// SPDX-FileCopyrightText: Copyright (c) 2024-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +//! Worker-side persistent multiplexed TCP response connection pool. + +use std::{ + collections::VecDeque, + sync::{ + Arc, Weak, + atomic::{AtomicBool, AtomicU64, AtomicUsize, Ordering}, + }, + time::{Duration, Instant}, +}; + +use anyhow::{Context, Result, anyhow}; +use dashmap::{DashMap, mapref::entry::Entry}; +use futures::{SinkExt, StreamExt}; +use parking_lot::{Mutex, RwLock}; +use tokio::{ + net::TcpStream, + sync::{OwnedSemaphorePermit, Semaphore, mpsc, oneshot}, +}; +use tokio_util::{ + codec::{FramedRead, FramedWrite}, + sync::CancellationToken, +}; +use uuid::Uuid; + +use super::{ + ConnectionHandshake, MuxCodec, MuxFrame, MuxFrameKind, RESPONSE_MUX_CONNECT_TIMEOUT_SECS, + RESPONSE_MUX_IDLE_TTL_SECS, RESPONSE_MUX_POOL_SIZE, RESPONSE_MUX_STREAM_WRITER_QUEUE, + RESPONSE_MUX_VERSION, RESPONSE_MUX_WRITER_QUEUE, ResponseMuxConfig, +}; +use crate::{ + engine::AsyncEngineContext, + metrics::response_mux, + pipeline::network::{ + ConnectionInfo, MultiplexedStreamSender, ResponseStreamPrologue, StreamSender, + codec::{TwoPartCodec, TwoPartMessage}, + egress::tcp_client::TcpWriteBuffer, + tcp::ResponseMuxConnectionInfo, + }, +}; + +struct WriterCommand { + frame: MuxFrame, + written: Option>>, + _writer_permit: Option, + _queued_byte_permit: Option, + priority_enqueued_at: Option, + enqueued_at: Instant, +} + +impl WriterCommand { + fn new(frame: MuxFrame, written: Option>>) -> Self { + Self { + frame, + written, + _writer_permit: None, + _queued_byte_permit: None, + priority_enqueued_at: None, + enqueued_at: Instant::now(), + } + } + + fn priority(frame: MuxFrame, written: Option>>) -> Self { + let mut command = Self::new(frame, written); + command.priority_enqueued_at = Some(Instant::now()); + command + } + + fn with_writer_permit(mut self, permit: OwnedSemaphorePermit) -> Self { + self._writer_permit = Some(permit); + self + } + + fn with_queued_byte_permit(mut self, permit: OwnedSemaphorePermit) -> Self { + self._queued_byte_permit = Some(permit); + self + } + + fn fail(mut self, reason: &str) { + if let Some(written) = self.written.take() { + let _ = written.send(Err(reason.to_string())); + } + } +} + +#[inline] +fn per_frame_metrics_enabled() -> bool { + true +} + +#[derive(Clone, Copy)] +struct PoolConfig { + pool_size: usize, + writer_queue: usize, + stream_writer_queue: usize, + initial_window: usize, + connection_window: usize, + batch_interval: Duration, + batch_max_bytes: usize, + batch_max_frames: usize, + packet_metrics: bool, + idle_ttl: Duration, + connect_timeout: Duration, +} + +impl PoolConfig { + fn from_runtime(config: ResponseMuxConfig) -> Self { + Self { + pool_size: RESPONSE_MUX_POOL_SIZE, + writer_queue: RESPONSE_MUX_WRITER_QUEUE, + stream_writer_queue: RESPONSE_MUX_STREAM_WRITER_QUEUE, + initial_window: config.stream_window_bytes, + connection_window: config.connection_window_bytes, + batch_interval: config.batch_interval, + batch_max_bytes: config.batch_max_bytes, + batch_max_frames: config.batch_max_frames, + packet_metrics: config.packet_metrics, + idle_ttl: Duration::from_secs(RESPONSE_MUX_IDLE_TTL_SECS), + connect_timeout: Duration::from_secs(RESPONSE_MUX_CONNECT_TIMEOUT_SECS), + } + } +} + +#[derive(Default)] +struct StreamWriterState { + pending: VecDeque, + scheduled: bool, +} + +struct WorkerStreamState { + context: Arc, + credits: Arc, + max_credits: usize, + writer_slots: Arc, + writer: Mutex, + closed: AtomicBool, + close_token: CancellationToken, +} + +type ScheduledStream = (Uuid, Arc); +type BlockedData = (WriterCommand, Option, Instant); + +impl WorkerStreamState { + fn replenish_credits(&self, credits: usize) -> usize { + if self.closed.load(Ordering::Acquire) || self.credits.is_closed() { + return 0; + } + let available = self.credits.available_permits(); + let replenished = credits.min(self.max_credits.saturating_sub(available)); + if replenished > 0 { + self.credits.add_permits(replenished); + } + replenished + } +} + +struct MuxConnection { + id: u64, + cancel: CancellationToken, + priority_tx: mpsc::Sender, + ready_tx: mpsc::UnboundedSender, + streams: DashMap>, + healthy: AtomicBool, + active_streams: AtomicUsize, + queued_frames: AtomicUsize, + queued_bytes: AtomicUsize, + max_queued_bytes: usize, + queued_byte_slots: Arc, + connection_credits: Arc, + max_connection_credits: usize, + sent_data_bytes: AtomicU64, + acknowledged_data_bytes: AtomicU64, + batch_interval: Duration, + batch_max_bytes: usize, + batch_max_frames: usize, +} + +impl MuxConnection { + async fn connect( + id: u64, + address: &str, + frontend_server_id: Uuid, + version: u8, + cancel: CancellationToken, + config: PoolConfig, + ) -> Result> { + let stream = tokio::time::timeout(config.connect_timeout, TcpStream::connect(address)) + .await + .map_err(|_| anyhow!("response mux connect timeout to {address}"))??; + stream.set_nodelay(true)?; + let packet_baseline = config + .packet_metrics + .then(|| crate::pipeline::network::tcp::tcp_data_segments_out(&stream)) + .flatten(); + + let (read_half, write_half) = stream.into_split(); + let mux_codec = + || TwoPartCodec::new(Some(crate::pipeline::network::get_tcp_max_message_size())); + let mut handshake_reader = FramedRead::new(read_half, mux_codec()); + let mut handshake_writer = FramedWrite::new(write_half, mux_codec()); + let handshake = ConnectionHandshake::ResponseMux { + version, + frontend_server_id, + connection_id: Uuid::new_v4(), + }; + let header = serde_json::to_vec(&handshake)?; + handshake_writer + .send(TwoPartMessage::from_header(header.into())) + .await + .context("failed to send response mux handshake")?; + let ack = tokio::time::timeout(config.connect_timeout, handshake_reader.next()) + .await + .map_err(|_| anyhow!("response mux handshake ack timeout from {address}"))? + .ok_or_else(|| anyhow!("frontend closed before response mux handshake ack"))??; + let ack = MuxFrame::try_from_two_part(ack)?; + if ack.kind != MuxFrameKind::ConnectionAck || ack.connection_ack_offset()? != 0 { + anyhow::bail!("frontend returned invalid response mux connection ack"); + } + let read_half = handshake_reader.into_inner(); + let write_half = handshake_writer.into_inner(); + + let (priority_tx, priority_rx) = mpsc::channel(config.writer_queue); + let (ready_tx, ready_rx) = mpsc::unbounded_channel(); + let cancel = cancel.child_token(); + let connection = Arc::new(Self { + id, + cancel: cancel.clone(), + priority_tx, + ready_tx, + streams: DashMap::new(), + healthy: AtomicBool::new(true), + active_streams: AtomicUsize::new(0), + queued_frames: AtomicUsize::new(0), + queued_bytes: AtomicUsize::new(0), + max_queued_bytes: config.connection_window, + queued_byte_slots: Arc::new(Semaphore::new(config.connection_window)), + connection_credits: Arc::new(Semaphore::new(config.connection_window)), + max_connection_credits: config.connection_window, + sent_data_bytes: AtomicU64::new(0), + acknowledged_data_bytes: AtomicU64::new(0), + batch_interval: config.batch_interval, + batch_max_bytes: config.batch_max_bytes, + batch_max_frames: config.batch_max_frames, + }); + response_mux::CONNECTIONS_TOTAL + .with_label_values(&["worker", "created"]) + .inc(); + response_mux::ACTIVE_CONNECTIONS + .with_label_values(&["worker"]) + .inc(); + + tokio::spawn(Self::writer_task( + Arc::downgrade(&connection), + write_half, + priority_rx, + ready_rx, + cancel.clone(), + )); + tokio::spawn(Self::reader_task( + Arc::downgrade(&connection), + FramedRead::new(read_half, MuxCodec::default()), + cancel, + packet_baseline, + )); + Ok(connection) + } + + fn is_healthy(&self) -> bool { + self.healthy.load(Ordering::Acquire) + } + + fn replenish_connection_credits(&self, credits: usize) -> usize { + if !self.is_healthy() || self.connection_credits.is_closed() { + return 0; + } + let available = self.connection_credits.available_permits(); + let replenished = credits.min(self.max_connection_credits.saturating_sub(available)); + if replenished > 0 { + self.connection_credits.add_permits(replenished); + } + replenished + } + + fn acknowledge_connection_credits(&self, acknowledged_bytes: u64) -> Result<()> { + let previous = self.acknowledged_data_bytes.load(Ordering::Acquire); + if acknowledged_bytes < previous { + anyhow::bail!( + "response mux connection ACK moved backwards from {previous} to {acknowledged_bytes}" + ); + } + if acknowledged_bytes == previous { + return Ok(()); + } + let sent = self.sent_data_bytes.load(Ordering::Acquire); + if acknowledged_bytes > sent { + anyhow::bail!( + "response mux connection ACK {acknowledged_bytes} exceeds sent offset {sent}" + ); + } + self.acknowledged_data_bytes + .store(acknowledged_bytes, Ordering::Release); + let delta = acknowledged_bytes.saturating_sub(previous) as usize; + self.replenish_connection_credits(delta.min(self.max_connection_credits)); + Ok(()) + } + + fn fail(&self, reason: &str) { + if !self.healthy.swap(false, Ordering::AcqRel) { + return; + } + tracing::warn!( + connection_id = self.id, + reason, + "response mux connection failed" + ); + self.cancel.cancel(); + self.connection_credits.close(); + self.queued_byte_slots.close(); + response_mux::CONNECTIONS_TOTAL + .with_label_values(&["worker", "failed"]) + .inc(); + response_mux::ACTIVE_CONNECTIONS + .with_label_values(&["worker"]) + .dec(); + let stream_ids: Vec = self.streams.iter().map(|entry| *entry.key()).collect(); + response_mux::CONNECTION_LOST_STREAMS_TOTAL.inc_by(stream_ids.len() as u64); + for stream_id in stream_ids { + self.remove_stream(stream_id, reason, true); + } + } + + fn close_stream_state(&self, state: &WorkerStreamState, reason: &str) { + if state.closed.swap(true, Ordering::AcqRel) { + return; + } + state.credits.close(); + state.writer_slots.close(); + state.close_token.cancel(); + let (pending, was_scheduled) = { + let mut writer = state.writer.lock(); + let pending = writer.pending.drain(..).collect::>(); + let was_scheduled = writer.scheduled; + writer.scheduled = false; + (pending, was_scheduled) + }; + self.queued_frames + .fetch_sub(pending.len(), Ordering::AcqRel); + let pending_bytes = pending + .iter() + .map(|command| command.frame.encoded_len()) + .sum::(); + self.queued_bytes.fetch_sub(pending_bytes, Ordering::AcqRel); + response_mux::QUEUED_BYTES + .with_label_values(&["worker"]) + .sub(pending_bytes as i64); + if was_scheduled { + response_mux::READY_STREAMS + .with_label_values(&["worker"]) + .dec(); + } + for command in pending { + command.fail(reason); + } + } + + fn remove_stream(&self, stream_id: Uuid, reason: &str, kill_context: bool) -> bool { + let Some((_, state)) = self.streams.remove(&stream_id) else { + return false; + }; + self.close_stream_state(&state, reason); + if kill_context { + state.context.kill(); + } + self.active_streams.fetch_sub(1, Ordering::AcqRel); + response_mux::ACTIVE_STREAMS + .with_label_values(&["worker"]) + .dec(); + true + } + + async fn send_priority_command(&self, command: WriterCommand) -> Result<()> { + if !self.is_healthy() { + command.fail("response mux connection is unhealthy"); + anyhow::bail!("response mux connection is unhealthy"); + } + self.queued_frames.fetch_add(1, Ordering::AcqRel); + match self.priority_tx.try_send(command) { + Ok(()) => Ok(()), + Err(mpsc::error::TrySendError::Full(command)) => { + if let Err(err) = self.priority_tx.send(command).await { + self.queued_frames.fetch_sub(1, Ordering::AcqRel); + err.0.fail("response mux priority writer stopped"); + anyhow::bail!("response mux priority writer stopped"); + } + Ok(()) + } + Err(mpsc::error::TrySendError::Closed(command)) => { + self.queued_frames.fetch_sub(1, Ordering::AcqRel); + command.fail("response mux priority writer stopped"); + anyhow::bail!("response mux priority writer stopped") + } + } + } + + fn try_send_priority_command(&self, command: WriterCommand) { + if !self.is_healthy() { + command.fail("response mux connection is unhealthy"); + return; + } + self.queued_frames.fetch_add(1, Ordering::AcqRel); + if let Err(err) = self.priority_tx.try_send(command) { + self.queued_frames.fetch_sub(1, Ordering::AcqRel); + let reason = match err { + mpsc::error::TrySendError::Full(command) => { + command.fail("response mux priority writer queue is full"); + "response mux priority writer queue is full" + } + mpsc::error::TrySendError::Closed(command) => { + command.fail("response mux priority writer is closed"); + "response mux priority writer is closed" + } + }; + self.fail(reason); + } + } + + fn enqueue_stream_command( + &self, + stream_id: Uuid, + state: &WorkerStreamState, + command: WriterCommand, + ) -> Result<()> { + if !self.is_healthy() || state.closed.load(Ordering::Acquire) { + command.fail("response mux stream is closed"); + anyhow::bail!("response mux stream is closed"); + } + + let encoded_len = command.frame.encoded_len(); + self.queued_bytes.fetch_add(encoded_len, Ordering::AcqRel); + response_mux::QUEUED_BYTES + .with_label_values(&["worker"]) + .add(encoded_len as i64); + + let mut writer = state.writer.lock(); + if state.closed.load(Ordering::Acquire) { + self.queued_bytes.fetch_sub(encoded_len, Ordering::AcqRel); + response_mux::QUEUED_BYTES + .with_label_values(&["worker"]) + .sub(encoded_len as i64); + command.fail("response mux stream is closed"); + anyhow::bail!("response mux stream is closed"); + } + writer.pending.push_back(command); + self.queued_frames.fetch_add(1, Ordering::AcqRel); + if per_frame_metrics_enabled() { + response_mux::STREAM_WRITER_QUEUE_OCCUPANCY.observe(writer.pending.len() as f64); + } + if !writer.scheduled { + writer.scheduled = true; + response_mux::READY_STREAMS + .with_label_values(&["worker"]) + .inc(); + if self.ready_tx.send(stream_id).is_err() { + writer.scheduled = false; + response_mux::READY_STREAMS + .with_label_values(&["worker"]) + .dec(); + self.queued_frames.fetch_sub(1, Ordering::AcqRel); + if let Some(command) = writer.pending.pop_back() { + self.queued_bytes.fetch_sub(encoded_len, Ordering::AcqRel); + response_mux::QUEUED_BYTES + .with_label_values(&["worker"]) + .sub(encoded_len as i64); + command.fail("response mux fair writer stopped"); + } + anyhow::bail!("response mux fair writer stopped"); + } + } + Ok(()) + } + + fn reschedule_stream(&self, stream_id: Uuid, state: &Arc) -> Result<()> { + response_mux::ROUND_ROBIN_TURNS_TOTAL.inc(); + let mut writer = state.writer.lock(); + if !state.closed.load(Ordering::Acquire) && !writer.pending.is_empty() { + self.ready_tx + .send(stream_id) + .map_err(|_| anyhow!("response mux fair writer stopped"))?; + } else if writer.scheduled { + writer.scheduled = false; + response_mux::READY_STREAMS + .with_label_values(&["worker"]) + .dec(); + } + Ok(()) + } + + async fn writer_task( + weak: Weak, + mut write_half: tokio::net::tcp::OwnedWriteHalf, + mut priority_rx: mpsc::Receiver, + mut ready_rx: mpsc::UnboundedReceiver, + cancel: CancellationToken, + ) { + let mut write_buf = TcpWriteBuffer::new(); + let queue_depth = response_mux::WRITER_QUEUE_DEPTH + .with_label_values(&["worker"]) + .clone(); + let frames_per_write = response_mux::FRAMES_PER_WRITE + .with_label_values(&["worker"]) + .clone(); + let frame_counters = response_mux::FrameCounters::for_direction("worker_to_frontend"); + let metrics_enabled = per_frame_metrics_enabled(); + let mut reported_queue_depth = 0_i64; + let mut blocked_data: Option = None; + let result: Result<()> = async { + loop { + let connection = weak + .upgrade() + .ok_or_else(|| anyhow!("response mux connection dropped"))?; + if metrics_enabled { + let current_queue_depth = + connection.queued_frames.load(Ordering::Acquire) as i64; + queue_depth.add(current_queue_depth - reported_queue_depth); + reported_queue_depth = current_queue_depth; + } + + let (command, scheduled_stream, connection_permit) = if let Some(( + blocked_command, + blocked_stream, + blocked_since, + )) = blocked_data.take() + { + if blocked_command.frame.kind != MuxFrameKind::Data { + (blocked_command, blocked_stream, None) + } else { + enum BlockedNext { + Priority(WriterCommand), + Credit(OwnedSemaphorePermit), + StreamClosed, + } + let blocked_close = blocked_stream + .as_ref() + .expect("Data commands are always stream-scheduled") + .1 + .close_token + .clone(); + let credits = connection.connection_credits.clone(); + let next = tokio::select! { + biased; + _ = cancel.cancelled() => return Ok(()), + _ = blocked_close.cancelled() => BlockedNext::StreamClosed, + Some(command) = priority_rx.recv() => BlockedNext::Priority(command), + permit = credits.acquire_many_owned( + blocked_command + .frame + .encoded_len() + .min(connection.max_connection_credits) as u32 + ) => BlockedNext::Credit( + permit.map_err(|_| anyhow!( + "response mux connection closed while writer awaited credits" + ))? + ), + }; + match next { + BlockedNext::Priority(command) => { + connection.queued_frames.fetch_sub(1, Ordering::AcqRel); + blocked_data = Some((blocked_command, blocked_stream, blocked_since)); + (command, None, None) + } + BlockedNext::Credit(permit) => { + if metrics_enabled { + response_mux::CONNECTION_FLOW_CONTROL_STALL_SECONDS + .observe(blocked_since.elapsed().as_secs_f64()); + } + (blocked_command, blocked_stream, Some(permit)) + } + BlockedNext::StreamClosed => { + blocked_command + .fail("response mux stream closed while writer awaited credits"); + continue; + } + } + } + } else { + let (command, scheduled_stream) = loop { + if let Ok(command) = priority_rx.try_recv() { + connection.queued_frames.fetch_sub(1, Ordering::AcqRel); + break (command, None); + } + + enum Next { + Priority(WriterCommand), + Stream(Uuid), + } + let next = tokio::select! { + biased; + _ = cancel.cancelled() => return Ok(()), + command = priority_rx.recv() => command.map(Next::Priority), + stream_id = ready_rx.recv() => stream_id.map(Next::Stream), + }; + let Some(next) = next else { + return Ok(()); + }; + match next { + Next::Priority(command) => { + connection.queued_frames.fetch_sub(1, Ordering::AcqRel); + break (command, None); + } + Next::Stream(stream_id) => { + let Some(state) = connection + .streams + .get(&stream_id) + .map(|entry| entry.value().clone()) + else { + continue; + }; + let command = state.writer.lock().pending.pop_front(); + if let Some(command) = command { + connection.queued_frames.fetch_sub(1, Ordering::AcqRel); + let encoded_len = command.frame.encoded_len(); + connection + .queued_bytes + .fetch_sub(encoded_len, Ordering::AcqRel); + response_mux::QUEUED_BYTES + .with_label_values(&["worker"]) + .sub(encoded_len as i64); + break (command, Some((stream_id, state))); + } + let mut writer = state.writer.lock(); + if writer.scheduled { + writer.scheduled = false; + response_mux::READY_STREAMS + .with_label_values(&["worker"]) + .dec(); + } + } + } + }; + if command.frame.kind == MuxFrameKind::Data { + let required = command + .frame + .encoded_len() + .min(connection.max_connection_credits) + as u32; + match connection + .connection_credits + .clone() + .try_acquire_many_owned(required) + { + Ok(permit) => (command, scheduled_stream, Some(permit)), + Err(tokio::sync::TryAcquireError::NoPermits) => { + blocked_data = Some((command, scheduled_stream, Instant::now())); + continue; + } + Err(tokio::sync::TryAcquireError::Closed) => { + return Err(anyhow!( + "response mux connection closed while writer acquired credits" + )); + } + } + } else { + (command, scheduled_stream, None) + } + }; + + let batching_started = Instant::now(); + let first_is_data = command.frame.kind == MuxFrameKind::Data; + let mut batch = vec![(command, connection_permit)]; + let mut batch_bytes = batch[0].0.frame.encoded_len(); + let mut force_flush = !first_is_data; + if let Some((stream_id, state)) = scheduled_stream { + let mut held_for_next_turn = false; + for _ in 1..super::RESPONSE_MUX_SCHEDULER_QUANTUM { + if force_flush + || batch.len() >= connection.batch_max_frames + || batch_bytes >= connection.batch_max_bytes + { + break; + } + let Some(next) = state.writer.lock().pending.pop_front() else { + break; + }; + connection.queued_frames.fetch_sub(1, Ordering::AcqRel); + let encoded_len = next.frame.encoded_len(); + connection + .queued_bytes + .fetch_sub(encoded_len, Ordering::AcqRel); + response_mux::QUEUED_BYTES + .with_label_values(&["worker"]) + .sub(encoded_len as i64); + if batch_bytes.saturating_add(encoded_len) > connection.batch_max_bytes { + blocked_data = Some((next, Some((stream_id, state.clone())), Instant::now())); + held_for_next_turn = true; + break; + } + let permit = if next.frame.kind == MuxFrameKind::Data { + let required = + encoded_len.min(connection.max_connection_credits) as u32; + match connection + .connection_credits + .clone() + .try_acquire_many_owned(required) + { + Ok(permit) => Some(permit), + Err(tokio::sync::TryAcquireError::NoPermits) => { + blocked_data = Some(( + next, + Some((stream_id, state.clone())), + Instant::now(), + )); + held_for_next_turn = true; + break; + } + Err(tokio::sync::TryAcquireError::Closed) => { + return Err(anyhow!( + "response mux connection credit window closed" + )); + } + } + } else { + None + }; + force_flush = next.frame.kind != MuxFrameKind::Data; + batch_bytes = batch_bytes.saturating_add(encoded_len); + batch.push((next, permit)); + } + if !held_for_next_turn { + connection.reschedule_stream(stream_id, &state)?; + } + } + let deadline = batching_started + connection.batch_interval; + + while first_is_data + && !force_flush + && batch.len() < connection.batch_max_frames + && batch_bytes < connection.batch_max_bytes + { + if let Ok(priority) = priority_rx.try_recv() { + connection.queued_frames.fetch_sub(1, Ordering::AcqRel); + batch_bytes = batch_bytes.saturating_add(priority.frame.encoded_len()); + batch.push((priority, None)); + break; + } + + let next_stream = match ready_rx.try_recv() { + Ok(stream_id) => Some(stream_id), + Err(mpsc::error::TryRecvError::Disconnected) => return Ok(()), + Err(mpsc::error::TryRecvError::Empty) + if connection.batch_interval.is_zero() => + { + None + } + Err(mpsc::error::TryRecvError::Empty) => { + tokio::select! { + biased; + _ = cancel.cancelled() => return Ok(()), + Some(priority) = priority_rx.recv() => { + connection.queued_frames.fetch_sub(1, Ordering::AcqRel); + batch_bytes = batch_bytes.saturating_add(priority.frame.encoded_len()); + batch.push((priority, None)); + break; + } + stream_id = ready_rx.recv() => stream_id, + _ = tokio::time::sleep_until(deadline.into()) => None, + } + } + }; + let Some(stream_id) = next_stream else { + break; + }; + let Some(state) = connection + .streams + .get(&stream_id) + .map(|entry| entry.value().clone()) + else { + continue; + }; + let Some(next) = state.writer.lock().pending.pop_front() else { + connection.reschedule_stream(stream_id, &state)?; + continue; + }; + connection.queued_frames.fetch_sub(1, Ordering::AcqRel); + let encoded_len = next.frame.encoded_len(); + connection + .queued_bytes + .fetch_sub(encoded_len, Ordering::AcqRel); + response_mux::QUEUED_BYTES + .with_label_values(&["worker"]) + .sub(encoded_len as i64); + if !batch.is_empty() + && (batch.len() + 1 > connection.batch_max_frames + || batch_bytes.saturating_add(encoded_len) + > connection.batch_max_bytes) + { + blocked_data = Some((next, Some((stream_id, state)), Instant::now())); + break; + } + let permit = if next.frame.kind == MuxFrameKind::Data { + let required = encoded_len.min(connection.max_connection_credits) as u32; + match connection + .connection_credits + .clone() + .try_acquire_many_owned(required) + { + Ok(permit) => Some(permit), + Err(tokio::sync::TryAcquireError::NoPermits) => { + blocked_data = + Some((next, Some((stream_id, state)), Instant::now())); + break; + } + Err(tokio::sync::TryAcquireError::Closed) => { + return Err(anyhow!("response mux connection credit window closed")); + } + } + } else { + None + }; + let urgent = next.frame.kind != MuxFrameKind::Data; + batch_bytes = batch_bytes.saturating_add(encoded_len); + batch.push((next, permit)); + connection.reschedule_stream(stream_id, &state)?; + if urgent { + break; + } + } + + for (command, _) in &batch { + let (header, payload) = command.frame.encode_parts()?; + write_buf.write(header); + write_buf.write(payload); + } + let observed_batch_wait = batching_started.elapsed(); + let data_bytes = batch + .iter() + .filter(|(command, _)| command.frame.kind == MuxFrameKind::Data) + .map(|(command, _)| command.frame.encoded_len() as u64) + .sum::(); + connection + .sent_data_bytes + .fetch_add(data_bytes, Ordering::AcqRel); + let mut write_calls = 0_u64; + let write_result: Result<()> = async { + let (_, calls) = write_buf.write_all_counted(&mut write_half).await?; + write_calls = calls; + Ok(()) + } + .await; + for (command, permit) in &mut batch { + if metrics_enabled { + response_mux::QUEUE_RESIDENCE_SECONDS + .with_label_values(&["worker"]) + .observe(command.enqueued_at.elapsed().as_secs_f64()); + if let Some(enqueued_at) = command.priority_enqueued_at { + response_mux::PRIORITY_QUEUE_RESIDENCE_SECONDS + .observe(enqueued_at.elapsed().as_secs_f64()); + } + frame_counters.inc(command.frame.kind.metric_label()); + } + if let Some(written) = command.written.take() { + let _ = written.send( + write_result + .as_ref() + .map(|_| ()) + .map_err(|err| err.to_string()), + ); + } + if let Some(permit) = permit.take() { + permit.forget(); + } + } + write_result?; + if metrics_enabled { + frames_per_write.observe(batch.len() as f64); + response_mux::BATCH_BYTES + .with_label_values(&["worker"]) + .observe(batch_bytes as f64); + response_mux::BATCH_WAIT_SECONDS + .with_label_values(&["worker"]) + .observe(observed_batch_wait.as_secs_f64()); + response_mux::WRITE_CALLS_TOTAL + .with_label_values(&["worker"]) + .inc_by(write_calls); + } + } + } + .await; + if metrics_enabled { + queue_depth.sub(reported_queue_depth); + } + + if let Some(connection) = weak.upgrade() { + connection.fail( + &result + .err() + .map(|err| err.to_string()) + .unwrap_or_else(|| "writer stopped".to_string()), + ); + } + } + + async fn reader_task( + weak: Weak, + mut reader: FramedRead, + cancel: CancellationToken, + mut reported_data_segments: Option, + ) { + let frame_counters = response_mux::FrameCounters::for_direction("frontend_to_worker"); + let window_updates = response_mux::WINDOW_UPDATES_TOTAL + .with_label_values(&["frontend_to_worker"]) + .clone(); + let connection_window_updates = response_mux::WINDOW_UPDATES_TOTAL + .with_label_values(&["connection_frontend_to_worker"]) + .clone(); + let metrics_enabled = per_frame_metrics_enabled(); + let mut packet_tick = tokio::time::interval(Duration::from_millis(100)); + packet_tick.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay); + let result: Result<()> = async { + loop { + let message = tokio::select! { + _ = cancel.cancelled() => break, + _ = packet_tick.tick(), if reported_data_segments.is_some() => { + if let Some(current) = crate::pipeline::network::tcp::tcp_data_segments_out(reader.get_ref().as_ref()) { + let previous = reported_data_segments.replace(current).unwrap_or(current); + response_mux::DATA_SEGMENTS_TOTAL + .with_label_values(&["mux", "worker"]) + .inc_by(current.saturating_sub(previous)); + } + continue; + } + message = reader.next() => match message { + Some(message) => message.map_err(|err| anyhow!(err.to_string()))?, + None => return Err(anyhow!("frontend closed response mux connection")), + }, + }; + let frame = message; + if metrics_enabled { + frame_counters.inc(frame.kind.metric_label()); + } + let connection = weak + .upgrade() + .ok_or_else(|| anyhow!("response mux connection dropped"))?; + if frame.kind == MuxFrameKind::ConnectionAck { + let offset = frame.connection_ack_offset()?; + if offset > 0 { + connection.acknowledge_connection_credits(offset)?; + if metrics_enabled { + connection_window_updates.inc(); + } + } + continue; + } + let Some(state) = connection.streams.get(&frame.stream_id) else { + if frame.kind != MuxFrameKind::Reset { + connection.try_send_priority_command(WriterCommand::priority( + MuxFrame::empty(MuxFrameKind::Reset, frame.stream_id), + None, + )); + } + continue; + }; + match frame.kind { + MuxFrameKind::WindowUpdate => { + let credits = frame.window_credits()? as usize; + if credits == 0 || credits > state.max_credits { + drop(state); + connection.try_send_priority_command(WriterCommand::priority( + MuxFrame::empty(MuxFrameKind::Reset, frame.stream_id), + None, + )); + connection.remove_stream( + frame.stream_id, + "invalid response mux window update", + true, + ); + continue; + } + state.replenish_credits(credits); + if metrics_enabled { + window_updates.inc(); + } + } + MuxFrameKind::Stop => state.context.stop(), + MuxFrameKind::Kill | MuxFrameKind::Reset => { + drop(state); + connection.remove_stream( + frame.stream_id, + "frontend closed response mux stream", + true, + ); + } + _ => { + drop(state); + connection.try_send_priority_command(WriterCommand::priority( + MuxFrame::empty(MuxFrameKind::Reset, frame.stream_id), + None, + )); + connection.remove_stream( + frame.stream_id, + "invalid frontend response mux frame", + true, + ); + } + } + } + Ok(()) + } + .await; + + if let Some(previous) = reported_data_segments + && let Some(current) = + crate::pipeline::network::tcp::tcp_data_segments_out(reader.get_ref().as_ref()) + { + response_mux::DATA_SEGMENTS_TOTAL + .with_label_values(&["mux", "worker"]) + .inc_by(current.saturating_sub(previous)); + } + + if let Some(connection) = weak.upgrade() { + connection.fail( + &result + .err() + .map(|err| err.to_string()) + .unwrap_or_else(|| "reader stopped".to_string()), + ); + } + } +} + +struct HostPool { + address: String, + frontend_server_id: Uuid, + version: u8, + connections: RwLock>>, + connect_lock: tokio::sync::Mutex<()>, + next_connection_id: AtomicU64, + warming: AtomicBool, + maintenance_started: AtomicBool, + last_used: parking_lot::Mutex, + cancel: CancellationToken, + config: PoolConfig, +} + +impl HostPool { + fn new( + address: String, + frontend_server_id: Uuid, + version: u8, + cancel: CancellationToken, + config: PoolConfig, + ) -> Arc { + Arc::new(Self { + address, + frontend_server_id, + version, + connections: RwLock::new(Vec::new()), + connect_lock: tokio::sync::Mutex::new(()), + next_connection_id: AtomicU64::new(1), + warming: AtomicBool::new(false), + maintenance_started: AtomicBool::new(false), + last_used: parking_lot::Mutex::new(Instant::now()), + cancel, + config, + }) + } + + fn healthy_connections(&self) -> Vec> { + self.connections + .read() + .iter() + .filter(|connection| connection.is_healthy()) + .cloned() + .collect() + } + + async fn ensure_first(&self) -> Result> { + let _guard = self.connect_lock.lock().await; + if let Some(connection) = self.healthy_connections().first().cloned() { + return Ok(connection); + } + self.connect_new().await + } + + async fn connect_additional(&self) -> Result> { + let _guard = self.connect_lock.lock().await; + if self.healthy_connections().len() >= self.config.pool_size { + return self + .healthy_connections() + .first() + .cloned() + .ok_or_else(|| anyhow!("response mux host pool has no healthy connection")); + } + self.connect_new().await + } + + async fn connect_new(&self) -> Result> { + let replacing_failed_connection = self + .connections + .read() + .iter() + .any(|connection| !connection.is_healthy()); + let id = self.next_connection_id.fetch_add(1, Ordering::Relaxed); + let connection = MuxConnection::connect( + id, + &self.address, + self.frontend_server_id, + self.version, + self.cancel.clone(), + self.config, + ) + .await?; + let mut connections = self.connections.write(); + connections.retain(|candidate| candidate.is_healthy()); + connections.push(connection.clone()); + if replacing_failed_connection { + response_mux::RECONNECTS_TOTAL + .with_label_values(&["worker"]) + .inc(); + } + Ok(connection) + } + + fn start_maintenance(self: &Arc) { + if self.maintenance_started.swap(true, Ordering::AcqRel) { + return; + } + let host = Arc::clone(self); + tokio::spawn(async move { + let mut interval = tokio::time::interval(Duration::from_secs(1)); + loop { + interval.tick().await; + if host.cancel.is_cancelled() { + break; + } + if host.healthy_connections().len() < host.config.pool_size { + host.warm(); + } + } + }); + } + + fn warm(self: &Arc) { + if self.warming.swap(true, Ordering::AcqRel) { + return; + } + let host = Arc::clone(self); + tokio::spawn(async move { + while host.healthy_connections().len() < host.config.pool_size + && !host.cancel.is_cancelled() + { + if let Err(err) = host.connect_additional().await { + tracing::warn!(address = %host.address, %err, "failed to warm response mux pool"); + break; + } + } + host.warming.store(false, Ordering::Release); + }); + } + + async fn connection(self: &Arc) -> Result> { + *self.last_used.lock() = Instant::now(); + let mut healthy = self.healthy_connections(); + if healthy.is_empty() { + healthy.push(self.ensure_first().await?); + } + self.start_maintenance(); + self.warm(); + let index = healthy + .iter() + .enumerate() + .min_by_key(|(_, connection)| { + ( + connection.active_streams.load(Ordering::Acquire), + connection.queued_bytes.load(Ordering::Acquire), + ) + }) + .map(|(index, _)| index) + .expect("healthy response mux connection list is non-empty"); + Ok(healthy.swap_remove(index)) + } + + fn is_idle(&self) -> bool { + self.healthy_connections() + .iter() + .all(|connection| connection.active_streams.load(Ordering::Acquire) == 0) + && self.last_used.lock().elapsed() >= self.config.idle_ttl + } +} + +pub struct ResponseMuxClientPool { + hosts: DashMap>, + cancel: CancellationToken, + enabled: bool, + config: PoolConfig, +} + +#[derive(Clone, Debug, PartialEq, Eq, Hash)] +struct HostKey { + address: String, + frontend_server_id: Uuid, + version: u8, +} + +impl ResponseMuxClientPool { + pub fn new(cancel: CancellationToken, runtime_config: ResponseMuxConfig) -> Arc { + response_mux::CONFIGURED_BATCH_INTERVAL_MS + .with_label_values(&["worker"]) + .set(runtime_config.batch_interval.as_millis() as i64); + let pool = Arc::new(Self { + hosts: DashMap::new(), + cancel, + enabled: runtime_config.enabled, + config: PoolConfig::from_runtime(runtime_config), + }); + Self::start_cleanup(&pool); + pool + } + + #[cfg(test)] + fn new_for_test_with_connection_window( + cancel: CancellationToken, + pool_size: usize, + writer_queue: usize, + initial_window: usize, + connection_window: usize, + idle_ttl: Duration, + connect_timeout: Duration, + ) -> Arc { + let pool = Arc::new(Self { + hosts: DashMap::new(), + cancel, + enabled: true, + config: PoolConfig { + pool_size: pool_size.max(1), + writer_queue: writer_queue.max(1), + stream_writer_queue: RESPONSE_MUX_STREAM_WRITER_QUEUE, + initial_window: initial_window.max(1), + connection_window: connection_window.max(1), + batch_interval: Duration::ZERO, + batch_max_bytes: 65_536, + batch_max_frames: 64, + packet_metrics: false, + idle_ttl, + connect_timeout, + }, + }); + Self::start_cleanup(&pool); + pool + } + + fn start_cleanup(pool: &Arc) { + let weak = Arc::downgrade(pool); + tokio::spawn(async move { + let mut interval = tokio::time::interval(Duration::from_secs(30)); + loop { + interval.tick().await; + let Some(pool) = weak.upgrade() else { + break; + }; + if pool.cancel.is_cancelled() { + break; + } + pool.hosts.retain(|_, host| { + let retain = !host.is_idle(); + if !retain { + host.cancel.cancel(); + } + retain + }); + } + }); + } + + pub async fn create_response_stream( + self: &Arc, + context: Arc, + info: ConnectionInfo, + ) -> Result { + if !self.enabled { + anyhow::bail!("response mux connection info received while mux mode is disabled"); + } + let info = ResponseMuxConnectionInfo::try_from(info) + .context("tcp-response-mux-connection-info-error")?; + if info.version != RESPONSE_MUX_VERSION { + anyhow::bail!( + "unsupported response mux version {}; expected {}", + info.version, + RESPONSE_MUX_VERSION + ); + } + if info.context != context.id() { + return Err(anyhow!( + "response mux context mismatch: expected {}, got {}", + context.id(), + info.context + )); + } + let stream_id = info.stream_id; + let host_key = HostKey { + address: info.address.clone(), + frontend_server_id: info.frontend_server_id, + version: info.version, + }; + let host = self + .hosts + .entry(host_key) + .or_insert_with(|| { + HostPool::new( + info.address.clone(), + info.frontend_server_id, + info.version, + self.cancel.child_token(), + self.config, + ) + }) + .clone(); + let connection = host.connection().await?; + let state = Arc::new(WorkerStreamState { + context, + credits: Arc::new(Semaphore::new(self.config.initial_window)), + max_credits: self.config.initial_window, + writer_slots: Arc::new(Semaphore::new(self.config.stream_writer_queue)), + writer: Mutex::new(StreamWriterState::default()), + closed: AtomicBool::new(false), + close_token: CancellationToken::new(), + }); + match connection.streams.entry(stream_id) { + Entry::Vacant(entry) => { + entry.insert(state.clone()); + } + Entry::Occupied(_) => { + return Err(anyhow!("duplicate response mux stream id {stream_id}")); + } + } + connection.active_streams.fetch_add(1, Ordering::AcqRel); + response_mux::ACTIVE_STREAMS + .with_label_values(&["worker"]) + .inc(); + let sender = StreamSender::multiplexed(Arc::new(MuxResponseStreamSender { + stream_id, + connection, + state, + finished: AtomicBool::new(false), + })); + Ok(sender) + } + + #[cfg(test)] + fn host_for_address(&self, address: &str) -> Option> { + self.hosts + .iter() + .find(|entry| entry.value().address == address) + .map(|entry| entry.value().clone()) + } + + #[cfg(test)] + pub(crate) fn healthy_connection_count(&self, address: &str) -> usize { + self.host_for_address(address) + .map(|host| host.healthy_connections().len()) + .unwrap_or(0) + } + + #[cfg(test)] + pub(crate) fn stream_connection_id(&self, address: &str, stream_id: Uuid) -> Option { + self.host_for_address(address).and_then(|host| { + host.connections + .read() + .iter() + .find(|connection| connection.streams.contains_key(&stream_id)) + .map(|connection| connection.id) + }) + } + + #[cfg(test)] + pub(crate) fn fail_connection(&self, address: &str, connection_id: u64) { + if let Some(host) = self.host_for_address(address) + && let Some(connection) = host + .connections + .read() + .iter() + .find(|connection| connection.id == connection_id) + { + connection.fail("test-injected physical connection failure"); + } + } +} + +struct MuxResponseStreamSender { + stream_id: Uuid, + connection: Arc, + state: Arc, + finished: AtomicBool, +} + +impl MuxResponseStreamSender { + async fn acquire_writer_permit( + &self, + state: &WorkerStreamState, + ) -> Result { + let slots = state.writer_slots.clone(); + match slots.try_acquire_owned() { + Ok(permit) => Ok(permit), + Err(tokio::sync::TryAcquireError::NoPermits) => { + let wait_start = Instant::now(); + let permit = state + .writer_slots + .clone() + .acquire_owned() + .await + .map_err(|_| anyhow!("response mux stream closed during writer admission"))?; + if per_frame_metrics_enabled() { + response_mux::WRITER_ADMISSION_STALL_SECONDS + .observe(wait_start.elapsed().as_secs_f64()); + } + Ok(permit) + } + Err(tokio::sync::TryAcquireError::Closed) => { + anyhow::bail!("response mux stream closed during writer admission") + } + } + } + + async fn enqueue_ordered( + &self, + frame: MuxFrame, + written: Option>>, + ) -> Result<()> { + self.enqueue_ordered_on(&self.connection, &self.state, frame, written) + .await + } + + async fn enqueue_ordered_on( + &self, + connection: &Arc, + state: &Arc, + frame: MuxFrame, + written: Option>>, + ) -> Result<()> { + if !connection.is_healthy() { + anyhow::bail!("response mux connection is unhealthy"); + } + let writer_permit = self.acquire_writer_permit(state).await?; + let required = frame.encoded_len().min(connection.max_queued_bytes) as u32; + let queued_byte_permit = match connection + .queued_byte_slots + .clone() + .try_acquire_many_owned(required) + { + Ok(permit) => permit, + Err(tokio::sync::TryAcquireError::NoPermits) => { + let wait_start = Instant::now(); + let slots = connection.queued_byte_slots.clone(); + let permit = tokio::select! { + _ = state.close_token.cancelled() => { + anyhow::bail!("response mux stream closed during byte-queue admission") + } + permit = slots.acquire_many_owned(required) => permit.map_err(|_| { + anyhow!("response mux connection closed during byte-queue admission") + })?, + }; + response_mux::QUEUED_BYTE_ADMISSION_STALL_SECONDS + .observe(wait_start.elapsed().as_secs_f64()); + permit + } + Err(tokio::sync::TryAcquireError::Closed) => { + anyhow::bail!("response mux connection closed during byte-queue admission") + } + }; + connection.enqueue_stream_command( + self.stream_id, + state, + WriterCommand::new(frame, written) + .with_writer_permit(writer_permit) + .with_queued_byte_permit(queued_byte_permit), + ) + } + + async fn enqueue_priority_and_wait(&self, frame: MuxFrame) -> Result<()> { + let (written_tx, written_rx) = oneshot::channel(); + self.connection + .send_priority_command(WriterCommand::priority(frame, Some(written_tx))) + .await?; + written_rx + .await + .map_err(|_| anyhow!("response mux priority writer dropped acknowledgement"))? + .map_err(anyhow::Error::msg) + } + + fn remove_stream(&self) { + self.connection + .remove_stream(self.stream_id, "response mux stream completed", false); + } +} + +#[async_trait::async_trait] +impl MultiplexedStreamSender for MuxResponseStreamSender { + async fn send_data(&self, data: bytes::Bytes) -> Result<()> { + let encoded_len = super::MUX_HEADER_LEN.saturating_add(data.len()); + let required = encoded_len.min(self.state.max_credits) as u32; + let permit = match self.state.credits.clone().try_acquire_many_owned(required) { + Ok(permit) => permit, + Err(tokio::sync::TryAcquireError::NoPermits) => { + let wait_start = Instant::now(); + let permit = self + .state + .credits + .clone() + .acquire_many_owned(required) + .await + .map_err(|_| anyhow!("response mux stream closed while waiting for credits"))?; + if per_frame_metrics_enabled() { + response_mux::FLOW_CONTROL_STALL_SECONDS + .observe(wait_start.elapsed().as_secs_f64()); + } + permit + } + Err(tokio::sync::TryAcquireError::Closed) => { + anyhow::bail!("response mux stream closed while waiting for credits") + } + }; + let result = self + .enqueue_ordered( + MuxFrame::new(MuxFrameKind::Data, self.stream_id, data), + None, + ) + .await; + if result.is_ok() { + permit.forget(); + } + result + } + + async fn send_prologue(&self, error: Option) -> Result<(), String> { + let terminal_error = error.is_some(); + let payload = + serde_json::to_vec(&ResponseStreamPrologue { error }).map_err(|err| err.to_string())?; + let frame = MuxFrame::new(MuxFrameKind::Prologue, self.stream_id, payload.into()); + let result = self.enqueue_priority_and_wait(frame).await; + if result.is_ok() && terminal_error { + self.finished.store(true, Ordering::Release); + self.remove_stream(); + } + result.map_err(|err| err.to_string()) + } + + async fn finish(&self) -> Result<()> { + if self.finished.swap(true, Ordering::AcqRel) { + return Ok(()); + } + let (written_tx, written_rx) = oneshot::channel(); + let end = MuxFrame::empty(MuxFrameKind::End, self.stream_id); + if let Err(err) = self.enqueue_ordered(end, Some(written_tx)).await { + self.remove_stream(); + return Err(err).context("response mux writer stopped before end"); + } + let result = written_rx + .await + .map_err(|_| anyhow!("response mux writer dropped end acknowledgement"))? + .map_err(anyhow::Error::msg); + self.remove_stream(); + result + } +} + +impl Drop for MuxResponseStreamSender { + fn drop(&mut self) { + if self.finished.swap(true, Ordering::AcqRel) { + return; + } + response_mux::RESETS_TOTAL + .with_label_values(&["worker", "publisher_drop"]) + .inc(); + self.connection + .try_send_priority_command(WriterCommand::priority( + MuxFrame::new( + MuxFrameKind::Reset, + self.stream_id, + bytes::Bytes::from_static(b"response sender dropped before finish"), + ), + None, + )); + self.remove_stream(); + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::pipeline::network::tcp::mux::ResponseMuxConfig; + use crate::{ + engine::AsyncEngineContextProvider, + pipeline::{ + Context, + network::{ResponseService, StreamOptions}, + }, + }; + use futures::{StreamExt as _, stream::FuturesUnordered}; + + const TEST_STREAM_WINDOW: usize = 64; + const TEST_STREAM_WINDOW_UPDATE: u32 = 32; + const TEST_CONNECTION_WINDOW: usize = 256; + + fn worker_stream_state(initial_credits: usize, writer_slots: usize) -> Arc { + let context = Context::new(()); + Arc::new(WorkerStreamState { + context: context.context(), + credits: Arc::new(Semaphore::new(initial_credits)), + max_credits: TEST_STREAM_WINDOW, + writer_slots: Arc::new(Semaphore::new(writer_slots)), + writer: Mutex::new(StreamWriterState::default()), + closed: AtomicBool::new(false), + close_token: CancellationToken::new(), + }) + } + + #[test] + fn credit_replenishment_never_exceeds_the_fixed_maximum() { + let state = worker_stream_state(TEST_STREAM_WINDOW, RESPONSE_MUX_STREAM_WRITER_QUEUE); + state + .credits + .clone() + .try_acquire_many_owned((TEST_STREAM_WINDOW - 16) as u32) + .expect("initial credits should be available") + .forget(); + + assert_eq!(state.credits.available_permits(), 16); + assert_eq!( + state.replenish_credits(TEST_STREAM_WINDOW_UPDATE as usize), + TEST_STREAM_WINDOW_UPDATE as usize + ); + assert_eq!( + state.credits.available_permits(), + 16 + TEST_STREAM_WINDOW_UPDATE as usize + ); + assert_eq!( + state.replenish_credits(TEST_STREAM_WINDOW), + TEST_STREAM_WINDOW - 16 - TEST_STREAM_WINDOW_UPDATE as usize + ); + assert_eq!(state.credits.available_permits(), TEST_STREAM_WINDOW); + assert_eq!( + state.replenish_credits(TEST_STREAM_WINDOW_UPDATE as usize), + 0 + ); + assert_eq!(state.credits.available_permits(), TEST_STREAM_WINDOW); + } + + #[test] + fn cumulative_connection_ack_replenishes_credits_without_exceeding_the_window() { + let (priority_tx, _priority_rx) = mpsc::channel(1); + let (ready_tx, _ready_rx) = mpsc::unbounded_channel(); + let remaining = 16; + let consumed = TEST_CONNECTION_WINDOW - remaining; + let connection = MuxConnection { + id: 1, + cancel: CancellationToken::new(), + priority_tx, + ready_tx, + streams: DashMap::new(), + healthy: AtomicBool::new(true), + active_streams: AtomicUsize::new(0), + queued_frames: AtomicUsize::new(0), + queued_bytes: AtomicUsize::new(0), + max_queued_bytes: TEST_CONNECTION_WINDOW, + queued_byte_slots: Arc::new(Semaphore::new(TEST_CONNECTION_WINDOW)), + connection_credits: Arc::new(Semaphore::new(TEST_CONNECTION_WINDOW)), + max_connection_credits: TEST_CONNECTION_WINDOW, + sent_data_bytes: AtomicU64::new(consumed as u64), + acknowledged_data_bytes: AtomicU64::new(0), + batch_interval: Duration::ZERO, + batch_max_bytes: 65_536, + batch_max_frames: 64, + }; + connection + .connection_credits + .clone() + .try_acquire_many_owned(consumed as u32) + .unwrap() + .forget(); + + assert_eq!(connection.connection_credits.available_permits(), remaining); + connection.acknowledge_connection_credits(128).unwrap(); + assert_eq!( + connection.connection_credits.available_permits(), + remaining + 128 + ); + connection + .acknowledge_connection_credits(consumed as u64) + .unwrap(); + assert_eq!( + connection.connection_credits.available_permits(), + TEST_CONNECTION_WINDOW + ); + connection + .acknowledge_connection_credits(consumed as u64) + .unwrap(); + assert_eq!( + connection.connection_credits.available_permits(), + TEST_CONNECTION_WINDOW + ); + assert!(connection.acknowledge_connection_credits(127).is_err()); + assert!( + connection + .acknowledge_connection_credits(consumed as u64 + 1) + .is_err() + ); + } + + #[tokio::test] + async fn closing_a_stream_wakes_credit_and_writer_admission_waiters() { + let state = worker_stream_state(0, 0); + let (written_tx, written_rx) = oneshot::channel(); + state.writer.lock().pending.push_back(WriterCommand::new( + MuxFrame::empty(MuxFrameKind::End, Uuid::new_v4()), + Some(written_tx), + )); + + let (priority_tx, _priority_rx) = mpsc::channel(1); + let (ready_tx, _ready_rx) = mpsc::unbounded_channel(); + let connection = MuxConnection { + id: 1, + cancel: CancellationToken::new(), + priority_tx, + ready_tx, + streams: DashMap::new(), + healthy: AtomicBool::new(true), + active_streams: AtomicUsize::new(0), + queued_frames: AtomicUsize::new(1), + queued_bytes: AtomicUsize::new(super::super::MUX_HEADER_LEN), + max_queued_bytes: TEST_CONNECTION_WINDOW, + queued_byte_slots: Arc::new(Semaphore::new(TEST_CONNECTION_WINDOW)), + connection_credits: Arc::new(Semaphore::new(TEST_CONNECTION_WINDOW)), + max_connection_credits: TEST_CONNECTION_WINDOW, + sent_data_bytes: AtomicU64::new(0), + acknowledged_data_bytes: AtomicU64::new(0), + batch_interval: Duration::ZERO, + batch_max_bytes: 65_536, + batch_max_frames: 64, + }; + + let credit_waiter = tokio::spawn({ + let credits = state.credits.clone(); + async move { credits.acquire_owned().await } + }); + let writer_waiter = tokio::spawn({ + let writer_slots = state.writer_slots.clone(); + async move { writer_slots.acquire_owned().await } + }); + tokio::task::yield_now().await; + + connection.close_stream_state(&state, "test stream closed"); + + assert!(credit_waiter.await.unwrap().is_err()); + assert!(writer_waiter.await.unwrap().is_err()); + assert_eq!(connection.queued_frames.load(Ordering::Acquire), 0); + assert_eq!(written_rx.await.unwrap().unwrap_err(), "test stream closed"); + assert_eq!(state.replenish_credits(1_024), 0); + assert_eq!(state.credits.available_permits(), 0); + } + + #[tokio::test] + async fn end_bypasses_exhausted_connection_credit_at_batch_byte_limit() { + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let address = listener.local_addr().unwrap(); + let client = TcpStream::connect(address).await.unwrap(); + let (server, _) = listener.accept().await.unwrap(); + let (_, write_half) = client.into_split(); + let (server_read, _) = server.into_split(); + + let stream_id = Uuid::new_v4(); + let (end_written_tx, end_written_rx) = oneshot::channel(); + let state = worker_stream_state(TEST_STREAM_WINDOW, RESPONSE_MUX_STREAM_WRITER_QUEUE); + { + let mut writer = state.writer.lock(); + writer.pending.push_back(WriterCommand::new( + MuxFrame::new( + MuxFrameKind::Data, + stream_id, + bytes::Bytes::from(vec![b'x'; 40]), + ), + None, + )); + writer.pending.push_back(WriterCommand::new( + MuxFrame::empty(MuxFrameKind::End, stream_id), + Some(end_written_tx), + )); + writer.scheduled = true; + } + + let (priority_tx, priority_rx) = mpsc::channel(1); + let (ready_tx, ready_rx) = mpsc::unbounded_channel(); + let cancel = CancellationToken::new(); + let connection = Arc::new(MuxConnection { + id: 1, + cancel: cancel.clone(), + priority_tx, + ready_tx: ready_tx.clone(), + streams: DashMap::new(), + healthy: AtomicBool::new(true), + active_streams: AtomicUsize::new(1), + queued_frames: AtomicUsize::new(2), + queued_bytes: AtomicUsize::new(88), + max_queued_bytes: TEST_CONNECTION_WINDOW, + queued_byte_slots: Arc::new(Semaphore::new(TEST_CONNECTION_WINDOW)), + connection_credits: Arc::new(Semaphore::new(64)), + max_connection_credits: 64, + sent_data_bytes: AtomicU64::new(0), + acknowledged_data_bytes: AtomicU64::new(0), + batch_interval: Duration::ZERO, + batch_max_bytes: 64, + batch_max_frames: 64, + }); + connection.streams.insert(stream_id, state); + ready_tx.send(stream_id).unwrap(); + + let writer_task = tokio::spawn(MuxConnection::writer_task( + Arc::downgrade(&connection), + write_half, + priority_rx, + ready_rx, + cancel, + )); + let reader_task = tokio::spawn(async move { + let mut reader = FramedRead::new(server_read, MuxCodec::default()); + let data = reader.next().await.unwrap().unwrap(); + let end = reader.next().await.unwrap().unwrap(); + (data.kind, end.kind) + }); + + tokio::time::timeout(Duration::from_millis(200), end_written_rx) + .await + .expect("End waited for exhausted Data credits") + .unwrap() + .unwrap(); + assert_eq!( + reader_task.await.unwrap(), + (MuxFrameKind::Data, MuxFrameKind::End) + ); + writer_task.abort(); + } + + fn integration_config() -> ResponseMuxConfig { + ResponseMuxConfig { + enabled: true, + packet_metrics: false, + batch_interval: Duration::from_millis(5), + batch_max_bytes: 65_536, + batch_max_frames: 64, + stream_window_bytes: 262_144, + connection_window_bytes: 262_144, + } + } + + async fn open_mux_stream( + server: Arc, + pool: Arc, + ) -> ( + Uuid, + Context<()>, + StreamSender, + crate::pipeline::network::StreamReceiver, + ) { + let context = Context::new(()); + let pending = server + .register( + StreamOptions::builder() + .context(context.context()) + .enable_request_stream(false) + .enable_response_stream(true) + .send_buffer_count(8) + .build() + .unwrap(), + ) + .await + .recv_stream + .unwrap(); + let (info, provider) = pending.into_parts(); + let stream_id = ResponseMuxConnectionInfo::try_from(info.clone()) + .unwrap() + .stream_id; + let mut sender = pool + .create_response_stream(context.context(), info) + .await + .unwrap(); + sender.send_prologue(None).await.unwrap(); + let receiver = provider.await.unwrap().unwrap(); + (stream_id, context, sender, receiver) + } + + async fn mux_address( + server: Arc, + ) -> String { + let probe = server + .register( + StreamOptions::builder() + .context(Context::new(()).context()) + .enable_request_stream(false) + .enable_response_stream(true) + .build() + .unwrap(), + ) + .await; + let info = probe.recv_stream.as_ref().unwrap().connection_info.clone(); + drop(probe); + ResponseMuxConnectionInfo::try_from(info).unwrap().address + } + + #[tokio::test] + async fn disabled_worker_rejects_mux_connection_info() { + let mut config = integration_config(); + config.enabled = false; + let cancel = CancellationToken::new(); + let pool = ResponseMuxClientPool::new(cancel.clone(), config); + let context = Context::new(()); + let info = ResponseMuxConnectionInfo { + address: "127.0.0.1:1".to_string(), + frontend_server_id: Uuid::new_v4(), + stream_id: Uuid::new_v4(), + context: context.context().id().to_string(), + version: RESPONSE_MUX_VERSION, + }; + + let error = pool + .create_response_stream(context.context(), info.into()) + .await + .err() + .expect("disabled mux mode must reject mux connection info"); + assert!(error.to_string().contains("mux mode is disabled")); + cancel.cancel(); + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] + async fn stream_credit_stall_does_not_block_another_stream() { + let config = ResponseMuxConfig { + enabled: true, + packet_metrics: false, + batch_interval: Duration::ZERO, + batch_max_bytes: 65_536, + batch_max_frames: 64, + stream_window_bytes: 64, + connection_window_bytes: 512, + }; + let server = crate::pipeline::network::tcp::server::TcpStreamServer::new_mux_for_test( + crate::pipeline::network::tcp::server::ServerOptions::default(), + config, + ) + .await + .unwrap(); + let cancel = CancellationToken::new(); + let pool = ResponseMuxClientPool::new_for_test_with_connection_window( + cancel.clone(), + 1, + 64, + 64, + 512, + Duration::from_secs(60), + Duration::from_secs(5), + ); + let (_, _context_a, sender_a, mut receiver_a) = + open_mux_stream(server.clone(), pool.clone()).await; + let (_, _context_b, sender_b, mut receiver_b) = open_mux_stream(server, pool.clone()).await; + + let full_window_payload = bytes::Bytes::from(vec![b'a'; 40]); + sender_a.send(full_window_payload.clone()).await.unwrap(); + let blocked_send = sender_a.send(full_window_payload.clone()); + tokio::pin!(blocked_send); + assert!( + tokio::time::timeout(Duration::from_millis(20), &mut blocked_send) + .await + .is_err(), + "second frame should wait for stream-local credits" + ); + + sender_b + .send(bytes::Bytes::from_static(b"healthy")) + .await + .unwrap(); + sender_b.finish().await.unwrap(); + assert_eq!( + receiver_b.recv().await.unwrap(), + bytes::Bytes::from_static(b"healthy") + ); + assert!(receiver_b.recv().await.is_none()); + + assert_eq!(receiver_a.recv().await.unwrap(), full_window_payload); + tokio::time::timeout(Duration::from_secs(1), &mut blocked_send) + .await + .expect("stream credit update did not unblock producer") + .unwrap(); + sender_a.finish().await.unwrap(); + assert_eq!(receiver_a.recv().await.unwrap(), full_window_payload); + assert!(receiver_a.recv().await.is_none()); + cancel.cancel(); + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 4)] + async fn connection_failure_is_scoped_and_pool_reconnects() { + let config = integration_config(); + let server = crate::pipeline::network::tcp::server::TcpStreamServer::new_mux_for_test( + crate::pipeline::network::tcp::server::ServerOptions::default(), + config, + ) + .await + .unwrap(); + let address = mux_address(server.clone()).await; + let cancel = CancellationToken::new(); + let pool = ResponseMuxClientPool::new(cancel.clone(), config); + + let mut streams = vec![open_mux_stream(server.clone(), pool.clone()).await]; + tokio::time::timeout(Duration::from_secs(5), async { + while pool.healthy_connection_count(&address) != RESPONSE_MUX_POOL_SIZE { + tokio::time::sleep(Duration::from_millis(10)).await; + } + }) + .await + .expect("response mux pool did not warm to four connections"); + for _ in 1..8 { + streams.push(open_mux_stream(server.clone(), pool.clone()).await); + } + + let target_connection = pool + .stream_connection_id(&address, streams[0].0) + .expect("stream must be assigned to a connection"); + let assignments = streams + .iter() + .map(|stream| { + pool.stream_connection_id(&address, stream.0) + .expect("stream must be assigned to a connection") + }) + .collect::>(); + assert!(assignments.iter().any(|id| *id != target_connection)); + + pool.fail_connection(&address, target_connection); + tokio::time::timeout(Duration::from_secs(1), async { + loop { + let correctly_scoped = + streams + .iter() + .zip(&assignments) + .all(|((_, context, _, _), connection_id)| { + context.context().is_killed() == (*connection_id == target_connection) + }); + if correctly_scoped { + break; + } + tokio::task::yield_now().await; + } + }) + .await + .expect("connection failure did not remain scoped to assigned streams"); + + let (_, _, replacement_sender, mut replacement_receiver) = + open_mux_stream(server, pool.clone()).await; + replacement_sender + .send(bytes::Bytes::from_static(b"replacement")) + .await + .unwrap(); + replacement_sender.finish().await.unwrap(); + assert_eq!( + replacement_receiver.recv().await.unwrap(), + bytes::Bytes::from_static(b"replacement") + ); + tokio::time::timeout(Duration::from_secs(5), async { + while pool.healthy_connection_count(&address) != RESPONSE_MUX_POOL_SIZE { + tokio::time::sleep(Duration::from_millis(10)).await; + } + }) + .await + .expect("response mux pool did not reconnect to four connections"); + cancel.cancel(); + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 4)] + async fn one_thousand_logical_streams_share_exactly_four_connections() { + let config = integration_config(); + let server = crate::pipeline::network::tcp::server::TcpStreamServer::new_mux_for_test( + crate::pipeline::network::tcp::server::ServerOptions::default(), + config, + ) + .await + .unwrap(); + let cancel = CancellationToken::new(); + let pool = ResponseMuxClientPool::new(cancel.clone(), config); + + async fn round_trip( + server: Arc, + pool: Arc, + value: usize, + ) -> String { + let context = Context::new(()); + let pending = server + .register( + StreamOptions::builder() + .context(context.context()) + .enable_request_stream(false) + .enable_response_stream(true) + .send_buffer_count(8) + .build() + .unwrap(), + ) + .await + .recv_stream + .unwrap(); + let (info, provider) = pending.into_parts(); + let mut sender = pool + .create_response_stream(context.context(), info) + .await + .unwrap(); + sender.send_prologue(None).await.unwrap(); + let mut receiver = provider.await.unwrap().unwrap(); + let expected = format!("response-{value}"); + sender.send(expected.clone().into()).await.unwrap(); + sender.finish().await.unwrap(); + let actual = receiver.recv().await.unwrap(); + assert!(receiver.recv().await.is_none()); + String::from_utf8(actual.to_vec()).unwrap() + } + + assert_eq!( + round_trip(server.clone(), pool.clone(), 0).await, + "response-0" + ); + let address = mux_address(server.clone()).await; + tokio::time::timeout(Duration::from_secs(5), async { + while pool.healthy_connection_count(&address) != RESPONSE_MUX_POOL_SIZE { + tokio::time::sleep(Duration::from_millis(10)).await; + } + }) + .await + .expect("response mux pool did not warm to four connections"); + + let mut tasks = FuturesUnordered::new(); + for value in 1..1_000 { + tasks.push(round_trip(server.clone(), pool.clone(), value)); + } + let mut completed = 1; + while let Some(actual) = tasks.next().await { + assert!(actual.starts_with("response-")); + completed += 1; + } + assert_eq!(completed, 1_000); + assert_eq!( + pool.healthy_connection_count(&address), + RESPONSE_MUX_POOL_SIZE + ); + cancel.cancel(); + } +} diff --git a/lib/runtime/src/pipeline/network/tcp/server.rs b/lib/runtime/src/pipeline/network/tcp/server.rs index 34156ae1d68b..433106e7de02 100644 --- a/lib/runtime/src/pipeline/network/tcp/server.rs +++ b/lib/runtime/src/pipeline/network/tcp/server.rs @@ -1,11 +1,12 @@ // SPDX-FileCopyrightText: Copyright (c) 2024-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. // SPDX-License-Identifier: Apache-2.0 +use dashmap::DashMap; use socket2::{Domain, SockAddr, Socket, Type}; use std::{ collections::{HashMap, HashSet}, net::{IpAddr, SocketAddr, TcpListener}, - os::fd::{AsFd, FromRawFd}, + os::fd::{AsFd, AsRawFd, FromRawFd}, sync::Arc, time::Duration, }; @@ -29,18 +30,27 @@ use tokio::{ sync::{mpsc, oneshot}, time, }; -use tokio_util::codec::{FramedRead, FramedWrite}; +use tokio_util::{ + codec::{FramedRead, FramedWrite}, + sync::CancellationToken, +}; use super::{ - CallHomeHandshake, ControlMessage, PendingConnections, RegisteredStream, StreamOptions, - StreamReceiver, StreamSender, TcpStreamConnectionInfo, TwoPartCodec, + CallHomeHandshake, ControlMessage, PendingConnections, RegisteredStream, + ResponseMuxConnectionInfo, StreamOptions, StreamReceiver, StreamSender, + TcpStreamConnectionInfo, TwoPartCodec, + mux::{ + ConnectionHandshake, MuxCodec, MuxFrame, MuxFrameKind, RESPONSE_MUX_CREDIT_UPDATE_BYTES, + RESPONSE_MUX_CREDIT_UPDATE_INTERVAL, RESPONSE_MUX_VERSION, RESPONSE_MUX_WRITER_QUEUE, + ResponseMuxConfig, initialize_response_mux_config, + }, }; use crate::discovery::EndpointInstanceId; use crate::engine::AsyncEngineContext; use crate::pipeline::{ PipelineError, network::{ - ResponseService, ResponseStreamPrologue, + ResponseService, ResponseStreamPrologue, StreamReceiverHooks, StreamRxItem, codec::{TwoPartMessage, TwoPartMessageType}, tcp::StreamType, }, @@ -90,7 +100,11 @@ impl ServerOptions { pub struct TcpStreamServer { local_ip: String, local_port: u16, + server_id: uuid::Uuid, + mux_config: ResponseMuxConfig, state: Arc>, + response_pending: Arc>, + response_active: Arc>, } // pub struct TcpStreamReceiver { @@ -116,6 +130,35 @@ struct RequestedRecvConnection { send_buffer_count: usize, } +struct RequestedMuxRecvConnection { + context: Arc, + connection: Mutex>>>, + send_buffer_count: usize, + registered_at: Instant, +} + +struct ActiveMuxResponseStream { + connection_id: uuid::Uuid, + context: Arc, + response_tx: mpsc::Sender, + control_tx: mpsc::Sender, + control_failed: CancellationToken, +} + +struct ResponseMuxSocket { + read_half: tokio::io::ReadHalf, + write_half: tokio::io::WriteHalf, + packet_socket: Option, +} + +impl Drop for ActiveMuxResponseStream { + fn drop(&mut self) { + crate::metrics::response_mux::ACTIVE_STREAMS + .with_label_values(&["frontend"]) + .dec(); + } +} + /// Build the per-stream data-plane mpsc channel that bridges the socket task /// and the engine producer/consumer. The capacity is driven by the /// registration options ([`StreamOptions::send_buffer_count`]) rather than a @@ -182,6 +225,16 @@ impl TcpStreamServer { pub async fn new_with_resolver( options: ServerOptions, resolver: R, + ) -> Result, PipelineError> { + let mux_config = initialize_response_mux_config() + .map_err(|err| PipelineError::Generic(err.to_string()))?; + Self::new_with_resolver_and_mux_config(options, resolver, mux_config).await + } + + async fn new_with_resolver_and_mux_config( + options: ServerOptions, + resolver: R, + mux_config: ResponseMuxConfig, ) -> Result, PipelineError> { let local_ip = match options.interface { Some(interface) => { @@ -225,22 +278,43 @@ impl TcpStreamServer { }; let state = Arc::new(Mutex::new(State::default())); - - let local_port = Self::start(local_ip.clone(), options.port, state.clone()) - .await - .map_err(|e| { - PipelineError::Generic(format!("Failed to start TcpStreamServer: {}", e)) - })?; + let response_pending = Arc::new(DashMap::new()); + let response_active = Arc::new(DashMap::new()); + let server_id = uuid::Uuid::new_v4(); + + let local_port = Self::start( + local_ip.clone(), + options.port, + state.clone(), + server_id, + mux_config, + response_pending.clone(), + response_active.clone(), + ) + .await + .map_err(|e| PipelineError::Generic(format!("Failed to start TcpStreamServer: {}", e)))?; tracing::debug!("tcp transport service on {local_ip}:{local_port}"); Ok(Arc::new(Self { local_ip, local_port, + server_id, + mux_config, state, + response_pending, + response_active, })) } + #[cfg(test)] + pub(crate) async fn new_mux_for_test( + options: ServerOptions, + mux_config: ResponseMuxConfig, + ) -> Result, PipelineError> { + Self::new_with_resolver_and_mux_config(options, DefaultIpResolver, mux_config).await + } + /// Associate one or both halves of a registration with a backend instance. /// /// `recv_subject` is the response-stream subject (always present on TCP); @@ -274,6 +348,9 @@ impl TcpStreamServer { "Cancelling subject immediately: instance already removed (tombstoned)" ); state.rx_subjects.remove(recv_subject); + if let Ok(stream_id) = uuid::Uuid::parse_str(recv_subject) { + self.cancel_mux_response_stream(stream_id, MuxFrameKind::Kill); + } if let Some(s) = send_subject { state.tx_subjects.remove(s); } @@ -298,6 +375,9 @@ impl TcpStreamServer { pub async fn cancel_recv_stream(&self, subject: &str) { let mut state = self.state.lock(); state.rx_subjects.remove(subject); + if let Ok(stream_id) = uuid::Uuid::parse_str(subject) { + self.cancel_mux_response_stream(stream_id, MuxFrameKind::Kill); + } if let Some(key) = state.subject_instance.remove(subject) && let Some(subjects) = state.instance_subjects.get_mut(&key) { @@ -344,6 +424,9 @@ impl TcpStreamServer { match kind { StreamType::Response => { state.rx_subjects.remove(subject); + if let Ok(stream_id) = uuid::Uuid::parse_str(subject) { + self.cancel_mux_response_stream(stream_id, MuxFrameKind::Kill); + } } StreamType::Request => { state.tx_subjects.remove(subject); @@ -361,7 +444,29 @@ impl TcpStreamServer { state.removed_instances.remove(id); } - async fn start(local_ip: String, local_port: u16, state: Arc>) -> Result { + fn cancel_mux_response_stream(&self, stream_id: uuid::Uuid, kind: MuxFrameKind) { + self.response_pending.remove(&stream_id); + if let Some((_, active)) = self.response_active.remove(&stream_id) { + active.context.kill(); + if active + .control_tx + .try_send(MuxFrame::empty(kind, stream_id)) + .is_err() + { + active.control_failed.cancel(); + } + } + } + + async fn start( + local_ip: String, + local_port: u16, + state: Arc>, + server_id: uuid::Uuid, + mux_config: ResponseMuxConfig, + response_pending: Arc>, + response_active: Arc>, + ) -> Result { let addr = format!("{}:{}", local_ip, local_port); let state_clone = state.clone(); let (ready_tx, ready_rx) = tokio::sync::oneshot::channel::>(); @@ -370,7 +475,15 @@ impl TcpStreamServer { if guard.handle.is_some() { panic!("TcpStreamServer already started"); } - guard.handle = Some(tokio::spawn(tcp_listener(addr, state_clone, ready_tx))); + guard.handle = Some(tokio::spawn(tcp_listener( + addr, + state_clone, + server_id, + mux_config, + response_pending, + response_active, + ready_tx, + ))); } let local_port = ready_rx.await??; Ok(local_port) @@ -494,30 +607,80 @@ impl ResponseService for TcpStreamServer { let recv_stream = if options.enable_response_stream { let (pending_recver_tx, pending_recver_rx) = oneshot::channel(); - let receiver_subject = uuid::Uuid::new_v4().to_string(); + let receiver_id = uuid::Uuid::new_v4(); + let receiver_subject = receiver_id.to_string(); let registry_subject = receiver_subject.clone(); - let connection_info = RequestedRecvConnection { - context: options.context.clone(), - connection: pending_recver_tx, - send_buffer_count: options.send_buffer_count, - }; - - let cleanup_subject = receiver_subject.clone(); - let cleanup_state = self.state.clone(); - let registered_stream = RegisteredStream::new( - TcpStreamConnectionInfo { - address: address.clone(), - subject: receiver_subject, - context: options.context.id().to_string(), - stream_type: StreamType::Response, - } - .into(), - pending_recver_rx, - ) - .with_cleanup(move || { - // Drop is sync; fire-and-forget the lock acquisition. - tokio::spawn(async move { + if self.mux_config.enabled { + self.response_pending.insert( + receiver_id, + RequestedMuxRecvConnection { + context: options.context.clone(), + connection: Mutex::new(Some(pending_recver_tx)), + send_buffer_count: options.send_buffer_count, + registered_at: Instant::now(), + }, + ); + + let cleanup_id = receiver_id; + let cleanup_subject = receiver_subject; + let cleanup_state = self.state.clone(); + let cleanup_pending = self.response_pending.clone(); + let cleanup_active = self.response_active.clone(); + let registered_stream = RegisteredStream::new( + ResponseMuxConnectionInfo { + address: address.clone(), + frontend_server_id: self.server_id, + stream_id: receiver_id, + context: options.context.id().to_string(), + version: RESPONSE_MUX_VERSION, + } + .into(), + pending_recver_rx, + ) + .with_cleanup(move || { + cleanup_pending.remove(&cleanup_id); + if let Some((_, active)) = cleanup_active.remove(&cleanup_id) { + active.context.kill(); + if active + .control_tx + .try_send(MuxFrame::empty(MuxFrameKind::Kill, cleanup_id)) + .is_err() + { + active.control_failed.cancel(); + } + } + let mut state = cleanup_state.lock(); + if let Some(key) = state.subject_instance.remove(&cleanup_subject) + && let Some(subjects) = state.instance_subjects.get_mut(&key) + { + subjects.remove(&(StreamType::Response, cleanup_subject.clone())); + if subjects.is_empty() { + state.instance_subjects.remove(&key); + } + } + }); + Some(registered_stream) + } else { + let connection_info = RequestedRecvConnection { + context: options.context.clone(), + connection: pending_recver_tx, + send_buffer_count: options.send_buffer_count, + }; + + let cleanup_subject = receiver_subject.clone(); + let cleanup_state = self.state.clone(); + let registered_stream = RegisteredStream::new( + TcpStreamConnectionInfo { + address: address.clone(), + subject: receiver_subject, + context: options.context.id().to_string(), + stream_type: StreamType::Response, + } + .into(), + pending_recver_rx, + ) + .with_cleanup(move || { let mut state = cleanup_state.lock(); state.rx_subjects.remove(&cleanup_subject); if let Some(key) = state.subject_instance.remove(&cleanup_subject) @@ -529,11 +692,11 @@ impl ResponseService for TcpStreamServer { } } }); - }); - self.insert_response_stream(registry_subject, connection_info); + self.insert_response_stream(registry_subject, connection_info); - Some(registered_stream) + Some(registered_stream) + } } else { None }; @@ -554,6 +717,10 @@ impl ResponseService for TcpStreamServer { async fn tcp_listener( addr: String, state: Arc>, + server_id: uuid::Uuid, + mux_config: ResponseMuxConfig, + response_pending: Arc>, + response_active: Arc>, read_tx: tokio::sync::oneshot::Sender>, ) -> Result<()> { let listener = tokio::net::TcpListener::bind(&addr) @@ -609,13 +776,35 @@ async fn tcp_listener( } } - tokio::spawn(handle_connection(stream, state.clone())); + tokio::spawn(handle_connection( + stream, + state.clone(), + server_id, + mux_config, + response_pending.clone(), + response_active.clone(), + )); } // #[instrument(level = "trace"), skip(state)] // todo - clone before spawn and trace process_stream - async fn handle_connection(stream: tokio::net::TcpStream, state: Arc>) { - let result = process_stream(stream, state).await; + async fn handle_connection( + stream: tokio::net::TcpStream, + state: Arc>, + server_id: uuid::Uuid, + mux_config: ResponseMuxConfig, + response_pending: Arc>, + response_active: Arc>, + ) { + let result = process_stream( + stream, + state, + server_id, + mux_config, + response_pending, + response_active, + ) + .await; match result { Ok(_) => tracing::trace!("successfully processed tcp connection"), Err(e) => { @@ -628,13 +817,23 @@ async fn tcp_listener( /// This method is responsible for the internal tcp stream handshake /// The handshake will specialize the stream as a request/sender or response/receiver stream - async fn process_stream(stream: tokio::net::TcpStream, state: Arc>) -> Result<()> { + async fn process_stream( + stream: tokio::net::TcpStream, + state: Arc>, + server_id: uuid::Uuid, + mux_config: ResponseMuxConfig, + response_pending: Arc>, + response_active: Arc>, + ) -> Result<()> { + let packet_socket = (mux_config.enabled && mux_config.packet_metrics) + .then(|| stream.as_fd().try_clone_to_owned().ok()) + .flatten(); // split the socket in to a reader and writer let (read_half, write_half) = tokio::io::split(stream); // attach the codec to the reader and writer to get framed readers and writers let mut framed_reader = FramedRead::new(read_half, TwoPartCodec::default()); - let framed_writer = FramedWrite::new(write_half, TwoPartCodec::default()); + let mut framed_writer = FramedWrite::new(write_half, TwoPartCodec::default()); // the internal tcp [`CallHomeHandshake`] connects the socket to the requester // here we await this first message as a raw bytes two part message @@ -645,17 +844,58 @@ async fn tcp_listener( // we await on the raw bytes which should come in as a header only message // todo - improve error handling - check for no data - let handshake: CallHomeHandshake = match first_message.header() { - Some(header) => serde_json::from_slice(header).map_err(|e| { - error!( - "Failed to deserialize the first message as a valid `CallHomeHandshake`: {e}", - ) - })?, + let header = match first_message.header() { + Some(header) => header, None => { return Err(error!("Expected ControlMessage, got DataMessage")); } }; + if let Ok(ConnectionHandshake::ResponseMux { + version, + frontend_server_id, + connection_id, + }) = serde_json::from_slice::(header) + { + if !mux_config.enabled { + anyhow::bail!("response mux handshake received while mux mode is disabled"); + } + if version != RESPONSE_MUX_VERSION { + anyhow::bail!( + "unsupported response mux version {version}; expected {RESPONSE_MUX_VERSION}" + ); + } + if frontend_server_id != server_id { + anyhow::bail!( + "response mux frontend UUID mismatch: got {frontend_server_id}, expected {server_id}" + ); + } + if connection_id.is_nil() { + anyhow::bail!("response mux physical connection UUID must not be nil"); + } + framed_writer + .send(MuxFrame::connection_ack(0).into_two_part()) + .await + .context("failed to send response mux connection ack")?; + return process_response_mux( + connection_id, + mux_config, + state, + response_pending, + response_active, + ResponseMuxSocket { + read_half: framed_reader.into_inner(), + write_half: framed_writer.into_inner(), + packet_socket, + }, + ) + .await; + } + + let handshake: CallHomeHandshake = serde_json::from_slice(header).map_err(|e| { + error!("Failed to deserialize the first message as a valid TCP handshake: {e}") + })?; + // branch here to handle sender stream or receiver stream match handshake.stream_type { StreamType::Request => { @@ -668,6 +908,349 @@ async fn tcp_listener( } } + async fn process_response_mux( + connection_id: uuid::Uuid, + mux_config: ResponseMuxConfig, + state: Arc>, + response_pending: Arc>, + response_active: Arc>, + socket: ResponseMuxSocket, + ) -> Result<()> { + let ResponseMuxSocket { + read_half, + write_half, + packet_socket, + } = socket; + let mut reader = FramedRead::new(read_half, MuxCodec::default()); + let mut writer = FramedWrite::new(write_half, MuxCodec::default()); + let (control_tx, mut control_rx) = mpsc::channel::(RESPONSE_MUX_WRITER_QUEUE); + let control_failed = CancellationToken::new(); + let mut reported_data_segments = packet_socket + .as_ref() + .and_then(|socket| super::tcp_data_segments_out_fd(socket.as_raw_fd())); + + crate::metrics::response_mux::CONNECTIONS_TOTAL + .with_label_values(&["frontend", "accepted"]) + .inc(); + crate::metrics::response_mux::ACTIVE_CONNECTIONS + .with_label_values(&["frontend"]) + .inc(); + + let writer_failed = control_failed.clone(); + let writer_task = tokio::spawn(async move { + let frame_counters = + crate::metrics::response_mux::FrameCounters::for_direction("frontend_to_worker"); + let write_calls = crate::metrics::response_mux::WRITE_CALLS_TOTAL + .with_label_values(&["frontend"]) + .clone(); + while let Some(frame) = control_rx.recv().await { + frame_counters.inc(frame.kind.metric_label()); + if let Err(err) = writer.send(frame).await { + writer_failed.cancel(); + return Err(err.into()); + } + write_calls.inc(); + } + Result::<()>::Ok(()) + }); + + let mut decoded_data_bytes = 0_u64; + let mut acknowledged_data_bytes = 0_u64; + let mut credit_tick = tokio::time::interval(RESPONSE_MUX_CREDIT_UPDATE_INTERVAL); + credit_tick.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay); + credit_tick.tick().await; + let mut packet_tick = tokio::time::interval(Duration::from_millis(100)); + packet_tick.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay); + packet_tick.tick().await; + let frame_counters = + crate::metrics::response_mux::FrameCounters::for_direction("worker_to_frontend"); + + let result: Result<()> = async { + loop { + let message = tokio::select! { + _ = control_failed.cancelled() => { + anyhow::bail!("frontend response mux control writer failed") + } + _ = packet_tick.tick(), if reported_data_segments.is_some() => { + if let Some(current) = packet_socket + .as_ref() + .and_then(|socket| super::tcp_data_segments_out_fd(socket.as_raw_fd())) + { + let previous = reported_data_segments.replace(current).unwrap_or(current); + crate::metrics::response_mux::DATA_SEGMENTS_TOTAL + .with_label_values(&["mux", "frontend"]) + .inc_by(current.saturating_sub(previous)); + } + continue; + } + _ = credit_tick.tick(), if decoded_data_bytes > acknowledged_data_bytes => { + acknowledged_data_bytes = decoded_data_bytes; + if control_tx + .try_send(MuxFrame::connection_ack(acknowledged_data_bytes)) + .is_err() + { + control_failed.cancel(); + } + continue; + } + message = reader.next() => match message { + Some(message) => message?, + None => anyhow::bail!("worker closed response mux connection"), + }, + }; + let frame = message; + frame_counters.inc(frame.kind.metric_label()); + let stream_id = frame.stream_id; + + match frame.kind { + MuxFrameKind::Prologue => { + let Some((_, pending)) = response_pending.remove(&stream_id) else { + crate::metrics::response_mux::RESETS_TOTAL + .with_label_values(&["frontend", "unknown_stream"]) + .inc(); + if control_tx + .try_send(MuxFrame::empty(MuxFrameKind::Reset, stream_id)) + .is_err() + { + control_failed.cancel(); + } + continue; + }; + let prologue: ResponseStreamPrologue = + serde_json::from_slice(&frame.payload).map_err(|err| { + error!("invalid response mux prologue for {stream_id}: {err}") + })?; + let Some(connection) = pending.connection.lock().take() else { + if control_tx + .try_send(MuxFrame::empty(MuxFrameKind::Reset, stream_id)) + .is_err() + { + control_failed.cancel(); + } + continue; + }; + crate::metrics::response_mux::SETUP_SECONDS + .observe(pending.registered_at.elapsed().as_secs_f64()); + if let Some(error) = prologue.error { + let _ = connection.send(Err(error)); + remove_response_association(&state, stream_id); + continue; + } + + let mailbox_frames = pending.send_buffer_count.max( + mux_config.stream_window_bytes + / crate::pipeline::network::tcp::mux::MUX_HEADER_LEN, + ); + let (response_tx, response_rx) = + data_plane_channel::(mailbox_frames); + let active_for_window = response_active.clone(); + let active_for_close = response_active.clone(); + let control_for_window = control_tx.clone(); + let control_for_close = control_tx.clone(); + let failed_for_window = control_failed.clone(); + let failed_for_close = control_failed.clone(); + let state_for_close = state.clone(); + let context = pending.context.clone(); + let hooks = StreamReceiverHooks { + context: pending.context, + window_update_threshold: RESPONSE_MUX_CREDIT_UPDATE_BYTES + .min(mux_config.stream_window_bytes), + on_window_update: Arc::new(move |mut credits| { + if !active_for_window.contains_key(&stream_id) { + return; + } + while credits > 0 { + let update = credits + .min(mux_config.stream_window_bytes.min(u32::MAX as usize)); + if control_for_window + .try_send(MuxFrame::window_update(stream_id, update as u32)) + .is_err() + { + failed_for_window.cancel(); + return; + } + credits -= update; + } + }), + on_close: Arc::new(move |control| match control { + ControlMessage::Stop => { + if active_for_close.contains_key(&stream_id) + && control_for_close + .try_send(MuxFrame::empty(MuxFrameKind::Stop, stream_id)) + .is_err() + { + failed_for_close.cancel(); + } + } + ControlMessage::Kill => { + if active_for_close.remove(&stream_id).is_some() { + if control_for_close + .try_send(MuxFrame::empty(MuxFrameKind::Kill, stream_id)) + .is_err() + { + failed_for_close.cancel(); + } + remove_response_association(&state_for_close, stream_id); + } + } + ControlMessage::Sentinel => {} + }), + }; + response_active.insert( + stream_id, + ActiveMuxResponseStream { + connection_id, + context, + response_tx, + control_tx: control_tx.clone(), + control_failed: control_failed.clone(), + }, + ); + crate::metrics::response_mux::ACTIVE_STREAMS + .with_label_values(&["frontend"]) + .inc(); + if connection + .send(Ok(StreamReceiver::multiplexed(response_rx, hooks))) + .is_err() + { + response_active.remove(&stream_id); + if control_tx + .try_send(MuxFrame::empty(MuxFrameKind::Kill, stream_id)) + .is_err() + { + control_failed.cancel(); + } + remove_response_association(&state, stream_id); + } + } + MuxFrameKind::Data => { + let encoded_len = frame.encoded_len(); + decoded_data_bytes = decoded_data_bytes.saturating_add(encoded_len as u64); + if decoded_data_bytes.saturating_sub(acknowledged_data_bytes) + >= RESPONSE_MUX_CREDIT_UPDATE_BYTES as u64 + { + acknowledged_data_bytes = decoded_data_bytes; + if control_tx + .try_send(MuxFrame::connection_ack(acknowledged_data_bytes)) + .is_err() + { + control_failed.cancel(); + } + } + let delivery_failure = match response_active.get(&stream_id) { + Some(active) if active.connection_id == connection_id => match active + .response_tx + .try_send(StreamRxItem::multiplexed(frame.payload, encoded_len)) + { + Ok(()) => None, + Err(mpsc::error::TrySendError::Full(_)) => Some("mailbox_full"), + Err(mpsc::error::TrySendError::Closed(_)) => { + Some("receiver_closed") + } + }, + _ => Some("unknown_stream"), + }; + if let Some(reason) = delivery_failure { + crate::metrics::response_mux::RESETS_TOTAL + .with_label_values(&["frontend", reason]) + .inc(); + response_active.remove(&stream_id); + response_pending.remove(&stream_id); + if control_tx + .try_send(MuxFrame::empty(MuxFrameKind::Reset, stream_id)) + .is_err() + { + control_failed.cancel(); + } + remove_response_association(&state, stream_id); + } + } + MuxFrameKind::End => { + if response_active + .get(&stream_id) + .is_some_and(|active| active.connection_id == connection_id) + { + response_active.remove(&stream_id); + remove_response_association(&state, stream_id); + } else { + if control_tx + .try_send(MuxFrame::empty(MuxFrameKind::Reset, stream_id)) + .is_err() + { + control_failed.cancel(); + } + } + } + MuxFrameKind::Reset => { + response_pending.remove(&stream_id); + if response_active + .get(&stream_id) + .is_some_and(|active| active.connection_id == connection_id) + { + response_active.remove(&stream_id); + } + remove_response_association(&state, stream_id); + } + MuxFrameKind::Stop + | MuxFrameKind::Kill + | MuxFrameKind::WindowUpdate + | MuxFrameKind::ConnectionAck => { + if control_tx + .try_send(MuxFrame::empty(MuxFrameKind::Reset, stream_id)) + .is_err() + { + control_failed.cancel(); + } + } + } + } + } + .await; + + if let Some(previous) = reported_data_segments + && let Some(current) = packet_socket + .as_ref() + .and_then(|socket| super::tcp_data_segments_out_fd(socket.as_raw_fd())) + { + crate::metrics::response_mux::DATA_SEGMENTS_TOTAL + .with_label_values(&["mux", "frontend"]) + .inc_by(current.saturating_sub(previous)); + } + + let affected: Vec = response_active + .iter() + .filter(|entry| entry.connection_id == connection_id) + .map(|entry| *entry.key()) + .collect(); + crate::metrics::response_mux::CONNECTION_LOST_STREAMS_TOTAL.inc_by(affected.len() as u64); + for stream_id in affected { + if let Some((_, active)) = response_active.remove(&stream_id) { + active.context.kill(); + } + remove_response_association(&state, stream_id); + } + crate::metrics::response_mux::ACTIVE_CONNECTIONS + .with_label_values(&["frontend"]) + .dec(); + drop(control_tx); + writer_task.abort(); + let _ = writer_task.await; + result + } + + fn remove_response_association(state: &Mutex, stream_id: uuid::Uuid) { + let subject = stream_id.to_string(); + let mut state = state.lock(); + if let Some(key) = state.subject_instance.remove(&subject) + && let Some(subjects) = state.instance_subjects.get_mut(&key) + { + subjects.remove(&(StreamType::Response, subject)); + if subjects.is_empty() { + state.instance_subjects.remove(&key); + } + } + } + /// Symmetric to [`process_response_stream`] for the upstream→downstream /// data direction: deliver the [`StreamSender`] half registered by the /// upstream to whoever awaits it, then pump every frame the upstream pushes @@ -706,12 +1289,12 @@ async fn tcp_listener( let (request_tx, request_rx) = data_plane_channel(send_buffer_count); if connection - .send(Ok(crate::pipeline::network::StreamSender { - tx: request_tx, + .send(Ok(crate::pipeline::network::StreamSender::dedicated( + request_tx, // Request streams don't carry a downstream-prologue today; the // upstream may begin sending immediately. - prologue: None, - })) + None, + ))) .is_err() { return Err(error!( @@ -854,9 +1437,9 @@ async fn tcp_listener( let (response_tx, response_rx) = data_plane_channel(send_buffer_count); if connection - .send(Ok(crate::pipeline::network::StreamReceiver { - rx: response_rx, - })) + .send(Ok(crate::pipeline::network::StreamReceiver::dedicated( + response_rx, + ))) .is_err() { return Err(error!( @@ -890,7 +1473,7 @@ async fn tcp_listener( async fn network_receive_handler( mut framed_reader: FramedRead, TwoPartCodec>, - response_tx: mpsc::Sender, + response_tx: mpsc::Sender, control_tx: mpsc::Sender, context: Arc, ) { @@ -962,7 +1545,10 @@ async fn tcp_listener( } if !data.is_empty() - && let Err(err) = response_tx.send(data).await { + && let Err(err) = response_tx + .send(crate::pipeline::network::StreamRxItem::dedicated(data)) + .await + { tracing::debug!(?err, "forwarding body/data to response channel failed"); let _ = control_tx.send(ControlMessage::Kill).await; break; @@ -2029,7 +2615,7 @@ mod tests { for (idx, expected, stream_provider) in pending_streams { let mut stream = stream_provider.await.unwrap().unwrap(); - let actual = stream.rx.recv().await.unwrap(); + let actual = stream.recv().await.unwrap(); assert_eq!(actual, expected, "payload mismatch for stream {idx}"); } }) From 06d22825e9111caec638c748f59d81e637747e3b Mon Sep 17 00:00:00 2001 From: jthomson04 Date: Wed, 22 Jul 2026 10:25:19 -0700 Subject: [PATCH 2/3] feat(runtime): make TCP response mux mandatory Signed-off-by: jthomson04 --- benchmarks/frontend/scripts/run_perf.sh | 67 +- docs/design-docs/request-plane.md | 17 + lib/runtime/src/config/environment_names.rs | 4 - lib/runtime/src/metrics/response_mux.rs | 238 ++-- .../pipeline/network/ingress/push_handler.rs | 47 +- lib/runtime/src/pipeline/network/tcp.rs | 126 +- .../src/pipeline/network/tcp/client.rs | 1251 +---------------- lib/runtime/src/pipeline/network/tcp/mux.rs | 14 +- .../src/pipeline/network/tcp/mux/client.rs | 1215 +++++++++++----- .../src/pipeline/network/tcp/server.rs | 1000 ++++--------- 10 files changed, 1381 insertions(+), 2598 deletions(-) diff --git a/benchmarks/frontend/scripts/run_perf.sh b/benchmarks/frontend/scripts/run_perf.sh index c5fe01971ff7..8d4cdc5105a1 100755 --- a/benchmarks/frontend/scripts/run_perf.sh +++ b/benchmarks/frontend/scripts/run_perf.sh @@ -64,12 +64,6 @@ BENCHMARK_DURATION="${BENCHMARK_DURATION:-}" # aiperf --benchmark-duration (sec REQUEST_RATE="${REQUEST_RATE:-}" # aiperf --request-rate (requests/sec) WARMUP_DURATION="${WARMUP_DURATION:-}" # aiperf --warmup-duration (seconds) WARMUP_COUNT="${WARMUP_COUNT:-}" # aiperf --warmup-request-count -FRONTEND_CORES="${FRONTEND_CORES:-}" # Optional taskset for frontend only -OTHER_CORES="${OTHER_CORES:-}" # Optional shorthand for mockers + aiperf -WORKER_CORES="${WORKER_CORES:-}" # Optional taskset for mockers -CLIENT_CORES="${CLIENT_CORES:-}" # Optional taskset for aiperf -AIPERF_WORKERS_MAX="${AIPERF_WORKERS_MAX:-}" # Optional client worker cap; aiperf auto-sizes by default -AIPERF_RECORD_PROCESSORS="${AIPERF_RECORD_PROCESSORS:-}" # Optional metrics process count; aiperf auto-sizes by default # Opt-out flags SKIP_BPF=false @@ -110,12 +104,6 @@ while [[ $# -gt 0 ]]; do --request-rate) REQUEST_RATE="$2"; shift 2 ;; --warmup-duration) WARMUP_DURATION="$2"; shift 2 ;; --warmup-count) WARMUP_COUNT="$2"; shift 2 ;; - --frontend-cores) FRONTEND_CORES="$2"; shift 2 ;; - --other-cores) OTHER_CORES="$2"; shift 2 ;; - --worker-cores) WORKER_CORES="$2"; shift 2 ;; - --client-cores) CLIENT_CORES="$2"; shift 2 ;; - --aiperf-workers-max) AIPERF_WORKERS_MAX="$2"; shift 2 ;; - --aiperf-record-processors) AIPERF_RECORD_PROCESSORS="$2"; shift 2 ;; --skip-bpf) SKIP_BPF=true; shift ;; --skip-nsys) SKIP_NSYS=true; shift ;; --skip-flamegraph) SKIP_FLAMEGRAPH=true; shift ;; @@ -149,13 +137,6 @@ Service Options: --request-rate N Target requests per second (aiperf --request-rate) --warmup-duration N aiperf warmup phase duration in seconds --warmup-count N aiperf warmup request count (default: concurrency) - --frontend-cores LIST Pin frontend to this taskset CPU list (for example 0-3) - --other-cores LIST Pin mockers and aiperf to this taskset CPU list (for example 4-23) - --worker-cores LIST Pin mockers to this taskset CPU list (overrides --other-cores) - --client-cores LIST Pin aiperf to this taskset CPU list (overrides --other-cores) - --aiperf-workers-max N Override aiperf's auto-sized client worker pool - --aiperf-record-processors N - Override aiperf's auto-sized metrics process pool Load Options: --concurrency N aiperf concurrency (default: 64) @@ -180,20 +161,6 @@ USAGE esac done -[[ -z "$WORKER_CORES" ]] && WORKER_CORES="$OTHER_CORES" -[[ -z "$CLIENT_CORES" ]] && CLIENT_CORES="$OTHER_CORES" -FRONTEND_CPU_CMD=() -WORKER_CPU_CMD=() -CLIENT_CPU_CMD=() -if [[ -n "$FRONTEND_CORES" || -n "$WORKER_CORES" || -n "$CLIENT_CORES" ]]; then - command -v taskset >/dev/null 2>&1 || { - echo "ERROR: taskset is required when CPU lists are configured"; exit 1; - } -fi -[[ -n "$FRONTEND_CORES" ]] && FRONTEND_CPU_CMD=(taskset -c "$FRONTEND_CORES") -[[ -n "$WORKER_CORES" ]] && WORKER_CPU_CMD=(taskset -c "$WORKER_CORES") -[[ -n "$CLIENT_CORES" ]] && CLIENT_CPU_CMD=(taskset -c "$CLIENT_CORES") - # Default model-name to model if not set [[ -z "$MODEL_NAME" ]] && MODEL_NAME="$MODEL" @@ -246,8 +213,6 @@ echo "╚═══════════════════════ echo "" echo "Output: $OUTPUT_DIR" echo "Tokenizer: ${TOKENIZER_BACKEND:-hf (default)}" -echo "CPU sets: frontend=${FRONTEND_CORES:-unrestricted} workers=${WORKER_CORES:-unrestricted} client=${CLIENT_CORES:-unrestricted}" -echo "Response: mux=${DYN_TCP_RESPONSE_MUX:-0} batch=${DYN_TCP_RESPONSE_BATCH_INTERVAL_MS:-5}ms window=${DYN_TCP_RESPONSE_CONNECTION_WINDOW_BYTES:-262144}B packet_metrics=${DYN_TCP_RESPONSE_PACKET_METRICS:-0}" echo "" # ─── Pre-flight: detect available tools ────────────────────────────────────── @@ -444,8 +409,7 @@ for MN in "${MODEL_NAMES[@]}"; do fi MN_SAFE="${MN//\//_}" - HF_HUB_OFFLINE=1 DYN_SYSTEM_PORT=$WORKER_PORT DYN_EVENT_PLANE="$EVENT_PLANE" \ - "${WORKER_CPU_CMD[@]}" python -m dynamo.mocker "${MOCKER_ARGS[@]}" \ + HF_HUB_OFFLINE=1 DYN_SYSTEM_PORT=$WORKER_PORT DYN_EVENT_PLANE="$EVENT_PLANE" python -m dynamo.mocker "${MOCKER_ARGS[@]}" \ > "$OUTPUT_DIR/logs/mocker_${MN_SAFE}_${i}.log" 2>&1 & ALL_PIDS+=($!) echo " Worker $WORKER_IDX ($MN #$i): PID ${ALL_PIDS[-1]}, port $WORKER_PORT" @@ -491,7 +455,7 @@ fi if [[ "$HAS_NSYS" == true ]]; then echo " (under nsys profiling)" - env "${FRONTEND_ENV[@]}" "${FRONTEND_CPU_CMD[@]}" \ + env "${FRONTEND_ENV[@]}" \ "$NSYS_CMD" profile \ --trace=osrt,nvtx \ --sample=cpu \ @@ -521,7 +485,7 @@ if [[ "$HAS_NSYS" == true ]]; then echo " nsys wrapper PID: $NSYS_WRAPPER_PID" echo " Frontend PID: $FRONTEND_PID" else - env "${FRONTEND_ENV[@]}" "${FRONTEND_CPU_CMD[@]}" python -m dynamo.frontend \ + env "${FRONTEND_ENV[@]}" python -m dynamo.frontend \ > "$OUTPUT_DIR/logs/frontend.log" 2>&1 & FRONTEND_PID=$! ALL_PIDS+=($FRONTEND_PID) @@ -761,14 +725,6 @@ else _WARMUP_ARGS+=(--warmup-request-count "$CONCURRENCY") fi -_AIPERF_WORKER_ARGS=() -if [[ -n "$AIPERF_WORKERS_MAX" ]]; then - _AIPERF_WORKER_ARGS+=(--workers-max "$AIPERF_WORKERS_MAX") -fi -if [[ -n "$AIPERF_RECORD_PROCESSORS" ]]; then - _AIPERF_WORKER_ARGS+=(--record-processors "$AIPERF_RECORD_PROCESSORS") -fi - # Build the list of models to target _AIPERF_MODELS=() if [[ "$AIPERF_TARGETS" == "all" && ${#MODEL_NAMES[@]} -gt 1 ]]; then @@ -795,7 +751,7 @@ for _AIPERF_MODEL in "${_AIPERF_MODELS[@]}"; do _AIPERF_TOK_ARGS=(--tokenizer "$MODEL") fi - HF_HUB_OFFLINE=1 "${CLIENT_CPU_CMD[@]}" aiperf profile --artifact-dir "$AIPERF_ARTIFACT_DIR" \ + HF_HUB_OFFLINE=1 aiperf profile --artifact-dir "$AIPERF_ARTIFACT_DIR" \ --model "$_AIPERF_MODEL" \ "${_AIPERF_TOK_ARGS[@]}" \ --endpoint-type chat \ @@ -816,7 +772,8 @@ for _AIPERF_MODEL in "${_AIPERF_MODELS[@]}"; do "${_WARMUP_ARGS[@]}" \ --num-dataset-entries 12800 \ --random-seed 100 \ - "${_AIPERF_WORKER_ARGS[@]}" \ + --workers-max "$CONCURRENCY" \ + --record-processors 32 \ --ui simple || echo "WARNING: aiperf failed for model ${_AIPERF_MODEL}" done @@ -956,18 +913,6 @@ cat > "$OUTPUT_DIR/config.json" < = Lazy::new(|| { .expect("response mux connection-lost stream counter") }); -static REGISTERED: OnceCell<()> = OnceCell::new(); +static REGISTERED: Lazy>>>> = + Lazy::new(|| Mutex::new(Vec::new())); pub fn ensure_registered(registry: &MetricsRegistry) { - let _ = REGISTERED.get_or_init(|| { - registry.add_metric_or_warn( - Box::new(ACTIVE_CONNECTIONS.clone()), - "response_mux_active_connections", - ); - registry.add_metric_or_warn( - Box::new(CONNECTIONS_TOTAL.clone()), - "response_mux_connections_total", - ); - registry.add_metric_or_warn( - Box::new(ACTIVE_STREAMS.clone()), - "response_mux_active_streams", - ); - registry.add_metric_or_warn( - Box::new(SETUP_SECONDS.clone()), - "response_mux_setup_seconds", - ); - registry.add_metric_or_warn(Box::new(FRAMES_TOTAL.clone()), "response_mux_frames_total"); - registry.add_metric_or_warn( - Box::new(WRITER_QUEUE_DEPTH.clone()), - "response_mux_writer_queue_depth", - ); - registry.add_metric_or_warn(Box::new(QUEUED_BYTES.clone()), "response_mux_queued_bytes"); - registry.add_metric_or_warn( - Box::new(FRAMES_PER_WRITE.clone()), - "response_mux_frames_per_write", - ); - registry.add_metric_or_warn(Box::new(BATCH_BYTES.clone()), "response_mux_batch_bytes"); - registry.add_metric_or_warn( - Box::new(BATCH_WAIT_SECONDS.clone()), - "response_mux_batch_wait_seconds", - ); - registry.add_metric_or_warn( - Box::new(CONFIGURED_BATCH_INTERVAL_MS.clone()), - "response_mux_configured_batch_interval_ms", - ); - registry.add_metric_or_warn( - Box::new(WRITE_CALLS_TOTAL.clone()), - "response_mux_write_calls_total", - ); - registry.add_metric_or_warn( - Box::new(DATA_SEGMENTS_TOTAL.clone()), - "response_data_segments_total", - ); - registry.add_metric_or_warn( - Box::new(QUEUE_RESIDENCE_SECONDS.clone()), - "response_mux_queue_residence_seconds", - ); - registry.add_metric_or_warn(Box::new(RESETS_TOTAL.clone()), "response_mux_resets_total"); - registry.add_metric_or_warn( - Box::new(RECONNECTS_TOTAL.clone()), - "response_mux_reconnects_total", - ); - registry.add_metric_or_warn( - Box::new(FLOW_CONTROL_STALL_SECONDS.clone()), - "response_mux_flow_control_stall_seconds", - ); - registry.add_metric_or_warn( - Box::new(CONNECTION_FLOW_CONTROL_STALL_SECONDS.clone()), - "response_mux_connection_flow_control_stall_seconds", - ); - registry.add_metric_or_warn( - Box::new(WRITER_ADMISSION_STALL_SECONDS.clone()), - "response_mux_writer_admission_stall_seconds", - ); - registry.add_metric_or_warn( - Box::new(QUEUED_BYTE_ADMISSION_STALL_SECONDS.clone()), - "response_mux_queued_byte_admission_stall_seconds", - ); - registry.add_metric_or_warn( - Box::new(STREAM_WRITER_QUEUE_OCCUPANCY.clone()), - "response_mux_stream_writer_queue_occupancy", - ); - registry.add_metric_or_warn( - Box::new(READY_STREAMS.clone()), - "response_mux_ready_streams", - ); - registry.add_metric_or_warn( - Box::new(PRIORITY_QUEUE_RESIDENCE_SECONDS.clone()), - "response_mux_priority_queue_residence_seconds", - ); - registry.add_metric_or_warn( - Box::new(ROUND_ROBIN_TURNS_TOTAL.clone()), - "response_mux_round_robin_turns_total", - ); - registry.add_metric_or_warn( - Box::new(WINDOW_UPDATES_TOTAL.clone()), - "response_mux_window_updates_total", - ); - registry.add_metric_or_warn( - Box::new(CONNECTION_LOST_STREAMS_TOTAL.clone()), - "response_mux_connection_lost_streams_total", - ); - }); + { + let mut registered = REGISTERED.lock().expect("response mux registry lock"); + registered.retain(|candidate| candidate.strong_count() > 0); + let identity = std::sync::Arc::downgrade(®istry.prometheus_registry); + if registered + .iter() + .any(|candidate| Weak::ptr_eq(candidate, &identity)) + { + return; + } + registered.push(identity); + } + + registry.add_metric_or_warn( + Box::new(ACTIVE_CONNECTIONS.clone()), + "response_mux_active_connections", + ); + registry.add_metric_or_warn( + Box::new(CONNECTIONS_TOTAL.clone()), + "response_mux_connections_total", + ); + registry.add_metric_or_warn( + Box::new(ACTIVE_STREAMS.clone()), + "response_mux_active_streams", + ); + registry.add_metric_or_warn( + Box::new(SETUP_SECONDS.clone()), + "response_mux_setup_seconds", + ); + registry.add_metric_or_warn(Box::new(FRAMES_TOTAL.clone()), "response_mux_frames_total"); + registry.add_metric_or_warn( + Box::new(WRITER_QUEUE_DEPTH.clone()), + "response_mux_writer_queue_depth", + ); + registry.add_metric_or_warn(Box::new(QUEUED_BYTES.clone()), "response_mux_queued_bytes"); + registry.add_metric_or_warn( + Box::new(FRAMES_PER_WRITE.clone()), + "response_mux_frames_per_write", + ); + registry.add_metric_or_warn(Box::new(BATCH_BYTES.clone()), "response_mux_batch_bytes"); + registry.add_metric_or_warn( + Box::new(BATCH_WAIT_SECONDS.clone()), + "response_mux_batch_wait_seconds", + ); + registry.add_metric_or_warn( + Box::new(CONFIGURED_BATCH_INTERVAL_MS.clone()), + "response_mux_configured_batch_interval_ms", + ); + registry.add_metric_or_warn( + Box::new(WRITE_CALLS_TOTAL.clone()), + "response_mux_write_calls_total", + ); + registry.add_metric_or_warn( + Box::new(DATA_SEGMENTS_TOTAL.clone()), + "response_data_segments_total", + ); + registry.add_metric_or_warn( + Box::new(QUEUE_RESIDENCE_SECONDS.clone()), + "response_mux_queue_residence_seconds", + ); + registry.add_metric_or_warn(Box::new(RESETS_TOTAL.clone()), "response_mux_resets_total"); + registry.add_metric_or_warn( + Box::new(RECONNECTS_TOTAL.clone()), + "response_mux_reconnects_total", + ); + registry.add_metric_or_warn( + Box::new(FLOW_CONTROL_STALL_SECONDS.clone()), + "response_mux_flow_control_stall_seconds", + ); + registry.add_metric_or_warn( + Box::new(CONNECTION_FLOW_CONTROL_STALL_SECONDS.clone()), + "response_mux_connection_flow_control_stall_seconds", + ); + registry.add_metric_or_warn( + Box::new(WRITER_ADMISSION_STALL_SECONDS.clone()), + "response_mux_writer_admission_stall_seconds", + ); + registry.add_metric_or_warn( + Box::new(QUEUED_BYTE_ADMISSION_STALL_SECONDS.clone()), + "response_mux_queued_byte_admission_stall_seconds", + ); + registry.add_metric_or_warn( + Box::new(STREAM_WRITER_QUEUE_OCCUPANCY.clone()), + "response_mux_stream_writer_queue_occupancy", + ); + registry.add_metric_or_warn( + Box::new(READY_STREAMS.clone()), + "response_mux_ready_streams", + ); + registry.add_metric_or_warn( + Box::new(PRIORITY_QUEUE_RESIDENCE_SECONDS.clone()), + "response_mux_priority_queue_residence_seconds", + ); + registry.add_metric_or_warn( + Box::new(ROUND_ROBIN_TURNS_TOTAL.clone()), + "response_mux_round_robin_turns_total", + ); + registry.add_metric_or_warn( + Box::new(WINDOW_UPDATES_TOTAL.clone()), + "response_mux_window_updates_total", + ); + registry.add_metric_or_warn( + Box::new(CONNECTION_LOST_STREAMS_TOTAL.clone()), + "response_mux_connection_lost_streams_total", + ); +} + +#[cfg(test)] +mod tests { + use super::*; + + fn contains_response_mux_metrics(registry: &MetricsRegistry) -> bool { + registry + .prometheus_registry + .read() + .expect("metrics registry read lock") + .gather() + .iter() + .any(|family| family.name() == "dynamo_tcp_response_mux_active_connections") + } + + #[test] + fn registration_is_deduplicated_by_registry_identity() { + let first = MetricsRegistry::new(); + let first_clone = first.clone(); + let second = MetricsRegistry::new(); + ACTIVE_CONNECTIONS + .with_label_values(&["registry_test"]) + .set(0); + + ensure_registered(&first); + ensure_registered(&first_clone); + ensure_registered(&second); + + assert!(contains_response_mux_metrics(&first)); + assert!(contains_response_mux_metrics(&second)); + } } diff --git a/lib/runtime/src/pipeline/network/ingress/push_handler.rs b/lib/runtime/src/pipeline/network/ingress/push_handler.rs index c016477e8cba..858cb6f6a72b 100644 --- a/lib/runtime/src/pipeline/network/ingress/push_handler.rs +++ b/lib/runtime/src/pipeline/network/ingress/push_handler.rs @@ -595,27 +595,22 @@ where WORK_HANDLER_NETWORK_TRANSIT_SECONDS.observe(transit_ns as f64 / 1_000_000_000.0); } - tracing::trace!("creating tcp response stream"); - let mut publisher = - if response_connection_info.transport == tcp::TCP_RESPONSE_MUX_TRANSPORT { - self.response_mux_client - .get() - .ok_or_else(|| { - PipelineError::Generic( - "response mux client was not initialized for endpoint ingress" - .to_string(), - ) - })? - .create_response_stream(request.context(), response_connection_info) - .await - } else { - tcp::client::TcpClient::create_response_stream( - request.context(), - response_connection_info, - self.metrics().map(|m| m.cancellation_total.clone()), + tracing::trace!("creating multiplexed TCP response stream"); + let response_context = request.context(); + let mut publisher = self + .response_mux_client + .get() + .ok_or_else(|| { + PipelineError::Generic( + "response mux client was not initialized for endpoint ingress".to_string(), ) - .await - } + })? + .create_response_stream( + response_context.clone(), + response_connection_info, + self.metrics().map(|m| m.cancellation_total.clone()), + ) + .await .map_err(|e| { if let Some(m) = self.metrics() { m.error_counter @@ -673,9 +668,15 @@ where self.pump_response_stream(stream, &publisher, payload_codec) .await; - publisher.finish().await.map_err(|err| { - PipelineError::Generic(format!("Failed to finish response stream: {err}")) - })?; + if let Err(err) = publisher.finish().await { + if response_context.is_killed() || response_context.is_stopped() { + tracing::debug!(%err, "response stream closed by frontend cancellation"); + } else { + return Err(PipelineError::Generic(format!( + "Failed to finish response stream: {err}" + ))); + } + } // Ensure the metrics guard is not dropped until the end of the function. // Drop fires "request completed" log via RAII. diff --git a/lib/runtime/src/pipeline/network/tcp.rs b/lib/runtime/src/pipeline/network/tcp.rs index dd75352ef530..e554368bc59d 100644 --- a/lib/runtime/src/pipeline/network/tcp.rs +++ b/lib/runtime/src/pipeline/network/tcp.rs @@ -12,30 +12,30 @@ //! The request plane (TCP, NATS, etc.) carries a two-part message whose header is a //! `RequestControlMessage` — embedding the [`ConnectionInfo`] that tells the worker where to call home //! — and whose data half is the serialized request body if request streaming is not needed. -//! All subsequent streaming bytes (responses, and request-stream) flow over the TCP socket established afterwards. +//! Subsequent request-stream bytes use a dedicated TCP socket. Responses use +//! persistent multiplexed TCP connections shared by logical response streams. //! //! For simplicity, if request streaming is needed, the `RequestControlMessage` should not contain //! the request body. Instead, all requests of the stream should be sent over the TCP socket. //! //! The TCP transport is the implementation that produces and consumes [`ConnectionInfo`] and -//! carries the streaming response (and, optionally, request-stream) bytes between two peers -//! on separate sockets from the initial request. +//! carries response and optional request-stream bytes between two peers on +//! sockets separate from the initial request. //! //! # Roles //! //! The TCP transport has two sides: //! //! - Request sender: The upstream that **initiates the transfer** runs [`server::TcpStreamServer`], -//! registers what it expects to receive, and listens. It publishes its address + a per-stream -//! subject UUID via [`TcpStreamConnectionInfo`], which is serialized into a [`ConnectionInfo`]. +//! registers what it expects to receive, and listens. It publishes dedicated +//! request-stream information or versioned response-mux information through [`ConnectionInfo`]. //! - Request receiver: The downstream that **acknowledges the transfer** runs [`client::TcpClient`], //! reads the connection info out of the request, dials the listener, and identifies itself with //! a `CallHomeHandshake` to the request sender. //! -//! Although TCP is bidirectional, we keep separate sockets for the request stream and the response -//! stream to match Dynamo's design principles. To establish both, the request receiver must receive -//! two [`ConnectionInfo`] objects and run two handshakes — each [`StreamType`] is its own TCP -//! connection with its own subject UUID. +//! Request streams retain one dedicated socket per logical stream. Response +//! streams are mandatory mux-v1 streams: each worker/frontend pair warms four +//! persistent connections and identifies logical responses by UUID. //! //! # Server-Client Interaction //! @@ -44,11 +44,9 @@ //! //! # Stream Types //! -//! [`StreamType::Response`] — worker pushes engine output back to the upstream. Server side is -//! `process_response_stream` (delivers a [`StreamReceiver`] to the awaiting registrant once the -//! client has sent its [`ResponseStreamPrologue`]). Client side is -//! [`client::TcpClient::create_response_stream`] (returns a [`StreamSender`]; spawns reader/writer -//! tasks plus a connection monitor that waits for the server's FIN). +//! [`StreamType::Response`] — worker pushes engine output through the response +//! mux pool. A `Prologue` activates the pending stream, `Data` carries output, +//! and `End` completes only that logical stream. //! //! [`StreamType::Request`] — upstream pushes the request body (or a stream of follow-up frames) //! into a downstream worker. Server side is `process_request_stream` (delivers a [`StreamSender`] @@ -69,27 +67,25 @@ //! 1. The returned [`RegisteredStream`] is RAII — dropping it without `into_parts()` removes the //! pending entry from the server's subject tables. This is typically used by the request sender //! up until the `RequestControlMessage` is sent and the stream is established. -//! 2. The server tracks `subject UUID → oneshot` in `tx_subjects` / `rx_subjects`. +//! 2. The server tracks dedicated request subjects and pending response-mux UUIDs. //! [`server::TcpStreamServer::associate_instance`] links one or both //! subjects to a discovery instance so [`server::TcpStreamServer::cancel_instance_streams`] can //! drop both halves' oneshots together when a worker disappears. Tombstones (`TOMBSTONE_TTL`) //! are the safety net that closes the cancel-vs-register race. //! -//! # CallHome Handshake +//! # Handshakes //! -//! The first message a [`client::TcpClient`] sends on a freshly-opened socket is a -//! `CallHomeHandshake` header-only frame carrying `{ subject, stream_type }`. The server pops -//! the matching entry out of `tx_subjects` (for [`StreamType::Request`]) or `rx_subjects` (for -//! [`StreamType::Response`]) and resolves the registrant's oneshot. After that the socket carries -//! framed [`TwoPartCodec`] messages: data frames in the natural direction, control frames -//! (and, for response streams, the prologue) interleaved. +//! A dedicated request socket starts with a `CallHomeHandshake` carrying its +//! subject and [`StreamType::Request`]. A response-mux connection instead uses +//! a versioned handshake containing the frontend and physical-connection UUIDs; +//! incompatible response versions are rejected without a legacy fallback. //! //! # Control / Shutdown Protocol //! -//! [`ControlMessage`] frames are header-only frames interleaved with data on either socket: +//! Dedicated request sockets use header-only [`ControlMessage`] frames: //! //! - [`ControlMessage::Sentinel`] — per-direction clean end-of-stream; the producing side emits -//! it before closing the socket. Used on both the request and response sockets. +//! it before closing the request socket. //! - [`ControlMessage::Stop`] — sender asks the receiver to cancel; the receiving side calls //! `context.stop()`. //! - [`ControlMessage::Kill`] — hard cancel; `context.kill()` and break out. @@ -103,10 +99,10 @@ //! The cancellation direction is fixed, but the two streams carry **data** in opposite //! directions, so the practical handling differs per stream: //! -//! ## Response stream (downstream → upstream) — bidirectional +//! ## Response mux (downstream → upstream) — bidirectional //! -//! - Upstream writes: `Stop` / `Kill` (any time, to cancel). -//! - Downstream writes: data frames, then `Sentinel` on clean close (skipped on kill/stop). +//! - Upstream writes: mux `Stop`, `Kill`, `WindowUpdate`, `ConnectionAck`, and `Reset` frames. +//! - Downstream writes: mux `Prologue`, `Data`, `End`, and `Reset` frames. //! //! ## Request stream (upstream → downstream) — unidirectional after the handshake //! @@ -145,6 +141,7 @@ pub(crate) fn tcp_data_segments_out(stream: &tokio::net::TcpStream) -> Option Option { #[repr(C)] + // Layout through tcpi_data_segs_out from Linux uapi/linux/tcp.h::tcp_info. struct LinuxTcpInfoThroughDataSegments { _header: [u8; 8], _metrics: [u32; 24], @@ -346,79 +343,4 @@ mod tests { drop(send_stream); assert!(recv_stream.recv().await.is_none()); } - - #[tokio::test] - async fn test_tcp_stream_client_server() { - // [server] start the server and register the response stream - let options = server::ServerOptions::builder().port(9124).build().unwrap(); - let server = server::TcpStreamServer::new(options).await.unwrap(); - - let context_rank0 = Context::new(()); - - let options = StreamOptions::builder() - .context(context_rank0.context()) - .enable_request_stream(false) - .enable_response_stream(true) - .build() - .unwrap(); - - let pending_connection = server.register(options).await; - - let connection_info = pending_connection - .recv_stream - .as_ref() - .unwrap() - .connection_info - .clone(); - - // [client] set up the other rank and create the response stream - let context_rank1 = - Context::with_id_and_metadata((), context_rank0.id().to_string(), Default::default()); - - let mut send_stream = client::TcpClient::create_response_stream( - context_rank1.context(), - connection_info, - None, - ) - .await - .unwrap(); - - // the client can now setup it's end of the stream and if it errors, it can send a message - // to the server to stop the stream - // - // this step must be done before the next step on the server can complete, i.e. - // the server's stream is now blocked on receiving the prologue message - // - // let's improve this and use an enum like Ok/Err; currently, None means good-to-go, and - // Some(String) means an error happened on this downstream node and we need to alert the - // upstream node that an error occurred - send_stream.send_prologue(None).await.unwrap(); - - // [server] After client sends the prologue, the server can pick up its `StreamReceiver` half. - let (_conn_info, stream_provider) = pending_connection.recv_stream.unwrap().into_parts(); - let mut recv_stream = stream_provider.await.unwrap(); - - // [client] The client can now send the response message to the server - let msg = TestMessage { - foo: "bar".to_string(), - }; - - let payload = serde_json::to_vec(&msg).unwrap(); - - send_stream.send(payload.into()).await.unwrap(); - - // [server] The server can now receive the response message from the client - - let data = recv_stream.as_mut().unwrap().recv().await.unwrap(); - - let recv_msg = serde_json::from_slice::(&data).unwrap(); - - assert_eq!(msg.foo, recv_msg.foo); - - drop(send_stream); - - // let data = recv_stream.rx.recv().await; - - // assert!(data.is_none()); - } } diff --git a/lib/runtime/src/pipeline/network/tcp/client.rs b/lib/runtime/src/pipeline/network/tcp/client.rs index cf61a1b89222..f914d5aa3bde 100644 --- a/lib/runtime/src/pipeline/network/tcp/client.rs +++ b/lib/runtime/src/pipeline/network/tcp/client.rs @@ -4,12 +4,7 @@ use std::sync::Arc; use futures::{SinkExt, StreamExt}; -use tokio::io::{AsyncReadExt, ReadHalf, WriteHalf}; -use tokio::{ - io::AsyncWriteExt, - net::TcpStream, - time::{self, Duration, Instant}, -}; +use tokio::{io::AsyncWriteExt, net::TcpStream}; use tokio_util::codec::{FramedRead, FramedWrite}; use prometheus::IntCounter; @@ -17,7 +12,7 @@ use prometheus::IntCounter; use super::{CallHomeHandshake, ControlMessage, TcpStreamConnectionInfo}; use crate::engine::AsyncEngineContext; use crate::pipeline::network::{ - ConnectionInfo, ResponseStreamPrologue, StreamReceiver, StreamRxItem, StreamSender, + ConnectionInfo, StreamReceiver, StreamRxItem, codec::{TwoPartCodec, TwoPartMessage, TwoPartMessageType}, tcp::StreamType, }; @@ -62,116 +57,7 @@ impl TcpClient { } } - pub async fn create_response_stream( - context: Arc, - info: ConnectionInfo, - cancellation_counter: Option, - ) -> Result { - let info = - TcpStreamConnectionInfo::try_from(info).context("tcp-stream-connection-info-error")?; - tracing::trace!("Creating response stream for {:?}", info); - - if info.stream_type != StreamType::Response { - return Err(error!( - "Invalid stream type; TcpClient requires the stream type to be `response`; however {:?} was passed", - info.stream_type - )); - } - - if info.context != context.id() { - return Err(error!( - "Invalid context; TcpClient requires the context to be {:?}; however {:?} was passed", - context.id(), - info.context - )); - } - - let stream = TcpClient::connect(&info.address).await?; - let packet_baseline = super::mux::response_packet_metrics_enabled() - .then(|| super::tcp_data_segments_out(&stream)) - .flatten(); - let peer_port = stream.peer_addr().ok().map(|addr| addr.port()); - let (read_half, write_half) = tokio::io::split(stream); - - let framed_reader = FramedRead::new(read_half, TwoPartCodec::default()); - let mut framed_writer = FramedWrite::new(write_half, TwoPartCodec::default()); - - // this is a oneshot channel that will be used to signal when the stream is closed - // when the stream sender is dropped, the bytes_rx will be closed and the forwarder task will exit - // the forwarder task will capture the alive_rx half of the oneshot channel; this will close the alive channel - // so the holder of the alive_tx half will be notified that the stream is closed; the alive_tx channel will be - // captured by the monitor task - let (alive_tx, alive_rx) = tokio::sync::oneshot::channel::<()>(); - - let reader_task = tokio::spawn(handle_reader( - framed_reader, - context.clone(), - alive_tx, - cancellation_counter, - )); - - // transport specific handshake message - let handshake = CallHomeHandshake { - subject: info.subject.clone(), - stream_type: StreamType::Response, - }; - - let handshake_bytes = match serde_json::to_vec(&handshake) { - Ok(hb) => hb, - Err(err) => { - return Err(error!( - "create_response_stream: Error converting CallHomeHandshake to JSON array: {err:#}" - )); - } - }; - let msg = TwoPartMessage::from_header(handshake_bytes.into()); - - // issue the the first tcp handshake message - framed_writer - .send(msg) - .await - .map_err(|e| error!("failed to send handshake: {:?}", e))?; - - // set up the channel to send bytes to the transport layer - let (bytes_tx, bytes_rx) = tokio::sync::mpsc::channel(64); - - // forwards the bytes send from this stream to the transport layer; hold the alive_rx half of the oneshot channel - let writer_context = context.clone(); - let writer_task = tokio::spawn(handle_writer( - framed_writer, - bytes_rx, - alive_rx, - writer_context, - )); - - let subject = info.subject.clone(); - let monitor_context = context; - // Spawn the connection monitor; errors are already logged inside - // wait_for_connection_tasks, so the Result is intentionally dropped. - tokio::spawn(async move { - let _ = wait_for_connection_tasks( - reader_task, - writer_task, - monitor_context, - peer_port, - subject, - packet_baseline, - ) - .await; - }); - - // set up the prologue for the stream - // this might have transport specific metadata in the future - let prologue = Some(ResponseStreamPrologue { error: None }); - - // create the stream sender - let stream_sender = StreamSender::dedicated(bytes_tx, prologue); - - Ok(stream_sender) - } - - /// Symmetric to [`Self::create_response_stream`] for the request-stream half: - /// dial the upstream TCP server with `StreamType::Request`, then return a + /// Dial the upstream TCP server with `StreamType::Request`, then return a /// [`StreamReceiver`] that yields the data frames the upstream pushes down. /// /// The request stream is unidirectional after the handshake: the write half @@ -367,1143 +253,18 @@ async fn handle_request_reader( drop(bytes_tx); } -async fn wait_for_connection_tasks( - reader_task: tokio::task::JoinHandle, TwoPartCodec>>, - writer_task: tokio::task::JoinHandle, TwoPartCodec>>>, - context: Arc, - peer_port: Option, - subject: String, - packet_baseline: Option, -) -> Result<()> { - // Await the reader first and abort the writer on reader Err — the - // writer parks on `bytes_rx.recv()` and won't wake on its own. - let reader = match reader_task.await { - Ok(reader) => reader, - Err(reader_err) => { - writer_task.abort(); - let _ = writer_task.await; - tracing::error!( - subject = %subject, - peer_port = ?peer_port, - err = ?reader_err, - "reader task failed to join" - ); - return Err(reader_err.into()); - } - }; - - let writer = match writer_task.await { - Ok(writer) => writer, - Err(writer_err) => { - tracing::error!( - subject = %subject, - peer_port = ?peer_port, - err = ?writer_err, - "writer task failed to join" - ); - return Err(writer_err.into()); - } - }; - - let reader = reader.into_inner(); - let writer = match writer { - Ok(writer) => writer.into_inner(), - Err(e) => { - tracing::error!( - subject = %subject, - peer_port = ?peer_port, - err = ?e, - "writer task returned error" - ); - return Err(e); - } - }; - - let stream = reader.unsplit(writer); - if let Some(baseline) = packet_baseline - && let Some(current) = super::tcp_data_segments_out(&stream) - { - crate::metrics::response_mux::DATA_SEGMENTS_TOTAL - .with_label_values(&["dedicated", "worker"]) - .inc_by(current.saturating_sub(baseline)); - } - wait_for_server_shutdown(stream, context).await -} - -async fn wait_for_server_shutdown( - mut stream: TcpStream, - context: Arc, -) -> Result<()> { - // `handle_writer` skips the closing sentinel on both `killed` and - // `stopped`, so the server has nothing to react to in either case; - // sitting in the read loop until the 10 s deadline would be dead time. - if context.is_killed() || context.is_stopped() { - tracing::debug!("stream context killed or stopped; skipping server FIN wait"); - return Ok(()); - } - - // Await the tcp server to shutdown the socket connection, bounded by a - // timeout so normal sentinel shutdown cannot hang indefinitely. - let mut buf = [0u8; 1024]; - let deadline = Instant::now() + Duration::from_secs(10); - loop { - let n = time::timeout_at(deadline, stream.read(&mut buf)) - .await - .inspect_err(|_| { - tracing::debug!("server did not close socket within the deadline"); - })? - .inspect_err(|e| { - tracing::debug!(err = ?e, "failed to read from stream"); - })?; - if n == 0 { - // Server has closed (FIN) - break; - } - } - - Ok(()) -} - -async fn handle_reader( - framed_reader: FramedRead, TwoPartCodec>, - context: Arc, - alive_tx: tokio::sync::oneshot::Sender<()>, - cancellation_counter: Option, -) -> FramedRead, TwoPartCodec> { - let mut framed_reader = framed_reader; - let mut alive_tx = alive_tx; - // Set on every cancellation arm; counted once after the loop. - let mut cancellation_seen = false; - loop { - tokio::select! { - msg = framed_reader.next() => { - match msg { - Some(Ok(two_part_msg)) => { - match two_part_msg.optional_parts() { - (Some(bytes), None) => { - let msg = match serde_json::from_slice::(bytes) { - Ok(msg) => msg, - Err(e) => { - tracing::warn!( - err = ?e, - "invalid control message, closing connection" - ); - cancellation_seen = true; - context.kill(); - break; - } - }; - - // Stop/Kill intentionally do not `break`: the - // reader keeps running so a later Kill can - // upgrade an earlier Stop (and vice versa). - // The loop still exits promptly via the - // `alive_tx.closed()` arm once `handle_writer` - // reacts to `context.stop()` / `context.kill()`. - match msg { - ControlMessage::Stop => { - cancellation_seen = true; - context.stop(); - } - ControlMessage::Kill => { - cancellation_seen = true; - context.kill(); - } - ControlMessage::Sentinel => { - tracing::warn!( - "unexpected sentinel on client reader, closing connection" - ); - cancellation_seen = true; - context.kill(); - break; - } - } - } - _ => { - tracing::warn!( - "unexpected non-control message on client reader, closing connection" - ); - cancellation_seen = true; - context.kill(); - break; - } - } - } - Some(Err(e)) => { - // Kill the engine context so the producer stops - // generating responses that can no longer be delivered. - tracing::warn!(err = ?e, "tcp stream read error, closing connection"); - cancellation_seen = true; - context.kill(); - break; - } - None => { - tracing::debug!("tcp stream closed by server"); - cancellation_seen = true; - break; - } - } - } - _ = alive_tx.closed() => { - break; - } - } - } - if cancellation_seen && let Some(counter) = &cancellation_counter { - counter.inc(); - } - framed_reader -} - -async fn handle_writer( - mut framed_writer: FramedWrite, TwoPartCodec>, - mut bytes_rx: tokio::sync::mpsc::Receiver, - alive_rx: tokio::sync::oneshot::Receiver<()>, - context: Arc, -) -> Result, TwoPartCodec>> { - // Keep one cancellation future per stream. Recreating these futures for every queued - // frame repeatedly clones the context's watch receivers and churns Notify state. - let killed = context.killed(); - let stopped = context.stopped(); - tokio::pin!(killed, stopped); - - // Only send sentinel for normal channel closure - let mut send_sentinel = true; - - loop { - let msg = tokio::select! { - biased; - - _ = &mut killed => { - tracing::trace!("context kill signal received; shutting down"); - send_sentinel = false; - break; - } - - _ = &mut stopped => { - tracing::trace!("context stop signal received; shutting down"); - send_sentinel = false; - break; - } - - msg = bytes_rx.recv() => { - match msg { - Some(msg) => msg, - None => { - tracing::trace!("response channel closed; shutting down"); - break; - } - } - } - }; - - if let Err(e) = framed_writer.send(msg).await { - tracing::trace!( - "failed to send message to network; possible disconnect: {:?}", - e - ); - send_sentinel = false; - break; - } - } - - // Send sentinel only on normal closure - if send_sentinel { - let message = serde_json::to_vec(&ControlMessage::Sentinel)?; - let msg = TwoPartMessage::from_header(message.into()); - framed_writer.send(msg).await?; - } - - drop(alive_rx); - Ok(framed_writer) -} - #[cfg(test)] mod tests { use super::*; use crate::pipeline::context::Controller; use crate::pipeline::network::tcp::test_utils::create_tcp_pair; use bytes::Bytes; - use futures::StreamExt; use std::sync::Arc; - use tokio::io::{AsyncReadExt, AsyncWriteExt}; - use tokio::net::TcpStream; - use tokio::sync::{mpsc, oneshot}; - use tokio_util::codec::FramedRead; - - struct WriterHarness { - server: tokio::net::TcpStream, - framed_writer: FramedWrite, TwoPartCodec>, - bytes_tx: mpsc::Sender, - bytes_rx: mpsc::Receiver, - alive_tx: oneshot::Sender<()>, - alive_rx: oneshot::Receiver<()>, - controller: Arc, - } - - /// Creates a reusable writer harness with paired TCP streams and test channels. - async fn writer_harness() -> WriterHarness { - let (client, server) = create_tcp_pair().await; - let (_, write_half) = tokio::io::split(client); - let framed_writer = FramedWrite::new(write_half, TwoPartCodec::default()); - - let (bytes_tx, bytes_rx) = mpsc::channel(64); - let (alive_tx, alive_rx) = oneshot::channel::<()>(); - let controller = Arc::new(Controller::default()); - - WriterHarness { - server, - framed_writer, - bytes_tx, - bytes_rx, - alive_tx, - alive_rx, - controller, - } - } - - async fn recv_msg(reader: &mut FramedRead) -> TwoPartMessage { - reader - .next() - .await - .expect("expected message") - .expect("failed to decode message") - } - - fn assert_data_only_message(msg: TwoPartMessage, expected: &[u8]) { - let (header, data) = msg.optional_parts(); - assert!(header.is_none(), "data-only message should not have header"); - assert_eq!( - data.expect("data payload missing").as_ref(), - expected, - "data payload should match" - ); - } - - fn assert_header_only_message(msg: TwoPartMessage, expected: &[u8]) { - let (header, data) = msg.optional_parts(); - assert!(data.is_none(), "header-only message should not carry data"); - assert_eq!( - header.expect("header missing").as_ref(), - expected, - "header payload should match" - ); - } - - fn assert_header_and_data_message( - msg: TwoPartMessage, - expected_header: &[u8], - expected_data: &[u8], - ) { - let (header, data) = msg.optional_parts(); - assert_eq!( - header.expect("header missing").as_ref(), - expected_header, - "header payload should match" - ); - assert_eq!( - data.expect("data missing").as_ref(), - expected_data, - "data payload should match" - ); - } - - fn assert_sentinel_message(msg: TwoPartMessage) { - let (header, data) = msg.optional_parts(); - assert!(data.is_none(), "sentinel should not include a data section"); - let expected_sentinel = serde_json::to_vec(&ControlMessage::Sentinel).unwrap(); - assert_eq!( - header.expect("sentinel header missing").as_ref(), - expected_sentinel.as_slice(), - "sentinel header should match serialized ControlMessage::Sentinel" - ); - } - - /// Test that handle_writer forwards messages from the channel to the framed writer - #[tokio::test] - async fn test_handle_writer_forwards_messages() { - let WriterHarness { - server, - framed_writer, - bytes_tx, - bytes_rx, - alive_rx, - controller, - .. - } = writer_harness().await; - - // Send test messages - let test_msg = TwoPartMessage::from_data(Bytes::from("test data")); - bytes_tx.send(test_msg).await.unwrap(); - - // Close the sender to trigger normal termination - drop(bytes_tx); - - let result = handle_writer(framed_writer, bytes_rx, alive_rx, controller).await; - - assert!(result.is_ok()); - - // Decode from server side to verify data and sentinel were sent - let mut reader = FramedRead::new(server, TwoPartCodec::default()); - - let msg = recv_msg(&mut reader).await; - assert_data_only_message(msg, b"test data"); - - let sentinel = recv_msg(&mut reader).await; - assert_sentinel_message(sentinel); - } - - /// Test that handle_writer sends sentinel on normal channel closure - #[tokio::test] - async fn test_handle_writer_sends_sentinel_on_normal_closure() { - let WriterHarness { - mut server, - framed_writer, - bytes_tx, - bytes_rx, - alive_rx, - controller, - .. - } = writer_harness().await; - - // Close the sender immediately to trigger normal termination - drop(bytes_tx); - - let result = handle_writer(framed_writer, bytes_rx, alive_rx, controller).await; - - assert!(result.is_ok()); - - // Read from server side to verify sentinel was sent - let mut buffer = vec![0u8; 1024]; - let n = server.read(&mut buffer).await.unwrap(); - - // Buffer should contain the sentinel message - assert!(n > 0, "Expected sentinel to be written to the TCP stream"); - - // Verify it contains the sentinel message by checking for the JSON - let sentinel_json = serde_json::to_vec(&ControlMessage::Sentinel).unwrap(); - assert!( - buffer[..n] - .windows(sentinel_json.len()) - .any(|w| w == sentinel_json.as_slice()), - "Buffer should contain sentinel message. Buffer: {:?}", - String::from_utf8_lossy(&buffer[..n]) - ); - } - - /// Test that handle_writer does NOT send sentinel when context is killed - #[tokio::test] - async fn test_handle_writer_no_sentinel_on_context_killed() { - let WriterHarness { - mut server, - framed_writer, - bytes_rx, - alive_rx, - controller, - .. - } = writer_harness().await; - - // Kill the context - controller.kill(); - - let result = handle_writer(framed_writer, bytes_rx, alive_rx, controller).await; - - assert!(result.is_ok()); - - // Drop the writer to close the connection, then try to read. Otherwise, - // the test will hang on `server.read()` - drop(result); - - // Read from server side - should get no sentinel - let mut buffer = vec![0u8; 1024]; - let n = server.read(&mut buffer).await.unwrap(); - - // Buffer should be empty (no sentinel sent) - let sentinel_json = serde_json::to_vec(&ControlMessage::Sentinel).unwrap(); - assert!( - n == 0 - || !buffer[..n] - .windows(sentinel_json.len()) - .any(|w| w == sentinel_json.as_slice()), - "Buffer should NOT contain sentinel message when context is killed" - ); - } - - /// Test that handle_writer does NOT send sentinel when context is stopped - #[tokio::test] - async fn test_handle_writer_no_sentinel_on_context_stopped() { - let WriterHarness { - mut server, - framed_writer, - bytes_rx, - alive_rx, - controller, - .. - } = writer_harness().await; - - // Stop the context - controller.stop(); - - let result = handle_writer(framed_writer, bytes_rx, alive_rx, controller).await; - - assert!(result.is_ok()); - - // Drop the writer to close the connection, then try to read. Otherwise, - // the test will hang on `server.read()` - drop(result); - - // Read from server side - should get no sentinel - let mut buffer = vec![0u8; 1024]; - let n = server.read(&mut buffer).await.unwrap(); - - // Buffer should be empty (no sentinel sent) - let sentinel_json = serde_json::to_vec(&ControlMessage::Sentinel).unwrap(); - assert!( - n == 0 - || !buffer[..n] - .windows(sentinel_json.len()) - .any(|w| w == sentinel_json.as_slice()), - "Buffer should NOT contain sentinel message when context is stopped" - ); - } - - /// Test that handle_writer handles multiple messages correctly - #[tokio::test] - async fn test_handle_writer_multiple_messages() { - let WriterHarness { - server, - framed_writer, - bytes_tx, - bytes_rx, - alive_rx, - controller, - .. - } = writer_harness().await; - - // Send multiple messages - for i in 0..5 { - let test_msg = TwoPartMessage::from_data(Bytes::from(format!("message {}", i))); - bytes_tx.send(test_msg).await.unwrap(); - } - - // Close the sender to trigger normal termination - drop(bytes_tx); - - let result = handle_writer(framed_writer, bytes_rx, alive_rx, controller).await; - - assert!(result.is_ok()); - - // Decode from server side to verify all messages plus sentinel - let mut reader = FramedRead::new(server, TwoPartCodec::default()); - for i in 0..5 { - let msg = recv_msg(&mut reader).await; - assert_data_only_message(msg, format!("message {}", i).as_bytes()); - } - - let sentinel = recv_msg(&mut reader).await; - assert_sentinel_message(sentinel); - } - - /// Test that alive_rx is dropped after handle_writer completes - #[tokio::test] - async fn test_handle_writer_drops_alive_rx() { - let WriterHarness { - framed_writer, - bytes_tx, - bytes_rx, - alive_tx, - alive_rx, - controller, - .. - } = writer_harness().await; - - // Close the sender to trigger normal termination - drop(bytes_tx); - - let result = handle_writer(framed_writer, bytes_rx, alive_rx, controller).await; - - assert!(result.is_ok()); - - // alive_tx should now be closed because alive_rx was dropped - assert!(alive_tx.is_closed()); - } - - /// Test handle_writer with header-only messages (control messages) - #[tokio::test] - async fn test_handle_writer_header_only_messages() { - let WriterHarness { - server, - framed_writer, - bytes_tx, - bytes_rx, - alive_rx, - controller, - .. - } = writer_harness().await; - - // Send a header-only message - let header_msg = TwoPartMessage::from_header(Bytes::from("header content")); - bytes_tx.send(header_msg).await.unwrap(); - - // Close the sender - drop(bytes_tx); - - let result = handle_writer(framed_writer, bytes_rx, alive_rx, controller).await; - - assert!(result.is_ok()); - - let mut reader = FramedRead::new(server, TwoPartCodec::default()); - - let header_msg = recv_msg(&mut reader).await; - assert_header_only_message(header_msg, b"header content"); - - let sentinel = recv_msg(&mut reader).await; - assert_sentinel_message(sentinel); - } - - /// Test handle_writer with mixed header and data messages - #[tokio::test] - async fn test_handle_writer_mixed_messages() { - let WriterHarness { - server, - framed_writer, - bytes_tx, - bytes_rx, - alive_rx, - controller, - .. - } = writer_harness().await; - - // Send mixed messages - bytes_tx - .send(TwoPartMessage::from_header(Bytes::from("header1"))) - .await - .unwrap(); - bytes_tx - .send(TwoPartMessage::from_data(Bytes::from("data1"))) - .await - .unwrap(); - bytes_tx - .send(TwoPartMessage::from_parts( - Bytes::from("header2"), - Bytes::from("data2"), - )) - .await - .unwrap(); - - // Close the sender - drop(bytes_tx); - - let result = handle_writer(framed_writer, bytes_rx, alive_rx, controller).await; - - assert!(result.is_ok()); - - let mut reader = FramedRead::new(server, TwoPartCodec::default()); - - let first = recv_msg(&mut reader).await; - assert_header_only_message(first, b"header1"); - - let second = recv_msg(&mut reader).await; - assert_data_only_message(second, b"data1"); - - let third = recv_msg(&mut reader).await; - assert_header_and_data_message(third, b"header2", b"data2"); - - let sentinel = recv_msg(&mut reader).await; - assert_sentinel_message(sentinel); - } - - /// Killed or stopped contexts skip the server FIN deadline. - #[tokio::test] - async fn test_wait_for_server_shutdown_skips_terminal_context() { - for action in [Controller::kill as fn(&Controller), Controller::stop] { - let (client, _server) = create_tcp_pair().await; - let controller = Arc::new(Controller::default()); - action(&controller); - - let context: Arc = controller; - let result = tokio::time::timeout( - std::time::Duration::from_millis(50), - wait_for_server_shutdown(client, context), - ) - .await; - - assert!(result.is_ok(), "terminal context should not wait for FIN"); - assert!( - result.unwrap().is_ok(), - "terminal context shutdown should succeed" - ); - } - } - - /// Read error in the connection monitor kills the context and skips the FIN wait. - #[tokio::test] - async fn test_connection_monitor_skips_fin_wait_after_read_error_kills_context() { - let (client, mut server) = create_tcp_pair().await; - let (read_half, write_half) = tokio::io::split(client); - let framed_reader = FramedRead::new(read_half, TwoPartCodec::default()); - let framed_writer = FramedWrite::new(write_half, TwoPartCodec::default()); - let (_bytes_tx, bytes_rx) = mpsc::channel(64); - let (alive_tx, alive_rx) = oneshot::channel::<()>(); - let controller = Arc::new(Controller::default()); - - let reader_context = controller.clone(); - let reader_task = tokio::spawn(async move { - handle_reader(framed_reader, reader_context, alive_tx, None).await - }); - let writer_context = controller.clone(); - let writer_task = tokio::spawn(async move { - handle_writer(framed_writer, bytes_rx, alive_rx, writer_context).await - }); - - // Bypass the codec and write a complete but invalid TwoPartCodec - // header. This drives the client reader into Some(Err(_)) without - // closing the server side of the socket. - server.write_all(&[0xFF; 24]).await.unwrap(); - - let monitor_context: Arc = controller.clone(); - let result = tokio::time::timeout( - std::time::Duration::from_millis(250), - wait_for_connection_tasks( - reader_task, - writer_task, - monitor_context, - None, - "test-subject".to_string(), - None, - ), - ) - .await; - - assert!( - result.is_ok(), - "connection monitor should not wait for the FIN deadline after read error" - ); - assert!(result.unwrap().is_ok(), "connection monitor should succeed"); - assert!( - controller.is_killed(), - "read error should kill the stream context" - ); - } - - /// Reader-side panic must abort the writer and return promptly rather than - /// hanging on `tokio::join!`. Locks in the fix added with this function's - /// sequential-await + writer-abort behavior. - /// - /// Setup: spawn a reader task that panics immediately (so - /// `reader_task.await` yields `Err(JoinError::panic)`), and a writer task - /// that parks indefinitely waiting for application bytes (so without the - /// abort, `tokio::join!` on the previous implementation would never wake). - /// Expect: `wait_for_connection_tasks` returns Err within the timeout. - #[tokio::test] - async fn test_connection_monitor_aborts_writer_when_reader_panics() { - // Reader task that panics immediately. The explicit JoinHandle type - // pins the inferred return type to the one wait_for_connection_tasks - // expects; `panic!` is type `!`, which coerces to that type. - let reader_task: tokio::task::JoinHandle< - FramedRead, TwoPartCodec>, - > = tokio::spawn(async { - panic!("simulated reader panic to trigger JoinError"); - }); - - // Writer task that would block indefinitely waiting on application - // bytes. Under the pre-fix `tokio::join!` implementation, this would - // prevent the function from returning when the reader panicked. - // After the fix, the abort drives this task to completion promptly. - let writer_task: tokio::task::JoinHandle< - Result, TwoPartCodec>>, - > = tokio::spawn(async { - std::future::pending::<()>().await; - unreachable!() - }); - - let controller = Arc::new(Controller::default()); - let context: Arc = controller.clone(); - - // 250 ms is generous — the abort + JoinHandle resolution should fire - // sub-millisecond. We are checking for "doesn't hang", not "fast". - let result = tokio::time::timeout( - std::time::Duration::from_millis(250), - wait_for_connection_tasks( - reader_task, - writer_task, - context, - None, - "test-reader-panic".to_string(), - None, - ), - ) - .await; - - // Outer timeout must not fire: the abort path must surface the reader - // JoinError before the writer would have produced any bytes. - assert!( - result.is_ok(), - "wait_for_connection_tasks must return after reader panic, \ - not hang waiting on the writer" - ); - - // The inner result must be Err — the reader's JoinError propagates. - assert!( - result.unwrap().is_err(), - "reader panic should propagate as Err from wait_for_connection_tasks" - ); - } - - // ==================== handle_reader tests ==================== - - struct ReaderHarness { - framed_server: FramedWrite, TwoPartCodec>, - framed_reader: FramedRead, TwoPartCodec>, - alive_tx: oneshot::Sender<()>, - alive_rx: oneshot::Receiver<()>, - controller: Arc, - } - - /// Creates a reusable reader harness with paired TCP streams and test channels. - async fn reader_harness() -> ReaderHarness { - let (client, server) = create_tcp_pair().await; - let (read_half, _write_half) = tokio::io::split(client); - let (_server_read, server_write) = tokio::io::split(server); - - let framed_reader = FramedRead::new(read_half, TwoPartCodec::default()); - let framed_server = FramedWrite::new(server_write, TwoPartCodec::default()); - let (alive_tx, alive_rx) = oneshot::channel::<()>(); - let controller = Arc::new(Controller::default()); - - ReaderHarness { - framed_server, - framed_reader, - alive_tx, - alive_rx, - controller, - } - } + use tokio::io::AsyncWriteExt; + use tokio::sync::mpsc; fn control_message(msg: &ControlMessage) -> TwoPartMessage { - let msg_bytes = serde_json::to_vec(msg).unwrap(); - TwoPartMessage::from_header(Bytes::from(msg_bytes)) - } - - /// Test that handle_reader handles Stop control message by calling context.stop() - #[tokio::test] - async fn test_handle_reader_stop_control_message() { - let ReaderHarness { - mut framed_server, - framed_reader, - alive_tx, - alive_rx: _alive_rx, - controller, - } = reader_harness().await; - - // Spawn the reader task - let controller_clone = controller.clone(); - let reader_handle = tokio::spawn(async move { - handle_reader(framed_reader, controller_clone, alive_tx, None).await - }); - - // Send Stop control message from server - framed_server - .send(control_message(&ControlMessage::Stop)) - .await - .unwrap(); - - // Close the framed server to signal EOF to the client - framed_server.close().await.unwrap(); - - // Wait for reader to finish - let _ = reader_handle.await.unwrap(); - - // Verify that stop was called on the controller - assert!( - controller.is_stopped(), - "Controller should be stopped after receiving Stop message" - ); - } - - /// Test that handle_reader handles Kill control message by calling context.kill() - #[tokio::test] - async fn test_handle_reader_kill_control_message() { - let ReaderHarness { - mut framed_server, - framed_reader, - alive_tx, - alive_rx: _alive_rx, - controller, - } = reader_harness().await; - - // Spawn the reader task - let controller_clone = controller.clone(); - let reader_handle = tokio::spawn(async move { - handle_reader(framed_reader, controller_clone, alive_tx, None).await - }); - - // Send Kill control message from server - framed_server - .send(control_message(&ControlMessage::Kill)) - .await - .unwrap(); - - // Close the framed server to signal EOF to the client - framed_server.close().await.unwrap(); - - // Wait for reader to finish - let _ = reader_handle.await.unwrap(); - - // Verify that kill was called on the controller - assert!( - controller.is_killed(), - "Controller should be killed after receiving Kill message" - ); - } - - /// Test that handle_reader exits when alive channel is closed - #[tokio::test] - async fn test_handle_reader_exits_on_alive_channel_closed() { - let ReaderHarness { - framed_reader, - alive_tx, - alive_rx, - controller, - .. - } = reader_harness().await; - - // Spawn the reader task - let reader_handle = - tokio::spawn( - async move { handle_reader(framed_reader, controller, alive_tx, None).await }, - ); - - // Drop the alive_rx to close the channel (simulating writer finishing) - drop(alive_rx); - - // Reader should exit due to alive channel closure - let result = reader_handle.await; - - assert!( - result.is_ok(), - "handle_reader should exit when alive channel is closed" - ); - } - - /// Test that handle_reader exits when TCP stream is closed - #[tokio::test] - async fn test_handle_reader_exits_on_stream_closed() { - let ReaderHarness { - mut framed_server, - framed_reader, - alive_tx, - alive_rx: _alive_rx, - controller, - } = reader_harness().await; - - // Spawn the reader task - let reader_handle = - tokio::spawn( - async move { handle_reader(framed_reader, controller, alive_tx, None).await }, - ); - - // Close the framed server to signal EOF to the client - framed_server.close().await.unwrap(); - - // Reader should exit due to stream closure - let result = tokio::time::timeout(std::time::Duration::from_secs(1), reader_handle).await; - - assert!( - result.is_ok(), - "handle_reader should exit when stream is closed" - ); - } - - /// Test that handle_reader handles multiple control messages in sequence - #[tokio::test] - async fn test_handle_reader_multiple_control_messages() { - let ReaderHarness { - mut framed_server, - framed_reader, - alive_tx, - alive_rx: _alive_rx, - controller, - } = reader_harness().await; - - // Spawn the reader task - let controller_clone = controller.clone(); - let reader_handle = tokio::spawn(async move { - handle_reader(framed_reader, controller_clone, alive_tx, None).await - }); - - // Send multiple Stop messages (first one will stop, subsequent ones are no-ops) - framed_server - .send(control_message(&ControlMessage::Stop)) - .await - .unwrap(); - framed_server - .send(control_message(&ControlMessage::Stop)) - .await - .unwrap(); - - // Close the framed server to signal EOF to the client - framed_server.close().await.unwrap(); - - // Wait for reader to finish - let _ = reader_handle.await.unwrap(); - - // Verify that stop was called - assert!( - controller.is_stopped(), - "Controller should be stopped after receiving Stop messages" - ); - } - - /// Test handle_reader with Stop followed by Kill - #[tokio::test] - async fn test_handle_reader_stop_then_kill() { - let ReaderHarness { - mut framed_server, - framed_reader, - alive_tx, - alive_rx: _alive_rx, - controller, - } = reader_harness().await; - - // Spawn the reader task - let controller_clone = controller.clone(); - let reader_handle = tokio::spawn(async move { - handle_reader(framed_reader, controller_clone, alive_tx, None).await - }); - - // Send Stop first, then Kill - framed_server - .send(control_message(&ControlMessage::Stop)) - .await - .unwrap(); - framed_server - .send(control_message(&ControlMessage::Kill)) - .await - .unwrap(); - - // Close the framed server to signal EOF to the client - framed_server.close().await.unwrap(); - - // Wait for reader to finish - let _ = reader_handle.await.unwrap(); - - // Verify that kill was called (which sets killed state) - assert!( - controller.is_killed(), - "Controller should be killed after receiving Kill message" - ); - } - - /// Read errors kill the context and are counted as cancellations. - #[tokio::test] - async fn test_handle_reader_increments_cancellation_counter_on_read_error() { - let ReaderHarness { - framed_server, - framed_reader, - alive_tx, - alive_rx: _alive_rx, - controller, - } = reader_harness().await; - let cancellation_counter = IntCounter::new( - "tcp_client_reader_read_error_cancellations_test", - "test cancellation counter", - ) - .unwrap(); - - let counter_clone = cancellation_counter.clone(); - let controller_clone = controller.clone(); - let reader_handle = tokio::spawn(async move { - handle_reader( - framed_reader, - controller_clone, - alive_tx, - Some(counter_clone), - ) - .await - }); - - let mut raw_writer = framed_server.into_inner(); - raw_writer.write_all(&[0u8; 8]).await.unwrap(); - raw_writer.shutdown().await.unwrap(); - - let _ = reader_handle.await.unwrap(); - - assert!( - controller.is_killed(), - "Controller should be killed after TCP stream read error" - ); - assert_eq!( - cancellation_counter.get(), - 1, - "read-error close should increment cancellation metric once" - ); - } - - /// Drives `handle_reader` against a single message and returns the - /// controller + cancellation counter for assertions. - async fn run_reader_with( - msg: TwoPartMessage, - counter_name: &str, - ) -> (Arc, IntCounter) { - let ReaderHarness { - mut framed_server, - framed_reader, - alive_tx, - alive_rx: _alive_rx, - controller, - } = reader_harness().await; - let counter = IntCounter::new(counter_name, "test counter").unwrap(); - - let counter_clone = counter.clone(); - let controller_clone = controller.clone(); - let reader_handle = tokio::spawn(async move { - handle_reader( - framed_reader, - controller_clone, - alive_tx, - Some(counter_clone), - ) - .await - }); - - framed_server.send(msg).await.unwrap(); - let _ = reader_handle.await.unwrap(); - - (controller, counter) - } - - /// Each protocol-violating message variant must kill only this stream - /// (controller killed, cancellation counted once) and never panic the - /// worker. Covers the three non-read-error panic arms in `handle_reader`: - /// undecodable control bytes, server-sent Sentinel, and non-control - /// (data-only) messages. - #[tokio::test] - async fn test_handle_reader_kills_on_protocol_violations() { - let cases: Vec<(&str, TwoPartMessage)> = vec![ - ( - "invalid control bytes", - TwoPartMessage::from_header(Bytes::from_static(b"not a valid control message")), - ), - ( - "sentinel from server", - control_message(&ControlMessage::Sentinel), - ), - ( - "non-control (data-only)", - TwoPartMessage::from_data(Bytes::from_static(b"unexpected payload")), - ), - ]; - - for (i, (label, msg)) in cases.into_iter().enumerate() { - let counter_name = format!("tcp_client_reader_protocol_violation_test_{i}"); - let (controller, counter) = run_reader_with(msg, &counter_name).await; - assert!( - controller.is_killed(), - "{label}: should kill stream context" - ); - assert_eq!(counter.get(), 1, "{label}: should be counted once"); - } + TwoPartMessage::from_header(serde_json::to_vec(msg).unwrap().into()) } // ==================== handle_request_reader tests ==================== diff --git a/lib/runtime/src/pipeline/network/tcp/mux.rs b/lib/runtime/src/pipeline/network/tcp/mux.rs index 6ca5a12c91bf..52b72d57157a 100644 --- a/lib/runtime/src/pipeline/network/tcp/mux.rs +++ b/lib/runtime/src/pipeline/network/tcp/mux.rs @@ -30,7 +30,7 @@ pub const RESPONSE_MUX_STREAM_WRITER_QUEUE: usize = 8; pub const RESPONSE_MUX_IDLE_TTL_SECS: u64 = 300; pub const RESPONSE_MUX_CONNECT_TIMEOUT_SECS: u64 = 5; -pub const RESPONSE_MUX_DEFAULT_BATCH_INTERVAL_MS: u64 = 5; +pub const RESPONSE_MUX_DEFAULT_BATCH_INTERVAL_MS: u64 = 1; pub const RESPONSE_MUX_MAX_BATCH_INTERVAL_MS: u64 = 100; pub const RESPONSE_MUX_DEFAULT_BATCH_MAX_BYTES: usize = 65_536; pub const RESPONSE_MUX_DEFAULT_BATCH_MAX_FRAMES: usize = 64; @@ -42,7 +42,6 @@ pub const RESPONSE_MUX_SCHEDULER_QUANTUM: usize = 8; #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub struct ResponseMuxConfig { - pub enabled: bool, pub packet_metrics: bool, pub batch_interval: Duration, pub batch_max_bytes: usize, @@ -77,7 +76,6 @@ impl ResponseMuxConfig { Some("1") | Some("true") => Ok(true), Some(value) => anyhow::bail!("invalid {name}={value:?}; expected 0, 1, false, or true"), }; - let enabled = parse_bool(env::DYN_TCP_RESPONSE_MUX, read(env::DYN_TCP_RESPONSE_MUX))?; let packet_metrics = parse_bool( env::DYN_TCP_RESPONSE_PACKET_METRICS, read(env::DYN_TCP_RESPONSE_PACKET_METRICS), @@ -146,7 +144,6 @@ impl ResponseMuxConfig { } } Ok(Self { - enabled, packet_metrics, batch_interval: Duration::from_millis(interval_ms), batch_max_bytes, @@ -184,8 +181,6 @@ pub const MUX_HEADER_LEN: usize = 24; #[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)] #[serde(tag = "kind", rename_all = "snake_case")] pub enum ConnectionHandshake { - /// Dedicated per-request upstream -> downstream request stream. - RequestStream { subject: String }, /// Persistent connection carrying many downstream -> upstream responses. ResponseMux { version: u8, @@ -658,11 +653,10 @@ mod tests { } #[test] - fn response_mux_config_defaults_to_disabled_and_five_ms() { + fn response_mux_config_defaults_to_one_ms() { let config = config(&[]).unwrap(); - assert!(!config.enabled); assert!(!config.packet_metrics); - assert_eq!(config.batch_interval, Duration::from_millis(5)); + assert_eq!(config.batch_interval, Duration::from_millis(1)); assert_eq!(config.batch_max_bytes, 65_536); assert_eq!(config.batch_max_frames, 64); assert_eq!(config.stream_window_bytes, 262_144); @@ -673,14 +667,12 @@ mod tests { fn response_mux_config_accepts_zero_delay_and_valid_overrides() { use crate::config::environment_names::tcp_response_stream as env; let config = config(&[ - (env::DYN_TCP_RESPONSE_MUX, "1"), (env::DYN_TCP_RESPONSE_PACKET_METRICS, "true"), (env::DYN_TCP_RESPONSE_BATCH_INTERVAL_MS, "0"), (env::DYN_TCP_RESPONSE_BATCH_MAX_BYTES, "8192"), (env::DYN_TCP_RESPONSE_BATCH_MAX_FRAMES, "8"), ]) .unwrap(); - assert!(config.enabled); assert!(config.packet_metrics); assert_eq!(config.batch_interval, Duration::ZERO); assert_eq!(config.batch_max_bytes, 8192); diff --git a/lib/runtime/src/pipeline/network/tcp/mux/client.rs b/lib/runtime/src/pipeline/network/tcp/mux/client.rs index 810cd48e87a1..75ce24695066 100644 --- a/lib/runtime/src/pipeline/network/tcp/mux/client.rs +++ b/lib/runtime/src/pipeline/network/tcp/mux/client.rs @@ -4,7 +4,7 @@ //! Worker-side persistent multiplexed TCP response connection pool. use std::{ - collections::VecDeque, + collections::{HashMap, VecDeque}, sync::{ Arc, Weak, atomic::{AtomicBool, AtomicU64, AtomicUsize, Ordering}, @@ -48,24 +48,32 @@ struct WriterCommand { _writer_permit: Option, _queued_byte_permit: Option, priority_enqueued_at: Option, - enqueued_at: Instant, + enqueued_at: Option, } impl WriterCommand { fn new(frame: MuxFrame, written: Option>>) -> Self { + Self::new_with_metrics(frame, written, per_frame_metrics_enabled()) + } + + fn new_with_metrics( + frame: MuxFrame, + written: Option>>, + metrics_enabled: bool, + ) -> Self { Self { frame, written, _writer_permit: None, _queued_byte_permit: None, priority_enqueued_at: None, - enqueued_at: Instant::now(), + enqueued_at: metrics_enabled.then(Instant::now), } } fn priority(frame: MuxFrame, written: Option>>) -> Self { let mut command = Self::new(frame, written); - command.priority_enqueued_at = Some(Instant::now()); + command.priority_enqueued_at = per_frame_metrics_enabled().then(Instant::now); command } @@ -88,7 +96,7 @@ impl WriterCommand { #[inline] fn per_frame_metrics_enabled() -> bool { - true + super::response_packet_metrics_enabled() } #[derive(Clone, Copy)] @@ -124,18 +132,13 @@ impl PoolConfig { } } -#[derive(Default)] -struct StreamWriterState { - pending: VecDeque, - scheduled: bool, -} - struct WorkerStreamState { context: Arc, + cancellation_counter: Option, + cancellation_recorded: AtomicBool, credits: Arc, max_credits: usize, writer_slots: Arc, - writer: Mutex, closed: AtomicBool, close_token: CancellationToken, } @@ -143,7 +146,40 @@ struct WorkerStreamState { type ScheduledStream = (Uuid, Arc); type BlockedData = (WriterCommand, Option, Instant); +struct StreamIngress { + stream_id: Uuid, + state: Arc, + command: WriterCommand, +} + +struct WriterStreamQueue { + state: Arc, + pending: VecDeque, + scheduled: bool, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum IngressPoll { + Ingested, + Empty, + Disconnected, +} + +#[derive(Default)] +struct WriterScheduler { + streams: HashMap, + ready: VecDeque, +} + impl WorkerStreamState { + fn record_cancellation(&self) { + if !self.cancellation_recorded.swap(true, Ordering::AcqRel) + && let Some(counter) = &self.cancellation_counter + { + counter.inc(); + } + } + fn replenish_credits(&self, credits: usize) -> usize { if self.closed.load(Ordering::Acquire) || self.credits.is_closed() { return 0; @@ -157,11 +193,175 @@ impl WorkerStreamState { } } +impl WriterScheduler { + fn account_removed(connection: &MuxConnection, command: &WriterCommand) { + connection.queued_frames.fetch_sub(1, Ordering::AcqRel); + let encoded_len = command.frame.encoded_len(); + connection + .queued_bytes + .fetch_sub(encoded_len, Ordering::AcqRel); + if per_frame_metrics_enabled() { + response_mux::QUEUED_BYTES + .with_label_values(&["worker"]) + .sub(encoded_len as i64); + } + } + + fn fail_commands( + connection: &MuxConnection, + commands: impl IntoIterator, + reason: &str, + ) { + for command in commands { + Self::account_removed(connection, &command); + command.fail(reason); + } + } + + fn ingest(&mut self, connection: &MuxConnection, ingress: StreamIngress) { + let StreamIngress { + stream_id, + state, + command, + } = ingress; + if state.closed.load(Ordering::Acquire) { + Self::account_removed(connection, &command); + command.fail("response mux stream is closed"); + return; + } + + let queue = self + .streams + .entry(stream_id) + .or_insert_with(|| WriterStreamQueue { + state, + pending: VecDeque::new(), + scheduled: false, + }); + queue.pending.push_back(command); + if per_frame_metrics_enabled() { + response_mux::STREAM_WRITER_QUEUE_OCCUPANCY.observe(queue.pending.len() as f64); + } + if !queue.scheduled { + queue.scheduled = true; + self.ready.push_back(stream_id); + if per_frame_metrics_enabled() { + response_mux::READY_STREAMS + .with_label_values(&["worker"]) + .inc(); + } + } + } + + fn ingest_one( + &mut self, + connection: &MuxConnection, + stream_rx: &mut mpsc::Receiver, + ) -> IngressPoll { + match stream_rx.try_recv() { + Ok(ingress) => { + self.ingest(connection, ingress); + IngressPoll::Ingested + } + Err(mpsc::error::TryRecvError::Empty) => IngressPoll::Empty, + Err(mpsc::error::TryRecvError::Disconnected) => IngressPoll::Disconnected, + } + } + + fn pop_ready( + &mut self, + connection: &MuxConnection, + ) -> Option<(WriterCommand, ScheduledStream)> { + while let Some(stream_id) = self.ready.pop_front() { + let Some(queue) = self.streams.get_mut(&stream_id) else { + continue; + }; + if queue.state.closed.load(Ordering::Acquire) { + let pending = queue.pending.drain(..).collect::>(); + if queue.scheduled { + queue.scheduled = false; + if per_frame_metrics_enabled() { + response_mux::READY_STREAMS + .with_label_values(&["worker"]) + .dec(); + } + } + Self::fail_commands(connection, pending, "response mux stream is closed"); + continue; + } + if let Some(command) = queue.pending.pop_front() { + Self::account_removed(connection, &command); + return Some((command, (stream_id, queue.state.clone()))); + } + if queue.scheduled { + queue.scheduled = false; + if per_frame_metrics_enabled() { + response_mux::READY_STREAMS + .with_label_values(&["worker"]) + .dec(); + } + } + } + None + } + + fn pop_same_stream( + &mut self, + connection: &MuxConnection, + stream_id: Uuid, + ) -> Option { + let queue = self.streams.get_mut(&stream_id)?; + if queue.state.closed.load(Ordering::Acquire) { + let pending = queue.pending.drain(..).collect::>(); + Self::fail_commands(connection, pending, "response mux stream is closed"); + return None; + } + let command = queue.pending.pop_front()?; + Self::account_removed(connection, &command); + Some(command) + } + + fn reschedule(&mut self, connection: &MuxConnection, stream_id: Uuid) { + if per_frame_metrics_enabled() { + response_mux::ROUND_ROBIN_TURNS_TOTAL.inc(); + } + let Some(queue) = self.streams.get_mut(&stream_id) else { + return; + }; + if queue.state.closed.load(Ordering::Acquire) { + let pending = queue.pending.drain(..).collect::>(); + Self::fail_commands(connection, pending, "response mux stream is closed"); + } + if !queue.state.closed.load(Ordering::Acquire) && !queue.pending.is_empty() { + self.ready.push_back(stream_id); + } else if queue.scheduled { + queue.scheduled = false; + if per_frame_metrics_enabled() { + response_mux::READY_STREAMS + .with_label_values(&["worker"]) + .dec(); + } + } + } + + fn fail_all(&mut self, connection: &MuxConnection, reason: &str) { + for (_, mut queue) in self.streams.drain() { + if queue.scheduled && per_frame_metrics_enabled() { + response_mux::READY_STREAMS + .with_label_values(&["worker"]) + .dec(); + } + Self::fail_commands(connection, queue.pending.drain(..), reason); + } + self.ready.clear(); + } +} + struct MuxConnection { id: u64, cancel: CancellationToken, priority_tx: mpsc::Sender, - ready_tx: mpsc::UnboundedSender, + stream_tx: mpsc::Sender, streams: DashMap>, healthy: AtomicBool, active_streams: AtomicUsize, @@ -219,17 +419,17 @@ impl MuxConnection { if ack.kind != MuxFrameKind::ConnectionAck || ack.connection_ack_offset()? != 0 { anyhow::bail!("frontend returned invalid response mux connection ack"); } - let read_half = handshake_reader.into_inner(); + let mux_reader = handshake_reader.map_decoder(|_| MuxCodec::default()); let write_half = handshake_writer.into_inner(); let (priority_tx, priority_rx) = mpsc::channel(config.writer_queue); - let (ready_tx, ready_rx) = mpsc::unbounded_channel(); + let (stream_tx, stream_rx) = mpsc::channel(config.writer_queue); let cancel = cancel.child_token(); let connection = Arc::new(Self { id, cancel: cancel.clone(), priority_tx, - ready_tx, + stream_tx, streams: DashMap::new(), healthy: AtomicBool::new(true), active_streams: AtomicUsize::new(0), @@ -256,12 +456,12 @@ impl MuxConnection { Arc::downgrade(&connection), write_half, priority_rx, - ready_rx, + stream_rx, cancel.clone(), )); tokio::spawn(Self::reader_task( Arc::downgrade(&connection), - FramedRead::new(read_half, MuxCodec::default()), + mux_reader, cancel, packet_baseline, )); @@ -336,34 +536,10 @@ impl MuxConnection { if state.closed.swap(true, Ordering::AcqRel) { return; } + tracing::trace!(reason, "closing response mux stream state"); state.credits.close(); state.writer_slots.close(); state.close_token.cancel(); - let (pending, was_scheduled) = { - let mut writer = state.writer.lock(); - let pending = writer.pending.drain(..).collect::>(); - let was_scheduled = writer.scheduled; - writer.scheduled = false; - (pending, was_scheduled) - }; - self.queued_frames - .fetch_sub(pending.len(), Ordering::AcqRel); - let pending_bytes = pending - .iter() - .map(|command| command.frame.encoded_len()) - .sum::(); - self.queued_bytes.fetch_sub(pending_bytes, Ordering::AcqRel); - response_mux::QUEUED_BYTES - .with_label_values(&["worker"]) - .sub(pending_bytes as i64); - if was_scheduled { - response_mux::READY_STREAMS - .with_label_values(&["worker"]) - .dec(); - } - for command in pending { - command.fail(reason); - } } fn remove_stream(&self, stream_id: Uuid, reason: &str, kill_context: bool) -> bool { @@ -427,10 +603,10 @@ impl MuxConnection { } } - fn enqueue_stream_command( + async fn enqueue_stream_command( &self, stream_id: Uuid, - state: &WorkerStreamState, + state: Arc, command: WriterCommand, ) -> Result<()> { if !self.is_healthy() || state.closed.load(Ordering::Acquire) { @@ -440,69 +616,42 @@ impl MuxConnection { let encoded_len = command.frame.encoded_len(); self.queued_bytes.fetch_add(encoded_len, Ordering::AcqRel); - response_mux::QUEUED_BYTES - .with_label_values(&["worker"]) - .add(encoded_len as i64); - - let mut writer = state.writer.lock(); - if state.closed.load(Ordering::Acquire) { - self.queued_bytes.fetch_sub(encoded_len, Ordering::AcqRel); + if per_frame_metrics_enabled() { response_mux::QUEUED_BYTES .with_label_values(&["worker"]) - .sub(encoded_len as i64); - command.fail("response mux stream is closed"); - anyhow::bail!("response mux stream is closed"); + .add(encoded_len as i64); } - writer.pending.push_back(command); + self.queued_frames.fetch_add(1, Ordering::AcqRel); - if per_frame_metrics_enabled() { - response_mux::STREAM_WRITER_QUEUE_OCCUPANCY.observe(writer.pending.len() as f64); - } - if !writer.scheduled { - writer.scheduled = true; - response_mux::READY_STREAMS - .with_label_values(&["worker"]) - .inc(); - if self.ready_tx.send(stream_id).is_err() { - writer.scheduled = false; - response_mux::READY_STREAMS - .with_label_values(&["worker"]) - .dec(); + match self + .stream_tx + .send(StreamIngress { + stream_id, + state, + command, + }) + .await + { + Ok(()) => Ok(()), + Err(err) => { self.queued_frames.fetch_sub(1, Ordering::AcqRel); - if let Some(command) = writer.pending.pop_back() { - self.queued_bytes.fetch_sub(encoded_len, Ordering::AcqRel); + self.queued_bytes.fetch_sub(encoded_len, Ordering::AcqRel); + if per_frame_metrics_enabled() { response_mux::QUEUED_BYTES .with_label_values(&["worker"]) .sub(encoded_len as i64); - command.fail("response mux fair writer stopped"); } - anyhow::bail!("response mux fair writer stopped"); + err.0.command.fail("response mux fair writer stopped"); + anyhow::bail!("response mux fair writer stopped") } } - Ok(()) - } - - fn reschedule_stream(&self, stream_id: Uuid, state: &Arc) -> Result<()> { - response_mux::ROUND_ROBIN_TURNS_TOTAL.inc(); - let mut writer = state.writer.lock(); - if !state.closed.load(Ordering::Acquire) && !writer.pending.is_empty() { - self.ready_tx - .send(stream_id) - .map_err(|_| anyhow!("response mux fair writer stopped"))?; - } else if writer.scheduled { - writer.scheduled = false; - response_mux::READY_STREAMS - .with_label_values(&["worker"]) - .dec(); - } - Ok(()) } async fn writer_task( weak: Weak, mut write_half: tokio::net::tcp::OwnedWriteHalf, mut priority_rx: mpsc::Receiver, - mut ready_rx: mpsc::UnboundedReceiver, + mut stream_rx: mpsc::Receiver, cancel: CancellationToken, ) { let mut write_buf = TcpWriteBuffer::new(); @@ -516,8 +665,11 @@ impl MuxConnection { let metrics_enabled = per_frame_metrics_enabled(); let mut reported_queue_depth = 0_i64; let mut blocked_data: Option = None; + let mut scheduler = WriterScheduler::default(); + let mut stream_input_open = true; + let result: Result<()> = async { - loop { + 'writer: loop { let connection = weak .upgrade() .ok_or_else(|| anyhow!("response mux connection dropped"))?; @@ -537,109 +689,104 @@ impl MuxConnection { if blocked_command.frame.kind != MuxFrameKind::Data { (blocked_command, blocked_stream, None) } else { - enum BlockedNext { - Priority(WriterCommand), - Credit(OwnedSemaphorePermit), - StreamClosed, - } - let blocked_close = blocked_stream - .as_ref() - .expect("Data commands are always stream-scheduled") - .1 - .close_token - .clone(); - let credits = connection.connection_credits.clone(); - let next = tokio::select! { - biased; - _ = cancel.cancelled() => return Ok(()), - _ = blocked_close.cancelled() => BlockedNext::StreamClosed, - Some(command) = priority_rx.recv() => BlockedNext::Priority(command), - permit = credits.acquire_many_owned( - blocked_command - .frame - .encoded_len() - .min(connection.max_connection_credits) as u32 - ) => BlockedNext::Credit( - permit.map_err(|_| anyhow!( - "response mux connection closed while writer awaited credits" - ))? - ), - }; - match next { - BlockedNext::Priority(command) => { - connection.queued_frames.fetch_sub(1, Ordering::AcqRel); - blocked_data = Some((blocked_command, blocked_stream, blocked_since)); - (command, None, None) + enum BlockedNext { + Priority(WriterCommand), + Credit(OwnedSemaphorePermit), + StreamClosed, } - BlockedNext::Credit(permit) => { - if metrics_enabled { - response_mux::CONNECTION_FLOW_CONTROL_STALL_SECONDS - .observe(blocked_since.elapsed().as_secs_f64()); + let blocked_close = blocked_stream + .as_ref() + .expect("Data commands are always stream-scheduled") + .1 + .close_token + .clone(); + let credits = connection.connection_credits.clone(); + let next = tokio::select! { + biased; + _ = cancel.cancelled() => return Ok(()), + _ = blocked_close.cancelled() => BlockedNext::StreamClosed, + Some(command) = priority_rx.recv() => BlockedNext::Priority(command), + permit = credits.acquire_many_owned( + blocked_command + .frame + .encoded_len() + .min(connection.max_connection_credits) as u32 + ) => BlockedNext::Credit( + permit.map_err(|_| anyhow!( + "response mux connection closed while writer awaited credits" + ))? + ), + }; + match next { + BlockedNext::Priority(command) => { + connection.queued_frames.fetch_sub(1, Ordering::AcqRel); + blocked_data = + Some((blocked_command, blocked_stream, blocked_since)); + (command, None, None) + } + BlockedNext::Credit(permit) => { + if metrics_enabled { + response_mux::CONNECTION_FLOW_CONTROL_STALL_SECONDS + .observe(blocked_since.elapsed().as_secs_f64()); + } + (blocked_command, blocked_stream, Some(permit)) + } + BlockedNext::StreamClosed => { + if let Some((stream_id, _)) = blocked_stream { + scheduler.reschedule(&connection, stream_id); + } + blocked_command.fail( + "response mux stream closed while writer awaited credits", + ); + continue 'writer; } - (blocked_command, blocked_stream, Some(permit)) - } - BlockedNext::StreamClosed => { - blocked_command - .fail("response mux stream closed while writer awaited credits"); - continue; } } - } } else { let (command, scheduled_stream) = loop { if let Ok(command) = priority_rx.try_recv() { connection.queued_frames.fetch_sub(1, Ordering::AcqRel); break (command, None); } + if let Some((command, scheduled_stream)) = scheduler.pop_ready(&connection) + { + break (command, Some(scheduled_stream)); + } + match scheduler.ingest_one(&connection, &mut stream_rx) { + IngressPoll::Ingested => continue, + IngressPoll::Empty => {} + IngressPoll::Disconnected => { + stream_input_open = false; + } + } enum Next { - Priority(WriterCommand), - Stream(Uuid), + Priority(Option), + Ingress(Option), } let next = tokio::select! { biased; _ = cancel.cancelled() => return Ok(()), - command = priority_rx.recv() => command.map(Next::Priority), - stream_id = ready_rx.recv() => stream_id.map(Next::Stream), - }; - let Some(next) = next else { - return Ok(()); + command = priority_rx.recv() => Next::Priority(command), + ingress = stream_rx.recv(), if stream_input_open => { + Next::Ingress(ingress) + } + else => return Ok(()), }; match next { - Next::Priority(command) => { + Next::Priority(Some(command)) => { connection.queued_frames.fetch_sub(1, Ordering::AcqRel); break (command, None); } - Next::Stream(stream_id) => { - let Some(state) = connection - .streams - .get(&stream_id) - .map(|entry| entry.value().clone()) - else { - continue; - }; - let command = state.writer.lock().pending.pop_front(); - if let Some(command) = command { - connection.queued_frames.fetch_sub(1, Ordering::AcqRel); - let encoded_len = command.frame.encoded_len(); - connection - .queued_bytes - .fetch_sub(encoded_len, Ordering::AcqRel); - response_mux::QUEUED_BYTES - .with_label_values(&["worker"]) - .sub(encoded_len as i64); - break (command, Some((stream_id, state))); - } - let mut writer = state.writer.lock(); - if writer.scheduled { - writer.scheduled = false; - response_mux::READY_STREAMS - .with_label_values(&["worker"]) - .dec(); - } + Next::Priority(None) if !stream_input_open => return Ok(()), + Next::Priority(None) => {} + Next::Ingress(Some(ingress)) => { + scheduler.ingest(&connection, ingress); } + Next::Ingress(None) => stream_input_open = false, } }; + if command.frame.kind == MuxFrameKind::Data { let required = command .frame @@ -653,8 +800,9 @@ impl MuxConnection { { Ok(permit) => (command, scheduled_stream, Some(permit)), Err(tokio::sync::TryAcquireError::NoPermits) => { + debug_assert!(blocked_data.is_none()); blocked_data = Some((command, scheduled_stream, Instant::now())); - continue; + continue 'writer; } Err(tokio::sync::TryAcquireError::Closed) => { return Err(anyhow!( @@ -672,8 +820,9 @@ impl MuxConnection { let mut batch = vec![(command, connection_permit)]; let mut batch_bytes = batch[0].0.frame.encoded_len(); let mut force_flush = !first_is_data; + let mut held_for_next_turn = false; + if let Some((stream_id, state)) = scheduled_stream { - let mut held_for_next_turn = false; for _ in 1..super::RESPONSE_MUX_SCHEDULER_QUANTUM { if force_flush || batch.len() >= connection.batch_max_frames @@ -681,19 +830,14 @@ impl MuxConnection { { break; } - let Some(next) = state.writer.lock().pending.pop_front() else { + let Some(next) = scheduler.pop_same_stream(&connection, stream_id) else { break; }; - connection.queued_frames.fetch_sub(1, Ordering::AcqRel); let encoded_len = next.frame.encoded_len(); - connection - .queued_bytes - .fetch_sub(encoded_len, Ordering::AcqRel); - response_mux::QUEUED_BYTES - .with_label_values(&["worker"]) - .sub(encoded_len as i64); if batch_bytes.saturating_add(encoded_len) > connection.batch_max_bytes { - blocked_data = Some((next, Some((stream_id, state.clone())), Instant::now())); + debug_assert!(blocked_data.is_none()); + blocked_data = + Some((next, Some((stream_id, state.clone())), Instant::now())); held_for_next_turn = true; break; } @@ -707,6 +851,7 @@ impl MuxConnection { { Ok(permit) => Some(permit), Err(tokio::sync::TryAcquireError::NoPermits) => { + debug_assert!(blocked_data.is_none()); blocked_data = Some(( next, Some((stream_id, state.clone())), @@ -729,13 +874,15 @@ impl MuxConnection { batch.push((next, permit)); } if !held_for_next_turn { - connection.reschedule_stream(stream_id, &state)?; + scheduler.reschedule(&connection, stream_id); } } - let deadline = batching_started + connection.batch_interval; + let deadline = batching_started + connection.batch_interval; while first_is_data && !force_flush + && !held_for_next_turn + && blocked_data.is_none() && batch.len() < connection.batch_max_frames && batch_bytes < connection.batch_max_bytes { @@ -746,56 +893,61 @@ impl MuxConnection { break; } - let next_stream = match ready_rx.try_recv() { - Ok(stream_id) => Some(stream_id), - Err(mpsc::error::TryRecvError::Disconnected) => return Ok(()), - Err(mpsc::error::TryRecvError::Empty) - if connection.batch_interval.is_zero() => - { - None + let mut scheduled = scheduler.pop_ready(&connection); + if scheduled.is_none() && stream_input_open { + match scheduler.ingest_one(&connection, &mut stream_rx) { + IngressPoll::Ingested => continue, + IngressPoll::Empty => {} + IngressPoll::Disconnected => { + stream_input_open = false; + } } - Err(mpsc::error::TryRecvError::Empty) => { - tokio::select! { - biased; - _ = cancel.cancelled() => return Ok(()), - Some(priority) = priority_rx.recv() => { - connection.queued_frames.fetch_sub(1, Ordering::AcqRel); - batch_bytes = batch_bytes.saturating_add(priority.frame.encoded_len()); - batch.push((priority, None)); - break; - } - stream_id = ready_rx.recv() => stream_id, - _ = tokio::time::sleep_until(deadline.into()) => None, + } + if scheduled.is_none() && !connection.batch_interval.is_zero() { + enum BatchNext { + Priority(WriterCommand), + Ingress(Option), + Deadline, + } + let next = tokio::select! { + biased; + _ = cancel.cancelled() => return Ok(()), + Some(priority) = priority_rx.recv() => { + BatchNext::Priority(priority) + } + ingress = stream_rx.recv(), if stream_input_open => { + BatchNext::Ingress(ingress) + } + _ = tokio::time::sleep_until(deadline.into()) => { + BatchNext::Deadline + } + }; + match next { + BatchNext::Priority(priority) => { + connection.queued_frames.fetch_sub(1, Ordering::AcqRel); + batch_bytes = + batch_bytes.saturating_add(priority.frame.encoded_len()); + batch.push((priority, None)); + break; + } + BatchNext::Ingress(Some(ingress)) => { + scheduler.ingest(&connection, ingress); + continue; } + BatchNext::Ingress(None) => { + stream_input_open = false; + continue; + } + BatchNext::Deadline => {} } - }; - let Some(stream_id) = next_stream else { + scheduled = scheduler.pop_ready(&connection); + } + let Some((next, (stream_id, state))) = scheduled else { break; }; - let Some(state) = connection - .streams - .get(&stream_id) - .map(|entry| entry.value().clone()) - else { - continue; - }; - let Some(next) = state.writer.lock().pending.pop_front() else { - connection.reschedule_stream(stream_id, &state)?; - continue; - }; - connection.queued_frames.fetch_sub(1, Ordering::AcqRel); let encoded_len = next.frame.encoded_len(); - connection - .queued_bytes - .fetch_sub(encoded_len, Ordering::AcqRel); - response_mux::QUEUED_BYTES - .with_label_values(&["worker"]) - .sub(encoded_len as i64); - if !batch.is_empty() - && (batch.len() + 1 > connection.batch_max_frames - || batch_bytes.saturating_add(encoded_len) - > connection.batch_max_bytes) - { + if batch_bytes.saturating_add(encoded_len) > connection.batch_max_bytes { + debug_assert!(blocked_data.is_none()); blocked_data = Some((next, Some((stream_id, state)), Instant::now())); break; } @@ -808,12 +960,15 @@ impl MuxConnection { { Ok(permit) => Some(permit), Err(tokio::sync::TryAcquireError::NoPermits) => { + debug_assert!(blocked_data.is_none()); blocked_data = Some((next, Some((stream_id, state)), Instant::now())); break; } Err(tokio::sync::TryAcquireError::Closed) => { - return Err(anyhow!("response mux connection credit window closed")); + return Err(anyhow!( + "response mux connection credit window closed" + )); } } } else { @@ -822,7 +977,7 @@ impl MuxConnection { let urgent = next.frame.kind != MuxFrameKind::Data; batch_bytes = batch_bytes.saturating_add(encoded_len); batch.push((next, permit)); - connection.reschedule_stream(stream_id, &state)?; + scheduler.reschedule(&connection, stream_id); if urgent { break; } @@ -842,18 +997,21 @@ impl MuxConnection { connection .sent_data_bytes .fetch_add(data_bytes, Ordering::AcqRel); - let mut write_calls = 0_u64; - let write_result: Result<()> = async { - let (_, calls) = write_buf.write_all_counted(&mut write_half).await?; - write_calls = calls; - Ok(()) - } - .await; + let write_result = write_buf.write_all_counted(&mut write_half).await; + let write_calls = write_result + .as_ref() + .map(|(_, calls)| *calls) + .unwrap_or_default(); + let write_result: Result<()> = + write_result.map(|_| ()).map_err(anyhow::Error::from); + for (command, permit) in &mut batch { if metrics_enabled { - response_mux::QUEUE_RESIDENCE_SECONDS - .with_label_values(&["worker"]) - .observe(command.enqueued_at.elapsed().as_secs_f64()); + if let Some(enqueued_at) = command.enqueued_at { + response_mux::QUEUE_RESIDENCE_SECONDS + .with_label_values(&["worker"]) + .observe(enqueued_at.elapsed().as_secs_f64()); + } if let Some(enqueued_at) = command.priority_enqueued_at { response_mux::PRIORITY_QUEUE_RESIDENCE_SECONDS .observe(enqueued_at.elapsed().as_secs_f64()); @@ -872,6 +1030,9 @@ impl MuxConnection { permit.forget(); } } + response_mux::WRITE_CALLS_TOTAL + .with_label_values(&["worker"]) + .inc_by(write_calls); write_result?; if metrics_enabled { frames_per_write.observe(batch.len() as f64); @@ -881,18 +1042,26 @@ impl MuxConnection { response_mux::BATCH_WAIT_SECONDS .with_label_values(&["worker"]) .observe(observed_batch_wait.as_secs_f64()); - response_mux::WRITE_CALLS_TOTAL - .with_label_values(&["worker"]) - .inc_by(write_calls); } } } .await; + if metrics_enabled { queue_depth.sub(reported_queue_depth); } - if let Some(connection) = weak.upgrade() { + if let Some((command, _, _)) = blocked_data.take() { + command.fail("response mux writer stopped"); + } + while let Ok(ingress) = stream_rx.try_recv() { + scheduler.ingest(&connection, ingress); + } + scheduler.fail_all(&connection, "response mux writer stopped"); + while let Ok(command) = priority_rx.try_recv() { + connection.queued_frames.fetch_sub(1, Ordering::AcqRel); + command.fail("response mux priority writer stopped"); + } connection.fail( &result .err() @@ -983,8 +1152,12 @@ impl MuxConnection { window_updates.inc(); } } - MuxFrameKind::Stop => state.context.stop(), + MuxFrameKind::Stop => { + state.record_cancellation(); + state.context.stop(); + } MuxFrameKind::Kill | MuxFrameKind::Reset => { + state.record_cancellation(); drop(state); connection.remove_stream( frame.stream_id, @@ -1039,11 +1212,29 @@ struct HostPool { next_connection_id: AtomicU64, warming: AtomicBool, maintenance_started: AtomicBool, - last_used: parking_lot::Mutex, + lifecycle: Mutex, cancel: CancellationToken, config: PoolConfig, } +struct HostLifecycle { + last_used: Instant, + retiring: bool, + openers: usize, +} + +struct HostOpenGuard { + host: Arc, +} + +impl Drop for HostOpenGuard { + fn drop(&mut self) { + let mut lifecycle = self.host.lifecycle.lock(); + lifecycle.openers = lifecycle.openers.saturating_sub(1); + lifecycle.last_used = Instant::now(); + } +} + impl HostPool { fn new( address: String, @@ -1061,7 +1252,11 @@ impl HostPool { next_connection_id: AtomicU64::new(1), warming: AtomicBool::new(false), maintenance_started: AtomicBool::new(false), - last_used: parking_lot::Mutex::new(Instant::now()), + lifecycle: Mutex::new(HostLifecycle { + last_used: Instant::now(), + retiring: false, + openers: 0, + }), cancel, config, }) @@ -1160,14 +1355,24 @@ impl HostPool { }); } - async fn connection(self: &Arc) -> Result> { - *self.last_used.lock() = Instant::now(); + async fn connection(self: &Arc) -> Result, HostOpenGuard)>> { + let opener = { + let mut lifecycle = self.lifecycle.lock(); + if lifecycle.retiring { + return Ok(None); + } + lifecycle.openers += 1; + lifecycle.last_used = Instant::now(); + HostOpenGuard { host: self.clone() } + }; let mut healthy = self.healthy_connections(); if healthy.is_empty() { healthy.push(self.ensure_first().await?); } self.start_maintenance(); - self.warm(); + if healthy.len() < self.config.pool_size { + self.warm(); + } let index = healthy .iter() .enumerate() @@ -1179,21 +1384,29 @@ impl HostPool { }) .map(|(index, _)| index) .expect("healthy response mux connection list is non-empty"); - Ok(healthy.swap_remove(index)) + Ok(Some((healthy.swap_remove(index), opener))) } - fn is_idle(&self) -> bool { - self.healthy_connections() - .iter() - .all(|connection| connection.active_streams.load(Ordering::Acquire) == 0) - && self.last_used.lock().elapsed() >= self.config.idle_ttl + fn try_retire(&self) -> bool { + let mut lifecycle = self.lifecycle.lock(); + if lifecycle.retiring + || lifecycle.openers != 0 + || lifecycle.last_used.elapsed() < self.config.idle_ttl + || !self + .healthy_connections() + .iter() + .all(|connection| connection.active_streams.load(Ordering::Acquire) == 0) + { + return false; + } + lifecycle.retiring = true; + true } } pub struct ResponseMuxClientPool { hosts: DashMap>, cancel: CancellationToken, - enabled: bool, config: PoolConfig, } @@ -1212,7 +1425,6 @@ impl ResponseMuxClientPool { let pool = Arc::new(Self { hosts: DashMap::new(), cancel, - enabled: runtime_config.enabled, config: PoolConfig::from_runtime(runtime_config), }); Self::start_cleanup(&pool); @@ -1232,7 +1444,6 @@ impl ResponseMuxClientPool { let pool = Arc::new(Self { hosts: DashMap::new(), cancel, - enabled: true, config: PoolConfig { pool_size: pool_size.max(1), writer_queue: writer_queue.max(1), @@ -1264,11 +1475,11 @@ impl ResponseMuxClientPool { break; } pool.hosts.retain(|_, host| { - let retain = !host.is_idle(); - if !retain { + let retiring = host.try_retire(); + if retiring { host.cancel.cancel(); } - retain + !retiring }); } }); @@ -1278,10 +1489,8 @@ impl ResponseMuxClientPool { self: &Arc, context: Arc, info: ConnectionInfo, + cancellation_counter: Option, ) -> Result { - if !self.enabled { - anyhow::bail!("response mux connection info received while mux mode is disabled"); - } let info = ResponseMuxConnectionInfo::try_from(info) .context("tcp-response-mux-connection-info-error")?; if info.version != RESPONSE_MUX_VERSION { @@ -1304,26 +1513,34 @@ impl ResponseMuxClientPool { frontend_server_id: info.frontend_server_id, version: info.version, }; - let host = self - .hosts - .entry(host_key) - .or_insert_with(|| { - HostPool::new( - info.address.clone(), - info.frontend_server_id, - info.version, - self.cancel.child_token(), - self.config, - ) - }) - .clone(); - let connection = host.connection().await?; + let (connection, _opener) = loop { + let host = self + .hosts + .entry(host_key.clone()) + .or_insert_with(|| { + HostPool::new( + info.address.clone(), + info.frontend_server_id, + info.version, + self.cancel.child_token(), + self.config, + ) + }) + .clone(); + if let Some(connection) = host.connection().await? { + break connection; + } + self.hosts + .remove_if(&host_key, |_, candidate| Arc::ptr_eq(candidate, &host)); + host.cancel.cancel(); + }; let state = Arc::new(WorkerStreamState { context, + cancellation_counter, + cancellation_recorded: AtomicBool::new(false), credits: Arc::new(Semaphore::new(self.config.initial_window)), max_credits: self.config.initial_window, writer_slots: Arc::new(Semaphore::new(self.config.stream_writer_queue)), - writer: Mutex::new(StreamWriterState::default()), closed: AtomicBool::new(false), close_token: CancellationToken::new(), }); @@ -1404,14 +1621,14 @@ impl MuxResponseStreamSender { match slots.try_acquire_owned() { Ok(permit) => Ok(permit), Err(tokio::sync::TryAcquireError::NoPermits) => { - let wait_start = Instant::now(); + let wait_start = per_frame_metrics_enabled().then(Instant::now); let permit = state .writer_slots .clone() .acquire_owned() .await .map_err(|_| anyhow!("response mux stream closed during writer admission"))?; - if per_frame_metrics_enabled() { + if let Some(wait_start) = wait_start { response_mux::WRITER_ADMISSION_STALL_SECONDS .observe(wait_start.elapsed().as_secs_f64()); } @@ -1451,7 +1668,7 @@ impl MuxResponseStreamSender { { Ok(permit) => permit, Err(tokio::sync::TryAcquireError::NoPermits) => { - let wait_start = Instant::now(); + let wait_start = per_frame_metrics_enabled().then(Instant::now); let slots = connection.queued_byte_slots.clone(); let permit = tokio::select! { _ = state.close_token.cancelled() => { @@ -1461,21 +1678,25 @@ impl MuxResponseStreamSender { anyhow!("response mux connection closed during byte-queue admission") })?, }; - response_mux::QUEUED_BYTE_ADMISSION_STALL_SECONDS - .observe(wait_start.elapsed().as_secs_f64()); + if let Some(wait_start) = wait_start { + response_mux::QUEUED_BYTE_ADMISSION_STALL_SECONDS + .observe(wait_start.elapsed().as_secs_f64()); + } permit } Err(tokio::sync::TryAcquireError::Closed) => { anyhow::bail!("response mux connection closed during byte-queue admission") } }; - connection.enqueue_stream_command( - self.stream_id, - state, - WriterCommand::new(frame, written) - .with_writer_permit(writer_permit) - .with_queued_byte_permit(queued_byte_permit), - ) + connection + .enqueue_stream_command( + self.stream_id, + state.clone(), + WriterCommand::new(frame, written) + .with_writer_permit(writer_permit) + .with_queued_byte_permit(queued_byte_permit), + ) + .await } async fn enqueue_priority_and_wait(&self, frame: MuxFrame) -> Result<()> { @@ -1503,7 +1724,7 @@ impl MultiplexedStreamSender for MuxResponseStreamSender { let permit = match self.state.credits.clone().try_acquire_many_owned(required) { Ok(permit) => permit, Err(tokio::sync::TryAcquireError::NoPermits) => { - let wait_start = Instant::now(); + let wait_start = per_frame_metrics_enabled().then(Instant::now); let permit = self .state .credits @@ -1511,7 +1732,7 @@ impl MultiplexedStreamSender for MuxResponseStreamSender { .acquire_many_owned(required) .await .map_err(|_| anyhow!("response mux stream closed while waiting for credits"))?; - if per_frame_metrics_enabled() { + if let Some(wait_start) = wait_start { response_mux::FLOW_CONTROL_STALL_SECONDS .observe(wait_start.elapsed().as_secs_f64()); } @@ -1607,10 +1828,11 @@ mod tests { let context = Context::new(()); Arc::new(WorkerStreamState { context: context.context(), + cancellation_counter: None, + cancellation_recorded: AtomicBool::new(false), credits: Arc::new(Semaphore::new(initial_credits)), max_credits: TEST_STREAM_WINDOW, writer_slots: Arc::new(Semaphore::new(writer_slots)), - writer: Mutex::new(StreamWriterState::default()), closed: AtomicBool::new(false), close_token: CancellationToken::new(), }) @@ -1647,17 +1869,88 @@ mod tests { assert_eq!(state.credits.available_permits(), TEST_STREAM_WINDOW); } + #[test] + fn detailed_metric_timestamps_are_opt_in() { + let stream_id = Uuid::new_v4(); + let disabled = WriterCommand::new_with_metrics( + MuxFrame::empty(MuxFrameKind::End, stream_id), + None, + false, + ); + let enabled = WriterCommand::new_with_metrics( + MuxFrame::empty(MuxFrameKind::End, stream_id), + None, + true, + ); + + assert!(disabled.enqueued_at.is_none()); + assert!(enabled.enqueued_at.is_some()); + } + + #[test] + fn ingress_admission_stops_after_one_frame_when_a_stream_becomes_ready() { + let stream_id = Uuid::new_v4(); + let state = worker_stream_state(TEST_STREAM_WINDOW, RESPONSE_MUX_STREAM_WRITER_QUEUE); + let (priority_tx, _priority_rx) = mpsc::channel(1); + let (stream_tx, mut stream_rx) = mpsc::channel(4); + let connection = MuxConnection { + id: 1, + cancel: CancellationToken::new(), + priority_tx, + stream_tx: stream_tx.clone(), + streams: DashMap::new(), + healthy: AtomicBool::new(true), + active_streams: AtomicUsize::new(1), + queued_frames: AtomicUsize::new(2), + queued_bytes: AtomicUsize::new(0), + max_queued_bytes: TEST_CONNECTION_WINDOW, + queued_byte_slots: Arc::new(Semaphore::new(TEST_CONNECTION_WINDOW)), + connection_credits: Arc::new(Semaphore::new(TEST_CONNECTION_WINDOW)), + max_connection_credits: TEST_CONNECTION_WINDOW, + sent_data_bytes: AtomicU64::new(0), + acknowledged_data_bytes: AtomicU64::new(0), + batch_interval: Duration::ZERO, + batch_max_bytes: 65_536, + batch_max_frames: 64, + }; + for payload in [b"first".as_slice(), b"second".as_slice()] { + stream_tx + .try_send(StreamIngress { + stream_id, + state: state.clone(), + command: WriterCommand::new( + MuxFrame::new( + MuxFrameKind::Data, + stream_id, + bytes::Bytes::copy_from_slice(payload), + ), + None, + ), + }) + .unwrap(); + } + + let mut scheduler = WriterScheduler::default(); + assert_eq!( + scheduler.ingest_one(&connection, &mut stream_rx), + IngressPoll::Ingested + ); + assert_eq!(stream_rx.len(), 1); + assert_eq!(scheduler.ready.len(), 1); + assert_eq!(scheduler.streams[&stream_id].pending.len(), 1); + } + #[test] fn cumulative_connection_ack_replenishes_credits_without_exceeding_the_window() { let (priority_tx, _priority_rx) = mpsc::channel(1); - let (ready_tx, _ready_rx) = mpsc::unbounded_channel(); + let (stream_tx, _stream_rx) = mpsc::channel(1); let remaining = 16; let consumed = TEST_CONNECTION_WINDOW - remaining; let connection = MuxConnection { id: 1, cancel: CancellationToken::new(), priority_tx, - ready_tx, + stream_tx, streams: DashMap::new(), healthy: AtomicBool::new(true), active_streams: AtomicUsize::new(0), @@ -1711,24 +2004,18 @@ mod tests { #[tokio::test] async fn closing_a_stream_wakes_credit_and_writer_admission_waiters() { let state = worker_stream_state(0, 0); - let (written_tx, written_rx) = oneshot::channel(); - state.writer.lock().pending.push_back(WriterCommand::new( - MuxFrame::empty(MuxFrameKind::End, Uuid::new_v4()), - Some(written_tx), - )); - let (priority_tx, _priority_rx) = mpsc::channel(1); - let (ready_tx, _ready_rx) = mpsc::unbounded_channel(); + let (stream_tx, _stream_rx) = mpsc::channel(1); let connection = MuxConnection { id: 1, cancel: CancellationToken::new(), priority_tx, - ready_tx, + stream_tx, streams: DashMap::new(), healthy: AtomicBool::new(true), active_streams: AtomicUsize::new(0), - queued_frames: AtomicUsize::new(1), - queued_bytes: AtomicUsize::new(super::super::MUX_HEADER_LEN), + queued_frames: AtomicUsize::new(0), + queued_bytes: AtomicUsize::new(0), max_queued_bytes: TEST_CONNECTION_WINDOW, queued_byte_slots: Arc::new(Semaphore::new(TEST_CONNECTION_WINDOW)), connection_credits: Arc::new(Semaphore::new(TEST_CONNECTION_WINDOW)), @@ -1754,8 +2041,6 @@ mod tests { assert!(credit_waiter.await.unwrap().is_err()); assert!(writer_waiter.await.unwrap().is_err()); - assert_eq!(connection.queued_frames.load(Ordering::Acquire), 0); - assert_eq!(written_rx.await.unwrap().unwrap_err(), "test stream closed"); assert_eq!(state.replenish_credits(1_024), 0); assert_eq!(state.credits.available_permits(), 0); } @@ -1772,36 +2057,20 @@ mod tests { let stream_id = Uuid::new_v4(); let (end_written_tx, end_written_rx) = oneshot::channel(); let state = worker_stream_state(TEST_STREAM_WINDOW, RESPONSE_MUX_STREAM_WRITER_QUEUE); - { - let mut writer = state.writer.lock(); - writer.pending.push_back(WriterCommand::new( - MuxFrame::new( - MuxFrameKind::Data, - stream_id, - bytes::Bytes::from(vec![b'x'; 40]), - ), - None, - )); - writer.pending.push_back(WriterCommand::new( - MuxFrame::empty(MuxFrameKind::End, stream_id), - Some(end_written_tx), - )); - writer.scheduled = true; - } let (priority_tx, priority_rx) = mpsc::channel(1); - let (ready_tx, ready_rx) = mpsc::unbounded_channel(); + let (stream_tx, stream_rx) = mpsc::channel(8); let cancel = CancellationToken::new(); let connection = Arc::new(MuxConnection { id: 1, cancel: cancel.clone(), priority_tx, - ready_tx: ready_tx.clone(), + stream_tx, streams: DashMap::new(), healthy: AtomicBool::new(true), active_streams: AtomicUsize::new(1), - queued_frames: AtomicUsize::new(2), - queued_bytes: AtomicUsize::new(88), + queued_frames: AtomicUsize::new(0), + queued_bytes: AtomicUsize::new(0), max_queued_bytes: TEST_CONNECTION_WINDOW, queued_byte_slots: Arc::new(Semaphore::new(TEST_CONNECTION_WINDOW)), connection_credits: Arc::new(Semaphore::new(64)), @@ -1812,16 +2081,41 @@ mod tests { batch_max_bytes: 64, batch_max_frames: 64, }); - connection.streams.insert(stream_id, state); - ready_tx.send(stream_id).unwrap(); + connection.streams.insert(stream_id, state.clone()); let writer_task = tokio::spawn(MuxConnection::writer_task( Arc::downgrade(&connection), write_half, priority_rx, - ready_rx, + stream_rx, cancel, )); + connection + .enqueue_stream_command( + stream_id, + state.clone(), + WriterCommand::new( + MuxFrame::new( + MuxFrameKind::Data, + stream_id, + bytes::Bytes::from(vec![b'x'; 40]), + ), + None, + ), + ) + .await + .unwrap(); + connection + .enqueue_stream_command( + stream_id, + state, + WriterCommand::new( + MuxFrame::empty(MuxFrameKind::End, stream_id), + Some(end_written_tx), + ), + ) + .await + .unwrap(); let reader_task = tokio::spawn(async move { let mut reader = FramedRead::new(server_read, MuxCodec::default()); let data = reader.next().await.unwrap().unwrap(); @@ -1841,11 +2135,206 @@ mod tests { writer_task.abort(); } + #[tokio::test] + async fn parked_frame_is_not_overwritten_by_another_ready_stream() { + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let address = listener.local_addr().unwrap(); + let client = TcpStream::connect(address).await.unwrap(); + let (server, _) = listener.accept().await.unwrap(); + let (_, write_half) = client.into_split(); + let (server_read, _) = server.into_split(); + + let stream_a = Uuid::new_v4(); + let stream_b = Uuid::new_v4(); + let state_a = worker_stream_state(TEST_STREAM_WINDOW, RESPONSE_MUX_STREAM_WRITER_QUEUE); + let state_b = worker_stream_state(TEST_STREAM_WINDOW, RESPONSE_MUX_STREAM_WRITER_QUEUE); + let (priority_tx, priority_rx) = mpsc::channel(1); + let (stream_tx, stream_rx) = mpsc::channel(8); + let cancel = CancellationToken::new(); + let connection = Arc::new(MuxConnection { + id: 1, + cancel: cancel.clone(), + priority_tx, + stream_tx, + streams: DashMap::new(), + healthy: AtomicBool::new(true), + active_streams: AtomicUsize::new(2), + queued_frames: AtomicUsize::new(0), + queued_bytes: AtomicUsize::new(0), + max_queued_bytes: TEST_CONNECTION_WINDOW, + queued_byte_slots: Arc::new(Semaphore::new(TEST_CONNECTION_WINDOW)), + connection_credits: Arc::new(Semaphore::new(TEST_CONNECTION_WINDOW)), + max_connection_credits: TEST_CONNECTION_WINDOW, + sent_data_bytes: AtomicU64::new(0), + acknowledged_data_bytes: AtomicU64::new(0), + batch_interval: Duration::ZERO, + batch_max_bytes: 64, + batch_max_frames: 64, + }); + connection.streams.insert(stream_a, state_a.clone()); + connection.streams.insert(stream_b, state_b.clone()); + + for (state, frame) in [ + ( + state_a.clone(), + MuxFrame::new( + MuxFrameKind::Data, + stream_a, + bytes::Bytes::from_static(b"a-small"), + ), + ), + ( + state_a, + MuxFrame::new( + MuxFrameKind::Data, + stream_a, + bytes::Bytes::from(vec![b'A'; 40]), + ), + ), + ( + state_b, + MuxFrame::new( + MuxFrameKind::Data, + stream_b, + bytes::Bytes::from_static(b"b-ready"), + ), + ), + ] { + connection + .enqueue_stream_command(frame.stream_id, state, WriterCommand::new(frame, None)) + .await + .unwrap(); + } + + let writer_task = tokio::spawn(MuxConnection::writer_task( + Arc::downgrade(&connection), + write_half, + priority_rx, + stream_rx, + cancel, + )); + let mut reader = FramedRead::new(server_read, MuxCodec::default()); + let first = reader.next().await.unwrap().unwrap(); + let second = reader.next().await.unwrap().unwrap(); + let third = reader.next().await.unwrap().unwrap(); + + assert_eq!(first.payload, bytes::Bytes::from_static(b"a-small")); + assert_eq!(second.payload, bytes::Bytes::from(vec![b'A'; 40])); + assert_eq!(third.payload, bytes::Bytes::from_static(b"b-ready")); + writer_task.abort(); + } + + #[tokio::test(start_paused = true)] + async fn one_ms_batch_interval_flushes_on_expiry() { + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let address = listener.local_addr().unwrap(); + let client = TcpStream::connect(address).await.unwrap(); + let (server, _) = listener.accept().await.unwrap(); + let (_, write_half) = client.into_split(); + let (server_read, _) = server.into_split(); + + let stream_id = Uuid::new_v4(); + let state = worker_stream_state(TEST_STREAM_WINDOW, RESPONSE_MUX_STREAM_WRITER_QUEUE); + let (priority_tx, priority_rx) = mpsc::channel(1); + let (stream_tx, stream_rx) = mpsc::channel(8); + let cancel = CancellationToken::new(); + let connection = Arc::new(MuxConnection { + id: 1, + cancel: cancel.clone(), + priority_tx, + stream_tx, + streams: DashMap::new(), + healthy: AtomicBool::new(true), + active_streams: AtomicUsize::new(1), + queued_frames: AtomicUsize::new(0), + queued_bytes: AtomicUsize::new(0), + max_queued_bytes: TEST_CONNECTION_WINDOW, + queued_byte_slots: Arc::new(Semaphore::new(TEST_CONNECTION_WINDOW)), + connection_credits: Arc::new(Semaphore::new(TEST_CONNECTION_WINDOW)), + max_connection_credits: TEST_CONNECTION_WINDOW, + sent_data_bytes: AtomicU64::new(0), + acknowledged_data_bytes: AtomicU64::new(0), + batch_interval: Duration::from_millis(1), + batch_max_bytes: 65_536, + batch_max_frames: 64, + }); + connection.streams.insert(stream_id, state.clone()); + connection + .enqueue_stream_command( + stream_id, + state, + WriterCommand::new( + MuxFrame::new( + MuxFrameKind::Data, + stream_id, + bytes::Bytes::from_static(b"batched"), + ), + None, + ), + ) + .await + .unwrap(); + + let writer_task = tokio::spawn(MuxConnection::writer_task( + Arc::downgrade(&connection), + write_half, + priority_rx, + stream_rx, + cancel, + )); + let (received_tx, mut received_rx) = oneshot::channel(); + tokio::spawn(async move { + let mut reader = FramedRead::new(server_read, MuxCodec::default()); + let _ = received_tx.send(reader.next().await.unwrap().unwrap()); + }); + + for _ in 0..4 { + tokio::task::yield_now().await; + } + assert!(matches!( + received_rx.try_recv(), + Err(oneshot::error::TryRecvError::Empty) + )); + tokio::time::advance(Duration::from_micros(999)).await; + tokio::task::yield_now().await; + assert!(matches!( + received_rx.try_recv(), + Err(oneshot::error::TryRecvError::Empty) + )); + tokio::time::advance(Duration::from_micros(1)).await; + assert_eq!( + received_rx.await.unwrap().payload, + bytes::Bytes::from_static(b"batched") + ); + writer_task.abort(); + } + + #[test] + fn idle_cleanup_cannot_retire_a_host_with_an_opener() { + let mut config = PoolConfig::from_runtime(integration_config()); + config.idle_ttl = Duration::ZERO; + let host = HostPool::new( + "127.0.0.1:1".to_string(), + Uuid::new_v4(), + RESPONSE_MUX_VERSION, + CancellationToken::new(), + config, + ); + let opener = { + let mut lifecycle = host.lifecycle.lock(); + lifecycle.openers += 1; + HostOpenGuard { host: host.clone() } + }; + + assert!(!host.try_retire()); + drop(opener); + assert!(host.try_retire()); + } + fn integration_config() -> ResponseMuxConfig { ResponseMuxConfig { - enabled: true, packet_metrics: false, - batch_interval: Duration::from_millis(5), + batch_interval: Duration::from_millis(1), batch_max_bytes: 65_536, batch_max_frames: 64, stream_window_bytes: 262_144, @@ -1881,7 +2370,7 @@ mod tests { .unwrap() .stream_id; let mut sender = pool - .create_response_stream(context.context(), info) + .create_response_stream(context.context(), info, None) .await .unwrap(); sender.send_prologue(None).await.unwrap(); @@ -1908,33 +2397,36 @@ mod tests { } #[tokio::test] - async fn disabled_worker_rejects_mux_connection_info() { - let mut config = integration_config(); - config.enabled = false; + async fn legacy_response_connection_info_is_rejected() { + use crate::pipeline::network::tcp::{StreamType, TcpStreamConnectionInfo}; + + let config = integration_config(); let cancel = CancellationToken::new(); let pool = ResponseMuxClientPool::new(cancel.clone(), config); let context = Context::new(()); - let info = ResponseMuxConnectionInfo { + let info = TcpStreamConnectionInfo { address: "127.0.0.1:1".to_string(), - frontend_server_id: Uuid::new_v4(), - stream_id: Uuid::new_v4(), context: context.context().id().to_string(), - version: RESPONSE_MUX_VERSION, + subject: Uuid::new_v4().to_string(), + stream_type: StreamType::Response, }; let error = pool - .create_response_stream(context.context(), info.into()) + .create_response_stream(context.context(), info.into(), None) .await .err() - .expect("disabled mux mode must reject mux connection info"); - assert!(error.to_string().contains("mux mode is disabled")); + .expect("legacy response connection info must be rejected"); + assert!( + error + .to_string() + .contains("tcp-response-mux-connection-info-error") + ); cancel.cancel(); } #[tokio::test(flavor = "multi_thread", worker_threads = 2)] async fn stream_credit_stall_does_not_block_another_stream() { let config = ResponseMuxConfig { - enabled: true, packet_metrics: false, batch_interval: Duration::ZERO, batch_max_bytes: 65_536, @@ -1995,6 +2487,35 @@ mod tests { cancel.cancel(); } + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] + async fn frontend_receiver_drop_removes_the_worker_stream() { + let config = integration_config(); + let server = crate::pipeline::network::tcp::server::TcpStreamServer::new_mux_for_test( + crate::pipeline::network::tcp::server::ServerOptions::default(), + config, + ) + .await + .unwrap(); + let address = mux_address(server.clone()).await; + let cancel = CancellationToken::new(); + let pool = ResponseMuxClientPool::new(cancel.clone(), config); + let (stream_id, context, sender, receiver) = open_mux_stream(server, pool.clone()).await; + + drop(receiver); + tokio::time::timeout(Duration::from_secs(1), async { + while !context.context().is_killed() + || pool.stream_connection_id(&address, stream_id).is_some() + { + tokio::task::yield_now().await; + } + }) + .await + .expect("frontend receiver drop did not remove the worker stream"); + + assert!(sender.finish().await.is_err()); + cancel.cancel(); + } + #[tokio::test(flavor = "multi_thread", worker_threads = 4)] async fn connection_failure_is_scoped_and_pool_reconnects() { let config = integration_config(); @@ -2105,7 +2626,7 @@ mod tests { .unwrap(); let (info, provider) = pending.into_parts(); let mut sender = pool - .create_response_stream(context.context(), info) + .create_response_stream(context.context(), info, None) .await .unwrap(); sender.send_prologue(None).await.unwrap(); diff --git a/lib/runtime/src/pipeline/network/tcp/server.rs b/lib/runtime/src/pipeline/network/tcp/server.rs index 433106e7de02..6331d660374d 100644 --- a/lib/runtime/src/pipeline/network/tcp/server.rs +++ b/lib/runtime/src/pipeline/network/tcp/server.rs @@ -94,9 +94,8 @@ impl ServerOptions { } } -/// A [`TcpStreamServer`] is a TCP service that listens on a port for incoming response connections. -/// A Response connection is a connection that is established by a client with the intention of sending -/// specific data back to the server. +/// A [`TcpStreamServer`] accepts dedicated request streams and persistent +/// multiplexed response connections. pub struct TcpStreamServer { local_ip: String, local_port: u16, @@ -104,7 +103,7 @@ pub struct TcpStreamServer { mux_config: ResponseMuxConfig, state: Arc>, response_pending: Arc>, - response_active: Arc>, + response_active: Arc>, } // pub struct TcpStreamReceiver { @@ -122,14 +121,6 @@ struct RequestedSendConnection { send_buffer_count: usize, } -struct RequestedRecvConnection { - context: Arc, - connection: oneshot::Sender>, - /// Capacity of the per-stream mpsc buffer between the socket task and the - /// engine consumer; carried from the registration [`StreamOptions`]. - send_buffer_count: usize, -} - struct RequestedMuxRecvConnection { context: Arc, connection: Mutex>>>, @@ -137,16 +128,21 @@ struct RequestedMuxRecvConnection { registered_at: Instant, } -struct ActiveMuxResponseStream { +struct ActiveMuxResponseControl { connection_id: uuid::Uuid, context: Arc, - response_tx: mpsc::Sender, control_tx: mpsc::Sender, + close_tx: mpsc::Sender, control_failed: CancellationToken, } +struct ActiveMuxResponseStream { + context: Arc, + response_tx: mpsc::Sender, +} + struct ResponseMuxSocket { - read_half: tokio::io::ReadHalf, + reader: FramedRead, MuxCodec>, write_half: tokio::io::WriteHalf, packet_socket: Option, } @@ -162,8 +158,8 @@ impl Drop for ActiveMuxResponseStream { /// Build the per-stream data-plane mpsc channel that bridges the socket task /// and the engine producer/consumer. The capacity is driven by the /// registration options ([`StreamOptions::send_buffer_count`]) rather than a -/// hard-coded constant; both `process_request_stream` and -/// `process_response_stream` size their channel through this helper. See #10293. +/// hard-coded constant; both the dedicated request path and response-mux +/// mailboxes use this helper. See #10293. fn data_plane_channel(send_buffer_count: usize) -> (mpsc::Sender, mpsc::Receiver) { // `tokio::sync::mpsc::channel` panics on a capacity of 0. Now that the value // is caller-configurable via `StreamOptions::send_buffer_count`, clamp to at @@ -191,14 +187,13 @@ fn data_plane_channel(send_buffer_count: usize) -> (mpsc::Sender, mpsc::Re #[derive(Default)] struct State { tx_subjects: HashMap, - rx_subjects: HashMap, /// subject UUID -> EndpointInstanceId. Full 4-field key isolates services /// that share an endpoint name across namespaces/components. subject_instance: HashMap, /// EndpointInstanceId -> tagged subject UUIDs, for batch cancellation on /// removal. The `StreamType` tag tells `cancel_instance_streams` which - /// of `rx_subjects` / `tx_subjects` holds the registration so both halves - /// of a bidirectional session get dropped together. + /// of the response mux registry / `tx_subjects` holds the registration so + /// both halves of a bidirectional session get dropped together. instance_subjects: HashMap>, /// Tombstones (instance -> insertion time) close the /// `cancel_instance_streams` vs `associate_instance` race; entries expire @@ -347,7 +342,6 @@ impl TcpStreamServer { instance_id = id.instance_id, "Cancelling subject immediately: instance already removed (tombstoned)" ); - state.rx_subjects.remove(recv_subject); if let Ok(stream_id) = uuid::Uuid::parse_str(recv_subject) { self.cancel_mux_response_stream(stream_id, MuxFrameKind::Kill); } @@ -374,7 +368,6 @@ impl TcpStreamServer { /// `oneshot::Sender` so the waiting receiver resolves with `RecvError`. pub async fn cancel_recv_stream(&self, subject: &str) { let mut state = self.state.lock(); - state.rx_subjects.remove(subject); if let Ok(stream_id) = uuid::Uuid::parse_str(subject) { self.cancel_mux_response_stream(stream_id, MuxFrameKind::Kill); } @@ -423,7 +416,6 @@ impl TcpStreamServer { for (kind, subject) in &subjects { match kind { StreamType::Response => { - state.rx_subjects.remove(subject); if let Ok(stream_id) = uuid::Uuid::parse_str(subject) { self.cancel_mux_response_stream(stream_id, MuxFrameKind::Kill); } @@ -452,6 +444,7 @@ impl TcpStreamServer { .control_tx .try_send(MuxFrame::empty(kind, stream_id)) .is_err() + || active.close_tx.try_send(stream_id).is_err() { active.control_failed.cancel(); } @@ -465,7 +458,7 @@ impl TcpStreamServer { server_id: uuid::Uuid, mux_config: ResponseMuxConfig, response_pending: Arc>, - response_active: Arc>, + response_active: Arc>, ) -> Result { let addr = format!("{}:{}", local_ip, local_port); let state_clone = state.clone(); @@ -493,10 +486,6 @@ impl TcpStreamServer { self.state.lock().tx_subjects.insert(subject, connection); } - fn insert_response_stream(&self, subject: String, connection: RequestedRecvConnection) { - self.state.lock().rx_subjects.insert(subject, connection); - } - fn take_request_stream(state: &Mutex, subject: &str) -> Option { let mut state = state.lock(); let connection = state.tx_subjects.remove(subject); @@ -510,23 +499,6 @@ impl TcpStreamServer { } connection } - - fn take_response_stream( - state: &Mutex, - subject: &str, - ) -> Option { - let mut state = state.lock(); - let connection = state.rx_subjects.remove(subject); - if let Some(key) = state.subject_instance.remove(subject) - && let Some(subjects) = state.instance_subjects.get_mut(&key) - { - subjects.remove(&(StreamType::Response, subject.to_string())); - if subjects.is_empty() { - state.instance_subjects.remove(&key); - } - } - connection - } } // todo - possible rename ResponseService to ResponseServer @@ -609,94 +581,56 @@ impl ResponseService for TcpStreamServer { let (pending_recver_tx, pending_recver_rx) = oneshot::channel(); let receiver_id = uuid::Uuid::new_v4(); let receiver_subject = receiver_id.to_string(); - let registry_subject = receiver_subject.clone(); - - if self.mux_config.enabled { - self.response_pending.insert( - receiver_id, - RequestedMuxRecvConnection { - context: options.context.clone(), - connection: Mutex::new(Some(pending_recver_tx)), - send_buffer_count: options.send_buffer_count, - registered_at: Instant::now(), - }, - ); - - let cleanup_id = receiver_id; - let cleanup_subject = receiver_subject; - let cleanup_state = self.state.clone(); - let cleanup_pending = self.response_pending.clone(); - let cleanup_active = self.response_active.clone(); - let registered_stream = RegisteredStream::new( - ResponseMuxConnectionInfo { - address: address.clone(), - frontend_server_id: self.server_id, - stream_id: receiver_id, - context: options.context.id().to_string(), - version: RESPONSE_MUX_VERSION, - } - .into(), - pending_recver_rx, - ) - .with_cleanup(move || { - cleanup_pending.remove(&cleanup_id); - if let Some((_, active)) = cleanup_active.remove(&cleanup_id) { - active.context.kill(); - if active - .control_tx - .try_send(MuxFrame::empty(MuxFrameKind::Kill, cleanup_id)) - .is_err() - { - active.control_failed.cancel(); - } - } - let mut state = cleanup_state.lock(); - if let Some(key) = state.subject_instance.remove(&cleanup_subject) - && let Some(subjects) = state.instance_subjects.get_mut(&key) - { - subjects.remove(&(StreamType::Response, cleanup_subject.clone())); - if subjects.is_empty() { - state.instance_subjects.remove(&key); - } - } - }); - Some(registered_stream) - } else { - let connection_info = RequestedRecvConnection { + self.response_pending.insert( + receiver_id, + RequestedMuxRecvConnection { context: options.context.clone(), - connection: pending_recver_tx, + connection: Mutex::new(Some(pending_recver_tx)), send_buffer_count: options.send_buffer_count, - }; + registered_at: Instant::now(), + }, + ); - let cleanup_subject = receiver_subject.clone(); - let cleanup_state = self.state.clone(); - let registered_stream = RegisteredStream::new( - TcpStreamConnectionInfo { - address: address.clone(), - subject: receiver_subject, - context: options.context.id().to_string(), - stream_type: StreamType::Response, - } - .into(), - pending_recver_rx, - ) - .with_cleanup(move || { - let mut state = cleanup_state.lock(); - state.rx_subjects.remove(&cleanup_subject); - if let Some(key) = state.subject_instance.remove(&cleanup_subject) - && let Some(subjects) = state.instance_subjects.get_mut(&key) + let cleanup_id = receiver_id; + let cleanup_subject = receiver_subject; + let cleanup_state = self.state.clone(); + let cleanup_pending = self.response_pending.clone(); + let cleanup_active = self.response_active.clone(); + let registered_stream = RegisteredStream::new( + ResponseMuxConnectionInfo { + address: address.clone(), + frontend_server_id: self.server_id, + stream_id: receiver_id, + context: options.context.id().to_string(), + version: RESPONSE_MUX_VERSION, + } + .into(), + pending_recver_rx, + ) + .with_cleanup(move || { + cleanup_pending.remove(&cleanup_id); + if let Some((_, active)) = cleanup_active.remove(&cleanup_id) { + active.context.kill(); + if active + .control_tx + .try_send(MuxFrame::empty(MuxFrameKind::Kill, cleanup_id)) + .is_err() + || active.close_tx.try_send(cleanup_id).is_err() { - subjects.remove(&(StreamType::Response, cleanup_subject.clone())); - if subjects.is_empty() { - state.instance_subjects.remove(&key); - } + active.control_failed.cancel(); } - }); - - self.insert_response_stream(registry_subject, connection_info); - - Some(registered_stream) - } + } + let mut state = cleanup_state.lock(); + if let Some(key) = state.subject_instance.remove(&cleanup_subject) + && let Some(subjects) = state.instance_subjects.get_mut(&key) + { + subjects.remove(&(StreamType::Response, cleanup_subject.clone())); + if subjects.is_empty() { + state.instance_subjects.remove(&key); + } + } + }); + Some(registered_stream) } else { None }; @@ -720,7 +654,7 @@ async fn tcp_listener( server_id: uuid::Uuid, mux_config: ResponseMuxConfig, response_pending: Arc>, - response_active: Arc>, + response_active: Arc>, read_tx: tokio::sync::oneshot::Sender>, ) -> Result<()> { let listener = tokio::net::TcpListener::bind(&addr) @@ -794,7 +728,7 @@ async fn tcp_listener( server_id: uuid::Uuid, mux_config: ResponseMuxConfig, response_pending: Arc>, - response_active: Arc>, + response_active: Arc>, ) { let result = process_stream( stream, @@ -823,9 +757,10 @@ async fn tcp_listener( server_id: uuid::Uuid, mux_config: ResponseMuxConfig, response_pending: Arc>, - response_active: Arc>, + response_active: Arc>, ) -> Result<()> { - let packet_socket = (mux_config.enabled && mux_config.packet_metrics) + let packet_socket = mux_config + .packet_metrics .then(|| stream.as_fd().try_clone_to_owned().ok()) .flatten(); // split the socket in to a reader and writer @@ -857,9 +792,6 @@ async fn tcp_listener( connection_id, }) = serde_json::from_slice::(header) { - if !mux_config.enabled { - anyhow::bail!("response mux handshake received while mux mode is disabled"); - } if version != RESPONSE_MUX_VERSION { anyhow::bail!( "unsupported response mux version {version}; expected {RESPONSE_MUX_VERSION}" @@ -884,7 +816,7 @@ async fn tcp_listener( response_pending, response_active, ResponseMuxSocket { - read_half: framed_reader.into_inner(), + reader: framed_reader.map_decoder(|_| MuxCodec::default()), write_half: framed_writer.into_inner(), packet_socket, }, @@ -901,10 +833,10 @@ async fn tcp_listener( StreamType::Request => { process_request_stream(handshake.subject, state, framed_reader, framed_writer).await } - StreamType::Response => { - process_response_stream(handshake.subject, state, framed_reader, framed_writer) - .await - } + StreamType::Response => anyhow::bail!( + "legacy dedicated TCP response streams are no longer supported; expected {}", + super::TCP_RESPONSE_MUX_TRANSPORT + ), } } @@ -913,17 +845,17 @@ async fn tcp_listener( mux_config: ResponseMuxConfig, state: Arc>, response_pending: Arc>, - response_active: Arc>, + response_active: Arc>, socket: ResponseMuxSocket, ) -> Result<()> { let ResponseMuxSocket { - read_half, + mut reader, write_half, packet_socket, } = socket; - let mut reader = FramedRead::new(read_half, MuxCodec::default()); let mut writer = FramedWrite::new(write_half, MuxCodec::default()); let (control_tx, mut control_rx) = mpsc::channel::(RESPONSE_MUX_WRITER_QUEUE); + let (close_tx, mut close_rx) = mpsc::channel::(RESPONSE_MUX_WRITER_QUEUE); let control_failed = CancellationToken::new(); let mut reported_data_segments = packet_socket .as_ref() @@ -937,6 +869,7 @@ async fn tcp_listener( .inc(); let writer_failed = control_failed.clone(); + let detailed_metrics = mux_config.packet_metrics; let writer_task = tokio::spawn(async move { let frame_counters = crate::metrics::response_mux::FrameCounters::for_direction("frontend_to_worker"); @@ -944,7 +877,9 @@ async fn tcp_listener( .with_label_values(&["frontend"]) .clone(); while let Some(frame) = control_rx.recv().await { - frame_counters.inc(frame.kind.metric_label()); + if detailed_metrics { + frame_counters.inc(frame.kind.metric_label()); + } if let Err(err) = writer.send(frame).await { writer_failed.cancel(); return Err(err.into()); @@ -964,6 +899,9 @@ async fn tcp_listener( packet_tick.tick().await; let frame_counters = crate::metrics::response_mux::FrameCounters::for_direction("worker_to_frontend"); + // Data routing is connection-local so the per-frame path does not take + // a shard lock in the global lifecycle registry. + let mut active_streams = HashMap::::new(); let result: Result<()> = async { loop { @@ -971,6 +909,18 @@ async fn tcp_listener( _ = control_failed.cancelled() => { anyhow::bail!("frontend response mux control writer failed") } + Some(stream_id) = close_rx.recv() => { + active_streams.remove(&stream_id); + response_pending.remove(&stream_id); + if response_active + .get(&stream_id) + .is_some_and(|active| active.connection_id == connection_id) + { + response_active.remove(&stream_id); + } + remove_response_association(&state, stream_id); + continue; + } _ = packet_tick.tick(), if reported_data_segments.is_some() => { if let Some(current) = packet_socket .as_ref() @@ -999,7 +949,9 @@ async fn tcp_listener( }, }; let frame = message; - frame_counters.inc(frame.kind.metric_label()); + if mux_config.packet_metrics { + frame_counters.inc(frame.kind.metric_label()); + } let stream_id = frame.stream_id; match frame.kind { @@ -1017,9 +969,28 @@ async fn tcp_listener( continue; }; let prologue: ResponseStreamPrologue = - serde_json::from_slice(&frame.payload).map_err(|err| { - error!("invalid response mux prologue for {stream_id}: {err}") - })?; + match serde_json::from_slice(&frame.payload) { + Ok(prologue) => prologue, + Err(err) => { + let reason = format!( + "invalid response mux prologue for {stream_id}: {err}" + ); + if let Some(connection) = pending.connection.lock().take() { + let _ = connection.send(Err(reason)); + } + crate::metrics::response_mux::RESETS_TOTAL + .with_label_values(&["frontend", "invalid_prologue"]) + .inc(); + if control_tx + .try_send(MuxFrame::empty(MuxFrameKind::Reset, stream_id)) + .is_err() + { + control_failed.cancel(); + } + remove_response_association(&state, stream_id); + continue; + } + }; let Some(connection) = pending.connection.lock().take() else { if control_tx .try_send(MuxFrame::empty(MuxFrameKind::Reset, stream_id)) @@ -1047,6 +1018,7 @@ async fn tcp_listener( let active_for_close = response_active.clone(); let control_for_window = control_tx.clone(); let control_for_close = control_tx.clone(); + let close_for_receiver = close_tx.clone(); let failed_for_window = control_failed.clone(); let failed_for_close = control_failed.clone(); let state_for_close = state.clone(); @@ -1087,6 +1059,7 @@ async fn tcp_listener( if control_for_close .try_send(MuxFrame::empty(MuxFrameKind::Kill, stream_id)) .is_err() + || close_for_receiver.try_send(stream_id).is_err() { failed_for_close.cancel(); } @@ -1098,14 +1071,21 @@ async fn tcp_listener( }; response_active.insert( stream_id, - ActiveMuxResponseStream { + ActiveMuxResponseControl { connection_id, - context, - response_tx, + context: context.clone(), control_tx: control_tx.clone(), + close_tx: close_tx.clone(), control_failed: control_failed.clone(), }, ); + active_streams.insert( + stream_id, + ActiveMuxResponseStream { + context, + response_tx, + }, + ); crate::metrics::response_mux::ACTIVE_STREAMS .with_label_values(&["frontend"]) .inc(); @@ -1114,6 +1094,7 @@ async fn tcp_listener( .is_err() { response_active.remove(&stream_id); + active_streams.remove(&stream_id); if control_tx .try_send(MuxFrame::empty(MuxFrameKind::Kill, stream_id)) .is_err() @@ -1137,8 +1118,8 @@ async fn tcp_listener( control_failed.cancel(); } } - let delivery_failure = match response_active.get(&stream_id) { - Some(active) if active.connection_id == connection_id => match active + let delivery_failure = match active_streams.get(&stream_id) { + Some(active) => match active .response_tx .try_send(StreamRxItem::multiplexed(frame.payload, encoded_len)) { @@ -1155,6 +1136,7 @@ async fn tcp_listener( .with_label_values(&["frontend", reason]) .inc(); response_active.remove(&stream_id); + active_streams.remove(&stream_id); response_pending.remove(&stream_id); if control_tx .try_send(MuxFrame::empty(MuxFrameKind::Reset, stream_id)) @@ -1166,10 +1148,7 @@ async fn tcp_listener( } } MuxFrameKind::End => { - if response_active - .get(&stream_id) - .is_some_and(|active| active.connection_id == connection_id) - { + if active_streams.remove(&stream_id).is_some() { response_active.remove(&stream_id); remove_response_association(&state, stream_id); } else { @@ -1183,12 +1162,8 @@ async fn tcp_listener( } MuxFrameKind::Reset => { response_pending.remove(&stream_id); - if response_active - .get(&stream_id) - .is_some_and(|active| active.connection_id == connection_id) - { - response_active.remove(&stream_id); - } + active_streams.remove(&stream_id); + response_active.remove(&stream_id); remove_response_association(&state, stream_id); } MuxFrameKind::Stop @@ -1217,16 +1192,18 @@ async fn tcp_listener( .inc_by(current.saturating_sub(previous)); } - let affected: Vec = response_active - .iter() - .filter(|entry| entry.connection_id == connection_id) - .map(|entry| *entry.key()) - .collect(); + let affected: Vec = active_streams.keys().copied().collect(); crate::metrics::response_mux::CONNECTION_LOST_STREAMS_TOTAL.inc_by(affected.len() as u64); for stream_id in affected { - if let Some((_, active)) = response_active.remove(&stream_id) { + if let Some(active) = active_streams.remove(&stream_id) { active.context.kill(); } + if response_active + .get(&stream_id) + .is_some_and(|active| active.connection_id == connection_id) + { + response_active.remove(&stream_id); + } remove_response_association(&state, stream_id); } crate::metrics::response_mux::ACTIVE_CONNECTIONS @@ -1251,8 +1228,8 @@ async fn tcp_listener( } } - /// Symmetric to [`process_response_stream`] for the upstream→downstream - /// data direction: deliver the [`StreamSender`] half registered by the + /// For the upstream→downstream data direction, deliver the [`StreamSender`] + /// half registered by the /// upstream to whoever awaits it, then pump every frame the upstream pushes /// into the now-connected TCP socket. /// @@ -1285,7 +1262,7 @@ async fn tcp_listener( // Buffer size is driven by the registration options // ([`StreamOptions::send_buffer_count`]) rather than hard-coded; the - // same applies to `process_response_stream`. See #10293. + // same applies to response-mux mailboxes. See #10293. let (request_tx, request_rx) = data_plane_channel(send_buffer_count); if connection @@ -1377,269 +1354,6 @@ async fn tcp_listener( tracing::trace!(?err, "request-stream socket shutdown failed"); } } - - async fn process_response_stream( - subject: String, - state: Arc>, - mut reader: FramedRead, TwoPartCodec>, - writer: FramedWrite, TwoPartCodec>, - ) -> Result<()> { - let response_stream = TcpStreamServer::take_response_stream(&state, &subject).ok_or_else(|| { - error!("Subject not found: {}; upstream publisher specified a subject unknown to the downsteam subscriber", subject) - })?; - - // unwrap response_stream - let RequestedRecvConnection { - context, - connection, - send_buffer_count, - } = response_stream; - - // the [`Prologue`] - // there must be a second control message it indicate the other segment's generate method was successful - let prologue = reader - .next() - .await - .ok_or(error!("Connection closed without a ControlMessge"))??; - - // deserialize prologue - let prologue = match prologue.into_message_type() { - TwoPartMessageType::HeaderOnly(header) => { - let prologue: ResponseStreamPrologue = serde_json::from_slice(&header) - .map_err(|e| error!("Failed to deserialize ControlMessage: {}", e))?; - prologue - } - _ => { - // Worker sent a non-HeaderOnly frame in the prologue slot - // (protocol violation, version skew, corruption). Notify the - // requester so the generate call chain fails cleanly, then - // return Err so the connection task ends without panicking. - let msg = "malformed prologue: expected HeaderOnly ControlMessage"; - let _ = connection.send(Err(msg.to_string())); - return Err(error!(msg)); - } - }; - - // await the control message of GTG or Error, if error, then connection.send(Err(String)), which should fail the - // generate call chain - // - // note: this second control message might be delayed, but the expensive part of setting up the connection - // is both complete and ready for data flow; awaiting here is not a performance hit or problem and it allows - // us to trace the initial setup time vs the time to prologue - if let Some(error) = &prologue.error { - let _ = connection.send(Err(error.clone())); - return Err(error!("Received error prologue: {}", error)); - } - - // Buffer size is driven by the registration options - // ([`StreamOptions::send_buffer_count`]) rather than hard-coded; the - // same applies to `process_request_stream`. See #10293. - let (response_tx, response_rx) = data_plane_channel(send_buffer_count); - - if connection - .send(Ok(crate::pipeline::network::StreamReceiver::dedicated( - response_rx, - ))) - .is_err() - { - return Err(error!( - "The requester of the stream has been dropped before the connection was established" - )); - } - - let (control_tx, control_rx) = mpsc::channel::(1); - - // sender task - // issues control messages to the sender and when finished shuts down the socket - // this should be the last task to finish and must - let send_task = tokio::spawn(network_send_handler(writer, control_rx)); - - // forward task - let recv_task = tokio::spawn(network_receive_handler( - reader, - response_tx, - control_tx, - context.clone(), - )); - - // check the results of each of the tasks - let (monitor_result, forward_result) = tokio::join!(send_task, recv_task); - - monitor_result?; - forward_result?; - - Ok(()) - } - - async fn network_receive_handler( - mut framed_reader: FramedRead, TwoPartCodec>, - response_tx: mpsc::Sender, - control_tx: mpsc::Sender, - context: Arc, - ) { - // These futures stay pending across frames. Constructing them inside the loop clones - // watch receivers and registers/drops notifications for every streamed token. - let response_closed = response_tx.closed(); - let killed = context.killed(); - let stopped = context.stopped(); - tokio::pin!(response_closed, killed, stopped); - - // loop over reading the tcp stream and checking if the writer is closed - let mut can_stop = true; - loop { - tokio::select! { - biased; - - _ = &mut response_closed => { - tracing::trace!("response channel closed before the client finished writing data"); - let _ = control_tx.send(ControlMessage::Kill).await; - break; - } - - _ = &mut killed => { - tracing::trace!("context kill signal received; shutting down"); - let _ = control_tx.send(ControlMessage::Kill).await; - break; - } - - _ = &mut stopped, if can_stop => { - tracing::trace!("context stop signal received; shutting down"); - // `stopped` is now complete; keep this branch disabled because polling - // the same completed async future again would panic. - can_stop = false; - let _ = control_tx.send(ControlMessage::Stop).await; - } - - msg = framed_reader.next() => { - match msg { - Some(Ok(msg)) => { - let (header, data) = msg.into_parts(); - - // received a control message - if !header.is_empty() { - match process_control_message(header) { - Ok(ControlAction::Continue) => {} - Ok(ControlAction::Shutdown) => { - if !data.is_empty() { - // Sentinel-with-data is a protocol - // violation; kill this stream, don't - // assert!() the process down. - tracing::warn!( - data_len = data.len(), - "client sent Sentinel with data (protocol violation); killing stream" - ); - let _ = control_tx.send(ControlMessage::Kill).await; - break; - } - tracing::trace!("received sentinel message; shutting down"); - break; - } - Err(e) => { - // Malformed control message — kill only - // this stream. - tracing::warn!(err = ?e, "malformed control message, closing connection"); - let _ = control_tx.send(ControlMessage::Kill).await; - break; - } - } - } - - if !data.is_empty() - && let Err(err) = response_tx - .send(crate::pipeline::network::StreamRxItem::dedicated(data)) - .await - { - tracing::debug!(?err, "forwarding body/data to response channel failed"); - let _ = control_tx.send(ControlMessage::Kill).await; - break; - }; - } - Some(Err(e)) => { - // TCP RST or decode error from worker — kill only - // this stream. - tracing::warn!(err = ?e, "tcp stream read error from worker, closing connection"); - let _ = control_tx.send(ControlMessage::Kill).await; - break; - } - None => { - // this is allowed but we try to avoid it - // the logic is that the client will tell us when its is done and the server - // will close the connection naturally when the sentinel message is received - // the client closing early represents a transport error outside the control of the - // transport library - tracing::trace!("tcp stream was closed by client"); - break; - } - } - } - - } - } - } - - async fn network_send_handler( - socket_tx: FramedWrite, TwoPartCodec>, - control_rx: mpsc::Receiver, - ) { - let mut socket_tx = socket_tx; - let mut control_rx = control_rx; - - while let Some(control_msg) = control_rx.recv().await { - // Sentinel is a worker→frontend message; receiving one here means - // a producer is buggy. Skip rather than asserting — a stream-level - // bug must not panic the worker. - if matches!(control_msg, ControlMessage::Sentinel) { - tracing::warn!("received sentinel on send-side control channel; dropping"); - continue; - } - let bytes = match serde_json::to_vec(&control_msg) { - Ok(b) => b, - Err(e) => { - // Closed enum of small variants; serialization shouldn't - // fail. If it ever does, log and skip rather than panic. - tracing::warn!(err = ?e, ?control_msg, "failed to serialize control message"); - continue; - } - }; - let message = TwoPartMessage::from_header(bytes.into()); - match socket_tx.send(message).await { - Ok(_) => tracing::debug!(?control_msg, "issued control message"), - Err(e) => { - tracing::debug!(err = ?e, ?control_msg, "failed to send control message") - } - } - } - - let mut inner = socket_tx.into_inner(); - if let Err(e) = inner.flush().await { - tracing::debug!("failed to flush socket: {e}"); - } - if let Err(e) = inner.shutdown().await { - tracing::debug!("failed to shutdown socket: {e}"); - } - } -} - -enum ControlAction { - Continue, - Shutdown, -} - -fn process_control_message(message: Bytes) -> Result { - match serde_json::from_slice::(&message)? { - ControlMessage::Sentinel => { - // the client issued a sentinel message - // it has finished writing data and is now awaiting the server to close the connection - tracing::trace!("sentinel received; shutting down"); - Ok(ControlAction::Shutdown) - } - ControlMessage::Kill | ControlMessage::Stop => { - // Worker→frontend control direction only carries Sentinel. Kill/Stop - // here is a protocol violation; the caller turns this Err into a - // stream-local Kill rather than a process-fatal event. - anyhow::bail!("unexpected control message on response stream"); - } - } } #[cfg(test)] @@ -1648,8 +1362,7 @@ mod tests { use crate::engine::AsyncEngineContextProvider; use crate::pipeline::Context; use crate::pipeline::network::DEFAULT_SEND_BUFFER_COUNT; - use crate::pipeline::network::tcp::client::TcpClient; - use tokio::io::{AsyncWriteExt, ReadHalf, WriteHalf}; + use tokio::io::AsyncWriteExt; use tokio::net::TcpStream; // Mock resolver that always fails to simulate the fallback scenario @@ -1698,8 +1411,8 @@ mod tests { .connection_info .clone(); - let tcp_info: TcpStreamConnectionInfo = connection_info.try_into().unwrap(); - let socket_addr = tcp_info.address.parse::().unwrap(); + let mux_info: ResponseMuxConnectionInfo = connection_info.try_into().unwrap(); + let socket_addr = mux_info.address.parse::().unwrap(); // Should have a valid port assigned assert!( @@ -1709,13 +1422,13 @@ mod tests { println!( "Server created successfully with address: {}", - tcp_info.address + mux_info.address ); } /// The data-plane channel helper sizes the mpsc buffer from - /// `send_buffer_count` — this is the value `process_request_stream` / - /// `process_response_stream` feed it. `max_capacity()` reflects the + /// `send_buffer_count` — this is the value the request stream and response + /// mux mailbox paths feed it. `max_capacity()` reflects the /// channel's configured buffer, so a custom value and the default both /// reach the channel. Guards against regressing back to a hard-coded 64. #[test] @@ -1732,9 +1445,8 @@ mod tests { } /// `register` must thread `StreamOptions::send_buffer_count` through to the - /// stored `RequestedSendConnection` / `RequestedRecvConnection` (the - /// registration structs `process_*_stream` later destructure to size the - /// channel). Verified here against the real registration path. + /// stored request and response registration records. Verified here against + /// the real registration path. #[tokio::test] async fn register_threads_send_buffer_count_into_connection_structs() { let server = TcpStreamServer::new(ServerOptions::default()) @@ -1753,14 +1465,21 @@ mod tests { let state = server.state.lock(); assert_eq!(state.tx_subjects.len(), 1, "one request stream registered"); - assert_eq!(state.rx_subjects.len(), 1, "one response stream registered"); + assert_eq!( + server.response_pending.len(), + 1, + "one response stream registered" + ); assert!( state.tx_subjects.values().all(|c| c.send_buffer_count == 7), "send_buffer_count must reach RequestedSendConnection" ); assert!( - state.rx_subjects.values().all(|c| c.send_buffer_count == 7), - "send_buffer_count must reach RequestedRecvConnection" + server + .response_pending + .iter() + .all(|entry| entry.send_buffer_count == 7), + "send_buffer_count must reach the response mux registration" ); } @@ -1797,8 +1516,8 @@ mod tests { .connection_info .clone(); - let tcp_info: TcpStreamConnectionInfo = connection_info.try_into().unwrap(); - let socket_addr = tcp_info.address.parse::().unwrap(); + let mux_info: ResponseMuxConnectionInfo = connection_info.try_into().unwrap(); + let socket_addr = mux_info.address.parse::().unwrap(); // With the failing resolver, fallback should ALWAYS be used let ip = socket_addr.ip(); @@ -1849,8 +1568,8 @@ mod tests { let pending = server.register(options).await; let recv_stream = pending.recv_stream.unwrap(); let (conn_info, provider) = recv_stream.into_parts(); - let tcp_info: TcpStreamConnectionInfo = conn_info.try_into().unwrap(); - (tcp_info.subject, provider) + let mux_info: ResponseMuxConnectionInfo = conn_info.try_into().unwrap(); + (mux_info.stream_id.to_string(), provider) } /// Convenience constructor so tests don't repeat the struct literal. @@ -1892,11 +1611,11 @@ mod tests { let (send_info, send_provider) = send_stream.into_parts(); let (recv_info, recv_provider) = recv_stream.into_parts(); let send_tcp_info: TcpStreamConnectionInfo = send_info.try_into().unwrap(); - let recv_tcp_info: TcpStreamConnectionInfo = recv_info.try_into().unwrap(); + let recv_mux_info: ResponseMuxConnectionInfo = recv_info.try_into().unwrap(); ( send_tcp_info.subject, send_provider, - recv_tcp_info.subject, + recv_mux_info.stream_id.to_string(), recv_provider, ) } @@ -2045,16 +1764,12 @@ mod tests { let pending = server.register(options).await; let recv_stream = pending.recv_stream.unwrap(); - // Get the subject before dropping - let tcp_info: TcpStreamConnectionInfo = + // Get the stream ID before dropping. + let mux_info: ResponseMuxConnectionInfo = recv_stream.connection_info.clone().try_into().unwrap(); - let subject = tcp_info.subject.clone(); + let stream_id = mux_info.stream_id; - // Verify it's in rx_subjects - { - let state = server.state.lock(); - assert!(state.rx_subjects.contains_key(&subject)); - } + assert!(server.response_pending.contains_key(&stream_id)); // Drop the RegisteredStream -- RAII cleanup should fire drop(recv_stream); @@ -2062,14 +1777,10 @@ mod tests { // Give the spawned cleanup task a moment to run tokio::time::sleep(std::time::Duration::from_millis(50)).await; - // Verify it's been removed from rx_subjects - { - let state = server.state.lock(); - assert!( - !state.rx_subjects.contains_key(&subject), - "RAII cleanup should have removed the rx_subjects entry" - ); - } + assert!( + !server.response_pending.contains_key(&stream_id), + "RAII cleanup should have removed the pending mux entry" + ); } #[tokio::test] @@ -2087,9 +1798,9 @@ mod tests { let pending = server.register(options).await; let recv_stream = pending.recv_stream.unwrap(); - let tcp_info: TcpStreamConnectionInfo = + let mux_info: ResponseMuxConnectionInfo = recv_stream.connection_info.clone().try_into().unwrap(); - let subject = tcp_info.subject.clone(); + let stream_id = mux_info.stream_id; // Call into_parts to disarm the cleanup let (_conn_info, _provider) = recv_stream.into_parts(); @@ -2097,14 +1808,10 @@ mod tests { // Give any potential cleanup a moment to run tokio::time::sleep(std::time::Duration::from_millis(50)).await; - // The entry should still be in rx_subjects (cleanup was disarmed) - { - let state = server.state.lock(); - assert!( - state.rx_subjects.contains_key(&subject), - "into_parts() should disarm the RAII cleanup" - ); - } + assert!( + server.response_pending.contains_key(&stream_id), + "into_parts() should disarm the RAII cleanup" + ); } #[tokio::test] @@ -2374,257 +2081,132 @@ mod tests { assert_eq!(server.cancel_instance_streams(&id_b).await, 1); } - type TestFramedRead = FramedRead, TwoPartCodec>; - type TestFramedWrite = FramedWrite, TwoPartCodec>; - type TestResponseStream = (TestFramedRead, TestFramedWrite, StreamReceiver); - - /// Stand up a TcpStreamServer, register a response stream, connect a - /// client, drive the handshake + prologue, and return the client-side - /// framed reader/writer along with the receiver. - async fn open_registered_response_stream() -> TestResponseStream { - let options = ServerOptions::builder().port(0).build().unwrap(); - let server = TcpStreamServer::new_with_resolver(options, FailingIpResolver) - .await - .unwrap(); + async fn register_mux_response( + server: &TcpStreamServer, + ) -> ( + ResponseMuxConnectionInfo, + tokio::sync::oneshot::Receiver>, + ) { let context = Context::new(()); - let stream_options = StreamOptions::builder() - .context(context.context()) - .enable_request_stream(false) - .enable_response_stream(true) - .build() - .unwrap(); - let pending_connection = server.register(stream_options).await; - let registered_stream = pending_connection.recv_stream.unwrap(); - let (connection_info, stream_provider) = registered_stream.into_parts(); - let tcp_info: TcpStreamConnectionInfo = connection_info.try_into().unwrap(); - - let stream = TcpStream::connect(&tcp_info.address).await.unwrap(); - let (read_half, write_half) = tokio::io::split(stream); - let framed_reader = FramedRead::new(read_half, TwoPartCodec::default()); - let mut framed_writer = FramedWrite::new(write_half, TwoPartCodec::default()); - - let handshake = CallHomeHandshake { - subject: tcp_info.subject, - stream_type: StreamType::Response, - }; - framed_writer - .send(TwoPartMessage::from_header( - serde_json::to_vec(&handshake).unwrap().into(), - )) - .await - .unwrap(); - framed_writer - .send(TwoPartMessage::from_header( - serde_json::to_vec(&ResponseStreamPrologue { error: None }) - .unwrap() - .into(), - )) + let pending = server + .register( + StreamOptions::builder() + .context(context.context()) + .enable_request_stream(false) + .enable_response_stream(true) + .build() + .unwrap(), + ) .await + .recv_stream .unwrap(); - - // SAFETY (test-only): healthy localhost handshake always resolves all - // three layers; a panic here means the harness is broken. - let receiver = tokio::time::timeout(std::time::Duration::from_secs(1), stream_provider) - .await - .expect("server should establish response stream within timeout") - .expect("stream provider should not be dropped") - .expect("response stream should be accepted"); - - (framed_reader, framed_writer, receiver) + let (info, provider) = pending.into_parts(); + (info.try_into().unwrap(), provider) } - async fn recv_control_message(framed_reader: &mut TestFramedRead) -> ControlMessage { - // SAFETY (test-only): a misbehaving server in any of these layers is - // exactly the harness failure we want surfaced as a test panic. - let message = tokio::time::timeout(std::time::Duration::from_secs(1), framed_reader.next()) - .await - .expect("server should send a control message within timeout") - .expect("server should not close before sending control") - .expect("control message should decode"); - let (header, data) = message.optional_parts(); - assert!(data.is_none(), "control message should not contain data"); - serde_json::from_slice(header.expect("control header missing").as_ref()).unwrap() - } - - /// Sending an unexpected control message (Stop or Kill from the data - /// direction) is a protocol violation. The server's - /// network_receive_handler must reply with ControlMessage::Kill on - /// that stream alone, not panic. #[tokio::test] - async fn test_tcp_stream_server_sends_kill_on_unexpected_control_message() { - let (mut framed_reader, mut framed_writer, _receiver) = - open_registered_response_stream().await; + async fn pipelined_mux_handshake_preserves_data_and_isolates_malformed_prologue() { + use bytes::BytesMut; + use tokio_util::codec::Encoder; - framed_writer - .send(TwoPartMessage::from_header( - serde_json::to_vec(&ControlMessage::Stop).unwrap().into(), - )) - .await + let server = test_server().await; + let (bad_info, bad_provider) = register_mux_response(&server).await; + let (good_info, good_provider) = register_mux_response(&server).await; + assert_eq!(bad_info.address, good_info.address); + assert_eq!(bad_info.frontend_server_id, good_info.frontend_server_id); + + let mut socket = TcpStream::connect(&bad_info.address).await.unwrap(); + let handshake = ConnectionHandshake::ResponseMux { + version: RESPONSE_MUX_VERSION, + frontend_server_id: bad_info.frontend_server_id, + connection_id: uuid::Uuid::new_v4(), + }; + let mut wire = BytesMut::new(); + TwoPartCodec::default() + .encode( + TwoPartMessage::from_header(serde_json::to_vec(&handshake).unwrap().into()), + &mut wire, + ) + .unwrap(); + let mut mux_codec = MuxCodec::default(); + mux_codec + .encode( + MuxFrame::new( + MuxFrameKind::Prologue, + bad_info.stream_id, + Bytes::from_static(b"{"), + ), + &mut wire, + ) + .unwrap(); + mux_codec + .encode( + MuxFrame::new( + MuxFrameKind::Prologue, + good_info.stream_id, + serde_json::to_vec(&ResponseStreamPrologue { error: None }) + .unwrap() + .into(), + ), + &mut wire, + ) + .unwrap(); + mux_codec + .encode( + MuxFrame::new( + MuxFrameKind::Data, + good_info.stream_id, + Bytes::from_static(b"preserved"), + ), + &mut wire, + ) + .unwrap(); + mux_codec + .encode( + MuxFrame::empty(MuxFrameKind::End, good_info.stream_id), + &mut wire, + ) .unwrap(); - assert_eq!( - recv_control_message(&mut framed_reader).await, - ControlMessage::Kill, - "unexpected control message should kill only this stream" - ); - } - - /// A framing/decode error from the worker side is unrecoverable for - /// this stream but must not panic the worker. Server should send Kill - /// and tear down only this connection. - #[tokio::test] - async fn test_tcp_stream_server_sends_kill_on_read_error() { - let (mut framed_reader, framed_writer, _receiver) = open_registered_response_stream().await; - - let mut raw_writer = framed_writer.into_inner(); - raw_writer.write_all(&[0u8; 8]).await.unwrap(); - raw_writer.shutdown().await.unwrap(); - - assert_eq!( - recv_control_message(&mut framed_reader).await, - ControlMessage::Kill, - "framing read error should kill only this stream" - ); - } - - /// Sentinel is supposed to be header-only. A misbehaving client that - /// attaches a data payload must not panic the worker via assert!(). - #[tokio::test] - async fn test_tcp_stream_server_sends_kill_on_sentinel_with_data() { - let (mut framed_reader, mut framed_writer, _receiver) = - open_registered_response_stream().await; - - let header = serde_json::to_vec(&ControlMessage::Sentinel) - .unwrap() - .into(); - framed_writer - .send(TwoPartMessage::from_parts( - header, - Bytes::from_static(b"unexpected payload"), - )) + // A single write makes handshake and mux frames available to the + // handshake decoder together, exercising decoder-buffer preservation. + socket.write_all(&wire).await.unwrap(); + let mut reader = FramedRead::new(socket, TwoPartCodec::default()); + let ack = tokio::time::timeout(Duration::from_secs(1), reader.next()) .await + .unwrap() + .unwrap() .unwrap(); - assert_eq!( - recv_control_message(&mut framed_reader).await, - ControlMessage::Kill, - "Sentinel with data should kill only this stream" + MuxFrame::try_from_two_part(ack).unwrap(), + MuxFrame::connection_ack(0) ); - } + let mut reader = reader.map_decoder(|_| MuxCodec::default()); - /// The prologue must be a HeaderOnly frame. A non-HeaderOnly prologue - /// (data-only or mixed) must surface as Err to the requester rather - /// than panic the worker. - #[tokio::test] - async fn test_tcp_stream_server_returns_error_on_invalid_prologue() { - let options = ServerOptions::builder().port(0).build().unwrap(); - let server = TcpStreamServer::new_with_resolver(options, FailingIpResolver) + let bad = tokio::time::timeout(Duration::from_secs(1), bad_provider) .await + .unwrap() .unwrap(); - let context = Context::new(()); - let stream_options = StreamOptions::builder() - .context(context.context()) - .enable_request_stream(false) - .enable_response_stream(true) - .build() - .unwrap(); - let pending_connection = server.register(stream_options).await; - let registered_stream = pending_connection.recv_stream.unwrap(); - let (connection_info, stream_provider) = registered_stream.into_parts(); - let tcp_info: TcpStreamConnectionInfo = connection_info.try_into().unwrap(); - - let stream = TcpStream::connect(&tcp_info.address).await.unwrap(); - let (_read_half, write_half) = tokio::io::split(stream); - let mut framed_writer = FramedWrite::new(write_half, TwoPartCodec::default()); + assert!( + bad.err() + .is_some_and(|error| error.contains("invalid response mux prologue")) + ); - let handshake = CallHomeHandshake { - subject: tcp_info.subject, - stream_type: StreamType::Response, - }; - framed_writer - .send(TwoPartMessage::from_header( - serde_json::to_vec(&handshake).unwrap().into(), - )) + let mut good = tokio::time::timeout(Duration::from_secs(1), good_provider) .await + .unwrap() + .unwrap() .unwrap(); + assert_eq!(good.recv().await.unwrap(), Bytes::from_static(b"preserved")); + assert!(good.recv().await.is_none()); - // Send a data-only frame in the prologue slot. - framed_writer - .send(TwoPartMessage::from_data(Bytes::from_static( - b"not a prologue", - ))) + let reset = tokio::time::timeout(Duration::from_secs(1), reader.next()) .await + .unwrap() + .unwrap() .unwrap(); - - let outcome = tokio::time::timeout(std::time::Duration::from_secs(1), stream_provider) - .await - .expect("stream provider should resolve quickly") - .expect("stream provider channel should not be dropped"); - // StreamReceiver doesn't impl Debug, so we can't use `.expect_err`. - match outcome { - Err(err) => assert!( - err.contains("malformed prologue"), - "expected malformed-prologue error, got: {err}" - ), - Ok(_) => panic!("invalid prologue should produce an error, but got Ok"), - } - } - - #[tokio::test(flavor = "multi_thread", worker_threads = 4)] - async fn test_concurrent_response_registration_and_call_home() { - const STREAMS: usize = 128; - - let result = time::timeout(Duration::from_secs(20), async { - let server = test_server().await; - let mut pending_streams = Vec::with_capacity(STREAMS); - let mut client_tasks = Vec::with_capacity(STREAMS); - - for idx in 0..STREAMS { - let context = Context::new(()); - let options = StreamOptions::builder() - .context(context.context()) - .enable_request_stream(false) - .enable_response_stream(true) - .build() - .unwrap(); - - let pending = server.register(options).await; - let registered_stream = pending.recv_stream.unwrap(); - let (connection_info, stream_provider) = registered_stream.into_parts(); - let client_context = - Context::with_id_and_metadata((), context.id().to_string(), Default::default()); - let payload = Bytes::from(format!("payload-{idx}")); - - pending_streams.push((idx, payload.clone(), stream_provider)); - client_tasks.push(tokio::spawn(async move { - let mut sender = TcpClient::create_response_stream( - client_context.context(), - connection_info, - None, - ) - .await - .unwrap(); - sender.send_prologue(None).await.unwrap(); - sender.send(payload).await.unwrap(); - })); - } - - for task in client_tasks { - task.await.unwrap(); - } - - for (idx, expected, stream_provider) in pending_streams { - let mut stream = stream_provider.await.unwrap().unwrap(); - let actual = stream.recv().await.unwrap(); - assert_eq!(actual, expected, "payload mismatch for stream {idx}"); - } - }) - .await; - - assert!( - result.is_ok(), - "concurrent response registration and call-home timed out" - ); + assert_eq!(reset.kind, MuxFrameKind::Reset); + assert_eq!(reset.stream_id, bad_info.stream_id); } // ==================== request_stream_send_handler integration tests ==================== From 4cae994e645d0b656464c6be7c35f7024fefeca1 Mon Sep 17 00:00:00 2001 From: jthomson04 Date: Wed, 22 Jul 2026 11:28:12 -0700 Subject: [PATCH 3/3] perf(runtime): simplify TCP response mux hot path Signed-off-by: jthomson04 --- docs/design-docs/request-plane.md | 9 +- lib/runtime/src/config/environment_names.rs | 5 - lib/runtime/src/metrics/response_mux.rs | 293 +-- lib/runtime/src/pipeline/network.rs | 75 +- lib/runtime/src/pipeline/network/tcp.rs | 17 +- lib/runtime/src/pipeline/network/tcp/mux.rs | 277 +-- .../src/pipeline/network/tcp/mux/client.rs | 2130 ++++------------- .../src/pipeline/network/tcp/server.rs | 699 +++--- 8 files changed, 856 insertions(+), 2649 deletions(-) diff --git a/docs/design-docs/request-plane.md b/docs/design-docs/request-plane.md index baad24b45533..b32ab4638bd0 100644 --- a/docs/design-docs/request-plane.md +++ b/docs/design-docs/request-plane.md @@ -103,7 +103,11 @@ Additional TCP-specific environment variables: The TCP response path always uses the versioned `tcp_response_mux_v1` protocol. Each worker process maintains four persistent response connections to each frontend it communicates with, and logical responses share those connections. For example, eight worker processes paired with one frontend maintain 32 physical response connections after warmup. Request streams remain on their existing dedicated sockets. -Data frames are batched for up to 1 ms by default to reduce small TCP writes and per-packet operating system overhead. Set `DYN_TCP_RESPONSE_BATCH_INTERVAL_MS=0` for opportunistic batching without an intentional delay. Values above 100 ms or malformed values prevent runtime initialization. +Response connections use the mux codec from the first byte. A worker sends a binary `ConnectionHello` containing the protocol version and frontend UUID, and the frontend returns an empty `ConnectionReady`. Every subsequent frame has a 21-byte header containing the payload length, frame kind, and logical-stream UUID. Successful stream prologues and reset frames have empty payloads; a failed prologue carries its error as UTF-8 text. + +Workers submit prologue and reset frames through an urgent writer lane and submit data and end frames through an ordered bounded lane. Data frames are batched for up to 1 ms by default to reduce small TCP writes and per-packet operating system overhead. Set `DYN_TCP_RESPONSE_BATCH_INTERVAL_MS=0` for opportunistic batching without an intentional delay. Values above 100 ms or malformed values prevent runtime initialization. + +Per-stream credits isolate slow consumers. Each logical stream can queue at most eight ordered frames, and a bounded connection queue limits total userspace memory. The mux does not add a connection-level credit window: TCP provides physical-connection backpressure. Host pools live for the worker-process lifetime, and a process-wide maintenance task warms new frontend pools to four connections and replaces failed connections. The response mux supports the following tuning variables: @@ -111,8 +115,7 @@ The response mux supports the following tuning variables: - `DYN_TCP_RESPONSE_BATCH_MAX_BYTES`: Maximum encoded bytes per batch (default: `65536`) - `DYN_TCP_RESPONSE_BATCH_MAX_FRAMES`: Maximum frames per batch (default: `64`) - `DYN_TCP_RESPONSE_STREAM_WINDOW_BYTES`: Initial per-stream flow-control window (default: `262144`) -- `DYN_TCP_RESPONSE_CONNECTION_WINDOW_BYTES`: Initial per-connection flow-control window (default: `262144`) -- `DYN_TCP_RESPONSE_PACKET_METRICS`: Enable detailed per-frame timing and Linux TCP segment diagnostics (`0` or `1`, default: `0`) +- `DYN_TCP_RESPONSE_PACKET_METRICS`: Enable Linux TCP segment diagnostics (`0` or `1`, default: `0`) Response mux versions do not fall back to the former dedicated response protocol. Upgrade frontends and workers together so both sides support `tcp_response_mux_v1`; mixed versions reject the response connection during its handshake. diff --git a/lib/runtime/src/config/environment_names.rs b/lib/runtime/src/config/environment_names.rs index c38f01914576..7916b68ee165 100644 --- a/lib/runtime/src/config/environment_names.rs +++ b/lib/runtime/src/config/environment_names.rs @@ -638,10 +638,6 @@ pub mod tcp_response_stream { /// Per-stream response flow-control window in bytes. pub const DYN_TCP_RESPONSE_STREAM_WINDOW_BYTES: &str = "DYN_TCP_RESPONSE_STREAM_WINDOW_BYTES"; - /// Per-connection response flow-control window in bytes. - pub const DYN_TCP_RESPONSE_CONNECTION_WINDOW_BYTES: &str = - "DYN_TCP_RESPONSE_CONNECTION_WINDOW_BYTES"; - /// Enables diagnostic TCP_INFO data-segment accounting for response sockets. pub const DYN_TCP_RESPONSE_PACKET_METRICS: &str = "DYN_TCP_RESPONSE_PACKET_METRICS"; } @@ -895,7 +891,6 @@ mod tests { tcp_response_stream::DYN_TCP_RESPONSE_BATCH_MAX_BYTES, tcp_response_stream::DYN_TCP_RESPONSE_BATCH_MAX_FRAMES, tcp_response_stream::DYN_TCP_RESPONSE_STREAM_WINDOW_BYTES, - tcp_response_stream::DYN_TCP_RESPONSE_CONNECTION_WINDOW_BYTES, tcp_response_stream::DYN_TCP_RESPONSE_PACKET_METRICS, // Event Plane event_plane::DYN_EVENT_PLANE, diff --git a/lib/runtime/src/metrics/response_mux.rs b/lib/runtime/src/metrics/response_mux.rs index c10f25b59b41..2ad5141dda64 100644 --- a/lib/runtime/src/metrics/response_mux.rs +++ b/lib/runtime/src/metrics/response_mux.rs @@ -58,86 +58,6 @@ pub static SETUP_SECONDS: Lazy = Lazy::new(|| { .expect("response mux setup histogram") }); -pub static FRAMES_TOTAL: Lazy = Lazy::new(|| { - IntCounterVec::new( - Opts::new( - "dynamo_tcp_response_mux_frames_total", - "Multiplexed response frames by wire direction and type", - ), - &["direction", "frame_type"], - ) - .expect("response mux frame counter") -}); - -/// Pre-bound frame counters for one wire direction. Keeping these handles next -/// to a connection avoids a label-map lookup for every generated token. -pub struct FrameCounters { - prologue: IntCounter, - data: IntCounter, - end: IntCounter, - stop: IntCounter, - kill: IntCounter, - window_update: IntCounter, - reset: IntCounter, - connection_ack: IntCounter, -} - -impl FrameCounters { - pub fn for_direction(direction: &str) -> Self { - let counter = |frame_type| { - FRAMES_TOTAL - .with_label_values(&[direction, frame_type]) - .clone() - }; - Self { - prologue: counter("prologue"), - data: counter("data"), - end: counter("end"), - stop: counter("stop"), - kill: counter("kill"), - window_update: counter("window_update"), - reset: counter("reset"), - connection_ack: counter("connection_ack"), - } - } - - pub fn inc(&self, frame_type: &str) { - match frame_type { - "prologue" => self.prologue.inc(), - "data" => self.data.inc(), - "end" => self.end.inc(), - "stop" => self.stop.inc(), - "kill" => self.kill.inc(), - "window_update" => self.window_update.inc(), - "reset" => self.reset.inc(), - "connection_ack" => self.connection_ack.inc(), - _ => debug_assert!(false, "unknown response mux frame type {frame_type}"), - } - } -} - -pub static WRITER_QUEUE_DEPTH: Lazy = Lazy::new(|| { - IntGaugeVec::new( - Opts::new( - "dynamo_tcp_response_mux_writer_queue_depth", - "Queued response-mux frames waiting for the shared writer", - ), - &["role"], - ) - .expect("response mux queue gauge") -}); - -pub static QUEUED_BYTES: Lazy = Lazy::new(|| { - IntGaugeVec::new( - Opts::new( - "dynamo_tcp_response_mux_queued_bytes", - "Encoded response bytes queued in connection writers", - ), - &["role"], - ) - .expect("response mux queued byte gauge") -}); - pub static FRAMES_PER_WRITE: Lazy = Lazy::new(|| { HistogramVec::new( HistogramOpts::new( @@ -150,32 +70,6 @@ pub static FRAMES_PER_WRITE: Lazy = Lazy::new(|| { .expect("response mux frames-per-write histogram") }); -pub static BATCH_BYTES: Lazy = Lazy::new(|| { - HistogramVec::new( - HistogramOpts::new( - "dynamo_tcp_response_mux_batch_bytes", - "Encoded response bytes in each physical TCP write", - ) - .buckets(vec![ - 64.0, 256.0, 1024.0, 4096.0, 16_384.0, 65_536.0, 262_144.0, - ]), - &["role"], - ) - .expect("response mux batch byte histogram") -}); - -pub static BATCH_WAIT_SECONDS: Lazy = Lazy::new(|| { - HistogramVec::new( - HistogramOpts::new( - "dynamo_tcp_response_mux_batch_wait_seconds", - "Observed userspace wait from first selected data frame to write", - ) - .buckets(vec![0.0, 0.0001, 0.0005, 0.001, 0.002, 0.005, 0.01, 0.1]), - &["role"], - ) - .expect("response mux batch wait histogram") -}); - pub static CONFIGURED_BATCH_INTERVAL_MS: Lazy = Lazy::new(|| { IntGaugeVec::new( Opts::new( @@ -202,32 +96,18 @@ pub static DATA_SEGMENTS_TOTAL: Lazy = Lazy::new(|| { IntCounterVec::new( Opts::new( "dynamo_tcp_response_data_segments_total", - "Kernel TCP data segments sent on response sockets when diagnostic packet metrics are enabled", + "Kernel TCP data segments sent on response sockets when packet metrics are enabled", ), &["transport", "role"], ) .expect("response TCP data segment counter") }); -pub static QUEUE_RESIDENCE_SECONDS: Lazy = Lazy::new(|| { - HistogramVec::new( - HistogramOpts::new( - "dynamo_tcp_response_mux_queue_residence_seconds", - "Time logical response frames wait before their physical write", - ) - .buckets(vec![ - 0.000001, 0.00001, 0.0001, 0.001, 0.005, 0.01, 0.1, 1.0, - ]), - &["role"], - ) - .expect("response mux queue residence histogram") -}); - pub static RESETS_TOTAL: Lazy = Lazy::new(|| { IntCounterVec::new( Opts::new( "dynamo_tcp_response_mux_resets_total", - "Logical response streams reset by role and low-cardinality reason", + "Logical response streams reset by role and reason", ), &["role", "reason"], ) @@ -245,100 +125,15 @@ pub static RECONNECTS_TOTAL: Lazy = Lazy::new(|| { .expect("response mux reconnect counter") }); -pub static FLOW_CONTROL_STALL_SECONDS: Lazy = Lazy::new(|| { - Histogram::with_opts( - HistogramOpts::new( - "dynamo_tcp_response_mux_flow_control_stall_seconds", - "Time response producers wait for stream-local credits", - ) - .buckets(vec![0.000001, 0.00001, 0.0001, 0.001, 0.01, 0.1, 1.0, 10.0]), - ) - .expect("response mux flow-control histogram") -}); - -pub static CONNECTION_FLOW_CONTROL_STALL_SECONDS: Lazy = Lazy::new(|| { - Histogram::with_opts( - HistogramOpts::new( - "dynamo_tcp_response_mux_connection_flow_control_stall_seconds", - "Time response producers wait for physical-connection Data credits", - ) - .buckets(vec![0.000001, 0.00001, 0.0001, 0.001, 0.01, 0.1, 1.0, 10.0]), - ) - .expect("response mux connection flow-control histogram") -}); - -pub static WRITER_ADMISSION_STALL_SECONDS: Lazy = Lazy::new(|| { - Histogram::with_opts( - HistogramOpts::new( - "dynamo_tcp_response_mux_writer_admission_stall_seconds", - "Time response producers wait for their stream-local writer queue", - ) - .buckets(vec![0.000001, 0.00001, 0.0001, 0.001, 0.01, 0.1, 1.0, 10.0]), - ) - .expect("response mux writer-admission histogram") -}); - -pub static QUEUED_BYTE_ADMISSION_STALL_SECONDS: Lazy = Lazy::new(|| { - Histogram::with_opts( - HistogramOpts::new( - "dynamo_tcp_response_mux_queued_byte_admission_stall_seconds", - "Time response producers wait for connection-wide queued-byte capacity", - ) - .buckets(vec![0.000001, 0.00001, 0.0001, 0.001, 0.01, 0.1, 1.0, 10.0]), - ) - .expect("response mux queued-byte admission histogram") -}); - -pub static STREAM_WRITER_QUEUE_OCCUPANCY: Lazy = Lazy::new(|| { - Histogram::with_opts( - HistogramOpts::new( - "dynamo_tcp_response_mux_stream_writer_queue_occupancy", - "Stream-local writer queue occupancy after admission", - ) - .buckets(vec![1.0, 2.0, 4.0, 8.0]), - ) - .expect("response mux stream writer queue occupancy histogram") -}); - -pub static READY_STREAMS: Lazy = Lazy::new(|| { - IntGaugeVec::new( - Opts::new( - "dynamo_tcp_response_mux_ready_streams", - "Logical streams currently scheduled on a fair connection writer", - ), - &["role"], - ) - .expect("response mux ready stream gauge") -}); - -pub static PRIORITY_QUEUE_RESIDENCE_SECONDS: Lazy = Lazy::new(|| { - Histogram::with_opts( - HistogramOpts::new( - "dynamo_tcp_response_mux_priority_queue_residence_seconds", - "Time prologue and reset frames wait in the priority writer lane", - ) - .buckets(vec![0.000001, 0.00001, 0.0001, 0.001, 0.01, 0.1, 1.0]), - ) - .expect("response mux priority queue residence histogram") -}); - -pub static ROUND_ROBIN_TURNS_TOTAL: Lazy = Lazy::new(|| { - IntCounter::new( - "dynamo_tcp_response_mux_round_robin_turns_total", - "Frames selected through the fair per-stream writer ring", - ) - .expect("response mux round-robin turn counter") -}); - -pub static WINDOW_UPDATES_TOTAL: Lazy = Lazy::new(|| { +pub static STALLS_TOTAL: Lazy = Lazy::new(|| { IntCounterVec::new( Opts::new( - "dynamo_tcp_response_mux_window_updates_total", - "Response-mux window update frames", + "dynamo_tcp_response_mux_stalls_total", + "Response producer stalls by bounded admission point", ), - &["direction"], + &["kind"], ) - .expect("response mux window-update counter") + .expect("response mux stall counter") }); pub static CONNECTION_LOST_STREAMS_TOTAL: Lazy = Lazy::new(|| { @@ -382,21 +177,10 @@ pub fn ensure_registered(registry: &MetricsRegistry) { Box::new(SETUP_SECONDS.clone()), "response_mux_setup_seconds", ); - registry.add_metric_or_warn(Box::new(FRAMES_TOTAL.clone()), "response_mux_frames_total"); - registry.add_metric_or_warn( - Box::new(WRITER_QUEUE_DEPTH.clone()), - "response_mux_writer_queue_depth", - ); - registry.add_metric_or_warn(Box::new(QUEUED_BYTES.clone()), "response_mux_queued_bytes"); registry.add_metric_or_warn( Box::new(FRAMES_PER_WRITE.clone()), "response_mux_frames_per_write", ); - registry.add_metric_or_warn(Box::new(BATCH_BYTES.clone()), "response_mux_batch_bytes"); - registry.add_metric_or_warn( - Box::new(BATCH_WAIT_SECONDS.clone()), - "response_mux_batch_wait_seconds", - ); registry.add_metric_or_warn( Box::new(CONFIGURED_BATCH_INTERVAL_MS.clone()), "response_mux_configured_batch_interval_ms", @@ -409,51 +193,12 @@ pub fn ensure_registered(registry: &MetricsRegistry) { Box::new(DATA_SEGMENTS_TOTAL.clone()), "response_data_segments_total", ); - registry.add_metric_or_warn( - Box::new(QUEUE_RESIDENCE_SECONDS.clone()), - "response_mux_queue_residence_seconds", - ); registry.add_metric_or_warn(Box::new(RESETS_TOTAL.clone()), "response_mux_resets_total"); registry.add_metric_or_warn( Box::new(RECONNECTS_TOTAL.clone()), "response_mux_reconnects_total", ); - registry.add_metric_or_warn( - Box::new(FLOW_CONTROL_STALL_SECONDS.clone()), - "response_mux_flow_control_stall_seconds", - ); - registry.add_metric_or_warn( - Box::new(CONNECTION_FLOW_CONTROL_STALL_SECONDS.clone()), - "response_mux_connection_flow_control_stall_seconds", - ); - registry.add_metric_or_warn( - Box::new(WRITER_ADMISSION_STALL_SECONDS.clone()), - "response_mux_writer_admission_stall_seconds", - ); - registry.add_metric_or_warn( - Box::new(QUEUED_BYTE_ADMISSION_STALL_SECONDS.clone()), - "response_mux_queued_byte_admission_stall_seconds", - ); - registry.add_metric_or_warn( - Box::new(STREAM_WRITER_QUEUE_OCCUPANCY.clone()), - "response_mux_stream_writer_queue_occupancy", - ); - registry.add_metric_or_warn( - Box::new(READY_STREAMS.clone()), - "response_mux_ready_streams", - ); - registry.add_metric_or_warn( - Box::new(PRIORITY_QUEUE_RESIDENCE_SECONDS.clone()), - "response_mux_priority_queue_residence_seconds", - ); - registry.add_metric_or_warn( - Box::new(ROUND_ROBIN_TURNS_TOTAL.clone()), - "response_mux_round_robin_turns_total", - ); - registry.add_metric_or_warn( - Box::new(WINDOW_UPDATES_TOTAL.clone()), - "response_mux_window_updates_total", - ); + registry.add_metric_or_warn(Box::new(STALLS_TOTAL.clone()), "response_mux_stalls_total"); registry.add_metric_or_warn( Box::new(CONNECTION_LOST_STREAMS_TOTAL.clone()), "response_mux_connection_lost_streams_total", @@ -464,30 +209,12 @@ pub fn ensure_registered(registry: &MetricsRegistry) { mod tests { use super::*; - fn contains_response_mux_metrics(registry: &MetricsRegistry) -> bool { - registry - .prometheus_registry - .read() - .expect("metrics registry read lock") - .gather() - .iter() - .any(|family| family.name() == "dynamo_tcp_response_mux_active_connections") - } - #[test] - fn registration_is_deduplicated_by_registry_identity() { + fn metrics_register_with_multiple_registries() { let first = MetricsRegistry::new(); - let first_clone = first.clone(); let second = MetricsRegistry::new(); - ACTIVE_CONNECTIONS - .with_label_values(&["registry_test"]) - .set(0); - ensure_registered(&first); - ensure_registered(&first_clone); + ensure_registered(&first); ensure_registered(&second); - - assert!(contains_response_mux_metrics(&first)); - assert!(contains_response_mux_metrics(&second)); } } diff --git a/lib/runtime/src/pipeline/network.rs b/lib/runtime/src/pipeline/network.rs index 53ff9d783548..6d03ba8c3b9c 100644 --- a/lib/runtime/src/pipeline/network.rs +++ b/lib/runtime/src/pipeline/network.rs @@ -179,17 +179,6 @@ pub enum ControlMessage { Sentinel, } -/// This is the first message in a `ResponseStream`. This is not a message that gets process -/// by the general pipeline, but is a control message that is awaited before the -/// [`AsyncEngine::generate`] method is allowed to return. -/// -/// If an error is present, the [`AsyncEngine::generate`] method will return the error instead -/// of returning the `ResponseStream`. -#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)] -pub struct ResponseStreamPrologue { - error: Option, -} - pub type StreamProvider = tokio::sync::oneshot::Receiver>; /// Owning `Drop` here (rather than on `RegisteredStream`) lets `into_parts()` @@ -354,46 +343,25 @@ mod registered_stream_tests { // } // } -// this probably needs to be come a ResponseStreamSender -// since the prologue in this scenario sender telling the receiver -// that all is good and it's ready to send -// -// in the RequestStreamSender, the prologue would be coming from the -// receiver, so the sender would have to await the prologue which if -// was not an error, would indicate the RequestStreamReceiver is read -// to receive data. -#[async_trait] -pub(crate) trait MultiplexedStreamSender: Send + Sync { - async fn send_data(&self, data: Bytes) -> Result<()>; - async fn send_prologue(&self, error: Option) -> Result<(), String>; - async fn finish(&self) -> Result<()>; -} - enum StreamSenderInner { Dedicated(tokio::sync::mpsc::Sender), - Multiplexed(Arc), + Multiplexed(tcp::mux::client::MuxResponseStreamSender), } pub struct StreamSender { inner: StreamSenderInner, - prologue: Option, } impl StreamSender { - pub(crate) fn dedicated( - tx: tokio::sync::mpsc::Sender, - prologue: Option, - ) -> Self { + pub(crate) fn dedicated(tx: tokio::sync::mpsc::Sender) -> Self { Self { inner: StreamSenderInner::Dedicated(tx), - prologue, } } - pub(crate) fn multiplexed(sender: Arc) -> Self { + pub(crate) fn multiplexed(sender: tcp::mux::client::MuxResponseStreamSender) -> Self { Self { inner: StreamSenderInner::Multiplexed(sender), - prologue: Some(ResponseStreamPrologue { error: None }), } } @@ -418,41 +386,20 @@ impl StreamSender { } } - #[allow(clippy::needless_update)] pub async fn send_prologue(&mut self, error: Option) -> Result<(), String> { - // leaving the original logic in place for now - // error overrides the dissolved prologue, but the only field on `ResponseStreamPrologue` is `error` - // so the second argument can never be used, and the value of error passed by the caller would always be used - if let Some(_prologue) = self.prologue.take() { - // let prologue = ResponseStreamPrologue { error, ..prologue }; - let prologue = ResponseStreamPrologue { error }; - let header_bytes: Bytes = match serde_json::to_vec(&prologue) { - Ok(b) => b.into(), - Err(err) => { - tracing::error!(%err, "send_prologue: ResponseStreamPrologue did not serialize to a JSON array"); - return Err("Invalid prologue".to_string()); - } - }; - match &self.inner { - StreamSenderInner::Dedicated(tx) => tx - .send(TwoPartMessage::from_header(header_bytes)) - .await - .map_err(|e| e.to_string())?, - StreamSenderInner::Multiplexed(sender) => { - sender.send_prologue(prologue.error).await? - } + match &mut self.inner { + StreamSenderInner::Dedicated(_) => { + Err("prologues are not valid on a dedicated request sender".to_string()) } - } else { - panic!("Prologue already sent; or not set; logic error"); + StreamSenderInner::Multiplexed(sender) => sender.send_prologue(error).await, } - Ok(()) } /// Finish one logical response stream without closing a shared physical /// connection. Dedicated request-stream senders retain their existing /// drop-driven lifecycle and therefore have no explicit finish action. - pub async fn finish(&self) -> Result<()> { - match &self.inner { + pub async fn finish(self) -> Result<()> { + match self.inner { StreamSenderInner::Dedicated(_) => Ok(()), StreamSenderInner::Multiplexed(sender) => sender.finish().await, } @@ -489,8 +436,8 @@ impl AsRef<[u8]> for StreamRxItem { pub(crate) struct StreamReceiverHooks { pub context: Arc, pub window_update_threshold: usize, - pub on_window_update: Arc, - pub on_close: Arc, + pub on_window_update: Box, + pub on_close: Box, } pub struct StreamReceiver { diff --git a/lib/runtime/src/pipeline/network/tcp.rs b/lib/runtime/src/pipeline/network/tcp.rs index e554368bc59d..28db6c29f5ba 100644 --- a/lib/runtime/src/pipeline/network/tcp.rs +++ b/lib/runtime/src/pipeline/network/tcp.rs @@ -39,8 +39,8 @@ //! //! # Server-Client Interaction //! -//! See the test cases below for detailed examples. Note that the response stream expects the client -//! to send a [`ResponseStreamPrologue`] in order to properly establish the stream. +//! See the test cases below for detailed examples. A response `Prologue` activates +//! the logical stream before any `Data` frames are sent. //! //! # Stream Types //! @@ -77,8 +77,10 @@ //! //! A dedicated request socket starts with a `CallHomeHandshake` carrying its //! subject and [`StreamType::Request`]. A response-mux connection instead uses -//! a versioned handshake containing the frontend and physical-connection UUIDs; -//! incompatible response versions are rejected without a legacy fallback. +//! [`mux::MuxCodec`] from the first byte: `ConnectionHello` carries the protocol +//! version and frontend UUID, and the frontend returns an empty +//! `ConnectionReady`. Incompatible response versions are rejected without a +//! legacy fallback. //! //! # Control / Shutdown Protocol //! @@ -101,7 +103,7 @@ //! //! ## Response mux (downstream → upstream) — bidirectional //! -//! - Upstream writes: mux `Stop`, `Kill`, `WindowUpdate`, `ConnectionAck`, and `Reset` frames. +//! - Upstream writes: mux `Stop`, `Kill`, `WindowUpdate`, and `Reset` frames. //! - Downstream writes: mux `Prologue`, `Data`, `End`, and `Reset` frames. //! //! ## Request stream (upstream → downstream) — unidirectional after the handshake @@ -121,8 +123,8 @@ use serde::{Deserialize, Serialize}; #[allow(unused_imports)] use super::{ - ConnectionInfo, PendingConnections, RegisteredStream, ResponseService, ResponseStreamPrologue, - StreamOptions, StreamReceiver, StreamSender, StreamType, codec::TwoPartCodec, + ConnectionInfo, PendingConnections, RegisteredStream, ResponseService, StreamOptions, + StreamReceiver, StreamSender, StreamType, codec::TwoPartCodec, }; const TCP_TRANSPORT: &str = "tcp_server"; @@ -228,7 +230,6 @@ pub struct ResponseMuxConnectionInfo { pub frontend_server_id: uuid::Uuid, pub stream_id: uuid::Uuid, pub context: String, - pub version: u8, } impl From for ConnectionInfo { diff --git a/lib/runtime/src/pipeline/network/tcp/mux.rs b/lib/runtime/src/pipeline/network/tcp/mux.rs index 52b72d57157a..13b1632cceef 100644 --- a/lib/runtime/src/pipeline/network/tcp/mux.rs +++ b/lib/runtime/src/pipeline/network/tcp/mux.rs @@ -3,22 +3,15 @@ //! Multiplexed TCP response-stream protocol. //! -//! A short [`TwoPartCodec`] handshake validates the version and frontend -//! identity. The persistent connection then switches to [`MuxCodec`], whose -//! compact fixed-width header carries the frame kind and logical stream UUID. -//! Connection writers can batch those frames without changing their wire -//! representation. +//! [`MuxCodec`] carries both the connection handshake and logical response +//! streams. Its compact fixed-width header lets connection writers batch frames +//! without copying their payloads. use std::{io, sync::OnceLock, time::Duration}; use bytes::{BufMut, Bytes, BytesMut}; -use serde::{Deserialize, Serialize}; use uuid::Uuid; -use crate::pipeline::{ - error::TwoPartCodecError, - network::codec::{TwoPartCodec, TwoPartMessage, TwoPartMessageType}, -}; use tokio_util::codec::{Decoder, Encoder}; pub mod client; @@ -27,7 +20,7 @@ pub const RESPONSE_MUX_VERSION: u8 = 1; pub const RESPONSE_MUX_POOL_SIZE: usize = 4; pub const RESPONSE_MUX_WRITER_QUEUE: usize = 4096; pub const RESPONSE_MUX_STREAM_WRITER_QUEUE: usize = 8; -pub const RESPONSE_MUX_IDLE_TTL_SECS: u64 = 300; +pub const RESPONSE_MUX_CONNECTION_QUEUE_BYTES: usize = 262_144; pub const RESPONSE_MUX_CONNECT_TIMEOUT_SECS: u64 = 5; pub const RESPONSE_MUX_DEFAULT_BATCH_INTERVAL_MS: u64 = 1; @@ -35,10 +28,7 @@ pub const RESPONSE_MUX_MAX_BATCH_INTERVAL_MS: u64 = 100; pub const RESPONSE_MUX_DEFAULT_BATCH_MAX_BYTES: usize = 65_536; pub const RESPONSE_MUX_DEFAULT_BATCH_MAX_FRAMES: usize = 64; pub const RESPONSE_MUX_DEFAULT_STREAM_WINDOW_BYTES: usize = 262_144; -pub const RESPONSE_MUX_DEFAULT_CONNECTION_WINDOW_BYTES: usize = 262_144; pub const RESPONSE_MUX_CREDIT_UPDATE_BYTES: usize = 65_536; -pub const RESPONSE_MUX_CREDIT_UPDATE_INTERVAL: Duration = Duration::from_millis(1); -pub const RESPONSE_MUX_SCHEDULER_QUANTUM: usize = 8; #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub struct ResponseMuxConfig { @@ -47,7 +37,6 @@ pub struct ResponseMuxConfig { pub batch_max_bytes: usize, pub batch_max_frames: usize, pub stream_window_bytes: usize, - pub connection_window_bytes: usize, } impl ResponseMuxConfig { @@ -108,11 +97,6 @@ impl ResponseMuxConfig { env::DYN_TCP_RESPONSE_STREAM_WINDOW_BYTES, RESPONSE_MUX_DEFAULT_STREAM_WINDOW_BYTES, )?; - let connection_window_bytes = parse( - &mut read, - env::DYN_TCP_RESPONSE_CONNECTION_WINDOW_BYTES, - RESPONSE_MUX_DEFAULT_CONNECTION_WINDOW_BYTES, - )?; for (name, value) in [ (env::DYN_TCP_RESPONSE_BATCH_MAX_BYTES, batch_max_bytes), (env::DYN_TCP_RESPONSE_BATCH_MAX_FRAMES, batch_max_frames), @@ -120,25 +104,15 @@ impl ResponseMuxConfig { env::DYN_TCP_RESPONSE_STREAM_WINDOW_BYTES, stream_window_bytes, ), - ( - env::DYN_TCP_RESPONSE_CONNECTION_WINDOW_BYTES, - connection_window_bytes, - ), ] { if value == 0 { anyhow::bail!("{name} must be greater than zero"); } } - for (name, value) in [ - ( - env::DYN_TCP_RESPONSE_STREAM_WINDOW_BYTES, - stream_window_bytes, - ), - ( - env::DYN_TCP_RESPONSE_CONNECTION_WINDOW_BYTES, - connection_window_bytes, - ), - ] { + for (name, value) in [( + env::DYN_TCP_RESPONSE_STREAM_WINDOW_BYTES, + stream_window_bytes, + )] { if value > u32::MAX as usize { anyhow::bail!("{name} must fit in an unsigned 32-bit credit update"); } @@ -149,7 +123,6 @@ impl ResponseMuxConfig { batch_max_bytes, batch_max_frames, stream_window_bytes, - connection_window_bytes, }) } @@ -169,25 +142,8 @@ pub fn initialize_response_mux_config() -> anyhow::Result { Ok(*RESPONSE_MUX_CONFIG.get().expect("response mux config set")) } -pub fn response_packet_metrics_enabled() -> bool { - RESPONSE_MUX_CONFIG - .get() - .is_some_and(|config| config.packet_metrics) -} - -pub const MUX_HEADER_LEN: usize = 24; - -/// First header-only frame on a newly accepted TCP stream. -#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)] -#[serde(tag = "kind", rename_all = "snake_case")] -pub enum ConnectionHandshake { - /// Persistent connection carrying many downstream -> upstream responses. - ResponseMux { - version: u8, - frontend_server_id: Uuid, - connection_id: Uuid, - }, -} +pub const MUX_HEADER_LEN: usize = 21; +const CONNECTION_HELLO_LEN: usize = 17; #[derive(Debug, Clone, Copy, PartialEq, Eq)] #[repr(u8)] @@ -199,7 +155,8 @@ pub enum MuxFrameKind { Kill = 5, WindowUpdate = 6, Reset = 7, - ConnectionAck = 8, + ConnectionHello = 8, + ConnectionReady = 9, } impl TryFrom for MuxFrameKind { @@ -214,7 +171,8 @@ impl TryFrom for MuxFrameKind { 5 => Ok(Self::Kill), 6 => Ok(Self::WindowUpdate), 7 => Ok(Self::Reset), - 8 => Ok(Self::ConnectionAck), + 8 => Ok(Self::ConnectionHello), + 9 => Ok(Self::ConnectionReady), _ => Err(io::Error::new( io::ErrorKind::InvalidData, format!("unknown response mux frame kind {value}"), @@ -223,21 +181,6 @@ impl TryFrom for MuxFrameKind { } } -impl MuxFrameKind { - pub const fn metric_label(self) -> &'static str { - match self { - Self::Prologue => "prologue", - Self::Data => "data", - Self::End => "end", - Self::Stop => "stop", - Self::Kill => "kill", - Self::WindowUpdate => "window_update", - Self::Reset => "reset", - Self::ConnectionAck => "connection_ack", - } - } -} - #[derive(Debug, Clone, PartialEq, Eq)] pub struct MuxFrame { pub kind: MuxFrameKind, @@ -264,27 +207,32 @@ impl MuxFrame { Self::new(MuxFrameKind::WindowUpdate, stream_id, payload.freeze()) } - pub fn connection_ack(decoded_bytes: u64) -> Self { - Self::new( - MuxFrameKind::ConnectionAck, - Uuid::nil(), - decoded_bytes.to_be_bytes().to_vec().into(), - ) + pub fn connection_hello(version: u8, frontend_server_id: Uuid) -> Self { + let mut payload = BytesMut::with_capacity(CONNECTION_HELLO_LEN); + payload.put_u8(version); + payload.extend_from_slice(frontend_server_id.as_bytes()); + Self::new(MuxFrameKind::ConnectionHello, Uuid::nil(), payload.freeze()) } - pub fn connection_ack_offset(&self) -> io::Result { - if self.kind != MuxFrameKind::ConnectionAck || self.payload.len() != 8 { + pub fn connection_ready() -> Self { + Self::empty(MuxFrameKind::ConnectionReady, Uuid::nil()) + } + + pub fn connection_identity(&self) -> io::Result<(u8, Uuid)> { + if self.kind != MuxFrameKind::ConnectionHello || self.payload.len() != CONNECTION_HELLO_LEN + { return Err(io::Error::new( io::ErrorKind::InvalidData, - "response mux connection ACK must contain eight bytes", + "response mux connection hello must contain a version and frontend UUID", )); } - Ok(u64::from_be_bytes( - self.payload - .as_ref() - .try_into() - .expect("validated connection ACK length"), - )) + let frontend_server_id = Uuid::from_slice(&self.payload[1..]).map_err(|err| { + io::Error::new( + io::ErrorKind::InvalidData, + format!("invalid response mux frontend UUID: {err}"), + ) + })?; + Ok((self.payload[0], frontend_server_id)) } pub fn encoded_len(&self) -> usize { @@ -301,15 +249,6 @@ impl MuxFrame { Ok(u32::from_be_bytes(self.payload[..4].try_into().unwrap())) } - pub fn into_two_part(self) -> TwoPartMessage { - let mut header = BytesMut::with_capacity(20); - header.put_u8(self.kind as u8); - header.put_u8(0); // flags, reserved for future protocol use - header.put_u16(0); - header.extend_from_slice(self.stream_id.as_bytes()); - TwoPartMessage::new(header.freeze(), self.payload) - } - /// Split the wire representation into a small header and the original /// payload allocation so connection writers can coalesce headers while /// retaining large payloads as `Bytes` for vectored I/O. @@ -319,54 +258,11 @@ impl MuxFrame { Ok((header.freeze(), self.payload.clone())) } - pub fn try_from_two_part(message: TwoPartMessage) -> io::Result { - let (header, payload) = match message.into_message_type() { - TwoPartMessageType::HeaderOnly(header) => (header, Bytes::new()), - TwoPartMessageType::HeaderAndData(header, data) => (header, data), - TwoPartMessageType::DataOnly(_) | TwoPartMessageType::Empty => { - return Err(io::Error::new( - io::ErrorKind::InvalidData, - "response mux frame is missing its fixed header", - )); - } - }; - - if header.len() != 20 { - return Err(io::Error::new( - io::ErrorKind::InvalidData, - format!( - "invalid response mux header length {}, expected 20", - header.len() - ), - )); - } - if header[1..4] != [0, 0, 0] { - return Err(io::Error::new( - io::ErrorKind::InvalidData, - "response mux frame has unsupported flags", - )); - } - - let kind = MuxFrameKind::try_from(header[0])?; - let stream_id = Uuid::from_slice(&header[4..20]).map_err(|err| { - io::Error::new( - io::ErrorKind::InvalidData, - format!("invalid response mux stream UUID: {err}"), - ) - })?; - let is_connection_frame = kind == MuxFrameKind::ConnectionAck; - if stream_id.is_nil() != is_connection_frame { - return Err(io::Error::new( - io::ErrorKind::InvalidData, - "only connection-level frames must use the nil stream UUID", - )); - } - - Self::validate(Self::new(kind, stream_id, payload)) - } - fn validate(frame: Self) -> io::Result { - let is_connection_frame = frame.kind == MuxFrameKind::ConnectionAck; + let is_connection_frame = matches!( + frame.kind, + MuxFrameKind::ConnectionHello | MuxFrameKind::ConnectionReady + ); if frame.stream_id.is_nil() != is_connection_frame { return Err(io::Error::new( io::ErrorKind::InvalidData, @@ -374,7 +270,9 @@ impl MuxFrame { )); } match frame.kind { - MuxFrameKind::Stop | MuxFrameKind::Kill if !frame.payload.is_empty() => { + MuxFrameKind::End | MuxFrameKind::Stop | MuxFrameKind::Kill | MuxFrameKind::Reset + if !frame.payload.is_empty() => + { Err(io::Error::new( io::ErrorKind::InvalidData, "control frame must not contain a payload", @@ -384,18 +282,21 @@ impl MuxFrame { frame.window_credits()?; Ok(frame) } - MuxFrameKind::ConnectionAck => { - frame.connection_ack_offset()?; + MuxFrameKind::ConnectionHello => { + frame.connection_identity()?; Ok(frame) } + MuxFrameKind::ConnectionReady if !frame.payload.is_empty() => Err(io::Error::new( + io::ErrorKind::InvalidData, + "response mux connection ready must be empty", + )), _ => Ok(frame), } } } -/// Compact response-mux framing used after the versioned connection -/// handshake. Each frame is `payload_len:u32`, kind, flags, reserved, UUID, -/// then payload. The fixed header is 24 bytes. +/// Compact response-mux framing. Each frame is `payload_len:u32`, kind, UUID, +/// then payload. #[derive(Clone, Debug)] pub struct MuxCodec { max_message_size: usize, @@ -434,8 +335,6 @@ impl MuxCodec { dst.reserve(MUX_HEADER_LEN); dst.put_u32(payload_len); dst.put_u8(frame.kind as u8); - dst.put_u8(0); - dst.put_u16(0); dst.extend_from_slice(frame.stream_id.as_bytes()); Ok(()) } @@ -479,14 +378,8 @@ impl Decoder for MuxCodec { src.reserve(encoded_len - src.len()); return Ok(None); } - if src[5] != 0 || src[6..8] != [0, 0] { - return Err(io::Error::new( - io::ErrorKind::InvalidData, - "response mux frame has unsupported flags or reserved bits", - )); - } let kind = MuxFrameKind::try_from(src[4])?; - let stream_id = Uuid::from_slice(&src[8..24]).map_err(|err| { + let stream_id = Uuid::from_slice(&src[5..21]).map_err(|err| { io::Error::new( io::ErrorKind::InvalidData, format!("invalid response mux stream UUID: {err}"), @@ -501,17 +394,8 @@ impl Decoder for MuxCodec { #[cfg(test)] mod tests { use super::*; - use crate::pipeline::network::codec::TwoPartCodec; use tokio_util::codec::{Decoder, Encoder}; - fn encode(frame: MuxFrame) -> BytesMut { - let mut bytes = BytesMut::new(); - TwoPartCodec::default() - .encode(frame.into_two_part(), &mut bytes) - .unwrap(); - bytes - } - fn encode_compact(frame: MuxFrame) -> BytesMut { let mut bytes = BytesMut::new(); MuxCodec::default().encode(frame, &mut bytes).unwrap(); @@ -526,10 +410,11 @@ mod tests { stream_id, Bytes::from_static(b"payload"), ); - let codec = TwoPartCodec::default(); - let encoded = codec.encode_message(frame.clone().into_two_part()).unwrap(); - let decoded = codec.decode_message(encoded).unwrap(); - assert_eq!(MuxFrame::try_from_two_part(decoded).unwrap(), frame); + let mut encoded = encode_compact(frame.clone()); + assert_eq!( + MuxCodec::default().decode(&mut encoded).unwrap(), + Some(frame) + ); } #[test] @@ -539,26 +424,29 @@ mod tests { Uuid::new_v4(), Bytes::from_static(&[1, 2]), ); - assert!(MuxFrame::try_from_two_part(frame.into_two_part()).is_err()); + let mut encoded = BytesMut::new(); + MuxCodec::default().encode(frame, &mut encoded).unwrap(); + assert!(MuxCodec::default().decode(&mut encoded).is_err()); } #[test] - fn connection_ack_round_trips_with_cumulative_byte_offset() { - let frame = MuxFrame::connection_ack(987_654); - let decoded = MuxFrame::try_from_two_part(frame.clone().into_two_part()).unwrap(); + fn connection_hello_round_trips_identity() { + let frontend_server_id = Uuid::new_v4(); + let frame = MuxFrame::connection_hello(RESPONSE_MUX_VERSION, frontend_server_id); + let mut encoded = encode_compact(frame.clone()); + let decoded = MuxCodec::default().decode(&mut encoded).unwrap().unwrap(); assert_eq!(decoded, frame); - assert_eq!(decoded.connection_ack_offset().unwrap(), 987_654); - assert_eq!(decoded.encoded_len(), 32); - } + assert_eq!( + decoded.connection_identity().unwrap(), + (RESPONSE_MUX_VERSION, frontend_server_id) + ); - #[test] - fn connection_ack_rejects_stream_uuid() { - let frame = MuxFrame::new( - MuxFrameKind::ConnectionAck, + let mut invalid = encode_compact(MuxFrame::new( + MuxFrameKind::ConnectionReady, Uuid::new_v4(), - 128_u64.to_be_bytes().to_vec().into(), - ); - assert!(MuxFrame::try_from_two_part(frame.into_two_part()).is_err()); + Bytes::new(), + )); + assert!(MuxCodec::default().decode(&mut invalid).is_err()); } #[test] @@ -624,26 +512,6 @@ mod tests { ); } - #[test] - fn malformed_outer_lengths_are_rejected() { - let mut input = BytesMut::new(); - input.put_u64(u64::MAX); - input.put_u64(1); - input.put_u64(0); - assert!(TwoPartCodec::default().decode(&mut input).is_err()); - } - - #[test] - fn compact_codec_rejects_flags_and_connection_uuid_mismatch() { - let mut flags = encode_compact(MuxFrame::empty(MuxFrameKind::End, Uuid::new_v4())); - flags[5] = 1; - assert!(MuxCodec::default().decode(&mut flags).is_err()); - - let mut invalid = encode_compact(MuxFrame::connection_ack(64)); - invalid[8..24].copy_from_slice(Uuid::new_v4().as_bytes()); - assert!(MuxCodec::default().decode(&mut invalid).is_err()); - } - fn config(values: &[(&str, &str)]) -> anyhow::Result { let values = values .iter() @@ -660,7 +528,6 @@ mod tests { assert_eq!(config.batch_max_bytes, 65_536); assert_eq!(config.batch_max_frames, 64); assert_eq!(config.stream_window_bytes, 262_144); - assert_eq!(config.connection_window_bytes, 262_144); } #[test] diff --git a/lib/runtime/src/pipeline/network/tcp/mux/client.rs b/lib/runtime/src/pipeline/network/tcp/mux/client.rs index 75ce24695066..416d55179a99 100644 --- a/lib/runtime/src/pipeline/network/tcp/mux/client.rs +++ b/lib/runtime/src/pipeline/network/tcp/mux/client.rs @@ -4,10 +4,10 @@ //! Worker-side persistent multiplexed TCP response connection pool. use std::{ - collections::{HashMap, VecDeque}, + cmp::Reverse, sync::{ Arc, Weak, - atomic::{AtomicBool, AtomicU64, AtomicUsize, Ordering}, + atomic::{AtomicBool, AtomicUsize, Ordering}, }, time::{Duration, Instant}, }; @@ -15,7 +15,7 @@ use std::{ use anyhow::{Context, Result, anyhow}; use dashmap::{DashMap, mapref::entry::Entry}; use futures::{SinkExt, StreamExt}; -use parking_lot::{Mutex, RwLock}; +use parking_lot::RwLock; use tokio::{ net::TcpStream, sync::{OwnedSemaphorePermit, Semaphore, mpsc, oneshot}, @@ -27,90 +27,30 @@ use tokio_util::{ use uuid::Uuid; use super::{ - ConnectionHandshake, MuxCodec, MuxFrame, MuxFrameKind, RESPONSE_MUX_CONNECT_TIMEOUT_SECS, - RESPONSE_MUX_IDLE_TTL_SECS, RESPONSE_MUX_POOL_SIZE, RESPONSE_MUX_STREAM_WRITER_QUEUE, + MuxCodec, MuxFrame, MuxFrameKind, RESPONSE_MUX_CONNECT_TIMEOUT_SECS, + RESPONSE_MUX_CONNECTION_QUEUE_BYTES, RESPONSE_MUX_POOL_SIZE, RESPONSE_MUX_STREAM_WRITER_QUEUE, RESPONSE_MUX_VERSION, RESPONSE_MUX_WRITER_QUEUE, ResponseMuxConfig, }; use crate::{ engine::AsyncEngineContext, metrics::response_mux, pipeline::network::{ - ConnectionInfo, MultiplexedStreamSender, ResponseStreamPrologue, StreamSender, - codec::{TwoPartCodec, TwoPartMessage}, - egress::tcp_client::TcpWriteBuffer, + ConnectionInfo, StreamSender, egress::tcp_client::TcpWriteBuffer, tcp::ResponseMuxConnectionInfo, }, }; -struct WriterCommand { - frame: MuxFrame, - written: Option>>, - _writer_permit: Option, - _queued_byte_permit: Option, - priority_enqueued_at: Option, - enqueued_at: Option, -} - -impl WriterCommand { - fn new(frame: MuxFrame, written: Option>>) -> Self { - Self::new_with_metrics(frame, written, per_frame_metrics_enabled()) - } - - fn new_with_metrics( - frame: MuxFrame, - written: Option>>, - metrics_enabled: bool, - ) -> Self { - Self { - frame, - written, - _writer_permit: None, - _queued_byte_permit: None, - priority_enqueued_at: None, - enqueued_at: metrics_enabled.then(Instant::now), - } - } - - fn priority(frame: MuxFrame, written: Option>>) -> Self { - let mut command = Self::new(frame, written); - command.priority_enqueued_at = per_frame_metrics_enabled().then(Instant::now); - command - } - - fn with_writer_permit(mut self, permit: OwnedSemaphorePermit) -> Self { - self._writer_permit = Some(permit); - self - } - - fn with_queued_byte_permit(mut self, permit: OwnedSemaphorePermit) -> Self { - self._queued_byte_permit = Some(permit); - self - } - - fn fail(mut self, reason: &str) { - if let Some(written) = self.written.take() { - let _ = written.send(Err(reason.to_string())); - } - } -} - -#[inline] -fn per_frame_metrics_enabled() -> bool { - super::response_packet_metrics_enabled() -} - #[derive(Clone, Copy)] struct PoolConfig { pool_size: usize, writer_queue: usize, stream_writer_queue: usize, initial_window: usize, - connection_window: usize, + queued_bytes: usize, batch_interval: Duration, batch_max_bytes: usize, batch_max_frames: usize, packet_metrics: bool, - idle_ttl: Duration, connect_timeout: Duration, } @@ -121,12 +61,11 @@ impl PoolConfig { writer_queue: RESPONSE_MUX_WRITER_QUEUE, stream_writer_queue: RESPONSE_MUX_STREAM_WRITER_QUEUE, initial_window: config.stream_window_bytes, - connection_window: config.connection_window_bytes, + queued_bytes: RESPONSE_MUX_CONNECTION_QUEUE_BYTES, batch_interval: config.batch_interval, batch_max_bytes: config.batch_max_bytes, batch_max_frames: config.batch_max_frames, packet_metrics: config.packet_metrics, - idle_ttl: Duration::from_secs(RESPONSE_MUX_IDLE_TTL_SECS), connect_timeout: Duration::from_secs(RESPONSE_MUX_CONNECT_TIMEOUT_SECS), } } @@ -143,34 +82,6 @@ struct WorkerStreamState { close_token: CancellationToken, } -type ScheduledStream = (Uuid, Arc); -type BlockedData = (WriterCommand, Option, Instant); - -struct StreamIngress { - stream_id: Uuid, - state: Arc, - command: WriterCommand, -} - -struct WriterStreamQueue { - state: Arc, - pending: VecDeque, - scheduled: bool, -} - -#[derive(Debug, Clone, Copy, PartialEq, Eq)] -enum IngressPoll { - Ingested, - Empty, - Disconnected, -} - -#[derive(Default)] -struct WriterScheduler { - streams: HashMap, - ready: VecDeque, -} - impl WorkerStreamState { fn record_cancellation(&self) { if !self.cancellation_recorded.swap(true, Ordering::AcqRel) @@ -180,199 +91,75 @@ impl WorkerStreamState { } } - fn replenish_credits(&self, credits: usize) -> usize { + fn replenish_credits(&self, credits: usize) { if self.closed.load(Ordering::Acquire) || self.credits.is_closed() { - return 0; + return; } let available = self.credits.available_permits(); let replenished = credits.min(self.max_credits.saturating_sub(available)); if replenished > 0 { self.credits.add_permits(replenished); } - replenished } } -impl WriterScheduler { - fn account_removed(connection: &MuxConnection, command: &WriterCommand) { - connection.queued_frames.fetch_sub(1, Ordering::AcqRel); - let encoded_len = command.frame.encoded_len(); - connection - .queued_bytes - .fetch_sub(encoded_len, Ordering::AcqRel); - if per_frame_metrics_enabled() { - response_mux::QUEUED_BYTES - .with_label_values(&["worker"]) - .sub(encoded_len as i64); - } - } - - fn fail_commands( - connection: &MuxConnection, - commands: impl IntoIterator, - reason: &str, - ) { - for command in commands { - Self::account_removed(connection, &command); - command.fail(reason); - } - } - - fn ingest(&mut self, connection: &MuxConnection, ingress: StreamIngress) { - let StreamIngress { - stream_id, - state, - command, - } = ingress; - if state.closed.load(Ordering::Acquire) { - Self::account_removed(connection, &command); - command.fail("response mux stream is closed"); - return; - } +struct UrgentCommand { + frame: MuxFrame, +} - let queue = self - .streams - .entry(stream_id) - .or_insert_with(|| WriterStreamQueue { - state, - pending: VecDeque::new(), - scheduled: false, - }); - queue.pending.push_back(command); - if per_frame_metrics_enabled() { - response_mux::STREAM_WRITER_QUEUE_OCCUPANCY.observe(queue.pending.len() as f64); - } - if !queue.scheduled { - queue.scheduled = true; - self.ready.push_back(stream_id); - if per_frame_metrics_enabled() { - response_mux::READY_STREAMS - .with_label_values(&["worker"]) - .inc(); - } - } - } +struct OrderedCommand { + frame: MuxFrame, + state: Arc, + written: Option>>, + _writer_permit: OwnedSemaphorePermit, + _queued_byte_permit: OwnedSemaphorePermit, +} - fn ingest_one( - &mut self, - connection: &MuxConnection, - stream_rx: &mut mpsc::Receiver, - ) -> IngressPoll { - match stream_rx.try_recv() { - Ok(ingress) => { - self.ingest(connection, ingress); - IngressPoll::Ingested - } - Err(mpsc::error::TryRecvError::Empty) => IngressPoll::Empty, - Err(mpsc::error::TryRecvError::Disconnected) => IngressPoll::Disconnected, +impl OrderedCommand { + fn fail(mut self, reason: &str) { + if let Some(written) = self.written.take() { + let _ = written.send(Err(reason.to_string())); } } +} - fn pop_ready( - &mut self, - connection: &MuxConnection, - ) -> Option<(WriterCommand, ScheduledStream)> { - while let Some(stream_id) = self.ready.pop_front() { - let Some(queue) = self.streams.get_mut(&stream_id) else { - continue; - }; - if queue.state.closed.load(Ordering::Acquire) { - let pending = queue.pending.drain(..).collect::>(); - if queue.scheduled { - queue.scheduled = false; - if per_frame_metrics_enabled() { - response_mux::READY_STREAMS - .with_label_values(&["worker"]) - .dec(); - } - } - Self::fail_commands(connection, pending, "response mux stream is closed"); - continue; - } - if let Some(command) = queue.pending.pop_front() { - Self::account_removed(connection, &command); - return Some((command, (stream_id, queue.state.clone()))); - } - if queue.scheduled { - queue.scheduled = false; - if per_frame_metrics_enabled() { - response_mux::READY_STREAMS - .with_label_values(&["worker"]) - .dec(); - } - } - } - None - } +enum WriteCommand { + Urgent(UrgentCommand), + Ordered(OrderedCommand), +} - fn pop_same_stream( - &mut self, - connection: &MuxConnection, - stream_id: Uuid, - ) -> Option { - let queue = self.streams.get_mut(&stream_id)?; - if queue.state.closed.load(Ordering::Acquire) { - let pending = queue.pending.drain(..).collect::>(); - Self::fail_commands(connection, pending, "response mux stream is closed"); - return None; +impl WriteCommand { + fn frame(&self) -> &MuxFrame { + match self { + Self::Urgent(command) => &command.frame, + Self::Ordered(command) => &command.frame, } - let command = queue.pending.pop_front()?; - Self::account_removed(connection, &command); - Some(command) } - fn reschedule(&mut self, connection: &MuxConnection, stream_id: Uuid) { - if per_frame_metrics_enabled() { - response_mux::ROUND_ROBIN_TURNS_TOTAL.inc(); - } - let Some(queue) = self.streams.get_mut(&stream_id) else { - return; - }; - if queue.state.closed.load(Ordering::Acquire) { - let pending = queue.pending.drain(..).collect::>(); - Self::fail_commands(connection, pending, "response mux stream is closed"); - } - if !queue.state.closed.load(Ordering::Acquire) && !queue.pending.is_empty() { - self.ready.push_back(stream_id); - } else if queue.scheduled { - queue.scheduled = false; - if per_frame_metrics_enabled() { - response_mux::READY_STREAMS - .with_label_values(&["worker"]) - .dec(); - } + fn fail(self, reason: &str) { + if let Self::Ordered(command) = self { + command.fail(reason); } } - fn fail_all(&mut self, connection: &MuxConnection, reason: &str) { - for (_, mut queue) in self.streams.drain() { - if queue.scheduled && per_frame_metrics_enabled() { - response_mux::READY_STREAMS - .with_label_values(&["worker"]) - .dec(); - } - Self::fail_commands(connection, queue.pending.drain(..), reason); + fn complete(&mut self, result: &Result<(), String>) { + if let Self::Ordered(command) = self + && let Some(written) = command.written.take() + { + let _ = written.send(result.clone()); } - self.ready.clear(); } } struct MuxConnection { - id: u64, cancel: CancellationToken, - priority_tx: mpsc::Sender, - stream_tx: mpsc::Sender, + urgent_tx: mpsc::Sender, + ordered_tx: mpsc::Sender, streams: DashMap>, healthy: AtomicBool, active_streams: AtomicUsize, - queued_frames: AtomicUsize, - queued_bytes: AtomicUsize, max_queued_bytes: usize, queued_byte_slots: Arc, - connection_credits: Arc, - max_connection_credits: usize, - sent_data_bytes: AtomicU64, - acknowledged_data_bytes: AtomicU64, batch_interval: Duration, batch_max_bytes: usize, batch_max_frames: usize, @@ -380,10 +167,8 @@ struct MuxConnection { impl MuxConnection { async fn connect( - id: u64, address: &str, frontend_server_id: Uuid, - version: u8, cancel: CancellationToken, config: PoolConfig, ) -> Result> { @@ -397,50 +182,36 @@ impl MuxConnection { .flatten(); let (read_half, write_half) = stream.into_split(); - let mux_codec = - || TwoPartCodec::new(Some(crate::pipeline::network::get_tcp_max_message_size())); - let mut handshake_reader = FramedRead::new(read_half, mux_codec()); - let mut handshake_writer = FramedWrite::new(write_half, mux_codec()); - let handshake = ConnectionHandshake::ResponseMux { - version, - frontend_server_id, - connection_id: Uuid::new_v4(), - }; - let header = serde_json::to_vec(&handshake)?; + let mut reader = FramedRead::new(read_half, MuxCodec::default()); + let mut handshake_writer = FramedWrite::new(write_half, MuxCodec::default()); handshake_writer - .send(TwoPartMessage::from_header(header.into())) + .send(MuxFrame::connection_hello( + RESPONSE_MUX_VERSION, + frontend_server_id, + )) .await - .context("failed to send response mux handshake")?; - let ack = tokio::time::timeout(config.connect_timeout, handshake_reader.next()) + .context("failed to send response mux connection hello")?; + let ready = tokio::time::timeout(config.connect_timeout, reader.next()) .await - .map_err(|_| anyhow!("response mux handshake ack timeout from {address}"))? - .ok_or_else(|| anyhow!("frontend closed before response mux handshake ack"))??; - let ack = MuxFrame::try_from_two_part(ack)?; - if ack.kind != MuxFrameKind::ConnectionAck || ack.connection_ack_offset()? != 0 { - anyhow::bail!("frontend returned invalid response mux connection ack"); + .map_err(|_| anyhow!("response mux connection-ready timeout from {address}"))? + .ok_or_else(|| anyhow!("frontend closed before response mux connection ready"))??; + if ready != MuxFrame::connection_ready() { + anyhow::bail!("frontend returned invalid response mux connection ready"); } - let mux_reader = handshake_reader.map_decoder(|_| MuxCodec::default()); let write_half = handshake_writer.into_inner(); - let (priority_tx, priority_rx) = mpsc::channel(config.writer_queue); - let (stream_tx, stream_rx) = mpsc::channel(config.writer_queue); + let (urgent_tx, urgent_rx) = mpsc::channel(config.writer_queue); + let (ordered_tx, ordered_rx) = mpsc::channel(config.writer_queue); let cancel = cancel.child_token(); let connection = Arc::new(Self { - id, cancel: cancel.clone(), - priority_tx, - stream_tx, + urgent_tx, + ordered_tx, streams: DashMap::new(), healthy: AtomicBool::new(true), active_streams: AtomicUsize::new(0), - queued_frames: AtomicUsize::new(0), - queued_bytes: AtomicUsize::new(0), - max_queued_bytes: config.connection_window, - queued_byte_slots: Arc::new(Semaphore::new(config.connection_window)), - connection_credits: Arc::new(Semaphore::new(config.connection_window)), - max_connection_credits: config.connection_window, - sent_data_bytes: AtomicU64::new(0), - acknowledged_data_bytes: AtomicU64::new(0), + max_queued_bytes: config.queued_bytes, + queued_byte_slots: Arc::new(Semaphore::new(config.queued_bytes)), batch_interval: config.batch_interval, batch_max_bytes: config.batch_max_bytes, batch_max_frames: config.batch_max_frames, @@ -455,13 +226,13 @@ impl MuxConnection { tokio::spawn(Self::writer_task( Arc::downgrade(&connection), write_half, - priority_rx, - stream_rx, + urgent_rx, + ordered_rx, cancel.clone(), )); tokio::spawn(Self::reader_task( Arc::downgrade(&connection), - mux_reader, + reader, cancel, packet_baseline, )); @@ -472,52 +243,12 @@ impl MuxConnection { self.healthy.load(Ordering::Acquire) } - fn replenish_connection_credits(&self, credits: usize) -> usize { - if !self.is_healthy() || self.connection_credits.is_closed() { - return 0; - } - let available = self.connection_credits.available_permits(); - let replenished = credits.min(self.max_connection_credits.saturating_sub(available)); - if replenished > 0 { - self.connection_credits.add_permits(replenished); - } - replenished - } - - fn acknowledge_connection_credits(&self, acknowledged_bytes: u64) -> Result<()> { - let previous = self.acknowledged_data_bytes.load(Ordering::Acquire); - if acknowledged_bytes < previous { - anyhow::bail!( - "response mux connection ACK moved backwards from {previous} to {acknowledged_bytes}" - ); - } - if acknowledged_bytes == previous { - return Ok(()); - } - let sent = self.sent_data_bytes.load(Ordering::Acquire); - if acknowledged_bytes > sent { - anyhow::bail!( - "response mux connection ACK {acknowledged_bytes} exceeds sent offset {sent}" - ); - } - self.acknowledged_data_bytes - .store(acknowledged_bytes, Ordering::Release); - let delta = acknowledged_bytes.saturating_sub(previous) as usize; - self.replenish_connection_credits(delta.min(self.max_connection_credits)); - Ok(()) - } - fn fail(&self, reason: &str) { if !self.healthy.swap(false, Ordering::AcqRel) { return; } - tracing::warn!( - connection_id = self.id, - reason, - "response mux connection failed" - ); + tracing::warn!(reason, "response mux connection failed"); self.cancel.cancel(); - self.connection_credits.close(); self.queued_byte_slots.close(); response_mux::CONNECTIONS_TOTAL .with_label_values(&["worker", "failed"]) @@ -557,511 +288,200 @@ impl MuxConnection { true } - async fn send_priority_command(&self, command: WriterCommand) -> Result<()> { + async fn send_urgent(&self, frame: MuxFrame) -> Result<()> { if !self.is_healthy() { - command.fail("response mux connection is unhealthy"); anyhow::bail!("response mux connection is unhealthy"); } - self.queued_frames.fetch_add(1, Ordering::AcqRel); - match self.priority_tx.try_send(command) { - Ok(()) => Ok(()), - Err(mpsc::error::TrySendError::Full(command)) => { - if let Err(err) = self.priority_tx.send(command).await { - self.queued_frames.fetch_sub(1, Ordering::AcqRel); - err.0.fail("response mux priority writer stopped"); - anyhow::bail!("response mux priority writer stopped"); - } - Ok(()) - } - Err(mpsc::error::TrySendError::Closed(command)) => { - self.queued_frames.fetch_sub(1, Ordering::AcqRel); - command.fail("response mux priority writer stopped"); - anyhow::bail!("response mux priority writer stopped") - } - } + self.urgent_tx + .send(UrgentCommand { frame }) + .await + .map_err(|_| anyhow!("response mux urgent writer stopped")) } - fn try_send_priority_command(&self, command: WriterCommand) { + fn try_send_urgent(&self, frame: MuxFrame) { if !self.is_healthy() { - command.fail("response mux connection is unhealthy"); return; } - self.queued_frames.fetch_add(1, Ordering::AcqRel); - if let Err(err) = self.priority_tx.try_send(command) { - self.queued_frames.fetch_sub(1, Ordering::AcqRel); - let reason = match err { - mpsc::error::TrySendError::Full(command) => { - command.fail("response mux priority writer queue is full"); - "response mux priority writer queue is full" - } - mpsc::error::TrySendError::Closed(command) => { - command.fail("response mux priority writer is closed"); - "response mux priority writer is closed" - } - }; - self.fail(reason); + if self.urgent_tx.try_send(UrgentCommand { frame }).is_err() { + self.fail("response mux urgent writer queue is unavailable"); } } - async fn enqueue_stream_command( - &self, - stream_id: Uuid, - state: Arc, - command: WriterCommand, - ) -> Result<()> { - if !self.is_healthy() || state.closed.load(Ordering::Acquire) { + async fn send_ordered(&self, command: OrderedCommand) -> Result<()> { + if !self.is_healthy() || command.state.closed.load(Ordering::Acquire) { command.fail("response mux stream is closed"); anyhow::bail!("response mux stream is closed"); } + self.ordered_tx.send(command).await.map_err(|err| { + err.0.fail("response mux ordered writer stopped"); + anyhow!("response mux ordered writer stopped") + }) + } - let encoded_len = command.frame.encoded_len(); - self.queued_bytes.fetch_add(encoded_len, Ordering::AcqRel); - if per_frame_metrics_enabled() { - response_mux::QUEUED_BYTES - .with_label_values(&["worker"]) - .add(encoded_len as i64); - } - - self.queued_frames.fetch_add(1, Ordering::AcqRel); - match self - .stream_tx - .send(StreamIngress { - stream_id, - state, - command, - }) - .await - { - Ok(()) => Ok(()), - Err(err) => { - self.queued_frames.fetch_sub(1, Ordering::AcqRel); - self.queued_bytes.fetch_sub(encoded_len, Ordering::AcqRel); - if per_frame_metrics_enabled() { - response_mux::QUEUED_BYTES - .with_label_values(&["worker"]) - .sub(encoded_len as i64); - } - err.0.command.fail("response mux fair writer stopped"); - anyhow::bail!("response mux fair writer stopped") - } - } + fn command_is_writable(command: &OrderedCommand) -> bool { + !command.state.closed.load(Ordering::Acquire) } async fn writer_task( weak: Weak, mut write_half: tokio::net::tcp::OwnedWriteHalf, - mut priority_rx: mpsc::Receiver, - mut stream_rx: mpsc::Receiver, + mut urgent_rx: mpsc::Receiver, + mut ordered_rx: mpsc::Receiver, cancel: CancellationToken, ) { let mut write_buf = TcpWriteBuffer::new(); - let queue_depth = response_mux::WRITER_QUEUE_DEPTH - .with_label_values(&["worker"]) - .clone(); let frames_per_write = response_mux::FRAMES_PER_WRITE .with_label_values(&["worker"]) .clone(); - let frame_counters = response_mux::FrameCounters::for_direction("worker_to_frontend"); - let metrics_enabled = per_frame_metrics_enabled(); - let mut reported_queue_depth = 0_i64; - let mut blocked_data: Option = None; - let mut scheduler = WriterScheduler::default(); - let mut stream_input_open = true; + let mut pending_urgent = None; + let mut pending_ordered = None; + let mut urgent_open = true; + let mut ordered_open = true; let result: Result<()> = async { - 'writer: loop { + loop { let connection = weak .upgrade() .ok_or_else(|| anyhow!("response mux connection dropped"))?; - if metrics_enabled { - let current_queue_depth = - connection.queued_frames.load(Ordering::Acquire) as i64; - queue_depth.add(current_queue_depth - reported_queue_depth); - reported_queue_depth = current_queue_depth; - } - - let (command, scheduled_stream, connection_permit) = if let Some(( - blocked_command, - blocked_stream, - blocked_since, - )) = blocked_data.take() - { - if blocked_command.frame.kind != MuxFrameKind::Data { - (blocked_command, blocked_stream, None) - } else { - enum BlockedNext { - Priority(WriterCommand), - Credit(OwnedSemaphorePermit), - StreamClosed, - } - let blocked_close = blocked_stream - .as_ref() - .expect("Data commands are always stream-scheduled") - .1 - .close_token - .clone(); - let credits = connection.connection_credits.clone(); - let next = tokio::select! { - biased; - _ = cancel.cancelled() => return Ok(()), - _ = blocked_close.cancelled() => BlockedNext::StreamClosed, - Some(command) = priority_rx.recv() => BlockedNext::Priority(command), - permit = credits.acquire_many_owned( - blocked_command - .frame - .encoded_len() - .min(connection.max_connection_credits) as u32 - ) => BlockedNext::Credit( - permit.map_err(|_| anyhow!( - "response mux connection closed while writer awaited credits" - ))? - ), - }; - match next { - BlockedNext::Priority(command) => { - connection.queued_frames.fetch_sub(1, Ordering::AcqRel); - blocked_data = - Some((blocked_command, blocked_stream, blocked_since)); - (command, None, None) - } - BlockedNext::Credit(permit) => { - if metrics_enabled { - response_mux::CONNECTION_FLOW_CONTROL_STALL_SECONDS - .observe(blocked_since.elapsed().as_secs_f64()); - } - (blocked_command, blocked_stream, Some(permit)) - } - BlockedNext::StreamClosed => { - if let Some((stream_id, _)) = blocked_stream { - scheduler.reschedule(&connection, stream_id); - } - blocked_command.fail( - "response mux stream closed while writer awaited credits", - ); - continue 'writer; - } - } + let first = loop { + if let Some(command) = pending_urgent.take() { + break WriteCommand::Urgent(command); } - } else { - let (command, scheduled_stream) = loop { - if let Ok(command) = priority_rx.try_recv() { - connection.queued_frames.fetch_sub(1, Ordering::AcqRel); - break (command, None); - } - if let Some((command, scheduled_stream)) = scheduler.pop_ready(&connection) - { - break (command, Some(scheduled_stream)); - } - match scheduler.ingest_one(&connection, &mut stream_rx) { - IngressPoll::Ingested => continue, - IngressPoll::Empty => {} - IngressPoll::Disconnected => { - stream_input_open = false; - } - } - - enum Next { - Priority(Option), - Ingress(Option), + if let Ok(command) = urgent_rx.try_recv() { + break WriteCommand::Urgent(command); + } + if let Some(command) = pending_ordered.take() { + if Self::command_is_writable(&command) { + break WriteCommand::Ordered(command); } - let next = tokio::select! { - biased; - _ = cancel.cancelled() => return Ok(()), - command = priority_rx.recv() => Next::Priority(command), - ingress = stream_rx.recv(), if stream_input_open => { - Next::Ingress(ingress) - } - else => return Ok(()), - }; - match next { - Next::Priority(Some(command)) => { - connection.queued_frames.fetch_sub(1, Ordering::AcqRel); - break (command, None); - } - Next::Priority(None) if !stream_input_open => return Ok(()), - Next::Priority(None) => {} - Next::Ingress(Some(ingress)) => { - scheduler.ingest(&connection, ingress); - } - Next::Ingress(None) => stream_input_open = false, + command.fail("response mux stream closed before write"); + continue; + } + if let Ok(command) = ordered_rx.try_recv() { + if Self::command_is_writable(&command) { + break WriteCommand::Ordered(command); } + command.fail("response mux stream closed before write"); + continue; + } + let next = tokio::select! { + biased; + _ = cancel.cancelled() => return Ok(()), + command = urgent_rx.recv(), if urgent_open => match command { + Some(command) => Some(WriteCommand::Urgent(command)), + None => { urgent_open = false; None } + }, + command = ordered_rx.recv(), if ordered_open => match command { + Some(command) => Some(WriteCommand::Ordered(command)), + None => { ordered_open = false; None } + }, + else => return Ok(()), }; - - if command.frame.kind == MuxFrameKind::Data { - let required = command - .frame - .encoded_len() - .min(connection.max_connection_credits) - as u32; - match connection - .connection_credits - .clone() - .try_acquire_many_owned(required) - { - Ok(permit) => (command, scheduled_stream, Some(permit)), - Err(tokio::sync::TryAcquireError::NoPermits) => { - debug_assert!(blocked_data.is_none()); - blocked_data = Some((command, scheduled_stream, Instant::now())); - continue 'writer; - } - Err(tokio::sync::TryAcquireError::Closed) => { - return Err(anyhow!( - "response mux connection closed while writer acquired credits" - )); - } + let Some(next) = next else { + continue; + }; + match next { + WriteCommand::Ordered(command) if !Self::command_is_writable(&command) => { + command.fail("response mux stream closed before write"); } - } else { - (command, scheduled_stream, None) + next => break next, } }; - let batching_started = Instant::now(); - let first_is_data = command.frame.kind == MuxFrameKind::Data; - let mut batch = vec![(command, connection_permit)]; - let mut batch_bytes = batch[0].0.frame.encoded_len(); - let mut force_flush = !first_is_data; - let mut held_for_next_turn = false; - - if let Some((stream_id, state)) = scheduled_stream { - for _ in 1..super::RESPONSE_MUX_SCHEDULER_QUANTUM { - if force_flush - || batch.len() >= connection.batch_max_frames + let first_is_data = first.frame().kind == MuxFrameKind::Data; + let mut batch = vec![first]; + let mut batch_bytes = batch[0].frame().encoded_len(); + if first_is_data { + let deadline = tokio::time::Instant::now() + connection.batch_interval; + loop { + if batch.len() >= connection.batch_max_frames || batch_bytes >= connection.batch_max_bytes { break; } - let Some(next) = scheduler.pop_same_stream(&connection, stream_id) else { - break; - }; - let encoded_len = next.frame.encoded_len(); - if batch_bytes.saturating_add(encoded_len) > connection.batch_max_bytes { - debug_assert!(blocked_data.is_none()); - blocked_data = - Some((next, Some((stream_id, state.clone())), Instant::now())); - held_for_next_turn = true; + if let Ok(command) = urgent_rx.try_recv() { + pending_urgent = Some(command); break; } - let permit = if next.frame.kind == MuxFrameKind::Data { - let required = - encoded_len.min(connection.max_connection_credits) as u32; - match connection - .connection_credits - .clone() - .try_acquire_many_owned(required) - { - Ok(permit) => Some(permit), - Err(tokio::sync::TryAcquireError::NoPermits) => { - debug_assert!(blocked_data.is_none()); - blocked_data = Some(( - next, - Some((stream_id, state.clone())), - Instant::now(), - )); - held_for_next_turn = true; - break; - } - Err(tokio::sync::TryAcquireError::Closed) => { - return Err(anyhow!( - "response mux connection credit window closed" - )); + + let next = if let Some(command) = pending_ordered.take() { + Some(command) + } else if let Ok(command) = ordered_rx.try_recv() { + Some(command) + } else if connection.batch_interval.is_zero() { + None + } else { + tokio::select! { + biased; + _ = cancel.cancelled() => return Ok(()), + command = urgent_rx.recv(), if urgent_open => { + match command { + Some(command) => pending_urgent = Some(command), + None => urgent_open = false, + } + None } + _ = tokio::time::sleep_until(deadline) => None, + command = ordered_rx.recv(), if ordered_open => match command { + Some(command) => Some(command), + None => { ordered_open = false; None } + }, } - } else { - None }; - force_flush = next.frame.kind != MuxFrameKind::Data; - batch_bytes = batch_bytes.saturating_add(encoded_len); - batch.push((next, permit)); - } - if !held_for_next_turn { - scheduler.reschedule(&connection, stream_id); - } - } - - let deadline = batching_started + connection.batch_interval; - while first_is_data - && !force_flush - && !held_for_next_turn - && blocked_data.is_none() - && batch.len() < connection.batch_max_frames - && batch_bytes < connection.batch_max_bytes - { - if let Ok(priority) = priority_rx.try_recv() { - connection.queued_frames.fetch_sub(1, Ordering::AcqRel); - batch_bytes = batch_bytes.saturating_add(priority.frame.encoded_len()); - batch.push((priority, None)); - break; - } - - let mut scheduled = scheduler.pop_ready(&connection); - if scheduled.is_none() && stream_input_open { - match scheduler.ingest_one(&connection, &mut stream_rx) { - IngressPoll::Ingested => continue, - IngressPoll::Empty => {} - IngressPoll::Disconnected => { - stream_input_open = false; - } - } - } - if scheduled.is_none() && !connection.batch_interval.is_zero() { - enum BatchNext { - Priority(WriterCommand), - Ingress(Option), - Deadline, - } - let next = tokio::select! { - biased; - _ = cancel.cancelled() => return Ok(()), - Some(priority) = priority_rx.recv() => { - BatchNext::Priority(priority) - } - ingress = stream_rx.recv(), if stream_input_open => { - BatchNext::Ingress(ingress) - } - _ = tokio::time::sleep_until(deadline.into()) => { - BatchNext::Deadline - } + let Some(command) = next else { + break; }; - match next { - BatchNext::Priority(priority) => { - connection.queued_frames.fetch_sub(1, Ordering::AcqRel); - batch_bytes = - batch_bytes.saturating_add(priority.frame.encoded_len()); - batch.push((priority, None)); - break; - } - BatchNext::Ingress(Some(ingress)) => { - scheduler.ingest(&connection, ingress); - continue; - } - BatchNext::Ingress(None) => { - stream_input_open = false; - continue; - } - BatchNext::Deadline => {} + if !Self::command_is_writable(&command) { + command.fail("response mux stream closed before write"); + continue; } - scheduled = scheduler.pop_ready(&connection); - } - let Some((next, (stream_id, state))) = scheduled else { - break; - }; - let encoded_len = next.frame.encoded_len(); - if batch_bytes.saturating_add(encoded_len) > connection.batch_max_bytes { - debug_assert!(blocked_data.is_none()); - blocked_data = Some((next, Some((stream_id, state)), Instant::now())); - break; - } - let permit = if next.frame.kind == MuxFrameKind::Data { - let required = encoded_len.min(connection.max_connection_credits) as u32; - match connection - .connection_credits - .clone() - .try_acquire_many_owned(required) - { - Ok(permit) => Some(permit), - Err(tokio::sync::TryAcquireError::NoPermits) => { - debug_assert!(blocked_data.is_none()); - blocked_data = - Some((next, Some((stream_id, state)), Instant::now())); - break; - } - Err(tokio::sync::TryAcquireError::Closed) => { - return Err(anyhow!( - "response mux connection credit window closed" - )); - } + let encoded_len = command.frame.encoded_len(); + if batch_bytes.saturating_add(encoded_len) > connection.batch_max_bytes { + pending_ordered = Some(command); + break; + } + let is_data = command.frame.kind == MuxFrameKind::Data; + batch_bytes = batch_bytes.saturating_add(encoded_len); + batch.push(WriteCommand::Ordered(command)); + if !is_data { + break; } - } else { - None - }; - let urgent = next.frame.kind != MuxFrameKind::Data; - batch_bytes = batch_bytes.saturating_add(encoded_len); - batch.push((next, permit)); - scheduler.reschedule(&connection, stream_id); - if urgent { - break; } } - for (command, _) in &batch { - let (header, payload) = command.frame.encode_parts()?; + for command in &batch { + let (header, payload) = command.frame().encode_parts()?; write_buf.write(header); write_buf.write(payload); } - let observed_batch_wait = batching_started.elapsed(); - let data_bytes = batch - .iter() - .filter(|(command, _)| command.frame.kind == MuxFrameKind::Data) - .map(|(command, _)| command.frame.encoded_len() as u64) - .sum::(); - connection - .sent_data_bytes - .fetch_add(data_bytes, Ordering::AcqRel); let write_result = write_buf.write_all_counted(&mut write_half).await; let write_calls = write_result .as_ref() .map(|(_, calls)| *calls) .unwrap_or_default(); - let write_result: Result<()> = - write_result.map(|_| ()).map_err(anyhow::Error::from); - - for (command, permit) in &mut batch { - if metrics_enabled { - if let Some(enqueued_at) = command.enqueued_at { - response_mux::QUEUE_RESIDENCE_SECONDS - .with_label_values(&["worker"]) - .observe(enqueued_at.elapsed().as_secs_f64()); - } - if let Some(enqueued_at) = command.priority_enqueued_at { - response_mux::PRIORITY_QUEUE_RESIDENCE_SECONDS - .observe(enqueued_at.elapsed().as_secs_f64()); - } - frame_counters.inc(command.frame.kind.metric_label()); - } - if let Some(written) = command.written.take() { - let _ = written.send( - write_result - .as_ref() - .map(|_| ()) - .map_err(|err| err.to_string()), - ); - } - if let Some(permit) = permit.take() { - permit.forget(); - } + let completion = write_result + .as_ref() + .map(|_| ()) + .map_err(|err| err.to_string()); + for command in &mut batch { + command.complete(&completion); } response_mux::WRITE_CALLS_TOTAL .with_label_values(&["worker"]) .inc_by(write_calls); + frames_per_write.observe(batch.len() as f64); write_result?; - if metrics_enabled { - frames_per_write.observe(batch.len() as f64); - response_mux::BATCH_BYTES - .with_label_values(&["worker"]) - .observe(batch_bytes as f64); - response_mux::BATCH_WAIT_SECONDS - .with_label_values(&["worker"]) - .observe(observed_batch_wait.as_secs_f64()); - } } } .await; - if metrics_enabled { - queue_depth.sub(reported_queue_depth); + if let Some(command) = pending_ordered.take() { + command.fail("response mux writer stopped"); + } + while let Ok(command) = ordered_rx.try_recv() { + command.fail("response mux writer stopped"); } if let Some(connection) = weak.upgrade() { - if let Some((command, _, _)) = blocked_data.take() { - command.fail("response mux writer stopped"); - } - while let Ok(ingress) = stream_rx.try_recv() { - scheduler.ingest(&connection, ingress); - } - scheduler.fail_all(&connection, "response mux writer stopped"); - while let Ok(command) = priority_rx.try_recv() { - connection.queued_frames.fetch_sub(1, Ordering::AcqRel); - command.fail("response mux priority writer stopped"); - } connection.fail( &result .err() @@ -1077,19 +497,11 @@ impl MuxConnection { cancel: CancellationToken, mut reported_data_segments: Option, ) { - let frame_counters = response_mux::FrameCounters::for_direction("frontend_to_worker"); - let window_updates = response_mux::WINDOW_UPDATES_TOTAL - .with_label_values(&["frontend_to_worker"]) - .clone(); - let connection_window_updates = response_mux::WINDOW_UPDATES_TOTAL - .with_label_values(&["connection_frontend_to_worker"]) - .clone(); - let metrics_enabled = per_frame_metrics_enabled(); let mut packet_tick = tokio::time::interval(Duration::from_millis(100)); packet_tick.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay); let result: Result<()> = async { loop { - let message = tokio::select! { + let frame = tokio::select! { _ = cancel.cancelled() => break, _ = packet_tick.tick(), if reported_data_segments.is_some() => { if let Some(current) = crate::pipeline::network::tcp::tcp_data_segments_out(reader.get_ref().as_ref()) { @@ -1105,28 +517,20 @@ impl MuxConnection { None => return Err(anyhow!("frontend closed response mux connection")), }, }; - let frame = message; - if metrics_enabled { - frame_counters.inc(frame.kind.metric_label()); - } let connection = weak .upgrade() .ok_or_else(|| anyhow!("response mux connection dropped"))?; - if frame.kind == MuxFrameKind::ConnectionAck { - let offset = frame.connection_ack_offset()?; - if offset > 0 { - connection.acknowledge_connection_credits(offset)?; - if metrics_enabled { - connection_window_updates.inc(); - } - } - continue; + if matches!( + frame.kind, + MuxFrameKind::ConnectionHello | MuxFrameKind::ConnectionReady + ) { + anyhow::bail!("unexpected response mux connection frame after handshake"); } let Some(state) = connection.streams.get(&frame.stream_id) else { if frame.kind != MuxFrameKind::Reset { - connection.try_send_priority_command(WriterCommand::priority( - MuxFrame::empty(MuxFrameKind::Reset, frame.stream_id), - None, + connection.try_send_urgent(MuxFrame::empty( + MuxFrameKind::Reset, + frame.stream_id, )); } continue; @@ -1136,9 +540,9 @@ impl MuxConnection { let credits = frame.window_credits()? as usize; if credits == 0 || credits > state.max_credits { drop(state); - connection.try_send_priority_command(WriterCommand::priority( - MuxFrame::empty(MuxFrameKind::Reset, frame.stream_id), - None, + connection.try_send_urgent(MuxFrame::empty( + MuxFrameKind::Reset, + frame.stream_id, )); connection.remove_stream( frame.stream_id, @@ -1148,9 +552,6 @@ impl MuxConnection { continue; } state.replenish_credits(credits); - if metrics_enabled { - window_updates.inc(); - } } MuxFrameKind::Stop => { state.record_cancellation(); @@ -1167,9 +568,9 @@ impl MuxConnection { } _ => { drop(state); - connection.try_send_priority_command(WriterCommand::priority( - MuxFrame::empty(MuxFrameKind::Reset, frame.stream_id), - None, + connection.try_send_urgent(MuxFrame::empty( + MuxFrameKind::Reset, + frame.stream_id, )); connection.remove_stream( frame.stream_id, @@ -1191,7 +592,6 @@ impl MuxConnection { .with_label_values(&["mux", "worker"]) .inc_by(current.saturating_sub(previous)); } - if let Some(connection) = weak.upgrade() { connection.fail( &result @@ -1206,57 +606,26 @@ impl MuxConnection { struct HostPool { address: String, frontend_server_id: Uuid, - version: u8, connections: RwLock>>, connect_lock: tokio::sync::Mutex<()>, - next_connection_id: AtomicU64, warming: AtomicBool, - maintenance_started: AtomicBool, - lifecycle: Mutex, cancel: CancellationToken, config: PoolConfig, } -struct HostLifecycle { - last_used: Instant, - retiring: bool, - openers: usize, -} - -struct HostOpenGuard { - host: Arc, -} - -impl Drop for HostOpenGuard { - fn drop(&mut self) { - let mut lifecycle = self.host.lifecycle.lock(); - lifecycle.openers = lifecycle.openers.saturating_sub(1); - lifecycle.last_used = Instant::now(); - } -} - impl HostPool { fn new( address: String, frontend_server_id: Uuid, - version: u8, cancel: CancellationToken, config: PoolConfig, ) -> Arc { Arc::new(Self { address, frontend_server_id, - version, connections: RwLock::new(Vec::new()), connect_lock: tokio::sync::Mutex::new(()), - next_connection_id: AtomicU64::new(1), warming: AtomicBool::new(false), - maintenance_started: AtomicBool::new(false), - lifecycle: Mutex::new(HostLifecycle { - last_used: Instant::now(), - retiring: false, - openers: 0, - }), cancel, config, }) @@ -1279,30 +648,23 @@ impl HostPool { self.connect_new().await } - async fn connect_additional(&self) -> Result> { + async fn connect_additional(&self) -> Result<()> { let _guard = self.connect_lock.lock().await; - if self.healthy_connections().len() >= self.config.pool_size { - return self - .healthy_connections() - .first() - .cloned() - .ok_or_else(|| anyhow!("response mux host pool has no healthy connection")); + if self.healthy_connections().len() < self.config.pool_size { + self.connect_new().await?; } - self.connect_new().await + Ok(()) } async fn connect_new(&self) -> Result> { - let replacing_failed_connection = self + let replacing_failed = self .connections .read() .iter() .any(|connection| !connection.is_healthy()); - let id = self.next_connection_id.fetch_add(1, Ordering::Relaxed); let connection = MuxConnection::connect( - id, &self.address, self.frontend_server_id, - self.version, self.cancel.clone(), self.config, ) @@ -1310,7 +672,7 @@ impl HostPool { let mut connections = self.connections.write(); connections.retain(|candidate| candidate.is_healthy()); connections.push(connection.clone()); - if replacing_failed_connection { + if replacing_failed { response_mux::RECONNECTS_TOTAL .with_label_values(&["worker"]) .inc(); @@ -1318,30 +680,11 @@ impl HostPool { Ok(connection) } - fn start_maintenance(self: &Arc) { - if self.maintenance_started.swap(true, Ordering::AcqRel) { - return; - } - let host = Arc::clone(self); - tokio::spawn(async move { - let mut interval = tokio::time::interval(Duration::from_secs(1)); - loop { - interval.tick().await; - if host.cancel.is_cancelled() { - break; - } - if host.healthy_connections().len() < host.config.pool_size { - host.warm(); - } - } - }); - } - fn warm(self: &Arc) { if self.warming.swap(true, Ordering::AcqRel) { return; } - let host = Arc::clone(self); + let host = self.clone(); tokio::spawn(async move { while host.healthy_connections().len() < host.config.pool_size && !host.cancel.is_cancelled() @@ -1355,21 +698,11 @@ impl HostPool { }); } - async fn connection(self: &Arc) -> Result, HostOpenGuard)>> { - let opener = { - let mut lifecycle = self.lifecycle.lock(); - if lifecycle.retiring { - return Ok(None); - } - lifecycle.openers += 1; - lifecycle.last_used = Instant::now(); - HostOpenGuard { host: self.clone() } - }; + async fn connection(self: &Arc) -> Result> { let mut healthy = self.healthy_connections(); if healthy.is_empty() { healthy.push(self.ensure_first().await?); } - self.start_maintenance(); if healthy.len() < self.config.pool_size { self.warm(); } @@ -1379,30 +712,14 @@ impl HostPool { .min_by_key(|(_, connection)| { ( connection.active_streams.load(Ordering::Acquire), - connection.queued_bytes.load(Ordering::Acquire), + Reverse(connection.queued_byte_slots.available_permits()), ) }) .map(|(index, _)| index) - .expect("healthy response mux connection list is non-empty"); - Ok(Some((healthy.swap_remove(index), opener))) + .ok_or_else(|| anyhow!("response mux host pool has no healthy connection"))?; + Ok(healthy.swap_remove(index)) } - - fn try_retire(&self) -> bool { - let mut lifecycle = self.lifecycle.lock(); - if lifecycle.retiring - || lifecycle.openers != 0 - || lifecycle.last_used.elapsed() < self.config.idle_ttl - || !self - .healthy_connections() - .iter() - .all(|connection| connection.active_streams.load(Ordering::Acquire) == 0) - { - return false; - } - lifecycle.retiring = true; - true - } -} +} pub struct ResponseMuxClientPool { hosts: DashMap>, @@ -1414,7 +731,6 @@ pub struct ResponseMuxClientPool { struct HostKey { address: String, frontend_server_id: Uuid, - version: u8, } impl ResponseMuxClientPool { @@ -1427,60 +743,51 @@ impl ResponseMuxClientPool { cancel, config: PoolConfig::from_runtime(runtime_config), }); - Self::start_cleanup(&pool); + Self::start_maintenance(&pool); pool } #[cfg(test)] - fn new_for_test_with_connection_window( + fn new_for_test( cancel: CancellationToken, + runtime_config: ResponseMuxConfig, pool_size: usize, - writer_queue: usize, - initial_window: usize, - connection_window: usize, - idle_ttl: Duration, - connect_timeout: Duration, + queued_bytes: usize, ) -> Arc { + let mut config = PoolConfig::from_runtime(runtime_config); + config.pool_size = pool_size.max(1); + config.queued_bytes = queued_bytes.max(1); let pool = Arc::new(Self { hosts: DashMap::new(), cancel, - config: PoolConfig { - pool_size: pool_size.max(1), - writer_queue: writer_queue.max(1), - stream_writer_queue: RESPONSE_MUX_STREAM_WRITER_QUEUE, - initial_window: initial_window.max(1), - connection_window: connection_window.max(1), - batch_interval: Duration::ZERO, - batch_max_bytes: 65_536, - batch_max_frames: 64, - packet_metrics: false, - idle_ttl, - connect_timeout, - }, + config, }); - Self::start_cleanup(&pool); + Self::start_maintenance(&pool); pool } - fn start_cleanup(pool: &Arc) { - let weak = Arc::downgrade(pool); + fn start_maintenance(pool: &Arc) { + let pool = Arc::downgrade(pool); tokio::spawn(async move { - let mut interval = tokio::time::interval(Duration::from_secs(30)); + let mut interval = tokio::time::interval(Duration::from_secs(1)); loop { interval.tick().await; - let Some(pool) = weak.upgrade() else { + let Some(pool) = pool.upgrade() else { break; }; if pool.cancel.is_cancelled() { break; } - pool.hosts.retain(|_, host| { - let retiring = host.try_retire(); - if retiring { - host.cancel.cancel(); + let hosts = pool + .hosts + .iter() + .map(|entry| entry.value().clone()) + .collect::>(); + for host in hosts { + if host.healthy_connections().len() < host.config.pool_size { + host.warm(); } - !retiring - }); + } } }); } @@ -1493,47 +800,31 @@ impl ResponseMuxClientPool { ) -> Result { let info = ResponseMuxConnectionInfo::try_from(info) .context("tcp-response-mux-connection-info-error")?; - if info.version != RESPONSE_MUX_VERSION { - anyhow::bail!( - "unsupported response mux version {}; expected {}", - info.version, - RESPONSE_MUX_VERSION - ); - } if info.context != context.id() { - return Err(anyhow!( + anyhow::bail!( "response mux context mismatch: expected {}, got {}", context.id(), info.context - )); + ); } let stream_id = info.stream_id; let host_key = HostKey { address: info.address.clone(), frontend_server_id: info.frontend_server_id, - version: info.version, - }; - let (connection, _opener) = loop { - let host = self - .hosts - .entry(host_key.clone()) - .or_insert_with(|| { - HostPool::new( - info.address.clone(), - info.frontend_server_id, - info.version, - self.cancel.child_token(), - self.config, - ) - }) - .clone(); - if let Some(connection) = host.connection().await? { - break connection; - } - self.hosts - .remove_if(&host_key, |_, candidate| Arc::ptr_eq(candidate, &host)); - host.cancel.cancel(); }; + let host = self + .hosts + .entry(host_key) + .or_insert_with(|| { + HostPool::new( + info.address, + info.frontend_server_id, + self.cancel.child_token(), + self.config, + ) + }) + .clone(); + let connection = host.connection().await?; let state = Arc::new(WorkerStreamState { context, cancellation_counter, @@ -1548,21 +839,19 @@ impl ResponseMuxClientPool { Entry::Vacant(entry) => { entry.insert(state.clone()); } - Entry::Occupied(_) => { - return Err(anyhow!("duplicate response mux stream id {stream_id}")); - } + Entry::Occupied(_) => anyhow::bail!("duplicate response mux stream id {stream_id}"), } connection.active_streams.fetch_add(1, Ordering::AcqRel); response_mux::ACTIVE_STREAMS .with_label_values(&["worker"]) .inc(); - let sender = StreamSender::multiplexed(Arc::new(MuxResponseStreamSender { + Ok(StreamSender::multiplexed(MuxResponseStreamSender { stream_id, connection, state, - finished: AtomicBool::new(false), - })); - Ok(sender) + prologue_sent: false, + finished: false, + })) } #[cfg(test)] @@ -1581,58 +870,52 @@ impl ResponseMuxClientPool { } #[cfg(test)] - pub(crate) fn stream_connection_id(&self, address: &str, stream_id: Uuid) -> Option { + pub(crate) fn stream_connection_id(&self, address: &str, stream_id: Uuid) -> Option { self.host_for_address(address).and_then(|host| { host.connections .read() .iter() .find(|connection| connection.streams.contains_key(&stream_id)) - .map(|connection| connection.id) + .map(|connection| Arc::as_ptr(connection) as usize) }) } #[cfg(test)] - pub(crate) fn fail_connection(&self, address: &str, connection_id: u64) { + pub(crate) fn fail_connection(&self, address: &str, connection_id: usize) { if let Some(host) = self.host_for_address(address) && let Some(connection) = host .connections .read() .iter() - .find(|connection| connection.id == connection_id) + .find(|connection| Arc::as_ptr(connection) as usize == connection_id) { connection.fail("test-injected physical connection failure"); } } } -struct MuxResponseStreamSender { +pub(crate) struct MuxResponseStreamSender { stream_id: Uuid, connection: Arc, state: Arc, - finished: AtomicBool, + prologue_sent: bool, + finished: bool, } impl MuxResponseStreamSender { - async fn acquire_writer_permit( - &self, - state: &WorkerStreamState, - ) -> Result { - let slots = state.writer_slots.clone(); - match slots.try_acquire_owned() { + async fn acquire_writer_permit(&self) -> Result { + match self.state.writer_slots.clone().try_acquire_owned() { Ok(permit) => Ok(permit), Err(tokio::sync::TryAcquireError::NoPermits) => { - let wait_start = per_frame_metrics_enabled().then(Instant::now); - let permit = state + response_mux::STALLS_TOTAL + .with_label_values(&["writer_admission"]) + .inc(); + self.state .writer_slots .clone() .acquire_owned() .await - .map_err(|_| anyhow!("response mux stream closed during writer admission"))?; - if let Some(wait_start) = wait_start { - response_mux::WRITER_ADMISSION_STALL_SECONDS - .observe(wait_start.elapsed().as_secs_f64()); - } - Ok(permit) + .map_err(|_| anyhow!("response mux stream closed during writer admission")) } Err(tokio::sync::TryAcquireError::Closed) => { anyhow::bail!("response mux stream closed during writer admission") @@ -1645,98 +928,65 @@ impl MuxResponseStreamSender { frame: MuxFrame, written: Option>>, ) -> Result<()> { - self.enqueue_ordered_on(&self.connection, &self.state, frame, written) - .await - } - - async fn enqueue_ordered_on( - &self, - connection: &Arc, - state: &Arc, - frame: MuxFrame, - written: Option>>, - ) -> Result<()> { - if !connection.is_healthy() { + if !self.connection.is_healthy() { anyhow::bail!("response mux connection is unhealthy"); } - let writer_permit = self.acquire_writer_permit(state).await?; - let required = frame.encoded_len().min(connection.max_queued_bytes) as u32; - let queued_byte_permit = match connection + let writer_permit = self.acquire_writer_permit().await?; + let required = frame.encoded_len().min(self.connection.max_queued_bytes) as u32; + let queued_byte_permit = match self + .connection .queued_byte_slots .clone() .try_acquire_many_owned(required) { Ok(permit) => permit, Err(tokio::sync::TryAcquireError::NoPermits) => { - let wait_start = per_frame_metrics_enabled().then(Instant::now); - let slots = connection.queued_byte_slots.clone(); - let permit = tokio::select! { - _ = state.close_token.cancelled() => { + response_mux::STALLS_TOTAL + .with_label_values(&["queued_byte_admission"]) + .inc(); + let slots = self.connection.queued_byte_slots.clone(); + tokio::select! { + _ = self.state.close_token.cancelled() => { anyhow::bail!("response mux stream closed during byte-queue admission") } permit = slots.acquire_many_owned(required) => permit.map_err(|_| { anyhow!("response mux connection closed during byte-queue admission") })?, - }; - if let Some(wait_start) = wait_start { - response_mux::QUEUED_BYTE_ADMISSION_STALL_SECONDS - .observe(wait_start.elapsed().as_secs_f64()); } - permit } Err(tokio::sync::TryAcquireError::Closed) => { anyhow::bail!("response mux connection closed during byte-queue admission") } }; - connection - .enqueue_stream_command( - self.stream_id, - state.clone(), - WriterCommand::new(frame, written) - .with_writer_permit(writer_permit) - .with_queued_byte_permit(queued_byte_permit), - ) - .await - } - - async fn enqueue_priority_and_wait(&self, frame: MuxFrame) -> Result<()> { - let (written_tx, written_rx) = oneshot::channel(); self.connection - .send_priority_command(WriterCommand::priority(frame, Some(written_tx))) - .await?; - written_rx + .send_ordered(OrderedCommand { + frame, + state: self.state.clone(), + written, + _writer_permit: writer_permit, + _queued_byte_permit: queued_byte_permit, + }) .await - .map_err(|_| anyhow!("response mux priority writer dropped acknowledgement"))? - .map_err(anyhow::Error::msg) - } - - fn remove_stream(&self) { - self.connection - .remove_stream(self.stream_id, "response mux stream completed", false); } -} -#[async_trait::async_trait] -impl MultiplexedStreamSender for MuxResponseStreamSender { - async fn send_data(&self, data: bytes::Bytes) -> Result<()> { + pub(crate) async fn send_data(&self, data: bytes::Bytes) -> Result<()> { + if !self.prologue_sent || self.finished { + anyhow::bail!("response mux data sent outside an active stream"); + } let encoded_len = super::MUX_HEADER_LEN.saturating_add(data.len()); let required = encoded_len.min(self.state.max_credits) as u32; let permit = match self.state.credits.clone().try_acquire_many_owned(required) { Ok(permit) => permit, Err(tokio::sync::TryAcquireError::NoPermits) => { - let wait_start = per_frame_metrics_enabled().then(Instant::now); - let permit = self - .state + response_mux::STALLS_TOTAL + .with_label_values(&["stream_credit"]) + .inc(); + self.state .credits .clone() .acquire_many_owned(required) .await - .map_err(|_| anyhow!("response mux stream closed while waiting for credits"))?; - if let Some(wait_start) = wait_start { - response_mux::FLOW_CONTROL_STALL_SECONDS - .observe(wait_start.elapsed().as_secs_f64()); - } - permit + .map_err(|_| anyhow!("response mux stream closed while waiting for credits"))? } Err(tokio::sync::TryAcquireError::Closed) => { anyhow::bail!("response mux stream closed while waiting for credits") @@ -1754,26 +1004,44 @@ impl MultiplexedStreamSender for MuxResponseStreamSender { result } - async fn send_prologue(&self, error: Option) -> Result<(), String> { + pub(crate) async fn send_prologue(&mut self, error: Option) -> Result<(), String> { + if self.prologue_sent || self.finished { + return Err("response mux prologue already sent".to_string()); + } let terminal_error = error.is_some(); - let payload = - serde_json::to_vec(&ResponseStreamPrologue { error }).map_err(|err| err.to_string())?; - let frame = MuxFrame::new(MuxFrameKind::Prologue, self.stream_id, payload.into()); - let result = self.enqueue_priority_and_wait(frame).await; - if result.is_ok() && terminal_error { - self.finished.store(true, Ordering::Release); + let payload = error.map(bytes::Bytes::from).unwrap_or_default(); + self.connection + .send_urgent(MuxFrame::new( + MuxFrameKind::Prologue, + self.stream_id, + payload, + )) + .await + .map_err(|err| err.to_string())?; + self.prologue_sent = true; + if terminal_error { + self.finished = true; self.remove_stream(); } - result.map_err(|err| err.to_string()) + Ok(()) } - async fn finish(&self) -> Result<()> { - if self.finished.swap(true, Ordering::AcqRel) { + pub(crate) async fn finish(mut self) -> Result<()> { + if self.finished { return Ok(()); } + if !self.prologue_sent { + anyhow::bail!("response mux stream finished before its prologue"); + } + self.finished = true; let (written_tx, written_rx) = oneshot::channel(); - let end = MuxFrame::empty(MuxFrameKind::End, self.stream_id); - if let Err(err) = self.enqueue_ordered(end, Some(written_tx)).await { + let result = self + .enqueue_ordered( + MuxFrame::empty(MuxFrameKind::End, self.stream_id), + Some(written_tx), + ) + .await; + if let Err(err) = result { self.remove_stream(); return Err(err).context("response mux writer stopped before end"); } @@ -1784,25 +1052,24 @@ impl MultiplexedStreamSender for MuxResponseStreamSender { self.remove_stream(); result } + + fn remove_stream(&self) { + self.connection + .remove_stream(self.stream_id, "response mux stream completed", false); + } } impl Drop for MuxResponseStreamSender { fn drop(&mut self) { - if self.finished.swap(true, Ordering::AcqRel) { + if self.finished { return; } + self.finished = true; response_mux::RESETS_TOTAL .with_label_values(&["worker", "publisher_drop"]) .inc(); self.connection - .try_send_priority_command(WriterCommand::priority( - MuxFrame::new( - MuxFrameKind::Reset, - self.stream_id, - bytes::Bytes::from_static(b"response sender dropped before finish"), - ), - None, - )); + .try_send_urgent(MuxFrame::empty(MuxFrameKind::Reset, self.stream_id)); self.remove_stream(); } } @@ -1810,7 +1077,6 @@ impl Drop for MuxResponseStreamSender { #[cfg(test)] mod tests { use super::*; - use crate::pipeline::network::tcp::mux::ResponseMuxConfig; use crate::{ engine::AsyncEngineContextProvider, pipeline::{ @@ -1820,517 +1086,6 @@ mod tests { }; use futures::{StreamExt as _, stream::FuturesUnordered}; - const TEST_STREAM_WINDOW: usize = 64; - const TEST_STREAM_WINDOW_UPDATE: u32 = 32; - const TEST_CONNECTION_WINDOW: usize = 256; - - fn worker_stream_state(initial_credits: usize, writer_slots: usize) -> Arc { - let context = Context::new(()); - Arc::new(WorkerStreamState { - context: context.context(), - cancellation_counter: None, - cancellation_recorded: AtomicBool::new(false), - credits: Arc::new(Semaphore::new(initial_credits)), - max_credits: TEST_STREAM_WINDOW, - writer_slots: Arc::new(Semaphore::new(writer_slots)), - closed: AtomicBool::new(false), - close_token: CancellationToken::new(), - }) - } - - #[test] - fn credit_replenishment_never_exceeds_the_fixed_maximum() { - let state = worker_stream_state(TEST_STREAM_WINDOW, RESPONSE_MUX_STREAM_WRITER_QUEUE); - state - .credits - .clone() - .try_acquire_many_owned((TEST_STREAM_WINDOW - 16) as u32) - .expect("initial credits should be available") - .forget(); - - assert_eq!(state.credits.available_permits(), 16); - assert_eq!( - state.replenish_credits(TEST_STREAM_WINDOW_UPDATE as usize), - TEST_STREAM_WINDOW_UPDATE as usize - ); - assert_eq!( - state.credits.available_permits(), - 16 + TEST_STREAM_WINDOW_UPDATE as usize - ); - assert_eq!( - state.replenish_credits(TEST_STREAM_WINDOW), - TEST_STREAM_WINDOW - 16 - TEST_STREAM_WINDOW_UPDATE as usize - ); - assert_eq!(state.credits.available_permits(), TEST_STREAM_WINDOW); - assert_eq!( - state.replenish_credits(TEST_STREAM_WINDOW_UPDATE as usize), - 0 - ); - assert_eq!(state.credits.available_permits(), TEST_STREAM_WINDOW); - } - - #[test] - fn detailed_metric_timestamps_are_opt_in() { - let stream_id = Uuid::new_v4(); - let disabled = WriterCommand::new_with_metrics( - MuxFrame::empty(MuxFrameKind::End, stream_id), - None, - false, - ); - let enabled = WriterCommand::new_with_metrics( - MuxFrame::empty(MuxFrameKind::End, stream_id), - None, - true, - ); - - assert!(disabled.enqueued_at.is_none()); - assert!(enabled.enqueued_at.is_some()); - } - - #[test] - fn ingress_admission_stops_after_one_frame_when_a_stream_becomes_ready() { - let stream_id = Uuid::new_v4(); - let state = worker_stream_state(TEST_STREAM_WINDOW, RESPONSE_MUX_STREAM_WRITER_QUEUE); - let (priority_tx, _priority_rx) = mpsc::channel(1); - let (stream_tx, mut stream_rx) = mpsc::channel(4); - let connection = MuxConnection { - id: 1, - cancel: CancellationToken::new(), - priority_tx, - stream_tx: stream_tx.clone(), - streams: DashMap::new(), - healthy: AtomicBool::new(true), - active_streams: AtomicUsize::new(1), - queued_frames: AtomicUsize::new(2), - queued_bytes: AtomicUsize::new(0), - max_queued_bytes: TEST_CONNECTION_WINDOW, - queued_byte_slots: Arc::new(Semaphore::new(TEST_CONNECTION_WINDOW)), - connection_credits: Arc::new(Semaphore::new(TEST_CONNECTION_WINDOW)), - max_connection_credits: TEST_CONNECTION_WINDOW, - sent_data_bytes: AtomicU64::new(0), - acknowledged_data_bytes: AtomicU64::new(0), - batch_interval: Duration::ZERO, - batch_max_bytes: 65_536, - batch_max_frames: 64, - }; - for payload in [b"first".as_slice(), b"second".as_slice()] { - stream_tx - .try_send(StreamIngress { - stream_id, - state: state.clone(), - command: WriterCommand::new( - MuxFrame::new( - MuxFrameKind::Data, - stream_id, - bytes::Bytes::copy_from_slice(payload), - ), - None, - ), - }) - .unwrap(); - } - - let mut scheduler = WriterScheduler::default(); - assert_eq!( - scheduler.ingest_one(&connection, &mut stream_rx), - IngressPoll::Ingested - ); - assert_eq!(stream_rx.len(), 1); - assert_eq!(scheduler.ready.len(), 1); - assert_eq!(scheduler.streams[&stream_id].pending.len(), 1); - } - - #[test] - fn cumulative_connection_ack_replenishes_credits_without_exceeding_the_window() { - let (priority_tx, _priority_rx) = mpsc::channel(1); - let (stream_tx, _stream_rx) = mpsc::channel(1); - let remaining = 16; - let consumed = TEST_CONNECTION_WINDOW - remaining; - let connection = MuxConnection { - id: 1, - cancel: CancellationToken::new(), - priority_tx, - stream_tx, - streams: DashMap::new(), - healthy: AtomicBool::new(true), - active_streams: AtomicUsize::new(0), - queued_frames: AtomicUsize::new(0), - queued_bytes: AtomicUsize::new(0), - max_queued_bytes: TEST_CONNECTION_WINDOW, - queued_byte_slots: Arc::new(Semaphore::new(TEST_CONNECTION_WINDOW)), - connection_credits: Arc::new(Semaphore::new(TEST_CONNECTION_WINDOW)), - max_connection_credits: TEST_CONNECTION_WINDOW, - sent_data_bytes: AtomicU64::new(consumed as u64), - acknowledged_data_bytes: AtomicU64::new(0), - batch_interval: Duration::ZERO, - batch_max_bytes: 65_536, - batch_max_frames: 64, - }; - connection - .connection_credits - .clone() - .try_acquire_many_owned(consumed as u32) - .unwrap() - .forget(); - - assert_eq!(connection.connection_credits.available_permits(), remaining); - connection.acknowledge_connection_credits(128).unwrap(); - assert_eq!( - connection.connection_credits.available_permits(), - remaining + 128 - ); - connection - .acknowledge_connection_credits(consumed as u64) - .unwrap(); - assert_eq!( - connection.connection_credits.available_permits(), - TEST_CONNECTION_WINDOW - ); - connection - .acknowledge_connection_credits(consumed as u64) - .unwrap(); - assert_eq!( - connection.connection_credits.available_permits(), - TEST_CONNECTION_WINDOW - ); - assert!(connection.acknowledge_connection_credits(127).is_err()); - assert!( - connection - .acknowledge_connection_credits(consumed as u64 + 1) - .is_err() - ); - } - - #[tokio::test] - async fn closing_a_stream_wakes_credit_and_writer_admission_waiters() { - let state = worker_stream_state(0, 0); - let (priority_tx, _priority_rx) = mpsc::channel(1); - let (stream_tx, _stream_rx) = mpsc::channel(1); - let connection = MuxConnection { - id: 1, - cancel: CancellationToken::new(), - priority_tx, - stream_tx, - streams: DashMap::new(), - healthy: AtomicBool::new(true), - active_streams: AtomicUsize::new(0), - queued_frames: AtomicUsize::new(0), - queued_bytes: AtomicUsize::new(0), - max_queued_bytes: TEST_CONNECTION_WINDOW, - queued_byte_slots: Arc::new(Semaphore::new(TEST_CONNECTION_WINDOW)), - connection_credits: Arc::new(Semaphore::new(TEST_CONNECTION_WINDOW)), - max_connection_credits: TEST_CONNECTION_WINDOW, - sent_data_bytes: AtomicU64::new(0), - acknowledged_data_bytes: AtomicU64::new(0), - batch_interval: Duration::ZERO, - batch_max_bytes: 65_536, - batch_max_frames: 64, - }; - - let credit_waiter = tokio::spawn({ - let credits = state.credits.clone(); - async move { credits.acquire_owned().await } - }); - let writer_waiter = tokio::spawn({ - let writer_slots = state.writer_slots.clone(); - async move { writer_slots.acquire_owned().await } - }); - tokio::task::yield_now().await; - - connection.close_stream_state(&state, "test stream closed"); - - assert!(credit_waiter.await.unwrap().is_err()); - assert!(writer_waiter.await.unwrap().is_err()); - assert_eq!(state.replenish_credits(1_024), 0); - assert_eq!(state.credits.available_permits(), 0); - } - - #[tokio::test] - async fn end_bypasses_exhausted_connection_credit_at_batch_byte_limit() { - let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); - let address = listener.local_addr().unwrap(); - let client = TcpStream::connect(address).await.unwrap(); - let (server, _) = listener.accept().await.unwrap(); - let (_, write_half) = client.into_split(); - let (server_read, _) = server.into_split(); - - let stream_id = Uuid::new_v4(); - let (end_written_tx, end_written_rx) = oneshot::channel(); - let state = worker_stream_state(TEST_STREAM_WINDOW, RESPONSE_MUX_STREAM_WRITER_QUEUE); - - let (priority_tx, priority_rx) = mpsc::channel(1); - let (stream_tx, stream_rx) = mpsc::channel(8); - let cancel = CancellationToken::new(); - let connection = Arc::new(MuxConnection { - id: 1, - cancel: cancel.clone(), - priority_tx, - stream_tx, - streams: DashMap::new(), - healthy: AtomicBool::new(true), - active_streams: AtomicUsize::new(1), - queued_frames: AtomicUsize::new(0), - queued_bytes: AtomicUsize::new(0), - max_queued_bytes: TEST_CONNECTION_WINDOW, - queued_byte_slots: Arc::new(Semaphore::new(TEST_CONNECTION_WINDOW)), - connection_credits: Arc::new(Semaphore::new(64)), - max_connection_credits: 64, - sent_data_bytes: AtomicU64::new(0), - acknowledged_data_bytes: AtomicU64::new(0), - batch_interval: Duration::ZERO, - batch_max_bytes: 64, - batch_max_frames: 64, - }); - connection.streams.insert(stream_id, state.clone()); - - let writer_task = tokio::spawn(MuxConnection::writer_task( - Arc::downgrade(&connection), - write_half, - priority_rx, - stream_rx, - cancel, - )); - connection - .enqueue_stream_command( - stream_id, - state.clone(), - WriterCommand::new( - MuxFrame::new( - MuxFrameKind::Data, - stream_id, - bytes::Bytes::from(vec![b'x'; 40]), - ), - None, - ), - ) - .await - .unwrap(); - connection - .enqueue_stream_command( - stream_id, - state, - WriterCommand::new( - MuxFrame::empty(MuxFrameKind::End, stream_id), - Some(end_written_tx), - ), - ) - .await - .unwrap(); - let reader_task = tokio::spawn(async move { - let mut reader = FramedRead::new(server_read, MuxCodec::default()); - let data = reader.next().await.unwrap().unwrap(); - let end = reader.next().await.unwrap().unwrap(); - (data.kind, end.kind) - }); - - tokio::time::timeout(Duration::from_millis(200), end_written_rx) - .await - .expect("End waited for exhausted Data credits") - .unwrap() - .unwrap(); - assert_eq!( - reader_task.await.unwrap(), - (MuxFrameKind::Data, MuxFrameKind::End) - ); - writer_task.abort(); - } - - #[tokio::test] - async fn parked_frame_is_not_overwritten_by_another_ready_stream() { - let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); - let address = listener.local_addr().unwrap(); - let client = TcpStream::connect(address).await.unwrap(); - let (server, _) = listener.accept().await.unwrap(); - let (_, write_half) = client.into_split(); - let (server_read, _) = server.into_split(); - - let stream_a = Uuid::new_v4(); - let stream_b = Uuid::new_v4(); - let state_a = worker_stream_state(TEST_STREAM_WINDOW, RESPONSE_MUX_STREAM_WRITER_QUEUE); - let state_b = worker_stream_state(TEST_STREAM_WINDOW, RESPONSE_MUX_STREAM_WRITER_QUEUE); - let (priority_tx, priority_rx) = mpsc::channel(1); - let (stream_tx, stream_rx) = mpsc::channel(8); - let cancel = CancellationToken::new(); - let connection = Arc::new(MuxConnection { - id: 1, - cancel: cancel.clone(), - priority_tx, - stream_tx, - streams: DashMap::new(), - healthy: AtomicBool::new(true), - active_streams: AtomicUsize::new(2), - queued_frames: AtomicUsize::new(0), - queued_bytes: AtomicUsize::new(0), - max_queued_bytes: TEST_CONNECTION_WINDOW, - queued_byte_slots: Arc::new(Semaphore::new(TEST_CONNECTION_WINDOW)), - connection_credits: Arc::new(Semaphore::new(TEST_CONNECTION_WINDOW)), - max_connection_credits: TEST_CONNECTION_WINDOW, - sent_data_bytes: AtomicU64::new(0), - acknowledged_data_bytes: AtomicU64::new(0), - batch_interval: Duration::ZERO, - batch_max_bytes: 64, - batch_max_frames: 64, - }); - connection.streams.insert(stream_a, state_a.clone()); - connection.streams.insert(stream_b, state_b.clone()); - - for (state, frame) in [ - ( - state_a.clone(), - MuxFrame::new( - MuxFrameKind::Data, - stream_a, - bytes::Bytes::from_static(b"a-small"), - ), - ), - ( - state_a, - MuxFrame::new( - MuxFrameKind::Data, - stream_a, - bytes::Bytes::from(vec![b'A'; 40]), - ), - ), - ( - state_b, - MuxFrame::new( - MuxFrameKind::Data, - stream_b, - bytes::Bytes::from_static(b"b-ready"), - ), - ), - ] { - connection - .enqueue_stream_command(frame.stream_id, state, WriterCommand::new(frame, None)) - .await - .unwrap(); - } - - let writer_task = tokio::spawn(MuxConnection::writer_task( - Arc::downgrade(&connection), - write_half, - priority_rx, - stream_rx, - cancel, - )); - let mut reader = FramedRead::new(server_read, MuxCodec::default()); - let first = reader.next().await.unwrap().unwrap(); - let second = reader.next().await.unwrap().unwrap(); - let third = reader.next().await.unwrap().unwrap(); - - assert_eq!(first.payload, bytes::Bytes::from_static(b"a-small")); - assert_eq!(second.payload, bytes::Bytes::from(vec![b'A'; 40])); - assert_eq!(third.payload, bytes::Bytes::from_static(b"b-ready")); - writer_task.abort(); - } - - #[tokio::test(start_paused = true)] - async fn one_ms_batch_interval_flushes_on_expiry() { - let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); - let address = listener.local_addr().unwrap(); - let client = TcpStream::connect(address).await.unwrap(); - let (server, _) = listener.accept().await.unwrap(); - let (_, write_half) = client.into_split(); - let (server_read, _) = server.into_split(); - - let stream_id = Uuid::new_v4(); - let state = worker_stream_state(TEST_STREAM_WINDOW, RESPONSE_MUX_STREAM_WRITER_QUEUE); - let (priority_tx, priority_rx) = mpsc::channel(1); - let (stream_tx, stream_rx) = mpsc::channel(8); - let cancel = CancellationToken::new(); - let connection = Arc::new(MuxConnection { - id: 1, - cancel: cancel.clone(), - priority_tx, - stream_tx, - streams: DashMap::new(), - healthy: AtomicBool::new(true), - active_streams: AtomicUsize::new(1), - queued_frames: AtomicUsize::new(0), - queued_bytes: AtomicUsize::new(0), - max_queued_bytes: TEST_CONNECTION_WINDOW, - queued_byte_slots: Arc::new(Semaphore::new(TEST_CONNECTION_WINDOW)), - connection_credits: Arc::new(Semaphore::new(TEST_CONNECTION_WINDOW)), - max_connection_credits: TEST_CONNECTION_WINDOW, - sent_data_bytes: AtomicU64::new(0), - acknowledged_data_bytes: AtomicU64::new(0), - batch_interval: Duration::from_millis(1), - batch_max_bytes: 65_536, - batch_max_frames: 64, - }); - connection.streams.insert(stream_id, state.clone()); - connection - .enqueue_stream_command( - stream_id, - state, - WriterCommand::new( - MuxFrame::new( - MuxFrameKind::Data, - stream_id, - bytes::Bytes::from_static(b"batched"), - ), - None, - ), - ) - .await - .unwrap(); - - let writer_task = tokio::spawn(MuxConnection::writer_task( - Arc::downgrade(&connection), - write_half, - priority_rx, - stream_rx, - cancel, - )); - let (received_tx, mut received_rx) = oneshot::channel(); - tokio::spawn(async move { - let mut reader = FramedRead::new(server_read, MuxCodec::default()); - let _ = received_tx.send(reader.next().await.unwrap().unwrap()); - }); - - for _ in 0..4 { - tokio::task::yield_now().await; - } - assert!(matches!( - received_rx.try_recv(), - Err(oneshot::error::TryRecvError::Empty) - )); - tokio::time::advance(Duration::from_micros(999)).await; - tokio::task::yield_now().await; - assert!(matches!( - received_rx.try_recv(), - Err(oneshot::error::TryRecvError::Empty) - )); - tokio::time::advance(Duration::from_micros(1)).await; - assert_eq!( - received_rx.await.unwrap().payload, - bytes::Bytes::from_static(b"batched") - ); - writer_task.abort(); - } - - #[test] - fn idle_cleanup_cannot_retire_a_host_with_an_opener() { - let mut config = PoolConfig::from_runtime(integration_config()); - config.idle_ttl = Duration::ZERO; - let host = HostPool::new( - "127.0.0.1:1".to_string(), - Uuid::new_v4(), - RESPONSE_MUX_VERSION, - CancellationToken::new(), - config, - ); - let opener = { - let mut lifecycle = host.lifecycle.lock(); - lifecycle.openers += 1; - HostOpenGuard { host: host.clone() } - }; - - assert!(!host.try_retire()); - drop(opener); - assert!(host.try_retire()); - } - fn integration_config() -> ResponseMuxConfig { ResponseMuxConfig { packet_metrics: false, @@ -2338,7 +1093,6 @@ mod tests { batch_max_bytes: 65_536, batch_max_frames: 64, stream_window_bytes: 262_144, - connection_window_bytes: 262_144, } } @@ -2400,9 +1154,8 @@ mod tests { async fn legacy_response_connection_info_is_rejected() { use crate::pipeline::network::tcp::{StreamType, TcpStreamConnectionInfo}; - let config = integration_config(); let cancel = CancellationToken::new(); - let pool = ResponseMuxClientPool::new(cancel.clone(), config); + let pool = ResponseMuxClientPool::new(cancel.clone(), integration_config()); let context = Context::new(()); let info = TcpStreamConnectionInfo { address: "127.0.0.1:1".to_string(), @@ -2410,7 +1163,6 @@ mod tests { subject: Uuid::new_v4().to_string(), stream_type: StreamType::Response, }; - let error = pool .create_response_stream(context.context(), info.into(), None) .await @@ -2425,15 +1177,8 @@ mod tests { } #[tokio::test(flavor = "multi_thread", worker_threads = 2)] - async fn stream_credit_stall_does_not_block_another_stream() { - let config = ResponseMuxConfig { - packet_metrics: false, - batch_interval: Duration::ZERO, - batch_max_bytes: 65_536, - batch_max_frames: 64, - stream_window_bytes: 64, - connection_window_bytes: 512, - }; + async fn duplicate_prologue_is_rejected_and_data_order_is_preserved() { + let config = integration_config(); let server = crate::pipeline::network::tcp::server::TcpStreamServer::new_mux_for_test( crate::pipeline::network::tcp::server::ServerOptions::default(), config, @@ -2441,83 +1186,78 @@ mod tests { .await .unwrap(); let cancel = CancellationToken::new(); - let pool = ResponseMuxClientPool::new_for_test_with_connection_window( - cancel.clone(), - 1, - 64, - 64, - 512, - Duration::from_secs(60), - Duration::from_secs(5), - ); - let (_, _context_a, sender_a, mut receiver_a) = - open_mux_stream(server.clone(), pool.clone()).await; - let (_, _context_b, sender_b, mut receiver_b) = open_mux_stream(server, pool.clone()).await; - - let full_window_payload = bytes::Bytes::from(vec![b'a'; 40]); - sender_a.send(full_window_payload.clone()).await.unwrap(); - let blocked_send = sender_a.send(full_window_payload.clone()); - tokio::pin!(blocked_send); - assert!( - tokio::time::timeout(Duration::from_millis(20), &mut blocked_send) + let pool = ResponseMuxClientPool::new(cancel.clone(), config); + let (_, _, mut sender, mut receiver) = open_mux_stream(server, pool).await; + assert!(sender.send_prologue(None).await.is_err()); + for value in 0_u16..128 { + sender + .send(bytes::Bytes::copy_from_slice(&value.to_be_bytes())) .await - .is_err(), - "second frame should wait for stream-local credits" - ); - - sender_b - .send(bytes::Bytes::from_static(b"healthy")) - .await - .unwrap(); - sender_b.finish().await.unwrap(); - assert_eq!( - receiver_b.recv().await.unwrap(), - bytes::Bytes::from_static(b"healthy") - ); - assert!(receiver_b.recv().await.is_none()); - - assert_eq!(receiver_a.recv().await.unwrap(), full_window_payload); - tokio::time::timeout(Duration::from_secs(1), &mut blocked_send) - .await - .expect("stream credit update did not unblock producer") - .unwrap(); - sender_a.finish().await.unwrap(); - assert_eq!(receiver_a.recv().await.unwrap(), full_window_payload); - assert!(receiver_a.recv().await.is_none()); + .unwrap(); + } + sender.finish().await.unwrap(); + for expected in 0_u16..128 { + let actual = receiver.recv().await.unwrap(); + assert_eq!(u16::from_be_bytes(actual[..].try_into().unwrap()), expected); + } + assert!(receiver.recv().await.is_none()); cancel.cancel(); } #[tokio::test(flavor = "multi_thread", worker_threads = 2)] - async fn frontend_receiver_drop_removes_the_worker_stream() { - let config = integration_config(); + async fn blocked_stream_does_not_block_another_on_the_same_connection() { + let config = ResponseMuxConfig { + stream_window_bytes: 64, + batch_interval: Duration::ZERO, + ..integration_config() + }; let server = crate::pipeline::network::tcp::server::TcpStreamServer::new_mux_for_test( crate::pipeline::network::tcp::server::ServerOptions::default(), config, ) .await .unwrap(); - let address = mux_address(server.clone()).await; let cancel = CancellationToken::new(); - let pool = ResponseMuxClientPool::new(cancel.clone(), config); - let (stream_id, context, sender, receiver) = open_mux_stream(server, pool.clone()).await; + let pool = ResponseMuxClientPool::new_for_test(cancel.clone(), config, 1, 512); + let (_, _, sender_a, mut receiver_a) = open_mux_stream(server.clone(), pool.clone()).await; + let (_, _, sender_b, mut receiver_b) = open_mux_stream(server, pool.clone()).await; - drop(receiver); - tokio::time::timeout(Duration::from_secs(1), async { - while !context.context().is_killed() - || pool.stream_connection_id(&address, stream_id).is_some() - { - tokio::task::yield_now().await; - } - }) - .await - .expect("frontend receiver drop did not remove the worker stream"); + let full_window_payload = bytes::Bytes::from(vec![b'a'; 43]); + sender_a.send(full_window_payload.clone()).await.unwrap(); + { + let blocked_send = sender_a.send(full_window_payload.clone()); + tokio::pin!(blocked_send); + assert!( + tokio::time::timeout(Duration::from_millis(20), &mut blocked_send) + .await + .is_err() + ); - assert!(sender.finish().await.is_err()); + sender_b + .send(bytes::Bytes::from_static(b"healthy")) + .await + .unwrap(); + sender_b.finish().await.unwrap(); + assert_eq!( + receiver_b.recv().await.unwrap(), + bytes::Bytes::from_static(b"healthy") + ); + assert!(receiver_b.recv().await.is_none()); + + assert_eq!(receiver_a.recv().await.unwrap(), full_window_payload); + tokio::time::timeout(Duration::from_secs(1), &mut blocked_send) + .await + .expect("stream credit update did not unblock producer") + .unwrap(); + } + sender_a.finish().await.unwrap(); + assert_eq!(receiver_a.recv().await.unwrap(), full_window_payload); + assert!(receiver_a.recv().await.is_none()); cancel.cancel(); } #[tokio::test(flavor = "multi_thread", worker_threads = 4)] - async fn connection_failure_is_scoped_and_pool_reconnects() { + async fn connection_failure_is_scoped_and_pool_repairs_to_four() { let config = integration_config(); let server = crate::pipeline::network::tcp::server::TcpStreamServer::new_mux_for_test( crate::pipeline::network::tcp::server::ServerOptions::default(), @@ -2528,7 +1268,6 @@ mod tests { let address = mux_address(server.clone()).await; let cancel = CancellationToken::new(); let pool = ResponseMuxClientPool::new(cancel.clone(), config); - let mut streams = vec![open_mux_stream(server.clone(), pool.clone()).await]; tokio::time::timeout(Duration::from_secs(5), async { while pool.healthy_connection_count(&address) != RESPONSE_MUX_POOL_SIZE { @@ -2541,60 +1280,41 @@ mod tests { streams.push(open_mux_stream(server.clone(), pool.clone()).await); } - let target_connection = pool - .stream_connection_id(&address, streams[0].0) - .expect("stream must be assigned to a connection"); + let target = pool.stream_connection_id(&address, streams[0].0).unwrap(); let assignments = streams .iter() - .map(|stream| { - pool.stream_connection_id(&address, stream.0) - .expect("stream must be assigned to a connection") - }) + .map(|stream| pool.stream_connection_id(&address, stream.0).unwrap()) .collect::>(); - assert!(assignments.iter().any(|id| *id != target_connection)); - - pool.fail_connection(&address, target_connection); + assert!(assignments.iter().any(|connection| *connection != target)); + pool.fail_connection(&address, target); tokio::time::timeout(Duration::from_secs(1), async { loop { - let correctly_scoped = - streams - .iter() - .zip(&assignments) - .all(|((_, context, _, _), connection_id)| { - context.context().is_killed() == (*connection_id == target_connection) - }); - if correctly_scoped { + if streams + .iter() + .zip(&assignments) + .all(|((_, context, _, _), connection)| { + context.context().is_killed() == (*connection == target) + }) + { break; } tokio::task::yield_now().await; } }) .await - .expect("connection failure did not remain scoped to assigned streams"); - - let (_, _, replacement_sender, mut replacement_receiver) = - open_mux_stream(server, pool.clone()).await; - replacement_sender - .send(bytes::Bytes::from_static(b"replacement")) - .await - .unwrap(); - replacement_sender.finish().await.unwrap(); - assert_eq!( - replacement_receiver.recv().await.unwrap(), - bytes::Bytes::from_static(b"replacement") - ); + .expect("connection failure escaped its assigned streams"); tokio::time::timeout(Duration::from_secs(5), async { while pool.healthy_connection_count(&address) != RESPONSE_MUX_POOL_SIZE { tokio::time::sleep(Duration::from_millis(10)).await; } }) .await - .expect("response mux pool did not reconnect to four connections"); + .expect("response mux pool did not repair to four connections"); cancel.cancel(); } #[tokio::test(flavor = "multi_thread", worker_threads = 4)] - async fn one_thousand_logical_streams_share_exactly_four_connections() { + async fn one_thousand_streams_share_exactly_four_connections() { let config = integration_config(); let server = crate::pipeline::network::tcp::server::TcpStreamServer::new_mux_for_test( crate::pipeline::network::tcp::server::ServerOptions::default(), @@ -2609,40 +1329,16 @@ mod tests { server: Arc, pool: Arc, value: usize, - ) -> String { - let context = Context::new(()); - let pending = server - .register( - StreamOptions::builder() - .context(context.context()) - .enable_request_stream(false) - .enable_response_stream(true) - .send_buffer_count(8) - .build() - .unwrap(), - ) - .await - .recv_stream - .unwrap(); - let (info, provider) = pending.into_parts(); - let mut sender = pool - .create_response_stream(context.context(), info, None) - .await - .unwrap(); - sender.send_prologue(None).await.unwrap(); - let mut receiver = provider.await.unwrap().unwrap(); - let expected = format!("response-{value}"); - sender.send(expected.clone().into()).await.unwrap(); + ) -> usize { + let (_, _, sender, mut receiver) = open_mux_stream(server, pool).await; + sender.send(value.to_string().into()).await.unwrap(); sender.finish().await.unwrap(); let actual = receiver.recv().await.unwrap(); assert!(receiver.recv().await.is_none()); - String::from_utf8(actual.to_vec()).unwrap() + std::str::from_utf8(&actual).unwrap().parse().unwrap() } - assert_eq!( - round_trip(server.clone(), pool.clone(), 0).await, - "response-0" - ); + assert_eq!(round_trip(server.clone(), pool.clone(), 0).await, 0); let address = mux_address(server.clone()).await; tokio::time::timeout(Duration::from_secs(5), async { while pool.healthy_connection_count(&address) != RESPONSE_MUX_POOL_SIZE { @@ -2657,8 +1353,8 @@ mod tests { tasks.push(round_trip(server.clone(), pool.clone(), value)); } let mut completed = 1; - while let Some(actual) = tasks.next().await { - assert!(actual.starts_with("response-")); + while let Some(value) = tasks.next().await { + assert!(value < 1_000); completed += 1; } assert_eq!(completed, 1_000); diff --git a/lib/runtime/src/pipeline/network/tcp/server.rs b/lib/runtime/src/pipeline/network/tcp/server.rs index 6331d660374d..b03255a4f82c 100644 --- a/lib/runtime/src/pipeline/network/tcp/server.rs +++ b/lib/runtime/src/pipeline/network/tcp/server.rs @@ -40,9 +40,8 @@ use super::{ ResponseMuxConnectionInfo, StreamOptions, StreamReceiver, StreamSender, TcpStreamConnectionInfo, TwoPartCodec, mux::{ - ConnectionHandshake, MuxCodec, MuxFrame, MuxFrameKind, RESPONSE_MUX_CREDIT_UPDATE_BYTES, - RESPONSE_MUX_CREDIT_UPDATE_INTERVAL, RESPONSE_MUX_VERSION, RESPONSE_MUX_WRITER_QUEUE, - ResponseMuxConfig, initialize_response_mux_config, + MuxCodec, MuxFrame, MuxFrameKind, RESPONSE_MUX_CREDIT_UPDATE_BYTES, RESPONSE_MUX_VERSION, + RESPONSE_MUX_WRITER_QUEUE, ResponseMuxConfig, initialize_response_mux_config, }, }; use crate::discovery::EndpointInstanceId; @@ -50,9 +49,7 @@ use crate::engine::AsyncEngineContext; use crate::pipeline::{ PipelineError, network::{ - ResponseService, ResponseStreamPrologue, StreamReceiverHooks, StreamRxItem, - codec::{TwoPartMessage, TwoPartMessageType}, - tcp::StreamType, + ResponseService, StreamReceiverHooks, StreamRxItem, codec::TwoPartMessage, tcp::StreamType, }, }; use anyhow::{Context, Result, anyhow as error}; @@ -102,8 +99,7 @@ pub struct TcpStreamServer { server_id: uuid::Uuid, mux_config: ResponseMuxConfig, state: Arc>, - response_pending: Arc>, - response_active: Arc>, + response_directory: ResponseDirectory, } // pub struct TcpStreamReceiver { @@ -121,29 +117,49 @@ struct RequestedSendConnection { send_buffer_count: usize, } -struct RequestedMuxRecvConnection { +struct PendingMuxResponse { context: Arc, - connection: Mutex>>>, + connection: oneshot::Sender>, send_buffer_count: usize, registered_at: Instant, } -struct ActiveMuxResponseControl { - connection_id: uuid::Uuid, +#[derive(Clone)] +struct ActiveMuxResponseRoute { context: Arc, - control_tx: mpsc::Sender, - close_tx: mpsc::Sender, + command_tx: mpsc::Sender, control_failed: CancellationToken, } +enum ResponseMuxEntry { + Pending(PendingMuxResponse), + Active(ActiveMuxResponseRoute), +} + +type ResponseDirectory = Arc>; + +enum ResponseMuxCommand { + WindowUpdate { + stream_id: uuid::Uuid, + credits: usize, + }, + Stop { + stream_id: uuid::Uuid, + }, + Close { + stream_id: uuid::Uuid, + kind: MuxFrameKind, + }, +} + struct ActiveMuxResponseStream { context: Arc, response_tx: mpsc::Sender, } struct ResponseMuxSocket { - reader: FramedRead, MuxCodec>, - write_half: tokio::io::WriteHalf, + reader: FramedRead, + write_half: tokio::net::tcp::OwnedWriteHalf, packet_socket: Option, } @@ -273,8 +289,7 @@ impl TcpStreamServer { }; let state = Arc::new(Mutex::new(State::default())); - let response_pending = Arc::new(DashMap::new()); - let response_active = Arc::new(DashMap::new()); + let response_directory = Arc::new(DashMap::new()); let server_id = uuid::Uuid::new_v4(); let local_port = Self::start( @@ -283,8 +298,7 @@ impl TcpStreamServer { state.clone(), server_id, mux_config, - response_pending.clone(), - response_active.clone(), + response_directory.clone(), ) .await .map_err(|e| PipelineError::Generic(format!("Failed to start TcpStreamServer: {}", e)))?; @@ -297,8 +311,7 @@ impl TcpStreamServer { server_id, mux_config, state, - response_pending, - response_active, + response_directory, })) } @@ -437,14 +450,14 @@ impl TcpStreamServer { } fn cancel_mux_response_stream(&self, stream_id: uuid::Uuid, kind: MuxFrameKind) { - self.response_pending.remove(&stream_id); - if let Some((_, active)) = self.response_active.remove(&stream_id) { + if let Some((_, ResponseMuxEntry::Active(active))) = + self.response_directory.remove(&stream_id) + { active.context.kill(); if active - .control_tx - .try_send(MuxFrame::empty(kind, stream_id)) + .command_tx + .try_send(ResponseMuxCommand::Close { stream_id, kind }) .is_err() - || active.close_tx.try_send(stream_id).is_err() { active.control_failed.cancel(); } @@ -457,8 +470,7 @@ impl TcpStreamServer { state: Arc>, server_id: uuid::Uuid, mux_config: ResponseMuxConfig, - response_pending: Arc>, - response_active: Arc>, + response_directory: ResponseDirectory, ) -> Result { let addr = format!("{}:{}", local_ip, local_port); let state_clone = state.clone(); @@ -473,8 +485,7 @@ impl TcpStreamServer { state_clone, server_id, mux_config, - response_pending, - response_active, + response_directory, ready_tx, ))); } @@ -555,19 +566,16 @@ impl ResponseService for TcpStreamServer { pending_sender_rx, ) .with_cleanup(move || { - // Drop is sync; fire-and-forget the lock acquisition. - tokio::spawn(async move { - let mut state = cleanup_state.lock(); - state.tx_subjects.remove(&cleanup_subject); - if let Some(key) = state.subject_instance.remove(&cleanup_subject) - && let Some(subjects) = state.instance_subjects.get_mut(&key) - { - subjects.remove(&(StreamType::Request, cleanup_subject.clone())); - if subjects.is_empty() { - state.instance_subjects.remove(&key); - } + let mut state = cleanup_state.lock(); + state.tx_subjects.remove(&cleanup_subject); + if let Some(key) = state.subject_instance.remove(&cleanup_subject) + && let Some(subjects) = state.instance_subjects.get_mut(&key) + { + subjects.remove(&(StreamType::Request, cleanup_subject.clone())); + if subjects.is_empty() { + state.instance_subjects.remove(&key); } - }); + } }); self.insert_request_stream(registry_subject, connection_info); @@ -581,41 +589,42 @@ impl ResponseService for TcpStreamServer { let (pending_recver_tx, pending_recver_rx) = oneshot::channel(); let receiver_id = uuid::Uuid::new_v4(); let receiver_subject = receiver_id.to_string(); - self.response_pending.insert( + self.response_directory.insert( receiver_id, - RequestedMuxRecvConnection { + ResponseMuxEntry::Pending(PendingMuxResponse { context: options.context.clone(), - connection: Mutex::new(Some(pending_recver_tx)), + connection: pending_recver_tx, send_buffer_count: options.send_buffer_count, registered_at: Instant::now(), - }, + }), ); let cleanup_id = receiver_id; let cleanup_subject = receiver_subject; let cleanup_state = self.state.clone(); - let cleanup_pending = self.response_pending.clone(); - let cleanup_active = self.response_active.clone(); + let cleanup_directory = self.response_directory.clone(); let registered_stream = RegisteredStream::new( ResponseMuxConnectionInfo { address: address.clone(), frontend_server_id: self.server_id, stream_id: receiver_id, context: options.context.id().to_string(), - version: RESPONSE_MUX_VERSION, } .into(), pending_recver_rx, ) .with_cleanup(move || { - cleanup_pending.remove(&cleanup_id); - if let Some((_, active)) = cleanup_active.remove(&cleanup_id) { + if let Some((_, ResponseMuxEntry::Active(active))) = + cleanup_directory.remove(&cleanup_id) + { active.context.kill(); if active - .control_tx - .try_send(MuxFrame::empty(MuxFrameKind::Kill, cleanup_id)) + .command_tx + .try_send(ResponseMuxCommand::Close { + stream_id: cleanup_id, + kind: MuxFrameKind::Kill, + }) .is_err() - || active.close_tx.try_send(cleanup_id).is_err() { active.control_failed.cancel(); } @@ -653,8 +662,7 @@ async fn tcp_listener( state: Arc>, server_id: uuid::Uuid, mux_config: ResponseMuxConfig, - response_pending: Arc>, - response_active: Arc>, + response_directory: ResponseDirectory, read_tx: tokio::sync::oneshot::Sender>, ) -> Result<()> { let listener = tokio::net::TcpListener::bind(&addr) @@ -715,8 +723,7 @@ async fn tcp_listener( state.clone(), server_id, mux_config, - response_pending.clone(), - response_active.clone(), + response_directory.clone(), )); } @@ -727,18 +734,10 @@ async fn tcp_listener( state: Arc>, server_id: uuid::Uuid, mux_config: ResponseMuxConfig, - response_pending: Arc>, - response_active: Arc>, + response_directory: ResponseDirectory, ) { - let result = process_stream( - stream, - state, - server_id, - mux_config, - response_pending, - response_active, - ) - .await; + let result = + process_connection(stream, state, server_id, mux_config, response_directory).await; match result { Ok(_) => tracing::trace!("successfully processed tcp connection"), Err(e) => { @@ -749,49 +748,62 @@ async fn tcp_listener( } } - /// This method is responsible for the internal tcp stream handshake - /// The handshake will specialize the stream as a request/sender or response/receiver stream - async fn process_stream( + fn remove_active_route( + directory: &ResponseDirectory, + command_tx: &mpsc::Sender, + stream_id: uuid::Uuid, + ) { + directory.remove_if(&stream_id, |_, entry| { + matches!( + entry, + ResponseMuxEntry::Active(route) + if route.command_tx.same_channel(command_tx) + ) + }); + } + + fn try_send_control( + control_tx: &mpsc::Sender, + control_failed: &CancellationToken, + frame: MuxFrame, + ) { + if control_tx.try_send(frame).is_err() { + control_failed.cancel(); + } + } + + async fn process_connection( stream: tokio::net::TcpStream, state: Arc>, server_id: uuid::Uuid, mux_config: ResponseMuxConfig, - response_pending: Arc>, - response_active: Arc>, + response_directory: ResponseDirectory, ) -> Result<()> { - let packet_socket = mux_config - .packet_metrics - .then(|| stream.as_fd().try_clone_to_owned().ok()) - .flatten(); - // split the socket in to a reader and writer - let (read_half, write_half) = tokio::io::split(stream); - - // attach the codec to the reader and writer to get framed readers and writers - let mut framed_reader = FramedRead::new(read_half, TwoPartCodec::default()); - let mut framed_writer = FramedWrite::new(write_half, TwoPartCodec::default()); - - // the internal tcp [`CallHomeHandshake`] connects the socket to the requester - // here we await this first message as a raw bytes two part message - let first_message = framed_reader - .next() - .await - .ok_or(error!("Connection closed without a ControlMessage"))??; - - // we await on the raw bytes which should come in as a header only message - // todo - improve error handling - check for no data - let header = match first_message.header() { - Some(header) => header, - None => { - return Err(error!("Expected ControlMessage, got DataMessage")); + let mut prefix = [0_u8; 5]; + loop { + let read = stream.peek(&mut prefix).await?; + if read == 0 { + anyhow::bail!("connection closed before its handshake"); } - }; + if read == prefix.len() { + break; + } + tokio::task::yield_now().await; + } - if let Ok(ConnectionHandshake::ResponseMux { - version, - frontend_server_id, - connection_id, - }) = serde_json::from_slice::(header) - { + if prefix[4] == MuxFrameKind::ConnectionHello as u8 { + let packet_socket = mux_config + .packet_metrics + .then(|| stream.as_fd().try_clone_to_owned().ok()) + .flatten(); + let (read_half, write_half) = stream.into_split(); + let mut reader = FramedRead::new(read_half, MuxCodec::default()); + let mut writer = FramedWrite::new(write_half, MuxCodec::default()); + let hello = reader + .next() + .await + .ok_or_else(|| error!("connection closed before response mux hello"))??; + let (version, frontend_server_id) = hello.connection_identity()?; if version != RESPONSE_MUX_VERSION { anyhow::bail!( "unsupported response mux version {version}; expected {RESPONSE_MUX_VERSION}" @@ -802,33 +814,35 @@ async fn tcp_listener( "response mux frontend UUID mismatch: got {frontend_server_id}, expected {server_id}" ); } - if connection_id.is_nil() { - anyhow::bail!("response mux physical connection UUID must not be nil"); - } - framed_writer - .send(MuxFrame::connection_ack(0).into_two_part()) + writer + .send(MuxFrame::connection_ready()) .await - .context("failed to send response mux connection ack")?; - return process_response_mux( - connection_id, + .context("failed to send response mux connection ready")?; + return process_response_mux_connection( mux_config, state, - response_pending, - response_active, + response_directory, ResponseMuxSocket { - reader: framed_reader.map_decoder(|_| MuxCodec::default()), - write_half: framed_writer.into_inner(), + reader, + write_half: writer.into_inner(), packet_socket, }, ) .await; } - let handshake: CallHomeHandshake = serde_json::from_slice(header).map_err(|e| { - error!("Failed to deserialize the first message as a valid TCP handshake: {e}") - })?; - - // branch here to handle sender stream or receiver stream + let (read_half, write_half) = tokio::io::split(stream); + let mut framed_reader = FramedRead::new(read_half, TwoPartCodec::default()); + let framed_writer = FramedWrite::new(write_half, TwoPartCodec::default()); + let first_message = framed_reader + .next() + .await + .ok_or_else(|| error!("connection closed without a request-stream handshake"))??; + let header = first_message + .header() + .ok_or_else(|| error!("expected request-stream handshake, got data"))?; + let handshake: CallHomeHandshake = serde_json::from_slice(header) + .map_err(|err| error!("failed to deserialize TCP request-stream handshake: {err}"))?; match handshake.stream_type { StreamType::Request => { process_request_stream(handshake.subject, state, framed_reader, framed_writer).await @@ -840,12 +854,10 @@ async fn tcp_listener( } } - async fn process_response_mux( - connection_id: uuid::Uuid, + async fn process_response_mux_connection( mux_config: ResponseMuxConfig, state: Arc>, - response_pending: Arc>, - response_active: Arc>, + response_directory: ResponseDirectory, socket: ResponseMuxSocket, ) -> Result<()> { let ResponseMuxSocket { @@ -855,7 +867,8 @@ async fn tcp_listener( } = socket; let mut writer = FramedWrite::new(write_half, MuxCodec::default()); let (control_tx, mut control_rx) = mpsc::channel::(RESPONSE_MUX_WRITER_QUEUE); - let (close_tx, mut close_rx) = mpsc::channel::(RESPONSE_MUX_WRITER_QUEUE); + let (command_tx, mut command_rx) = + mpsc::channel::(RESPONSE_MUX_WRITER_QUEUE); let control_failed = CancellationToken::new(); let mut reported_data_segments = packet_socket .as_ref() @@ -869,17 +882,11 @@ async fn tcp_listener( .inc(); let writer_failed = control_failed.clone(); - let detailed_metrics = mux_config.packet_metrics; let writer_task = tokio::spawn(async move { - let frame_counters = - crate::metrics::response_mux::FrameCounters::for_direction("frontend_to_worker"); let write_calls = crate::metrics::response_mux::WRITE_CALLS_TOTAL .with_label_values(&["frontend"]) .clone(); while let Some(frame) = control_rx.recv().await { - if detailed_metrics { - frame_counters.inc(frame.kind.metric_label()); - } if let Err(err) = writer.send(frame).await { writer_failed.cancel(); return Err(err.into()); @@ -889,36 +896,63 @@ async fn tcp_listener( Result::<()>::Ok(()) }); - let mut decoded_data_bytes = 0_u64; - let mut acknowledged_data_bytes = 0_u64; - let mut credit_tick = tokio::time::interval(RESPONSE_MUX_CREDIT_UPDATE_INTERVAL); - credit_tick.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay); - credit_tick.tick().await; let mut packet_tick = tokio::time::interval(Duration::from_millis(100)); packet_tick.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay); packet_tick.tick().await; - let frame_counters = - crate::metrics::response_mux::FrameCounters::for_direction("worker_to_frontend"); - // Data routing is connection-local so the per-frame path does not take - // a shard lock in the global lifecycle registry. let mut active_streams = HashMap::::new(); let result: Result<()> = async { loop { - let message = tokio::select! { + let frame = tokio::select! { _ = control_failed.cancelled() => { anyhow::bail!("frontend response mux control writer failed") } - Some(stream_id) = close_rx.recv() => { - active_streams.remove(&stream_id); - response_pending.remove(&stream_id); - if response_active - .get(&stream_id) - .is_some_and(|active| active.connection_id == connection_id) - { - response_active.remove(&stream_id); + command = command_rx.recv() => { + let Some(command) = command else { + anyhow::bail!("frontend response mux command channel closed") + }; + match command { + ResponseMuxCommand::WindowUpdate { stream_id, mut credits } => { + if !active_streams.contains_key(&stream_id) { + continue; + } + while credits > 0 { + let update = credits.min( + mux_config.stream_window_bytes.min(u32::MAX as usize), + ); + try_send_control( + &control_tx, + &control_failed, + MuxFrame::window_update(stream_id, update as u32), + ); + credits -= update; + } + } + ResponseMuxCommand::Stop { stream_id } => { + if active_streams.contains_key(&stream_id) { + try_send_control( + &control_tx, + &control_failed, + MuxFrame::empty(MuxFrameKind::Stop, stream_id), + ); + } + } + ResponseMuxCommand::Close { stream_id, kind } => { + if active_streams.remove(&stream_id).is_some() { + remove_active_route( + &response_directory, + &command_tx, + stream_id, + ); + remove_response_association(&state, stream_id); + try_send_control( + &control_tx, + &control_failed, + MuxFrame::empty(kind, stream_id), + ); + } + } } - remove_response_association(&state, stream_id); continue; } _ = packet_tick.tick(), if reported_data_segments.is_some() => { @@ -933,191 +967,161 @@ async fn tcp_listener( } continue; } - _ = credit_tick.tick(), if decoded_data_bytes > acknowledged_data_bytes => { - acknowledged_data_bytes = decoded_data_bytes; - if control_tx - .try_send(MuxFrame::connection_ack(acknowledged_data_bytes)) - .is_err() - { - control_failed.cancel(); - } - continue; - } message = reader.next() => match message { Some(message) => message?, None => anyhow::bail!("worker closed response mux connection"), }, }; - let frame = message; - if mux_config.packet_metrics { - frame_counters.inc(frame.kind.metric_label()); - } let stream_id = frame.stream_id; match frame.kind { - MuxFrameKind::Prologue => { - let Some((_, pending)) = response_pending.remove(&stream_id) else { - crate::metrics::response_mux::RESETS_TOTAL - .with_label_values(&["frontend", "unknown_stream"]) - .inc(); - if control_tx - .try_send(MuxFrame::empty(MuxFrameKind::Reset, stream_id)) - .is_err() - { - control_failed.cancel(); - } - continue; - }; - let prologue: ResponseStreamPrologue = - match serde_json::from_slice(&frame.payload) { - Ok(prologue) => prologue, - Err(err) => { - let reason = format!( - "invalid response mux prologue for {stream_id}: {err}" + MuxFrameKind::Prologue if frame.payload.is_empty() => { + let pending = { + let Some(mut entry) = response_directory.get_mut(&stream_id) else { + crate::metrics::response_mux::RESETS_TOTAL + .with_label_values(&["frontend", "unknown_stream"]) + .inc(); + try_send_control( + &control_tx, + &control_failed, + MuxFrame::empty(MuxFrameKind::Reset, stream_id), + ); + continue; + }; + let route = match entry.value() { + ResponseMuxEntry::Pending(pending) => ActiveMuxResponseRoute { + context: pending.context.clone(), + command_tx: command_tx.clone(), + control_failed: control_failed.clone(), + }, + ResponseMuxEntry::Active(_) => { + try_send_control( + &control_tx, + &control_failed, + MuxFrame::empty(MuxFrameKind::Reset, stream_id), ); - if let Some(connection) = pending.connection.lock().take() { - let _ = connection.send(Err(reason)); - } - crate::metrics::response_mux::RESETS_TOTAL - .with_label_values(&["frontend", "invalid_prologue"]) - .inc(); - if control_tx - .try_send(MuxFrame::empty(MuxFrameKind::Reset, stream_id)) - .is_err() - { - control_failed.cancel(); - } - remove_response_association(&state, stream_id); continue; } }; - let Some(connection) = pending.connection.lock().take() else { - if control_tx - .try_send(MuxFrame::empty(MuxFrameKind::Reset, stream_id)) - .is_err() - { - control_failed.cancel(); + match std::mem::replace( + entry.value_mut(), + ResponseMuxEntry::Active(route), + ) { + ResponseMuxEntry::Pending(pending) => pending, + ResponseMuxEntry::Active(active) => { + *entry = ResponseMuxEntry::Active(active); + try_send_control( + &control_tx, + &control_failed, + MuxFrame::empty(MuxFrameKind::Reset, stream_id), + ); + continue; + } } - continue; }; crate::metrics::response_mux::SETUP_SECONDS .observe(pending.registered_at.elapsed().as_secs_f64()); - if let Some(error) = prologue.error { - let _ = connection.send(Err(error)); - remove_response_association(&state, stream_id); - continue; - } - let mailbox_frames = pending.send_buffer_count.max( mux_config.stream_window_bytes / crate::pipeline::network::tcp::mux::MUX_HEADER_LEN, ); let (response_tx, response_rx) = data_plane_channel::(mailbox_frames); - let active_for_window = response_active.clone(); - let active_for_close = response_active.clone(); - let control_for_window = control_tx.clone(); - let control_for_close = control_tx.clone(); - let close_for_receiver = close_tx.clone(); - let failed_for_window = control_failed.clone(); + let update_tx = command_tx.clone(); + let close_tx = command_tx.clone(); + let failed_for_update = control_failed.clone(); let failed_for_close = control_failed.clone(); - let state_for_close = state.clone(); - let context = pending.context.clone(); let hooks = StreamReceiverHooks { - context: pending.context, + context: pending.context.clone(), window_update_threshold: RESPONSE_MUX_CREDIT_UPDATE_BYTES .min(mux_config.stream_window_bytes), - on_window_update: Arc::new(move |mut credits| { - if !active_for_window.contains_key(&stream_id) { - return; - } - while credits > 0 { - let update = credits - .min(mux_config.stream_window_bytes.min(u32::MAX as usize)); - if control_for_window - .try_send(MuxFrame::window_update(stream_id, update as u32)) - .is_err() - { - failed_for_window.cancel(); - return; - } - credits -= update; + on_window_update: Box::new(move |credits| { + if update_tx + .try_send(ResponseMuxCommand::WindowUpdate { + stream_id, + credits, + }) + .is_err() + { + failed_for_update.cancel(); } }), - on_close: Arc::new(move |control| match control { - ControlMessage::Stop => { - if active_for_close.contains_key(&stream_id) - && control_for_close - .try_send(MuxFrame::empty(MuxFrameKind::Stop, stream_id)) - .is_err() - { - failed_for_close.cancel(); - } - } - ControlMessage::Kill => { - if active_for_close.remove(&stream_id).is_some() { - if control_for_close - .try_send(MuxFrame::empty(MuxFrameKind::Kill, stream_id)) - .is_err() - || close_for_receiver.try_send(stream_id).is_err() - { - failed_for_close.cancel(); - } - remove_response_association(&state_for_close, stream_id); - } + on_close: Box::new(move |control| { + let command = match control { + ControlMessage::Stop => ResponseMuxCommand::Stop { stream_id }, + ControlMessage::Kill => ResponseMuxCommand::Close { + stream_id, + kind: MuxFrameKind::Kill, + }, + ControlMessage::Sentinel => return, + }; + if close_tx.try_send(command).is_err() { + failed_for_close.cancel(); } - ControlMessage::Sentinel => {} }), }; - response_active.insert( - stream_id, - ActiveMuxResponseControl { - connection_id, - context: context.clone(), - control_tx: control_tx.clone(), - close_tx: close_tx.clone(), - control_failed: control_failed.clone(), - }, - ); active_streams.insert( stream_id, ActiveMuxResponseStream { - context, + context: pending.context, response_tx, }, ); crate::metrics::response_mux::ACTIVE_STREAMS .with_label_values(&["frontend"]) .inc(); - if connection + if pending + .connection .send(Ok(StreamReceiver::multiplexed(response_rx, hooks))) .is_err() { - response_active.remove(&stream_id); active_streams.remove(&stream_id); - if control_tx - .try_send(MuxFrame::empty(MuxFrameKind::Kill, stream_id)) - .is_err() - { - control_failed.cancel(); - } + remove_active_route(&response_directory, &command_tx, stream_id); + try_send_control( + &control_tx, + &control_failed, + MuxFrame::empty(MuxFrameKind::Kill, stream_id), + ); remove_response_association(&state, stream_id); } } - MuxFrameKind::Data => { - let encoded_len = frame.encoded_len(); - decoded_data_bytes = decoded_data_bytes.saturating_add(encoded_len as u64); - if decoded_data_bytes.saturating_sub(acknowledged_data_bytes) - >= RESPONSE_MUX_CREDIT_UPDATE_BYTES as u64 - { - acknowledged_data_bytes = decoded_data_bytes; - if control_tx - .try_send(MuxFrame::connection_ack(acknowledged_data_bytes)) - .is_err() - { - control_failed.cancel(); + MuxFrameKind::Prologue => { + let error = std::str::from_utf8(&frame.payload) + .map(str::to_owned) + .map_err(|err| { + format!("invalid response mux prologue for {stream_id}: {err}") + }); + let pending = match response_directory.remove(&stream_id) { + Some((_, ResponseMuxEntry::Pending(pending))) => pending, + _ => { + try_send_control( + &control_tx, + &control_failed, + MuxFrame::empty(MuxFrameKind::Reset, stream_id), + ); + continue; + } + }; + match error { + Ok(error) => { + let _ = pending.connection.send(Err(error)); + } + Err(reason) => { + let _ = pending.connection.send(Err(reason)); + crate::metrics::response_mux::RESETS_TOTAL + .with_label_values(&["frontend", "invalid_prologue"]) + .inc(); + try_send_control( + &control_tx, + &control_failed, + MuxFrame::empty(MuxFrameKind::Reset, stream_id), + ); } } + remove_response_association(&state, stream_id); + } + MuxFrameKind::Data => { + let encoded_len = frame.encoded_len(); let delivery_failure = match active_streams.get(&stream_id) { Some(active) => match active .response_tx @@ -1129,53 +1133,48 @@ async fn tcp_listener( Some("receiver_closed") } }, - _ => Some("unknown_stream"), + None => Some("unknown_stream"), }; if let Some(reason) = delivery_failure { crate::metrics::response_mux::RESETS_TOTAL .with_label_values(&["frontend", reason]) .inc(); - response_active.remove(&stream_id); active_streams.remove(&stream_id); - response_pending.remove(&stream_id); - if control_tx - .try_send(MuxFrame::empty(MuxFrameKind::Reset, stream_id)) - .is_err() - { - control_failed.cancel(); - } + response_directory.remove(&stream_id); + try_send_control( + &control_tx, + &control_failed, + MuxFrame::empty(MuxFrameKind::Reset, stream_id), + ); remove_response_association(&state, stream_id); } } MuxFrameKind::End => { if active_streams.remove(&stream_id).is_some() { - response_active.remove(&stream_id); + remove_active_route(&response_directory, &command_tx, stream_id); remove_response_association(&state, stream_id); } else { - if control_tx - .try_send(MuxFrame::empty(MuxFrameKind::Reset, stream_id)) - .is_err() - { - control_failed.cancel(); - } + try_send_control( + &control_tx, + &control_failed, + MuxFrame::empty(MuxFrameKind::Reset, stream_id), + ); } } MuxFrameKind::Reset => { - response_pending.remove(&stream_id); + response_directory.remove(&stream_id); active_streams.remove(&stream_id); - response_active.remove(&stream_id); remove_response_association(&state, stream_id); } - MuxFrameKind::Stop - | MuxFrameKind::Kill - | MuxFrameKind::WindowUpdate - | MuxFrameKind::ConnectionAck => { - if control_tx - .try_send(MuxFrame::empty(MuxFrameKind::Reset, stream_id)) - .is_err() - { - control_failed.cancel(); - } + MuxFrameKind::ConnectionHello | MuxFrameKind::ConnectionReady => { + anyhow::bail!("unexpected response mux connection frame after handshake") + } + MuxFrameKind::Stop | MuxFrameKind::Kill | MuxFrameKind::WindowUpdate => { + try_send_control( + &control_tx, + &control_failed, + MuxFrame::empty(MuxFrameKind::Reset, stream_id), + ); } } } @@ -1191,19 +1190,13 @@ async fn tcp_listener( .with_label_values(&["mux", "frontend"]) .inc_by(current.saturating_sub(previous)); } - - let affected: Vec = active_streams.keys().copied().collect(); + let affected = active_streams.keys().copied().collect::>(); crate::metrics::response_mux::CONNECTION_LOST_STREAMS_TOTAL.inc_by(affected.len() as u64); for stream_id in affected { if let Some(active) = active_streams.remove(&stream_id) { active.context.kill(); } - if response_active - .get(&stream_id) - .is_some_and(|active| active.connection_id == connection_id) - { - response_active.remove(&stream_id); - } + remove_active_route(&response_directory, &command_tx, stream_id); remove_response_association(&state, stream_id); } crate::metrics::response_mux::ACTIVE_CONNECTIONS @@ -1215,6 +1208,8 @@ async fn tcp_listener( result } + /// This method is responsible for the internal tcp stream handshake + /// The handshake will specialize the stream as a request/sender or response/receiver stream fn remove_response_association(state: &Mutex, stream_id: uuid::Uuid) { let subject = stream_id.to_string(); let mut state = state.lock(); @@ -1268,9 +1263,6 @@ async fn tcp_listener( if connection .send(Ok(crate::pipeline::network::StreamSender::dedicated( request_tx, - // Request streams don't carry a downstream-prologue today; the - // upstream may begin sending immediately. - None, ))) .is_err() { @@ -1466,7 +1458,7 @@ mod tests { let state = server.state.lock(); assert_eq!(state.tx_subjects.len(), 1, "one request stream registered"); assert_eq!( - server.response_pending.len(), + server.response_directory.len(), 1, "one response stream registered" ); @@ -1475,10 +1467,10 @@ mod tests { "send_buffer_count must reach RequestedSendConnection" ); assert!( - server - .response_pending - .iter() - .all(|entry| entry.send_buffer_count == 7), + server.response_directory.iter().all(|entry| matches!( + entry.value(), + ResponseMuxEntry::Pending(pending) if pending.send_buffer_count == 7 + )), "send_buffer_count must reach the response mux registration" ); } @@ -1769,16 +1761,13 @@ mod tests { recv_stream.connection_info.clone().try_into().unwrap(); let stream_id = mux_info.stream_id; - assert!(server.response_pending.contains_key(&stream_id)); + assert!(server.response_directory.contains_key(&stream_id)); // Drop the RegisteredStream -- RAII cleanup should fire drop(recv_stream); - // Give the spawned cleanup task a moment to run - tokio::time::sleep(std::time::Duration::from_millis(50)).await; - assert!( - !server.response_pending.contains_key(&stream_id), + !server.response_directory.contains_key(&stream_id), "RAII cleanup should have removed the pending mux entry" ); } @@ -1805,11 +1794,8 @@ mod tests { // Call into_parts to disarm the cleanup let (_conn_info, _provider) = recv_stream.into_parts(); - // Give any potential cleanup a moment to run - tokio::time::sleep(std::time::Duration::from_millis(50)).await; - assert!( - server.response_pending.contains_key(&stream_id), + server.response_directory.contains_key(&stream_id), "into_parts() should disarm the RAII cleanup" ); } @@ -2116,38 +2102,27 @@ mod tests { assert_eq!(bad_info.frontend_server_id, good_info.frontend_server_id); let mut socket = TcpStream::connect(&bad_info.address).await.unwrap(); - let handshake = ConnectionHandshake::ResponseMux { - version: RESPONSE_MUX_VERSION, - frontend_server_id: bad_info.frontend_server_id, - connection_id: uuid::Uuid::new_v4(), - }; let mut wire = BytesMut::new(); - TwoPartCodec::default() + let mut mux_codec = MuxCodec::default(); + mux_codec .encode( - TwoPartMessage::from_header(serde_json::to_vec(&handshake).unwrap().into()), + MuxFrame::connection_hello(RESPONSE_MUX_VERSION, bad_info.frontend_server_id), &mut wire, ) .unwrap(); - let mut mux_codec = MuxCodec::default(); mux_codec .encode( MuxFrame::new( MuxFrameKind::Prologue, bad_info.stream_id, - Bytes::from_static(b"{"), + Bytes::from_static(&[0xff]), ), &mut wire, ) .unwrap(); mux_codec .encode( - MuxFrame::new( - MuxFrameKind::Prologue, - good_info.stream_id, - serde_json::to_vec(&ResponseStreamPrologue { error: None }) - .unwrap() - .into(), - ), + MuxFrame::new(MuxFrameKind::Prologue, good_info.stream_id, Bytes::new()), &mut wire, ) .unwrap(); @@ -2171,17 +2146,13 @@ mod tests { // A single write makes handshake and mux frames available to the // handshake decoder together, exercising decoder-buffer preservation. socket.write_all(&wire).await.unwrap(); - let mut reader = FramedRead::new(socket, TwoPartCodec::default()); + let mut reader = FramedRead::new(socket, MuxCodec::default()); let ack = tokio::time::timeout(Duration::from_secs(1), reader.next()) .await .unwrap() .unwrap() .unwrap(); - assert_eq!( - MuxFrame::try_from_two_part(ack).unwrap(), - MuxFrame::connection_ack(0) - ); - let mut reader = reader.map_decoder(|_| MuxCodec::default()); + assert_eq!(ack, MuxFrame::connection_ready()); let bad = tokio::time::timeout(Duration::from_secs(1), bad_provider) .await