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