1use crate::core::{
2 self, ProxyHealthcheckResponse, ProxyHttpError, ProxyMeta, forward_response, join_path,
3 outbound_authorization,
4};
5use crate::error::ProxyError;
6use axum::body::Body;
7use axum::extract::{Json, State};
8use axum::http::{HeaderMap, Response, StatusCode};
9use axum::routing::{get, post};
10use axum::{Router, serve};
11use serde::Serialize;
12use serde_json::Value;
13use std::sync::Arc;
14use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
15use std::time::Duration;
16use tokio::net::TcpListener;
17use tokio::sync::RwLock;
18
19pub const ID: &str = "inferlab-vllm-mooncake-proxy";
21pub const VERSION: u32 = 1;
23
24pub fn meta() -> ProxyMeta {
26 ProxyMeta {
27 id: ID,
28 version: VERSION,
29 }
30}
31
32#[derive(Clone, Debug)]
33pub struct Config {
34 pub host: String,
35 pub port: u16,
36 pub prefill: Vec<PrefillTarget>,
37 pub decode: Vec<String>,
38}
39
40#[derive(Clone, Debug)]
41pub struct PrefillTarget {
42 pub url: String,
43 pub bootstrap_url: String,
44}
45
46pub fn run(config: Config) -> Result<(), ProxyError> {
47 core::run(|| run_async(config))
48}
49
50pub async fn run_async(config: Config) -> Result<(), ProxyError> {
51 let host = config.host.clone();
52 let port = config.port;
53 let state = ProxyState::new(config)?;
54 tokio::spawn(discover_prefillers(state.clone()));
55 let app = router(state);
56 let listener = TcpListener::bind((host.as_str(), port))
57 .await
58 .map_err(|error| ProxyError::Io {
59 message: format!("failed to bind vLLM Mooncake proxy on {host}:{port}: {error}"),
60 })?;
61 serve(listener, app).await.map_err(|error| ProxyError::Io {
62 message: format!("vLLM Mooncake proxy server failed: {error}"),
63 })
64}
65
66fn router(state: ProxyState) -> Router {
67 Router::new()
68 .route("/healthcheck", get(healthcheck))
69 .route("/v1/models", get(models))
70 .route("/v1/completions", post(completions))
71 .route("/v1/chat/completions", post(chat_completions))
72 .with_state(state)
73}
74
75#[derive(Clone)]
76struct ProxyState {
77 inner: Arc<ProxyStateInner>,
78}
79
80struct ProxyStateInner {
81 client: reqwest::Client,
82 prefill: Vec<PrefillClient>,
83 decode: Vec<String>,
84 ready: AtomicBool,
85 prefill_cursor: AtomicUsize,
86 decode_cursor: AtomicUsize,
87 request_counter: AtomicUsize,
88}
89
90#[derive(Clone)]
91struct PrefillClient {
92 url: String,
93 bootstrap_addr: String,
94 engine_ids: Arc<RwLock<Vec<String>>>,
95}
96
97#[derive(Clone)]
98struct SelectedPrefill {
99 url: String,
100 bootstrap_addr: String,
101 dp_rank: usize,
102 engine_id: String,
103}
104
105impl ProxyState {
106 fn new(config: Config) -> Result<Self, ProxyError> {
107 if config.prefill.is_empty() {
108 return Err(ProxyError::Invalid {
109 message: "vLLM Mooncake proxy requires at least one prefill endpoint".to_owned(),
110 });
111 }
112 if config.decode.is_empty() {
113 return Err(ProxyError::Invalid {
114 message: "vLLM Mooncake proxy requires at least one decode endpoint".to_owned(),
115 });
116 }
117 let client = core::build_pooled_client().map_err(|error| ProxyError::Io {
118 message: format!("failed to create vLLM Mooncake proxy HTTP client: {error}"),
119 })?;
120 let prefill = config
121 .prefill
122 .into_iter()
123 .map(PrefillClient::from_target)
124 .collect::<Result<Vec<_>, ProxyError>>()?;
125 Ok(Self {
126 inner: Arc::new(ProxyStateInner {
127 client,
128 prefill,
129 decode: config.decode,
130 ready: AtomicBool::new(false),
131 prefill_cursor: AtomicUsize::new(0),
132 decode_cursor: AtomicUsize::new(0),
133 request_counter: AtomicUsize::new(0),
134 }),
135 })
136 }
137
138 fn client(&self) -> reqwest::Client {
139 self.inner.client.clone()
140 }
141
142 fn ready(&self) -> bool {
143 self.inner.ready.load(Ordering::SeqCst)
144 }
145
146 fn set_ready(&self) {
147 self.inner.ready.store(true, Ordering::SeqCst);
148 }
149
150 async fn next_prefill(&self) -> Result<SelectedPrefill, ProxyHttpError> {
151 let mut candidates = Vec::new();
152 for prefill in &self.inner.prefill {
153 let engine_ids = prefill.engine_ids.read().await;
154 for (dp_rank, engine_id) in engine_ids.iter().enumerate() {
155 candidates.push(SelectedPrefill {
156 url: prefill.url.clone(),
157 bootstrap_addr: prefill.bootstrap_addr.clone(),
158 dp_rank,
159 engine_id: engine_id.clone(),
160 });
161 }
162 }
163 if candidates.is_empty() {
164 return Err(ProxyHttpError::status(
165 StatusCode::SERVICE_UNAVAILABLE,
166 "no ready prefill data-parallel engines",
167 ));
168 }
169 let index = core::round_robin_index(&self.inner.prefill_cursor, candidates.len());
170 Ok(candidates.swap_remove(index))
171 }
172
173 fn next_decode_url(&self) -> Result<String, ProxyHttpError> {
174 if self.inner.decode.is_empty() {
175 return Err(ProxyHttpError::status(
176 StatusCode::SERVICE_UNAVAILABLE,
177 "no decode endpoints configured",
178 ));
179 }
180 let index = core::round_robin_index(&self.inner.decode_cursor, self.inner.decode.len());
181 Ok(self.inner.decode[index].clone())
182 }
183
184 fn request_id(&self) -> String {
185 core::next_request_id(&self.inner.request_counter)
186 }
187}
188
189impl PrefillClient {
190 fn from_target(target: PrefillTarget) -> Result<Self, ProxyError> {
191 Ok(Self {
192 url: target.url,
193 bootstrap_addr: target.bootstrap_url,
194 engine_ids: Arc::new(RwLock::new(Vec::new())),
195 })
196 }
197}
198
199async fn discover_prefillers(state: ProxyState) {
200 for prefill in &state.inner.prefill {
201 loop {
202 if discover_prefiller(&state.client(), prefill).await.is_ok() {
203 break;
204 }
205 tokio::time::sleep(Duration::from_secs(1)).await;
206 }
207 }
208 state.set_ready();
209}
210
211async fn discover_prefiller(
212 client: &reqwest::Client,
213 prefill: &PrefillClient,
214) -> Result<(), ProxyError> {
215 let health = client
216 .get(join_path(&prefill.url, "/health"))
217 .send()
218 .await
219 .map_err(|error| ProxyError::ExternalTool {
220 message: format!("prefill health request failed: {error}"),
221 })?;
222 if !health.status().is_success() {
223 return Err(ProxyError::ExternalTool {
224 message: format!("prefill health returned HTTP {}", health.status()),
225 });
226 }
227 let response = client
228 .get(join_path(&prefill.bootstrap_addr, "/query"))
229 .send()
230 .await
231 .map_err(|error| ProxyError::ExternalTool {
232 message: format!("prefill bootstrap query failed: {error}"),
233 })?;
234 if !response.status().is_success() {
235 return Err(ProxyError::ExternalTool {
236 message: format!(
237 "prefill bootstrap query returned HTTP {}",
238 response.status()
239 ),
240 });
241 }
242 let body = response
243 .json::<Value>()
244 .await
245 .map_err(|error| ProxyError::ExternalTool {
246 message: format!("prefill bootstrap query returned invalid JSON: {error}"),
247 })?;
248 let engine_ids = parse_engine_ids(&body)?;
249 *prefill.engine_ids.write().await = engine_ids;
250 Ok(())
251}
252
253fn parse_engine_ids(body: &Value) -> Result<Vec<String>, ProxyError> {
254 let object = body.as_object().ok_or_else(|| ProxyError::ExternalTool {
255 message: "prefill bootstrap query JSON must be an object".to_owned(),
256 })?;
257 if object.is_empty() {
258 return Err(ProxyError::ExternalTool {
259 message: "prefill bootstrap query returned no data-parallel engines".to_owned(),
260 });
261 }
262 let mut ranks = Vec::new();
263 for (rank_text, entry) in object {
264 let rank = rank_text
265 .parse::<usize>()
266 .map_err(|error| ProxyError::ExternalTool {
267 message: format!("invalid data-parallel rank {rank_text:?}: {error}"),
268 })?;
269 let engine_id = entry
270 .get("engine_id")
271 .and_then(Value::as_str)
272 .ok_or_else(|| ProxyError::ExternalTool {
273 message: format!("missing engine_id for data-parallel rank {rank_text:?}"),
274 })?;
275 ranks.push((rank, engine_id.to_owned()));
276 }
277 ranks.sort_by_key(|(rank, _engine_id)| *rank);
278 for (expected, (rank, _engine_id)) in ranks.iter().enumerate() {
279 if expected != *rank {
280 return Err(ProxyError::ExternalTool {
281 message: "prefill bootstrap query ranks must be contiguous from 0".to_owned(),
282 });
283 }
284 }
285 Ok(ranks
286 .into_iter()
287 .map(|(_rank, engine_id)| engine_id)
288 .collect())
289}
290
291async fn healthcheck(
292 State(state): State<ProxyState>,
293) -> (StatusCode, Json<ProxyHealthcheckResponse>) {
294 let ready = state.ready();
295 let status = if ready {
296 StatusCode::OK
297 } else {
298 StatusCode::SERVICE_UNAVAILABLE
299 };
300 (
301 status,
302 Json(ProxyHealthcheckResponse {
303 ready,
304 prefill_instances: state.inner.prefill.len(),
305 decode_instances: state.inner.decode.len(),
306 }),
307 )
308}
309
310async fn models(State(state): State<ProxyState>) -> Result<Response<Body>, ProxyHttpError> {
311 if !state.ready() {
312 return Err(ProxyHttpError::status(
313 StatusCode::SERVICE_UNAVAILABLE,
314 "proxy is not ready",
315 ));
316 }
317 let decode_url = state.next_decode_url()?;
318 let response = state
319 .client()
320 .get(join_path(&decode_url, "/v1/models"))
321 .send()
322 .await
323 .map_err(|error| ProxyHttpError::upstream("decode /v1/models request failed", error))?;
324 forward_response(response).await
325}
326
327async fn completions(
328 State(state): State<ProxyState>,
329 headers: HeaderMap,
330 Json(body): Json<Value>,
331) -> Result<Response<Body>, ProxyHttpError> {
332 completion_route(state, headers, body, "/v1/completions").await
333}
334
335async fn chat_completions(
336 State(state): State<ProxyState>,
337 headers: HeaderMap,
338 Json(body): Json<Value>,
339) -> Result<Response<Body>, ProxyHttpError> {
340 completion_route(state, headers, body, "/v1/chat/completions").await
341}
342
343async fn completion_route(
344 state: ProxyState,
345 headers: HeaderMap,
346 body: Value,
347 path: &'static str,
348) -> Result<Response<Body>, ProxyHttpError> {
349 if !state.ready() {
350 return Err(ProxyHttpError::status(
351 StatusCode::SERVICE_UNAVAILABLE,
352 "proxy is not ready",
353 ));
354 }
355 let selected_prefill = state.next_prefill().await?;
356 let decode_url = state.next_decode_url()?;
357 let request_id = state.request_id();
358 let authorization = outbound_authorization(&headers);
359 let client = state.client();
360 let prefill_body = prefill_body(&body, &request_id)?;
361 let decode_body = decode_body(&body, &selected_prefill, &request_id)?;
362 let prefill_task = tokio::spawn(send_prefill_request(
363 client.clone(),
364 selected_prefill.clone(),
365 path,
366 prefill_body,
367 request_id.clone(),
368 authorization.clone(),
369 ));
370 let decode_response = core::send_json_post(
371 client,
372 join_path(&decode_url, path),
373 &decode_body,
374 Some(&request_id),
375 authorization.as_deref(),
376 &[],
377 "decode request",
378 )
379 .await?;
380 core::stream_decode_response(decode_response, prefill_task)
381}
382
383fn prefill_body(body: &Value, request_id: &str) -> Result<Value, ProxyHttpError> {
384 let mut body = body.clone();
385 let Some(object) = body.as_object_mut() else {
386 return Err(ProxyHttpError::status(
387 StatusCode::BAD_REQUEST,
388 "OpenAI completion request body must be a JSON object",
389 ));
390 };
391 object.insert(
392 "kv_transfer_params".to_owned(),
393 MooncakePrefillKvTransferParams::new(request_id).into_protocol_value()?,
394 );
395 object.insert("stream".to_owned(), Value::Bool(false));
396 object.insert("max_tokens".to_owned(), Value::from(1_u8));
397 if object
398 .get("min_tokens")
399 .and_then(Value::as_u64)
400 .is_some_and(|min_tokens| min_tokens > 1)
401 {
402 object.insert("min_tokens".to_owned(), Value::from(1_u8));
403 }
404 if object.contains_key("max_completion_tokens") {
405 object.insert("max_completion_tokens".to_owned(), Value::from(1_u8));
406 }
407 object.remove("stream_options");
408 Ok(body)
409}
410
411fn decode_body(
412 body: &Value,
413 selected_prefill: &SelectedPrefill,
414 request_id: &str,
415) -> Result<Value, ProxyHttpError> {
416 let mut body = body.clone();
417 let Some(object) = body.as_object_mut() else {
418 return Err(ProxyHttpError::status(
419 StatusCode::BAD_REQUEST,
420 "OpenAI completion request body must be a JSON object",
421 ));
422 };
423 object.insert(
424 "kv_transfer_params".to_owned(),
425 MooncakeDecodeKvTransferParams::new(selected_prefill, request_id).into_protocol_value()?,
426 );
427 Ok(body)
428}
429
430#[derive(Serialize)]
431struct MooncakePrefillKvTransferParams {
432 do_remote_decode: bool,
433 do_remote_prefill: bool,
434 transfer_id: String,
435}
436
437impl MooncakePrefillKvTransferParams {
438 fn new(request_id: &str) -> Self {
439 Self {
440 do_remote_decode: true,
441 do_remote_prefill: false,
442 transfer_id: format!("xfer-{request_id}"),
443 }
444 }
445
446 fn into_protocol_value(self) -> Result<Value, ProxyHttpError> {
447 serde_json::to_value(self).map_err(|error| {
448 ProxyHttpError::internal(format!(
449 "failed to serialize vLLM Mooncake prefill transfer params: {error}"
450 ))
451 })
452 }
453}
454
455#[derive(Serialize)]
456struct MooncakeDecodeKvTransferParams<'a> {
457 do_remote_decode: bool,
458 do_remote_prefill: bool,
459 remote_bootstrap_addr: &'a str,
460 remote_engine_id: &'a str,
461 transfer_id: String,
462}
463
464impl<'a> MooncakeDecodeKvTransferParams<'a> {
465 fn new(selected_prefill: &'a SelectedPrefill, request_id: &str) -> Self {
466 Self {
467 do_remote_decode: false,
468 do_remote_prefill: true,
469 remote_bootstrap_addr: &selected_prefill.bootstrap_addr,
470 remote_engine_id: &selected_prefill.engine_id,
471 transfer_id: format!("xfer-{request_id}"),
472 }
473 }
474
475 fn into_protocol_value(self) -> Result<Value, ProxyHttpError> {
476 serde_json::to_value(self).map_err(|error| {
477 ProxyHttpError::internal(format!(
478 "failed to serialize vLLM Mooncake decode transfer params: {error}"
479 ))
480 })
481 }
482}
483
484async fn send_prefill_request(
485 client: reqwest::Client,
486 selected_prefill: SelectedPrefill,
487 path: &'static str,
488 body: Value,
489 request_id: String,
490 authorization: Option<String>,
491) -> Result<(), ProxyHttpError> {
492 let response = core::send_json_post(
493 client,
494 join_path(&selected_prefill.url, path),
495 &body,
496 Some(&request_id),
497 authorization.as_deref(),
498 &[("X-data-parallel-rank", selected_prefill.dp_rank.to_string())],
499 "prefill request",
500 )
501 .await?;
502 response
503 .bytes()
504 .await
505 .map_err(|error| ProxyHttpError::upstream("prefill response drain failed", error))?;
506 Ok(())
507}
508
509#[cfg(test)]
510mod tests {
511 use super::*;
512 use anyhow::Result;
513 use serde_json::json;
514
515 #[test]
516 fn meta_exports_byte_stable_proxy_identity() {
517 assert_eq!(ID, "inferlab-vllm-mooncake-proxy");
521 assert_eq!(VERSION, 1);
522 assert_eq!(meta().id, ID);
523 assert_eq!(meta().version, VERSION);
524 }
525
526 #[tokio::test]
527 async fn healthcheck_response_reports_readiness_and_configured_instances() -> Result<()> {
528 let state = ProxyState::new(Config {
529 host: "127.0.0.1".to_owned(),
530 port: 8000,
531 prefill: vec![PrefillTarget {
532 url: "http://127.0.0.1:8010".to_owned(),
533 bootstrap_url: "http://127.0.0.1:8998".to_owned(),
534 }],
535 decode: vec![
536 "http://127.0.0.1:8020".to_owned(),
537 "http://127.0.0.1:8021".to_owned(),
538 ],
539 })?;
540 let (status, Json(response)) = healthcheck(State(state.clone())).await;
541 assert_eq!(status, StatusCode::SERVICE_UNAVAILABLE);
542 assert!(!response.ready);
543
544 state.set_ready();
545 let (status, Json(response)) = healthcheck(State(state)).await;
546 let value = serde_json::to_value(response)?;
547
548 assert_eq!(status, StatusCode::OK);
549 assert_eq!(value.get("ready").and_then(Value::as_bool), Some(true));
550 assert_eq!(
551 value.get("prefill_instances").and_then(Value::as_u64),
552 Some(1)
553 );
554 assert_eq!(
555 value.get("decode_instances").and_then(Value::as_u64),
556 Some(2)
557 );
558 Ok(())
559 }
560
561 #[test]
562 fn prefill_body_forces_single_token_non_streaming_transfer() -> Result<()> {
563 let body = json!({
564 "model": "m",
565 "prompt": "hello",
566 "stream": true,
567 "stream_options": {"include_usage": true},
568 "max_tokens": 64,
569 "max_completion_tokens": 64,
570 "min_tokens": 64,
571 });
572 let lowered =
573 prefill_body(&body, "request-1").map_err(|error| anyhow::anyhow!(error.to_string()))?;
574 assert_eq!(
575 lowered.pointer("/kv_transfer_params/do_remote_decode"),
576 Some(&Value::Bool(true))
577 );
578 assert_eq!(
579 lowered.pointer("/kv_transfer_params/do_remote_prefill"),
580 Some(&Value::Bool(false))
581 );
582 assert_eq!(
583 lowered
584 .pointer("/kv_transfer_params/transfer_id")
585 .and_then(Value::as_str),
586 Some("xfer-request-1")
587 );
588 assert_eq!(lowered.get("stream"), Some(&Value::Bool(false)));
589 assert_eq!(lowered.get("max_tokens").and_then(Value::as_u64), Some(1));
590 assert_eq!(lowered.get("min_tokens").and_then(Value::as_u64), Some(1));
591 assert_eq!(
592 lowered.get("max_completion_tokens").and_then(Value::as_u64),
593 Some(1)
594 );
595 assert!(lowered.get("stream_options").is_none());
596 Ok(())
597 }
598
599 #[test]
600 fn decode_body_attaches_remote_prefill_identity() -> Result<()> {
601 let selected = SelectedPrefill {
602 url: "http://127.0.0.1:8010".to_owned(),
603 bootstrap_addr: "http://127.0.0.1:8998".to_owned(),
604 dp_rank: 0,
605 engine_id: "engine-a".to_owned(),
606 };
607 let lowered = decode_body(&json!({"model": "m"}), &selected, "request-2")
608 .map_err(|error| anyhow::anyhow!(error.to_string()))?;
609 assert_eq!(
610 lowered.pointer("/kv_transfer_params/do_remote_decode"),
611 Some(&Value::Bool(false))
612 );
613 assert_eq!(
614 lowered.pointer("/kv_transfer_params/do_remote_prefill"),
615 Some(&Value::Bool(true))
616 );
617 assert_eq!(
618 lowered
619 .pointer("/kv_transfer_params/remote_bootstrap_addr")
620 .and_then(Value::as_str),
621 Some("http://127.0.0.1:8998")
622 );
623 assert_eq!(
624 lowered
625 .pointer("/kv_transfer_params/remote_engine_id")
626 .and_then(Value::as_str),
627 Some("engine-a")
628 );
629 assert_eq!(
630 lowered
631 .pointer("/kv_transfer_params/transfer_id")
632 .and_then(Value::as_str),
633 Some("xfer-request-2")
634 );
635 Ok(())
636 }
637
638 #[test]
639 fn prefill_client_uses_explicit_bootstrap_url() -> Result<()> {
640 let client = PrefillClient::from_target(PrefillTarget {
641 url: "http://10.0.0.1:8010".to_owned(),
642 bootstrap_url: "http://192.0.2.10:8998".to_owned(),
643 })?;
644 assert_eq!(client.url, "http://10.0.0.1:8010");
645 assert_eq!(client.bootstrap_addr, "http://192.0.2.10:8998");
646 Ok(())
647 }
648
649 #[tokio::test]
650 async fn static_backend_selection_round_robins_prefill_engines_and_decode_urls() -> Result<()> {
651 let state = ProxyState::new(Config {
652 host: "127.0.0.1".to_owned(),
653 port: 8000,
654 prefill: vec![
655 PrefillTarget {
656 url: "http://127.0.0.1:8010".to_owned(),
657 bootstrap_url: "http://127.0.0.1:8998".to_owned(),
658 },
659 PrefillTarget {
660 url: "http://127.0.0.1:8011".to_owned(),
661 bootstrap_url: "http://127.0.0.1:8999".to_owned(),
662 },
663 ],
664 decode: vec![
665 "http://127.0.0.1:8020".to_owned(),
666 "http://127.0.0.1:8021".to_owned(),
667 ],
668 })?;
669 *state.inner.prefill[0].engine_ids.write().await =
670 vec!["p0-r0".to_owned(), "p0-r1".to_owned()];
671 *state.inner.prefill[1].engine_ids.write().await = vec!["p1-r0".to_owned()];
672
673 let prefill0 = state
674 .next_prefill()
675 .await
676 .map_err(|error| anyhow::anyhow!(error.to_string()))?;
677 let prefill1 = state
678 .next_prefill()
679 .await
680 .map_err(|error| anyhow::anyhow!(error.to_string()))?;
681 let prefill2 = state
682 .next_prefill()
683 .await
684 .map_err(|error| anyhow::anyhow!(error.to_string()))?;
685 assert_eq!(prefill0.engine_id, "p0-r0");
686 assert_eq!(prefill1.engine_id, "p0-r1");
687 assert_eq!(prefill2.engine_id, "p1-r0");
688 assert_eq!(
689 state
690 .next_prefill()
691 .await
692 .map_err(|error| anyhow::anyhow!(error.to_string()))?
693 .engine_id,
694 "p0-r0"
695 );
696
697 assert_eq!(state.next_decode_url()?, "http://127.0.0.1:8020");
698 assert_eq!(state.next_decode_url()?, "http://127.0.0.1:8021");
699 assert_eq!(state.next_decode_url()?, "http://127.0.0.1:8020");
700 Ok(())
701 }
702}