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 serde_json::json;
304
305    #[test]
306    fn meta_exports_byte_stable_proxy_identity() {
307        // AC4: the NIXL proxy owns and exports its own id+version. These exact
308        // strings/numbers are persisted in BuiltinProxy evidence, so they must stay
309        // byte-stable.
310        assert_eq!(ID, "inferlab-vllm-nixl-proxy");
311        assert_eq!(VERSION, 1);
312        assert_eq!(meta().id, ID);
313        assert_eq!(meta().version, VERSION);
314    }
315
316    #[tokio::test]
317    async fn healthcheck_response_reports_configured_instances() -> Result<()> {
318        let state = ProxyState::new(Config {
319            host: "127.0.0.1".to_owned(),
320            port: 8000,
321            prefill: vec![
322                "http://127.0.0.1:8010".to_owned(),
323                "http://127.0.0.1:8011".to_owned(),
324            ],
325            decode: vec!["http://127.0.0.1:8020".to_owned()],
326        })?;
327
328        let Json(response) = healthcheck(State(state)).await;
329        let value = serde_json::to_value(response)?;
330
331        assert_eq!(value.get("ready").and_then(Value::as_bool), Some(true));
332        assert_eq!(
333            value.get("prefill_instances").and_then(Value::as_u64),
334            Some(2)
335        );
336        assert_eq!(
337            value.get("decode_instances").and_then(Value::as_u64),
338            Some(1)
339        );
340        Ok(())
341    }
342
343    #[test]
344    fn prefill_body_sets_nixl_prefill_transfer_params() -> Result<()> {
345        let body = json!({
346            "model": "m",
347            "prompt": "hello",
348            "stream": true,
349            "stream_options": {"include_usage": true},
350            "max_tokens": 64,
351            "max_completion_tokens": 64,
352            "min_tokens": 4,
353        });
354
355        let lowered =
356            prefill_body(&body, "request-1").map_err(|error| anyhow::anyhow!(error.to_string()))?;
357
358        assert_eq!(
359            lowered.pointer("/kv_transfer_params/do_remote_decode"),
360            Some(&Value::Bool(true))
361        );
362        assert_eq!(
363            lowered.pointer("/kv_transfer_params/do_remote_prefill"),
364            Some(&Value::Bool(false))
365        );
366        assert_eq!(
367            lowered
368                .pointer("/kv_transfer_params/transfer_id")
369                .and_then(Value::as_str),
370            Some("xfer-request-1")
371        );
372        assert_eq!(lowered.get("stream"), Some(&Value::Bool(false)));
373        assert_eq!(lowered.get("max_tokens").and_then(Value::as_u64), Some(1));
374        assert_eq!(
375            lowered.get("max_completion_tokens").and_then(Value::as_u64),
376            Some(1)
377        );
378        assert!(lowered.get("stream_options").is_none());
379        assert!(lowered.get("min_tokens").is_none());
380        Ok(())
381    }
382
383    #[test]
384    fn decode_body_forwards_prefill_kv_transfer_params() -> Result<()> {
385        let kv_transfer_params = json!({
386            "remote_engine_id": "engine-p",
387            "remote_host": "10.0.0.1",
388            "remote_port": 5600,
389        });
390
391        let lowered = decode_body(&json!({"model": "m"}), kv_transfer_params.clone())
392            .map_err(|error| anyhow::anyhow!(error.to_string()))?;
393
394        assert_eq!(lowered.get("kv_transfer_params"), Some(&kv_transfer_params));
395        Ok(())
396    }
397
398    #[test]
399    fn proxy_state_requires_prefill_and_decode_targets() -> Result<()> {
400        let result = ProxyState::new(Config {
401            host: "127.0.0.1".to_owned(),
402            port: 8000,
403            prefill: Vec::new(),
404            decode: vec!["http://127.0.0.1:8020".to_owned()],
405        });
406        let error = match result {
407            Ok(_) => bail!("empty prefill targets should fail"),
408            Err(error) => error,
409        };
410        assert!(error.to_string().contains("at least one prefill endpoint"));
411        Ok(())
412    }
413}