Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
48 changes: 39 additions & 9 deletions rust/src/server/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -35,9 +35,24 @@ use crate::routes::build_router;
use crate::server_info::ServerInfoSnapshot;
use crate::state::AppState;

/// Resolve the public model names accepted by the frontend.
fn effective_served_model_names(model: &str, served_model_name: &[String]) -> Vec<String> {
if served_model_name.is_empty() {
vec![model.to_string()]
} else {
served_model_name.to_vec()
}
}

/// Build the shared application state for one configured model and one engine
/// client.
async fn build_state(config: &Config) -> Result<Arc<AppState>> {
// If no served names are specified, fall back to the backend model path so
// that the API always has at least one valid model ID. Use the same primary
// public name for frontend-side metrics labels.
let served_model_names = effective_served_model_names(&config.model, &config.served_model_name);
let metrics_model_name = served_model_names[0].clone();

// Load both backends from the same model metadata so they stay in sync.
let loaded = load_model_backends(
&config.model,
Expand Down Expand Up @@ -68,7 +83,7 @@ async fn build_state(config: &Config) -> Result<Arc<AppState>> {
let client = EngineCoreClient::connect(EngineCoreClientConfig {
transport_mode: config.transport_mode.clone(),
coordinator_mode,
model_name: config.model.clone(),
model_name: metrics_model_name,
client_index: 0,
})
.await
Expand All @@ -81,14 +96,6 @@ async fn build_state(config: &Config) -> Result<Arc<AppState>> {
.with_tool_call_parser(config.tool_call_parser.clone())
.with_reasoning_parser(config.reasoning_parser.clone());

// If no served names are specified, fall back to the backend model path so
// that the API always has at least one valid model ID.
let served_model_names = if config.served_model_name.is_empty() {
vec![config.model.clone()]
} else {
config.served_model_name.clone()
};

Ok(Arc::new(
AppState::new(served_model_names, chat)
.with_api_server_options(config.api_server_options)
Expand Down Expand Up @@ -258,3 +265,26 @@ where
.unwrap_or_else(|| Instant::now() + config.shutdown_timeout);
state.shutdown(shutdown_deadline).await
}

#[cfg(test)]
mod tests {
use super::*;

#[test]
fn effective_served_model_names_falls_back_to_backend_model() {
assert_eq!(
effective_served_model_names("backend-model", &[]),
vec!["backend-model"]
);
}

#[test]
fn effective_served_model_names_preserves_public_names() {
let served_names = vec!["public-model".to_string(), "public-alias".to_string()];

assert_eq!(
effective_served_model_names("backend-model", &served_names),
served_names
);
}
}
79 changes: 79 additions & 0 deletions rust/src/server/src/routes/tests.rs
Original file line number Diff line number Diff line change
Expand Up @@ -1677,6 +1677,85 @@ async fn http_metrics_record_list_models_requests() {
);
}

#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
#[serial]
async fn request_metrics_use_served_model_name_label() {
let ipc = IpcNamespace::new().expect("create ipc namespace");
let handshake_address = ipc.handshake_endpoint();
let engine_id = b"engine-openai-served-model-metrics".to_vec();

let engine_task = MockEngineTask::new(spawn_mock_engine_task(
handshake_address.clone(),
engine_id.clone(),
|dealer, push| {
boxed_test_future(async move {
let add = recv_engine_message(dealer).await;
let request: EngineCoreRequest =
rmp_serde::from_slice(&add[1]).expect("decode request");
send_outputs(
push,
engine_outputs_for_request(&request.request_id, default_stream_output_specs()),
)
.await;
})
},
));

let client = EngineCoreClient::connect(
EngineCoreClientConfig::new_single(handshake_address)
.with_model_name("served-model-metrics")
.with_local_input_output_addresses(
Some(ipc.input_endpoint()),
Some(ipc.output_endpoint()),
),
)
.await
.expect("connect client");
let chat = ChatLlm::from_shared_backend(test_llm(client), Arc::new(FakeChatBackend::new()));
let mut app = build_router(Arc::new(AppState::new(
vec![
"served-model-metrics".to_string(),
"served-model-alias".to_string(),
],
chat,
)));
let before = METRICS.render().unwrap();

let response = app
.call(
Request::builder()
.method("POST")
.uri("/v1/chat/completions")
.header("content-type", "application/json")
.body(Body::from(
json!({
"model": "served-model-alias",
"stream": false,
"messages": [{"role": "user", "content": "hello"}]
})
.to_string(),
))
.expect("build request"),
)
.await
.expect("call app");

assert_eq!(response.status(), StatusCode::OK);
let _ = to_bytes(response.into_body(), usize::MAX).await.unwrap();

let after = METRICS.render().unwrap();
assert_eq!(
metric_delta(
&before,
&after,
"vllm:request_success_total",
Some("model_name=\"served-model-metrics\",engine=\"0\",finished_reason=\"stop\""),
),
1.0
);
engine_task.await.expect("mock engine task");
}

#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
#[serial]
async fn wrong_model_returns_not_found() {
Expand Down
Loading