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
4 changes: 2 additions & 2 deletions model_gateway/src/routers/openai/chat.rs
Original file line number Diff line number Diff line change
Expand Up @@ -30,7 +30,7 @@ use crate::{
};

/// Shared context passed to chat routing functions.
pub(super) struct RouterContext<'a> {
pub(super) struct ChatRouterContext<'a> {
pub worker_registry: &'a WorkerRegistry,
pub provider_registry: &'a ProviderRegistry,
pub shared_components: &'a Arc<SharedComponents>,
Expand All @@ -40,7 +40,7 @@ pub(super) struct RouterContext<'a> {

/// Route a chat completion request to the appropriate upstream worker.
pub(super) async fn route_chat(
deps: &RouterContext<'_>,
deps: &ChatRouterContext<'_>,
headers: Option<&HeaderMap>,
body: &ChatCompletionRequest,
model_id: Option<&str>,
Expand Down
1 change: 1 addition & 0 deletions model_gateway/src/routers/openai/responses/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@ mod accumulator;
mod common;
pub(crate) mod history;
mod non_streaming;
pub(crate) mod route;
mod streaming;
mod utils;

Expand Down
195 changes: 195 additions & 0 deletions model_gateway/src/routers/openai/responses/route.rs
Original file line number Diff line number Diff line change
@@ -0,0 +1,195 @@
//! Responses API routing orchestration.
//!
//! Mirrors the delegation pattern in `chat.rs`: the `RouterTrait` method in
//! `router.rs` packs borrowed references into [`ResponsesRouterContext`] and
//! delegates to [`route_responses`].

use std::{sync::Arc, time::Instant};

use axum::{http::HeaderMap, response::Response};
use openai_protocol::responses::{ResponseInput, ResponseInputOutputItem, ResponsesRequest};
use serde_json::to_value;

use super::{
super::{
context::{
ComponentRefs, PayloadState, RequestContext, ResponsesComponents, WorkerSelection,
},
provider::ProviderRegistry,
router::resolve_provider,
},
handle_non_streaming_response, handle_streaming_response,
};
use crate::{
core::{Endpoint, ProviderType, WorkerRegistry},
observability::metrics::{bool_to_static_str, metrics_labels, Metrics},
routers::{
error,
worker_selection::{SelectWorkerRequest, WorkerSelector},
},
};

/// Shared context passed to responses routing functions.
pub(in crate::routers::openai) struct ResponsesRouterContext<'a> {
pub worker_registry: &'a WorkerRegistry,
pub provider_registry: &'a ProviderRegistry,
pub responses_components: &'a Arc<ResponsesComponents>,
pub client: &'a reqwest::Client,

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

medium

The client field in ResponsesRouterContext is redundant, as the reqwest::Client can be accessed via deps.responses_components.shared.client. Removing this field will simplify the context struct.

}

/// Route a responses API request to the appropriate upstream worker.
pub(in crate::routers::openai) async fn route_responses(
deps: &ResponsesRouterContext<'_>,
headers: Option<&HeaderMap>,
body: &ResponsesRequest,
model_id: Option<&str>,
) -> Response {
let start = Instant::now();
let model = model_id.unwrap_or(body.model.as_str());
let streaming = body.stream.unwrap_or(false);

Metrics::record_router_request(
metrics_labels::ROUTER_OPENAI,
metrics_labels::BACKEND_EXTERNAL,
metrics_labels::CONNECTION_HTTP,
model,
metrics_labels::ENDPOINT_RESPONSES,
bool_to_static_str(streaming),
);

let worker = match WorkerSelector::new(deps.worker_registry, deps.client)

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

medium

To go along with the removal of the redundant client field from ResponsesRouterContext, please update this to access the client via deps.responses_components.shared.client.

Suggested change
let worker = match WorkerSelector::new(deps.worker_registry, deps.client)
let worker = match WorkerSelector::new(deps.worker_registry, &deps.responses_components.shared.client)

.select_worker(&SelectWorkerRequest {
model_id: model,
headers,
provider: Some(ProviderType::OpenAI),
..Default::default()
})
.await
{
Ok(w) => w,
Err(response) => {
Metrics::record_router_error(
metrics_labels::ROUTER_OPENAI,
metrics_labels::BACKEND_EXTERNAL,
metrics_labels::CONNECTION_HTTP,
model,
metrics_labels::ENDPOINT_RESPONSES,
metrics_labels::ERROR_NO_WORKERS,
);
return response;
}
};

// Validate mutual exclusivity of conversation and previous_response_id
// Treat empty strings as unset to match other metadata paths
let has_conversation = body.conversation.as_ref().is_some_and(|s| !s.is_empty());
let has_previous_response = body
.previous_response_id
.as_ref()
.is_some_and(|s| !s.is_empty());
if has_conversation && has_previous_response {
Metrics::record_router_error(
metrics_labels::ROUTER_OPENAI,
metrics_labels::BACKEND_EXTERNAL,
metrics_labels::CONNECTION_HTTP,
model,
metrics_labels::ENDPOINT_RESPONSES,
metrics_labels::ERROR_VALIDATION,
);
return error::bad_request(
"invalid_request",
"Cannot specify both 'conversation' and 'previous_response_id'".to_string(),
);
}

let mut request_body = body.clone();
if let Some(model) = model_id {
request_body.model = model.to_string();
}
request_body.conversation = None;

let original_previous_response_id = match super::history::load_input_history(
deps.responses_components,
body,
&mut request_body,
model,
)
.await
{
Ok(id) => id,
Err(response) => return response,
};

request_body.store = Some(false);
if let ResponseInput::Items(ref mut items) = request_body.input {
items.retain(|item| !matches!(item, ResponseInputOutputItem::Reasoning { .. }));
}

let mut payload = match to_value(&request_body) {
Ok(v) => v,
Err(e) => {
Metrics::record_router_error(
metrics_labels::ROUTER_OPENAI,
metrics_labels::BACKEND_EXTERNAL,
metrics_labels::CONNECTION_HTTP,
model,
metrics_labels::ENDPOINT_RESPONSES,
metrics_labels::ERROR_VALIDATION,
);
return error::bad_request(
"invalid_request",
format!("Failed to serialize request: {e}"),
);
}
};

let provider = resolve_provider(deps.provider_registry, worker.as_ref(), model);
if let Err(e) = provider.transform_request(&mut payload, Endpoint::Responses) {
Metrics::record_router_error(
metrics_labels::ROUTER_OPENAI,
metrics_labels::BACKEND_EXTERNAL,
metrics_labels::CONNECTION_HTTP,
model,
metrics_labels::ENDPOINT_RESPONSES,
metrics_labels::ERROR_VALIDATION,
);
return error::bad_request("invalid_request", format!("Provider transform error: {e}"));
}

let mut ctx = RequestContext::for_responses(
Arc::new(body.clone()),
headers.cloned(),
model_id.map(String::from),
ComponentRefs::Responses(Arc::clone(deps.responses_components)),
);

ctx.state.worker = Some(WorkerSelection {
worker: Arc::clone(&worker),
provider: Arc::clone(&provider),
});

ctx.state.payload = Some(PayloadState {
json: payload,
url: format!("{}/v1/responses", worker.url()),
previous_response_id: original_previous_response_id,
});

let response = if ctx.is_streaming() {
handle_streaming_response(ctx).await
} else {
handle_non_streaming_response(ctx).await
};

if response.status().is_success() {
Metrics::record_router_duration(
metrics_labels::ROUTER_OPENAI,
metrics_labels::BACKEND_EXTERNAL,
metrics_labels::CONNECTION_HTTP,
model,
metrics_labels::ENDPOINT_RESPONSES,
start.elapsed(),
);
}

response
}
Loading
Loading