1use crate::core::{
2 self, OnClientDrop, ProxyHealthcheckResponse, ProxyHttpError, forward_response, join_path,
3 outbound_authorization,
4};
5use crate::error::ProxyError;
6use axum::Router;
7use axum::body::Body;
8use axum::extract::{Json, State};
9use axum::http::{HeaderMap, Response, StatusCode};
10use axum::response::IntoResponse;
11use axum::routing::{get, post};
12use serde::Serialize;
13use serde_json::Value;
14use std::sync::Arc;
15use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
16use tokio::sync::RwLock;
17
18pub const VERSION: u32 = 1;
19
20pub const HEALTHCHECK_PATH: &str = "/healthcheck";
21pub const RESET_PREFIX_CACHE_PATH: &str = "/reset_prefix_cache";
22pub const PRIME_PREFIX_CACHE_PATH: &str = "/prime_prefix_cache";
23
24pub const COMPLETIONS_PATH: &str = "/v1/completions";
25pub const CHAT_COMPLETIONS_PATH: &str = "/v1/chat/completions";
26
27const PROXY_NAME: &str = "vLLM Mooncake proxy";
29
30#[derive(Clone, Debug)]
31pub struct Config {
32 pub host: String,
33 pub port: u16,
34 pub prefill: Vec<PrefillTarget>,
35 pub decode: Vec<String>,
36}
37
38#[derive(Clone, Debug)]
39pub struct PrefillTarget {
40 pub url: String,
41 pub bootstrap_url: String,
42}
43
44pub fn run(config: Config) -> Result<(), ProxyError> {
45 core::run(|| run_async(config))
46}
47
48pub async fn run_async(config: Config) -> Result<(), ProxyError> {
49 let host = config.host.clone();
50 let port = config.port;
51 let state = ProxyState::new(config)?;
52 tokio::spawn(discover_prefillers(state.clone()));
53 core::serve_router(PROXY_NAME, &host, port, router(state)).await
54}
55
56fn router(state: ProxyState) -> Router {
57 Router::new()
58 .route(HEALTHCHECK_PATH, get(healthcheck))
59 .route("/v1/models", get(models))
60 .route(COMPLETIONS_PATH, post(completions))
61 .route(CHAT_COMPLETIONS_PATH, post(chat_completions))
62 .route(RESET_PREFIX_CACHE_PATH, post(reset_prefix_cache))
63 .route(PRIME_PREFIX_CACHE_PATH, post(prime_prefix_cache))
64 .with_state(state)
65}
66
67#[derive(Clone)]
68struct ProxyState {
69 inner: Arc<ProxyStateInner>,
70}
71
72struct ProxyStateInner {
73 client: reqwest::Client,
74 prefill: Vec<PrefillClient>,
75 decode: Vec<String>,
76 ready: AtomicBool,
77 prefill_cursor: AtomicUsize,
78 decode_cursor: AtomicUsize,
79 request_counter: AtomicUsize,
80}
81
82#[derive(Clone)]
83struct PrefillClient {
84 url: String,
85 bootstrap_addr: String,
86 engine_ids: Arc<RwLock<Vec<String>>>,
87}
88
89#[derive(Clone)]
90struct SelectedPrefill {
91 url: String,
92 bootstrap_addr: String,
93 dp_rank: usize,
94 engine_id: String,
95}
96
97impl ProxyState {
98 fn new(config: Config) -> Result<Self, ProxyError> {
99 core::require_endpoints(
100 PROXY_NAME,
101 config.prefill.is_empty(),
102 config.decode.is_empty(),
103 )?;
104 let prefill = config
105 .prefill
106 .into_iter()
107 .map(PrefillClient::from_target)
108 .collect::<Result<Vec<_>, ProxyError>>()?;
109 Ok(Self {
110 inner: Arc::new(ProxyStateInner {
111 client: core::pooled_client(PROXY_NAME)?,
112 prefill,
113 decode: config.decode,
114 ready: AtomicBool::new(false),
115 prefill_cursor: AtomicUsize::new(0),
116 decode_cursor: AtomicUsize::new(0),
117 request_counter: AtomicUsize::new(0),
118 }),
119 })
120 }
121
122 fn client(&self) -> reqwest::Client {
123 self.inner.client.clone()
124 }
125
126 fn ready(&self) -> bool {
127 self.inner.ready.load(Ordering::SeqCst)
128 }
129
130 fn set_ready(&self) {
131 self.inner.ready.store(true, Ordering::SeqCst);
132 }
133
134 async fn next_prefill(&self) -> Result<SelectedPrefill, ProxyHttpError> {
135 let mut candidates = Vec::new();
136 for prefill in &self.inner.prefill {
137 let engine_ids = prefill.engine_ids.read().await;
138 for (dp_rank, engine_id) in engine_ids.iter().enumerate() {
139 candidates.push(SelectedPrefill {
140 url: prefill.url.clone(),
141 bootstrap_addr: prefill.bootstrap_addr.clone(),
142 dp_rank,
143 engine_id: engine_id.clone(),
144 });
145 }
146 }
147 if candidates.is_empty() {
148 return Err(ProxyHttpError::status(
149 StatusCode::SERVICE_UNAVAILABLE,
150 "no ready prefill data-parallel engines",
151 ));
152 }
153 let index = core::round_robin_index(&self.inner.prefill_cursor, candidates.len());
154 Ok(candidates.swap_remove(index))
155 }
156
157 fn next_decode_url(&self) -> String {
158 let index = core::round_robin_index(&self.inner.decode_cursor, self.inner.decode.len());
159 self.inner.decode[index].clone()
160 }
161
162 fn request_id(&self) -> String {
163 core::next_request_id(&self.inner.request_counter)
164 }
165}
166
167impl PrefillClient {
168 fn from_target(target: PrefillTarget) -> Result<Self, ProxyError> {
169 Ok(Self {
170 url: target.url,
171 bootstrap_addr: target.bootstrap_url,
172 engine_ids: Arc::new(RwLock::new(Vec::new())),
173 })
174 }
175}
176
177async fn discover_prefillers(state: ProxyState) {
178 for prefill in &state.inner.prefill {
179 loop {
180 if discover_prefiller(&state.client(), prefill).await.is_ok() {
181 break;
182 }
183 tokio::time::sleep(crate::core::BACKEND_RETRY_INTERVAL).await;
184 }
185 }
186 state.set_ready();
187}
188
189async fn discover_prefiller(
190 client: &reqwest::Client,
191 prefill: &PrefillClient,
192) -> Result<(), ProxyError> {
193 let health = client
194 .get(join_path(&prefill.url, "/health"))
195 .send()
196 .await
197 .map_err(|error| ProxyError::ExternalTool {
198 message: format!("prefill health request failed: {error}"),
199 })?;
200 if !health.status().is_success() {
201 return Err(ProxyError::ExternalTool {
202 message: format!("prefill health returned HTTP {}", health.status()),
203 });
204 }
205 let response = client
206 .get(join_path(&prefill.bootstrap_addr, "/query"))
207 .send()
208 .await
209 .map_err(|error| ProxyError::ExternalTool {
210 message: format!("prefill bootstrap query failed: {error}"),
211 })?;
212 if !response.status().is_success() {
213 return Err(ProxyError::ExternalTool {
214 message: format!(
215 "prefill bootstrap query returned HTTP {}",
216 response.status()
217 ),
218 });
219 }
220 let body = response
221 .json::<Value>()
222 .await
223 .map_err(|error| ProxyError::ExternalTool {
224 message: format!("prefill bootstrap query returned invalid JSON: {error}"),
225 })?;
226 let engine_ids = parse_engine_ids(&body)?;
227 *prefill.engine_ids.write().await = engine_ids;
228 Ok(())
229}
230
231fn parse_engine_ids(body: &Value) -> Result<Vec<String>, ProxyError> {
232 let object = body.as_object().ok_or_else(|| ProxyError::ExternalTool {
233 message: "prefill bootstrap query JSON must be an object".to_owned(),
234 })?;
235 if object.is_empty() {
236 return Err(ProxyError::ExternalTool {
237 message: "prefill bootstrap query returned no data-parallel engines".to_owned(),
238 });
239 }
240 let mut ranks = Vec::new();
241 for (rank_text, entry) in object {
242 let rank = rank_text
243 .parse::<usize>()
244 .map_err(|error| ProxyError::ExternalTool {
245 message: format!("invalid data-parallel rank {rank_text:?}: {error}"),
246 })?;
247 let engine_id = entry
248 .get("engine_id")
249 .and_then(Value::as_str)
250 .ok_or_else(|| ProxyError::ExternalTool {
251 message: format!("missing engine_id for data-parallel rank {rank_text:?}"),
252 })?;
253 ranks.push((rank, engine_id.to_owned()));
254 }
255 ranks.sort_by_key(|(rank, _engine_id)| *rank);
256 for (expected, (rank, _engine_id)) in ranks.iter().enumerate() {
257 if expected != *rank {
258 return Err(ProxyError::ExternalTool {
259 message: "prefill bootstrap query ranks must be contiguous from 0".to_owned(),
260 });
261 }
262 }
263 Ok(ranks
264 .into_iter()
265 .map(|(_rank, engine_id)| engine_id)
266 .collect())
267}
268
269async fn healthcheck(
270 State(state): State<ProxyState>,
271) -> (StatusCode, Json<ProxyHealthcheckResponse>) {
272 core::healthcheck_response(
273 state.ready(),
274 state.inner.prefill.len(),
275 state.inner.decode.len(),
276 )
277}
278
279async fn models(State(state): State<ProxyState>) -> Result<Response<Body>, ProxyHttpError> {
280 if !state.ready() {
281 return Err(ProxyHttpError::status(
282 StatusCode::SERVICE_UNAVAILABLE,
283 "proxy is not ready",
284 ));
285 }
286 let decode_url = state.next_decode_url();
287 let response = state
288 .client()
289 .get(join_path(&decode_url, "/v1/models"))
290 .send()
291 .await
292 .map_err(|error| ProxyHttpError::upstream("decode /v1/models request failed", error))?;
293 forward_response(response).await
294}
295
296async fn completions(
297 State(state): State<ProxyState>,
298 headers: HeaderMap,
299 Json(body): Json<Value>,
300) -> Result<Response<Body>, ProxyHttpError> {
301 completion_route(state, headers, body, COMPLETIONS_PATH).await
302}
303
304async fn chat_completions(
305 State(state): State<ProxyState>,
306 headers: HeaderMap,
307 Json(body): Json<Value>,
308) -> Result<Response<Body>, ProxyHttpError> {
309 completion_route(state, headers, body, CHAT_COMPLETIONS_PATH).await
310}
311
312async fn completion_route(
313 state: ProxyState,
314 headers: HeaderMap,
315 body: Value,
316 path: &'static str,
317) -> Result<Response<Body>, ProxyHttpError> {
318 if !state.ready() {
319 return Err(ProxyHttpError::status(
320 StatusCode::SERVICE_UNAVAILABLE,
321 "proxy is not ready",
322 ));
323 }
324 let selected_prefill = state.next_prefill().await?;
325 let decode_url = state.next_decode_url();
326 let request_id = state.request_id();
327 let authorization = outbound_authorization(&headers);
328 let client = state.client();
329 let prefill_body = prefill_body(&body, &request_id)?;
330 let decode_body = decode_body(&body, &selected_prefill, &request_id)?;
331 let prefill_task = tokio::spawn(send_prefill_request(
332 client.clone(),
333 selected_prefill.clone(),
334 path,
335 prefill_body,
336 request_id.clone(),
337 authorization.clone(),
338 ));
339 let decode_response = core::send_json_post(
340 client,
341 join_path(&decode_url, path),
342 &decode_body,
343 Some(&request_id),
344 authorization.as_deref(),
345 &[],
346 "decode request",
347 )
348 .await?;
349 core::stream_decode_response(decode_response, prefill_task, OnClientDrop::Abort)
350}
351
352fn prefill_body(body: &Value, request_id: &str) -> Result<Value, ProxyHttpError> {
353 let mut body = body.clone();
354 let Some(object) = body.as_object_mut() else {
355 return Err(ProxyHttpError::status(
356 StatusCode::BAD_REQUEST,
357 "OpenAI completion request body must be a JSON object",
358 ));
359 };
360 object.insert(
361 "kv_transfer_params".to_owned(),
362 MooncakePrefillKvTransferParams::new(request_id).into_protocol_value()?,
363 );
364 object.insert("stream".to_owned(), Value::Bool(false));
365 object.insert("max_tokens".to_owned(), Value::from(1_u8));
366 if object
367 .get("min_tokens")
368 .and_then(Value::as_u64)
369 .is_some_and(|min_tokens| min_tokens > 1)
370 {
371 object.insert("min_tokens".to_owned(), Value::from(1_u8));
372 }
373 if object.contains_key("max_completion_tokens") {
374 object.insert("max_completion_tokens".to_owned(), Value::from(1_u8));
375 }
376 object.remove("stream_options");
377 Ok(body)
378}
379
380fn decode_body(
381 body: &Value,
382 selected_prefill: &SelectedPrefill,
383 request_id: &str,
384) -> Result<Value, ProxyHttpError> {
385 let mut body = body.clone();
386 let Some(object) = body.as_object_mut() else {
387 return Err(ProxyHttpError::status(
388 StatusCode::BAD_REQUEST,
389 "OpenAI completion request body must be a JSON object",
390 ));
391 };
392 object.insert(
393 "kv_transfer_params".to_owned(),
394 MooncakeDecodeKvTransferParams::new(selected_prefill, request_id).into_protocol_value()?,
395 );
396 Ok(body)
397}
398
399#[derive(Serialize)]
400struct MooncakePrefillKvTransferParams {
401 do_remote_decode: bool,
402 do_remote_prefill: bool,
403 transfer_id: String,
404}
405
406impl MooncakePrefillKvTransferParams {
407 fn new(request_id: &str) -> Self {
408 Self {
409 do_remote_decode: true,
410 do_remote_prefill: false,
411 transfer_id: format!("xfer-{request_id}"),
412 }
413 }
414
415 fn into_protocol_value(self) -> Result<Value, ProxyHttpError> {
416 serde_json::to_value(self).map_err(|error| {
417 ProxyHttpError::internal(format!(
418 "failed to serialize vLLM Mooncake prefill transfer params: {error}"
419 ))
420 })
421 }
422}
423
424#[derive(Serialize)]
425struct MooncakeDecodeKvTransferParams<'a> {
426 do_remote_decode: bool,
427 do_remote_prefill: bool,
428 remote_bootstrap_addr: &'a str,
429 remote_engine_id: &'a str,
430 transfer_id: String,
431}
432
433impl<'a> MooncakeDecodeKvTransferParams<'a> {
434 fn new(selected_prefill: &'a SelectedPrefill, request_id: &str) -> Self {
435 Self {
436 do_remote_decode: false,
437 do_remote_prefill: true,
438 remote_bootstrap_addr: &selected_prefill.bootstrap_addr,
439 remote_engine_id: &selected_prefill.engine_id,
440 transfer_id: format!("xfer-{request_id}"),
441 }
442 }
443
444 fn into_protocol_value(self) -> Result<Value, ProxyHttpError> {
445 serde_json::to_value(self).map_err(|error| {
446 ProxyHttpError::internal(format!(
447 "failed to serialize vLLM Mooncake decode transfer params: {error}"
448 ))
449 })
450 }
451}
452
453async fn send_prefill_request(
454 client: reqwest::Client,
455 selected_prefill: SelectedPrefill,
456 path: &'static str,
457 body: Value,
458 request_id: String,
459 authorization: Option<String>,
460) -> Result<(), ProxyHttpError> {
461 let response = core::send_json_post(
462 client,
463 join_path(&selected_prefill.url, path),
464 &body,
465 Some(&request_id),
466 authorization.as_deref(),
467 &[("X-data-parallel-rank", selected_prefill.dp_rank.to_string())],
468 "prefill request",
469 )
470 .await?;
471 response
472 .bytes()
473 .await
474 .map_err(|error| ProxyHttpError::upstream("prefill response drain failed", error))?;
475 Ok(())
476}
477
478async fn prime_flow(
483 state: &ProxyState,
484 selected_prefill: &SelectedPrefill,
485 authorization: Option<String>,
486 body: &Value,
487) -> Result<u16, core::PrimeFlowFailure> {
488 use core::PrimeFlowFailure;
489 let request_id = state.request_id();
490 let prefill_body = prefill_body(body, &request_id).map_err(PrimeFlowFailure::transport)?;
491 let decode_body =
492 decode_body(body, selected_prefill, &request_id).map_err(PrimeFlowFailure::transport)?;
493 let client = state.client();
494 let prefill_response = core::send_json_post_status(
495 client.clone(),
496 join_path(&selected_prefill.url, COMPLETIONS_PATH),
497 &prefill_body,
498 Some(&request_id),
499 authorization.as_deref(),
500 &[("X-data-parallel-rank", selected_prefill.dp_rank.to_string())],
501 "prefill conditioning request",
502 )
503 .await
504 .map_err(PrimeFlowFailure::transport)?;
505 let (prefill_status, _) = core::expect_2xx("prefill conditioning", prefill_response).await?;
506 let decode_url = state.next_decode_url();
507 let decode_response = core::send_json_post_status(
508 client,
509 join_path(&decode_url, COMPLETIONS_PATH),
510 &decode_body,
511 Some(&request_id),
512 authorization.as_deref(),
513 &[],
514 "decode conditioning request",
515 )
516 .await
517 .map_err(PrimeFlowFailure::transport)?;
518 core::expect_2xx("decode conditioning", decode_response).await?;
519 Ok(prefill_status)
520}
521
522impl core::PrimeFanoutTarget for SelectedPrefill {
523 fn url(&self) -> &str {
524 &self.url
525 }
526
527 fn rank(&self) -> u32 {
528 self.dp_rank as u32
529 }
530}
531
532async fn prime_prefix_cache(
533 State(state): State<ProxyState>,
534 headers: HeaderMap,
535 Json(body): Json<Value>,
536) -> Response<Body> {
537 if !state.ready() {
538 return ProxyHttpError::status(StatusCode::SERVICE_UNAVAILABLE, "proxy is not ready")
539 .into_response();
540 }
541 let authorization = outbound_authorization(&headers);
542 let mut targets = Vec::new();
543 for prefill in &state.inner.prefill {
544 let engine_ids = prefill.engine_ids.read().await.clone();
545 for (dp_rank, engine_id) in engine_ids.iter().enumerate() {
546 targets.push(SelectedPrefill {
547 url: prefill.url.clone(),
548 bootstrap_addr: prefill.bootstrap_addr.clone(),
549 dp_rank,
550 engine_id: engine_id.clone(),
551 });
552 }
553 }
554 core::run_prime_fanout("prefix cache conditioning", targets, |selected| {
555 let state = state.clone();
556 let authorization = authorization.clone();
557 let body = body.clone();
558 async move { prime_flow(&state, &selected, authorization, &body).await }
559 })
560 .await
561}
562
563async fn reset_prefix_cache(State(state): State<ProxyState>, headers: HeaderMap) -> Response<Body> {
564 let authorization = outbound_authorization(&headers);
565 let targets = core::fanout_target_urls(
566 state
567 .inner
568 .prefill
569 .iter()
570 .map(|prefill| prefill.url.as_str()),
571 state.inner.decode.iter().map(String::as_str),
572 );
573 core::run_sweep_fanout(
574 state.client(),
575 "prefix cache reset",
576 "/reset_prefix_cache",
577 targets,
578 authorization,
579 )
580 .await
581}
582
583#[cfg(test)]
584mod tests {
585 use super::*;
586 use anyhow::{Context, Result};
587 use async_stream::stream;
588 use axum::body::to_bytes;
589 use axum::http::{HeaderValue, header};
590 use axum::response::IntoResponse;
591 use axum::serve;
592 use bytes::Bytes;
593 use futures_util::StreamExt;
594 use serde_json::json;
595 use std::sync::atomic::AtomicU16;
596 use std::time::Duration;
597 use tokio::net::TcpListener;
598 use tokio::sync::{Mutex, Notify};
599 use tokio::task::JoinHandle;
600
601 #[tokio::test]
602 async fn healthcheck_response_reports_readiness_and_configured_instances() -> Result<()> {
603 let state = ProxyState::new(Config {
604 host: "127.0.0.1".to_owned(),
605 port: 8000,
606 prefill: vec![PrefillTarget {
607 url: "http://127.0.0.1:8010".to_owned(),
608 bootstrap_url: "http://127.0.0.1:8998".to_owned(),
609 }],
610 decode: vec![
611 "http://127.0.0.1:8020".to_owned(),
612 "http://127.0.0.1:8021".to_owned(),
613 ],
614 })?;
615 let (status, Json(response)) = healthcheck(State(state.clone())).await;
616 assert_eq!(status, StatusCode::SERVICE_UNAVAILABLE);
617 assert!(!response.ready);
618
619 state.set_ready();
620 let (status, Json(response)) = healthcheck(State(state)).await;
621 let value = serde_json::to_value(response)?;
622
623 assert_eq!(status, StatusCode::OK);
624 assert_eq!(value.get("ready").and_then(Value::as_bool), Some(true));
625 assert_eq!(
626 value.get("prefill_instances").and_then(Value::as_u64),
627 Some(1)
628 );
629 assert_eq!(
630 value.get("decode_instances").and_then(Value::as_u64),
631 Some(2)
632 );
633 Ok(())
634 }
635
636 #[test]
637 fn prefill_body_forces_single_token_non_streaming_transfer() -> Result<()> {
638 let body = json!({
639 "model": "m",
640 "prompt": "hello",
641 "stream": true,
642 "stream_options": {"include_usage": true},
643 "max_tokens": 64,
644 "max_completion_tokens": 64,
645 "min_tokens": 64,
646 });
647 let lowered =
648 prefill_body(&body, "request-1").map_err(|error| anyhow::anyhow!(error.to_string()))?;
649 assert_eq!(
650 lowered.pointer("/kv_transfer_params/do_remote_decode"),
651 Some(&Value::Bool(true))
652 );
653 assert_eq!(
654 lowered.pointer("/kv_transfer_params/do_remote_prefill"),
655 Some(&Value::Bool(false))
656 );
657 assert_eq!(
658 lowered
659 .pointer("/kv_transfer_params/transfer_id")
660 .and_then(Value::as_str),
661 Some("xfer-request-1")
662 );
663 assert_eq!(lowered.get("stream"), Some(&Value::Bool(false)));
664 assert_eq!(lowered.get("max_tokens").and_then(Value::as_u64), Some(1));
665 assert_eq!(lowered.get("min_tokens").and_then(Value::as_u64), Some(1));
666 assert_eq!(
667 lowered.get("max_completion_tokens").and_then(Value::as_u64),
668 Some(1)
669 );
670 assert!(lowered.get("stream_options").is_none());
671 Ok(())
672 }
673
674 #[test]
675 fn decode_body_attaches_remote_prefill_identity() -> Result<()> {
676 let selected = SelectedPrefill {
677 url: "http://127.0.0.1:8010".to_owned(),
678 bootstrap_addr: "http://127.0.0.1:8998".to_owned(),
679 dp_rank: 0,
680 engine_id: "engine-a".to_owned(),
681 };
682 let lowered = decode_body(&json!({"model": "m"}), &selected, "request-2")
683 .map_err(|error| anyhow::anyhow!(error.to_string()))?;
684 assert_eq!(
685 lowered.pointer("/kv_transfer_params/do_remote_decode"),
686 Some(&Value::Bool(false))
687 );
688 assert_eq!(
689 lowered.pointer("/kv_transfer_params/do_remote_prefill"),
690 Some(&Value::Bool(true))
691 );
692 assert_eq!(
693 lowered
694 .pointer("/kv_transfer_params/remote_bootstrap_addr")
695 .and_then(Value::as_str),
696 Some("http://127.0.0.1:8998")
697 );
698 assert_eq!(
699 lowered
700 .pointer("/kv_transfer_params/remote_engine_id")
701 .and_then(Value::as_str),
702 Some("engine-a")
703 );
704 assert_eq!(
705 lowered
706 .pointer("/kv_transfer_params/transfer_id")
707 .and_then(Value::as_str),
708 Some("xfer-request-2")
709 );
710 Ok(())
711 }
712
713 #[test]
714 fn prefill_client_uses_explicit_bootstrap_url() -> Result<()> {
715 let client = PrefillClient::from_target(PrefillTarget {
716 url: "http://10.0.0.1:8010".to_owned(),
717 bootstrap_url: "http://192.0.2.10:8998".to_owned(),
718 })?;
719 assert_eq!(client.url, "http://10.0.0.1:8010");
720 assert_eq!(client.bootstrap_addr, "http://192.0.2.10:8998");
721 Ok(())
722 }
723
724 #[tokio::test]
725 async fn chat_dispatch_preserves_route_messages_and_unowned_fields() -> Result<()> {
726 let prefill_backend = MockBackend::default();
727 let decode_backend = MockBackend::default();
728 let (prefill, prefill_server) = spawn_backend(prefill_backend.clone()).await?;
729 let (decode, decode_server) = spawn_backend(decode_backend.clone()).await?;
730 let state = ProxyState::new(Config {
731 host: "127.0.0.1".to_owned(),
732 port: 8000,
733 prefill: vec![PrefillTarget {
734 url: prefill,
735 bootstrap_url: "http://127.0.0.1:8998".to_owned(),
736 }],
737 decode: vec![decode],
738 })?;
739 *state.inner.prefill[0].engine_ids.write().await = vec!["prefill-0".to_owned()];
740 state.set_ready();
741 let request = json!({
742 "model": "m",
743 "messages": [{"role": "user", "content": "hello"}],
744 "temperature": 1.0,
745 "reasoning_effort": "high",
746 "chat_template_kwargs": {"enable_thinking": true}
747 });
748
749 let response = completion_route(
750 state,
751 HeaderMap::new(),
752 request.clone(),
753 CHAT_COMPLETIONS_PATH,
754 )
755 .await
756 .map_err(|error| anyhow::anyhow!(error.to_string()))?;
757 assert_eq!(response.status(), StatusCode::OK);
758 let _body = to_bytes(response.into_body(), usize::MAX).await?;
759
760 let prefill_requests = prefill_backend.requests.lock().await;
761 let decode_requests = decode_backend.requests.lock().await;
762 assert_eq!(prefill_requests.len(), 1);
763 assert_eq!(decode_requests.len(), 1);
764 for key in [
765 "messages",
766 "temperature",
767 "reasoning_effort",
768 "chat_template_kwargs",
769 ] {
770 assert_eq!(prefill_requests[0][key], request[key]);
771 assert_eq!(decode_requests[0][key], request[key]);
772 }
773 prefill_server.abort();
774 decode_server.abort();
775 Ok(())
776 }
777
778 #[tokio::test]
779 async fn streaming_decode_reaches_both_public_routes_before_terminal_event() -> Result<()> {
780 for path in [COMPLETIONS_PATH, CHAT_COMPLETIONS_PATH] {
781 let terminal_gate = Arc::new(Notify::new());
782 let prefill_backend = StreamingBackend::prefill(terminal_gate.clone());
783 let decode_backend = StreamingBackend::decode(terminal_gate.clone());
784 let (prefill, prefill_server) =
785 spawn_streaming_backend(prefill_backend.clone()).await?;
786 let (decode, decode_server) = spawn_streaming_backend(decode_backend).await?;
787 let state = ProxyState::new(Config {
788 host: "127.0.0.1".to_owned(),
789 port: 8000,
790 prefill: vec![PrefillTarget {
791 url: prefill,
792 bootstrap_url: "http://127.0.0.1:8998".to_owned(),
793 }],
794 decode: vec![decode],
795 })?;
796 *state.inner.prefill[0].engine_ids.write().await = vec!["prefill-0".to_owned()];
797 state.set_ready();
798 let request = if path == COMPLETIONS_PATH {
799 json!({"model": "m", "prompt": "hello", "stream": true})
800 } else {
801 json!({
802 "model": "m",
803 "messages": [{"role": "user", "content": "hello"}],
804 "stream": true
805 })
806 };
807
808 let response = tokio::time::timeout(
809 Duration::from_secs(1),
810 completion_route(state, HeaderMap::new(), request, path),
811 )
812 .await
813 .context("Gateway waited for decode completion before returning response headers")?
814 .map_err(|error| anyhow::anyhow!(error.to_string()))?;
815 assert_eq!(response.status(), StatusCode::ACCEPTED);
816 assert_eq!(
817 response.headers().get(header::CONTENT_TYPE),
818 Some(&HeaderValue::from_static("text/event-stream"))
819 );
820 let mut stream = response.into_body().into_data_stream();
821 let first = tokio::time::timeout(Duration::from_secs(1), stream.next())
822 .await
823 .context("Gateway buffered the first SSE event until decode completion")?
824 .context("decode stream ended before its first SSE event")??;
825 assert_eq!(first, Bytes::from_static(b"data: first\n\n"));
826
827 terminal_gate.notify_one();
828 let terminal = stream
829 .next()
830 .await
831 .context("decode stream ended before its terminal SSE event")??;
832 assert_eq!(terminal, Bytes::from_static(b"data: [DONE]\n\n"));
833 prefill_server.abort();
834 decode_server.abort();
835 }
836 Ok(())
837 }
838
839 #[tokio::test]
840 async fn static_backend_selection_round_robins_prefill_engines_and_decode_urls() -> Result<()> {
841 let state = ProxyState::new(Config {
842 host: "127.0.0.1".to_owned(),
843 port: 8000,
844 prefill: vec![
845 PrefillTarget {
846 url: "http://127.0.0.1:8010".to_owned(),
847 bootstrap_url: "http://127.0.0.1:8998".to_owned(),
848 },
849 PrefillTarget {
850 url: "http://127.0.0.1:8011".to_owned(),
851 bootstrap_url: "http://127.0.0.1:8999".to_owned(),
852 },
853 ],
854 decode: vec![
855 "http://127.0.0.1:8020".to_owned(),
856 "http://127.0.0.1:8021".to_owned(),
857 ],
858 })?;
859 *state.inner.prefill[0].engine_ids.write().await =
860 vec!["p0-r0".to_owned(), "p0-r1".to_owned()];
861 *state.inner.prefill[1].engine_ids.write().await = vec!["p1-r0".to_owned()];
862
863 let prefill0 = state
864 .next_prefill()
865 .await
866 .map_err(|error| anyhow::anyhow!(error.to_string()))?;
867 let prefill1 = state
868 .next_prefill()
869 .await
870 .map_err(|error| anyhow::anyhow!(error.to_string()))?;
871 let prefill2 = state
872 .next_prefill()
873 .await
874 .map_err(|error| anyhow::anyhow!(error.to_string()))?;
875 assert_eq!(prefill0.engine_id, "p0-r0");
876 assert_eq!(prefill1.engine_id, "p0-r1");
877 assert_eq!(prefill2.engine_id, "p1-r0");
878 assert_eq!(
879 state
880 .next_prefill()
881 .await
882 .map_err(|error| anyhow::anyhow!(error.to_string()))?
883 .engine_id,
884 "p0-r0"
885 );
886
887 assert_eq!(state.next_decode_url(), "http://127.0.0.1:8020");
888 assert_eq!(state.next_decode_url(), "http://127.0.0.1:8021");
889 assert_eq!(state.next_decode_url(), "http://127.0.0.1:8020");
890 Ok(())
891 }
892
893 #[tokio::test]
894 async fn reset_prefix_cache_attempts_all_targets_and_reports_partial_failure() -> Result<()> {
895 let prefill_backend = ResetBackend::new();
896 let decode_backend = ResetBackend::new();
897 let (prefill, prefill_server) = spawn_reset_backend(prefill_backend.clone()).await?;
898 let (decode, decode_server) = spawn_reset_backend(decode_backend.clone()).await?;
899 let state = ProxyState::new(Config {
900 host: "127.0.0.1".to_owned(),
901 port: 8000,
902 prefill: vec![PrefillTarget {
903 url: prefill,
904 bootstrap_url: "http://127.0.0.1:8998".to_owned(),
905 }],
906 decode: vec![decode],
907 })?;
908
909 let all_succeeded = reset_prefix_cache(State(state.clone()), HeaderMap::new()).await;
910 assert_eq!(all_succeeded.status(), StatusCode::OK);
911
912 decode_backend
913 .status
914 .store(StatusCode::PARTIAL_CONTENT.as_u16(), Ordering::SeqCst);
915 let partial = reset_prefix_cache(State(state), HeaderMap::new()).await;
916 assert_eq!(partial.status(), StatusCode::PARTIAL_CONTENT);
917 let body: Value =
918 serde_json::from_slice(&to_bytes(partial.into_body(), usize::MAX).await?)?;
919 assert_eq!(body["successful"].as_array().map(Vec::len), Some(1));
920 assert_eq!(body["failed"].as_array().map(Vec::len), Some(1));
921 assert_eq!(prefill_backend.requests.load(Ordering::SeqCst), 2);
922 assert_eq!(decode_backend.requests.load(Ordering::SeqCst), 2);
923 prefill_server.abort();
924 decode_server.abort();
925 Ok(())
926 }
927
928 #[derive(Clone)]
929 struct ResetBackend {
930 status: Arc<AtomicU16>,
931 requests: Arc<AtomicUsize>,
932 }
933
934 impl ResetBackend {
935 fn new() -> Self {
936 Self {
937 status: Arc::new(AtomicU16::new(StatusCode::OK.as_u16())),
938 requests: Arc::new(AtomicUsize::new(0)),
939 }
940 }
941 }
942
943 async fn mock_reset(State(state): State<ResetBackend>) -> Response<Body> {
944 state.requests.fetch_add(1, Ordering::SeqCst);
945 let status = StatusCode::from_u16(state.status.load(Ordering::SeqCst))
946 .unwrap_or(StatusCode::INTERNAL_SERVER_ERROR);
947 (status, "reset").into_response()
948 }
949
950 async fn spawn_reset_backend(state: ResetBackend) -> Result<(String, JoinHandle<()>)> {
951 let app = Router::new()
952 .route("/reset_prefix_cache", post(mock_reset))
953 .with_state(state);
954 let listener = TcpListener::bind("127.0.0.1:0").await?;
955 let address = listener.local_addr()?;
956 let server = tokio::spawn(async move {
957 let _result = serve(listener, app).await;
958 });
959 Ok((format!("http://{address}"), server))
960 }
961
962 #[derive(Clone, Default)]
963 struct MockBackend {
964 requests: Arc<Mutex<Vec<Value>>>,
965 }
966 async fn mock_chat(
967 State(state): State<MockBackend>,
968 Json(body): Json<Value>,
969 ) -> Response<Body> {
970 state.requests.lock().await.push(body);
971 Json(json!({"object": "chat.completion", "choices": []})).into_response()
972 }
973
974 async fn spawn_backend(state: MockBackend) -> Result<(String, JoinHandle<()>)> {
975 let app = Router::new()
976 .route("/v1/chat/completions", post(mock_chat))
977 .with_state(state);
978 let listener = TcpListener::bind("127.0.0.1:0").await?;
979 let address = listener.local_addr()?;
980 let server = tokio::spawn(async move {
981 let _result = serve(listener, app).await;
982 });
983 Ok((format!("http://{address}"), server))
984 }
985
986 #[derive(Clone)]
987 struct StreamingBackend {
988 prefill: bool,
989 terminal_gate: Arc<Notify>,
990 }
991
992 impl StreamingBackend {
993 fn prefill(terminal_gate: Arc<Notify>) -> Self {
994 Self {
995 prefill: true,
996 terminal_gate,
997 }
998 }
999
1000 fn decode(terminal_gate: Arc<Notify>) -> Self {
1001 Self {
1002 prefill: false,
1003 terminal_gate,
1004 }
1005 }
1006 }
1007
1008 async fn streaming_response(
1009 State(state): State<StreamingBackend>,
1010 Json(_body): Json<Value>,
1011 ) -> Response<Body> {
1012 if state.prefill {
1013 return Json(json!({"status": "prefill-complete"})).into_response();
1014 }
1015
1016 let terminal_gate = state.terminal_gate;
1017 let body = Body::from_stream(stream! {
1018 yield Ok::<Bytes, std::io::Error>(Bytes::from_static(b"data: first\n\n"));
1019 terminal_gate.notified().await;
1020 yield Ok(Bytes::from_static(b"data: [DONE]\n\n"));
1021 });
1022 (
1023 StatusCode::ACCEPTED,
1024 [(header::CONTENT_TYPE, "text/event-stream")],
1025 body,
1026 )
1027 .into_response()
1028 }
1029
1030 async fn spawn_streaming_backend(state: StreamingBackend) -> Result<(String, JoinHandle<()>)> {
1031 let app = Router::new()
1032 .route("/v1/completions", post(streaming_response))
1033 .route("/v1/chat/completions", post(streaming_response))
1034 .with_state(state);
1035 let listener = TcpListener::bind("127.0.0.1:0").await?;
1036 let address = listener.local_addr()?;
1037 let server = tokio::spawn(async move {
1038 let _result = serve(listener, app).await;
1039 });
1040 Ok((format!("http://{address}"), server))
1041 }
1042
1043 #[tokio::test]
1044 async fn prime_prefix_cache_fans_out_to_each_discovered_rank_and_reports_partial_failure()
1045 -> Result<()> {
1046 let prefill_backend = PrimeBackend::default();
1047 let decode_backend = PrimeBackend::default();
1048 let (prefill, prefill_server) = spawn_prime_backend(prefill_backend.clone()).await?;
1049 let (decode, decode_server) = spawn_prime_backend(decode_backend.clone()).await?;
1050 let state = ProxyState::new(Config {
1051 host: "127.0.0.1".to_owned(),
1052 port: 8000,
1053 prefill: vec![PrefillTarget {
1054 url: prefill.clone(),
1055 bootstrap_url: "http://127.0.0.1:8998".to_owned(),
1056 }],
1057 decode: vec![decode],
1058 })?;
1059 *state.inner.prefill[0].engine_ids.write().await =
1060 vec!["prefill-r0".to_owned(), "prefill-r1".to_owned()];
1061 state.set_ready();
1062 let conditioning =
1063 || Json(json!({"model": "m", "prompt": "canonical prefix", "max_tokens": 1}));
1064
1065 let response =
1066 prime_prefix_cache(State(state.clone()), HeaderMap::new(), conditioning()).await;
1067 assert_eq!(response.status(), StatusCode::OK);
1068 let body: Value =
1069 serde_json::from_slice(&to_bytes(response.into_body(), usize::MAX).await?)?;
1070 let targets = body["targets"]
1071 .as_array()
1072 .ok_or_else(|| anyhow::anyhow!("fan-out response has no targets"))?;
1073 assert_eq!(targets.len(), 2);
1074 for (rank, target) in targets.iter().enumerate() {
1075 assert_eq!(target["url"].as_str(), Some(prefill.as_str()));
1076 assert_eq!(target["rank"].as_u64(), Some(rank as u64));
1077 assert_eq!(target["http_status"].as_u64(), Some(200));
1078 assert!(target["error"].is_null());
1079 }
1080 let prefill_requests = prefill_backend.requests.lock().await;
1081 assert_eq!(prefill_requests.len(), 2);
1082 assert_eq!(prefill_requests[0].0.as_deref(), Some("0"));
1083 assert_eq!(prefill_requests[1].0.as_deref(), Some("1"));
1084 assert_eq!(
1086 prefill_requests[0]
1087 .1
1088 .pointer("/kv_transfer_params/do_remote_decode"),
1089 Some(&Value::Bool(true))
1090 );
1091 assert_eq!(
1092 prefill_requests[1]
1093 .1
1094 .pointer("/kv_transfer_params/do_remote_decode"),
1095 Some(&Value::Bool(true))
1096 );
1097 let decode_requests = decode_backend.requests.lock().await;
1098 assert_eq!(decode_requests.len(), 2);
1099 assert_eq!(
1100 decode_requests[0]
1101 .1
1102 .pointer("/kv_transfer_params/remote_engine_id")
1103 .and_then(Value::as_str),
1104 Some("prefill-r0")
1105 );
1106 assert_eq!(
1107 decode_requests[1]
1108 .1
1109 .pointer("/kv_transfer_params/remote_engine_id")
1110 .and_then(Value::as_str),
1111 Some("prefill-r1")
1112 );
1113 drop(prefill_requests);
1114 drop(decode_requests);
1115
1116 prefill_backend.set_fail_rank(Some("1".to_owned())).await;
1117 let partial = prime_prefix_cache(State(state), HeaderMap::new(), conditioning()).await;
1118 assert_eq!(partial.status(), StatusCode::PARTIAL_CONTENT);
1119 let body: Value =
1120 serde_json::from_slice(&to_bytes(partial.into_body(), usize::MAX).await?)?;
1121 let targets = body["targets"]
1122 .as_array()
1123 .ok_or_else(|| anyhow::anyhow!("fan-out response has no targets"))?;
1124 assert_eq!(targets.len(), 2);
1125 assert!(targets[0]["error"].is_null());
1126 assert_eq!(targets[1]["rank"].as_u64(), Some(1));
1127 assert_eq!(targets[1]["http_status"].as_u64(), Some(500));
1128 assert!(
1129 targets[1]["error"]
1130 .as_str()
1131 .is_some_and(|error| error.contains("HTTP 500"))
1132 );
1133 prefill_server.abort();
1134 decode_server.abort();
1135 Ok(())
1136 }
1137
1138 #[tokio::test]
1142 async fn prime_prefix_cache_rejects_an_empty_target_set() -> Result<()> {
1143 let state = ProxyState::new(Config {
1144 host: "127.0.0.1".to_owned(),
1145 port: 8000,
1146 prefill: vec![PrefillTarget {
1147 url: "http://127.0.0.1:8010".to_owned(),
1148 bootstrap_url: "http://127.0.0.1:8998".to_owned(),
1149 }],
1150 decode: vec!["http://127.0.0.1:8020".to_owned()],
1151 })?;
1152 state.set_ready();
1153
1154 let response = prime_prefix_cache(
1155 State(state),
1156 HeaderMap::new(),
1157 Json(json!({"model": "m", "prompt": "canonical prefix", "max_tokens": 1})),
1158 )
1159 .await;
1160 assert_eq!(response.status(), StatusCode::BAD_GATEWAY);
1161 let body: Value =
1162 serde_json::from_slice(&to_bytes(response.into_body(), usize::MAX).await?)?;
1163 assert!(
1164 body["error"]
1165 .as_str()
1166 .is_some_and(|error| error.contains("no targets")),
1167 "got {body}"
1168 );
1169 Ok(())
1170 }
1171
1172 type PrimeRequests = Arc<Mutex<Vec<(Option<String>, Value)>>>;
1173
1174 #[derive(Clone, Default)]
1175 struct PrimeBackend {
1176 requests: PrimeRequests,
1177 fail_rank: Arc<Mutex<Option<String>>>,
1178 }
1179
1180 impl PrimeBackend {
1181 async fn set_fail_rank(&self, rank: Option<String>) {
1182 *self.fail_rank.lock().await = rank;
1183 }
1184 }
1185
1186 async fn mock_prime(
1187 State(state): State<PrimeBackend>,
1188 headers: HeaderMap,
1189 Json(body): Json<Value>,
1190 ) -> Response<Body> {
1191 let rank = headers
1192 .get("x-data-parallel-rank")
1193 .and_then(|value| value.to_str().ok())
1194 .map(str::to_owned);
1195 let fail_rank = state.fail_rank.lock().await.clone();
1196 state.requests.lock().await.push((rank.clone(), body));
1197 if fail_rank.is_some() && fail_rank == rank {
1198 return StatusCode::INTERNAL_SERVER_ERROR.into_response();
1199 }
1200 Json(json!({"object": "text_completion", "choices": []})).into_response()
1201 }
1202
1203 async fn spawn_prime_backend(state: PrimeBackend) -> Result<(String, JoinHandle<()>)> {
1204 let app = Router::new()
1205 .route("/v1/completions", post(mock_prime))
1206 .with_state(state);
1207 let listener = TcpListener::bind("127.0.0.1:0").await?;
1208 let address = listener.local_addr()?;
1209 let server = tokio::spawn(async move {
1210 let _result = serve(listener, app).await;
1211 });
1212 Ok((format!("http://{address}"), server))
1213 }
1214}