1use crate::core::{
2 self, ProxyHealthcheckResponse, ProxyHttpError, ProxyMeta, forward_response, join_path,
3 outbound_authorization,
4};
5use crate::error::ProxyError;
6use axum::body::Body;
7use axum::extract::{Json, State};
8use axum::http::{HeaderMap, Response, StatusCode};
9use axum::routing::{get, post};
10use axum::{Router, serve};
11use serde::Serialize;
12use serde_json::Value;
13use std::sync::Arc;
14use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
15use std::time::Duration;
16use tokio::net::TcpListener;
17use tokio::sync::RwLock;
18
19pub const ID: &str = "inferlab-vllm-mooncake-proxy";
21pub const VERSION: u32 = 1;
23
24pub fn meta() -> ProxyMeta {
26 ProxyMeta {
27 id: ID,
28 version: VERSION,
29 }
30}
31
32#[derive(Clone, Debug)]
33pub struct Config {
34 pub host: String,
35 pub port: u16,
36 pub prefill: Vec<PrefillTarget>,
37 pub decode: Vec<String>,
38}
39
40#[derive(Clone, Debug)]
41pub struct PrefillTarget {
42 pub url: String,
43 pub bootstrap_url: String,
44}
45
46pub fn run(config: Config) -> Result<(), ProxyError> {
47 core::run(|| run_async(config))
48}
49
50pub async fn run_async(config: Config) -> Result<(), ProxyError> {
51 let host = config.host.clone();
52 let port = config.port;
53 let state = ProxyState::new(config)?;
54 tokio::spawn(discover_prefillers(state.clone()));
55 let app = router(state);
56 let listener = TcpListener::bind((host.as_str(), port))
57 .await
58 .map_err(|error| ProxyError::Io {
59 message: format!("failed to bind vLLM Mooncake proxy on {host}:{port}: {error}"),
60 })?;
61 serve(listener, app).await.map_err(|error| ProxyError::Io {
62 message: format!("vLLM Mooncake proxy server failed: {error}"),
63 })
64}
65
66fn router(state: ProxyState) -> Router {
67 Router::new()
68 .route("/healthcheck", get(healthcheck))
69 .route("/v1/models", get(models))
70 .route("/v1/completions", post(completions))
71 .route("/v1/chat/completions", post(chat_completions))
72 .with_state(state)
73}
74
75#[derive(Clone)]
76struct ProxyState {
77 inner: Arc<ProxyStateInner>,
78}
79
80struct ProxyStateInner {
81 client: reqwest::Client,
82 prefill: Vec<PrefillClient>,
83 decode: Vec<String>,
84 ready: AtomicBool,
85 prefill_cursor: AtomicUsize,
86 decode_cursor: AtomicUsize,
87 request_counter: AtomicUsize,
88}
89
90#[derive(Clone)]
91struct PrefillClient {
92 url: String,
93 bootstrap_addr: String,
94 engine_ids: Arc<RwLock<Vec<String>>>,
95}
96
97#[derive(Clone)]
98struct SelectedPrefill {
99 url: String,
100 bootstrap_addr: String,
101 dp_rank: usize,
102 engine_id: String,
103}
104
105impl ProxyState {
106 fn new(config: Config) -> Result<Self, ProxyError> {
107 if config.prefill.is_empty() {
108 return Err(ProxyError::Invalid {
109 message: "vLLM Mooncake proxy requires at least one prefill endpoint".to_owned(),
110 });
111 }
112 if config.decode.is_empty() {
113 return Err(ProxyError::Invalid {
114 message: "vLLM Mooncake proxy requires at least one decode endpoint".to_owned(),
115 });
116 }
117 let client = core::build_pooled_client().map_err(|error| ProxyError::Io {
118 message: format!("failed to create vLLM Mooncake proxy HTTP client: {error}"),
119 })?;
120 let prefill = config
121 .prefill
122 .into_iter()
123 .map(PrefillClient::from_target)
124 .collect::<Result<Vec<_>, ProxyError>>()?;
125 Ok(Self {
126 inner: Arc::new(ProxyStateInner {
127 client,
128 prefill,
129 decode: config.decode,
130 ready: AtomicBool::new(false),
131 prefill_cursor: AtomicUsize::new(0),
132 decode_cursor: AtomicUsize::new(0),
133 request_counter: AtomicUsize::new(0),
134 }),
135 })
136 }
137
138 fn client(&self) -> reqwest::Client {
139 self.inner.client.clone()
140 }
141
142 fn ready(&self) -> bool {
143 self.inner.ready.load(Ordering::SeqCst)
144 }
145
146 fn set_ready(&self) {
147 self.inner.ready.store(true, Ordering::SeqCst);
148 }
149
150 async fn next_prefill(&self) -> Result<SelectedPrefill, ProxyHttpError> {
151 let mut candidates = Vec::new();
152 for prefill in &self.inner.prefill {
153 let engine_ids = prefill.engine_ids.read().await;
154 for (dp_rank, engine_id) in engine_ids.iter().enumerate() {
155 candidates.push(SelectedPrefill {
156 url: prefill.url.clone(),
157 bootstrap_addr: prefill.bootstrap_addr.clone(),
158 dp_rank,
159 engine_id: engine_id.clone(),
160 });
161 }
162 }
163 if candidates.is_empty() {
164 return Err(ProxyHttpError::status(
165 StatusCode::SERVICE_UNAVAILABLE,
166 "no ready prefill data-parallel engines",
167 ));
168 }
169 let index = core::round_robin_index(&self.inner.prefill_cursor, candidates.len());
170 Ok(candidates.swap_remove(index))
171 }
172
173 fn next_decode_url(&self) -> Result<String, ProxyHttpError> {
174 if self.inner.decode.is_empty() {
175 return Err(ProxyHttpError::status(
176 StatusCode::SERVICE_UNAVAILABLE,
177 "no decode endpoints configured",
178 ));
179 }
180 let index = core::round_robin_index(&self.inner.decode_cursor, self.inner.decode.len());
181 Ok(self.inner.decode[index].clone())
182 }
183
184 fn request_id(&self) -> String {
185 core::next_request_id(&self.inner.request_counter)
186 }
187}
188
189impl PrefillClient {
190 fn from_target(target: PrefillTarget) -> Result<Self, ProxyError> {
191 Ok(Self {
192 url: target.url,
193 bootstrap_addr: target.bootstrap_url,
194 engine_ids: Arc::new(RwLock::new(Vec::new())),
195 })
196 }
197}
198
199async fn discover_prefillers(state: ProxyState) {
200 for prefill in &state.inner.prefill {
201 loop {
202 if discover_prefiller(&state.client(), prefill).await.is_ok() {
203 break;
204 }
205 tokio::time::sleep(Duration::from_secs(1)).await;
206 }
207 }
208 state.set_ready();
209}
210
211async fn discover_prefiller(
212 client: &reqwest::Client,
213 prefill: &PrefillClient,
214) -> Result<(), ProxyError> {
215 let health = client
216 .get(join_path(&prefill.url, "/health"))
217 .send()
218 .await
219 .map_err(|error| ProxyError::ExternalTool {
220 message: format!("prefill health request failed: {error}"),
221 })?;
222 if !health.status().is_success() {
223 return Err(ProxyError::ExternalTool {
224 message: format!("prefill health returned HTTP {}", health.status()),
225 });
226 }
227 let response = client
228 .get(join_path(&prefill.bootstrap_addr, "/query"))
229 .send()
230 .await
231 .map_err(|error| ProxyError::ExternalTool {
232 message: format!("prefill bootstrap query failed: {error}"),
233 })?;
234 if !response.status().is_success() {
235 return Err(ProxyError::ExternalTool {
236 message: format!(
237 "prefill bootstrap query returned HTTP {}",
238 response.status()
239 ),
240 });
241 }
242 let body = response
243 .json::<Value>()
244 .await
245 .map_err(|error| ProxyError::ExternalTool {
246 message: format!("prefill bootstrap query returned invalid JSON: {error}"),
247 })?;
248 let engine_ids = parse_engine_ids(&body)?;
249 *prefill.engine_ids.write().await = engine_ids;
250 Ok(())
251}
252
253fn parse_engine_ids(body: &Value) -> Result<Vec<String>, ProxyError> {
254 let object = body.as_object().ok_or_else(|| ProxyError::ExternalTool {
255 message: "prefill bootstrap query JSON must be an object".to_owned(),
256 })?;
257 if object.is_empty() {
258 return Err(ProxyError::ExternalTool {
259 message: "prefill bootstrap query returned no data-parallel engines".to_owned(),
260 });
261 }
262 let mut ranks = Vec::new();
263 for (rank_text, entry) in object {
264 let rank = rank_text
265 .parse::<usize>()
266 .map_err(|error| ProxyError::ExternalTool {
267 message: format!("invalid data-parallel rank {rank_text:?}: {error}"),
268 })?;
269 let engine_id = entry
270 .get("engine_id")
271 .and_then(Value::as_str)
272 .ok_or_else(|| ProxyError::ExternalTool {
273 message: format!("missing engine_id for data-parallel rank {rank_text:?}"),
274 })?;
275 ranks.push((rank, engine_id.to_owned()));
276 }
277 ranks.sort_by_key(|(rank, _engine_id)| *rank);
278 for (expected, (rank, _engine_id)) in ranks.iter().enumerate() {
279 if expected != *rank {
280 return Err(ProxyError::ExternalTool {
281 message: "prefill bootstrap query ranks must be contiguous from 0".to_owned(),
282 });
283 }
284 }
285 Ok(ranks
286 .into_iter()
287 .map(|(_rank, engine_id)| engine_id)
288 .collect())
289}
290
291async fn healthcheck(
292 State(state): State<ProxyState>,
293) -> (StatusCode, Json<ProxyHealthcheckResponse>) {
294 let ready = state.ready();
295 let status = if ready {
296 StatusCode::OK
297 } else {
298 StatusCode::SERVICE_UNAVAILABLE
299 };
300 (
301 status,
302 Json(ProxyHealthcheckResponse {
303 ready,
304 prefill_instances: state.inner.prefill.len(),
305 decode_instances: state.inner.decode.len(),
306 }),
307 )
308}
309
310async fn models(State(state): State<ProxyState>) -> Result<Response<Body>, ProxyHttpError> {
311 if !state.ready() {
312 return Err(ProxyHttpError::status(
313 StatusCode::SERVICE_UNAVAILABLE,
314 "proxy is not ready",
315 ));
316 }
317 let decode_url = state.next_decode_url()?;
318 let response = state
319 .client()
320 .get(join_path(&decode_url, "/v1/models"))
321 .send()
322 .await
323 .map_err(|error| ProxyHttpError::upstream("decode /v1/models request failed", error))?;
324 forward_response(response).await
325}
326
327async fn completions(
328 State(state): State<ProxyState>,
329 headers: HeaderMap,
330 Json(body): Json<Value>,
331) -> Result<Response<Body>, ProxyHttpError> {
332 completion_route(state, headers, body, "/v1/completions").await
333}
334
335async fn chat_completions(
336 State(state): State<ProxyState>,
337 headers: HeaderMap,
338 Json(body): Json<Value>,
339) -> Result<Response<Body>, ProxyHttpError> {
340 completion_route(state, headers, body, "/v1/chat/completions").await
341}
342
343async fn completion_route(
344 state: ProxyState,
345 headers: HeaderMap,
346 body: Value,
347 path: &'static str,
348) -> Result<Response<Body>, ProxyHttpError> {
349 if !state.ready() {
350 return Err(ProxyHttpError::status(
351 StatusCode::SERVICE_UNAVAILABLE,
352 "proxy is not ready",
353 ));
354 }
355 let selected_prefill = state.next_prefill().await?;
356 let decode_url = state.next_decode_url()?;
357 let request_id = state.request_id();
358 let authorization = outbound_authorization(&headers);
359 let client = state.client();
360 let prefill_body = prefill_body(&body, &request_id)?;
361 let decode_body = decode_body(&body, &selected_prefill, &request_id)?;
362 let prefill_task = tokio::spawn(send_prefill_request(
363 client.clone(),
364 selected_prefill.clone(),
365 path,
366 prefill_body,
367 request_id.clone(),
368 authorization.clone(),
369 ));
370 let decode_response = core::send_json_post(
371 client,
372 join_path(&decode_url, path),
373 &decode_body,
374 Some(&request_id),
375 authorization.as_deref(),
376 &[],
377 "decode request",
378 )
379 .await?;
380 core::stream_decode_response(decode_response, prefill_task)
381}
382
383fn prefill_body(body: &Value, request_id: &str) -> Result<Value, ProxyHttpError> {
384 let mut body = body.clone();
385 let Some(object) = body.as_object_mut() else {
386 return Err(ProxyHttpError::status(
387 StatusCode::BAD_REQUEST,
388 "OpenAI completion request body must be a JSON object",
389 ));
390 };
391 object.insert(
392 "kv_transfer_params".to_owned(),
393 MooncakePrefillKvTransferParams::new(request_id).into_protocol_value()?,
394 );
395 object.insert("stream".to_owned(), Value::Bool(false));
396 object.insert("max_tokens".to_owned(), Value::from(1_u8));
397 if object
398 .get("min_tokens")
399 .and_then(Value::as_u64)
400 .is_some_and(|min_tokens| min_tokens > 1)
401 {
402 object.insert("min_tokens".to_owned(), Value::from(1_u8));
403 }
404 if object.contains_key("max_completion_tokens") {
405 object.insert("max_completion_tokens".to_owned(), Value::from(1_u8));
406 }
407 object.remove("stream_options");
408 Ok(body)
409}
410
411fn decode_body(
412 body: &Value,
413 selected_prefill: &SelectedPrefill,
414 request_id: &str,
415) -> Result<Value, ProxyHttpError> {
416 let mut body = body.clone();
417 let Some(object) = body.as_object_mut() else {
418 return Err(ProxyHttpError::status(
419 StatusCode::BAD_REQUEST,
420 "OpenAI completion request body must be a JSON object",
421 ));
422 };
423 object.insert(
424 "kv_transfer_params".to_owned(),
425 MooncakeDecodeKvTransferParams::new(selected_prefill, request_id).into_protocol_value()?,
426 );
427 Ok(body)
428}
429
430#[derive(Serialize)]
431struct MooncakePrefillKvTransferParams {
432 do_remote_decode: bool,
433 do_remote_prefill: bool,
434 transfer_id: String,
435}
436
437impl MooncakePrefillKvTransferParams {
438 fn new(request_id: &str) -> Self {
439 Self {
440 do_remote_decode: true,
441 do_remote_prefill: false,
442 transfer_id: format!("xfer-{request_id}"),
443 }
444 }
445
446 fn into_protocol_value(self) -> Result<Value, ProxyHttpError> {
447 serde_json::to_value(self).map_err(|error| {
448 ProxyHttpError::internal(format!(
449 "failed to serialize vLLM Mooncake prefill transfer params: {error}"
450 ))
451 })
452 }
453}
454
455#[derive(Serialize)]
456struct MooncakeDecodeKvTransferParams<'a> {
457 do_remote_decode: bool,
458 do_remote_prefill: bool,
459 remote_bootstrap_addr: &'a str,
460 remote_engine_id: &'a str,
461 transfer_id: String,
462}
463
464impl<'a> MooncakeDecodeKvTransferParams<'a> {
465 fn new(selected_prefill: &'a SelectedPrefill, request_id: &str) -> Self {
466 Self {
467 do_remote_decode: false,
468 do_remote_prefill: true,
469 remote_bootstrap_addr: &selected_prefill.bootstrap_addr,
470 remote_engine_id: &selected_prefill.engine_id,
471 transfer_id: format!("xfer-{request_id}"),
472 }
473 }
474
475 fn into_protocol_value(self) -> Result<Value, ProxyHttpError> {
476 serde_json::to_value(self).map_err(|error| {
477 ProxyHttpError::internal(format!(
478 "failed to serialize vLLM Mooncake decode transfer params: {error}"
479 ))
480 })
481 }
482}
483
484async fn send_prefill_request(
485 client: reqwest::Client,
486 selected_prefill: SelectedPrefill,
487 path: &'static str,
488 body: Value,
489 request_id: String,
490 authorization: Option<String>,
491) -> Result<(), ProxyHttpError> {
492 let response = core::send_json_post(
493 client,
494 join_path(&selected_prefill.url, path),
495 &body,
496 Some(&request_id),
497 authorization.as_deref(),
498 &[("X-data-parallel-rank", selected_prefill.dp_rank.to_string())],
499 "prefill request",
500 )
501 .await?;
502 response
503 .bytes()
504 .await
505 .map_err(|error| ProxyHttpError::upstream("prefill response drain failed", error))?;
506 Ok(())
507}
508
509#[cfg(test)]
510mod tests {
511 use super::*;
512 use anyhow::Result;
513 use 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 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}