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::{Map, Value};
13use std::sync::Arc;
14use std::sync::atomic::AtomicUsize;
15use tokio::net::TcpListener;
16
17pub const ID: &str = "inferlab-vllm-nixl-proxy";
19pub const VERSION: u32 = 1;
21
22pub fn meta() -> ProxyMeta {
24 ProxyMeta {
25 id: ID,
26 version: VERSION,
27 }
28}
29
30#[derive(Clone, Debug)]
31pub struct Config {
32 pub host: String,
33 pub port: u16,
34 pub prefill: Vec<String>,
35 pub decode: Vec<String>,
36}
37
38pub fn run(config: Config) -> Result<(), ProxyError> {
39 core::run(|| run_async(config))
40}
41
42pub async fn run_async(config: Config) -> Result<(), ProxyError> {
43 let host = config.host.clone();
44 let port = config.port;
45 let state = ProxyState::new(config)?;
46 let app = router(state);
47 let listener = TcpListener::bind((host.as_str(), port))
48 .await
49 .map_err(|error| ProxyError::Io {
50 message: format!("failed to bind vLLM NIXL proxy on {host}:{port}: {error}"),
51 })?;
52 serve(listener, app).await.map_err(|error| ProxyError::Io {
53 message: format!("vLLM NIXL proxy server failed: {error}"),
54 })
55}
56
57fn router(state: ProxyState) -> Router {
58 Router::new()
59 .route("/healthcheck", get(healthcheck))
60 .route("/v1/models", get(models))
61 .route("/v1/completions", post(completions))
62 .route("/v1/chat/completions", post(chat_completions))
63 .with_state(state)
64}
65
66#[derive(Clone)]
67struct ProxyState {
68 inner: Arc<ProxyStateInner>,
69}
70
71struct ProxyStateInner {
72 client: reqwest::Client,
73 prefill: Vec<String>,
74 decode: Vec<String>,
75 prefill_cursor: AtomicUsize,
76 decode_cursor: AtomicUsize,
77 request_counter: AtomicUsize,
78}
79
80impl ProxyState {
81 fn new(config: Config) -> Result<Self, ProxyError> {
82 if config.prefill.is_empty() {
83 return Err(ProxyError::Invalid {
84 message: "vLLM NIXL proxy requires at least one prefill endpoint".to_owned(),
85 });
86 }
87 if config.decode.is_empty() {
88 return Err(ProxyError::Invalid {
89 message: "vLLM NIXL proxy requires at least one decode endpoint".to_owned(),
90 });
91 }
92 let client = core::build_pooled_client().map_err(|error| ProxyError::Io {
93 message: format!("failed to create vLLM NIXL proxy HTTP client: {error}"),
94 })?;
95 Ok(Self {
96 inner: Arc::new(ProxyStateInner {
97 client,
98 prefill: config.prefill,
99 decode: config.decode,
100 prefill_cursor: AtomicUsize::new(0),
101 decode_cursor: AtomicUsize::new(0),
102 request_counter: AtomicUsize::new(0),
103 }),
104 })
105 }
106
107 fn client(&self) -> reqwest::Client {
108 self.inner.client.clone()
109 }
110
111 fn next_prefill_url(&self) -> Result<String, ProxyHttpError> {
112 let index = core::round_robin_index(&self.inner.prefill_cursor, self.inner.prefill.len());
113 Ok(self.inner.prefill[index].clone())
114 }
115
116 fn next_decode_url(&self) -> Result<String, ProxyHttpError> {
117 let index = core::round_robin_index(&self.inner.decode_cursor, self.inner.decode.len());
118 Ok(self.inner.decode[index].clone())
119 }
120
121 fn request_id(&self) -> String {
122 core::next_request_id(&self.inner.request_counter)
123 }
124}
125
126async fn healthcheck(State(state): State<ProxyState>) -> Json<ProxyHealthcheckResponse> {
127 Json(ProxyHealthcheckResponse {
128 ready: true,
129 prefill_instances: state.inner.prefill.len(),
130 decode_instances: state.inner.decode.len(),
131 })
132}
133
134async fn models(State(state): State<ProxyState>) -> Result<Response<Body>, ProxyHttpError> {
135 let decode_url = state.next_decode_url()?;
136 let response = state
137 .client()
138 .get(join_path(&decode_url, "/v1/models"))
139 .send()
140 .await
141 .map_err(|error| ProxyHttpError::upstream("decode /v1/models request failed", error))?;
142 forward_response(response).await
143}
144
145async fn completions(
146 State(state): State<ProxyState>,
147 headers: HeaderMap,
148 Json(body): Json<Value>,
149) -> Result<Response<Body>, ProxyHttpError> {
150 completion_route(state, headers, body, "/v1/completions").await
151}
152
153async fn chat_completions(
154 State(state): State<ProxyState>,
155 headers: HeaderMap,
156 Json(body): Json<Value>,
157) -> Result<Response<Body>, ProxyHttpError> {
158 completion_route(state, headers, body, "/v1/chat/completions").await
159}
160
161async fn completion_route(
162 state: ProxyState,
163 headers: HeaderMap,
164 body: Value,
165 path: &'static str,
166) -> Result<Response<Body>, ProxyHttpError> {
167 let prefill_url = state.next_prefill_url()?;
168 let decode_url = state.next_decode_url()?;
169 let request_id = state.request_id();
170 let authorization = outbound_authorization(&headers);
171 let client = state.client();
172 let prefill_body = prefill_body(&body, &request_id)?;
173 let prefill_response = send_prefill_request(
174 client.clone(),
175 &prefill_url,
176 path,
177 prefill_body,
178 &request_id,
179 authorization.as_deref(),
180 )
181 .await?;
182 let decode_body = decode_body(&body, prefill_response.kv_transfer_params)?;
183 let decode_response = core::send_json_post(
184 client,
185 join_path(&decode_url, path),
186 &decode_body,
187 Some(&request_id),
188 authorization.as_deref(),
189 &[],
190 "decode request",
191 )
192 .await?;
193 forward_response(decode_response).await
194}
195
196#[derive(Debug)]
197struct PrefillResponse {
198 kv_transfer_params: Value,
199}
200
201fn prefill_body(body: &Value, request_id: &str) -> Result<Value, ProxyHttpError> {
202 let mut body = body.clone();
203 let object = object_mut(&mut body)?;
204 object.insert(
205 "kv_transfer_params".to_owned(),
206 NixlPrefillKvTransferParams::new(request_id).into_protocol_value()?,
207 );
208 object.insert("stream".to_owned(), Value::Bool(false));
209 object.insert("max_tokens".to_owned(), Value::from(1_u8));
210 if object.contains_key("max_completion_tokens") {
211 object.insert("max_completion_tokens".to_owned(), Value::from(1_u8));
212 }
213 object.remove("stream_options");
214 object.remove("min_tokens");
215 object.remove("min_completion_tokens");
216 Ok(body)
217}
218
219#[derive(Serialize)]
220struct NixlPrefillKvTransferParams {
221 do_remote_decode: bool,
222 do_remote_prefill: bool,
223 remote_engine_id: Option<String>,
224 remote_block_ids: Option<Vec<u64>>,
225 remote_host: Option<String>,
226 remote_port: Option<u16>,
227 transfer_id: String,
228}
229
230impl NixlPrefillKvTransferParams {
231 fn new(request_id: &str) -> Self {
232 Self {
233 do_remote_decode: true,
234 do_remote_prefill: false,
235 remote_engine_id: None,
236 remote_block_ids: None,
237 remote_host: None,
238 remote_port: None,
239 transfer_id: format!("xfer-{request_id}"),
240 }
241 }
242
243 fn into_protocol_value(self) -> Result<Value, ProxyHttpError> {
244 serde_json::to_value(self).map_err(|error| {
245 ProxyHttpError::internal(format!(
246 "failed to serialize vLLM NIXL prefill transfer params: {error}"
247 ))
248 })
249 }
250}
251
252fn decode_body(body: &Value, kv_transfer_params: Value) -> Result<Value, ProxyHttpError> {
253 let mut body = body.clone();
254 let object = object_mut(&mut body)?;
255 object.insert("kv_transfer_params".to_owned(), kv_transfer_params);
256 Ok(body)
257}
258
259fn object_mut(body: &mut Value) -> Result<&mut Map<String, Value>, ProxyHttpError> {
260 body.as_object_mut().ok_or_else(|| {
261 ProxyHttpError::status(
262 StatusCode::BAD_REQUEST,
263 "OpenAI completion request body must be a JSON object",
264 )
265 })
266}
267
268async fn send_prefill_request(
269 client: reqwest::Client,
270 prefill_url: &str,
271 path: &'static str,
272 body: Value,
273 request_id: &str,
274 authorization: Option<&str>,
275) -> Result<PrefillResponse, ProxyHttpError> {
276 let response = core::send_json_post(
277 client,
278 join_path(prefill_url, path),
279 &body,
280 Some(request_id),
281 authorization,
282 &[],
283 "prefill request",
284 )
285 .await?;
286 let body = response
287 .json::<Value>()
288 .await
289 .map_err(|error| ProxyHttpError::upstream("prefill response JSON read failed", error))?;
290 let kv_transfer_params = body.get("kv_transfer_params").cloned().ok_or_else(|| {
291 ProxyHttpError::status(
292 StatusCode::BAD_GATEWAY,
293 "prefill response did not include kv_transfer_params",
294 )
295 })?;
296 Ok(PrefillResponse { kv_transfer_params })
297}
298
299#[cfg(test)]
300mod tests {
301 use super::*;
302 use anyhow::{Result, bail};
303 use axum::body::to_bytes;
304 use axum::response::IntoResponse;
305 use serde_json::json;
306 use tokio::sync::Mutex;
307 use tokio::task::JoinHandle;
308
309 #[test]
310 fn meta_exports_byte_stable_proxy_identity() {
311 assert_eq!(ID, "inferlab-vllm-nixl-proxy");
315 assert_eq!(VERSION, 1);
316 assert_eq!(meta().id, ID);
317 assert_eq!(meta().version, VERSION);
318 }
319
320 #[tokio::test]
321 async fn healthcheck_response_reports_configured_instances() -> Result<()> {
322 let state = ProxyState::new(Config {
323 host: "127.0.0.1".to_owned(),
324 port: 8000,
325 prefill: vec![
326 "http://127.0.0.1:8010".to_owned(),
327 "http://127.0.0.1:8011".to_owned(),
328 ],
329 decode: vec!["http://127.0.0.1:8020".to_owned()],
330 })?;
331
332 let Json(response) = healthcheck(State(state)).await;
333 let value = serde_json::to_value(response)?;
334
335 assert_eq!(value.get("ready").and_then(Value::as_bool), Some(true));
336 assert_eq!(
337 value.get("prefill_instances").and_then(Value::as_u64),
338 Some(2)
339 );
340 assert_eq!(
341 value.get("decode_instances").and_then(Value::as_u64),
342 Some(1)
343 );
344 Ok(())
345 }
346
347 #[test]
348 fn prefill_body_sets_nixl_prefill_transfer_params() -> Result<()> {
349 let body = json!({
350 "model": "m",
351 "prompt": "hello",
352 "stream": true,
353 "stream_options": {"include_usage": true},
354 "max_tokens": 64,
355 "max_completion_tokens": 64,
356 "min_tokens": 4,
357 });
358
359 let lowered =
360 prefill_body(&body, "request-1").map_err(|error| anyhow::anyhow!(error.to_string()))?;
361
362 assert_eq!(
363 lowered.pointer("/kv_transfer_params/do_remote_decode"),
364 Some(&Value::Bool(true))
365 );
366 assert_eq!(
367 lowered.pointer("/kv_transfer_params/do_remote_prefill"),
368 Some(&Value::Bool(false))
369 );
370 assert_eq!(
371 lowered
372 .pointer("/kv_transfer_params/transfer_id")
373 .and_then(Value::as_str),
374 Some("xfer-request-1")
375 );
376 assert_eq!(lowered.get("stream"), Some(&Value::Bool(false)));
377 assert_eq!(lowered.get("max_tokens").and_then(Value::as_u64), Some(1));
378 assert_eq!(
379 lowered.get("max_completion_tokens").and_then(Value::as_u64),
380 Some(1)
381 );
382 assert!(lowered.get("stream_options").is_none());
383 assert!(lowered.get("min_tokens").is_none());
384 Ok(())
385 }
386
387 #[test]
388 fn decode_body_forwards_prefill_kv_transfer_params() -> Result<()> {
389 let kv_transfer_params = json!({
390 "remote_engine_id": "engine-p",
391 "remote_host": "10.0.0.1",
392 "remote_port": 5600,
393 });
394
395 let lowered = decode_body(&json!({"model": "m"}), kv_transfer_params.clone())
396 .map_err(|error| anyhow::anyhow!(error.to_string()))?;
397
398 assert_eq!(lowered.get("kv_transfer_params"), Some(&kv_transfer_params));
399 Ok(())
400 }
401
402 #[tokio::test]
403 async fn chat_dispatch_preserves_route_messages_and_unowned_fields() -> Result<()> {
404 let prefill_backend = MockBackend::new(true);
405 let decode_backend = MockBackend::new(false);
406 let (prefill, prefill_server) = spawn_backend(prefill_backend.clone()).await?;
407 let (decode, decode_server) = spawn_backend(decode_backend.clone()).await?;
408 let state = ProxyState::new(Config {
409 host: "127.0.0.1".to_owned(),
410 port: 8000,
411 prefill: vec![prefill],
412 decode: vec![decode],
413 })?;
414 let request = json!({
415 "model": "m",
416 "messages": [{"role": "user", "content": "hello"}],
417 "temperature": 1.0,
418 "reasoning_effort": "high",
419 "chat_template_kwargs": {"enable_thinking": true}
420 });
421
422 let response = completion_route(
423 state,
424 HeaderMap::new(),
425 request.clone(),
426 "/v1/chat/completions",
427 )
428 .await
429 .map_err(|error| anyhow::anyhow!(error.to_string()))?;
430 assert_eq!(response.status(), StatusCode::OK);
431 let _body = to_bytes(response.into_body(), usize::MAX).await?;
432
433 let prefill_requests = prefill_backend.requests.lock().await;
434 let decode_requests = decode_backend.requests.lock().await;
435 assert_eq!(prefill_requests.len(), 1);
436 assert_eq!(decode_requests.len(), 1);
437 for key in [
438 "messages",
439 "temperature",
440 "reasoning_effort",
441 "chat_template_kwargs",
442 ] {
443 assert_eq!(prefill_requests[0][key], request[key]);
444 assert_eq!(decode_requests[0][key], request[key]);
445 }
446 prefill_server.abort();
447 decode_server.abort();
448 Ok(())
449 }
450
451 #[test]
452 fn proxy_state_requires_prefill_and_decode_targets() -> Result<()> {
453 let result = ProxyState::new(Config {
454 host: "127.0.0.1".to_owned(),
455 port: 8000,
456 prefill: Vec::new(),
457 decode: vec!["http://127.0.0.1:8020".to_owned()],
458 });
459 let error = match result {
460 Ok(_) => bail!("empty prefill targets should fail"),
461 Err(error) => error,
462 };
463 assert!(error.to_string().contains("at least one prefill endpoint"));
464 Ok(())
465 }
466
467 #[derive(Clone)]
468 struct MockBackend {
469 requests: Arc<Mutex<Vec<Value>>>,
470 prefill: bool,
471 }
472
473 impl MockBackend {
474 fn new(prefill: bool) -> Self {
475 Self {
476 requests: Arc::new(Mutex::new(Vec::new())),
477 prefill,
478 }
479 }
480 }
481
482 async fn mock_chat(
483 State(state): State<MockBackend>,
484 Json(body): Json<Value>,
485 ) -> Response<Body> {
486 state.requests.lock().await.push(body);
487 if state.prefill {
488 Json(json!({
489 "kv_transfer_params": {
490 "remote_engine_id": "prefill-0",
491 "remote_host": "127.0.0.1",
492 "remote_port": 5600
493 }
494 }))
495 .into_response()
496 } else {
497 Json(json!({"object": "chat.completion", "choices": []})).into_response()
498 }
499 }
500
501 async fn spawn_backend(state: MockBackend) -> Result<(String, JoinHandle<()>)> {
502 let app = Router::new()
503 .route("/v1/chat/completions", post(mock_chat))
504 .with_state(state);
505 let listener = TcpListener::bind("127.0.0.1:0").await?;
506 let address = listener.local_addr()?;
507 let server = tokio::spawn(async move {
508 let _result = serve(listener, app).await;
509 });
510 Ok((format!("http://{address}"), server))
511 }
512}