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:
co-authored by
Simo Lin
Chang Su
parent
d646cf6347
commit
4ef4390540
@@ -33,6 +33,7 @@ fn get_bootstrap_info(worker: &BasicWorker) -> (String, Option<u16>) {
|
|||||||
fn default_generate_request() -> GenerateRequest {
|
fn default_generate_request() -> GenerateRequest {
|
||||||
GenerateRequest {
|
GenerateRequest {
|
||||||
text: None,
|
text: None,
|
||||||
|
model: None,
|
||||||
input_ids: None,
|
input_ids: None,
|
||||||
input_embeds: None,
|
input_embeds: None,
|
||||||
image_data: None,
|
image_data: None,
|
||||||
|
|||||||
@@ -21,6 +21,8 @@ pub struct GenerateRequest {
|
|||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
#[serde(skip_serializing_if = "Option::is_none")]
|
||||||
pub text: Option<String>,
|
pub text: Option<String>,
|
||||||
|
|
||||||
|
pub model: Option<String>,
|
||||||
|
|
||||||
/// Input IDs for tokenized input
|
/// Input IDs for tokenized input
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
#[serde(skip_serializing_if = "Option::is_none")]
|
||||||
pub input_ids: Option<InputIds>,
|
pub input_ids: Option<InputIds>,
|
||||||
@@ -201,9 +203,13 @@ impl GenerationRequest for GenerateRequest {
|
|||||||
}
|
}
|
||||||
|
|
||||||
fn get_model(&self) -> Option<&str> {
|
fn get_model(&self) -> Option<&str> {
|
||||||
// Generate requests typically don't have a model field
|
// Generate requests have an optional model field
|
||||||
|
if let Some(s) = &self.model {
|
||||||
|
Some(s.as_str())
|
||||||
|
} else {
|
||||||
None
|
None
|
||||||
}
|
}
|
||||||
|
}
|
||||||
|
|
||||||
fn extract_text_for_routing(&self) -> String {
|
fn extract_text_for_routing(&self) -> String {
|
||||||
// Check fields in priority order: text, input_ids
|
// Check fields in priority order: text, input_ids
|
||||||
|
|||||||
@@ -350,12 +350,12 @@ impl RouterTrait for RouterManager {
|
|||||||
&self,
|
&self,
|
||||||
headers: Option<&HeaderMap>,
|
headers: Option<&HeaderMap>,
|
||||||
body: &GenerateRequest,
|
body: &GenerateRequest,
|
||||||
_model_id: Option<&str>,
|
model_id: Option<&str>,
|
||||||
) -> Response {
|
) -> 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 {
|
if let Some(router) = router {
|
||||||
router.route_generate(headers, body, None).await
|
router.route_generate(headers, body, model_id).await
|
||||||
} else {
|
} else {
|
||||||
(
|
(
|
||||||
StatusCode::NOT_FOUND,
|
StatusCode::NOT_FOUND,
|
||||||
@@ -369,12 +369,12 @@ impl RouterTrait for RouterManager {
|
|||||||
&self,
|
&self,
|
||||||
headers: Option<&HeaderMap>,
|
headers: Option<&HeaderMap>,
|
||||||
body: &ChatCompletionRequest,
|
body: &ChatCompletionRequest,
|
||||||
_model_id: Option<&str>,
|
model_id: Option<&str>,
|
||||||
) -> Response {
|
) -> 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 {
|
if let Some(router) = router {
|
||||||
router.route_chat(headers, body, Some(&body.model)).await
|
router.route_chat(headers, body, model_id).await
|
||||||
} else {
|
} else {
|
||||||
(
|
(
|
||||||
StatusCode::NOT_FOUND,
|
StatusCode::NOT_FOUND,
|
||||||
@@ -388,14 +388,12 @@ impl RouterTrait for RouterManager {
|
|||||||
&self,
|
&self,
|
||||||
headers: Option<&HeaderMap>,
|
headers: Option<&HeaderMap>,
|
||||||
body: &CompletionRequest,
|
body: &CompletionRequest,
|
||||||
_model_id: Option<&str>,
|
model_id: Option<&str>,
|
||||||
) -> Response {
|
) -> 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 {
|
if let Some(router) = router {
|
||||||
router
|
router.route_completion(headers, body, model_id).await
|
||||||
.route_completion(headers, body, Some(&body.model))
|
|
||||||
.await
|
|
||||||
} else {
|
} else {
|
||||||
(
|
(
|
||||||
StatusCode::NOT_FOUND,
|
StatusCode::NOT_FOUND,
|
||||||
@@ -487,14 +485,12 @@ impl RouterTrait for RouterManager {
|
|||||||
&self,
|
&self,
|
||||||
headers: Option<&HeaderMap>,
|
headers: Option<&HeaderMap>,
|
||||||
body: &EmbeddingRequest,
|
body: &EmbeddingRequest,
|
||||||
_model_id: Option<&str>,
|
model_id: Option<&str>,
|
||||||
) -> Response {
|
) -> 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 {
|
if let Some(router) = router {
|
||||||
router
|
router.route_embeddings(headers, body, model_id).await
|
||||||
.route_embeddings(headers, body, Some(&body.model))
|
|
||||||
.await
|
|
||||||
} else {
|
} else {
|
||||||
(
|
(
|
||||||
StatusCode::NOT_FOUND,
|
StatusCode::NOT_FOUND,
|
||||||
@@ -510,7 +506,7 @@ impl RouterTrait for RouterManager {
|
|||||||
body: &RerankRequest,
|
body: &RerankRequest,
|
||||||
model_id: Option<&str>,
|
model_id: Option<&str>,
|
||||||
) -> Response {
|
) -> 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 {
|
if let Some(router) = router {
|
||||||
router.route_rerank(headers, body, model_id).await
|
router.route_rerank(headers, body, model_id).await
|
||||||
@@ -529,7 +525,7 @@ impl RouterTrait for RouterManager {
|
|||||||
body: &ClassifyRequest,
|
body: &ClassifyRequest,
|
||||||
model_id: Option<&str>,
|
model_id: Option<&str>,
|
||||||
) -> Response {
|
) -> 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 {
|
if let Some(router) = router {
|
||||||
router.route_classify(headers, body, model_id).await
|
router.route_classify(headers, body, model_id).await
|
||||||
|
|||||||
@@ -136,9 +136,10 @@ async fn generate(
|
|||||||
headers: http::HeaderMap,
|
headers: http::HeaderMap,
|
||||||
Json(body): Json<GenerateRequest>,
|
Json(body): Json<GenerateRequest>,
|
||||||
) -> Response {
|
) -> Response {
|
||||||
|
let model_id = body.model.as_deref();
|
||||||
state
|
state
|
||||||
.router
|
.router
|
||||||
.route_generate(Some(&headers), &body, None)
|
.route_generate(Some(&headers), &body, model_id)
|
||||||
.await
|
.await
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -147,7 +148,10 @@ async fn v1_chat_completions(
|
|||||||
headers: http::HeaderMap,
|
headers: http::HeaderMap,
|
||||||
ValidatedJson(body): ValidatedJson<ChatCompletionRequest>,
|
ValidatedJson(body): ValidatedJson<ChatCompletionRequest>,
|
||||||
) -> Response {
|
) -> 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(
|
async fn v1_completions(
|
||||||
@@ -157,7 +161,7 @@ async fn v1_completions(
|
|||||||
) -> Response {
|
) -> Response {
|
||||||
state
|
state
|
||||||
.router
|
.router
|
||||||
.route_completion(Some(&headers), &body, None)
|
.route_completion(Some(&headers), &body, Some(&body.model))
|
||||||
.await
|
.await
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -166,7 +170,10 @@ async fn rerank(
|
|||||||
headers: http::HeaderMap,
|
headers: http::HeaderMap,
|
||||||
ValidatedJson(body): ValidatedJson<RerankRequest>,
|
ValidatedJson(body): ValidatedJson<RerankRequest>,
|
||||||
) -> Response {
|
) -> 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(
|
async fn v1_rerank(
|
||||||
@@ -174,9 +181,10 @@ async fn v1_rerank(
|
|||||||
headers: http::HeaderMap,
|
headers: http::HeaderMap,
|
||||||
Json(body): Json<V1RerankReqInput>,
|
Json(body): Json<V1RerankReqInput>,
|
||||||
) -> Response {
|
) -> Response {
|
||||||
|
let rerank_body = &body.into();
|
||||||
state
|
state
|
||||||
.router
|
.router
|
||||||
.route_rerank(Some(&headers), &body.into(), None)
|
.route_rerank(Some(&headers), rerank_body, Some(&rerank_body.model))
|
||||||
.await
|
.await
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -187,7 +195,7 @@ async fn v1_responses(
|
|||||||
) -> Response {
|
) -> Response {
|
||||||
state
|
state
|
||||||
.router
|
.router
|
||||||
.route_responses(Some(&headers), &body, None)
|
.route_responses(Some(&headers), &body, Some(&body.model))
|
||||||
.await
|
.await
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -198,7 +206,7 @@ async fn v1_embeddings(
|
|||||||
) -> Response {
|
) -> Response {
|
||||||
state
|
state
|
||||||
.router
|
.router
|
||||||
.route_embeddings(Some(&headers), &body, None)
|
.route_embeddings(Some(&headers), &body, Some(&body.model))
|
||||||
.await
|
.await
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -209,7 +217,7 @@ async fn v1_classify(
|
|||||||
) -> Response {
|
) -> Response {
|
||||||
state
|
state
|
||||||
.router
|
.router
|
||||||
.route_classify(Some(&headers), &body, None)
|
.route_classify(Some(&headers), &body, Some(&body.model))
|
||||||
.await
|
.await
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -602,6 +602,7 @@ async fn test_unsupported_endpoints() {
|
|||||||
|
|
||||||
let generate_request = GenerateRequest {
|
let generate_request = GenerateRequest {
|
||||||
text: Some("Hello world".to_string()),
|
text: Some("Hello world".to_string()),
|
||||||
|
model: None,
|
||||||
input_ids: None,
|
input_ids: None,
|
||||||
input_embeds: None,
|
input_embeds: None,
|
||||||
image_data: None,
|
image_data: None,
|
||||||
|
|||||||
Reference in New Issue
Block a user