use std::{
collections::{HashMap, HashSet},
fmt::Display,
sync::Arc,
time::{SystemTime, UNIX_EPOCH},
};
use axum::{
Json, Router,
body::Body,
extract::State,
http::Request,
http::{HeaderMap, StatusCode},
middleware::{self, Next},
response::{
IntoResponse, Response,
sse::{Event, KeepAlive, Sse},
},
routing::{get, post},
};
use base64::Engine as _;
use bytes::Bytes;
use dynamo_runtime::config::environment_names::llm as env_llm;
use dynamo_runtime::{
pipeline::{AsyncEngineContextProvider, Context},
protocols::annotated::AnnotationsProvider,
};
use futures::{StreamExt, stream};
use serde::{Deserialize, Serialize, de::DeserializeOwned};
use super::{
RouteDoc,
disconnect::{ConnectionHandle, create_connection_monitor, monitor_for_disconnects},
error::HttpError,
metadata::{attach_x_request_id, extract_metadata_from_http},
metrics::{
CancellationLabels, Endpoint, ErrorType, EventConverter,
process_chat_response_and_observe_metrics,
process_chat_response_using_event_converter_and_observe_metrics,
process_response_and_observe_metrics,
process_response_using_event_converter_and_observe_metrics,
},
service_v2,
};
use crate::engines::ValidateRequest;
use crate::preprocessor::PRESERVE_OMITTED_MAX_TOKENS_CONTEXT_KEY;
use crate::protocols::common::extensions::{
AGENT_CONTEXT_CONTEXT_KEY, AgentContext, SESSION_AFFINITY_CONTEXT_KEY, SessionAffinityId,
agent_context_from_headers, apply_header_routing_overrides, session_affinity_from_headers,
};
use crate::protocols::openai::chat_completions::aggregator::ChatCompletionAggregator;
use crate::protocols::openai::{
audios::{NvAudioSpeechResponse, NvCreateAudioSpeechRequest},
chat_completions::{
NvCreateChatCompletionRequest, NvCreateChatCompletionResponse,
NvCreateChatCompletionStreamResponse,
},
completions::{NvCreateCompletionRequest, NvCreateCompletionResponse},
embeddings::{NvCreateEmbeddingRequest, NvCreateEmbeddingResponse},
images::{NvCreateImageRequest, NvImagesResponse},
responses::{NvCreateResponse, NvResponse, ResponseParams, chat_completion_to_response},
videos::{NvCreateVideoRequest, NvVideosResponse},
};
use crate::protocols::unified::UnifiedRequest;
use crate::request_template::{RequestTemplate, resolve_request_model};
use crate::types::Annotated;
use dynamo_protocols::types::ChatCompletionMessageContent;
use dynamo_protocols::types::ChatCompletionMessageToolCallChunk;
use dynamo_protocols::types::ChatCompletionStreamResponseDelta;
use dynamo_protocols::types::Choice;
use dynamo_runtime::logging::get_distributed_tracing_context;
use tracing::Instrument;
pub const DYNAMO_REQUEST_ID_HEADER: &str = "x-dynamo-request-id";
pub const ANNOTATION_REQUEST_ID: &str = "request_id";
const VALIDATION_PREFIX: &str = "Validation: ";
use super::error::{SanitizedError, overload_status_code};
pub(super) fn rl_router(
drt: Arc<dynamo_runtime::DistributedRuntime>,
) -> anyhow::Result<axum::Router> {
let config = dynamo_rl::RlDiscoveryConfig::from_env(drt);
let state = dynamo_rl::RlDiscoveryState::new(config);
Ok(dynamo_rl::rl_router(state))
}
pub(super) fn get_body_limit() -> usize {
std::env::var(env_llm::DYN_HTTP_BODY_LIMIT_MB)
.ok()
.and_then(|s| s.parse::<usize>().ok())
.map(|mb| mb * 1024 * 1024)
.unwrap_or(45 * 1024 * 1024)
}
pub type ErrorResponse = (StatusCode, Json<ErrorMessage>);
#[derive(Serialize, Deserialize, Debug)]
pub(crate) struct ErrorMessage {
message: String,
#[serde(rename = "type")]
error_type: String,
code: u16,
#[serde(skip_serializing_if = "Option::is_none")]
details: Option<Box<serde_json::Value>>,
}
fn map_error_code_to_error_type(code: StatusCode) -> String {
match code.canonical_reason() {
Some(reason) => reason.to_string(),
None if code.as_u16() == 529 => "Overloaded".to_string(),
None if code.as_u16() == 499 => "Client Closed Request".to_string(),
None => "UnknownError".to_string(),
}
}
fn classify_error_for_metrics(code: StatusCode, message: &str) -> ErrorType {
match code {
StatusCode::BAD_REQUEST => {
if message.starts_with("Validation:") {
ErrorType::Validation
} else {
ErrorType::Internal
}
}
StatusCode::NOT_FOUND => ErrorType::NotFound, StatusCode::NOT_IMPLEMENTED => ErrorType::NotImplemented, StatusCode::TOO_MANY_REQUESTS => ErrorType::Overload, StatusCode::SERVICE_UNAVAILABLE => ErrorType::Unavailable, StatusCode::INTERNAL_SERVER_ERROR => ErrorType::Internal, _ if code.as_u16() == 529 => ErrorType::Overload, _ if code.as_u16() == 499 => ErrorType::Cancelled, _ if code.is_client_error() => ErrorType::Validation, _ => ErrorType::Internal, }
}
fn extract_error_type_from_response(response: &ErrorResponse) -> ErrorType {
classify_error_for_metrics(response.0, &response.1.message)
}
fn find_invalid_argument_in_chain<'a>(
err: &'a (dyn std::error::Error + 'static),
) -> Option<&'a dynamo_runtime::error::DynamoError> {
use dynamo_runtime::error::{BackendError, ErrorType};
let mut current = Some(err);
while let Some(e) = current {
if let Some(dynamo_err) = e.downcast_ref::<dynamo_runtime::error::DynamoError>()
&& matches!(
dynamo_err.error_type(),
ErrorType::InvalidArgument | ErrorType::Backend(BackendError::InvalidArgument)
)
{
return Some(dynamo_err);
}
current = e.source();
}
None
}
fn find_queue_rejection_in_chain<'a>(
err: &'a (dyn std::error::Error + 'static),
) -> Option<&'a dynamo_kv_router::scheduling::QueueRejection> {
let mut current = Some(err);
while let Some(error) = current {
if let Some(rejection) =
error.downcast_ref::<dynamo_kv_router::scheduling::QueueRejection>()
{
return Some(rejection);
}
current = error.source();
}
None
}
impl ErrorMessage {
pub fn model_not_found() -> ErrorResponse {
let code = StatusCode::NOT_FOUND;
let error_type = map_error_code_to_error_type(code);
(
code,
Json(ErrorMessage {
message: "Model not found".to_string(),
error_type,
code: code.as_u16(),
details: None,
}),
)
}
pub fn from_model_error(e: &crate::discovery::ModelManagerError) -> ErrorResponse {
match e {
crate::discovery::ModelManagerError::ModelUnavailable(model) => {
Self::service_unavailable_with_body(model_not_ready_message(model))
}
_ => Self::model_not_found(),
}
}
pub fn _service_unavailable() -> ErrorResponse {
let code = StatusCode::SERVICE_UNAVAILABLE;
let error_type = map_error_code_to_error_type(code);
(
code,
Json(ErrorMessage {
message: "Service is not ready".to_string(),
error_type,
code: code.as_u16(),
details: None,
}),
)
}
pub fn service_unavailable_with_body(message: String) -> ErrorResponse {
let code = StatusCode::SERVICE_UNAVAILABLE;
let error_type = map_error_code_to_error_type(code);
(
code,
Json(ErrorMessage {
message,
error_type,
code: code.as_u16(),
details: None,
}),
)
}
pub fn internal_server_error(msg: &str) -> ErrorResponse {
tracing::error!("Internal server error: {msg}");
let code = StatusCode::INTERNAL_SERVER_ERROR;
let error_type = map_error_code_to_error_type(code);
(
code,
Json(ErrorMessage {
message: msg.to_string(),
error_type,
code: code.as_u16(),
details: None,
}),
)
}
pub fn internal_server_error_with_details(
public_msg: &str,
details: impl std::fmt::Display,
) -> ErrorResponse {
tracing::error!("Internal server error: {public_msg}: {details}");
let code = StatusCode::INTERNAL_SERVER_ERROR;
let error_type = map_error_code_to_error_type(code);
(
code,
Json(ErrorMessage {
message: public_msg.to_string(),
error_type,
code: code.as_u16(),
details: None,
}),
)
}
pub fn sanitized_with_details(
err: SanitizedError,
details: impl std::fmt::Display,
) -> ErrorResponse {
let status = err.status();
if err.log_as_error() {
tracing::error!(status = %status, "{err}: {details}");
} else {
tracing::debug!(status = %status, "{err}: {details}");
}
(
status,
Json(ErrorMessage {
message: err.to_string(),
error_type: map_error_code_to_error_type(status),
code: status.as_u16(),
details: None,
}),
)
}
pub fn not_implemented_error<T: Display>(msg: T) -> ErrorResponse {
tracing::error!("Not Implemented error: {msg}");
let code = StatusCode::NOT_IMPLEMENTED;
let error_type = map_error_code_to_error_type(code);
(
code,
Json(ErrorMessage {
message: msg.to_string(),
error_type,
code: code.as_u16(),
details: None,
}),
)
}
pub fn request_headers_too_large(msg: &str) -> ErrorResponse {
let code = StatusCode::REQUEST_HEADER_FIELDS_TOO_LARGE;
let error_type = map_error_code_to_error_type(code);
(
code,
Json(ErrorMessage {
message: msg.to_string(),
error_type,
code: code.as_u16(),
details: None,
}),
)
}
pub fn from_anyhow(err: anyhow::Error, alt_msg: &str) -> ErrorResponse {
if let Some(rejection) = find_queue_rejection_in_chain(err.as_ref()) {
let code = overload_status_code();
return (
code,
Json(ErrorMessage {
message: rejection.to_string(),
error_type: map_error_code_to_error_type(code),
code: code.as_u16(),
details: serde_json::to_value(rejection).ok().map(Box::new),
}),
);
}
if super::metrics::request_was_rejected(err.as_ref()) {
return ErrorMessage::sanitized_with_details(
SanitizedError::Overloaded,
format!("{err:#}"),
);
}
if super::metrics::request_was_unavailable(err.as_ref()) {
return ErrorMessage::sanitized_with_details(
SanitizedError::Unavailable,
format!("{err:#}"),
);
}
if let Some(dynamo_err) = find_invalid_argument_in_chain(err.as_ref()) {
return (
StatusCode::BAD_REQUEST,
Json(ErrorMessage {
message: dynamo_err.message().to_string(),
error_type: map_error_code_to_error_type(StatusCode::BAD_REQUEST),
code: StatusCode::BAD_REQUEST.as_u16(),
details: None,
}),
);
}
if super::metrics::request_was_cancelled(err.as_ref()) {
return ErrorMessage::sanitized_with_details(
SanitizedError::Cancelled,
format!("{err:#}"),
);
}
match err.downcast::<HttpError>() {
Ok(http_error) => ErrorMessage::from_http_error(http_error),
Err(err) => {
ErrorMessage::internal_server_error_with_details(alt_msg, format!("{err:#}"))
}
}
}
pub fn from_http_error(err: HttpError) -> ErrorResponse {
if err.code == 499 {
return ErrorMessage::sanitized_with_details(SanitizedError::Cancelled, err.message);
}
if err.code < 400 || err.code >= 500 {
return ErrorMessage::sanitized_with_details(SanitizedError::Internal, err.message);
}
match StatusCode::from_u16(err.code) {
Ok(code) => (
code,
Json(ErrorMessage {
message: err.message,
error_type: map_error_code_to_error_type(code),
code: code.as_u16(),
details: None,
}),
),
Err(_) => ErrorMessage::sanitized_with_details(SanitizedError::Internal, err.message),
}
}
}
impl From<HttpError> for ErrorMessage {
fn from(err: HttpError) -> Self {
ErrorMessage {
message: err.message,
error_type: map_error_code_to_error_type(
StatusCode::from_u16(err.code).unwrap_or(StatusCode::INTERNAL_SERVER_ERROR),
),
code: err.code,
details: None,
}
}
}
pub async fn smart_json_error_middleware(request: Request<Body>, next: Next) -> Response {
let response = next.run(request).await;
if response.status() == StatusCode::UNPROCESSABLE_ENTITY {
let (_parts, body) = response.into_parts();
let body_bytes = axum::body::to_bytes(body, get_body_limit())
.await
.unwrap_or_default();
let error_message = String::from_utf8_lossy(&body_bytes).to_string();
(
StatusCode::BAD_REQUEST,
Json(ErrorMessage {
message: error_message,
error_type: map_error_code_to_error_type(StatusCode::BAD_REQUEST),
code: StatusCode::BAD_REQUEST.as_u16(),
details: None,
}),
)
.into_response()
} else {
response
}
}
pub(super) fn get_or_create_request_id(headers: &HeaderMap) -> String {
let validated_header = if let Some(raw) = headers.get(DYNAMO_REQUEST_ID_HEADER) {
tracing::warn!(
"{} header is deprecated (DEP #7812); server-generated request IDs should be used instead",
DYNAMO_REQUEST_ID_HEADER
);
match raw.to_str() {
Err(_) => {
tracing::warn!(
"{} header must be a valid UTF-8 string",
DYNAMO_REQUEST_ID_HEADER
);
None
}
Ok(s) if uuid::Uuid::parse_str(s).is_err() => {
tracing::warn!(
"{} header must be a valid UUID, got: {}",
DYNAMO_REQUEST_ID_HEADER,
s
);
None
}
Ok(s) => Some(s.to_string()),
}
} else {
None
};
if let Some(trace_context) = get_distributed_tracing_context()
&& let Some(request_id) = trace_context.request_id
{
return request_id;
}
validated_header.unwrap_or_else(|| uuid::Uuid::new_v4().to_string())
}
fn context_from_headers<T: Send + Sync + 'static>(
request: T,
request_id: String,
headers: &HeaderMap,
) -> Result<Context<T>, ErrorResponse> {
let metadata = extract_metadata_from_http(headers)
.map_err(|err| ErrorMessage::request_headers_too_large(&err.to_string()))?;
let mut request = Context::with_id_and_metadata(request, request_id, metadata);
attach_x_request_id(&mut request, headers);
if let Some(agent_context) = agent_context_from_headers(headers) {
request.insert(AGENT_CONTEXT_CONTEXT_KEY, agent_context);
}
if let Some(session_affinity) = session_affinity_from_headers(headers) {
request.insert(SESSION_AFFINITY_CONTEXT_KEY, session_affinity);
}
Ok(request)
}
fn copy_context_metadata<T: Send + Sync + 'static, U: Send + Sync + 'static>(
source: &Context<T>,
target: &mut Context<U>,
) {
if crate::request_trace::is_enabled()
&& let Ok(x_request_id) =
source.get::<String>(crate::request_trace::X_REQUEST_ID_CONTEXT_KEY)
{
target.insert(
crate::request_trace::X_REQUEST_ID_CONTEXT_KEY,
x_request_id.as_ref().clone(),
);
}
if let Ok(agent_context) = source.get::<AgentContext>(AGENT_CONTEXT_CONTEXT_KEY) {
target.insert(AGENT_CONTEXT_CONTEXT_KEY, agent_context.as_ref().clone());
}
if let Ok(session_affinity) = source.get::<SessionAffinityId>(SESSION_AFFINITY_CONTEXT_KEY) {
target.insert(
SESSION_AFFINITY_CONTEXT_KEY,
session_affinity.as_ref().clone(),
);
}
}
fn warn_nvext_disabled(endpoint: &str, nvext_present: bool, headers: &HeaderMap) {
use crate::protocols::common::extensions::{
HEADER_DATA_PARALLEL_RANK_ALIAS, HEADER_DP_RANK, HEADER_DP_RANK_ALIAS,
HEADER_PREFILL_DP_RANK, HEADER_PREFILL_DP_RANK_ALIAS, HEADER_PREFILL_INSTANCE_ID,
HEADER_PREFILL_INSTANCE_ID_ALIAS, HEADER_REQUEST_PRIORITY, HEADER_REQUEST_STRICT_PRIORITY,
HEADER_WORKER_INSTANCE_ID, HEADER_WORKER_INSTANCE_ID_ALIAS,
};
let header_present = [
HEADER_WORKER_INSTANCE_ID,
HEADER_WORKER_INSTANCE_ID_ALIAS,
HEADER_PREFILL_INSTANCE_ID,
HEADER_PREFILL_INSTANCE_ID_ALIAS,
HEADER_DP_RANK,
HEADER_DP_RANK_ALIAS,
HEADER_DATA_PARALLEL_RANK_ALIAS,
HEADER_PREFILL_DP_RANK,
HEADER_PREFILL_DP_RANK_ALIAS,
HEADER_REQUEST_PRIORITY,
HEADER_REQUEST_STRICT_PRIORITY,
]
.iter()
.any(|h| headers.contains_key(*h));
if nvext_present || header_present {
tracing::warn!(
endpoint,
"request carried nvext data but the nvext extension is disabled on this frontend; dropping it"
);
}
}
async fn handler_completions(
State(state): State<Arc<service_v2::State>>,
headers: HeaderMap,
body: Bytes,
) -> Result<Response, ErrorResponse> {
ensure_json_content_type(&headers)?;
let mut request: NvCreateCompletionRequest = parse_json_request("completions", &body)?;
check_ready(&state)?;
check_model_serving_ready(&state, &request.inner.model)?;
request.nvext = if state.nvext_enabled() {
apply_header_routing_overrides(request.nvext.take(), &headers)
} else {
warn_nvext_disabled("completions", request.nvext.is_some(), &headers);
None
};
let request_id = get_or_create_request_id(&headers);
let streaming = request.inner.stream.unwrap_or(false);
let cancellation_labels = CancellationLabels {
model: state
.manager()
.metric_model_for(&request.inner.model)
.to_string(),
endpoint: Endpoint::Completions.to_string(),
request_type: if streaming { "stream" } else { "unary" }.to_string(),
};
let request = context_from_headers(request, request_id, &headers)?;
let context = request.context();
let (mut connection_handle, stream_handle) = create_connection_monitor(
context.clone(),
Some(state.metrics_clone()),
cancellation_labels,
)
.await;
let response = tokio::spawn(completions(state, request, stream_handle).in_current_span())
.await
.map_err(|e| {
ErrorMessage::internal_server_error_with_details(
"Failed to await chat completions task",
format!("{e:?}"),
)
})?;
connection_handle.disarm();
response
}
#[tracing::instrument(skip_all)]
async fn completions(
state: Arc<service_v2::State>,
request: Context<NvCreateCompletionRequest>,
stream_handle: ConnectionHandle,
) -> Result<Response, ErrorResponse> {
use crate::protocols::openai::completions::get_prompt_batch_size;
check_ready(&state)?;
validate_completion_stream_options(&request)?;
validate_completion_fields_generic(&request)?;
let batch_size = get_prompt_batch_size(&request.inner.prompt);
let n = request.inner.n.unwrap_or(1);
if batch_size == 1 {
return completions_single(state, request, stream_handle).await;
}
completions_batch(state, request, stream_handle, batch_size, n).await
}
#[tracing::instrument(skip_all)]
async fn completions_single(
state: Arc<service_v2::State>,
request: Context<NvCreateCompletionRequest>,
stream_handle: ConnectionHandle,
) -> Result<Response, ErrorResponse> {
let request_id = request.id().to_string();
let streaming = request.inner.stream.unwrap_or(false);
let model = request.inner.model.clone();
let metric_model = state.manager().metric_model_for(&model).to_string();
let mut inflight_guard = state.metrics_clone().create_inflight_guard(
&metric_model,
Endpoint::Completions,
streaming,
&request_id,
);
let http_queue_guard = state.metrics_clone().create_http_queue_guard(&metric_model);
let (engine, parsing_options) = state
.manager()
.get_completions_engine_with_parsing(&model)
.map_err(|e| {
let err_response = ErrorMessage::from_model_error(&e);
inflight_guard.mark_error(extract_error_type_from_response(&err_response));
err_response
})?;
let mut response_collector = state
.metrics_clone()
.create_response_collector(&metric_model);
let annotations = request.annotations();
let stream = engine.generate(request).await.map_err(|e| {
if super::metrics::request_was_rejected(e.as_ref()) {
state
.metrics_clone()
.inc_rejection(&model, super::metrics::Endpoint::Completions);
}
let err_response = ErrorMessage::from_anyhow(e, "Failed to generate completions");
inflight_guard.mark_error(extract_error_type_from_response(&err_response));
err_response
})?;
let ctx = stream.context();
let annotations = annotations.map_or(Vec::new(), |annotations| {
annotations
.iter()
.filter_map(|annotation| {
if annotation == ANNOTATION_REQUEST_ID {
Annotated::<NvCreateCompletionResponse>::from_annotation(
ANNOTATION_REQUEST_ID,
&request_id,
)
.ok()
} else {
None
}
})
.collect::<Vec<_>>()
});
let stream = stream::iter(annotations).chain(stream);
if streaming {
let mut http_queue_guard = Some(http_queue_guard);
let stream = stream
.filter(|r| {
futures::future::ready(
!r.data
.as_ref()
.is_some_and(is_empty_completion_stream_response),
)
})
.map(move |response| {
process_response_using_event_converter_and_observe_metrics(
EventConverter::from(response),
&mut response_collector,
&mut http_queue_guard,
)
})
.filter_map(|result| {
use futures::future;
future::ready(result.transpose())
});
let stream = monitor_for_disconnects(stream, ctx, inflight_guard, stream_handle);
let mut sse_stream = Sse::new(stream);
if let Some(keep_alive) = state.sse_keep_alive() {
sse_stream = sse_stream.keep_alive(KeepAlive::default().interval(keep_alive));
}
Ok(sse_stream.into_response())
} else {
let mut http_queue_guard = Some(http_queue_guard);
let stream = stream.inspect(move |response| {
process_response_and_observe_metrics(
response,
&mut response_collector,
&mut http_queue_guard,
);
});
let response = NvCreateCompletionResponse::from_annotated_stream(stream, parsing_options)
.await
.map_err(|e| {
tracing::error!(
"Failed to fold completions stream for {}: {:?}",
request_id,
e
);
let err_response = ErrorMessage::internal_server_error(&format!(
"Failed to fold completions stream for {request_id}"
));
inflight_guard.mark_error(extract_error_type_from_response(&err_response));
err_response
})?;
inflight_guard.mark_ok();
if ctx.is_killed() {
inflight_guard.mark_error(ErrorType::Cancelled);
}
Ok(Json(response).into_response())
}
}
#[tracing::instrument(skip_all)]
async fn completions_batch(
state: Arc<service_v2::State>,
request: Context<NvCreateCompletionRequest>,
stream_handle: ConnectionHandle,
batch_size: usize,
n: u8,
) -> Result<Response, ErrorResponse> {
use crate::protocols::openai::completions::extract_single_prompt;
use futures::stream::{self, StreamExt};
let request_id = request.id().to_string();
let streaming = request.inner.stream.unwrap_or(false);
let model = request.inner.model.clone();
let metric_model = state.manager().metric_model_for(&model).to_string();
let mut inflight_guard = state.metrics_clone().create_inflight_guard(
&metric_model,
Endpoint::Completions,
streaming,
&request_id,
);
let http_queue_guard = state.metrics_clone().create_http_queue_guard(&metric_model);
let (engine, parsing_options) = state
.manager()
.get_completions_engine_with_parsing(&model)
.map_err(|e| {
let err_response = ErrorMessage::from_model_error(&e);
inflight_guard.mark_error(extract_error_type_from_response(&err_response));
err_response
})?;
let mut response_collector = state
.metrics_clone()
.create_response_collector(&metric_model);
let annotations = request.annotations();
let mut all_streams = Vec::new();
let mut first_ctx = None;
for prompt_idx in 0..batch_size {
let single_prompt = extract_single_prompt(&request.inner.prompt, prompt_idx);
let mut single_request = request.content().clone();
single_request.inner.prompt = single_prompt;
let unique_request_id = format!("{}-{}", request.id(), prompt_idx);
let mut single_request_context = Context::with_id_and_metadata(
single_request,
unique_request_id,
request.metadata().clone(),
);
copy_context_metadata(&request, &mut single_request_context);
let stream = engine.generate(single_request_context).await.map_err(|e| {
if super::metrics::request_was_rejected(e.as_ref()) {
state
.metrics_clone()
.inc_rejection(&model, super::metrics::Endpoint::Completions);
}
let err_response = ErrorMessage::from_anyhow(e, "Failed to generate completions");
inflight_guard.mark_error(extract_error_type_from_response(&err_response));
err_response
})?;
if first_ctx.is_none() {
first_ctx = Some(stream.context());
}
let prompt_idx_u32 = prompt_idx as u32;
let n_u32 = n as u32;
let remapped_stream = stream.map(move |mut response| {
if let Some(ref mut data) = response.data {
for choice in &mut data.inner.choices {
choice.index += prompt_idx_u32 * n_u32;
}
}
response
});
all_streams.push(remapped_stream);
}
let merged_stream = stream::select_all(all_streams);
let ctx = first_ctx.expect("At least one stream should be generated");
let annotations_vec = annotations.map_or(Vec::new(), |annotations| {
annotations
.iter()
.filter_map(|annotation| {
if annotation == ANNOTATION_REQUEST_ID {
Annotated::<NvCreateCompletionResponse>::from_annotation(
ANNOTATION_REQUEST_ID,
&request_id,
)
.ok()
} else {
None
}
})
.collect::<Vec<_>>()
});
let merged_stream = stream::iter(annotations_vec).chain(merged_stream);
if streaming {
let mut http_queue_guard = Some(http_queue_guard);
let stream = merged_stream
.filter(|r| {
futures::future::ready(
!r.data
.as_ref()
.is_some_and(is_empty_completion_stream_response),
)
})
.map(move |response| {
process_response_using_event_converter_and_observe_metrics(
EventConverter::from(response),
&mut response_collector,
&mut http_queue_guard,
)
})
.filter_map(|result| {
use futures::future;
future::ready(result.transpose())
});
let stream = monitor_for_disconnects(stream, ctx, inflight_guard, stream_handle);
let mut sse_stream = Sse::new(stream);
if let Some(keep_alive) = state.sse_keep_alive() {
sse_stream = sse_stream.keep_alive(KeepAlive::default().interval(keep_alive));
}
Ok(sse_stream.into_response())
} else {
let mut http_queue_guard = Some(http_queue_guard);
let stream = merged_stream.inspect(move |response| {
process_response_and_observe_metrics(
response,
&mut response_collector,
&mut http_queue_guard,
);
});
let response = NvCreateCompletionResponse::from_annotated_stream(stream, parsing_options)
.await
.map_err(|e| {
tracing::error!(
"Failed to fold completions stream for {}: {:?}",
request_id,
e
);
let err_response = ErrorMessage::internal_server_error(&format!(
"Failed to fold completions stream for {request_id}"
));
inflight_guard.mark_error(extract_error_type_from_response(&err_response));
err_response
})?;
inflight_guard.mark_ok();
if ctx.is_killed() {
inflight_guard.mark_error(ErrorType::Cancelled);
}
Ok(Json(response).into_response())
}
}
#[tracing::instrument(skip_all)]
async fn embeddings(
State(state): State<Arc<service_v2::State>>,
headers: HeaderMap,
Json(mut request): Json<NvCreateEmbeddingRequest>,
) -> Result<Response, ErrorResponse> {
check_ready(&state)?;
check_model_serving_ready(&state, &request.inner.model)?;
if !state.nvext_enabled() {
warn_nvext_disabled("embeddings", request.nvext.is_some(), &headers);
request.nvext = None;
}
let request_id = get_or_create_request_id(&headers);
let request = context_from_headers(request, request_id, &headers)?;
let request_id = request.id().to_string();
let client_wants_float = !matches!(
request.inner.encoding_format.as_ref(),
Some(dynamo_protocols::types::EncodingFormat::Base64)
);
let streaming = false;
let model = &request.inner.model;
let metric_model = state.manager().metric_model_for(model).to_string();
let embedding_start = std::time::Instant::now();
let mut inflight = state.metrics_clone().create_inflight_guard(
&metric_model,
Endpoint::Embeddings,
streaming,
&request_id,
);
let http_queue_guard = state.metrics_clone().create_http_queue_guard(&metric_model);
let engine = state.manager().get_embeddings_engine(model).map_err(|e| {
let err_response = ErrorMessage::from_model_error(&e);
inflight.mark_error(extract_error_type_from_response(&err_response));
err_response
})?;
let mut response_collector = state
.metrics_clone()
.create_response_collector(&metric_model);
let model_name = model.to_string();
let stream = engine.generate(request).await.map_err(|e| {
if super::metrics::request_was_rejected(e.as_ref()) {
state
.metrics_clone()
.inc_rejection(&model_name, super::metrics::Endpoint::Embeddings);
}
let err_response = ErrorMessage::from_anyhow(e, "Failed to generate embeddings");
inflight.mark_error(extract_error_type_from_response(&err_response));
err_response
})?;
let mut http_queue_guard = Some(http_queue_guard);
let stream = stream.inspect(move |response| {
process_response_and_observe_metrics(
response,
&mut response_collector,
&mut http_queue_guard,
);
});
let mut response = NvCreateEmbeddingResponse::from_annotated_stream(stream)
.await
.map_err(|e| {
tracing::error!(
"Failed to fold embeddings stream for {}: {:?}",
request_id,
e
);
let err_response =
ErrorMessage::internal_server_error("Failed to fold embeddings stream");
inflight.mark_error(extract_error_type_from_response(&err_response));
err_response
})?;
if client_wants_float {
for embedding_obj in response.inner.data.iter_mut() {
if let dynamo_protocols::types::EmbeddingVector::Base64(s) = &embedding_obj.embedding {
match decode_base64_embedding_to_floats(s) {
Ok(floats) => {
embedding_obj.embedding =
dynamo_protocols::types::EmbeddingVector::Float(floats);
}
Err(e) => {
tracing::error!(
"Failed to decode base64 embedding for request {}: {:?}",
request_id,
e
);
let err_response = ErrorMessage::internal_server_error(
"Failed to decode embedding payload",
);
inflight.mark_error(extract_error_type_from_response(&err_response));
return Err(err_response);
}
}
}
}
}
state
.metrics_clone()
.observe_embedding_latency(&model_name, embedding_start.elapsed().as_secs_f64());
inflight.mark_ok();
Ok(Json(response).into_response())
}
fn decode_base64_embedding_to_floats(s: &str) -> Result<Vec<f32>, anyhow::Error> {
use base64::{Engine as _, engine::general_purpose::STANDARD};
let bytes = STANDARD.decode(s)?;
if bytes.len() % std::mem::size_of::<f32>() != 0 {
anyhow::bail!(
"base64-decoded byte length {} is not a multiple of 4",
bytes.len()
);
}
let mut floats = Vec::with_capacity(bytes.len() / std::mem::size_of::<f32>());
for chunk in bytes.chunks_exact(std::mem::size_of::<f32>()) {
floats.push(f32::from_le_bytes([chunk[0], chunk[1], chunk[2], chunk[3]]));
}
Ok(floats)
}
async fn handler_chat_completions(
State((state, template)): State<(Arc<service_v2::State>, Option<RequestTemplate>)>,
headers: HeaderMap,
body: Bytes,
) -> Result<Response, ErrorResponse> {
ensure_json_content_type(&headers)?;
let mut request: NvCreateChatCompletionRequest = parse_json_request("chat completions", &body)?;
check_ready(&state)?;
let resolved_model = resolve_request_model(&request.inner.model, template.as_ref());
if !resolved_model.is_empty() {
check_model_serving_ready(&state, resolved_model)?;
}
request.nvext = if state.nvext_enabled() {
apply_header_routing_overrides(request.nvext.take(), &headers)
} else {
warn_nvext_disabled("chat_completions", request.nvext.is_some(), &headers);
None
};
let request_id = get_or_create_request_id(&headers);
let streaming = request.inner.stream.unwrap_or(false);
let resolved_model = resolve_request_model(&request.inner.model, template.as_ref());
let cancellation_labels = CancellationLabels {
model: state.manager().metric_model_for(resolved_model).to_string(),
endpoint: Endpoint::ChatCompletions.to_string(),
request_type: if streaming { "stream" } else { "unary" }.to_string(),
};
let request = context_from_headers(request, request_id, &headers)?;
let context = request.context();
let (mut connection_handle, stream_handle) = create_connection_monitor(
context.clone(),
Some(state.metrics_clone()),
cancellation_labels,
)
.await;
let response =
tokio::spawn(chat_completions(state, template, request, stream_handle).in_current_span())
.await
.map_err(|e| {
ErrorMessage::internal_server_error_with_details(
"Failed to await chat completions task",
format!("{e:?}"),
)
})?;
connection_handle.disarm();
response
}
fn parse_json_request<T>(endpoint: &'static str, body: &[u8]) -> Result<T, ErrorResponse>
where
T: DeserializeOwned,
{
match serde_json::from_slice(body) {
Ok(request) => Ok(request),
Err(original_error) => {
if let Some(escaped_body) = escape_json_string_control_chars(body) {
match serde_json::from_slice(&escaped_body) {
Ok(request) => {
tracing::warn!(
endpoint,
"Accepted request after escaping unescaped control characters in JSON strings"
);
Ok(request)
}
Err(_) => parse_json_request_lossy(endpoint, body)
.map_err(|_| json_deserialize_error(original_error)),
}
} else {
parse_json_request_lossy(endpoint, body)
.map_err(|_| json_deserialize_error(original_error))
}
}
}
}
fn parse_json_request_lossy<T>(endpoint: &'static str, body: &[u8]) -> Result<T, serde_json::Error>
where
T: DeserializeOwned,
{
let lossy_body = String::from_utf8_lossy(body);
if lossy_body.as_bytes() == body {
return serde_json::from_slice(body);
}
let escaped_body = escape_json_string_control_chars(lossy_body.as_bytes())
.unwrap_or_else(|| lossy_body.into_owned().into_bytes());
let request = serde_json::from_slice(&escaped_body)?;
tracing::warn!(
endpoint,
"Accepted request after replacing invalid UTF-8 and escaping unescaped control characters in JSON strings"
);
Ok(request)
}
fn json_deserialize_error(error: serde_json::Error) -> ErrorResponse {
let code = StatusCode::BAD_REQUEST;
(
code,
Json(ErrorMessage {
message: format!("Failed to deserialize the JSON body into the target type: {error}"),
error_type: map_error_code_to_error_type(code),
code: code.as_u16(),
details: None,
}),
)
}
fn ensure_json_content_type(headers: &HeaderMap) -> Result<(), ErrorResponse> {
let Some(content_type) = headers
.get(axum::http::header::CONTENT_TYPE)
.and_then(|value| value.to_str().ok())
else {
return Err(unsupported_media_type_error());
};
if is_json_content_type(content_type) {
Ok(())
} else {
Err(unsupported_media_type_error())
}
}
fn unsupported_media_type_error() -> ErrorResponse {
let code = StatusCode::UNSUPPORTED_MEDIA_TYPE;
(
code,
Json(ErrorMessage {
message: "Expected request with Content-Type application/json".to_string(),
error_type: map_error_code_to_error_type(code),
code: code.as_u16(),
details: None,
}),
)
}
fn is_json_content_type(content_type: &str) -> bool {
let media_type = content_type.split(';').next().unwrap_or_default().trim();
let Some((media_type, subtype)) = media_type.split_once('/') else {
return false;
};
media_type.eq_ignore_ascii_case("application")
&& (subtype.eq_ignore_ascii_case("json")
|| subtype
.to_ascii_lowercase()
.rsplit_once('+')
.is_some_and(|(_, suffix)| suffix == "json"))
}
fn escape_json_string_control_chars(body: &[u8]) -> Option<Vec<u8>> {
let mut out = Vec::with_capacity(body.len());
let mut in_string = false;
let mut escaped = false;
let mut changed = false;
for &byte in body {
if in_string && byte <= 0x1f {
const HEX: &[u8; 16] = b"0123456789abcdef";
if escaped {
out.extend_from_slice(b"\\\\u00");
escaped = false;
} else {
out.extend_from_slice(b"\\u00");
}
out.push(HEX[(byte >> 4) as usize]);
out.push(HEX[(byte & 0x0f) as usize]);
changed = true;
continue;
}
out.push(byte);
if escaped {
escaped = false;
} else if in_string && byte == b'\\' {
escaped = true;
} else if byte == b'"' {
in_string = !in_string;
}
}
changed.then_some(out)
}
fn extract_backend_error_if_present<T: serde::Serialize>(
event: &Annotated<T>,
) -> Option<(String, StatusCode)> {
#[derive(serde::Deserialize)]
struct ErrorPayload {
message: Option<String>,
code: Option<u16>,
}
if let Some(event_type) = &event.event
&& event_type == "error"
{
use dynamo_runtime::error::{BackendError, ErrorType};
let invalid_argument = event.error.as_ref().filter(|error| {
matches!(
error.error_type(),
ErrorType::InvalidArgument | ErrorType::Backend(BackendError::InvalidArgument)
)
});
let error_str = if let Some(ref dynamo_err) = event.error {
let mut parts = Vec::new();
let mut current: Option<&dyn std::error::Error> = Some(dynamo_err);
while let Some(e) = current {
if let Some(de) = e.downcast_ref::<dynamo_runtime::error::DynamoError>() {
parts.push(de.message().to_string());
} else {
parts.push(e.to_string());
}
current = e.source();
}
parts.join(", ")
} else {
event
.comment
.as_ref()
.map(|c| c.join(", "))
.unwrap_or_else(|| "Unknown error".to_string())
};
let status_message = event
.error
.as_ref()
.map(|error| error.message())
.unwrap_or(&error_str);
if let Ok(error_payload) = serde_json::from_str::<ErrorPayload>(status_message) {
let code = match error_payload.code {
Some(code) => {
StatusCode::from_u16(code).unwrap_or(StatusCode::INTERNAL_SERVER_ERROR)
}
None if invalid_argument.is_some() => StatusCode::BAD_REQUEST,
None => StatusCode::INTERNAL_SERVER_ERROR,
};
let message = error_payload
.message
.unwrap_or_else(|| status_message.to_string());
return Some((message, code));
}
if let Some(invalid_argument) = invalid_argument {
return Some((
invalid_argument.message().to_string(),
StatusCode::BAD_REQUEST,
));
}
return Some((error_str, StatusCode::INTERNAL_SERVER_ERROR));
}
if let Some(data) = &event.data
&& let Ok(json_value) = serde_json::to_value(data)
&& let Ok(error_payload) = serde_json::from_value::<ErrorPayload>(json_value.clone())
&& let Some(code_num) = error_payload.code
&& code_num >= 400
{
let code = StatusCode::from_u16(code_num).unwrap_or(StatusCode::INTERNAL_SERVER_ERROR);
let message = error_payload
.message
.unwrap_or_else(|| json_value.to_string());
return Some((message, code));
}
if let Some(comments) = &event.comment
&& !comments.is_empty()
{
let comment_str = comments.join(", ");
if let Ok(error_payload) = serde_json::from_str::<ErrorPayload>(&comment_str)
&& let Some(code_num) = error_payload.code
&& code_num >= 400
{
let code = StatusCode::from_u16(code_num).unwrap_or(StatusCode::INTERNAL_SERVER_ERROR);
let message = error_payload.message.unwrap_or(comment_str);
return Some((message, code));
}
if event.data.is_none() && event.event.is_none() {
return Some((comment_str, StatusCode::INTERNAL_SERVER_ERROR));
}
}
None
}
fn is_annotation_frame<T>(e: &Annotated<T>) -> bool {
e.data.is_none()
&& e.error.is_none()
&& matches!(e.event.as_deref(), Some(tag) if tag != "error")
}
const MAX_LEADING_ANNOTATIONS: usize = 16;
pub(super) async fn check_for_backend_error(
mut stream: impl futures::Stream<Item = Annotated<NvCreateChatCompletionStreamResponse>>
+ Send
+ Unpin
+ 'static,
) -> Result<
impl futures::Stream<Item = Annotated<NvCreateChatCompletionStreamResponse>> + Send,
ErrorResponse,
> {
use futures::stream::StreamExt;
let mut buffered: Vec<Annotated<NvCreateChatCompletionStreamResponse>> = Vec::new();
while let Some(event) = stream.next().await {
if is_annotation_frame(&event) && buffered.len() < MAX_LEADING_ANNOTATIONS {
buffered.push(event);
continue;
}
if let Some((error_msg, status_code)) = extract_backend_error_if_present(&event) {
return Err(match SanitizedError::for_backend_status(status_code) {
Some(variant) => ErrorMessage::sanitized_with_details(variant, error_msg),
None => (
status_code,
Json(ErrorMessage {
message: error_msg,
error_type: map_error_code_to_error_type(status_code),
code: status_code.as_u16(),
details: None,
}),
),
});
}
buffered.push(event);
break;
}
Ok(futures::stream::iter(buffered).chain(stream))
}
#[derive(Serialize)]
struct ToolCallDispatchPayload<'a> {
choice_index: u32,
tool_call: &'a ChatCompletionMessageToolCallChunk,
}
#[derive(Serialize)]
struct ReasoningDispatchPayload<'a> {
index: u32,
reasoning_content: &'a str,
}
fn push_dispatch_event(
event_name: &str,
payload: &impl serde::Serialize,
out: &mut Vec<Result<Event, axum::Error>>,
) {
match serde_json::to_string(payload) {
Ok(json) => out.push(Ok(Event::default().event(event_name).data(json))),
Err(e) => {
tracing::warn!("streaming_{event_name}: failed to serialize: {e}");
}
}
}
fn is_empty_stream_response(resp: &NvCreateChatCompletionStreamResponse) -> bool {
if resp.nvext.is_some() {
return false;
}
resp.inner.usage.is_none()
&& resp.inner.choices.iter().all(|c| {
let ChatCompletionStreamResponseDelta {
content,
function_call,
tool_calls,
role: _,
refusal,
reasoning_content,
} = &c.delta;
let content_empty = match content {
None => true,
Some(ChatCompletionMessageContent::Text(t)) => t.is_empty(),
Some(ChatCompletionMessageContent::Parts(p)) => p.is_empty(),
};
c.finish_reason.is_none()
&& c.logprobs.is_none()
&& content_empty
&& function_call.is_none()
&& tool_calls.is_none()
&& refusal.is_none()
&& reasoning_content.is_none()
})
}
fn is_empty_completion_stream_response(resp: &NvCreateCompletionResponse) -> bool {
if resp.nvext.is_some() {
return false;
}
resp.inner.usage.is_none()
&& resp.inner.choices.iter().all(|c| {
let Choice {
text,
index: _,
logprobs,
finish_reason,
} = c;
text.is_empty() && finish_reason.is_none() && logprobs.is_none()
})
}
fn streaming_tool_dispatch_events(
response: &crate::types::Annotated<NvCreateChatCompletionStreamResponse>,
dispatched_ids: &mut HashSet<String>,
out: &mut Vec<Result<Event, axum::Error>>,
) {
let Some(data) = &response.data else {
return;
};
for choice in &data.inner.choices {
let Some(tool_calls) = &choice.delta.tool_calls else {
continue;
};
for chunk in tool_calls {
let has_name_and_args = chunk
.function
.as_ref()
.is_some_and(|f| f.name.is_some() && f.arguments.is_some());
if let (true, Some(id)) = (has_name_and_args, &chunk.id) {
if !dispatched_ids.insert(id.clone()) {
continue;
}
let payload = ToolCallDispatchPayload {
choice_index: choice.index,
tool_call: chunk,
};
push_dispatch_event("tool_call_dispatch", &payload, out);
}
}
}
}
fn accumulate_reasoning_dispatch(
response: &crate::types::Annotated<NvCreateChatCompletionStreamResponse>,
buffers: &mut HashMap<u32, String>,
out: &mut Vec<Result<Event, axum::Error>>,
) {
let Some(data) = &response.data else {
return;
};
for choice in &data.inner.choices {
let buffer = buffers.entry(choice.index).or_default();
let has_reasoning = choice
.delta
.reasoning_content
.as_ref()
.is_some_and(|r| !r.is_empty());
if has_reasoning {
buffer.push_str(choice.delta.reasoning_content.as_ref().unwrap());
}
if !buffer.is_empty() && (!has_reasoning || choice.finish_reason.is_some()) {
let payload = ReasoningDispatchPayload {
index: choice.index,
reasoning_content: buffer.as_str(),
};
push_dispatch_event("reasoning_dispatch", &payload, out);
buffer.clear();
}
}
}
async fn chat_completions(
state: Arc<service_v2::State>,
template: Option<RequestTemplate>,
mut request: Context<NvCreateChatCompletionRequest>,
mut stream_handle: ConnectionHandle,
) -> Result<Response, ErrorResponse> {
check_ready(&state)?;
let request_id = request.id().to_string();
let streaming = request.inner.stream.unwrap_or(false);
if let Some(template) = template {
if request.inner.model.is_empty() {
request.inner.model = template.model.clone();
}
if request.inner.temperature.unwrap_or(0.0) == 0.0 {
request.inner.temperature = Some(template.temperature);
}
if request.inner.max_completion_tokens.unwrap_or(0) == 0 {
request.inner.max_completion_tokens = Some(template.max_completion_tokens);
}
}
let model = request.inner.model.clone();
let metric_model = state.manager().metric_model_for(&model).to_string();
tracing::trace!("Received chat completions request: {:?}", request.content());
let mut inflight_guard = state.metrics_clone().create_inflight_guard(
&metric_model,
Endpoint::ChatCompletions,
streaming,
&request_id,
);
if let Err(err_response) = normalize_chat_reasoning_template_args(&mut request) {
inflight_guard.mark_error(extract_error_type_from_response(&err_response));
return Err(err_response);
}
if let Err(err_response) = validate_chat_completion_unsupported_fields(&request) {
inflight_guard.mark_error(extract_error_type_from_response(&err_response));
return Err(err_response);
}
if let Err(err_response) = validate_chat_completion_required_fields(&request) {
inflight_guard.mark_error(extract_error_type_from_response(&err_response));
return Err(err_response);
}
if let Err(err_response) = validate_chat_completion_stream_options(&request) {
inflight_guard.mark_error(extract_error_type_from_response(&err_response));
return Err(err_response);
}
if let Err(err_response) = validate_chat_completion_fields_generic(&request) {
inflight_guard.mark_error(extract_error_type_from_response(&err_response));
return Err(err_response);
}
let http_queue_guard = state.metrics_clone().create_http_queue_guard(&metric_model);
if request.inner.max_completion_tokens.is_none() {
request.insert(PRESERVE_OMITTED_MAX_TOKENS_CONTEXT_KEY, true);
}
tracing::trace!("Getting chat completions engine for model: {}", model);
let (engine, parsing_options) = state
.manager()
.get_chat_completions_engine_with_parsing(&model)
.map_err(|e| {
let err_response = ErrorMessage::from_model_error(&e);
inflight_guard.mark_error(extract_error_type_from_response(&err_response));
err_response
})?;
let parsing_options = parsing_options.with_experimental_v2_batch_eligible(
crate::protocols::openai::chat_completions::tool_parser_v2::batch_tool_choice_eligible(
request.inner.tool_choice.as_ref(),
),
);
let mut response_collector = state
.metrics_clone()
.create_response_collector(&metric_model);
let annotations = request.annotations();
let stream = engine.generate(request).await.map_err(|e| {
if super::metrics::request_was_rejected(e.as_ref()) {
state
.metrics_clone()
.inc_rejection(&model, super::metrics::Endpoint::ChatCompletions);
}
let err_response = ErrorMessage::from_anyhow(e, "Failed to generate completions");
inflight_guard.mark_error(extract_error_type_from_response(&err_response));
err_response
})?;
let ctx = stream.context();
let annotations = annotations.map_or(Vec::new(), |annotations| {
annotations
.iter()
.filter_map(|annotation| {
if annotation == ANNOTATION_REQUEST_ID {
Annotated::from_annotation(ANNOTATION_REQUEST_ID, &request_id).ok()
} else {
None
}
})
.collect::<Vec<_>>()
});
let stream = stream::iter(annotations).chain(stream);
if streaming {
stream_handle.arm();
let mut http_queue_guard = Some(http_queue_guard);
let tool_dispatch_enabled = state.streaming_tool_dispatch_enabled();
let reasoning_dispatch_enabled = state.streaming_reasoning_dispatch_enabled();
let mut reasoning_buffer: HashMap<u32, String> = HashMap::new();
let mut dispatched_tool_ids: HashSet<String> = HashSet::new();
let stream = async_stream::stream! {
let mut stream = Box::pin(stream);
let mut events: Vec<Result<Event, axum::Error>> = Vec::with_capacity(4);
while let Some(response) = stream.next().await {
events.clear();
if response.data.as_ref().is_some_and(is_empty_stream_response) {
continue;
}
if tool_dispatch_enabled {
streaming_tool_dispatch_events(
&response,
&mut dispatched_tool_ids,
&mut events,
);
}
if reasoning_dispatch_enabled {
accumulate_reasoning_dispatch(
&response,
&mut reasoning_buffer,
&mut events,
);
}
let sse_result = process_chat_response_using_event_converter_and_observe_metrics(
EventConverter::from(response),
&mut response_collector,
&mut http_queue_guard,
);
match sse_result {
Ok(Some(ev)) => events.push(Ok(ev)),
Ok(None) => {}
Err(e) => events.push(Err(e)),
}
events.reverse();
while let Some(event) = events.pop() {
yield event;
}
}
};
let stream = monitor_for_disconnects(stream, ctx, inflight_guard, stream_handle);
let mut sse_stream = Sse::new(stream);
if let Some(keep_alive) = state.sse_keep_alive() {
sse_stream = sse_stream.keep_alive(KeepAlive::default().interval(keep_alive));
}
Ok(sse_stream.into_response())
} else {
let stream_with_check =
check_for_backend_error(stream)
.await
.map_err(|error_response| {
tracing::error!(request_id, "Backend error detected: {:?}", error_response);
inflight_guard.mark_error(extract_error_type_from_response(&error_response));
error_response
})?;
let mut http_queue_guard = Some(http_queue_guard);
let stream = stream_with_check.inspect(move |response| {
process_chat_response_and_observe_metrics(
response,
&mut response_collector,
&mut http_queue_guard,
);
});
let response =
NvCreateChatCompletionResponse::from_annotated_stream(stream, parsing_options.clone())
.await
.map_err(|e| {
tracing::error!(
request_id,
"Failed to parse chat completion response: {:?}",
e
);
let err_response = ErrorMessage::internal_server_error(
"Failed to parse chat completion response",
);
inflight_guard.mark_error(extract_error_type_from_response(&err_response));
err_response
})?;
inflight_guard.mark_ok();
if ctx.is_killed() {
inflight_guard.mark_error(ErrorType::Cancelled);
}
Ok(Json(response).into_response())
}
}
#[allow(deprecated)]
pub fn validate_chat_completion_unsupported_fields(
request: &NvCreateChatCompletionRequest,
) -> Result<(), ErrorResponse> {
let inner = &request.inner;
if inner.function_call.is_some() {
return Err(ErrorMessage::not_implemented_error(
VALIDATION_PREFIX.to_string()
+ "`function_call` is deprecated. Please migrate to use `tool_choice` instead.",
));
}
if inner.functions.is_some() {
return Err(ErrorMessage::not_implemented_error(
VALIDATION_PREFIX.to_string()
+ "`functions` is deprecated. Please migrate to use `tools` instead.",
));
}
Ok(())
}
fn normalize_chat_reasoning_template_args(
request: &mut NvCreateChatCompletionRequest,
) -> Result<(), ErrorResponse> {
request.normalize_reasoning_template_args().map_err(|e| {
ErrorMessage::from_http_error(HttpError {
code: 400,
message: VALIDATION_PREFIX.to_string() + &e.to_string(),
})
})
}
pub fn validate_chat_completion_required_fields(
request: &NvCreateChatCompletionRequest,
) -> Result<(), ErrorResponse> {
let inner = &request.inner;
if inner.messages.is_empty() {
return Err(ErrorMessage::from_http_error(HttpError {
code: 400,
message: VALIDATION_PREFIX.to_string()
+ "The 'messages' field cannot be empty. At least one message is required.",
}));
}
Ok(())
}
pub fn validate_chat_completion_stream_options(
request: &NvCreateChatCompletionRequest,
) -> Result<(), ErrorResponse> {
let inner = &request.inner;
let streaming = inner.stream.unwrap_or(false);
if !streaming && inner.stream_options.is_some() {
return Err(ErrorMessage::from_http_error(HttpError {
code: 400,
message: VALIDATION_PREFIX.to_string()
+ "The 'stream_options' field is only allowed when 'stream' is set to true.",
}));
}
Ok(())
}
pub fn validate_chat_completion_fields_generic(
request: &NvCreateChatCompletionRequest,
) -> Result<(), ErrorResponse> {
request.validate().map_err(|e| {
ErrorMessage::from_http_error(HttpError {
code: 400,
message: VALIDATION_PREFIX.to_string() + &e.to_string(),
})
})
}
pub fn validate_completion_stream_options(
request: &NvCreateCompletionRequest,
) -> Result<(), ErrorResponse> {
let inner = &request.inner;
let streaming = inner.stream.unwrap_or(false);
if !streaming && inner.stream_options.is_some() {
return Err(ErrorMessage::from_http_error(HttpError {
code: 400,
message: VALIDATION_PREFIX.to_string()
+ "The 'stream_options' field is only allowed when 'stream' is set to true.",
}));
}
Ok(())
}
pub fn validate_completion_fields_generic(
request: &NvCreateCompletionRequest,
) -> Result<(), ErrorResponse> {
request.validate().map_err(|e| {
ErrorMessage::from_http_error(HttpError {
code: 400,
message: VALIDATION_PREFIX.to_string() + &e.to_string(),
})
})
}
async fn handler_responses(
State((state, template)): State<(Arc<service_v2::State>, Option<RequestTemplate>)>,
headers: HeaderMap,
Json(mut request): Json<NvCreateResponse>,
) -> Result<Response, ErrorResponse> {
check_ready(&state)?;
let resolved_model = resolve_request_model(
request.inner.model.as_deref().unwrap_or(""),
template.as_ref(),
);
if !resolved_model.is_empty() {
check_model_serving_ready(&state, resolved_model)?;
}
request.nvext = if state.nvext_enabled() {
apply_header_routing_overrides(request.nvext.take(), &headers)
} else {
warn_nvext_disabled("responses", request.nvext.is_some(), &headers);
None
};
let request_id = get_or_create_request_id(&headers);
let streaming = request.inner.stream.unwrap_or(false);
let raw_model = request.inner.model.as_deref().unwrap_or("");
let resolved_model = resolve_request_model(raw_model, template.as_ref());
let cancellation_labels = CancellationLabels {
model: state.manager().metric_model_for(resolved_model).to_string(),
endpoint: Endpoint::Responses.to_string(),
request_type: if streaming { "stream" } else { "unary" }.to_string(),
};
let request = context_from_headers(request, request_id, &headers)?;
let context = request.context();
let (mut connection_handle, stream_handle) = create_connection_monitor(
context.clone(),
Some(state.metrics_clone()),
cancellation_labels,
)
.await;
let response =
tokio::spawn(responses(state, template, request, stream_handle).in_current_span())
.await
.map_err(|e| {
ErrorMessage::internal_server_error_with_details(
"Failed to await responses task",
format!("{e:?}"),
)
})?;
connection_handle.disarm();
response
}
#[tracing::instrument(level = "debug", skip_all, fields(request_id = %request.id()))]
async fn responses(
state: Arc<service_v2::State>,
template: Option<RequestTemplate>,
mut request: Context<NvCreateResponse>,
mut stream_handle: ConnectionHandle,
) -> Result<Response, ErrorResponse> {
check_ready(&state)?;
if let Some(template) = template {
if request.inner.model.as_deref().unwrap_or("").is_empty() {
request.inner.model = Some(template.model.clone());
}
if request.inner.temperature.is_none() {
request.inner.temperature = Some(template.temperature);
}
if request.inner.max_output_tokens.is_none() {
request.inner.max_output_tokens = Some(template.max_completion_tokens);
}
}
tracing::trace!("Received responses request: {:?}", request.inner);
let model = request.inner.model.clone().unwrap_or_default();
let streaming = request.inner.stream.unwrap_or(false);
let metric_model = state.manager().metric_model_for(&model).to_string();
let http_queue_guard = state.metrics_clone().create_http_queue_guard(&metric_model);
let mut inflight_guard = state.metrics_clone().create_inflight_guard(
&metric_model,
Endpoint::Responses,
streaming,
request.id(),
);
if let Some(resp) = validate_response_unsupported_fields(&request) {
inflight_guard.mark_error(ErrorType::NotImplemented);
return Ok(resp.into_response());
}
let response_params = ResponseParams {
model: request.inner.model.clone(),
temperature: request.inner.temperature,
top_p: request.inner.top_p,
max_output_tokens: request.inner.max_output_tokens,
parallel_tool_calls: request.inner.parallel_tool_calls,
store: request.inner.store,
tools: request.inner.tools.clone(),
tool_choice: request.inner.tool_choice.clone(),
instructions: request.inner.instructions.clone(),
reasoning: request.inner.reasoning.clone(),
text: request.inner.text.clone(),
service_tier: request.inner.service_tier,
include: request.inner.include.clone(),
truncation: request.inner.truncation,
presence_penalty: None,
frequency_penalty: None,
prompt_cache_key: request.inner.prompt_cache_key.clone(),
prompt_cache_retention: request.inner.prompt_cache_retention,
safety_identifier: request.inner.safety_identifier.clone(),
};
let request_id = request.id().to_string();
let (orig_request, context) = request.into_parts();
let unified_request: UnifiedRequest = orig_request.try_into().map_err(|e: anyhow::Error| {
tracing::error!(
request_id,
error = %e,
"Failed to convert NvCreateResponse to UnifiedRequest",
);
let err_response = ErrorMessage::not_implemented_error(
VALIDATION_PREFIX.to_string()
+ "Failed to convert responses request: "
+ &e.to_string(),
);
inflight_guard.mark_error(extract_error_type_from_response(&err_response));
err_response
})?;
let responses_ctx = unified_request.responses_context().cloned();
let mut chat_request = unified_request.into_inner();
if let Err(err_response) = normalize_chat_reasoning_template_args(&mut chat_request) {
inflight_guard.mark_error(extract_error_type_from_response(&err_response));
return Err(err_response);
}
chat_request.inner.stream = Some(true);
chat_request.inner.stream_options =
Some(dynamo_protocols::types::ChatCompletionStreamOptions {
include_usage: true,
continuous_usage_stats: false,
});
let mut request = context.map(|mut _req| chat_request);
if response_params.max_output_tokens.is_none() {
request.insert(PRESERVE_OMITTED_MAX_TOKENS_CONTEXT_KEY, true);
}
tracing::trace!("Getting chat completions engine for model: {}", model);
let (engine, parsing_options) = state
.manager()
.get_chat_completions_engine_with_parsing(&model)
.map_err(|e| {
let err_response = ErrorMessage::from_model_error(&e);
inflight_guard.mark_error(extract_error_type_from_response(&err_response));
err_response
})?;
let parsing_options = parsing_options.with_experimental_v2_batch_eligible(
crate::protocols::openai::chat_completions::tool_parser_v2::batch_tool_choice_eligible(
request.inner.tool_choice.as_ref(),
),
);
let mut response_collector = state
.metrics_clone()
.create_response_collector(&metric_model);
tracing::trace!("Issuing generate call for responses");
let engine_stream = engine.generate(request).await.map_err(|e| {
if super::metrics::request_was_rejected(e.as_ref()) {
state
.metrics_clone()
.inc_rejection(&model, super::metrics::Endpoint::Responses);
}
let err_response = ErrorMessage::from_anyhow(e, "Failed to generate completions");
inflight_guard.mark_error(extract_error_type_from_response(&err_response));
err_response
})?;
let ctx = engine_stream.context();
if streaming {
stream_handle.arm();
use crate::protocols::openai::responses::stream_converter::ResponseStreamConverter;
let mut converter = match responses_ctx {
Some(ctx) => ResponseStreamConverter::with_context(model.clone(), response_params, ctx),
None => ResponseStreamConverter::new(model.clone(), response_params),
};
let mut http_queue_guard = Some(http_queue_guard);
let mut engine_stream = Box::pin(engine_stream);
let full_stream = async_stream::stream! {
let mut events = Vec::with_capacity(4);
converter.append_start_events(&mut events);
for event in events.drain(..) {
yield event.map_err(axum::Error::new);
}
let mut saw_error = false;
while let Some(annotated_chunk) = engine_stream.next().await {
process_chat_response_and_observe_metrics(
&annotated_chunk,
&mut response_collector,
&mut http_queue_guard,
);
if extract_backend_error_if_present(&annotated_chunk).is_some() {
saw_error = true;
continue;
}
let Some(stream_resp) = annotated_chunk.data else {
continue;
};
converter.append_chunk_events(&stream_resp, &mut events);
for event in events.drain(..) {
yield event.map_err(axum::Error::new);
}
}
if saw_error {
converter.append_error_events(&mut events);
} else {
converter.append_end_events(&mut events);
}
for event in events.drain(..) {
yield event.map_err(axum::Error::new);
}
};
let stream = monitor_for_disconnects(full_stream, ctx, inflight_guard, stream_handle);
let mut sse_stream = Sse::new(stream);
if let Some(keep_alive) = state.sse_keep_alive() {
sse_stream = sse_stream.keep_alive(KeepAlive::default().interval(keep_alive));
}
Ok(sse_stream.into_response())
} else {
let stream_with_check =
check_for_backend_error(engine_stream)
.await
.map_err(|error_response| {
tracing::error!(request_id, "Backend error detected: {:?}", error_response);
inflight_guard.mark_error(extract_error_type_from_response(&error_response));
error_response
})?;
let mut http_queue_guard = Some(http_queue_guard);
let stream = stream_with_check.inspect(move |response| {
process_chat_response_and_observe_metrics(
response,
&mut response_collector,
&mut http_queue_guard,
);
});
let response =
NvCreateChatCompletionResponse::from_annotated_stream(stream, parsing_options.clone())
.await
.map_err(|e| {
tracing::error!(request_id, "Failed to fold responses stream: {:?}", e);
let err_response =
ErrorMessage::internal_server_error("Failed to fold responses stream");
inflight_guard.mark_error(extract_error_type_from_response(&err_response));
err_response
})?;
let response: NvResponse =
chat_completion_to_response(response, &response_params, responses_ctx.as_ref())
.map_err(|e| {
tracing::error!(
request_id,
"Failed to convert NvCreateChatCompletionResponse to NvResponse: {:?}",
e
);
let err_response =
ErrorMessage::internal_server_error("Failed to convert internal response");
inflight_guard.mark_error(extract_error_type_from_response(&err_response));
err_response
})?;
inflight_guard.mark_ok();
if ctx.is_killed() {
inflight_guard.mark_error(ErrorType::Cancelled);
}
Ok(Json(response).into_response())
}
}
pub fn validate_response_unsupported_fields(
request: &NvCreateResponse,
) -> Option<impl IntoResponse> {
let inner = &request.inner;
if let Some(field) = request
.nvext
.as_ref()
.and_then(|nvext| nvext.extra_fields.as_ref())
.and_then(|fields| {
fields
.iter()
.find(|field| matches!(field.as_str(), "completion_token_ids" | "prompt_logprobs"))
})
{
return Some(ErrorMessage::not_implemented_error(format!(
"{VALIDATION_PREFIX}`nvext.extra_fields=[\"{field}\"]` is not supported by the Responses API."
)));
}
if inner.background == Some(true) {
return Some(ErrorMessage::not_implemented_error(
VALIDATION_PREFIX.to_string() + "`background: true` is not supported.",
));
}
if inner.previous_response_id.is_some() {
return Some(ErrorMessage::not_implemented_error(
VALIDATION_PREFIX.to_string() + "`previous_response_id` is not supported.",
));
}
if inner.prompt.is_some() {
return Some(ErrorMessage::not_implemented_error(
VALIDATION_PREFIX.to_string() + "`prompt` is not supported.",
));
}
if inner.max_tool_calls.is_some() {
return Some(ErrorMessage::not_implemented_error(
VALIDATION_PREFIX.to_string() + "`max_tool_calls` is not supported.",
));
}
None
}
pub(crate) fn check_ready(state: &Arc<service_v2::State>) -> Result<(), ErrorResponse> {
if !state.is_ready() {
return Err(ErrorMessage::_service_unavailable());
}
Ok(())
}
pub(crate) fn model_not_ready_message(model_name: &str) -> String {
format!(
"Model `{model_name}` is not ready to serve requests yet. \
The deployment may still be starting up or is not fully provisioned. \
Please retry shortly."
)
}
pub(crate) fn check_model_serving_ready(
state: &Arc<service_v2::State>,
model_name: &str,
) -> Result<(), ErrorResponse> {
let Some(model) = state.manager().get_model(model_name) else {
return Ok(());
};
if model.has_ready_workers() {
return Ok(());
}
Err(ErrorMessage::service_unavailable_with_body(
model_not_ready_message(model_name),
))
}
async fn list_models_openai(
State(state): State<Arc<service_v2::State>>,
) -> Result<Response, ErrorResponse> {
check_ready(&state)?;
let created = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_secs();
let cards = state.manager().get_model_cards();
let card_map: HashMap<String, u32> = cards
.iter()
.map(|c| (c.display_name.clone(), c.effective_context_length()))
.collect();
let cw_override: Option<u64> = std::env::var("DYN_CONTEXT_WINDOW")
.ok()
.and_then(|v| v.parse().ok());
let mot_override: Option<u64> = std::env::var("DYN_MAX_OUTPUT_TOKENS")
.ok()
.and_then(|v| v.parse().ok());
let mut data = Vec::new();
let models: HashSet<String> = state.manager().serving_ready_display_names();
for model_name in models {
let context_window = cw_override.or_else(|| card_map.get(&model_name).map(|&cl| cl as u64));
data.push(ModelListing {
id: model_name.clone(),
object: "model",
created,
owned_by: "nvidia".to_string(),
context_window,
max_output_tokens: mot_override,
});
}
let out = ListModelOpenAI {
object: "list",
data,
};
Ok(Json(out).into_response())
}
#[derive(Serialize)]
struct ListModelOpenAI {
object: &'static str, data: Vec<ModelListing>,
}
#[derive(Serialize)]
struct ModelListing {
id: String,
object: &'static str, created: u64, owned_by: String,
#[serde(skip_serializing_if = "Option::is_none")]
context_window: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
max_output_tokens: Option<u64>,
}
pub fn completions_router(
state: Arc<service_v2::State>,
path: Option<String>,
) -> (Vec<RouteDoc>, Router) {
let path = path.unwrap_or("/v1/completions".to_string());
let doc = RouteDoc::new(axum::http::Method::POST, &path);
let router = Router::new()
.route(&path, post(handler_completions))
.layer(middleware::from_fn(smart_json_error_middleware))
.layer(axum::extract::DefaultBodyLimit::max(get_body_limit()))
.with_state(state);
(vec![doc], router)
}
pub fn chat_completions_router(
state: Arc<service_v2::State>,
template: Option<RequestTemplate>,
path: Option<String>,
) -> (Vec<RouteDoc>, Router) {
let path = path.unwrap_or("/v1/chat/completions".to_string());
let doc = RouteDoc::new(axum::http::Method::POST, &path);
let router = Router::new()
.route(&path, post(handler_chat_completions))
.layer(middleware::from_fn(smart_json_error_middleware))
.layer(axum::extract::DefaultBodyLimit::max(get_body_limit()))
.with_state((state, template));
(vec![doc], router)
}
pub fn embeddings_router(
state: Arc<service_v2::State>,
path: Option<String>,
) -> (Vec<RouteDoc>, Router) {
let path = path.unwrap_or("/v1/embeddings".to_string());
let doc = RouteDoc::new(axum::http::Method::POST, &path);
let router = Router::new()
.route(&path, post(embeddings))
.layer(middleware::from_fn(smart_json_error_middleware))
.layer(axum::extract::DefaultBodyLimit::max(get_body_limit()))
.with_state(state);
(vec![doc], router)
}
pub fn list_models_router(
state: Arc<service_v2::State>,
path: Option<String>,
) -> (Vec<RouteDoc>, Router) {
let openai_path = path.unwrap_or("/v1/models".to_string());
let retrieve_path = format!("{}/{{*model_id}}", openai_path);
let doc_for_openai = RouteDoc::new(axum::http::Method::GET, &openai_path);
let doc_for_retrieve = RouteDoc::new(axum::http::Method::GET, &retrieve_path);
let doc_for_readiness = RouteDoc::new(
axum::http::Method::GET,
format!("{}/{{model_id}}/ready", openai_path),
);
let router = Router::new()
.route(&openai_path, get(list_models_openai))
.route(&retrieve_path, get(get_model_openai))
.with_state(state);
(
vec![doc_for_openai, doc_for_retrieve, doc_for_readiness],
router,
)
}
async fn get_model_openai(
State(state): State<Arc<service_v2::State>>,
axum::extract::Path(model_id): axum::extract::Path<String>,
) -> Result<Response, ErrorResponse> {
check_ready(&state)?;
let model_id = model_id.strip_prefix('/').unwrap_or(&model_id);
if state.manager().get_model(model_id).is_some() {
return get_model_retrieve(&state, model_id);
}
if let Some(base) = model_id.strip_suffix("/ready")
&& state.manager().get_model(base).is_some()
{
return get_model_readiness(&state, base);
}
Err(ErrorMessage::model_not_found())
}
fn get_model_retrieve(
state: &Arc<service_v2::State>,
model_id: &str,
) -> Result<Response, ErrorResponse> {
check_model_serving_ready(state, model_id)?;
let created = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_secs();
let cards = state.manager().get_model_cards();
let context_length = cards
.iter()
.find(|c| c.display_name == model_id)
.map(|c| c.effective_context_length() as u64);
let context_window: Option<u64> = std::env::var("DYN_CONTEXT_WINDOW")
.ok()
.and_then(|v| v.parse().ok())
.or(context_length);
let max_output_tokens: Option<u64> = std::env::var("DYN_MAX_OUTPUT_TOKENS")
.ok()
.and_then(|v| v.parse().ok());
Ok(Json(ModelListing {
id: model_id.to_string(),
object: "model",
created,
owned_by: "nvidia".to_string(),
context_window,
max_output_tokens,
})
.into_response())
}
fn get_model_readiness(
state: &Arc<service_v2::State>,
model_id: &str,
) -> Result<Response, ErrorResponse> {
let model = state
.manager()
.get_model(model_id)
.ok_or_else(ErrorMessage::model_not_found)?;
Ok(Json(model.namespace_readiness()).into_response())
}
pub fn responses_router(
state: Arc<service_v2::State>,
template: Option<RequestTemplate>,
path: Option<String>,
) -> (Vec<RouteDoc>, Router) {
let path = path.unwrap_or("/v1/responses".to_string());
let doc = RouteDoc::new(axum::http::Method::POST, &path);
let router = Router::new()
.route(&path, post(handler_responses))
.layer(middleware::from_fn(smart_json_error_middleware))
.layer(axum::extract::DefaultBodyLimit::max(get_body_limit()))
.with_state((state, template));
(vec![doc], router)
}
async fn images(
State(state): State<Arc<service_v2::State>>,
headers: HeaderMap,
Json(request): Json<NvCreateImageRequest>,
) -> Result<Response, ErrorResponse> {
check_ready(&state)?;
let request_id = get_or_create_request_id(&headers);
let request = context_from_headers(request, request_id, &headers)?;
let request_id = request.id().to_string();
let streaming = false;
let model = request
.inner
.model
.as_ref()
.map(|m| match m {
dynamo_protocols::types::ImageModel::DallE2 => "dall-e-2".to_string(),
dynamo_protocols::types::ImageModel::DallE3 => "dall-e-3".to_string(),
dynamo_protocols::types::ImageModel::GptImage1 => "gpt-image-1".to_string(),
dynamo_protocols::types::ImageModel::GptImage1dot5 => "gpt-image-1.5".to_string(),
dynamo_protocols::types::ImageModel::GptImage1Mini => "gpt-image-1-mini".to_string(),
dynamo_protocols::types::ImageModel::Other(s) => s.clone(),
})
.unwrap_or_else(|| "diffusion".to_string());
check_model_serving_ready(&state, &model)?;
let metric_model = state.manager().metric_model_for(&model).to_string();
let http_queue_guard = state.metrics_clone().create_http_queue_guard(&metric_model);
let engine = state
.manager()
.get_images_engine(&model)
.map_err(|e| ErrorMessage::from_model_error(&e))?;
let mut inflight = state.metrics_clone().create_inflight_guard(
&model,
Endpoint::Images,
streaming,
&request_id,
);
let mut response_collector = state.metrics_clone().create_response_collector(&model);
let stream = engine.generate(request).await.map_err(|e| {
if super::metrics::request_was_rejected(e.as_ref()) {
state
.metrics_clone()
.inc_rejection(&model, super::metrics::Endpoint::Images);
}
let err_response = ErrorMessage::from_anyhow(e, "Failed to generate images");
inflight.mark_error(extract_error_type_from_response(&err_response));
err_response
})?;
let mut http_queue_guard = Some(http_queue_guard);
let stream = stream.inspect(move |response| {
process_response_and_observe_metrics(
response,
&mut response_collector,
&mut http_queue_guard,
);
});
let response = NvImagesResponse::from_annotated_stream(stream)
.await
.map_err(|e| {
tracing::error!("Failed to fold images stream for {}: {:?}", request_id, e);
let err_response = ErrorMessage::internal_server_error("Failed to fold images stream");
inflight.mark_error(extract_error_type_from_response(&err_response));
err_response
})?;
inflight.mark_ok();
Ok(Json(response).into_response())
}
async fn images_edits(
state: State<Arc<service_v2::State>>,
headers: HeaderMap,
Json(request): Json<NvCreateImageRequest>,
) -> Result<Response, ErrorResponse> {
if request.input_reference.is_none() {
let code = StatusCode::BAD_REQUEST;
return Err((
code,
Json(ErrorMessage {
message: "input_reference is required for /v1/images/edits".to_string(),
error_type: map_error_code_to_error_type(code),
code: code.as_u16(),
details: None,
}),
));
}
images(state, headers, Json(request)).await
}
pub fn images_router(
state: Arc<service_v2::State>,
path: Option<String>,
) -> (Vec<RouteDoc>, Router) {
let generations_path = path.unwrap_or("/v1/images/generations".to_string());
let edits_path = generations_path.replace("/generations", "/edits");
let doc = RouteDoc::new(axum::http::Method::POST, &generations_path);
let edits_doc = RouteDoc::new(axum::http::Method::POST, &edits_path);
let router = Router::new()
.route(&generations_path, post(images))
.route(&edits_path, post(images_edits))
.layer(middleware::from_fn(smart_json_error_middleware))
.layer(axum::extract::DefaultBodyLimit::max(get_body_limit()))
.with_state(state);
(vec![doc, edits_doc], router)
}
async fn videos(
State(state): State<Arc<service_v2::State>>,
headers: HeaderMap,
Json(request): Json<NvCreateVideoRequest>,
) -> Result<Response, ErrorResponse> {
check_ready(&state)?;
check_model_serving_ready(&state, &request.model)?;
let request_id = get_or_create_request_id(&headers);
let request = context_from_headers(request, request_id, &headers)?;
let request_id = request.id().to_string();
let streaming = request.stream.unwrap_or(false);
let model = request.model.clone();
let metric_model = state.manager().metric_model_for(&model).to_string();
let http_queue_guard = state.metrics_clone().create_http_queue_guard(&metric_model);
let engine = state
.manager()
.get_videos_engine(&model)
.map_err(|e| ErrorMessage::from_model_error(&e))?;
let mut inflight = state.metrics_clone().create_inflight_guard(
&model,
Endpoint::Videos,
streaming,
&request_id,
);
let mut response_collector = state.metrics_clone().create_response_collector(&model);
let stream = engine.generate(request).await.map_err(|e| {
if super::metrics::request_was_rejected(e.as_ref()) {
state
.metrics_clone()
.inc_rejection(&model, super::metrics::Endpoint::Videos);
}
let err_response = ErrorMessage::from_anyhow(e, "Failed to generate videos");
inflight.mark_error(extract_error_type_from_response(&err_response));
err_response
})?;
let mut http_queue_guard = Some(http_queue_guard);
if streaming {
let ctx = stream.context();
let (mut connection_handle, stream_handle) = create_connection_monitor(
ctx.clone(),
Some(state.metrics_clone()),
CancellationLabels {
model: model.clone(),
endpoint: Endpoint::Videos.to_string(),
request_type: "stream".to_string(),
},
)
.await;
let stream = stream.flat_map(move |response| {
let sse_result = process_response_using_event_converter_and_observe_metrics(
EventConverter::from(response),
&mut response_collector,
&mut http_queue_guard,
);
match sse_result {
Ok(Some(ev)) => stream::iter(vec![Ok(ev)]),
Ok(None) => stream::iter(vec![]),
Err(e) => stream::iter(vec![Err(e)]),
}
});
let stream = monitor_for_disconnects(stream, ctx, inflight, stream_handle);
let mut sse_stream = Sse::new(stream);
if let Some(keep_alive) = state.sse_keep_alive() {
sse_stream = sse_stream.keep_alive(KeepAlive::default().interval(keep_alive));
}
connection_handle.disarm();
Ok(sse_stream.into_response())
} else {
let stream = stream.inspect(move |response| {
process_response_and_observe_metrics(
response,
&mut response_collector,
&mut http_queue_guard,
);
});
let response = NvVideosResponse::from_annotated_stream(stream)
.await
.map_err(|e| {
tracing::error!("Failed to fold videos stream for {}: {:?}", request_id, e);
let err_response =
ErrorMessage::internal_server_error("Failed to fold videos stream");
inflight.mark_error(extract_error_type_from_response(&err_response));
err_response
})?;
inflight.mark_ok();
Ok(Json(response).into_response())
}
}
async fn video_stream(
State(state): State<Arc<service_v2::State>>,
headers: HeaderMap,
Json(request): Json<NvCreateVideoRequest>,
) -> Result<Response, ErrorResponse> {
check_ready(&state)?;
check_model_serving_ready(&state, &request.model)?;
let request_id = get_or_create_request_id(&headers);
let request = context_from_headers(request, request_id, &headers)?;
let model = request.model.clone();
let metric_model = state.manager().metric_model_for(&model).to_string();
let http_queue_guard = state.metrics_clone().create_http_queue_guard(&metric_model);
let engine = state
.manager()
.get_videos_engine(&model)
.map_err(|e| ErrorMessage::from_model_error(&e))?;
let mut inflight =
state
.metrics_clone()
.create_inflight_guard(&model, Endpoint::Videos, true, request.id());
let mut response_collector = state.metrics_clone().create_response_collector(&model);
let stream = engine.generate(request).await.map_err(|e| {
if super::metrics::request_was_rejected(e.as_ref()) {
state
.metrics_clone()
.inc_rejection(&model, super::metrics::Endpoint::Videos);
}
let err_response = ErrorMessage::from_anyhow(e, "Failed to start video stream");
inflight.mark_error(extract_error_type_from_response(&err_response));
err_response
})?;
let ctx = stream.context();
let (mut connection_handle, mut stream_handle) = create_connection_monitor(
ctx.clone(),
Some(state.metrics_clone()),
CancellationLabels {
model: model.clone(),
endpoint: Endpoint::Videos.to_string(),
request_type: "stream".to_string(),
},
)
.await;
connection_handle.disarm();
let mut http_queue_guard = Some(http_queue_guard);
let stream = stream.inspect(move |response| {
process_response_and_observe_metrics(
response,
&mut response_collector,
&mut http_queue_guard,
);
});
let mjpeg_stream = stream.filter_map(|annotated| async move {
let ann = match annotated.ok() {
Ok(a) => a,
Err(e) => {
tracing::error!("Video stream error: {e}");
return None;
}
};
let response = ann.data?;
let frame = response.data.into_iter().next()?;
let b64 = frame.b64_json?;
let jpeg_bytes = match base64::prelude::BASE64_STANDARD.decode(&b64) {
Ok(b) => b,
Err(e) => {
tracing::warn!("Failed to decode frame base64: {e}");
return None;
}
};
let header = format!(
"--frame\r\nContent-Type: image/jpeg\r\nContent-Length: {}\r\n\r\n",
jpeg_bytes.len()
);
let mut chunk = Vec::with_capacity(header.len() + jpeg_bytes.len() + 2);
chunk.extend_from_slice(header.as_bytes());
chunk.extend_from_slice(&jpeg_bytes);
chunk.extend_from_slice(b"\r\n");
Some(Ok::<Bytes, std::convert::Infallible>(Bytes::from(chunk)))
});
stream_handle.arm();
let monitored_stream = async_stream::stream! {
tokio::pin!(mjpeg_stream);
loop {
tokio::select! {
frame = mjpeg_stream.next() => {
match frame {
Some(item) => yield item,
None => {
inflight.mark_ok();
stream_handle.disarm();
break;
}
}
}
_ = ctx.stopped() => {
tracing::trace!("Context stopped; breaking MJPEG stream");
inflight.mark_error(ErrorType::Cancelled);
break;
}
}
}
};
axum::http::Response::builder()
.status(axum::http::StatusCode::OK)
.header(
axum::http::header::CONTENT_TYPE,
"multipart/x-mixed-replace; boundary=frame",
)
.body(Body::from_stream(monitored_stream))
.map(|r| r.into_response())
.map_err(|e| {
ErrorMessage::internal_server_error_with_details(
"Failed to build MJPEG response",
format!("{e}"),
)
})
}
pub fn videos_router(
state: Arc<service_v2::State>,
path: Option<String>,
) -> (Vec<RouteDoc>, Router) {
let path = path.unwrap_or("/v1/videos".to_string());
let stream_path = format!("{}/stream", path);
let doc = RouteDoc::new(axum::http::Method::POST, &path);
let stream_doc = RouteDoc::new(axum::http::Method::POST, &stream_path);
let router = Router::new()
.route(&path, post(videos))
.route(&stream_path, post(video_stream))
.layer(middleware::from_fn(smart_json_error_middleware))
.layer(axum::extract::DefaultBodyLimit::max(get_body_limit()))
.with_state(state);
(vec![doc, stream_doc], router)
}
async fn audio_speech(
State(state): State<Arc<service_v2::State>>,
headers: HeaderMap,
Json(request): Json<NvCreateAudioSpeechRequest>,
) -> Result<Response, ErrorResponse> {
check_ready(&state)?;
let request_id = get_or_create_request_id(&headers);
let request = context_from_headers(request, request_id, &headers)?;
let request_id = request.id().to_string();
let streaming = false;
let model = request.model.clone().unwrap_or_else(|| {
state
.manager()
.serving_ready_display_names()
.into_iter()
.next()
.unwrap_or_default()
});
let metric_model = state.manager().metric_model_for(&model).to_string();
check_model_serving_ready(&state, &model)?;
let http_queue_guard = state.metrics_clone().create_http_queue_guard(&metric_model);
let engine = state
.manager()
.get_audios_engine(&model)
.map_err(|e| ErrorMessage::from_model_error(&e))?;
let mut inflight = state.metrics_clone().create_inflight_guard(
&model,
Endpoint::Audios,
streaming,
&request_id,
);
let mut response_collector = state.metrics_clone().create_response_collector(&model);
let stream = engine
.generate(request)
.await
.map_err(|e| ErrorMessage::from_anyhow(e, "Failed to generate audio"))?;
let mut http_queue_guard = Some(http_queue_guard);
let stream = stream.inspect(move |response| {
process_response_and_observe_metrics(
response,
&mut response_collector,
&mut http_queue_guard,
);
});
let response = NvAudioSpeechResponse::from_annotated_stream(stream)
.await
.map_err(|e| {
tracing::error!("Failed to fold audio stream for {}: {:?}", request_id, e);
ErrorMessage::internal_server_error("Failed to fold audio stream")
})?;
if response.status == "failed" {
return Ok((axum::http::StatusCode::BAD_REQUEST, Json(response)).into_response());
}
inflight.mark_ok();
if let Some(first) = response.data.first()
&& let Some(b64) = &first.b64_json
&& let Ok(audio_bytes) = base64::engine::general_purpose::STANDARD.decode(b64)
{
let content_type = match first.output_format.as_str() {
"mp3" => "audio/mpeg",
"flac" => "audio/flac",
"pcm" => "audio/pcm",
"aac" => "audio/aac",
"opus" => "audio/ogg; codecs=opus",
_ => "audio/wav",
};
return Ok(Response::builder()
.header("content-type", content_type)
.body(axum::body::Body::from(audio_bytes))
.unwrap());
}
Ok(Json(response).into_response())
}
pub fn audios_router(
state: Arc<service_v2::State>,
path: Option<String>,
) -> (Vec<RouteDoc>, Router) {
let path = path.unwrap_or("/v1/audio/speech".to_string());
let doc = RouteDoc::new(axum::http::Method::POST, &path);
let router = Router::new()
.route(&path, post(audio_speech))
.layer(middleware::from_fn(smart_json_error_middleware))
.layer(axum::extract::DefaultBodyLimit::max(get_body_limit()))
.with_state(state);
(vec![doc], router)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::discovery::ModelManagerError;
use crate::protocols::common::extensions::NvExt;
use crate::protocols::openai::chat_completions::NvCreateChatCompletionRequest;
use crate::protocols::openai::common_ext::CommonExt;
use crate::protocols::openai::completions::NvCreateCompletionRequest;
use crate::protocols::openai::responses::NvCreateResponse;
use dynamo_protocols::types::responses::{CreateResponse, Input, PromptConfig};
use dynamo_protocols::types::{
ChatCompletionRequestMessage, ChatCompletionRequestUserMessage,
ChatCompletionRequestUserMessageContent, CreateChatCompletionRequest,
CreateCompletionRequest, Prompt,
};
const BACKUP_ERROR_MESSAGE: &str = "Failed to generate completions";
#[test]
fn test_is_json_content_type() {
assert!(is_json_content_type("application/json"));
assert!(is_json_content_type("application/json; charset=utf-8"));
assert!(is_json_content_type("Application/JSON"));
assert!(is_json_content_type("application/vnd.dynamo+json"));
assert!(!is_json_content_type("text/plain"));
assert!(!is_json_content_type("application/json-patch"));
assert!(!is_json_content_type("application"));
}
#[test]
fn test_ensure_json_content_type_rejects_missing_or_non_json() {
let headers = HeaderMap::new();
let err = ensure_json_content_type(&headers).expect_err("missing content type should fail");
assert_eq!(err.0, StatusCode::UNSUPPORTED_MEDIA_TYPE);
let mut headers = HeaderMap::new();
headers.insert(
axum::http::header::CONTENT_TYPE,
"text/plain".parse().unwrap(),
);
let err =
ensure_json_content_type(&headers).expect_err("non-json content type should fail");
assert_eq!(err.0, StatusCode::UNSUPPORTED_MEDIA_TYPE);
}
#[test]
fn test_parse_chat_completion_request_escapes_control_chars_in_strings() {
let body = b"{\"model\":\"test-model\",\"messages\":[{\"role\":\"user\",\"content\":\"log \x1b[33mPK\x03\x04\"}]}";
let request: NvCreateChatCompletionRequest =
parse_json_request("chat completions", body).expect("request should parse");
let message = request
.inner
.messages
.first()
.expect("message should exist");
let ChatCompletionRequestMessage::User(user_message) = message else {
panic!("expected user message");
};
let ChatCompletionRequestUserMessageContent::Text(content) = &user_message.content else {
panic!("expected text content");
};
assert_eq!(content, "log \u{1b}[33mPK\u{3}\u{4}");
}
#[test]
fn test_parse_chat_completion_request_replaces_invalid_utf8_in_strings() {
let body = b"{\"model\":\"test-model\",\"messages\":[{\"role\":\"user\",\"content\":\"raw \xff data\"}]}";
let request: NvCreateChatCompletionRequest =
parse_json_request("chat completions", body).expect("request should parse");
let message = request
.inner
.messages
.first()
.expect("message should exist");
let ChatCompletionRequestMessage::User(user_message) = message else {
panic!("expected user message");
};
let ChatCompletionRequestUserMessageContent::Text(content) = &user_message.content else {
panic!("expected text content");
};
assert_eq!(content, "raw \u{fffd} data");
}
#[test]
fn test_parse_chat_completion_request_escapes_control_char_after_backslash() {
let body = b"{\"model\":\"test-model\",\"messages\":[{\"role\":\"user\",\"content\":\"slash \\\nnext\"}]}";
let request: NvCreateChatCompletionRequest =
parse_json_request("chat completions", body).expect("request should parse");
let message = request
.inner
.messages
.first()
.expect("message should exist");
let ChatCompletionRequestMessage::User(user_message) = message else {
panic!("expected user message");
};
let ChatCompletionRequestUserMessageContent::Text(content) = &user_message.content else {
panic!("expected text content");
};
assert_eq!(content, "slash \\\nnext");
}
#[test]
fn test_parse_chat_completion_request_keeps_schema_errors() {
let body = br#"{"model":"test-model","messages":[{"role":"assistant","content":[{"type":"thinking","thinking":"working"}]}]}"#;
let err =
match parse_json_request::<NvCreateChatCompletionRequest>("chat completions", body) {
Ok(_) => panic!("schema should still fail"),
Err(err) => err,
};
assert_eq!(err.0, StatusCode::BAD_REQUEST);
assert!(
err.1
.message
.contains("ChatCompletionRequestAssistantMessageContent"),
"unexpected error: {}",
err.1.message
);
}
#[test]
fn test_parse_completion_request_escapes_control_chars_in_prompt() {
let body =
b"{\"model\":\"test-model\",\"prompt\":\"log \x1b[33mPK\x03\x04\",\"max_tokens\":1}";
let request: NvCreateCompletionRequest =
parse_json_request("completions", body).expect("request should parse");
let Prompt::String(prompt) = &request.inner.prompt else {
panic!("expected string prompt");
};
assert_eq!(prompt, "log \u{1b}[33mPK\u{3}\u{4}");
}
#[test]
fn test_parse_completion_request_replaces_invalid_utf8_in_prompt() {
let body = b"{\"model\":\"test-model\",\"prompt\":\"raw \xff data\",\"max_tokens\":1}";
let request: NvCreateCompletionRequest =
parse_json_request("completions", body).expect("request should parse");
let Prompt::String(prompt) = &request.inner.prompt else {
panic!("expected string prompt");
};
assert_eq!(prompt, "raw \u{fffd} data");
}
fn http_error_from_engine(code: u16) -> Result<(), anyhow::Error> {
Err(HttpError {
code,
message: "custom error message".to_string(),
})?
}
fn other_error_from_engine() -> Result<(), anyhow::Error> {
Err(ModelManagerError::ModelNotFound("foo".to_string()))?
}
fn make_base_request() -> NvCreateResponse {
NvCreateResponse {
inner: CreateResponse {
input: Input::Text("hello".into()),
model: Some("test-model".into()),
..Default::default()
},
nvext: None,
}
}
#[test]
fn test_openai_nvext_rejects_agent_context() {
let err = serde_json::from_value::<NvExt>(serde_json::json!({
"agent_context": {
"session_id": "run-123"
}
}))
.unwrap_err();
assert!(err.to_string().contains("unknown field `agent_context`"));
}
#[test]
fn test_copy_context_metadata_preserves_agent_context() {
let mut source = Context::new(());
source.insert(
AGENT_CONTEXT_CONTEXT_KEY,
AgentContext {
session_id: "session-123".to_string(),
parent_session_id: Some("parent-456".to_string()),
session_final: Some(true),
kv_hints: None,
},
);
let mut target = Context::new(());
copy_context_metadata(&source, &mut target);
let agent_context = target
.get::<AgentContext>(AGENT_CONTEXT_CONTEXT_KEY)
.expect("agent context copied");
assert_eq!(agent_context.session_id, "session-123");
assert_eq!(
agent_context.parent_session_id.as_deref(),
Some("parent-456")
);
assert_eq!(agent_context.session_final, Some(true));
}
#[test]
fn test_context_metadata_preserves_session_affinity() {
let mut headers = HeaderMap::new();
headers.insert("x-dynamo-session-id", "session-123".parse().unwrap());
let source = context_from_headers((), "request-1".to_string(), &headers).unwrap();
let affinity = source
.get::<SessionAffinityId>(SESSION_AFFINITY_CONTEXT_KEY)
.expect("session affinity attached");
assert_eq!(affinity.as_str(), "session-123");
let mut target = Context::new(());
copy_context_metadata(&source, &mut target);
let affinity = target
.get::<SessionAffinityId>(SESSION_AFFINITY_CONTEXT_KEY)
.expect("session affinity copied");
assert_eq!(affinity.as_str(), "session-123");
}
#[test]
fn test_http_error_response_from_anyhow() {
let err = http_error_from_engine(400).unwrap_err();
let response = ErrorMessage::from_anyhow(err, BACKUP_ERROR_MESSAGE);
assert_eq!(response.0, StatusCode::BAD_REQUEST);
assert_eq!(response.1.message, "custom error message");
}
#[test]
fn test_check_ready_rejects_draining_service() {
let service = service_v2::HttpService::builder().build().unwrap();
let state = service.state_clone();
assert!(check_ready(&state).is_ok());
state.start_draining();
let response = check_ready(&state).unwrap_err();
assert_eq!(response.0, StatusCode::SERVICE_UNAVAILABLE);
}
#[test]
fn test_error_response_from_anyhow_out_of_range() {
for code in [399u16, 500, 501] {
let err = http_error_from_engine(code).unwrap_err();
let response = ErrorMessage::from_anyhow(err, BACKUP_ERROR_MESSAGE);
assert_eq!(response.0, StatusCode::INTERNAL_SERVER_ERROR);
assert_eq!(response.1.message, "Internal server error");
assert!(
!response.1.message.contains("custom error message"),
"client response must not include the backend-supplied HttpError message"
);
}
}
#[test]
fn test_from_http_error_sanitizes_499_message() {
let err = HttpError {
code: 499,
message: "session abc-123 cancelled at /srv/queue.py:42".to_string(),
};
let response = ErrorMessage::from_http_error(err);
assert_eq!(response.0.as_u16(), 499);
assert_eq!(response.1.code, 499);
assert_eq!(response.1.message, "Request cancelled");
assert!(!response.1.message.contains("abc-123"));
assert!(!response.1.message.contains("/srv/queue.py"));
}
#[test]
fn test_other_error_response_from_anyhow() {
let err = other_error_from_engine().unwrap_err();
let leaked_chain = format!("{err:#}");
let response = ErrorMessage::from_anyhow(err, BACKUP_ERROR_MESSAGE);
assert_eq!(response.0, StatusCode::INTERNAL_SERVER_ERROR);
assert_eq!(response.1.message, BACKUP_ERROR_MESSAGE);
assert!(
!response.1.message.contains(&leaked_chain),
"client response must not contain the anyhow error chain"
);
}
#[test]
fn test_resource_exhausted_error_response_from_anyhow() {
use dynamo_runtime::error::{DynamoError, ErrorType};
use dynamo_runtime::pipeline::error::PipelineError;
let cause = PipelineError::ServiceOverloaded(
"All workers are busy, please retry later".to_string(),
);
let err: anyhow::Error = DynamoError::builder()
.error_type(ErrorType::ResourceExhausted)
.message("All workers are busy, please retry later")
.cause(cause)
.build()
.into();
let response = ErrorMessage::from_anyhow(err, BACKUP_ERROR_MESSAGE);
assert_eq!(response.0.as_u16(), 529);
assert_eq!(response.1.code, 529);
assert_eq!(response.1.error_type, "Overloaded");
assert_eq!(response.1.message, "Service temporarily overloaded");
assert!(
!response.1.message.contains("All workers are busy"),
"client response must not include the underlying engine message"
);
}
#[test]
fn unavailable_error_response_from_anyhow() {
use dynamo_runtime::error::{DynamoError, ErrorType};
let err: anyhow::Error = DynamoError::builder()
.error_type(ErrorType::Unavailable)
.message("No workers available for endpoint test/worker/generate")
.build()
.into();
let response = ErrorMessage::from_anyhow(err, BACKUP_ERROR_MESSAGE);
assert_eq!(response.0, StatusCode::SERVICE_UNAVAILABLE);
assert_eq!(response.1.code, StatusCode::SERVICE_UNAVAILABLE.as_u16());
assert_eq!(response.1.message, "Service temporarily unavailable");
}
#[test]
fn queue_rejection_maps_to_structured_http_529() {
use dynamo_kv_router::scheduling::{QueueLimitKind, QueueRejection};
let rejection = QueueRejection {
policy_class: "latency".to_string(),
limit_kind: QueueLimitKind::CachedTokens,
current: 2048,
limit: 1024,
};
let response =
ErrorMessage::from_anyhow(anyhow::Error::new(rejection), BACKUP_ERROR_MESSAGE);
assert_eq!(response.0.as_u16(), 529);
assert_eq!(response.1.code, 529);
assert_eq!(response.1.error_type, "Overloaded");
assert_eq!(
response.1.details.as_deref(),
Some(&serde_json::json!({
"policy_class": "latency",
"limit_kind": "cached_tokens",
"current": 2048,
"limit": 1024,
}))
);
}
#[test]
fn test_nested_invalid_argument_response_from_anyhow() {
use dynamo_runtime::error::{DynamoError, ErrorType};
#[derive(Debug)]
struct WrappedError {
source: DynamoError,
}
impl std::fmt::Display for WrappedError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "outer routing failure")
}
}
impl std::error::Error for WrappedError {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
Some(&self.source)
}
}
let source = DynamoError::builder()
.error_type(ErrorType::InvalidArgument)
.message(
"Request payload is too large for this deployment. Reduce the input size or metadata size and retry.",
)
.build();
let err: anyhow::Error = WrappedError { source }.into();
let response = ErrorMessage::from_anyhow(err, BACKUP_ERROR_MESSAGE);
assert_eq!(response.0, StatusCode::BAD_REQUEST);
assert_eq!(response.1.code, StatusCode::BAD_REQUEST.as_u16());
assert!(response.1.message.contains("Request payload is too large"));
assert!(!response.1.message.contains("NATS"));
assert!(!response.1.message.contains("payload_bytes"));
}
#[test]
fn test_backend_invalid_argument_surfaces_as_400() {
use dynamo_runtime::error::{BackendError, DynamoError, ErrorType};
let err: anyhow::Error = DynamoError::builder()
.error_type(ErrorType::Backend(BackendError::InvalidArgument))
.message("Dynamo's SGLang backend does not currently support logprobs >= 1")
.build()
.into();
let response = ErrorMessage::from_anyhow(err, BACKUP_ERROR_MESSAGE);
assert_eq!(response.0, StatusCode::BAD_REQUEST);
assert_eq!(response.1.code, StatusCode::BAD_REQUEST.as_u16());
assert!(response.1.message.contains("does not currently support"));
}
#[test]
fn test_cancelled_error_response_from_anyhow() {
use dynamo_runtime::error::{DynamoError, ErrorType};
let err: anyhow::Error = DynamoError::builder()
.error_type(ErrorType::Cancelled)
.message("Context id abc-123 is stopped or killed")
.build()
.into();
let response = ErrorMessage::from_anyhow(err, BACKUP_ERROR_MESSAGE);
assert_eq!(
response.0.as_u16(),
499,
"Cancelled errors should return HTTP 499"
);
assert_eq!(response.1.code, 499);
assert_eq!(response.1.error_type, "Client Closed Request");
assert_eq!(response.1.message, "Request cancelled");
assert!(!response.1.message.contains("abc-123"));
assert!(!response.1.message.contains("stopped or killed"));
}
#[test]
fn test_cancelled_error_metrics_classification() {
let error_type =
classify_error_for_metrics(StatusCode::from_u16(499).unwrap(), "cancelled request");
assert_eq!(
error_type,
ErrorType::Cancelled,
"HTTP 499 should map to ErrorType::Cancelled in metrics"
);
}
#[test]
fn test_validate_unsupported_fields_accepts_clean_request() {
let request = make_base_request();
let result = validate_response_unsupported_fields(&request);
assert!(result.is_none());
}
#[test]
fn test_validate_unsupported_fields_accepts_parallel_tool_calls() {
let mut request = make_base_request();
request.inner.parallel_tool_calls = Some(true);
let result = validate_response_unsupported_fields(&request);
assert!(result.is_none(), "parallel_tool_calls should be supported");
}
#[test]
fn test_validate_unsupported_fields_accepts_store() {
let mut request = make_base_request();
request.inner.store = Some(true);
let result = validate_response_unsupported_fields(&request);
assert!(
result.is_none(),
"store should be supported for audit opt-in"
);
}
#[tokio::test]
async fn test_validate_unsupported_fields_rejects_rl_nvext_fields() {
for field in ["completion_token_ids", "prompt_logprobs"] {
for stream in [false, true] {
let mut request = make_base_request();
request.inner.stream = Some(stream);
request.nvext = Some(
NvExt::builder()
.extra_fields(vec![field.to_string()])
.build()
.unwrap(),
);
let response = validate_response_unsupported_fields(&request)
.expect("RL nvext response field should be rejected")
.into_response();
assert_eq!(response.status(), StatusCode::NOT_IMPLEMENTED);
let body = axum::body::to_bytes(response.into_body(), get_body_limit())
.await
.unwrap();
let error: ErrorMessage = serde_json::from_slice(&body).unwrap();
assert_eq!(
error.message,
format!(
"{VALIDATION_PREFIX}`nvext.extra_fields=[\"{field}\"]` is not supported by the Responses API."
)
);
}
}
}
#[test]
fn test_validate_unsupported_fields_rejects_mixed_nvext_fields() {
let mut request = make_base_request();
request.nvext = Some(
NvExt::builder()
.extra_fields(vec![
"timing".to_string(),
"completion_token_ids".to_string(),
])
.build()
.unwrap(),
);
assert!(validate_response_unsupported_fields(&request).is_some());
}
#[test]
fn test_validate_unsupported_fields_accepts_supported_nvext_fields() {
let mut request = make_base_request();
request.nvext = Some(
NvExt::builder()
.extra_fields(vec!["timing".to_string(), "worker_id".to_string()])
.build()
.unwrap(),
);
assert!(validate_response_unsupported_fields(&request).is_none());
}
#[test]
fn test_validate_unsupported_fields_detects_flags() {
#[allow(clippy::type_complexity)]
let unsupported_cases: Vec<(&str, Box<dyn FnOnce(&mut CreateResponse)>)> = vec![
("background", Box::new(|r| r.background = Some(true))),
(
"previous_response_id",
Box::new(|r| r.previous_response_id = Some("prev-id".into())),
),
(
"prompt",
Box::new(|r| {
r.prompt = Some(PromptConfig {
id: "template-id".into(),
version: None,
variables: None,
})
}),
),
("max_tool_calls", Box::new(|r| r.max_tool_calls = Some(5))),
];
for (field, set_field) in unsupported_cases {
let mut req = make_base_request();
(set_field)(&mut req.inner);
let result = validate_response_unsupported_fields(&req);
assert!(result.is_some(), "Expected rejection for `{field}`");
}
}
#[test]
fn test_validate_unsupported_fields_accepts_passthrough_metadata() {
#[allow(clippy::type_complexity)]
let passthrough_cases: Vec<(&str, Box<dyn FnOnce(&mut CreateResponse)>)> = vec![
(
"prompt_cache_key",
Box::new(|r| r.prompt_cache_key = Some("ck-1".into())),
),
(
"prompt_cache_retention",
Box::new(|r| {
r.prompt_cache_retention =
Some(dynamo_protocols::types::responses::PromptCacheRetention::InMemory)
}),
),
(
"safety_identifier",
Box::new(|r| r.safety_identifier = Some("user-hash".into())),
),
];
for (field, set_field) in passthrough_cases {
let mut req = make_base_request();
(set_field)(&mut req.inner);
let result = validate_response_unsupported_fields(&req);
assert!(
result.is_none(),
"Expected `{field}` to be accepted as pass-through metadata"
);
}
}
#[test]
fn test_validate_chat_completion_required_fields_empty_messages() {
let request = NvCreateChatCompletionRequest {
inner: CreateChatCompletionRequest {
model: "test-model".to_string(),
messages: vec![],
..Default::default()
},
common: Default::default(),
nvext: None,
chat_template_args: None,
thinking: None,
media_io_kwargs: None,
return_tokens_as_token_ids: None,
unsupported_fields: Default::default(),
};
let result = validate_chat_completion_required_fields(&request);
assert!(result.is_err());
if let Err(error_response) = result {
assert_eq!(error_response.0, StatusCode::BAD_REQUEST);
assert_eq!(
error_response.1.message,
format!(
"{VALIDATION_PREFIX}The 'messages' field cannot be empty. At least one message is required."
)
);
}
}
#[test]
fn test_validate_chat_completion_required_fields_with_messages() {
let request = NvCreateChatCompletionRequest {
inner: CreateChatCompletionRequest {
model: "test-model".to_string(),
messages: vec![ChatCompletionRequestMessage::User(
ChatCompletionRequestUserMessage {
content: ChatCompletionRequestUserMessageContent::Text("Hello".to_string()),
name: None,
},
)],
..Default::default()
},
common: Default::default(),
nvext: None,
chat_template_args: None,
thinking: None,
media_io_kwargs: None,
return_tokens_as_token_ids: None,
unsupported_fields: Default::default(),
};
let result = validate_chat_completion_required_fields(&request);
assert!(result.is_ok());
}
#[test]
fn test_normalize_chat_reasoning_template_args_error_response() {
let mut request: NvCreateChatCompletionRequest =
serde_json::from_value(serde_json::json!({
"model": "test-model",
"messages": [{"role": "user", "content": "Hello"}],
"thinking": {"type": "auto"}
}))
.unwrap();
let result = normalize_chat_reasoning_template_args(&mut request);
assert!(result.is_err());
if let Err(error_response) = result {
assert_eq!(error_response.0, StatusCode::BAD_REQUEST);
assert_eq!(
error_response.1.message,
format!(
"{VALIDATION_PREFIX}`thinking.type` must be `enabled`, `disabled`, or `adaptive`"
)
);
}
}
#[test]
fn test_bad_base_request_for_completion() {
let request = NvCreateCompletionRequest {
inner: CreateCompletionRequest {
model: "test-model".to_string(),
prompt: "Hello".into(),
frequency_penalty: Some(-3.0),
..Default::default()
},
common: Default::default(),
nvext: None,
metadata: None,
return_tokens_as_token_ids: None,
unsupported_fields: Default::default(),
};
let result = validate_completion_fields_generic(&request);
assert!(result.is_err());
if let Err(error_response) = result {
assert_eq!(error_response.0, StatusCode::BAD_REQUEST);
assert_eq!(
error_response.1.message,
format!("{VALIDATION_PREFIX}Frequency penalty must be between -2 and 2, got -3")
);
}
let request = NvCreateCompletionRequest {
inner: CreateCompletionRequest {
model: "test-model".to_string(),
prompt: "Hello".into(),
presence_penalty: Some(-3.0),
..Default::default()
},
common: Default::default(),
nvext: None,
metadata: None,
return_tokens_as_token_ids: None,
unsupported_fields: Default::default(),
};
let result = validate_completion_fields_generic(&request);
assert!(result.is_err());
if let Err(error_response) = result {
assert_eq!(error_response.0, StatusCode::BAD_REQUEST);
assert_eq!(
error_response.1.message,
format!("{VALIDATION_PREFIX}Presence penalty must be between -2 and 2, got -3")
);
}
let request = NvCreateCompletionRequest {
inner: CreateCompletionRequest {
model: "test-model".to_string(),
prompt: "Hello".into(),
temperature: Some(-3.0),
..Default::default()
},
common: Default::default(),
nvext: None,
metadata: None,
return_tokens_as_token_ids: None,
unsupported_fields: Default::default(),
};
let result = validate_completion_fields_generic(&request);
assert!(result.is_err());
if let Err(error_response) = result {
assert_eq!(error_response.0, StatusCode::BAD_REQUEST);
assert_eq!(
error_response.1.message,
format!("{VALIDATION_PREFIX}Temperature must be between 0 and 2, got -3")
);
}
let request = NvCreateCompletionRequest {
inner: CreateCompletionRequest {
model: "test-model".to_string(),
prompt: "Hello".into(),
top_p: Some(-3.0),
..Default::default()
},
common: Default::default(),
nvext: None,
metadata: None,
return_tokens_as_token_ids: None,
unsupported_fields: Default::default(),
};
let result = validate_completion_fields_generic(&request);
assert!(result.is_err());
if let Err(error_response) = result {
assert_eq!(error_response.0, StatusCode::BAD_REQUEST);
assert_eq!(
error_response.1.message,
format!("{VALIDATION_PREFIX}Top_p must be between 0 and 1, got -3")
);
}
let request = NvCreateCompletionRequest {
inner: CreateCompletionRequest {
model: "test-model".to_string(),
prompt: "Hello".into(),
..Default::default()
},
common: CommonExt::builder()
.repetition_penalty(-3.0)
.build()
.unwrap(),
nvext: None,
metadata: None,
return_tokens_as_token_ids: None,
unsupported_fields: Default::default(),
};
let result = validate_completion_fields_generic(&request);
assert!(result.is_err());
if let Err(error_response) = result {
assert_eq!(error_response.0, StatusCode::BAD_REQUEST);
assert_eq!(
error_response.1.message,
format!("{VALIDATION_PREFIX}Repetition penalty must be between 0 and 2, got -3")
);
}
let request = NvCreateCompletionRequest {
inner: CreateCompletionRequest {
model: "test-model".to_string(),
prompt: "Hello".into(),
logprobs: Some(6),
..Default::default()
},
common: Default::default(),
nvext: None,
metadata: None,
return_tokens_as_token_ids: None,
unsupported_fields: Default::default(),
};
let result = validate_completion_fields_generic(&request);
assert!(result.is_err());
if let Err(error_response) = result {
assert_eq!(error_response.0, StatusCode::BAD_REQUEST);
assert_eq!(
error_response.1.message,
format!("{VALIDATION_PREFIX}Logprobs must be between 0 and 5, got 6")
);
}
}
#[test]
fn test_metadata_field_nested() {
use serde_json::json;
let request = NvCreateCompletionRequest {
inner: CreateCompletionRequest {
model: "test-model".to_string(),
prompt: "Hello".into(),
..Default::default()
},
common: Default::default(),
nvext: None,
metadata: json!({
"user": {"id": 1, "name": "user-1"},
"session": {"id": "session-1", "timestamp": 1640995200}
})
.into(),
return_tokens_as_token_ids: None,
unsupported_fields: Default::default(),
};
let result = validate_completion_fields_generic(&request);
assert!(result.is_ok());
assert!(request.metadata.is_some());
assert_eq!(request.metadata.as_ref().unwrap()["user"]["id"], 1);
}
#[test]
fn test_bad_base_request_for_chatcompletion() {
let request = NvCreateChatCompletionRequest {
inner: CreateChatCompletionRequest {
model: "test-model".to_string(),
messages: vec![ChatCompletionRequestMessage::User(
ChatCompletionRequestUserMessage {
content: ChatCompletionRequestUserMessageContent::Text("Hello".to_string()),
name: None,
},
)],
frequency_penalty: Some(-3.0),
..Default::default()
},
common: Default::default(),
nvext: None,
chat_template_args: None,
thinking: None,
media_io_kwargs: None,
return_tokens_as_token_ids: None,
unsupported_fields: Default::default(),
};
let result = validate_chat_completion_fields_generic(&request);
assert!(result.is_err());
if let Err(error_response) = result {
assert_eq!(error_response.0, StatusCode::BAD_REQUEST);
assert_eq!(
error_response.1.message,
format!("{VALIDATION_PREFIX}Frequency penalty must be between -2 and 2, got -3")
);
}
let request = NvCreateChatCompletionRequest {
inner: CreateChatCompletionRequest {
model: "test-model".to_string(),
messages: vec![ChatCompletionRequestMessage::User(
ChatCompletionRequestUserMessage {
content: ChatCompletionRequestUserMessageContent::Text("Hello".to_string()),
name: None,
},
)],
presence_penalty: Some(-3.0),
..Default::default()
},
common: Default::default(),
nvext: None,
chat_template_args: None,
thinking: None,
media_io_kwargs: None,
return_tokens_as_token_ids: None,
unsupported_fields: Default::default(),
};
let result = validate_chat_completion_fields_generic(&request);
assert!(result.is_err());
if let Err(error_response) = result {
assert_eq!(error_response.0, StatusCode::BAD_REQUEST);
assert_eq!(
error_response.1.message,
format!("{VALIDATION_PREFIX}Presence penalty must be between -2 and 2, got -3")
);
}
let request = NvCreateChatCompletionRequest {
inner: CreateChatCompletionRequest {
model: "test-model".to_string(),
messages: vec![ChatCompletionRequestMessage::User(
ChatCompletionRequestUserMessage {
content: ChatCompletionRequestUserMessageContent::Text("Hello".to_string()),
name: None,
},
)],
temperature: Some(-3.0),
..Default::default()
},
common: Default::default(),
nvext: None,
chat_template_args: None,
thinking: None,
media_io_kwargs: None,
return_tokens_as_token_ids: None,
unsupported_fields: Default::default(),
};
let result = validate_chat_completion_fields_generic(&request);
assert!(result.is_err());
if let Err(error_response) = result {
assert_eq!(error_response.0, StatusCode::BAD_REQUEST);
assert_eq!(
error_response.1.message,
format!("{VALIDATION_PREFIX}Temperature must be between 0 and 2, got -3")
);
}
let request = NvCreateChatCompletionRequest {
inner: CreateChatCompletionRequest {
model: "test-model".to_string(),
messages: vec![ChatCompletionRequestMessage::User(
ChatCompletionRequestUserMessage {
content: ChatCompletionRequestUserMessageContent::Text("Hello".to_string()),
name: None,
},
)],
top_p: Some(-3.0),
..Default::default()
},
common: Default::default(),
nvext: None,
chat_template_args: None,
thinking: None,
media_io_kwargs: None,
return_tokens_as_token_ids: None,
unsupported_fields: Default::default(),
};
let result = validate_chat_completion_fields_generic(&request);
assert!(result.is_err());
if let Err(error_response) = result {
assert_eq!(error_response.0, StatusCode::BAD_REQUEST);
assert_eq!(
error_response.1.message,
format!("{VALIDATION_PREFIX}Top_p must be between 0 and 1, got -3")
);
}
let request = NvCreateChatCompletionRequest {
inner: CreateChatCompletionRequest {
model: "test-model".to_string(),
messages: vec![ChatCompletionRequestMessage::User(
ChatCompletionRequestUserMessage {
content: ChatCompletionRequestUserMessageContent::Text("Hello".to_string()),
name: None,
},
)],
..Default::default()
},
common: CommonExt::builder()
.repetition_penalty(-3.0)
.build()
.unwrap(),
nvext: None,
chat_template_args: None,
thinking: None,
media_io_kwargs: None,
return_tokens_as_token_ids: None,
unsupported_fields: Default::default(),
};
let result = validate_chat_completion_fields_generic(&request);
assert!(result.is_err());
if let Err(error_response) = result {
assert_eq!(error_response.0, StatusCode::BAD_REQUEST);
assert_eq!(
error_response.1.message,
format!("{VALIDATION_PREFIX}Repetition penalty must be between 0 and 2, got -3")
);
}
let request = NvCreateChatCompletionRequest {
inner: CreateChatCompletionRequest {
model: "test-model".to_string(),
messages: vec![ChatCompletionRequestMessage::User(
ChatCompletionRequestUserMessage {
content: ChatCompletionRequestUserMessageContent::Text("Hello".to_string()),
name: None,
},
)],
top_logprobs: Some(25),
..Default::default()
},
common: Default::default(),
nvext: None,
chat_template_args: None,
thinking: None,
media_io_kwargs: None,
return_tokens_as_token_ids: None,
unsupported_fields: Default::default(),
};
let result = validate_chat_completion_fields_generic(&request);
assert!(result.is_err());
if let Err(error_response) = result {
assert_eq!(error_response.0, StatusCode::BAD_REQUEST);
assert_eq!(
error_response.1.message,
format!("{VALIDATION_PREFIX}Top_logprobs must be between 0 and 20, got 25")
);
}
}
#[test]
fn test_chat_completions_unknown_fields_rejected() {
let json = r#"{
"messages": [{"role": "user", "content": "Hello"}],
"model": "test-model",
"add_special_tokens": true,
"documents": ["doc1"],
"chat_template": "custom"
}"#;
let request: NvCreateChatCompletionRequest = serde_json::from_str(json).unwrap();
assert!(
request
.unsupported_fields
.contains_key("add_special_tokens")
);
assert!(request.unsupported_fields.contains_key("documents"));
assert!(request.unsupported_fields.contains_key("chat_template"));
let result = validate_chat_completion_fields_generic(&request);
assert!(result.is_err());
if let Err(error_response) = result {
assert_eq!(error_response.0, StatusCode::BAD_REQUEST);
let msg = &error_response.1.message;
assert!(msg.contains("Unsupported parameter"));
assert!(msg.contains("add_special_tokens"));
assert!(msg.contains("documents"));
assert!(msg.contains("chat_template"));
}
}
#[test]
fn test_completions_unsupported_fields_rejected() {
let json = r#"{
"model": "test-model",
"prompt": "Hello",
"add_special_tokens": true,
"response_format": {"type": "json_object"}
}"#;
let request: NvCreateCompletionRequest = serde_json::from_str(json).unwrap();
assert!(
request
.unsupported_fields
.contains_key("add_special_tokens")
);
assert!(request.unsupported_fields.contains_key("response_format"));
let result = validate_completion_fields_generic(&request);
assert!(result.is_err());
if let Err(error_response) = result {
assert_eq!(error_response.0, StatusCode::BAD_REQUEST);
let msg = &error_response.1.message;
assert!(msg.contains("Unsupported parameter"));
assert!(msg.contains("add_special_tokens"));
assert!(msg.contains("response_format"));
}
}
#[tokio::test]
async fn test_check_for_backend_error_with_error_event() {
use crate::types::openai::chat_completions::NvCreateChatCompletionStreamResponse;
use futures::stream;
let error_event = Annotated::<NvCreateChatCompletionStreamResponse> {
data: None,
id: None,
event: Some("error".to_string()),
comment: Some(vec!["Backend service unavailable".to_string()]),
error: None,
};
let test_stream = stream::iter(vec![error_event]);
let result = check_for_backend_error(test_stream).await;
assert!(result.is_err());
if let Err(error_response) = result {
assert_eq!(error_response.0, StatusCode::INTERNAL_SERVER_ERROR);
assert_eq!(error_response.1.message, "Internal server error");
assert!(
!error_response
.1
.message
.contains("Backend service unavailable")
);
}
}
#[tokio::test]
async fn test_check_for_backend_error_with_typed_invalid_argument() {
use crate::types::openai::chat_completions::NvCreateChatCompletionStreamResponse;
use dynamo_runtime::error::{BackendError, DynamoError, ErrorType};
use futures::stream;
for error_type in [
ErrorType::InvalidArgument,
ErrorType::Backend(BackendError::InvalidArgument),
] {
let error_event = Annotated::<NvCreateChatCompletionStreamResponse> {
data: None,
id: None,
event: Some("error".to_string()),
comment: None,
error: Some(
DynamoError::builder()
.error_type(error_type)
.message("unsupported JSON schema keyword")
.build(),
),
};
let result = check_for_backend_error(stream::iter(vec![error_event])).await;
let error_response = match result {
Err(error_response) => error_response,
Ok(_) => panic!("typed invalid argument must fail"),
};
assert_eq!(error_response.0, StatusCode::BAD_REQUEST);
assert_eq!(error_response.1.code, StatusCode::BAD_REQUEST.as_u16());
assert_eq!(error_response.1.error_type, "Bad Request");
assert_eq!(error_response.1.message, "unsupported JSON schema keyword");
}
}
#[tokio::test]
async fn test_check_for_backend_error_with_json_error_and_code() {
use crate::types::openai::chat_completions::NvCreateChatCompletionStreamResponse;
use futures::stream;
let error_json =
r#"{"message":"prompt > max_seq_len","type":"Internal Server Error","code":500}"#;
let error_event = Annotated::<NvCreateChatCompletionStreamResponse> {
data: None,
id: None,
event: Some("error".to_string()),
comment: Some(vec![error_json.to_string()]),
error: None,
};
let test_stream = stream::iter(vec![error_event]);
let result = check_for_backend_error(test_stream).await;
assert!(result.is_err());
if let Err(error_response) = result {
assert_eq!(error_response.0, StatusCode::INTERNAL_SERVER_ERROR);
assert_eq!(error_response.1.message, "Internal server error");
assert_eq!(error_response.1.code, 500);
assert!(!error_response.1.message.contains("prompt > max_seq_len"));
}
}
#[tokio::test]
async fn test_check_for_backend_error_with_non_client_error_code() {
use crate::types::openai::chat_completions::NvCreateChatCompletionStreamResponse;
use futures::stream;
let error_json =
r#"{"message":"panic at /srv/model.py:42","type":"Backend Error","code":399}"#;
let error_event = Annotated::<NvCreateChatCompletionStreamResponse> {
data: None,
id: None,
event: Some("error".to_string()),
comment: Some(vec![error_json.to_string()]),
error: None,
};
let test_stream = stream::iter(vec![error_event]);
let result = check_for_backend_error(test_stream).await;
assert!(result.is_err());
if let Err(error_response) = result {
assert_eq!(error_response.0, StatusCode::INTERNAL_SERVER_ERROR);
assert_eq!(error_response.1.code, 500);
assert_eq!(error_response.1.message, "Internal server error");
assert!(!error_response.1.message.contains("/srv/model.py"));
assert!(!error_response.1.message.contains("panic"));
}
}
#[tokio::test]
async fn test_check_for_backend_error_with_503_preserves_status() {
use crate::types::openai::chat_completions::NvCreateChatCompletionStreamResponse;
use futures::stream;
let error_json = r#"{"message":"engine pool exhausted at /srv/engine.py:88","code":503}"#;
let error_event = Annotated::<NvCreateChatCompletionStreamResponse> {
data: None,
id: None,
event: Some("error".to_string()),
comment: Some(vec![error_json.to_string()]),
error: None,
};
let test_stream = stream::iter(vec![error_event]);
let result = check_for_backend_error(test_stream).await;
assert!(result.is_err());
if let Err(error_response) = result {
assert_eq!(error_response.0, StatusCode::SERVICE_UNAVAILABLE);
assert_eq!(error_response.1.code, 503);
assert_eq!(error_response.1.message, "Internal server error");
assert!(!error_response.1.message.contains("engine pool"));
assert!(!error_response.1.message.contains("/srv/engine.py"));
}
}
#[tokio::test]
async fn test_check_for_backend_error_with_499_sanitizes_cancellation() {
use crate::types::openai::chat_completions::NvCreateChatCompletionStreamResponse;
use futures::stream;
let error_json =
r#"{"message":"Context id abc-123 cancelled at /srv/queue.py:42","code":499}"#;
let error_event = Annotated::<NvCreateChatCompletionStreamResponse> {
data: None,
id: None,
event: Some("error".to_string()),
comment: Some(vec![error_json.to_string()]),
error: None,
};
let test_stream = stream::iter(vec![error_event]);
let result = check_for_backend_error(test_stream).await;
assert!(result.is_err());
if let Err(error_response) = result {
assert_eq!(error_response.0.as_u16(), 499);
assert_eq!(error_response.1.code, 499);
assert_eq!(error_response.1.message, "Request cancelled");
assert!(!error_response.1.message.contains("abc-123"));
assert!(!error_response.1.message.contains("/srv/queue.py"));
}
}
#[tokio::test]
async fn test_check_for_backend_error_skips_leading_annotation_frames() {
use crate::types::openai::chat_completions::NvCreateChatCompletionStreamResponse;
use futures::stream;
let annotation = Annotated::<NvCreateChatCompletionStreamResponse>::from_annotation(
ANNOTATION_REQUEST_ID,
&"req-123".to_string(),
)
.expect("annotation construction should succeed");
let error_event = Annotated::<NvCreateChatCompletionStreamResponse> {
data: None,
id: None,
event: Some("error".to_string()),
comment: Some(vec![
r#"{"message":"bad input from client","code":400}"#.to_string(),
]),
error: None,
};
let test_stream = stream::iter(vec![annotation, error_event]);
let result = check_for_backend_error(test_stream).await;
assert!(
result.is_err(),
"annotation followed by an error event must still be detected as an error"
);
if let Err(error_response) = result {
assert_eq!(error_response.0, StatusCode::BAD_REQUEST);
assert_eq!(error_response.1.code, 400);
assert_eq!(error_response.1.message, "bad input from client");
}
}
#[tokio::test]
async fn test_check_for_backend_error_replays_leading_annotation_frames() {
use crate::types::openai::chat_completions::NvCreateChatCompletionStreamResponse;
use dynamo_protocols::types::CreateChatCompletionStreamResponse;
use futures::stream::{self, StreamExt};
let annotation = Annotated::<NvCreateChatCompletionStreamResponse>::from_annotation(
ANNOTATION_REQUEST_ID,
&"req-123".to_string(),
)
.expect("annotation construction should succeed");
let normal_event = Annotated::<NvCreateChatCompletionStreamResponse> {
data: Some(NvCreateChatCompletionStreamResponse {
inner: CreateChatCompletionStreamResponse {
id: "test-id".to_string(),
choices: vec![],
created: 0,
model: "test-model".to_string(),
system_fingerprint: None,
object: "chat.completion.chunk".to_string(),
service_tier: None,
usage: None,
},
nvext: None,
llm_metrics: None,
}),
id: Some("msg-1".to_string()),
event: None,
comment: None,
error: None,
};
let test_stream = stream::iter(vec![annotation, normal_event]);
let result = check_for_backend_error(test_stream).await;
assert!(result.is_ok());
let mut returned: Vec<_> = result.unwrap().collect().await;
assert_eq!(returned.len(), 2, "annotation + data event must replay");
let first = returned.remove(0);
assert_eq!(first.event.as_deref(), Some(ANNOTATION_REQUEST_ID));
let second = returned.remove(0);
assert_eq!(second.id, Some("msg-1".to_string()));
}
#[tokio::test]
async fn test_check_for_backend_error_with_normal_event() {
use crate::types::openai::chat_completions::NvCreateChatCompletionStreamResponse;
use dynamo_protocols::types::CreateChatCompletionStreamResponse;
use futures::stream::{self, StreamExt};
let normal_event = Annotated::<NvCreateChatCompletionStreamResponse> {
data: Some(NvCreateChatCompletionStreamResponse {
inner: CreateChatCompletionStreamResponse {
id: "test-id".to_string(),
choices: vec![],
created: 0,
model: "test-model".to_string(),
system_fingerprint: None,
object: "chat.completion.chunk".to_string(),
service_tier: None,
usage: None,
},
nvext: None,
llm_metrics: None,
}),
id: Some("msg-1".to_string()),
event: None,
comment: None,
error: None,
};
let test_stream = stream::iter(vec![normal_event.clone()]);
let result = check_for_backend_error(test_stream).await;
assert!(result.is_ok());
let mut returned_stream = result.unwrap();
let first = returned_stream.next().await;
assert!(first.is_some());
let first_event = first.unwrap();
assert_eq!(first_event.id, Some("msg-1".to_string()));
}
#[tokio::test]
async fn test_check_for_backend_error_with_empty_stream() {
use crate::types::openai::chat_completions::NvCreateChatCompletionStreamResponse;
use futures::stream::{self, StreamExt};
let test_stream =
stream::iter::<Vec<Annotated<NvCreateChatCompletionStreamResponse>>>(vec![]);
let result = check_for_backend_error(test_stream).await;
assert!(result.is_ok());
let mut returned_stream = result.unwrap();
let first = returned_stream.next().await;
assert!(first.is_none());
}
#[tokio::test]
async fn test_check_for_backend_error_with_comment_but_no_event_type() {
use crate::types::openai::chat_completions::NvCreateChatCompletionStreamResponse;
use futures::stream;
let error_event = Annotated::<NvCreateChatCompletionStreamResponse> {
data: None,
id: None,
event: None,
comment: Some(vec!["Connection timeout".to_string()]),
error: None,
};
let test_stream = stream::iter(vec![error_event]);
let result = check_for_backend_error(test_stream).await;
assert!(result.is_err());
if let Err(error_response) = result {
assert_eq!(error_response.0, StatusCode::INTERNAL_SERVER_ERROR);
assert_eq!(error_response.1.message, "Internal server error");
assert!(!error_response.1.message.contains("Connection timeout"));
}
}
#[test]
fn test_classify_error_for_metrics_validation() {
let error_type =
classify_error_for_metrics(StatusCode::BAD_REQUEST, "Validation: Invalid parameter");
assert_eq!(error_type, ErrorType::Validation);
let error_type = classify_error_for_metrics(StatusCode::BAD_REQUEST, "Some other error");
assert_eq!(error_type, ErrorType::Internal);
}
#[test]
fn test_classify_error_for_metrics_status_codes() {
assert_eq!(
classify_error_for_metrics(StatusCode::NOT_FOUND, "Model not found"),
ErrorType::NotFound
);
assert_eq!(
classify_error_for_metrics(StatusCode::NOT_IMPLEMENTED, "Feature not supported"),
ErrorType::NotImplemented
);
assert_eq!(
classify_error_for_metrics(StatusCode::TOO_MANY_REQUESTS, "Rate limit exceeded"),
ErrorType::Overload
);
assert_eq!(
classify_error_for_metrics(StatusCode::SERVICE_UNAVAILABLE, "Unavailable"),
ErrorType::Unavailable
);
assert_eq!(
classify_error_for_metrics(overload_status_code(), "Overloaded"),
ErrorType::Overload
);
assert_eq!(
classify_error_for_metrics(StatusCode::INTERNAL_SERVER_ERROR, "Panic"),
ErrorType::Internal
);
}
#[test]
fn test_classify_error_for_metrics_client_errors() {
assert_eq!(
classify_error_for_metrics(StatusCode::UNAUTHORIZED, "Unauthorized"),
ErrorType::Validation
);
assert_eq!(
classify_error_for_metrics(StatusCode::FORBIDDEN, "Forbidden"),
ErrorType::Validation
);
}
#[test]
fn test_extract_error_type_from_response_validation() {
let response = ErrorMessage::from_http_error(HttpError {
code: 400,
message: "Validation: bad input".to_string(),
});
assert_eq!(
extract_error_type_from_response(&response),
ErrorType::Validation
);
}
#[test]
fn test_extract_error_type_from_response_not_found() {
let response = ErrorMessage::model_not_found();
assert_eq!(
extract_error_type_from_response(&response),
ErrorType::NotFound
);
}
#[test]
fn test_extract_error_type_from_response_unavailable() {
let response =
ErrorMessage::from_model_error(&ModelManagerError::ModelUnavailable("x".to_string()));
assert_eq!(
extract_error_type_from_response(&response),
ErrorType::Unavailable
);
}
#[test]
fn test_from_model_error_maps_correctly() {
let not_found = ModelManagerError::ModelNotFound("x".to_string());
assert_eq!(
ErrorMessage::from_model_error(¬_found).0,
StatusCode::NOT_FOUND
);
let unavailable = ModelManagerError::ModelUnavailable("x".to_string());
assert_eq!(
ErrorMessage::from_model_error(&unavailable).0,
StatusCode::SERVICE_UNAVAILABLE
);
}
#[test]
fn test_model_not_ready_message_hides_internals() {
let msg = model_not_ready_message("my-model").to_lowercase();
for leak in [
"prefill",
"decode",
"encode",
"worker",
"namespace",
"needs",
] {
assert!(
!msg.contains(leak),
"not-ready message leaks internal term `{leak}`: {msg}"
);
}
assert!(model_not_ready_message("my-model").contains("my-model"));
assert!(msg.contains("retry"));
}
#[test]
fn test_unavailable_paths_share_one_message() {
let backstop = ErrorMessage::from_model_error(&ModelManagerError::ModelUnavailable(
"my-model".to_string(),
));
assert_eq!(backstop.0, StatusCode::SERVICE_UNAVAILABLE);
assert_eq!(backstop.1.message, model_not_ready_message("my-model"));
let gate = ErrorMessage::service_unavailable_with_body(model_not_ready_message("my-model"));
assert_eq!(gate.1.message, backstop.1.message);
}
#[test]
fn test_extract_error_type_from_response_internal() {
let response = ErrorMessage::internal_server_error("Something went wrong");
assert_eq!(
extract_error_type_from_response(&response),
ErrorType::Internal
);
}
#[test]
fn test_extract_error_type_from_response_not_implemented() {
let response = ErrorMessage::not_implemented_error("Feature not available");
assert_eq!(
extract_error_type_from_response(&response),
ErrorType::NotImplemented
);
}
use std::collections::{HashMap, HashSet};
use dynamo_protocols::types::{
ChatChoiceStream, ChatCompletionMessageToolCallChunk, ChatCompletionStreamResponseDelta,
ChatCompletionStreamResponseDeltaFunctionCall, CreateChatCompletionStreamResponse,
FinishReason, FunctionCallStream, FunctionType, Role,
};
use dynamo_runtime::protocols::annotated::Annotated;
fn extract_sse_data_json(event: &axum::response::sse::Event) -> serde_json::Value {
let debug = format!("{:?}", event);
let data_marker = "data: ";
let after_data = debug
.find(data_marker)
.map(|p| p + data_marker.len())
.expect("no 'data: ' in Event debug output");
let rest = &debug[after_data..];
let json_start = rest.find('{').expect("no JSON object after data:");
let mut depth = 0i32;
let mut json_end = 0;
for (i, b) in rest[json_start..].bytes().enumerate() {
match b {
b'{' => depth += 1,
b'}' => {
depth -= 1;
if depth == 0 {
json_end = json_start + i + 1;
break;
}
}
_ => {}
}
}
let raw = &rest[json_start..json_end];
let s = raw
.replace("\\\\\\\"", "\x00NESTED\x00")
.replace("\\\"", "\"")
.replace("\x00NESTED\x00", "\\\"");
let mut result = Vec::new();
let sbytes = s.as_bytes();
let mut idx = 0;
while idx < sbytes.len() {
if idx + 3 < sbytes.len()
&& sbytes[idx] == b'\\'
&& sbytes[idx + 1] == b'x'
&& let Ok(v) = u8::from_str_radix(
std::str::from_utf8(&sbytes[idx + 2..idx + 4]).unwrap_or(""),
16,
)
{
result.push(v);
idx += 4;
continue;
}
result.push(sbytes[idx]);
idx += 1;
}
let final_str = String::from_utf8_lossy(&result);
serde_json::from_str(&final_str).unwrap_or_else(|e| {
panic!(
"failed to parse JSON from Event: {e}\nraw: {raw}\nunescaped: {s}\nfinal: {final_str}"
)
})
}
fn assert_event_type(event: &axum::response::sse::Event, expected: &str) {
let debug = format!("{:?}", event);
let pattern = format!("event: {expected}\\n");
assert!(
debug.contains(&pattern),
"expected event type '{expected}' not found in: {debug}"
);
}
fn make_stream_response(
choices: Vec<ChatChoiceStream>,
) -> Annotated<NvCreateChatCompletionStreamResponse> {
let response = NvCreateChatCompletionStreamResponse {
inner: CreateChatCompletionStreamResponse {
id: "test-id".to_string(),
choices,
created: 0,
model: "test-model".to_string(),
system_fingerprint: None,
object: "chat.completion.chunk".to_string(),
usage: None,
service_tier: None,
},
nvext: None,
llm_metrics: None,
};
Annotated {
id: Some("test-id".to_string()),
data: Some(response),
event: None,
comment: None,
error: None,
}
}
fn collect_tool_dispatch_events(
response: &Annotated<NvCreateChatCompletionStreamResponse>,
dispatched_ids: &mut HashSet<String>,
) -> Vec<Result<Event, axum::Error>> {
let mut events = Vec::new();
streaming_tool_dispatch_events(response, dispatched_ids, &mut events);
events
}
fn collect_reasoning_dispatch_events(
response: &Annotated<NvCreateChatCompletionStreamResponse>,
buffers: &mut HashMap<u32, String>,
) -> Vec<Result<Event, axum::Error>> {
let mut events = Vec::new();
accumulate_reasoning_dispatch(response, buffers, &mut events);
events
}
fn make_choice_with_reasoning(
index: u32,
reasoning: Option<&str>,
finish: Option<FinishReason>,
) -> ChatChoiceStream {
#[allow(deprecated)]
ChatChoiceStream {
index,
delta: ChatCompletionStreamResponseDelta {
content: None,
function_call: None,
tool_calls: None,
role: None,
refusal: None,
reasoning_content: reasoning.map(|s| s.to_string()),
},
finish_reason: finish,
logprobs: None,
}
}
fn make_choice_with_tool_call(
index: u32,
id: Option<&str>,
name: Option<&str>,
arguments: Option<&str>,
) -> ChatChoiceStream {
let tool_call = ChatCompletionMessageToolCallChunk {
index: 0,
id: id.map(|s| s.to_string()),
r#type: Some(FunctionType::Function),
function: Some(FunctionCallStream {
name: name.map(|s| s.to_string()),
arguments: arguments.map(|s| s.to_string()),
}),
};
#[allow(deprecated)]
ChatChoiceStream {
index,
delta: ChatCompletionStreamResponseDelta {
content: None,
function_call: None,
tool_calls: Some(vec![tool_call]),
role: None,
refusal: None,
reasoning_content: None,
},
finish_reason: None,
logprobs: None,
}
}
#[test]
fn test_tool_dispatch_emits_event_for_complete_tool_call() {
let response = make_stream_response(vec![make_choice_with_tool_call(
0,
Some("call_123"),
Some("get_weather"),
Some(r#"{"city":"Paris"}"#),
)]);
let events = collect_tool_dispatch_events(&response, &mut HashSet::new());
assert_eq!(events.len(), 1);
let event = events[0].as_ref().unwrap();
assert_event_type(event, "tool_call_dispatch");
let json = extract_sse_data_json(event);
assert_eq!(json["choice_index"], 0);
assert_eq!(json["tool_call"]["id"], "call_123");
assert_eq!(json["tool_call"]["function"]["name"], "get_weather");
assert_eq!(
json["tool_call"]["function"]["arguments"],
r#"{"city":"Paris"}"#
);
}
#[test]
fn test_tool_dispatch_skips_incomplete_tool_call_no_id() {
let response = make_stream_response(vec![make_choice_with_tool_call(
0,
None, Some("get_weather"),
Some(r#"{"city":"Paris"}"#),
)]);
let events = collect_tool_dispatch_events(&response, &mut HashSet::new());
assert!(events.is_empty(), "should not dispatch without id");
}
#[test]
fn test_tool_dispatch_skips_incomplete_tool_call_no_name() {
let response = make_stream_response(vec![make_choice_with_tool_call(
0,
Some("call_123"),
None, Some(r#"{"city":"Paris"}"#),
)]);
let events = collect_tool_dispatch_events(&response, &mut HashSet::new());
assert!(events.is_empty(), "should not dispatch without name");
}
#[test]
fn test_tool_dispatch_skips_incomplete_tool_call_no_arguments() {
let response = make_stream_response(vec![make_choice_with_tool_call(
0,
Some("call_123"),
Some("get_weather"),
None, )]);
let events = collect_tool_dispatch_events(&response, &mut HashSet::new());
assert!(events.is_empty(), "should not dispatch without arguments");
}
#[test]
fn test_tool_dispatch_multiple_tool_calls() {
let tc1 = ChatCompletionMessageToolCallChunk {
index: 0,
id: Some("call_1".to_string()),
r#type: Some(FunctionType::Function),
function: Some(FunctionCallStream {
name: Some("get_weather".to_string()),
arguments: Some(r#"{"city":"Paris"}"#.to_string()),
}),
};
let tc2 = ChatCompletionMessageToolCallChunk {
index: 1,
id: Some("call_2".to_string()),
r#type: Some(FunctionType::Function),
function: Some(FunctionCallStream {
name: Some("get_time".to_string()),
arguments: Some(r#"{"tz":"UTC"}"#.to_string()),
}),
};
#[allow(deprecated)]
let choice = ChatChoiceStream {
index: 0,
delta: ChatCompletionStreamResponseDelta {
content: None,
function_call: None,
tool_calls: Some(vec![tc1, tc2]),
role: None,
refusal: None,
reasoning_content: None,
},
finish_reason: None,
logprobs: None,
};
let response = make_stream_response(vec![choice]);
let events = collect_tool_dispatch_events(&response, &mut HashSet::new());
assert_eq!(events.len(), 2, "should dispatch both tool calls");
let json0 = extract_sse_data_json(events[0].as_ref().unwrap());
assert_eq!(json0["tool_call"]["id"], "call_1");
assert_eq!(json0["tool_call"]["function"]["name"], "get_weather");
let json1 = extract_sse_data_json(events[1].as_ref().unwrap());
assert_eq!(json1["tool_call"]["id"], "call_2");
assert_eq!(json1["tool_call"]["function"]["name"], "get_time");
}
#[test]
fn test_tool_dispatch_no_data() {
let response: Annotated<NvCreateChatCompletionStreamResponse> = Annotated {
id: Some("test".to_string()),
data: None,
event: None,
comment: None,
error: None,
};
let events = collect_tool_dispatch_events(&response, &mut HashSet::new());
assert!(events.is_empty());
}
#[test]
fn test_tool_dispatch_empty_choices() {
let response = make_stream_response(vec![]);
let events = collect_tool_dispatch_events(&response, &mut HashSet::new());
assert!(events.is_empty());
}
#[test]
fn test_tool_dispatch_mixed_complete_and_incomplete() {
let complete = ChatCompletionMessageToolCallChunk {
index: 0,
id: Some("call_complete".to_string()),
r#type: Some(FunctionType::Function),
function: Some(FunctionCallStream {
name: Some("get_weather".to_string()),
arguments: Some(r#"{"city":"Paris"}"#.to_string()),
}),
};
let incomplete = ChatCompletionMessageToolCallChunk {
index: 1,
id: Some("call_partial".to_string()),
r#type: Some(FunctionType::Function),
function: Some(FunctionCallStream {
name: Some("search".to_string()),
arguments: None, }),
};
#[allow(deprecated)]
let choice = ChatChoiceStream {
index: 0,
delta: ChatCompletionStreamResponseDelta {
content: None,
function_call: None,
tool_calls: Some(vec![complete, incomplete]),
role: None,
refusal: None,
reasoning_content: None,
},
finish_reason: None,
logprobs: None,
};
let response = make_stream_response(vec![choice]);
let events = collect_tool_dispatch_events(&response, &mut HashSet::new());
assert_eq!(
events.len(),
1,
"only the complete tool call should dispatch"
);
let json = extract_sse_data_json(events[0].as_ref().unwrap());
assert_eq!(json["tool_call"]["id"], "call_complete");
}
#[test]
fn test_tool_dispatch_function_none() {
let tool_call = ChatCompletionMessageToolCallChunk {
index: 0,
id: Some("call_999".to_string()),
r#type: Some(FunctionType::Function),
function: None,
};
#[allow(deprecated)]
let choice = ChatChoiceStream {
index: 0,
delta: ChatCompletionStreamResponseDelta {
content: None,
function_call: None,
tool_calls: Some(vec![tool_call]),
role: None,
refusal: None,
reasoning_content: None,
},
finish_reason: None,
logprobs: None,
};
let response = make_stream_response(vec![choice]);
let events = collect_tool_dispatch_events(&response, &mut HashSet::new());
assert!(events.is_empty(), "function: None should not dispatch");
}
#[test]
fn test_tool_dispatch_empty_arguments_still_dispatches() {
let response = make_stream_response(vec![make_choice_with_tool_call(
0,
Some("call_empty"),
Some("no_params_tool"),
Some(""),
)]);
let events = collect_tool_dispatch_events(&response, &mut HashSet::new());
assert_eq!(events.len(), 1, "empty arguments should still dispatch");
let json = extract_sse_data_json(events[0].as_ref().unwrap());
assert_eq!(json["tool_call"]["id"], "call_empty");
assert_eq!(json["tool_call"]["function"]["name"], "no_params_tool");
assert_eq!(json["tool_call"]["function"]["arguments"], "");
}
#[test]
fn test_tool_dispatch_n_greater_than_1_includes_choice_index() {
let choice_0 = make_choice_with_tool_call(
0,
Some("call_a"),
Some("get_weather"),
Some(r#"{"city":"Paris"}"#),
);
let choice_1 = make_choice_with_tool_call(
1,
Some("call_b"),
Some("get_time"),
Some(r#"{"tz":"UTC"}"#),
);
let response = make_stream_response(vec![choice_0, choice_1]);
let events = collect_tool_dispatch_events(&response, &mut HashSet::new());
assert_eq!(events.len(), 2, "should dispatch from both choices");
let json0 = extract_sse_data_json(events[0].as_ref().unwrap());
assert_eq!(json0["choice_index"], 0);
assert_eq!(json0["tool_call"]["id"], "call_a");
let json1 = extract_sse_data_json(events[1].as_ref().unwrap());
assert_eq!(json1["choice_index"], 1);
assert_eq!(json1["tool_call"]["id"], "call_b");
}
#[test]
fn test_tool_dispatch_dedup_skips_already_dispatched_id() {
let response = make_stream_response(vec![make_choice_with_tool_call(
0,
Some("call_dup"),
Some("get_weather"),
Some(r#"{"city":"Paris"}"#),
)]);
let mut dispatched = HashSet::new();
let events = collect_tool_dispatch_events(&response, &mut dispatched);
assert_eq!(events.len(), 1);
let events = collect_tool_dispatch_events(&response, &mut dispatched);
assert!(events.is_empty(), "duplicate id should not dispatch twice");
}
#[test]
fn test_reasoning_dispatch_accumulates_and_emits_once() {
let mut buffers: HashMap<u32, String> = HashMap::new();
let r1 = make_stream_response(vec![make_choice_with_reasoning(0, Some("Let me"), None)]);
let events = collect_reasoning_dispatch_events(&r1, &mut buffers);
assert!(
events.is_empty(),
"should not emit yet — still accumulating"
);
assert_eq!(buffers.get(&0).map(|s| s.as_str()), Some("Let me"));
let r2 = make_stream_response(vec![make_choice_with_reasoning(0, Some(" think"), None)]);
let events = collect_reasoning_dispatch_events(&r2, &mut buffers);
assert!(
events.is_empty(),
"should not emit yet — still accumulating"
);
assert_eq!(buffers.get(&0).map(|s| s.as_str()), Some("Let me think"));
let r3 = make_stream_response(vec![make_choice_with_reasoning(0, None, None)]);
let events = collect_reasoning_dispatch_events(&r3, &mut buffers);
assert_eq!(events.len(), 1, "should emit single reasoning_dispatch");
let event = events[0].as_ref().unwrap();
assert_event_type(event, "reasoning_dispatch");
let json = extract_sse_data_json(event);
assert_eq!(json["reasoning_content"], "Let me think");
assert_eq!(json["index"], 0);
assert!(
buffers.get(&0).is_none_or(|s| s.is_empty()),
"buffer should be cleared after emit"
);
}
#[test]
fn test_reasoning_dispatch_flushes_on_finish_reason() {
let mut buffers: HashMap<u32, String> = HashMap::new();
let r1 = make_stream_response(vec![make_choice_with_reasoning(
0,
Some("Thinking..."),
None,
)]);
collect_reasoning_dispatch_events(&r1, &mut buffers);
let r2 = make_stream_response(vec![make_choice_with_reasoning(
0,
Some(" more"),
Some(FinishReason::Length),
)]);
let events = collect_reasoning_dispatch_events(&r2, &mut buffers);
assert_eq!(events.len(), 1, "should flush on finish_reason");
let json = extract_sse_data_json(events[0].as_ref().unwrap());
assert_eq!(json["reasoning_content"], "Thinking... more");
}
#[test]
fn test_reasoning_dispatch_flushes_on_stop() {
let mut buffers: HashMap<u32, String> = HashMap::new();
let r1 = make_stream_response(vec![make_choice_with_reasoning(
0,
Some("Analysis complete"),
None,
)]);
collect_reasoning_dispatch_events(&r1, &mut buffers);
let r2 = make_stream_response(vec![make_choice_with_reasoning(
0,
Some("."),
Some(FinishReason::Stop),
)]);
let events = collect_reasoning_dispatch_events(&r2, &mut buffers);
assert_eq!(events.len(), 1, "should flush on FinishReason::Stop");
let json = extract_sse_data_json(events[0].as_ref().unwrap());
assert_eq!(json["reasoning_content"], "Analysis complete.");
}
#[test]
fn test_reasoning_dispatch_no_reasoning_no_event() {
let mut buffers: HashMap<u32, String> = HashMap::new();
let r = make_stream_response(vec![make_choice_with_reasoning(0, None, None)]);
let events = collect_reasoning_dispatch_events(&r, &mut buffers);
assert!(events.is_empty(), "no reasoning content = no event");
}
#[test]
fn test_reasoning_dispatch_empty_string_not_accumulated() {
let mut buffers: HashMap<u32, String> = HashMap::new();
let r = make_stream_response(vec![make_choice_with_reasoning(0, Some(""), None)]);
let events = collect_reasoning_dispatch_events(&r, &mut buffers);
assert!(events.is_empty());
assert!(
buffers.get(&0).is_none_or(|s| s.is_empty()),
"empty string should not accumulate"
);
}
#[test]
fn test_reasoning_dispatch_no_data() {
let mut buffers: HashMap<u32, String> = HashMap::new();
let response: Annotated<NvCreateChatCompletionStreamResponse> = Annotated {
id: Some("test".to_string()),
data: None,
event: None,
comment: None,
error: None,
};
let events = collect_reasoning_dispatch_events(&response, &mut buffers);
assert!(events.is_empty());
}
#[test]
fn test_reasoning_dispatch_empty_choices() {
let mut buffers: HashMap<u32, String> = HashMap::new();
let response = make_stream_response(vec![]);
let events = collect_reasoning_dispatch_events(&response, &mut buffers);
assert!(events.is_empty());
}
#[test]
fn test_reasoning_dispatch_multi_choice_independent_buffers() {
let mut buffers: HashMap<u32, String> = HashMap::new();
let r1 = make_stream_response(vec![
make_choice_with_reasoning(0, Some("Thinking A"), None),
make_choice_with_reasoning(1, Some("Thinking B"), None),
]);
let events = collect_reasoning_dispatch_events(&r1, &mut buffers);
assert!(events.is_empty(), "both still accumulating");
assert_eq!(buffers.get(&0).map(|s| s.as_str()), Some("Thinking A"));
assert_eq!(buffers.get(&1).map(|s| s.as_str()), Some("Thinking B"));
let r2 = make_stream_response(vec![
make_choice_with_reasoning(0, None, None),
make_choice_with_reasoning(1, Some(" more"), None),
]);
let events = collect_reasoning_dispatch_events(&r2, &mut buffers);
assert_eq!(events.len(), 1, "only choice 0 should emit");
let json = extract_sse_data_json(events[0].as_ref().unwrap());
assert_eq!(json["reasoning_content"], "Thinking A");
assert_eq!(json["index"], 0);
let r3 = make_stream_response(vec![make_choice_with_reasoning(1, None, None)]);
let events = collect_reasoning_dispatch_events(&r3, &mut buffers);
assert_eq!(events.len(), 1, "choice 1 should emit");
let json = extract_sse_data_json(events[0].as_ref().unwrap());
assert_eq!(json["reasoning_content"], "Thinking B more");
assert_eq!(json["index"], 1);
}
#[test]
fn test_reasoning_dispatch_multiple_blocks() {
let mut buffers: HashMap<u32, String> = HashMap::new();
let r1 = make_stream_response(vec![make_choice_with_reasoning(0, Some("First"), None)]);
collect_reasoning_dispatch_events(&r1, &mut buffers);
let r2 = make_stream_response(vec![make_choice_with_reasoning(0, None, None)]);
let events = collect_reasoning_dispatch_events(&r2, &mut buffers);
assert_eq!(events.len(), 1);
let json = extract_sse_data_json(events[0].as_ref().unwrap());
assert_eq!(json["reasoning_content"], "First");
let r3 = make_stream_response(vec![make_choice_with_reasoning(0, Some("Second"), None)]);
collect_reasoning_dispatch_events(&r3, &mut buffers);
let r4 = make_stream_response(vec![make_choice_with_reasoning(0, None, None)]);
let events = collect_reasoning_dispatch_events(&r4, &mut buffers);
assert_eq!(events.len(), 1);
let json = extract_sse_data_json(events[0].as_ref().unwrap());
assert_eq!(
json["reasoning_content"], "Second",
"second emit should only contain second block's content"
);
}
#[test]
fn test_reasoning_dispatch_unicode() {
let mut buffers: HashMap<u32, String> = HashMap::new();
let r1 = make_stream_response(vec![make_choice_with_reasoning(
0,
Some("让我想想 🤔"),
None,
)]);
collect_reasoning_dispatch_events(&r1, &mut buffers);
let r2 = make_stream_response(vec![make_choice_with_reasoning(
0,
Some(" 分析完成 ✅"),
None,
)]);
collect_reasoning_dispatch_events(&r2, &mut buffers);
let r3 = make_stream_response(vec![make_choice_with_reasoning(0, None, None)]);
let events = collect_reasoning_dispatch_events(&r3, &mut buffers);
assert_eq!(events.len(), 1);
let json = extract_sse_data_json(events[0].as_ref().unwrap());
assert_eq!(json["reasoning_content"], "让我想想 🤔 分析完成 ✅");
}
#[allow(clippy::too_many_arguments)]
fn make_delta(
content: Option<&str>,
reasoning: Option<&str>,
tool_calls: Option<Vec<ChatCompletionMessageToolCallChunk>>,
finish: Option<FinishReason>,
usage: Option<dynamo_protocols::types::CompletionUsage>,
role: Option<Role>,
refusal: Option<&str>,
function_call: Option<ChatCompletionStreamResponseDeltaFunctionCall>,
) -> NvCreateChatCompletionStreamResponse {
use dynamo_protocols::types::ChatCompletionMessageContent;
#[allow(deprecated)]
let choice = ChatChoiceStream {
index: 0,
delta: ChatCompletionStreamResponseDelta {
content: content.map(|s| ChatCompletionMessageContent::Text(s.to_string())),
function_call,
tool_calls,
role,
refusal: refusal.map(|s| s.to_string()),
reasoning_content: reasoning.map(|s| s.to_string()),
},
finish_reason: finish,
logprobs: None,
};
NvCreateChatCompletionStreamResponse {
inner: CreateChatCompletionStreamResponse {
id: "test".to_string(),
choices: vec![choice],
created: 0,
model: "m".to_string(),
system_fingerprint: None,
object: "chat.completion.chunk".to_string(),
usage,
service_tier: None,
},
nvext: None,
llm_metrics: None,
}
}
#[test]
fn test_is_empty_stream_response() {
assert!(
is_empty_stream_response(&make_delta(None, None, None, None, None, None, None, None)),
"all-None delta → empty",
);
assert!(
!is_empty_stream_response(&make_delta(
Some("hi"),
None,
None,
None,
None,
None,
None,
None
)),
"content present → not empty",
);
assert!(
!is_empty_stream_response(&make_delta(
None,
Some("thinking"),
None,
None,
None,
None,
None,
None
)),
"reasoning present → not empty",
);
assert!(
!is_empty_stream_response(&make_delta(
None,
None,
None,
Some(FinishReason::Stop),
None,
None,
None,
None,
)),
"finish_reason → not empty",
);
let tc = vec![ChatCompletionMessageToolCallChunk {
index: 0,
id: Some("call_1".to_string()),
r#type: Some(FunctionType::Function),
function: Some(FunctionCallStream {
name: Some("f".to_string()),
arguments: Some("{}".to_string()),
}),
}];
assert!(
!is_empty_stream_response(&make_delta(
None,
None,
Some(tc),
None,
None,
None,
None,
None
)),
"tool_calls present → not empty",
);
let usage = dynamo_protocols::types::CompletionUsage {
prompt_tokens: 10,
completion_tokens: 5,
total_tokens: 15,
prompt_tokens_details: None,
completion_tokens_details: None,
};
assert!(
!is_empty_stream_response(&make_delta(
None,
None,
None,
None,
Some(usage),
None,
None,
None
)),
"usage present → not empty",
);
assert!(
is_empty_stream_response(&make_delta(
None,
None,
None,
None,
None,
Some(Role::Assistant),
None,
None,
)),
"role-only → empty",
);
assert!(
!is_empty_stream_response(&make_delta(
None,
None,
None,
None,
None,
None,
Some("I can't help with that"),
None,
)),
"refusal present → not empty",
);
assert!(
!is_empty_stream_response(&make_delta(
None,
None,
None,
None,
None,
None,
None,
Some(ChatCompletionStreamResponseDeltaFunctionCall {
name: Some("my_fn".to_string()),
arguments: Some("{}".to_string()),
}),
)),
"function_call present → not empty",
);
}
#[test]
fn test_chat_predicate_filters_text_empty_string() {
use dynamo_protocols::types::{
ChatChoiceLogprobs, ChatCompletionMessageContent, ChatCompletionResponseContentPart,
ChatCompletionResponseContentPartText, ChatCompletionTokenLogprob,
};
let resp = make_delta(Some(""), None, None, None, None, None, None, None);
assert!(
is_empty_stream_response(&resp),
"Text(\"\") delta should be filtered as empty",
);
let mut resp = make_delta(None, None, None, None, None, None, None, None);
resp.inner.choices[0].delta.content = Some(ChatCompletionMessageContent::Parts(Vec::new()));
assert!(
is_empty_stream_response(&resp),
"Parts(vec![]) delta should be filtered as empty",
);
let mut resp = make_delta(None, None, None, None, None, None, None, None);
resp.inner.choices[0].delta.content = Some(ChatCompletionMessageContent::Parts(vec![
ChatCompletionResponseContentPart::Text(ChatCompletionResponseContentPartText {
text: "hi".to_string(),
}),
]));
assert!(
!is_empty_stream_response(&resp),
"Parts with content must not be filtered",
);
let resp = make_delta(
Some(""),
None,
None,
Some(FinishReason::Stop),
None,
None,
None,
None,
);
assert!(
!is_empty_stream_response(&resp),
"Text(\"\") + finish_reason must not be filtered",
);
let mut resp = make_delta(Some(""), None, None, None, None, None, None, None);
resp.inner.choices[0].logprobs = Some(ChatChoiceLogprobs {
content: Some(vec![ChatCompletionTokenLogprob {
token: "h".to_string(),
logprob: -0.5,
bytes: Some(vec![104]),
top_logprobs: vec![],
}]),
refusal: None,
});
assert!(
!is_empty_stream_response(&resp),
"Text(\"\") + logprobs must not be filtered",
);
}
use dynamo_protocols::types::{Choice, CompletionFinishReason, CreateCompletionResponse};
fn make_completion_chunk(
text: &str,
finish: Option<CompletionFinishReason>,
usage: Option<dynamo_protocols::types::CompletionUsage>,
) -> NvCreateCompletionResponse {
let choice = Choice {
text: text.to_string(),
index: 0,
logprobs: None,
finish_reason: finish,
};
NvCreateCompletionResponse {
inner: CreateCompletionResponse {
id: "test".to_string(),
choices: vec![choice],
created: 0,
model: "m".to_string(),
system_fingerprint: None,
object: "text_completion".to_string(),
usage,
},
nvext: None,
}
}
#[test]
fn test_is_empty_completion_stream_response() {
assert!(
is_empty_completion_stream_response(&make_completion_chunk("", None, None)),
"empty text, no finish → empty",
);
assert!(
!is_empty_completion_stream_response(&make_completion_chunk("hi", None, None)),
"text present → not empty",
);
assert!(
!is_empty_completion_stream_response(&make_completion_chunk(
"",
Some(CompletionFinishReason::Stop),
None,
)),
"finish_reason → not empty",
);
let usage = dynamo_protocols::types::CompletionUsage {
prompt_tokens: 10,
completion_tokens: 5,
total_tokens: 15,
prompt_tokens_details: None,
completion_tokens_details: None,
};
assert!(
!is_empty_completion_stream_response(&make_completion_chunk("", None, Some(usage))),
"usage present → not empty",
);
}
#[test]
fn decode_base64_embedding_to_floats_round_trips_little_endian_f32() {
use base64::Engine as _;
let floats: Vec<f32> = vec![0.0, 1.0, -1.0, 2.5, -42.5, f32::MIN, f32::MAX];
let mut bytes: Vec<u8> = Vec::with_capacity(floats.len() * 4);
for f in &floats {
bytes.extend_from_slice(&f.to_le_bytes());
}
let encoded = base64::engine::general_purpose::STANDARD.encode(&bytes);
let decoded = decode_base64_embedding_to_floats(&encoded)
.expect("valid base64 of f32 bytes should decode");
assert_eq!(decoded, floats);
}
#[test]
fn decode_base64_embedding_to_floats_rejects_invalid_base64() {
let result = decode_base64_embedding_to_floats("not!valid!base64");
assert!(
result.is_err(),
"non-base64 input should fail decode, got Ok({:?})",
result.ok()
);
}
#[test]
fn decode_base64_embedding_to_floats_rejects_non_multiple_of_4_byte_length() {
use base64::Engine as _;
let bytes: Vec<u8> = vec![1, 2, 3, 4, 5];
let encoded = base64::engine::general_purpose::STANDARD.encode(&bytes);
let result = decode_base64_embedding_to_floats(&encoded);
assert!(result.is_err(), "5-byte payload must fail, got Ok");
let err_msg = result.unwrap_err().to_string();
assert!(
err_msg.contains("not a multiple of 4"),
"error should mention the multiple-of-4 check, got: {err_msg}"
);
}
}