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 = 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 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}