Skip to main content

inferlab_proxy/
vllm_mooncake.rs

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