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