1use 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 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}