Skip to main content

inferlab_proxy/
trtllm.rs

1//! Built-in routing for TensorRT-LLM prefill/decode serving under
2//! [[RFC-0003:C-TENSORRT-LLM-PREFILL-DECODE]].
3
4use 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
27/// Display name used in lifecycle/validation error messages.
28const 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}