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