Skip to main content

inferlab_proxy/
sglang.rs

1//! Built-in routing for SGLang prefill/decode serving under
2//! [[RFC-0003:C-SGLANG-PREFILL-DECODE]].
3
4use crate::core::{
5    self, OnClientDrop, ProxyHealthcheckResponse, ProxyHttpError, forward_response, join_path,
6    outbound_authorization,
7};
8use crate::error::ProxyError;
9use async_stream::stream;
10use axum::Json;
11use axum::body::Body;
12use axum::extract::State;
13use axum::http::{HeaderMap, Response, StatusCode};
14use axum::response::IntoResponse;
15use axum::routing::{Router, get, post};
16use bytes::Bytes;
17use futures_util::{Stream, StreamExt};
18use serde_json::Value;
19use std::fmt;
20use std::sync::Arc;
21use std::sync::atomic::{AtomicBool, AtomicU64, AtomicUsize, Ordering};
22use std::time::{SystemTime, UNIX_EPOCH};
23use tokio::task::JoinHandle;
24
25pub const VERSION: u32 = 2;
26
27pub const HEALTHCHECK_PATH: &str = "/healthcheck";
28/// Prefix-cache reset route this proxy serves; SGLang's cache-clearing
29/// vocabulary is `flush_cache`.
30pub const RESET_PREFIX_CACHE_PATH: &str = "/flush_cache";
31pub const PRIME_PREFIX_CACHE_PATH: &str = "/prime_prefix_cache";
32
33pub const COMPLETIONS_PATH: &str = "/v1/completions";
34pub const CHAT_COMPLETIONS_PATH: &str = "/v1/chat/completions";
35
36/// Display name used in lifecycle/validation error messages.
37const PROXY_NAME: &str = "SGLang proxy";
38
39#[derive(Clone, Debug)]
40pub struct Config {
41    pub host: String,
42    pub port: u16,
43    pub prefill: Vec<PrefillTarget>,
44    pub decode: Vec<String>,
45}
46
47#[derive(Clone, Debug)]
48pub struct PrefillTarget {
49    pub url: String,
50    pub bootstrap_host: String,
51    pub bootstrap_port: u16,
52    /// Effective attention data-parallel size of this prefill replica, issued
53    /// by the control plane at launch: the conditioning fan-out primes each
54    /// rank ([[RFC-0004:C-BENCH-CACHE-STATE]]).
55    pub data_parallel_size: u32,
56}
57
58pub fn run(config: Config) -> Result<(), ProxyError> {
59    core::run(|| run_async(config))
60}
61
62pub async fn run_async(config: Config) -> Result<(), ProxyError> {
63    let host = config.host.clone();
64    let port = config.port;
65    let state = ProxyState::new(config)?;
66    tokio::spawn(await_backends(state.clone()));
67    core::serve_router(PROXY_NAME, &host, port, router(state)).await
68}
69
70fn router(state: ProxyState) -> Router {
71    Router::new()
72        .route(HEALTHCHECK_PATH, get(healthcheck))
73        .route(COMPLETIONS_PATH, post(completions))
74        .route(CHAT_COMPLETIONS_PATH, post(chat_completions))
75        .route(RESET_PREFIX_CACHE_PATH, post(flush_cache))
76        .route(PRIME_PREFIX_CACHE_PATH, post(prime_prefix_cache))
77        .with_state(state)
78}
79
80#[derive(Clone)]
81struct ProxyState {
82    inner: Arc<ProxyStateInner>,
83}
84
85struct ProxyStateInner {
86    client: reqwest::Client,
87    prefill: Vec<PrefillTarget>,
88    decode: Vec<String>,
89    ready: AtomicBool,
90    prefill_cursor: AtomicUsize,
91    decode_cursor: AtomicUsize,
92    room_seed: u64,
93    room_counter: AtomicU64,
94}
95
96impl ProxyState {
97    fn new(config: Config) -> Result<Self, ProxyError> {
98        core::require_endpoints(
99            PROXY_NAME,
100            config.prefill.is_empty(),
101            config.decode.is_empty(),
102        )?;
103        Ok(Self {
104            inner: Arc::new(ProxyStateInner {
105                client: core::pooled_client(PROXY_NAME)?,
106                prefill: config.prefill,
107                decode: config.decode,
108                ready: AtomicBool::new(false),
109                prefill_cursor: AtomicUsize::new(0),
110                decode_cursor: AtomicUsize::new(0),
111                room_seed: room_seed(),
112                room_counter: AtomicU64::new(0),
113            }),
114        })
115    }
116
117    fn client(&self) -> reqwest::Client {
118        self.inner.client.clone()
119    }
120
121    fn ready(&self) -> bool {
122        self.inner.ready.load(Ordering::SeqCst)
123    }
124
125    fn set_ready(&self) {
126        self.inner.ready.store(true, Ordering::SeqCst);
127    }
128
129    fn next_prefill(&self) -> PrefillTarget {
130        let index = core::round_robin_index(&self.inner.prefill_cursor, self.inner.prefill.len());
131        self.inner.prefill[index].clone()
132    }
133
134    fn next_decode(&self) -> String {
135        let index = core::round_robin_index(&self.inner.decode_cursor, self.inner.decode.len());
136        self.inner.decode[index].clone()
137    }
138
139    fn next_room(&self) -> u64 {
140        let counter = self.inner.room_counter.fetch_add(1, Ordering::SeqCst);
141        self.inner.room_seed.wrapping_add(counter) & ((1_u64 << 63) - 1)
142    }
143
144    /// The readiness-wait and reset/flush sweep targets: every prefill
145    /// replica URL followed by every decode URL.
146    fn fanout_target_urls(&self) -> Vec<String> {
147        core::fanout_target_urls(
148            self.inner.prefill.iter().map(|target| target.url.as_str()),
149            self.inner.decode.iter().map(String::as_str),
150        )
151    }
152}
153
154fn room_seed() -> u64 {
155    let nanos = SystemTime::now()
156        .duration_since(UNIX_EPOCH)
157        .map_or(0, |elapsed| elapsed.as_nanos() as u64);
158    (nanos ^ (u64::from(std::process::id()) << 32)) & ((1_u64 << 63) - 1)
159}
160
161async fn await_backends(state: ProxyState) {
162    let urls = state.fanout_target_urls();
163    core::await_backends(state.client(), urls, "/v1/models").await;
164    state.set_ready();
165}
166
167async fn healthcheck(
168    State(state): State<ProxyState>,
169) -> (StatusCode, Json<ProxyHealthcheckResponse>) {
170    core::healthcheck_response(
171        state.ready(),
172        state.inner.prefill.len(),
173        state.inner.decode.len(),
174    )
175}
176
177async fn completions(
178    State(state): State<ProxyState>,
179    headers: HeaderMap,
180    Json(body): Json<Value>,
181) -> Result<Response<Body>, ProxyHttpError> {
182    request_route(state, headers, body, COMPLETIONS_PATH).await
183}
184
185async fn chat_completions(
186    State(state): State<ProxyState>,
187    headers: HeaderMap,
188    Json(body): Json<Value>,
189) -> Result<Response<Body>, ProxyHttpError> {
190    request_route(state, headers, body, CHAT_COMPLETIONS_PATH).await
191}
192
193async fn request_route(
194    state: ProxyState,
195    headers: HeaderMap,
196    body: Value,
197    path: &'static str,
198) -> Result<Response<Body>, ProxyHttpError> {
199    if !state.ready() {
200        return Err(ProxyHttpError::status(
201            StatusCode::SERVICE_UNAVAILABLE,
202            "proxy is not ready",
203        ));
204    }
205
206    let stream = body.get("stream").and_then(Value::as_bool).unwrap_or(false);
207    let prefill = state.next_prefill();
208    let decode = state.next_decode();
209    let request_body = bootstrap_body(&body, &prefill, state.next_room(), path)?;
210    let authorization = outbound_authorization(&headers);
211    let client = state.client();
212
213    let prefill_url = join_path(&prefill.url, path);
214    let prefill_body = request_body.clone();
215    let prefill_authorization = authorization.clone();
216    let prefill_client = client.clone();
217    let prefill_task = tokio::spawn(async move {
218        let response = core::send_json_post(
219            prefill_client,
220            prefill_url,
221            &prefill_body,
222            None,
223            prefill_authorization.as_deref(),
224            &[],
225            "prefill request",
226        )
227        .await?;
228        drain_response(response, "prefill").await
229    });
230    let decode_result = core::send_json_post(
231        client,
232        join_path(&decode, path),
233        &request_body,
234        None,
235        authorization.as_deref(),
236        &[],
237        "decode request",
238    )
239    .await;
240    let decode_response = match decode_result {
241        Ok(response) => response,
242        Err(error) => {
243            // Dropping a Tokio join handle detaches the task. The public
244            // failure returns promptly while the prefill response continues
245            // to be drained in the background.
246            drop(prefill_task);
247            return Err(error);
248        }
249    };
250
251    if stream {
252        if prefill_task.is_finished() {
253            await_prefill(prefill_task).await?;
254            core::stream_response(decode_response)
255        } else if is_text_event_stream(&decode_response) {
256            stream_sse_decode_response(decode_response, prefill_task)
257        } else {
258            // Detach policy: the SGLang prefill/decode pair is coordinated over
259            // a bootstrap room, so aborting the prefill HTTP request mid-flight
260            // (on client disconnect) risks leaving the decode-side engine
261            // request waiting for KV that never arrives. Draining the detached
262            // prefill to completion lets both engine requests conclude cleanly.
263            core::stream_decode_response(decode_response, prefill_task, OnClientDrop::Detach)
264        }
265    } else {
266        let (prefill_result, decode_result) = tokio::join!(
267            await_prefill(prefill_task),
268            forward_response(decode_response)
269        );
270        prefill_result?;
271        decode_result
272    }
273}
274
275fn is_text_event_stream(response: &reqwest::Response) -> bool {
276    response
277        .headers()
278        .get(reqwest::header::CONTENT_TYPE)
279        .and_then(|value| value.to_str().ok())
280        .and_then(|value| value.split(';').next())
281        .is_some_and(|value| value.trim().eq_ignore_ascii_case("text/event-stream"))
282}
283
284fn stream_sse_decode_response(
285    response: reqwest::Response,
286    prefill_task: JoinHandle<Result<(), ProxyHttpError>>,
287) -> Result<Response<Body>, ProxyHttpError> {
288    let builder = core::upstream_response_builder(&response)?;
289    let stream = sse_decode_response_stream(response.bytes_stream(), prefill_task);
290    core::response_body(builder, Body::from_stream(stream))
291}
292
293fn sse_decode_response_stream<S, E>(
294    decode_stream: S,
295    prefill_task: JoinHandle<Result<(), ProxyHttpError>>,
296) -> impl Stream<Item = Result<Bytes, std::io::Error>>
297where
298    S: Stream<Item = Result<Bytes, E>> + Unpin,
299    E: fmt::Display,
300{
301    stream! {
302        let mut decode_stream = decode_stream;
303        let mut prefill_task = prefill_task;
304        let mut prefill_done = false;
305        let mut scanner = SseTerminalScanner::default();
306
307        loop {
308            if !prefill_done && prefill_task.is_finished() {
309                prefill_done = true;
310                let outcome = prefill_stream_outcome((&mut prefill_task).await);
311                if let Err(error) = outcome {
312                    yield Err(error);
313                    return;
314                }
315            }
316
317            tokio::select! {
318                prefill = &mut prefill_task, if !prefill_done => {
319                    prefill_done = true;
320                    let outcome = prefill_stream_outcome(prefill);
321                    if let Err(error) = outcome {
322                        yield Err(error);
323                        return;
324                    }
325                }
326                item = decode_stream.next() => match item {
327                    Some(Ok(chunk)) => {
328                        let scan = scanner.push(&chunk);
329                        if let Some(safe) = scan.safe {
330                            yield Ok(safe);
331                        }
332                        if let Some(terminal) = scan.terminal {
333                            let mut held = terminal.to_vec();
334                            loop {
335                                if !prefill_done && prefill_task.is_finished() {
336                                    prefill_done = true;
337                                    let outcome = prefill_stream_outcome((&mut prefill_task).await);
338                                    if let Err(error) = outcome {
339                                        yield Err(error);
340                                        return;
341                                    }
342                                }
343
344                                tokio::select! {
345                                    prefill = &mut prefill_task, if !prefill_done => {
346                                        prefill_done = true;
347                                        let outcome = prefill_stream_outcome(prefill);
348                                        if let Err(error) = outcome {
349                                            yield Err(error);
350                                            return;
351                                        }
352                                    }
353                                    item = decode_stream.next() => match item {
354                                        Some(Ok(chunk)) => held.extend_from_slice(&chunk),
355                                        Some(Err(error)) => {
356                                            yield Err(std::io::Error::other(format!(
357                                                "decode stream failed: {error}"
358                                            )));
359                                            return;
360                                        }
361                                        None => break,
362                                    }
363                                }
364                            }
365
366                            if !prefill_done {
367                                let outcome = prefill_stream_outcome((&mut prefill_task).await);
368                                if let Err(error) = outcome {
369                                    yield Err(error);
370                                    return;
371                                }
372                            }
373                            yield Ok(Bytes::from(held));
374                            return;
375                        }
376                    }
377                    Some(Err(error)) => {
378                        yield Err(std::io::Error::other(format!("decode stream failed: {error}")));
379                        return;
380                    }
381                    None => {
382                        let scan = scanner.finish();
383                        if let Some(safe) = scan.safe {
384                            yield Ok(safe);
385                        }
386                        if !prefill_done {
387                            let outcome = prefill_stream_outcome((&mut prefill_task).await);
388                            if let Err(error) = outcome {
389                                yield Err(error);
390                                return;
391                            }
392                        }
393                        if let Some(terminal) = scan.terminal {
394                            yield Ok(terminal);
395                        }
396                        return;
397                    }
398                }
399            }
400        }
401    }
402}
403
404fn prefill_stream_outcome(
405    outcome: Result<Result<(), ProxyHttpError>, tokio::task::JoinError>,
406) -> Result<(), std::io::Error> {
407    outcome
408        .map_err(|error| std::io::Error::other(format!("prefill task failed: {error}")))?
409        .map_err(|error| std::io::Error::other(error.to_string()))
410}
411
412#[derive(Default)]
413struct SseTerminalScanner {
414    pending: Vec<u8>,
415}
416
417impl SseTerminalScanner {
418    fn push(&mut self, chunk: &[u8]) -> SseScan {
419        self.pending.extend_from_slice(chunk);
420        let mut event_start = 0;
421        while let Some(relative_end) = sse_event_end(&self.pending[event_start..]) {
422            let event_end = event_start + relative_end;
423            if is_terminal_sse_event(&self.pending[event_start..event_end]) {
424                let terminal = self.pending.split_off(event_start);
425                let safe = std::mem::take(&mut self.pending);
426                return SseScan::new(safe, terminal);
427            }
428            event_start = event_end;
429        }
430        if event_start == 0 {
431            return SseScan::default();
432        }
433        let incomplete = self.pending.split_off(event_start);
434        let safe = std::mem::replace(&mut self.pending, incomplete);
435        SseScan::safe(safe)
436    }
437
438    fn finish(&mut self) -> SseScan {
439        let pending = std::mem::take(&mut self.pending);
440        if pending.is_empty() {
441            SseScan::default()
442        } else if is_terminal_sse_event(&pending) {
443            SseScan::terminal(pending)
444        } else {
445            SseScan::safe(pending)
446        }
447    }
448}
449
450#[derive(Default)]
451struct SseScan {
452    safe: Option<Bytes>,
453    terminal: Option<Bytes>,
454}
455
456impl SseScan {
457    fn new(safe: Vec<u8>, terminal: Vec<u8>) -> Self {
458        Self {
459            safe: (!safe.is_empty()).then(|| Bytes::from(safe)),
460            terminal: Some(Bytes::from(terminal)),
461        }
462    }
463
464    fn safe(bytes: Vec<u8>) -> Self {
465        Self {
466            safe: Some(Bytes::from(bytes)),
467            terminal: None,
468        }
469    }
470
471    fn terminal(bytes: Vec<u8>) -> Self {
472        Self {
473            safe: None,
474            terminal: Some(Bytes::from(bytes)),
475        }
476    }
477}
478
479fn sse_event_end(bytes: &[u8]) -> Option<usize> {
480    let mut line_start = 0;
481    for (index, byte) in bytes.iter().enumerate() {
482        if *byte != b'\n' {
483            continue;
484        }
485        let line_end = if index > line_start && bytes[index - 1] == b'\r' {
486            index - 1
487        } else {
488            index
489        };
490        if line_end == line_start {
491            return Some(index + 1);
492        }
493        line_start = index + 1;
494    }
495    None
496}
497
498fn is_terminal_sse_event(event: &[u8]) -> bool {
499    let mut data = Vec::new();
500    let mut saw_data = false;
501    for raw_line in event.split(|byte| *byte == b'\n') {
502        let line = raw_line.strip_suffix(b"\r").unwrap_or(raw_line);
503        let value = if line == b"data" {
504            Some(&b""[..])
505        } else if let Some(value) = line.strip_prefix(b"data:") {
506            Some(value.strip_prefix(b" ").unwrap_or(value))
507        } else {
508            None
509        };
510        if let Some(value) = value {
511            if saw_data {
512                data.push(b'\n');
513            }
514            data.extend_from_slice(value);
515            saw_data = true;
516        }
517    }
518    saw_data && data == b"[DONE]"
519}
520
521fn bootstrap_body(
522    body: &Value,
523    prefill: &PrefillTarget,
524    room: u64,
525    path: &'static str,
526) -> Result<Value, ProxyHttpError> {
527    let mut body = body.clone();
528    let object = body.as_object_mut().ok_or_else(|| {
529        ProxyHttpError::status(
530            StatusCode::BAD_REQUEST,
531            "OpenAI completion request body must be a JSON object",
532        )
533    })?;
534    if path == COMPLETIONS_PATH && object.get("prompt").is_some_and(Value::is_array) {
535        return Err(ProxyHttpError::status(
536            StatusCode::BAD_REQUEST,
537            "SGLang built-in proxy does not support prompt arrays",
538        ));
539    }
540    object.insert(
541        "bootstrap_host".to_owned(),
542        Value::String(prefill.bootstrap_host.clone()),
543    );
544    object.insert(
545        "bootstrap_port".to_owned(),
546        Value::from(prefill.bootstrap_port),
547    );
548    object.insert("bootstrap_room".to_owned(), Value::from(room));
549    Ok(body)
550}
551
552async fn drain_response(
553    response: reqwest::Response,
554    role: &'static str,
555) -> Result<(), ProxyHttpError> {
556    response.bytes().await.map_err(|error| {
557        ProxyHttpError::upstream(&format!("{role} response drain failed"), error)
558    })?;
559    Ok(())
560}
561
562async fn await_prefill(task: JoinHandle<Result<(), ProxyHttpError>>) -> Result<(), ProxyHttpError> {
563    task.await
564        .map_err(|error| ProxyHttpError::internal(format!("prefill task failed: {error}")))?
565}
566
567// Prefill replicas carry a config-issued static data-parallel size, so the
568// conditioning fan-out enumerates (replica, rank) targets over it.
569impl core::PrimeReplica for PrefillTarget {
570    fn url(&self) -> &str {
571        &self.url
572    }
573
574    fn data_parallel_size(&self) -> u32 {
575        self.data_parallel_size
576    }
577}
578
579/// One prefill/decode conditioning flow through the ordinary bootstrap
580/// pairing with the prefill request pinned to `rank`; the decode side rides
581/// the ordinary round-robin selection and is incidental coverage
582/// ([[RFC-0004:C-BENCH-CACHE-STATE]]).
583async fn prime_flow(
584    state: &ProxyState,
585    prefill: &PrefillTarget,
586    rank: u32,
587    authorization: Option<String>,
588    body: &Value,
589) -> Result<u16, core::PrimeFlowFailure> {
590    use core::PrimeFlowFailure;
591    let request_body = bootstrap_body(body, prefill, state.next_room(), COMPLETIONS_PATH)
592        .map_err(PrimeFlowFailure::transport)?;
593    let decode = state.next_decode();
594    let client = state.client();
595    let prefill_response = core::send_json_post_status(
596        client.clone(),
597        join_path(&prefill.url, COMPLETIONS_PATH),
598        &request_body,
599        None,
600        authorization.as_deref(),
601        &[("X-data-parallel-rank", rank.to_string())],
602        "prefill conditioning request",
603    )
604    .await
605    .map_err(PrimeFlowFailure::transport)?;
606    let (prefill_status, _prefill_text) =
607        core::expect_2xx("prefill conditioning", prefill_response).await?;
608    let decode_response = core::send_json_post_status(
609        client,
610        join_path(&decode, COMPLETIONS_PATH),
611        &request_body,
612        None,
613        authorization.as_deref(),
614        &[],
615        "decode conditioning request",
616    )
617    .await
618    .map_err(PrimeFlowFailure::transport)?;
619    core::expect_2xx("decode conditioning", decode_response).await?;
620    Ok(prefill_status)
621}
622
623async fn prime_prefix_cache(
624    State(state): State<ProxyState>,
625    headers: HeaderMap,
626    Json(body): Json<Value>,
627) -> Response<Body> {
628    if !state.ready() {
629        return ProxyHttpError::status(StatusCode::SERVICE_UNAVAILABLE, "proxy is not ready")
630            .into_response();
631    }
632    let authorization = outbound_authorization(&headers);
633    let targets = core::ranked_prime_targets(&state.inner.prefill);
634    core::run_prime_fanout("prefix cache conditioning", targets, |target| {
635        let state = state.clone();
636        let authorization = authorization.clone();
637        let body = body.clone();
638        async move { prime_flow(&state, &target.replica, target.rank, authorization, &body).await }
639    })
640    .await
641}
642
643async fn flush_cache(State(state): State<ProxyState>, headers: HeaderMap) -> Response<Body> {
644    let authorization = outbound_authorization(&headers);
645    let targets = state.fanout_target_urls();
646    core::run_sweep_fanout(
647        state.client(),
648        "cache flush",
649        "/flush_cache",
650        targets,
651        authorization,
652    )
653    .await
654}
655
656#[cfg(test)]
657mod tests {
658    use super::*;
659    use anyhow::{Context, Result, bail};
660    use axum::body::{Body, to_bytes};
661    use axum::extract::{Json, State};
662    use axum::http::{HeaderMap, HeaderValue, Response, StatusCode, header};
663    use axum::response::IntoResponse;
664    use axum::routing::{get, post};
665    use axum::{Router, serve};
666    use bytes::Bytes;
667    use futures_util::StreamExt;
668    use serde_json::{Value, json};
669    use std::sync::Arc;
670    use std::sync::atomic::{AtomicBool, AtomicU16, AtomicUsize, Ordering};
671    use std::time::Duration;
672    use tokio::net::TcpListener;
673    use tokio::sync::{Mutex, Notify};
674    use tokio::task::JoinHandle;
675
676    #[derive(Clone)]
677    struct BackendState {
678        completion_requests: Arc<Mutex<Vec<Value>>>,
679        chat_requests: Arc<Mutex<Vec<Value>>>,
680        completion_status: Arc<AtomicU16>,
681        completion_content_type: &'static str,
682        completion_chunks: Vec<Bytes>,
683        notify_on_request: Option<Arc<Notify>>,
684        wait_before_response: Option<Arc<Notify>>,
685        gate_after_first_chunk: Option<Arc<Notify>>,
686        body_error: bool,
687        body_polled: Arc<AtomicBool>,
688        flush_status: Arc<AtomicU16>,
689        flush_requests: Arc<AtomicUsize>,
690    }
691
692    impl BackendState {
693        fn new(body: &'static [u8]) -> Self {
694            Self {
695                completion_requests: Arc::new(Mutex::new(Vec::new())),
696                chat_requests: Arc::new(Mutex::new(Vec::new())),
697                completion_status: Arc::new(AtomicU16::new(StatusCode::OK.as_u16())),
698                completion_content_type: "application/json",
699                completion_chunks: vec![Bytes::from_static(body)],
700                notify_on_request: None,
701                wait_before_response: None,
702                gate_after_first_chunk: None,
703                body_error: false,
704                body_polled: Arc::new(AtomicBool::new(false)),
705                flush_status: Arc::new(AtomicU16::new(StatusCode::OK.as_u16())),
706                flush_requests: Arc::new(AtomicUsize::new(0)),
707            }
708        }
709    }
710
711    async fn mock_completion(
712        State(state): State<BackendState>,
713        Json(body): Json<Value>,
714    ) -> Response<Body> {
715        state.completion_requests.lock().await.push(body);
716        if let Some(notify) = &state.notify_on_request {
717            notify.notify_one();
718        }
719        if let Some(wait) = &state.wait_before_response {
720            wait.notified().await;
721        }
722        let chunks = state.completion_chunks.clone();
723        let body_polled = state.body_polled.clone();
724        let gate = state.gate_after_first_chunk.clone();
725        let stream = async_stream::stream! {
726            body_polled.store(true, Ordering::SeqCst);
727            for (index, chunk) in chunks.into_iter().enumerate() {
728                yield Ok::<Bytes, std::io::Error>(chunk);
729                if index == 0
730                    && let Some(gate) = &gate
731                {
732                    gate.notified().await;
733                }
734            }
735            if state.body_error {
736                yield Err(std::io::Error::other("mock response body failed"));
737            }
738        };
739        let status = StatusCode::from_u16(state.completion_status.load(Ordering::SeqCst))
740            .unwrap_or(StatusCode::INTERNAL_SERVER_ERROR);
741        let mut response = Response::new(Body::from_stream(stream));
742        *response.status_mut() = status;
743        response.headers_mut().insert(
744            header::CONTENT_TYPE,
745            HeaderValue::from_static(state.completion_content_type),
746        );
747        response
748    }
749
750    async fn mock_chat_completion(
751        State(state): State<BackendState>,
752        Json(body): Json<Value>,
753    ) -> Response<Body> {
754        state.chat_requests.lock().await.push(body);
755        (StatusCode::OK, Json(json!({"route": "chat"}))).into_response()
756    }
757
758    async fn mock_flush(State(state): State<BackendState>) -> Response<Body> {
759        state.flush_requests.fetch_add(1, Ordering::SeqCst);
760        let status = StatusCode::from_u16(state.flush_status.load(Ordering::SeqCst))
761            .unwrap_or(StatusCode::INTERNAL_SERVER_ERROR);
762        (status, "flush").into_response()
763    }
764
765    async fn spawn_backend(state: BackendState) -> Result<(String, JoinHandle<()>)> {
766        let app = Router::new()
767            .route(
768                "/health",
769                get(|| async { StatusCode::INTERNAL_SERVER_ERROR }),
770            )
771            .route("/v1/models", get(|| async { StatusCode::OK }))
772            .route("/v1/completions", post(mock_completion))
773            .route("/v1/chat/completions", post(mock_chat_completion))
774            .route("/flush_cache", post(mock_flush))
775            .with_state(state);
776        let listener = TcpListener::bind(("127.0.0.1", 0)).await?;
777        let address = listener.local_addr()?;
778        let handle = tokio::spawn(async move {
779            let _result = serve(listener, app).await;
780        });
781        Ok((format!("http://{address}"), handle))
782    }
783
784    fn proxy_state(prefill_url: String, decode_url: String) -> Result<ProxyState> {
785        ProxyState::new(Config {
786            host: "127.0.0.1".to_owned(),
787            port: 8000,
788            prefill: vec![PrefillTarget {
789                url: prefill_url,
790                bootstrap_host: "10.0.0.7".to_owned(),
791                bootstrap_port: 8998,
792                data_parallel_size: 1,
793            }],
794            decode: vec![decode_url],
795        })
796        .map_err(Into::into)
797    }
798
799    #[tokio::test]
800    async fn non_streaming_completion_dispatches_both_roles_and_drains_prefill() -> Result<()> {
801        let decode_seen = Arc::new(Notify::new());
802        let mut prefill_backend = BackendState::new(br#"{"prefill":true}"#);
803        prefill_backend.wait_before_response = Some(decode_seen.clone());
804        let mut decode_backend = BackendState::new(br#"{"decode":true}"#);
805        decode_backend.notify_on_request = Some(decode_seen);
806        decode_backend
807            .completion_status
808            .store(StatusCode::CREATED.as_u16(), Ordering::SeqCst);
809
810        let (prefill_url, prefill_server) = spawn_backend(prefill_backend.clone()).await?;
811        let (decode_url, decode_server) = spawn_backend(decode_backend.clone()).await?;
812        let state = proxy_state(prefill_url, decode_url)?;
813        state.set_ready();
814
815        let response = tokio::time::timeout(
816            std::time::Duration::from_secs(2),
817            request_route(
818                state,
819                HeaderMap::new(),
820                json!({"model": "m", "prompt": "hello"}),
821                COMPLETIONS_PATH,
822            ),
823        )
824        .await
825        .context("prefill waited for decode instead of both requests being initiated")?
826        .map_err(|error| anyhow::anyhow!(error.to_string()))?;
827
828        assert_eq!(response.status(), StatusCode::CREATED);
829        assert_eq!(
830            response.headers().get(header::CONTENT_TYPE),
831            Some(&HeaderValue::from_static("application/json"))
832        );
833        assert_eq!(
834            to_bytes(response.into_body(), usize::MAX).await?,
835            Bytes::from_static(br#"{"decode":true}"#)
836        );
837        assert!(prefill_backend.body_polled.load(Ordering::SeqCst));
838
839        let prefill_requests = prefill_backend.completion_requests.lock().await;
840        let decode_requests = decode_backend.completion_requests.lock().await;
841        assert_eq!(prefill_requests.len(), 1);
842        assert_eq!(*prefill_requests, *decode_requests);
843        assert_eq!(
844            prefill_requests[0]
845                .get("bootstrap_host")
846                .and_then(Value::as_str),
847            Some("10.0.0.7")
848        );
849        assert_eq!(
850            prefill_requests[0]
851                .get("bootstrap_port")
852                .and_then(Value::as_u64),
853            Some(8998)
854        );
855        assert!(
856            prefill_requests[0]
857                .get("bootstrap_room")
858                .and_then(Value::as_u64)
859                .is_some()
860        );
861        prefill_server.abort();
862        decode_server.abort();
863        Ok(())
864    }
865
866    #[tokio::test]
867    async fn chat_dispatch_preserves_messages_and_unowned_fields_on_both_roles() -> Result<()> {
868        let prefill_backend = BackendState::new(b"prefill");
869        let decode_backend = BackendState::new(b"decode");
870        let (prefill_url, prefill_server) = spawn_backend(prefill_backend.clone()).await?;
871        let (decode_url, decode_server) = spawn_backend(decode_backend.clone()).await?;
872        let state = proxy_state(prefill_url, decode_url)?;
873        state.set_ready();
874        let request = json!({
875            "model": "m",
876            "messages": [{"role": "user", "content": "hello"}],
877            "temperature": 1.0,
878            "reasoning_effort": "high",
879            "chat_template_kwargs": {"enable_thinking": true}
880        });
881
882        let response = request_route(
883            state,
884            HeaderMap::new(),
885            request.clone(),
886            CHAT_COMPLETIONS_PATH,
887        )
888        .await
889        .map_err(|error| anyhow::anyhow!(error.to_string()))?;
890        assert_eq!(response.status(), StatusCode::OK);
891
892        let prefill_requests = prefill_backend.chat_requests.lock().await;
893        let decode_requests = decode_backend.chat_requests.lock().await;
894        assert_eq!(prefill_requests.len(), 1);
895        assert_eq!(*prefill_requests, *decode_requests);
896        for (key, value) in request.as_object().context("request was not an object")? {
897            assert_eq!(prefill_requests[0].get(key), Some(value), "changed {key}");
898        }
899        assert!(prefill_requests[0]["bootstrap_room"].is_u64());
900        assert!(prefill_backend.completion_requests.lock().await.is_empty());
901        assert!(decode_backend.completion_requests.lock().await.is_empty());
902        prefill_server.abort();
903        decode_server.abort();
904        Ok(())
905    }
906
907    #[tokio::test]
908    async fn prompt_array_is_rejected_before_role_dispatch() -> Result<()> {
909        let prefill_backend = BackendState::new(b"prefill");
910        let decode_backend = BackendState::new(b"decode");
911        let (prefill_url, prefill_server) = spawn_backend(prefill_backend.clone()).await?;
912        let (decode_url, decode_server) = spawn_backend(decode_backend.clone()).await?;
913        let state = proxy_state(prefill_url, decode_url)?;
914        state.set_ready();
915
916        let result = request_route(
917            state,
918            HeaderMap::new(),
919            json!({"model": "m", "prompt": ["one", "two"]}),
920            COMPLETIONS_PATH,
921        )
922        .await;
923        let error = match result {
924            Ok(_) => bail!("prompt arrays must fail"),
925            Err(error) => error,
926        };
927
928        assert_eq!(error.into_response().status(), StatusCode::BAD_REQUEST);
929        assert!(prefill_backend.completion_requests.lock().await.is_empty());
930        assert!(decode_backend.completion_requests.lock().await.is_empty());
931        prefill_server.abort();
932        decode_server.abort();
933        Ok(())
934    }
935
936    #[tokio::test]
937    async fn streaming_completion_relays_decode_chunks_incrementally() -> Result<()> {
938        let release_prefill = Arc::new(Notify::new());
939        let mut prefill_backend = BackendState::new(b"prefill");
940        prefill_backend.wait_before_response = Some(release_prefill.clone());
941        let second_chunk = Arc::new(Notify::new());
942        let mut decode_backend = BackendState::new(b"");
943        decode_backend.completion_content_type = "text/event-stream";
944        decode_backend.completion_chunks = vec![
945            Bytes::from_static(b"data: first\n\n"),
946            Bytes::from_static(b"data: second\n\n"),
947            Bytes::from_static(b"data: [DO"),
948            Bytes::from_static(b"NE]\r"),
949            Bytes::from_static(b"\n\r"),
950            Bytes::from_static(b"\n"),
951        ];
952        decode_backend.gate_after_first_chunk = Some(second_chunk.clone());
953        let (prefill_url, prefill_server) = spawn_backend(prefill_backend).await?;
954        let (decode_url, decode_server) = spawn_backend(decode_backend).await?;
955        let state = proxy_state(prefill_url, decode_url)?;
956        state.set_ready();
957
958        let response = tokio::time::timeout(
959            std::time::Duration::from_secs(1),
960            request_route(
961                state,
962                HeaderMap::new(),
963                json!({"model": "m", "prompt": "hello", "stream": true}),
964                COMPLETIONS_PATH,
965            ),
966        )
967        .await
968        .context("a delayed prefill response blocked decode streaming")?
969        .map_err(|error| anyhow::anyhow!(error.to_string()))?;
970        assert_eq!(
971            response.headers().get(header::CONTENT_TYPE),
972            Some(&HeaderValue::from_static("text/event-stream"))
973        );
974        let mut stream = response.into_body().into_data_stream();
975        let first = tokio::time::timeout(std::time::Duration::from_secs(1), stream.next())
976            .await?
977            .context("decode stream ended before its first chunk")??;
978        assert_eq!(first, Bytes::from_static(b"data: first\n\n"));
979        release_prefill.notify_one();
980        second_chunk.notify_one();
981        let second = stream
982            .next()
983            .await
984            .context("decode stream ended before its second chunk")??;
985        assert_eq!(second, Bytes::from_static(b"data: second\n\n"));
986        let terminal = stream
987            .next()
988            .await
989            .context("decode stream ended before its terminal event")??;
990        assert_eq!(terminal, Bytes::from_static(b"data: [DONE]\r\n\r\n"));
991        prefill_server.abort();
992        decode_server.abort();
993        Ok(())
994    }
995
996    #[tokio::test]
997    async fn late_prefill_failure_prevents_terminal_sse_event() -> Result<()> {
998        let release_prefill_failure = Arc::new(Notify::new());
999        let mut prefill_backend = BackendState::new(b"prefill");
1000        prefill_backend.body_error = true;
1001        prefill_backend.gate_after_first_chunk = Some(release_prefill_failure.clone());
1002
1003        let release_terminal = Arc::new(Notify::new());
1004        let mut decode_backend = BackendState::new(b"");
1005        decode_backend.completion_content_type = "text/event-stream; charset=utf-8";
1006        decode_backend.completion_chunks = vec![
1007            Bytes::from_static(b"data: first\r\n\r\n"),
1008            Bytes::from_static(b"data: [DO"),
1009            Bytes::from_static(b"NE]\r"),
1010            Bytes::from_static(b"\n\r"),
1011            Bytes::from_static(b"\n"),
1012        ];
1013        decode_backend.gate_after_first_chunk = Some(release_terminal.clone());
1014
1015        let (prefill_url, prefill_server) = spawn_backend(prefill_backend).await?;
1016        let (decode_url, decode_server) = spawn_backend(decode_backend).await?;
1017        let state = proxy_state(prefill_url, decode_url)?;
1018        state.set_ready();
1019
1020        let response = request_route(
1021            state,
1022            HeaderMap::new(),
1023            json!({"model": "m", "prompt": "hello", "stream": true}),
1024            COMPLETIONS_PATH,
1025        )
1026        .await
1027        .map_err(|error| anyhow::anyhow!(error.to_string()))?;
1028        assert_eq!(response.status(), StatusCode::OK);
1029        assert_eq!(
1030            response.headers().get(header::CONTENT_TYPE),
1031            Some(&HeaderValue::from_static(
1032                "text/event-stream; charset=utf-8"
1033            ))
1034        );
1035
1036        let mut stream = response.into_body().into_data_stream();
1037        let first = stream
1038            .next()
1039            .await
1040            .context("decode stream ended before its first event")??;
1041        assert_eq!(first, Bytes::from_static(b"data: first\r\n\r\n"));
1042
1043        release_terminal.notify_one();
1044        assert!(
1045            tokio::time::timeout(Duration::from_millis(50), stream.next())
1046                .await
1047                .is_err(),
1048            "terminal SSE bytes were forwarded before prefill completed"
1049        );
1050
1051        release_prefill_failure.notify_one();
1052        let result = stream
1053            .next()
1054            .await
1055            .context("stream completed after a late prefill failure")?;
1056        let error = match result {
1057            Ok(bytes) => bail!(
1058                "late prefill failure forwarded terminal bytes: {:?}",
1059                String::from_utf8_lossy(&bytes)
1060            ),
1061            Err(error) => error,
1062        };
1063        assert!(error.to_string().contains("prefill response drain failed"));
1064
1065        prefill_server.abort();
1066        decode_server.abort();
1067        Ok(())
1068    }
1069
1070    #[tokio::test]
1071    async fn prefill_failure_before_headers_returns_non_success() -> Result<()> {
1072        let prefill_backend = BackendState::new(b"prefill failed");
1073        prefill_backend
1074            .completion_status
1075            .store(StatusCode::INTERNAL_SERVER_ERROR.as_u16(), Ordering::SeqCst);
1076        let release_decode = Arc::new(Notify::new());
1077        let mut decode_backend = BackendState::new(b"decode");
1078        decode_backend.wait_before_response = Some(release_decode.clone());
1079        let (prefill_url, prefill_server) = spawn_backend(prefill_backend.clone()).await?;
1080        let (decode_url, decode_server) = spawn_backend(decode_backend.clone()).await?;
1081        let state = proxy_state(prefill_url, decode_url)?;
1082        state.set_ready();
1083
1084        let request = tokio::spawn(request_route(
1085            state,
1086            HeaderMap::new(),
1087            json!({"model": "m", "prompt": "hello", "stream": true}),
1088            COMPLETIONS_PATH,
1089        ));
1090        wait_until(&prefill_backend.body_polled).await?;
1091        tokio::task::yield_now().await;
1092        release_decode.notify_one();
1093        let error = match request.await? {
1094            Ok(_) => bail!("prefill failure must fail before public headers"),
1095            Err(error) => error,
1096        };
1097        assert_eq!(error.into_response().status(), StatusCode::BAD_GATEWAY);
1098        assert_eq!(decode_backend.completion_requests.lock().await.len(), 1);
1099        prefill_server.abort();
1100        decode_server.abort();
1101        Ok(())
1102    }
1103
1104    #[tokio::test]
1105    async fn prefill_body_failure_before_headers_returns_non_success() -> Result<()> {
1106        let mut prefill_backend = BackendState::new(b"prefill");
1107        prefill_backend.body_error = true;
1108        let release_decode = Arc::new(Notify::new());
1109        let mut decode_backend = BackendState::new(b"decode");
1110        decode_backend.wait_before_response = Some(release_decode.clone());
1111        let (prefill_url, prefill_server) = spawn_backend(prefill_backend.clone()).await?;
1112        let (decode_url, decode_server) = spawn_backend(decode_backend).await?;
1113        let state = proxy_state(prefill_url, decode_url)?;
1114        state.set_ready();
1115
1116        let request = tokio::spawn(request_route(
1117            state,
1118            HeaderMap::new(),
1119            json!({"model": "m", "prompt": "hello", "stream": true}),
1120            COMPLETIONS_PATH,
1121        ));
1122        wait_until(&prefill_backend.body_polled).await?;
1123        tokio::task::yield_now().await;
1124        release_decode.notify_one();
1125        let error = match request.await? {
1126            Ok(_) => bail!("prefill body failure must fail before public headers"),
1127            Err(error) => error,
1128        };
1129        assert_eq!(error.into_response().status(), StatusCode::BAD_GATEWAY);
1130        prefill_server.abort();
1131        decode_server.abort();
1132        Ok(())
1133    }
1134
1135    #[tokio::test]
1136    async fn decode_failure_does_not_wait_for_slow_prefill() -> Result<()> {
1137        let prefill_seen = Arc::new(Notify::new());
1138        let release_prefill = Arc::new(Notify::new());
1139        let mut prefill_backend = BackendState::new(b"prefill");
1140        prefill_backend.notify_on_request = Some(prefill_seen.clone());
1141        prefill_backend.wait_before_response = Some(release_prefill.clone());
1142        let mut decode_backend = BackendState::new(b"decode failed");
1143        decode_backend.wait_before_response = Some(prefill_seen);
1144        decode_backend
1145            .completion_status
1146            .store(StatusCode::INTERNAL_SERVER_ERROR.as_u16(), Ordering::SeqCst);
1147        let (prefill_url, prefill_server) = spawn_backend(prefill_backend.clone()).await?;
1148        let (decode_url, decode_server) = spawn_backend(decode_backend).await?;
1149        let state = proxy_state(prefill_url, decode_url)?;
1150        state.set_ready();
1151
1152        let result = tokio::time::timeout(
1153            std::time::Duration::from_secs(1),
1154            request_route(
1155                state,
1156                HeaderMap::new(),
1157                json!({"model": "m", "prompt": "hello", "stream": true}),
1158                COMPLETIONS_PATH,
1159            ),
1160        )
1161        .await
1162        .context("decode failure waited for a slow prefill response")?;
1163        let error = match result {
1164            Ok(_) => bail!("decode failure must return a non-success response"),
1165            Err(error) => error,
1166        };
1167        assert_eq!(error.into_response().status(), StatusCode::BAD_GATEWAY);
1168        assert_eq!(prefill_backend.completion_requests.lock().await.len(), 1);
1169        release_prefill.notify_one();
1170        prefill_server.abort();
1171        decode_server.abort();
1172        Ok(())
1173    }
1174
1175    #[tokio::test]
1176    async fn dropping_public_stream_detaches_prefill_drain() -> Result<()> {
1177        let release_prefill = Arc::new(Notify::new());
1178        let prefill_drained = Arc::new(AtomicBool::new(false));
1179        let drained = prefill_drained.clone();
1180        let release = release_prefill.clone();
1181        let prefill = tokio::spawn(async move {
1182            release.notified().await;
1183            drained.store(true, Ordering::SeqCst);
1184            Ok::<(), ProxyHttpError>(())
1185        });
1186        let decode = Box::pin(
1187            futures_util::stream::once(async {
1188                Ok::<Bytes, std::io::Error>(Bytes::from_static(b"data: first\n\n"))
1189            })
1190            .chain(futures_util::stream::pending()),
1191        );
1192        let mut stream = Box::pin(sse_decode_response_stream(decode, prefill));
1193
1194        assert!(matches!(stream.next().await, Some(Ok(_))));
1195        drop(stream);
1196        release_prefill.notify_one();
1197        wait_until(&prefill_drained).await?;
1198        Ok(())
1199    }
1200
1201    #[tokio::test]
1202    async fn decode_failure_after_streaming_starts_fails_the_public_stream() -> Result<()> {
1203        let prefill_backend = BackendState::new(b"prefill");
1204        let release_error = Arc::new(Notify::new());
1205        let mut decode_backend = BackendState::new(b"");
1206        decode_backend.completion_content_type = "text/event-stream";
1207        decode_backend.completion_chunks = vec![
1208            Bytes::from_static(b"data: first\n\n"),
1209            Bytes::from_static(b"data: [DONE]\n\n"),
1210        ];
1211        decode_backend.body_error = true;
1212        decode_backend.gate_after_first_chunk = Some(release_error.clone());
1213        let (prefill_url, prefill_server) = spawn_backend(prefill_backend).await?;
1214        let (decode_url, decode_server) = spawn_backend(decode_backend).await?;
1215        let state = proxy_state(prefill_url, decode_url)?;
1216        state.set_ready();
1217
1218        let response = request_route(
1219            state,
1220            HeaderMap::new(),
1221            json!({"model": "m", "prompt": "hello", "stream": true}),
1222            COMPLETIONS_PATH,
1223        )
1224        .await
1225        .map_err(|error| anyhow::anyhow!(error.to_string()))?;
1226        let mut stream = response.into_body().into_data_stream();
1227        assert!(matches!(stream.next().await, Some(Ok(_))));
1228        release_error.notify_one();
1229        let result = stream
1230            .next()
1231            .await
1232            .context("decode stream ended successfully after an upstream body failure")?;
1233        let error = match result {
1234            Ok(_) => bail!("decode body failure must fail the public stream"),
1235            Err(error) => error,
1236        };
1237        assert!(error.to_string().contains("decode stream failed"));
1238        prefill_server.abort();
1239        decode_server.abort();
1240        Ok(())
1241    }
1242
1243    #[tokio::test]
1244    async fn terminal_sse_waits_for_clean_decode_eof() -> Result<()> {
1245        let decode = Box::pin(futures_util::stream::iter(vec![
1246            Ok::<Bytes, std::io::Error>(Bytes::from_static(b"data: first\n\n")),
1247            Ok(Bytes::from_static(b"data: [DONE]\n\n")),
1248            Err(std::io::Error::other("decode failed after terminal event")),
1249        ]));
1250        let prefill = tokio::spawn(async { Ok::<(), ProxyHttpError>(()) });
1251        let mut stream = Box::pin(sse_decode_response_stream(decode, prefill));
1252
1253        let first = stream
1254            .next()
1255            .await
1256            .context("stream ended before its first event")??;
1257        assert_eq!(first, Bytes::from_static(b"data: first\n\n"));
1258        let result = stream
1259            .next()
1260            .await
1261            .context("stream completed after a late decode failure")?;
1262        let error = match result {
1263            Ok(bytes) => bail!(
1264                "decode failure forwarded terminal bytes: {:?}",
1265                String::from_utf8_lossy(&bytes)
1266            ),
1267            Err(error) => error,
1268        };
1269        assert!(error.to_string().contains("decode stream failed"));
1270        Ok(())
1271    }
1272
1273    async fn wait_until(flag: &AtomicBool) -> Result<()> {
1274        tokio::time::timeout(std::time::Duration::from_secs(1), async {
1275            while !flag.load(Ordering::SeqCst) {
1276                tokio::task::yield_now().await;
1277            }
1278        })
1279        .await
1280        .context("expected asynchronous condition was not observed")?;
1281        Ok(())
1282    }
1283
1284    #[tokio::test]
1285    async fn flush_cache_attempts_all_targets_and_reports_partial_failure() -> Result<()> {
1286        let prefill_backend = BackendState::new(b"prefill");
1287        let decode_backend = BackendState::new(b"decode");
1288        let (prefill_url, prefill_server) = spawn_backend(prefill_backend.clone()).await?;
1289        let (decode_url, decode_server) = spawn_backend(decode_backend.clone()).await?;
1290        let state = proxy_state(prefill_url, decode_url)?;
1291
1292        let all_succeeded = flush_cache(State(state.clone()), HeaderMap::new()).await;
1293        assert_eq!(all_succeeded.status(), StatusCode::OK);
1294
1295        decode_backend
1296            .flush_status
1297            .store(StatusCode::PARTIAL_CONTENT.as_u16(), Ordering::SeqCst);
1298        let partial = flush_cache(State(state), HeaderMap::new()).await;
1299        assert_eq!(partial.status(), StatusCode::PARTIAL_CONTENT);
1300        let body: Value =
1301            serde_json::from_slice(&to_bytes(partial.into_body(), usize::MAX).await?)?;
1302        assert_eq!(body["successful"].as_array().map(Vec::len), Some(1));
1303        assert_eq!(body["failed"].as_array().map(Vec::len), Some(1));
1304        assert_eq!(prefill_backend.flush_requests.load(Ordering::SeqCst), 2);
1305        assert_eq!(decode_backend.flush_requests.load(Ordering::SeqCst), 2);
1306        prefill_server.abort();
1307        decode_server.abort();
1308        Ok(())
1309    }
1310
1311    #[tokio::test]
1312    async fn healthcheck_is_unsuccessful_until_all_backends_are_ready() -> Result<()> {
1313        let prefill_backend = BackendState::new(b"prefill");
1314        let decode_backend = BackendState::new(b"decode");
1315        let (prefill_url, prefill_server) = spawn_backend(prefill_backend).await?;
1316        let (decode_url, decode_server) = spawn_backend(decode_backend).await?;
1317        let state = proxy_state(prefill_url, decode_url)?;
1318        let (status, Json(body)) = healthcheck(State(state.clone())).await;
1319        assert_eq!(status, StatusCode::SERVICE_UNAVAILABLE);
1320        assert!(!body.ready);
1321
1322        tokio::time::timeout(
1323            Duration::from_secs(1),
1324            tokio::spawn(await_backends(state.clone())),
1325        )
1326        .await
1327        .context("backend readiness did not use the responsive model endpoint")??;
1328        let (status, Json(body)) = healthcheck(State(state)).await;
1329        assert_eq!(status, StatusCode::OK);
1330        assert!(body.ready);
1331        prefill_server.abort();
1332        decode_server.abort();
1333        Ok(())
1334    }
1335
1336    #[tokio::test]
1337    async fn prime_prefix_cache_fans_out_to_each_prefill_rank_and_reports_partial_failure()
1338    -> Result<()> {
1339        let prefill_backend = PrimeBackend::default();
1340        let decode_backend = PrimeBackend::default();
1341        let (prefill_url, prefill_server) = spawn_prime_backend(prefill_backend.clone()).await?;
1342        let (decode_url, decode_server) = spawn_prime_backend(decode_backend.clone()).await?;
1343        let state = ProxyState::new(Config {
1344            host: "127.0.0.1".to_owned(),
1345            port: 8000,
1346            prefill: vec![PrefillTarget {
1347                url: prefill_url.clone(),
1348                bootstrap_host: "10.0.0.7".to_owned(),
1349                bootstrap_port: 8998,
1350                data_parallel_size: 2,
1351            }],
1352            decode: vec![decode_url],
1353        })?;
1354        state.set_ready();
1355        let conditioning =
1356            || Json(json!({"model": "m", "prompt": "canonical prefix", "max_tokens": 1}));
1357
1358        let response =
1359            prime_prefix_cache(State(state.clone()), HeaderMap::new(), conditioning()).await;
1360        assert_eq!(response.status(), StatusCode::OK);
1361        let body: Value =
1362            serde_json::from_slice(&to_bytes(response.into_body(), usize::MAX).await?)?;
1363        let targets = body["targets"]
1364            .as_array()
1365            .ok_or_else(|| anyhow::anyhow!("fan-out response has no targets"))?;
1366        assert_eq!(targets.len(), 2);
1367        for (rank, target) in targets.iter().enumerate() {
1368            assert_eq!(target["url"].as_str(), Some(prefill_url.as_str()));
1369            assert_eq!(target["rank"].as_u64(), Some(rank as u64));
1370            assert_eq!(target["http_status"].as_u64(), Some(200));
1371            assert!(target["error"].is_null());
1372        }
1373        let prefill_requests = prefill_backend.requests.lock().await;
1374        assert_eq!(prefill_requests.len(), 2);
1375        assert_eq!(prefill_requests[0].0.as_deref(), Some("0"));
1376        assert_eq!(prefill_requests[1].0.as_deref(), Some("1"));
1377        // Each flow rides the ordinary bootstrap pairing.
1378        assert_eq!(
1379            prefill_requests[0].1["bootstrap_host"].as_str(),
1380            Some("10.0.0.7")
1381        );
1382        let decode_requests = decode_backend.requests.lock().await;
1383        assert_eq!(decode_requests.len(), 2);
1384        assert!(decode_requests.iter().all(|(rank, _)| rank.is_none()));
1385        drop(prefill_requests);
1386        drop(decode_requests);
1387
1388        prefill_backend.set_fail_rank(Some("1".to_owned())).await;
1389        let partial = prime_prefix_cache(State(state), HeaderMap::new(), conditioning()).await;
1390        assert_eq!(partial.status(), StatusCode::PARTIAL_CONTENT);
1391        let body: Value =
1392            serde_json::from_slice(&to_bytes(partial.into_body(), usize::MAX).await?)?;
1393        let targets = body["targets"]
1394            .as_array()
1395            .ok_or_else(|| anyhow::anyhow!("fan-out response has no targets"))?;
1396        assert_eq!(targets.len(), 2);
1397        assert!(targets[0]["error"].is_null());
1398        assert_eq!(targets[1]["rank"].as_u64(), Some(1));
1399        assert_eq!(targets[1]["http_status"].as_u64(), Some(500));
1400        assert!(
1401            targets[1]["error"]
1402                .as_str()
1403                .is_some_and(|error| error.contains("HTTP 500"))
1404        );
1405        prefill_server.abort();
1406        decode_server.abort();
1407        Ok(())
1408    }
1409
1410    type PrimeRequests = Arc<Mutex<Vec<(Option<String>, Value)>>>;
1411
1412    #[derive(Clone, Default)]
1413    struct PrimeBackend {
1414        requests: PrimeRequests,
1415        fail_rank: Arc<Mutex<Option<String>>>,
1416    }
1417
1418    impl PrimeBackend {
1419        async fn set_fail_rank(&self, rank: Option<String>) {
1420            *self.fail_rank.lock().await = rank;
1421        }
1422    }
1423
1424    async fn mock_prime(
1425        State(state): State<PrimeBackend>,
1426        headers: HeaderMap,
1427        Json(body): Json<Value>,
1428    ) -> Response<Body> {
1429        let rank = headers
1430            .get("x-data-parallel-rank")
1431            .and_then(|value| value.to_str().ok())
1432            .map(str::to_owned);
1433        let fail_rank = state.fail_rank.lock().await.clone();
1434        state.requests.lock().await.push((rank.clone(), body));
1435        if fail_rank.is_some() && fail_rank == rank {
1436            return StatusCode::INTERNAL_SERVER_ERROR.into_response();
1437        }
1438        Json(json!({"object": "text_completion", "choices": []})).into_response()
1439    }
1440
1441    async fn spawn_prime_backend(state: PrimeBackend) -> Result<(String, JoinHandle<()>)> {
1442        let app = Router::new()
1443            .route("/v1/completions", post(mock_prime))
1444            .with_state(state);
1445        let listener = TcpListener::bind("127.0.0.1:0").await?;
1446        let address = listener.local_addr()?;
1447        let server = tokio::spawn(async move {
1448            let _result = serve(listener, app).await;
1449        });
1450        Ok((format!("http://{address}"), server))
1451    }
1452}