bugfix: multi-model routing for /generate api (#12979)

Co-authored-by: Simo Lin <linsimo.mark@gmail.com>
Co-authored-by: Chang Su <chang.s.su@oracle.com>
This commit is contained in:
Siyuan Chen
2025-11-13 03:13:27 +08:00
committed by GitHub
parent d646cf6347
commit 4ef4390540
5 changed files with 40 additions and 28 deletions

View File

@@ -33,6 +33,7 @@ fn get_bootstrap_info(worker: &BasicWorker) -> (String, Option<u16>) {
fn default_generate_request() -> GenerateRequest {
GenerateRequest {
text: None,
model: None,
input_ids: None,
input_embeds: None,
image_data: None,

View File

@@ -21,6 +21,8 @@ pub struct GenerateRequest {
#[serde(skip_serializing_if = "Option::is_none")]
pub text: Option<String>,
pub model: Option<String>,
/// Input IDs for tokenized input
#[serde(skip_serializing_if = "Option::is_none")]
pub input_ids: Option<InputIds>,
@@ -201,8 +203,12 @@ impl GenerationRequest for GenerateRequest {
}
fn get_model(&self) -> Option<&str> {
// Generate requests typically don't have a model field
None
// Generate requests have an optional model field
if let Some(s) = &self.model {
Some(s.as_str())
} else {
None
}
}
fn extract_text_for_routing(&self) -> String {

View File

@@ -350,12 +350,12 @@ impl RouterTrait for RouterManager {
&self,
headers: Option<&HeaderMap>,
body: &GenerateRequest,
_model_id: Option<&str>,
model_id: Option<&str>,
) -> Response {
let router = self.select_router_for_request(headers, None);
let router = self.select_router_for_request(headers, model_id);
if let Some(router) = router {
router.route_generate(headers, body, None).await
router.route_generate(headers, body, model_id).await
} else {
(
StatusCode::NOT_FOUND,
@@ -369,12 +369,12 @@ impl RouterTrait for RouterManager {
&self,
headers: Option<&HeaderMap>,
body: &ChatCompletionRequest,
_model_id: Option<&str>,
model_id: Option<&str>,
) -> Response {
let router = self.select_router_for_request(headers, Some(&body.model));
let router = self.select_router_for_request(headers, model_id);
if let Some(router) = router {
router.route_chat(headers, body, Some(&body.model)).await
router.route_chat(headers, body, model_id).await
} else {
(
StatusCode::NOT_FOUND,
@@ -388,14 +388,12 @@ impl RouterTrait for RouterManager {
&self,
headers: Option<&HeaderMap>,
body: &CompletionRequest,
_model_id: Option<&str>,
model_id: Option<&str>,
) -> Response {
let router = self.select_router_for_request(headers, Some(&body.model));
let router = self.select_router_for_request(headers, model_id);
if let Some(router) = router {
router
.route_completion(headers, body, Some(&body.model))
.await
router.route_completion(headers, body, model_id).await
} else {
(
StatusCode::NOT_FOUND,
@@ -487,14 +485,12 @@ impl RouterTrait for RouterManager {
&self,
headers: Option<&HeaderMap>,
body: &EmbeddingRequest,
_model_id: Option<&str>,
model_id: Option<&str>,
) -> Response {
let router = self.select_router_for_request(headers, Some(&body.model));
let router = self.select_router_for_request(headers, model_id);
if let Some(router) = router {
router
.route_embeddings(headers, body, Some(&body.model))
.await
router.route_embeddings(headers, body, model_id).await
} else {
(
StatusCode::NOT_FOUND,
@@ -510,7 +506,7 @@ impl RouterTrait for RouterManager {
body: &RerankRequest,
model_id: Option<&str>,
) -> Response {
let router = self.select_router_for_request(headers, None);
let router = self.select_router_for_request(headers, model_id);
if let Some(router) = router {
router.route_rerank(headers, body, model_id).await
@@ -529,7 +525,7 @@ impl RouterTrait for RouterManager {
body: &ClassifyRequest,
model_id: Option<&str>,
) -> Response {
let router = self.select_router_for_request(headers, Some(&body.model));
let router = self.select_router_for_request(headers, model_id);
if let Some(router) = router {
router.route_classify(headers, body, model_id).await

View File

@@ -136,9 +136,10 @@ async fn generate(
headers: http::HeaderMap,
Json(body): Json<GenerateRequest>,
) -> Response {
let model_id = body.model.as_deref();
state
.router
.route_generate(Some(&headers), &body, None)
.route_generate(Some(&headers), &body, model_id)
.await
}
@@ -147,7 +148,10 @@ async fn v1_chat_completions(
headers: http::HeaderMap,
ValidatedJson(body): ValidatedJson<ChatCompletionRequest>,
) -> Response {
state.router.route_chat(Some(&headers), &body, None).await
state
.router
.route_chat(Some(&headers), &body, Some(&body.model))
.await
}
async fn v1_completions(
@@ -157,7 +161,7 @@ async fn v1_completions(
) -> Response {
state
.router
.route_completion(Some(&headers), &body, None)
.route_completion(Some(&headers), &body, Some(&body.model))
.await
}
@@ -166,7 +170,10 @@ async fn rerank(
headers: http::HeaderMap,
ValidatedJson(body): ValidatedJson<RerankRequest>,
) -> Response {
state.router.route_rerank(Some(&headers), &body, None).await
state
.router
.route_rerank(Some(&headers), &body, Some(&body.model))
.await
}
async fn v1_rerank(
@@ -174,9 +181,10 @@ async fn v1_rerank(
headers: http::HeaderMap,
Json(body): Json<V1RerankReqInput>,
) -> Response {
let rerank_body = &body.into();
state
.router
.route_rerank(Some(&headers), &body.into(), None)
.route_rerank(Some(&headers), rerank_body, Some(&rerank_body.model))
.await
}
@@ -187,7 +195,7 @@ async fn v1_responses(
) -> Response {
state
.router
.route_responses(Some(&headers), &body, None)
.route_responses(Some(&headers), &body, Some(&body.model))
.await
}
@@ -198,7 +206,7 @@ async fn v1_embeddings(
) -> Response {
state
.router
.route_embeddings(Some(&headers), &body, None)
.route_embeddings(Some(&headers), &body, Some(&body.model))
.await
}
@@ -209,7 +217,7 @@ async fn v1_classify(
) -> Response {
state
.router
.route_classify(Some(&headers), &body, None)
.route_classify(Some(&headers), &body, Some(&body.model))
.await
}

View File

@@ -602,6 +602,7 @@ async fn test_unsupported_endpoints() {
let generate_request = GenerateRequest {
text: Some("Hello world".to_string()),
model: None,
input_ids: None,
input_embeds: None,
image_data: None,