Skip to main content

inferlab_proxy/
vllm_mooncake.rs

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
27/// Display name used in lifecycle/validation error messages.
28const 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
478/// One paired prefill/decode conditioning flow with the prefill request
479/// pinned to `selected_prefill.dp_rank`; the decode side rides the ordinary
480/// round-robin pairing and is incidental coverage
481/// ([[RFC-0004:C-BENCH-CACHE-STATE]]).
482async 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        // Each flow rides the ordinary prefill/decode pairing.
1085        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    /// Readiness without any discovered data-parallel engine must not read
1139    /// as a successful fan-out: with nothing to prime, the endpoint answers
1140    /// 502 instead of a 200 over an empty target set.
1141    #[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}