Skip to main content

inferlab_proxy/
vllm_nixl.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::{Map, Value};
13use std::sync::Arc;
14use std::sync::atomic::AtomicUsize;
15use tokio::net::TcpListener;
16
17/// Identity recorded in `BuiltinProxy` evidence for the NIXL proxy.
18pub const ID: &str = "inferlab-vllm-nixl-proxy";
19/// Evidence version for the NIXL proxy identity.
20pub const VERSION: u32 = 1;
21
22/// Owned identity of the built-in NIXL proxy.
23pub 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        // AC4: the NIXL proxy owns and exports its own id+version. These exact
312        // strings/numbers are persisted in BuiltinProxy evidence, so they must stay
313        // byte-stable.
314        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}