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