Skip to main content

inferlab_proxy/
vllm_mooncake.rs

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
19/// Identity recorded in `BuiltinProxy` evidence for the Mooncake proxy.
20pub const ID: &str = "inferlab-vllm-mooncake-proxy";
21/// Evidence version for the Mooncake proxy identity.
22pub const VERSION: u32 = 1;
23
24/// Owned identity of the built-in Mooncake proxy.
25pub 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        // AC4: the Mooncake proxy owns and exports its own id+version. These exact
518        // strings/numbers are persisted in BuiltinProxy evidence, so they must stay
519        // byte-stable.
520        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}