1use crate::core::{
5 self, ProxyHealthcheckResponse, ProxyHttpError, forward_response, join_path,
6 outbound_authorization,
7};
8use crate::error::ProxyError;
9use axum::Json;
10use axum::Router;
11use axum::body::Body;
12use axum::extract::State;
13use axum::http::{HeaderMap, Response, StatusCode, header};
14use axum::routing::{get, post};
15use serde_json::{Map, Value};
16use std::sync::Arc;
17use std::sync::atomic::{AtomicBool, AtomicU64, AtomicUsize, Ordering};
18use std::time::{SystemTime, UNIX_EPOCH};
19
20pub const VERSION: u32 = 2;
21
22pub const HEALTHCHECK_PATH: &str = "/healthcheck";
23
24pub const COMPLETIONS_PATH: &str = "/v1/completions";
25pub const CHAT_COMPLETIONS_PATH: &str = "/v1/chat/completions";
26
27const PROXY_NAME: &str = "TensorRT-LLM proxy";
29
30const MIN_REQUEST_ID: u64 = 1_u64 << 42;
31const CONTEXT_FIRST_SCHEDULE_STYLE: u64 = 0;
32const TERMINAL_SSE: &[u8] = b"data: [DONE]\n\n";
33
34#[derive(Clone, Debug)]
35pub struct Config {
36 pub host: String,
37 pub port: u16,
38 pub prefill: Vec<String>,
39 pub decode: Vec<String>,
40}
41
42pub fn run(config: Config) -> Result<(), ProxyError> {
43 core::run(|| run_async(config))
44}
45
46pub async fn run_async(config: Config) -> Result<(), ProxyError> {
47 let host = config.host.clone();
48 let port = config.port;
49 let state = ProxyState::new(config)?;
50 tokio::spawn(await_backends(state.clone()));
51 core::serve_router(PROXY_NAME, &host, port, router(state)).await
52}
53
54fn router(state: ProxyState) -> Router {
55 Router::new()
56 .route(HEALTHCHECK_PATH, get(healthcheck))
57 .route(COMPLETIONS_PATH, post(completions))
58 .route(CHAT_COMPLETIONS_PATH, post(chat_completions))
59 .with_state(state)
60}
61
62#[derive(Clone, Copy)]
63enum RequestFamily {
64 Completions,
65 ChatCompletions,
66}
67
68impl RequestFamily {
69 fn path(self) -> &'static str {
70 match self {
71 Self::Completions => COMPLETIONS_PATH,
72 Self::ChatCompletions => CHAT_COMPLETIONS_PATH,
73 }
74 }
75}
76
77#[derive(Clone)]
78struct ProxyState {
79 inner: Arc<ProxyStateInner>,
80}
81
82struct ProxyStateInner {
83 client: reqwest::Client,
84 prefill: Vec<String>,
85 decode: Vec<String>,
86 ready: AtomicBool,
87 prefill_cursor: AtomicUsize,
88 decode_cursor: AtomicUsize,
89 request_counter: AtomicU64,
90}
91
92impl ProxyState {
93 fn new(config: Config) -> Result<Self, ProxyError> {
94 core::require_endpoints(
95 PROXY_NAME,
96 config.prefill.is_empty(),
97 config.decode.is_empty(),
98 )?;
99 Ok(Self {
100 inner: Arc::new(ProxyStateInner {
101 client: core::pooled_client(PROXY_NAME)?,
102 prefill: config.prefill,
103 decode: config.decode,
104 ready: AtomicBool::new(false),
105 prefill_cursor: AtomicUsize::new(0),
106 decode_cursor: AtomicUsize::new(0),
107 request_counter: AtomicU64::new(request_id_seed()),
108 }),
109 })
110 }
111
112 fn client(&self) -> reqwest::Client {
113 self.inner.client.clone()
114 }
115
116 fn ready(&self) -> bool {
117 self.inner.ready.load(Ordering::SeqCst)
118 }
119
120 fn set_ready(&self) {
121 self.inner.ready.store(true, Ordering::SeqCst);
122 }
123
124 fn next_prefill(&self) -> String {
125 let index = core::round_robin_index(&self.inner.prefill_cursor, self.inner.prefill.len());
126 self.inner.prefill[index].clone()
127 }
128
129 fn next_decode(&self) -> String {
130 let index = core::round_robin_index(&self.inner.decode_cursor, self.inner.decode.len());
131 self.inner.decode[index].clone()
132 }
133
134 fn next_request_id(&self) -> u64 {
135 self.inner.request_counter.fetch_add(1, Ordering::SeqCst)
136 }
137}
138
139fn request_id_seed() -> u64 {
140 const SEED_CEILING: u64 = 1_u64 << 61;
141 let nanos = SystemTime::now()
142 .duration_since(UNIX_EPOCH)
143 .map_or(0, |elapsed| elapsed.as_nanos() as u64);
144 let entropy = nanos ^ (u64::from(std::process::id()) << 32);
145 MIN_REQUEST_ID + entropy % (SEED_CEILING - MIN_REQUEST_ID)
146}
147
148async fn await_backends(state: ProxyState) {
149 let urls = core::fanout_target_urls(
150 state.inner.prefill.iter().map(String::as_str),
151 state.inner.decode.iter().map(String::as_str),
152 );
153 core::await_backends(state.client(), urls, "/health").await;
154 state.set_ready();
155}
156
157async fn healthcheck(
158 State(state): State<ProxyState>,
159) -> (StatusCode, Json<ProxyHealthcheckResponse>) {
160 core::healthcheck_response(
161 state.ready(),
162 state.inner.prefill.len(),
163 state.inner.decode.len(),
164 )
165}
166
167async fn completions(
168 State(state): State<ProxyState>,
169 headers: HeaderMap,
170 Json(body): Json<Value>,
171) -> Result<Response<Body>, ProxyHttpError> {
172 request_route(state, headers, body, RequestFamily::Completions).await
173}
174
175async fn chat_completions(
176 State(state): State<ProxyState>,
177 headers: HeaderMap,
178 Json(body): Json<Value>,
179) -> Result<Response<Body>, ProxyHttpError> {
180 request_route(state, headers, body, RequestFamily::ChatCompletions).await
181}
182
183async fn request_route(
184 state: ProxyState,
185 headers: HeaderMap,
186 body: Value,
187 family: RequestFamily,
188) -> Result<Response<Body>, ProxyHttpError> {
189 let stream = validate_public_request(&body, family)?;
190 if !state.ready() {
191 return Err(ProxyHttpError::status(
192 StatusCode::SERVICE_UNAVAILABLE,
193 "proxy is not ready",
194 ));
195 }
196
197 let prefill = state.next_prefill();
198 let decode = state.next_decode();
199 let request_id = state.next_request_id();
200 let request_id_header = request_id.to_string();
201 let authorization = outbound_authorization(&headers);
202 let context_body = context_body(&body, request_id)?;
203 let context = send_context_request(
204 state.client(),
205 prefill,
206 context_body,
207 family.path(),
208 &request_id_header,
209 authorization.as_deref(),
210 )
211 .await?;
212
213 match context_outcome(context, request_id, family)? {
214 ContextOutcome::Complete(context) => complete_context_response(context, stream, family),
215 ContextOutcome::Handoff(handoff) => {
216 let generation_body = generation_body(&body, handoff, family)?;
217 let response = core::send_json_post(
218 state.client(),
219 join_path(&decode, family.path()),
220 &generation_body,
221 Some(&request_id_header),
222 authorization.as_deref(),
223 &[],
224 "decode request",
225 )
226 .await?;
227 if stream {
228 core::stream_response(response)
229 } else {
230 forward_response(response).await
231 }
232 }
233 }
234}
235
236fn validate_public_request(body: &Value, family: RequestFamily) -> Result<bool, ProxyHttpError> {
237 let object = body.as_object().ok_or_else(|| {
238 ProxyHttpError::status(
239 StatusCode::BAD_REQUEST,
240 "OpenAI request body must be a JSON object",
241 )
242 })?;
243 match family {
244 RequestFamily::Completions => match object.get("prompt") {
245 Some(Value::String(_)) => {}
246 Some(Value::Array(_)) => {
247 return Err(ProxyHttpError::status(
248 StatusCode::BAD_REQUEST,
249 "TensorRT-LLM built-in proxy does not support prompt arrays",
250 ));
251 }
252 _ => {
253 return Err(ProxyHttpError::status(
254 StatusCode::BAD_REQUEST,
255 "TensorRT-LLM built-in proxy requires a scalar string prompt",
256 ));
257 }
258 },
259 RequestFamily::ChatCompletions => {
260 if !object.get("messages").is_some_and(Value::is_array) {
261 return Err(ProxyHttpError::status(
262 StatusCode::BAD_REQUEST,
263 "TensorRT-LLM built-in proxy requires structured chat messages",
264 ));
265 }
266 }
267 }
268 if object
269 .get("n")
270 .is_some_and(|count| count.as_u64() != Some(1))
271 {
272 return Err(ProxyHttpError::status(
273 StatusCode::BAD_REQUEST,
274 "TensorRT-LLM built-in proxy supports only n=1",
275 ));
276 }
277 Ok(object
278 .get("stream")
279 .and_then(Value::as_bool)
280 .unwrap_or(false))
281}
282
283fn context_body(body: &Value, request_id: u64) -> Result<Value, ProxyHttpError> {
284 let mut body = body.clone();
285 let object = body.as_object_mut().ok_or_else(|| {
286 ProxyHttpError::status(
287 StatusCode::BAD_REQUEST,
288 "OpenAI completion request body must be a JSON object",
289 )
290 })?;
291 object.insert("stream".to_owned(), Value::Bool(false));
292 object.remove("stream_options");
293 object.insert(
294 "disaggregated_params".to_owned(),
295 Value::Object(Map::from_iter([
296 (
297 "request_type".to_owned(),
298 Value::String("context_only".to_owned()),
299 ),
300 ("disagg_request_id".to_owned(), Value::from(request_id)),
301 (
302 "schedule_style".to_owned(),
303 Value::from(CONTEXT_FIRST_SCHEDULE_STYLE),
304 ),
305 ])),
306 );
307 Ok(body)
308}
309
310struct ContextResponse {
311 status: StatusCode,
312 content_type: Option<String>,
313 body: Value,
314}
315
316async fn send_context_request(
317 client: reqwest::Client,
318 prefill: String,
319 body: Value,
320 path: &'static str,
321 request_id: &str,
322 authorization: Option<&str>,
323) -> Result<ContextResponse, ProxyHttpError> {
324 let response = core::send_json_post(
325 client,
326 join_path(&prefill, path),
327 &body,
328 Some(request_id),
329 authorization,
330 &[],
331 "context request",
332 )
333 .await?;
334 let status = core::status_code(response.status())?;
335 let content_type = response
336 .headers()
337 .get(reqwest::header::CONTENT_TYPE)
338 .and_then(|value| value.to_str().ok())
339 .map(str::to_owned);
340 let bytes = response
341 .bytes()
342 .await
343 .map_err(|error| ProxyHttpError::upstream("context response body read failed", error))?;
344 let body = serde_json::from_slice(&bytes).map_err(|error| {
345 ProxyHttpError::status(
346 StatusCode::BAD_GATEWAY,
347 format!("context response was not valid JSON: {error}"),
348 )
349 })?;
350 Ok(ContextResponse {
351 status,
352 content_type,
353 body,
354 })
355}
356
357enum ContextOutcome {
358 Complete(ContextResponse),
359 Handoff(Handoff),
360}
361
362struct Handoff {
363 prompt_token_ids: PromptTokenIds,
364 usage: Value,
365 disaggregated_params: Map<String, Value>,
366}
367
368enum PromptTokenIds {
369 Array(Value),
370 Base64(String),
371}
372
373fn context_outcome(
374 mut response: ContextResponse,
375 request_id: u64,
376 family: RequestFamily,
377) -> Result<ContextOutcome, ProxyHttpError> {
378 let first = response
379 .body
380 .get("choices")
381 .and_then(Value::as_array)
382 .and_then(|choices| choices.first())
383 .and_then(Value::as_object)
384 .ok_or_else(|| {
385 ProxyHttpError::status(
386 StatusCode::BAD_GATEWAY,
387 "context response did not include a first choice",
388 )
389 })?;
390 let needs_generation = first
391 .get("finish_reason")
392 .and_then(Value::as_str)
393 .is_some_and(|reason| matches!(reason, "length" | "not_finished"));
394 if !needs_generation {
395 sanitize_context_response(&mut response.body);
396 return Ok(ContextOutcome::Complete(response));
397 }
398
399 let prompt_token_ids = match family {
400 RequestFamily::Completions => response
401 .body
402 .get("prompt_token_ids")
403 .filter(|tokens| is_scalar_token_array(tokens))
404 .cloned()
405 .map(PromptTokenIds::Array)
406 .ok_or_else(|| handoff_error("prompt_token_ids must be a scalar token array"))?,
407 RequestFamily::ChatCompletions => {
408 if let Some(tokens) = response
409 .body
410 .get("prompt_token_ids_b64")
411 .and_then(Value::as_str)
412 {
413 PromptTokenIds::Base64(tokens.to_owned())
414 } else {
415 response
416 .body
417 .get("prompt_token_ids")
418 .filter(|tokens| is_scalar_token_array(tokens))
419 .cloned()
420 .map(PromptTokenIds::Array)
421 .ok_or_else(|| {
422 handoff_error(
423 "chat handoff requires prompt_token_ids_b64 or a scalar prompt_token_ids array",
424 )
425 })?
426 }
427 }
428 };
429 let usage = response
430 .body
431 .get("usage")
432 .filter(|usage| usage.is_object())
433 .cloned()
434 .ok_or_else(|| handoff_error("usage is missing"))?;
435 let params = first
436 .get("disaggregated_params")
437 .and_then(Value::as_object)
438 .cloned()
439 .ok_or_else(|| handoff_error("disaggregated_params is missing"))?;
440 if params.get("ctx_request_id").is_none_or(Value::is_null) {
441 return Err(handoff_error("ctx_request_id is null"));
442 }
443 if params.get("disagg_request_id").and_then(Value::as_u64) != Some(request_id) {
444 return Err(handoff_error(
445 "disagg_request_id does not match the assigned request",
446 ));
447 }
448 if params.get("first_gen_tokens").is_none_or(Value::is_null) {
449 return Err(handoff_error("first_gen_tokens is missing"));
450 }
451 Ok(ContextOutcome::Handoff(Handoff {
452 prompt_token_ids,
453 usage,
454 disaggregated_params: params,
455 }))
456}
457
458fn is_scalar_token_array(value: &Value) -> bool {
459 value.as_array().is_some_and(|tokens| {
460 tokens.iter().all(|token| {
461 token
462 .as_number()
463 .is_some_and(|number| number.is_i64() || number.is_u64())
464 })
465 })
466}
467
468fn handoff_error(detail: &str) -> ProxyHttpError {
469 ProxyHttpError::status(
470 StatusCode::BAD_GATEWAY,
471 format!("invalid TensorRT-LLM context handoff: {detail}"),
472 )
473}
474
475fn sanitize_context_response(body: &mut Value) {
476 if let Some(choices) = body.get_mut("choices").and_then(Value::as_array_mut) {
477 for choice in choices {
478 if let Some(choice) = choice.as_object_mut() {
479 choice.remove("disaggregated_params");
480 }
481 }
482 }
483}
484
485fn complete_context_response(
486 context: ContextResponse,
487 stream: bool,
488 family: RequestFamily,
489) -> Result<Response<Body>, ProxyHttpError> {
490 if stream {
491 let event = context_stream_event(&context.body, family)?;
492 let mut body = b"data: ".to_vec();
493 body.extend(serde_json::to_vec(&event).map_err(|error| {
494 ProxyHttpError::internal(format!("failed to serialize context stream event: {error}"))
495 })?);
496 body.extend_from_slice(b"\n\n");
497 body.extend_from_slice(TERMINAL_SSE);
498 return Response::builder()
499 .status(context.status)
500 .header(header::CONTENT_TYPE, "text/event-stream")
501 .body(Body::from(body))
502 .map_err(|error| {
503 ProxyHttpError::internal(format!(
504 "failed to build terminal context response: {error}"
505 ))
506 });
507 }
508 let body = serde_json::to_vec(&context.body).map_err(|error| {
509 ProxyHttpError::internal(format!("failed to serialize context response: {error}"))
510 })?;
511 let mut builder = Response::builder().status(context.status);
512 if let Some(content_type) = context.content_type {
513 builder = builder.header(header::CONTENT_TYPE, content_type);
514 }
515 builder.body(Body::from(body)).map_err(|error| {
516 ProxyHttpError::internal(format!("failed to build context response: {error}"))
517 })
518}
519
520fn context_stream_event(body: &Value, family: RequestFamily) -> Result<Value, ProxyHttpError> {
521 let mut event = body.clone();
522 let object = event.as_object_mut().ok_or_else(|| {
523 ProxyHttpError::status(
524 StatusCode::BAD_GATEWAY,
525 "context response body must be a JSON object",
526 )
527 })?;
528 match family {
529 RequestFamily::Completions => {
530 object.insert(
531 "object".to_owned(),
532 Value::String("text_completion".to_owned()),
533 );
534 }
535 RequestFamily::ChatCompletions => {
536 object.insert(
537 "object".to_owned(),
538 Value::String("chat.completion.chunk".to_owned()),
539 );
540 if let Some(choices) = object.get_mut("choices").and_then(Value::as_array_mut) {
541 for choice in choices {
542 if let Some(choice) = choice.as_object_mut()
543 && let Some(message) = choice.remove("message")
544 {
545 choice.insert("delta".to_owned(), message);
546 }
547 }
548 }
549 }
550 }
551 Ok(event)
552}
553
554fn generation_body(
555 body: &Value,
556 handoff: Handoff,
557 family: RequestFamily,
558) -> Result<Value, ProxyHttpError> {
559 let mut body = body.clone();
560 let object = body.as_object_mut().ok_or_else(|| {
561 ProxyHttpError::status(
562 StatusCode::BAD_REQUEST,
563 "OpenAI request body must be a JSON object",
564 )
565 })?;
566 let mut params = handoff.disaggregated_params;
567 params.insert(
568 "request_type".to_owned(),
569 Value::String("generation_only".to_owned()),
570 );
571 params.insert(
572 "schedule_style".to_owned(),
573 Value::from(CONTEXT_FIRST_SCHEDULE_STYLE),
574 );
575 params.insert("ctx_usage".to_owned(), handoff.usage);
576 match (family, handoff.prompt_token_ids) {
577 (RequestFamily::Completions, PromptTokenIds::Array(tokens)) => {
578 object.insert("prompt".to_owned(), tokens);
579 }
580 (RequestFamily::ChatCompletions, PromptTokenIds::Base64(tokens)) => {
581 object.remove("prompt_token_ids");
582 object.insert("prompt_token_ids_b64".to_owned(), Value::String(tokens));
583 }
584 (RequestFamily::ChatCompletions, PromptTokenIds::Array(tokens)) => {
585 object.remove("prompt_token_ids_b64");
586 object.insert("prompt_token_ids".to_owned(), tokens);
587 }
588 (RequestFamily::Completions, PromptTokenIds::Base64(_)) => {
589 return Err(handoff_error(
590 "completion handoff cannot use prompt_token_ids_b64",
591 ));
592 }
593 }
594 object.insert("disaggregated_params".to_owned(), Value::Object(params));
595 Ok(body)
596}
597
598#[cfg(test)]
599mod tests {
600 use super::*;
601 use anyhow::{Context, Result, bail};
602 use async_stream::stream;
603 use axum::body::{Body, to_bytes};
604 use axum::http::{HeaderValue, header};
605 use axum::response::IntoResponse;
606 use axum::routing::{get, post};
607 use axum::serve;
608 use bytes::Bytes;
609 use futures_util::StreamExt;
610 use serde_json::json;
611 use std::sync::atomic::AtomicUsize;
612 use std::time::Duration;
613 use tokio::net::TcpListener;
614 use tokio::sync::{Mutex, Notify};
615 use tokio::task::JoinHandle;
616
617 #[test]
618 fn context_request_is_non_streaming_context_first_with_large_integer_id() -> Result<()> {
619 let state = proxy_state(
620 vec!["http://prefill".to_owned()],
621 vec!["http://decode".to_owned()],
622 )?;
623 let first = state.next_request_id();
624 let second = state.next_request_id();
625 assert!(first >= MIN_REQUEST_ID);
626 assert_eq!(second, first + 1);
627
628 let lowered = context_body(
629 &json!({
630 "model": "m",
631 "prompt": "hello",
632 "stream": true,
633 "stream_options": {"include_usage": true},
634 "opaque": "preserved"
635 }),
636 first,
637 )
638 .map_err(|error| anyhow::anyhow!(error.to_string()))?;
639 assert_eq!(lowered["stream"], Value::Bool(false));
640 assert!(lowered.get("stream_options").is_none());
641 assert_eq!(lowered["opaque"], Value::String("preserved".to_owned()));
642 assert_eq!(
643 lowered["disaggregated_params"]["request_type"],
644 "context_only"
645 );
646 assert_eq!(
647 lowered["disaggregated_params"]["schedule_style"],
648 Value::from(0)
649 );
650 assert_eq!(lowered["disaggregated_params"]["disagg_request_id"], first);
651 Ok(())
652 }
653
654 #[test]
655 fn handoff_preserves_opaque_params_and_replaces_only_owned_fields() -> Result<()> {
656 let request_id = MIN_REQUEST_ID + 7;
657 let context = context_response(json!({
658 "choices": [{
659 "finish_reason": "not_finished",
660 "disaggregated_params": {
661 "request_type": "context_only",
662 "schedule_style": 1,
663 "ctx_usage": {"stale": true},
664 "ctx_request_id": 91,
665 "disagg_request_id": request_id,
666 "first_gen_tokens": [8],
667 "opaque_future_field": {"endpoint": "nixl://ctx"}
668 }
669 }],
670 "prompt_token_ids": [10, 11, 12],
671 "usage": {"prompt_tokens": 3, "completion_tokens": 1}
672 }));
673 let handoff = match context_outcome(context, request_id, RequestFamily::Completions)
674 .map_err(|error| anyhow::anyhow!(error.to_string()))?
675 {
676 ContextOutcome::Handoff(handoff) => handoff,
677 ContextOutcome::Complete(_) => bail!("not_finished must require generation"),
678 };
679 let generated = generation_body(
680 &json!({
681 "model": "m",
682 "prompt": "hello",
683 "stream": true,
684 "temperature": 0.25,
685 "opaque_request_field": [1, 2]
686 }),
687 handoff,
688 RequestFamily::Completions,
689 )
690 .map_err(|error| anyhow::anyhow!(error.to_string()))?;
691
692 assert_eq!(generated["prompt"], json!([10, 11, 12]));
693 assert_eq!(generated["stream"], Value::Bool(true));
694 assert_eq!(generated["temperature"], json!(0.25));
695 assert_eq!(generated["opaque_request_field"], json!([1, 2]));
696 let params = &generated["disaggregated_params"];
697 assert_eq!(params["request_type"], "generation_only");
698 assert_eq!(params["schedule_style"], CONTEXT_FIRST_SCHEDULE_STYLE);
699 assert_eq!(
700 params["ctx_usage"],
701 json!({"prompt_tokens": 3, "completion_tokens": 1})
702 );
703 assert_eq!(params["ctx_request_id"], 91);
704 assert_eq!(params["disagg_request_id"], request_id);
705 assert_eq!(params["first_gen_tokens"], json!([8]));
706 assert_eq!(
707 params["opaque_future_field"],
708 json!({"endpoint": "nixl://ctx"})
709 );
710 Ok(())
711 }
712
713 #[test]
714 fn malformed_required_handoff_metadata_is_rejected() -> Result<()> {
715 let request_id = MIN_REQUEST_ID + 9;
716 let cases = [
717 (
718 "missing choices",
719 json!({"choices": [], "prompt_token_ids": [1], "usage": {}}),
720 ),
721 (
722 "nested prompt tokens",
723 handoff_response(
724 request_id,
725 json!([[1, 2]]),
726 json!({}),
727 valid_params(request_id),
728 ),
729 ),
730 (
731 "missing usage",
732 handoff_response(
733 request_id,
734 json!([1, 2]),
735 Value::Null,
736 valid_params(request_id),
737 ),
738 ),
739 (
740 "missing disaggregated params",
741 handoff_response(request_id, json!([1, 2]), json!({}), Value::Null),
742 ),
743 (
744 "null context id",
745 handoff_response(
746 request_id,
747 json!([1, 2]),
748 json!({}),
749 json!({
750 "ctx_request_id": null,
751 "disagg_request_id": request_id,
752 "first_gen_tokens": [3]
753 }),
754 ),
755 ),
756 (
757 "mismatched request id",
758 handoff_response(
759 request_id,
760 json!([1, 2]),
761 json!({}),
762 json!({
763 "ctx_request_id": 1,
764 "disagg_request_id": request_id + 1,
765 "first_gen_tokens": [3]
766 }),
767 ),
768 ),
769 (
770 "missing first token",
771 handoff_response(
772 request_id,
773 json!([1, 2]),
774 json!({}),
775 json!({"ctx_request_id": 1, "disagg_request_id": request_id}),
776 ),
777 ),
778 ];
779 for (label, body) in cases {
780 let result = context_outcome(
781 context_response(body),
782 request_id,
783 RequestFamily::Completions,
784 );
785 assert!(result.is_err(), "{label} was accepted");
786 }
787 Ok(())
788 }
789
790 #[test]
791 fn prefill_and_decode_round_robin_are_independent() -> Result<()> {
792 let state = proxy_state(
793 vec!["p0".to_owned(), "p1".to_owned()],
794 vec!["d0".to_owned(), "d1".to_owned(), "d2".to_owned()],
795 )?;
796 assert_eq!(state.next_prefill(), "p0");
797 assert_eq!(state.next_decode(), "d0");
798 assert_eq!(state.next_decode(), "d1");
799 assert_eq!(state.next_prefill(), "p1");
800 assert_eq!(state.next_decode(), "d2");
801 assert_eq!(state.next_prefill(), "p0");
802 Ok(())
803 }
804
805 #[tokio::test]
806 async fn invalid_public_shapes_are_rejected_before_dispatch() -> Result<()> {
807 let context_backend = ContextBackend::default();
808 let decode_backend = DecodeBackend::default();
809 let (prefill, prefill_server) = spawn_context_backend(context_backend.clone()).await?;
810 let (decode, decode_server) = spawn_decode_backend(decode_backend.clone()).await?;
811 let state = proxy_state(vec![prefill], vec![decode])?;
812 state.set_ready();
813
814 for request in [
815 json!({"model": "m", "prompt": ["hello"]}),
816 json!({"model": "m", "prompt": "hello", "n": 2}),
817 ] {
818 let error = match request_route(
819 state.clone(),
820 HeaderMap::new(),
821 request,
822 RequestFamily::Completions,
823 )
824 .await
825 {
826 Ok(_) => bail!("invalid public request was dispatched"),
827 Err(error) => error,
828 };
829 assert_eq!(error.into_response().status(), StatusCode::BAD_REQUEST);
830 }
831 assert!(context_backend.requests.lock().await.is_empty());
832 assert!(decode_backend.requests.lock().await.is_empty());
833 prefill_server.abort();
834 decode_server.abort();
835 Ok(())
836 }
837
838 #[tokio::test]
839 async fn context_completion_skips_decode_and_returns_public_shape() -> Result<()> {
840 let context_backend = ContextBackend::default();
841 let decode_backend = DecodeBackend::default();
842 let (prefill, prefill_server) = spawn_context_backend(context_backend).await?;
843 let (decode, decode_server) = spawn_decode_backend(decode_backend.clone()).await?;
844 let state = proxy_state(vec![prefill], vec![decode])?;
845 state.set_ready();
846
847 let response = request_route(
848 state.clone(),
849 HeaderMap::new(),
850 json!({"model": "m", "prompt": "hello", "mode": "complete"}),
851 RequestFamily::Completions,
852 )
853 .await
854 .map_err(|error| anyhow::anyhow!(error.to_string()))?;
855 assert_eq!(response.status(), StatusCode::CREATED);
856 let returned: Value =
857 serde_json::from_slice(&to_bytes(response.into_body(), usize::MAX).await?)?;
858 assert_eq!(returned["opaque"], "kept");
859 assert!(returned["choices"][0].get("disaggregated_params").is_none());
860
861 let response = request_route(
862 state,
863 HeaderMap::new(),
864 json!({
865 "model": "m",
866 "prompt": "hello",
867 "mode": "complete",
868 "stream": true
869 }),
870 RequestFamily::Completions,
871 )
872 .await
873 .map_err(|error| anyhow::anyhow!(error.to_string()))?;
874 assert_eq!(
875 response.headers().get(header::CONTENT_TYPE),
876 Some(&HeaderValue::from_static("text/event-stream"))
877 );
878 let bytes = to_bytes(response.into_body(), usize::MAX).await?;
879 let stream = std::str::from_utf8(&bytes)?;
880 let event = first_sse_event(stream)?;
881 assert_eq!(event["object"], "text_completion");
882 assert_eq!(event["choices"][0]["index"], 0);
883 assert_eq!(event["choices"][0]["text"], "answer");
884 assert_eq!(event["choices"][0]["finish_reason"], "stop");
885 assert!(stream.ends_with("data: [DONE]\n\n"));
886 assert!(decode_backend.requests.lock().await.is_empty());
887 prefill_server.abort();
888 decode_server.abort();
889 Ok(())
890 }
891
892 #[tokio::test]
893 async fn generation_handoff_reuses_id_auth_and_forwards_both_response_modes() -> Result<()> {
894 let context_backend = ContextBackend::default();
895 let decode_backend = DecodeBackend::default();
896 let stream_gate = decode_backend.stream_gate.clone();
897 let (prefill, prefill_server) = spawn_context_backend(context_backend.clone()).await?;
898 let (decode, decode_server) = spawn_decode_backend(decode_backend.clone()).await?;
899 let state = proxy_state(vec![prefill], vec![decode])?;
900 state.set_ready();
901 let mut headers = HeaderMap::new();
902 headers.insert(header::AUTHORIZATION, "Bearer inbound".parse()?);
903
904 let response = request_route(
905 state.clone(),
906 headers.clone(),
907 json!({"model": "m", "prompt": "hello", "mode": "generate"}),
908 RequestFamily::Completions,
909 )
910 .await
911 .map_err(|error| anyhow::anyhow!(error.to_string()))?;
912 assert_eq!(response.status(), StatusCode::CREATED);
913 assert_eq!(
914 response.headers().get(header::CONTENT_TYPE),
915 Some(&HeaderValue::from_static("application/x-inferlab-test"))
916 );
917 assert_eq!(
918 to_bytes(response.into_body(), usize::MAX).await?,
919 Bytes::from_static(b"decode-complete")
920 );
921
922 let response = request_route(
923 state,
924 headers,
925 json!({
926 "model": "m",
927 "prompt": "hello",
928 "mode": "generate",
929 "stream": true
930 }),
931 RequestFamily::Completions,
932 )
933 .await
934 .map_err(|error| anyhow::anyhow!(error.to_string()))?;
935 assert_eq!(response.status(), StatusCode::ACCEPTED);
936 assert_eq!(
937 response.headers().get(header::CONTENT_TYPE),
938 Some(&HeaderValue::from_static("text/event-stream"))
939 );
940 let mut stream = response.into_body().into_data_stream();
941 let first = tokio::time::timeout(Duration::from_secs(1), stream.next())
942 .await?
943 .context("decode stream ended before the first event")??;
944 assert_eq!(first, Bytes::from_static(b"data: first\n\n"));
945 stream_gate.notify_one();
946 let second = stream
947 .next()
948 .await
949 .context("decode stream ended before the terminal event")??;
950 assert_eq!(second, Bytes::from_static(TERMINAL_SSE));
951
952 let context_requests = context_backend.requests.lock().await;
953 let decode_requests = decode_backend.requests.lock().await;
954 assert_eq!(context_requests.len(), 2);
955 assert_eq!(decode_requests.len(), 2);
956 for (context, decode) in context_requests.iter().zip(decode_requests.iter()) {
957 let assigned = context.body["disaggregated_params"]["disagg_request_id"]
958 .as_u64()
959 .context("context request lacked an integer disagg_request_id")?;
960 assert!(assigned >= MIN_REQUEST_ID);
961 assert_eq!(
962 decode.body["disaggregated_params"]["disagg_request_id"],
963 assigned
964 );
965 assert_eq!(decode.body["prompt"], json!([10, 11, 12]));
966 assert_eq!(
967 context.headers.get(header::AUTHORIZATION),
968 Some(&HeaderValue::from_static("Bearer inbound"))
969 );
970 assert_eq!(
971 decode.headers.get(header::AUTHORIZATION),
972 Some(&HeaderValue::from_static("Bearer inbound"))
973 );
974 let assigned_header = assigned.to_string();
975 let context_header = context
976 .headers
977 .get("x-request-id")
978 .and_then(|value| value.to_str().ok());
979 let decode_header = decode
980 .headers
981 .get("x-request-id")
982 .and_then(|value| value.to_str().ok());
983 assert_eq!(context_header, Some(assigned_header.as_str()));
984 assert_eq!(decode_header, context_header);
985 }
986 drop(context_requests);
987 drop(decode_requests);
988 prefill_server.abort();
989 decode_server.abort();
990 Ok(())
991 }
992
993 #[tokio::test]
994 async fn chat_uses_chat_handoff_and_emits_route_specific_context_stream() -> Result<()> {
995 let context_backend = ContextBackend::default();
996 let decode_backend = DecodeBackend::default();
997 let (prefill, prefill_server) = spawn_context_backend(context_backend.clone()).await?;
998 let (decode, decode_server) = spawn_decode_backend(decode_backend.clone()).await?;
999 let state = proxy_state(vec![prefill], vec![decode])?;
1000 state.set_ready();
1001 let messages = json!([{"role": "user", "content": "hello"}]);
1002
1003 let response = request_route(
1004 state.clone(),
1005 HeaderMap::new(),
1006 json!({
1007 "model": "m",
1008 "messages": messages,
1009 "mode": "generate",
1010 "temperature": 1.0,
1011 "reasoning_effort": "high",
1012 "chat_template_kwargs": {"enable_thinking": true}
1013 }),
1014 RequestFamily::ChatCompletions,
1015 )
1016 .await
1017 .map_err(|error| anyhow::anyhow!(error.to_string()))?;
1018 assert_eq!(response.status(), StatusCode::CREATED);
1019
1020 let context_requests = context_backend.requests.lock().await;
1021 let decode_requests = decode_backend.requests.lock().await;
1022 assert_eq!(context_requests.len(), 1);
1023 assert_eq!(decode_requests.len(), 1);
1024 assert_eq!(context_requests[0].path, CHAT_COMPLETIONS_PATH);
1025 assert_eq!(decode_requests[0].path, CHAT_COMPLETIONS_PATH);
1026 assert_eq!(context_requests[0].body["messages"], messages);
1027 assert_eq!(decode_requests[0].body["messages"], messages);
1028 assert_eq!(decode_requests[0].body["prompt_token_ids_b64"], "encoded");
1029 assert!(decode_requests[0].body.get("prompt").is_none());
1030 for key in ["temperature", "reasoning_effort", "chat_template_kwargs"] {
1031 assert_eq!(decode_requests[0].body[key], context_requests[0].body[key]);
1032 }
1033 drop(context_requests);
1034 drop(decode_requests);
1035
1036 let response = request_route(
1037 state,
1038 HeaderMap::new(),
1039 json!({
1040 "model": "m",
1041 "messages": messages,
1042 "mode": "complete",
1043 "stream": true
1044 }),
1045 RequestFamily::ChatCompletions,
1046 )
1047 .await
1048 .map_err(|error| anyhow::anyhow!(error.to_string()))?;
1049 let bytes = to_bytes(response.into_body(), usize::MAX).await?;
1050 let stream = std::str::from_utf8(&bytes)?;
1051 let event = first_sse_event(stream)?;
1052 assert_eq!(event["object"], "chat.completion.chunk");
1053 assert_eq!(event["choices"][0]["index"], 0);
1054 assert_eq!(event["choices"][0]["delta"]["content"], "answer");
1055 assert_eq!(event["choices"][0]["finish_reason"], "stop");
1056 assert!(stream.ends_with("data: [DONE]\n\n"));
1057 assert_eq!(decode_backend.requests.lock().await.len(), 1);
1058 prefill_server.abort();
1059 decode_server.abort();
1060 Ok(())
1061 }
1062
1063 #[tokio::test]
1064 async fn upstream_failures_remain_failures_before_and_after_headers() -> Result<()> {
1065 let context_backend = ContextBackend::default();
1066 let decode_backend = DecodeBackend::default();
1067 let stream_gate = decode_backend.stream_gate.clone();
1068 let (prefill, prefill_server) = spawn_context_backend(context_backend).await?;
1069 let (decode, decode_server) = spawn_decode_backend(decode_backend).await?;
1070 let state = proxy_state(vec![prefill], vec![decode])?;
1071 state.set_ready();
1072
1073 for mode in ["context-fail", "decode-fail"] {
1074 let result = request_route(
1075 state.clone(),
1076 HeaderMap::new(),
1077 json!({"model": "m", "prompt": "hello", "mode": mode}),
1078 RequestFamily::Completions,
1079 )
1080 .await;
1081 let error = match result {
1082 Ok(_) => bail!("{mode} returned a successful public response"),
1083 Err(error) => error,
1084 };
1085 assert_eq!(error.into_response().status(), StatusCode::BAD_GATEWAY);
1086 }
1087
1088 let response = request_route(
1089 state,
1090 HeaderMap::new(),
1091 json!({
1092 "model": "m",
1093 "prompt": "hello",
1094 "mode": "stream-error",
1095 "stream": true
1096 }),
1097 RequestFamily::Completions,
1098 )
1099 .await
1100 .map_err(|error| anyhow::anyhow!(error.to_string()))?;
1101 assert_eq!(response.status(), StatusCode::ACCEPTED);
1102 let mut stream = response.into_body().into_data_stream();
1103 assert!(matches!(stream.next().await, Some(Ok(_))));
1104 stream_gate.notify_one();
1105 let result = stream
1106 .next()
1107 .await
1108 .context("decode stream ended cleanly after an upstream body failure")?;
1109 let error = match result {
1110 Ok(_) => bail!("decode body failure was returned as successful bytes"),
1111 Err(error) => error,
1112 };
1113 assert!(error.to_string().contains("decode stream failed"));
1114 prefill_server.abort();
1115 decode_server.abort();
1116 Ok(())
1117 }
1118
1119 #[tokio::test]
1120 async fn healthcheck_waits_for_every_configured_worker() -> Result<()> {
1121 let context_backend = ContextBackend::default();
1122 let decode_backend = DecodeBackend::default();
1123 let (prefill, prefill_server) = spawn_context_backend(context_backend.clone()).await?;
1124 let (decode, decode_server) = spawn_decode_backend(decode_backend.clone()).await?;
1125 let state = proxy_state(vec![prefill], vec![decode])?;
1126 let (status, Json(body)) = healthcheck(State(state.clone())).await;
1127 assert_eq!(status, StatusCode::SERVICE_UNAVAILABLE);
1128 assert!(!body.ready);
1129
1130 tokio::time::timeout(
1131 Duration::from_secs(1),
1132 tokio::spawn(await_backends(state.clone())),
1133 )
1134 .await
1135 .context("worker-aware health did not observe both backends")??;
1136 let (status, Json(body)) = healthcheck(State(state)).await;
1137 assert_eq!(status, StatusCode::OK);
1138 assert!(body.ready);
1139 assert_eq!(context_backend.health_requests.load(Ordering::SeqCst), 1);
1140 assert_eq!(decode_backend.health_requests.load(Ordering::SeqCst), 1);
1141 prefill_server.abort();
1142 decode_server.abort();
1143 Ok(())
1144 }
1145
1146 fn proxy_state(prefill: Vec<String>, decode: Vec<String>) -> Result<ProxyState> {
1147 ProxyState::new(Config {
1148 host: "127.0.0.1".to_owned(),
1149 port: 8000,
1150 prefill,
1151 decode,
1152 })
1153 .map_err(Into::into)
1154 }
1155
1156 fn context_response(body: Value) -> ContextResponse {
1157 ContextResponse {
1158 status: StatusCode::CREATED,
1159 content_type: Some("application/json".to_owned()),
1160 body,
1161 }
1162 }
1163
1164 fn valid_params(request_id: u64) -> Value {
1165 json!({
1166 "ctx_request_id": 1,
1167 "disagg_request_id": request_id,
1168 "first_gen_tokens": [3]
1169 })
1170 }
1171
1172 fn handoff_response(
1173 request_id: u64,
1174 prompt_token_ids: Value,
1175 usage: Value,
1176 params: Value,
1177 ) -> Value {
1178 json!({
1179 "choices": [{
1180 "finish_reason": "length",
1181 "disaggregated_params": params
1182 }],
1183 "prompt_token_ids": prompt_token_ids,
1184 "usage": usage,
1185 "assigned_for_fixture": request_id
1186 })
1187 }
1188
1189 fn first_sse_event(stream: &str) -> Result<Value> {
1190 let event = stream
1191 .strip_prefix("data: ")
1192 .and_then(|stream| stream.split_once("\n\n"))
1193 .map(|(event, _)| event)
1194 .context("response lacked an SSE data event")?;
1195 serde_json::from_str(event).map_err(Into::into)
1196 }
1197
1198 #[derive(Clone)]
1199 struct ObservedRequest {
1200 headers: HeaderMap,
1201 body: Value,
1202 path: &'static str,
1203 }
1204
1205 #[derive(Clone, Default)]
1206 struct ContextBackend {
1207 requests: Arc<Mutex<Vec<ObservedRequest>>>,
1208 health_requests: Arc<AtomicUsize>,
1209 }
1210
1211 #[derive(Clone)]
1212 struct DecodeBackend {
1213 requests: Arc<Mutex<Vec<ObservedRequest>>>,
1214 health_requests: Arc<AtomicUsize>,
1215 stream_gate: Arc<Notify>,
1216 }
1217
1218 impl Default for DecodeBackend {
1219 fn default() -> Self {
1220 Self {
1221 requests: Arc::new(Mutex::new(Vec::new())),
1222 health_requests: Arc::new(AtomicUsize::new(0)),
1223 stream_gate: Arc::new(Notify::new()),
1224 }
1225 }
1226 }
1227
1228 async fn spawn_context_backend(state: ContextBackend) -> Result<(String, JoinHandle<()>)> {
1229 let app = Router::new()
1230 .route("/health", get(context_health))
1231 .route("/v1/completions", post(context_completion))
1232 .route("/v1/chat/completions", post(context_chat_completion))
1233 .with_state(state);
1234 spawn_router(app).await
1235 }
1236
1237 async fn spawn_decode_backend(state: DecodeBackend) -> Result<(String, JoinHandle<()>)> {
1238 let app = Router::new()
1239 .route("/health", get(decode_health))
1240 .route("/v1/completions", post(decode_completion))
1241 .route("/v1/chat/completions", post(decode_chat_completion))
1242 .with_state(state);
1243 spawn_router(app).await
1244 }
1245
1246 async fn spawn_router(app: Router) -> Result<(String, JoinHandle<()>)> {
1247 let listener = TcpListener::bind("127.0.0.1:0").await?;
1248 let address = listener.local_addr()?;
1249 let server = tokio::spawn(async move {
1250 let _ = serve(listener, app).await;
1251 });
1252 Ok((format!("http://{address}"), server))
1253 }
1254
1255 async fn context_health(State(state): State<ContextBackend>) -> StatusCode {
1256 state.health_requests.fetch_add(1, Ordering::SeqCst);
1257 StatusCode::OK
1258 }
1259
1260 async fn decode_health(State(state): State<DecodeBackend>) -> StatusCode {
1261 state.health_requests.fetch_add(1, Ordering::SeqCst);
1262 StatusCode::OK
1263 }
1264
1265 async fn context_completion(
1266 State(state): State<ContextBackend>,
1267 headers: HeaderMap,
1268 Json(body): Json<Value>,
1269 ) -> Response<Body> {
1270 context_request(state, headers, body, RequestFamily::Completions).await
1271 }
1272
1273 async fn context_chat_completion(
1274 State(state): State<ContextBackend>,
1275 headers: HeaderMap,
1276 Json(body): Json<Value>,
1277 ) -> Response<Body> {
1278 context_request(state, headers, body, RequestFamily::ChatCompletions).await
1279 }
1280
1281 async fn context_request(
1282 state: ContextBackend,
1283 headers: HeaderMap,
1284 body: Value,
1285 family: RequestFamily,
1286 ) -> Response<Body> {
1287 state.requests.lock().await.push(ObservedRequest {
1288 headers,
1289 body: body.clone(),
1290 path: family.path(),
1291 });
1292 if body.get("mode").and_then(Value::as_str) == Some("context-fail") {
1293 return (StatusCode::INTERNAL_SERVER_ERROR, "context failed").into_response();
1294 }
1295 let request_id = body["disaggregated_params"]["disagg_request_id"].clone();
1296 let finish_reason = if body.get("mode").and_then(Value::as_str) == Some("complete") {
1297 "stop"
1298 } else {
1299 "length"
1300 };
1301 let mut choice = match family {
1302 RequestFamily::Completions => json!({"text": "answer"}),
1303 RequestFamily::ChatCompletions => {
1304 json!({"message": {"role": "assistant", "content": "answer"}})
1305 }
1306 };
1307 choice["finish_reason"] = Value::String(finish_reason.to_owned());
1308 choice["index"] = Value::from(0);
1309 choice["disaggregated_params"] = json!({
1310 "request_type": "context_only",
1311 "ctx_request_id": 91,
1312 "disagg_request_id": request_id,
1313 "first_gen_tokens": [8],
1314 "opaque_future_field": {"endpoint": "nixl://ctx"}
1315 });
1316 let mut response = json!({
1317 "id": "cmpl-context",
1318 "choices": [choice],
1319 "prompt_token_ids": [10, 11, 12],
1320 "usage": {"prompt_tokens": 3, "completion_tokens": 1},
1321 "opaque": "kept"
1322 });
1323 if matches!(family, RequestFamily::ChatCompletions) {
1324 response["object"] = Value::String("chat.completion".to_owned());
1325 response["prompt_token_ids_b64"] = Value::String("encoded".to_owned());
1326 }
1327 (StatusCode::CREATED, Json(response)).into_response()
1328 }
1329
1330 async fn decode_completion(
1331 State(state): State<DecodeBackend>,
1332 headers: HeaderMap,
1333 Json(body): Json<Value>,
1334 ) -> Response<Body> {
1335 decode_request(state, headers, body, RequestFamily::Completions).await
1336 }
1337
1338 async fn decode_chat_completion(
1339 State(state): State<DecodeBackend>,
1340 headers: HeaderMap,
1341 Json(body): Json<Value>,
1342 ) -> Response<Body> {
1343 decode_request(state, headers, body, RequestFamily::ChatCompletions).await
1344 }
1345
1346 async fn decode_request(
1347 state: DecodeBackend,
1348 headers: HeaderMap,
1349 body: Value,
1350 family: RequestFamily,
1351 ) -> Response<Body> {
1352 state.requests.lock().await.push(ObservedRequest {
1353 headers,
1354 body: body.clone(),
1355 path: family.path(),
1356 });
1357 let mode = body.get("mode").and_then(Value::as_str);
1358 if mode == Some("decode-fail") {
1359 return (StatusCode::INTERNAL_SERVER_ERROR, "decode failed").into_response();
1360 }
1361 if body.get("stream").and_then(Value::as_bool) == Some(true) {
1362 let gate = state.stream_gate.clone();
1363 let fail = mode == Some("stream-error");
1364 let body = Body::from_stream(stream! {
1365 yield Ok::<Bytes, std::io::Error>(Bytes::from_static(b"data: first\n\n"));
1366 gate.notified().await;
1367 if fail {
1368 yield Err(std::io::Error::other("decode body failed"));
1369 } else {
1370 yield Ok(Bytes::from_static(TERMINAL_SSE));
1371 }
1372 });
1373 return (
1374 StatusCode::ACCEPTED,
1375 [(header::CONTENT_TYPE, "text/event-stream")],
1376 body,
1377 )
1378 .into_response();
1379 }
1380 (
1381 StatusCode::CREATED,
1382 [(header::CONTENT_TYPE, "application/x-inferlab-test")],
1383 "decode-complete",
1384 )
1385 .into_response()
1386 }
1387}