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