Skip to main content

inferlab_proxy/
core.rs

1//! Shared HTTP mechanics for the built-in disaggregated-serving proxies.
2//!
3//! Proxy-specific protocol bodies remain in their owning modules.
4
5use crate::error::ProxyError as ProxyLifecycleError;
6use async_stream::try_stream;
7use axum::Json;
8use axum::body::Body;
9use axum::http::{HeaderMap, Response, StatusCode, header};
10use axum::response::IntoResponse;
11use bytes::Bytes;
12use futures_util::{FutureExt, Stream, StreamExt};
13use serde::{Deserialize, Serialize};
14use serde_json::Value;
15use std::env;
16use std::fmt;
17use std::future::Future;
18use std::sync::atomic::{AtomicUsize, Ordering};
19use std::time::Duration;
20use tokio::task::JoinHandle;
21
22/// Build a multi-threaded Tokio runtime and drive `run_async` to completion.
23pub fn run<F, Fut>(run_async: F) -> Result<(), ProxyLifecycleError>
24where
25    F: FnOnce() -> Fut,
26    Fut: Future<Output = Result<(), ProxyLifecycleError>>,
27{
28    let runtime = tokio::runtime::Builder::new_multi_thread()
29        .enable_all()
30        .build()
31        .map_err(|error| ProxyLifecycleError::Lifecycle {
32            message: format!("failed to create proxy tokio runtime: {error}"),
33        })?;
34    runtime.block_on(run_async())
35}
36
37/// Healthcheck response body shared by the proxies.
38#[derive(Serialize)]
39pub struct ProxyHealthcheckResponse {
40    pub ready: bool,
41    pub prefill_instances: usize,
42    pub decode_instances: usize,
43}
44
45/// The shared `/healthcheck` payload: 200 once ready, 503 before, with the configured instance counts.
46pub(crate) fn healthcheck_response(
47    ready: bool,
48    prefill_instances: usize,
49    decode_instances: usize,
50) -> (StatusCode, Json<ProxyHealthcheckResponse>) {
51    let status = if ready {
52        StatusCode::OK
53    } else {
54        StatusCode::SERVICE_UNAVAILABLE
55    };
56    (
57        status,
58        Json(ProxyHealthcheckResponse {
59            ready,
60            prefill_instances,
61            decode_instances,
62        }),
63    )
64}
65
66/// Validate that both role endpoint lists are non-empty, naming the proxy in the validation message.
67pub(crate) fn require_endpoints(
68    proxy_name: &'static str,
69    prefill_is_empty: bool,
70    decode_is_empty: bool,
71) -> Result<(), ProxyLifecycleError> {
72    if prefill_is_empty {
73        return Err(ProxyLifecycleError::Invalid {
74            message: format!("{proxy_name} requires at least one prefill endpoint"),
75        });
76    }
77    if decode_is_empty {
78        return Err(ProxyLifecycleError::Invalid {
79            message: format!("{proxy_name} requires at least one decode endpoint"),
80        });
81    }
82    Ok(())
83}
84
85/// Build the shared pooled client, naming the proxy in the failure message.
86pub(crate) fn pooled_client(
87    proxy_name: &'static str,
88) -> Result<reqwest::Client, ProxyLifecycleError> {
89    build_pooled_client().map_err(|error| ProxyLifecycleError::Io {
90        message: format!("failed to create {proxy_name} HTTP client: {error}"),
91    })
92}
93
94/// Bind `host:port` and serve `router`, naming the proxy in bind/serve failure messages.
95pub(crate) async fn serve_router(
96    proxy_name: &'static str,
97    host: &str,
98    port: u16,
99    router: axum::Router,
100) -> Result<(), ProxyLifecycleError> {
101    let listener = tokio::net::TcpListener::bind((host, port))
102        .await
103        .map_err(|error| ProxyLifecycleError::Io {
104            message: format!("failed to bind {proxy_name} on {host}:{port}: {error}"),
105        })?;
106    axum::serve(listener, router)
107        .await
108        .map_err(|error| ProxyLifecycleError::Io {
109            message: format!("{proxy_name} server failed: {error}"),
110        })
111}
112
113/// Poll every backend `url` at `path` once per second until each answers with a success status.
114pub(crate) async fn await_backends(client: reqwest::Client, urls: Vec<String>, path: &'static str) {
115    let waits = urls
116        .into_iter()
117        .map(|url| await_backend(client.clone(), url, path));
118    futures_util::future::join_all(waits).await;
119}
120
121async fn await_backend(client: reqwest::Client, url: String, path: &'static str) {
122    loop {
123        if client
124            .get(join_path(&url, path))
125            .send()
126            .await
127            .is_ok_and(|response| response.status().is_success())
128        {
129            return;
130        }
131        tokio::time::sleep(BACKEND_RETRY_INTERVAL).await;
132    }
133}
134
135/// Pause between readiness-gate attempts against a backend that is not up
136/// yet; the owning server readiness budget ends the wait.
137pub(crate) const BACKEND_RETRY_INTERVAL: Duration = Duration::from_secs(1);
138
139/// Collect the fan-out target URLs shared by the reset/flush sweeps and the
140/// readiness wait: every prefill replica URL followed by every decode URL.
141pub(crate) fn fanout_target_urls<'a>(
142    prefill_urls: impl IntoIterator<Item = &'a str>,
143    decode_urls: impl IntoIterator<Item = &'a str>,
144) -> Vec<String> {
145    prefill_urls
146        .into_iter()
147        .chain(decode_urls)
148        .map(str::to_owned)
149        .collect()
150}
151
152/// Start a proxy response builder that mirrors the upstream response's
153/// status and (when present) content-type.
154pub(crate) fn upstream_response_builder(
155    response: &reqwest::Response,
156) -> Result<axum::http::response::Builder, ProxyHttpError> {
157    let mut builder = Response::builder().status(status_code(response.status())?);
158    if let Some(content_type) = response
159        .headers()
160        .get(reqwest::header::CONTENT_TYPE)
161        .and_then(|value| value.to_str().ok())
162    {
163        builder = builder.header(header::CONTENT_TYPE, content_type);
164    }
165    Ok(builder)
166}
167
168/// Finish a proxy response builder with `body`.
169pub(crate) fn response_body(
170    builder: axum::http::response::Builder,
171    body: Body,
172) -> Result<Response<Body>, ProxyHttpError> {
173    builder.body(body).map_err(|error| {
174        ProxyHttpError::internal(format!("failed to build proxy response: {error}"))
175    })
176}
177
178/// Forward an upstream response body verbatim, preserving status and
179/// content-type.
180pub async fn forward_response(
181    response: reqwest::Response,
182) -> Result<Response<Body>, ProxyHttpError> {
183    let builder = upstream_response_builder(&response)?;
184    let bytes = response
185        .bytes()
186        .await
187        .map_err(|error| ProxyHttpError::upstream("upstream response body read failed", error))?;
188    response_body(builder, Body::from(bytes))
189}
190
191/// Convert an unsuccessful upstream response into a `502 Bad Gateway`
192/// [`ProxyHttpError`] that captures the upstream status and body.
193pub async fn upstream_status_error(context: &str, response: reqwest::Response) -> ProxyHttpError {
194    let status = response.status();
195    let body = match response.text().await {
196        Ok(text) => text,
197        Err(error) => format!("<failed to read upstream error body: {error}>"),
198    };
199    ProxyHttpError::status(
200        StatusCode::BAD_GATEWAY,
201        format!("{context} returned HTTP {status}: {body}"),
202    )
203}
204
205/// Resolve the outbound `Authorization` header from the inbound request or the
206/// `OPENAI_API_KEY` environment variable.
207pub fn outbound_authorization(headers: &HeaderMap) -> Option<String> {
208    headers
209        .get(header::AUTHORIZATION)
210        .and_then(|value| value.to_str().ok())
211        .map(str::to_owned)
212        .or_else(|| {
213            env::var("OPENAI_API_KEY")
214                .ok()
215                .map(|key| format!("Bearer {key}"))
216        })
217}
218
219/// Join a base URL with a path, normalizing a single trailing slash on the
220/// base.
221pub fn join_path(base: &str, path: &str) -> String {
222    format!("{}{}", base.trim_end_matches('/'), path)
223}
224
225/// Convert a `reqwest` status code into an `axum`/`http` status code.
226pub fn status_code(status: reqwest::StatusCode) -> Result<StatusCode, ProxyHttpError> {
227    StatusCode::from_u16(status.as_u16())
228        .map_err(|error| ProxyHttpError::internal(format!("invalid upstream status code: {error}")))
229}
230
231/// Advance a round-robin cursor and return the selected index into a non-empty
232/// target list. Shared by all proxies' prefill/decode selection; each proxy
233/// keeps its own cursor and target list (which stay local).
234pub(crate) fn round_robin_index(cursor: &AtomicUsize, len: usize) -> usize {
235    cursor.fetch_add(1, Ordering::SeqCst) % len
236}
237
238/// Build and send a JSON POST to `url` with an optional `X-Request-Id`, any
239/// `extra_headers`, and an optional `Authorization`, returning the response or a
240/// [`ProxyHttpError`] on transport or non-success status. `context` names the
241/// call in error messages (e.g. "decode request"). Owns the transport for every
242/// built-in proxy POST. Proxy-specific request ids and headers are optional.
243pub(crate) async fn send_json_post(
244    client: reqwest::Client,
245    url: String,
246    body: &Value,
247    request_id: Option<&str>,
248    authorization: Option<&str>,
249    extra_headers: &[(&str, String)],
250    context: &'static str,
251) -> Result<reqwest::Response, ProxyHttpError> {
252    let response = send_json_post_status(
253        client,
254        url,
255        body,
256        request_id,
257        authorization,
258        extra_headers,
259        context,
260    )
261    .await?;
262    if !response.status().is_success() {
263        return Err(upstream_status_error(context, response).await);
264    }
265    Ok(response)
266}
267
268/// Like [`send_json_post`], but returns the response for any upstream status:
269/// fan-out callers record per-target statuses instead of failing fast on the
270/// first non-success upstream response.
271pub(crate) async fn send_json_post_status(
272    client: reqwest::Client,
273    url: String,
274    body: &Value,
275    request_id: Option<&str>,
276    authorization: Option<&str>,
277    extra_headers: &[(&str, String)],
278    context: &'static str,
279) -> Result<reqwest::Response, ProxyHttpError> {
280    let mut request = client.post(url).json(body);
281    if let Some(request_id) = request_id {
282        request = request.header("X-Request-Id", request_id);
283    }
284    // Extra headers precede `Authorization`: header insertion order reaches
285    // the wire, and Mooncake's prefill always sent its rank header first.
286    for (name, value) in extra_headers {
287        request = request.header(*name, value);
288    }
289    if let Some(authorization) = authorization {
290        request = request.header(reqwest::header::AUTHORIZATION, authorization);
291    }
292    request
293        .send()
294        .await
295        .map_err(|error| ProxyHttpError::upstream(&format!("{context} failed"), error))
296}
297
298/// A per-process monotonic request id, `"{pid}-{n}"`, drawn from a proxy-owned
299/// counter. Shared by the vLLM proxies so the id scheme has one home.
300pub(crate) fn next_request_id(counter: &AtomicUsize) -> String {
301    let value = counter.fetch_add(1, Ordering::SeqCst);
302    format!("{}-{value}", std::process::id())
303}
304
305/// Build the outbound HTTP client shared by the proxies, with the pool tuning
306/// (unbounded idle connections per host) both require. Returns the raw
307/// `reqwest` error so each proxy keeps its own construction-failure message.
308pub(crate) fn build_pooled_client() -> reqwest::Result<reqwest::Client> {
309    reqwest::Client::builder()
310        .pool_max_idle_per_host(usize::MAX)
311        .build()
312}
313
314/// Failure detail of one reset/flush fan-out target.
315#[derive(Debug, Deserialize, Serialize)]
316pub struct FanoutFailure {
317    pub url: String,
318    pub error: String,
319}
320
321/// Aggregated response of the cache reset/flush fan-out endpoints. SGLang's
322/// `flush_cache` and the vLLM proxies' `reset_prefix_cache` share this wire
323/// contract.
324#[derive(Debug, Deserialize, Serialize)]
325pub struct ResetPrefixCacheResponse {
326    pub successful: Vec<String>,
327    pub failed: Vec<FanoutFailure>,
328}
329
330/// Aggregated response of the prefix-cache conditioning fan-out endpoint.
331/// The control plane deserializes this exact shape, so the fields are the
332/// cross-process contract.
333#[derive(Debug, Deserialize, Serialize)]
334pub struct PrimePrefixCacheResponse {
335    pub targets: Vec<PrimePrefixCacheTarget>,
336}
337
338/// One fanned-out conditioning flow: the prefill replica URL and the pinned
339/// data-parallel rank, with the observed status or the failure detail.
340#[derive(Debug, Deserialize, Serialize)]
341pub struct PrimePrefixCacheTarget {
342    pub url: String,
343    pub rank: u32,
344    pub http_status: Option<u16>,
345    pub elapsed_ms: u64,
346    pub error: Option<String>,
347}
348
349/// Failure of one fanned-out conditioning flow: the upstream status when a
350/// response was observed, plus the failure detail.
351pub(crate) struct PrimeFlowFailure {
352    pub http_status: Option<u16>,
353    pub error: String,
354}
355
356impl PrimeFlowFailure {
357    pub(crate) fn transport(error: ProxyHttpError) -> Self {
358        Self {
359            http_status: None,
360            error: error.to_string(),
361        }
362    }
363
364    pub(crate) fn status(status: u16, detail: String) -> Self {
365        Self {
366            http_status: Some(status),
367            error: detail,
368        }
369    }
370}
371
372/// Read a fanned-out conditioning response to text, requiring a 2xx status.
373/// The body is captured either way so a non-2xx status reports it as the
374/// failure detail; a body read failure is a transport failure. Returns the
375/// upstream status and body on success.
376pub(crate) async fn expect_2xx(
377    context: &'static str,
378    response: reqwest::Response,
379) -> Result<(u16, String), PrimeFlowFailure> {
380    let status = response.status().as_u16();
381    let text = response.text().await.map_err(|error| {
382        PrimeFlowFailure::transport(ProxyHttpError::upstream(
383            &format!("{context} response read failed"),
384            error,
385        ))
386    })?;
387    if !(200..300).contains(&status) {
388        return Err(PrimeFlowFailure::status(
389            status,
390            format!("{context} returned HTTP {status}: {text}"),
391        ));
392    }
393    Ok((status, text))
394}
395
396/// Identity of one prime fan-out target: the prefill replica URL and the
397/// data-parallel rank the flow pins.
398pub(crate) trait PrimeFanoutTarget {
399    fn url(&self) -> &str;
400    fn rank(&self) -> u32;
401}
402
403/// A prefill replica with a static (config-issued) data-parallel size, as
404/// opposed to Mooncake's discovered per-rank engines.
405pub(crate) trait PrimeReplica {
406    fn url(&self) -> &str;
407    fn data_parallel_size(&self) -> u32;
408}
409
410/// One prime fan-out target over a static-size replica: the replica and the
411/// pinned data-parallel rank.
412pub(crate) struct RankedPrimeTarget<R> {
413    pub replica: R,
414    pub rank: u32,
415}
416
417impl<R: PrimeReplica> PrimeFanoutTarget for RankedPrimeTarget<R> {
418    fn url(&self) -> &str {
419        self.replica.url()
420    }
421
422    fn rank(&self) -> u32 {
423        self.rank
424    }
425}
426
427/// Enumerate the prime fan-out targets for static-size replicas, expanding
428/// each replica over its data-parallel ranks (at least one).
429pub(crate) fn ranked_prime_targets<R: PrimeReplica + Clone>(
430    replicas: &[R],
431) -> Vec<RankedPrimeTarget<R>> {
432    let mut targets = Vec::new();
433    for replica in replicas {
434        for rank in 0..replica.data_parallel_size().max(1) {
435            targets.push(RankedPrimeTarget {
436                replica: replica.clone(),
437                rank,
438            });
439        }
440    }
441    targets
442}
443
444/// Run the reset/flush fan-out skeleton: the engine module enumerates the
445/// target base URLs and names its endpoint (`path`) and operation; target
446/// execution and the response aggregation (200 when every target succeeded,
447/// 206 on partial failure) live here. The caller's measurement-case request
448/// deadline bounds the whole fan-out; the proxy adds no shorter cap
449/// ([[RFC-0009:C-MEASUREMENT-CASE-BUDGETS]]). An empty target set is a 502 —
450/// "no targets" must not be conflated with success.
451pub(crate) async fn run_sweep_fanout(
452    client: reqwest::Client,
453    operation: &'static str,
454    path: &'static str,
455    targets: Vec<String>,
456    authorization: Option<String>,
457) -> Response<Body> {
458    if targets.is_empty() {
459        return empty_fanout_failure(operation);
460    }
461    let attempts = targets
462        .into_iter()
463        .map(|url| sweep_target(client.clone(), operation, path, url, authorization.clone()));
464    let mut successful = Vec::new();
465    let mut failed = Vec::new();
466    for result in futures_util::future::join_all(attempts).await {
467        match result {
468            Ok(url) => successful.push(url),
469            Err(failure) => failed.push(failure),
470        }
471    }
472    let status = if failed.is_empty() {
473        StatusCode::OK
474    } else {
475        StatusCode::PARTIAL_CONTENT
476    };
477    (
478        status,
479        Json(ResetPrefixCacheResponse { successful, failed }),
480    )
481        .into_response()
482}
483
484async fn sweep_target(
485    client: reqwest::Client,
486    operation: &'static str,
487    path: &'static str,
488    url: String,
489    authorization: Option<String>,
490) -> Result<String, FanoutFailure> {
491    let endpoint = join_path(&url, path);
492    let mut request = client.post(endpoint);
493    if let Some(authorization) = authorization {
494        request = request.header(reqwest::header::AUTHORIZATION, authorization);
495    }
496    let response = request.send().await.map_err(|error| FanoutFailure {
497        url: url.clone(),
498        error: format!("{operation} request failed: {error}"),
499    })?;
500    // A 206 from an upstream that is itself an aggregating frontend reports
501    // partial failure, not success.
502    if response.status().is_success() && response.status() != reqwest::StatusCode::PARTIAL_CONTENT {
503        Ok(url)
504    } else {
505        let status = response.status();
506        let detail = response
507            .text()
508            .await
509            .unwrap_or_else(|error| format!("failed to read response body: {error}"));
510        Err(FanoutFailure {
511            url,
512            error: format!("HTTP {status}: {detail}"),
513        })
514    }
515}
516
517/// Run the prefix-cache conditioning fan-out skeleton: the engine module
518/// enumerates the (replica, rank) targets and supplies the per-target
519/// conditioning flow; sequential target execution and the response
520/// aggregation (200 when every flow succeeded, 206 on partial failure) live
521/// here, bounded by the caller's request deadline like the sweep. An empty
522/// target set is a 502 — "no targets" must not be conflated with success.
523pub(crate) async fn run_prime_fanout<T, F, Fut>(
524    operation: &'static str,
525    targets: Vec<T>,
526    mut execute: F,
527) -> Response<Body>
528where
529    T: PrimeFanoutTarget,
530    F: FnMut(T) -> Fut,
531    Fut: Future<Output = Result<u16, PrimeFlowFailure>>,
532{
533    if targets.is_empty() {
534        return empty_fanout_failure(operation);
535    }
536    let mut results = Vec::new();
537    for target in targets {
538        let url = target.url().to_owned();
539        let rank = target.rank();
540        let started = std::time::Instant::now();
541        let outcome = execute(target).await;
542        let elapsed_ms = u64::try_from(started.elapsed().as_millis()).unwrap_or(u64::MAX);
543        results.push(match outcome {
544            Ok(status) => PrimePrefixCacheTarget {
545                url,
546                rank,
547                http_status: Some(status),
548                elapsed_ms,
549                error: None,
550            },
551            Err(failure) => PrimePrefixCacheTarget {
552                url,
553                rank,
554                http_status: failure.http_status,
555                elapsed_ms,
556                error: Some(failure.error),
557            },
558        });
559    }
560    let status = if results.iter().all(|target| target.error.is_none()) {
561        StatusCode::OK
562    } else {
563        StatusCode::PARTIAL_CONTENT
564    };
565    (status, Json(PrimePrefixCacheResponse { targets: results })).into_response()
566}
567
568/// "No targets" is a proxy-side failure (502 with an explicit error), never
569/// a 200/206 aggregate: an empty fan-out primes or resets nothing.
570fn empty_fanout_failure(operation: &str) -> Response<Body> {
571    ProxyHttpError::status(
572        StatusCode::BAD_GATEWAY,
573        format!("{operation} fan-out has no targets: no prefill replica or data-parallel rank is available"),
574    )
575    .into_response()
576}
577
578/// Per-request error of the HTTP handlers: an HTTP status plus a message.
579#[derive(Debug)]
580pub struct ProxyHttpError {
581    status: StatusCode,
582    message: String,
583}
584
585impl ProxyHttpError {
586    pub fn status(status: StatusCode, message: impl Into<String>) -> Self {
587        Self {
588            status,
589            message: message.into(),
590        }
591    }
592
593    pub fn upstream(context: &str, error: reqwest::Error) -> Self {
594        Self::status(StatusCode::BAD_GATEWAY, format!("{context}: {error}"))
595    }
596
597    pub fn internal(message: impl Into<String>) -> Self {
598        Self::status(StatusCode::INTERNAL_SERVER_ERROR, message)
599    }
600}
601
602impl fmt::Display for ProxyHttpError {
603    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
604        write!(formatter, "{}", self.message)
605    }
606}
607
608impl std::error::Error for ProxyHttpError {}
609
610impl IntoResponse for ProxyHttpError {
611    fn into_response(self) -> axum::response::Response {
612        let body = Json(ProxyErrorResponse {
613            error: self.message,
614        });
615        (self.status, body).into_response()
616    }
617}
618
619#[derive(Serialize)]
620pub struct ProxyErrorResponse {
621    pub error: String,
622}
623
624/// What happens to the prefill task when the client drops the decode
625/// response stream before prefill completes.
626#[derive(Clone, Copy, Debug)]
627pub(crate) enum OnClientDrop {
628    /// Abort the prefill task (via [`AbortOnDrop`]). The default: once the
629    /// client is gone, the orphaned prefill request is cancelled.
630    Abort,
631    /// Leave the prefill task running: it drains to completion in the
632    /// background. Required when aborting prefill mid-flight would strand
633    /// the paired decode-side engine request (for example the SGLang
634    /// prefill/decode bootstrap room, where the decode engine waits for KV
635    /// that a cancelled prefill would never deliver).
636    Detach,
637}
638
639/// Stream a decode response body while a concurrently-running prefill task
640/// completes. Used by proxies whose backend protocol starts both roles
641/// together; the vLLM NIXL proxy instead forwards them sequentially.
642/// `on_client_drop` selects the prefill task's fate when the client drops
643/// the response before prefill finishes (see [`OnClientDrop`]); once prefill
644/// completes any armed abort is disarmed, and a prefill failure surfaces as
645/// a stream error.
646pub(crate) fn stream_decode_response(
647    response: reqwest::Response,
648    prefill_task: JoinHandle<Result<(), ProxyHttpError>>,
649    on_client_drop: OnClientDrop,
650) -> Result<Response<Body>, ProxyHttpError> {
651    let builder = upstream_response_builder(&response)?;
652    let stream = decode_response_stream(response.bytes_stream(), prefill_task, on_client_drop);
653    response_body(builder, Body::from_stream(stream))
654}
655
656/// Stream one successful upstream response without waiting for another role.
657pub(crate) fn stream_response(
658    response: reqwest::Response,
659) -> Result<Response<Body>, ProxyHttpError> {
660    let builder = upstream_response_builder(&response)?;
661    let stream = response
662        .bytes_stream()
663        .map(|chunk| chunk.map_err(|error| stream_error(format!("decode stream failed: {error}"))));
664    response_body(builder, Body::from_stream(stream))
665}
666
667/// The decode byte stream, generic over the decode stream and its error type so
668/// it can be exercised without a live `reqwest::Response`. Yields decode bytes
669/// in arrival order; concurrently drives `prefill_task` to completion and, per
670/// `on_client_drop`, aborts or detaches it if the consumer drops the stream
671/// before prefill finishes. On decode EOF before prefill completes, prefill is
672/// awaited and its error (if any) surfaces.
673pub(crate) fn decode_response_stream<S, E>(
674    decode_stream: S,
675    prefill_task: JoinHandle<Result<(), ProxyHttpError>>,
676    on_client_drop: OnClientDrop,
677) -> impl Stream<Item = std::result::Result<Bytes, std::io::Error>>
678where
679    S: Stream<Item = std::result::Result<Bytes, E>> + Unpin,
680    E: fmt::Display,
681{
682    let prefill_abort = prefill_task.abort_handle();
683    try_stream! {
684        let mut decode_stream = decode_stream;
685        let mut prefill_task = prefill_task;
686        // Under `OnClientDrop::Detach` no abort guard is armed at all: dropping
687        // the stream (client disconnect) leaves the prefill task draining to
688        // completion in the background.
689        let mut prefill_abort = match on_client_drop {
690            OnClientDrop::Abort => Some(AbortOnDrop::new(prefill_abort)),
691            OnClientDrop::Detach => None,
692        };
693        let mut prefill_done = false;
694        loop {
695            match next_stream_event(&mut prefill_task, &mut decode_stream, prefill_done).await {
696                StreamEvent::Prefill(prefill) => {
697                    prefill_done = true;
698                    // One-time tie-break: if a decode item was already ready at the
699                    // instant prefill completed, handle that single item before
700                    // surfacing the prefill outcome. `now_or_never()` polls (and so
701                    // consumes) the item, so it must be matched exhaustively: deliver
702                    // a ready chunk, and PROPAGATE a ready decode error (an Ok-only
703                    // match would drop it and truncate the response into a clean 200).
704                    // EOF / not-ready fall through to the
705                    // prefill outcome; decode is never indefinitely preferred.
706                    match decode_stream.next().now_or_never() {
707                        Some(Some(Ok(bytes))) => yield bytes,
708                        Some(Some(Err(error))) => {
709                            Err(stream_error(format!("decode stream failed: {error}")))?;
710                        }
711                        Some(None) | None => {}
712                    }
713                    prefill
714                        .map_err(join_error)?
715                        .map_err(|error| stream_error(error.to_string()))?;
716                    if let Some(abort) = &mut prefill_abort {
717                        abort.disarm();
718                    }
719                }
720                StreamEvent::Decode(Some(Ok(bytes))) => yield bytes,
721                StreamEvent::Decode(Some(Err(error))) => {
722                    Err(stream_error(format!("decode stream failed: {error}")))?;
723                }
724                StreamEvent::Decode(None) => break,
725            }
726        }
727        if !prefill_done {
728            prefill_task
729                .await
730                .map_err(join_error)?
731                .map_err(|error| stream_error(error.to_string()))?;
732            if let Some(abort) = &mut prefill_abort {
733                abort.disarm();
734            }
735        }
736    }
737}
738
739enum StreamEvent<E> {
740    Prefill(std::result::Result<Result<(), ProxyHttpError>, tokio::task::JoinError>),
741    Decode(Option<std::result::Result<Bytes, E>>),
742}
743
744async fn next_stream_event<S, E>(
745    prefill_task: &mut JoinHandle<Result<(), ProxyHttpError>>,
746    decode_stream: &mut S,
747    prefill_done: bool,
748) -> StreamEvent<E>
749where
750    S: Stream<Item = std::result::Result<Bytes, E>> + Unpin,
751{
752    // Surface the prefill outcome promptly once the prefill task has COMPLETED, so
753    // a continuously-ready decode stream cannot defer (and thereby suppress) a
754    // prefill failure indefinitely. The caller delivers one already-ready decode
755    // chunk before propagating a prefill error (a one-time tie-break), so a chunk
756    // that was ready at the instant prefill finished is not dropped — but decode
757    // is NOT permanently prioritized.
758    if !prefill_done && prefill_task.is_finished() {
759        return StreamEvent::Prefill(prefill_task.await);
760    }
761    // Prefill is still running: deliver decode bytes as they arrive, and otherwise
762    // await the prefill task's completion (picked up by the `is_finished` check on
763    // the next call). An unbiased race is fine here — there is no completed prefill
764    // outcome to drop, and a ready decode chunk taken by its own branch is yielded,
765    // not lost.
766    tokio::select! {
767        prefill = prefill_task, if !prefill_done => StreamEvent::Prefill(prefill),
768        chunk = decode_stream.next() => StreamEvent::Decode(chunk),
769    }
770}
771
772fn join_error(error: tokio::task::JoinError) -> std::io::Error {
773    stream_error(format!("prefill task failed: {error}"))
774}
775
776fn stream_error(message: String) -> std::io::Error {
777    std::io::Error::other(message)
778}
779
780/// Aborts the held task when dropped unless [`disarm`](AbortOnDrop::disarm)ed.
781struct AbortOnDrop {
782    handle: tokio::task::AbortHandle,
783    armed: bool,
784}
785
786impl AbortOnDrop {
787    fn new(handle: tokio::task::AbortHandle) -> Self {
788        Self {
789            handle,
790            armed: true,
791        }
792    }
793
794    fn disarm(&mut self) {
795        self.armed = false;
796    }
797}
798
799impl Drop for AbortOnDrop {
800    fn drop(&mut self) {
801        if self.armed {
802            self.handle.abort();
803        }
804    }
805}
806
807#[cfg(test)]
808mod tests {
809    use super::*;
810    use anyhow::{Context, Result};
811
812    #[test]
813    fn join_path_normalizes_single_trailing_slash() {
814        assert_eq!(
815            join_path("http://h:1/", "/v1/models"),
816            "http://h:1/v1/models"
817        );
818        assert_eq!(
819            join_path("http://h:1", "/v1/models"),
820            "http://h:1/v1/models"
821        );
822    }
823
824    #[test]
825    fn status_code_maps_reqwest_status() -> Result<()> {
826        let mapped = status_code(reqwest::StatusCode::OK)
827            .map_err(|error| anyhow::anyhow!(error.to_string()))?;
828        assert_eq!(mapped, StatusCode::OK);
829        Ok(())
830    }
831
832    #[test]
833    fn outbound_authorization_prefers_inbound_header() -> Result<()> {
834        let mut headers = HeaderMap::new();
835        headers.insert(header::AUTHORIZATION, "Bearer inbound".parse()?);
836        assert_eq!(
837            outbound_authorization(&headers),
838            Some("Bearer inbound".to_owned())
839        );
840        Ok(())
841    }
842
843    #[test]
844    fn proxy_error_internal_uses_500() {
845        let error = ProxyHttpError::internal("boom");
846        assert_eq!(error.status, StatusCode::INTERNAL_SERVER_ERROR);
847        assert_eq!(error.to_string(), "boom");
848    }
849
850    struct StaticPrimeTarget {
851        url: &'static str,
852        rank: u32,
853    }
854
855    impl PrimeFanoutTarget for StaticPrimeTarget {
856        fn url(&self) -> &str {
857            self.url
858        }
859
860        fn rank(&self) -> u32 {
861            self.rank
862        }
863    }
864
865    /// An empty prime fan-out must be a 502 with an explicit error: a 200
866    /// over zero targets would record a primed cache that primed nothing.
867    #[test]
868    fn prime_fanout_rejects_an_empty_target_set() -> Result<()> {
869        let runtime = proxy_test_runtime()?;
870        let response = runtime.block_on(run_prime_fanout(
871            "prefix cache conditioning",
872            Vec::<StaticPrimeTarget>::new(),
873            |_target| async { Ok::<u16, PrimeFlowFailure>(200) },
874        ));
875        assert_eq!(response.status(), StatusCode::BAD_GATEWAY);
876        let body = runtime.block_on(axum::body::to_bytes(response.into_body(), usize::MAX))?;
877        let value: Value = serde_json::from_slice(&body)?;
878        assert!(
879            value["error"]
880                .as_str()
881                .is_some_and(|error| error.contains("no targets")),
882            "got {value}"
883        );
884        Ok(())
885    }
886
887    /// Same guard for the reset/flush sweep: zero targets is a proxy
888    /// failure, never a clean sweep.
889    #[test]
890    fn sweep_fanout_rejects_an_empty_target_set() -> Result<()> {
891        let runtime = proxy_test_runtime()?;
892        let client = build_pooled_client().map_err(|error| anyhow::anyhow!(error.to_string()))?;
893        let response = runtime.block_on(run_sweep_fanout(
894            client,
895            "prefix cache reset",
896            "/reset_prefix_cache",
897            Vec::new(),
898            None,
899        ));
900        assert_eq!(response.status(), StatusCode::BAD_GATEWAY);
901        let body = runtime.block_on(axum::body::to_bytes(response.into_body(), usize::MAX))?;
902        let value: Value = serde_json::from_slice(&body)?;
903        assert!(
904            value["error"]
905                .as_str()
906                .is_some_and(|error| error.contains("no targets")),
907            "got {value}"
908        );
909        Ok(())
910    }
911
912    /// The proxy imposes no per-target cap of its own: the caller's
913    /// measurement-case deadline bounds a fan-out, and a caller that gives up
914    /// cancels the in-flight target work ([[RFC-0009:C-MEASUREMENT-CASE-BUDGETS]]).
915    #[test]
916    fn prime_fanout_is_cancelled_when_the_caller_gives_up() -> Result<()> {
917        let runtime = proxy_test_runtime()?;
918        let dropped = Arc::new(AtomicBool::new(false));
919        let flag = dropped.clone();
920        let observed = dropped.clone();
921        runtime.block_on(async move {
922            let app = axum::Router::new().route(
923                "/prime",
924                axum::routing::post(move || {
925                    let flag = flag.clone();
926                    async move {
927                        run_prime_fanout(
928                            "prefix cache conditioning",
929                            vec![StaticPrimeTarget {
930                                url: "http://127.0.0.1:1",
931                                rank: 0,
932                            }],
933                            move |_target| {
934                                let guard = SetOnDrop(flag.clone());
935                                async move {
936                                    let _guard = guard;
937                                    futures_util::future::pending::<()>().await;
938                                    Ok::<u16, PrimeFlowFailure>(200)
939                                }
940                            },
941                        )
942                        .await
943                    }
944                }),
945            );
946            let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await?;
947            let addr = listener.local_addr()?;
948            tokio::spawn(async move { axum::serve(listener, app).await });
949            let result = reqwest::Client::new()
950                .post(format!("http://{addr}/prime"))
951                .timeout(Duration::from_millis(200))
952                .send()
953                .await;
954            assert!(result.is_err(), "the pending target must hold the response");
955            for _ in 0..50 {
956                if observed.load(Ordering::SeqCst) {
957                    break;
958                }
959                tokio::time::sleep(Duration::from_millis(20)).await;
960            }
961            anyhow::Ok(())
962        })?;
963        assert!(
964            dropped.load(Ordering::SeqCst),
965            "target work outlived the caller's deadline"
966        );
967        Ok(())
968    }
969
970    use std::sync::Arc;
971    use std::sync::atomic::{AtomicBool, Ordering};
972
973    /// Flips a shared flag when dropped — used to observe that an aborted
974    /// prefill task is actually cancelled (its future is dropped).
975    struct SetOnDrop(Arc<AtomicBool>);
976
977    impl Drop for SetOnDrop {
978        fn drop(&mut self) {
979            self.0.store(true, Ordering::SeqCst);
980        }
981    }
982
983    fn proxy_test_runtime() -> Result<tokio::runtime::Runtime> {
984        tokio::runtime::Builder::new_multi_thread()
985            .enable_all()
986            .build()
987            .map_err(|error| anyhow::anyhow!(error.to_string()))
988    }
989
990    #[test]
991    fn streamed_decode_yields_bytes_in_order_when_prefill_succeeds() -> Result<()> {
992        let runtime = proxy_test_runtime()?;
993        let bytes = runtime.block_on(async {
994            let decode = Box::pin(futures_util::stream::iter(vec![
995                std::result::Result::<Bytes, std::io::Error>::Ok(Bytes::from_static(b"hello")),
996                Ok(Bytes::from_static(b" world")),
997            ]));
998            let prefill = tokio::spawn(async { Ok::<(), ProxyHttpError>(()) });
999            let mut stream = Box::pin(decode_response_stream(decode, prefill, OnClientDrop::Abort));
1000            let mut out = Vec::new();
1001            while let Some(item) = stream.next().await {
1002                out.push(item.map_err(|error| anyhow::anyhow!(error.to_string()))?);
1003            }
1004            anyhow::Ok(out)
1005        })?;
1006        let joined: Vec<u8> = bytes.into_iter().flatten().collect();
1007        assert_eq!(joined, b"hello world");
1008        Ok(())
1009    }
1010
1011    #[test]
1012    fn streamed_decode_surfaces_prefill_error_after_decode_ends() -> Result<()> {
1013        let runtime = proxy_test_runtime()?;
1014        let (bytes, error) = runtime.block_on(async {
1015            let decode = Box::pin(futures_util::stream::iter(vec![std::result::Result::<
1016                Bytes,
1017                std::io::Error,
1018            >::Ok(
1019                Bytes::from_static(b"partial"),
1020            )]));
1021            let prefill = tokio::spawn(async {
1022                Err::<(), ProxyHttpError>(ProxyHttpError::internal("prefill boom"))
1023            });
1024            let mut stream = Box::pin(decode_response_stream(decode, prefill, OnClientDrop::Abort));
1025            let mut bytes = Vec::new();
1026            let mut error = None;
1027            while let Some(item) = stream.next().await {
1028                match item {
1029                    Ok(chunk) => bytes.extend_from_slice(&chunk),
1030                    Err(stream_error) => {
1031                        error = Some(stream_error.to_string());
1032                        break;
1033                    }
1034                }
1035            }
1036            anyhow::Ok((bytes, error))
1037        })?;
1038        assert_eq!(bytes, b"partial");
1039        let error = error.context("expected a prefill error to surface after decode ended")?;
1040        assert!(error.contains("prefill boom"), "got {error}");
1041        Ok(())
1042    }
1043
1044    /// A prefill failure must surface even while the decode stream stays
1045    /// continuously ready. The one-time tie-break delivers an already-ready chunk
1046    /// but must NOT let an always-ready decode stream defer the prefill error
1047    /// indefinitely (a permanent decode bias would suppress it).
1048    #[test]
1049    fn prefill_error_surfaces_even_while_decode_stays_ready() -> Result<()> {
1050        let runtime = proxy_test_runtime()?;
1051        let error = runtime.block_on(async {
1052            // An unbounded, always-synchronously-ready decode stream.
1053            let decode = Box::pin(futures_util::stream::repeat_with(|| {
1054                std::result::Result::<Bytes, std::io::Error>::Ok(Bytes::from_static(b"x"))
1055            }));
1056            let prefill = tokio::spawn(async {
1057                Err::<(), ProxyHttpError>(ProxyHttpError::internal("prefill boom"))
1058            });
1059            let mut stream = Box::pin(decode_response_stream(decode, prefill, OnClientDrop::Abort));
1060            let mut chunks = 0usize;
1061            let mut error = None;
1062            while let Some(item) = stream.next().await {
1063                match item {
1064                    Ok(_) => {
1065                        chunks += 1;
1066                        // Bound: prove the error is not suppressed forever. The fix
1067                        // surfaces it within a handful of chunks; a regression that
1068                        // permanently prefers decode would never break out here.
1069                        assert!(
1070                            chunks < 100_000,
1071                            "prefill error was suppressed by a continuously-ready decode stream"
1072                        );
1073                    }
1074                    Err(stream_error) => {
1075                        error = Some(stream_error.to_string());
1076                        break;
1077                    }
1078                }
1079            }
1080            anyhow::Ok(error)
1081        })?;
1082        let error = error.context("a prefill error must surface even while decode stays ready")?;
1083        assert!(error.contains("prefill boom"), "got {error}");
1084        Ok(())
1085    }
1086
1087    /// A decode error that is ALREADY ready at the instant prefill completes must be
1088    /// propagated by the one-time tie-break, not silently dropped. `now_or_never()`
1089    /// polls (and thus consumes) that ready item, so an Ok-only match would discard
1090    /// the error; with a successful prefill the stream would then end cleanly,
1091    /// turning a decode failure into a truncated 200.
1092    #[test]
1093    fn decode_error_ready_at_tiebreak_is_not_swallowed() -> Result<()> {
1094        let runtime = proxy_test_runtime()?;
1095        let error = runtime.block_on(async {
1096            // Force prefill to be FINISHED (Ok) so the first stream event is the
1097            // prefill outcome and the tie-break is what polls the decode stream.
1098            let prefill = tokio::spawn(async { Ok::<(), ProxyHttpError>(()) });
1099            while !prefill.is_finished() {
1100                tokio::task::yield_now().await;
1101            }
1102            // A synchronously-ready decode Err waiting at the tie-break instant.
1103            let decode = Box::pin(futures_util::stream::iter(vec![std::result::Result::<
1104                Bytes,
1105                std::io::Error,
1106            >::Err(
1107                std::io::Error::other("decode boom"),
1108            )]));
1109            let mut stream = Box::pin(decode_response_stream(decode, prefill, OnClientDrop::Abort));
1110            let mut error = None;
1111            while let Some(item) = stream.next().await {
1112                if let Err(stream_error) = item {
1113                    error = Some(stream_error.to_string());
1114                    break;
1115                }
1116            }
1117            anyhow::Ok(error)
1118        })?;
1119        let error =
1120            error.context("a decode error ready at the tie-break must surface, not truncate")?;
1121        assert!(error.contains("decode boom"), "got {error}");
1122        Ok(())
1123    }
1124
1125    #[test]
1126    fn dropping_the_stream_before_prefill_finishes_aborts_prefill() -> Result<()> {
1127        let runtime = proxy_test_runtime()?;
1128        let aborted = Arc::new(AtomicBool::new(false));
1129        let flag = aborted.clone();
1130        let cancelled = runtime.block_on(async move {
1131            let (started_tx, started_rx) = tokio::sync::oneshot::channel::<()>();
1132            // Prefill never completes; once aborted, its future is dropped and
1133            // SetOnDrop flips the flag. It signals `started` only AFTER the guard
1134            // is constructed, so the drop-abort is observed deterministically (no
1135            // race on whether the task was polled before the abort fired).
1136            let prefill = tokio::spawn(async move {
1137                let _guard = SetOnDrop(flag);
1138                let _ = started_tx.send(());
1139                futures_util::future::pending::<()>().await;
1140                Ok::<(), ProxyHttpError>(())
1141            });
1142            let _ = started_rx.await;
1143            // Decode yields one chunk then stays pending, so the loop neither
1144            // breaks (decode EOF) nor selects prefill — leaving prefill in flight.
1145            let decode = Box::pin(
1146                futures_util::stream::once(async {
1147                    std::result::Result::<Bytes, std::io::Error>::Ok(Bytes::from_static(b"a"))
1148                })
1149                .chain(futures_util::stream::pending::<
1150                    std::result::Result<Bytes, std::io::Error>,
1151                >()),
1152            );
1153            let mut stream = Box::pin(decode_response_stream(decode, prefill, OnClientDrop::Abort));
1154            assert!(matches!(stream.next().await, Some(Ok(_))));
1155            drop(stream);
1156            for _ in 0..200 {
1157                if aborted.load(Ordering::SeqCst) {
1158                    return true;
1159                }
1160                tokio::time::sleep(std::time::Duration::from_millis(10)).await;
1161            }
1162            false
1163        });
1164        assert!(
1165            cancelled,
1166            "prefill task was not aborted when the response stream was dropped"
1167        );
1168        Ok(())
1169    }
1170
1171    /// The counterpart pin for `OnClientDrop::Detach`: dropping the consumer
1172    /// must NOT abort the prefill task — it keeps running in the background
1173    /// and completes on its own.
1174    #[test]
1175    fn dropping_the_stream_before_prefill_finishes_detaches_prefill() -> Result<()> {
1176        let runtime = proxy_test_runtime()?;
1177        let completed = Arc::new(AtomicBool::new(false));
1178        let flag = completed.clone();
1179        let finished = runtime.block_on(async move {
1180            let (started_tx, started_rx) = tokio::sync::oneshot::channel::<()>();
1181            // Prefill completes after a short delay and flips the flag; if the
1182            // stream drop aborted it, the flag would never be set.
1183            let prefill = tokio::spawn(async move {
1184                let _ = started_tx.send(());
1185                tokio::time::sleep(std::time::Duration::from_millis(50)).await;
1186                flag.store(true, Ordering::SeqCst);
1187                Ok::<(), ProxyHttpError>(())
1188            });
1189            let _ = started_rx.await;
1190            // Decode yields one chunk then stays pending, so prefill is still in
1191            // flight when the stream is dropped.
1192            let decode = Box::pin(
1193                futures_util::stream::once(async {
1194                    std::result::Result::<Bytes, std::io::Error>::Ok(Bytes::from_static(b"a"))
1195                })
1196                .chain(futures_util::stream::pending::<
1197                    std::result::Result<Bytes, std::io::Error>,
1198                >()),
1199            );
1200            let mut stream = Box::pin(decode_response_stream(
1201                decode,
1202                prefill,
1203                OnClientDrop::Detach,
1204            ));
1205            assert!(matches!(stream.next().await, Some(Ok(_))));
1206            drop(stream);
1207            for _ in 0..200 {
1208                if completed.load(Ordering::SeqCst) {
1209                    return true;
1210                }
1211                tokio::time::sleep(std::time::Duration::from_millis(10)).await;
1212            }
1213            false
1214        });
1215        assert!(
1216            finished,
1217            "prefill task was aborted instead of detached when the response stream was dropped"
1218        );
1219        Ok(())
1220    }
1221}