[model-gateway] code clean up on oai router (#14850)

This commit is contained in:
Simo Lin
2025-12-10 15:34:27 -08:00
committed by GitHub
parent a4992873d4
commit ccf2602773
+166 -230
View File
@@ -65,9 +65,79 @@ impl std::fmt::Debug for OpenAIRouter {
} }
} }
/// Error response helpers for consistent API error formatting
mod error_responses {
use axum::{
http::StatusCode,
response::{IntoResponse, Response},
Json,
};
use serde_json::json;
pub fn bad_request(message: impl Into<String>) -> Response {
(StatusCode::BAD_REQUEST, message.into()).into_response()
}
pub fn not_found(resource: &str, id: &str) -> Response {
(
StatusCode::NOT_FOUND,
Json(json!({
"error": {
"message": format!("No {} found with id '{}'", resource, id),
"type": "invalid_request_error",
"param": null,
"code": "not_found"
}
})),
)
.into_response()
}
pub fn internal_error(message: impl Into<String>) -> Response {
(
StatusCode::INTERNAL_SERVER_ERROR,
Json(json!({
"error": {
"message": message.into(),
"type": "internal_error",
"param": null,
"code": "storage_error"
}
})),
)
.into_response()
}
pub fn service_unavailable(message: impl Into<String>) -> Response {
(StatusCode::SERVICE_UNAVAILABLE, message.into()).into_response()
}
pub fn model_not_found(model: &str) -> Response {
(
StatusCode::NOT_FOUND,
Json(json!({
"error": {
"message": format!("No worker available for model '{}'", model),
"type": "model_not_found",
}
})),
)
.into_response()
}
}
impl OpenAIRouter { impl OpenAIRouter {
const MAX_CONVERSATION_HISTORY_ITEMS: usize = 100; const MAX_CONVERSATION_HISTORY_ITEMS: usize = 100;
/// Get all external workers from the registry
fn external_workers(&self) -> Vec<Arc<dyn Worker>> {
self.worker_registry
.get_all()
.into_iter()
.filter(|w| w.metadata().runtime_type == RuntimeType::External)
.collect()
}
fn shared_components(&self) -> Arc<SharedComponents> { fn shared_components(&self) -> Arc<SharedComponents> {
Arc::clone(&self.shared_components) Arc::clone(&self.shared_components)
} }
@@ -204,53 +274,73 @@ impl OpenAIRouter {
join_all(futures).await; join_all(futures).await;
} }
async fn select_worker_for_model( /// Find workers that can handle the given model and select the least loaded one
&self, fn find_best_worker_for_model(&self, model_id: &str) -> Option<Arc<dyn Worker>> {
model_id: &str,
auth_header: Option<&HeaderValue>,
) -> Result<Arc<dyn Worker>, Box<Response>> {
let find_candidates = || {
self.worker_registry self.worker_registry
.get_workers_filtered(None, None, None, Some(RuntimeType::External), true) .get_workers_filtered(None, None, None, Some(RuntimeType::External), true)
.into_iter() .into_iter()
.filter(|w| w.supports_model(model_id) && w.circuit_breaker().can_execute()) .filter(|w| w.supports_model(model_id) && w.circuit_breaker().can_execute())
.collect::<Vec<_>>()
};
let candidates = find_candidates();
if !candidates.is_empty() {
return Ok(candidates
.into_iter()
.min_by_key(|w| w.load()) .min_by_key(|w| w.load())
.expect("candidates is not empty"));
} }
async fn select_worker_for_model(
&self,
model_id: &str,
auth_header: Option<&HeaderValue>,
) -> Result<Arc<dyn Worker>, Response> {
// Try to find a worker immediately
if let Some(worker) = self.find_best_worker_for_model(model_id) {
return Ok(worker);
}
// Refresh external models and try again
tracing::debug!( tracing::debug!(
"No worker found for model '{}', refreshing external worker models", "No worker found for model '{}', refreshing external worker models",
model_id model_id
); );
self.refresh_external_models(auth_header).await; self.refresh_external_models(auth_header).await;
let candidates = find_candidates(); self.find_best_worker_for_model(model_id)
if !candidates.is_empty() { .ok_or_else(|| error_responses::model_not_found(model_id))
return Ok(candidates
.into_iter()
.min_by_key(|w| w.load())
.expect("candidates is not empty"));
} }
Err(Box::new( /// Deserialize ResponseInputOutputItems from a JSON array value
( fn deserialize_items_from_array(array: &Value) -> Vec<ResponseInputOutputItem> {
StatusCode::NOT_FOUND, array
Json(json!({ .as_array()
"error": { .map(|arr| {
"message": format!("No worker available for model '{}'", model_id), arr.iter()
"type": "model_not_found", .filter_map(|item| {
serde_json::from_value::<ResponseInputOutputItem>(item.clone())
.map_err(|e| warn!("Failed to deserialize item: {}. Item: {}", e, item))
.ok()
})
.collect()
})
.unwrap_or_default()
}
/// Append current request input to items list, creating a user message if needed
fn append_current_input(
items: &mut Vec<ResponseInputOutputItem>,
input: &ResponseInput,
id_suffix: &str,
) {
match input {
ResponseInput::Text(text) => {
items.push(ResponseInputOutputItem::Message {
id: format!("msg_u_{}", id_suffix),
role: "user".to_string(),
content: vec![ResponseContentPart::InputText { text: text.clone() }],
status: Some("completed".to_string()),
});
}
ResponseInput::Items(current_items) => {
for item in current_items {
items.push(crate::protocols::responses::normalize_input_item(item));
}
}
} }
})),
)
.into_response(),
))
} }
async fn handle_non_streaming_response(&self, mut ctx: RequestContext) -> Response { async fn handle_non_streaming_response(&self, mut ctx: RequestContext) -> Response {
@@ -383,65 +473,38 @@ impl crate::routers::RouterTrait for OpenAIRouter {
} }
async fn health_generate(&self, _req: Request<Body>) -> Response { async fn health_generate(&self, _req: Request<Body>) -> Response {
let external_workers: Vec<_> = self let external_workers = self.external_workers();
.worker_registry
.get_all()
.into_iter()
.filter(|w| w.metadata().runtime_type == RuntimeType::External)
.collect();
if external_workers.is_empty() { if external_workers.is_empty() {
return ( return error_responses::service_unavailable("No external workers registered");
StatusCode::SERVICE_UNAVAILABLE,
"No external workers registered",
)
.into_response();
} }
let mut healthy_count = 0; let (healthy, unhealthy): (Vec<_>, Vec<_>) =
let mut unhealthy_workers = Vec::new(); external_workers.iter().partition(|w| w.is_healthy());
for worker in &external_workers { if unhealthy.is_empty() {
if worker.is_healthy() {
healthy_count += 1;
} else {
unhealthy_workers.push(format!("{} ({})", worker.model_id(), worker.url()));
}
}
if unhealthy_workers.is_empty() {
( (
StatusCode::OK, StatusCode::OK,
format!("OK - {} workers healthy", healthy_count), format!("OK - {} workers healthy", healthy.len()),
) )
.into_response() .into_response()
} else { } else {
( let unhealthy_info: Vec<_> = unhealthy
StatusCode::SERVICE_UNAVAILABLE, .iter()
format!( .map(|w| format!("{} ({})", w.model_id(), w.url()))
.collect();
error_responses::service_unavailable(format!(
"{}/{} workers unhealthy: {}", "{}/{} workers unhealthy: {}",
unhealthy_workers.len(), unhealthy.len(),
external_workers.len(), external_workers.len(),
unhealthy_workers.join(", ") unhealthy_info.join(", ")
), ))
)
.into_response()
} }
} }
async fn get_server_info(&self, _req: Request<Body>) -> Response { async fn get_server_info(&self, _req: Request<Body>) -> Response {
let stats = self.worker_registry.stats(); let stats = self.worker_registry.stats();
let external_workers: Vec<_> = self let external_workers = self.external_workers();
.worker_registry let worker_urls: Vec<_> = external_workers.iter().map(|w| w.url()).collect();
.get_all()
.into_iter()
.filter(|w| w.metadata().runtime_type == RuntimeType::External)
.collect();
let worker_urls: Vec<String> = external_workers
.iter()
.map(|w| w.url().to_string())
.collect();
let info = json!({ let info = json!({
"router_type": "openai", "router_type": "openai",
@@ -455,19 +518,9 @@ impl crate::routers::RouterTrait for OpenAIRouter {
} }
async fn get_models(&self, req: Request<Body>) -> Response { async fn get_models(&self, req: Request<Body>) -> Response {
let external_workers: Vec<_> = self let external_workers = self.external_workers();
.worker_registry
.get_all()
.into_iter()
.filter(|w| w.metadata().runtime_type == RuntimeType::External)
.collect();
if external_workers.is_empty() { if external_workers.is_empty() {
return ( return error_responses::service_unavailable("No external workers registered");
StatusCode::SERVICE_UNAVAILABLE,
"No external workers registered",
)
.into_response();
} }
let auth_header = extract_auth_header(Some(req.headers()), &None); let auth_header = extract_auth_header(Some(req.headers()), &None);
@@ -530,27 +583,19 @@ impl crate::routers::RouterTrait for OpenAIRouter {
.await .await
{ {
Ok(w) => w, Ok(w) => w,
Err(response) => return *response, Err(response) => return response,
}; };
let mut payload = match to_value(body) { let mut payload = match to_value(body) {
Ok(v) => v, Ok(v) => v,
Err(e) => { Err(e) => {
return ( return error_responses::bad_request(format!("Failed to serialize request: {}", e))
StatusCode::BAD_REQUEST,
format!("Failed to serialize request: {}", e),
)
.into_response();
} }
}; };
let provider = self.get_provider_arc_for_worker(worker.as_ref(), model_id); let provider = self.get_provider_arc_for_worker(worker.as_ref(), model_id);
if let Err(e) = provider.transform_request(&mut payload, Endpoint::Chat) { if let Err(e) = provider.transform_request(&mut payload, Endpoint::Chat) {
return ( return error_responses::bad_request(format!("Provider transform error: {}", e));
StatusCode::BAD_REQUEST,
format!("Provider transform error: {}", e),
)
.into_response();
} }
let mut ctx = RequestContext::for_chat( let mut ctx = RequestContext::for_chat(
@@ -659,7 +704,7 @@ impl crate::routers::RouterTrait for OpenAIRouter {
.await .await
{ {
Ok(w) => w, Ok(w) => w,
Err(response) => return *response, Err(response) => return response,
}; };
let mut request_body = body.clone(); let mut request_body = body.clone();
@@ -670,8 +715,9 @@ impl crate::routers::RouterTrait for OpenAIRouter {
let original_previous_response_id = request_body.previous_response_id.clone(); let original_previous_response_id = request_body.previous_response_id.clone();
// Load items from previous response chain if specified
let mut conversation_items: Option<Vec<ResponseInputOutputItem>> = None; let mut conversation_items: Option<Vec<ResponseInputOutputItem>> = None;
if let Some(prev_id_str) = request_body.previous_response_id.clone() { if let Some(prev_id_str) = request_body.previous_response_id.take() {
let prev_id = ResponseId::from(prev_id_str.as_str()); let prev_id = ResponseId::from(prev_id_str.as_str());
match self match self
.responses_components .responses_components
@@ -680,43 +726,16 @@ impl crate::routers::RouterTrait for OpenAIRouter {
.await .await
{ {
Ok(chain) => { Ok(chain) => {
let mut items = Vec::new(); let items: Vec<ResponseInputOutputItem> = chain
for stored in chain.responses.iter() { .responses
if let Some(input_arr) = stored.input.as_array() { .iter()
for item in input_arr { .flat_map(|stored| {
match serde_json::from_value::<ResponseInputOutputItem>( Self::deserialize_items_from_array(&stored.input)
item.clone(), .into_iter()
) { .chain(Self::deserialize_items_from_array(&stored.output))
Ok(input_item) => { })
items.push(input_item); .collect();
}
Err(e) => {
warn!(
"Failed to deserialize stored input item: {}. Item: {}",
e, item
);
}
}
}
}
if let Some(output_arr) = stored.output.as_array() {
for item in output_arr {
match serde_json::from_value::<ResponseInputOutputItem>(
item.clone(),
) {
Ok(output_item) => {
items.push(output_item);
}
Err(e) => {
warn!("Failed to deserialize stored output item: {}. Item: {}", e, item);
}
}
}
}
}
conversation_items = Some(items); conversation_items = Some(items);
request_body.previous_response_id = None;
} }
Err(e) => { Err(e) => {
warn!( warn!(
@@ -736,11 +755,7 @@ impl crate::routers::RouterTrait for OpenAIRouter {
.get_conversation(&conv_id) .get_conversation(&conv_id)
.await .await
{ {
return ( return error_responses::not_found("conversation", &conv_id.0);
StatusCode::NOT_FOUND,
Json(json!({"error": "Conversation not found"})),
)
.into_response();
} }
let params = ListParams { let params = ListParams {
@@ -825,26 +840,7 @@ impl crate::routers::RouterTrait for OpenAIRouter {
} }
} }
match &request_body.input { Self::append_current_input(&mut items, &request_body.input, &conv_id.0);
ResponseInput::Text(text) => {
items.push(ResponseInputOutputItem::Message {
id: format!("msg_u_{}", conv_id.0),
role: "user".to_string(),
content: vec![ResponseContentPart::InputText {
text: text.clone(),
}],
status: Some("completed".to_string()),
});
}
ResponseInput::Items(current_items) => {
for item in current_items.iter() {
let normalized =
crate::protocols::responses::normalize_input_item(item);
items.push(normalized);
}
}
}
request_body.input = ResponseInput::Items(items); request_body.input = ResponseInput::Items(items);
} }
Err(e) => { Err(e) => {
@@ -853,29 +849,10 @@ impl crate::routers::RouterTrait for OpenAIRouter {
} }
} }
// Apply previous response chain items if loaded
if let Some(mut items) = conversation_items { if let Some(mut items) = conversation_items {
match &request_body.input { let id_suffix = original_previous_response_id.as_deref().unwrap_or("new");
ResponseInput::Text(text) => { Self::append_current_input(&mut items, &request_body.input, id_suffix);
items.push(ResponseInputOutputItem::Message {
id: format!(
"msg_u_{}",
original_previous_response_id
.as_ref()
.unwrap_or(&"new".to_string())
),
role: "user".to_string(),
content: vec![ResponseContentPart::InputText { text: text.clone() }],
status: Some("completed".to_string()),
});
}
ResponseInput::Items(current_items) => {
for item in current_items.iter() {
let normalized = crate::protocols::responses::normalize_input_item(item);
items.push(normalized);
}
}
}
request_body.input = ResponseInput::Items(items); request_body.input = ResponseInput::Items(items);
} }
@@ -887,21 +864,13 @@ impl crate::routers::RouterTrait for OpenAIRouter {
let mut payload = match to_value(&request_body) { let mut payload = match to_value(&request_body) {
Ok(v) => v, Ok(v) => v,
Err(e) => { Err(e) => {
return ( return error_responses::bad_request(format!("Failed to serialize request: {}", e))
StatusCode::BAD_REQUEST,
format!("Failed to serialize request: {}", e),
)
.into_response();
} }
}; };
let provider = self.get_provider_arc_for_worker(worker.as_ref(), model_id); let provider = self.get_provider_arc_for_worker(worker.as_ref(), model_id);
if let Err(e) = provider.transform_request(&mut payload, Endpoint::Responses) { if let Err(e) = provider.transform_request(&mut payload, Endpoint::Responses) {
return ( return error_responses::bad_request(format!("Provider transform error: {}", e));
StatusCode::BAD_REQUEST,
format!("Provider transform error: {}", e),
)
.into_response();
} }
let mut ctx = RequestContext::for_responses( let mut ctx = RequestContext::for_responses(
@@ -949,16 +918,8 @@ impl crate::routers::RouterTrait for OpenAIRouter {
} }
(StatusCode::OK, Json(response_json)).into_response() (StatusCode::OK, Json(response_json)).into_response()
} }
Ok(None) => ( Ok(None) => error_responses::not_found("response", response_id),
StatusCode::NOT_FOUND, Err(e) => error_responses::internal_error(format!("Failed to get response: {}", e)),
Json(json!({"error": "Response not found"})),
)
.into_response(),
Err(e) => (
StatusCode::INTERNAL_SERVER_ERROR,
Json(json!({ "error": format!("Failed to get response: {}", e) })),
)
.into_response(),
} }
} }
@@ -976,10 +937,7 @@ impl crate::routers::RouterTrait for OpenAIRouter {
.await .await
{ {
Ok(Some(stored)) => { Ok(Some(stored)) => {
let items = match &stored.input { let items = stored.input.as_array().cloned().unwrap_or_default();
Value::Array(arr) => arr.clone(),
_ => vec![],
};
let items_with_ids: Vec<Value> = items let items_with_ids: Vec<Value> = items
.into_iter() .into_iter()
@@ -1003,32 +961,10 @@ impl crate::routers::RouterTrait for OpenAIRouter {
(StatusCode::OK, Json(response_body)).into_response() (StatusCode::OK, Json(response_body)).into_response()
} }
Ok(None) => ( Ok(None) => error_responses::not_found("response", response_id),
StatusCode::NOT_FOUND,
Json(json!({
"error": {
"message": format!("No response found with id '{}'", response_id),
"type": "invalid_request_error",
"param": Value::Null,
"code": "not_found"
}
})),
)
.into_response(),
Err(e) => { Err(e) => {
warn!("Failed to retrieve input items for {}: {}", response_id, e); warn!("Failed to retrieve input items for {}: {}", response_id, e);
( error_responses::internal_error(format!("Failed to retrieve input items: {}", e))
StatusCode::INTERNAL_SERVER_ERROR,
Json(json!({
"error": {
"message": format!("Failed to retrieve input items: {}", e),
"type": "internal_error",
"param": Value::Null,
"code": "storage_error"
}
})),
)
.into_response()
} }
} }
} }