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 serde_json::json;
304
305 #[test]
306 fn meta_exports_byte_stable_proxy_identity() {
307 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}